From bb76949fc0d9c954d1a03bc0f4dd7f35f2a5dcb2 Mon Sep 17 00:00:00 2001 From: boojack Date: Thu, 4 Jun 2026 22:37:41 +0800 Subject: [PATCH] chore(server): centralize CORS policy --- server/cors.go | 43 ++++++ server/cors_test.go | 122 ++++++++++++++++++ .../api/v1/connect_interceptors_test.go | 26 ---- server/router/api/v1/v1.go | 40 +----- server/server.go | 1 + 5 files changed, 167 insertions(+), 65 deletions(-) create mode 100644 server/cors.go create mode 100644 server/cors_test.go diff --git a/server/cors.go b/server/cors.go new file mode 100644 index 00000000..3eb3db7b --- /dev/null +++ b/server/cors.go @@ -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) +} diff --git a/server/cors_test.go b/server/cors_test.go new file mode 100644 index 00000000..fe9fd9ae --- /dev/null +++ b/server/cors_test.go @@ -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) + } + }) +} diff --git a/server/router/api/v1/connect_interceptors_test.go b/server/router/api/v1/connect_interceptors_test.go index 62925610..b4f8c79b 100644 --- a/server/router/api/v1/connect_interceptors_test.go +++ b/server/router/api/v1/connect_interceptors_test.go @@ -2,16 +2,11 @@ package v1 import ( "context" - "net/http" - "net/http/httptest" "testing" "connectrpc.com/connect" - "github.com/labstack/echo/v5" "google.golang.org/grpc/metadata" "google.golang.org/protobuf/types/known/emptypb" - - "github.com/usememos/memos/internal/profile" ) func TestMetadataInterceptorForwardsSecurityHeaders(t *testing.T) { @@ -42,24 +37,3 @@ func TestMetadataInterceptorForwardsSecurityHeaders(t *testing.T) { 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") - } -} diff --git a/server/router/api/v1/v1.go b/server/router/api/v1/v1.go index 17269da2..83c26fd2 100644 --- a/server/router/api/v1/v1.go +++ b/server/router/api/v1/v1.go @@ -3,13 +3,10 @@ package v1 import ( "context" "net/http" - "net/url" - "strings" "connectrpc.com/connect" "github.com/grpc-ecosystem/grpc-gateway/v2/runtime" "github.com/labstack/echo/v5" - "github.com/labstack/echo/v5/middleware" "golang.org/x/sync/semaphore" "github.com/usememos/memos/internal/markdown" @@ -127,9 +124,6 @@ func (s *APIV1Service) RegisterGateway(ctx context.Context, echoServer *echo.Ech return err } gwGroup := echoServer.Group("") - gwGroup.Use(middleware.CORSWithConfig(middleware.CORSConfig{ - AllowOrigins: []string{"*"}, - })) // Register SSE endpoint with same CORS as rest of /api/v1. RegisterSSERoutes(gwGroup, s.SSEHub, s.Store, s.Secret) 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.RegisterConnectHandlers(connectMux, connectInterceptors, connect.WithReadMaxBytes(maxAPIRequestBytes)) - // Wrap with CORS for browser access - 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 := echoServer.Group("") connectGroup.Any("/memos.api.v1.*", echo.WrapHandler(http.MaxBytesHandler(connectMux, maxAPIRequestBytes))) 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) -} diff --git a/server/server.go b/server/server.go index 8c7e47dd..dac5079a 100644 --- a/server/server.go +++ b/server/server.go @@ -49,6 +49,7 @@ func NewServer(ctx context.Context, profile *profile.Profile, store *store.Store echoServer := echo.New() echoServer.Use(middleware.Recover()) + echoServer.Use(newCORSMiddleware(profile)) s.echoServer = echoServer instanceBasicSetting, err := s.getOrUpsertInstanceBasicSetting(ctx)