chore(server): centralize CORS policy

This commit is contained in:
boojack 2026-06-04 22:37:41 +08:00
parent d69f1aab27
commit bb76949fc0
5 changed files with 167 additions and 65 deletions

43
server/cors.go Normal file
View 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
View 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)
}
})
}

View file

@ -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")
}
}

View file

@ -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)
}

View file

@ -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)