From 73f83d41cddc84367b0aa8ebf9c090224165ad77 Mon Sep 17 00:00:00 2001 From: SaladDay <1203511142@qq.com> Date: Wed, 30 Sep 2026 16:31:36 +0000 Subject: [PATCH 1/2] Add the vaults domain and its PostgreSQL adapter --- .../postgres/vaultpg/audit_test.go | 220 +++++++ .../postgres/vaultpg/credentials.go | 201 +++++++ .../postgres/vaultpg/credentials_test.go | 497 +++++++++++++++ .../postgres/vaultpg/fixture_test.go | 136 +++++ .../persistence/postgres/vaultpg/oauth.go | 87 +++ .../postgres/vaultpg/oauth_test.go | 564 ++++++++++++++++++ .../persistence/postgres/vaultpg/selection.go | 64 ++ .../postgres/vaultpg/selection_test.go | 127 ++++ .../persistence/postgres/vaultpg/vaultpg.go | 191 ++++++ .../postgres/vaultpg/vaults_test.go | 348 +++++++++++ services/core/internal/vaults/credential.go | 37 ++ services/core/internal/vaults/doc.go | 5 + services/core/internal/vaults/errors.go | 54 ++ services/core/internal/vaults/oauth.go | 158 +++++ services/core/internal/vaults/page.go | 38 ++ services/core/internal/vaults/rules_test.go | 279 +++++++++ services/core/internal/vaults/selection.go | 122 ++++ services/core/internal/vaults/service.go | 414 +++++++++++++ services/core/internal/vaults/service_test.go | 519 ++++++++++++++++ services/core/internal/vaults/storage.go | 115 ++++ services/core/internal/vaults/vault.go | 38 ++ 21 files changed, 4214 insertions(+) create mode 100644 services/core/internal/persistence/postgres/vaultpg/audit_test.go create mode 100644 services/core/internal/persistence/postgres/vaultpg/credentials.go create mode 100644 services/core/internal/persistence/postgres/vaultpg/credentials_test.go create mode 100644 services/core/internal/persistence/postgres/vaultpg/fixture_test.go create mode 100644 services/core/internal/persistence/postgres/vaultpg/oauth.go create mode 100644 services/core/internal/persistence/postgres/vaultpg/oauth_test.go create mode 100644 services/core/internal/persistence/postgres/vaultpg/selection.go create mode 100644 services/core/internal/persistence/postgres/vaultpg/selection_test.go create mode 100644 services/core/internal/persistence/postgres/vaultpg/vaultpg.go create mode 100644 services/core/internal/persistence/postgres/vaultpg/vaults_test.go create mode 100644 services/core/internal/vaults/credential.go create mode 100644 services/core/internal/vaults/doc.go create mode 100644 services/core/internal/vaults/errors.go create mode 100644 services/core/internal/vaults/oauth.go create mode 100644 services/core/internal/vaults/page.go create mode 100644 services/core/internal/vaults/rules_test.go create mode 100644 services/core/internal/vaults/selection.go create mode 100644 services/core/internal/vaults/service.go create mode 100644 services/core/internal/vaults/service_test.go create mode 100644 services/core/internal/vaults/storage.go create mode 100644 services/core/internal/vaults/vault.go 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 00000000..329d2b1f --- /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 00000000..bd938d24 --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/credentials.go @@ -0,0 +1,201 @@ +package vaultpg + +import ( + "context" + "encoding/json" + "errors" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "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, "credential creation failed", 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 errors.Is(err, pgx.ErrNoRows) { + return vaults.Credential{}, vaults.ErrNotFound + } + if err != nil { + return vaults.Credential{}, errors.New("credential lookup failed") + } + 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, "credential list failed", 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 errors.New("credential list failed") + } + 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, "credential update failed", 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, "credential deletion failed", 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 00000000..95036d26 --- /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 updateErr == nil || updateErr.Error() != "credential update failed" { + t.Fatal("database write failure was accepted or exposed", 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 deletionErr == nil || deletionErr.Error() != "credential deletion failed" { + t.Fatal("failed mutation was accepted or exposed") + } + 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 00000000..e0cef08a --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/fixture_test.go @@ -0,0 +1,136 @@ +package vaultpg_test + +import ( + "context" + "encoding/json" + "errors" + "testing" + + "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)) +} + +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 00000000..a9422556 --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/oauth.go @@ -0,0 +1,87 @@ +package vaultpg + +import ( + "context" + "errors" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgtype" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/adminaudit" + "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" +) + +// 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)} + var applied error + err := s.pool.Transaction(ctx, func(ctx context.Context, t pgx.Tx) error { + tx.q = sqlc.New(t) + applied = apply(tx) + return applied + }) + if err != nil && applied == nil { + return errors.New("credential transaction failed") + } + return err +} + +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 errors.Is(err, pgx.ErrNoRows) { + return vaults.Credential{}, nil, vaults.ErrNotFound + } + if err != nil { + return vaults.Credential{}, nil, 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 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 + } + err = auditpg.RecordWriteAudit(ctx, t.q, t.tenantID, "update", "credential", updated.ID, updated.VaultID) + if errors.Is(err, writeaudit.ErrInvalidSource) || errors.Is(err, adminaudit.ErrInvalidSource) { + return vaults.Credential{}, err + } + if err != nil { + return vaults.Credential{}, errors.New("credential update failed") + } + 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 errors.Is(err, pgx.ErrNoRows) { + return vaults.Credential{}, vaults.ErrNotFound + } + if err != nil { + return vaults.Credential{}, errors.New("credential update failed") + } + 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 00000000..41463c88 --- /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 00000000..e8f69a2e --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/selection.go @@ -0,0 +1,64 @@ +package vaultpg + +import ( + "context" + "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/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, errors.New("cannot resolve attached Vaults") + } + 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, errors.New("cannot resolve MCP credential") + } + 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 errors.Is(err, pgx.ErrNoRows) { + return nil, vaults.ErrNotFound + } + if err != nil { + return nil, errors.New("cannot read MCP credential") + } + 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 00000000..5b6ee8fd --- /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 00000000..0b620a21 --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/vaultpg.go @@ -0,0 +1,191 @@ +// 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/adminaudit" + "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: no +// row or a missing parent is ErrNotFound, audit provenance and unstorable-text rejections keep +// their shared errors, and any other failure is the operation's opaque +// failure, never database text. +func (s *Store) write(ctx context.Context, failure string, apply func(context.Context, *sqlc.Queries) error) error { + err := s.pool.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + return apply(ctx, sqlc.New(tx)) + }) + switch { + case err == nil: + return nil + case errors.Is(err, pgx.ErrNoRows), errors.Is(err, vaults.ErrNotFound): + return vaults.ErrNotFound + case errors.Is(err, writeaudit.ErrInvalidSource), errors.Is(err, adminaudit.ErrInvalidSource): + return err + case pgunit.IsUnstorableText(err): + return textvalue.ErrUnstorable + default: + return errors.New(failure) + } +} + +// read runs apply in one snapshot. apply returns translated errors; a failure +// to begin or commit is the operation's opaque failure. +func (s *Store) read(ctx context.Context, failure string, apply func(context.Context, *sqlc.Queries) error) error { + var applied error + err := s.pool.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { + applied = apply(ctx, sqlc.New(tx)) + return applied + }) + if err != nil && applied == nil { + return errors.New(failure) + } + 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, "vault creation failed", 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 errors.Is(err, pgx.ErrNoRows) { + return vaults.Vault{}, vaults.ErrNotFound + } + if err != nil { + return vaults.Vault{}, errors.New("vault lookup failed") + } + 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, "vault list failed", 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 errors.New("vault list failed") + } + 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, "vault deletion failed", 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 00000000..5d4f82f5 --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/vaults_test.go @@ -0,0 +1,348 @@ +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 deletionErr == nil || deletionErr.Error() != "vault deletion failed" { + t.Fatal("failed mutation was accepted or exposed", 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/vaults/credential.go b/services/core/internal/vaults/credential.go new file mode 100644 index 00000000..5a28ddfd --- /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 00000000..87ec6298 --- /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 00000000..75a95ba3 --- /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 00000000..d455ef48 --- /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 00000000..bd3cd0f7 --- /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 00000000..2cc10cfe --- /dev/null +++ b/services/core/internal/vaults/rules_test.go @@ -0,0 +1,279 @@ +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 00000000..5a98c203 --- /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 00000000..2c2d1b3d --- /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 00000000..3201837f --- /dev/null +++ b/services/core/internal/vaults/service_test.go @@ -0,0 +1,519 @@ +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 00000000..683d0a4e --- /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 00000000..177bbce7 --- /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 +} From 968dc18531470e2c9881f81cae1e756849cba702 Mon Sep 17 00:00:00 2001 From: SaladDay <1203511142@qq.com> Date: Wed, 30 Sep 2026 17:03:25 +0000 Subject: [PATCH 2/2] Serve Vaults and Credentials from the vaults domain The api handlers, the Session credential binding and the Dispatcher's MCP bearer lookup use the vaults service and the vaultpg adapter. The Vault and Credential code and tests leave store. --- services/core/IMPLEMENTATION.md | 5 +- services/core/cmd/server/http_routes_test.go | 3 +- services/core/cmd/server/main.go | 14 +- services/core/internal/api/claude_mcp_test.go | 10 +- services/core/internal/api/configuration.go | 10 +- services/core/internal/api/credentials.go | 23 +- .../core/internal/api/credentials_delete.go | 5 +- .../internal/api/credentials_delete_test.go | 8 +- .../core/internal/api/credentials_list.go | 6 +- .../internal/api/credentials_list_test.go | 17 +- .../core/internal/api/credentials_oauth.go | 40 +-- .../internal/api/credentials_oauth_test.go | 31 +- .../core/internal/api/credentials_test.go | 30 +- .../core/internal/api/credentials_update.go | 15 +- .../internal/api/credentials_update_test.go | 14 +- services/core/internal/api/dependencies.go | 4 +- .../core/internal/api/dependencies_test.go | 7 +- services/core/internal/api/errors.go | 9 - services/core/internal/api/errors_vaults.go | 33 ++ services/core/internal/api/fakes_test.go | 117 +++---- services/core/internal/api/handler.go | 2 +- .../core/internal/api/session_credentials.go | 12 +- .../internal/api/session_credentials_test.go | 9 +- .../internal/api/validation_errors_test.go | 5 +- .../core/internal/api/vault_pagination.go | 12 +- .../internal/api/vault_pagination_test.go | 16 +- services/core/internal/api/vaults.go | 53 ++-- services/core/internal/api/vaults_delete.go | 5 +- .../core/internal/api/vaults_delete_test.go | 8 +- services/core/internal/api/vaults_list.go | 6 +- .../core/internal/api/vaults_list_test.go | 19 +- services/core/internal/api/vaults_test.go | 24 +- .../core/internal/execution/dispatcher.go | 19 +- .../internal/execution/mcp_credentials.go | 6 +- .../execution/mcp_credentials_test.go | 6 +- .../core/internal/execution/mcp_support.go | 8 +- .../internal/execution/mcp_support_test.go | 48 ++- .../core/internal/execution/owner_test.go | 8 +- services/core/internal/execution/request.go | 6 +- services/core/internal/execution/worker.go | 3 + .../postgres/vaultpg/credentials.go | 16 +- .../postgres/vaultpg/credentials_test.go | 8 +- .../postgres/vaultpg/fixture_test.go | 8 + .../persistence/postgres/vaultpg/oauth.go | 33 +- .../persistence/postgres/vaultpg/selection.go | 11 +- .../persistence/postgres/vaultpg/vaultpg.go | 59 ++-- .../postgres/vaultpg/vaults_test.go | 8 +- .../providers/configuration_flow_test.go | 3 +- .../internal/store/admin_delete_audit_test.go | 8 +- .../core/internal/store/credential_cipher.go | 14 + .../internal/store/credential_oauth_secret.go | 82 ----- .../internal/store/credential_oauth_types.go | 38 --- .../core/internal/store/fixture_db_test.go | 11 +- .../internal/store/function_worker_test.go | 12 +- .../core/internal/store/mcp_credentials.go | 196 ------------ .../internal/store/mcp_credentials_oauth.go | 79 ----- .../store/mcp_credentials_oauth_test.go | 293 ------------------ .../internal/store/mcp_credentials_test.go | 138 --------- .../store/public_handler_fixture_test.go | 9 +- .../store/remote_mcp_credentials_test.go | 7 +- .../core/internal/store/remote_mcp_test.go | 11 +- services/core/internal/store/sessions.go | 6 - .../core/internal/store/vault_credentials.go | 127 -------- .../store/vault_credentials_delete.go | 44 --- .../store/vault_credentials_delete_test.go | 151 --------- .../internal/store/vault_credentials_list.go | 64 ---- .../store/vault_credentials_list_test.go | 150 --------- .../internal/store/vault_credentials_oauth.go | 189 ----------- .../store/vault_credentials_oauth_test.go | 280 ----------------- .../internal/store/vault_credentials_test.go | 145 --------- .../store/vault_credentials_update.go | 74 ----- .../store/vault_credentials_update_test.go | 184 ----------- services/core/internal/store/vaults.go | 111 ------- services/core/internal/store/vaults_delete.go | 40 --- .../core/internal/store/vaults_delete_test.go | 201 ------------ .../internal/store/vaults_fixture_test.go | 20 ++ services/core/internal/store/vaults_list.go | 60 ---- .../core/internal/store/vaults_list_test.go | 122 -------- services/core/internal/store/vaults_test.go | 87 ------ .../store/write_audit_resources_test.go | 70 +---- services/core/internal/vaults/rules_test.go | 8 +- services/core/internal/vaults/service_test.go | 4 +- 82 files changed, 531 insertions(+), 3326 deletions(-) create mode 100644 services/core/internal/api/errors_vaults.go create mode 100644 services/core/internal/store/credential_cipher.go delete mode 100644 services/core/internal/store/credential_oauth_secret.go delete mode 100644 services/core/internal/store/credential_oauth_types.go delete mode 100644 services/core/internal/store/mcp_credentials.go delete mode 100644 services/core/internal/store/mcp_credentials_oauth.go delete mode 100644 services/core/internal/store/mcp_credentials_oauth_test.go delete mode 100644 services/core/internal/store/mcp_credentials_test.go delete mode 100644 services/core/internal/store/vault_credentials.go delete mode 100644 services/core/internal/store/vault_credentials_delete.go delete mode 100644 services/core/internal/store/vault_credentials_delete_test.go delete mode 100644 services/core/internal/store/vault_credentials_list.go delete mode 100644 services/core/internal/store/vault_credentials_list_test.go delete mode 100644 services/core/internal/store/vault_credentials_oauth.go delete mode 100644 services/core/internal/store/vault_credentials_oauth_test.go delete mode 100644 services/core/internal/store/vault_credentials_test.go delete mode 100644 services/core/internal/store/vault_credentials_update.go delete mode 100644 services/core/internal/store/vault_credentials_update_test.go delete mode 100644 services/core/internal/store/vaults.go delete mode 100644 services/core/internal/store/vaults_delete.go delete mode 100644 services/core/internal/store/vaults_delete_test.go create mode 100644 services/core/internal/store/vaults_fixture_test.go delete mode 100644 services/core/internal/store/vaults_list.go delete mode 100644 services/core/internal/store/vaults_list_test.go delete mode 100644 services/core/internal/store/vaults_test.go diff --git a/services/core/IMPLEMENTATION.md b/services/core/IMPLEMENTATION.md index 13d442c7..da293e57 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 2fd15735..e6c96873 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 6f40a95f..9e4a32f4 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 640e34df..542f42e9 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 386b8e5b..7150241f 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 639f018e..c90948ab 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 758b26be..27a44210 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 781cf7c7..4749304a 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 dbca6786..6f2f335d 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 5896245b..78c15941 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 936ea119..4e384f66 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 d11d4ecb..779bbfc2 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 d8224533..0c2ace58 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 7b77aca7..58323a2a 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 0d0ff646..2cd88305 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 9ea03325..6f1f3c08 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 153c9e92..8ff1a324 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 1c832d4f..5f1ba4ea 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 00000000..430e2272 --- /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 180b892f..759849c6 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 e6a4f300..1f88dd83 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 954e4bfd..ac03df82 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 ff410772..fde0ac9d 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 8c0410ba..3b693261 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 08fbede3..5e46b3f8 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 d743af7a..24eec708 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 4bef9348..bd2eeb14 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 0d9274bf..3e666bf8 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 57e2e21e..1f6c8585 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 ab565ca4..72ebed7d 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 0c8d2bc8..302504fd 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 3b84c48d..e52faa0d 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 392fefea..6bf2faa8 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 edf6280a..0350d07f 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 c9c80127..a4cde8c9 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 3f197436..bc36584f 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 f522dbe5..d02cd278 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 bdd5ae67..d68a749b 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 35319705..dcd23cac 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 97e91ad2..fd9b6a42 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/credentials.go b/services/core/internal/persistence/postgres/vaultpg/credentials.go index bd938d24..dbed4d95 100644 --- a/services/core/internal/persistence/postgres/vaultpg/credentials.go +++ b/services/core/internal/persistence/postgres/vaultpg/credentials.go @@ -6,7 +6,6 @@ import ( "errors" "github.com/google/uuid" - "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgconn" "github.com/jackc/pgx/v5/pgtype" @@ -26,7 +25,7 @@ func (s *Store) CreateCredential(ctx context.Context, credential vaults.NewCrede return vaults.Credential{}, vaults.ErrInvalidInput } var created vaults.Credential - err := s.write(ctx, "credential creation failed", func(ctx context.Context, q *sqlc.Queries) error { + 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 @@ -81,11 +80,8 @@ func (s *Store) GetCredential(ctx context.Context, tenantID, vaultID, credential 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 errors.Is(err, pgx.ErrNoRows) { - return vaults.Credential{}, vaults.ErrNotFound - } if err != nil { - return vaults.Credential{}, errors.New("credential lookup failed") + return vaults.Credential{}, translate(err) } return credentialFromRow(row) } @@ -99,7 +95,7 @@ func (s *Store) ListCredentials(ctx context.Context, tenantID, vaultID string, q } vault := pgunit.PathID(vaultID) var page vaults.CredentialPage - err = s.read(ctx, "credential list failed", func(ctx context.Context, q *sqlc.Queries) error { + 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 @@ -121,7 +117,7 @@ func (s *Store) ListCredentials(ctx context.Context, tenantID, vaultID string, q } rows, err := q.ListCredentials(ctx, params) if err != nil { - return errors.New("credential list failed") + return err } page = vaults.CredentialPage{Credentials: make([]vaults.Credential, 0, min(query.Limit, len(rows)))} if len(rows) > query.Limit { @@ -147,7 +143,7 @@ func (s *Store) ListCredentials(ctx context.Context, tenantID, vaultID string, q // 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, "credential update failed", func(ctx context.Context, q *sqlc.Queries) error { + 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, @@ -171,7 +167,7 @@ func (s *Store) ReplaceStaticToken(ctx context.Context, replacement vaults.Stati func (s *Store) DeleteCredential(ctx context.Context, key vaults.CredentialKey) (string, error) { vault := pgunit.PathID(key.VaultID) var deleted string - err := s.write(ctx, "credential deletion failed", func(ctx context.Context, q *sqlc.Queries) error { + 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 diff --git a/services/core/internal/persistence/postgres/vaultpg/credentials_test.go b/services/core/internal/persistence/postgres/vaultpg/credentials_test.go index 95036d26..fafd9b86 100644 --- a/services/core/internal/persistence/postgres/vaultpg/credentials_test.go +++ b/services/core/internal/persistence/postgres/vaultpg/credentials_test.go @@ -344,8 +344,8 @@ func TestStaticCredentialUpdatePreservesBindingsAndReplacesCurrentSecret(t *test } // 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 updateErr == nil || updateErr.Error() != "credential update failed" { - t.Fatal("database write failure was accepted or exposed", updateErr) + 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. @@ -423,8 +423,8 @@ func TestCredentialDeletionScopeBindingAndRestart(t *testing.T) { } // 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 deletionErr == nil || deletionErr.Error() != "credential deletion failed" { - t.Fatal("failed mutation was accepted or exposed") + 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) diff --git a/services/core/internal/persistence/postgres/vaultpg/fixture_test.go b/services/core/internal/persistence/postgres/vaultpg/fixture_test.go index e0cef08a..e0fdda7f 100644 --- a/services/core/internal/persistence/postgres/vaultpg/fixture_test.go +++ b/services/core/internal/persistence/postgres/vaultpg/fixture_test.go @@ -6,6 +6,7 @@ import ( "errors" "testing" + "github.com/jackc/pgx/v5/pgconn" "github.com/jackc/pgx/v5/pgxpool" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" @@ -36,6 +37,13 @@ func readOnlyStore(t *testing.T, pool *pgxpool.Pool) *vaultpg.Store { 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) diff --git a/services/core/internal/persistence/postgres/vaultpg/oauth.go b/services/core/internal/persistence/postgres/vaultpg/oauth.go index a9422556..13af0461 100644 --- a/services/core/internal/persistence/postgres/vaultpg/oauth.go +++ b/services/core/internal/persistence/postgres/vaultpg/oauth.go @@ -2,33 +2,24 @@ package vaultpg import ( "context" - "errors" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/adminaudit" "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" ) // 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)} - var applied error - err := s.pool.Transaction(ctx, func(ctx context.Context, t pgx.Tx) error { + return translate(s.pool.Transaction(ctx, func(ctx context.Context, t pgx.Tx) error { tx.q = sqlc.New(t) - applied = apply(tx) - return applied - }) - if err != nil && applied == nil { - return errors.New("credential transaction failed") - } - return err + return apply(tx) + })) } type oauthTx struct { @@ -39,11 +30,8 @@ type oauthTx struct { 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 errors.Is(err, pgx.ErrNoRows) { - return vaults.Credential{}, nil, vaults.ErrNotFound - } if err != nil { - return vaults.Credential{}, nil, errors.New("credential lookup failed") + 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}) @@ -63,12 +51,8 @@ func (t *oauthTx) ApplyOAuthReplacement(ctx context.Context, sealed vaults.Seale if err != nil { return vaults.Credential{}, err } - err = auditpg.RecordWriteAudit(ctx, t.q, t.tenantID, "update", "credential", updated.ID, updated.VaultID) - if errors.Is(err, writeaudit.ErrInvalidSource) || errors.Is(err, adminaudit.ErrInvalidSource) { - return vaults.Credential{}, err - } - if err != nil { - return vaults.Credential{}, errors.New("credential update failed") + 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 } @@ -77,11 +61,8 @@ func (t *oauthTx) ApplyOAuthReplacement(ctx context.Context, sealed vaults.Seale 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 errors.Is(err, pgx.ErrNoRows) { - return vaults.Credential{}, vaults.ErrNotFound - } if err != nil { - return vaults.Credential{}, errors.New("credential update failed") + return vaults.Credential{}, translate(err) } return credentialFromRow(sqlc.GetCredentialRow(row)) } diff --git a/services/core/internal/persistence/postgres/vaultpg/selection.go b/services/core/internal/persistence/postgres/vaultpg/selection.go index e8f69a2e..f45e2b1d 100644 --- a/services/core/internal/persistence/postgres/vaultpg/selection.go +++ b/services/core/internal/persistence/postgres/vaultpg/selection.go @@ -2,10 +2,8 @@ package vaultpg import ( "context" - "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" @@ -16,7 +14,7 @@ import ( 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, errors.New("cannot resolve attached Vaults") + return 0, translate(err) } return len(owned), nil } @@ -29,7 +27,7 @@ func (s *Store) FindMCPCredentials(ctx context.Context, query vaults.MCPCredenti 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, errors.New("cannot resolve MCP credential") + return nil, translate(err) } matches := make([]vaults.MCPCredentialMatch, 0, len(rows)) for _, row := range rows { @@ -46,11 +44,8 @@ func (s *Store) StaticTokenCiphertext(ctx context.Context, query vaults.StaticTo TenantID: pgunit.PathID(query.TenantID), VaultIds: pathIDs(query.VaultIDs), VaultID: pgunit.PathID(query.VaultID), CredentialID: pgunit.PathID(query.CredentialID), McpServerUrl: query.MCPServerURL, }) - if errors.Is(err, pgx.ErrNoRows) { - return nil, vaults.ErrNotFound - } if err != nil { - return nil, errors.New("cannot read MCP credential") + return nil, translate(err) } return ciphertext, nil } diff --git a/services/core/internal/persistence/postgres/vaultpg/vaultpg.go b/services/core/internal/persistence/postgres/vaultpg/vaultpg.go index 0b620a21..b77115e4 100644 --- a/services/core/internal/persistence/postgres/vaultpg/vaultpg.go +++ b/services/core/internal/persistence/postgres/vaultpg/vaultpg.go @@ -10,7 +10,6 @@ import ( "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/adminaudit" "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" @@ -28,38 +27,29 @@ 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: no -// row or a missing parent is ErrNotFound, audit provenance and unstorable-text rejections keep -// their shared errors, and any other failure is the operation's opaque -// failure, never database text. -func (s *Store) write(ctx context.Context, failure string, apply func(context.Context, *sqlc.Queries) error) error { - err := s.pool.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { +// 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 err == nil: - return nil - case errors.Is(err, pgx.ErrNoRows), errors.Is(err, vaults.ErrNotFound): + case errors.Is(err, pgx.ErrNoRows): return vaults.ErrNotFound - case errors.Is(err, writeaudit.ErrInvalidSource), errors.Is(err, adminaudit.ErrInvalidSource): - return err case pgunit.IsUnstorableText(err): return textvalue.ErrUnstorable - default: - return errors.New(failure) - } -} - -// read runs apply in one snapshot. apply returns translated errors; a failure -// to begin or commit is the operation's opaque failure. -func (s *Store) read(ctx context.Context, failure string, apply func(context.Context, *sqlc.Queries) error) error { - var applied error - err := s.pool.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { - applied = apply(ctx, sqlc.New(tx)) - return applied - }) - if err != nil && applied == nil { - return errors.New(failure) } return err } @@ -74,7 +64,7 @@ func (s *Store) CreateVault(ctx context.Context, vault vaults.NewVault) (vaults. name = pgtype.Text{String: *vault.Name, Valid: true} } var created vaults.Vault - err = s.write(ctx, "vault creation failed", func(ctx context.Context, q *sqlc.Queries) error { + 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 @@ -105,11 +95,8 @@ func (s *Store) GetVault(ctx context.Context, tenantID, vaultID string) (vaults. 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 errors.Is(err, pgx.ErrNoRows) { - return vaults.Vault{}, vaults.ErrNotFound - } if err != nil { - return vaults.Vault{}, errors.New("vault lookup failed") + return vaults.Vault{}, translate(err) } return vaultFromRow(row) } @@ -125,7 +112,7 @@ func (s *Store) ListVaults(ctx context.Context, tenantID string, query vaults.Pa return vaults.VaultPage{}, err } var page vaults.VaultPage - err = s.read(ctx, "vault list failed", func(ctx context.Context, q *sqlc.Queries) error { + 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. @@ -138,7 +125,7 @@ func (s *Store) ListVaults(ctx context.Context, tenantID string, query vaults.Pa } rows, err := q.ListVaults(ctx, params) if err != nil { - return errors.New("vault list failed") + return err } page = vaults.VaultPage{Vaults: make([]vaults.Vault, 0, min(query.Limit, len(rows)))} if len(rows) > query.Limit { @@ -163,7 +150,7 @@ func (s *Store) ListVaults(ctx context.Context, tenantID string, query vaults.Pa // 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, "vault deletion failed", func(ctx context.Context, q *sqlc.Queries) error { + 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 diff --git a/services/core/internal/persistence/postgres/vaultpg/vaults_test.go b/services/core/internal/persistence/postgres/vaultpg/vaults_test.go index 5d4f82f5..cfe6d6ce 100644 --- a/services/core/internal/persistence/postgres/vaultpg/vaults_test.go +++ b/services/core/internal/persistence/postgres/vaultpg/vaults_test.go @@ -226,8 +226,8 @@ func TestVaultDeletionCascadeBindingAndRestart(t *testing.T) { } } _, deletionErr := newService(t, readOnlyStore(t, pool), nil, nil).DeleteVault(t.Context(), vaults.DeleteVault{TenantID: tenant, VaultID: vault.ID}) - if deletionErr == nil || deletionErr.Error() != "vault deletion failed" { - t.Fatal("failed mutation was accepted or exposed", deletionErr) + if !isReadOnlyFailure(deletionErr) { + t.Fatal("failed mutation was accepted or translated", deletionErr) } tx, err := pool.Begin(t.Context()) if err != nil { @@ -308,7 +308,9 @@ func TestVaultDeletionConcurrentChildMutations(t *testing.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) } + 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") diff --git a/services/core/internal/sandbox/providers/configuration_flow_test.go b/services/core/internal/sandbox/providers/configuration_flow_test.go index 1611a17a..74f6adfb 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 b68342b4..4bf59c64 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 00000000..3051915f --- /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 943b14eb..00000000 --- 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 1eb674f0..00000000 --- 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 ac144563..25636872 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 65d9d16d..0279a72a 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 3e635f8a..00000000 --- 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 4abdb853..00000000 --- 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 1d176013..00000000 --- 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 cb76b56e..00000000 --- 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 b3da7c69..03d3714a 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 60eea2c4..fb3f3f1d 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 0c324b1e..841b76f8 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 af0f8a0f..c81dcc9a 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 870085af..00000000 --- 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 8fb28bb8..00000000 --- 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 1adcf0a1..00000000 --- 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 082dfa5a..00000000 --- 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 57bd8d9c..00000000 --- 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 c2016132..00000000 --- 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 f41c9744..00000000 --- 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 3faafd09..00000000 --- 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 e9692947..00000000 --- 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 efef86bc..00000000 --- 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 4f66c594..00000000 --- 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 811240a2..00000000 --- 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 3a02e5ab..00000000 --- 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 00000000..ef742202 --- /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 3015a8c6..00000000 --- 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 cd574a4b..00000000 --- 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 09e50a09..00000000 --- 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 434bccbc..f30ebf11 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/rules_test.go b/services/core/internal/vaults/rules_test.go index 2cc10cfe..e6f9110f 100644 --- a/services/core/internal/vaults/rules_test.go +++ b/services/core/internal/vaults/rules_test.go @@ -59,7 +59,9 @@ func TestMCPCredentialSelectionRules(t *testing.T) { 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} } + named := func(id string) MCPCredentialRequest { + return MCPCredentialRequest{ServerLabel: "tools", ServerURL: url, CredentialID: &id} + } for _, tc := range []struct { request MCPCredentialRequest attached []string @@ -160,7 +162,9 @@ func TestOAuthCreationRules(t *testing.T) { }{ {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, 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}, diff --git a/services/core/internal/vaults/service_test.go b/services/core/internal/vaults/service_test.go index 3201837f..6f67e77a 100644 --- a/services/core/internal/vaults/service_test.go +++ b/services/core/internal/vaults/service_test.go @@ -395,7 +395,9 @@ func newOAuthScenario(t *testing.T, expiresAt time.Time, refresher oauthrefresh. t.Fatalf("unexpected key %+v", key) } err := apply(&fakeOAuthTx{t: t, - load: func(context.Context) (Credential, []byte, error) { return scenario.credential, scenario.ciphertext, nil }, + 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