diff --git a/configs/config.example.yaml b/configs/config.example.yaml index a0f101d..cf53517 100644 --- a/configs/config.example.yaml +++ b/configs/config.example.yaml @@ -31,6 +31,7 @@ auth: - id: "oauth-client" client_id: "" client_secret: "" + superadmin: false description: "optional OAuth client credentials" database: diff --git a/internal/app/admin_identity.go b/internal/app/admin_identity.go index 0d6beb9..db4bca1 100644 --- a/internal/app/admin_identity.go +++ b/internal/app/admin_identity.go @@ -17,13 +17,21 @@ import ( ) type identityAdmin struct { - pool *pgxpool.Pool - keyring *auth.Keyring - logger *slog.Logger + pool *pgxpool.Pool + keyring *auth.Keyring + oauthRegistry *auth.OAuthRegistry + logger *slog.Logger } -func newIdentityAdmin(pool *pgxpool.Pool, keyring *auth.Keyring, logger *slog.Logger) *identityAdmin { - return &identityAdmin{pool: pool, keyring: keyring, logger: logger} +func newIdentityAdmin(pool *pgxpool.Pool, keyring *auth.Keyring, oauthRegistry *auth.OAuthRegistry, logger *slog.Logger) *identityAdmin { + return &identityAdmin{pool: pool, keyring: keyring, oauthRegistry: oauthRegistry, logger: logger} +} + +func (a *identityAdmin) isSuperadmin(keyID string) bool { + if a.keyring != nil && a.keyring.IsSuperadmin(keyID) { + return true + } + return a.oauthRegistry != nil && a.oauthRegistry.IsSuperadmin(keyID) } func loadIdentityKeyring(ctx context.Context, pool *pgxpool.Pool, keyring *auth.Keyring) error { @@ -71,7 +79,7 @@ func (a *identityAdmin) handler() http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.Method != http.MethodGet { keyID, ok := auth.KeyIDFromContext(r.Context()) - if !ok || a.keyring == nil || !a.keyring.IsSuperadmin(keyID) { + if !ok || !a.isSuperadmin(keyID) { writeJSON(w, http.StatusForbidden, map[string]string{"error": "superadmin API key required"}) return } diff --git a/internal/app/app.go b/internal/app/app.go index 96c6183..39467fe 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -200,7 +200,7 @@ func routes(logger *slog.Logger, cfg *config.Config, info buildinfo.Info, db *st enrichmentRetryer := tools.NewEnrichmentRetryer(context.Background(), db, bgMetadata, cfg.Capture, cfg.AI.Metadata.Timeout, activeProjects, logger) backfillTool := tools.NewBackfillTool(db, bgEmbeddings, activeProjects, logger) adminActions := newAdminActions(backfillTool, enrichmentRetryer, logger) - identityAdmin := newIdentityAdmin(db.Pool(), keyring, logger) + identityAdmin := newIdentityAdmin(db.Pool(), keyring, oauthRegistry, logger) toolSet := mcpserver.ToolSet{ Capture: tools.NewCaptureTool(db, embeddings, cfg.Capture, activeProjects, enrichmentRetryer, backfillTool), diff --git a/internal/auth/oauth_registry.go b/internal/auth/oauth_registry.go index e512bd4..1bc3a48 100644 --- a/internal/auth/oauth_registry.go +++ b/internal/auth/oauth_registry.go @@ -31,3 +31,18 @@ func (o *OAuthRegistry) Lookup(clientID string, clientSecret string) (string, bo } return "", false } + +// IsSuperadmin reports whether keyID (as returned by Lookup) belongs to an +// OAuth client configured with superadmin: true. +func (o *OAuthRegistry) IsSuperadmin(keyID string) bool { + for _, client := range o.clients { + id := client.ID + if id == "" { + id = client.ClientID + } + if id == keyID { + return client.Superadmin + } + } + return false +} diff --git a/internal/auth/oauth_registry_test.go b/internal/auth/oauth_registry_test.go index 7f8ba37..d4ea782 100644 --- a/internal/auth/oauth_registry_test.go +++ b/internal/auth/oauth_registry_test.go @@ -32,6 +32,26 @@ func TestNewOAuthRegistryAndLookup(t *testing.T) { } } +func TestOAuthRegistryIsSuperadmin(t *testing.T) { + registry, err := NewOAuthRegistry([]config.OAuthClient{ + {ID: "oauth-admin", ClientID: "admin-id", ClientSecret: "admin-secret", Superadmin: true}, + {ID: "oauth-client", ClientID: "client-id", ClientSecret: "client-secret"}, + }) + if err != nil { + t.Fatalf("NewOAuthRegistry() error = %v", err) + } + + if !registry.IsSuperadmin("oauth-admin") { + t.Fatal("IsSuperadmin(oauth-admin) = false, want true") + } + if registry.IsSuperadmin("oauth-client") { + t.Fatal("IsSuperadmin(oauth-client) = true, want false") + } + if registry.IsSuperadmin("unknown") { + t.Fatal("IsSuperadmin(unknown) = true, want false") + } +} + func TestMiddlewareAllowsOAuthBasicAuthAndSetsContext(t *testing.T) { oauthRegistry, err := NewOAuthRegistry([]config.OAuthClient{{ ID: "oauth-client", diff --git a/internal/config/config.go b/internal/config/config.go index 8f8dad0..21c918a 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -66,6 +66,7 @@ type OAuthClient struct { ID string `yaml:"id"` ClientID string `yaml:"client_id"` ClientSecret string `yaml:"client_secret"` + Superadmin bool `yaml:"superadmin"` Description string `yaml:"description"` }