chore(server): centralize CORS policy
This commit is contained in:
parent
d69f1aab27
commit
bb76949fc0
5 changed files with 167 additions and 65 deletions
43
server/cors.go
Normal file
43
server/cors.go
Normal file
|
|
@ -0,0 +1,43 @@
|
||||||
|
package server
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/url"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/labstack/echo/v5"
|
||||||
|
"github.com/labstack/echo/v5/middleware"
|
||||||
|
|
||||||
|
"github.com/usememos/memos/internal/profile"
|
||||||
|
)
|
||||||
|
|
||||||
|
func newCORSMiddleware(profile *profile.Profile) echo.MiddlewareFunc {
|
||||||
|
return middleware.CORSWithConfig(middleware.CORSConfig{
|
||||||
|
UnsafeAllowOriginFunc: func(c *echo.Context, origin string) (string, bool, error) {
|
||||||
|
if isAllowedCORSOrigin(profile, c.Request().Host, origin) {
|
||||||
|
return origin, true, nil
|
||||||
|
}
|
||||||
|
return "", false, nil
|
||||||
|
},
|
||||||
|
AllowCredentials: true,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func isAllowedCORSOrigin(profile *profile.Profile, requestHost, origin string) bool {
|
||||||
|
originURL, err := url.Parse(origin)
|
||||||
|
if err != nil || originURL.Scheme == "" || originURL.Host == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
if strings.EqualFold(originURL.Host, requestHost) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
if profile == nil || profile.InstanceURL == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
instanceURL, err := url.Parse(profile.InstanceURL)
|
||||||
|
if err != nil || instanceURL.Scheme == "" || instanceURL.Host == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return strings.EqualFold(originURL.Scheme, instanceURL.Scheme) && strings.EqualFold(originURL.Host, instanceURL.Host)
|
||||||
|
}
|
||||||
122
server/cors_test.go
Normal file
122
server/cors_test.go
Normal file
|
|
@ -0,0 +1,122 @@
|
||||||
|
package server
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/labstack/echo/v5"
|
||||||
|
|
||||||
|
"github.com/usememos/memos/internal/profile"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAllowedCORSOrigin(t *testing.T) {
|
||||||
|
p := &profile.Profile{InstanceURL: "https://memos.example"}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
requestHost string
|
||||||
|
origin string
|
||||||
|
allowed bool
|
||||||
|
}{
|
||||||
|
{
|
||||||
|
name: "same host",
|
||||||
|
requestHost: "localhost",
|
||||||
|
origin: "http://localhost",
|
||||||
|
allowed: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "instance URL",
|
||||||
|
requestHost: "localhost",
|
||||||
|
origin: "https://memos.example",
|
||||||
|
allowed: true,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "usememos apex",
|
||||||
|
requestHost: "localhost",
|
||||||
|
origin: "https://usememos.com",
|
||||||
|
allowed: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "usememos subdomain",
|
||||||
|
requestHost: "localhost",
|
||||||
|
origin: "https://demo.usememos.com",
|
||||||
|
allowed: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "nested usememos subdomain",
|
||||||
|
requestHost: "localhost",
|
||||||
|
origin: "https://preview.demo.usememos.com",
|
||||||
|
allowed: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "usememos subdomain with port",
|
||||||
|
requestHost: "localhost",
|
||||||
|
origin: "http://localhost.usememos.com:3001",
|
||||||
|
allowed: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "unknown origin",
|
||||||
|
requestHost: "localhost",
|
||||||
|
origin: "https://evil.example",
|
||||||
|
allowed: false,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
name: "lookalike domain",
|
||||||
|
requestHost: "localhost",
|
||||||
|
origin: "https://evilusememos.com",
|
||||||
|
allowed: false,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, test := range tests {
|
||||||
|
t.Run(test.name, func(t *testing.T) {
|
||||||
|
if allowed := isAllowedCORSOrigin(p, test.requestHost, test.origin); allowed != test.allowed {
|
||||||
|
t.Fatalf("expected allowed=%t, got %t", test.allowed, allowed)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCORSMiddleware(t *testing.T) {
|
||||||
|
e := echo.New()
|
||||||
|
e.Use(newCORSMiddleware(&profile.Profile{InstanceURL: "https://memos.example"}))
|
||||||
|
e.POST("/api/v1/test", func(c *echo.Context) error {
|
||||||
|
return c.NoContent(http.StatusOK)
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("allows instance URL origin on preflight", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodOptions, "/api/v1/test", nil)
|
||||||
|
req.Header.Set("Origin", "https://memos.example")
|
||||||
|
req.Header.Set("Access-Control-Request-Method", http.MethodPost)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
|
e.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusNoContent {
|
||||||
|
t.Fatalf("expected status %d, got %d", http.StatusNoContent, rec.Code)
|
||||||
|
}
|
||||||
|
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "https://memos.example" {
|
||||||
|
t.Fatalf("unexpected Access-Control-Allow-Origin: %q", got)
|
||||||
|
}
|
||||||
|
if got := rec.Header().Get("Access-Control-Allow-Credentials"); got != "true" {
|
||||||
|
t.Fatalf("unexpected Access-Control-Allow-Credentials: %q", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
t.Run("omits CORS headers for unknown origin preflight", func(t *testing.T) {
|
||||||
|
req := httptest.NewRequest(http.MethodOptions, "/api/v1/test", nil)
|
||||||
|
req.Header.Set("Origin", "https://evil.example")
|
||||||
|
req.Header.Set("Access-Control-Request-Method", http.MethodPost)
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
|
||||||
|
e.ServeHTTP(rec, req)
|
||||||
|
|
||||||
|
if rec.Code != http.StatusNoContent {
|
||||||
|
t.Fatalf("expected status %d, got %d", http.StatusNoContent, rec.Code)
|
||||||
|
}
|
||||||
|
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "" {
|
||||||
|
t.Fatalf("expected no Access-Control-Allow-Origin, got %q", got)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
@ -2,16 +2,11 @@ package v1
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"connectrpc.com/connect"
|
"connectrpc.com/connect"
|
||||||
"github.com/labstack/echo/v5"
|
|
||||||
"google.golang.org/grpc/metadata"
|
"google.golang.org/grpc/metadata"
|
||||||
"google.golang.org/protobuf/types/known/emptypb"
|
"google.golang.org/protobuf/types/known/emptypb"
|
||||||
|
|
||||||
"github.com/usememos/memos/internal/profile"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestMetadataInterceptorForwardsSecurityHeaders(t *testing.T) {
|
func TestMetadataInterceptorForwardsSecurityHeaders(t *testing.T) {
|
||||||
|
|
@ -42,24 +37,3 @@ func TestMetadataInterceptorForwardsSecurityHeaders(t *testing.T) {
|
||||||
t.Fatalf("metadata interceptor returned error: %v", err)
|
t.Fatalf("metadata interceptor returned error: %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAllowedConnectOrigin(t *testing.T) {
|
|
||||||
service := &APIV1Service{
|
|
||||||
Profile: &profile.Profile{InstanceURL: "https://memos.example"},
|
|
||||||
}
|
|
||||||
e := echo.New()
|
|
||||||
req := httptest.NewRequest(http.MethodOptions, "http://localhost/memos.api.v1.AuthService/SignIn", nil)
|
|
||||||
req.Host = "localhost"
|
|
||||||
rec := httptest.NewRecorder()
|
|
||||||
ctx := e.NewContext(req, rec)
|
|
||||||
|
|
||||||
if !service.isAllowedConnectOrigin(ctx, "http://localhost") {
|
|
||||||
t.Fatal("expected same host origin to be allowed")
|
|
||||||
}
|
|
||||||
if !service.isAllowedConnectOrigin(ctx, "https://memos.example") {
|
|
||||||
t.Fatal("expected instance URL origin to be allowed")
|
|
||||||
}
|
|
||||||
if service.isAllowedConnectOrigin(ctx, "https://evil.example") {
|
|
||||||
t.Fatal("expected unknown origin to be denied")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -3,13 +3,10 @@ package v1
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
|
||||||
"strings"
|
|
||||||
|
|
||||||
"connectrpc.com/connect"
|
"connectrpc.com/connect"
|
||||||
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
"github.com/grpc-ecosystem/grpc-gateway/v2/runtime"
|
||||||
"github.com/labstack/echo/v5"
|
"github.com/labstack/echo/v5"
|
||||||
"github.com/labstack/echo/v5/middleware"
|
|
||||||
"golang.org/x/sync/semaphore"
|
"golang.org/x/sync/semaphore"
|
||||||
|
|
||||||
"github.com/usememos/memos/internal/markdown"
|
"github.com/usememos/memos/internal/markdown"
|
||||||
|
|
@ -127,9 +124,6 @@ func (s *APIV1Service) RegisterGateway(ctx context.Context, echoServer *echo.Ech
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
gwGroup := echoServer.Group("")
|
gwGroup := echoServer.Group("")
|
||||||
gwGroup.Use(middleware.CORSWithConfig(middleware.CORSConfig{
|
|
||||||
AllowOrigins: []string{"*"},
|
|
||||||
}))
|
|
||||||
// Register SSE endpoint with same CORS as rest of /api/v1.
|
// Register SSE endpoint with same CORS as rest of /api/v1.
|
||||||
RegisterSSERoutes(gwGroup, s.SSEHub, s.Store, s.Secret)
|
RegisterSSERoutes(gwGroup, s.SSEHub, s.Store, s.Secret)
|
||||||
handler := echo.WrapHandler(http.MaxBytesHandler(gwMux, maxAPIRequestBytes))
|
handler := echo.WrapHandler(http.MaxBytesHandler(gwMux, maxAPIRequestBytes))
|
||||||
|
|
@ -149,40 +143,8 @@ func (s *APIV1Service) RegisterGateway(ctx context.Context, echoServer *echo.Ech
|
||||||
connectHandler := NewConnectServiceHandler(s)
|
connectHandler := NewConnectServiceHandler(s)
|
||||||
connectHandler.RegisterConnectHandlers(connectMux, connectInterceptors, connect.WithReadMaxBytes(maxAPIRequestBytes))
|
connectHandler.RegisterConnectHandlers(connectMux, connectInterceptors, connect.WithReadMaxBytes(maxAPIRequestBytes))
|
||||||
|
|
||||||
// Wrap with CORS for browser access
|
connectGroup := echoServer.Group("")
|
||||||
corsHandler := middleware.CORSWithConfig(middleware.CORSConfig{
|
|
||||||
UnsafeAllowOriginFunc: func(c *echo.Context, origin string) (string, bool, error) {
|
|
||||||
if s.isAllowedConnectOrigin(c, origin) {
|
|
||||||
return origin, true, nil
|
|
||||||
}
|
|
||||||
return "", false, nil
|
|
||||||
},
|
|
||||||
AllowMethods: []string{http.MethodGet, http.MethodPost, http.MethodOptions},
|
|
||||||
AllowHeaders: []string{"*"},
|
|
||||||
AllowCredentials: true,
|
|
||||||
})
|
|
||||||
connectGroup := echoServer.Group("", corsHandler)
|
|
||||||
connectGroup.Any("/memos.api.v1.*", echo.WrapHandler(http.MaxBytesHandler(connectMux, maxAPIRequestBytes)))
|
connectGroup.Any("/memos.api.v1.*", echo.WrapHandler(http.MaxBytesHandler(connectMux, maxAPIRequestBytes)))
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *APIV1Service) isAllowedConnectOrigin(c *echo.Context, origin string) bool {
|
|
||||||
originURL, err := url.Parse(origin)
|
|
||||||
if err != nil || originURL.Scheme == "" || originURL.Host == "" {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
if strings.EqualFold(originURL.Host, c.Request().Host) {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
if s.Profile == nil || s.Profile.InstanceURL == "" {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
instanceURL, err := url.Parse(s.Profile.InstanceURL)
|
|
||||||
if err != nil || instanceURL.Scheme == "" || instanceURL.Host == "" {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
return strings.EqualFold(originURL.Scheme, instanceURL.Scheme) && strings.EqualFold(originURL.Host, instanceURL.Host)
|
|
||||||
}
|
|
||||||
|
|
|
||||||
|
|
@ -49,6 +49,7 @@ func NewServer(ctx context.Context, profile *profile.Profile, store *store.Store
|
||||||
|
|
||||||
echoServer := echo.New()
|
echoServer := echo.New()
|
||||||
echoServer.Use(middleware.Recover())
|
echoServer.Use(middleware.Recover())
|
||||||
|
echoServer.Use(newCORSMiddleware(profile))
|
||||||
s.echoServer = echoServer
|
s.echoServer = echoServer
|
||||||
|
|
||||||
instanceBasicSetting, err := s.getOrUpsertInstanceBasicSetting(ctx)
|
instanceBasicSetting, err := s.getOrUpsertInstanceBasicSetting(ctx)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue