From 9902f61b3a907625772b12408144d3b90c60d91a Mon Sep 17 00:00:00 2001 From: wangzhengzhuo05 <175673456+wangzhengzhuo05@users.noreply.github.com> Date: Sun, 13 Sep 2026 19:38:48 +0800 Subject: [PATCH] fix(api): accept service account keys on the Bearer auth path The Authorization: Bearer fallback in WithAuth only validated user API keys, so a service account key that works in X-API-Key got a 401 as a bearer token. Extract the service-account validation into a shared helper and fall through to it on the bearer path too. Fixes #253 --- internal/api/v1/common/middleware.go | 64 ++++++--- .../common/middleware_serviceaccount_test.go | 132 ++++++++++++++++++ 2 files changed, 174 insertions(+), 22 deletions(-) create mode 100644 internal/api/v1/common/middleware_serviceaccount_test.go diff --git a/internal/api/v1/common/middleware.go b/internal/api/v1/common/middleware.go index 09ae3724..2f9ac8b6 100644 --- a/internal/api/v1/common/middleware.go +++ b/internal/api/v1/common/middleware.go @@ -94,18 +94,8 @@ func WithAuth(userService user.Service, authService auth.Service, cfg *config.Co return } - if errors.Is(err, user.ErrInvalidAPIKey) && globalServiceAccountService != nil { - sa, saErr := globalServiceAccountService.ValidateAPIKey(r.Context(), apiKey) - if saErr == nil { - roleNames := make([]string, 0, len(sa.Roles)) - permKeys := make([]string, 0) - for _, r := range sa.Roles { - roleNames = append(roleNames, r.Name) - for _, p := range r.Permissions { - permKeys = append(permKeys, p.ResourceType+":"+p.Action) - } - } - principal := auth.NewServiceAccountPrincipal(sa.ID, sa.Name, roleNames, permKeys) + if errors.Is(err, user.ErrInvalidAPIKey) { + if principal, ok := serviceAccountPrincipal(r.Context(), apiKey); ok { ctx := setPrincipalContext(r.Context(), principal) next(w, r.WithContext(ctx)) return @@ -176,19 +166,28 @@ func WithAuth(userService user.Service, authService auth.Service, cfg *config.Co } } - // Fall back to API key in Bearer header + // Fall back to API key in Bearer header, then to a service-account key u, err := userService.ValidateAPIKey(r.Context(), tokenString) - if err != nil { - log.Error().Err(err). - Str("endpoint", r.URL.Path). - Str("method", r.Method). - Msg("Failed to validate bearer token as JWT or API key") - setWWWAuthenticate(w, cfg) - RespondError(w, http.StatusUnauthorized, "Invalid token") + if err == nil { + ctx := setPrincipalContext(r.Context(), auth.NewUserPrincipal(u)) + next(w, r.WithContext(ctx)) return } - ctx := setPrincipalContext(r.Context(), auth.NewUserPrincipal(u)) - next(w, r.WithContext(ctx)) + + if errors.Is(err, user.ErrInvalidAPIKey) { + if principal, ok := serviceAccountPrincipal(r.Context(), tokenString); ok { + ctx := setPrincipalContext(r.Context(), principal) + next(w, r.WithContext(ctx)) + return + } + } + + log.Error().Err(err). + Str("endpoint", r.URL.Path). + Str("method", r.Method). + Msg("Failed to validate bearer token as JWT or API key") + setWWWAuthenticate(w, cfg) + RespondError(w, http.StatusUnauthorized, "Invalid token") return } @@ -212,6 +211,27 @@ func WithAuth(userService user.Service, authService auth.Service, cfg *config.Co } } +// serviceAccountPrincipal validates apiKey as a service-account key and returns the +// matching principal. The bool reports whether the key was a valid service-account key. +func serviceAccountPrincipal(ctx context.Context, apiKey string) (auth.Principal, bool) { + if globalServiceAccountService == nil { + return nil, false + } + sa, err := globalServiceAccountService.ValidateAPIKey(ctx, apiKey) + if err != nil { + return nil, false + } + roleNames := make([]string, 0, len(sa.Roles)) + permKeys := make([]string, 0) + for _, r := range sa.Roles { + roleNames = append(roleNames, r.Name) + for _, p := range r.Permissions { + permKeys = append(permKeys, p.ResourceType+":"+p.Action) + } + } + return auth.NewServiceAccountPrincipal(sa.ID, sa.Name, roleNames, permKeys), true +} + func setPrincipalContext(ctx context.Context, p auth.Principal) context.Context { if p == nil { return ctx diff --git a/internal/api/v1/common/middleware_serviceaccount_test.go b/internal/api/v1/common/middleware_serviceaccount_test.go new file mode 100644 index 00000000..3ec73e0c --- /dev/null +++ b/internal/api/v1/common/middleware_serviceaccount_test.go @@ -0,0 +1,132 @@ +package common + +import ( + "context" + "errors" + "net/http" + "net/http/httptest" + "testing" + "time" + + "github.com/marmotdata/marmot/internal/core/auth" + "github.com/marmotdata/marmot/internal/core/role" + "github.com/marmotdata/marmot/internal/core/serviceaccount" + "github.com/marmotdata/marmot/internal/core/user" + "github.com/marmotdata/marmot/pkg/config" +) + +type mockServiceAccountService struct{} + +func (m *mockServiceAccountService) Create(_ context.Context, _ serviceaccount.CreateInput, _ *string) (*serviceaccount.ServiceAccount, error) { + return nil, nil +} + +func (m *mockServiceAccountService) Get(_ context.Context, _ string) (*serviceaccount.ServiceAccount, error) { + return nil, nil +} + +func (m *mockServiceAccountService) List(_ context.Context) ([]*serviceaccount.ServiceAccount, error) { + return nil, nil +} + +func (m *mockServiceAccountService) Update(_ context.Context, _ string, _ serviceaccount.UpdateInput) (*serviceaccount.ServiceAccount, error) { + return nil, nil +} + +func (m *mockServiceAccountService) Delete(_ context.Context, _ string) error { return nil } + +func (m *mockServiceAccountService) CreateAPIKey(_ context.Context, _ string, _ string, _ *time.Duration) (*serviceaccount.APIKey, error) { + return nil, nil +} + +func (m *mockServiceAccountService) ListAPIKeys(_ context.Context, _ string) ([]*serviceaccount.APIKey, error) { + return nil, nil +} + +func (m *mockServiceAccountService) DeleteAPIKey(_ context.Context, _, _ string) error { + return nil +} + +func (m *mockServiceAccountService) ValidateAPIKey(_ context.Context, apiKey string) (*serviceaccount.ServiceAccount, error) { + if apiKey == "sa-valid-key" { + return &serviceaccount.ServiceAccount{ + ID: "sa1", + Name: "ci-bot", + Active: true, + Roles: []*role.Role{{Name: "viewer"}}, + }, nil + } + return nil, errors.New("invalid service account key") +} + +func TestWithAuth_BearerServiceAccountKey(t *testing.T) { + old := globalServiceAccountService + defer func() { globalServiceAccountService = old }() + + userSvc := &mockUserService{ + validateAPIKeyFn: func(_ context.Context, key string) (*user.User, error) { + if key == "user-valid-key" { + return &user.User{ID: "user1", Username: "alice", Active: true}, nil + } + return nil, user.ErrInvalidAPIKey + }, + } + authSvc := &mockAuthService{} + cfg := &config.Config{} + + tests := []struct { + name string + bearerKey string + withSA bool + wantStatus int + }{ + {name: "service account key on bearer path", bearerKey: "sa-valid-key", withSA: true, wantStatus: http.StatusOK}, + {name: "service account key with no service registered", bearerKey: "sa-valid-key", withSA: false, wantStatus: http.StatusUnauthorized}, + {name: "unknown service account key", bearerKey: "sa-unknown-key", withSA: true, wantStatus: http.StatusUnauthorized}, + {name: "valid user key on bearer path", bearerKey: "user-valid-key", withSA: true, wantStatus: http.StatusOK}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.withSA { + globalServiceAccountService = &mockServiceAccountService{} + } else { + globalServiceAccountService = nil + } + + var captured auth.Principal + handler := WithAuth(userSvc, authSvc, cfg)(func(w http.ResponseWriter, r *http.Request) { + if p, ok := r.Context().Value(PrincipalContextKey).(auth.Principal); ok { + captured = p + } + w.WriteHeader(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodGet, "/api/v1/assets", nil) + req.Header.Set("Authorization", "Bearer "+tt.bearerKey) + rec := httptest.NewRecorder() + + handler(rec, req) + + if rec.Code != tt.wantStatus { + t.Fatalf("expected %d, got %d", tt.wantStatus, rec.Code) + } + + if tt.wantStatus == http.StatusOK { + if captured == nil { + t.Fatal("expected principal in context, got nil") + } + switch tt.bearerKey { + case "sa-valid-key": + if captured.Type() != auth.PrincipalTypeServiceAccount { + t.Fatalf("expected principal type %q, got %q", auth.PrincipalTypeServiceAccount, captured.Type()) + } + case "user-valid-key": + if captured.Type() != auth.PrincipalTypeUser { + t.Fatalf("expected principal type %q, got %q", auth.PrincipalTypeUser, captured.Type()) + } + } + } + }) + } +}