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 (
|
||||
"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")
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in a new issue