diff --git a/services/core/IMPLEMENTATION.md b/services/core/IMPLEMENTATION.md index 13d442c73..da293e57c 100644 --- a/services/core/IMPLEMENTATION.md +++ b/services/core/IMPLEMENTATION.md @@ -24,6 +24,7 @@ Domain owners, each with its PostgreSQL adapter under `internal/persistence/post - `agents` (`agentpg`): saved Agents, their configuration merge and bounds, and the encrypted model-provider bundle bound to each Agent. - `files` (`filepg`): source Files. +- `vaults` (`vaultpg`): Vaults and Credentials, the encryption of Credential secrets, OAuth access-token refresh, and the MCP credential selection that Session creation freezes and the Dispatcher's `Credentials` resolves into a bearer token. ## Request handling @@ -113,7 +114,7 @@ Provider input validation uses the adapter rules in `internal/harnessconfig`: on ## Vaults and credentials -[Vaults and credentials](../../contracts/agents-api/vaults.md) describes the resources, selection rules, refresh and deletion. The store implements them under these rules: +[Vaults and credentials](../../contracts/agents-api/vaults.md) describes the resources, selection rules, refresh and deletion. `vaults` implements them, with `vaultpg` as its storage, under these rules: - Credentials are children of tenant-owned Vaults. Creation admits the owner in the same SQL statement as the insert; retrieval joins the owning Vault; listing enforces Project and Vault ownership on the parent, cursor and row query. Metadata queries never select ciphertext and need no encryption key. - Secret values are encrypted before they reach SQL, with Core's separately configured random 32-byte key and the standard library's random-nonce AES-GCM. The versioned authenticated binding covers tenant, Vault, Credential, authentication purpose and exact destination. Never reuse daemon transport encryption for this storage. A missing key disables credential writes with `credentialcrypto.ErrUnavailable`; a malformed configured key fails startup. @@ -124,7 +125,7 @@ Provider input validation uses the adapter rules in `internal/harnessconfig`: on - `credentialcrypto` ciphertext is a format version byte followed by the standard AEAD nonce, ciphertext and tag. The authenticated data holds a fixed domain and version plus the binding (tenant, Vault, Credential, auth type, exact destination). Keep the domain string unchanged: existing rows must still decrypt. - Random-nonce GCM allows at most 2^32 encryptions per key. `secrets/credential.key` also seals model providers, the E2B key, Skills, initial files and environment setup, so every sealed write counts toward that bound; there is no rotation or re-encryption path. -- OAuth dispatch refresh holds the Credential row lock and the external exchange under one 20-second context (`store.oauthRefreshTimeout`). The refresh HTTP client has a 10-second overall timeout and 5-second TLS handshake and response-header timeouts, uses no proxy and treats any redirect as failure. +- OAuth dispatch refresh holds the Credential row lock and the external exchange under one 20-second context (`vaults.oauthRefreshTimeout`). The refresh HTTP client has a 10-second overall timeout and 5-second TLS handshake and response-header timeouts, uses no proxy and treats any redirect as failure. ## MCP diff --git a/services/core/cmd/server/http_routes_test.go b/services/core/cmd/server/http_routes_test.go index 2fd157351..e6c968731 100644 --- a/services/core/cmd/server/http_routes_test.go +++ b/services/core/cmd/server/http_routes_test.go @@ -87,7 +87,8 @@ func daemonComposition(t testing.TB) http.Handler { } apiHandler, err := api.NewHandler(api.Dependencies{ Engine: "codex", CoreKeys: admin, InstallationBindings: struct{ api.InstallationBindings }{}, - Projects: trapProjects{keys: keys}, Vaults: struct{ api.Vaults }{}, ModelProviders: struct{ api.ModelProviders }{}, + Projects: trapProjects{keys: keys}, ModelProviders: struct{ api.ModelProviders }{}, + Vaults: struct{ api.Vaults }{}, VaultsReader: struct{ api.VaultsReader }{}, Files: struct{ api.Files }{}, FilesReader: struct{ api.FilesReader }{}, Skills: struct{ api.Skills }{}, EnvironmentTemplates: struct{ api.EnvironmentTemplates }{}, Agents: struct{ api.Agents }{}, AgentsReader: struct{ api.AgentsReader }{}, Sessions: struct{ api.Sessions }{}, SessionEvents: struct{ api.SessionEvents }{}, SessionHistory: struct{ api.SessionHistory }{}, Subagents: struct{ api.Subagents }{}, Artifacts: struct{ api.Artifacts }{}, diff --git a/services/core/cmd/server/main.go b/services/core/cmd/server/main.go index 6f40a95fb..9e4a32f48 100644 --- a/services/core/cmd/server/main.go +++ b/services/core/cmd/server/main.go @@ -42,6 +42,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/filepg" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/vaultpg" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtime" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimeenrollment" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimegateway" @@ -50,6 +51,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimeobs" observationstoreresolver "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimeobs/storeresolver" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/jackc/pgx/v5/pgxpool" ) @@ -109,7 +111,7 @@ func run() error { if err != nil { return err } - executionStore := store.NewWithCredentialCipherAndOAuthRefresh(pool, credentialKey, oauthClient) + executionStore := store.NewWithCredentialCipher(pool, credentialKey) executionStore.SetPublicURL(public) units := pgunit.NewPool(pool) auditStore := auditpg.New(units) @@ -118,6 +120,11 @@ func run() error { if err != nil { return err } + vaultStore := vaultpg.New(units) + vaultService, err := vaults.NewService(vaultStore, credentialKey, oauthClient) + if err != nil { + return err + } installation, err := installationFacts(public) if err != nil { return err @@ -237,7 +244,7 @@ func run() error { } } if registry != nil { - dispatcher := &execution.Dispatcher{Store: executionStore, Registry: registry, + dispatcher := &execution.Dispatcher{Store: executionStore, Registry: registry, Credentials: vaultService, ManagedRuntimes: managed, MaxConcurrentExecutions: concurrency} lease, err := pgunit.AcquireLease(ctx, pool) if err != nil { @@ -309,7 +316,8 @@ func run() error { deps := api.Dependencies{ Engine: engine, Harnesses: kinds, CoreKeys: keyAdmin, Installation: installation, InstallationBindings: executionStore, - Projects: executionStore, Vaults: executionStore, ModelProviders: executionStore, + Projects: executionStore, ModelProviders: executionStore, + Vaults: vaultService, VaultsReader: vaultStore, Skills: executionStore, EnvironmentTemplates: executionStore, Files: fileService, FilesReader: fileStore, Agents: agentService, AgentsReader: agentStore, diff --git a/services/core/internal/api/claude_mcp_test.go b/services/core/internal/api/claude_mcp_test.go index 640e34dfc..542f42e9d 100644 --- a/services/core/internal/api/claude_mcp_test.go +++ b/services/core/internal/api/claude_mcp_test.go @@ -6,25 +6,25 @@ import ( "strings" "testing" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) type claudeCredentialStore struct { - binding store.MCPCredentialBinding + binding vaults.MCPCredentialBinding calls int } -func (s *claudeCredentialStore) ResolveMCPCredentials(_ context.Context, _ string, _ []string, _ []store.MCPCredentialRequest) ([]store.MCPCredentialBinding, error) { +func (s *claudeCredentialStore) ResolveMCPCredentials(context.Context, vaults.ResolveMCPCredentials) ([]vaults.MCPCredentialBinding, error) { s.calls++ - return []store.MCPCredentialBinding{s.binding}, nil + return []vaults.MCPCredentialBinding{s.binding}, nil } func TestClaudeMCPAdmitsResolvedCredentials(t *testing.T) { for _, selection := range []string{"implicit", "explicit", "unmatched"} { t.Run(selection, func(t *testing.T) { vault, credential := uuid.NewString(), uuid.NewString() - s := &claudeCredentialStore{binding: store.MCPCredentialBinding{ServerLabel: "records", ServerURL: "https://mcp.example.test/tools", VaultID: vault, CredentialID: credential, AuthType: "static_bearer"}} + s := &claudeCredentialStore{binding: vaults.MCPCredentialBinding{ServerLabel: "records", ServerURL: "https://mcp.example.test/tools", VaultID: vault, CredentialID: credential, AuthType: "static_bearer"}} tool := publicMCP if selection == "explicit" { tool = strings.TrimSuffix(tool, "}") + `,"credential_id":"` + credential + `"}` diff --git a/services/core/internal/api/configuration.go b/services/core/internal/api/configuration.go index 386b8e5ba..7150241f3 100644 --- a/services/core/internal/api/configuration.go +++ b/services/core/internal/api/configuration.go @@ -6,15 +6,15 @@ import ( v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/metadata" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) type configuration struct { - Agent v1.Agent `json:"agent"` - Environment v1.Environment `json:"environment"` - VaultIDs []string `json:"vault_ids,omitempty"` - MCPCredentials []store.MCPCredentialBinding `json:"mcp_credentials,omitempty"` + Agent v1.Agent `json:"agent"` + Environment v1.Environment `json:"environment"` + VaultIDs []string `json:"vault_ids,omitempty"` + MCPCredentials []vaults.MCPCredentialBinding `json:"mcp_credentials,omitempty"` } func resolve(input sessionRequest, tenant, key string, saved *v1.SavedAgent) (json.RawMessage, error) { diff --git a/services/core/internal/api/credentials.go b/services/core/internal/api/credentials.go index 639f018e2..c90948ab9 100644 --- a/services/core/internal/api/credentials.go +++ b/services/core/internal/api/credentials.go @@ -5,7 +5,7 @@ import ( "net/http" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/go-chi/chi/v5" "github.com/google/uuid" ) @@ -41,7 +41,7 @@ func (h *Handler) createCredential(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusBadRequest, "invalid_request", err.Error()) return } - var credential store.Credential + var credential vaults.Credential switch credentialAuthType(request.Auth) { case "static_bearer": var auth v1.CredentialAuthInput @@ -49,20 +49,21 @@ func (h *Handler) createCredential(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusBadRequest, "invalid_request", "static_bearer requires a nonempty string token and an absolute HTTPS mcp_server_url without userinfo or a fragment.") return } - credential, err = h.Vaults.CreateStaticCredential(r.Context(), tenantID(r), vaultID, store.CreateStaticCredentialInput{Name: name, MCPServerURL: *auth.MCPServerURL, Token: *auth.Token}) + credential, err = h.Vaults.CreateStaticCredential(r.Context(), vaults.CreateStaticCredential{TenantID: tenantID(r), VaultID: vaultID, Name: name, MCPServerURL: *auth.MCPServerURL, Token: *auth.Token}) case "mcp_oauth": - input, parseErr := oauthCredentialCreate(request.Auth, name) + command, parseErr := oauthCredentialCreate(request.Auth, name) if parseErr != nil { - writeStoreError(w, r, parseErr) + writeVaultsError(w, r, parseErr) return } - credential, err = h.Vaults.CreateOAuthCredential(r.Context(), tenantID(r), vaultID, input) + command.TenantID, command.VaultID = tenantID(r), vaultID + credential, err = h.Vaults.CreateOAuthCredential(r.Context(), command) default: writeError(w, http.StatusBadRequest, "invalid_request", "auth requires type static_bearer or mcp_oauth.") return } if err != nil { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) return } writeJSON(w, http.StatusCreated, credentialResponse(credential)) @@ -88,9 +89,9 @@ func (h *Handler) getCredential(w http.ResponseWriter, r *http.Request) { if !ok { return } - credential, err := h.Vaults.GetCredential(r.Context(), tenantID(r), vaultID, id) + credential, err := h.VaultsReader.GetCredential(r.Context(), tenantID(r), vaultID, id) if err != nil { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) return } writeJSON(w, http.StatusOK, credentialResponse(credential)) @@ -101,13 +102,13 @@ func (h *Handler) getCredential(w http.ResponseWriter, r *http.Request) { func credentialResourceID(w http.ResponseWriter, r *http.Request, param string) (string, bool) { id, err := uuid.Parse(chi.URLParam(r, param)) if err != nil || id == uuid.Nil { - writeStoreError(w, r, store.ErrNotFound) + writeVaultsError(w, r, vaults.ErrNotFound) return "", false } return id.String(), true } -func credentialResponse(c store.Credential) v1.Credential { +func credentialResponse(c vaults.Credential) v1.Credential { auth := v1.CredentialAuth{Type: c.AuthType, MCPServerURL: c.MCPServerURL} if c.OAuth != nil { auth.ExpiresAt = c.OAuth.ExpiresAt diff --git a/services/core/internal/api/credentials_delete.go b/services/core/internal/api/credentials_delete.go index 758b26beb..27a44210a 100644 --- a/services/core/internal/api/credentials_delete.go +++ b/services/core/internal/api/credentials_delete.go @@ -5,6 +5,7 @@ import ( "net/http" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) // @Summary Delete a Vault Credential @@ -35,9 +36,9 @@ func (h *Handler) deleteCredential(w http.ResponseWriter, r *http.Request) { if !ok { return } - deleted, err := h.Vaults.DeleteCredential(r.Context(), tenantID(r), vaultID, id) + deleted, err := h.Vaults.DeleteCredential(r.Context(), vaults.DeleteCredential{TenantID: tenantID(r), VaultID: vaultID, CredentialID: id}) if err != nil { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) return } writeJSON(w, http.StatusOK, v1.CredentialDeleted{ID: deleted, Deleted: true, Object: "vault.credential.deleted"}) diff --git a/services/core/internal/api/credentials_delete_test.go b/services/core/internal/api/credentials_delete_test.go index 781cf7c73..4749304ad 100644 --- a/services/core/internal/api/credentials_delete_test.go +++ b/services/core/internal/api/credentials_delete_test.go @@ -10,12 +10,12 @@ import ( "strings" "testing" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) -func (f *credentialFixture) DeleteCredential(_ context.Context, tenant, vault, id string) (string, error) { - f.tenant, f.vault, f.id, f.calls = tenant, vault, id, f.calls+1 +func (f *credentialFixture) DeleteCredential(_ context.Context, command vaults.DeleteCredential) (string, error) { + f.tenant, f.vault, f.id, f.calls = command.TenantID, command.VaultID, command.CredentialID, f.calls+1 return f.credential.ID, f.err } @@ -74,7 +74,7 @@ func TestCredentialDeletionRejectsBeforeMutation(t *testing.T) { for _, tc := range []struct { err error status int - }{{store.ErrNotFound, 404}, {errors.New("credential-canary"), 500}} { + }{{vaults.ErrNotFound, 404}, {errors.New("credential-canary"), 500}} { h, f, _ := credentialHandler(t) f.err = tc.err w := credentialRequest(h, "DELETE", "/v1/vaults/"+f.credential.VaultID+"/credentials/"+f.credential.ID, "") diff --git a/services/core/internal/api/credentials_list.go b/services/core/internal/api/credentials_list.go index dbca67869..6f2f335df 100644 --- a/services/core/internal/api/credentials_list.go +++ b/services/core/internal/api/credentials_list.go @@ -24,13 +24,13 @@ import ( // @Router /vaults/{vault_id}/credentials [get] func (h *Handler) listCredentials(w http.ResponseWriter, r *http.Request) { vaultID := chi.URLParam(r, "vault_id") - options, statuses, ok := readVaultPage(w, r) + query, ok := readVaultPage(w, r) if !ok { return } - page, err := h.Vaults.ListCredentials(r.Context(), tenantID(r), vaultID, options.after, options.limit, options.ascending, statuses) + page, err := h.VaultsReader.ListCredentials(r.Context(), tenantID(r), vaultID, query) if err != nil { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) return } response := v1.CredentialList{Object: "list", Data: make([]v1.Credential, 0, len(page.Credentials)), HasMore: page.NextCursor != ""} diff --git a/services/core/internal/api/credentials_list_test.go b/services/core/internal/api/credentials_list_test.go index 5896245b8..78c15941f 100644 --- a/services/core/internal/api/credentials_list_test.go +++ b/services/core/internal/api/credentials_list_test.go @@ -7,12 +7,11 @@ import ( "testing" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) -func (f *credentialFixture) ListCredentials(_ context.Context, tenant, vault, after string, limit int, ascending bool, statuses []string) (store.CredentialPage, error) { - f.tenant, f.vault, f.calls = tenant, vault, f.calls+1 - f.options, f.statuses = pageOptions{after: after, limit: limit, ascending: ascending}, statuses +func (f *credentialFixture) ListCredentials(_ context.Context, tenant, vault string, query vaults.PageQuery) (vaults.CredentialPage, error) { + f.tenant, f.vault, f.query, f.calls = tenant, vault, query, f.calls+1 return f.page, f.err } @@ -29,13 +28,13 @@ func TestCredentialListScopeProjectionAndParameters(t *testing.T) { {"?limit=-3&tenant_id=foreign&unknown=1", 1, nil}, } { h, f, tenant := credentialHandler(t) - f.page = store.CredentialPage{Credentials: []store.Credential{f.credential}, NextCursor: f.credential.ID} + f.page = vaults.CredentialPage{Credentials: []vaults.Credential{f.credential}, NextCursor: f.credential.ID} w := credentialRequest(h, "GET", "/v1/vaults/"+f.credential.VaultID+"/credentials"+tc.query, "") var body v1.CredentialList if w.Code != 200 || json.Unmarshal(w.Body.Bytes(), &body) != nil { t.Fatal(w.Code, w.Body.String()) } - if f.calls != 1 || f.tenant != tenant || f.vault != f.credential.VaultID || f.options.limit != tc.limit || !reflect.DeepEqual(f.statuses, tc.statuses) { + if f.calls != 1 || f.tenant != tenant || f.vault != f.credential.VaultID || f.query.Limit != tc.limit || !reflect.DeepEqual(f.query.Statuses, tc.statuses) { t.Fatal("list scope or parameters changed") } if !reflect.DeepEqual(body.Data, []v1.Credential{credentialResponse(f.credential)}) || !body.HasMore || body.FirstID == nil || *body.FirstID != f.credential.ID || body.LastID == nil || *body.LastID != f.credential.ID { @@ -46,10 +45,10 @@ func TestCredentialListScopeProjectionAndParameters(t *testing.T) { path := "/v1/vaults/" + f.credential.VaultID + "/credentials" w := credentialRequest(h, "GET", path+"?order=asc&after="+f.credential.ID, "") var empty map[string]any - if w.Code != 200 || json.Unmarshal(w.Body.Bytes(), &empty) != nil || !reflect.DeepEqual(empty, map[string]any{"object": "list", "data": []any{}, "has_more": false, "first_id": nil, "last_id": nil}) || !f.options.ascending || f.options.after != f.credential.ID { + if w.Code != 200 || json.Unmarshal(w.Body.Bytes(), &empty) != nil || !reflect.DeepEqual(empty, map[string]any{"object": "list", "data": []any{}, "has_more": false, "first_id": nil, "last_id": nil}) || !f.query.Ascending || f.query.After != f.credential.ID { t.Fatal("empty page or cursor parsing changed") } - f.err = store.ErrNotFound + f.err = vaults.ErrNotFound if w = credentialRequest(h, "GET", path, ""); w.Code != 404 { t.Fatal("missing parent must not be an empty collection") } @@ -66,7 +65,7 @@ func TestCredentialListRejectsInvalidInputBeforeStorage(t *testing.T) { // A malformed parent reaches storage unchanged after query validation, and // storage reports it as a missing Vault. h, f, _ := credentialHandler(t) - f.err = store.ErrNotFound + f.err = vaults.ErrNotFound if w := credentialRequest(h, "GET", "/v1/vaults/invalid/credentials", ""); w.Code != 404 || f.vault != "invalid" { t.Fatal("invalid parent was not resolved as a missing Vault", w.Code, f.vault) } diff --git a/services/core/internal/api/credentials_oauth.go b/services/core/internal/api/credentials_oauth.go index 936ea119d..4e384f66a 100644 --- a/services/core/internal/api/credentials_oauth.go +++ b/services/core/internal/api/credentials_oauth.go @@ -7,7 +7,7 @@ import ( "time" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) func credentialHTTPSURL(value *string) bool { @@ -41,17 +41,19 @@ func credentialAuthType(raw json.RawMessage) string { return value } -func oauthCredentialCreate(raw json.RawMessage, name string) (store.CreateOAuthCredentialInput, error) { +// oauthCredentialCreate parses an mcp_oauth creation; the caller sets the +// tenant and Vault. +func oauthCredentialCreate(raw json.RawMessage, name string) (vaults.CreateOAuthCredential, error) { var auth v1.CredentialAuthInput - input := store.CreateOAuthCredentialInput{Name: name} + input := vaults.CreateOAuthCredential{Name: name} if decodeInputObject(raw, &auth, "type", "mcp_server_url", "access_token", "expires_at", "refresh") != nil || auth.Type != "mcp_oauth" || auth.AccessToken == nil || *auth.AccessToken == "" || !credentialHTTPSURL(auth.MCPServerURL) || !credentialExpiry(auth.ExpiresAt) { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } input.MCPServerURL, input.AccessToken = *auth.MCPServerURL, *auth.AccessToken input.OAuth.ExpiresAt = auth.ExpiresAt if r := auth.Refresh; r != nil { if r.ClientID == nil || r.RefreshToken == nil || !credentialHTTPSURL(r.TokenEndpoint) || r.TokenEndpointAuth == nil { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } a := r.TokenEndpointAuth switch a.Type { @@ -65,56 +67,58 @@ func oauthCredentialCreate(raw json.RawMessage, name string) (store.CreateOAuthC if json.Unmarshal(raw, &fields) != nil || decodeInputObject(fields.Refresh.Auth, &struct { Type string `json:"type"` }{}, "type") != nil { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } case "client_secret_basic", "client_secret_post": if a.ClientSecret == nil { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } input.ClientSecret = *a.ClientSecret default: - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } input.RefreshToken = *r.RefreshToken - input.OAuth.Refresh = &store.OAuthRefreshMetadata{ClientID: *r.ClientID, TokenEndpoint: *r.TokenEndpoint, TokenEndpointAuth: a.Type, Resource: r.Resource, Scope: r.Scope} + input.OAuth.Refresh = &vaults.OAuthRefreshMetadata{ClientID: *r.ClientID, TokenEndpoint: *r.TokenEndpoint, TokenEndpointAuth: a.Type, Resource: r.Resource, Scope: r.Scope} } return input, nil } -func oauthCredentialUpdate(raw json.RawMessage) (store.UpdateOAuthCredentialInput, error) { +// oauthCredentialUpdate parses an mcp_oauth patch; the caller sets the +// Credential's identity. +func oauthCredentialUpdate(raw json.RawMessage) (vaults.UpdateOAuthCredential, error) { var auth v1.CredentialAuthReplacement - var input store.UpdateOAuthCredentialInput + var input vaults.UpdateOAuthCredential if decodeInputObject(raw, &auth, "type", "access_token", "expires_at", "refresh") != nil || auth.Type != "mcp_oauth" { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } input.AccessToken = auth.AccessToken if input.AccessToken != nil && *input.AccessToken == "" { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } if len(auth.ExpiresAt) > 0 { input.ExpiresAtSet = true if json.Unmarshal(auth.ExpiresAt, &input.ExpiresAt) != nil || !credentialExpiry(input.ExpiresAt) { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } } if r := auth.Refresh; r != nil { - input.Refresh = &store.OAuthRefreshUpdate{RefreshToken: r.RefreshToken} + input.Refresh = &vaults.OAuthRefreshUpdate{RefreshToken: r.RefreshToken} if len(r.Scope) > 0 { input.Refresh.ScopeSet = true if json.Unmarshal(r.Scope, &input.Refresh.Scope) != nil { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } } if a := r.TokenEndpointAuth; a != nil { if a.Type != "client_secret_basic" && a.Type != "client_secret_post" { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } input.Refresh.TokenEndpointAuthType, input.Refresh.ClientSecret = a.Type, a.ClientSecret } } if input.AccessToken == nil && !input.ExpiresAtSet && (input.Refresh == nil || (input.Refresh.RefreshToken == nil && input.Refresh.ClientSecret == nil && !input.Refresh.ScopeSet)) { - return input, store.ErrInvalidInput + return input, vaults.ErrInvalidInput } return input, nil } diff --git a/services/core/internal/api/credentials_oauth_test.go b/services/core/internal/api/credentials_oauth_test.go index d11d4ecbd..779bbfc2a 100644 --- a/services/core/internal/api/credentials_oauth_test.go +++ b/services/core/internal/api/credentials_oauth_test.go @@ -9,17 +9,17 @@ import ( "testing" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) -func (f *credentialFixture) CreateOAuthCredential(_ context.Context, tenant, vault string, input store.CreateOAuthCredentialInput) (store.Credential, error) { - f.tenant, f.vault, f.oauthInput, f.calls = tenant, vault, input, f.calls+1 +func (f *credentialFixture) CreateOAuthCredential(_ context.Context, input vaults.CreateOAuthCredential) (vaults.Credential, error) { + f.tenant, f.vault, f.oauthInput, f.calls = input.TenantID, input.VaultID, input, f.calls+1 f.credential.Name, f.credential.MCPServerURL, f.credential.AuthType, f.credential.OAuth = input.Name, input.MCPServerURL, "mcp_oauth", &input.OAuth return f.credential, f.err } -func (f *credentialFixture) UpdateOAuthCredential(_ context.Context, tenant, vault, id string, input store.UpdateOAuthCredentialInput) (store.Credential, error) { - f.tenant, f.vault, f.id, f.oauthUpdate, f.calls = tenant, vault, id, input, f.calls+1 +func (f *credentialFixture) UpdateOAuthCredential(_ context.Context, input vaults.UpdateOAuthCredential) (vaults.Credential, error) { + f.tenant, f.vault, f.id, f.oauthUpdate, f.calls = input.TenantID, input.VaultID, input.CredentialID, input, f.calls+1 return f.credential, f.err } @@ -65,7 +65,7 @@ func TestOAuthCredentialVariantsAndSafeResourceReads(t *testing.T) { if read.Code != 200 || read.Body.String() != w.Body.String() { t.Fatal("OAuth retrieval changed safe metadata") } - f.page = store.CredentialPage{Credentials: []store.Credential{f.credential}} + f.page = vaults.CredentialPage{Credentials: []vaults.Credential{f.credential}} list := credentialRequest(h, "GET", path, "") var page struct { Data []map[string]any `json:"data"` @@ -107,18 +107,19 @@ func TestOAuthUpdateRetainsPresenceAndSecretPointers(t *testing.T) { text := func(s string) *string { return &s } for _, tc := range []struct { auth string - want store.UpdateOAuthCredentialInput + want vaults.UpdateOAuthCredential }{ - {`{"type":"mcp_oauth","access_token":" \t"}`, store.UpdateOAuthCredentialInput{AccessToken: text(" \t")}}, - {`{"type":"mcp_oauth","expires_at":null}`, store.UpdateOAuthCredentialInput{ExpiresAtSet: true}}, - {`{"type":"mcp_oauth","expires_at":"2026-09-22T12:30:00.123+08:00"}`, store.UpdateOAuthCredentialInput{ExpiresAtSet: true, ExpiresAt: text("2026-09-22T12:30:00.123+08:00")}}, - {`{"type":"mcp_oauth","refresh":{"scope":null}}`, store.UpdateOAuthCredentialInput{Refresh: &store.OAuthRefreshUpdate{ScopeSet: true}}}, - {`{"type":"mcp_oauth","refresh":{"scope":"","refresh_token":"refresh-canary","token_endpoint_auth":{"type":"client_secret_post","client_secret":"client-canary"}}}`, store.UpdateOAuthCredentialInput{Refresh: &store.OAuthRefreshUpdate{Scope: text(""), ScopeSet: true, RefreshToken: text("refresh-canary"), TokenEndpointAuthType: "client_secret_post", ClientSecret: text("client-canary")}}}, + {`{"type":"mcp_oauth","access_token":" \t"}`, vaults.UpdateOAuthCredential{AccessToken: text(" \t")}}, + {`{"type":"mcp_oauth","expires_at":null}`, vaults.UpdateOAuthCredential{ExpiresAtSet: true}}, + {`{"type":"mcp_oauth","expires_at":"2026-09-22T12:30:00.123+08:00"}`, vaults.UpdateOAuthCredential{ExpiresAtSet: true, ExpiresAt: text("2026-09-22T12:30:00.123+08:00")}}, + {`{"type":"mcp_oauth","refresh":{"scope":null}}`, vaults.UpdateOAuthCredential{Refresh: &vaults.OAuthRefreshUpdate{ScopeSet: true}}}, + {`{"type":"mcp_oauth","refresh":{"scope":"","refresh_token":"refresh-canary","token_endpoint_auth":{"type":"client_secret_post","client_secret":"client-canary"}}}`, vaults.UpdateOAuthCredential{Refresh: &vaults.OAuthRefreshUpdate{Scope: text(""), ScopeSet: true, RefreshToken: text("refresh-canary"), TokenEndpointAuthType: "client_secret_post", ClientSecret: text("client-canary")}}}, } { h, f, tenant := credentialHandler(t) - f.credential.AuthType, f.credential.OAuth = "mcp_oauth", &store.OAuthMetadata{} + f.credential.AuthType, f.credential.OAuth = "mcp_oauth", &vaults.OAuthMetadata{} w := credentialRequest(h, "POST", "/v1/vaults/"+f.credential.VaultID+"/credentials/"+f.credential.ID, `{"auth":`+tc.auth+`}`) - if w.Code != 200 || f.tenant != tenant || f.vault != f.credential.VaultID || f.id != f.credential.ID || !reflect.DeepEqual(f.oauthUpdate, tc.want) || strings.Contains(w.Body.String(), "canary") { + tc.want.TenantID, tc.want.VaultID, tc.want.CredentialID = tenant, f.credential.VaultID, f.credential.ID + if w.Code != 200 || !reflect.DeepEqual(f.oauthUpdate, tc.want) || strings.Contains(w.Body.String(), "canary") { t.Fatal("OAuth update lost scope or nullable field intent", tc.auth, w.Code) } } @@ -148,7 +149,7 @@ func TestOAuthCredentialStoreFailuresUseSafeExistingErrors(t *testing.T) { for _, tc := range []struct { err error code int - }{{store.ErrNotFound, 404}, {store.ErrInvalidInput, 400}, {credentialcrypto.ErrUnavailable, 503}, {errors.New("access-canary"), 500}} { + }{{vaults.ErrNotFound, 404}, {vaults.ErrInvalidInput, 400}, {credentialcrypto.ErrUnavailable, 503}, {errors.New("access-canary"), 500}} { for _, update := range []bool{false, true} { h, f, _ := credentialHandler(t) f.err = tc.err diff --git a/services/core/internal/api/credentials_test.go b/services/core/internal/api/credentials_test.go index d8224533a..0c2ace584 100644 --- a/services/core/internal/api/credentials_test.go +++ b/services/core/internal/api/credentials_test.go @@ -12,31 +12,30 @@ import ( "time" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) type credentialFixture struct { - credential store.Credential - input store.CreateStaticCredentialInput - replacement store.UpdateStaticCredentialInput - oauthInput store.CreateOAuthCredentialInput - oauthUpdate store.UpdateOAuthCredentialInput + credential vaults.Credential + input vaults.CreateStaticCredential + replacement vaults.UpdateStaticCredential + oauthInput vaults.CreateOAuthCredential + oauthUpdate vaults.UpdateOAuthCredential tenant, vault, id string calls int err error - page store.CredentialPage - options pageOptions - statuses []string + page vaults.CredentialPage + query vaults.PageQuery } -func (f *credentialFixture) CreateStaticCredential(_ context.Context, tenant, vault string, input store.CreateStaticCredentialInput) (store.Credential, error) { - f.tenant, f.vault, f.input, f.calls = tenant, vault, input, f.calls+1 +func (f *credentialFixture) CreateStaticCredential(_ context.Context, input vaults.CreateStaticCredential) (vaults.Credential, error) { + f.tenant, f.vault, f.input, f.calls = input.TenantID, input.VaultID, input, f.calls+1 f.credential.Name, f.credential.MCPServerURL = input.Name, input.MCPServerURL return f.credential, f.err } -func (f *credentialFixture) GetCredential(_ context.Context, tenant, vault, id string) (store.Credential, error) { +func (f *credentialFixture) GetCredential(_ context.Context, tenant, vault, id string) (vaults.Credential, error) { f.tenant, f.vault, f.id, f.calls = tenant, vault, id, f.calls+1 return f.credential, f.err } @@ -44,13 +43,14 @@ func (f *credentialFixture) GetCredential(_ context.Context, tenant, vault, id s // serve answers the Vault credential operations from f. func (f *credentialFixture) serve(_ *Dependencies, fakes *testFakes) { v := fakes.vaults - v.createStaticCredential, v.updateStaticCredential, v.getCredential, v.listCredentials, v.deleteCredential = f.CreateStaticCredential, f.UpdateStaticCredential, f.GetCredential, f.ListCredentials, f.DeleteCredential + v.createStaticCredential, v.updateStaticCredential, v.deleteCredential = f.CreateStaticCredential, f.UpdateStaticCredential, f.DeleteCredential v.createOAuthCredential, v.updateOAuthCredential = f.CreateOAuthCredential, f.UpdateOAuthCredential + fakes.vaultsReader.getCredential, fakes.vaultsReader.listCredentials = f.GetCredential, f.ListCredentials } func credentialHandler(t *testing.T) (http.Handler, *credentialFixture, string) { t.Helper() - f := &credentialFixture{credential: store.Credential{ID: uuid.NewString(), VaultID: uuid.NewString(), AuthType: "static_bearer", CreatedAt: time.Unix(1700000000, 0), UpdatedAt: time.Unix(1700000000, 0)}} + f := &credentialFixture{credential: vaults.Credential{ID: uuid.NewString(), VaultID: uuid.NewString(), AuthType: "static_bearer", CreatedAt: time.Unix(1700000000, 0), UpdatedAt: time.Unix(1700000000, 0)}} h, _, tenant := testHandler(t, f.serve) return h, f, tenant } @@ -113,7 +113,7 @@ func TestCredentialStorageErrorsStaySafe(t *testing.T) { for _, test := range []struct { err error status int - }{{store.ErrNotFound, 404}, {credentialcrypto.ErrUnavailable, 503}, {errors.New("credential-canary"), 500}} { + }{{vaults.ErrNotFound, 404}, {credentialcrypto.ErrUnavailable, 503}, {errors.New("credential-canary"), 500}} { h, f, _ := credentialHandler(t) f.err = test.err w := credentialRequest(h, "POST", "/v1/vaults/"+f.credential.VaultID+"/credentials", `{"name":"n","auth":{"type":"static_bearer","mcp_server_url":"https://example.invalid","token":"credential-canary"}}`) diff --git a/services/core/internal/api/credentials_update.go b/services/core/internal/api/credentials_update.go index 7b77aca7a..58323a2ab 100644 --- a/services/core/internal/api/credentials_update.go +++ b/services/core/internal/api/credentials_update.go @@ -5,7 +5,7 @@ import ( "net/http" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/go-chi/chi/v5" ) @@ -35,7 +35,7 @@ func (h *Handler) updateCredential(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusBadRequest, "invalid_request", "auth with a supported type is required.") return } - var credential store.Credential + var credential vaults.Credential var err error switch credentialAuthType(request.Auth) { case "static_bearer": @@ -44,20 +44,21 @@ func (h *Handler) updateCredential(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusBadRequest, "invalid_request", "static_bearer auth requires a nonempty string token.") return } - credential, err = h.Vaults.UpdateStaticCredential(r.Context(), tenantID(r), vaultID, id, store.UpdateStaticCredentialInput{Token: *auth.Token}) + credential, err = h.Vaults.UpdateStaticCredential(r.Context(), vaults.UpdateStaticCredential{TenantID: tenantID(r), VaultID: vaultID, CredentialID: id, Token: *auth.Token}) case "mcp_oauth": - input, parseErr := oauthCredentialUpdate(request.Auth) + command, parseErr := oauthCredentialUpdate(request.Auth) if parseErr != nil { - writeStoreError(w, r, parseErr) + writeVaultsError(w, r, parseErr) return } - credential, err = h.Vaults.UpdateOAuthCredential(r.Context(), tenantID(r), vaultID, id, input) + command.TenantID, command.VaultID, command.CredentialID = tenantID(r), vaultID, id + credential, err = h.Vaults.UpdateOAuthCredential(r.Context(), command) default: writeError(w, http.StatusBadRequest, "invalid_request", "auth requires type static_bearer or mcp_oauth.") return } if err != nil { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) return } writeJSON(w, http.StatusOK, credentialResponse(credential)) diff --git a/services/core/internal/api/credentials_update_test.go b/services/core/internal/api/credentials_update_test.go index 0d0ff6461..2cd883056 100644 --- a/services/core/internal/api/credentials_update_test.go +++ b/services/core/internal/api/credentials_update_test.go @@ -11,12 +11,12 @@ import ( "testing" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) -func (f *credentialFixture) UpdateStaticCredential(_ context.Context, tenant, vault, id string, input store.UpdateStaticCredentialInput) (store.Credential, error) { - f.tenant, f.vault, f.id, f.replacement, f.calls = tenant, vault, id, input, f.calls+1 +func (f *credentialFixture) UpdateStaticCredential(_ context.Context, input vaults.UpdateStaticCredential) (vaults.Credential, error) { + f.tenant, f.vault, f.id, f.replacement, f.calls = input.TenantID, input.VaultID, input.CredentialID, input, f.calls+1 return f.credential, f.err } @@ -70,11 +70,11 @@ func TestCredentialUpdateUsesExistingBoundariesAndSafeErrors(t *testing.T) { h, f, _ := credentialHandler(t) path := "/v1/vaults/" + f.credential.VaultID + "/credentials/" + f.credential.ID method, status := "POST", http.StatusBadRequest - // Malformed identifiers reach storage unchanged, after body validation, - // and storage reports them like a well-formed missing identifier. + // Malformed identifiers reach the use case unchanged, after body + // validation, and resolve exactly like a well-formed missing identifier. malformed := strings.Contains(mode, "invalid") || strings.Contains(mode, "zero") if malformed { - f.err = store.ErrNotFound + f.err = vaults.ErrNotFound } switch mode { case "method": @@ -107,7 +107,7 @@ func TestCredentialUpdateUsesExistingBoundariesAndSafeErrors(t *testing.T) { for _, tc := range []struct { err error status int - }{{store.ErrNotFound, 404}, {credentialcrypto.ErrUnavailable, 503}, {errors.New("credential-canary"), 500}} { + }{{vaults.ErrNotFound, 404}, {credentialcrypto.ErrUnavailable, 503}, {errors.New("credential-canary"), 500}} { h, f, _ := credentialHandler(t) f.err = tc.err w := credentialRequest(h, "POST", "/v1/vaults/"+f.credential.VaultID+"/credentials/"+f.credential.ID, body) diff --git a/services/core/internal/api/dependencies.go b/services/core/internal/api/dependencies.go index 9ea033253..6f1f3c08e 100644 --- a/services/core/internal/api/dependencies.go +++ b/services/core/internal/api/dependencies.go @@ -31,6 +31,7 @@ type Dependencies struct { Projects Projects Vaults Vaults + VaultsReader VaultsReader ModelProviders ModelProviders Files Files FilesReader FilesReader @@ -111,7 +112,8 @@ func (d Dependencies) validate() error { return errors.New("api: CoreKeys is required") } if err := required( - field{"InstallationBindings", d.InstallationBindings}, field{"Projects", d.Projects}, field{"Vaults", d.Vaults}, + field{"InstallationBindings", d.InstallationBindings}, field{"Projects", d.Projects}, + field{"Vaults", d.Vaults}, field{"VaultsReader", d.VaultsReader}, field{"ModelProviders", d.ModelProviders}, field{"Skills", d.Skills}, field{"Files", d.Files}, field{"FilesReader", d.FilesReader}, field{"Agents", d.Agents}, field{"AgentsReader", d.AgentsReader}, diff --git a/services/core/internal/api/dependencies_test.go b/services/core/internal/api/dependencies_test.go index 153c9e92b..8ff1a324a 100644 --- a/services/core/internal/api/dependencies_test.go +++ b/services/core/internal/api/dependencies_test.go @@ -17,6 +17,7 @@ const testExecutorURL = "wss://core.example/api/v1/agent-daemon/ws" type testFakes struct { projects *fakeProjects vaults *fakeVaults + vaultsReader *fakeVaultsReader modelProviders *fakeModelProviders files *fakeFiles filesReader *fakeFilesReader @@ -55,7 +56,8 @@ type testFakes struct { func testDependencies(t testing.TB) (Dependencies, *testFakes) { t.Helper() f := &testFakes{ - projects: &fakeProjects{t: t}, vaults: &fakeVaults{t: t}, modelProviders: &fakeModelProviders{t: t}, + projects: &fakeProjects{t: t}, modelProviders: &fakeModelProviders{t: t}, + vaults: &fakeVaults{t: t}, vaultsReader: &fakeVaultsReader{t: t}, skills: &fakeSkills{t: t}, environmentTemplates: &fakeEnvironmentTemplates{t: t}, files: &fakeFiles{t: t}, filesReader: &fakeFilesReader{t: t}, agents: &fakeAgents{t: t}, agentsReader: &fakeAgentsReader{t: t}, @@ -69,7 +71,8 @@ func testDependencies(t testing.TB) (Dependencies, *testFakes) { } return Dependencies{ Engine: "codex", CoreKeys: coreKeys(t, "admin"), InstallationBindings: f.installationBindings, - Projects: f.projects, Vaults: f.vaults, ModelProviders: f.modelProviders, Skills: f.skills, + Projects: f.projects, ModelProviders: f.modelProviders, Skills: f.skills, + Vaults: f.vaults, VaultsReader: f.vaultsReader, Files: f.files, FilesReader: f.filesReader, Agents: f.agents, AgentsReader: f.agentsReader, EnvironmentTemplates: f.environmentTemplates, Sessions: f.sessions, SessionEvents: f.sessionEvents, diff --git a/services/core/internal/api/errors.go b/services/core/internal/api/errors.go index 1c832d4fc..5f1ba4eab 100644 --- a/services/core/internal/api/errors.go +++ b/services/core/internal/api/errors.go @@ -116,7 +116,6 @@ func writeStoreError(w http.ResponseWriter, r *http.Request, err error, notFound var configuration *sandbox.ConfigurationError var unsupported *providercontract.UnsupportedError var cursor *store.InvalidCursorError - var selection *store.MCPCredentialSelectionError var sandboxConfiguration *store.SandboxConfigurationError var stale *store.SandboxGenerationStaleError var resetRequired *store.SandboxResetRequiredError @@ -214,14 +213,6 @@ func writeStoreError(w http.ResponseWriter, r *http.Request, err error, notFound } else { writeError(w, http.StatusBadRequest, "invalid_request_error", cursor.Message) } - case errors.As(err, &selection): - // Observed official fields for Session MCP credential selection (MV-03), - // all with a null param. - if selection.Conflict { - writeError(w, http.StatusConflict, "conflict_error", selection.Message) - } else { - writeError(w, http.StatusBadRequest, "invalid_request_error", selection.Message) - } case errors.Is(err, store.ErrNotFound): code := "not_found_error" // Skills retain their non-beta error envelope. diff --git a/services/core/internal/api/errors_vaults.go b/services/core/internal/api/errors_vaults.go new file mode 100644 index 000000000..430e2272d --- /dev/null +++ b/services/core/internal/api/errors_vaults.go @@ -0,0 +1,33 @@ +package api + +import ( + "errors" + "net/http" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +// writeVaultsError maps a Vault, Credential or Session credential selection +// error to its public error response. +func writeVaultsError(w http.ResponseWriter, r *http.Request, err error) { + if writeAuditSourceError(w, r, err) || writeTextValueError(w, r, err) || writeCredentialUnavailableError(w, r, err) { + return + } + var selection *vaults.MCPCredentialSelectionError + switch { + case errors.As(err, &selection): + // Observed official fields for Session MCP credential selection (MV-03), + // all with a null param. + if selection.Conflict { + writeError(w, http.StatusConflict, "conflict_error", selection.Message) + } else { + writeError(w, http.StatusBadRequest, "invalid_request_error", selection.Message) + } + case errors.Is(err, vaults.ErrNotFound): + writeError(w, http.StatusNotFound, "not_found_error", "Resource not found.") + case errors.Is(err, vaults.ErrInvalidInput): + writeError(w, http.StatusBadRequest, "invalid_request", invalidInputMessage) + default: + writeInternalError(w, r) + } +} diff --git a/services/core/internal/api/fakes_test.go b/services/core/internal/api/fakes_test.go index 180b892fc..759849c65 100644 --- a/services/core/internal/api/fakes_test.go +++ b/services/core/internal/api/fakes_test.go @@ -19,6 +19,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sandbox" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/writeaudit" ) @@ -1012,102 +1013,106 @@ func (f *fakeSubagents) ListSubagentTurnItems(a0 context.Context, a1 string, a2 type fakeVaults struct { t testing.TB - createVault func(context.Context, string, store.CreateVaultInput) (store.Vault, error) - getVault func(context.Context, string, string) (store.Vault, error) - deleteVault func(context.Context, string, string) (string, error) - listVaults func(context.Context, string, string, int, bool, []string) (store.VaultPage, error) - createOAuthCredential func(context.Context, string, string, store.CreateOAuthCredentialInput) (store.Credential, error) - updateOAuthCredential func(context.Context, string, string, string, store.UpdateOAuthCredentialInput) (store.Credential, error) - createStaticCredential func(context.Context, string, string, store.CreateStaticCredentialInput) (store.Credential, error) - updateStaticCredential func(context.Context, string, string, string, store.UpdateStaticCredentialInput) (store.Credential, error) - getCredential func(context.Context, string, string, string) (store.Credential, error) - deleteCredential func(context.Context, string, string, string) (string, error) - listCredentials func(context.Context, string, string, string, int, bool, []string) (store.CredentialPage, error) - resolveMCPCredentials func(context.Context, string, []string, []store.MCPCredentialRequest) ([]store.MCPCredentialBinding, error) -} - -func (f *fakeVaults) CreateVault(a0 context.Context, a1 string, a2 store.CreateVaultInput) (store.Vault, error) { + createVault func(context.Context, vaults.CreateVault) (vaults.Vault, error) + deleteVault func(context.Context, vaults.DeleteVault) (string, error) + createStaticCredential func(context.Context, vaults.CreateStaticCredential) (vaults.Credential, error) + updateStaticCredential func(context.Context, vaults.UpdateStaticCredential) (vaults.Credential, error) + createOAuthCredential func(context.Context, vaults.CreateOAuthCredential) (vaults.Credential, error) + updateOAuthCredential func(context.Context, vaults.UpdateOAuthCredential) (vaults.Credential, error) + deleteCredential func(context.Context, vaults.DeleteCredential) (string, error) + resolveMCPCredentials func(context.Context, vaults.ResolveMCPCredentials) ([]vaults.MCPCredentialBinding, error) +} + +func (f *fakeVaults) CreateVault(a0 context.Context, a1 vaults.CreateVault) (vaults.Vault, error) { if f.createVault == nil { unexpectedCall(f.t, "CreateVault") } - return f.createVault(a0, a1, a2) + return f.createVault(a0, a1) } -func (f *fakeVaults) GetVault(a0 context.Context, a1 string, a2 string) (store.Vault, error) { - if f.getVault == nil { - unexpectedCall(f.t, "GetVault") +func (f *fakeVaults) DeleteVault(a0 context.Context, a1 vaults.DeleteVault) (string, error) { + if f.deleteVault == nil { + unexpectedCall(f.t, "DeleteVault") } - return f.getVault(a0, a1, a2) + return f.deleteVault(a0, a1) } -func (f *fakeVaults) DeleteVault(a0 context.Context, a1 string, a2 string) (string, error) { - if f.deleteVault == nil { - unexpectedCall(f.t, "DeleteVault") +func (f *fakeVaults) CreateStaticCredential(a0 context.Context, a1 vaults.CreateStaticCredential) (vaults.Credential, error) { + if f.createStaticCredential == nil { + unexpectedCall(f.t, "CreateStaticCredential") } - return f.deleteVault(a0, a1, a2) + return f.createStaticCredential(a0, a1) } -func (f *fakeVaults) ListVaults(a0 context.Context, a1 string, a2 string, a3 int, a4 bool, a5 []string) (store.VaultPage, error) { - if f.listVaults == nil { - unexpectedCall(f.t, "ListVaults") +func (f *fakeVaults) UpdateStaticCredential(a0 context.Context, a1 vaults.UpdateStaticCredential) (vaults.Credential, error) { + if f.updateStaticCredential == nil { + unexpectedCall(f.t, "UpdateStaticCredential") } - return f.listVaults(a0, a1, a2, a3, a4, a5) + return f.updateStaticCredential(a0, a1) } -func (f *fakeVaults) CreateOAuthCredential(a0 context.Context, a1 string, a2 string, a3 store.CreateOAuthCredentialInput) (store.Credential, error) { +func (f *fakeVaults) CreateOAuthCredential(a0 context.Context, a1 vaults.CreateOAuthCredential) (vaults.Credential, error) { if f.createOAuthCredential == nil { unexpectedCall(f.t, "CreateOAuthCredential") } - return f.createOAuthCredential(a0, a1, a2, a3) + return f.createOAuthCredential(a0, a1) } -func (f *fakeVaults) UpdateOAuthCredential(a0 context.Context, a1 string, a2 string, a3 string, a4 store.UpdateOAuthCredentialInput) (store.Credential, error) { +func (f *fakeVaults) UpdateOAuthCredential(a0 context.Context, a1 vaults.UpdateOAuthCredential) (vaults.Credential, error) { if f.updateOAuthCredential == nil { unexpectedCall(f.t, "UpdateOAuthCredential") } - return f.updateOAuthCredential(a0, a1, a2, a3, a4) + return f.updateOAuthCredential(a0, a1) } -func (f *fakeVaults) CreateStaticCredential(a0 context.Context, a1 string, a2 string, a3 store.CreateStaticCredentialInput) (store.Credential, error) { - if f.createStaticCredential == nil { - unexpectedCall(f.t, "CreateStaticCredential") +func (f *fakeVaults) DeleteCredential(a0 context.Context, a1 vaults.DeleteCredential) (string, error) { + if f.deleteCredential == nil { + unexpectedCall(f.t, "DeleteCredential") } - return f.createStaticCredential(a0, a1, a2, a3) + return f.deleteCredential(a0, a1) } -func (f *fakeVaults) UpdateStaticCredential(a0 context.Context, a1 string, a2 string, a3 string, a4 store.UpdateStaticCredentialInput) (store.Credential, error) { - if f.updateStaticCredential == nil { - unexpectedCall(f.t, "UpdateStaticCredential") +func (f *fakeVaults) ResolveMCPCredentials(a0 context.Context, a1 vaults.ResolveMCPCredentials) ([]vaults.MCPCredentialBinding, error) { + if f.resolveMCPCredentials == nil { + unexpectedCall(f.t, "ResolveMCPCredentials") } - return f.updateStaticCredential(a0, a1, a2, a3, a4) + return f.resolveMCPCredentials(a0, a1) } -func (f *fakeVaults) GetCredential(a0 context.Context, a1 string, a2 string, a3 string) (store.Credential, error) { - if f.getCredential == nil { - unexpectedCall(f.t, "GetCredential") +type fakeVaultsReader struct { + t testing.TB + getVault func(context.Context, string, string) (vaults.Vault, error) + listVaults func(context.Context, string, vaults.PageQuery) (vaults.VaultPage, error) + getCredential func(context.Context, string, string, string) (vaults.Credential, error) + listCredentials func(context.Context, string, string, vaults.PageQuery) (vaults.CredentialPage, error) +} + +func (f *fakeVaultsReader) GetVault(a0 context.Context, a1 string, a2 string) (vaults.Vault, error) { + if f.getVault == nil { + unexpectedCall(f.t, "GetVault") } - return f.getCredential(a0, a1, a2, a3) + return f.getVault(a0, a1, a2) } -func (f *fakeVaults) DeleteCredential(a0 context.Context, a1 string, a2 string, a3 string) (string, error) { - if f.deleteCredential == nil { - unexpectedCall(f.t, "DeleteCredential") +func (f *fakeVaultsReader) ListVaults(a0 context.Context, a1 string, a2 vaults.PageQuery) (vaults.VaultPage, error) { + if f.listVaults == nil { + unexpectedCall(f.t, "ListVaults") } - return f.deleteCredential(a0, a1, a2, a3) + return f.listVaults(a0, a1, a2) } -func (f *fakeVaults) ListCredentials(a0 context.Context, a1 string, a2 string, a3 string, a4 int, a5 bool, a6 []string) (store.CredentialPage, error) { - if f.listCredentials == nil { - unexpectedCall(f.t, "ListCredentials") +func (f *fakeVaultsReader) GetCredential(a0 context.Context, a1 string, a2 string, a3 string) (vaults.Credential, error) { + if f.getCredential == nil { + unexpectedCall(f.t, "GetCredential") } - return f.listCredentials(a0, a1, a2, a3, a4, a5, a6) + return f.getCredential(a0, a1, a2, a3) } -func (f *fakeVaults) ResolveMCPCredentials(a0 context.Context, a1 string, a2 []string, a3 []store.MCPCredentialRequest) ([]store.MCPCredentialBinding, error) { - if f.resolveMCPCredentials == nil { - unexpectedCall(f.t, "ResolveMCPCredentials") +func (f *fakeVaultsReader) ListCredentials(a0 context.Context, a1 string, a2 string, a3 vaults.PageQuery) (vaults.CredentialPage, error) { + if f.listCredentials == nil { + unexpectedCall(f.t, "ListCredentials") } - return f.resolveMCPCredentials(a0, a1, a2, a3) + return f.listCredentials(a0, a1, a2, a3) } type fakeWriteAudit struct { diff --git a/services/core/internal/api/handler.go b/services/core/internal/api/handler.go index e6a4f3001..1f88dd830 100644 --- a/services/core/internal/api/handler.go +++ b/services/core/internal/api/handler.go @@ -196,7 +196,7 @@ func (h *Handler) createSession(w http.ResponseWriter, r *http.Request) { configuration, err = h.bindSessionCredentials(r.Context(), tenantID(r), configuration) if err != nil { if !h.recoverSessionCreation(w, r, key, creationRequest, input.Stream) { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) } return } diff --git a/services/core/internal/api/session_credentials.go b/services/core/internal/api/session_credentials.go index 954e4bfd9..ac03df820 100644 --- a/services/core/internal/api/session_credentials.go +++ b/services/core/internal/api/session_credentials.go @@ -6,31 +6,31 @@ import ( "encoding/json" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) func (h *Handler) bindSessionCredentials(ctx context.Context, tenant string, raw json.RawMessage) (json.RawMessage, error) { var cfg configuration if json.Unmarshal(raw, &cfg) != nil { - return nil, store.ErrInvalidInput + return nil, vaults.ErrInvalidInput } - var requests []store.MCPCredentialRequest + var requests []vaults.MCPCredentialRequest required := len(cfg.VaultIDs) > 0 for _, rawTool := range cfg.Agent.Tools { var tool v1.MCPTool if json.Unmarshal(rawTool, &tool) != nil { - return nil, store.ErrInvalidInput + return nil, vaults.ErrInvalidInput } if tool.Type == "mcp" { - requests = append(requests, store.MCPCredentialRequest{ServerLabel: tool.ServerLabel, ServerURL: tool.Transport.ServerURL, CredentialID: tool.CredentialID}) + requests = append(requests, vaults.MCPCredentialRequest{ServerLabel: tool.ServerLabel, ServerURL: tool.Transport.ServerURL, CredentialID: tool.CredentialID}) required = required || tool.CredentialID != nil } } if !required { return raw, nil } - bindings, err := h.Vaults.ResolveMCPCredentials(ctx, tenant, cfg.VaultIDs, requests) + bindings, err := h.Vaults.ResolveMCPCredentials(ctx, vaults.ResolveMCPCredentials{TenantID: tenant, VaultIDs: cfg.VaultIDs, Requests: requests}) if err != nil { return nil, err } diff --git a/services/core/internal/api/session_credentials_test.go b/services/core/internal/api/session_credentials_test.go index ff4107727..fde0ac9d7 100644 --- a/services/core/internal/api/session_credentials_test.go +++ b/services/core/internal/api/session_credentials_test.go @@ -8,6 +8,7 @@ import ( v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) @@ -74,8 +75,8 @@ func TestSessionProjectionShowsSelectedMCPCredential(t *testing.T) { encoded, _ := json.Marshal(value) return encoded } - binding := func(label, credentialID string) store.MCPCredentialBinding { - b := store.MCPCredentialBinding{ServerLabel: label, ServerURL: "https://mcp.example.test/" + label} + binding := func(label, credentialID string) vaults.MCPCredentialBinding { + b := vaults.MCPCredentialBinding{ServerLabel: label, ServerURL: "https://mcp.example.test/" + label} if credentialID != "" { b.VaultID, b.CredentialID, b.AuthType = vault, credentialID, "static_bearer" } @@ -85,7 +86,7 @@ func TestSessionProjectionShowsSelectedMCPCredential(t *testing.T) { explicit := uuid.NewString() cfg := configuration{Agent: v1.Agent{ID: "agent", Model: "model", Tools: []json.RawMessage{tool("implicit", ""), tool("anonymous", ""), tool("explicit", strings.ToUpper(explicit)), function}}, Environment: v1.Environment{Type: "none"}, VaultIDs: []string{strings.ToUpper(vault)}, - MCPCredentials: []store.MCPCredentialBinding{binding("implicit", credential), binding("anonymous", ""), binding("explicit", explicit)}} + MCPCredentials: []vaults.MCPCredentialBinding{binding("implicit", credential), binding("anonymous", ""), binding("explicit", explicit)}} // The stored caller intent is checked against PostgreSQL by the storedNull // guard in TestMCPCredentialSelectionPublicPostgres. raw, _ := json.Marshal(cfg) @@ -128,7 +129,7 @@ func TestSessionProjectionShowsSelectedMCPCredential(t *testing.T) { func(c *configuration) { c.MCPCredentials[0].ServerLabel = "other" }, } { changed := cfg - changed.MCPCredentials = append([]store.MCPCredentialBinding(nil), cfg.MCPCredentials...) + changed.MCPCredentials = append([]vaults.MCPCredentialBinding(nil), cfg.MCPCredentials...) change(&changed) raw, _ := json.Marshal(changed) response, err := sessionResponse(store.Session{Configuration: raw}, "") diff --git a/services/core/internal/api/validation_errors_test.go b/services/core/internal/api/validation_errors_test.go index 8c0410ba4..3b693261f 100644 --- a/services/core/internal/api/validation_errors_test.go +++ b/services/core/internal/api/validation_errors_test.go @@ -13,6 +13,7 @@ import ( v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/agents" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" "github.com/jackc/pgx/v5/pgconn" ) @@ -33,9 +34,9 @@ func (s *validationStore) UpdateAgent(_ context.Context, command agents.UpdateCo return agents.Agent{ID: command.AgentID, TenantID: command.TenantID, Configuration: json.RawMessage(`{"model":"validation-model"}`), Metadata: map[string]string{}}, nil } -func (s *validationStore) CreateVault(_ context.Context, tenant string, input store.CreateVaultInput) (store.Vault, error) { +func (s *validationStore) CreateVault(_ context.Context, input vaults.CreateVault) (vaults.Vault, error) { s.writes++ - return store.Vault{ID: uuid.NewString(), TenantID: tenant, Name: input.Name, Metadata: input.Metadata}, nil + return vaults.Vault{ID: uuid.NewString(), TenantID: input.TenantID, Name: input.Name, Metadata: input.Metadata}, nil } func (s *validationStore) UpdateSessionMetadata(_ context.Context, tenant, id string, metadata map[string]string) (store.Session, error) { diff --git a/services/core/internal/api/vault_pagination.go b/services/core/internal/api/vault_pagination.go index 08fbede36..5e46b3f8b 100644 --- a/services/core/internal/api/vault_pagination.go +++ b/services/core/internal/api/vault_pagination.go @@ -5,20 +5,22 @@ import ( "net/http" "slices" "strconv" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) -func readVaultPage(w http.ResponseWriter, r *http.Request) (pageOptions, []string, bool) { +func readVaultPage(w http.ResponseWriter, r *http.Request) (vaults.PageQuery, bool) { q := r.URL.Query() if len(q["status"]) > 1 { writeListDuplicateError(w, r, "status", nil) - return pageOptions{}, nil, false + return vaults.PageQuery{}, false } // A scalar status and status[] entries filter by their union. statuses := slices.Concat(q["status"], q["status[]"]) for _, status := range statuses { - if status != "active" && status != "archived" { + if status != vaults.StatusActive && status != vaults.StatusArchived { writeError(w, http.StatusBadRequest, "invalid_request_error", "Failed to deserialize query string: status: data did not match any variant of untagged enum VaultStatusFilterParam") - return pageOptions{}, nil, false + return vaults.PageQuery{}, false } } // Vault and Credential limits also clamp negative and overflowing integers; @@ -29,5 +31,5 @@ func readVaultPage(w http.ResponseWriter, r *http.Request) (pageOptions, []strin } } options, ok := readPageQuery(w, r, q, true) - return options, statuses, ok + return vaults.PageQuery{After: options.after, Limit: options.limit, Ascending: options.ascending, Statuses: statuses}, ok } diff --git a/services/core/internal/api/vault_pagination_test.go b/services/core/internal/api/vault_pagination_test.go index d743af7ad..24eec708c 100644 --- a/services/core/internal/api/vault_pagination_test.go +++ b/services/core/internal/api/vault_pagination_test.go @@ -13,7 +13,7 @@ func TestVaultStatusErrorEnvelopes(t *testing.T) { t.Run(path+"?"+query, func(t *testing.T) { w := httptest.NewRecorder() r := httptest.NewRequest(http.MethodGet, path+"?"+query, nil) - if _, _, ok := readVaultPage(w, r); ok { + if _, ok := readVaultPage(w, r); ok { t.Fatal("invalid status accepted") } assertListQueryError(t, w, "invalid_request_error", nil, "Failed to deserialize query string: status: data did not match any variant of untagged enum VaultStatusFilterParam") @@ -35,9 +35,9 @@ func TestVaultStatusUnionAndDuplicates(t *testing.T) { } { t.Run(path+"?"+test.query, func(t *testing.T) { w := httptest.NewRecorder() - _, statuses, ok := readVaultPage(w, httptest.NewRequest(http.MethodGet, path+"?"+test.query, nil)) - if !ok || !reflect.DeepEqual(statuses, test.statuses) || w.Body.Len() != 0 { - t.Fatalf("statuses=%v ok=%t response=%s", statuses, ok, w.Body.String()) + query, ok := readVaultPage(w, httptest.NewRequest(http.MethodGet, path+"?"+test.query, nil)) + if !ok || !reflect.DeepEqual(query.Statuses, test.statuses) || w.Body.Len() != 0 { + t.Fatalf("statuses=%v ok=%t response=%s", query.Statuses, ok, w.Body.String()) } }) } @@ -46,7 +46,7 @@ func TestVaultStatusUnionAndDuplicates(t *testing.T) { w := httptest.NewRecorder() values := map[string]string{"status": "active", "limit": "5", "after": "x", "order": "asc"} query := key + "=" + values[key] + "&" + key + "=" + values[key] - if _, _, ok := readVaultPage(w, httptest.NewRequest(http.MethodGet, path+"?"+query, nil)); ok { + if _, ok := readVaultPage(w, httptest.NewRequest(http.MethodGet, path+"?"+query, nil)); ok { t.Fatal("repeated key accepted") } assertListQueryError(t, w, "invalid_request_error", nil, "Failed to deserialize query string: duplicate field `"+key+"`") @@ -63,8 +63,8 @@ func TestVaultLimitClamp(t *testing.T) { } { t.Run(path+"?"+query, func(t *testing.T) { w := httptest.NewRecorder() - page, _, ok := readVaultPage(w, httptest.NewRequest(http.MethodGet, path+"?"+query, nil)) - if !ok || page.limit != limit || w.Body.Len() != 0 { + page, ok := readVaultPage(w, httptest.NewRequest(http.MethodGet, path+"?"+query, nil)) + if !ok || page.Limit != limit || w.Body.Len() != 0 { t.Fatalf("page=%+v ok=%t response=%s", page, ok, w.Body.String()) } }) @@ -72,7 +72,7 @@ func TestVaultLimitClamp(t *testing.T) { for _, query := range []string{"limit=abc", "limit=1.5", "limit=", "limit=null"} { t.Run(path+"?"+query, func(t *testing.T) { w := httptest.NewRecorder() - if _, _, ok := readVaultPage(w, httptest.NewRequest(http.MethodGet, path+"?"+query, nil)); ok { + if _, ok := readVaultPage(w, httptest.NewRequest(http.MethodGet, path+"?"+query, nil)); ok { t.Fatal("non-integer limit accepted") } assertListQueryError(t, w, "invalid_request_error", nil, invalidDigit) diff --git a/services/core/internal/api/vaults.go b/services/core/internal/api/vaults.go index 4bef9348a..bd2eeb14a 100644 --- a/services/core/internal/api/vaults.go +++ b/services/core/internal/api/vaults.go @@ -9,26 +9,28 @@ import ( v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/metadata" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" - "github.com/go-chi/chi/v5" - "github.com/google/uuid" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) -// Vaults manages Vaults and their Credentials, and selects the Credentials a -// Session's MCP servers use at creation. +// Vaults runs the Vault and Credential use cases, and selects the Credentials +// a Session's MCP servers use at creation. type Vaults interface { - CreateVault(context.Context, string, store.CreateVaultInput) (store.Vault, error) - GetVault(context.Context, string, string) (store.Vault, error) - DeleteVault(context.Context, string, string) (string, error) - ListVaults(context.Context, string, string, int, bool, []string) (store.VaultPage, error) - CreateOAuthCredential(context.Context, string, string, store.CreateOAuthCredentialInput) (store.Credential, error) - UpdateOAuthCredential(context.Context, string, string, string, store.UpdateOAuthCredentialInput) (store.Credential, error) - CreateStaticCredential(context.Context, string, string, store.CreateStaticCredentialInput) (store.Credential, error) - UpdateStaticCredential(context.Context, string, string, string, store.UpdateStaticCredentialInput) (store.Credential, error) - GetCredential(context.Context, string, string, string) (store.Credential, error) - DeleteCredential(context.Context, string, string, string) (string, error) - ListCredentials(context.Context, string, string, string, int, bool, []string) (store.CredentialPage, error) - ResolveMCPCredentials(context.Context, string, []string, []store.MCPCredentialRequest) ([]store.MCPCredentialBinding, error) + CreateVault(context.Context, vaults.CreateVault) (vaults.Vault, error) + DeleteVault(context.Context, vaults.DeleteVault) (string, error) + CreateStaticCredential(context.Context, vaults.CreateStaticCredential) (vaults.Credential, error) + UpdateStaticCredential(context.Context, vaults.UpdateStaticCredential) (vaults.Credential, error) + CreateOAuthCredential(context.Context, vaults.CreateOAuthCredential) (vaults.Credential, error) + UpdateOAuthCredential(context.Context, vaults.UpdateOAuthCredential) (vaults.Credential, error) + DeleteCredential(context.Context, vaults.DeleteCredential) (string, error) + ResolveMCPCredentials(context.Context, vaults.ResolveMCPCredentials) ([]vaults.MCPCredentialBinding, error) +} + +// VaultsReader reads Vaults and the public metadata of their Credentials. +type VaultsReader interface { + GetVault(ctx context.Context, tenantID, vaultID string) (vaults.Vault, error) + ListVaults(ctx context.Context, tenantID string, query vaults.PageQuery) (vaults.VaultPage, error) + GetCredential(ctx context.Context, tenantID, vaultID, credentialID string) (vaults.Credential, error) + ListCredentials(ctx context.Context, tenantID, vaultID string, query vaults.PageQuery) (vaults.CredentialPage, error) } // @Summary Create a Vault @@ -58,7 +60,7 @@ func (h *Handler) createVault(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusBadRequest, "invalid_request", "Request must be a JSON object containing supported fields.") return } - input := store.CreateVaultInput{} + input := vaults.CreateVault{TenantID: tenantID(r)} if len(request.Name) > 0 { var name *string if json.Unmarshal(request.Name, &name) != nil || name == nil { @@ -84,9 +86,9 @@ func (h *Handler) createVault(w http.ResponseWriter, r *http.Request) { } return } - vault, err := h.Vaults.CreateVault(r.Context(), tenantID(r), input) + vault, err := h.Vaults.CreateVault(r.Context(), input) if err != nil { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) return } writeJSON(w, http.StatusCreated, vaultResponse(vault)) @@ -103,20 +105,19 @@ func (h *Handler) createVault(w http.ResponseWriter, r *http.Request) { // @Failure 400,401,404,500 {object} v1.ErrorResponse // @Router /vaults/{vault_id} [get] func (h *Handler) getVault(w http.ResponseWriter, r *http.Request) { - id := chi.URLParam(r, "vault_id") - if parsed, err := uuid.Parse(id); err != nil || parsed == uuid.Nil { - writeStoreError(w, r, store.ErrNotFound) + id, ok := credentialResourceID(w, r, "vault_id") + if !ok { return } - vault, err := h.Vaults.GetVault(r.Context(), tenantID(r), id) + vault, err := h.VaultsReader.GetVault(r.Context(), tenantID(r), id) if err != nil { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) return } writeJSON(w, http.StatusOK, vaultResponse(vault)) } -func vaultResponse(vault store.Vault) v1.Vault { +func vaultResponse(vault vaults.Vault) v1.Vault { return v1.Vault{ID: vault.ID, Object: "vault", CreatedAt: vault.CreatedAt.Unix(), Name: vault.Name, Metadata: vault.Metadata} } diff --git a/services/core/internal/api/vaults_delete.go b/services/core/internal/api/vaults_delete.go index 0d9274bf2..3e666bf80 100644 --- a/services/core/internal/api/vaults_delete.go +++ b/services/core/internal/api/vaults_delete.go @@ -5,6 +5,7 @@ import ( "net/http" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) // @Summary Delete a Vault and all its Credentials @@ -30,9 +31,9 @@ func (h *Handler) deleteVault(w http.ResponseWriter, r *http.Request) { if !ok { return } - deleted, err := h.Vaults.DeleteVault(r.Context(), tenantID(r), id) + deleted, err := h.Vaults.DeleteVault(r.Context(), vaults.DeleteVault{TenantID: tenantID(r), VaultID: id}) if err != nil { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) return } writeJSON(w, http.StatusOK, v1.VaultDeleted{ID: deleted, Deleted: true, Object: "vault.deleted"}) diff --git a/services/core/internal/api/vaults_delete_test.go b/services/core/internal/api/vaults_delete_test.go index 57e2e21e4..1f6c85853 100644 --- a/services/core/internal/api/vaults_delete_test.go +++ b/services/core/internal/api/vaults_delete_test.go @@ -9,12 +9,12 @@ import ( "strings" "testing" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) -func (f *vaultResourceFixture) DeleteVault(_ context.Context, tenant, id string) (string, error) { - f.tenant, f.id, f.calls = tenant, id, f.calls+1 +func (f *vaultResourceFixture) DeleteVault(_ context.Context, command vaults.DeleteVault) (string, error) { + f.tenant, f.id, f.calls = command.TenantID, command.VaultID, f.calls+1 return f.vault.ID, f.err } @@ -70,7 +70,7 @@ func TestVaultDeletionRejectsBeforeMutation(t *testing.T) { for _, tc := range []struct { err error status int - }{{store.ErrNotFound, 404}, {errors.New("vault-delete-canary"), 500}} { + }{{vaults.ErrNotFound, 404}, {errors.New("vault-delete-canary"), 500}} { h, f := vaultResourceHandler(t) f.err = tc.err w := vaultRequest(h, "DELETE", "/v1/vaults/"+f.vault.ID, "") diff --git a/services/core/internal/api/vaults_list.go b/services/core/internal/api/vaults_list.go index ab565ca4b..72ebed7da 100644 --- a/services/core/internal/api/vaults_list.go +++ b/services/core/internal/api/vaults_list.go @@ -21,13 +21,13 @@ import ( // @Failure 400,401,404,500 {object} v1.ErrorResponse // @Router /vaults [get] func (h *Handler) listVaults(w http.ResponseWriter, r *http.Request) { - options, statuses, ok := readVaultPage(w, r) + query, ok := readVaultPage(w, r) if !ok { return } - page, err := h.Vaults.ListVaults(r.Context(), tenantID(r), options.after, options.limit, options.ascending, statuses) + page, err := h.VaultsReader.ListVaults(r.Context(), tenantID(r), query) if err != nil { - writeStoreError(w, r, err) + writeVaultsError(w, r, err) return } response := v1.VaultList{Object: "list", Data: make([]v1.Vault, 0, len(page.Vaults)), HasMore: page.NextCursor != ""} diff --git a/services/core/internal/api/vaults_list_test.go b/services/core/internal/api/vaults_list_test.go index 0c8d2bc88..302504fd5 100644 --- a/services/core/internal/api/vaults_list_test.go +++ b/services/core/internal/api/vaults_list_test.go @@ -8,12 +8,11 @@ import ( "testing" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) -func (f *vaultResourceFixture) ListVaults(_ context.Context, tenant, after string, limit int, ascending bool, statuses []string) (store.VaultPage, error) { - f.tenant, f.calls = tenant, f.calls+1 - f.options, f.statuses = pageOptions{after: after, limit: limit, ascending: ascending}, statuses +func (f *vaultResourceFixture) ListVaults(_ context.Context, tenant string, query vaults.PageQuery) (vaults.VaultPage, error) { + f.tenant, f.query, f.calls = tenant, query, f.calls+1 return f.page, f.err } @@ -34,8 +33,8 @@ func TestVaultListParameters(t *testing.T) { t.Run(tc.query, func(t *testing.T) { h, f := vaultResourceHandler(t) w := vaultRequest(h, http.MethodGet, "/v1/vaults"+tc.query, "") - if w.Code != 200 || f.calls != 1 || f.tenant != f.vault.TenantID || f.options.limit != tc.limit || !reflect.DeepEqual(f.statuses, tc.statuses) { - t.Fatalf("list: %d %s; options %+v statuses %v owner %s", w.Code, w.Body.String(), f.options, f.statuses, f.tenant) + if w.Code != 200 || f.calls != 1 || f.tenant != f.vault.TenantID || f.query.Limit != tc.limit || !reflect.DeepEqual(f.query.Statuses, tc.statuses) { + t.Fatalf("list: %d %s; query %+v owner %s", w.Code, w.Body.String(), f.query, f.tenant) } var body map[string]any if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil { @@ -58,16 +57,16 @@ func TestVaultListParameters(t *testing.T) { func TestVaultListSafeProjectionAndCursor(t *testing.T) { h, f := vaultResourceHandler(t) - f.page = store.VaultPage{Vaults: []store.Vault{f.vault}, NextCursor: f.vault.ID} + f.page = vaults.VaultPage{Vaults: []vaults.Vault{f.vault}, NextCursor: f.vault.ID} w := vaultRequest(h, http.MethodGet, "/v1/vaults?order=asc&after="+f.vault.ID, "") var body v1.VaultList if err := json.Unmarshal(w.Body.Bytes(), &body); err != nil || w.Code != 200 { t.Fatal(w.Code, w.Body.String(), err) } - if !reflect.DeepEqual(body.Data, []v1.Vault{vaultResponse(f.vault)}) || !body.HasMore || body.FirstID == nil || *body.FirstID != f.vault.ID || body.LastID == nil || *body.LastID != f.vault.ID || !f.options.ascending || f.options.after != f.vault.ID { - t.Fatalf("page changed: %+v, %+v", body, f.options) + if !reflect.DeepEqual(body.Data, []v1.Vault{vaultResponse(f.vault)}) || !body.HasMore || body.FirstID == nil || *body.FirstID != f.vault.ID || body.LastID == nil || *body.LastID != f.vault.ID || !f.query.Ascending || f.query.After != f.vault.ID { + t.Fatalf("page changed: %+v, %+v", body, f.query) } - f.err = store.ErrNotFound + f.err = vaults.ErrNotFound if w = vaultRequest(h, http.MethodGet, "/v1/vaults?after="+f.vault.ID, ""); w.Code != 404 { t.Fatal(w.Code, w.Body.String()) } diff --git a/services/core/internal/api/vaults_test.go b/services/core/internal/api/vaults_test.go index 3b84c48d1..e52faa0dc 100644 --- a/services/core/internal/api/vaults_test.go +++ b/services/core/internal/api/vaults_test.go @@ -13,27 +13,26 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/agents" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimedevice" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) type vaultResourceFixture struct { - vault store.Vault + vault vaults.Vault err error tenant, id string calls int - page store.VaultPage - options pageOptions - statuses []string + page vaults.VaultPage + query vaults.PageQuery } -func (f *vaultResourceFixture) CreateVault(_ context.Context, tenant string, input store.CreateVaultInput) (store.Vault, error) { - f.tenant, f.calls = tenant, f.calls+1 +func (f *vaultResourceFixture) CreateVault(_ context.Context, input vaults.CreateVault) (vaults.Vault, error) { + f.tenant, f.calls = input.TenantID, f.calls+1 f.vault.Name, f.vault.Metadata = input.Name, input.Metadata return f.vault, f.err } -func (f *vaultResourceFixture) GetVault(_ context.Context, tenant, id string) (store.Vault, error) { +func (f *vaultResourceFixture) GetVault(_ context.Context, tenant, id string) (vaults.Vault, error) { f.tenant, f.id, f.calls = tenant, id, f.calls+1 return f.vault, f.err } @@ -46,11 +45,12 @@ func (f *vaultResourceFixture) UpdateAgent(context.Context, agents.UpdateCommand func vaultResourceHandler(t *testing.T) (http.Handler, *vaultResourceFixture) { t.Helper() - f := &vaultResourceFixture{vault: store.Vault{ID: uuid.NewString(), TenantID: uuid.NewString(), Metadata: map[string]string{}, CreatedAt: time.Unix(1700000000, 0)}} + f := &vaultResourceFixture{vault: vaults.Vault{ID: uuid.NewString(), TenantID: uuid.NewString(), Metadata: map[string]string{}, CreatedAt: time.Unix(1700000000, 0)}} deps, fakes := testDependencies(t) deps.Engine = "fake_alpha" fakes.projects.resolveProjectAPIKey = projectKeys(t, APIKey{OrganizationID: "vault-org", ProjectID: "vault-project", SubjectKind: "user", SubjectID: "vault-owner", TokenSHA256: runtimedevice.HashCredential("vault-key"), TenantID: f.vault.TenantID}).ResolveProjectAPIKey - fakes.vaults.createVault, fakes.vaults.getVault, fakes.vaults.listVaults, fakes.vaults.deleteVault = f.CreateVault, f.GetVault, f.ListVaults, f.DeleteVault + fakes.vaults.createVault, fakes.vaults.deleteVault = f.CreateVault, f.DeleteVault + fakes.vaultsReader.getVault, fakes.vaultsReader.listVaults = f.GetVault, f.ListVaults fakes.agents.update = f.UpdateAgent return newTestHandler(t, deps), f } @@ -141,7 +141,7 @@ func TestVaultResourceIgnoresUnknownQueryKeys(t *testing.T) { t.Fatal(test.method, w.Code, f.calls, f.tenant) } // A missing or foreign Vault stays indistinguishable with the same query. - f.err = store.ErrNotFound + f.err = vaults.ErrNotFound missing := vaultRequest(h, "GET", "/v1/vaults/"+uuid.NewString()+"?tenant_id=foreign", "") plain := vaultRequest(h, "GET", "/v1/vaults/"+uuid.NewString(), "") if missing.Code != 404 || missing.Body.String() != plain.Body.String() || f.tenant != f.vault.TenantID { @@ -174,7 +174,7 @@ func TestVaultResourceUsesSharedAuthenticationAndErrors(t *testing.T) { for _, test := range []struct { err error status int - }{{store.ErrNotFound, 404}, {errors.New("private-vault-backend"), 500}} { + }{{vaults.ErrNotFound, 404}, {errors.New("private-vault-backend"), 500}} { h, f := vaultResourceHandler(t) f.err = test.err w := vaultRequest(h, "GET", "/v1/vaults/"+f.vault.ID, "") diff --git a/services/core/internal/execution/dispatcher.go b/services/core/internal/execution/dispatcher.go index 392fefeac..6bf2faa84 100644 --- a/services/core/internal/execution/dispatcher.go +++ b/services/core/internal/execution/dispatcher.go @@ -13,15 +13,16 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimegateway" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) // Snapshot is the Session configuration frozen at creation. type Snapshot struct { - ModelProviderConfigured bool `json:"model_provider_configured,omitempty"` - Agent v1.Agent `json:"agent"` - Environment *v1.Environment `json:"environment"` - VaultIDs []string `json:"vault_ids,omitempty"` - MCPCredentials []store.MCPCredentialBinding `json:"mcp_credentials,omitempty"` + ModelProviderConfigured bool `json:"model_provider_configured,omitempty"` + Agent v1.Agent `json:"agent"` + Environment *v1.Environment `json:"environment"` + VaultIDs []string `json:"vault_ids,omitempty"` + MCPCredentials []vaults.MCPCredentialBinding `json:"mcp_credentials,omitempty"` } type Dispatcher struct { @@ -29,6 +30,8 @@ type Dispatcher struct { Policy Store *store.Store Registry *runtimegateway.Registry + // Credentials opens the bearer tokens of authenticated MCP servers. + Credentials Credentials // ManagedRuntimes is optional internal provisioning; it does not admit hosted API requests. ManagedRuntimes *RuntimeProvider // MaxConcurrentExecutions bounds work admitted by this Core execution owner. @@ -36,6 +39,12 @@ type Dispatcher struct { MaxConcurrentExecutions int } +// Credentials opens the bearer token of a Session's frozen MCP credential +// binding. +type Credentials interface { + MCPBearerToken(context.Context, vaults.MCPBearerToken) (string, error) +} + type Result struct { EngineErrorCode string `json:"engine_error_code,omitempty"` EngineHTTPStatus *int `json:"engine_http_status,omitempty"` diff --git a/services/core/internal/execution/mcp_credentials.go b/services/core/internal/execution/mcp_credentials.go index edf6280af..0350d07f4 100644 --- a/services/core/internal/execution/mcp_credentials.go +++ b/services/core/internal/execution/mcp_credentials.go @@ -6,13 +6,13 @@ import ( "net/url" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) // Resolve only the frozen decision. Never search current Vault contents during // dispatch: a later credential must not change an anonymous or selected server. -func selectedMCPCredentials(snapshot Snapshot) (map[string]store.MCPCredentialBinding, error) { +func selectedMCPCredentials(snapshot Snapshot) (map[string]vaults.MCPCredentialBinding, error) { invalid := errors.New("invalid frozen MCP credential binding") tools := map[string]v1.MCPTool{} for _, raw := range snapshot.Agent.Tools { @@ -33,7 +33,7 @@ func selectedMCPCredentials(snapshot Snapshot) (map[string]store.MCPCredentialBi attached[id] = true } seen := map[string]bool{} - selected := map[string]store.MCPCredentialBinding{} + selected := map[string]vaults.MCPCredentialBinding{} for _, binding := range snapshot.MCPCredentials { tool, exists := tools[binding.ServerLabel] if !exists || seen[binding.ServerLabel] || tool.Transport.ServerURL != binding.ServerURL { diff --git a/services/core/internal/execution/mcp_credentials_test.go b/services/core/internal/execution/mcp_credentials_test.go index c9c801275..a4cde8c93 100644 --- a/services/core/internal/execution/mcp_credentials_test.go +++ b/services/core/internal/execution/mcp_credentials_test.go @@ -6,7 +6,7 @@ import ( "testing" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) @@ -15,7 +15,7 @@ func TestMCPFrozenCredentialAdmission(t *testing.T) { for _, mode := range []string{"implicit", "explicit", "oauth implicit", "oauth explicit", "anonymous", "missing", "unattached", "wrong URL", "wrong auth", "changed selection", "HTTP", "self-hosted explicit", "self-hosted implicit", "self-hosted anonymous"} { t.Run(mode, func(t *testing.T) { tool := v1.MCPTool{Type: "mcp", ServerLabel: "tools", ConnectionOrigin: "service", Transport: v1.MCPHTTPTransport{Type: "http", ServerURL: "https://mcp.example/tools"}} - binding := store.MCPCredentialBinding{ServerLabel: "tools", ServerURL: tool.Transport.ServerURL, VaultID: vault, CredentialID: credential, AuthType: "static_bearer"} + binding := vaults.MCPCredentialBinding{ServerLabel: "tools", ServerURL: tool.Transport.ServerURL, VaultID: vault, CredentialID: credential, AuthType: "static_bearer"} snapshot := Snapshot{Agent: v1.Agent{Model: "model"}, Environment: &v1.Environment{Type: "none"}, VaultIDs: []string{vault}} switch mode { case "explicit", "oauth explicit", "missing", "changed selection", "self-hosted explicit": @@ -40,7 +40,7 @@ func TestMCPFrozenCredentialAdmission(t *testing.T) { if mode == "changed selection" { binding.CredentialID = uuid.NewString() } - snapshot.MCPCredentials = []store.MCPCredentialBinding{binding} + snapshot.MCPCredentials = []vaults.MCPCredentialBinding{binding} if mode == "missing" { snapshot.MCPCredentials = nil } diff --git a/services/core/internal/execution/mcp_support.go b/services/core/internal/execution/mcp_support.go index 3f1974366..bc36584f5 100644 --- a/services/core/internal/execution/mcp_support.go +++ b/services/core/internal/execution/mcp_support.go @@ -5,10 +5,10 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimedevice" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) -func (p Policy) mcpCredentialBindings(engine string, snapshot Snapshot) (map[string]store.MCPCredentialBinding, error) { +func (p Policy) mcpCredentialBindings(engine string, snapshot Snapshot) (map[string]vaults.MCPCredentialBinding, error) { selected, err := selectedMCPCredentials(snapshot) if err != nil { return nil, err @@ -22,8 +22,8 @@ func (p Policy) mcpCredentialBindings(engine string, snapshot Snapshot) (map[str // Selection, final preclaim and request construction use the same combination // checks. This function never reads plaintext credentials or native configuration. -func (p Policy) mcpExecutionCredentials(engine string, snapshot Snapshot, servers []proto.MCPHTTPServer, caps runtimedevice.KindCapabilities) (map[string]store.MCPCredentialBinding, error) { - fail := func(message string) (map[string]store.MCPCredentialBinding, error) { +func (p Policy) mcpExecutionCredentials(engine string, snapshot Snapshot, servers []proto.MCPHTTPServer, caps runtimedevice.KindCapabilities) (map[string]vaults.MCPCredentialBinding, error) { + fail := func(message string) (map[string]vaults.MCPCredentialBinding, error) { return nil, errors.New(message) } profile, _ := p.Engines.Lookup(engine) diff --git a/services/core/internal/execution/mcp_support_test.go b/services/core/internal/execution/mcp_support_test.go index f522dbe5d..d02cd278e 100644 --- a/services/core/internal/execution/mcp_support_test.go +++ b/services/core/internal/execution/mcp_support_test.go @@ -1,22 +1,36 @@ package execution import ( + "context" "encoding/json" + "reflect" "testing" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimedevice" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) +// recordingCredentials records each bearer-token lookup and answers with token. +type recordingCredentials struct { + token string + requests []vaults.MCPBearerToken +} + +func (c *recordingCredentials) MCPBearerToken(_ context.Context, command vaults.MCPBearerToken) (string, error) { + c.requests = append(c.requests, command) + return c.token, nil +} + func mcpSupportFixture(t *testing.T) (Snapshot, []proto.MCPHTTPServer, runtimedevice.KindCapabilities) { t.Helper() vault, credential := uuid.NewString(), uuid.NewString() tool := json.RawMessage(`{"type":"mcp","server_label":"tickets","connection_origin":"service","transport":{"type":"http","server_url":"https://mcp.example/tools"}}`) snapshot := Snapshot{Agent: v1.Agent{Model: "model", Tools: []json.RawMessage{tool}}, Environment: &v1.Environment{Type: "none"}, VaultIDs: []string{vault}, - MCPCredentials: []store.MCPCredentialBinding{{ServerLabel: "tickets", ServerURL: "https://mcp.example/tools", VaultID: vault, CredentialID: credential, AuthType: "static_bearer"}}} + MCPCredentials: []vaults.MCPCredentialBinding{{ServerLabel: "tickets", ServerURL: "https://mcp.example/tools", VaultID: vault, CredentialID: credential, AuthType: "static_bearer"}}} tools, err := executionTools(snapshot.Agent.Tools) if err != nil { t.Fatal(err) @@ -41,15 +55,22 @@ func TestMCPPublicBearerPolicyIsIndependentOfRuntimeCapabilities(t *testing.T) { if _, err := (Policy{}).mcpExecutionCredentials(engine, snapshot, servers, caps); (err == nil) != allowed { t.Fatal("runtime capabilities widened public admission", err) } - request, err := (&Dispatcher{}).executionRequest(t.Context(), store.Session{Engine: engine}, snapshot, caps, store.SessionExecutionBinding{}) - if err == nil || request.MCPHTTPServers != nil { - t.Fatal("credential execution without a store was admitted") + credentials := &recordingCredentials{token: "scoped-token"} + session := store.Session{TenantID: uuid.NewString(), Engine: engine} + request, err := (&Dispatcher{Credentials: credentials}).executionRequest(t.Context(), session, snapshot, caps, store.SessionExecutionBinding{}) + if !allowed { + if err == nil || err.Error() != "The configured engine does not support this MCP connection origin." || request.MCPHTTPServers != nil || len(credentials.requests) != 0 { + t.Fatal("unverified profile bypassed public policy", err) + } + return } - if allowed && err.Error() != "authenticated MCP execution is unavailable" { - t.Fatal("accepted profile did not reach scoped credential lookup", err) + // The lookup carries exactly the Session's tenant, attached Vaults and frozen binding. + want := []vaults.MCPBearerToken{{TenantID: session.TenantID, VaultIDs: snapshot.VaultIDs, Binding: snapshot.MCPCredentials[0]}} + if err != nil || !reflect.DeepEqual(credentials.requests, want) { + t.Fatal("accepted profile did not reach scoped credential lookup", err, credentials.requests) } - if !allowed && err.Error() != "The configured engine does not support this MCP connection origin." { - t.Fatal("unverified profile bypassed public policy", err) + if servers := request.MCPHTTPServers; servers == nil || len(*servers) != 1 || (*servers)[0].BearerToken == nil || *(*servers)[0].BearerToken != "scoped-token" { + t.Fatal("looked-up token did not reach its server") } }) } @@ -89,8 +110,9 @@ func TestMCPExecutionChecksRequireVerifiedCapabilityCombinations(t *testing.T) { t.Fatal("incorrect combined MCP capability decision", err) } if !allowed { - request, requestErr := (&Dispatcher{}).executionRequest(t.Context(), store.Session{Engine: "codex"}, snapshot, caps, store.SessionExecutionBinding{}) - if requestErr == nil || request.MCPHTTPServers != nil || requestErr.Error() == "authenticated MCP execution is unavailable" { + credentials := &recordingCredentials{token: "scoped-token"} + request, requestErr := (&Dispatcher{Credentials: credentials}).executionRequest(t.Context(), store.Session{Engine: "codex"}, snapshot, caps, store.SessionExecutionBinding{}) + if requestErr == nil || request.MCPHTTPServers != nil || len(credentials.requests) != 0 { t.Fatal("request bypassed capability checks before credential lookup", requestErr) } } @@ -120,7 +142,11 @@ func TestMCPAnonymousExecutionPreservesFrozenDecision(t *testing.T) { if err := (Policy{}).ValidateSessionConfiguration(engine, raw); (err == nil) != allowed || (Policy{}).canAdmitInputs(engine, raw) != allowed { t.Fatal("anonymous binding validation changed", err) } - request, err := (&Dispatcher{}).executionRequest(t.Context(), store.Session{Engine: engine}, snapshot, caps, store.SessionExecutionBinding{}) + credentials := &recordingCredentials{token: "scoped-token"} + request, err := (&Dispatcher{Credentials: credentials}).executionRequest(t.Context(), store.Session{Engine: engine}, snapshot, caps, store.SessionExecutionBinding{}) + if len(credentials.requests) != 0 { + t.Fatal("anonymous or invalid binding reached credential lookup") + } if !allowed { if err == nil || request.MCPHTTPServers != nil { t.Fatal("invalid frozen binding reached dispatch") diff --git a/services/core/internal/execution/owner_test.go b/services/core/internal/execution/owner_test.go index bdd5ae675..d68a749b2 100644 --- a/services/core/internal/execution/owner_test.go +++ b/services/core/internal/execution/owner_test.go @@ -68,6 +68,10 @@ func TestStartWorkerFailureClosesLeaseOnce(t *testing.T) { return err }, "missing Store": func(t *testing.T, lease *closeCountingLease) error { + _, err := StartWorker(canceled, &Dispatcher{Credentials: &recordingCredentials{}}, Owner{Lease: lease}) + return err + }, + "missing Credentials": func(t *testing.T, lease *closeCountingLease) error { _, err := StartWorker(canceled, &Dispatcher{}, Owner{Lease: lease}) return err }, @@ -75,7 +79,7 @@ func TestStartWorkerFailureClosesLeaseOnce(t *testing.T) { s, owner := resetManagerStore(t) lease.inner = owner.Lease id := uuid.NewString() - dispatcher := &Dispatcher{Store: s, Registry: runtimegateway.NewRegistry(), ManagedRuntimes: NewDeferredRuntimeProvider(id, func(context.Context) (*RuntimeProvider, error) { return nil, nil })} + dispatcher := &Dispatcher{Store: s, Registry: runtimegateway.NewRegistry(), Credentials: &recordingCredentials{}, ManagedRuntimes: NewDeferredRuntimeProvider(id, func(context.Context) (*RuntimeProvider, error) { return nil, nil })} _, err := StartWorker(canceled, dispatcher, Owner{Lease: lease, Store: owner.Store}) if ping := owner.Lease.CheckOwnership(t.Context()); !errors.Is(ping, pgunit.ErrLeaseClosed) { t.Error("failed start kept the database lease", ping) @@ -100,7 +104,7 @@ func TestWorkerRunClosesLeaseAfterDrain(t *testing.T) { s, owner := resetManagerStore(t) lease := &closeCountingLease{t: t, inner: owner.Lease} id := uuid.NewString() - dispatcher := &Dispatcher{Store: s, Registry: runtimegateway.NewRegistry(), ManagedRuntimes: NewDeferredRuntimeProvider(id, func(context.Context) (*RuntimeProvider, error) { return nil, nil })} + dispatcher := &Dispatcher{Store: s, Registry: runtimegateway.NewRegistry(), Credentials: &recordingCredentials{}, ManagedRuntimes: NewDeferredRuntimeProvider(id, func(context.Context) (*RuntimeProvider, error) { return nil, nil })} worker, err := StartWorker(t.Context(), dispatcher, Owner{Lease: lease, Store: owner.Store}) if err != nil { t.Fatal(err) diff --git a/services/core/internal/execution/request.go b/services/core/internal/execution/request.go index 353197059..dcd23cace 100644 --- a/services/core/internal/execution/request.go +++ b/services/core/internal/execution/request.go @@ -9,6 +9,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimedevice" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) func (d *Dispatcher) executionRequest(ctx context.Context, session store.Session, snapshot Snapshot, caps runtimedevice.KindCapabilities, bound store.SessionExecutionBinding) (proto.PromptRequestPayload, error) { @@ -60,12 +61,9 @@ func (d *Dispatcher) executionRequest(ctx context.Context, session store.Session if err != nil { return proto.PromptRequestPayload{}, err } - if len(selected) > 0 && d.Store == nil { - return proto.PromptRequestPayload{}, errors.New("authenticated MCP execution is unavailable") - } for i := range tools.MCP { if binding, ok := selected[tools.MCP[i].ServerLabel]; ok { - token, err := d.Store.MCPBearerToken(ctx, session.TenantID, snapshot.VaultIDs, binding) + token, err := d.Credentials.MCPBearerToken(ctx, vaults.MCPBearerToken{TenantID: session.TenantID, VaultIDs: snapshot.VaultIDs, Binding: binding}) if err != nil { return proto.PromptRequestPayload{}, err } diff --git a/services/core/internal/execution/worker.go b/services/core/internal/execution/worker.go index 97e91ad22..fd9b6a423 100644 --- a/services/core/internal/execution/worker.go +++ b/services/core/internal/execution/worker.go @@ -49,6 +49,9 @@ func StartWorker(ctx context.Context, dispatcher *Dispatcher, owner Owner) (_ *W if dispatcher.MaxConcurrentExecutions < 0 || dispatcher.MaxConcurrentExecutions > 1024 { return nil, errors.New("execution concurrency must be between 1 and 1024, or zero for the default") } + if dispatcher.Credentials == nil { + return nil, errors.New("execution worker requires MCP Credentials") + } if owner.Store == nil { return nil, errors.New("execution worker requires the execution Store") } diff --git a/services/core/internal/persistence/postgres/vaultpg/audit_test.go b/services/core/internal/persistence/postgres/vaultpg/audit_test.go new file mode 100644 index 000000000..329d2b1fc --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/audit_test.go @@ -0,0 +1,220 @@ +package vaultpg_test + +import ( + "bytes" + "context" + "reflect" + "strings" + "testing" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/adminaudit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/vaultpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/writeaudit" +) + +const rejectedAudit = "reject-vault-audit" + +// auditMutation is one audited write, prepared against its own tenant. +type auditMutation struct { + action, kind, parent string + owners int + run func(context.Context) (string, error) +} + +func publicAuditContext(ctx context.Context, tenant, request string) context.Context { + return writeaudit.WithSource(ctx, writeaudit.Source{ + KeyID: "static:" + strings.Repeat("a", 64), Name: "vault audit fixture", Prefix: "aaaaaaaa", + Kind: "static", TenantID: tenant, RequestID: request, TraceID: "vault-audit-trace", + }) +} + +// adminAuditContext keeps an inherited public source, which must not turn an +// administrator operation into a user-key operation. +func adminAuditContext(ctx context.Context, tenant, request string) context.Context { + return adminaudit.WithSource(publicAuditContext(ctx, tenant, request), adminaudit.Source{ + CredentialID: "87654321", ActorLabel: "administrator fixture", ProjectID: tenant, RequestID: request, TraceID: "admin-mutation-trace", + }) +} + +// rejectAudits fails every audit insertion made for the rejected request. +// The sequence survives rollback, so it proves the insertion was reached even +// though the Store reports only a sanitized failure. +func rejectAudits(t *testing.T, pool *pgxpool.Pool) { + t.Helper() + _, err := pool.Exec(t.Context(), `CREATE SEQUENCE vault_audit_rejections; + CREATE FUNCTION reject_vault_audit() RETURNS trigger LANGUAGE plpgsql AS $$ + BEGIN IF NEW.request_id = '`+rejectedAudit+`' THEN PERFORM nextval('vault_audit_rejections'); RAISE EXCEPTION 'forced audit insertion failure'; END IF; RETURN NEW; END $$; + CREATE TRIGGER reject_vault_audit BEFORE INSERT ON write_audit_operations FOR EACH ROW EXECUTE FUNCTION reject_vault_audit(); + CREATE TRIGGER reject_vault_audit BEFORE INSERT ON admin_audit_log FOR EACH ROW EXECUTE FUNCTION reject_vault_audit()`) + if err != nil { + t.Fatal(err) + } +} + +func auditRejections(t *testing.T, pool *pgxpool.Pool) int64 { + t.Helper() + var count int64 + if err := pool.QueryRow(t.Context(), "SELECT CASE WHEN is_called THEN last_value ELSE 0 END FROM vault_audit_rejections").Scan(&count); err != nil { + t.Fatal(err) + } + return count +} + +// auditSnapshot covers the isolated database's complete rows, including +// ciphertext and timestamps. +func auditSnapshot(t *testing.T, pool *pgxpool.Pool) map[string]string { + t.Helper() + result := map[string]string{} + for _, table := range []string{"vaults", "vault_credentials", "write_audit_operations", "write_audit_owners", "admin_audit_log"} { + var rows string + if err := pool.QueryRow(t.Context(), "SELECT COALESCE(jsonb_agg(to_jsonb(r) ORDER BY to_jsonb(r)::text)::text,'[]') FROM "+pgx.Identifier{table}.Sanitize()+" r").Scan(&rows); err != nil { + t.Fatalf("snapshot %s: %v", table, err) + } + result[table] = rows + } + return result +} + +func prepareAuditMutation(t *testing.T, service *vaults.Service, tenant, name string) auditMutation { + t.Helper() + ctx := t.Context() + if name == "vault_create" { + return auditMutation{action: "create", kind: "vault", owners: 1, run: func(ctx context.Context) (string, error) { + v, err := service.CreateVault(ctx, vaults.CreateVault{TenantID: tenant}) + return v.ID, err + }} + } + vault := createVault(t, service, tenant) + static := vaults.CreateStaticCredential{TenantID: tenant, VaultID: vault.ID, Name: "fixture", MCPServerURL: "https://mcp.example/", Token: "audit-private-token"} + remove := func(id string) auditMutation { + return auditMutation{action: "delete", kind: "credential", parent: vault.ID, run: func(ctx context.Context) (string, error) { + return service.DeleteCredential(ctx, vaults.DeleteCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: id}) + }} + } + switch name { + case "credential_create": + return auditMutation{action: "create", kind: "credential", parent: vault.ID, owners: 1, run: func(ctx context.Context) (string, error) { + v, err := service.CreateStaticCredential(ctx, static) + return v.ID, err + }} + case "oauth_create", "oauth_update", "oauth_delete": + command := vaults.CreateOAuthCredential{TenantID: tenant, VaultID: vault.ID, Name: "fixture", MCPServerURL: static.MCPServerURL, AccessToken: static.Token} + if name == "oauth_create" { + return auditMutation{action: "create", kind: "credential", parent: vault.ID, owners: 1, run: func(ctx context.Context) (string, error) { + v, err := service.CreateOAuthCredential(ctx, command) + return v.ID, err + }} + } + credential, err := service.CreateOAuthCredential(ctx, command) + if err != nil { + t.Fatal(err) + } + if name == "oauth_delete" { + return remove(credential.ID) + } + return auditMutation{action: "update", kind: "credential", parent: vault.ID, run: func(ctx context.Context) (string, error) { + v, err := service.UpdateOAuthCredential(ctx, vaults.UpdateOAuthCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: credential.ID, AccessToken: ptr("audit-private-replacement")}) + return v.ID, err + }} + } + credential := createStatic(t, service, tenant, vault.ID, static.Name, static.MCPServerURL, static.Token) + switch name { + case "vault_delete": + return auditMutation{action: "delete", kind: "vault", run: func(ctx context.Context) (string, error) { + return service.DeleteVault(ctx, vaults.DeleteVault{TenantID: tenant, VaultID: vault.ID}) + }} + case "credential_update": + return auditMutation{action: "update", kind: "credential", parent: vault.ID, run: func(ctx context.Context) (string, error) { + v, err := service.UpdateStaticCredential(ctx, vaults.UpdateStaticCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: credential.ID, Token: "audit-private-replacement"}) + return v.ID, err + }} + } + return remove(credential.ID) +} + +// A trigger fails the final audit insertion after each real mutation. Complete +// table snapshots prove the rollback of ciphertext, timestamps and cascades. +func TestVaultMutationsRollBackWithTheirAudit(t *testing.T) { + pool := pgtest.OpenIsolated(t, nil) + service := newService(t, vaultpg.New(pgunit.NewPool(pool)), newCipher(t, bytes.Repeat([]byte{91}, 32)), nil) + rejectAudits(t, pool) + secrets := []string{"audit-private-token", "audit-private-replacement"} + for _, provenance := range []string{"public", "admin"} { + names := []string{"vault_create", "vault_delete", "credential_create", "credential_update", "credential_delete", "oauth_create", "oauth_update", "oauth_delete"} + source := publicAuditContext + if provenance == "admin" { + // Administrators only delete Vaults and Credentials. + names, source = []string{"vault_delete", "credential_delete", "oauth_delete"}, adminAuditContext + } + for _, name := range names { + t.Run(provenance+"/"+name, func(t *testing.T) { + tenant := uuid.NewString() + mutation := prepareAuditMutation(t, service, tenant, name) + if provenance == "admin" { + // The administrator audit references the Project, which + // references its execution scope. + if _, err := pool.Exec(t.Context(), "INSERT INTO execution_project_scopes(tenant_id,organization_id,project_id) VALUES($1,'admin-delete',$2)", tenant, tenant); err != nil { + t.Fatal(err) + } + if _, err := pool.Exec(t.Context(), "INSERT INTO projects(id,name,tenant_id,subject_kind,subject_id) VALUES($1,'Delete fixture',$1,'service_account',$2)", tenant, "project:"+tenant); err != nil { + t.Fatal(err) + } + } + before, rejections := auditSnapshot(t, pool), auditRejections(t, pool) + if _, err := mutation.run(source(t.Context(), tenant, rejectedAudit)); err == nil || auditRejections(t, pool) != rejections+1 { + t.Fatal("mutation did not reach the failing audit insertion", err) + } + if !reflect.DeepEqual(before, auditSnapshot(t, pool)) { + t.Fatal("audit failure left business or audit changes") + } + request := uuid.NewString() + id, err := mutation.run(source(t.Context(), tenant, request)) + if err != nil { + t.Fatal(err) + } + var public, admin, owners int + if err := pool.QueryRow(t.Context(), `SELECT (SELECT count(*) FROM write_audit_operations WHERE tenant_id=$1 AND request_id=$2), + (SELECT count(*) FROM admin_audit_log WHERE tenant_id=$1 AND request_id=$2), (SELECT count(*) FROM write_audit_owners WHERE tenant_id=$1)`, tenant, request).Scan(&public, &admin, &owners); err != nil { + t.Fatal(err) + } + var action, kind, gotID, parent, raw string + if provenance == "public" { + if public != 1 || admin != 0 || owners != mutation.owners { + t.Fatalf("public audit rows %d, admin rows %d, owners %d", public, admin, owners) + } + if err := pool.QueryRow(t.Context(), `SELECT action,resource_type,resource_id,COALESCE(parent_id,''),to_jsonb(o)::text FROM write_audit_operations o WHERE tenant_id=$1 AND request_id=$2`, tenant, request).Scan(&action, &kind, &gotID, &parent, &raw); err != nil { + t.Fatal(err) + } + } else { + // An administrator deletion never records public-key provenance. + if public != 0 || admin != 1 || owners != 0 { + t.Fatalf("public audit rows %d, admin rows %d, owners %d", public, admin, owners) + } + var credential, actor, project, trace, results string + if err := pool.QueryRow(t.Context(), `SELECT admin_credential_id,actor_label,project_id,trace_id,result_ids::text,action,resource_type,resource_id,to_jsonb(a)::text FROM admin_audit_log a WHERE tenant_id=$1 AND request_id=$2`, tenant, request).Scan(&credential, &actor, &project, &trace, &results, &action, &kind, &gotID, &raw); err != nil { + t.Fatal(err) + } + if credential != "87654321" || actor != "administrator fixture" || project != tenant || trace != "admin-mutation-trace" || results != "[]" { + t.Fatal("administrator audit identity differs") + } + parent = mutation.parent + } + if action != mutation.action || kind != mutation.kind || gotID != id || parent != mutation.parent { + t.Fatalf("wrong operation identity: %s %s %s %s", action, kind, gotID, parent) + } + for _, secret := range secrets { + if strings.Contains(raw, secret) { + t.Fatal("audit contains a secret") + } + } + }) + } + } +} diff --git a/services/core/internal/persistence/postgres/vaultpg/credentials.go b/services/core/internal/persistence/postgres/vaultpg/credentials.go new file mode 100644 index 000000000..dbed4d95f --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/credentials.go @@ -0,0 +1,197 @@ +package vaultpg + +import ( + "context" + "encoding/json" + "errors" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgconn" + "github.com/jackc/pgx/v5/pgtype" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/writeaudit" +) + +// CreateCredential admits the owning Vault in the insert itself, so a missing +// or foreign Vault stores nothing and is ErrNotFound, as is a Vault deleted +// while the insert waits on it. A malformed new Credential ID is +// ErrInvalidInput. +func (s *Store) CreateCredential(ctx context.Context, credential vaults.NewCredential) (vaults.Credential, error) { + if _, err := pgunit.ParseID(credential.CredentialID); err != nil { + return vaults.Credential{}, vaults.ErrInvalidInput + } + var created vaults.Credential + err := s.write(ctx, func(ctx context.Context, q *sqlc.Queries) error { + row, err := insertCredential(ctx, q, credential) + // A foreign-key violation means the Vault was deleted after the insert read it. + var constraint *pgconn.PgError + if errors.As(err, &constraint) && constraint.Code == "23503" { + return vaults.ErrNotFound + } + if err != nil { + return err + } + if created, err = credentialFromRow(row); err != nil { + return err + } + return auditpg.RecordWriteAudit(ctx, q, credential.TenantID, "create", "credential", created.ID, created.VaultID, + writeaudit.Resource{Type: "credential", ID: created.ID, ParentID: created.VaultID}) + }) + if err != nil { + return vaults.Credential{}, err + } + return created, nil +} + +func insertCredential(ctx context.Context, q *sqlc.Queries, credential vaults.NewCredential) (sqlc.GetCredentialRow, error) { + id, tenant, vault := pgunit.PathID(credential.CredentialID), pgunit.PathID(credential.TenantID), pgunit.PathID(credential.VaultID) + switch credential.AuthType { + case vaults.AuthStaticBearer: + row, err := q.CreateStaticCredential(ctx, sqlc.CreateStaticCredentialParams{ID: id, TenantID: tenant, VaultID: vault, + Name: credential.Name, McpServerUrl: credential.MCPServerURL, TokenCiphertext: credential.Ciphertext}) + return sqlc.GetCredentialRow(row), err + case vaults.AuthMCPOAuth: + row, err := q.CreateOAuthCredential(ctx, sqlc.CreateOAuthCredentialParams{ID: id, TenantID: tenant, VaultID: vault, + Name: credential.Name, McpServerUrl: credential.MCPServerURL, OauthMetadata: credential.OAuthMetadata, TokenCiphertext: credential.Ciphertext}) + return sqlc.GetCredentialRow(row), err + } + return sqlc.GetCredentialRow{}, errors.New("unknown credential authentication type") +} + +func (s *Store) GetCredential(ctx context.Context, tenantID, vaultID, credentialID string) (vaults.Credential, error) { + tenant, err := pgunit.ParseID(tenantID) + if err != nil { + return vaults.Credential{}, vaults.ErrInvalidInput + } + vault, err := pgunit.ParseID(vaultID) + if err != nil { + return vaults.Credential{}, vaults.ErrInvalidInput + } + id, err := pgunit.ParseID(credentialID) + if err != nil { + return vaults.Credential{}, vaults.ErrInvalidInput + } + return getCredential(ctx, s.pool.Queries(), tenant, vault, id) +} + +func getCredential(ctx context.Context, q *sqlc.Queries, tenant, vault, id pgtype.UUID) (vaults.Credential, error) { + row, err := q.GetCredential(ctx, sqlc.GetCredentialParams{TenantID: tenant, VaultID: vault, ID: id}) + if err != nil { + return vaults.Credential{}, translate(err) + } + return credentialFromRow(row) +} + +// ListCredentials reads the parent Vault, the cursor and the page from one +// snapshot. A parent ID that cannot name a Vault is ErrNotFound. +func (s *Store) ListCredentials(ctx context.Context, tenantID, vaultID string, query vaults.PageQuery) (vaults.CredentialPage, error) { + tenant, err := pgunit.ParseID(tenantID) + if err != nil { + return vaults.CredentialPage{}, vaults.ErrInvalidInput + } + vault := pgunit.PathID(vaultID) + var page vaults.CredentialPage + err = s.read(ctx, func(ctx context.Context, q *sqlc.Queries) error { + // An inaccessible parent is not an authorized empty collection. + if _, err := getVault(ctx, q, tenant, vault); err != nil { + return err + } + statuses, err := query.Validate() + if err != nil { + return err + } + params := sqlc.ListCredentialsParams{TenantID: tenant, VaultID: vault, PageLimit: int32(query.Limit + 1), AfterID: pgtype.UUID{Valid: true}, Ascending: query.Ascending, Statuses: statuses} + if query.After != "" { + // A cursor that cannot name a Credential of this Vault follows the + // missing-cursor path. + after, err := getCredential(ctx, q, tenant, vault, pgunit.PathID(query.After)) + if err != nil { + return err + } + params.AfterCreated = pgtype.Timestamptz{Time: after.CreatedAt, Valid: true} + params.AfterID = pgunit.PathID(after.ID) + } + rows, err := q.ListCredentials(ctx, params) + if err != nil { + return err + } + page = vaults.CredentialPage{Credentials: make([]vaults.Credential, 0, min(query.Limit, len(rows)))} + if len(rows) > query.Limit { + page.NextCursor = uuid.UUID(rows[query.Limit-1].ID.Bytes).String() + rows = rows[:query.Limit] + } + for _, row := range rows { + credential, err := credentialFromRow(sqlc.GetCredentialRow(row)) + if err != nil { + return err + } + page.Credentials = append(page.Credentials, credential) + } + return nil + }) + if err != nil { + return vaults.CredentialPage{}, err + } + return page, nil +} + +// ReplaceStaticToken matches the destination the token is sealed to, so a +// concurrent change of scope stores nothing. +func (s *Store) ReplaceStaticToken(ctx context.Context, replacement vaults.StaticTokenReplacement) (vaults.Credential, error) { + var updated vaults.Credential + err := s.write(ctx, func(ctx context.Context, q *sqlc.Queries) error { + row, err := q.UpdateStaticCredential(ctx, sqlc.UpdateStaticCredentialParams{ + TenantID: pgunit.PathID(replacement.TenantID), VaultID: pgunit.PathID(replacement.VaultID), ID: pgunit.PathID(replacement.CredentialID), + McpServerUrl: replacement.MCPServerURL, TokenCiphertext: replacement.Ciphertext, + }) + if err != nil { + return err + } + updated, err = credentialFromRow(sqlc.GetCredentialRow(row)) + if err != nil { + return err + } + return auditpg.RecordWriteAudit(ctx, q, replacement.TenantID, "update", "credential", updated.ID, updated.VaultID) + }) + if err != nil { + return vaults.Credential{}, err + } + return updated, nil +} + +// DeleteCredential removes the sealed secret without reading it. +func (s *Store) DeleteCredential(ctx context.Context, key vaults.CredentialKey) (string, error) { + vault := pgunit.PathID(key.VaultID) + var deleted string + err := s.write(ctx, func(ctx context.Context, q *sqlc.Queries) error { + id, err := q.DeleteCredential(ctx, sqlc.DeleteCredentialParams{TenantID: pgunit.PathID(key.TenantID), VaultID: vault, ID: pgunit.PathID(key.CredentialID)}) + if err != nil { + return err + } + deleted = uuid.UUID(id.Bytes).String() + return auditpg.RecordWriteAudit(ctx, q, key.TenantID, "delete", "credential", deleted, uuid.UUID(vault.Bytes).String()) + }) + if err != nil { + return "", err + } + return deleted, nil +} + +func credentialFromRow(row sqlc.GetCredentialRow) (vaults.Credential, error) { + credential := vaults.Credential{ + ID: uuid.UUID(row.ID.Bytes).String(), VaultID: uuid.UUID(row.VaultID.Bytes).String(), + Name: row.Name, AuthType: row.AuthType, MCPServerURL: row.McpServerUrl, + CreatedAt: row.CreatedAt.Time, UpdatedAt: row.UpdatedAt.Time, + } + if row.AuthType == vaults.AuthMCPOAuth { + credential.OAuth = &vaults.OAuthMetadata{} + if err := json.Unmarshal(row.OauthMetadata, credential.OAuth); err != nil { + return vaults.Credential{}, errors.New("invalid stored OAuth metadata") + } + } + return credential, nil +} diff --git a/services/core/internal/persistence/postgres/vaultpg/credentials_test.go b/services/core/internal/persistence/postgres/vaultpg/credentials_test.go new file mode 100644 index 000000000..fafd9b862 --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/credentials_test.go @@ -0,0 +1,497 @@ +package vaultpg_test + +import ( + "bytes" + "cmp" + "crypto/rand" + "encoding/hex" + "encoding/json" + "errors" + "reflect" + "slices" + "strings" + "testing" + "time" + + "github.com/google/uuid" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +func TestStaticCredentialsPersistEncryptedAndRemainScoped(t *testing.T) { + store, pool := openStore(t) + ctx := t.Context() + tenant, foreignTenant := uuid.NewString(), uuid.NewString() + keyless := newService(t, store, nil, nil) + var owned []vaults.Vault + for _, owner := range []string{tenant, tenant, foreignTenant} { + owned = append(owned, createVault(t, keyless, owner)) + } + key, randomToken := make([]byte, 32), make([]byte, 32) + if _, err := rand.Read(key); err != nil { + t.Fatal(err) + } + if _, err := rand.Read(randomToken); err != nil { + t.Fatal(err) + } + service := newService(t, store, newCipher(t, key), nil) + canary := hex.EncodeToString(randomToken) + opaque := " \t" + canary + " 凭据\n" + strings.Repeat("x", 300) + " " + tokens := []string{opaque, opaque, ""} + var records []vaults.Credential + before := time.Now().Add(-time.Second) + for _, token := range tokens { + command := vaults.CreateStaticCredential{TenantID: tenant, VaultID: owned[0].ID, Name: "MCP credential", MCPServerURL: "https://mcp.example/tools", Token: token} + record, err := service.CreateStaticCredential(ctx, command) + if err != nil { + t.Fatal(err) + } + if _, err := uuid.Parse(record.ID); err != nil || record.VaultID != owned[0].ID || record.Name != command.Name || record.AuthType != vaults.AuthStaticBearer || record.MCPServerURL != command.MCPServerURL || record.CreatedAt.Before(before) || record.CreatedAt.After(time.Now().Add(time.Second)) || !record.CreatedAt.Equal(record.UpdatedAt) || record.OAuth != nil { + t.Fatal("credential metadata or database timestamps differ") + } + metadata, err := json.Marshal(record) + if err != nil || bytes.Contains(metadata, []byte(canary)) { + t.Fatal("credential metadata contains token plaintext") + } + records = append(records, record) + } + if records[0].ID == records[1].ID { + t.Fatal("separate creates reused a credential identity") + } + valid := vaults.CreateStaticCredential{TenantID: tenant, VaultID: owned[0].ID, Name: "Rejected", MCPServerURL: "https://mcp.example/tools", Token: opaque} + if _, err := keyless.CreateStaticCredential(ctx, valid); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("missing encryption key did not fail writes closed") + } + for _, target := range []struct{ tenant, vault string }{{tenant, owned[2].ID}, {foreignTenant, owned[0].ID}, {tenant, uuid.NewString()}, {tenant, "invalid"}} { + command := valid + command.TenantID, command.VaultID = target.tenant, target.vault + if _, err := service.CreateStaticCredential(ctx, command); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("creation admitted an unowned or missing Vault") + } + } + for _, target := range []struct{ tenant, vault, credential string }{ + {tenant, owned[1].ID, records[0].ID}, {foreignTenant, owned[0].ID, records[0].ID}, + {tenant, owned[2].ID, records[0].ID}, {tenant, uuid.NewString(), records[0].ID}, {tenant, owned[0].ID, uuid.NewString()}, + } { + if _, err := store.GetCredential(ctx, target.tenant, target.vault, target.credential); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("unowned, wrong-Vault or missing credential was disclosed") + } + } + pool.Close() + store, pool = openStore(t) + restartedCipher := newCipher(t, bytes.Clone(key)) + wrongKey := bytes.Clone(key) + wrongKey[0] ^= 1 + wrongCipher := newCipher(t, wrongKey) + var ciphertexts [][]byte + for i, record := range records { + got, err := store.GetCredential(ctx, tenant, record.VaultID, record.ID) + if err != nil || !reflect.DeepEqual(got, record) { + t.Fatal("safe metadata recovery changed", err) + } + var ciphertext []byte + var storageType string + if err := pool.QueryRow(ctx, "SELECT token_ciphertext, pg_typeof(token_ciphertext)::text FROM vault_credentials WHERE id=$1 AND vault_id=$2", record.ID, record.VaultID).Scan(&ciphertext, &storageType); err != nil || storageType != "bytea" || len(ciphertext) < 29 || bytes.Contains(ciphertext, []byte(canary)) { + t.Fatal("credential ciphertext was not stored as private bytea", err) + } + binding := credentialcrypto.Binding{TenantID: tenant, VaultID: record.VaultID, CredentialID: record.ID, AuthType: record.AuthType, Destination: record.MCPServerURL} + plaintext, err := restartedCipher.Open(ciphertext, binding) + if err != nil || !bytes.Equal(plaintext, []byte(tokens[i])) { + t.Fatal("private restart decryption did not preserve token bytes", err) + } + if plaintext, err := wrongCipher.Open(ciphertext, binding); err == nil || plaintext != nil { + t.Fatal("wrong key decrypted persisted ciphertext") + } + ciphertexts = append(ciphertexts, ciphertext) + } + // Version 1 prefixes the standard library's 12-byte random nonce. + if bytes.Equal(ciphertexts[0], ciphertexts[1]) || bytes.Equal(ciphertexts[0][1:13], ciphertexts[1][1:13]) { + t.Fatal("same token in distinct records reused ciphertext or nonce") + } + // Change one persisted binding field. Metadata reads still work, while + // private decryption must reject the altered stored destination. + if _, err := pool.Exec(ctx, "UPDATE vault_credentials SET mcp_server_url=$1 WHERE id=$2", "https://other.example/tools", records[0].ID); err != nil { + t.Fatal(err) + } + changed, err := store.GetCredential(ctx, tenant, owned[0].ID, records[0].ID) + if err != nil || changed.MCPServerURL != "https://other.example/tools" { + t.Fatal("metadata read unexpectedly required decryption", err) + } + binding := credentialcrypto.Binding{TenantID: tenant, VaultID: changed.VaultID, CredentialID: changed.ID, AuthType: changed.AuthType, Destination: changed.MCPServerURL} + if plaintext, err := restartedCipher.Open(ciphertexts[0], binding); err == nil || plaintext != nil { + t.Fatal("persisted destination substitution authenticated") + } + if _, err := pool.Exec(ctx, "UPDATE vault_credentials SET mcp_server_url=$1 WHERE id=$2", records[0].MCPServerURL, records[0].ID); err != nil { + t.Fatal(err) + } + for _, vault := range owned { + got, err := store.GetVault(ctx, vault.TenantID, vault.ID) + if err != nil || !reflect.DeepEqual(got, vault) { + t.Fatal("credential operations changed an owning Vault", err) + } + } + var count, sessions int + if err := pool.QueryRow(ctx, "SELECT (SELECT count(*) FROM vault_credentials WHERE vault_id=ANY($1::uuid[])), (SELECT count(*) FROM sessions WHERE tenant_id=ANY($2::uuid[]))", []string{owned[0].ID, owned[1].ID, owned[2].ID}, []string{tenant, foreignTenant}).Scan(&count, &sessions); err != nil || count != len(records) || sessions != 0 { + t.Fatal("rejected requests wrote rows or credential operations created Sessions", err) + } +} + +func TestCredentialListFilteringOwnershipAndKeylessReconnect(t *testing.T) { + store, pool := openStore(t) + ctx := t.Context() + tenant, foreign := uuid.NewString(), uuid.NewString() + service := newService(t, store, newCipher(t, make([]byte, 32)), nil) + var owned []vaults.Vault + for _, owner := range []string{tenant, tenant, foreign, tenant} { + owned = append(owned, createVault(t, service, owner)) + } + // Parent classification does not classify its Credentials. + if _, err := pool.Exec(ctx, "UPDATE vaults SET status='archived' WHERE id=$1", owned[0].ID); err != nil { + t.Fatal(err) + } + create := func(owner, vault string) vaults.Credential { + t.Helper() + return createStatic(t, service, owner, vault, "List fixture", "https://example.invalid/mcp", "synthetic-token-not-public") + } + var all []vaults.Credential + archived := map[string]bool{} + for i := range 105 { + c := create(tenant, owned[0].ID) + status := vaults.StatusActive + if i%3 == 0 { + status, archived[c.ID] = vaults.StatusArchived, true + } + if _, err := pool.Exec(ctx, "UPDATE vault_credentials SET status=$1, created_at=$2 WHERE id=$3", status, time.Unix(1700000000+int64(i%2), 0), c.ID); err != nil { + t.Fatal(err) + } + c, err := store.GetCredential(ctx, tenant, owned[0].ID, c.ID) + if err != nil { + t.Fatal(err) + } + all = append(all, c) + } + slices.SortFunc(all, func(a, b vaults.Credential) int { + if c := a.CreatedAt.Compare(b.CreatedAt); c != 0 { + return c + } + return cmp.Compare(a.ID, b.ID) + }) + otherVault, otherProject := create(tenant, owned[1].ID), create(foreign, owned[2].ID) + snapshot := func() string { + t.Helper() + var value string + if err := pool.QueryRow(ctx, "SELECT jsonb_agg(to_jsonb(c) ORDER BY id)::text FROM vault_credentials c WHERE vault_id=$1", owned[0].ID).Scan(&value); err != nil { + t.Fatal(err) + } + return value + } + before := snapshot() + read := func(reader vaults.Reader, ascending bool, statuses []string, size int) []vaults.Credential { + t.Helper() + actual := []vaults.Credential{} + cursor := "" + for { + page, err := reader.ListCredentials(ctx, tenant, owned[0].ID, vaults.PageQuery{After: cursor, Limit: size, Ascending: ascending, Statuses: statuses}) + if err != nil || len(page.Credentials) == 0 || len(page.Credentials) > size { + t.Fatal("invalid page", err) + } + actual = append(actual, page.Credentials...) + if len(actual) > len(all) { + t.Fatal("repeated pagination") + } + if page.NextCursor == "" { + break + } + if page.NextCursor != page.Credentials[len(page.Credentials)-1].ID { + t.Fatal("cursor is not last included Credential") + } + cursor = page.NextCursor + } + return actual + } + for _, ascending := range []bool{true, false} { + for _, statuses := range [][]string{nil, {vaults.StatusActive}, {vaults.StatusArchived}, {vaults.StatusActive, vaults.StatusArchived}} { + want := []vaults.Credential{} + for _, c := range all { + if len(statuses) != 1 || archived[c.ID] == (statuses[0] == vaults.StatusArchived) { + want = append(want, c) + } + } + if !ascending { + slices.Reverse(want) + } + for _, size := range []int{20, 100} { + if got := read(store, ascending, statuses, size); !reflect.DeepEqual(got, want) { + t.Fatalf("metadata/filter/order mismatch: ascending=%t statuses=%v size=%d", ascending, statuses, size) + } + } + } + } + // A parent or cursor ID that cannot name a resource of this list is a + // missing one. + for _, tc := range []struct{ owner, vault, cursor string }{ + {tenant, owned[0].ID, otherVault.ID}, {tenant, owned[0].ID, otherProject.ID}, {tenant, owned[0].ID, uuid.NewString()}, {tenant, owned[0].ID, "invalid"}, + {tenant, owned[2].ID, ""}, {foreign, owned[0].ID, ""}, {tenant, uuid.NewString(), ""}, {tenant, "invalid", ""}, + } { + if _, err := store.ListCredentials(ctx, tc.owner, tc.vault, vaults.PageQuery{After: tc.cursor, Limit: 20}); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("unowned/unknown parent or cursor accepted", err) + } + } + for _, tc := range []struct{ vault, cursor string }{{owned[0].ID, all[len(all)-1].ID}, {owned[3].ID, ""}} { + page, err := store.ListCredentials(ctx, tenant, tc.vault, vaults.PageQuery{After: tc.cursor, Limit: 20, Ascending: true}) + if err != nil || page.Credentials == nil || len(page.Credentials) != 0 || page.NextCursor != "" { + t.Fatal("empty/terminal page", err) + } + } + for _, tc := range []struct { + owner string + query vaults.PageQuery + }{{"invalid", vaults.PageQuery{Limit: 20}}, {tenant, vaults.PageQuery{Limit: 0}}, {tenant, vaults.PageQuery{Limit: 101}}, {tenant, vaults.PageQuery{Limit: 20, Statuses: []string{"deleted"}}}} { + if _, err := store.ListCredentials(ctx, tc.owner, owned[0].ID, tc.query); !errors.Is(err, vaults.ErrInvalidInput) { + t.Fatal("invalid internal query accepted", err) + } + } + pool.Close() + reopened, pool := openStore(t) + if got := read(reopened, true, nil, 20); !reflect.DeepEqual(got, all) || snapshot() != before { + t.Fatal("keyless reads/restart changed metadata, classification or ciphertext") + } +} + +func TestStaticCredentialUpdatePreservesBindingsAndReplacesCurrentSecret(t *testing.T) { + store, pool := openStore(t) + tenant, foreign := uuid.NewString(), uuid.NewString() + key := make([]byte, 32) + if _, err := rand.Read(key); err != nil { + t.Fatal(err) + } + cipher := newCipher(t, key) + service := newService(t, store, cipher, nil) + var owned []vaults.Vault + for _, owner := range []string{tenant, tenant, foreign} { + owned = append(owned, createVault(t, service, owner)) + } + firstToken, endpoint := uuid.NewString(), "https://mcp.example/tools" + original := createStatic(t, service, tenant, owned[0].ID, "Retained name", endpoint, firstToken) + unrelated := createStatic(t, service, tenant, owned[1].ID, "Unrelated", endpoint, firstToken) + attached := []string{owned[0].ID} + // An implicit and an explicit selection freeze the same Credential. + bindings := []vaults.MCPCredentialBinding{ + resolve(t, service, tenant, attached, vaults.MCPCredentialRequest{ServerLabel: "tools", ServerURL: endpoint})[0], + resolve(t, service, tenant, attached, vaults.MCPCredentialRequest{ServerLabel: "tools", ServerURL: endpoint, CredentialID: &original.ID})[0], + } + readCiphertext := func(id string) []byte { + t.Helper() + var ciphertext []byte + if err := pool.QueryRow(t.Context(), "SELECT token_ciphertext FROM vault_credentials WHERE id=$1", id).Scan(&ciphertext); err != nil { + t.Fatal("private ciphertext observation failed") + } + return ciphertext + } + update := func(service *vaults.Service, tenant, vault, id, token string) (vaults.Credential, error) { + return service.UpdateStaticCredential(t.Context(), vaults.UpdateStaticCredential{TenantID: tenant, VaultID: vault, CredentialID: id, Token: token}) + } + prior, unrelatedCiphertext := readCiphertext(original.ID), readCiphertext(unrelated.ID) + lastToken := " \t" + uuid.NewString() + " 雪\n" + for _, token := range []string{"", lastToken, lastToken} { + updated, err := update(service, tenant, original.VaultID, original.ID, token) + if err != nil { + t.Fatal("token replacement failed") + } + want := original + want.UpdatedAt = updated.UpdatedAt + if !reflect.DeepEqual(updated, want) || updated.UpdatedAt.Before(original.UpdatedAt) { + t.Fatal("replacement changed immutable metadata") + } + current := readCiphertext(original.ID) + if bytes.Equal(prior, current) || len(token) > 0 && bytes.Contains(current, []byte(token)) { + t.Fatal("replacement reused ciphertext or stored plaintext") + } + for _, binding := range bindings { + got, err := bearerToken(t.Context(), service, tenant, attached, binding) + if err != nil || got != token { + t.Fatal("existing selection did not read the exact committed replacement") + } + } + prior = current + } + before, err := store.GetCredential(t.Context(), tenant, original.VaultID, original.ID) + if err != nil { + t.Fatal(err) + } + assertUnchanged := func() { + t.Helper() + after, err := store.GetCredential(t.Context(), tenant, original.VaultID, original.ID) + if err != nil || !reflect.DeepEqual(after, before) || !bytes.Equal(readCiphertext(original.ID), prior) { + t.Fatal("failed replacement changed the existing row") + } + } + for _, scope := range []struct{ tenant, vault, id string }{ + {foreign, original.VaultID, original.ID}, {tenant, owned[1].ID, original.ID}, + {tenant, owned[2].ID, original.ID}, {tenant, original.VaultID, uuid.NewString()}, {tenant, "invalid", original.ID}, + } { + if _, err := update(service, scope.tenant, scope.vault, scope.id, "rejected"); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("unowned or invalid replacement was admitted") + } + assertUnchanged() + } + for _, unusable := range []*credentialcrypto.Cipher{nil, {}} { + if _, err := update(newService(t, store, unusable, nil), tenant, original.VaultID, original.ID, "rejected"); err == nil { + t.Fatal("missing or unusable cipher admitted replacement") + } + assertUnchanged() + } + // A real PostgreSQL mutation failure must preserve both ciphertext and time. + _, updateErr := update(newService(t, readOnlyStore(t, pool), cipher, nil), tenant, original.VaultID, original.ID, "rejected") + if !isReadOnlyFailure(updateErr) { + t.Fatal("database write failure was accepted or translated", updateErr) + } + assertUnchanged() + // A stale destination from a prior metadata read cannot authorize the write. + _, err = store.ReplaceStaticToken(t.Context(), vaults.StaticTokenReplacement{CredentialKey: vaults.CredentialKey{TenantID: tenant, VaultID: original.VaultID, CredentialID: original.ID}, MCPServerURL: endpoint + "/other", Ciphertext: prior}) + if !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("mutation failed to recheck immutable destination", err) + } + assertUnchanged() + // Replacing a damaged old payload needs no old-token decryption. + if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET token_ciphertext=set_byte(token_ciphertext, 15, get_byte(token_ciphertext,15) # 1) WHERE id=$1", original.ID); err != nil { + t.Fatal(err) + } + if _, err := update(service, tenant, original.VaultID, original.ID, lastToken); err != nil { + t.Fatal("replacement tried to decrypt the old token") + } + pool.Close() + store, pool = openStore(t) + service = newService(t, store, newCipher(t, bytes.Clone(key)), nil) + for _, binding := range bindings { + current, err := bearerToken(t.Context(), service, tenant, attached, binding) + if err != nil || current != lastToken { + t.Fatal("reopened dispatch lookup lost the replacement") + } + } + // Competing whole-secret replacements may win in either order, never tear. + left, right := uuid.NewString()+strings.Repeat("L", 513), uuid.NewString()+strings.Repeat("R", 1025) + start, results := make(chan struct{}), make(chan error, 2) + for _, token := range []string{left, right} { + go func() { + <-start + _, err := update(service, tenant, original.VaultID, original.ID, token) + results <- err + }() + } + close(start) + for range 2 { + if err := <-results; err != nil { + t.Fatal("concurrent replacement failed") + } + } + current, err := bearerToken(t.Context(), service, tenant, attached, bindings[0]) + if err != nil || current != left && current != right { + t.Fatal("concurrent replacements produced an incomplete secret") + } + if !bytes.Equal(readCiphertext(unrelated.ID), unrelatedCiphertext) { + t.Fatal("replacement changed an unrelated Credential") + } +} + +func TestCredentialDeletionScopeBindingAndRestart(t *testing.T) { + store, pool := openStore(t) + tenant, foreign := uuid.NewString(), uuid.NewString() + key := bytes.Repeat([]byte{41}, 32) + service := newService(t, store, newCipher(t, key), nil) + keyless := newService(t, store, nil, nil) + vault, wrong := createVault(t, service, tenant), createVault(t, service, tenant) + original := createStatic(t, service, tenant, vault.ID, "original", "https://mcp.example/tools", "original-secret") + attached := []string{vault.ID} + selected := resolve(t, service, tenant, attached, vaults.MCPCredentialRequest{ServerLabel: "tools", ServerURL: original.MCPServerURL}) + retained, err := bearerToken(t.Context(), service, tenant, attached, selected[0]) + if err != nil || retained != "original-secret" { + t.Fatal("pre-delete dispatch lookup failed") + } + sibling := createStatic(t, service, tenant, vault.ID, "sibling", "https://mcp.example/tools", "sibling-secret") + remove := func(service *vaults.Service, tenant, vault, id string) (string, error) { + return service.DeleteCredential(t.Context(), vaults.DeleteCredential{TenantID: tenant, VaultID: vault, CredentialID: id}) + } + for _, scope := range []struct{ tenant, vault, id string }{ + {foreign, vault.ID, original.ID}, {tenant, wrong.ID, original.ID}, + {tenant, vault.ID, uuid.NewString()}, {tenant, "invalid", original.ID}, {tenant, vault.ID, "invalid"}, + } { + if _, err := remove(keyless, scope.tenant, scope.vault, scope.id); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("foreign or invalid delete was accepted", err) + } + } + // An actual database write failure must leave the resource and token intact. + _, deletionErr := remove(newService(t, readOnlyStore(t, pool), nil, nil), tenant, vault.ID, original.ID) + if !isReadOnlyFailure(deletionErr) { + t.Fatal("failed mutation was accepted or translated", deletionErr) + } + if value, err := store.GetCredential(t.Context(), tenant, vault.ID, original.ID); err != nil || !reflect.DeepEqual(value, original) { + t.Fatal("rejected deletion changed the resource", err) + } + if token, err := bearerToken(t.Context(), service, tenant, attached, selected[0]); err != nil || token != retained { + t.Fatal("rejected deletion changed the stored token") + } + // Delete without a key, even if the stored payload is damaged. + if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET token_ciphertext=decode('00','hex') WHERE id=$1", original.ID); err != nil { + t.Fatal(err) + } + if id, err := remove(keyless, tenant, vault.ID, original.ID); err != nil || id != original.ID { + t.Fatal("keyless deletion failed", err) + } + var count int + if err := pool.QueryRow(t.Context(), "SELECT count(*) FROM vault_credentials WHERE id=$1", original.ID).Scan(&count); err != nil || count != 0 { + t.Fatal("deleted row or ciphertext remains") + } + pool.Close() + store, _ = openStore(t) + service = newService(t, store, newCipher(t, bytes.Clone(key)), nil) + if _, err := remove(service, tenant, vault.ID, original.ID); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("repeat deletion did not stay absent") + } + if _, err := store.GetCredential(t.Context(), tenant, vault.ID, original.ID); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("deleted metadata reappeared after restart") + } + if _, err := service.UpdateStaticCredential(t.Context(), vaults.UpdateStaticCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: original.ID, Token: "replacement"}); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("replacement resurrected a deleted credential") + } + if _, err := bearerToken(t.Context(), service, tenant, attached, selected[0]); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("frozen selection fell back to another token") + } + _, err = service.ResolveMCPCredentials(t.Context(), vaults.ResolveMCPCredentials{TenantID: tenant, VaultIDs: attached, Requests: []vaults.MCPCredentialRequest{{ServerLabel: "tools", ServerURL: original.MCPServerURL, CredentialID: &original.ID}}}) + if !isSelectionError(err, false, "MCP credential_id "+original.ID+" was not found in an attached vault") { + t.Fatal("deleted explicit selection was admitted", err) + } + page, err := store.ListCredentials(t.Context(), tenant, vault.ID, vaults.PageQuery{Limit: 100, Ascending: true}) + if err != nil || len(page.Credentials) != 1 || !reflect.DeepEqual(page.Credentials[0], sibling) { + t.Fatal("deletion changed a sibling or list membership", err) + } + if value, err := store.GetVault(t.Context(), tenant, vault.ID); err != nil || !reflect.DeepEqual(value, vault) { + t.Fatal("deletion changed its parent Vault") + } +} + +func TestCredentialDeletionConcurrentReplacementCannotResurrect(t *testing.T) { + store, _ := openStore(t) + tenant := uuid.NewString() + service := newService(t, store, newCipher(t, bytes.Repeat([]byte{42}, 32)), nil) + vault := createVault(t, service, tenant) + for range 8 { + value := createStatic(t, service, tenant, vault.ID, "competing", "https://mcp.example/tools", "before") + start, updated := make(chan struct{}), make(chan error, 1) + go func() { + <-start + _, err := service.UpdateStaticCredential(t.Context(), vaults.UpdateStaticCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: value.ID, Token: "after"}) + updated <- err + }() + close(start) + _, deleted := service.DeleteCredential(t.Context(), vaults.DeleteCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: value.ID}) + updateErr := <-updated + if deleted != nil || updateErr != nil && !errors.Is(updateErr, vaults.ErrNotFound) { + t.Fatal("competing update/delete failed unexpectedly", deleted, updateErr) + } + if _, err := store.GetCredential(t.Context(), tenant, vault.ID, value.ID); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("concurrent update resurrected deleted metadata") + } + } +} diff --git a/services/core/internal/persistence/postgres/vaultpg/fixture_test.go b/services/core/internal/persistence/postgres/vaultpg/fixture_test.go new file mode 100644 index 000000000..e0fdda7fa --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/fixture_test.go @@ -0,0 +1,144 @@ +package vaultpg_test + +import ( + "context" + "encoding/json" + "errors" + "testing" + + "github.com/jackc/pgx/v5/pgconn" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/vaultpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +// openStore opens the shared test database, as a restarted Core would. +func openStore(t *testing.T) (*vaultpg.Store, *pgxpool.Pool) { + t.Helper() + pool := pgtest.Open(t) + return vaultpg.New(pgunit.NewPool(pool)), pool +} + +// readOnlyStore fails every write with a real PostgreSQL error. +func readOnlyStore(t *testing.T, pool *pgxpool.Pool) *vaultpg.Store { + t.Helper() + config := pool.Config().Copy() + config.ConnConfig.RuntimeParams["default_transaction_read_only"] = "on" + readOnly, err := pgxpool.NewWithConfig(t.Context(), config) + if err != nil { + t.Fatal(err) + } + t.Cleanup(readOnly.Close) + return vaultpg.New(pgunit.NewPool(readOnly)) +} + +// isReadOnlyFailure reports the unexpected failure of a write on a +// readOnlyStore, which the Store returns as is. +func isReadOnlyFailure(err error) bool { + var failure *pgconn.PgError + return errors.As(err, &failure) && failure.Code == "25006" +} + +func newCipher(t *testing.T, key []byte) *credentialcrypto.Cipher { + t.Helper() + cipher, err := credentialcrypto.New(key) + if err != nil { + t.Fatal(err) + } + return cipher +} + +// newService serves storage; a nil cipher is a Core without a credential key. +func newService(t *testing.T, storage vaults.Storage, cipher *credentialcrypto.Cipher, refresher oauthrefresh.Refresher) *vaults.Service { + t.Helper() + if refresher == nil { + refresher = noRefresh(t) + } + service, err := vaults.NewService(storage, cipher, refresher) + if err != nil { + t.Fatal(err) + } + return service +} + +type refreshFunc func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) + +func (f refreshFunc) Refresh(ctx context.Context, request oauthrefresh.Request) (oauthrefresh.Token, error) { + return f(ctx, request) +} + +// noRefresh fails the test if an exchange is attempted. +func noRefresh(t *testing.T) refreshFunc { + return func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + t.Error("unexpected OAuth refresh") + return oauthrefresh.Token{}, errors.New("unexpected OAuth refresh") + } +} + +func createVault(t *testing.T, service *vaults.Service, tenant string) vaults.Vault { + t.Helper() + vault, err := service.CreateVault(t.Context(), vaults.CreateVault{TenantID: tenant}) + if err != nil { + t.Fatal(err) + } + return vault +} + +func createStatic(t *testing.T, service *vaults.Service, tenant, vault, name, url, token string) vaults.Credential { + t.Helper() + credential, err := service.CreateStaticCredential(t.Context(), vaults.CreateStaticCredential{TenantID: tenant, VaultID: vault, Name: name, MCPServerURL: url, Token: token}) + if err != nil { + t.Fatal(err) + } + return credential +} + +func resolve(t *testing.T, service *vaults.Service, tenant string, attached []string, requests ...vaults.MCPCredentialRequest) []vaults.MCPCredentialBinding { + t.Helper() + bindings, err := service.ResolveMCPCredentials(t.Context(), vaults.ResolveMCPCredentials{TenantID: tenant, VaultIDs: attached, Requests: requests}) + if err != nil || len(bindings) != len(requests) { + t.Fatal("selection failed", err) + } + return bindings +} + +func bearerToken(ctx context.Context, service *vaults.Service, tenant string, attached []string, binding vaults.MCPCredentialBinding) (string, error) { + return service.MCPBearerToken(ctx, vaults.MCPBearerToken{TenantID: tenant, VaultIDs: attached, Binding: binding}) +} + +func isSelectionError(err error, conflict bool, message string) bool { + var selection *vaults.MCPCredentialSelectionError + return errors.As(err, &selection) && selection.Conflict == conflict && selection.Message == message +} + +// storedGrant is the sealed plaintext of an mcp_oauth Credential. +type storedGrant struct { + Version int `json:"version"` + Metadata vaults.OAuthMetadata `json:"metadata"` + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + ClientSecret string `json:"client_secret"` +} + +// readGrant opens a stored OAuth grant independently of the service. +func readGrant(t *testing.T, pool *pgxpool.Pool, cipher *credentialcrypto.Cipher, tenant string, credential vaults.Credential) storedGrant { + t.Helper() + var ciphertext []byte + if err := pool.QueryRow(t.Context(), "SELECT token_ciphertext FROM vault_credentials WHERE id=$1", credential.ID).Scan(&ciphertext); err != nil { + t.Fatal(err) + } + plaintext, err := cipher.Open(ciphertext, credentialcrypto.Binding{TenantID: tenant, VaultID: credential.VaultID, CredentialID: credential.ID, AuthType: vaults.AuthMCPOAuth, Destination: credential.MCPServerURL}) + if err != nil { + t.Fatal("open stored grant", err) + } + var grant storedGrant + if err := json.Unmarshal(plaintext, &grant); err != nil || grant.Version != 1 { + t.Fatal("decode stored grant", err) + } + return grant +} diff --git a/services/core/internal/persistence/postgres/vaultpg/oauth.go b/services/core/internal/persistence/postgres/vaultpg/oauth.go new file mode 100644 index 000000000..13af0461f --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/oauth.go @@ -0,0 +1,68 @@ +package vaultpg + +import ( + "context" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgtype" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +// WithOAuthCredential holds the Credential's row lock, once loaded, until +// apply returns: through an external refresh too. +func (s *Store) WithOAuthCredential(ctx context.Context, key vaults.CredentialKey, apply func(vaults.OAuthTx) error) error { + tx := &oauthTx{tenantID: key.TenantID, tenant: pgunit.PathID(key.TenantID), vault: pgunit.PathID(key.VaultID), id: pgunit.PathID(key.CredentialID)} + return translate(s.pool.Transaction(ctx, func(ctx context.Context, t pgx.Tx) error { + tx.q = sqlc.New(t) + return apply(tx) + })) +} + +type oauthTx struct { + q *sqlc.Queries + tenantID string + tenant, vault, id pgtype.UUID +} + +func (t *oauthTx) LoadOAuthCredential(ctx context.Context) (vaults.Credential, []byte, error) { + row, err := t.q.GetOAuthCredentialForUpdate(ctx, sqlc.GetOAuthCredentialForUpdateParams{TenantID: t.tenant, VaultID: t.vault, ID: t.id}) + if err != nil { + return vaults.Credential{}, nil, translate(err) + } + credential, err := credentialFromRow(sqlc.GetCredentialRow{ID: row.ID, VaultID: row.VaultID, Name: row.Name, AuthType: row.AuthType, + McpServerUrl: row.McpServerUrl, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt, OauthMetadata: row.OauthMetadata}) + if err != nil { + return vaults.Credential{}, nil, err + } + return credential, row.TokenCiphertext, nil +} + +func (t *oauthTx) ApplyOAuthRefresh(ctx context.Context, sealed vaults.SealedOAuth) error { + _, err := t.update(ctx, sealed) + return err +} + +func (t *oauthTx) ApplyOAuthReplacement(ctx context.Context, sealed vaults.SealedOAuth) (vaults.Credential, error) { + updated, err := t.update(ctx, sealed) + if err != nil { + return vaults.Credential{}, err + } + if err := auditpg.RecordWriteAudit(ctx, t.q, t.tenantID, "update", "credential", updated.ID, updated.VaultID); err != nil { + return vaults.Credential{}, translate(err) + } + return updated, nil +} + +// update matches the destination the grant is sealed to. +func (t *oauthTx) update(ctx context.Context, sealed vaults.SealedOAuth) (vaults.Credential, error) { + row, err := t.q.UpdateOAuthCredential(ctx, sqlc.UpdateOAuthCredentialParams{TenantID: t.tenant, VaultID: t.vault, ID: t.id, + McpServerUrl: sealed.MCPServerURL, OauthMetadata: sealed.Metadata, TokenCiphertext: sealed.Ciphertext}) + if err != nil { + return vaults.Credential{}, translate(err) + } + return credentialFromRow(sqlc.GetCredentialRow(row)) +} diff --git a/services/core/internal/persistence/postgres/vaultpg/oauth_test.go b/services/core/internal/persistence/postgres/vaultpg/oauth_test.go new file mode 100644 index 000000000..41463c88b --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/oauth_test.go @@ -0,0 +1,564 @@ +package vaultpg_test + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "reflect" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/vaultpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +// oauthFixture is one tenant's Vault served with a fixed credential key. Its +// command creates an expired grant with refresh configuration. +type oauthFixture struct { + store *vaultpg.Store + pool *pgxpool.Pool + cipher *credentialcrypto.Cipher + service *vaults.Service + tenant string + vault vaults.Vault + command vaults.CreateOAuthCredential +} + +// newOAuthFixture fails the test on any exchange when refresher is nil. +func newOAuthFixture(t *testing.T, refresher oauthrefresh.Refresher) *oauthFixture { + t.Helper() + store, pool := openStore(t) + cipher := newCipher(t, bytes.Repeat([]byte{17}, 32)) + service := newService(t, store, cipher, refresher) + tenant := uuid.NewString() + vault := createVault(t, service, tenant) + return &oauthFixture{store: store, pool: pool, cipher: cipher, service: service, tenant: tenant, vault: vault, + command: vaults.CreateOAuthCredential{TenantID: tenant, VaultID: vault.ID, Name: "OAuth fixture", MCPServerURL: "https://mcp.example/tools", + AccessToken: "private-access-canary", RefreshToken: "private-refresh-canary", ClientSecret: "private-client-canary", + OAuth: vaults.OAuthMetadata{ExpiresAt: ptr(time.Now().Add(-time.Hour).UTC().Format(time.RFC3339Nano)), + Refresh: &vaults.OAuthRefreshMetadata{ClientID: "test-client", TokenEndpoint: "https://issuer.example/token", + TokenEndpointAuth: "client_secret_basic", Resource: ptr("https://mcp.example/tools"), Scope: ptr("read write")}}}} +} + +// otherService serves the fixture's database through a separate Store, so +// competing callers meet only in PostgreSQL. +func (f *oauthFixture) otherService(t *testing.T, cipher *credentialcrypto.Cipher, refresher oauthrefresh.Refresher) *vaults.Service { + t.Helper() + return newService(t, vaultpg.New(pgunit.NewPool(f.pool)), cipher, refresher) +} + +func (f *oauthFixture) create(t *testing.T, command vaults.CreateOAuthCredential) vaults.Credential { + t.Helper() + credential, err := f.service.CreateOAuthCredential(t.Context(), command) + if err != nil { + t.Fatal("create OAuth fixture", err) + } + return credential +} + +func (f *oauthFixture) token(ctx context.Context, service *vaults.Service, credential vaults.Credential) (string, error) { + return bearerToken(ctx, service, f.tenant, []string{f.vault.ID}, oauthBinding(credential)) +} + +func (f *oauthFixture) update(ctx context.Context, service *vaults.Service, credential vaults.Credential, patch vaults.UpdateOAuthCredential) (vaults.Credential, error) { + patch.TenantID, patch.VaultID, patch.CredentialID = f.tenant, credential.VaultID, credential.ID + return service.UpdateOAuthCredential(ctx, patch) +} + +func (f *oauthFixture) grant(t *testing.T, credential vaults.Credential) storedGrant { + t.Helper() + return readGrant(t, f.pool, f.cipher, f.tenant, credential) +} + +func oauthBinding(credential vaults.Credential) vaults.MCPCredentialBinding { + return vaults.MCPCredentialBinding{ServerLabel: "test", ServerURL: credential.MCPServerURL, + VaultID: credential.VaultID, CredentialID: credential.ID, AuthType: credential.AuthType} +} + +func ptr[T any](value T) *T { return &value } + +func TestOAuthCredentialMetadataEncryptionAndScope(t *testing.T) { + var exchanges atomic.Int32 + counting := refreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + exchanges.Add(1) + return oauthrefresh.Token{}, errors.New("unexpected exchange") + }) + f := newOAuthFixture(t, counting) + input := f.command + credential := f.create(t, input) + got, err := f.store.GetCredential(t.Context(), f.tenant, f.vault.ID, credential.ID) + if err != nil || !reflect.DeepEqual(got, credential) { + t.Fatal("keyless safe metadata changed", err) + } + page, err := f.store.ListCredentials(t.Context(), f.tenant, f.vault.ID, vaults.PageQuery{Limit: 20, Ascending: true}) + if err != nil || len(page.Credentials) != 1 || !reflect.DeepEqual(page.Credentials[0], credential) { + t.Fatal("keyless listing failed", err) + } + encoded, _ := json.Marshal(page) + var ciphertext []byte + if err := f.pool.QueryRow(t.Context(), "SELECT token_ciphertext FROM vault_credentials WHERE id=$1", credential.ID).Scan(&ciphertext); err != nil { + t.Fatal(err) + } + for _, secret := range []string{input.AccessToken, input.RefreshToken, input.ClientSecret} { + if bytes.Contains(encoded, []byte(secret)) || bytes.Contains(ciphertext, []byte(secret)) { + t.Fatal("metadata disclosed a secret or plaintext grant persisted") + } + } + if grant := readGrant(t, f.pool, newCipher(t, bytes.Repeat([]byte{17}, 32)), f.tenant, credential); grant.AccessToken != input.AccessToken || grant.RefreshToken != input.RefreshToken || grant.ClientSecret != input.ClientSecret { + t.Fatal("restart lost grant material") + } + binding := oauthBinding(credential) + for _, target := range []struct { + tenant string + vaults []string + binding vaults.MCPCredentialBinding + }{ + {uuid.NewString(), []string{f.vault.ID}, binding}, {f.tenant, nil, binding}, + {f.tenant, []string{f.vault.ID}, vaults.MCPCredentialBinding{ServerLabel: "test", ServerURL: input.MCPServerURL + "/other", VaultID: f.vault.ID, CredentialID: credential.ID, AuthType: vaults.AuthMCPOAuth}}, + } { + if token, err := bearerToken(t.Context(), f.service, target.tenant, target.vaults, target.binding); !errors.Is(err, vaults.ErrNotFound) || token != "" { + t.Fatal("foreign or mismatched scope admitted", err) + } + } + if _, err := f.service.UpdateOAuthCredential(t.Context(), vaults.UpdateOAuthCredential{TenantID: uuid.NewString(), VaultID: f.vault.ID, CredentialID: credential.ID, AccessToken: ptr("replacement")}); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("foreign update admitted") + } + foreign := input + foreign.TenantID = uuid.NewString() + if _, err := f.service.CreateOAuthCredential(t.Context(), foreign); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("foreign creation admitted") + } + keyless := f.otherService(t, nil, counting) + if _, err := keyless.CreateOAuthCredential(t.Context(), input); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("keyless creation admitted") + } + if token, err := f.token(t.Context(), keyless, credential); !errors.Is(err, credentialcrypto.ErrUnavailable) || token != "" { + t.Fatal("keyless execution admitted") + } + wrong := f.otherService(t, newCipher(t, bytes.Repeat([]byte{18}, 32)), counting) + if token, err := f.token(t.Context(), wrong, credential); err == nil || token != "" { + t.Fatal("wrong key executed") + } + if exchanges.Load() != 0 { + t.Fatal("resource operations or invalid scopes contacted provider") + } +} + +func TestOAuthCredentialUpdatesPreservePinnedSemantics(t *testing.T) { + f := newOAuthFixture(t, nil) + input := f.command + credential := f.create(t, input) + update := func(patch vaults.UpdateOAuthCredential) vaults.Credential { + t.Helper() + got, err := f.update(t.Context(), f.service, credential, patch) + if err != nil { + t.Fatal(err) + } + if got.ID != credential.ID || got.Name != credential.Name || got.AuthType != credential.AuthType || got.MCPServerURL != credential.MCPServerURL || !got.CreatedAt.Equal(credential.CreatedAt) { + t.Fatal("immutable metadata changed") + } + return got + } + unchanged := update(vaults.UpdateOAuthCredential{}) + if !reflect.DeepEqual(unchanged.OAuth, credential.OAuth) { + t.Fatal("omitted values changed") + } + replacement := update(vaults.UpdateOAuthCredential{AccessToken: ptr("new-access")}) + if replacement.OAuth.ExpiresAt != nil { + t.Fatal("new access token retained old expiry") + } + expiry := time.Now().Add(time.Hour).UTC().Format(time.RFC3339Nano) + replacement = update(vaults.UpdateOAuthCredential{ExpiresAtSet: true, ExpiresAt: &expiry, Refresh: &vaults.OAuthRefreshUpdate{ScopeSet: true, Scope: nil}}) + if replacement.OAuth.ExpiresAt == nil || *replacement.OAuth.ExpiresAt != expiry || replacement.OAuth.Refresh.Scope != nil { + t.Fatal("expiry or null scope semantics failed") + } + grant := f.grant(t, replacement) + if grant.AccessToken != "new-access" || grant.RefreshToken != input.RefreshToken || grant.ClientSecret != input.ClientSecret { + t.Fatal("omitted secrets were replaced") + } + replacement = update(vaults.UpdateOAuthCredential{ExpiresAtSet: true, Refresh: &vaults.OAuthRefreshUpdate{RefreshToken: ptr("new-refresh"), TokenEndpointAuthType: "client_secret_basic", ClientSecret: ptr("new-secret"), ScopeSet: true, Scope: ptr("read")}}) + grant = f.grant(t, replacement) + if grant.Metadata.ExpiresAt != nil || grant.RefreshToken != "new-refresh" || grant.ClientSecret != "new-secret" || *grant.Metadata.Refresh.Scope != "read" { + t.Fatal("replacement fields not persisted") + } + for _, patch := range []vaults.UpdateOAuthCredential{ + {ExpiresAtSet: true, ExpiresAt: ptr("not-a-date")}, + {Refresh: &vaults.OAuthRefreshUpdate{TokenEndpointAuthType: "client_secret_post"}}, + } { + if _, err := f.update(t.Context(), f.service, credential, patch); !errors.Is(err, vaults.ErrInvalidInput) { + t.Fatal("invalid mutation admitted") + } + } + if got := f.grant(t, replacement); !reflect.DeepEqual(got, grant) { + t.Fatal("rejected update changed the grant") + } + input.OAuth.Refresh = nil + input.RefreshToken = "" + input.ClientSecret = "" + withoutRefresh := f.create(t, input) + if _, err := f.update(t.Context(), f.service, withoutRefresh, vaults.UpdateOAuthCredential{Refresh: &vaults.OAuthRefreshUpdate{RefreshToken: ptr("cannot-add")}}); !errors.Is(err, vaults.ErrInvalidInput) { + t.Fatal("added missing refresh configuration") + } +} + +func TestOAuthMetadataTamperingNeverReachesProvider(t *testing.T) { + var exchanges atomic.Int32 + f := newOAuthFixture(t, refreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + exchanges.Add(1) + return oauthrefresh.Token{}, errors.New("unexpected refresh") + })) + for _, mutation := range []string{ + `jsonb_set(oauth_metadata,'{refresh,token_endpoint}','"https://attacker.example/token"')`, + `jsonb_set(oauth_metadata,'{refresh,client_id}','"other-client"')`, + `jsonb_set(oauth_metadata,'{refresh,scope}','"all"')`, + `jsonb_set(oauth_metadata,'{expires_at}','null')`, + } { + credential := f.create(t, f.command) + if _, err := f.pool.Exec(t.Context(), "UPDATE vault_credentials SET oauth_metadata="+mutation+" WHERE id=$1", credential.ID); err != nil { + t.Fatal(err) + } + if token, err := f.token(t.Context(), f.service, credential); err == nil || token != "" { + t.Fatal("metadata substitution executed") + } + if _, err := f.update(t.Context(), f.service, credential, vaults.UpdateOAuthCredential{AccessToken: ptr("new")}); err == nil { + t.Fatal("update authenticated substituted metadata") + } + } + if exchanges.Load() != 0 { + t.Fatal("tampered metadata reached token endpoint") + } +} + +func TestOAuthCredentialSelectionIncludesBothAuthTypes(t *testing.T) { + f := newOAuthFixture(t, nil) + credential := f.create(t, f.command) + requests := []vaults.MCPCredentialRequest{{ServerLabel: "test", ServerURL: f.command.MCPServerURL}} + bindings := resolve(t, f.service, f.tenant, []string{f.vault.ID}, requests...) + if bindings[0].AuthType != vaults.AuthMCPOAuth || bindings[0].CredentialID != credential.ID { + t.Fatal("OAuth was not selected") + } + createStatic(t, f.service, f.tenant, f.vault.ID, "Static", f.command.MCPServerURL, "static") + _, err := f.service.ResolveMCPCredentials(t.Context(), vaults.ResolveMCPCredentials{TenantID: f.tenant, VaultIDs: []string{f.vault.ID}, Requests: requests}) + if !isSelectionError(err, true, "multiple attached vault credentials match MCP server_url "+f.command.MCPServerURL+"; specify credential_id") { + t.Fatal("ambiguous mixed credentials selected", err) + } + requests[0].CredentialID = &credential.ID + if selected := resolve(t, f.service, f.tenant, []string{f.vault.ID}, requests...); selected[0] != bindings[0] { + t.Fatal("explicit OAuth identity changed") + } +} + +func TestOAuthRefreshErrorsAreSafeAndPreserveGrant(t *testing.T) { + for _, mode := range []string{"provider_error", "empty_access", "expired_response", "no_refresh"} { + t.Run(mode, func(t *testing.T) { + f := newOAuthFixture(t, refreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + switch mode { + case "provider_error": + return oauthrefresh.Token{}, errors.New("private-refresh-canary: provider body") + case "expired_response": + past := time.Now().Add(-time.Second) + return oauthrefresh.Token{AccessToken: "new", ExpiresAt: &past}, nil + } + return oauthrefresh.Token{}, nil + })) + input := f.command + if mode == "no_refresh" { + input.OAuth.Refresh = nil + input.RefreshToken = "" + input.ClientSecret = "" + } + credential := f.create(t, input) + before := f.grant(t, credential) + token, err := f.token(t.Context(), f.service, credential) + if err == nil || token != "" || strings.Contains(err.Error(), "private-refresh-canary") { + t.Fatal("refresh failed unsafely") + } + if after := f.grant(t, credential); !reflect.DeepEqual(before, after) { + t.Fatal("failed refresh changed grant") + } + }) + } +} + +func TestOAuthRefreshPersistsRotatedGrantAndRequest(t *testing.T) { + for _, method := range []string{"none", "client_secret_basic", "client_secret_post"} { + t.Run(method, func(t *testing.T) { + var request oauthrefresh.Request + expiry := time.Now().Add(time.Hour).UTC() + f := newOAuthFixture(t, refreshFunc(func(_ context.Context, r oauthrefresh.Request) (oauthrefresh.Token, error) { + request = r + return oauthrefresh.Token{AccessToken: "renewed-access", RefreshToken: "rotated-refresh", ExpiresAt: &expiry}, nil + })) + input := f.command + input.OAuth.Refresh.TokenEndpointAuth = method + if method == "none" { + input.ClientSecret = "" + } + credential := f.create(t, input) + got, err := f.token(t.Context(), f.service, credential) + if err != nil || got != "renewed-access" { + t.Fatal("refresh did not return committed access", err) + } + want := oauthrefresh.Request{TokenEndpoint: input.OAuth.Refresh.TokenEndpoint, ClientID: input.OAuth.Refresh.ClientID, AuthMethod: method, ClientSecret: input.ClientSecret, RefreshToken: input.RefreshToken, Resource: input.OAuth.Refresh.Resource, Scope: input.OAuth.Refresh.Scope} + if !reflect.DeepEqual(request, want) { + t.Fatal("refresh request lost grant fields") + } + restartedStore, restartedPool := openStore(t) + after := readGrant(t, restartedPool, f.cipher, f.tenant, credential) + if after.AccessToken != "renewed-access" || after.RefreshToken != "rotated-refresh" || after.Metadata.ExpiresAt == nil || *after.Metadata.ExpiresAt != expiry.Format(time.RFC3339Nano) { + t.Fatal("refreshed grant not durable") + } + restarted := newService(t, restartedStore, f.cipher, nil) + got, err = f.token(t.Context(), restarted, credential) + if err != nil || got != "renewed-access" { + t.Fatal("fresh grant needed another refresh after restart", err) + } + metadata, err := restartedStore.GetCredential(t.Context(), f.tenant, f.vault.ID, credential.ID) + if err != nil || !reflect.DeepEqual(metadata.OAuth, &after.Metadata) { + t.Fatal("safe expiry metadata did not follow refresh", err) + } + }) + } +} + +func TestOAuthRefreshPreservesRefreshTokenWhenOmitted(t *testing.T) { + f := newOAuthFixture(t, refreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + return oauthrefresh.Token{AccessToken: "renewed-access"}, nil + })) + credential := f.create(t, f.command) + if _, err := f.token(t.Context(), f.service, credential); err != nil { + t.Fatal(err) + } + if grant := f.grant(t, credential); grant.RefreshToken != f.command.RefreshToken || grant.Metadata.ExpiresAt != nil { + t.Fatal("omitted refresh token or unknown expiry changed incorrectly") + } +} + +func TestOAuthFreshAndUnknownExpiryDoNotRefresh(t *testing.T) { + f := newOAuthFixture(t, nil) + input := f.command + for _, expiry := range []*string{nil, ptr(time.Now().Add(time.Hour).UTC().Format(time.RFC3339Nano))} { + input.OAuth.ExpiresAt = expiry + credential := f.create(t, input) + if token, err := f.token(t.Context(), f.service, credential); err != nil || token != input.AccessToken { + t.Fatal("usable token was not returned", err) + } + } + input.AccessToken = "" + input.OAuth.ExpiresAt = nil + credential := f.create(t, input) + if token, err := f.token(t.Context(), f.service, credential); err == nil || token != "" { + t.Fatal("empty access token admitted") + } +} + +func TestOAuthConcurrentRefreshUsesOneCommittedGrant(t *testing.T) { + var requests atomic.Int32 + expiry := time.Now().Add(time.Hour) + refresher := refreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + requests.Add(1) + return oauthrefresh.Token{AccessToken: "concurrent-access", RefreshToken: "single-use-next", ExpiresAt: &expiry}, nil + }) + f := newOAuthFixture(t, refresher) + credential := f.create(t, f.command) + // Separate Stores and services exercise PostgreSQL serialization, not a + // local lock. + var callers []*vaults.Service + for range 12 { + callers = append(callers, f.otherService(t, f.cipher, refresher)) + } + var workers sync.WaitGroup + errorsFound := make(chan error, len(callers)) + start := make(chan struct{}) + for _, caller := range callers { + workers.Go(func() { + <-start + token, err := f.token(t.Context(), caller, credential) + if err == nil && token != "concurrent-access" { + err = errors.New("concurrent lookup returned stale token") + } + errorsFound <- err + }) + } + close(start) + workers.Wait() + close(errorsFound) + for err := range errorsFound { + if err != nil { + t.Fatal(err) + } + } + if requests.Load() != 1 { + t.Fatal("one expiring grant triggered duplicate provider exchanges") + } +} + +func TestOAuthRefreshSerializesReplacementAndDeletion(t *testing.T) { + for _, mutation := range []string{"replacement", "credential-delete", "vault-delete"} { + t.Run(mutation, func(t *testing.T) { + entered := make(chan struct{}) + release := make(chan struct{}) + var releaseOnce sync.Once + unblock := func() { releaseOnce.Do(func() { close(release) }) } + defer unblock() + expiry := time.Now().Add(time.Hour) + f := newOAuthFixture(t, refreshFunc(func(ctx context.Context, _ oauthrefresh.Request) (oauthrefresh.Token, error) { + close(entered) + select { + case <-release: + return oauthrefresh.Token{AccessToken: "refreshed-access", RefreshToken: "rotated-refresh", ExpiresAt: &expiry}, nil + case <-ctx.Done(): + return oauthrefresh.Token{}, ctx.Err() + } + })) + credential := f.create(t, f.command) + other := f.otherService(t, f.cipher, nil) + refreshed := make(chan error, 1) + go func() { + _, err := f.token(t.Context(), f.service, credential) + refreshed <- err + }() + select { + case <-entered: + case <-time.After(3 * time.Second): + t.Fatal("refresh never reached provider") + } + mutated := make(chan error, 1) + go func() { + var err error + switch mutation { + case "replacement": + _, err = f.update(t.Context(), other, credential, vaults.UpdateOAuthCredential{AccessToken: ptr("manual-access"), Refresh: &vaults.OAuthRefreshUpdate{RefreshToken: ptr("manual-refresh")}}) + case "credential-delete": + _, err = other.DeleteCredential(t.Context(), vaults.DeleteCredential{TenantID: f.tenant, VaultID: f.vault.ID, CredentialID: credential.ID}) + case "vault-delete": + _, err = other.DeleteVault(t.Context(), vaults.DeleteVault{TenantID: f.tenant, VaultID: f.vault.ID}) + } + mutated <- err + }() + select { + case <-mutated: + t.Fatal("mutation bypassed pending refresh ownership") + case <-time.After(75 * time.Millisecond): + } + unblock() + for _, completed := range []chan error{refreshed, mutated} { + select { + case err := <-completed: + if err != nil { + t.Fatal("serialized operation failed", err) + } + case <-time.After(3 * time.Second): + t.Fatal("serialized operation deadlocked") + } + } + if mutation == "replacement" { + grant := f.grant(t, credential) + if grant.AccessToken != "manual-access" || grant.RefreshToken != "manual-refresh" || grant.Metadata.ExpiresAt != nil { + t.Fatal("late refresh overwrote manual replacement") + } + return + } + if _, err := f.store.GetCredential(t.Context(), f.tenant, f.vault.ID, credential.ID); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("deleted credential resurrected") + } + if token, err := f.token(t.Context(), f.service, credential); !errors.Is(err, vaults.ErrNotFound) || token != "" { + t.Fatal("deleted grant remained usable") + } + }) + } +} + +func TestOAuthCancelledRefreshRollsBackAndAllowsReplacement(t *testing.T) { + entered := make(chan struct{}) + f := newOAuthFixture(t, refreshFunc(func(ctx context.Context, _ oauthrefresh.Request) (oauthrefresh.Token, error) { + close(entered) + <-ctx.Done() + return oauthrefresh.Token{}, ctx.Err() + })) + credential := f.create(t, f.command) + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + done := make(chan error, 1) + go func() { + _, err := f.token(ctx, f.service, credential) + done <- err + }() + select { + case <-entered: + case <-time.After(3 * time.Second): + t.Fatal("refresh not started") + } + cancel() + select { + case err := <-done: + if err == nil { + t.Fatal("cancelled refresh succeeded") + } + case <-time.After(3 * time.Second): + t.Fatal("cancel did not release refresh") + } + if _, err := f.update(t.Context(), f.service, credential, vaults.UpdateOAuthCredential{AccessToken: ptr("manual-after-cancel")}); err != nil { + t.Fatal("cancelled refresh retained ownership", err) + } + if token, err := f.token(t.Context(), f.service, credential); err != nil || token != "manual-after-cancel" { + t.Fatal("replacement after cancellation unusable") + } +} + +func TestOAuthDeletionCannotReselectAnotherCredential(t *testing.T) { + f := newOAuthFixture(t, nil) + input := f.command + input.OAuth.ExpiresAt = nil + credential := f.create(t, input) + bindings := resolve(t, f.service, f.tenant, []string{f.vault.ID}, vaults.MCPCredentialRequest{ServerLabel: "test", ServerURL: input.MCPServerURL}) + if _, err := f.service.DeleteCredential(t.Context(), vaults.DeleteCredential{TenantID: f.tenant, VaultID: f.vault.ID, CredentialID: credential.ID}); err != nil { + t.Fatal(err) + } + input.AccessToken = uuid.NewString() + f.create(t, input) + if token, err := bearerToken(t.Context(), f.service, f.tenant, []string{f.vault.ID}, bindings[0]); !errors.Is(err, vaults.ErrNotFound) || token != "" { + t.Fatal("frozen identity fell back after deletion") + } +} + +func TestOAuthRefreshCommitFailureDoesNotReturnUncommittedToken(t *testing.T) { + expiry := time.Now().Add(time.Hour) + f := newOAuthFixture(t, refreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + return oauthrefresh.Token{AccessToken: "uncommitted-access", RefreshToken: "uncommitted-refresh", ExpiresAt: &expiry}, nil + })) + credential := f.create(t, f.command) + before := f.grant(t, credential) + suffix := strings.ReplaceAll(uuid.NewString(), "-", "") + function, trigger := "oauth_commit_fail_"+suffix, "oauth_commit_fail_"+suffix + // A deferred trigger fails the commit after the UPDATE has returned + // metadata. The condition confines the failure to this test's grant. + if _, err := f.pool.Exec(t.Context(), "CREATE FUNCTION "+function+"() RETURNS trigger LANGUAGE plpgsql AS $$ BEGIN RAISE EXCEPTION 'private-refresh-canary'; END $$"); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if _, err := f.pool.Exec(context.Background(), "DROP FUNCTION "+function+"() CASCADE"); err != nil { + t.Error(err) + } + }) + if _, err := f.pool.Exec(t.Context(), "CREATE CONSTRAINT TRIGGER "+trigger+" AFTER UPDATE ON vault_credentials DEFERRABLE INITIALLY DEFERRED FOR EACH ROW WHEN (NEW.id='"+credential.ID+"'::uuid) EXECUTE FUNCTION "+function+"()"); err != nil { + t.Fatal(err) + } + token, err := f.token(t.Context(), f.service, credential) + if err == nil || token != "" || err.Error() != "OAuth credential refresh commit failed" { + t.Fatal("commit failure returned an uncommitted grant or unsafe error", err) + } + if after := f.grant(t, credential); !reflect.DeepEqual(before, after) { + t.Fatal("failed commit modified the grant") + } +} diff --git a/services/core/internal/persistence/postgres/vaultpg/selection.go b/services/core/internal/persistence/postgres/vaultpg/selection.go new file mode 100644 index 000000000..f45e2b1da --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/selection.go @@ -0,0 +1,59 @@ +package vaultpg + +import ( + "context" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgtype" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +func (s *Store) CountOwnedVaults(ctx context.Context, tenantID string, vaultIDs []string) (int, error) { + owned, err := s.pool.Queries().GetAttachedVaultIDs(ctx, sqlc.GetAttachedVaultIDsParams{TenantID: pgunit.PathID(tenantID), VaultIds: pathIDs(vaultIDs)}) + if err != nil { + return 0, translate(err) + } + return len(owned), nil +} + +func (s *Store) FindMCPCredentials(ctx context.Context, query vaults.MCPCredentialQuery) ([]vaults.MCPCredentialMatch, error) { + var credential pgtype.UUID + if query.CredentialID != "" { + credential = pgunit.PathID(query.CredentialID) + } + rows, err := s.pool.Queries().FindMCPCredentials(ctx, sqlc.FindMCPCredentialsParams{TenantID: pgunit.PathID(query.TenantID), + VaultIds: pathIDs(query.VaultIDs), McpServerUrl: query.ServerURL, CredentialID: credential}) + if err != nil { + return nil, translate(err) + } + matches := make([]vaults.MCPCredentialMatch, 0, len(rows)) + for _, row := range rows { + matches = append(matches, vaults.MCPCredentialMatch{VaultID: uuid.UUID(row.VaultID.Bytes).String(), CredentialID: uuid.UUID(row.ID.Bytes).String(), + AuthType: row.AuthType, MCPServerURL: row.McpServerUrl}) + } + return matches, nil +} + +// StaticTokenCiphertext is the only read of a static token. Resource reads +// never select ciphertext. +func (s *Store) StaticTokenCiphertext(ctx context.Context, query vaults.StaticTokenQuery) ([]byte, error) { + ciphertext, err := s.pool.Queries().GetMCPStaticCredentialCiphertext(ctx, sqlc.GetMCPStaticCredentialCiphertextParams{ + TenantID: pgunit.PathID(query.TenantID), VaultIds: pathIDs(query.VaultIDs), VaultID: pgunit.PathID(query.VaultID), + CredentialID: pgunit.PathID(query.CredentialID), McpServerUrl: query.MCPServerURL, + }) + if err != nil { + return nil, translate(err) + } + return ciphertext, nil +} + +func pathIDs(ids []string) []pgtype.UUID { + result := make([]pgtype.UUID, 0, len(ids)) + for _, id := range ids { + result = append(result, pgunit.PathID(id)) + } + return result +} diff --git a/services/core/internal/persistence/postgres/vaultpg/selection_test.go b/services/core/internal/persistence/postgres/vaultpg/selection_test.go new file mode 100644 index 000000000..5b6ee8fdd --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/selection_test.go @@ -0,0 +1,127 @@ +package vaultpg_test + +import ( + "bytes" + "crypto/rand" + "encoding/json" + "errors" + "strings" + "testing" + + "github.com/google/uuid" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +func TestMCPCredentialSelectionAndScopedDecryption(t *testing.T) { + store, pool := openStore(t) + tenant, foreign := uuid.NewString(), uuid.NewString() + key := make([]byte, 32) + if _, err := rand.Read(key); err != nil { + t.Fatal(err) + } + service := newService(t, store, newCipher(t, key), nil) + keyless := newService(t, store, nil, nil) + var owned []vaults.Vault + for _, owner := range []string{tenant, tenant, foreign} { + owned = append(owned, createVault(t, service, owner)) + } + token := " \t" + uuid.NewString() + "雪\n" + destination := "https://mcp.example/tools" + create := func(vault vaults.Vault) vaults.Credential { + t.Helper() + return createStatic(t, service, vault.TenantID, vault.ID, "private", destination, token) + } + first, foreignCredential := create(owned[0]), create(owned[2]) + attached := []string{owned[0].ID, owned[1].ID} + requests := []vaults.MCPCredentialRequest{{ServerLabel: "tools", ServerURL: destination}, {ServerLabel: "anonymous", ServerURL: "https://anonymous.example/mcp"}} + selectFor := func(service *vaults.Service, tenant string, attached []string, requests []vaults.MCPCredentialRequest) ([]vaults.MCPCredentialBinding, error) { + return service.ResolveMCPCredentials(t.Context(), vaults.ResolveMCPCredentials{TenantID: tenant, VaultIDs: attached, Requests: requests}) + } + // Selection reads metadata only, so a keyless Core can select. + bindings, err := selectFor(keyless, tenant, attached, requests) + if err != nil || len(bindings) != 2 || bindings[0].CredentialID != first.ID || bindings[0].AuthType != vaults.AuthStaticBearer || bindings[1].CredentialID != "" { + t.Fatal("metadata selection or frozen anonymous decision differs", err) + } + encoded, err := json.Marshal(bindings) + if err != nil || bytes.Contains(encoded, []byte(token)) || strings.Contains(string(encoded), "ciphertext") { + t.Fatal("private binding contains secret material") + } + second := create(owned[1]) + if _, err := selectFor(keyless, tenant, attached, requests); !isSelectionError(err, true, "multiple attached vault credentials match MCP server_url "+destination+"; specify credential_id") { + t.Fatal("ambiguous selection was admitted", err) + } + requests[0].CredentialID = &second.ID + explicit, err := selectFor(keyless, tenant, attached, requests) + if err != nil || explicit[0].CredentialID != second.ID || requests[1].CredentialID != nil { + t.Fatal("explicit selection did not disambiguate", err) + } + notAttached := func(id string) string { return "MCP credential_id " + id + " was not found in an attached vault" } + for _, tc := range []struct { + owner string + vaults []string + id, url string + message string // Empty for the unchanged Vault 404. + }{ + {tenant, attached, foreignCredential.ID, destination, notAttached(foreignCredential.ID)}, + {tenant, []string{owned[1].ID}, first.ID, destination, notAttached(first.ID)}, + {tenant, attached, "not-a-credential", destination, notAttached("not-a-credential")}, + {tenant, attached, first.ID, destination + "/other", "MCP credential_id " + first.ID + " does not match server_url " + destination + "/other"}, + {foreign, attached, first.ID, destination, ""}, + {tenant, []string{owned[0].ID, owned[2].ID}, first.ID, destination, ""}, + {tenant, []string{uuid.NewString()}, first.ID, destination, ""}, + } { + _, err := selectFor(keyless, tc.owner, tc.vaults, []vaults.MCPCredentialRequest{{ServerLabel: "tools", ServerURL: tc.url, CredentialID: &tc.id}}) + if tc.message == "" && !errors.Is(err, vaults.ErrNotFound) || tc.message != "" && !isSelectionError(err, false, tc.message) { + t.Fatal("unowned, unattached or wrong-destination selection was admitted", err) + } + } + if _, err := selectFor(keyless, tenant, nil, []vaults.MCPCredentialRequest{{ServerLabel: "tools", ServerURL: destination, CredentialID: &first.ID}}); !isSelectionError(err, false, "MCP credential_id requires an attached vault") { + t.Fatal("a reference without attachments was admitted", err) + } + pool.Close() + store, pool = openStore(t) + service = newService(t, store, newCipher(t, bytes.Clone(key)), nil) + keyless = newService(t, store, nil, nil) + got, err := bearerToken(t.Context(), service, tenant, attached, bindings[0]) + if err != nil || got != token { + t.Fatal("frozen selection or opaque bytes changed across restart", err) + } + if got, err := bearerToken(t.Context(), keyless, tenant, attached, bindings[0]); !errors.Is(err, credentialcrypto.ErrUnavailable) || got != "" { + t.Fatal("missing key did not fail execution closed") + } + key[0] ^= 1 + if got, err := bearerToken(t.Context(), newService(t, store, newCipher(t, key), nil), tenant, attached, bindings[0]); err == nil || got != "" || strings.Contains(err.Error(), token) { + t.Fatal("wrong key leaked or decrypted a credential") + } + for _, mutate := range []func(*vaults.MCPCredentialBinding){ + func(b *vaults.MCPCredentialBinding) { b.VaultID = owned[1].ID }, + func(b *vaults.MCPCredentialBinding) { b.CredentialID = foreignCredential.ID }, + func(b *vaults.MCPCredentialBinding) { b.ServerURL += "/other" }, + func(b *vaults.MCPCredentialBinding) { b.AuthType = "other" }, + } { + binding := bindings[0] + mutate(&binding) + if got, err := bearerToken(t.Context(), service, tenant, attached, binding); !errors.Is(err, vaults.ErrNotFound) || got != "" { + t.Fatal("substituted frozen authorization was decrypted") + } + } + for _, scope := range []struct { + owner string + vaults []string + }{{foreign, attached}, {tenant, []string{owned[1].ID}}} { + if got, err := bearerToken(t.Context(), service, scope.owner, scope.vaults, bindings[0]); !errors.Is(err, vaults.ErrNotFound) || got != "" { + t.Fatal("tenant or attachment authorization was bypassed") + } + } + if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET token_ciphertext=set_byte(token_ciphertext, 15, get_byte(token_ciphertext,15) # 1) WHERE id=$1", first.ID); err != nil { + t.Fatal(err) + } + if got, err := bearerToken(t.Context(), service, tenant, attached, bindings[0]); err == nil || got != "" { + t.Fatal("tampered ciphertext decrypted") + } + if _, err := store.GetCredential(t.Context(), tenant, first.VaultID, first.ID); err != nil { + t.Fatal("safe metadata lookup depended on ciphertext", err) + } +} diff --git a/services/core/internal/persistence/postgres/vaultpg/vaultpg.go b/services/core/internal/persistence/postgres/vaultpg/vaultpg.go new file mode 100644 index 000000000..b77115e42 --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/vaultpg.go @@ -0,0 +1,178 @@ +// Package vaultpg stores Vaults and Credentials in PostgreSQL. +package vaultpg + +import ( + "context" + "encoding/json" + "errors" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgtype" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/textvalue" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/writeaudit" +) + +// Store runs on pooled connections. Credential operations never need the +// execution lease: an OAuth refresh holds its Credential's row lock for up to +// the refresh bound and must not hold up execution-owner work. +type Store struct{ pool *pgunit.Pool } + +var _ vaults.Storage = (*Store)(nil) + +func New(pool *pgunit.Pool) *Store { return &Store{pool: pool} } + +// write runs apply in one pooled transaction and translates its outcome. +func (s *Store) write(ctx context.Context, apply func(context.Context, *sqlc.Queries) error) error { + return translate(s.pool.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + return apply(ctx, sqlc.New(tx)) + })) +} + +// read runs apply in one snapshot and translates its outcome. +func (s *Store) read(ctx context.Context, apply func(context.Context, *sqlc.Queries) error) error { + return translate(s.pool.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { + return apply(ctx, sqlc.New(tx)) + })) +} + +// translate turns the database outcomes a caller acts on into domain errors: +// no row is ErrNotFound and text PostgreSQL cannot store is +// textvalue.ErrUnstorable. Any other error returns as is. +func translate(err error) error { + switch { + case errors.Is(err, pgx.ErrNoRows): + return vaults.ErrNotFound + case pgunit.IsUnstorableText(err): + return textvalue.ErrUnstorable + } + return err +} + +func (s *Store) CreateVault(ctx context.Context, vault vaults.NewVault) (vaults.Vault, error) { + tenant, err := pgunit.ParseID(vault.TenantID) + if err != nil { + return vaults.Vault{}, vaults.ErrInvalidInput + } + var name pgtype.Text + if vault.Name != nil { + name = pgtype.Text{String: *vault.Name, Valid: true} + } + var created vaults.Vault + err = s.write(ctx, func(ctx context.Context, q *sqlc.Queries) error { + row, err := q.CreateVault(ctx, sqlc.CreateVaultParams{ID: newID(), TenantID: tenant, Name: name, Metadata: vault.Metadata}) + if err != nil { + return err + } + created, err = vaultFromRow(row) + if err != nil { + return err + } + return auditpg.RecordWriteAudit(ctx, q, vault.TenantID, "create", "vault", created.ID, "", writeaudit.Resource{Type: "vault", ID: created.ID}) + }) + if err != nil { + return vaults.Vault{}, err + } + return created, nil +} + +func (s *Store) GetVault(ctx context.Context, tenantID, vaultID string) (vaults.Vault, error) { + tenant, err := pgunit.ParseID(tenantID) + if err != nil { + return vaults.Vault{}, vaults.ErrInvalidInput + } + id, err := pgunit.ParseID(vaultID) + if err != nil { + return vaults.Vault{}, vaults.ErrInvalidInput + } + return getVault(ctx, s.pool.Queries(), tenant, id) +} + +func getVault(ctx context.Context, q *sqlc.Queries, tenant, id pgtype.UUID) (vaults.Vault, error) { + row, err := q.GetVault(ctx, sqlc.GetVaultParams{TenantID: tenant, ID: id}) + if err != nil { + return vaults.Vault{}, translate(err) + } + return vaultFromRow(row) +} + +// ListVaults reads the cursor and the page from one snapshot. +func (s *Store) ListVaults(ctx context.Context, tenantID string, query vaults.PageQuery) (vaults.VaultPage, error) { + tenant, err := pgunit.ParseID(tenantID) + if err != nil { + return vaults.VaultPage{}, vaults.ErrInvalidInput + } + statuses, err := query.Validate() + if err != nil { + return vaults.VaultPage{}, err + } + var page vaults.VaultPage + err = s.read(ctx, func(ctx context.Context, q *sqlc.Queries) error { + params := sqlc.ListVaultsParams{TenantID: tenant, PageLimit: int32(query.Limit + 1), AfterID: pgtype.UUID{Valid: true}, Ascending: query.Ascending, Statuses: statuses} + if query.After != "" { + // A cursor that cannot name a Vault follows the missing-cursor path. + after, err := getVault(ctx, q, tenant, pgunit.PathID(query.After)) + if err != nil { + return err + } + params.AfterCreated = pgtype.Timestamptz{Time: after.CreatedAt, Valid: true} + params.AfterID = pgunit.PathID(after.ID) + } + rows, err := q.ListVaults(ctx, params) + if err != nil { + return err + } + page = vaults.VaultPage{Vaults: make([]vaults.Vault, 0, min(query.Limit, len(rows)))} + if len(rows) > query.Limit { + page.NextCursor = uuid.UUID(rows[query.Limit-1].ID.Bytes).String() + rows = rows[:query.Limit] + } + for _, row := range rows { + vault, err := vaultFromRow(row) + if err != nil { + return err + } + page.Vaults = append(page.Vaults, vault) + } + return nil + }) + if err != nil { + return vaults.VaultPage{}, err + } + return page, nil +} + +// DeleteVault relies on the owning foreign key to remove every stored Credential. +func (s *Store) DeleteVault(ctx context.Context, tenantID, vaultID string) (string, error) { + var deleted string + err := s.write(ctx, func(ctx context.Context, q *sqlc.Queries) error { + id, err := q.DeleteVault(ctx, sqlc.DeleteVaultParams{TenantID: pgunit.PathID(tenantID), ID: pgunit.PathID(vaultID)}) + if err != nil { + return err + } + deleted = uuid.UUID(id.Bytes).String() + return auditpg.RecordWriteAudit(ctx, q, tenantID, "delete", "vault", deleted, "") + }) + if err != nil { + return "", err + } + return deleted, nil +} + +func vaultFromRow(row sqlc.Vault) (vaults.Vault, error) { + vault := vaults.Vault{ID: uuid.UUID(row.ID.Bytes).String(), TenantID: uuid.UUID(row.TenantID.Bytes).String(), CreatedAt: row.CreatedAt.Time} + if row.Name.Valid { + vault.Name = &row.Name.String + } + if err := json.Unmarshal(row.Metadata, &vault.Metadata); err != nil { + return vaults.Vault{}, errors.New("invalid stored vault metadata") + } + return vault, nil +} + +func newID() pgtype.UUID { return pgtype.UUID{Bytes: uuid.New(), Valid: true} } diff --git a/services/core/internal/persistence/postgres/vaultpg/vaults_test.go b/services/core/internal/persistence/postgres/vaultpg/vaults_test.go new file mode 100644 index 000000000..cfe6d6cef --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/vaults_test.go @@ -0,0 +1,350 @@ +package vaultpg_test + +import ( + "bytes" + "cmp" + "errors" + "reflect" + "slices" + "strings" + "testing" + "time" + + "github.com/google/uuid" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/vaultpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +func TestVaultsPersistAndStayTenantScoped(t *testing.T) { + store, pool := openStore(t) + service := newService(t, store, nil, nil) + ctx := t.Context() + tenantA, tenantB := uuid.NewString(), uuid.NewString() + before := time.Now().Add(-time.Second) + unnamed, err := service.CreateVault(ctx, vaults.CreateVault{TenantID: tenantA}) + if err != nil { + t.Fatal(err) + } + if _, err := uuid.Parse(unnamed.ID); err != nil || unnamed.TenantID != tenantA || unnamed.Name != nil || unnamed.Metadata == nil || len(unnamed.Metadata) != 0 || unnamed.CreatedAt.Before(before) || unnamed.CreatedAt.After(time.Now().Add(time.Second)) { + t.Fatalf("unexpected unnamed vault: %+v, %v", unnamed, err) + } + // Validate the byte boundary with multibyte text, without Session metadata + // count or character limits. Public name trimming belongs to the API layer. + name := strings.Repeat("é", 128) + command := vaults.CreateVault{TenantID: tenantA, Name: &name, Metadata: map[string]string{"": "", "purpose": "保存 configuration"}} + named, err := service.CreateVault(ctx, command) + if err != nil || named.ID == unnamed.ID || named.Name == nil || *named.Name != name || !reflect.DeepEqual(named.Metadata, command.Metadata) { + t.Fatalf("unexpected named vault: %+v, %v", named, err) + } + for _, tenant := range []string{tenantA, tenantB} { + command.TenantID = tenant + other, err := service.CreateVault(ctx, command) + if err != nil || other.ID == named.ID || other.TenantID != tenant { + t.Fatalf("distinct resource creation: %+v, %v", other, err) + } + } + for _, lookup := range []struct{ tenant, id string }{{tenantB, named.ID}, {tenantA, uuid.NewString()}} { + if _, err := store.GetVault(ctx, lookup.tenant, lookup.id); !errors.Is(err, vaults.ErrNotFound) { + t.Fatalf("unowned/unknown vault lookup: %v", err) + } + } + // Recreate the pool and Store as a restarted standalone service would. + pool.Close() + reopened, pool := openStore(t) + for _, want := range []vaults.Vault{unnamed, named} { + got, err := reopened.GetVault(ctx, tenantA, want.ID) + if err != nil || !reflect.DeepEqual(got, want) { + t.Fatalf("durable read: %+v, %v; want %+v", got, err, want) + } + } + var sessions int + if err := pool.QueryRow(ctx, "SELECT count(*) FROM sessions WHERE tenant_id = $1", tenantA).Scan(&sessions); err != nil || sessions != 0 { + t.Fatalf("Vault creation produced %d Sessions: %v", sessions, err) + } +} + +func TestVaultsRejectInvalidInputWithoutWrites(t *testing.T) { + store, pool := openStore(t) + service := newService(t, store, nil, nil) + ctx := t.Context() + tenant := uuid.NewString() + for _, name := range []string{"", strings.Repeat("x", 257), strings.Repeat("é", 129), string([]byte{0xff})} { + if _, err := service.CreateVault(ctx, vaults.CreateVault{TenantID: tenant, Name: &name}); !errors.Is(err, vaults.ErrInvalidInput) { + t.Fatalf("invalid name length %d: %v", len(name), err) + } + } + for _, invalid := range []string{"", "not-a-uuid", uuid.Nil.String()} { + if _, err := service.CreateVault(ctx, vaults.CreateVault{TenantID: invalid}); !errors.Is(err, vaults.ErrInvalidInput) { + t.Fatalf("invalid create tenant accepted: %v", err) + } + if _, err := store.GetVault(ctx, invalid, uuid.NewString()); !errors.Is(err, vaults.ErrInvalidInput) { + t.Fatalf("invalid read tenant accepted: %v", err) + } + if _, err := store.GetVault(ctx, tenant, invalid); !errors.Is(err, vaults.ErrInvalidInput) { + t.Fatalf("invalid vault ID accepted: %v", err) + } + } + if _, err := service.CreateVault(ctx, vaults.CreateVault{TenantID: tenant, Metadata: map[string]string{"large": strings.Repeat("x", 64*1024)}}); !errors.Is(err, vaults.ErrInvalidInput) { + t.Fatalf("oversized metadata accepted: %v", err) + } + var count int + if err := pool.QueryRow(ctx, "SELECT count(*) FROM vaults WHERE tenant_id = $1", tenant).Scan(&count); err != nil || count != 0 { + t.Fatalf("invalid input created %d Vaults: %v", count, err) + } +} + +func TestVaultListFilteringPaginationAndReconnect(t *testing.T) { + store, pool := openStore(t) + service := newService(t, store, nil, nil) + ctx := t.Context() + tenant, other := uuid.NewString(), uuid.NewString() + empty, err := store.ListVaults(ctx, tenant, vaults.PageQuery{Limit: 20}) + if err != nil || empty.Vaults == nil || len(empty.Vaults) != 0 || empty.NextCursor != "" { + t.Fatalf("empty page: %+v, %v", empty, err) + } + var all []vaults.Vault + archived := map[string]bool{} + for i := range 105 { + vault, err := service.CreateVault(ctx, vaults.CreateVault{TenantID: tenant, Metadata: map[string]string{"purpose": "safe list fixture"}}) + if err != nil { + t.Fatal(err) + } + status := vaults.StatusActive + if i%3 == 0 { + status, archived[vault.ID] = vaults.StatusArchived, true + } + // Synthetic classifications exercise reads, not a public archive lifecycle. + if _, err := pool.Exec(ctx, "UPDATE vaults SET created_at=$1, status=$2 WHERE tenant_id=$3 AND id=$4", time.Unix(1700000000+int64(i%2), 0).UTC(), status, tenant, vault.ID); err != nil { + t.Fatal(err) + } + vault, err = store.GetVault(ctx, tenant, vault.ID) + if err != nil { + t.Fatal(err) + } + all = append(all, vault) + } + slices.SortFunc(all, func(a, b vaults.Vault) int { + if c := a.CreatedAt.Compare(b.CreatedAt); c != 0 { + return c + } + return cmp.Compare(a.ID, b.ID) + }) + foreign := createVault(t, service, other) + read := func(reader vaults.Reader, ascending bool, statuses []string, size int) []vaults.Vault { + t.Helper() + actual := []vaults.Vault{} + cursor := "" + for { + page, err := reader.ListVaults(ctx, tenant, vaults.PageQuery{After: cursor, Limit: size, Ascending: ascending, Statuses: statuses}) + if err != nil || len(page.Vaults) == 0 || len(page.Vaults) > size { + t.Fatalf("page: %+v, %v", page, err) + } + actual = append(actual, page.Vaults...) + if len(actual) > len(all) { + t.Fatal("pagination repeated records") + } + if page.NextCursor == "" { + break + } + if page.NextCursor != page.Vaults[len(page.Vaults)-1].ID { + t.Fatal("cursor is not the last included resource") + } + cursor = page.NextCursor + } + return actual + } + for _, ascending := range []bool{true, false} { + for _, statuses := range [][]string{nil, {vaults.StatusActive}, {vaults.StatusArchived}, {vaults.StatusActive, vaults.StatusArchived}} { + want := []vaults.Vault{} + for _, vault := range all { + if len(statuses) != 1 || archived[vault.ID] == (statuses[0] == vaults.StatusArchived) { + want = append(want, vault) + } + } + if !ascending { + slices.Reverse(want) + } + for _, size := range []int{20, 100} { + if got := read(store, ascending, statuses, size); !reflect.DeepEqual(got, want) { + t.Fatalf("filtered ordering/projection mismatch: ascending=%t statuses=%v size=%d got=%d want=%d", ascending, statuses, size, len(got), len(want)) + } + } + } + } + for _, cursor := range []string{foreign.ID, uuid.NewString(), "invalid"} { + if _, err := store.ListVaults(ctx, tenant, vaults.PageQuery{After: cursor, Limit: 20, Ascending: true, Statuses: []string{vaults.StatusArchived}}); !errors.Is(err, vaults.ErrNotFound) { + t.Fatalf("foreign/unknown/malformed cursor: %v", err) + } + } + for _, tc := range []struct { + tenant string + query vaults.PageQuery + }{{"invalid", vaults.PageQuery{Limit: 20}}, {tenant, vaults.PageQuery{Limit: 0}}, {tenant, vaults.PageQuery{Limit: 101}}, {tenant, vaults.PageQuery{Limit: 20, Statuses: []string{"deleted"}}}} { + if _, err := store.ListVaults(ctx, tc.tenant, tc.query); !errors.Is(err, vaults.ErrInvalidInput) { + t.Fatalf("invalid store query: %v", err) + } + } + tail, err := store.ListVaults(ctx, tenant, vaults.PageQuery{After: all[len(all)-1].ID, Limit: 100, Ascending: true}) + if err != nil || tail.Vaults == nil || len(tail.Vaults) != 0 || tail.NextCursor != "" { + t.Fatalf("terminal page: %+v, %v", tail, err) + } + page, err := store.ListVaults(ctx, other, vaults.PageQuery{Limit: 100}) + if err != nil || !reflect.DeepEqual(page.Vaults, []vaults.Vault{foreign}) || page.NextCursor != "" { + t.Fatalf("project isolation: %+v, %v", page, err) + } + pool.Close() + reopened, _ := openStore(t) + if got := read(reopened, true, nil, 20); !reflect.DeepEqual(got, all) { + t.Fatal("listing changed after reconnect") + } +} + +func TestVaultDeletionCascadeBindingAndRestart(t *testing.T) { + store, pool := openStore(t) + tenant, foreign := uuid.NewString(), uuid.NewString() + key := bytes.Repeat([]byte{43}, 32) + service := newService(t, store, newCipher(t, key), nil) + keyless := newService(t, store, nil, nil) + vault, retained, empty := createVault(t, service, tenant), createVault(t, service, tenant), createVault(t, service, tenant) + original := createStatic(t, service, tenant, vault.ID, "original", "https://mcp.example/tools", "original-secret") + attached := []string{vault.ID, retained.ID} + selected := resolve(t, service, tenant, attached, vaults.MCPCredentialRequest{ServerLabel: "tools", ServerURL: original.MCPServerURL}) + extra := createStatic(t, service, tenant, vault.ID, "archived", "https://mcp.example/other", "archived-secret") + sibling := createStatic(t, service, tenant, retained.ID, "sibling", original.MCPServerURL, "sibling-secret") + if _, err := pool.Exec(t.Context(), "UPDATE vaults SET status='archived' WHERE id=$1", vault.ID); err != nil { + t.Fatal(err) + } + if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET status='archived' WHERE id=$1", extra.ID); err != nil { + t.Fatal(err) + } + for _, scope := range []struct{ tenant, id string }{{foreign, vault.ID}, {tenant, uuid.NewString()}, {tenant, "invalid"}, {"invalid", vault.ID}} { + if _, err := keyless.DeleteVault(t.Context(), vaults.DeleteVault{TenantID: scope.tenant, VaultID: scope.id}); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("foreign or invalid deletion was accepted", err) + } + } + _, deletionErr := newService(t, readOnlyStore(t, pool), nil, nil).DeleteVault(t.Context(), vaults.DeleteVault{TenantID: tenant, VaultID: vault.ID}) + if !isReadOnlyFailure(deletionErr) { + t.Fatal("failed mutation was accepted or translated", deletionErr) + } + tx, err := pool.Begin(t.Context()) + if err != nil { + t.Fatal(err) + } + defer func() { _ = tx.Rollback(t.Context()) }() + // Verify the database cascade independently inside an explicit transaction. + if _, err := sqlc.New(tx).DeleteVault(t.Context(), sqlc.DeleteVaultParams{TenantID: pgunit.PathID(tenant), ID: pgunit.PathID(vault.ID)}); err != nil { + t.Fatal(err) + } + var count int + if err := tx.QueryRow(t.Context(), "SELECT (SELECT count(*) FROM vaults WHERE id=$1)+(SELECT count(*) FROM vault_credentials WHERE vault_id=$1)", vault.ID).Scan(&count); err != nil || count != 0 { + t.Fatal("cascade was not visible in the deletion transaction", err) + } + if err := tx.Rollback(t.Context()); err != nil { + t.Fatal(err) + } + if value, err := store.GetVault(t.Context(), tenant, vault.ID); err != nil || !reflect.DeepEqual(value, vault) { + t.Fatal("rollback changed the Vault", err) + } + for _, expected := range []vaults.Credential{original, extra} { + if value, err := store.GetCredential(t.Context(), tenant, vault.ID, expected.ID); err != nil || !reflect.DeepEqual(value, expected) { + t.Fatal("rollback changed a child", err) + } + } + if token, err := bearerToken(t.Context(), service, tenant, attached, selected[0]); err != nil || token != "original-secret" { + t.Fatal("rejected deletion changed the stored token") + } + if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET token_ciphertext=decode('00','hex') WHERE vault_id=$1", vault.ID); err != nil { + t.Fatal(err) + } + for _, target := range []vaults.Vault{empty, vault} { + if id, err := keyless.DeleteVault(t.Context(), vaults.DeleteVault{TenantID: tenant, VaultID: target.ID}); err != nil || id != target.ID { + t.Fatal("keyless deletion failed", err) + } + } + if err := pool.QueryRow(t.Context(), "SELECT (SELECT count(*) FROM vaults WHERE id=$1)+(SELECT count(*) FROM vault_credentials WHERE vault_id=$1)", vault.ID).Scan(&count); err != nil || count != 0 { + t.Fatal("committed parent or encrypted children remain", err) + } + pool.Close() + store, _ = openStore(t) + service = newService(t, store, newCipher(t, bytes.Clone(key)), nil) + for _, target := range []vaults.Vault{empty, vault} { + if _, err := store.GetVault(t.Context(), tenant, target.ID); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("deleted Vault reappeared after restart") + } + if _, err := service.DeleteVault(t.Context(), vaults.DeleteVault{TenantID: tenant, VaultID: target.ID}); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("repeated deletion did not remain absent") + } + } + for _, child := range []vaults.Credential{original, extra} { + if _, err := store.GetCredential(t.Context(), tenant, vault.ID, child.ID); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("deleted child reappeared") + } + if _, err := service.UpdateStaticCredential(t.Context(), vaults.UpdateStaticCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: child.ID, Token: "replacement"}); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("replacement recreated a deleted child") + } + } + if _, err := service.CreateStaticCredential(t.Context(), vaults.CreateStaticCredential{TenantID: tenant, VaultID: vault.ID, Name: "late", MCPServerURL: original.MCPServerURL, Token: "late"}); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("new child was admitted under a deleted Vault") + } + if _, err := bearerToken(t.Context(), service, tenant, attached, selected[0]); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("frozen binding reselected a credential in another attached Vault") + } + if _, err := store.ListCredentials(t.Context(), tenant, vault.ID, vaults.PageQuery{Limit: 100, Ascending: true}); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("deleted parent remained listable") + } + if value, err := store.GetVault(t.Context(), tenant, retained.ID); err != nil || !reflect.DeepEqual(value, retained) { + t.Fatal("deletion changed another Vault") + } + if value, err := store.GetCredential(t.Context(), tenant, retained.ID, sibling.ID); err != nil || !reflect.DeepEqual(value, sibling) { + t.Fatal("deletion changed another Vault's credential") + } +} + +func TestVaultDeletionConcurrentChildMutations(t *testing.T) { + store, pool := openStore(t) + tenant := uuid.NewString() + service := newService(t, store, newCipher(t, bytes.Repeat([]byte{44}, 32)), nil) + // Each operation runs on its own Store, so only PostgreSQL orders them. + other := func() *vaults.Service { + return newService(t, vaultpg.New(pgunit.NewPool(pool)), newCipher(t, bytes.Repeat([]byte{44}, 32)), nil) + } + for range 8 { + vault := createVault(t, service, tenant) + value := createStatic(t, service, tenant, vault.ID, "competing", "https://mcp.example/tools", "before") + start, created, updated, removed := make(chan struct{}), make(chan error, 1), make(chan error, 1), make(chan error, 1) + creator, updater, remover := other(), other(), other() + go func() { + <-start + _, err := creator.CreateStaticCredential(t.Context(), vaults.CreateStaticCredential{TenantID: tenant, VaultID: vault.ID, Name: "competing", MCPServerURL: "https://mcp.example/tools", Token: "before"}) + created <- err + }() + go func() { + <-start + _, err := updater.UpdateStaticCredential(t.Context(), vaults.UpdateStaticCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: value.ID, Token: "after"}) + updated <- err + }() + go func() { + <-start + _, err := remover.DeleteCredential(t.Context(), vaults.DeleteCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: value.ID}) + removed <- err + }() + close(start) + _, deleted := service.DeleteVault(t.Context(), vaults.DeleteVault{TenantID: tenant, VaultID: vault.ID}) + createErr, updateErr, removeErr := <-created, <-updated, <-removed + for _, err := range []error{createErr, updateErr, removeErr} { + if err != nil && !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("competing mutation failed unexpectedly", err) + } + } + if deleted != nil { + t.Fatal("Vault deletion failed", deleted) + } + var count int + if err := pool.QueryRow(t.Context(), "SELECT (SELECT count(*) FROM vaults WHERE id=$1)+(SELECT count(*) FROM vault_credentials WHERE vault_id=$1)", vault.ID).Scan(&count); err != nil || count != 0 { + t.Fatal("concurrent mutation resurrected deleted resources", err) + } + } +} diff --git a/services/core/internal/sandbox/providers/configuration_flow_test.go b/services/core/internal/sandbox/providers/configuration_flow_test.go index 1611a17a0..74f6adfbe 100644 --- a/services/core/internal/sandbox/providers/configuration_flow_test.go +++ b/services/core/internal/sandbox/providers/configuration_flow_test.go @@ -103,7 +103,8 @@ func TestAdditionalConfigurationProviderUsesCommonAPIAndStore(t *testing.T) { // The flow reaches only the store areas and the deployment setup; every // other dependency panics if called. h, err := api.NewHandler(api.Dependencies{ - Engine: "codex", CoreKeys: auth, InstallationBindings: s, Projects: s, Vaults: s, ModelProviders: s, Skills: s, + Engine: "codex", CoreKeys: auth, InstallationBindings: s, Projects: s, ModelProviders: s, Skills: s, + Vaults: struct{ api.Vaults }{}, VaultsReader: struct{ api.VaultsReader }{}, Files: struct{ api.Files }{}, FilesReader: struct{ api.FilesReader }{}, Agents: struct{ api.Agents }{}, AgentsReader: struct{ api.AgentsReader }{}, EnvironmentTemplates: s, Sessions: s, SessionEvents: s, SessionHistory: s, Subagents: s, Artifacts: s, diff --git a/services/core/internal/store/admin_delete_audit_test.go b/services/core/internal/store/admin_delete_audit_test.go index b68342b40..4bf59c643 100644 --- a/services/core/internal/store/admin_delete_audit_test.go +++ b/services/core/internal/store/admin_delete_audit_test.go @@ -105,8 +105,8 @@ func TestAdminDeleteResourceAuditTransactions(t *testing.T) { s := NewWithCredentialCipher(pool, cipher) rejectAdminAuditInsert(t, s) archive := skillArchive(t, "admin-private-archive") - tables := []string{"agents", "agent_model_execution", "environment_templates", "skills", "skill_versions", "vaults", "vault_credentials", "sessions", "turns", "environments", "session_artifacts", "admin_audit_log", "write_audit_operations", "write_audit_owners", "pg_largeobject_metadata", "pg_largeobject"} - for _, name := range []string{"template_delete", "skill_delete", "version_delete", "version_delete_last", "vault_delete", "credential_delete", "oauth_delete", "session_delete", "artifact_delete"} { + tables := []string{"agents", "agent_model_execution", "environment_templates", "skills", "skill_versions", "sessions", "turns", "environments", "session_artifacts", "admin_audit_log", "write_audit_operations", "write_audit_owners", "pg_largeobject_metadata", "pg_largeobject"} + for _, name := range []string{"template_delete", "skill_delete", "version_delete", "version_delete_last", "session_delete", "artifact_delete"} { t.Run(name, func(t *testing.T) { tenant := uuid.NewString() var mutation resourceAuditMutation @@ -216,10 +216,6 @@ func assertAdminDeletedResource(t *testing.T, s *Store, tenant string, mutation t.Fatal("skill version survived deletion", queryErr) } return - case "vault": - _, err = s.GetVault(t.Context(), tenant, id) - case "credential": - _, err = s.GetCredential(t.Context(), tenant, mutation.parent, id) case "session": _, err = s.GetSession(t.Context(), tenant, id) case "artifact": diff --git a/services/core/internal/store/credential_cipher.go b/services/core/internal/store/credential_cipher.go new file mode 100644 index 000000000..3051915f3 --- /dev/null +++ b/services/core/internal/store/credential_cipher.go @@ -0,0 +1,14 @@ +package store + +import ( + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/jackc/pgx/v5/pgxpool" +) + +// NewWithCredentialCipher configures immutable credential encryption before the +// Store is published. A nil cipher leaves non-secret resource operations available. +func NewWithCredentialCipher(pool *pgxpool.Pool, cipher *credentialcrypto.Cipher) *Store { + s := New(pool) + s.credentialCipher = cipher + return s +} diff --git a/services/core/internal/store/credential_oauth_secret.go b/services/core/internal/store/credential_oauth_secret.go deleted file mode 100644 index 943b14eb5..000000000 --- a/services/core/internal/store/credential_oauth_secret.go +++ /dev/null @@ -1,82 +0,0 @@ -package store - -import ( - "encoding/json" - "errors" - "reflect" - "time" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" -) - -// The encrypted copy authenticates every public setting used for refresh, including -// the token endpoint. Substituting database metadata must never redirect a grant. -type oauthSecret struct { - Version int `json:"version"` - Metadata OAuthMetadata `json:"metadata"` - AccessToken string `json:"access_token"` - RefreshToken string `json:"refresh_token"` - ClientSecret string `json:"client_secret"` -} - -func validOAuthMetadata(metadata OAuthMetadata) bool { - if metadata.ExpiresAt != nil { - if _, err := time.Parse(time.RFC3339Nano, *metadata.ExpiresAt); err != nil { - return false - } - } - if refresh := metadata.Refresh; refresh != nil { - if refresh.ClientID == "" || refresh.TokenEndpoint == "" { - return false - } - switch refresh.TokenEndpointAuth { - case "none", "client_secret_basic", "client_secret_post": - default: - return false - } - } - return true -} - -func oauthBinding(tenantID string, credential Credential) credentialcrypto.Binding { - return credentialcrypto.Binding{TenantID: tenantID, VaultID: credential.VaultID, - CredentialID: credential.ID, AuthType: "mcp_oauth", Destination: credential.MCPServerURL} -} - -func (s *Store) sealOAuth(tenantID string, credential Credential, secret oauthSecret) ([]byte, []byte, error) { - if s.credentialCipher == nil { - return nil, nil, credentialcrypto.ErrUnavailable - } - if !validOAuthMetadata(secret.Metadata) { - return nil, nil, ErrInvalidInput - } - metadata, err := json.Marshal(secret.Metadata) - if err != nil { - return nil, nil, errors.New("credential encoding failed") - } - plaintext, err := json.Marshal(secret) - if err != nil { - return nil, nil, errors.New("credential encoding failed") - } - ciphertext, err := s.credentialCipher.Seal(plaintext, oauthBinding(tenantID, credential)) - if err != nil { - return nil, nil, errors.New("credential encryption failed") - } - return metadata, ciphertext, nil -} - -func (s *Store) openOAuth(tenantID string, credential Credential, ciphertext []byte) (oauthSecret, error) { - if s.credentialCipher == nil { - return oauthSecret{}, credentialcrypto.ErrUnavailable - } - plaintext, err := s.credentialCipher.Open(ciphertext, oauthBinding(tenantID, credential)) - if err != nil { - return oauthSecret{}, errors.New("OAuth credential decryption failed") - } - var secret oauthSecret - if json.Unmarshal(plaintext, &secret) != nil || secret.Version != 1 || credential.OAuth == nil || - !reflect.DeepEqual(secret.Metadata, *credential.OAuth) || !validOAuthMetadata(secret.Metadata) { - return oauthSecret{}, errors.New("OAuth credential authentication failed") - } - return secret, nil -} diff --git a/services/core/internal/store/credential_oauth_types.go b/services/core/internal/store/credential_oauth_types.go deleted file mode 100644 index 1eb674f0b..000000000 --- a/services/core/internal/store/credential_oauth_types.go +++ /dev/null @@ -1,38 +0,0 @@ -package store - -// OAuthMetadata is the safe projection of a stored OAuth grant. -type OAuthMetadata struct { - ExpiresAt *string `json:"expires_at"` - Refresh *OAuthRefreshMetadata `json:"refresh"` -} - -type OAuthRefreshMetadata struct { - ClientID string `json:"client_id"` - TokenEndpoint string `json:"token_endpoint"` - TokenEndpointAuth string `json:"token_endpoint_auth"` - Resource *string `json:"resource"` - Scope *string `json:"scope"` -} - -// CreateOAuthCredentialInput keeps write-only secrets separate from metadata. -type CreateOAuthCredentialInput struct { - Name, MCPServerURL, AccessToken string - OAuth OAuthMetadata - RefreshToken, ClientSecret string -} - -// ExpiresAtSet and ScopeSet preserve omitted-versus-null update semantics. -type UpdateOAuthCredentialInput struct { - AccessToken *string - ExpiresAt *string - ExpiresAtSet bool - Refresh *OAuthRefreshUpdate -} - -type OAuthRefreshUpdate struct { - RefreshToken *string - Scope *string - ScopeSet bool - TokenEndpointAuthType string - ClientSecret *string -} diff --git a/services/core/internal/store/fixture_db_test.go b/services/core/internal/store/fixture_db_test.go index ac1445632..256368724 100644 --- a/services/core/internal/store/fixture_db_test.go +++ b/services/core/internal/store/fixture_db_test.go @@ -39,7 +39,8 @@ func newManagedTestStoreDB(t *testing.T) (*store.Store, fixtureDB) { // startWorker starts the execution Worker as cmd/server does: it acquires the // execution lease on db and hands it, with the execution writer built on it, to -// the Worker, which closes it when Run exits. +// the Worker, which closes it when Run exits. The Worker opens MCP bearer tokens +// through the vaults service on db. func startWorker(t testing.TB, ctx context.Context, db fixtureDB, dispatcher *execution.Dispatcher) *execution.Worker { t.Helper() worker, err := startWorkerErr(ctx, db, dispatcher) @@ -51,11 +52,17 @@ func startWorker(t testing.TB, ctx context.Context, db fixtureDB, dispatcher *ex // startWorkerErr is startWorker for tests that assert a startup failure. func startWorkerErr(ctx context.Context, db fixtureDB, dispatcher *execution.Dispatcher) (*execution.Worker, error) { + _, credentials, err := fixtureVaults(db) + if err != nil { + return nil, err + } lease, err := pgunit.AcquireLease(ctx, db.pool) if err != nil { return nil, err } - return execution.StartWorker(ctx, dispatcher, execution.Owner{Lease: lease, Store: store.NewExecution(dispatcher.Store, lease)}) + owned := *dispatcher + owned.Credentials = credentials + return execution.StartWorker(ctx, &owned, execution.Owner{Lease: lease, Store: store.NewExecution(dispatcher.Store, lease)}) } // executionOwner acquires the execution lease on db and builds s's execution diff --git a/services/core/internal/store/function_worker_test.go b/services/core/internal/store/function_worker_test.go index 65d9d16d4..0279a72a1 100644 --- a/services/core/internal/store/function_worker_test.go +++ b/services/core/internal/store/function_worker_test.go @@ -14,6 +14,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/execution" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" ) @@ -127,12 +128,16 @@ func mcpBearerWorkerConfiguration(t *testing.T, h *dispatchHarness) (string, str } h.s, h.db = store.NewWithCredentialCipher(pool, cipher), fixtureDB{pool: pool, cipher: cipher} h.d.Store = h.s - vault, err := h.s.CreateVault(t.Context(), h.tenant, store.CreateVaultInput{}) + _, service, err := fixtureVaults(h.db) + if err != nil { + t.Fatal(err) + } + vault, err := service.CreateVault(t.Context(), vaults.CreateVault{TenantID: h.tenant}) if err != nil { t.Fatal(err) } token, endpoint := uuid.NewString(), "https://mcp.example/tools" - _, err = h.s.CreateStaticCredential(t.Context(), h.tenant, vault.ID, store.CreateStaticCredentialInput{Name: "worker", MCPServerURL: endpoint, Token: token}) + _, err = service.CreateStaticCredential(t.Context(), vaults.CreateStaticCredential{TenantID: h.tenant, VaultID: vault.ID, Name: "worker", MCPServerURL: endpoint, Token: token}) if err != nil { t.Fatal(err) } @@ -141,7 +146,8 @@ func mcpBearerWorkerConfiguration(t *testing.T, h *dispatchHarness) (string, str t.Fatal("invalid worker fixture") } snapshot.VaultIDs = []string{vault.ID} - snapshot.MCPCredentials, err = h.s.ResolveMCPCredentials(t.Context(), h.tenant, snapshot.VaultIDs, []store.MCPCredentialRequest{{ServerLabel: "tickets", ServerURL: endpoint}}) + snapshot.MCPCredentials, err = service.ResolveMCPCredentials(t.Context(), vaults.ResolveMCPCredentials{TenantID: h.tenant, VaultIDs: snapshot.VaultIDs, + Requests: []vaults.MCPCredentialRequest{{ServerLabel: "tickets", ServerURL: endpoint}}}) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/mcp_credentials.go b/services/core/internal/store/mcp_credentials.go deleted file mode 100644 index 3e635f8a7..000000000 --- a/services/core/internal/store/mcp_credentials.go +++ /dev/null @@ -1,196 +0,0 @@ -package store - -import ( - "context" - "errors" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/echotext" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgtype" -) - -type MCPCredentialRequest struct { - ServerLabel, ServerURL string - CredentialID *string -} - -// MCPCredentialBinding freezes a non-secret selection, including anonymous -// servers. It is private execution configuration, not the public MCP tool shape. -type MCPCredentialBinding struct { - ServerLabel string `json:"server_label"` - ServerURL string `json:"server_url"` - VaultID string `json:"vault_id,omitempty"` - CredentialID string `json:"credential_id,omitempty"` - AuthType string `json:"auth_type,omitempty"` -} - -// MCPCredentialSelectionError rejects a Session MCP credential selection with -// the observed official message (MV-03); the API reports a Conflict as 409 -// conflict_error and any other as 400 invalid_request_error. Selection searches -// only the attached Vaults, which the caller owns, so a missing, foreign-tenant, -// unattached or malformed reference produces the same error, and only a -// credential of an attached Vault can report a server_url mismatch. -type MCPCredentialSelectionError struct { - Conflict bool - Message string -} - -func (e *MCPCredentialSelectionError) Error() string { return e.Message } - -// echoed repeats a caller-supplied value in a selection message only within -// the shared bound; otherwise the message leaves it out. -func echoed(value string) string { - if !echotext.Allowed(value) { - return "" - } - return " " + value -} - -func mcpCredentialRequiresVault() error { - return &MCPCredentialSelectionError{Message: "MCP credential_id requires an attached vault"} -} - -func mcpCredentialNotAttached(id string) error { - return &MCPCredentialSelectionError{Message: "MCP credential_id" + echoed(id) + " was not found in an attached vault"} -} - -func mcpCredentialURLMismatch(id, url string) error { - return &MCPCredentialSelectionError{Message: "MCP credential_id" + echoed(id) + " does not match server_url" + echoed(url)} -} - -func mcpCredentialAmbiguous(url string) error { - return &MCPCredentialSelectionError{Conflict: true, Message: "multiple attached vault credentials match MCP server_url" + echoed(url) + "; specify credential_id"} -} - -func attachedVaultIDs(ids []string) ([]pgtype.UUID, error) { - result := make([]pgtype.UUID, 0, len(ids)) - seen := map[pgtype.UUID]bool{} - for _, raw := range ids { - id, err := parseID(raw) - if err != nil { - return nil, ErrNotFound - } - if !seen[id] { - result = append(result, id) - seen[id] = true - } - } - return result, nil -} - -// ResolveMCPCredentials reads metadata only. Resource changes after this read do -// not reselect credentials for an accepted Session or its creation retries. -func (s *Store) ResolveMCPCredentials(ctx context.Context, tenantID string, vaultIDs []string, requests []MCPCredentialRequest) ([]MCPCredentialBinding, error) { - tenant, err := parseID(tenantID) - if err != nil { - return nil, err - } - vaults, err := attachedVaultIDs(vaultIDs) - if err != nil { - return nil, err - } - owned, err := s.queries.GetAttachedVaultIDs(ctx, sqlc.GetAttachedVaultIDsParams{TenantID: tenant, VaultIds: vaults}) - if err != nil { - return nil, errors.New("cannot resolve attached Vaults") - } - if len(owned) != len(vaults) { - return nil, ErrNotFound - } - bindings := make([]MCPCredentialBinding, 0, len(requests)) - for _, request := range requests { - if request.ServerLabel == "" || request.ServerURL == "" { - return nil, ErrInvalidInput - } - var id pgtype.UUID - if request.CredentialID != nil { - if len(vaults) == 0 { - return nil, mcpCredentialRequiresVault() - } - id, err = parseID(*request.CredentialID) - if err != nil { - return nil, mcpCredentialNotAttached(*request.CredentialID) - } - } - rows, err := s.queries.FindMCPCredentials(ctx, sqlc.FindMCPCredentialsParams{ - TenantID: tenant, VaultIds: vaults, McpServerUrl: request.ServerURL, CredentialID: id, - }) - if err != nil { - return nil, errors.New("cannot resolve MCP credential") - } - if request.CredentialID != nil { - if len(rows) == 0 { - return nil, mcpCredentialNotAttached(*request.CredentialID) - } - if rows[0].McpServerUrl != request.ServerURL { - return nil, mcpCredentialURLMismatch(*request.CredentialID, request.ServerURL) - } - } - if len(rows) > 1 { - return nil, mcpCredentialAmbiguous(request.ServerURL) - } - binding := MCPCredentialBinding{ServerLabel: request.ServerLabel, ServerURL: request.ServerURL} - if len(rows) == 1 { - binding.VaultID, binding.CredentialID = uuid.UUID(rows[0].VaultID.Bytes).String(), uuid.UUID(rows[0].ID.Bytes).String() - binding.AuthType = rows[0].AuthType - } - bindings = append(bindings, binding) - } - return bindings, nil -} - -// MCPBearerToken is execution-only: recheck the complete frozen authorization -// before decrypting. Never persist or log its result, or downgrade failure to an -// anonymous request. Public resource queries do not select ciphertext. -func (s *Store) MCPBearerToken(ctx context.Context, tenantID string, vaultIDs []string, binding MCPCredentialBinding) (string, error) { - tenant, err := parseID(tenantID) - if err != nil { - return "", ErrNotFound - } - vaults, err := attachedVaultIDs(vaultIDs) - if err != nil { - return "", err - } - vault, err := parseID(binding.VaultID) - if err != nil { - return "", ErrNotFound - } - id, err := parseID(binding.CredentialID) - if err != nil || (binding.AuthType != "static_bearer" && binding.AuthType != "mcp_oauth") || binding.ServerURL == "" { - return "", ErrNotFound - } - if binding.AuthType == "mcp_oauth" { - attached := false - for _, candidate := range vaults { - if candidate == vault { - attached = true - } - } - if !attached { - return "", ErrNotFound - } - return s.oauthBearerToken(ctx, tenantID, binding) - } - ciphertext, err := s.queries.GetMCPStaticCredentialCiphertext(ctx, sqlc.GetMCPStaticCredentialCiphertextParams{ - TenantID: tenant, VaultIds: vaults, VaultID: vault, CredentialID: id, McpServerUrl: binding.ServerURL, - }) - if errors.Is(err, pgx.ErrNoRows) { - return "", ErrNotFound - } - if err != nil { - return "", errors.New("cannot read MCP credential") - } - if s.credentialCipher == nil { - return "", credentialcrypto.ErrUnavailable - } - plaintext, err := s.credentialCipher.Open(ciphertext, credentialcrypto.Binding{ - TenantID: uuid.UUID(tenant.Bytes).String(), VaultID: uuid.UUID(vault.Bytes).String(), CredentialID: uuid.UUID(id.Bytes).String(), - AuthType: binding.AuthType, Destination: binding.ServerURL, - }) - if err != nil { - return "", errors.New("MCP credential decryption failed") - } - return string(plaintext), nil -} diff --git a/services/core/internal/store/mcp_credentials_oauth.go b/services/core/internal/store/mcp_credentials_oauth.go deleted file mode 100644 index 4abdb853f..000000000 --- a/services/core/internal/store/mcp_credentials_oauth.go +++ /dev/null @@ -1,79 +0,0 @@ -package store - -import ( - "context" - "errors" - "time" - - "github.com/jackc/pgx/v5" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" -) - -const oauthRefreshTimeout = 20 * time.Second - -func (s *Store) oauthBearerToken(ctx context.Context, tenantID string, binding MCPCredentialBinding) (string, error) { - // Bound both lock contention and the external exchange, so a slow provider cannot - // hold a credential indefinitely. The HTTP client supplies its own tighter bound. - ctx, cancel := context.WithTimeout(ctx, oauthRefreshTimeout) - defer cancel() - var bearer string - err := s.withOAuth(ctx, tenantID, binding.VaultID, binding.CredentialID, binding.ServerURL, "OAuth credential refresh commit failed", func(ctx context.Context, tx pgx.Tx, credential Credential, secret oauthSecret) error { - if credential.MCPServerURL != binding.ServerURL { - return ErrNotFound - } - if secret.Metadata.ExpiresAt == nil { - var err error - bearer, err = oauthAccessToken(secret.AccessToken) - return err - } - expiry, err := time.Parse(time.RFC3339Nano, *secret.Metadata.ExpiresAt) - if err != nil { - return errors.New("invalid OAuth token expiry") - } - if time.Now().Before(expiry) { - bearer, err = oauthAccessToken(secret.AccessToken) - return err - } - refresh := secret.Metadata.Refresh - if refresh == nil || secret.RefreshToken == "" || s.oauthRefresher == nil { - return errors.New("expired OAuth credential cannot be refreshed") - } - token, err := s.oauthRefresher.Refresh(ctx, oauthrefresh.Request{ - TokenEndpoint: refresh.TokenEndpoint, ClientID: refresh.ClientID, - AuthMethod: refresh.TokenEndpointAuth, ClientSecret: secret.ClientSecret, - RefreshToken: secret.RefreshToken, Resource: refresh.Resource, Scope: refresh.Scope, - }) - if err != nil { - return errors.New("OAuth credential refresh failed") - } - if token.AccessToken == "" || token.ExpiresAt != nil && !time.Now().Before(*token.ExpiresAt) { - return errors.New("OAuth refresh returned an unusable token") - } - secret.AccessToken = token.AccessToken - if token.RefreshToken != "" { - secret.RefreshToken = token.RefreshToken - } - secret.Metadata.ExpiresAt = nil - if token.ExpiresAt != nil { - value := token.ExpiresAt.UTC().Format(time.RFC3339Nano) - secret.Metadata.ExpiresAt = &value - } - if _, err := s.saveOAuth(ctx, tx, tenantID, credential, secret); err != nil { - return err - } - bearer = token.AccessToken - return nil - }) - if err != nil { - return "", err - } - return bearer, nil -} - -func oauthAccessToken(token string) (string, error) { - if token == "" { - return "", errors.New("OAuth access token is missing") - } - return token, nil -} diff --git a/services/core/internal/store/mcp_credentials_oauth_test.go b/services/core/internal/store/mcp_credentials_oauth_test.go deleted file mode 100644 index 1d176013e..000000000 --- a/services/core/internal/store/mcp_credentials_oauth_test.go +++ /dev/null @@ -1,293 +0,0 @@ -package store - -import ( - "context" - "errors" - "reflect" - "strings" - "sync" - "sync/atomic" - "testing" - "time" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" - "github.com/google/uuid" -) - -func TestOAuthRefreshPersistsRotatedGrantAndRequest(t *testing.T) { - for _, method := range []string{"none", "client_secret_basic", "client_secret_post"} { - t.Run(method, func(t *testing.T) { - var request oauthrefresh.Request - expiry := time.Now().Add(time.Hour).UTC() - s, pool, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(_ context.Context, r oauthrefresh.Request) (oauthrefresh.Token, error) { - request = r - return oauthrefresh.Token{AccessToken: "renewed-access", RefreshToken: "rotated-refresh", ExpiresAt: &expiry}, nil - })) - input.OAuth.Refresh.TokenEndpointAuth = method - if method == "none" { - input.ClientSecret = "" - } - credential := createOAuthFixture(t, s, tenant, vault, input) - binding := oauthFixtureBinding(credential) - got, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, binding) - if err != nil || got != "renewed-access" { - t.Fatal("refresh did not return committed access", err) - } - want := oauthrefresh.Request{TokenEndpoint: input.OAuth.Refresh.TokenEndpoint, ClientID: input.OAuth.Refresh.ClientID, AuthMethod: method, ClientSecret: input.ClientSecret, RefreshToken: input.RefreshToken, Resource: input.OAuth.Refresh.Resource, Scope: input.OAuth.Refresh.Scope} - if !reflect.DeepEqual(request, want) { - t.Fatal("refresh request lost grant fields") - } - restarted := NewWithCredentialCipherAndOAuthRefresh(pool, s.credentialCipher, nil) - after := storedOAuthSecret(t, restarted, tenant, credential) - if after.AccessToken != "renewed-access" || after.RefreshToken != "rotated-refresh" || after.Metadata.ExpiresAt == nil || *after.Metadata.ExpiresAt != expiry.Format(time.RFC3339Nano) { - t.Fatal("refreshed grant not durable") - } - got, err = restarted.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, binding) - if err != nil || got != "renewed-access" { - t.Fatal("fresh grant needed another refresh after restart", err) - } - metadata, err := New(pool).GetCredential(t.Context(), tenant, vault.ID, credential.ID) - if err != nil || !reflect.DeepEqual(metadata.OAuth, &after.Metadata) { - t.Fatal("safe expiry metadata did not follow refresh", err) - } - }) - } -} - -func TestOAuthRefreshPreservesRefreshTokenWhenOmitted(t *testing.T) { - s, _, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { - return oauthrefresh.Token{AccessToken: "renewed-access"}, nil - })) - credential := createOAuthFixture(t, s, tenant, vault, input) - if _, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)); err != nil { - t.Fatal(err) - } - secret := storedOAuthSecret(t, s, tenant, credential) - if secret.RefreshToken != input.RefreshToken || secret.Metadata.ExpiresAt != nil { - t.Fatal("omitted refresh token or unknown expiry changed incorrectly") - } -} - -func TestOAuthFreshAndUnknownExpiryDoNotRefresh(t *testing.T) { - var requests atomic.Int32 - s, _, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { - requests.Add(1) - return oauthrefresh.Token{}, errors.New("unexpected refresh") - })) - for _, expiry := range []*string{nil, oauthString(time.Now().Add(time.Hour).UTC().Format(time.RFC3339Nano))} { - input.OAuth.ExpiresAt = expiry - credential := createOAuthFixture(t, s, tenant, vault, input) - if token, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)); err != nil || token != input.AccessToken { - t.Fatal("usable token was not returned", err) - } - } - if requests.Load() != 0 { - t.Fatal("non-expired grant refreshed") - } - input.AccessToken = "" - input.OAuth.ExpiresAt = nil - credential := createOAuthFixture(t, s, tenant, vault, input) - if token, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)); err == nil || token != "" { - t.Fatal("empty access token admitted") - } -} - -func TestOAuthConcurrentRefreshUsesOneCommittedGrant(t *testing.T) { - var requests atomic.Int32 - expiry := time.Now().Add(time.Hour) - s, pool, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { - requests.Add(1) - return oauthrefresh.Token{AccessToken: "concurrent-access", RefreshToken: "single-use-next", ExpiresAt: &expiry}, nil - })) - credential := createOAuthFixture(t, s, tenant, vault, input) - var workers sync.WaitGroup - errorsFound := make(chan error, 12) - start := make(chan struct{}) - for i := 0; i < 12; i++ { - workers.Add(1) - go func() { - defer workers.Done() - <-start - // Separate Store instances exercise PostgreSQL serialization, not a local lock. - reader := NewWithCredentialCipherAndOAuthRefresh(pool, s.credentialCipher, s.oauthRefresher) - token, err := reader.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)) - if err == nil && token != "concurrent-access" { - err = errors.New("concurrent lookup returned stale token") - } - errorsFound <- err - }() - } - close(start) - workers.Wait() - close(errorsFound) - for err := range errorsFound { - if err != nil { - t.Fatal(err) - } - } - if requests.Load() != 1 { - t.Fatal("one expiring grant triggered duplicate provider exchanges") - } -} - -func TestOAuthRefreshSerializesReplacementAndDeletion(t *testing.T) { - for _, mutation := range []string{"replacement", "credential-delete", "vault-delete"} { - t.Run(mutation, func(t *testing.T) { - entered := make(chan struct{}) - release := make(chan struct{}) - var releaseOnce sync.Once - unblock := func() { releaseOnce.Do(func() { close(release) }) } - defer unblock() - expiry := time.Now().Add(time.Hour) - s, pool, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(ctx context.Context, _ oauthrefresh.Request) (oauthrefresh.Token, error) { - close(entered) - select { - case <-release: - return oauthrefresh.Token{AccessToken: "refreshed-access", RefreshToken: "rotated-refresh", ExpiresAt: &expiry}, nil - case <-ctx.Done(): - return oauthrefresh.Token{}, ctx.Err() - } - })) - credential := createOAuthFixture(t, s, tenant, vault, input) - refreshed := make(chan error, 1) - go func() { - _, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)) - refreshed <- err - }() - select { - case <-entered: - case <-time.After(3 * time.Second): - t.Fatal("refresh never reached provider") - } - mutated := make(chan error, 1) - go func() { - var err error - switch mutation { - case "replacement": - _, err = s.UpdateOAuthCredential(t.Context(), tenant, vault.ID, credential.ID, UpdateOAuthCredentialInput{AccessToken: oauthString("manual-access"), Refresh: &OAuthRefreshUpdate{RefreshToken: oauthString("manual-refresh")}}) - case "credential-delete": - _, err = New(pool).DeleteCredential(t.Context(), tenant, vault.ID, credential.ID) - case "vault-delete": - _, err = New(pool).DeleteVault(t.Context(), tenant, vault.ID) - } - mutated <- err - }() - select { - case <-mutated: - t.Fatal("mutation bypassed pending refresh ownership") - case <-time.After(75 * time.Millisecond): - } - unblock() - for _, completed := range []chan error{refreshed, mutated} { - select { - case err := <-completed: - if err != nil { - t.Fatal("serialized operation failed", err) - } - case <-time.After(3 * time.Second): - t.Fatal("serialized operation deadlocked") - } - } - if mutation == "replacement" { - secret := storedOAuthSecret(t, s, tenant, credential) - if secret.AccessToken != "manual-access" || secret.RefreshToken != "manual-refresh" || secret.Metadata.ExpiresAt != nil { - t.Fatal("late refresh overwrote manual replacement") - } - } else { - if _, err := s.GetCredential(t.Context(), tenant, vault.ID, credential.ID); !errors.Is(err, ErrNotFound) { - t.Fatal("deleted credential resurrected") - } - if token, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)); !errors.Is(err, ErrNotFound) || token != "" { - t.Fatal("deleted grant remained usable") - } - } - }) - } -} - -func TestOAuthCancelledRefreshRollsBackAndAllowsReplacement(t *testing.T) { - entered := make(chan struct{}) - s, _, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(ctx context.Context, _ oauthrefresh.Request) (oauthrefresh.Token, error) { - close(entered) - <-ctx.Done() - return oauthrefresh.Token{}, ctx.Err() - })) - credential := createOAuthFixture(t, s, tenant, vault, input) - ctx, cancel := context.WithCancel(t.Context()) - defer cancel() - done := make(chan error, 1) - go func() { - _, err := s.MCPBearerToken(ctx, tenant, []string{vault.ID}, oauthFixtureBinding(credential)) - done <- err - }() - select { - case <-entered: - case <-time.After(3 * time.Second): - t.Fatal("refresh not started") - } - cancel() - select { - case err := <-done: - if err == nil { - t.Fatal("cancelled refresh succeeded") - } - case <-time.After(3 * time.Second): - t.Fatal("cancel did not release refresh") - } - if _, err := s.UpdateOAuthCredential(t.Context(), tenant, vault.ID, credential.ID, UpdateOAuthCredentialInput{AccessToken: oauthString("manual-after-cancel")}); err != nil { - t.Fatal("cancelled refresh retained ownership", err) - } - if token, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)); err != nil || token != "manual-after-cancel" { - t.Fatal("replacement after cancellation unusable") - } -} - -func TestOAuthDeletionCannotReselectAnotherCredential(t *testing.T) { - s, _, tenant, vault, input := oauthFixture(t, nil) - input.OAuth.ExpiresAt = nil - credential := createOAuthFixture(t, s, tenant, vault, input) - bindings, err := s.ResolveMCPCredentials(t.Context(), tenant, []string{vault.ID}, []MCPCredentialRequest{{ServerLabel: "test", ServerURL: input.MCPServerURL}}) - if err != nil { - t.Fatal(err) - } - if _, err := s.DeleteCredential(t.Context(), tenant, vault.ID, credential.ID); err != nil { - t.Fatal(err) - } - input.AccessToken = uuid.NewString() - createOAuthFixture(t, s, tenant, vault, input) - if token, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, bindings[0]); !errors.Is(err, ErrNotFound) || token != "" { - t.Fatal("frozen identity fell back after deletion") - } -} - -func TestOAuthRefreshCommitFailureDoesNotReturnUncommittedToken(t *testing.T) { - expiry := time.Now().Add(time.Hour) - s, pool, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { - return oauthrefresh.Token{AccessToken: "uncommitted-access", RefreshToken: "uncommitted-refresh", ExpiresAt: &expiry}, nil - })) - credential := createOAuthFixture(t, s, tenant, vault, input) - before := storedOAuthSecret(t, s, tenant, credential) - suffix := strings.ReplaceAll(uuid.NewString(), "-", "") - function, trigger := "oauth_commit_fail_"+suffix, "oauth_commit_fail_"+suffix - // A deferred trigger exercises failure after the UPDATE has returned metadata. - // The condition confines the synthetic failure to this test's unique grant. - if _, err := pool.Exec(t.Context(), "CREATE FUNCTION "+function+"() RETURNS trigger LANGUAGE plpgsql AS $$ BEGIN RAISE EXCEPTION 'private-refresh-canary'; END $$"); err != nil { - t.Fatal(err) - } - t.Cleanup(func() { - _, err := pool.Exec(context.Background(), "DROP FUNCTION "+function+"() CASCADE") - if err != nil { - t.Error(err) - } - }) - if _, err := pool.Exec(t.Context(), "CREATE CONSTRAINT TRIGGER "+trigger+" AFTER UPDATE ON vault_credentials DEFERRABLE INITIALLY DEFERRED FOR EACH ROW WHEN (NEW.id='"+credential.ID+"'::uuid) EXECUTE FUNCTION "+function+"()"); err != nil { - t.Fatal(err) - } - token, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)) - if err == nil || token != "" || strings.Contains(err.Error(), "private-refresh-canary") { - t.Fatal("commit failure returned an uncommitted grant or unsafe error") - } - if after := storedOAuthSecret(t, s, tenant, credential); !reflect.DeepEqual(before, after) { - t.Fatal("failed commit modified the grant") - } -} diff --git a/services/core/internal/store/mcp_credentials_test.go b/services/core/internal/store/mcp_credentials_test.go deleted file mode 100644 index cb76b56e5..000000000 --- a/services/core/internal/store/mcp_credentials_test.go +++ /dev/null @@ -1,138 +0,0 @@ -package store - -import ( - "bytes" - "crypto/rand" - "encoding/json" - "errors" - "strings" - "testing" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/google/uuid" -) - -func TestMCPCredentialSelectionAndScopedDecryption(t *testing.T) { - public, pool := testStore(t) - tenant, foreign := uuid.NewString(), uuid.NewString() - key := make([]byte, 32) - if _, err := rand.Read(key); err != nil { - t.Fatal(err) - } - cipher, err := credentialcrypto.New(key) - if err != nil { - t.Fatal(err) - } - s := NewWithCredentialCipher(pool, cipher) - var vaults []Vault - for _, owner := range []string{tenant, tenant, foreign} { - vault, err := s.CreateVault(t.Context(), owner, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - vaults = append(vaults, vault) - } - token := " \t" + uuid.NewString() + "雪\n" - destination := "https://mcp.example/tools" - create := func(vault Vault) Credential { - t.Helper() - value, err := s.CreateStaticCredential(t.Context(), vault.TenantID, vault.ID, CreateStaticCredentialInput{Name: "private", MCPServerURL: destination, Token: token}) - if err != nil { - t.Fatal(err) - } - return value - } - first, foreignCredential := create(vaults[0]), create(vaults[2]) - attached := []string{vaults[0].ID, vaults[1].ID} - requests := []MCPCredentialRequest{{ServerLabel: "tools", ServerURL: destination}, {ServerLabel: "anonymous", ServerURL: "https://anonymous.example/mcp"}} - bindings, err := public.ResolveMCPCredentials(t.Context(), tenant, attached, requests) - if err != nil || len(bindings) != 2 || bindings[0].CredentialID != first.ID || bindings[0].AuthType != "static_bearer" || bindings[1].CredentialID != "" { - t.Fatal("metadata selection or frozen anonymous decision differs", err) - } - encoded, err := json.Marshal(bindings) - if err != nil || bytes.Contains(encoded, []byte(token)) || strings.Contains(string(encoded), "ciphertext") { - t.Fatal("private binding contains secret material") - } - second := create(vaults[1]) - if _, err := public.ResolveMCPCredentials(t.Context(), tenant, attached, requests); !isSelectionError(err, true, "multiple attached vault credentials match MCP server_url "+destination+"; specify credential_id") { - t.Fatal("ambiguous selection was admitted", err) - } - requests[0].CredentialID = &second.ID - explicit, err := public.ResolveMCPCredentials(t.Context(), tenant, attached, requests) - if err != nil || explicit[0].CredentialID != second.ID || requests[1].CredentialID != nil { - t.Fatal("explicit selection did not disambiguate", err) - } - notAttached := func(id string) string { return "MCP credential_id " + id + " was not found in an attached vault" } - for _, tc := range []struct { - owner string - vaults []string - id, url string - message string // Empty for the unchanged Vault 404. - }{ - {tenant, attached, foreignCredential.ID, destination, notAttached(foreignCredential.ID)}, - {tenant, []string{vaults[1].ID}, first.ID, destination, notAttached(first.ID)}, - {tenant, attached, "not-a-credential", destination, notAttached("not-a-credential")}, - {tenant, attached, first.ID, destination + "/other", "MCP credential_id " + first.ID + " does not match server_url " + destination + "/other"}, - {foreign, attached, first.ID, destination, ""}, - {tenant, []string{vaults[0].ID, vaults[2].ID}, first.ID, destination, ""}, - {tenant, []string{uuid.NewString()}, first.ID, destination, ""}, - } { - _, err := public.ResolveMCPCredentials(t.Context(), tc.owner, tc.vaults, []MCPCredentialRequest{{ServerLabel: "tools", ServerURL: tc.url, CredentialID: &tc.id}}) - if tc.message == "" && !errors.Is(err, ErrNotFound) || tc.message != "" && !isSelectionError(err, false, tc.message) { - t.Fatal("unowned, unattached or wrong-destination selection was admitted", err) - } - } - if _, err := public.ResolveMCPCredentials(t.Context(), tenant, nil, []MCPCredentialRequest{{ServerLabel: "tools", ServerURL: destination, CredentialID: &first.ID}}); !isSelectionError(err, false, "MCP credential_id requires an attached vault") { - t.Fatal("a reference without attachments was admitted", err) - } - pool.Close() - public, pool = testStore(t) - cipher, _ = credentialcrypto.New(bytes.Clone(key)) - s = NewWithCredentialCipher(pool, cipher) - got, err := s.MCPBearerToken(t.Context(), tenant, attached, bindings[0]) - if err != nil || got != token { - t.Fatal("frozen selection or opaque bytes changed across restart", err) - } - if got, err := public.MCPBearerToken(t.Context(), tenant, attached, bindings[0]); !errors.Is(err, credentialcrypto.ErrUnavailable) || got != "" { - t.Fatal("missing key did not fail execution closed") - } - key[0] ^= 1 - wrong, _ := credentialcrypto.New(key) - if got, err := NewWithCredentialCipher(pool, wrong).MCPBearerToken(t.Context(), tenant, attached, bindings[0]); err == nil || got != "" || strings.Contains(err.Error(), token) { - t.Fatal("wrong key leaked or decrypted a credential") - } - for _, mutate := range []func(*MCPCredentialBinding){ - func(b *MCPCredentialBinding) { b.VaultID = vaults[1].ID }, - func(b *MCPCredentialBinding) { b.CredentialID = foreignCredential.ID }, - func(b *MCPCredentialBinding) { b.ServerURL += "/other" }, - func(b *MCPCredentialBinding) { b.AuthType = "other" }, - } { - binding := bindings[0] - mutate(&binding) - if got, err := s.MCPBearerToken(t.Context(), tenant, attached, binding); !errors.Is(err, ErrNotFound) || got != "" { - t.Fatal("substituted frozen authorization was decrypted") - } - } - for _, scope := range []struct { - owner string - vaults []string - }{{foreign, attached}, {tenant, []string{vaults[1].ID}}} { - if got, err := s.MCPBearerToken(t.Context(), scope.owner, scope.vaults, bindings[0]); !errors.Is(err, ErrNotFound) || got != "" { - t.Fatal("tenant or attachment authorization was bypassed") - } - } - if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET token_ciphertext=set_byte(token_ciphertext, 15, get_byte(token_ciphertext,15) # 1) WHERE id=$1", first.ID); err != nil { - t.Fatal(err) - } - if got, err := s.MCPBearerToken(t.Context(), tenant, attached, bindings[0]); err == nil || got != "" { - t.Fatal("tampered ciphertext decrypted") - } - if _, err := public.GetCredential(t.Context(), tenant, first.VaultID, first.ID); err != nil { - t.Fatal("safe metadata lookup depended on ciphertext", err) - } -} - -func isSelectionError(err error, conflict bool, message string) bool { - var selection *MCPCredentialSelectionError - return errors.As(err, &selection) && selection.Conflict == conflict && selection.Message == message -} diff --git a/services/core/internal/store/public_handler_fixture_test.go b/services/core/internal/store/public_handler_fixture_test.go index b3da7c695..03d3714a9 100644 --- a/services/core/internal/store/public_handler_fixture_test.go +++ b/services/core/internal/store/public_handler_fixture_test.go @@ -29,7 +29,7 @@ const testExecutorURL = "wss://core.example/api/v1/agent-daemon/ws" // publicHandler serves s through api.NewHandler. s backs every area the Store // implements, and db is the database and credential key that built s; the -// audit reads, Agents and Files come from db. keys authenticate as Project +// audit reads, Agents, Files and Vaults come from db. keys authenticate as Project // keys and "admin" as the Core key. Metrics, Runtime observation and history, // and executor connections are strict stand-ins. Execution and Sandboxes stay // disabled unless configure sets them. @@ -43,9 +43,14 @@ func publicHandler(t testing.TB, s *store.Store, db fixtureDB, keys fixtureKeyRe audit := auditpg.New(pgunit.NewPool(db.pool)) agentStore, agentService := fixtureAgents(t, db) fileStore, fileService := fixtureFiles(t, db) + vaultStore, vaultService, err := fixtureVaults(db) + if err != nil { + return nil, err + } deps := api.Dependencies{ Engine: engine, CoreKeys: admin, InstallationBindings: s, - Projects: fixtureProjects{Store: s, keys: keys}, Vaults: s, ModelProviders: s, Skills: s, + Projects: fixtureProjects{Store: s, keys: keys}, ModelProviders: s, Skills: s, + Vaults: vaultService, VaultsReader: vaultStore, Files: fileService, FilesReader: fileStore, Agents: agentService, AgentsReader: agentStore, EnvironmentTemplates: s, Sessions: s, SessionEvents: s, SessionHistory: s, Subagents: s, diff --git a/services/core/internal/store/remote_mcp_credentials_test.go b/services/core/internal/store/remote_mcp_credentials_test.go index 60eea2c44..fb3f3f1d8 100644 --- a/services/core/internal/store/remote_mcp_credentials_test.go +++ b/services/core/internal/store/remote_mcp_credentials_test.go @@ -8,6 +8,7 @@ import ( "testing" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) func TestSelfHostedServiceMCPRejectionDoesNotRequireCredentialDecryption(t *testing.T) { @@ -18,7 +19,11 @@ func TestSelfHostedServiceMCPRejectionDoesNotRequireCredentialDecryption(t *test case "missing key": s, db = store.New(db.pool), fixtureDB{pool: db.pool} case "deleted": - if _, err := s.DeleteCredential(t.Context(), tenant, vault.ID, credential.ID); err != nil { + _, service, err := fixtureVaults(db) + if err != nil { + t.Fatal(err) + } + if _, err := service.DeleteCredential(t.Context(), vaults.DeleteCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: credential.ID}); err != nil { t.Fatal(err) } case "tampered": diff --git a/services/core/internal/store/remote_mcp_test.go b/services/core/internal/store/remote_mcp_test.go index 0c324b1ef..841b76f82 100644 --- a/services/core/internal/store/remote_mcp_test.go +++ b/services/core/internal/store/remote_mcp_test.go @@ -10,6 +10,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimedevice" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" "github.com/google/uuid" "github.com/jackc/pgx/v5/pgxpool" ) @@ -72,7 +73,7 @@ func TestSelfHostedServiceMCPRejectedWithoutWrites(t *testing.T) { } } -func selfHostedMCPAdmissionFixture(t *testing.T) (*store.Store, fixtureDB, string, store.Vault, store.Credential) { +func selfHostedMCPAdmissionFixture(t *testing.T) (*store.Store, fixtureDB, string, vaults.Vault, vaults.Credential) { t.Helper() _, pool := store.NewTestStore(t) cipher, err := credentialcrypto.New([]byte(strings.Repeat("k", 32))) @@ -80,12 +81,16 @@ func selfHostedMCPAdmissionFixture(t *testing.T) (*store.Store, fixtureDB, strin t.Fatal(err) } s, db := store.NewWithCredentialCipher(pool, cipher), fixtureDB{pool: pool, cipher: cipher} + _, service, err := fixtureVaults(db) + if err != nil { + t.Fatal(err) + } tenant := uuid.NewString() - vault, err := s.CreateVault(t.Context(), tenant, store.CreateVaultInput{}) + vault, err := service.CreateVault(t.Context(), vaults.CreateVault{TenantID: tenant}) if err != nil { t.Fatal(err) } - credential, err := s.CreateStaticCredential(t.Context(), tenant, vault.ID, store.CreateStaticCredentialInput{Name: "test", MCPServerURL: "https://tools.example/mcp", Token: "synthetic-token"}) + credential, err := service.CreateStaticCredential(t.Context(), vaults.CreateStaticCredential{TenantID: tenant, VaultID: vault.ID, Name: "test", MCPServerURL: "https://tools.example/mcp", Token: "synthetic-token"}) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/sessions.go b/services/core/internal/store/sessions.go index af0f8a0f5..c81dcc9af 100644 --- a/services/core/internal/store/sessions.go +++ b/services/core/internal/store/sessions.go @@ -24,7 +24,6 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/identity" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/jsonobject" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/metadata" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" ) @@ -95,7 +94,6 @@ type Store struct { writer transactor lease *pgunit.Lease credentialCipher *credentialcrypto.Cipher - oauthRefresher oauthrefresh.Refresher // publicURL is OAC_PUBLIC_URL. Core derives every address it gives // nodes, sandboxes and administrators from it. publicURL string @@ -293,10 +291,6 @@ func parseID(value string) (pgtype.UUID, error) { return id, nil } -// pathID resolves a caller-supplied path identifier with pgunit.PathID for an -// operation that passes it on as a string. -func pathID(value string) string { return uuid.UUID(pgunit.PathID(value).Bytes).String() } - func sessionFromRow(row sqlc.Session) (Session, error) { session := Session{ID: uuid.UUID(row.ID.Bytes).String(), TenantID: uuid.UUID(row.TenantID.Bytes).String(), Engine: row.Engine, CreatedAt: row.CreatedAt.Time, RequiredActions: []v1.FunctionCallAction{}} creator, err := sessionCreator(row.CreatorKind, row.CreatorID) diff --git a/services/core/internal/store/vault_credentials.go b/services/core/internal/store/vault_credentials.go deleted file mode 100644 index 870085af1..000000000 --- a/services/core/internal/store/vault_credentials.go +++ /dev/null @@ -1,127 +0,0 @@ -package store - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "time" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/writeaudit" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgtype" - "github.com/jackc/pgx/v5/pgxpool" -) - -// Credential contains only public metadata. Secret ciphertext is never selected -// by resource reads; decryption belongs to scoped execution lookup only. -type Credential struct { - ID, VaultID, Name, AuthType, MCPServerURL string - CreatedAt, UpdatedAt time.Time - OAuth *OAuthMetadata -} - -type CreateStaticCredentialInput struct { - Name, MCPServerURL, Token string -} - -// NewWithCredentialCipher configures immutable credential encryption before the -// Store is published. A nil cipher leaves non-secret resource operations available. -func NewWithCredentialCipher(pool *pgxpool.Pool, cipher *credentialcrypto.Cipher) *Store { - refresher, _ := oauthrefresh.NewClient(nil) - return NewWithCredentialCipherAndOAuthRefresh(pool, cipher, refresher) -} - -func (s *Store) CreateStaticCredential(ctx context.Context, tenantID, vaultID string, input CreateStaticCredentialInput) (Credential, error) { - tenant, err := parseID(tenantID) - if err != nil { - return Credential{}, err - } - vault := pgunit.PathID(vaultID) - if !validVaultName(input.Name) || input.MCPServerURL == "" { - return Credential{}, ErrInvalidInput - } - if s.credentialCipher == nil { - return Credential{}, credentialcrypto.ErrUnavailable - } - id := uuid.New() - binding := credentialcrypto.Binding{TenantID: uuid.UUID(tenant.Bytes).String(), VaultID: uuid.UUID(vault.Bytes).String(), CredentialID: id.String(), AuthType: "static_bearer", Destination: input.MCPServerURL} - ciphertext, err := s.credentialCipher.Seal([]byte(input.Token), binding) - if err != nil { - return Credential{}, errors.New("credential encryption failed") - } - var created Credential - err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { - q := s.queries.WithTx(tx) - row, err := q.CreateStaticCredential(ctx, sqlc.CreateStaticCredentialParams{ - ID: pgtype.UUID{Bytes: id, Valid: true}, TenantID: tenant, VaultID: vault, - Name: input.Name, McpServerUrl: input.MCPServerURL, TokenCiphertext: ciphertext, - }) - if err != nil { - return err - } - created, err = credentialFromRow(sqlc.GetCredentialRow(row)) - if err != nil { - return err - } - return auditpg.RecordWriteAudit(ctx, q, tenantID, "create", "credential", created.ID, created.VaultID, writeaudit.Resource{Type: "credential", ID: created.ID, ParentID: created.VaultID}) - }) - if errors.Is(err, pgx.ErrNoRows) { - return Credential{}, ErrNotFound - } - if err != nil { - return Credential{}, fmt.Errorf("create credential: %w", err) - } - return created, nil -} - -func (s *Store) GetCredential(ctx context.Context, tenantID, vaultID, credentialID string) (Credential, error) { - tenant, err := parseID(tenantID) - if err != nil { - return Credential{}, err - } - vault, err := parseID(vaultID) - if err != nil { - return Credential{}, err - } - id, err := parseID(credentialID) - if err != nil { - return Credential{}, err - } - row, err := s.queries.GetCredential(ctx, sqlc.GetCredentialParams{TenantID: tenant, VaultID: vault, ID: id}) - if errors.Is(err, pgx.ErrNoRows) { - return Credential{}, ErrNotFound - } - if err != nil { - return Credential{}, fmt.Errorf("get credential: %w", err) - } - return credentialFromRow(row) -} - -func credentialFromRow(row sqlc.GetCredentialRow) (Credential, error) { - result := Credential{ - ID: uuid.UUID(row.ID.Bytes).String(), VaultID: uuid.UUID(row.VaultID.Bytes).String(), - Name: row.Name, AuthType: row.AuthType, MCPServerURL: row.McpServerUrl, - CreatedAt: row.CreatedAt.Time, UpdatedAt: row.UpdatedAt.Time, - } - if row.AuthType == "mcp_oauth" { - result.OAuth = &OAuthMetadata{} - if err := json.Unmarshal(row.OauthMetadata, result.OAuth); err != nil { - return Credential{}, errors.New("invalid stored OAuth metadata") - } - } - return result, nil -} - -// NewWithCredentialCipherAndOAuthRefresh configures the execution-only refresh boundary. -func NewWithCredentialCipherAndOAuthRefresh(pool *pgxpool.Pool, cipher *credentialcrypto.Cipher, refresher oauthrefresh.Refresher) *Store { - s := New(pool) - s.credentialCipher, s.oauthRefresher = cipher, refresher - return s -} diff --git a/services/core/internal/store/vault_credentials_delete.go b/services/core/internal/store/vault_credentials_delete.go deleted file mode 100644 index 8fb28bb8c..000000000 --- a/services/core/internal/store/vault_credentials_delete.go +++ /dev/null @@ -1,44 +0,0 @@ -package store - -import ( - "context" - "errors" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" -) - -// DeleteCredential removes the stored secret without reading or decrypting it. -func (s *Store) DeleteCredential(ctx context.Context, tenantID, vaultID, credentialID string) (string, error) { - tenant, err := parseID(tenantID) - if err != nil { - return "", ErrNotFound - } - vault, err := parseID(vaultID) - if err != nil { - return "", ErrNotFound - } - id, err := parseID(credentialID) - if err != nil { - return "", ErrNotFound - } - var deletedID string - err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { - q := s.queries.WithTx(tx) - deleted, err := q.DeleteCredential(ctx, sqlc.DeleteCredentialParams{TenantID: tenant, VaultID: vault, ID: id}) - if err != nil { - return err - } - deletedID = uuid.UUID(deleted.Bytes).String() - return auditpg.RecordWriteAudit(ctx, q, tenantID, "delete", "credential", deletedID, uuid.UUID(vault.Bytes).String()) - }) - if errors.Is(err, pgx.ErrNoRows) { - return "", ErrNotFound - } - if err != nil { - return "", errors.New("credential deletion failed") - } - return deletedID, nil -} diff --git a/services/core/internal/store/vault_credentials_delete_test.go b/services/core/internal/store/vault_credentials_delete_test.go deleted file mode 100644 index 1adcf0a10..000000000 --- a/services/core/internal/store/vault_credentials_delete_test.go +++ /dev/null @@ -1,151 +0,0 @@ -package store - -import ( - "bytes" - "encoding/json" - "errors" - "reflect" - "testing" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/google/uuid" -) - -func TestCredentialDeletionScopeBindingAndRestart(t *testing.T) { - public, pool := testStore(t) - tenant, foreign := uuid.NewString(), uuid.NewString() - key := bytes.Repeat([]byte{41}, 32) - cipher, err := credentialcrypto.New(key) - if err != nil { - t.Fatal(err) - } - s := NewWithCredentialCipher(pool, cipher) - vault, err := s.CreateVault(t.Context(), tenant, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - wrong, err := s.CreateVault(t.Context(), tenant, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - create := func(name string) Credential { - t.Helper() - value, err := s.CreateStaticCredential(t.Context(), tenant, vault.ID, CreateStaticCredentialInput{Name: name, MCPServerURL: "https://mcp.example/tools", Token: name + "-secret"}) - if err != nil { - t.Fatal(err) - } - return value - } - original := create("original") - attached := []string{vault.ID} - selected, err := public.ResolveMCPCredentials(t.Context(), tenant, attached, []MCPCredentialRequest{{ServerLabel: "tools", ServerURL: original.MCPServerURL}}) - if err != nil || len(selected) != 1 { - t.Fatal("initial automatic selection failed", err) - } - configuration, _ := json.Marshal(map[string]any{"agent": map[string]string{"model": "model"}, "vault_ids": attached, "mcp_credentials": selected}) - input := CreateSessionInput{Creator: FixtureCreator(), Engine: "codex", IdempotencyKey: "retained", Configuration: configuration} - session, err := s.CreateSession(t.Context(), tenant, input) - if err != nil { - t.Fatal(err) - } - retained, err := s.MCPBearerToken(t.Context(), tenant, attached, selected[0]) - if err != nil || retained != "original-secret" { - t.Fatal("pre-delete dispatch lookup failed") - } - sibling := create("sibling") - for _, scope := range []struct{ tenant, vault, id string }{ - {foreign, vault.ID, original.ID}, {tenant, wrong.ID, original.ID}, - {tenant, vault.ID, uuid.NewString()}, {tenant, "invalid", original.ID}, {tenant, vault.ID, "invalid"}, - } { - if _, err := public.DeleteCredential(t.Context(), scope.tenant, scope.vault, scope.id); !errors.Is(err, ErrNotFound) { - t.Fatal("foreign or invalid delete was accepted", err) - } - } - // An actual database write failure must leave the resource and token intact. - readOnly := readOnlyResourceStore(t, pool) - _, deletionErr := readOnly.DeleteCredential(t.Context(), tenant, vault.ID, original.ID) - if deletionErr == nil || deletionErr.Error() != "credential deletion failed" { - t.Fatal("failed mutation was accepted or exposed") - } - if value, err := public.GetCredential(t.Context(), tenant, vault.ID, original.ID); err != nil || !reflect.DeepEqual(value, original) { - t.Fatal("rejected deletion changed the resource", err) - } - if token, err := s.MCPBearerToken(t.Context(), tenant, attached, selected[0]); err != nil || token != retained { - t.Fatal("rejected deletion changed the stored token") - } - // Delete through a keyless Store, even if the stored payload is damaged. - if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET token_ciphertext=decode('00','hex') WHERE id=$1", original.ID); err != nil { - t.Fatal(err) - } - if id, err := public.DeleteCredential(t.Context(), tenant, vault.ID, original.ID); err != nil || id != original.ID { - t.Fatal("keyless deletion failed", err) - } - var count int - if err := pool.QueryRow(t.Context(), "SELECT count(*) FROM vault_credentials WHERE id=$1", original.ID).Scan(&count); err != nil || count != 0 { - t.Fatal("deleted row or ciphertext remains") - } - pool.Close() - public, pool = testStore(t) - cipher, _ = credentialcrypto.New(bytes.Clone(key)) - s = NewWithCredentialCipher(pool, cipher) - if _, err := public.DeleteCredential(t.Context(), tenant, vault.ID, original.ID); !errors.Is(err, ErrNotFound) { - t.Fatal("repeat deletion did not stay absent") - } - if _, err := public.GetCredential(t.Context(), tenant, vault.ID, original.ID); !errors.Is(err, ErrNotFound) { - t.Fatal("deleted metadata reappeared after restart") - } - if _, err := s.UpdateStaticCredential(t.Context(), tenant, vault.ID, original.ID, UpdateStaticCredentialInput{Token: "replacement"}); !errors.Is(err, ErrNotFound) { - t.Fatal("replacement resurrected a deleted credential") - } - if _, err := s.MCPBearerToken(t.Context(), tenant, attached, selected[0]); !errors.Is(err, ErrNotFound) { - t.Fatal("frozen selection fell back to another token") - } - if _, err := s.ResolveMCPCredentials(t.Context(), tenant, attached, []MCPCredentialRequest{{ServerLabel: "tools", ServerURL: original.MCPServerURL, CredentialID: &original.ID}}); !isSelectionError(err, false, "MCP credential_id "+original.ID+" was not found in an attached vault") { - t.Fatal("deleted explicit selection was admitted", err) - } - page, err := public.ListCredentials(t.Context(), tenant, vault.ID, "", 100, true, []string{"active", "archived"}) - if err != nil || len(page.Credentials) != 1 || !reflect.DeepEqual(page.Credentials[0], sibling) { - t.Fatal("deletion changed a sibling or list membership", err) - } - if value, err := public.GetVault(t.Context(), tenant, vault.ID); err != nil || !reflect.DeepEqual(value, vault) { - t.Fatal("deletion changed its parent Vault") - } - if retry, err := public.CreateSession(t.Context(), tenant, input); err != nil || retry.ID != session.ID || !bytes.Equal(retry.Configuration, session.Configuration) { - t.Fatal("deletion changed frozen creation identity", err) - } -} - -func TestCredentialDeletionConcurrentReplacementCannotResurrect(t *testing.T) { - public, pool := testStore(t) - tenant := uuid.NewString() - cipher, err := credentialcrypto.New(bytes.Repeat([]byte{42}, 32)) - if err != nil { - t.Fatal(err) - } - s := NewWithCredentialCipher(pool, cipher) - vault, err := s.CreateVault(t.Context(), tenant, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - for range 8 { - value, err := s.CreateStaticCredential(t.Context(), tenant, vault.ID, CreateStaticCredentialInput{Name: "competing", MCPServerURL: "https://mcp.example/tools", Token: "before"}) - if err != nil { - t.Fatal(err) - } - start, updated := make(chan struct{}), make(chan error, 1) - go func() { - <-start - _, err := s.UpdateStaticCredential(t.Context(), tenant, vault.ID, value.ID, UpdateStaticCredentialInput{Token: "after"}) - updated <- err - }() - close(start) - _, deleted := public.DeleteCredential(t.Context(), tenant, vault.ID, value.ID) - updateErr := <-updated - if deleted != nil || updateErr != nil && !errors.Is(updateErr, ErrNotFound) { - t.Fatal("competing update/delete failed unexpectedly", deleted, updateErr) - } - if _, err := public.GetCredential(t.Context(), tenant, vault.ID, value.ID); !errors.Is(err, ErrNotFound) { - t.Fatal("concurrent update resurrected deleted metadata") - } - } -} diff --git a/services/core/internal/store/vault_credentials_list.go b/services/core/internal/store/vault_credentials_list.go deleted file mode 100644 index 082dfa5aa..000000000 --- a/services/core/internal/store/vault_credentials_list.go +++ /dev/null @@ -1,64 +0,0 @@ -package store - -import ( - "context" - "fmt" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" - "github.com/google/uuid" - "github.com/jackc/pgx/v5/pgtype" -) - -type CredentialPage struct { - Credentials []Credential - NextCursor string -} - -func (s *Store) ListCredentials(ctx context.Context, tenantID, vaultID, cursor string, limit int, ascending bool, statuses []string) (CredentialPage, error) { - vaultID = pathID(vaultID) - // An inaccessible parent is not an authorized empty collection. - vault, err := s.GetVault(ctx, tenantID, vaultID) - if err != nil { - return CredentialPage{}, err - } - if limit < 1 || limit > 100 { - return CredentialPage{}, fmt.Errorf("%w: internal page size must be 1..100", ErrInvalidInput) - } - if len(statuses) == 0 { - statuses = []string{"active", "archived"} - } - for _, status := range statuses { - if status != "active" && status != "archived" { - return CredentialPage{}, fmt.Errorf("%w: invalid Credential status", ErrInvalidInput) - } - } - tenant, _ := parseID(vault.TenantID) - parent, _ := parseID(vault.ID) - params := sqlc.ListCredentialsParams{TenantID: tenant, VaultID: parent, PageLimit: int32(limit + 1), AfterID: pgtype.UUID{Valid: true}, Ascending: ascending, Statuses: statuses} - if cursor != "" { - after, err := s.GetCredential(ctx, tenantID, vaultID, pgunit.LookupCursor(cursor)) - if err != nil { - return CredentialPage{}, err - } - params.AfterCreated = pgtype.Timestamptz{Time: after.CreatedAt, Valid: true} - params.AfterID, _ = parseID(after.ID) - } - rows, err := s.queries.ListCredentials(ctx, params) - if err != nil { - return CredentialPage{}, fmt.Errorf("list credentials: %w", err) - } - page := CredentialPage{Credentials: make([]Credential, 0, min(limit, len(rows)))} - if len(rows) > limit { - page.NextCursor = uuid.UUID(rows[limit-1].ID.Bytes).String() - rows = rows[:limit] - } - for _, row := range rows { - credential, err := credentialFromRow(sqlc.GetCredentialRow(row)) - if err != nil { - return CredentialPage{}, err - } - page.Credentials = append(page.Credentials, credential) - } - return page, nil -} diff --git a/services/core/internal/store/vault_credentials_list_test.go b/services/core/internal/store/vault_credentials_list_test.go deleted file mode 100644 index 57bd8d9c4..000000000 --- a/services/core/internal/store/vault_credentials_list_test.go +++ /dev/null @@ -1,150 +0,0 @@ -package store - -import ( - "cmp" - "errors" - "reflect" - "slices" - "testing" - "time" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/google/uuid" -) - -func TestCredentialListFilteringOwnershipAndKeylessReconnect(t *testing.T) { - reader, pool := testStore(t) - ctx := t.Context() - tenant, foreign := uuid.NewString(), uuid.NewString() - var vaults []Vault - for _, owner := range []string{tenant, tenant, foreign, tenant} { - vault, err := reader.CreateVault(ctx, owner, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - vaults = append(vaults, vault) - } - // Parent classification does not classify its Credentials. - if _, err := pool.Exec(ctx, "UPDATE vaults SET status='archived' WHERE id=$1", vaults[0].ID); err != nil { - t.Fatal(err) - } - cipher, err := credentialcrypto.New(make([]byte, 32)) - if err != nil { - t.Fatal(err) - } - writer := NewWithCredentialCipher(pool, cipher) - create := func(owner, vault string) Credential { - t.Helper() - c, err := writer.CreateStaticCredential(ctx, owner, vault, CreateStaticCredentialInput{Name: "List fixture", MCPServerURL: "https://example.invalid/mcp", Token: "synthetic-token-not-public"}) - if err != nil { - t.Fatal(err) - } - return c - } - var all []Credential - archived := map[string]bool{} - for i := range 105 { - c := create(tenant, vaults[0].ID) - status := "active" - if i%3 == 0 { - status, archived[c.ID] = "archived", true - } - if _, err := pool.Exec(ctx, "UPDATE vault_credentials SET status=$1, created_at=$2 WHERE id=$3", status, time.Unix(1700000000+int64(i%2), 0), c.ID); err != nil { - t.Fatal(err) - } - c, err = reader.GetCredential(ctx, tenant, vaults[0].ID, c.ID) - if err != nil { - t.Fatal(err) - } - all = append(all, c) - } - slices.SortFunc(all, func(a, b Credential) int { - if c := a.CreatedAt.Compare(b.CreatedAt); c != 0 { - return c - } - return cmp.Compare(a.ID, b.ID) - }) - otherVault, otherProject := create(tenant, vaults[1].ID), create(foreign, vaults[2].ID) - snapshot := func(s *Store) string { - t.Helper() - var value string - if err := s.pool.QueryRow(ctx, "SELECT jsonb_agg(to_jsonb(c) ORDER BY id)::text FROM vault_credentials c WHERE vault_id=$1", vaults[0].ID).Scan(&value); err != nil { - t.Fatal(err) - } - return value - } - before := snapshot(reader) - read := func(s *Store, ascending bool, statuses []string, size int) []Credential { - t.Helper() - actual := []Credential{} - cursor := "" - for { - page, err := s.ListCredentials(ctx, tenant, vaults[0].ID, cursor, size, ascending, statuses) - if err != nil || len(page.Credentials) == 0 || len(page.Credentials) > size { - t.Fatal("invalid page", err) - } - actual = append(actual, page.Credentials...) - if len(actual) > len(all) { - t.Fatal("repeated pagination") - } - if page.NextCursor == "" { - break - } - if page.NextCursor != page.Credentials[len(page.Credentials)-1].ID { - t.Fatal("cursor is not last included Credential") - } - cursor = page.NextCursor - } - return actual - } - for _, ascending := range []bool{true, false} { - for _, statuses := range [][]string{nil, {"active"}, {"archived"}, {"active", "archived"}} { - want := []Credential{} - for _, c := range all { - if len(statuses) != 1 || archived[c.ID] == (statuses[0] == "archived") { - want = append(want, c) - } - } - if !ascending { - slices.Reverse(want) - } - for _, size := range []int{20, 100} { - if got := read(reader, ascending, statuses, size); !reflect.DeepEqual(got, want) { - t.Fatalf("metadata/filter/order mismatch: ascending=%t statuses=%v size=%d", ascending, statuses, size) - } - } - } - } - for _, tc := range []struct{ owner, vault, cursor string }{ - {tenant, vaults[0].ID, otherVault.ID}, {tenant, vaults[0].ID, otherProject.ID}, {tenant, vaults[0].ID, uuid.NewString()}, {tenant, vaults[0].ID, "invalid"}, - {tenant, vaults[2].ID, ""}, {foreign, vaults[0].ID, ""}, {tenant, uuid.NewString(), ""}, - } { - if _, err := reader.ListCredentials(ctx, tc.owner, tc.vault, tc.cursor, 20, false, nil); !errors.Is(err, ErrNotFound) { - t.Fatal("unowned/unknown parent or cursor accepted", err) - } - } - for _, tc := range []struct{ vault, cursor string }{{vaults[0].ID, all[len(all)-1].ID}, {vaults[3].ID, ""}} { - page, err := reader.ListCredentials(ctx, tenant, tc.vault, tc.cursor, 20, true, nil) - if err != nil || page.Credentials == nil || len(page.Credentials) != 0 || page.NextCursor != "" { - t.Fatal("empty/terminal page", err) - } - } - for _, tc := range []struct { - owner, vault, cursor string - limit int - statuses []string - }{{"invalid", vaults[0].ID, "", 20, nil}, {tenant, vaults[0].ID, "", 0, nil}, {tenant, vaults[0].ID, "", 101, nil}, {tenant, vaults[0].ID, "", 20, []string{"deleted"}}} { - if _, err := reader.ListCredentials(ctx, tc.owner, tc.vault, tc.cursor, tc.limit, false, tc.statuses); !errors.Is(err, ErrInvalidInput) { - t.Fatal("invalid internal query accepted", err) - } - } - // A malformed Vault path identifier follows the missing-Vault path. - if _, err := reader.ListCredentials(ctx, tenant, "invalid", "", 20, false, nil); !errors.Is(err, ErrNotFound) { - t.Fatal("malformed Vault was not missing", err) - } - pool.Close() - reopened, _ := testStore(t) - if got := read(reopened, true, nil, 20); !reflect.DeepEqual(got, all) || snapshot(reopened) != before { - t.Fatal("keyless reads/restart changed metadata, classification or ciphertext") - } -} diff --git a/services/core/internal/store/vault_credentials_oauth.go b/services/core/internal/store/vault_credentials_oauth.go deleted file mode 100644 index c2016132f..000000000 --- a/services/core/internal/store/vault_credentials_oauth.go +++ /dev/null @@ -1,189 +0,0 @@ -package store - -import ( - "context" - "errors" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/writeaudit" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgtype" -) - -func (s *Store) CreateOAuthCredential(ctx context.Context, tenantID, vaultID string, input CreateOAuthCredentialInput) (Credential, error) { - tenant, err := parseID(tenantID) - if err != nil { - return Credential{}, ErrNotFound - } - vault := pgunit.PathID(vaultID) - if !validVaultName(input.Name) || input.MCPServerURL == "" || !validOAuthMetadata(input.OAuth) { - return Credential{}, ErrInvalidInput - } - if input.OAuth.Refresh == nil && (input.RefreshToken != "" || input.ClientSecret != "") || - input.OAuth.Refresh != nil && input.OAuth.Refresh.TokenEndpointAuth == "none" && input.ClientSecret != "" { - return Credential{}, ErrInvalidInput - } - credential := Credential{ID: uuid.NewString(), VaultID: uuid.UUID(vault.Bytes).String(), - Name: input.Name, AuthType: "mcp_oauth", MCPServerURL: input.MCPServerURL} - secret := oauthSecret{Version: 1, Metadata: input.OAuth, AccessToken: input.AccessToken, - RefreshToken: input.RefreshToken, ClientSecret: input.ClientSecret} - metadata, ciphertext, err := s.sealOAuth(uuid.UUID(tenant.Bytes).String(), credential, secret) - if err != nil { - return Credential{}, err - } - var created Credential - err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { - q := s.queries.WithTx(tx) - row, err := q.CreateOAuthCredential(ctx, sqlc.CreateOAuthCredentialParams{ - ID: pgtype.UUID{Bytes: uuid.MustParse(credential.ID), Valid: true}, TenantID: tenant, VaultID: vault, - Name: input.Name, McpServerUrl: input.MCPServerURL, OauthMetadata: metadata, TokenCiphertext: ciphertext, - }) - if err != nil { - return err - } - created, err = credentialFromRow(sqlc.GetCredentialRow(row)) - if err != nil { - return err - } - return auditpg.RecordWriteAudit(ctx, q, tenantID, "create", "credential", created.ID, created.VaultID, writeaudit.Resource{Type: "credential", ID: created.ID, ParentID: created.VaultID}) - }) - if errors.Is(err, pgx.ErrNoRows) { - return Credential{}, ErrNotFound - } - if err != nil { - return Credential{}, errors.New("credential creation failed") - } - return created, nil -} - -func (s *Store) UpdateOAuthCredential(ctx context.Context, tenantID, vaultID, credentialID string, input UpdateOAuthCredentialInput) (Credential, error) { - vaultID, credentialID = pathID(vaultID), pathID(credentialID) - current, err := s.GetCredential(ctx, tenantID, vaultID, credentialID) - if err != nil { - return Credential{}, err - } - if current.AuthType != "mcp_oauth" { - return Credential{}, ErrInvalidInput - } - var updated Credential - err = s.withOAuth(ctx, tenantID, vaultID, credentialID, "", "credential update failed", func(ctx context.Context, tx pgx.Tx, credential Credential, secret oauthSecret) error { - if input.AccessToken != nil { - secret.AccessToken = *input.AccessToken - secret.Metadata.ExpiresAt = nil - } - if input.ExpiresAtSet { - secret.Metadata.ExpiresAt = input.ExpiresAt - } - if update := input.Refresh; update != nil { - refresh := secret.Metadata.Refresh - if refresh == nil { - return ErrInvalidInput - } - if update.TokenEndpointAuthType != "" && update.TokenEndpointAuthType != refresh.TokenEndpointAuth { - return ErrInvalidInput - } - if update.ClientSecret != nil { - if refresh.TokenEndpointAuth == "none" { - return ErrInvalidInput - } - secret.ClientSecret = *update.ClientSecret - } - if update.RefreshToken != nil { - secret.RefreshToken = *update.RefreshToken - } - if update.ScopeSet { - refresh.Scope = update.Scope - } - } - var err error - updated, err = s.saveOAuth(ctx, tx, tenantID, credential, secret) - if err != nil { - return err - } - if err := auditpg.RecordWriteAudit(ctx, s.queries.WithTx(tx), tenantID, "update", "credential", updated.ID, updated.VaultID); err != nil { - return errors.New("credential update failed") - } - return nil - }) - if err != nil { - return Credential{}, err - } - return updated, nil -} - -// withOAuth locks the credential row for the whole of apply, including an -// external refresh, and commits only when apply succeeds. Row ownership lasts -// through refresh or manual replacement. PostgreSQL serializes competing updates -// and deletes, including a parent Vault's cascading deletion. A failure to begin -// or commit is reported as failure, never with database error text. -func (s *Store) withOAuth(ctx context.Context, tenantID, vaultID, credentialID, destination, failure string, apply func(context.Context, pgx.Tx, Credential, oauthSecret) error) error { - tenant, e1 := parseID(tenantID) - vault, e2 := parseID(vaultID) - id, e3 := parseID(credentialID) - if e1 != nil || e2 != nil || e3 != nil { - return ErrNotFound - } - var applied error - err := s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { - credential, secret, err := s.lockOAuth(ctx, tx, tenant, vault, id, destination) - if err == nil { - err = apply(ctx, tx, credential, secret) - } - applied = err - return err - }) - if err != nil && applied == nil { - return errors.New(failure) - } - return err -} - -func (s *Store) lockOAuth(ctx context.Context, tx pgx.Tx, tenant, vault, id pgtype.UUID, destination string) (Credential, oauthSecret, error) { - row, err := sqlc.New(tx).GetOAuthCredentialForUpdate(ctx, sqlc.GetOAuthCredentialForUpdateParams{ - TenantID: tenant, VaultID: vault, ID: id, - }) - if errors.Is(err, pgx.ErrNoRows) { - return Credential{}, oauthSecret{}, ErrNotFound - } - if err != nil { - return Credential{}, oauthSecret{}, errors.New("credential lookup failed") - } - credential, err := credentialFromRow(sqlc.GetCredentialRow{ID: row.ID, VaultID: row.VaultID, - Name: row.Name, AuthType: row.AuthType, McpServerUrl: row.McpServerUrl, - CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt, OauthMetadata: row.OauthMetadata}) - if err != nil { - return Credential{}, oauthSecret{}, err - } - if destination != "" && credential.MCPServerURL != destination { - return Credential{}, oauthSecret{}, ErrNotFound - } - secret, err := s.openOAuth(uuid.UUID(tenant.Bytes).String(), credential, row.TokenCiphertext) - if err != nil { - return Credential{}, oauthSecret{}, err - } - return credential, secret, nil -} - -func (s *Store) saveOAuth(ctx context.Context, tx pgx.Tx, tenantID string, credential Credential, secret oauthSecret) (Credential, error) { - tenant, _ := parseID(tenantID) - metadata, ciphertext, err := s.sealOAuth(uuid.UUID(tenant.Bytes).String(), credential, secret) - if err != nil { - return Credential{}, err - } - vault, _ := parseID(credential.VaultID) - id, _ := parseID(credential.ID) - row, err := sqlc.New(tx).UpdateOAuthCredential(ctx, sqlc.UpdateOAuthCredentialParams{ - TenantID: tenant, VaultID: vault, ID: id, McpServerUrl: credential.MCPServerURL, - OauthMetadata: metadata, TokenCiphertext: ciphertext, - }) - if errors.Is(err, pgx.ErrNoRows) { - return Credential{}, ErrNotFound - } - if err != nil { - return Credential{}, errors.New("credential update failed") - } - return credentialFromRow(sqlc.GetCredentialRow(row)) -} diff --git a/services/core/internal/store/vault_credentials_oauth_test.go b/services/core/internal/store/vault_credentials_oauth_test.go deleted file mode 100644 index f41c9744e..000000000 --- a/services/core/internal/store/vault_credentials_oauth_test.go +++ /dev/null @@ -1,280 +0,0 @@ -package store - -import ( - "bytes" - "context" - "encoding/json" - "errors" - "reflect" - "strings" - "sync/atomic" - "testing" - "time" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgxpool" -) - -type oauthRefreshFunc func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) - -func (f oauthRefreshFunc) Refresh(ctx context.Context, request oauthrefresh.Request) (oauthrefresh.Token, error) { - return f(ctx, request) -} -func oauthString(value string) *string { return &value } - -func oauthFixture(t *testing.T, refresher oauthrefresh.Refresher) (*Store, *pgxpool.Pool, string, Vault, CreateOAuthCredentialInput) { - t.Helper() - _, pool := testStore(t) - cipher, err := credentialcrypto.New(bytes.Repeat([]byte{17}, 32)) - if err != nil { - t.Fatal(err) - } - s := NewWithCredentialCipherAndOAuthRefresh(pool, cipher, refresher) - tenant := uuid.NewString() - vault, err := s.CreateVault(t.Context(), tenant, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - input := CreateOAuthCredentialInput{Name: "OAuth fixture", MCPServerURL: "https://mcp.example/tools", - AccessToken: "private-access-canary", RefreshToken: "private-refresh-canary", ClientSecret: "private-client-canary", - OAuth: OAuthMetadata{ExpiresAt: oauthString(time.Now().Add(-time.Hour).UTC().Format(time.RFC3339Nano)), - Refresh: &OAuthRefreshMetadata{ClientID: "test-client", TokenEndpoint: "https://issuer.example/token", - TokenEndpointAuth: "client_secret_basic", Resource: oauthString("https://mcp.example/tools"), Scope: oauthString("read write")}}} - return s, pool, tenant, vault, input -} - -func createOAuthFixture(t *testing.T, s *Store, tenant string, vault Vault, input CreateOAuthCredentialInput) Credential { - t.Helper() - credential, err := s.CreateOAuthCredential(t.Context(), tenant, vault.ID, input) - if err != nil { - t.Fatal("create OAuth fixture", err) - } - return credential -} - -func oauthFixtureBinding(credential Credential) MCPCredentialBinding { - return MCPCredentialBinding{ServerLabel: "test", ServerURL: credential.MCPServerURL, - VaultID: credential.VaultID, CredentialID: credential.ID, AuthType: credential.AuthType} -} - -func storedOAuthSecret(t *testing.T, s *Store, tenant string, credential Credential) oauthSecret { - t.Helper() - var secret oauthSecret - err := s.withOAuth(t.Context(), tenant, credential.VaultID, credential.ID, "", "credential read failed", func(_ context.Context, _ pgx.Tx, _ Credential, stored oauthSecret) error { - secret = stored - return nil - }) - if err != nil { - t.Fatal("read private test grant", err) - } - return secret -} - -func TestOAuthCredentialMetadataEncryptionAndScope(t *testing.T) { - var exchanges atomic.Int32 - s, pool, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { - exchanges.Add(1) - return oauthrefresh.Token{}, errors.New("unexpected exchange") - })) - credential := createOAuthFixture(t, s, tenant, vault, input) - for _, reader := range []*Store{New(pool), s} { - got, err := reader.GetCredential(t.Context(), tenant, vault.ID, credential.ID) - if err != nil || !reflect.DeepEqual(got, credential) { - t.Fatal("keyless safe metadata changed", err) - } - page, err := reader.ListCredentials(t.Context(), tenant, vault.ID, "", 20, true, nil) - if err != nil || len(page.Credentials) != 1 || !reflect.DeepEqual(page.Credentials[0], credential) { - t.Fatal("keyless listing failed", err) - } - encoded, _ := json.Marshal(page) - for _, secret := range []string{input.AccessToken, input.RefreshToken, input.ClientSecret} { - if bytes.Contains(encoded, []byte(secret)) { - t.Fatal("metadata disclosed a secret") - } - } - } - var ciphertext []byte - if err := pool.QueryRow(t.Context(), "SELECT token_ciphertext FROM vault_credentials WHERE id=$1", credential.ID).Scan(&ciphertext); err != nil { - t.Fatal(err) - } - for _, secret := range []string{input.AccessToken, input.RefreshToken, input.ClientSecret} { - if bytes.Contains(ciphertext, []byte(secret)) { - t.Fatal("plaintext grant persisted") - } - } - restored := NewWithCredentialCipherAndOAuthRefresh(pool, s.credentialCipher, s.oauthRefresher) - if secret := storedOAuthSecret(t, restored, tenant, credential); secret.AccessToken != input.AccessToken || secret.RefreshToken != input.RefreshToken || secret.ClientSecret != input.ClientSecret { - t.Fatal("restart lost grant material") - } - binding := oauthFixtureBinding(credential) - for _, target := range []struct { - tenant string - vaults []string - binding MCPCredentialBinding - }{ - {uuid.NewString(), []string{vault.ID}, binding}, {tenant, nil, binding}, - {tenant, []string{vault.ID}, MCPCredentialBinding{ServerLabel: "test", ServerURL: input.MCPServerURL + "/other", VaultID: vault.ID, CredentialID: credential.ID, AuthType: "mcp_oauth"}}, - } { - if token, err := s.MCPBearerToken(t.Context(), target.tenant, target.vaults, target.binding); !errors.Is(err, ErrNotFound) || token != "" { - t.Fatal("foreign or mismatched scope admitted", err) - } - } - if _, err := s.UpdateOAuthCredential(t.Context(), uuid.NewString(), vault.ID, credential.ID, UpdateOAuthCredentialInput{AccessToken: oauthString("replacement")}); !errors.Is(err, ErrNotFound) { - t.Fatal("foreign update admitted") - } - if _, err := s.CreateOAuthCredential(t.Context(), uuid.NewString(), vault.ID, input); !errors.Is(err, ErrNotFound) { - t.Fatal("foreign creation admitted") - } - if _, err := New(pool).CreateOAuthCredential(t.Context(), tenant, vault.ID, input); !errors.Is(err, credentialcrypto.ErrUnavailable) { - t.Fatal("keyless creation admitted") - } - if token, err := New(pool).MCPBearerToken(t.Context(), tenant, []string{vault.ID}, binding); !errors.Is(err, credentialcrypto.ErrUnavailable) || token != "" { - t.Fatal("keyless execution admitted") - } - wrongCipher, _ := credentialcrypto.New(bytes.Repeat([]byte{18}, 32)) - wrong := NewWithCredentialCipherAndOAuthRefresh(pool, wrongCipher, s.oauthRefresher) - if token, err := wrong.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, binding); err == nil || token != "" { - t.Fatal("wrong key executed") - } - if exchanges.Load() != 0 { - t.Fatal("resource operations or invalid scopes contacted provider") - } -} - -func TestOAuthCredentialUpdatesPreservePinnedSemantics(t *testing.T) { - s, _, tenant, vault, input := oauthFixture(t, nil) - credential := createOAuthFixture(t, s, tenant, vault, input) - update := func(patch UpdateOAuthCredentialInput) Credential { - t.Helper() - got, err := s.UpdateOAuthCredential(t.Context(), tenant, vault.ID, credential.ID, patch) - if err != nil { - t.Fatal(err) - } - if got.ID != credential.ID || got.Name != credential.Name || got.AuthType != credential.AuthType || got.MCPServerURL != credential.MCPServerURL || !got.CreatedAt.Equal(credential.CreatedAt) { - t.Fatal("immutable metadata changed") - } - return got - } - unchanged := update(UpdateOAuthCredentialInput{}) - if !reflect.DeepEqual(unchanged.OAuth, credential.OAuth) { - t.Fatal("omitted values changed") - } - replacement := update(UpdateOAuthCredentialInput{AccessToken: oauthString("new-access")}) - if replacement.OAuth.ExpiresAt != nil { - t.Fatal("new access token retained old expiry") - } - expiry := time.Now().Add(time.Hour).UTC().Format(time.RFC3339Nano) - replacement = update(UpdateOAuthCredentialInput{ExpiresAtSet: true, ExpiresAt: &expiry, Refresh: &OAuthRefreshUpdate{ScopeSet: true, Scope: nil}}) - if replacement.OAuth.ExpiresAt == nil || *replacement.OAuth.ExpiresAt != expiry || replacement.OAuth.Refresh.Scope != nil { - t.Fatal("expiry or null scope semantics failed") - } - secret := storedOAuthSecret(t, s, tenant, replacement) - if secret.AccessToken != "new-access" || secret.RefreshToken != input.RefreshToken || secret.ClientSecret != input.ClientSecret { - t.Fatal("omitted secrets were replaced") - } - replacement = update(UpdateOAuthCredentialInput{ExpiresAtSet: true, Refresh: &OAuthRefreshUpdate{RefreshToken: oauthString("new-refresh"), TokenEndpointAuthType: "client_secret_basic", ClientSecret: oauthString("new-secret"), ScopeSet: true, Scope: oauthString("read")}}) - secret = storedOAuthSecret(t, s, tenant, replacement) - if secret.Metadata.ExpiresAt != nil || secret.RefreshToken != "new-refresh" || secret.ClientSecret != "new-secret" || *secret.Metadata.Refresh.Scope != "read" { - t.Fatal("replacement fields not persisted") - } - for _, patch := range []UpdateOAuthCredentialInput{ - {ExpiresAtSet: true, ExpiresAt: oauthString("not-a-date")}, - {Refresh: &OAuthRefreshUpdate{TokenEndpointAuthType: "client_secret_post"}}, - } { - if _, err := s.UpdateOAuthCredential(t.Context(), tenant, vault.ID, credential.ID, patch); !errors.Is(err, ErrInvalidInput) { - t.Fatal("invalid mutation admitted") - } - } - if got := storedOAuthSecret(t, s, tenant, replacement); !reflect.DeepEqual(got, secret) { - t.Fatal("rejected update changed the grant") - } - input.OAuth.Refresh = nil - input.RefreshToken = "" - input.ClientSecret = "" - noRefresh := createOAuthFixture(t, s, tenant, vault, input) - if _, err := s.UpdateOAuthCredential(t.Context(), tenant, vault.ID, noRefresh.ID, UpdateOAuthCredentialInput{Refresh: &OAuthRefreshUpdate{RefreshToken: oauthString("cannot-add")}}); !errors.Is(err, ErrInvalidInput) { - t.Fatal("added missing refresh configuration") - } -} - -func TestOAuthMetadataTamperingNeverReachesProvider(t *testing.T) { - var exchanges atomic.Int32 - s, pool, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { - exchanges.Add(1) - return oauthrefresh.Token{}, errors.New("unexpected refresh") - })) - for _, mutation := range []string{ - `jsonb_set(oauth_metadata,'{refresh,token_endpoint}','"https://attacker.example/token"')`, - `jsonb_set(oauth_metadata,'{refresh,client_id}','"other-client"')`, - `jsonb_set(oauth_metadata,'{refresh,scope}','"all"')`, - `jsonb_set(oauth_metadata,'{expires_at}','null')`, - } { - credential := createOAuthFixture(t, s, tenant, vault, input) - if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET oauth_metadata="+mutation+" WHERE id=$1", credential.ID); err != nil { - t.Fatal(err) - } - if token, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)); err == nil || token != "" { - t.Fatal("metadata substitution executed") - } - if _, err := s.UpdateOAuthCredential(t.Context(), tenant, vault.ID, credential.ID, UpdateOAuthCredentialInput{AccessToken: oauthString("new")}); err == nil { - t.Fatal("update authenticated substituted metadata") - } - } - if exchanges.Load() != 0 { - t.Fatal("tampered metadata reached token endpoint") - } -} - -func TestOAuthCredentialSelectionIncludesBothAuthTypes(t *testing.T) { - s, _, tenant, vault, input := oauthFixture(t, nil) - credential := createOAuthFixture(t, s, tenant, vault, input) - requests := []MCPCredentialRequest{{ServerLabel: "test", ServerURL: input.MCPServerURL}} - bindings, err := s.ResolveMCPCredentials(t.Context(), tenant, []string{vault.ID}, requests) - if err != nil || len(bindings) != 1 || bindings[0].AuthType != "mcp_oauth" || bindings[0].CredentialID != credential.ID { - t.Fatal("OAuth was not selected", err) - } - if _, err := s.CreateStaticCredential(t.Context(), tenant, vault.ID, CreateStaticCredentialInput{Name: "Static", MCPServerURL: input.MCPServerURL, Token: "static"}); err != nil { - t.Fatal(err) - } - if _, err := s.ResolveMCPCredentials(t.Context(), tenant, []string{vault.ID}, requests); !isSelectionError(err, true, "multiple attached vault credentials match MCP server_url "+input.MCPServerURL+"; specify credential_id") { - t.Fatal("ambiguous mixed credentials selected", err) - } - requests[0].CredentialID = &credential.ID - if selected, err := s.ResolveMCPCredentials(t.Context(), tenant, []string{vault.ID}, requests); err != nil || selected[0] != bindings[0] { - t.Fatal("explicit OAuth identity changed") - } -} - -func TestOAuthRefreshErrorsAreSafeAndPreserveGrant(t *testing.T) { - for _, mode := range []string{"provider_error", "empty_access", "expired_response", "no_refresh"} { - t.Run(mode, func(t *testing.T) { - s, _, tenant, vault, input := oauthFixture(t, oauthRefreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { - switch mode { - case "provider_error": - return oauthrefresh.Token{}, errors.New("private-refresh-canary: provider body") - case "expired_response": - past := time.Now().Add(-time.Second) - return oauthrefresh.Token{AccessToken: "new", ExpiresAt: &past}, nil - } - return oauthrefresh.Token{}, nil - })) - if mode == "no_refresh" { - input.OAuth.Refresh = nil - input.RefreshToken = "" - input.ClientSecret = "" - } - credential := createOAuthFixture(t, s, tenant, vault, input) - before := storedOAuthSecret(t, s, tenant, credential) - token, err := s.MCPBearerToken(t.Context(), tenant, []string{vault.ID}, oauthFixtureBinding(credential)) - if err == nil || token != "" || strings.Contains(err.Error(), "private-refresh-canary") { - t.Fatal("refresh failed unsafely") - } - if after := storedOAuthSecret(t, s, tenant, credential); !reflect.DeepEqual(before, after) { - t.Fatal("failed refresh changed grant") - } - }) - } -} diff --git a/services/core/internal/store/vault_credentials_test.go b/services/core/internal/store/vault_credentials_test.go deleted file mode 100644 index 3faafd097..000000000 --- a/services/core/internal/store/vault_credentials_test.go +++ /dev/null @@ -1,145 +0,0 @@ -package store - -import ( - "bytes" - "crypto/rand" - "encoding/hex" - "encoding/json" - "errors" - "reflect" - "strings" - "testing" - "time" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/google/uuid" -) - -func TestStaticCredentialsPersistEncryptedAndRemainScoped(t *testing.T) { - withoutKey, pool := testStore(t) - ctx := t.Context() - tenant, foreignTenant := uuid.NewString(), uuid.NewString() - var vaults []Vault - for _, owner := range []string{tenant, tenant, foreignTenant} { - vault, err := withoutKey.CreateVault(ctx, owner, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - vaults = append(vaults, vault) - } - key, randomToken := make([]byte, 32), make([]byte, 32) - if _, err := rand.Read(key); err != nil { - t.Fatal(err) - } - if _, err := rand.Read(randomToken); err != nil { - t.Fatal(err) - } - newCipher := func(key []byte) *credentialcrypto.Cipher { - c, err := credentialcrypto.New(key) - if err != nil { - t.Fatal(err) - } - return c - } - s := NewWithCredentialCipher(pool, newCipher(key)) - canary := hex.EncodeToString(randomToken) - opaque := " \t" + canary + " 凭据\n" + strings.Repeat("x", 300) + " " - tokens := []string{opaque, opaque, ""} - var records []Credential - before := time.Now().Add(-time.Second) - for _, token := range tokens { - input := CreateStaticCredentialInput{Name: "MCP credential", MCPServerURL: "https://mcp.example/tools", Token: token} - record, err := s.CreateStaticCredential(ctx, tenant, vaults[0].ID, input) - if err != nil { - t.Fatal(err) - } - if _, err := uuid.Parse(record.ID); err != nil || record.VaultID != vaults[0].ID || record.Name != input.Name || record.AuthType != "static_bearer" || record.MCPServerURL != input.MCPServerURL || record.CreatedAt.Before(before) || record.CreatedAt.After(time.Now().Add(time.Second)) || !record.CreatedAt.Equal(record.UpdatedAt) { - t.Fatal("credential metadata or database timestamps differ") - } - metadata, err := json.Marshal(record) - if err != nil || bytes.Contains(metadata, []byte(canary)) { - t.Fatal("credential metadata contains token plaintext") - } - records = append(records, record) - } - if records[0].ID == records[1].ID { - t.Fatal("separate creates reused a credential identity") - } - valid := CreateStaticCredentialInput{Name: "Rejected", MCPServerURL: "https://mcp.example/tools", Token: opaque} - if _, err := withoutKey.CreateStaticCredential(ctx, tenant, vaults[0].ID, valid); !errors.Is(err, credentialcrypto.ErrUnavailable) { - t.Fatal("missing encryption key did not fail writes closed") - } - for _, target := range []struct{ tenant, vault string }{{tenant, vaults[2].ID}, {foreignTenant, vaults[0].ID}, {tenant, uuid.NewString()}} { - if _, err := s.CreateStaticCredential(ctx, target.tenant, target.vault, valid); !errors.Is(err, ErrNotFound) { - t.Fatal("creation admitted an unowned or missing Vault") - } - } - for _, target := range []struct{ tenant, vault, credential string }{ - {tenant, vaults[1].ID, records[0].ID}, {foreignTenant, vaults[0].ID, records[0].ID}, - {tenant, vaults[2].ID, records[0].ID}, {tenant, uuid.NewString(), records[0].ID}, {tenant, vaults[0].ID, uuid.NewString()}, - } { - if _, err := s.GetCredential(ctx, target.tenant, target.vault, target.credential); !errors.Is(err, ErrNotFound) { - t.Fatal("unowned, wrong-Vault or missing credential was disclosed") - } - } - pool.Close() - withoutKey, pool = testStore(t) - restartedCipher := newCipher(bytes.Clone(key)) - wrongKey := bytes.Clone(key) - wrongKey[0] ^= 1 - wrongCipher := newCipher(wrongKey) - s = NewWithCredentialCipher(pool, restartedCipher) - var ciphertexts [][]byte - for i, record := range records { - for _, reader := range []*Store{withoutKey, s, NewWithCredentialCipher(pool, wrongCipher)} { - got, err := reader.GetCredential(ctx, tenant, record.VaultID, record.ID) - if err != nil || !reflect.DeepEqual(got, record) { - t.Fatal("safe metadata recovery depended on the encryption key", err) - } - } - var ciphertext []byte - var storageType string - if err := pool.QueryRow(ctx, "SELECT token_ciphertext, pg_typeof(token_ciphertext)::text FROM vault_credentials WHERE id=$1 AND vault_id=$2", record.ID, record.VaultID).Scan(&ciphertext, &storageType); err != nil || storageType != "bytea" || len(ciphertext) < 29 || bytes.Contains(ciphertext, []byte(canary)) { - t.Fatal("credential ciphertext was not stored as private bytea", err) - } - binding := credentialcrypto.Binding{TenantID: tenant, VaultID: record.VaultID, CredentialID: record.ID, AuthType: record.AuthType, Destination: record.MCPServerURL} - plaintext, err := restartedCipher.Open(ciphertext, binding) - if err != nil || !bytes.Equal(plaintext, []byte(tokens[i])) { - t.Fatal("private restart decryption did not preserve token bytes", err) - } - if plaintext, err := wrongCipher.Open(ciphertext, binding); err == nil || plaintext != nil { - t.Fatal("wrong key decrypted persisted ciphertext") - } - ciphertexts = append(ciphertexts, ciphertext) - } - // Version 1 prefixes the standard library's 12-byte random nonce. - if bytes.Equal(ciphertexts[0], ciphertexts[1]) || bytes.Equal(ciphertexts[0][1:13], ciphertexts[1][1:13]) { - t.Fatal("same token in distinct records reused ciphertext or nonce") - } - // Change one persisted binding field. Metadata GET must still work without a - // key, while private decryption must reject the altered stored destination. - if _, err := pool.Exec(ctx, "UPDATE vault_credentials SET mcp_server_url=$1 WHERE id=$2", "https://other.example/tools", records[0].ID); err != nil { - t.Fatal(err) - } - changed, err := withoutKey.GetCredential(ctx, tenant, vaults[0].ID, records[0].ID) - if err != nil || changed.MCPServerURL != "https://other.example/tools" { - t.Fatal("metadata GET unexpectedly required decryption", err) - } - binding := credentialcrypto.Binding{TenantID: tenant, VaultID: changed.VaultID, CredentialID: changed.ID, AuthType: changed.AuthType, Destination: changed.MCPServerURL} - if plaintext, err := restartedCipher.Open(ciphertexts[0], binding); err == nil || plaintext != nil { - t.Fatal("persisted destination substitution authenticated") - } - if _, err := pool.Exec(ctx, "UPDATE vault_credentials SET mcp_server_url=$1 WHERE id=$2", records[0].MCPServerURL, records[0].ID); err != nil { - t.Fatal(err) - } - for _, vault := range vaults { - got, err := withoutKey.GetVault(ctx, vault.TenantID, vault.ID) - if err != nil || !reflect.DeepEqual(got, vault) { - t.Fatal("credential operations changed an owning Vault", err) - } - } - var count, sessions int - if err := pool.QueryRow(ctx, "SELECT (SELECT count(*) FROM vault_credentials WHERE vault_id=ANY($1::uuid[])), (SELECT count(*) FROM sessions WHERE tenant_id=ANY($2::uuid[]))", []string{vaults[0].ID, vaults[1].ID, vaults[2].ID}, []string{tenant, foreignTenant}).Scan(&count, &sessions); err != nil || count != len(records) || sessions != 0 { - t.Fatal("rejected requests wrote rows or credential operations created Sessions", err) - } -} diff --git a/services/core/internal/store/vault_credentials_update.go b/services/core/internal/store/vault_credentials_update.go deleted file mode 100644 index e96929470..000000000 --- a/services/core/internal/store/vault_credentials_update.go +++ /dev/null @@ -1,74 +0,0 @@ -package store - -import ( - "context" - "errors" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" -) - -type UpdateStaticCredentialInput struct { - Token string -} - -// UpdateStaticCredential replaces only the secret and update time. Safe metadata -// supplies immutable AAD; the mutation independently checks that same scope. -// A subsequent dispatch reads the replacement through the existing frozen binding. -func (s *Store) UpdateStaticCredential(ctx context.Context, tenantID, vaultID, credentialID string, input UpdateStaticCredentialInput) (Credential, error) { - vaultID, credentialID = pathID(vaultID), pathID(credentialID) - tenant, err := parseID(tenantID) - if err != nil { - return Credential{}, ErrNotFound - } - vault, err := parseID(vaultID) - if err != nil { - return Credential{}, ErrNotFound - } - id, err := parseID(credentialID) - if err != nil { - return Credential{}, ErrNotFound - } - current, err := s.GetCredential(ctx, tenantID, vaultID, credentialID) - if err != nil { - return Credential{}, err - } - if current.AuthType != "static_bearer" { - return Credential{}, ErrInvalidInput - } - if s.credentialCipher == nil { - return Credential{}, credentialcrypto.ErrUnavailable - } - ciphertext, err := s.credentialCipher.Seal([]byte(input.Token), credentialcrypto.Binding{ - TenantID: uuid.UUID(tenant.Bytes).String(), VaultID: current.VaultID, - CredentialID: current.ID, AuthType: current.AuthType, Destination: current.MCPServerURL, - }) - if err != nil { - return Credential{}, errors.New("credential encryption failed") - } - var updated Credential - err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { - q := s.queries.WithTx(tx) - row, err := q.UpdateStaticCredential(ctx, sqlc.UpdateStaticCredentialParams{ - TenantID: tenant, VaultID: vault, ID: id, McpServerUrl: current.MCPServerURL, TokenCiphertext: ciphertext, - }) - if err != nil { - return err - } - updated, err = credentialFromRow(sqlc.GetCredentialRow(row)) - if err != nil { - return err - } - return auditpg.RecordWriteAudit(ctx, q, tenantID, "update", "credential", updated.ID, updated.VaultID) - }) - if errors.Is(err, pgx.ErrNoRows) { - return Credential{}, ErrNotFound - } - if err != nil { - return Credential{}, errors.New("credential update failed") - } - return updated, nil -} diff --git a/services/core/internal/store/vault_credentials_update_test.go b/services/core/internal/store/vault_credentials_update_test.go deleted file mode 100644 index efef86bc0..000000000 --- a/services/core/internal/store/vault_credentials_update_test.go +++ /dev/null @@ -1,184 +0,0 @@ -package store - -import ( - "bytes" - "crypto/rand" - "encoding/json" - "errors" - "reflect" - "strings" - "testing" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" -) - -func TestStaticCredentialUpdatePreservesBindingsAndReplacesCurrentSecret(t *testing.T) { - public, pool := testStore(t) - tenant, foreign := uuid.NewString(), uuid.NewString() - key := make([]byte, 32) - if _, err := rand.Read(key); err != nil { - t.Fatal(err) - } - cipher, err := credentialcrypto.New(key) - if err != nil { - t.Fatal(err) - } - s := NewWithCredentialCipher(pool, cipher) - var vaults []Vault - for _, owner := range []string{tenant, tenant, foreign} { - vault, err := s.CreateVault(t.Context(), owner, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - vaults = append(vaults, vault) - } - firstToken, endpoint := uuid.NewString(), "https://mcp.example/tools" - original, err := s.CreateStaticCredential(t.Context(), tenant, vaults[0].ID, CreateStaticCredentialInput{Name: "Retained name", MCPServerURL: endpoint, Token: firstToken}) - if err != nil { - t.Fatal(err) - } - unrelated, err := s.CreateStaticCredential(t.Context(), tenant, vaults[1].ID, CreateStaticCredentialInput{Name: "Unrelated", MCPServerURL: endpoint, Token: firstToken}) - if err != nil { - t.Fatal(err) - } - attached := []string{vaults[0].ID} - var sessions []Session - var bindings []MCPCredentialBinding - for index, id := range []*string{nil, &original.ID} { - selected, err := public.ResolveMCPCredentials(t.Context(), tenant, attached, []MCPCredentialRequest{{ServerLabel: "tools", ServerURL: endpoint, CredentialID: id}}) - if err != nil || len(selected) != 1 { - t.Fatal("binding setup failed", err) - } - configuration, _ := json.Marshal(map[string]any{ - "agent": map[string]any{"model": "model", "tools": []any{map[string]any{"type": "mcp", "server_label": "tools", "transport": map[string]string{"type": "http", "server_url": endpoint}, "connection_origin": "service", "credential_id": id}}}, - "environment": map[string]string{"type": "none"}, "vault_ids": attached, "mcp_credentials": selected, - }) - session, err := s.CreateSession(t.Context(), tenant, CreateSessionInput{Creator: FixtureCreator(), Engine: "codex", IdempotencyKey: []string{"implicit", "explicit"}[index], Configuration: configuration}) - if err != nil { - t.Fatal(err) - } - sessions, bindings = append(sessions, session), append(bindings, selected[0]) - } - readCiphertext := func(id string) []byte { - t.Helper() - var ciphertext []byte - if err := pool.QueryRow(t.Context(), "SELECT token_ciphertext FROM vault_credentials WHERE id=$1", id).Scan(&ciphertext); err != nil { - t.Fatal("private ciphertext observation failed") - } - return ciphertext - } - prior, unrelatedCiphertext := readCiphertext(original.ID), readCiphertext(unrelated.ID) - lastToken := " \t" + uuid.NewString() + " 雪\n" - for _, token := range []string{"", lastToken, lastToken} { - updated, err := s.UpdateStaticCredential(t.Context(), tenant, original.VaultID, original.ID, UpdateStaticCredentialInput{Token: token}) - if err != nil { - t.Fatal("token replacement failed") - } - want := original - want.UpdatedAt = updated.UpdatedAt - if !reflect.DeepEqual(updated, want) || updated.UpdatedAt.Before(original.UpdatedAt) { - t.Fatal("replacement changed immutable metadata") - } - current := readCiphertext(original.ID) - if bytes.Equal(prior, current) || len(token) > 0 && bytes.Contains(current, []byte(token)) { - t.Fatal("replacement reused ciphertext or stored plaintext") - } - for _, binding := range bindings { - got, err := s.MCPBearerToken(t.Context(), tenant, attached, binding) - if err != nil || got != token { - t.Fatal("existing selection did not read the exact committed replacement") - } - } - prior = current - } - before, err := public.GetCredential(t.Context(), tenant, original.VaultID, original.ID) - if err != nil { - t.Fatal(err) - } - assertUnchanged := func() { - t.Helper() - after, err := public.GetCredential(t.Context(), tenant, original.VaultID, original.ID) - if err != nil || !reflect.DeepEqual(after, before) || !bytes.Equal(readCiphertext(original.ID), prior) { - t.Fatal("failed replacement changed the existing row") - } - } - for _, scope := range []struct{ tenant, vault, id string }{ - {foreign, original.VaultID, original.ID}, {tenant, vaults[1].ID, original.ID}, - {tenant, vaults[2].ID, original.ID}, {tenant, original.VaultID, uuid.NewString()}, {tenant, "invalid", original.ID}, - } { - if _, err := s.UpdateStaticCredential(t.Context(), scope.tenant, scope.vault, scope.id, UpdateStaticCredentialInput{Token: "rejected"}); !errors.Is(err, ErrNotFound) { - t.Fatal("unowned or invalid replacement was admitted") - } - assertUnchanged() - } - for _, writer := range []*Store{public, NewWithCredentialCipher(pool, &credentialcrypto.Cipher{})} { - if _, err := writer.UpdateStaticCredential(t.Context(), tenant, original.VaultID, original.ID, UpdateStaticCredentialInput{Token: "rejected"}); err == nil { - t.Fatal("missing or unusable cipher admitted replacement") - } - assertUnchanged() - } - // A real PostgreSQL mutation failure must preserve both ciphertext and time. - readOnly := readOnlyResourceStore(t, pool) - readOnly.credentialCipher = cipher - _, updateErr := readOnly.UpdateStaticCredential(t.Context(), tenant, original.VaultID, original.ID, UpdateStaticCredentialInput{Token: "rejected"}) - if updateErr == nil || updateErr.Error() != "credential update failed" { - t.Fatal("database write failure was accepted or exposed") - } - assertUnchanged() - // A stale destination from a prior metadata read cannot authorize the UPDATE. - tenantID, _ := parseID(tenant) - vaultID, _ := parseID(original.VaultID) - credentialID, _ := parseID(original.ID) - _, err = s.queries.UpdateStaticCredential(t.Context(), sqlc.UpdateStaticCredentialParams{TenantID: tenantID, VaultID: vaultID, ID: credentialID, McpServerUrl: endpoint + "/other", TokenCiphertext: prior}) - if !errors.Is(err, pgx.ErrNoRows) { - t.Fatal("mutation failed to recheck immutable destination") - } - assertUnchanged() - // Replacing a damaged old payload needs no old-token decryption. - if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET token_ciphertext=set_byte(token_ciphertext, 15, get_byte(token_ciphertext,15) # 1) WHERE id=$1", original.ID); err != nil { - t.Fatal(err) - } - if _, err := s.UpdateStaticCredential(t.Context(), tenant, original.VaultID, original.ID, UpdateStaticCredentialInput{Token: lastToken}); err != nil { - t.Fatal("replacement tried to decrypt the old token") - } - pool.Close() - public, pool = testStore(t) - cipher, _ = credentialcrypto.New(bytes.Clone(key)) - s = NewWithCredentialCipher(pool, cipher) - for index, session := range sessions { - got, err := s.GetSession(t.Context(), tenant, session.ID) - if err != nil || !bytes.Equal(got.Configuration, session.Configuration) || got.LastTurn != nil { - t.Fatal("replacement changed an existing Session") - } - current, err := s.MCPBearerToken(t.Context(), tenant, attached, bindings[index]) - if err != nil || current != lastToken { - t.Fatal("reopened dispatch lookup lost the replacement") - } - } - // Competing whole-secret replacements may win in either order, never tear. - left, right := uuid.NewString()+strings.Repeat("L", 513), uuid.NewString()+strings.Repeat("R", 1025) - start, results := make(chan struct{}), make(chan error, 2) - for _, token := range []string{left, right} { - go func() { - <-start - _, err := s.UpdateStaticCredential(t.Context(), tenant, original.VaultID, original.ID, UpdateStaticCredentialInput{Token: token}) - results <- err - }() - } - close(start) - for range 2 { - if err := <-results; err != nil { - t.Fatal("concurrent replacement failed") - } - } - current, err := s.MCPBearerToken(t.Context(), tenant, attached, bindings[0]) - if err != nil || current != left && current != right { - t.Fatal("concurrent replacements produced an incomplete secret") - } - if !bytes.Equal(readCiphertext(unrelated.ID), unrelatedCiphertext) { - t.Fatal("replacement changed an unrelated Credential") - } -} diff --git a/services/core/internal/store/vaults.go b/services/core/internal/store/vaults.go deleted file mode 100644 index 4f66c594c..000000000 --- a/services/core/internal/store/vaults.go +++ /dev/null @@ -1,111 +0,0 @@ -package store - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "time" - "unicode/utf8" - - "github.com/google/uuid" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgtype" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/metadata" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/writeaudit" -) - -// Vault is a tenant-owned resource, independent of Sessions and engine execution. -type Vault struct { - ID string - TenantID string - Name *string - Metadata map[string]string - CreatedAt time.Time -} - -type CreateVaultInput struct { - Name *string - Metadata map[string]string -} - -// CreateVault persists the public layer's normalized name. Each call creates a -// distinct resource; this primitive does not define create retry semantics. -func (s *Store) CreateVault(ctx context.Context, tenantID string, input CreateVaultInput) (Vault, error) { - tenant, err := parseID(tenantID) - if err != nil { - return Vault{}, err - } - var name pgtype.Text - if input.Name != nil { - if !validVaultName(*input.Name) { - return Vault{}, fmt.Errorf("%w: vault name must contain 1–256 UTF-8 bytes", ErrInvalidInput) - } - name = pgtype.Text{String: *input.Name, Valid: true} - } - encodedMetadata, err := metadata.Encode(input.Metadata) - if err != nil { - return Vault{}, fmt.Errorf("%w: %w", ErrInvalidInput, err) - } - var created Vault - err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { - q := s.queries.WithTx(tx) - row, err := q.CreateVault(ctx, sqlc.CreateVaultParams{ - ID: pgtype.UUID{Bytes: uuid.New(), Valid: true}, TenantID: tenant, - Name: name, Metadata: encodedMetadata, - }) - if err != nil { - return err - } - created, err = vaultFromRow(row) - if err != nil { - return err - } - return auditpg.RecordWriteAudit(ctx, q, tenantID, "create", "vault", created.ID, "", writeaudit.Resource{Type: "vault", ID: created.ID, ParentID: ""}) - }) - if err != nil { - return Vault{}, fmt.Errorf("create vault: %w", err) - } - return created, nil -} - -func validVaultName(name string) bool { - return len(name) >= 1 && len(name) <= 256 && utf8.ValidString(name) -} - -// GetVault scopes every lookup to the authenticated caller's tenant. -func (s *Store) GetVault(ctx context.Context, tenantID, vaultID string) (Vault, error) { - tenant, err := parseID(tenantID) - if err != nil { - return Vault{}, err - } - id, err := parseID(vaultID) - if err != nil { - return Vault{}, err - } - row, err := s.queries.GetVault(ctx, sqlc.GetVaultParams{TenantID: tenant, ID: id}) - if errors.Is(err, pgx.ErrNoRows) { - return Vault{}, ErrNotFound - } - if err != nil { - return Vault{}, fmt.Errorf("get vault: %w", err) - } - return vaultFromRow(row) -} - -func vaultFromRow(row sqlc.Vault) (Vault, error) { - vault := Vault{ - ID: uuid.UUID(row.ID.Bytes).String(), TenantID: uuid.UUID(row.TenantID.Bytes).String(), - CreatedAt: row.CreatedAt.Time, - } - if row.Name.Valid { - vault.Name = &row.Name.String - } - if err := json.Unmarshal(row.Metadata, &vault.Metadata); err != nil { - return Vault{}, fmt.Errorf("decode vault metadata: %w", err) - } - return vault, nil -} diff --git a/services/core/internal/store/vaults_delete.go b/services/core/internal/store/vaults_delete.go deleted file mode 100644 index 811240a21..000000000 --- a/services/core/internal/store/vaults_delete.go +++ /dev/null @@ -1,40 +0,0 @@ -package store - -import ( - "context" - "errors" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" -) - -// DeleteVault relies on the owning foreign key to remove every stored Credential. -func (s *Store) DeleteVault(ctx context.Context, tenantID, vaultID string) (string, error) { - tenant, err := parseID(tenantID) - if err != nil { - return "", ErrNotFound - } - id, err := parseID(vaultID) - if err != nil { - return "", ErrNotFound - } - var deletedID string - err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { - q := s.queries.WithTx(tx) - deleted, err := q.DeleteVault(ctx, sqlc.DeleteVaultParams{TenantID: tenant, ID: id}) - if err != nil { - return err - } - deletedID = uuid.UUID(deleted.Bytes).String() - return auditpg.RecordWriteAudit(ctx, q, tenantID, "delete", "vault", deletedID, "") - }) - if errors.Is(err, pgx.ErrNoRows) { - return "", ErrNotFound - } - if err != nil { - return "", errors.New("vault deletion failed") - } - return deletedID, nil -} diff --git a/services/core/internal/store/vaults_delete_test.go b/services/core/internal/store/vaults_delete_test.go deleted file mode 100644 index 3a02e5abe..000000000 --- a/services/core/internal/store/vaults_delete_test.go +++ /dev/null @@ -1,201 +0,0 @@ -package store - -import ( - "bytes" - "encoding/json" - "errors" - "reflect" - "testing" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/google/uuid" - "github.com/jackc/pgx/v5/pgconn" -) - -func TestVaultDeletionCascadeBindingAndRestart(t *testing.T) { - public, pool := testStore(t) - tenant, foreign := uuid.NewString(), uuid.NewString() - key := bytes.Repeat([]byte{43}, 32) - cipher, err := credentialcrypto.New(key) - if err != nil { - t.Fatal(err) - } - s := NewWithCredentialCipher(pool, cipher) - createVault := func() Vault { - t.Helper() - v, err := public.CreateVault(t.Context(), tenant, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - return v - } - vault, retained, empty := createVault(), createVault(), createVault() - create := func(id, name, url string) Credential { - t.Helper() - v, err := s.CreateStaticCredential(t.Context(), tenant, id, CreateStaticCredentialInput{Name: name, MCPServerURL: url, Token: name + "-secret"}) - if err != nil { - t.Fatal(err) - } - return v - } - original := create(vault.ID, "original", "https://mcp.example/tools") - attached := []string{vault.ID, retained.ID} - selected, err := public.ResolveMCPCredentials(t.Context(), tenant, attached, []MCPCredentialRequest{{ServerLabel: "tools", ServerURL: original.MCPServerURL}}) - if err != nil || len(selected) != 1 { - t.Fatal("initial unique selection failed", err) - } - configuration, _ := json.Marshal(map[string]any{"agent": map[string]string{"model": "model"}, "vault_ids": attached, "mcp_credentials": selected}) - input := CreateSessionInput{Creator: FixtureCreator(), Engine: "codex", IdempotencyKey: "retained", Configuration: configuration} - session, err := s.CreateSession(t.Context(), tenant, input) - if err != nil { - t.Fatal(err) - } - extra := create(vault.ID, "archived", "https://mcp.example/other") - sibling := create(retained.ID, "sibling", original.MCPServerURL) - if _, err := pool.Exec(t.Context(), "UPDATE vaults SET status='archived' WHERE id=$1", vault.ID); err != nil { - t.Fatal(err) - } - if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET status='archived' WHERE id=$1", extra.ID); err != nil { - t.Fatal(err) - } - for _, scope := range []struct{ tenant, id string }{{foreign, vault.ID}, {tenant, uuid.NewString()}, {tenant, "invalid"}, {"invalid", vault.ID}} { - if _, err := public.DeleteVault(t.Context(), scope.tenant, scope.id); !errors.Is(err, ErrNotFound) { - t.Fatal("foreign or invalid deletion was accepted", err) - } - } - readOnly := readOnlyResourceStore(t, pool) - _, deletionErr := readOnly.DeleteVault(t.Context(), tenant, vault.ID) - if deletionErr == nil || deletionErr.Error() != "vault deletion failed" { - t.Fatal("failed mutation was accepted or exposed") - } - tx, err := pool.Begin(t.Context()) - if err != nil { - t.Fatal(err) - } - defer func() { _ = tx.Rollback(t.Context()) }() - // Verify the database cascade independently inside an explicit transaction. - tenantID, _ := parseID(tenant) - vaultID, _ := parseID(vault.ID) - if _, err := public.queries.WithTx(tx).DeleteVault(t.Context(), sqlc.DeleteVaultParams{TenantID: tenantID, ID: vaultID}); err != nil { - t.Fatal(err) - } - var count int - if err := tx.QueryRow(t.Context(), "SELECT (SELECT count(*) FROM vaults WHERE id=$1)+(SELECT count(*) FROM vault_credentials WHERE vault_id=$1)", vault.ID).Scan(&count); err != nil || count != 0 { - t.Fatal("cascade was not visible in the deletion transaction", err) - } - if err := tx.Rollback(t.Context()); err != nil { - t.Fatal(err) - } - if value, err := public.GetVault(t.Context(), tenant, vault.ID); err != nil || !reflect.DeepEqual(value, vault) { - t.Fatal("rollback changed the Vault", err) - } - for _, expected := range []Credential{original, extra} { - if value, err := public.GetCredential(t.Context(), tenant, vault.ID, expected.ID); err != nil || !reflect.DeepEqual(value, expected) { - t.Fatal("rollback changed a child", err) - } - } - if token, err := s.MCPBearerToken(t.Context(), tenant, attached, selected[0]); err != nil || token != "original-secret" { - t.Fatal("rejected deletion changed the stored token") - } - if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET token_ciphertext=decode('00','hex') WHERE vault_id=$1", vault.ID); err != nil { - t.Fatal(err) - } - for _, target := range []Vault{empty, vault} { - if id, err := public.DeleteVault(t.Context(), tenant, target.ID); err != nil || id != target.ID { - t.Fatal("keyless deletion failed", err) - } - } - if err := pool.QueryRow(t.Context(), "SELECT (SELECT count(*) FROM vaults WHERE id=$1)+(SELECT count(*) FROM vault_credentials WHERE vault_id=$1)", vault.ID).Scan(&count); err != nil || count != 0 { - t.Fatal("committed parent or encrypted children remain", err) - } - pool.Close() - public, pool = testStore(t) - cipher, _ = credentialcrypto.New(bytes.Clone(key)) - s = NewWithCredentialCipher(pool, cipher) - for _, target := range []Vault{empty, vault} { - if _, err := public.GetVault(t.Context(), tenant, target.ID); !errors.Is(err, ErrNotFound) { - t.Fatal("deleted Vault reappeared after restart") - } - if _, err := public.DeleteVault(t.Context(), tenant, target.ID); !errors.Is(err, ErrNotFound) { - t.Fatal("repeated deletion did not remain absent") - } - } - for _, child := range []Credential{original, extra} { - if _, err := public.GetCredential(t.Context(), tenant, vault.ID, child.ID); !errors.Is(err, ErrNotFound) { - t.Fatal("deleted child reappeared") - } - if _, err := s.UpdateStaticCredential(t.Context(), tenant, vault.ID, child.ID, UpdateStaticCredentialInput{Token: "replacement"}); !errors.Is(err, ErrNotFound) { - t.Fatal("replacement recreated a deleted child") - } - } - if _, err := s.CreateStaticCredential(t.Context(), tenant, vault.ID, CreateStaticCredentialInput{Name: "late", MCPServerURL: original.MCPServerURL, Token: "late"}); !errors.Is(err, ErrNotFound) { - t.Fatal("new child was admitted under a deleted Vault") - } - if _, err := s.MCPBearerToken(t.Context(), tenant, attached, selected[0]); !errors.Is(err, ErrNotFound) { - t.Fatal("frozen binding reselected a credential in another attached Vault") - } - if _, err := public.ListCredentials(t.Context(), tenant, vault.ID, "", 100, true, []string{"active", "archived"}); !errors.Is(err, ErrNotFound) { - t.Fatal("deleted parent remained listable") - } - if value, err := public.GetVault(t.Context(), tenant, retained.ID); err != nil || !reflect.DeepEqual(value, retained) { - t.Fatal("deletion changed another Vault") - } - if value, err := public.GetCredential(t.Context(), tenant, retained.ID, sibling.ID); err != nil || !reflect.DeepEqual(value, sibling) { - t.Fatal("deletion changed another Vault's credential") - } - if retry, err := public.CreateSession(t.Context(), tenant, input); err != nil || retry.ID != session.ID || !bytes.Equal(retry.Configuration, session.Configuration) { - t.Fatal("deletion changed frozen creation identity", err) - } -} - -func TestVaultDeletionConcurrentChildMutations(t *testing.T) { - public, pool := testStore(t) - tenant := uuid.NewString() - cipher, err := credentialcrypto.New(bytes.Repeat([]byte{44}, 32)) - if err != nil { - t.Fatal(err) - } - s := NewWithCredentialCipher(pool, cipher) - for range 8 { - vault, err := s.CreateVault(t.Context(), tenant, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - input := CreateStaticCredentialInput{Name: "competing", MCPServerURL: "https://mcp.example/tools", Token: "before"} - value, err := s.CreateStaticCredential(t.Context(), tenant, vault.ID, input) - if err != nil { - t.Fatal(err) - } - start, created, updated, removed := make(chan struct{}), make(chan error, 1), make(chan error, 1), make(chan error, 1) - go func() { - <-start - _, err := s.CreateStaticCredential(t.Context(), tenant, vault.ID, input) - created <- err - }() - go func() { - <-start - _, err := s.UpdateStaticCredential(t.Context(), tenant, vault.ID, value.ID, UpdateStaticCredentialInput{Token: "after"}) - updated <- err - }() - go func() { - <-start - _, err := public.DeleteCredential(t.Context(), tenant, vault.ID, value.ID) - removed <- err - }() - close(start) - _, deleted := public.DeleteVault(t.Context(), tenant, vault.ID) - createErr, updateErr, removeErr := <-created, <-updated, <-removed - var constraint *pgconn.PgError - if createErr != nil && !errors.Is(createErr, ErrNotFound) && !(errors.As(createErr, &constraint) && constraint.Code == "23503") { - t.Fatal("competing creation failed unexpectedly", createErr) - } - if deleted != nil || updateErr != nil && !errors.Is(updateErr, ErrNotFound) || removeErr != nil && !errors.Is(removeErr, ErrNotFound) { - t.Fatal("competing mutation failed unexpectedly", deleted, updateErr, removeErr) - } - var count int - if err := pool.QueryRow(t.Context(), "SELECT (SELECT count(*) FROM vaults WHERE id=$1)+(SELECT count(*) FROM vault_credentials WHERE vault_id=$1)", vault.ID).Scan(&count); err != nil || count != 0 { - t.Fatal("concurrent mutation resurrected deleted resources", err) - } - } -} diff --git a/services/core/internal/store/vaults_fixture_test.go b/services/core/internal/store/vaults_fixture_test.go new file mode 100644 index 000000000..ef7422022 --- /dev/null +++ b/services/core/internal/store/vaults_fixture_test.go @@ -0,0 +1,20 @@ +package store_test + +import ( + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/vaultpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +// fixtureVaults builds the Vault adapter and service on db, as cmd/server does. +// A keyless db leaves the operations that need no credential key available. +func fixtureVaults(db fixtureDB) (*vaultpg.Store, *vaults.Service, error) { + refresher, err := oauthrefresh.NewClient(nil) + if err != nil { + return nil, nil, err + } + vaultStore := vaultpg.New(pgunit.NewPool(db.pool)) + vaultService, err := vaults.NewService(vaultStore, db.cipher, refresher) + return vaultStore, vaultService, err +} diff --git a/services/core/internal/store/vaults_list.go b/services/core/internal/store/vaults_list.go deleted file mode 100644 index 3015a8c60..000000000 --- a/services/core/internal/store/vaults_list.go +++ /dev/null @@ -1,60 +0,0 @@ -package store - -import ( - "context" - "fmt" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" - "github.com/google/uuid" - "github.com/jackc/pgx/v5/pgtype" -) - -type VaultPage struct { - Vaults []Vault - NextCursor string -} - -func (s *Store) ListVaults(ctx context.Context, tenantID, cursor string, limit int, ascending bool, statuses []string) (VaultPage, error) { - tenant, err := parseID(tenantID) - if err != nil { - return VaultPage{}, err - } - if limit < 1 || limit > 100 { - return VaultPage{}, fmt.Errorf("%w: internal page size must be 1..100", ErrInvalidInput) - } - if len(statuses) == 0 { - statuses = []string{"active", "archived"} - } - for _, status := range statuses { - if status != "active" && status != "archived" { - return VaultPage{}, fmt.Errorf("%w: invalid Vault status", ErrInvalidInput) - } - } - params := sqlc.ListVaultsParams{TenantID: tenant, PageLimit: int32(limit + 1), AfterID: pgtype.UUID{Valid: true}, Ascending: ascending, Statuses: statuses} - if cursor != "" { - after, err := s.GetVault(ctx, tenantID, pgunit.LookupCursor(cursor)) - if err != nil { - return VaultPage{}, err - } - params.AfterCreated = pgtype.Timestamptz{Time: after.CreatedAt, Valid: true} - params.AfterID, _ = parseID(after.ID) - } - rows, err := s.queries.ListVaults(ctx, params) - if err != nil { - return VaultPage{}, fmt.Errorf("list vaults: %w", err) - } - page := VaultPage{Vaults: make([]Vault, 0, min(limit, len(rows)))} - if len(rows) > limit { - page.NextCursor = uuid.UUID(rows[limit-1].ID.Bytes).String() - rows = rows[:limit] - } - for _, row := range rows { - vault, err := vaultFromRow(row) - if err != nil { - return VaultPage{}, err - } - page.Vaults = append(page.Vaults, vault) - } - return page, nil -} diff --git a/services/core/internal/store/vaults_list_test.go b/services/core/internal/store/vaults_list_test.go deleted file mode 100644 index cd574a4b0..000000000 --- a/services/core/internal/store/vaults_list_test.go +++ /dev/null @@ -1,122 +0,0 @@ -package store - -import ( - "cmp" - "errors" - "reflect" - "slices" - "testing" - "time" - - "github.com/google/uuid" -) - -func TestVaultListFilteringPaginationAndReconnect(t *testing.T) { - s, pool := testStore(t) - ctx := t.Context() - tenant, other := uuid.NewString(), uuid.NewString() - empty, err := s.ListVaults(ctx, tenant, "", 20, false, nil) - if err != nil || empty.Vaults == nil || len(empty.Vaults) != 0 || empty.NextCursor != "" { - t.Fatalf("empty page: %+v, %v", empty, err) - } - var all []Vault - archived := map[string]bool{} - for i := range 105 { - vault, err := s.CreateVault(ctx, tenant, CreateVaultInput{Metadata: map[string]string{"purpose": "safe list fixture"}}) - if err != nil { - t.Fatal(err) - } - status := "active" - if i%3 == 0 { - status, archived[vault.ID] = "archived", true - } - vault.CreatedAt = time.Unix(1700000000+int64(i%2), 0).UTC() - // Synthetic classifications exercise reads, not a public archive lifecycle. - if _, err := pool.Exec(ctx, "UPDATE vaults SET created_at=$1, status=$2 WHERE tenant_id=$3 AND id=$4", vault.CreatedAt, status, tenant, vault.ID); err != nil { - t.Fatal(err) - } - vault, err = s.GetVault(ctx, tenant, vault.ID) - if err != nil { - t.Fatal(err) - } - all = append(all, vault) - } - slices.SortFunc(all, func(a, b Vault) int { - if c := a.CreatedAt.Compare(b.CreatedAt); c != 0 { - return c - } - return cmp.Compare(a.ID, b.ID) - }) - foreign, err := s.CreateVault(ctx, other, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - read := func(s *Store, ascending bool, statuses []string, size int) []Vault { - t.Helper() - actual := []Vault{} - cursor := "" - for { - page, err := s.ListVaults(ctx, tenant, cursor, size, ascending, statuses) - if err != nil || len(page.Vaults) == 0 || len(page.Vaults) > size { - t.Fatalf("page: %+v, %v", page, err) - } - actual = append(actual, page.Vaults...) - if len(actual) > len(all) { - t.Fatal("pagination repeated records") - } - if page.NextCursor == "" { - break - } - if page.NextCursor != page.Vaults[len(page.Vaults)-1].ID { - t.Fatal("cursor is not the last included resource") - } - cursor = page.NextCursor - } - return actual - } - for _, ascending := range []bool{true, false} { - for _, statuses := range [][]string{nil, {"active"}, {"archived"}, {"active", "archived"}} { - want := []Vault{} - for _, vault := range all { - if len(statuses) != 1 || archived[vault.ID] == (statuses[0] == "archived") { - want = append(want, vault) - } - } - if !ascending { - slices.Reverse(want) - } - for _, size := range []int{20, 100} { - if got := read(s, ascending, statuses, size); !reflect.DeepEqual(got, want) { - t.Fatalf("filtered ordering/projection mismatch: ascending=%t statuses=%v size=%d got=%d want=%d", ascending, statuses, size, len(got), len(want)) - } - } - } - } - for _, cursor := range []string{foreign.ID, uuid.NewString(), "invalid"} { - if _, err := s.ListVaults(ctx, tenant, cursor, 20, true, []string{"archived"}); !errors.Is(err, ErrNotFound) { - t.Fatalf("foreign/unknown/malformed cursor: %v", err) - } - } - for _, tc := range []struct { - tenant, cursor string - limit int - statuses []string - }{{"invalid", "", 20, nil}, {tenant, "", 0, nil}, {tenant, "", 101, nil}, {tenant, "", 20, []string{"deleted"}}} { - if _, err := s.ListVaults(ctx, tc.tenant, tc.cursor, tc.limit, true, tc.statuses); !errors.Is(err, ErrInvalidInput) { - t.Fatalf("invalid store query: %v", err) - } - } - tail, err := s.ListVaults(ctx, tenant, all[len(all)-1].ID, 100, true, nil) - if err != nil || tail.Vaults == nil || len(tail.Vaults) != 0 || tail.NextCursor != "" { - t.Fatalf("terminal page: %+v, %v", tail, err) - } - page, err := s.ListVaults(ctx, other, "", 100, false, nil) - if err != nil || !reflect.DeepEqual(page.Vaults, []Vault{foreign}) || page.NextCursor != "" { - t.Fatalf("project isolation: %+v, %v", page, err) - } - pool.Close() - reopened, _ := testStore(t) - if got := read(reopened, true, nil, 20); !reflect.DeepEqual(got, all) { - t.Fatal("listing changed after reconnect") - } -} diff --git a/services/core/internal/store/vaults_test.go b/services/core/internal/store/vaults_test.go deleted file mode 100644 index 09e50a090..000000000 --- a/services/core/internal/store/vaults_test.go +++ /dev/null @@ -1,87 +0,0 @@ -package store - -import ( - "context" - "errors" - "reflect" - "strings" - "testing" - "time" - - "github.com/google/uuid" -) - -func TestVaultsPersistAndStayTenantScoped(t *testing.T) { - s, pool := testStore(t) - ctx := context.Background() - tenantA, tenantB := uuid.NewString(), uuid.NewString() - before := time.Now().Add(-time.Second) - unnamed, err := s.CreateVault(ctx, tenantA, CreateVaultInput{}) - if err != nil { - t.Fatal(err) - } - if _, err := uuid.Parse(unnamed.ID); err != nil || unnamed.TenantID != tenantA || unnamed.Name != nil || unnamed.Metadata == nil || len(unnamed.Metadata) != 0 || unnamed.CreatedAt.Before(before) || unnamed.CreatedAt.After(time.Now().Add(time.Second)) { - t.Fatalf("unexpected unnamed vault: %+v, %v", unnamed, err) - } - // Validate the byte boundary with multibyte text, without Session metadata - // count or character limits. Public name trimming belongs to the API layer. - name := strings.Repeat("é", 128) - input := CreateVaultInput{Name: &name, Metadata: map[string]string{"": "", "purpose": "保存 configuration"}} - named, err := s.CreateVault(ctx, tenantA, input) - if err != nil || named.ID == unnamed.ID || named.Name == nil || *named.Name != name || !reflect.DeepEqual(named.Metadata, input.Metadata) { - t.Fatalf("unexpected named vault: %+v, %v", named, err) - } - for _, tenant := range []string{tenantA, tenantB} { - other, err := s.CreateVault(ctx, tenant, input) - if err != nil || other.ID == named.ID || other.TenantID != tenant { - t.Fatalf("distinct resource creation: %+v, %v", other, err) - } - } - for _, lookup := range []struct{ tenant, id string }{{tenantB, named.ID}, {tenantA, uuid.NewString()}} { - if _, err := s.GetVault(ctx, lookup.tenant, lookup.id); !errors.Is(err, ErrNotFound) { - t.Fatalf("unowned/unknown vault lookup: %v", err) - } - } - // Recreate the pool and Store as a restarted standalone service would. - pool.Close() - reopened, _ := testStore(t) - for _, want := range []Vault{unnamed, named} { - got, err := reopened.GetVault(ctx, tenantA, want.ID) - if err != nil || !reflect.DeepEqual(got, want) { - t.Fatalf("durable read: %+v, %v; want %+v", got, err, want) - } - } - var sessions int - if err := reopened.pool.QueryRow(ctx, "SELECT count(*) FROM sessions WHERE tenant_id = $1", tenantA).Scan(&sessions); err != nil || sessions != 0 { - t.Fatalf("Vault creation produced %d Sessions: %v", sessions, err) - } -} - -func TestVaultsRejectInvalidStoreInputWithoutWrites(t *testing.T) { - s, _ := testStore(t) - ctx := context.Background() - tenant := uuid.NewString() - for _, name := range []string{"", strings.Repeat("x", 257), strings.Repeat("é", 129), string([]byte{0xff})} { - if _, err := s.CreateVault(ctx, tenant, CreateVaultInput{Name: &name}); !errors.Is(err, ErrInvalidInput) { - t.Fatalf("invalid name length %d: %v", len(name), err) - } - } - for _, invalid := range []string{"", "not-a-uuid", uuid.Nil.String()} { - if _, err := s.CreateVault(ctx, invalid, CreateVaultInput{}); !errors.Is(err, ErrInvalidInput) { - t.Fatalf("invalid create tenant accepted: %v", err) - } - if _, err := s.GetVault(ctx, invalid, uuid.NewString()); !errors.Is(err, ErrInvalidInput) { - t.Fatalf("invalid read tenant accepted: %v", err) - } - if _, err := s.GetVault(ctx, tenant, invalid); !errors.Is(err, ErrInvalidInput) { - t.Fatalf("invalid vault ID accepted: %v", err) - } - } - if _, err := s.CreateVault(ctx, tenant, CreateVaultInput{Metadata: map[string]string{"large": strings.Repeat("x", 64*1024)}}); !errors.Is(err, ErrInvalidInput) { - t.Fatalf("oversized metadata accepted: %v", err) - } - var count int - if err := s.pool.QueryRow(ctx, "SELECT count(*) FROM vaults WHERE tenant_id = $1", tenant).Scan(&count); err != nil || count != 0 { - t.Fatalf("invalid input created %d Vaults: %v", count, err) - } -} diff --git a/services/core/internal/store/write_audit_resources_test.go b/services/core/internal/store/write_audit_resources_test.go index 434bccbcb..f30ebf118 100644 --- a/services/core/internal/store/write_audit_resources_test.go +++ b/services/core/internal/store/write_audit_resources_test.go @@ -11,7 +11,6 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/writeaudit" "github.com/google/uuid" "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgxpool" ) type resourceAuditMutation struct { @@ -48,19 +47,14 @@ func TestWriteAuditStandaloneResourceTransactions(t *testing.T) { for _, name := range []string{ "template_create", "template_update", "template_delete", "skill_create", "skill_upload_version", "skill_update_default", "skill_delete", "version_delete", "version_delete_last", - "vault_create", "vault_delete", "credential_create", "credential_update", "credential_delete", - "oauth_create", "oauth_update", "oauth_delete", } { t.Run(name, func(t *testing.T) { tenant := uuid.NewString() mutation := prepareResourceAuditMutation(t, s, tenant, name, archive) snapshot := func() map[string]string { result := make(map[string]string) - for _, table := range []string{"agents", "environment_templates", "skills", "skill_versions", "vaults", "vault_credentials", "write_audit_operations", "write_audit_owners"} { + for _, table := range []string{"agents", "environment_templates", "skills", "skill_versions", "write_audit_operations", "write_audit_owners"} { query := "SELECT COALESCE(jsonb_agg(to_jsonb(r) ORDER BY to_jsonb(r)::text)::text, '[]') FROM " + pgx.Identifier{table}.Sanitize() + " r WHERE tenant_id=$1" - if table == "vault_credentials" { - query = "SELECT COALESCE(jsonb_agg(to_jsonb(r) ORDER BY r.id)::text, '[]') FROM vault_credentials r JOIN vaults v ON v.id=r.vault_id WHERE v.tenant_id=$1" - } var value string if err := pool.QueryRow(ctx, query, tenant).Scan(&value); err != nil { t.Fatalf("snapshot %s: %v", table, err) @@ -170,64 +164,6 @@ func prepareResourceAuditMutation(t *testing.T, s *Store, tenant, name string, a return v.ID, e }} } - if name == "vault_create" { - return resourceAuditMutation{action: "create", kind: "vault", owners: 1, run: func(ctx context.Context) (string, error) { - v, e := s.CreateVault(ctx, tenant, CreateVaultInput{}) - return v.ID, e - }} - } - vault, err := s.CreateVault(ctx, tenant, CreateVaultInput{}) - must(err) - static := CreateStaticCredentialInput{Name: "fixture", MCPServerURL: "https://mcp.example/", Token: "audit-private-token"} - if name == "credential_create" { - return resourceAuditMutation{action: "create", kind: "credential", parent: vault.ID, owners: 1, run: func(ctx context.Context) (string, error) { - v, e := s.CreateStaticCredential(ctx, tenant, vault.ID, static) - return v.ID, e - }} - } - if strings.HasPrefix(name, "oauth_") { - input := CreateOAuthCredentialInput{Name: "fixture", MCPServerURL: static.MCPServerURL, AccessToken: static.Token} - if name == "oauth_create" { - return resourceAuditMutation{action: "create", kind: "credential", parent: vault.ID, owners: 1, run: func(ctx context.Context) (string, error) { - v, e := s.CreateOAuthCredential(ctx, tenant, vault.ID, input) - return v.ID, e - }} - } - v, err := s.CreateOAuthCredential(ctx, tenant, vault.ID, input) - must(err) - if name == "oauth_update" { - return resourceAuditMutation{action: "update", kind: "credential", parent: vault.ID, run: func(ctx context.Context) (string, error) { - token := "audit-private-replacement" - v, e := s.UpdateOAuthCredential(ctx, tenant, vault.ID, v.ID, UpdateOAuthCredentialInput{AccessToken: &token}) - return v.ID, e - }} - } - return resourceAuditMutation{action: "delete", kind: "credential", parent: vault.ID, run: func(ctx context.Context) (string, error) { return s.DeleteCredential(ctx, tenant, vault.ID, v.ID) }} - } - v, err := s.CreateStaticCredential(ctx, tenant, vault.ID, static) - must(err) - if name == "vault_delete" { - return resourceAuditMutation{action: "delete", kind: "vault", run: func(ctx context.Context) (string, error) { return s.DeleteVault(ctx, tenant, vault.ID) }} - } - if name == "credential_update" { - return resourceAuditMutation{action: "update", kind: "credential", parent: vault.ID, run: func(ctx context.Context) (string, error) { - v, e := s.UpdateStaticCredential(ctx, tenant, vault.ID, v.ID, UpdateStaticCredentialInput{Token: "audit-private-replacement"}) - return v.ID, e - }} - } - return resourceAuditMutation{action: "delete", kind: "credential", parent: vault.ID, run: func(ctx context.Context) (string, error) { return s.DeleteCredential(ctx, tenant, vault.ID, v.ID) }} -} - -// A public mutation owns its transaction, so inject database failures through -// the connection configuration rather than replacing a Store query wrapper. -func readOnlyResourceStore(t *testing.T, pool *pgxpool.Pool) *Store { - t.Helper() - config := pool.Config().Copy() - config.ConnConfig.RuntimeParams["default_transaction_read_only"] = "on" - readOnly, err := pgxpool.NewWithConfig(t.Context(), config) - if err != nil { - t.Fatal(err) - } - t.Cleanup(readOnly.Close) - return New(readOnly) + t.Fatal("unknown resource audit mutation", name) + return resourceAuditMutation{} } diff --git a/services/core/internal/vaults/credential.go b/services/core/internal/vaults/credential.go new file mode 100644 index 000000000..5a28ddfd9 --- /dev/null +++ b/services/core/internal/vaults/credential.go @@ -0,0 +1,37 @@ +package vaults + +import "time" + +// Credential authentication types. +const ( + AuthStaticBearer = "static_bearer" + AuthMCPOAuth = "mcp_oauth" +) + +// Credential contains only public metadata. Resource reads never select the +// sealed secret; only the bearer-token lookup for execution opens it. +type Credential struct { + ID, VaultID, Name, AuthType, MCPServerURL string + CreatedAt, UpdatedAt time.Time + // OAuth is set for mcp_oauth Credentials only. + OAuth *OAuthMetadata +} + +type CredentialPage struct { + Credentials []Credential + NextCursor string +} + +// OAuthMetadata is the safe projection of a stored OAuth grant. +type OAuthMetadata struct { + ExpiresAt *string `json:"expires_at"` + Refresh *OAuthRefreshMetadata `json:"refresh"` +} + +type OAuthRefreshMetadata struct { + ClientID string `json:"client_id"` + TokenEndpoint string `json:"token_endpoint"` + TokenEndpointAuth string `json:"token_endpoint_auth"` + Resource *string `json:"resource"` + Scope *string `json:"scope"` +} diff --git a/services/core/internal/vaults/doc.go b/services/core/internal/vaults/doc.go new file mode 100644 index 000000000..87ec62989 --- /dev/null +++ b/services/core/internal/vaults/doc.go @@ -0,0 +1,5 @@ +// Package vaults owns Vaults and their Credentials: the resources, the +// encryption of Credential secrets, OAuth access-token refresh, and the MCP +// credential selection that Session creation freezes and execution resolves +// into a bearer token. +package vaults diff --git a/services/core/internal/vaults/errors.go b/services/core/internal/vaults/errors.go new file mode 100644 index 000000000..75a95ba31 --- /dev/null +++ b/services/core/internal/vaults/errors.go @@ -0,0 +1,54 @@ +package vaults + +import ( + "errors" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/echotext" +) + +var ( + // ErrNotFound reports a Vault or Credential that does not exist in the + // caller's tenant, including one another tenant owns and an ID that cannot + // name one. + ErrNotFound = errors.New("vault resource not found") + // ErrInvalidInput rejects a request the operation cannot accept. + ErrInvalidInput = errors.New("invalid vault request") +) + +// MCPCredentialSelectionError rejects a Session MCP credential selection with +// the observed official message (MV-03); the API reports a Conflict as 409 +// conflict_error and any other as 400 invalid_request_error. Selection searches +// only the attached Vaults, which the caller owns, so a missing, foreign-tenant, +// unattached or malformed reference produces the same error, and only a +// credential of an attached Vault can report a server_url mismatch. +type MCPCredentialSelectionError struct { + Conflict bool + Message string +} + +func (e *MCPCredentialSelectionError) Error() string { return e.Message } + +// echoed repeats a caller-supplied value in a selection message only within +// the shared bound; otherwise the message leaves it out. +func echoed(value string) string { + if !echotext.Allowed(value) { + return "" + } + return " " + value +} + +func mcpCredentialRequiresVault() error { + return &MCPCredentialSelectionError{Message: "MCP credential_id requires an attached vault"} +} + +func mcpCredentialNotAttached(id string) error { + return &MCPCredentialSelectionError{Message: "MCP credential_id" + echoed(id) + " was not found in an attached vault"} +} + +func mcpCredentialURLMismatch(id, url string) error { + return &MCPCredentialSelectionError{Message: "MCP credential_id" + echoed(id) + " does not match server_url" + echoed(url)} +} + +func mcpCredentialAmbiguous(url string) error { + return &MCPCredentialSelectionError{Conflict: true, Message: "multiple attached vault credentials match MCP server_url" + echoed(url) + "; specify credential_id"} +} diff --git a/services/core/internal/vaults/oauth.go b/services/core/internal/vaults/oauth.go new file mode 100644 index 000000000..d455ef484 --- /dev/null +++ b/services/core/internal/vaults/oauth.go @@ -0,0 +1,158 @@ +package vaults + +import ( + "errors" + "time" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" +) + +// OAuthRefreshUpdate patches a stored refresh configuration. ScopeSet keeps +// omitted apart from null. +type OAuthRefreshUpdate struct { + RefreshToken *string + Scope *string + ScopeSet bool + TokenEndpointAuthType string + ClientSecret *string +} + +// oauthSecret is the sealed plaintext of an mcp_oauth Credential. The +// encrypted copy authenticates every public setting used for refresh, +// including the token endpoint, so substituting stored metadata can never +// redirect a grant. +type oauthSecret struct { + Version int `json:"version"` + Metadata OAuthMetadata `json:"metadata"` + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + ClientSecret string `json:"client_secret"` +} + +func validOAuthMetadata(metadata OAuthMetadata) bool { + if metadata.ExpiresAt != nil { + if _, err := time.Parse(time.RFC3339Nano, *metadata.ExpiresAt); err != nil { + return false + } + } + if refresh := metadata.Refresh; refresh != nil { + if refresh.ClientID == "" || refresh.TokenEndpoint == "" { + return false + } + switch refresh.TokenEndpointAuth { + case "none", "client_secret_basic", "client_secret_post": + default: + return false + } + } + return true +} + +// validOAuthCreation checks a new grant: secrets that the refresh +// configuration cannot use are rejected rather than stored. +func validOAuthCreation(command CreateOAuthCredential) bool { + if !validName(command.Name) || command.MCPServerURL == "" || !validOAuthMetadata(command.OAuth) { + return false + } + refresh := command.OAuth.Refresh + if refresh == nil { + return command.RefreshToken == "" && command.ClientSecret == "" + } + return refresh.TokenEndpointAuth != "none" || command.ClientSecret == "" +} + +// oauthBinding seals an OAuth secret to its tenant, Vault, Credential and +// destination. +func oauthBinding(tenantID string, credential Credential) credentialcrypto.Binding { + return credentialcrypto.Binding{TenantID: tenantID, VaultID: credential.VaultID, + CredentialID: credential.ID, AuthType: AuthMCPOAuth, Destination: credential.MCPServerURL} +} + +// applyOAuthUpdate patches a stored grant. A new access token clears an +// omitted expiry. A refresh patch cannot add configuration or change the +// authentication method, and a client secret needs a method that uses one. +func applyOAuthUpdate(secret oauthSecret, update UpdateOAuthCredential) (oauthSecret, error) { + if update.AccessToken != nil { + secret.AccessToken = *update.AccessToken + secret.Metadata.ExpiresAt = nil + } + if update.ExpiresAtSet { + secret.Metadata.ExpiresAt = update.ExpiresAt + } + patch := update.Refresh + if patch == nil { + return secret, nil + } + if secret.Metadata.Refresh == nil { + return oauthSecret{}, ErrInvalidInput + } + refresh := *secret.Metadata.Refresh + if patch.TokenEndpointAuthType != "" && patch.TokenEndpointAuthType != refresh.TokenEndpointAuth { + return oauthSecret{}, ErrInvalidInput + } + if patch.ClientSecret != nil { + if refresh.TokenEndpointAuth == "none" { + return oauthSecret{}, ErrInvalidInput + } + secret.ClientSecret = *patch.ClientSecret + } + if patch.RefreshToken != nil { + secret.RefreshToken = *patch.RefreshToken + } + if patch.ScopeSet { + refresh.Scope = patch.Scope + } + secret.Metadata.Refresh = &refresh + return secret, nil +} + +// currentAccessToken returns the stored access token while it is usable at +// now. expired reports that the grant must be refreshed first; a grant without +// an expiry never expires. +func currentAccessToken(secret oauthSecret, now time.Time) (token string, expired bool, err error) { + if secret.Metadata.ExpiresAt != nil { + expiry, err := time.Parse(time.RFC3339Nano, *secret.Metadata.ExpiresAt) + if err != nil { + return "", false, errors.New("invalid OAuth token expiry") + } + if !now.Before(expiry) { + return "", true, nil + } + } + if secret.AccessToken == "" { + return "", false, errors.New("OAuth access token is missing") + } + return secret.AccessToken, false, nil +} + +// refreshRequest is the exchange that renews an expired grant. +func refreshRequest(secret oauthSecret) (oauthrefresh.Request, error) { + refresh := secret.Metadata.Refresh + if refresh == nil || secret.RefreshToken == "" { + return oauthrefresh.Request{}, errors.New("expired OAuth credential cannot be refreshed") + } + return oauthrefresh.Request{ + TokenEndpoint: refresh.TokenEndpoint, ClientID: refresh.ClientID, + AuthMethod: refresh.TokenEndpointAuth, ClientSecret: secret.ClientSecret, + RefreshToken: secret.RefreshToken, Resource: refresh.Resource, Scope: refresh.Scope, + }, nil +} + +// applyRefreshedToken stores a refresh result that is usable at now. An +// omitted refresh token keeps the stored one. +func applyRefreshedToken(secret oauthSecret, token oauthrefresh.Token, now time.Time) (oauthSecret, error) { + if token.AccessToken == "" || token.ExpiresAt != nil && !now.Before(*token.ExpiresAt) { + return oauthSecret{}, errors.New("OAuth refresh returned an unusable token") + } + secret.AccessToken = token.AccessToken + if token.RefreshToken != "" { + secret.RefreshToken = token.RefreshToken + } + secret.Metadata.ExpiresAt = nil + if token.ExpiresAt != nil { + value := token.ExpiresAt.UTC().Format(time.RFC3339Nano) + secret.Metadata.ExpiresAt = &value + } + return secret, nil +} diff --git a/services/core/internal/vaults/page.go b/services/core/internal/vaults/page.go new file mode 100644 index 000000000..bd3cd0f70 --- /dev/null +++ b/services/core/internal/vaults/page.go @@ -0,0 +1,38 @@ +package vaults + +import "fmt" + +// Vault and Credential statuses. +const ( + StatusActive = "active" + StatusArchived = "archived" +) + +// PageQuery selects one page of Vaults, or of one Vault's Credentials, ordered +// by creation time and then ID. +type PageQuery struct { + // After is the last ID of the previous page, or empty for the first page. + // An After that names no resource of the list is ErrNotFound. + After string + Limit int + Ascending bool + // Statuses filters by status; empty selects every status. + Statuses []string +} + +// Validate checks the page size and statuses and returns the statuses to +// select. +func (q PageQuery) Validate() ([]string, error) { + if q.Limit < 1 || q.Limit > 100 { + return nil, fmt.Errorf("%w: internal page size must be 1..100", ErrInvalidInput) + } + if len(q.Statuses) == 0 { + return []string{StatusActive, StatusArchived}, nil + } + for _, status := range q.Statuses { + if status != StatusActive && status != StatusArchived { + return nil, fmt.Errorf("%w: invalid status", ErrInvalidInput) + } + } + return q.Statuses, nil +} diff --git a/services/core/internal/vaults/rules_test.go b/services/core/internal/vaults/rules_test.go new file mode 100644 index 000000000..e6f9110fe --- /dev/null +++ b/services/core/internal/vaults/rules_test.go @@ -0,0 +1,283 @@ +package vaults + +import ( + "errors" + "reflect" + "strings" + "testing" + "time" + + "github.com/google/uuid" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" +) + +func TestNamesAndIDs(t *testing.T) { + for name, valid := range map[string]bool{ + "x": true, strings.Repeat("é", 128): true, strings.Repeat("x", 256): true, + "": false, strings.Repeat("x", 257): false, strings.Repeat("é", 129): false, string([]byte{0xff}): false, + } { + if validName(name) != valid { + t.Fatalf("validName(%d bytes) = %t", len(name), !valid) + } + } + id := uuid.New() + for raw, want := range map[string]string{ + id.String(): id.String(), strings.ToUpper(id.String()): id.String(), "{" + id.String() + "}": id.String(), + "": "", "invalid": "", uuid.Nil.String(): "", + } { + got, ok := canonicalID(raw) + if got != want || ok != (want != "") { + t.Fatalf("canonicalID(%q) = %q, %t", raw, got, ok) + } + } +} + +func TestPageQueryValidate(t *testing.T) { + statuses, err := PageQuery{Limit: 1}.Validate() + if err != nil || !reflect.DeepEqual(statuses, []string{StatusActive, StatusArchived}) { + t.Fatal("empty statuses do not select every status", statuses, err) + } + statuses, err = PageQuery{Limit: 100, Statuses: []string{StatusArchived}}.Validate() + if err != nil || !reflect.DeepEqual(statuses, []string{StatusArchived}) { + t.Fatal("status filter changed", statuses, err) + } + for _, query := range []PageQuery{{Limit: 0}, {Limit: 101}, {Limit: 20, Statuses: []string{"deleted"}}} { + if _, err := query.Validate(); !errors.Is(err, ErrInvalidInput) { + t.Fatalf("invalid query %+v accepted: %v", query, err) + } + } +} + +func TestMCPCredentialSelectionRules(t *testing.T) { + first, second := uuid.NewString(), uuid.NewString() + attached, err := attachedVaultIDs([]string{first, strings.ToUpper(first), second}) + if err != nil || !reflect.DeepEqual(attached, []string{first, second}) { + t.Fatal("attached Vaults were not canonical and deduplicated", attached, err) + } + if _, err := attachedVaultIDs([]string{first, "invalid"}); !errors.Is(err, ErrNotFound) { + t.Fatal("a malformed attached Vault named one", err) + } + url := "https://mcp.example/tools" + named := func(id string) MCPCredentialRequest { + return MCPCredentialRequest{ServerLabel: "tools", ServerURL: url, CredentialID: &id} + } + for _, tc := range []struct { + request MCPCredentialRequest + attached []string + want string + err error + }{ + {MCPCredentialRequest{ServerLabel: "tools", ServerURL: url}, nil, "", nil}, + {named(strings.ToUpper(second)), attached, second, nil}, + {MCPCredentialRequest{ServerURL: url}, attached, "", ErrInvalidInput}, + {MCPCredentialRequest{ServerLabel: "tools"}, attached, "", ErrInvalidInput}, + {named(second), nil, "", mcpCredentialRequiresVault()}, + {named("not-a-credential"), attached, "", mcpCredentialNotAttached("not-a-credential")}, + } { + got, err := mcpCredentialLookup(tc.request, tc.attached) + if got != tc.want || !sameError(err, tc.err) { + t.Fatalf("lookup %+v = %q, %v", tc.request, got, err) + } + } + match := MCPCredentialMatch{VaultID: first, CredentialID: second, AuthType: AuthStaticBearer, MCPServerURL: url} + bound := MCPCredentialBinding{ServerLabel: "tools", ServerURL: url, VaultID: first, CredentialID: second, AuthType: AuthStaticBearer} + for _, tc := range []struct { + request MCPCredentialRequest + matches []MCPCredentialMatch + want MCPCredentialBinding + err error + }{ + {MCPCredentialRequest{ServerLabel: "tools", ServerURL: url}, nil, MCPCredentialBinding{ServerLabel: "tools", ServerURL: url}, nil}, + {MCPCredentialRequest{ServerLabel: "tools", ServerURL: url}, []MCPCredentialMatch{match}, bound, nil}, + {named(second), []MCPCredentialMatch{match}, bound, nil}, + {MCPCredentialRequest{ServerLabel: "tools", ServerURL: url}, []MCPCredentialMatch{match, match}, MCPCredentialBinding{}, mcpCredentialAmbiguous(url)}, + {named(second), nil, MCPCredentialBinding{}, mcpCredentialNotAttached(second)}, + {MCPCredentialRequest{ServerLabel: "tools", ServerURL: url + "/other", CredentialID: &second}, []MCPCredentialMatch{match}, MCPCredentialBinding{}, mcpCredentialURLMismatch(second, url+"/other")}, + } { + got, err := selectMCPCredential(tc.request, tc.matches) + if got != tc.want || !sameError(err, tc.err) { + t.Fatalf("select %+v = %+v, %v", tc.request, got, err) + } + } +} + +func TestMCPCredentialSelectionMessages(t *testing.T) { + id, url := uuid.NewString(), "https://mcp.example/tools" + for _, tc := range []struct { + err error + conflict bool + message string + }{ + {mcpCredentialRequiresVault(), false, "MCP credential_id requires an attached vault"}, + {mcpCredentialNotAttached(id), false, "MCP credential_id " + id + " was not found in an attached vault"}, + {mcpCredentialURLMismatch(id, url), false, "MCP credential_id " + id + " does not match server_url " + url}, + {mcpCredentialAmbiguous(url), true, "multiple attached vault credentials match MCP server_url " + url + "; specify credential_id"}, + // A value beyond the shared echo bound is left out. + {mcpCredentialNotAttached(strings.Repeat("x", 4096)), false, "MCP credential_id was not found in an attached vault"}, + } { + var selection *MCPCredentialSelectionError + if !errors.As(tc.err, &selection) || selection.Conflict != tc.conflict || selection.Message != tc.message { + t.Fatalf("selection error %v, want %q", tc.err, tc.message) + } + } +} + +func TestBearerTokenScope(t *testing.T) { + tenant, vault, credential := uuid.New(), uuid.New(), uuid.New() + binding := MCPCredentialBinding{ServerLabel: "tools", ServerURL: "https://mcp.example/tools", VaultID: strings.ToUpper(vault.String()), CredentialID: credential.String(), AuthType: AuthStaticBearer} + scope, err := bearerTokenScope(MCPBearerToken{TenantID: tenant.String(), VaultIDs: []string{vault.String()}, Binding: binding}) + want := bearerScope{tenantID: tenant.String(), vaultID: vault.String(), credentialID: credential.String(), attached: []string{vault.String()}} + if err != nil || !reflect.DeepEqual(scope, want) || !scope.vaultAttached() { + t.Fatal("frozen scope was not canonical", scope, err) + } + for _, mutate := range []func(*MCPBearerToken){ + func(c *MCPBearerToken) { c.TenantID = "invalid" }, + func(c *MCPBearerToken) { c.VaultIDs = []string{"invalid"} }, + func(c *MCPBearerToken) { c.Binding.VaultID = "" }, + func(c *MCPBearerToken) { c.Binding.CredentialID = uuid.Nil.String() }, + func(c *MCPBearerToken) { c.Binding.AuthType = "other" }, + func(c *MCPBearerToken) { c.Binding.ServerURL = "" }, + } { + command := MCPBearerToken{TenantID: tenant.String(), VaultIDs: []string{vault.String()}, Binding: binding} + mutate(&command) + if _, err := bearerTokenScope(command); !errors.Is(err, ErrNotFound) { + t.Fatal("an incomplete frozen binding was admitted", command, err) + } + } + scope, err = bearerTokenScope(MCPBearerToken{TenantID: tenant.String(), Binding: binding}) + if err != nil || scope.vaultAttached() { + t.Fatal("an unattached Vault was reported attached", err) + } +} + +func TestOAuthCreationRules(t *testing.T) { + refresh := func(method string) *OAuthRefreshMetadata { + return &OAuthRefreshMetadata{ClientID: "client", TokenEndpoint: "https://issuer.example/token", TokenEndpointAuth: method} + } + valid := CreateOAuthCredential{Name: "OAuth", MCPServerURL: "https://mcp.example/tools"} + for _, tc := range []struct { + mutate func(*CreateOAuthCredential) + valid bool + }{ + {func(*CreateOAuthCredential) {}, true}, + {func(c *CreateOAuthCredential) { c.OAuth.ExpiresAt = ptr("2030-01-02T03:04:05.123456789Z") }, true}, + {func(c *CreateOAuthCredential) { + c.OAuth.Refresh, c.RefreshToken, c.ClientSecret = refresh("client_secret_post"), "r", "s" + }, true}, + {func(c *CreateOAuthCredential) { c.OAuth.Refresh, c.RefreshToken = refresh("none"), "r" }, true}, + {func(c *CreateOAuthCredential) { c.Name = "" }, false}, + {func(c *CreateOAuthCredential) { c.MCPServerURL = "" }, false}, + {func(c *CreateOAuthCredential) { c.OAuth.ExpiresAt = ptr("not-a-date") }, false}, + {func(c *CreateOAuthCredential) { c.OAuth.Refresh = refresh("private_key_jwt") }, false}, + {func(c *CreateOAuthCredential) { c.OAuth.Refresh = &OAuthRefreshMetadata{TokenEndpointAuth: "none"} }, false}, + {func(c *CreateOAuthCredential) { c.RefreshToken = "r" }, false}, + {func(c *CreateOAuthCredential) { c.ClientSecret = "s" }, false}, + {func(c *CreateOAuthCredential) { c.OAuth.Refresh, c.ClientSecret = refresh("none"), "s" }, false}, + } { + command := valid + tc.mutate(&command) + if validOAuthCreation(command) != tc.valid { + t.Fatalf("validOAuthCreation(%+v) = %t", command, !tc.valid) + } + } +} + +func TestOAuthUpdateRules(t *testing.T) { + expiry := "2030-01-02T03:04:05Z" + stored := oauthSecret{Version: 1, AccessToken: "access", RefreshToken: "refresh", ClientSecret: "secret", + Metadata: OAuthMetadata{ExpiresAt: &expiry, Refresh: &OAuthRefreshMetadata{ClientID: "client", TokenEndpoint: "https://issuer.example/token", TokenEndpointAuth: "client_secret_basic", Scope: ptr("read write")}}} + got, err := applyOAuthUpdate(stored, UpdateOAuthCredential{}) + if err != nil || !reflect.DeepEqual(got, stored) { + t.Fatal("an empty patch changed the grant", err) + } + got, err = applyOAuthUpdate(stored, UpdateOAuthCredential{AccessToken: ptr("new")}) + if err != nil || got.AccessToken != "new" || got.Metadata.ExpiresAt != nil || got.RefreshToken != "refresh" || got.ClientSecret != "secret" { + t.Fatal("a new access token kept the old expiry or changed other secrets", err) + } + got, err = applyOAuthUpdate(stored, UpdateOAuthCredential{AccessToken: ptr("new"), ExpiresAtSet: true, ExpiresAt: ptr("2031-01-01T00:00:00Z")}) + if err != nil || *got.Metadata.ExpiresAt != "2031-01-01T00:00:00Z" { + t.Fatal("an explicit expiry was not kept", err) + } + got, err = applyOAuthUpdate(stored, UpdateOAuthCredential{Refresh: &OAuthRefreshUpdate{RefreshToken: ptr("r2"), ClientSecret: ptr("s2"), ScopeSet: true, TokenEndpointAuthType: "client_secret_basic"}}) + if err != nil || got.RefreshToken != "r2" || got.ClientSecret != "s2" || got.Metadata.Refresh.Scope != nil || *stored.Metadata.Refresh.Scope != "read write" { + t.Fatal("refresh patch was not applied, or it changed the stored copy", err) + } + none := stored + none.Metadata.Refresh = &OAuthRefreshMetadata{ClientID: "client", TokenEndpoint: "https://issuer.example/token", TokenEndpointAuth: "none"} + withoutRefresh := stored + withoutRefresh.Metadata.Refresh = nil + for _, tc := range []struct { + secret oauthSecret + patch OAuthRefreshUpdate + }{ + {withoutRefresh, OAuthRefreshUpdate{RefreshToken: ptr("cannot-add")}}, + {stored, OAuthRefreshUpdate{TokenEndpointAuthType: "client_secret_post"}}, + {none, OAuthRefreshUpdate{ClientSecret: ptr("unused")}}, + } { + if _, err := applyOAuthUpdate(tc.secret, UpdateOAuthCredential{Refresh: &tc.patch}); !errors.Is(err, ErrInvalidInput) { + t.Fatal("an invalid refresh patch was applied", tc.patch, err) + } + } +} + +func TestOAuthAccessTokenAndRefreshRules(t *testing.T) { + now := time.Date(2030, 1, 1, 0, 0, 0, 0, time.UTC) + at := func(when time.Time) *string { value := when.Format(time.RFC3339Nano); return &value } + for _, tc := range []struct { + expiresAt *string + access string + token string + expired bool + failed bool + }{ + {nil, "access", "access", false, false}, + {at(now.Add(time.Second)), "access", "access", false, false}, + {at(now), "access", "", true, false}, + {at(now.Add(-time.Hour)), "", "", true, false}, + {nil, "", "", false, true}, + {ptr("not-a-date"), "access", "", false, true}, + } { + token, expired, err := currentAccessToken(oauthSecret{AccessToken: tc.access, Metadata: OAuthMetadata{ExpiresAt: tc.expiresAt}}, now) + if token != tc.token || expired != tc.expired || (err != nil) != tc.failed { + t.Fatalf("currentAccessToken(%v, %q) = %q, %t, %v", tc.expiresAt, tc.access, token, expired, err) + } + } + secret := oauthSecret{RefreshToken: "refresh", ClientSecret: "secret", Metadata: OAuthMetadata{Refresh: &OAuthRefreshMetadata{ + ClientID: "client", TokenEndpoint: "https://issuer.example/token", TokenEndpointAuth: "client_secret_post", Resource: ptr("https://mcp.example/tools"), Scope: ptr("read")}}} + request, err := refreshRequest(secret) + want := oauthrefresh.Request{TokenEndpoint: "https://issuer.example/token", ClientID: "client", AuthMethod: "client_secret_post", ClientSecret: "secret", RefreshToken: "refresh", Resource: ptr("https://mcp.example/tools"), Scope: ptr("read")} + if err != nil || !reflect.DeepEqual(request, want) { + t.Fatal("refresh request lost grant fields", err) + } + for _, unusable := range []oauthSecret{{RefreshToken: "refresh"}, {Metadata: secret.Metadata}} { + if _, err := refreshRequest(unusable); err == nil { + t.Fatal("a grant without refresh configuration or token was refreshed") + } + } + later := now.Add(time.Hour) + refreshed, err := applyRefreshedToken(secret, oauthrefresh.Token{AccessToken: "renewed", ExpiresAt: &later}, now) + if err != nil || refreshed.AccessToken != "renewed" || refreshed.RefreshToken != "refresh" || *refreshed.Metadata.ExpiresAt != later.Format(time.RFC3339Nano) { + t.Fatal("refreshed grant lost the stored refresh token or the expiry", err) + } + refreshed, err = applyRefreshedToken(refreshed, oauthrefresh.Token{AccessToken: "rotated", RefreshToken: "next"}, now) + if err != nil || refreshed.RefreshToken != "next" || refreshed.Metadata.ExpiresAt != nil { + t.Fatal("rotated refresh token or unknown expiry was not stored", err) + } + for _, token := range []oauthrefresh.Token{{}, {AccessToken: "stale", ExpiresAt: &now}} { + if _, err := applyRefreshedToken(secret, token, now); err == nil { + t.Fatal("an unusable refresh result was stored") + } + } +} + +func sameError(got, want error) bool { + var selection *MCPCredentialSelectionError + if errors.As(want, &selection) { + var actual *MCPCredentialSelectionError + return errors.As(got, &actual) && *actual == *selection + } + return errors.Is(got, want) +} + +func ptr(value string) *string { return &value } diff --git a/services/core/internal/vaults/selection.go b/services/core/internal/vaults/selection.go new file mode 100644 index 000000000..5a98c203f --- /dev/null +++ b/services/core/internal/vaults/selection.go @@ -0,0 +1,122 @@ +package vaults + +// MCPCredentialRequest is one MCP server of a Session being created, with the +// Credential the caller named for it, if any. +type MCPCredentialRequest struct { + ServerLabel, ServerURL string + CredentialID *string +} + +// MCPCredentialBinding freezes a non-secret selection, including anonymous +// servers. It is private execution configuration stored in the Session, not +// the public MCP tool shape, so its JSON names are stable. +type MCPCredentialBinding struct { + ServerLabel string `json:"server_label"` + ServerURL string `json:"server_url"` + VaultID string `json:"vault_id,omitempty"` + CredentialID string `json:"credential_id,omitempty"` + AuthType string `json:"auth_type,omitempty"` +} + +// MCPCredentialMatch is the non-secret identity of a Credential that an MCP +// credential lookup found. +type MCPCredentialMatch struct { + VaultID, CredentialID, AuthType, MCPServerURL string +} + +// attachedVaultIDs returns the canonical, deduplicated IDs of a Session's +// attached Vaults. An ID that cannot name a Vault is ErrNotFound. +func attachedVaultIDs(ids []string) ([]string, error) { + result := make([]string, 0, len(ids)) + seen := map[string]bool{} + for _, raw := range ids { + id, ok := canonicalID(raw) + if !ok { + return nil, ErrNotFound + } + if !seen[id] { + result = append(result, id) + seen[id] = true + } + } + return result, nil +} + +// mcpCredentialLookup checks a request before its lookup and returns the +// canonical Credential ID the caller named, or "" to select by destination. +func mcpCredentialLookup(request MCPCredentialRequest, attached []string) (string, error) { + if request.ServerLabel == "" || request.ServerURL == "" { + return "", ErrInvalidInput + } + if request.CredentialID == nil { + return "", nil + } + if len(attached) == 0 { + return "", mcpCredentialRequiresVault() + } + id, ok := canonicalID(*request.CredentialID) + if !ok { + return "", mcpCredentialNotAttached(*request.CredentialID) + } + return id, nil +} + +// selectMCPCredential binds a request to the Credentials its lookup found: a +// named Credential must exist in an attached Vault and match the destination, +// and a destination may match at most one Credential. +func selectMCPCredential(request MCPCredentialRequest, matches []MCPCredentialMatch) (MCPCredentialBinding, error) { + if request.CredentialID != nil { + if len(matches) == 0 { + return MCPCredentialBinding{}, mcpCredentialNotAttached(*request.CredentialID) + } + if matches[0].MCPServerURL != request.ServerURL { + return MCPCredentialBinding{}, mcpCredentialURLMismatch(*request.CredentialID, request.ServerURL) + } + } + if len(matches) > 1 { + return MCPCredentialBinding{}, mcpCredentialAmbiguous(request.ServerURL) + } + binding := MCPCredentialBinding{ServerLabel: request.ServerLabel, ServerURL: request.ServerURL} + if len(matches) == 1 { + binding.VaultID, binding.CredentialID, binding.AuthType = matches[0].VaultID, matches[0].CredentialID, matches[0].AuthType + } + return binding, nil +} + +// bearerScope is a frozen binding's complete authorization, in canonical IDs. +type bearerScope struct { + tenantID, vaultID, credentialID string + attached []string +} + +// bearerTokenScope rechecks a frozen binding before any lookup. Every failure +// is ErrNotFound: execution never downgrades to an anonymous request. +func bearerTokenScope(command MCPBearerToken) (bearerScope, error) { + tenant, ok := canonicalID(command.TenantID) + if !ok { + return bearerScope{}, ErrNotFound + } + attached, err := attachedVaultIDs(command.VaultIDs) + if err != nil { + return bearerScope{}, err + } + binding := command.Binding + vault, ok := canonicalID(binding.VaultID) + if !ok { + return bearerScope{}, ErrNotFound + } + id, ok := canonicalID(binding.CredentialID) + if !ok || (binding.AuthType != AuthStaticBearer && binding.AuthType != AuthMCPOAuth) || binding.ServerURL == "" { + return bearerScope{}, ErrNotFound + } + return bearerScope{tenantID: tenant, vaultID: vault, credentialID: id, attached: attached}, nil +} + +func (s bearerScope) vaultAttached() bool { + for _, candidate := range s.attached { + if candidate == s.vaultID { + return true + } + } + return false +} diff --git a/services/core/internal/vaults/service.go b/services/core/internal/vaults/service.go new file mode 100644 index 000000000..2c2d1b3d5 --- /dev/null +++ b/services/core/internal/vaults/service.go @@ -0,0 +1,414 @@ +package vaults + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "reflect" + "time" + + "github.com/google/uuid" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/metadata" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" +) + +// oauthRefreshTimeout bounds both lock contention and the external exchange, +// so a slow provider cannot hold a Credential indefinitely. The refresh client +// has its own tighter bound. +const oauthRefreshTimeout = 20 * time.Second + +// Service runs the Vault and Credential operations. +type Service struct { + storage Storage + cipher *credentialcrypto.Cipher + refresher oauthrefresh.Refresher +} + +// NewService requires storage and the OAuth refresher. A nil cipher means this +// Core has no credential key: operations that seal or open a secret then +// return credentialcrypto.ErrUnavailable, and the others keep working. +func NewService(storage Storage, cipher *credentialcrypto.Cipher, refresher oauthrefresh.Refresher) (*Service, error) { + if storage == nil { + return nil, errors.New("vaults: storage is required") + } + if refresher == nil { + return nil, errors.New("vaults: OAuth refresher is required") + } + return &Service{storage: storage, cipher: cipher, refresher: refresher}, nil +} + +type CreateVault struct { + TenantID string + // Name is the public layer's trimmed name, or nil for none. + Name *string + Metadata map[string]string +} + +// CreateVault stores a new Vault. Each call creates a distinct Vault. +func (s *Service) CreateVault(ctx context.Context, command CreateVault) (Vault, error) { + if command.Name != nil && !validName(*command.Name) { + return Vault{}, fmt.Errorf("%w: vault name must contain 1–256 UTF-8 bytes", ErrInvalidInput) + } + encoded, err := metadata.Encode(command.Metadata) + if err != nil { + return Vault{}, fmt.Errorf("%w: %w", ErrInvalidInput, err) + } + return s.storage.CreateVault(ctx, NewVault{TenantID: command.TenantID, Name: command.Name, Metadata: encoded}) +} + +type DeleteVault struct{ TenantID, VaultID string } + +// DeleteVault deletes a Vault with its Credentials. Sessions keep their frozen +// bindings, which then resolve to ErrNotFound. +func (s *Service) DeleteVault(ctx context.Context, command DeleteVault) (string, error) { + return s.storage.DeleteVault(ctx, command.TenantID, command.VaultID) +} + +type CreateStaticCredential struct { + TenantID, VaultID string + Name, MCPServerURL, Token string +} + +// CreateStaticCredential seals a bearer token to its tenant, Vault, new +// Credential ID and destination, and stores it. +func (s *Service) CreateStaticCredential(ctx context.Context, command CreateStaticCredential) (Credential, error) { + tenant, ok := canonicalID(command.TenantID) + if !ok { + return Credential{}, ErrInvalidInput + } + if !validName(command.Name) || command.MCPServerURL == "" { + return Credential{}, ErrInvalidInput + } + if s.cipher == nil { + return Credential{}, credentialcrypto.ErrUnavailable + } + // A malformed Vault ID follows the missing-Vault path, after validation. + vault, ok := canonicalID(command.VaultID) + if !ok { + return Credential{}, ErrNotFound + } + key := CredentialKey{TenantID: tenant, VaultID: vault, CredentialID: uuid.NewString()} + ciphertext, err := s.cipher.Seal([]byte(command.Token), credentialcrypto.Binding{TenantID: tenant, VaultID: vault, + CredentialID: key.CredentialID, AuthType: AuthStaticBearer, Destination: command.MCPServerURL}) + if err != nil { + return Credential{}, errors.New("credential encryption failed") + } + return s.storage.CreateCredential(ctx, NewCredential{CredentialKey: key, Name: command.Name, + AuthType: AuthStaticBearer, MCPServerURL: command.MCPServerURL, Ciphertext: ciphertext}) +} + +type UpdateStaticCredential struct { + TenantID, VaultID, CredentialID string + Token string +} + +// UpdateStaticCredential replaces only the token and the update time. The +// stored metadata supplies the seal's scope, which the write checks again, so +// the existing frozen bindings read the replacement. +func (s *Service) UpdateStaticCredential(ctx context.Context, command UpdateStaticCredential) (Credential, error) { + key, ok := credentialKey(command.TenantID, command.VaultID, command.CredentialID) + if !ok { + return Credential{}, ErrNotFound + } + current, err := s.storage.GetCredential(ctx, key.TenantID, key.VaultID, key.CredentialID) + if err != nil { + return Credential{}, err + } + if current.AuthType != AuthStaticBearer { + return Credential{}, ErrInvalidInput + } + if s.cipher == nil { + return Credential{}, credentialcrypto.ErrUnavailable + } + ciphertext, err := s.cipher.Seal([]byte(command.Token), credentialcrypto.Binding{TenantID: key.TenantID, VaultID: current.VaultID, + CredentialID: current.ID, AuthType: current.AuthType, Destination: current.MCPServerURL}) + if err != nil { + return Credential{}, errors.New("credential encryption failed") + } + return s.storage.ReplaceStaticToken(ctx, StaticTokenReplacement{CredentialKey: key, MCPServerURL: current.MCPServerURL, Ciphertext: ciphertext}) +} + +// CreateOAuthCredential keeps the write-only secrets apart from the metadata. +type CreateOAuthCredential struct { + TenantID, VaultID string + Name, MCPServerURL, AccessToken string + OAuth OAuthMetadata + RefreshToken, ClientSecret string +} + +// CreateOAuthCredential seals a grant, with the metadata it is refreshed by, +// to its tenant, Vault, new Credential ID and destination, and stores it. +func (s *Service) CreateOAuthCredential(ctx context.Context, command CreateOAuthCredential) (Credential, error) { + tenant, ok := canonicalID(command.TenantID) + if !ok { + return Credential{}, ErrNotFound + } + if !validOAuthCreation(command) { + return Credential{}, ErrInvalidInput + } + if s.cipher == nil { + return Credential{}, credentialcrypto.ErrUnavailable + } + // A malformed Vault ID follows the missing-Vault path, after validation. + vault, ok := canonicalID(command.VaultID) + if !ok { + return Credential{}, ErrNotFound + } + credential := Credential{ID: uuid.NewString(), VaultID: vault, Name: command.Name, AuthType: AuthMCPOAuth, MCPServerURL: command.MCPServerURL} + sealed, err := s.sealOAuth(tenant, credential, oauthSecret{Version: 1, Metadata: command.OAuth, AccessToken: command.AccessToken, + RefreshToken: command.RefreshToken, ClientSecret: command.ClientSecret}) + if err != nil { + return Credential{}, err + } + return s.storage.CreateCredential(ctx, NewCredential{CredentialKey: CredentialKey{TenantID: tenant, VaultID: vault, CredentialID: credential.ID}, + Name: command.Name, AuthType: AuthMCPOAuth, MCPServerURL: command.MCPServerURL, OAuthMetadata: sealed.Metadata, Ciphertext: sealed.Ciphertext}) +} + +// UpdateOAuthCredential patches an OAuth grant. ExpiresAtSet keeps omitted +// apart from null; a nil AccessToken keeps the stored one. +type UpdateOAuthCredential struct { + TenantID, VaultID, CredentialID string + AccessToken *string + ExpiresAt *string + ExpiresAtSet bool + Refresh *OAuthRefreshUpdate +} + +// UpdateOAuthCredential applies the patch to the locked grant, so it +// serializes with refreshes. +func (s *Service) UpdateOAuthCredential(ctx context.Context, command UpdateOAuthCredential) (Credential, error) { + key, ok := credentialKey(command.TenantID, command.VaultID, command.CredentialID) + if !ok { + return Credential{}, ErrNotFound + } + current, err := s.storage.GetCredential(ctx, key.TenantID, key.VaultID, key.CredentialID) + if err != nil { + return Credential{}, err + } + if current.AuthType != AuthMCPOAuth { + return Credential{}, ErrInvalidInput + } + var updated Credential + err = s.withOAuth(ctx, key, "", "credential update failed", func(tx OAuthTx, credential Credential, secret oauthSecret) error { + secret, err := applyOAuthUpdate(secret, command) + if err != nil { + return err + } + sealed, err := s.sealOAuth(key.TenantID, credential, secret) + if err != nil { + return err + } + updated, err = tx.ApplyOAuthReplacement(ctx, sealed) + return err + }) + if err != nil { + return Credential{}, err + } + return updated, nil +} + +type DeleteCredential struct{ TenantID, VaultID, CredentialID string } + +// DeleteCredential removes the Credential without reading or opening its secret. +func (s *Service) DeleteCredential(ctx context.Context, command DeleteCredential) (string, error) { + return s.storage.DeleteCredential(ctx, CredentialKey{TenantID: command.TenantID, VaultID: command.VaultID, CredentialID: command.CredentialID}) +} + +type ResolveMCPCredentials struct { + TenantID string + VaultIDs []string + Requests []MCPCredentialRequest +} + +// ResolveMCPCredentials selects each MCP server's Credential from the attached +// Vaults, reading metadata only. Later resource changes do not reselect +// Credentials for an accepted Session or its creation retries. +func (s *Service) ResolveMCPCredentials(ctx context.Context, command ResolveMCPCredentials) ([]MCPCredentialBinding, error) { + tenant, ok := canonicalID(command.TenantID) + if !ok { + return nil, ErrInvalidInput + } + attached, err := attachedVaultIDs(command.VaultIDs) + if err != nil { + return nil, err + } + owned, err := s.storage.CountOwnedVaults(ctx, tenant, attached) + if err != nil { + return nil, err + } + if owned != len(attached) { + return nil, ErrNotFound + } + bindings := make([]MCPCredentialBinding, 0, len(command.Requests)) + for _, request := range command.Requests { + id, err := mcpCredentialLookup(request, attached) + if err != nil { + return nil, err + } + matches, err := s.storage.FindMCPCredentials(ctx, MCPCredentialQuery{TenantID: tenant, VaultIDs: attached, ServerURL: request.ServerURL, CredentialID: id}) + if err != nil { + return nil, err + } + binding, err := selectMCPCredential(request, matches) + if err != nil { + return nil, err + } + bindings = append(bindings, binding) + } + return bindings, nil +} + +// MCPBearerToken asks for the bearer token of a Session's frozen binding. +type MCPBearerToken struct { + TenantID string + VaultIDs []string + Binding MCPCredentialBinding +} + +// MCPBearerToken is for execution only. It rechecks the complete frozen +// authorization before opening a secret, and refreshes an expired OAuth grant +// while holding its lock. Never persist or log the token, or downgrade a +// failure to an anonymous request. +func (s *Service) MCPBearerToken(ctx context.Context, command MCPBearerToken) (string, error) { + scope, err := bearerTokenScope(command) + if err != nil { + return "", err + } + if command.Binding.AuthType == AuthMCPOAuth { + if !scope.vaultAttached() { + return "", ErrNotFound + } + return s.oauthBearerToken(ctx, CredentialKey{TenantID: scope.tenantID, VaultID: scope.vaultID, CredentialID: scope.credentialID}, command.Binding.ServerURL) + } + ciphertext, err := s.storage.StaticTokenCiphertext(ctx, StaticTokenQuery{TenantID: scope.tenantID, VaultIDs: scope.attached, + VaultID: scope.vaultID, CredentialID: scope.credentialID, MCPServerURL: command.Binding.ServerURL}) + if err != nil { + return "", err + } + if s.cipher == nil { + return "", credentialcrypto.ErrUnavailable + } + plaintext, err := s.cipher.Open(ciphertext, credentialcrypto.Binding{TenantID: scope.tenantID, VaultID: scope.vaultID, + CredentialID: scope.credentialID, AuthType: AuthStaticBearer, Destination: command.Binding.ServerURL}) + if err != nil { + return "", errors.New("MCP credential decryption failed") + } + return string(plaintext), nil +} + +func (s *Service) oauthBearerToken(ctx context.Context, key CredentialKey, destination string) (string, error) { + ctx, cancel := context.WithTimeout(ctx, oauthRefreshTimeout) + defer cancel() + var bearer string + err := s.withOAuth(ctx, key, destination, "OAuth credential refresh commit failed", func(tx OAuthTx, credential Credential, secret oauthSecret) error { + token, expired, err := currentAccessToken(secret, time.Now()) + if err != nil || !expired { + bearer = token + return err + } + request, err := refreshRequest(secret) + if err != nil { + return err + } + refreshed, err := s.refresher.Refresh(ctx, request) + if err != nil { + return errors.New("OAuth credential refresh failed") + } + secret, err = applyRefreshedToken(secret, refreshed, time.Now()) + if err != nil { + return err + } + sealed, err := s.sealOAuth(key.TenantID, credential, secret) + if err != nil { + return err + } + if err := tx.ApplyOAuthRefresh(ctx, sealed); err != nil { + return err + } + bearer = secret.AccessToken + return nil + }) + if err != nil { + return "", err + } + return bearer, nil +} + +// withOAuth locks the Credential for the whole of apply, including an +// external refresh, and commits only when apply succeeds. PostgreSQL +// serializes competing updates and deletions, including the parent Vault's +// cascade. A non-empty destination must match the stored one. A failure to +// begin or commit is reported as failure, never with database error text. +func (s *Service) withOAuth(ctx context.Context, key CredentialKey, destination, failure string, apply func(OAuthTx, Credential, oauthSecret) error) error { + var applied error + err := s.storage.WithOAuthCredential(ctx, key, func(tx OAuthTx) error { + credential, ciphertext, err := tx.LoadOAuthCredential(ctx) + if err == nil && destination != "" && credential.MCPServerURL != destination { + err = ErrNotFound + } + var secret oauthSecret + if err == nil { + secret, err = s.openOAuth(key.TenantID, credential, ciphertext) + } + if err == nil { + err = apply(tx, credential, secret) + } + applied = err + return err + }) + if err != nil && applied == nil { + return errors.New(failure) + } + return err +} + +func (s *Service) sealOAuth(tenantID string, credential Credential, secret oauthSecret) (SealedOAuth, error) { + if s.cipher == nil { + return SealedOAuth{}, credentialcrypto.ErrUnavailable + } + if !validOAuthMetadata(secret.Metadata) { + return SealedOAuth{}, ErrInvalidInput + } + encoded, err := json.Marshal(secret.Metadata) + if err != nil { + return SealedOAuth{}, errors.New("credential encoding failed") + } + plaintext, err := json.Marshal(secret) + if err != nil { + return SealedOAuth{}, errors.New("credential encoding failed") + } + ciphertext, err := s.cipher.Seal(plaintext, oauthBinding(tenantID, credential)) + if err != nil { + return SealedOAuth{}, errors.New("credential encryption failed") + } + return SealedOAuth{MCPServerURL: credential.MCPServerURL, Metadata: encoded, Ciphertext: ciphertext}, nil +} + +// openOAuth authenticates the stored metadata against the sealed copy, so a +// changed stored setting never reaches a provider. +func (s *Service) openOAuth(tenantID string, credential Credential, ciphertext []byte) (oauthSecret, error) { + if s.cipher == nil { + return oauthSecret{}, credentialcrypto.ErrUnavailable + } + plaintext, err := s.cipher.Open(ciphertext, oauthBinding(tenantID, credential)) + if err != nil { + return oauthSecret{}, errors.New("OAuth credential decryption failed") + } + var secret oauthSecret + if json.Unmarshal(plaintext, &secret) != nil || secret.Version != 1 || credential.OAuth == nil || + !reflect.DeepEqual(secret.Metadata, *credential.OAuth) || !validOAuthMetadata(secret.Metadata) { + return oauthSecret{}, errors.New("OAuth credential authentication failed") + } + return secret, nil +} + +// credentialKey canonicalizes a Credential's IDs. +func credentialKey(tenantID, vaultID, credentialID string) (CredentialKey, bool) { + tenant, ok1 := canonicalID(tenantID) + vault, ok2 := canonicalID(vaultID) + id, ok3 := canonicalID(credentialID) + return CredentialKey{TenantID: tenant, VaultID: vault, CredentialID: id}, ok1 && ok2 && ok3 +} diff --git a/services/core/internal/vaults/service_test.go b/services/core/internal/vaults/service_test.go new file mode 100644 index 000000000..6f67e77a1 --- /dev/null +++ b/services/core/internal/vaults/service_test.go @@ -0,0 +1,521 @@ +package vaults + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "reflect" + "strings" + "testing" + "time" + + "github.com/google/uuid" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" +) + +// fakeStorage fails the test on any call whose func the test did not set. +type fakeStorage struct { + t testing.TB + getVault func(context.Context, string, string) (Vault, error) + listVaults func(context.Context, string, PageQuery) (VaultPage, error) + getCredential func(context.Context, string, string, string) (Credential, error) + listCredentials func(context.Context, string, string, PageQuery) (CredentialPage, error) + createVault func(context.Context, NewVault) (Vault, error) + deleteVault func(context.Context, string, string) (string, error) + createCredential func(context.Context, NewCredential) (Credential, error) + replaceStaticToken func(context.Context, StaticTokenReplacement) (Credential, error) + deleteCredential func(context.Context, CredentialKey) (string, error) + withOAuthCredential func(context.Context, CredentialKey, func(OAuthTx) error) error + countOwnedVaults func(context.Context, string, []string) (int, error) + findMCPCredentials func(context.Context, MCPCredentialQuery) ([]MCPCredentialMatch, error) + staticTokenCiphertxt func(context.Context, StaticTokenQuery) ([]byte, error) +} + +func (f *fakeStorage) unexpected(method string) { + f.t.Helper() + f.t.Fatalf("unexpected call to %s", method) +} + +func (f *fakeStorage) GetVault(ctx context.Context, tenantID, vaultID string) (Vault, error) { + if f.getVault == nil { + f.unexpected("GetVault") + } + return f.getVault(ctx, tenantID, vaultID) +} + +func (f *fakeStorage) ListVaults(ctx context.Context, tenantID string, query PageQuery) (VaultPage, error) { + if f.listVaults == nil { + f.unexpected("ListVaults") + } + return f.listVaults(ctx, tenantID, query) +} + +func (f *fakeStorage) GetCredential(ctx context.Context, tenantID, vaultID, credentialID string) (Credential, error) { + if f.getCredential == nil { + f.unexpected("GetCredential") + } + return f.getCredential(ctx, tenantID, vaultID, credentialID) +} + +func (f *fakeStorage) ListCredentials(ctx context.Context, tenantID, vaultID string, query PageQuery) (CredentialPage, error) { + if f.listCredentials == nil { + f.unexpected("ListCredentials") + } + return f.listCredentials(ctx, tenantID, vaultID, query) +} + +func (f *fakeStorage) CreateVault(ctx context.Context, vault NewVault) (Vault, error) { + if f.createVault == nil { + f.unexpected("CreateVault") + } + return f.createVault(ctx, vault) +} + +func (f *fakeStorage) DeleteVault(ctx context.Context, tenantID, vaultID string) (string, error) { + if f.deleteVault == nil { + f.unexpected("DeleteVault") + } + return f.deleteVault(ctx, tenantID, vaultID) +} + +func (f *fakeStorage) CreateCredential(ctx context.Context, credential NewCredential) (Credential, error) { + if f.createCredential == nil { + f.unexpected("CreateCredential") + } + return f.createCredential(ctx, credential) +} + +func (f *fakeStorage) ReplaceStaticToken(ctx context.Context, replacement StaticTokenReplacement) (Credential, error) { + if f.replaceStaticToken == nil { + f.unexpected("ReplaceStaticToken") + } + return f.replaceStaticToken(ctx, replacement) +} + +func (f *fakeStorage) DeleteCredential(ctx context.Context, key CredentialKey) (string, error) { + if f.deleteCredential == nil { + f.unexpected("DeleteCredential") + } + return f.deleteCredential(ctx, key) +} + +func (f *fakeStorage) WithOAuthCredential(ctx context.Context, key CredentialKey, apply func(OAuthTx) error) error { + if f.withOAuthCredential == nil { + f.unexpected("WithOAuthCredential") + } + return f.withOAuthCredential(ctx, key, apply) +} + +func (f *fakeStorage) CountOwnedVaults(ctx context.Context, tenantID string, vaultIDs []string) (int, error) { + if f.countOwnedVaults == nil { + f.unexpected("CountOwnedVaults") + } + return f.countOwnedVaults(ctx, tenantID, vaultIDs) +} + +func (f *fakeStorage) FindMCPCredentials(ctx context.Context, query MCPCredentialQuery) ([]MCPCredentialMatch, error) { + if f.findMCPCredentials == nil { + f.unexpected("FindMCPCredentials") + } + return f.findMCPCredentials(ctx, query) +} + +func (f *fakeStorage) StaticTokenCiphertext(ctx context.Context, query StaticTokenQuery) ([]byte, error) { + if f.staticTokenCiphertxt == nil { + f.unexpected("StaticTokenCiphertext") + } + return f.staticTokenCiphertxt(ctx, query) +} + +// fakeOAuthTx fails the test on any call whose func the test did not set. +type fakeOAuthTx struct { + t testing.TB + load func(context.Context) (Credential, []byte, error) + refresh func(context.Context, SealedOAuth) error + replacement func(context.Context, SealedOAuth) (Credential, error) +} + +func (f *fakeOAuthTx) LoadOAuthCredential(ctx context.Context) (Credential, []byte, error) { + if f.load == nil { + f.t.Fatal("unexpected call to LoadOAuthCredential") + } + return f.load(ctx) +} + +func (f *fakeOAuthTx) ApplyOAuthRefresh(ctx context.Context, sealed SealedOAuth) error { + if f.refresh == nil { + f.t.Fatal("unexpected call to ApplyOAuthRefresh") + } + return f.refresh(ctx, sealed) +} + +func (f *fakeOAuthTx) ApplyOAuthReplacement(ctx context.Context, sealed SealedOAuth) (Credential, error) { + if f.replacement == nil { + f.t.Fatal("unexpected call to ApplyOAuthReplacement") + } + return f.replacement(ctx, sealed) +} + +type refreshFunc func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) + +func (f refreshFunc) Refresh(ctx context.Context, request oauthrefresh.Request) (oauthrefresh.Token, error) { + return f(ctx, request) +} + +func unexpectedRefresh(t testing.TB) refreshFunc { + return func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + t.Fatal("unexpected call to Refresh") + return oauthrefresh.Token{}, nil + } +} + +func testCipher(t testing.TB) *credentialcrypto.Cipher { + t.Helper() + cipher, err := credentialcrypto.New(bytes.Repeat([]byte{7}, 32)) + if err != nil { + t.Fatal(err) + } + return cipher +} + +func testService(t testing.TB, storage *fakeStorage, cipher *credentialcrypto.Cipher, refresher oauthrefresh.Refresher) *Service { + t.Helper() + storage.t = t + service, err := NewService(storage, cipher, refresher) + if err != nil { + t.Fatal(err) + } + return service +} + +func TestNewServiceRequiresStorageAndRefresher(t *testing.T) { + if _, err := NewService(nil, nil, unexpectedRefresh(t)); err == nil { + t.Fatal("missing storage accepted") + } + if _, err := NewService(&fakeStorage{t: t}, nil, nil); err == nil { + t.Fatal("missing refresher accepted") + } + if _, err := NewService(&fakeStorage{t: t}, nil, unexpectedRefresh(t)); err != nil { + t.Fatal("a keyless service was rejected", err) + } +} + +func TestCreateVaultValidatesBeforeStorage(t *testing.T) { + service := testService(t, &fakeStorage{}, nil, unexpectedRefresh(t)) + long := strings.Repeat("x", 257) + for _, command := range []CreateVault{ + {TenantID: uuid.NewString(), Name: &long}, + {TenantID: uuid.NewString(), Metadata: map[string]string{"large": strings.Repeat("x", 64*1024)}}, + } { + if _, err := service.CreateVault(t.Context(), command); !errors.Is(err, ErrInvalidInput) { + t.Fatal("invalid Vault reached storage", err) + } + } + tenant, name := uuid.NewString(), "named" + stored := Vault{ID: uuid.NewString(), TenantID: tenant} + service = testService(t, &fakeStorage{createVault: func(_ context.Context, vault NewVault) (Vault, error) { + if vault.TenantID != tenant || vault.Name == nil || *vault.Name != name || string(vault.Metadata) != "{}" { + t.Fatalf("unexpected new Vault %+v", vault) + } + return stored, nil + }}, nil, unexpectedRefresh(t)) + if got, err := service.CreateVault(t.Context(), CreateVault{TenantID: tenant, Name: &name}); err != nil || !reflect.DeepEqual(got, stored) { + t.Fatal("Vault creation did not return the stored Vault", err) + } +} + +// Credential creation checks the body, then the credential key, then the +// Vault ID, and never reaches storage when one fails. +func TestCredentialCreationValidationOrder(t *testing.T) { + tenant, url := uuid.NewString(), "https://mcp.example/tools" + create := map[string]func(*Service, string, string) error{ + "static": func(s *Service, name, vault string) error { + _, err := s.CreateStaticCredential(t.Context(), CreateStaticCredential{TenantID: tenant, VaultID: vault, Name: name, MCPServerURL: url, Token: "token"}) + return err + }, + "oauth": func(s *Service, name, vault string) error { + _, err := s.CreateOAuthCredential(t.Context(), CreateOAuthCredential{TenantID: tenant, VaultID: vault, Name: name, MCPServerURL: url, AccessToken: "token"}) + return err + }, + } + for kind, run := range create { + for _, tc := range []struct { + name string + cipher *credentialcrypto.Cipher + want error + }{ + {"", nil, ErrInvalidInput}, + {"valid", nil, credentialcrypto.ErrUnavailable}, + {"valid", testCipher(t), ErrNotFound}, + } { + service := testService(t, &fakeStorage{}, tc.cipher, unexpectedRefresh(t)) + if err := run(service, tc.name, "not-a-vault"); !errors.Is(err, tc.want) { + t.Fatalf("%s creation: got %v, want %v", kind, err, tc.want) + } + } + } +} + +func TestCreateStaticCredentialSealsToCanonicalIDs(t *testing.T) { + tenant, vault, url := uuid.New(), uuid.New(), "https://mcp.example/tools" + cipher := testCipher(t) + var stored NewCredential + service := testService(t, &fakeStorage{createCredential: func(_ context.Context, credential NewCredential) (Credential, error) { + stored = credential + return Credential{ID: credential.CredentialID}, nil + }}, cipher, unexpectedRefresh(t)) + created, err := service.CreateStaticCredential(t.Context(), CreateStaticCredential{TenantID: strings.ToUpper(tenant.String()), VaultID: strings.ToUpper(vault.String()), Name: "static", MCPServerURL: url, Token: "private-token"}) + if err != nil || created.ID != stored.CredentialID { + t.Fatal(err) + } + if stored.TenantID != tenant.String() || stored.VaultID != vault.String() || stored.AuthType != AuthStaticBearer || stored.OAuthMetadata != nil || bytes.Contains(stored.Ciphertext, []byte("private-token")) { + t.Fatalf("unexpected stored Credential %+v", stored) + } + plaintext, err := cipher.Open(stored.Ciphertext, credentialcrypto.Binding{TenantID: tenant.String(), VaultID: vault.String(), CredentialID: stored.CredentialID, AuthType: AuthStaticBearer, Destination: url}) + if err != nil || string(plaintext) != "private-token" { + t.Fatal("token was not sealed to its canonical scope", err) + } +} + +func TestUpdateStaticCredential(t *testing.T) { + tenant, vault, id, url := uuid.NewString(), uuid.NewString(), uuid.NewString(), "https://mcp.example/tools" + command := UpdateStaticCredential{TenantID: tenant, VaultID: vault, CredentialID: id, Token: "replacement"} + for _, malformed := range []UpdateStaticCredential{{TenantID: "x", VaultID: vault, CredentialID: id}, {TenantID: tenant, VaultID: "x", CredentialID: id}, {TenantID: tenant, VaultID: vault, CredentialID: "x"}} { + if _, err := testService(t, &fakeStorage{}, testCipher(t), unexpectedRefresh(t)).UpdateStaticCredential(t.Context(), malformed); !errors.Is(err, ErrNotFound) { + t.Fatal("a malformed ID named a Credential", err) + } + } + current := func(authType string) func(context.Context, string, string, string) (Credential, error) { + return func(context.Context, string, string, string) (Credential, error) { + return Credential{ID: id, VaultID: vault, AuthType: authType, MCPServerURL: url}, nil + } + } + if _, err := testService(t, &fakeStorage{getCredential: current(AuthMCPOAuth)}, testCipher(t), unexpectedRefresh(t)).UpdateStaticCredential(t.Context(), command); !errors.Is(err, ErrInvalidInput) { + t.Fatal("an OAuth Credential took a static token", err) + } + if _, err := testService(t, &fakeStorage{getCredential: current(AuthStaticBearer)}, nil, unexpectedRefresh(t)).UpdateStaticCredential(t.Context(), command); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("a keyless service replaced a token", err) + } + cipher := testCipher(t) + service := testService(t, &fakeStorage{getCredential: current(AuthStaticBearer), replaceStaticToken: func(_ context.Context, replacement StaticTokenReplacement) (Credential, error) { + plaintext, err := cipher.Open(replacement.Ciphertext, credentialcrypto.Binding{TenantID: tenant, VaultID: vault, CredentialID: id, AuthType: AuthStaticBearer, Destination: url}) + if replacement.CredentialKey != (CredentialKey{tenant, vault, id}) || replacement.MCPServerURL != url || err != nil || string(plaintext) != "replacement" { + t.Fatalf("unexpected replacement %+v: %v", replacement, err) + } + return Credential{ID: id}, nil + }}, cipher, unexpectedRefresh(t)) + if _, err := service.UpdateStaticCredential(t.Context(), command); err != nil { + t.Fatal(err) + } +} + +func TestResolveMCPCredentials(t *testing.T) { + tenant, vault, id, url := uuid.NewString(), uuid.NewString(), uuid.NewString(), "https://mcp.example/tools" + command := ResolveMCPCredentials{TenantID: tenant, VaultIDs: []string{vault, vault}, Requests: []MCPCredentialRequest{ + {ServerLabel: "tools", ServerURL: url}, {ServerLabel: "anonymous", ServerURL: "https://anonymous.example/mcp"}, + }} + owned := func(count int) func(context.Context, string, []string) (int, error) { + return func(_ context.Context, gotTenant string, ids []string) (int, error) { + if gotTenant != tenant || !reflect.DeepEqual(ids, []string{vault}) { + t.Fatalf("unexpected ownership check %s %v", gotTenant, ids) + } + return count, nil + } + } + if _, err := testService(t, &fakeStorage{countOwnedVaults: owned(0)}, nil, unexpectedRefresh(t)).ResolveMCPCredentials(t.Context(), command); !errors.Is(err, ErrNotFound) { + t.Fatal("an unowned attached Vault was accepted", err) + } + service := testService(t, &fakeStorage{countOwnedVaults: owned(1), findMCPCredentials: func(_ context.Context, query MCPCredentialQuery) ([]MCPCredentialMatch, error) { + if query.TenantID != tenant || !reflect.DeepEqual(query.VaultIDs, []string{vault}) || query.CredentialID != "" { + t.Fatalf("unexpected lookup %+v", query) + } + if query.ServerURL != url { + return nil, nil + } + return []MCPCredentialMatch{{VaultID: vault, CredentialID: id, AuthType: AuthStaticBearer, MCPServerURL: url}}, nil + }}, nil, unexpectedRefresh(t)) + bindings, err := service.ResolveMCPCredentials(t.Context(), command) + want := []MCPCredentialBinding{{ServerLabel: "tools", ServerURL: url, VaultID: vault, CredentialID: id, AuthType: AuthStaticBearer}, {ServerLabel: "anonymous", ServerURL: "https://anonymous.example/mcp"}} + if err != nil || !reflect.DeepEqual(bindings, want) { + t.Fatal("selection or the frozen anonymous decision differs", bindings, err) + } +} + +func TestStaticBearerToken(t *testing.T) { + tenant, vault, id, url := uuid.NewString(), uuid.NewString(), uuid.NewString(), "https://mcp.example/tools" + cipher := testCipher(t) + sealed, err := cipher.Seal([]byte("private-token"), credentialcrypto.Binding{TenantID: tenant, VaultID: vault, CredentialID: id, AuthType: AuthStaticBearer, Destination: url}) + if err != nil { + t.Fatal(err) + } + command := MCPBearerToken{TenantID: tenant, VaultIDs: []string{vault}, Binding: MCPCredentialBinding{ServerLabel: "tools", ServerURL: url, VaultID: vault, CredentialID: id, AuthType: AuthStaticBearer}} + lookup := func(ciphertext []byte) func(context.Context, StaticTokenQuery) ([]byte, error) { + return func(_ context.Context, query StaticTokenQuery) ([]byte, error) { + if !reflect.DeepEqual(query, StaticTokenQuery{TenantID: tenant, VaultIDs: []string{vault}, VaultID: vault, CredentialID: id, MCPServerURL: url}) { + t.Fatalf("unexpected scope %+v", query) + } + return ciphertext, nil + } + } + if token, err := testService(t, &fakeStorage{staticTokenCiphertxt: lookup(sealed)}, cipher, unexpectedRefresh(t)).MCPBearerToken(t.Context(), command); err != nil || token != "private-token" { + t.Fatal("static token was not opened", err) + } + if token, err := testService(t, &fakeStorage{staticTokenCiphertxt: lookup(sealed)}, nil, unexpectedRefresh(t)).MCPBearerToken(t.Context(), command); !errors.Is(err, credentialcrypto.ErrUnavailable) || token != "" { + t.Fatal("a keyless service opened a token", err) + } + other := command + other.Binding.ServerURL += "/other" + if token, err := testService(t, &fakeStorage{staticTokenCiphertxt: func(context.Context, StaticTokenQuery) ([]byte, error) { return sealed, nil }}, cipher, unexpectedRefresh(t)).MCPBearerToken(t.Context(), other); err == nil || token != "" || strings.Contains(err.Error(), "private-token") { + t.Fatal("a token opened outside its sealed destination", err) + } +} + +// oauthScenario is one mcp_oauth Credential sealed by service, served by a +// fake transaction. +type oauthScenario struct { + service *Service + credential Credential + command MCPBearerToken + ciphertext []byte + refreshes []SealedOAuth + committed error +} + +func newOAuthScenario(t *testing.T, expiresAt time.Time, refresher oauthrefresh.Refresher) *oauthScenario { + t.Helper() + tenant, url := uuid.NewString(), "https://mcp.example/tools" + expiry := expiresAt.UTC().Format(time.RFC3339Nano) + metadata := OAuthMetadata{ExpiresAt: &expiry, Refresh: &OAuthRefreshMetadata{ClientID: "client", TokenEndpoint: "https://issuer.example/token", TokenEndpointAuth: "client_secret_basic"}} + scenario := &oauthScenario{credential: Credential{ID: uuid.NewString(), VaultID: uuid.NewString(), Name: "OAuth", AuthType: AuthMCPOAuth, MCPServerURL: url, OAuth: &metadata}} + storage := &fakeStorage{withOAuthCredential: func(ctx context.Context, key CredentialKey, apply func(OAuthTx) error) error { + if key != (CredentialKey{tenant, scenario.credential.VaultID, scenario.credential.ID}) { + t.Fatalf("unexpected key %+v", key) + } + err := apply(&fakeOAuthTx{t: t, + load: func(context.Context) (Credential, []byte, error) { + return scenario.credential, scenario.ciphertext, nil + }, + refresh: func(_ context.Context, sealed SealedOAuth) error { + scenario.refreshes = append(scenario.refreshes, sealed) + return nil + }, + }) + if err != nil { + return err + } + return scenario.committed + }} + scenario.service = testService(t, storage, testCipher(t), refresher) + sealed, err := scenario.service.sealOAuth(tenant, scenario.credential, oauthSecret{Version: 1, Metadata: metadata, AccessToken: "stored-access", RefreshToken: "stored-refresh", ClientSecret: "private-client"}) + if err != nil { + t.Fatal(err) + } + scenario.ciphertext = sealed.Ciphertext + scenario.command = MCPBearerToken{TenantID: tenant, VaultIDs: []string{scenario.credential.VaultID}, Binding: MCPCredentialBinding{ServerLabel: "tools", ServerURL: url, + VaultID: scenario.credential.VaultID, CredentialID: scenario.credential.ID, AuthType: AuthMCPOAuth}} + return scenario +} + +func TestOAuthBearerTokenRefreshesExpiredGrantUnderLock(t *testing.T) { + later := time.Now().Add(time.Hour) + var requests []oauthrefresh.Request + scenario := newOAuthScenario(t, time.Now().Add(-time.Minute), refreshFunc(func(_ context.Context, request oauthrefresh.Request) (oauthrefresh.Token, error) { + requests = append(requests, request) + return oauthrefresh.Token{AccessToken: "renewed-access", ExpiresAt: &later}, nil + })) + token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command) + if err != nil || token != "renewed-access" || len(requests) != 1 || requests[0].RefreshToken != "stored-refresh" || requests[0].ClientSecret != "private-client" || len(scenario.refreshes) != 1 { + t.Fatal("expired grant was not refreshed once", token, err) + } + sealed := scenario.refreshes[0] + var metadata OAuthMetadata + if sealed.MCPServerURL != scenario.credential.MCPServerURL || json.Unmarshal(sealed.Metadata, &metadata) != nil || *metadata.ExpiresAt != later.UTC().Format(time.RFC3339Nano) { + t.Fatal("refreshed metadata was not stored with the grant", string(sealed.Metadata)) + } + scenario.credential.OAuth, scenario.ciphertext = &metadata, sealed.Ciphertext + if token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command); err != nil || token != "renewed-access" || len(requests) != 1 { + t.Fatal("the refreshed grant was not used", token, err) + } +} + +func TestOAuthBearerTokenFailuresReturnNoToken(t *testing.T) { + past := time.Now().Add(-time.Minute) + t.Run("fresh grant", func(t *testing.T) { + scenario := newOAuthScenario(t, time.Now().Add(time.Hour), unexpectedRefresh(t)) + if token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command); err != nil || token != "stored-access" { + t.Fatal("a fresh grant was not returned as stored", err) + } + }) + t.Run("unattached vault", func(t *testing.T) { + scenario := newOAuthScenario(t, past, unexpectedRefresh(t)) + scenario.command.VaultIDs = []string{uuid.NewString()} + scenario.service.storage = &fakeStorage{t: t} + if token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command); !errors.Is(err, ErrNotFound) || token != "" { + t.Fatal("an unattached Vault's grant was opened", err) + } + }) + t.Run("changed destination", func(t *testing.T) { + scenario := newOAuthScenario(t, past, unexpectedRefresh(t)) + scenario.command.Binding.ServerURL += "/other" + if token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command); !errors.Is(err, ErrNotFound) || token != "" { + t.Fatal("a grant was used for another destination", err) + } + }) + t.Run("substituted metadata", func(t *testing.T) { + scenario := newOAuthScenario(t, past, unexpectedRefresh(t)) + scenario.credential.OAuth = &OAuthMetadata{ExpiresAt: scenario.credential.OAuth.ExpiresAt, Refresh: &OAuthRefreshMetadata{ClientID: "client", TokenEndpoint: "https://attacker.example/token", TokenEndpointAuth: "client_secret_basic"}} + if token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command); err == nil || token != "" { + t.Fatal("substituted stored metadata was authenticated", err) + } + }) + t.Run("provider error", func(t *testing.T) { + scenario := newOAuthScenario(t, past, refreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + return oauthrefresh.Token{}, errors.New("stored-refresh: provider body") + })) + if token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command); err == nil || token != "" || strings.Contains(err.Error(), "stored-refresh") || len(scenario.refreshes) != 0 { + t.Fatal("a failed refresh was unsafe or stored", err) + } + }) + t.Run("commit failure", func(t *testing.T) { + scenario := newOAuthScenario(t, past, refreshFunc(func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) { + return oauthrefresh.Token{AccessToken: "uncommitted-access"}, nil + })) + scenario.committed = errors.New("private database text") + token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command) + if err == nil || err.Error() != "OAuth credential refresh commit failed" || token != "" { + t.Fatal("an uncommitted grant was returned", token, err) + } + }) +} + +func TestUpdateOAuthCredentialReplacesLockedGrant(t *testing.T) { + scenario := newOAuthScenario(t, time.Now().Add(-time.Minute), unexpectedRefresh(t)) + storage := scenario.service.storage.(*fakeStorage) + storage.getCredential = func(context.Context, string, string, string) (Credential, error) { return scenario.credential, nil } + lock := storage.withOAuthCredential + var replaced SealedOAuth + storage.withOAuthCredential = func(ctx context.Context, key CredentialKey, apply func(OAuthTx) error) error { + return lock(ctx, key, func(tx OAuthTx) error { + return apply(&fakeOAuthTx{t: t, load: tx.LoadOAuthCredential, replacement: func(_ context.Context, sealed SealedOAuth) (Credential, error) { + replaced = sealed + return scenario.credential, nil + }}) + }) + } + command := UpdateOAuthCredential{TenantID: scenario.command.TenantID, VaultID: scenario.credential.VaultID, CredentialID: scenario.credential.ID, AccessToken: ptr("manual-access")} + if _, err := scenario.service.UpdateOAuthCredential(t.Context(), command); err != nil { + t.Fatal(err) + } + secret, err := scenario.service.openOAuth(command.TenantID, Credential{ID: scenario.credential.ID, VaultID: scenario.credential.VaultID, MCPServerURL: scenario.credential.MCPServerURL, OAuth: &OAuthMetadata{Refresh: scenario.credential.OAuth.Refresh}}, replaced.Ciphertext) + if err != nil || secret.AccessToken != "manual-access" || secret.RefreshToken != "stored-refresh" || secret.Metadata.ExpiresAt != nil || !strings.Contains(string(replaced.Metadata), `"expires_at":null`) { + t.Fatal("replacement was not sealed with its patched metadata", err) + } + command.ExpiresAtSet, command.ExpiresAt = true, ptr("not-a-date") + replaced = SealedOAuth{} + if _, err := scenario.service.UpdateOAuthCredential(t.Context(), command); !errors.Is(err, ErrInvalidInput) || replaced.Ciphertext != nil { + t.Fatal("an invalid expiry was stored", err) + } +} diff --git a/services/core/internal/vaults/storage.go b/services/core/internal/vaults/storage.go new file mode 100644 index 000000000..683d0a4e9 --- /dev/null +++ b/services/core/internal/vaults/storage.go @@ -0,0 +1,115 @@ +package vaults + +import "context" + +// Reader reads Vaults and Credential metadata within one tenant. A malformed +// tenant, Vault or Credential ID is ErrInvalidInput, except the parent Vault +// of ListCredentials, which is ErrNotFound; a missing or foreign resource is +// ErrNotFound. +type Reader interface { + GetVault(ctx context.Context, tenantID, vaultID string) (Vault, error) + ListVaults(ctx context.Context, tenantID string, query PageQuery) (VaultPage, error) + GetCredential(ctx context.Context, tenantID, vaultID, credentialID string) (Credential, error) + // ListCredentials resolves the parent Vault first: an inaccessible parent + // is ErrNotFound, never an empty page. + ListCredentials(ctx context.Context, tenantID, vaultID string, query PageQuery) (CredentialPage, error) +} + +// Storage persists Vaults and Credentials. Every operation is scoped to the +// tenant it names, and every caller-visible write records its audit row in +// the same transaction. In a write, a malformed ID of an existing resource +// names none and is ErrNotFound; a malformed tenant of a new Vault or a +// malformed new Credential ID is ErrInvalidInput. +type Storage interface { + Reader + + // CreateVault stores a new Vault and allocates its ID. + CreateVault(ctx context.Context, vault NewVault) (Vault, error) + // DeleteVault deletes a Vault with all of its Credentials and returns its ID. + DeleteVault(ctx context.Context, tenantID, vaultID string) (string, error) + // CreateCredential stores a sealed Credential in a Vault of the tenant. + CreateCredential(ctx context.Context, credential NewCredential) (Credential, error) + // ReplaceStaticToken replaces a static_bearer Credential's sealed token. + ReplaceStaticToken(ctx context.Context, replacement StaticTokenReplacement) (Credential, error) + // DeleteCredential deletes a Credential and its sealed secret and returns its ID. + DeleteCredential(ctx context.Context, key CredentialKey) (string, error) + + // WithOAuthCredential runs apply in one transaction and commits only when + // apply returns nil. It returns apply's error unchanged. + WithOAuthCredential(ctx context.Context, key CredentialKey, apply func(OAuthTx) error) error + + // CountOwnedVaults counts how many of the given Vaults the tenant owns. + CountOwnedVaults(ctx context.Context, tenantID string, vaultIDs []string) (int, error) + // FindMCPCredentials returns at most two Credentials of the attached + // Vaults that the query selects, ordered by ID. + FindMCPCredentials(ctx context.Context, query MCPCredentialQuery) ([]MCPCredentialMatch, error) + // StaticTokenCiphertext returns a static_bearer Credential's sealed token + // when the complete frozen scope still names it. + StaticTokenCiphertext(ctx context.Context, query StaticTokenQuery) ([]byte, error) +} + +// OAuthTx is one mcp_oauth Credential inside a WithOAuthCredential +// transaction. +type OAuthTx interface { + // LoadOAuthCredential locks the Credential until the transaction ends, so + // competing refreshes, replacements and deletions, including the parent + // Vault's, wait. It returns the metadata and the sealed secret. + LoadOAuthCredential(ctx context.Context) (Credential, []byte, error) + // ApplyOAuthRefresh stores a refreshed grant. Execution refreshes are not + // caller writes and record no audit row. + ApplyOAuthRefresh(ctx context.Context, sealed SealedOAuth) error + // ApplyOAuthReplacement stores a caller's replacement and audits it. + ApplyOAuthReplacement(ctx context.Context, sealed SealedOAuth) (Credential, error) +} + +// CredentialKey names one Credential of one Vault of one tenant. +type CredentialKey struct { + TenantID, VaultID, CredentialID string +} + +// NewVault is a Vault ready to store. Metadata is its encoded JSON object. +type NewVault struct { + TenantID string + Name *string + Metadata []byte +} + +// NewCredential is a Credential with its ID, sealed to that ID, ready to store. +type NewCredential struct { + CredentialKey + Name, AuthType, MCPServerURL string + // OAuthMetadata is the encoded OAuthMetadata of an mcp_oauth Credential. + OAuthMetadata []byte + Ciphertext []byte +} + +// StaticTokenReplacement is a static_bearer token sealed to the destination +// the write matches. +type StaticTokenReplacement struct { + CredentialKey + MCPServerURL string + Ciphertext []byte +} + +// SealedOAuth is an OAuth grant sealed to the destination the write matches, +// with the metadata the seal authenticates. +type SealedOAuth struct { + MCPServerURL string + Metadata, Ciphertext []byte +} + +// MCPCredentialQuery selects a named Credential by ID alone, so its +// destination can be compared, and otherwise selects by exact destination. +type MCPCredentialQuery struct { + TenantID string + VaultIDs []string + ServerURL string + CredentialID string +} + +// StaticTokenQuery is a frozen binding's complete scope. +type StaticTokenQuery struct { + TenantID string + VaultIDs []string + VaultID, CredentialID, MCPServerURL string +} diff --git a/services/core/internal/vaults/vault.go b/services/core/internal/vaults/vault.go new file mode 100644 index 000000000..177bbce7f --- /dev/null +++ b/services/core/internal/vaults/vault.go @@ -0,0 +1,38 @@ +package vaults + +import ( + "time" + "unicode/utf8" + + "github.com/google/uuid" +) + +// Vault is a tenant-owned resource, independent of Sessions and engine execution. +type Vault struct { + ID string + TenantID string + Name *string + Metadata map[string]string + CreatedAt time.Time +} + +type VaultPage struct { + Vaults []Vault + NextCursor string +} + +// validName checks a Vault or Credential name the public layer has already +// trimmed. +func validName(name string) bool { + return len(name) >= 1 && len(name) <= 256 && utf8.ValidString(name) +} + +// canonicalID returns the canonical form of a nonzero UUID. Credential +// secrets are sealed to canonical IDs, so every ID in a binding goes through it. +func canonicalID(value string) (string, bool) { + id, err := uuid.Parse(value) + if err != nil || id == uuid.Nil { + return "", false + } + return id.String(), true +}