Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 42 additions & 22 deletions internal/api/v1/common/middleware.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
}

Expand All @@ -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
Expand Down
132 changes: 132 additions & 0 deletions internal/api/v1/common/middleware_serviceaccount_test.go
Original file line number Diff line number Diff line change
@@ -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())
}
}
}
})
}
}