diff --git a/services/core/IMPLEMENTATION.md b/services/core/IMPLEMENTATION.md index d6adda08..33a277d8 100644 --- a/services/core/IMPLEMENTATION.md +++ b/services/core/IMPLEMENTATION.md @@ -10,6 +10,8 @@ These are the code-level rules of `services/core` that no contract states. Contr `internal/persistence/postgres/auditpg` is the one adapter that other adapters call directly. Audit rows are written inside the business transaction, so each adapter calls `RecordWriteAudit`, `RecordAdminMutation` or `RecordDeploymentMutation` with its own transaction's queries; the audit provenance travels in the context. `writeaudit` and `adminaudit` own the sources, their validation, `ErrInvalidSource` and the read models, and `auditpg.Store` serves the audit reads. +The adapter that stores a secret seals and opens it with the credential key `cmd/server` builds it with: domain storage interfaces carry plaintext, domain service constructors never take the key, a missing key is `credentialcrypto.ErrUnavailable`, and a ciphertext that fails to open or authenticate is an internal error, never a missing key, value or row. + Two errors are shared across domains, each with one `api` helper: `textvalue.ErrUnstorable` (400, `writeTextValueError`) for text PostgreSQL cannot store, and `credentialcrypto.ErrUnavailable` (503, `writeCredentialUnavailableError`) for a missing credential key. The audit `ErrInvalidSource` errors pass through adapters unchanged, and `writeAuditSourceError` maps both to 400. Shared vocabulary has one owner each, and domains use it rather than copy it. `internal/environmentconfig` owns Environment setup, Skills, Plugins and initial files with their validation and public metadata; `Setup.Validate` checks requested configuration, where a Skill may be an unresolved reference, and `Setup.ValidateInstalled` checks frozen, installable configuration. `internal/skills` owns `ParseVersion`, the canonical positive decimal Skill version. `internal/metadata` owns the metadata rules: `Validate` for the pair, key and value limits and U+0000, `ValidateStorable` for U+0000 alone, and `Encode` with its 64 KiB bound. `internal/jsonobject` owns `Normalize`, the stable encoding of stored JSON objects that snapshots and retry identities compare. These packages import no persistence. @@ -108,7 +110,7 @@ Reusable Agents are tenant-scoped rows independent of Session snapshots and engi Saved execution defaults keep a model-provider bundle whole at every replacement boundary: endpoint, key, protocol and limits are never inherited separately. Agent JSON holds only the safe provider fields and an output-only configured flag; the complete bundle is encrypted separately with a tenant and Agent binding and its own purpose, and configuration and secret changes commit together under the Agent row lock. Model-only edits need no key. Merged harness, protocol and limits are validated without reading keys. Session creation reads safe defaults and ciphertext in one snapshot, and a complete Session override does not decrypt the inherited bundle. The resolved bundle is frozen in an encrypted Session-owned row, and dispatch fails closed when that snapshot is missing or cannot be decrypted; later Agent edits, default changes, restarts and suspension never resolve it again. -`modelconfiguration` owns each Harness's deployment default: its service validates a replacement through the Harness declaration and seals the complete bundle, and `modelconfigurationpg` stores it under a private revision UUID generated on every PUT, identical replacements included, with its audit row in the same transaction. Session creation resolves the ciphertext and revision together and freezes them; retries and older Sessions never gain or replace revision metadata. After a successful root terminal commit, the Dispatcher's required `modelconfiguration.Observer` runs one independent pool operation with at most one second to update the matching current revision. Only completed Turns and native provider failures with `engine_failed` count; cancelled work, Core or Runtime errors and input-policy classifications never do. `ShouldObserveProvider` skips the round trip for outcomes that cannot count, and the SQL stays authoritative: it verifies the tenant, root Turn and committed outcome. Both check the same shared cases. The metadata-only transaction sets statement and lock timeouts within the remaining budget and issues one UPDATE that locks only the default and samples database time after the lock. Errors throttle for 30 seconds, ordinary successes throttle for 30 seconds with one immediate recovery write after each accepted error, and an unchanged revision has at most three effective writes in any 30-second window of nondecreasing database time. Observations never change `updated_at`, readiness or execution truth, and can be lost or stale; there is no queue, retry, probe or backfill. +`modelconfiguration` owns each Harness's deployment default: its service validates a replacement through the Harness declaration, and `modelconfigurationpg` seals the complete bundle to the Harness and stores it under a private revision UUID generated on every PUT, identical replacements included, with its audit row in the same transaction. Session creation resolves the opened bundle and its revision together and freezes them; retries and older Sessions never gain or replace revision metadata. After a successful root terminal commit, the Dispatcher's required `modelconfiguration.Observer` runs one independent pool operation with at most one second to update the matching current revision. Only completed Turns and native provider failures with `engine_failed` count; cancelled work, Core or Runtime errors and input-policy classifications never do. `ShouldObserveProvider` skips the round trip for outcomes that cannot count, and the SQL stays authoritative: it verifies the tenant, root Turn and committed outcome. Both check the same shared cases. The metadata-only transaction sets statement and lock timeouts within the remaining budget and issues one UPDATE that locks only the default and samples database time after the lock. Errors throttle for 30 seconds, ordinary successes throttle for 30 seconds with one immediate recovery write after each accepted error, and an unchanged revision has at most three effective writes in any 30-second window of nondecreasing database time. Observations never change `updated_at`, readiness or execution truth, and can be lost or stale; there is no queue, retry, probe or backfill. Session execution-configuration reads use a separate immutable safe projection written with its provenance in the Session's creation transaction. It reads no ciphertext, never recomputes sources from current Agents or defaults, never touches activity or wakes a sandbox, and does not affect retry identity. @@ -121,7 +123,7 @@ Provider input validation uses the adapter rules in `internal/harnessconfig`: on [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. +- `vaultpg` seals secret values 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. - A static replacement is one SQL mutation scoped by tenant, Vault, Credential, auth type and destination, reusing the safe metadata for the immutable binding; it never decrypts the previous token, and a failed write keeps the old row. - OAuth refresh and replacement serialize on the Credential row lock, authenticate the stored grant metadata against its encrypted copy before using an endpoint, and persist the refreshed grant before returning an access token, so a stale refresh cannot undo a deletion or Vault cascade. - Credential deletion is one mutation checked by tenant, Vault and ID; Vault deletion removes the parent and its Credentials through the foreign-key cascade in one SQL statement, without decrypting, needing the key or calling providers. diff --git a/services/core/cmd/server/main.go b/services/core/cmd/server/main.go index f9990f8d..b55976a4 100644 --- a/services/core/cmd/server/main.go +++ b/services/core/cmd/server/main.go @@ -128,8 +128,8 @@ func run() error { if err != nil { return err } - vaultStore := vaultpg.New(units) - vaultService, err := vaults.NewService(vaultStore, credentialKey, oauthClient) + vaultStore := vaultpg.New(units, credentialKey) + vaultService, err := vaults.NewService(vaultStore, oauthClient) if err != nil { return err } @@ -138,8 +138,8 @@ func run() error { if err != nil { return err } - modelConfigurationStore := modelconfigurationpg.New(units) - modelConfigurationService, err := modelconfiguration.NewService(modelConfigurationStore, credentialKey) + modelConfigurationStore := modelconfigurationpg.New(units, credentialKey) + modelConfigurationService, err := modelconfiguration.NewService(modelConfigurationStore) if err != nil { return err } diff --git a/services/core/internal/agents/storage.go b/services/core/internal/agents/storage.go index 16969e87..7203ea76 100644 --- a/services/core/internal/agents/storage.go +++ b/services/core/internal/agents/storage.go @@ -56,7 +56,8 @@ type Reader interface { GetAgent(ctx context.Context, tenantID, agentID string) (Agent, error) ListAgents(context.Context, ListQuery) (Page, error) // GetAgentWithModelProvider reads the Agent and its opened model provider - // bundle from one snapshot. The bundle is nil when the Agent has none; a - // bundle that cannot be opened is credentialcrypto.ErrUnavailable. + // bundle from one snapshot. The bundle is nil when the Agent has none. + // Opening one without a credential key is credentialcrypto.ErrUnavailable; + // a bundle that fails to open is an internal error. GetAgentWithModelProvider(ctx context.Context, tenantID, agentID string) (Agent, *v1.ModelProviderInput, error) } diff --git a/services/core/internal/api/core_model_provider_validation_test.go b/services/core/internal/api/core_model_provider_validation_test.go index d059dc24..f4d01ac4 100644 --- a/services/core/internal/api/core_model_provider_validation_test.go +++ b/services/core/internal/api/core_model_provider_validation_test.go @@ -1,7 +1,6 @@ package api import ( - "bytes" "context" "encoding/json" "net/http" @@ -9,7 +8,6 @@ import ( "strings" "testing" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/modelconfiguration" ) @@ -30,17 +28,13 @@ func (s *coreProviderValidationStore) Delete(context.Context, string) error { return nil } -func (s *coreProviderValidationStore) LoadSealed(context.Context, string) (modelconfiguration.Sealed, error) { - unexpectedCall(s.t, "LoadSealed") - return modelconfiguration.Sealed{}, nil +func (s *coreProviderValidationStore) LoadBundle(context.Context, string) (modelconfiguration.Bundle, error) { + unexpectedCall(s.t, "LoadBundle") + return modelconfiguration.Bundle{}, nil } func (s *coreProviderValidationStore) configure(d *Dependencies, _ *testFakes) { - cipher, err := credentialcrypto.New(bytes.Repeat([]byte{3}, 32)) - if err != nil { - s.t.Fatal(err) - } - service, err := modelconfiguration.NewService(s, cipher) + service, err := modelconfiguration.NewService(s) if err != nil { s.t.Fatal(err) } diff --git a/services/core/internal/execution/deployment_provider_observations_test.go b/services/core/internal/execution/deployment_provider_observations_test.go index f16ed807..ac245d70 100644 --- a/services/core/internal/execution/deployment_provider_observations_test.go +++ b/services/core/internal/execution/deployment_provider_observations_test.go @@ -51,8 +51,8 @@ func newFinishObservationFixture(t *testing.T, maxConnections int32) finishObser if err != nil { t.Fatal(err) } - defaults := modelconfigurationpg.New(pgunit.NewPool(pool)) - service, err := modelconfiguration.NewService(defaults, cipher) + defaults := modelconfigurationpg.New(pgunit.NewPool(pool), cipher) + service, err := modelconfiguration.NewService(defaults) if err != nil { t.Fatal(err) } @@ -194,7 +194,7 @@ func TestFinishRunObservationLockTimeoutAndFailureKeepLease(t *testing.T) { if err != nil { t.Fatal(err) } - d := Dispatcher{Observer: modelconfigurationpg.New(pgunit.NewPool(pool))} + d := Dispatcher{Observer: modelconfigurationpg.New(pgunit.NewPool(pool), nil)} started := time.Now() d.observeDeploymentProvider(f.tenant, f.session.ID, turn) if elapsed := time.Since(started); elapsed < 900*time.Millisecond || elapsed > 2*time.Second { diff --git a/services/core/internal/modelconfiguration/configuration.go b/services/core/internal/modelconfiguration/configuration.go index 8ad3066e..f4b9e776 100644 --- a/services/core/internal/modelconfiguration/configuration.go +++ b/services/core/internal/modelconfiguration/configuration.go @@ -40,18 +40,18 @@ type Snapshot struct { } // Record is a validated deployment default ready to store: its safe columns -// and the sealed complete bundle. Storage assigns a new revision to each -// Record it stores. +// and the complete bundle, provider key included, which storage seals to the +// Harness. Storage assigns a new revision to each Record it stores. type Record struct { Harness string Provider v1.ModelProviderView Model string HarnessConfig json.RawMessage - Sealed []byte + Configuration v1.ModelConfigurationInput } -// Sealed is a stored bundle and the revision read with it. -type Sealed struct { - Bundle []byte - Revision uuid.UUID +// Bundle is an opened stored bundle and the revision read with it. +type Bundle struct { + Configuration v1.ModelConfigurationInput + Revision uuid.UUID } diff --git a/services/core/internal/modelconfiguration/errors.go b/services/core/internal/modelconfiguration/errors.go index ef2aa95d..1c4719fc 100644 --- a/services/core/internal/modelconfiguration/errors.go +++ b/services/core/internal/modelconfiguration/errors.go @@ -4,9 +4,9 @@ import "errors" // Replace and Resolve report a configuration the Harness declaration rejects // with the contract's *v1.ModelProviderError, which names the field, and a -// bundle they cannot seal or open with credentialcrypto.ErrUnavailable. Storage -// passes textvalue.ErrUnstorable and adminaudit.ErrInvalidSource through -// unchanged. +// missing credential key with credentialcrypto.ErrUnavailable. A bundle that +// fails to open is an internal error. Storage passes textvalue.ErrUnstorable +// and adminaudit.ErrInvalidSource through unchanged. var ( // ErrNotFound reports a Harness without a deployment default. ErrNotFound = errors.New("the harness has no deployment default model configuration") diff --git a/services/core/internal/modelconfiguration/service.go b/services/core/internal/modelconfiguration/service.go index 2925873f..fa79b9b6 100644 --- a/services/core/internal/modelconfiguration/service.go +++ b/services/core/internal/modelconfiguration/service.go @@ -2,47 +2,35 @@ package modelconfiguration import ( "context" - "encoding/json" "errors" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" ) -// Service replaces, removes and resolves deployment defaults. The complete -// bundle, key included, is sealed to its Harness before it reaches storage. +// Service replaces, removes and resolves deployment defaults. Storage seals +// the complete bundle, key included, to its Harness. type Service struct { storage Storage - cipher *credentialcrypto.Cipher } -// NewService requires storage. A nil cipher means credential encryption is not -// configured: Replace and Resolve then return credentialcrypto.ErrUnavailable. -func NewService(storage Storage, cipher *credentialcrypto.Cipher) (*Service, error) { +// NewService requires storage. +func NewService(storage Storage) (*Service, error) { if storage == nil { return nil, errors.New("model configuration requires storage") } - return &Service{storage: storage, cipher: cipher}, nil + return &Service{storage: storage}, nil } // Replace validates the complete configuration through the Harness -// declaration, seals it and stores it under a new revision. Sessions that -// already froze a default keep theirs. +// declaration and stores it under a new revision. Sessions that already froze +// a default keep theirs. func (s *Service) Replace(ctx context.Context, replacement Replacement) (Configuration, error) { configuration := replacement.Configuration if err := configuration.ValidateHarness(replacement.Harness); err != nil { return Configuration{}, err } - raw, err := json.Marshal(configuration) - if err != nil { - return Configuration{}, err - } - sealed, err := s.cipher.SealDeploymentModelProvider(raw, replacement.Harness) - if err != nil { - return Configuration{}, credentialcrypto.ErrUnavailable - } view := configuration.SafeView() - return s.storage.Replace(ctx, Record{Harness: replacement.Harness, Provider: *view.ModelProvider, Model: view.Model, HarnessConfig: view.HarnessConfig, Sealed: sealed}) + return s.storage.Replace(ctx, Record{Harness: replacement.Harness, Provider: *view.ModelProvider, Model: view.Model, HarnessConfig: view.HarnessConfig, Configuration: configuration}) } // Delete removes the Harness's default. It is idempotent. Sessions that @@ -52,28 +40,21 @@ func (s *Service) Delete(ctx context.Context, harness string) error { } // Resolve opens the Harness's default for Session creation. It returns nil -// when the Harness has none, and credentialcrypto.ErrUnavailable when the -// bundle does not open. +// when the Harness has none, and credentialcrypto.ErrUnavailable without a +// credential key. func (s *Service) Resolve(ctx context.Context, harness string) (*Snapshot, error) { - sealed, err := s.storage.LoadSealed(ctx, harness) + bundle, err := s.storage.LoadBundle(ctx, harness) if errors.Is(err, ErrNotFound) { return nil, nil } if err != nil { return nil, err } - raw, err := s.cipher.OpenDeploymentModelProvider(sealed.Bundle, harness) - if err != nil { - return nil, credentialcrypto.ErrUnavailable - } - var configuration v1.ModelConfigurationInput - if json.Unmarshal(raw, &configuration) != nil { - return nil, credentialcrypto.ErrUnavailable - } + configuration := bundle.Configuration // A bundle that opens but no longer validates is not a credential failure. // Its stored row stays intact so the operator can inspect and replace it. if err := configuration.ValidateHarness(harness); err != nil { return nil, err } - return &Snapshot{Provider: &configuration.ModelProvider, Model: configuration.Model, HarnessConfig: v1.ResolvedHarnessConfig(configuration.HarnessConfig), Revision: sealed.Revision}, nil + return &Snapshot{Provider: &configuration.ModelProvider, Model: configuration.Model, HarnessConfig: v1.ResolvedHarnessConfig(configuration.HarnessConfig), Revision: bundle.Revision}, nil } diff --git a/services/core/internal/modelconfiguration/service_test.go b/services/core/internal/modelconfiguration/service_test.go index 7ec0988f..3f23c5ad 100644 --- a/services/core/internal/modelconfiguration/service_test.go +++ b/services/core/internal/modelconfiguration/service_test.go @@ -1,10 +1,9 @@ package modelconfiguration import ( - "bytes" "context" - "encoding/json" "errors" + "reflect" "testing" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" @@ -17,7 +16,7 @@ type fakeStorage struct { t *testing.T replace func(context.Context, Record) (Configuration, error) delete func(context.Context, string) error - loadSealed func(context.Context, string) (Sealed, error) + loadBundle func(context.Context, string) (Bundle, error) } func (f *fakeStorage) Replace(ctx context.Context, record Record) (Configuration, error) { @@ -34,29 +33,20 @@ func (f *fakeStorage) Delete(ctx context.Context, harness string) error { return f.delete(ctx, harness) } -func (f *fakeStorage) LoadSealed(ctx context.Context, harness string) (Sealed, error) { - if f.loadSealed == nil { - f.t.Fatal("unexpected call to LoadSealed") +func (f *fakeStorage) LoadBundle(ctx context.Context, harness string) (Bundle, error) { + if f.loadBundle == nil { + f.t.Fatal("unexpected call to LoadBundle") } - return f.loadSealed(ctx, harness) -} - -func testCipher(t *testing.T, fill byte) *credentialcrypto.Cipher { - t.Helper() - cipher, err := credentialcrypto.New(bytes.Repeat([]byte{fill}, 32)) - if err != nil { - t.Fatal(err) - } - return cipher + return f.loadBundle(ctx, harness) } func validConfiguration() v1.ModelConfigurationInput { return v1.ModelConfigurationInput{ModelProvider: v1.ModelProviderInput{Protocol: "responses", BaseURL: "https://model.example/v1", APIKey: "secret-key"}, Model: "fixture-model"} } -func newService(t *testing.T, storage Storage, cipher *credentialcrypto.Cipher) *Service { +func newService(t *testing.T, storage Storage) *Service { t.Helper() - service, err := NewService(storage, cipher) + service, err := NewService(storage) if err != nil { t.Fatal(err) } @@ -64,20 +54,19 @@ func newService(t *testing.T, storage Storage, cipher *credentialcrypto.Cipher) } func TestNewServiceRequiresStorage(t *testing.T) { - if _, err := NewService(nil, testCipher(t, 1)); err == nil { + if _, err := NewService(nil); err == nil { t.Fatal("nil storage accepted") } } -func TestReplaceSealsTheCompleteBundleForItsHarness(t *testing.T) { - cipher := testCipher(t, 1) +func TestReplacePassesTheCompleteBundleWithItsSafeColumns(t *testing.T) { stored := Configuration{Harness: "codex", Model: "fixture-model"} var record Record storage := &fakeStorage{t: t, replace: func(_ context.Context, r Record) (Configuration, error) { record = r return stored, nil }} - result, err := newService(t, storage, cipher).Replace(t.Context(), Replacement{Harness: "codex", Configuration: validConfiguration()}) + result, err := newService(t, storage).Replace(t.Context(), Replacement{Harness: "codex", Configuration: validConfiguration()}) if err != nil || result.Harness != stored.Harness || result.Model != stored.Model { t.Fatal(result, err) } @@ -85,16 +74,14 @@ func TestReplaceSealsTheCompleteBundleForItsHarness(t *testing.T) { record.Provider != (v1.ModelProviderView{Protocol: "responses", BaseURL: "https://model.example/v1", APIKeyConfigured: true}) { t.Fatalf("safe columns: %+v", record) } - if bytes.Contains(record.Sealed, []byte("secret-key")) { - t.Fatal("key stored in plaintext") + if !reflect.DeepEqual(record.Configuration, validConfiguration()) { + t.Fatalf("bundle: %+v", record.Configuration) } - if _, err := cipher.OpenDeploymentModelProvider(record.Sealed, "claude_code"); err == nil { - t.Fatal("bundle opens for another harness") + storage.replace = func(context.Context, Record) (Configuration, error) { + return Configuration{}, credentialcrypto.ErrUnavailable } - raw, err := cipher.OpenDeploymentModelProvider(record.Sealed, "codex") - var opened v1.ModelConfigurationInput - if err != nil || json.Unmarshal(raw, &opened) != nil || opened.ModelProvider.APIKey != "secret-key" || opened.Model != "fixture-model" { - t.Fatal("sealed bundle incomplete", err) + if _, err := newService(t, storage).Replace(t.Context(), Replacement{Harness: "codex", Configuration: validConfiguration()}); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("missing credential key", err) } } @@ -102,12 +89,9 @@ func TestReplaceRejectsBeforeStorage(t *testing.T) { invalid := validConfiguration() invalid.ModelProvider.Protocol = "anthropic-unknown" var field *v1.ModelProviderError - if _, err := newService(t, &fakeStorage{t: t}, testCipher(t, 1)).Replace(t.Context(), Replacement{Harness: "codex", Configuration: invalid}); !errors.As(err, &field) || field.Param != "protocol" { + if _, err := newService(t, &fakeStorage{t: t}).Replace(t.Context(), Replacement{Harness: "codex", Configuration: invalid}); !errors.As(err, &field) || field.Param != "protocol" { t.Fatal("invalid configuration", err) } - if _, err := newService(t, &fakeStorage{t: t}, nil).Replace(t.Context(), Replacement{Harness: "codex", Configuration: validConfiguration()}); !errors.Is(err, credentialcrypto.ErrUnavailable) { - t.Fatal("missing credential key", err) - } } func TestDeletePassesTheHarness(t *testing.T) { @@ -118,55 +102,37 @@ func TestDeletePassesTheHarness(t *testing.T) { } return failure }} - if err := newService(t, storage, nil).Delete(t.Context(), "codex"); !errors.Is(err, failure) { + if err := newService(t, storage).Delete(t.Context(), "codex"); !errors.Is(err, failure) { t.Fatal(err) } } -func TestResolveOpensTheStoredBundle(t *testing.T) { - cipher := testCipher(t, 1) - seal := func(harness string, configuration any) []byte { - raw, _ := json.Marshal(configuration) - sealed, err := cipher.SealDeploymentModelProvider(raw, harness) - if err != nil { - t.Fatal(err) - } - return sealed - } +func TestResolveValidatesTheOpenedBundle(t *testing.T) { revision := uuid.New() - loaded := func(bundle []byte, err error) *fakeStorage { - return &fakeStorage{t: t, loadSealed: func(_ context.Context, harness string) (Sealed, error) { + loaded := func(configuration v1.ModelConfigurationInput, err error) *fakeStorage { + return &fakeStorage{t: t, loadBundle: func(_ context.Context, harness string) (Bundle, error) { if harness != "codex" { t.Fatal(harness) } - return Sealed{Bundle: bundle, Revision: revision}, err + return Bundle{Configuration: configuration, Revision: revision}, err }} } - snapshot, err := newService(t, loaded(seal("codex", validConfiguration()), nil), cipher).Resolve(t.Context(), "codex") + snapshot, err := newService(t, loaded(validConfiguration(), nil)).Resolve(t.Context(), "codex") if err != nil || snapshot == nil || snapshot.Revision != revision || snapshot.Provider.APIKey != "secret-key" || snapshot.Model != "fixture-model" || string(snapshot.HarnessConfig) != "{}" { t.Fatal("snapshot", snapshot, err) } - if snapshot, err := newService(t, loaded(nil, ErrNotFound), cipher).Resolve(t.Context(), "codex"); snapshot != nil || err != nil { + if snapshot, err := newService(t, loaded(v1.ModelConfigurationInput{}, ErrNotFound)).Resolve(t.Context(), "codex"); snapshot != nil || err != nil { t.Fatal("missing default", snapshot, err) } - failure := errors.New("storage failed") - if _, err := newService(t, loaded(nil, failure), cipher).Resolve(t.Context(), "codex"); !errors.Is(err, failure) { - t.Fatal("storage failure", err) + for _, failure := range []error{errors.New("storage failed"), credentialcrypto.ErrUnavailable} { + if _, err := newService(t, loaded(v1.ModelConfigurationInput{}, failure)).Resolve(t.Context(), "codex"); !errors.Is(err, failure) { + t.Fatal("storage failure", err) + } } invalid := validConfiguration() invalid.Model = "" var field *v1.ModelProviderError - if _, err := newService(t, loaded(seal("codex", invalid), nil), cipher).Resolve(t.Context(), "codex"); !errors.As(err, &field) { + if _, err := newService(t, loaded(invalid, nil)).Resolve(t.Context(), "codex"); !errors.As(err, &field) { t.Fatal("unsupported stored configuration", err) } - for name, service := range map[string]*Service{ - "no key": newService(t, loaded(seal("codex", validConfiguration()), nil), nil), - "wrong key": newService(t, loaded(seal("codex", validConfiguration()), nil), testCipher(t, 2)), - "other harness": newService(t, loaded(seal("claude_code", validConfiguration()), nil), cipher), - "not JSON": newService(t, loaded(seal("codex", "not an object"), nil), cipher), - } { - if _, err := service.Resolve(t.Context(), "codex"); !errors.Is(err, credentialcrypto.ErrUnavailable) { - t.Fatal(name, err) - } - } } diff --git a/services/core/internal/modelconfiguration/storage.go b/services/core/internal/modelconfiguration/storage.go index fcddf24c..b132e80c 100644 --- a/services/core/internal/modelconfiguration/storage.go +++ b/services/core/internal/modelconfiguration/storage.go @@ -2,19 +2,21 @@ package modelconfiguration import "context" -// Storage keeps one deployment default per Harness. Replace and Delete record -// the administrator mutation in the same transaction, from the provenance the -// context carries; an audit failure aborts the change. +// Storage keeps one deployment default per Harness and seals and opens its +// bundle; without a credential key, Replace and LoadBundle return +// credentialcrypto.ErrUnavailable. Replace and Delete record the administrator +// mutation in the same transaction, from the provenance the context carries; +// an audit failure aborts the change. type Storage interface { - // Replace stores record under a new revision and clears the observations - // of the revision it replaces. + // Replace seals and stores record under a new revision and clears the + // observations of the revision it replaces. Replace(ctx context.Context, record Record) (Configuration, error) // Delete removes the Harness's default. Removing a missing default succeeds // and is audited too. Delete(ctx context.Context, harness string) error - // LoadSealed returns the Harness's sealed bundle and its revision, or - // ErrNotFound. - LoadSealed(ctx context.Context, harness string) (Sealed, error) + // LoadBundle opens the Harness's bundle and returns it with its revision. + // A Harness without a default is ErrNotFound, checked before the key. + LoadBundle(ctx context.Context, harness string) (Bundle, error) } // Reader lists the configured defaults without opening any bundle. diff --git a/services/core/internal/persistence/postgres/agentpg/store.go b/services/core/internal/persistence/postgres/agentpg/store.go index 2978328a..a52b917a 100644 --- a/services/core/internal/persistence/postgres/agentpg/store.go +++ b/services/core/internal/persistence/postgres/agentpg/store.go @@ -29,7 +29,8 @@ type Store struct { // New returns the Agent store. cipher is nil when Core has no credential key; // saving or opening a model provider bundle then fails with -// credentialcrypto.ErrUnavailable, as does opening a bundle the key cannot open. +// credentialcrypto.ErrUnavailable. A bundle the key cannot open is an internal +// error. func New(pool *pgunit.Pool, cipher *credentialcrypto.Cipher) *Store { return &Store{pool: pool, cipher: cipher} } @@ -246,13 +247,16 @@ func (s *Store) GetAgentWithModelProvider(ctx context.Context, tenantID, agentID if row.EncryptedConfig == nil { return agent, nil, nil } + if s.cipher == nil { + return agents.Agent{}, nil, credentialcrypto.ErrUnavailable + } raw, err := s.cipher.OpenAgentModelExecution(row.EncryptedConfig, agent.TenantID, agent.ID) if err != nil { - return agents.Agent{}, nil, credentialcrypto.ErrUnavailable + return agents.Agent{}, nil, errors.New("agent model provider decryption failed") } var provider v1.ModelProviderInput if json.Unmarshal(raw, &provider) != nil || provider.Validate() != nil { - return agents.Agent{}, nil, credentialcrypto.ErrUnavailable + return agents.Agent{}, nil, errors.New("invalid stored agent model provider") } return agent, &provider, nil } diff --git a/services/core/internal/persistence/postgres/agentpg/store_test.go b/services/core/internal/persistence/postgres/agentpg/store_test.go index 7e8ad43b..f0f4bdce 100644 --- a/services/core/internal/persistence/postgres/agentpg/store_test.go +++ b/services/core/internal/persistence/postgres/agentpg/store_test.go @@ -355,10 +355,11 @@ func TestAgentModelExecutionAtomicEncryptedSnapshot(t *testing.T) { if _, err := keylessStore.GetAgent(ctx, tenant, agent.ID); err != nil { t.Fatal("plain read required Agent decryption", err) } - for name, s := range map[string]*agentpg.Store{"missing key": keylessStore, "wrong key": agentpg.New(pgunit.NewPool(pool), testCipher(t, 99))} { - if _, _, err := s.GetAgentWithModelProvider(ctx, tenant, agent.ID); !errors.Is(err, credentialcrypto.ErrUnavailable) { - t.Fatalf("%s opened the bundle: %v", name, err) - } + if _, _, err := keylessStore.GetAgentWithModelProvider(ctx, tenant, agent.ID); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("missing key opened the bundle", err) + } + if _, _, err := agentpg.New(pgunit.NewPool(pool), testCipher(t, 99)).GetAgentWithModelProvider(ctx, tenant, agent.ID); err == nil || err.Error() != "agent model provider decryption failed" { + t.Fatal("wrong key was not a decryption failure", err) } if _, err := keyless.Create(ctx, create); !errors.Is(err, credentialcrypto.ErrUnavailable) { t.Fatal("unencrypted Agent create accepted", err) @@ -430,6 +431,34 @@ func TestAgentModelExecutionAtomicEncryptedSnapshot(t *testing.T) { } } +// A bundle sealed to another Agent is a decryption failure, never a missing +// key or a missing bundle. +func TestAgentBundleSealedToAnotherAgentDoesNotOpen(t *testing.T) { + pool := pgtest.Open(t) + c := testCipher(t, 32) + store, service := open(t, pool, c) + ctx, tenant := t.Context(), uuid.NewString() + provider := providerFixture(0) + agent, err := service.Create(ctx, agents.CreateCommand{TenantID: tenant, Configuration: providerConfiguration(t, provider, "codex"), ModelProvider: provider}) + if err != nil { + t.Fatal(err) + } + raw, err := json.Marshal(provider) + if err != nil { + t.Fatal(err) + } + sealed, err := c.SealAgentModelExecution(raw, tenant, uuid.NewString()) + if err != nil { + t.Fatal(err) + } + if _, err := pool.Exec(ctx, "UPDATE agent_model_execution SET encrypted_config=$2 WHERE agent_id=$1", agent.ID, sealed); err != nil { + t.Fatal(err) + } + if _, inherited, err := store.GetAgentWithModelProvider(ctx, tenant, agent.ID); err == nil || err.Error() != "agent model provider decryption failed" || inherited != nil { + t.Fatal("a wrong binding was not a decryption failure", err) + } +} + func TestAgentModelExecutionConcurrentSnapshots(t *testing.T) { pool := pgtest.Open(t) store, service := open(t, pool, testCipher(t, 32)) diff --git a/services/core/internal/persistence/postgres/modelconfigurationpg/fixture_test.go b/services/core/internal/persistence/postgres/modelconfigurationpg/fixture_test.go index 699df747..e054af58 100644 --- a/services/core/internal/persistence/postgres/modelconfigurationpg/fixture_test.go +++ b/services/core/internal/persistence/postgres/modelconfigurationpg/fixture_test.go @@ -37,8 +37,8 @@ func newFixture(t *testing.T) fixture { if err != nil { t.Fatal(err) } - adapter := modelconfigurationpg.New(pgunit.NewPool(pool)) - service, err := modelconfiguration.NewService(adapter, cipher) + adapter := modelconfigurationpg.New(pgunit.NewPool(pool), cipher) + service, err := modelconfiguration.NewService(adapter) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/persistence/postgres/modelconfigurationpg/seal_test.go b/services/core/internal/persistence/postgres/modelconfigurationpg/seal_test.go new file mode 100644 index 00000000..b482438d --- /dev/null +++ b/services/core/internal/persistence/postgres/modelconfigurationpg/seal_test.go @@ -0,0 +1,71 @@ +package modelconfigurationpg_test + +import ( + "encoding/json" + "errors" + "reflect" + "testing" + + v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/modelconfiguration" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/modelconfigurationpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" +) + +// The Store seals the bundle to its Harness as Core always has, so a bundle +// sealed before the Store owned the key still opens. A missing key is +// credentialcrypto.ErrUnavailable, and a wrong binding is a decryption +// failure, never a missing key or default. +func TestBundlesKeepTheirSealedFormat(t *testing.T) { + f := newFixture(t) + keyless := modelconfigurationpg.New(pgunit.NewPool(f.pool), nil) + configuration := v1.ModelConfigurationInput{ModelProvider: fixtureProvider, Model: "fixture"} + if _, err := keyless.LoadBundle(t.Context(), "codex"); !errors.Is(err, modelconfiguration.ErrNotFound) { + t.Fatal("a missing default was not reported before the key", err) + } + record := modelconfiguration.Record{Harness: "codex", Provider: *fixtureProvider.SafeView(), Model: "fixture", HarnessConfig: json.RawMessage(`{}`), Configuration: configuration} + if _, err := keyless.Replace(admin(t), record); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("a keyless Store sealed a bundle", err) + } + if listed := f.list(t); len(listed) != 0 { + t.Fatal("a keyless replacement was stored", listed) + } + if _, err := f.adapter.Replace(admin(t), record); err != nil { + t.Fatal(err) + } + loaded, err := f.adapter.LoadBundle(t.Context(), "codex") + if err != nil || !reflect.DeepEqual(loaded.Configuration, configuration) { + t.Fatal("the bundle did not round-trip", err) + } + if _, err := keyless.LoadBundle(t.Context(), "codex"); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("a keyless Store opened a bundle", err) + } + // write stores a bundle sealed the way Core sealed it before this Store + // owned the key. + write := func(harness string, bundle v1.ModelConfigurationInput) { + t.Helper() + raw, err := json.Marshal(bundle) + if err != nil { + t.Fatal(err) + } + sealed, err := f.cipher.SealDeploymentModelProvider(raw, harness) + if err != nil { + t.Fatal(err) + } + if _, err := f.pool.Exec(t.Context(), "UPDATE deployment_model_providers SET encrypted_config=$1 WHERE harness='codex'", sealed); err != nil { + t.Fatal(err) + } + } + old := configuration + old.ModelProvider.APIKey = "sealed-before-the-move" + write("codex", old) + if loaded, err := f.adapter.LoadBundle(t.Context(), "codex"); err != nil || !reflect.DeepEqual(loaded.Configuration, old) { + t.Fatal("a bundle sealed before the move did not open", err) + } + write("claude_code", old) + if _, err := f.adapter.LoadBundle(t.Context(), "codex"); err == nil || err.Error() != "deployment model configuration decryption failed" { + t.Fatal("a wrong binding was not a decryption failure", err) + } +} diff --git a/services/core/internal/persistence/postgres/modelconfigurationpg/store.go b/services/core/internal/persistence/postgres/modelconfigurationpg/store.go index 74909c6c..1b6db598 100644 --- a/services/core/internal/persistence/postgres/modelconfigurationpg/store.go +++ b/services/core/internal/persistence/postgres/modelconfigurationpg/store.go @@ -14,6 +14,7 @@ import ( "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" + "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/modelconfiguration" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" @@ -29,9 +30,12 @@ const auditResource = "deployment_model_provider" // connection and for the default's row lock. const observationBudget = time.Second -// Store keeps deployment defaults on pooled connections. It never uses the -// execution lease. -type Store struct{ pool *pgunit.Pool } +// Store keeps deployment defaults on pooled connections and seals and opens +// their bundles. It never uses the execution lease. +type Store struct { + pool *pgunit.Pool + cipher *credentialcrypto.Cipher +} var ( _ modelconfiguration.Storage = (*Store)(nil) @@ -39,7 +43,12 @@ var ( _ modelconfiguration.Observer = (*Store)(nil) ) -func New(pool *pgunit.Pool) *Store { return &Store{pool: pool} } +// New returns a Store. Without a credential key (cipher nil), Replace and +// LoadBundle fail with credentialcrypto.ErrUnavailable; List, Delete and +// observations keep working. +func New(pool *pgunit.Pool, cipher *credentialcrypto.Cipher) *Store { + return &Store{pool: pool, cipher: cipher} +} // List reads every default in Harness order without its sealed bundle. func (s *Store) List(ctx context.Context) ([]modelconfiguration.Configuration, error) { @@ -54,16 +63,28 @@ func (s *Store) List(ctx context.Context) ([]modelconfiguration.Configuration, e return result, nil } -// Replace upserts the record under a new revision, which clears the replaced -// revision's observations, and audits the write in the same transaction. +// Replace seals the complete bundle to its Harness, upserts the record under a +// new revision, which clears the replaced revision's observations, and audits +// the write in the same transaction. func (s *Store) Replace(ctx context.Context, record modelconfiguration.Record) (modelconfiguration.Configuration, error) { + if s.cipher == nil { + return modelconfiguration.Configuration{}, credentialcrypto.ErrUnavailable + } + raw, err := json.Marshal(record.Configuration) + if err != nil { + return modelconfiguration.Configuration{}, err + } + sealed, err := s.cipher.SealDeploymentModelProvider(raw, record.Harness) + if err != nil { + return modelconfiguration.Configuration{}, credentialcrypto.ErrUnavailable + } var result modelconfiguration.Configuration - err := s.pool.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + err = s.pool.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := sqlc.New(tx) row, err := q.UpsertDeploymentModelProvider(ctx, sqlc.UpsertDeploymentModelProviderParams{ Harness: record.Harness, Protocol: record.Provider.Protocol, BaseUrl: record.Provider.BaseURL, ContextWindow: record.Provider.ContextWindow, MaxOutputTokens: record.Provider.MaxOutputTokens, - Model: record.Model, HarnessConfig: record.HarnessConfig, EncryptedConfig: record.Sealed, + Model: record.Model, HarnessConfig: record.HarnessConfig, EncryptedConfig: sealed, Revision: pgtype.UUID{Bytes: uuid.New(), Valid: true}, }) if err != nil { @@ -90,17 +111,28 @@ func (s *Store) Delete(ctx context.Context, harness string) error { })) } -// LoadSealed reads the sealed bundle and its revision in one statement, so the -// pair always belongs to the same replacement. -func (s *Store) LoadSealed(ctx context.Context, harness string) (modelconfiguration.Sealed, error) { +// LoadBundle reads the sealed bundle and its revision in one statement, so the +// pair always belongs to the same replacement, and opens the bundle. +func (s *Store) LoadBundle(ctx context.Context, harness string) (modelconfiguration.Bundle, error) { row, err := s.pool.Queries().GetDeploymentModelProviderSecret(ctx, harness) if errors.Is(err, pgx.ErrNoRows) { - return modelconfiguration.Sealed{}, modelconfiguration.ErrNotFound + return modelconfiguration.Bundle{}, modelconfiguration.ErrNotFound + } + if err != nil { + return modelconfiguration.Bundle{}, translate(err) + } + if s.cipher == nil { + return modelconfiguration.Bundle{}, credentialcrypto.ErrUnavailable } + raw, err := s.cipher.OpenDeploymentModelProvider(row.EncryptedConfig, harness) if err != nil { - return modelconfiguration.Sealed{}, translate(err) + return modelconfiguration.Bundle{}, errors.New("deployment model configuration decryption failed") + } + var configuration v1.ModelConfigurationInput + if json.Unmarshal(raw, &configuration) != nil { + return modelconfiguration.Bundle{}, errors.New("invalid stored deployment model configuration") } - return modelconfiguration.Sealed{Bundle: row.EncryptedConfig, Revision: uuid.UUID(row.Revision.Bytes)}, nil + return modelconfiguration.Bundle{Configuration: configuration, Revision: uuid.UUID(row.Revision.Bytes)}, nil } // ObserveDeploymentModelProvider runs one metadata UPDATE in its own pooled diff --git a/services/core/internal/persistence/postgres/vaultpg/audit_test.go b/services/core/internal/persistence/postgres/vaultpg/audit_test.go index 329d2b1f..8e5b5e06 100644 --- a/services/core/internal/persistence/postgres/vaultpg/audit_test.go +++ b/services/core/internal/persistence/postgres/vaultpg/audit_test.go @@ -13,8 +13,6 @@ import ( "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" ) @@ -143,7 +141,7 @@ func prepareAuditMutation(t *testing.T, service *vaults.Service, tenant, name st // 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) + service := newService(t, 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"} { diff --git a/services/core/internal/persistence/postgres/vaultpg/credentials.go b/services/core/internal/persistence/postgres/vaultpg/credentials.go index dbed4d95..478325dd 100644 --- a/services/core/internal/persistence/postgres/vaultpg/credentials.go +++ b/services/core/internal/persistence/postgres/vaultpg/credentials.go @@ -9,6 +9,7 @@ import ( "github.com/jackc/pgx/v5/pgconn" "github.com/jackc/pgx/v5/pgtype" + "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/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" @@ -16,17 +17,24 @@ import ( "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 +// CreateCredential admits the owning Vault in the insert itself, so a missing, +// foreign or malformed 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 { + id, err := pgunit.ParseID(credential.CredentialID) + if err != nil { return vaults.Credential{}, vaults.ErrInvalidInput } + // A malformed Vault ID names none, so the insert finds no Vault. + scope := binding(pgunit.PathID(credential.TenantID), pgunit.PathID(credential.VaultID), id, credential.AuthType, credential.MCPServerURL) + metadata, ciphertext, err := s.sealSecret(scope, credential) + if err != nil { + return vaults.Credential{}, err + } var created vaults.Credential - err := s.write(ctx, func(ctx context.Context, q *sqlc.Queries) error { - row, err := insertCredential(ctx, q, credential) + err = s.write(ctx, func(ctx context.Context, q *sqlc.Queries) error { + row, err := insertCredential(ctx, q, credential, metadata, ciphertext) // 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" { @@ -47,16 +55,29 @@ func (s *Store) CreateCredential(ctx context.Context, credential vaults.NewCrede return created, nil } -func insertCredential(ctx context.Context, q *sqlc.Queries, credential vaults.NewCredential) (sqlc.GetCredentialRow, error) { +// sealSecret seals a new Credential's secret. An mcp_oauth Credential also +// gets the encoded metadata the seal authenticates. +func (s *Store) sealSecret(scope credentialcrypto.Binding, credential vaults.NewCredential) (metadata, ciphertext []byte, err error) { + switch credential.AuthType { + case vaults.AuthStaticBearer: + ciphertext, err = s.sealStatic(scope, credential.Token) + return nil, ciphertext, err + case vaults.AuthMCPOAuth: + return s.sealOAuth(scope, credential.OAuth) + } + return nil, nil, errors.New("unknown credential authentication type") +} + +func insertCredential(ctx context.Context, q *sqlc.Queries, credential vaults.NewCredential, metadata, ciphertext []byte) (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}) + Name: credential.Name, McpServerUrl: credential.MCPServerURL, TokenCiphertext: 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}) + Name: credential.Name, McpServerUrl: credential.MCPServerURL, OauthMetadata: metadata, TokenCiphertext: ciphertext}) return sqlc.GetCredentialRow(row), err } return sqlc.GetCredentialRow{}, errors.New("unknown credential authentication type") @@ -142,11 +163,15 @@ func (s *Store) ListCredentials(ctx context.Context, tenantID, vaultID string, q // 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) { + tenant, vault, id := pgunit.PathID(replacement.TenantID), pgunit.PathID(replacement.VaultID), pgunit.PathID(replacement.CredentialID) + ciphertext, err := s.sealStatic(binding(tenant, vault, id, vaults.AuthStaticBearer, replacement.MCPServerURL), replacement.Token) + if err != nil { + return vaults.Credential{}, err + } var updated vaults.Credential - err := s.write(ctx, 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, + TenantID: tenant, VaultID: vault, ID: id, McpServerUrl: replacement.MCPServerURL, TokenCiphertext: ciphertext, }) 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 fafd9b86..a299087b 100644 --- a/services/core/internal/persistence/postgres/vaultpg/credentials_test.go +++ b/services/core/internal/persistence/postgres/vaultpg/credentials_test.go @@ -23,7 +23,7 @@ func TestStaticCredentialsPersistEncryptedAndRemainScoped(t *testing.T) { store, pool := openStore(t) ctx := t.Context() tenant, foreignTenant := uuid.NewString(), uuid.NewString() - keyless := newService(t, store, nil, nil) + keyless := newService(t, pool, nil, nil) var owned []vaults.Vault for _, owner := range []string{tenant, tenant, foreignTenant} { owned = append(owned, createVault(t, keyless, owner)) @@ -35,7 +35,7 @@ func TestStaticCredentialsPersistEncryptedAndRemainScoped(t *testing.T) { if _, err := rand.Read(randomToken); err != nil { t.Fatal(err) } - service := newService(t, store, newCipher(t, key), nil) + service := newService(t, pool, newCipher(t, key), nil) canary := hex.EncodeToString(randomToken) opaque := " \t" + canary + " 凭据\n" + strings.Repeat("x", 300) + " " tokens := []string{opaque, opaque, ""} @@ -141,7 +141,7 @@ 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) + service := newService(t, pool, 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)) @@ -267,7 +267,7 @@ func TestStaticCredentialUpdatePreservesBindingsAndReplacesCurrentSecret(t *test t.Fatal(err) } cipher := newCipher(t, key) - service := newService(t, store, cipher, nil) + service := newService(t, pool, cipher, nil) var owned []vaults.Vault for _, owner := range []string{tenant, tenant, foreign} { owned = append(owned, createVault(t, service, owner)) @@ -337,19 +337,19 @@ func TestStaticCredentialUpdatePreservesBindingsAndReplacesCurrentSecret(t *test assertUnchanged() } for _, unusable := range []*credentialcrypto.Cipher{nil, {}} { - if _, err := update(newService(t, store, unusable, nil), tenant, original.VaultID, original.ID, "rejected"); err == nil { + if _, err := update(newService(t, pool, 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") + _, updateErr := update(newService(t, readOnlyPool(t, pool), cipher, nil), tenant, original.VaultID, original.ID, "rejected") if !isReadOnlyFailure(updateErr) { t.Fatal("database write failure was accepted or translated", updateErr) } assertUnchanged() // A stale destination from a prior metadata read cannot authorize the write. - _, err = store.ReplaceStaticToken(t.Context(), vaults.StaticTokenReplacement{CredentialKey: vaults.CredentialKey{TenantID: tenant, VaultID: original.VaultID, CredentialID: original.ID}, MCPServerURL: endpoint + "/other", Ciphertext: prior}) + _, err = keyedStore(pool, cipher).ReplaceStaticToken(t.Context(), vaults.StaticTokenReplacement{CredentialKey: vaults.CredentialKey{TenantID: tenant, VaultID: original.VaultID, CredentialID: original.ID}, MCPServerURL: endpoint + "/other", Token: "rejected"}) if !errors.Is(err, vaults.ErrNotFound) { t.Fatal("mutation failed to recheck immutable destination", err) } @@ -363,7 +363,7 @@ func TestStaticCredentialUpdatePreservesBindingsAndReplacesCurrentSecret(t *test } pool.Close() store, pool = openStore(t) - service = newService(t, store, newCipher(t, bytes.Clone(key)), nil) + service = newService(t, pool, newCipher(t, bytes.Clone(key)), nil) for _, binding := range bindings { current, err := bearerToken(t.Context(), service, tenant, attached, binding) if err != nil || current != lastToken { @@ -399,8 +399,8 @@ 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) + service := newService(t, pool, newCipher(t, key), nil) + keyless := newService(t, pool, 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} @@ -422,7 +422,7 @@ 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) + _, deletionErr := remove(newService(t, readOnlyPool(t, pool), nil, nil), tenant, vault.ID, original.ID) if !isReadOnlyFailure(deletionErr) { t.Fatal("failed mutation was accepted or translated", deletionErr) } @@ -444,8 +444,8 @@ func TestCredentialDeletionScopeBindingAndRestart(t *testing.T) { t.Fatal("deleted row or ciphertext remains") } pool.Close() - store, _ = openStore(t) - service = newService(t, store, newCipher(t, bytes.Clone(key)), nil) + store, pool = openStore(t) + service = newService(t, pool, 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") } @@ -472,9 +472,9 @@ func TestCredentialDeletionScopeBindingAndRestart(t *testing.T) { } func TestCredentialDeletionConcurrentReplacementCannotResurrect(t *testing.T) { - store, _ := openStore(t) + store, pool := openStore(t) tenant := uuid.NewString() - service := newService(t, store, newCipher(t, bytes.Repeat([]byte{42}, 32)), nil) + service := newService(t, pool, 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") diff --git a/services/core/internal/persistence/postgres/vaultpg/fixture_test.go b/services/core/internal/persistence/postgres/vaultpg/fixture_test.go index e0fdda7f..6d23ff73 100644 --- a/services/core/internal/persistence/postgres/vaultpg/fixture_test.go +++ b/services/core/internal/persistence/postgres/vaultpg/fixture_test.go @@ -17,15 +17,16 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" ) -// openStore opens the shared test database, as a restarted Core would. +// openStore opens the shared test database, as a restarted Core would. The +// Store has no credential key, which reads never need. func openStore(t *testing.T) (*vaultpg.Store, *pgxpool.Pool) { t.Helper() pool := pgtest.Open(t) - return vaultpg.New(pgunit.NewPool(pool)), pool + return vaultpg.New(pgunit.NewPool(pool), nil), pool } -// readOnlyStore fails every write with a real PostgreSQL error. -func readOnlyStore(t *testing.T, pool *pgxpool.Pool) *vaultpg.Store { +// readOnlyPool fails every write with a real PostgreSQL error. +func readOnlyPool(t *testing.T, pool *pgxpool.Pool) *pgxpool.Pool { t.Helper() config := pool.Config().Copy() config.ConnConfig.RuntimeParams["default_transaction_read_only"] = "on" @@ -34,11 +35,11 @@ func readOnlyStore(t *testing.T, pool *pgxpool.Pool) *vaultpg.Store { t.Fatal(err) } t.Cleanup(readOnly.Close) - return vaultpg.New(pgunit.NewPool(readOnly)) + return readOnly } // isReadOnlyFailure reports the unexpected failure of a write on a -// readOnlyStore, which the Store returns as is. +// readOnlyPool, which the Store returns as is. func isReadOnlyFailure(err error) bool { var failure *pgconn.PgError return errors.As(err, &failure) && failure.Code == "25006" @@ -53,13 +54,19 @@ func newCipher(t *testing.T, key []byte) *credentialcrypto.Cipher { 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 { +// keyedStore is a new Store on pool; a nil cipher is a Core without a +// credential key. +func keyedStore(pool *pgxpool.Pool, cipher *credentialcrypto.Cipher) *vaultpg.Store { + return vaultpg.New(pgunit.NewPool(pool), cipher) +} + +// newService serves a keyedStore, as cmd/server wires it. +func newService(t *testing.T, pool *pgxpool.Pool, cipher *credentialcrypto.Cipher, refresher oauthrefresh.Refresher) *vaults.Service { t.Helper() if refresher == nil { refresher = noRefresh(t) } - service, err := vaults.NewService(storage, cipher, refresher) + service, err := vaults.NewService(keyedStore(pool, cipher), refresher) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/persistence/postgres/vaultpg/oauth.go b/services/core/internal/persistence/postgres/vaultpg/oauth.go index 13af0461..07746fe9 100644 --- a/services/core/internal/persistence/postgres/vaultpg/oauth.go +++ b/services/core/internal/persistence/postgres/vaultpg/oauth.go @@ -2,10 +2,12 @@ package vaultpg import ( "context" + "errors" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" + "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/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" @@ -15,7 +17,7 @@ import ( // 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)} + tx := &oauthTx{store: s, tenantID: key.TenantID, tenant: pgunit.PathID(key.TenantID), vault: pgunit.PathID(key.VaultID), id: pgunit.PathID(key.CredentialID)} return translate(s.pool.Transaction(ctx, func(ctx context.Context, t pgx.Tx) error { tx.q = sqlc.New(t) return apply(tx) @@ -23,31 +25,42 @@ func (s *Store) WithOAuthCredential(ctx context.Context, key vaults.CredentialKe } type oauthTx struct { + store *Store q *sqlc.Queries tenantID string tenant, vault, id pgtype.UUID + // loaded is the locked Credential, which the applied grant is sealed to. + loaded *vaults.Credential } -func (t *oauthTx) LoadOAuthCredential(ctx context.Context) (vaults.Credential, []byte, error) { +func (t *oauthTx) LoadOAuthGrant(ctx context.Context, destination string) (vaults.OAuthGrant, error) { row, err := t.q.GetOAuthCredentialForUpdate(ctx, sqlc.GetOAuthCredentialForUpdateParams{TenantID: t.tenant, VaultID: t.vault, ID: t.id}) if err != nil { - return vaults.Credential{}, nil, translate(err) + return vaults.OAuthGrant{}, translate(err) } credential, err := credentialFromRow(sqlc.GetCredentialRow{ID: row.ID, VaultID: row.VaultID, Name: row.Name, AuthType: row.AuthType, McpServerUrl: row.McpServerUrl, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt, OauthMetadata: row.OauthMetadata}) if err != nil { - return vaults.Credential{}, nil, err + return vaults.OAuthGrant{}, err } - return credential, row.TokenCiphertext, nil + if destination != "" && credential.MCPServerURL != destination { + return vaults.OAuthGrant{}, vaults.ErrNotFound + } + grant, err := t.store.openOAuth(t.scope(credential), credential.OAuth, row.TokenCiphertext) + if err != nil { + return vaults.OAuthGrant{}, err + } + t.loaded = &credential + return grant, nil } -func (t *oauthTx) ApplyOAuthRefresh(ctx context.Context, sealed vaults.SealedOAuth) error { - _, err := t.update(ctx, sealed) +func (t *oauthTx) ApplyOAuthRefresh(ctx context.Context, grant vaults.OAuthGrant) error { + _, err := t.update(ctx, grant) return err } -func (t *oauthTx) ApplyOAuthReplacement(ctx context.Context, sealed vaults.SealedOAuth) (vaults.Credential, error) { - updated, err := t.update(ctx, sealed) +func (t *oauthTx) ApplyOAuthReplacement(ctx context.Context, grant vaults.OAuthGrant) (vaults.Credential, error) { + updated, err := t.update(ctx, grant) if err != nil { return vaults.Credential{}, err } @@ -57,12 +70,25 @@ func (t *oauthTx) ApplyOAuthReplacement(ctx context.Context, sealed vaults.Seale 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) { +// update seals the grant to the loaded Credential and matches the destination +// it is sealed to. +func (t *oauthTx) update(ctx context.Context, grant vaults.OAuthGrant) (vaults.Credential, error) { + if t.loaded == nil { + return vaults.Credential{}, errors.New("OAuth credential was not loaded") + } + metadata, ciphertext, err := t.store.sealOAuth(t.scope(*t.loaded), grant) + if err != nil { + return vaults.Credential{}, err + } 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}) + McpServerUrl: t.loaded.MCPServerURL, OauthMetadata: metadata, TokenCiphertext: ciphertext}) if err != nil { return vaults.Credential{}, translate(err) } return credentialFromRow(sqlc.GetCredentialRow(row)) } + +// scope is the seal scope of the locked Credential's grant. +func (t *oauthTx) scope(credential vaults.Credential) credentialcrypto.Binding { + return binding(t.tenant, pgunit.PathID(credential.VaultID), pgunit.PathID(credential.ID), vaults.AuthMCPOAuth, credential.MCPServerURL) +} diff --git a/services/core/internal/persistence/postgres/vaultpg/oauth_test.go b/services/core/internal/persistence/postgres/vaultpg/oauth_test.go index 41463c88..5b4c47ef 100644 --- a/services/core/internal/persistence/postgres/vaultpg/oauth_test.go +++ b/services/core/internal/persistence/postgres/vaultpg/oauth_test.go @@ -17,7 +17,6 @@ import ( "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" ) @@ -39,7 +38,7 @@ func newOAuthFixture(t *testing.T, refresher oauthrefresh.Refresher) *oauthFixtu t.Helper() store, pool := openStore(t) cipher := newCipher(t, bytes.Repeat([]byte{17}, 32)) - service := newService(t, store, cipher, refresher) + service := newService(t, pool, cipher, refresher) tenant := uuid.NewString() vault := createVault(t, service, tenant) return &oauthFixture{store: store, pool: pool, cipher: cipher, service: service, tenant: tenant, vault: vault, @@ -54,7 +53,7 @@ func newOAuthFixture(t *testing.T, refresher oauthrefresh.Refresher) *oauthFixtu // 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) + return newService(t, f.pool, cipher, refresher) } func (f *oauthFixture) create(t *testing.T, command vaults.CreateOAuthCredential) vaults.Credential { @@ -318,7 +317,7 @@ func TestOAuthRefreshPersistsRotatedGrantAndRequest(t *testing.T) { 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) + restarted := newService(t, restartedPool, 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) diff --git a/services/core/internal/persistence/postgres/vaultpg/seal.go b/services/core/internal/persistence/postgres/vaultpg/seal.go new file mode 100644 index 00000000..58e73e2d --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/seal.go @@ -0,0 +1,95 @@ +package vaultpg + +import ( + "encoding/json" + "errors" + "reflect" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgtype" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/vaults" +) + +// binding scopes a Credential secret to its canonical tenant, Vault and +// Credential IDs, its authentication type and its destination. Every stored +// ciphertext is sealed to exactly these fields. +func binding(tenant, vault, id pgtype.UUID, authType, destination string) credentialcrypto.Binding { + return credentialcrypto.Binding{TenantID: uuid.UUID(tenant.Bytes).String(), VaultID: uuid.UUID(vault.Bytes).String(), + CredentialID: uuid.UUID(id.Bytes).String(), AuthType: authType, Destination: destination} +} + +// oauthPayload is the sealed plaintext of an mcp_oauth Credential. Its sealed +// copy of the metadata authenticates the stored one. +type oauthPayload 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"` +} + +// sealStatic seals a static_bearer token. +func (s *Store) sealStatic(scope credentialcrypto.Binding, token string) ([]byte, error) { + if s.cipher == nil { + return nil, credentialcrypto.ErrUnavailable + } + ciphertext, err := s.cipher.Seal([]byte(token), scope) + if err != nil { + return nil, errors.New("credential encryption failed") + } + return ciphertext, nil +} + +// openStatic opens a static_bearer token. +func (s *Store) openStatic(scope credentialcrypto.Binding, ciphertext []byte) (string, error) { + if s.cipher == nil { + return "", credentialcrypto.ErrUnavailable + } + plaintext, err := s.cipher.Open(ciphertext, scope) + if err != nil { + return "", errors.New("MCP credential decryption failed") + } + return string(plaintext), nil +} + +// sealOAuth encodes the grant's metadata for its column and seals the whole +// grant. +func (s *Store) sealOAuth(scope credentialcrypto.Binding, grant vaults.OAuthGrant) (metadata, ciphertext []byte, err error) { + if s.cipher == nil { + return nil, nil, credentialcrypto.ErrUnavailable + } + metadata, err = json.Marshal(grant.Metadata) + if err != nil { + return nil, nil, errors.New("credential encoding failed") + } + plaintext, err := json.Marshal(oauthPayload{Version: 1, Metadata: grant.Metadata, AccessToken: grant.AccessToken, + RefreshToken: grant.RefreshToken, ClientSecret: grant.ClientSecret}) + if err != nil { + return nil, nil, errors.New("credential encoding failed") + } + ciphertext, err = s.cipher.Seal(plaintext, scope) + if err != nil { + return nil, nil, errors.New("credential encryption failed") + } + return metadata, ciphertext, nil +} + +// openOAuth opens a grant and authenticates the stored metadata against the +// sealed copy. +func (s *Store) openOAuth(scope credentialcrypto.Binding, stored *vaults.OAuthMetadata, ciphertext []byte) (vaults.OAuthGrant, error) { + if s.cipher == nil { + return vaults.OAuthGrant{}, credentialcrypto.ErrUnavailable + } + plaintext, err := s.cipher.Open(ciphertext, scope) + if err != nil { + return vaults.OAuthGrant{}, errors.New("OAuth credential decryption failed") + } + var payload oauthPayload + if json.Unmarshal(plaintext, &payload) != nil || payload.Version != 1 || stored == nil || !reflect.DeepEqual(payload.Metadata, *stored) { + return vaults.OAuthGrant{}, errors.New("OAuth credential authentication failed") + } + return vaults.OAuthGrant{Metadata: payload.Metadata, AccessToken: payload.AccessToken, + RefreshToken: payload.RefreshToken, ClientSecret: payload.ClientSecret}, nil +} diff --git a/services/core/internal/persistence/postgres/vaultpg/seal_test.go b/services/core/internal/persistence/postgres/vaultpg/seal_test.go new file mode 100644 index 00000000..e04ff428 --- /dev/null +++ b/services/core/internal/persistence/postgres/vaultpg/seal_test.go @@ -0,0 +1,125 @@ +package vaultpg_test + +import ( + "bytes" + "encoding/json" + "errors" + "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" +) + +// The Store seals to the canonical binding Core has always used, so secrets +// sealed before the Store owned the key still open. A missing key is +// credentialcrypto.ErrUnavailable, and a wrong binding is a decryption +// failure, never a missing key, row or token. +func TestCredentialSecretsKeepTheirSealedFormat(t *testing.T) { + _, pool := openStore(t) + cipher := newCipher(t, bytes.Repeat([]byte{61}, 32)) + service, keyless := newService(t, pool, cipher, nil), newService(t, pool, nil, nil) + tenant, url := uuid.NewString(), "https://mcp.example/tools" + vault := createVault(t, service, tenant) + expiry := time.Now().Add(time.Hour).UTC().Format(time.RFC3339Nano) + metadata := vaults.OAuthMetadata{ExpiresAt: &expiry, Refresh: &vaults.OAuthRefreshMetadata{ClientID: "client", TokenEndpoint: "https://issuer.example/token", TokenEndpointAuth: "client_secret_basic"}} + createOAuth := func(service *vaults.Service, vaultID string) (vaults.Credential, error) { + return service.CreateOAuthCredential(t.Context(), vaults.CreateOAuthCredential{TenantID: tenant, VaultID: vaultID, Name: "oauth", MCPServerURL: url, + AccessToken: "new-access", OAuth: metadata, RefreshToken: "refresh"}) + } + createToken := func(service *vaults.Service, vaultID string) (vaults.Credential, error) { + return service.CreateStaticCredential(t.Context(), vaults.CreateStaticCredential{TenantID: tenant, VaultID: vaultID, Name: "static", MCPServerURL: url, Token: "new-static"}) + } + // A non-canonical Vault ID seals to the canonical one. + static, err := createToken(service, strings.ToUpper(vault.ID)) + if err != nil { + t.Fatal(err) + } + oauth, err := createOAuth(service, strings.ToUpper(vault.ID)) + if err != nil { + t.Fatal(err) + } + scope := func(credential vaults.Credential) credentialcrypto.Binding { + return credentialcrypto.Binding{TenantID: tenant, VaultID: vault.ID, CredentialID: credential.ID, AuthType: credential.AuthType, Destination: url} + } + var ciphertext []byte + if err := pool.QueryRow(t.Context(), "SELECT token_ciphertext FROM vault_credentials WHERE id=$1", static.ID).Scan(&ciphertext); err != nil { + t.Fatal(err) + } + if plaintext, err := cipher.Open(ciphertext, scope(static)); err != nil || string(plaintext) != "new-static" { + t.Fatal("the token was not sealed to its canonical binding", err) + } + if grant := readGrant(t, pool, cipher, tenant, oauth); grant.AccessToken != "new-access" || grant.RefreshToken != "refresh" { + t.Fatal("the grant was not sealed to its canonical binding") + } + // write stores a secret sealed the way Core sealed it before this Store + // owned the key. + write := func(credential vaults.Credential, plaintext []byte, binding credentialcrypto.Binding) { + t.Helper() + ciphertext, err := cipher.Seal(plaintext, binding) + if err != nil { + t.Fatal(err) + } + if _, err := pool.Exec(t.Context(), "UPDATE vault_credentials SET token_ciphertext=$2 WHERE id=$1", credential.ID, ciphertext); err != nil { + t.Fatal(err) + } + } + legacy, err := json.Marshal(storedGrant{Version: 1, Metadata: metadata, AccessToken: "old-access", RefreshToken: "refresh"}) + if err != nil { + t.Fatal(err) + } + token := func(service *vaults.Service, credential vaults.Credential) (string, error) { + return bearerToken(t.Context(), service, tenant, []string{vault.ID}, oauthBinding(credential)) + } + write(static, []byte("old-static"), scope(static)) + write(oauth, legacy, scope(oauth)) + for _, tc := range []struct { + credential vaults.Credential + want string + }{{static, "old-static"}, {oauth, "old-access"}} { + if got, err := token(service, tc.credential); err != nil || got != tc.want { + t.Fatal("a secret sealed before the move did not open", err) + } + if got, err := token(keyless, tc.credential); !errors.Is(err, credentialcrypto.ErrUnavailable) || got != "" { + t.Fatal("a keyless Store opened a secret", err) + } + } + if _, err := service.UpdateOAuthCredential(t.Context(), vaults.UpdateOAuthCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: oauth.ID, AccessToken: ptr("patched-access")}); err != nil { + t.Fatal("a grant sealed before the move was not replaced", err) + } + if grant := readGrant(t, pool, cipher, tenant, oauth); grant.AccessToken != "patched-access" || grant.RefreshToken != "refresh" { + t.Fatal("the replaced grant lost its material") + } + // Without a key, every write that seals fails first, before a malformed Vault. + for _, create := range []func(*vaults.Service, string) (vaults.Credential, error){createToken, createOAuth} { + if _, err := create(keyless, "not-a-vault"); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("a keyless creation did not fail closed", err) + } + if _, err := create(service, "not-a-vault"); !errors.Is(err, vaults.ErrNotFound) { + t.Fatal("a malformed Vault named one", err) + } + } + if _, err := keyless.UpdateStaticCredential(t.Context(), vaults.UpdateStaticCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: static.ID, Token: "rejected"}); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("a keyless replacement did not fail closed", err) + } + if _, err := keyless.UpdateOAuthCredential(t.Context(), vaults.UpdateOAuthCredential{TenantID: tenant, VaultID: vault.ID, CredentialID: oauth.ID, AccessToken: ptr("rejected")}); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("a keyless OAuth replacement did not fail closed", err) + } + // A secret sealed to another Credential or destination does not open. + wrongID, wrongDestination := scope(static), scope(oauth) + wrongID.CredentialID = oauth.ID + wrongDestination.Destination += "/other" + write(static, []byte("old-static"), wrongID) + write(oauth, legacy, wrongDestination) + for _, tc := range []struct { + credential vaults.Credential + message string + }{{static, "MCP credential decryption failed"}, {oauth, "OAuth credential decryption failed"}} { + if got, err := token(service, tc.credential); err == nil || err.Error() != tc.message || got != "" { + t.Fatal("a wrong binding was not a decryption failure", err) + } + } +} diff --git a/services/core/internal/persistence/postgres/vaultpg/selection.go b/services/core/internal/persistence/postgres/vaultpg/selection.go index f45e2b1d..480bd2fe 100644 --- a/services/core/internal/persistence/postgres/vaultpg/selection.go +++ b/services/core/internal/persistence/postgres/vaultpg/selection.go @@ -37,17 +37,17 @@ func (s *Store) FindMCPCredentials(ctx context.Context, query vaults.MCPCredenti 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) { +// StaticToken is the only read of a static token. Resource reads never select +// ciphertext. +func (s *Store) StaticToken(ctx context.Context, query vaults.StaticTokenQuery) (string, error) { + tenant, vault, id := pgunit.PathID(query.TenantID), pgunit.PathID(query.VaultID), pgunit.PathID(query.CredentialID) 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, + TenantID: tenant, VaultIds: pathIDs(query.VaultIDs), VaultID: vault, CredentialID: id, McpServerUrl: query.MCPServerURL, }) if err != nil { - return nil, translate(err) + return "", translate(err) } - return ciphertext, nil + return s.openStatic(binding(tenant, vault, id, vaults.AuthStaticBearer, query.MCPServerURL), ciphertext) } func pathIDs(ids []string) []pgtype.UUID { diff --git a/services/core/internal/persistence/postgres/vaultpg/selection_test.go b/services/core/internal/persistence/postgres/vaultpg/selection_test.go index 5b6ee8fd..81f0a658 100644 --- a/services/core/internal/persistence/postgres/vaultpg/selection_test.go +++ b/services/core/internal/persistence/postgres/vaultpg/selection_test.go @@ -21,8 +21,8 @@ func TestMCPCredentialSelectionAndScopedDecryption(t *testing.T) { if _, err := rand.Read(key); err != nil { t.Fatal(err) } - service := newService(t, store, newCipher(t, key), nil) - keyless := newService(t, store, nil, nil) + service := newService(t, pool, newCipher(t, key), nil) + keyless := newService(t, pool, nil, nil) var owned []vaults.Vault for _, owner := range []string{tenant, tenant, foreign} { owned = append(owned, createVault(t, service, owner)) @@ -82,8 +82,8 @@ func TestMCPCredentialSelectionAndScopedDecryption(t *testing.T) { } pool.Close() store, pool = openStore(t) - service = newService(t, store, newCipher(t, bytes.Clone(key)), nil) - keyless = newService(t, store, nil, nil) + service = newService(t, pool, newCipher(t, bytes.Clone(key)), nil) + keyless = newService(t, pool, 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) @@ -92,7 +92,7 @@ func TestMCPCredentialSelectionAndScopedDecryption(t *testing.T) { 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) { + if got, err := bearerToken(t.Context(), newService(t, pool, 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){ diff --git a/services/core/internal/persistence/postgres/vaultpg/vaultpg.go b/services/core/internal/persistence/postgres/vaultpg/vaultpg.go index b77115e4..f8863309 100644 --- a/services/core/internal/persistence/postgres/vaultpg/vaultpg.go +++ b/services/core/internal/persistence/postgres/vaultpg/vaultpg.go @@ -10,6 +10,7 @@ import ( "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" + "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/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" @@ -18,14 +19,23 @@ import ( "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 } +// Store runs on pooled connections and seals and opens the Credential +// secrets it stores. 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 + cipher *credentialcrypto.Cipher +} var _ vaults.Storage = (*Store)(nil) -func New(pool *pgunit.Pool) *Store { return &Store{pool: pool} } +// New returns a Store. Without a credential key (cipher nil), operations that +// seal or open a secret fail with credentialcrypto.ErrUnavailable; Vaults, +// Credential metadata, selection and deletion keep working. +func New(pool *pgunit.Pool, cipher *credentialcrypto.Cipher) *Store { + return &Store{pool: pool, cipher: cipher} +} // 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 { diff --git a/services/core/internal/persistence/postgres/vaultpg/vaults_test.go b/services/core/internal/persistence/postgres/vaultpg/vaults_test.go index cfe6d6ce..f9528265 100644 --- a/services/core/internal/persistence/postgres/vaultpg/vaults_test.go +++ b/services/core/internal/persistence/postgres/vaultpg/vaults_test.go @@ -14,13 +14,12 @@ import ( "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) + service := newService(t, pool, nil, nil) ctx := t.Context() tenantA, tenantB := uuid.NewString(), uuid.NewString() before := time.Now().Add(-time.Second) @@ -68,7 +67,7 @@ func TestVaultsPersistAndStayTenantScoped(t *testing.T) { func TestVaultsRejectInvalidInputWithoutWrites(t *testing.T) { store, pool := openStore(t) - service := newService(t, store, nil, nil) + service := newService(t, pool, nil, nil) ctx := t.Context() tenant := uuid.NewString() for _, name := range []string{"", strings.Repeat("x", 257), strings.Repeat("é", 129), string([]byte{0xff})} { @@ -98,7 +97,7 @@ func TestVaultsRejectInvalidInputWithoutWrites(t *testing.T) { func TestVaultListFilteringPaginationAndReconnect(t *testing.T) { store, pool := openStore(t) - service := newService(t, store, nil, nil) + service := newService(t, pool, nil, nil) ctx := t.Context() tenant, other := uuid.NewString(), uuid.NewString() empty, err := store.ListVaults(ctx, tenant, vaults.PageQuery{Limit: 20}) @@ -206,8 +205,8 @@ 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) + service := newService(t, pool, newCipher(t, key), nil) + keyless := newService(t, pool, 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} @@ -225,7 +224,7 @@ func TestVaultDeletionCascadeBindingAndRestart(t *testing.T) { 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}) + _, deletionErr := newService(t, readOnlyPool(t, pool), nil, nil).DeleteVault(t.Context(), vaults.DeleteVault{TenantID: tenant, VaultID: vault.ID}) if !isReadOnlyFailure(deletionErr) { t.Fatal("failed mutation was accepted or translated", deletionErr) } @@ -268,8 +267,8 @@ func TestVaultDeletionCascadeBindingAndRestart(t *testing.T) { t.Fatal("committed parent or encrypted children remain", err) } pool.Close() - store, _ = openStore(t) - service = newService(t, store, newCipher(t, bytes.Clone(key)), nil) + store, pool = openStore(t) + service = newService(t, pool, 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") @@ -304,12 +303,12 @@ func TestVaultDeletionCascadeBindingAndRestart(t *testing.T) { } func TestVaultDeletionConcurrentChildMutations(t *testing.T) { - store, pool := openStore(t) + _, pool := openStore(t) tenant := uuid.NewString() - service := newService(t, store, newCipher(t, bytes.Repeat([]byte{44}, 32)), nil) + service := newService(t, pool, 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) + return newService(t, pool, newCipher(t, bytes.Repeat([]byte{44}, 32)), nil) } for range 8 { vault := createVault(t, service, tenant) diff --git a/services/core/internal/store/archive_cancellation_test.go b/services/core/internal/store/archive_cancellation_test.go index 80a50531..28b8ffc9 100644 --- a/services/core/internal/store/archive_cancellation_test.go +++ b/services/core/internal/store/archive_cancellation_test.go @@ -99,7 +99,7 @@ func TestArchiveWaitingCancellationReceipts(t *testing.T) { t.Fatal(err) } t.Cleanup(func() { conn.Close() }) - h := &dispatchHarness{t: t, s: s, db: db, lease: leased.Lease, tenant: project.TenantID, session: session, conn: conn, registry: registry, d: &execution.Dispatcher{Store: writer, Registry: registry, Observer: modelconfigurationpg.New(pgunit.NewPool(db.pool))}} + h := &dispatchHarness{t: t, s: s, db: db, lease: leased.Lease, tenant: project.TenantID, session: session, conn: conn, registry: registry, d: &execution.Dispatcher{Store: writer, Registry: registry, Observer: modelconfigurationpg.New(pgunit.NewPool(db.pool), db.cipher)}} capabilities := workerEnvironmentCapabilities() capabilities.FunctionTools = proto.CapabilitySupported h.write("", proto.TypeHeartbeat, proto.HeartbeatPayload{SupportedAgentKinds: []proto.SupportedAgentKind{{Kind: "codex", Available: true, Capabilities: capabilities}}}) diff --git a/services/core/internal/store/deployment_model_providers_http_test.go b/services/core/internal/store/deployment_model_providers_http_test.go index b8c3b417..6a180761 100644 --- a/services/core/internal/store/deployment_model_providers_http_test.go +++ b/services/core/internal/store/deployment_model_providers_http_test.go @@ -433,7 +433,7 @@ func TestDeploymentProviderResolutionFixtureIsolation(t *testing.T) { // as the Core routes do. func deploymentDefaults(t *testing.T, db fixtureDB) *modelconfiguration.Service { t.Helper() - service, err := modelconfiguration.NewService(modelconfigurationpg.New(pgunit.NewPool(db.pool)), db.cipher) + service, err := modelconfiguration.NewService(modelconfigurationpg.New(pgunit.NewPool(db.pool), db.cipher)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/dispatch_test.go b/services/core/internal/store/dispatch_test.go index ca0dc023..a1ad3777 100644 --- a/services/core/internal/store/dispatch_test.go +++ b/services/core/internal/store/dispatch_test.go @@ -114,7 +114,7 @@ func newDispatchHarnessForSession(t *testing.T, configuration []byte, local bool } time.Sleep(10 * time.Millisecond) } - h.d = &execution.Dispatcher{Store: s, Registry: h.registry, Observer: modelconfigurationpg.New(pgunit.NewPool(db.pool))} + h.d = &execution.Dispatcher{Store: s, Registry: h.registry, Observer: modelconfigurationpg.New(pgunit.NewPool(db.pool), db.cipher)} return h } diff --git a/services/core/internal/store/fixture_db_test.go b/services/core/internal/store/fixture_db_test.go index 81ad3077..a60f21d9 100644 --- a/services/core/internal/store/fixture_db_test.go +++ b/services/core/internal/store/fixture_db_test.go @@ -64,7 +64,7 @@ func startWorkerErr(ctx context.Context, db fixtureDB, dispatcher *execution.Dis } owned := *dispatcher owned.Credentials = credentials - owned.Observer = modelconfigurationpg.New(pgunit.NewPool(db.pool)) + owned.Observer = modelconfigurationpg.New(pgunit.NewPool(db.pool), db.cipher) return execution.StartWorker(ctx, &owned, execution.Owner{Lease: lease, Store: store.NewExecution(dispatcher.Store, lease)}) } diff --git a/services/core/internal/store/public_handler_fixture_test.go b/services/core/internal/store/public_handler_fixture_test.go index 48a56841..01566aa6 100644 --- a/services/core/internal/store/public_handler_fixture_test.go +++ b/services/core/internal/store/public_handler_fixture_test.go @@ -58,8 +58,8 @@ func publicHandler(t testing.TB, s *store.Store, db fixtureDB, keys fixtureKeyRe if err != nil { return nil, err } - modelConfigurationStore := modelconfigurationpg.New(pgunit.NewPool(db.pool)) - modelConfigurationService, err := modelconfiguration.NewService(modelConfigurationStore, db.cipher) + modelConfigurationStore := modelconfigurationpg.New(pgunit.NewPool(db.pool), db.cipher) + modelConfigurationService, err := modelconfiguration.NewService(modelConfigurationStore) if err != nil { return nil, err } diff --git a/services/core/internal/store/session_creation_identity_test.go b/services/core/internal/store/session_creation_identity_test.go index cab1fcc5..3f31b416 100644 --- a/services/core/internal/store/session_creation_identity_test.go +++ b/services/core/internal/store/session_creation_identity_test.go @@ -163,8 +163,8 @@ func TestSessionCreationKeepsItsResolvedDeploymentRevision(t *testing.T) { if err != nil { t.Fatal(err) } - defaults := modelconfigurationpg.New(pgunit.NewPool(pool)) - service, err := modelconfiguration.NewService(defaults, cipher) + defaults := modelconfigurationpg.New(pgunit.NewPool(pool), cipher) + service, err := modelconfiguration.NewService(defaults) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/vaults_fixture_test.go b/services/core/internal/store/vaults_fixture_test.go index ef742202..17a110e8 100644 --- a/services/core/internal/store/vaults_fixture_test.go +++ b/services/core/internal/store/vaults_fixture_test.go @@ -14,7 +14,7 @@ func fixtureVaults(db fixtureDB) (*vaultpg.Store, *vaults.Service, error) { if err != nil { return nil, nil, err } - vaultStore := vaultpg.New(pgunit.NewPool(db.pool)) - vaultService, err := vaults.NewService(vaultStore, db.cipher, refresher) + vaultStore := vaultpg.New(pgunit.NewPool(db.pool), db.cipher) + vaultService, err := vaults.NewService(vaultStore, refresher) return vaultStore, vaultService, err } diff --git a/services/core/internal/vaults/doc.go b/services/core/internal/vaults/doc.go index 87ec6298..431c0e53 100644 --- a/services/core/internal/vaults/doc.go +++ b/services/core/internal/vaults/doc.go @@ -1,5 +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 owns Vaults and their Credentials: the resources, OAuth +// access-token refresh, and the MCP credential selection that Session creation +// freezes and execution resolves into a bearer token. Storage seals and opens +// the Credential secrets. package vaults diff --git a/services/core/internal/vaults/oauth.go b/services/core/internal/vaults/oauth.go index d455ef48..9dd1a546 100644 --- a/services/core/internal/vaults/oauth.go +++ b/services/core/internal/vaults/oauth.go @@ -4,7 +4,6 @@ import ( "errors" "time" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" ) @@ -18,16 +17,16 @@ type OAuthRefreshUpdate struct { 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 +// OAuthGrant is an mcp_oauth Credential's secret with the metadata it is +// refreshed by. Storage keeps the metadata readable and seals the whole grant, +// so the sealed copy authenticates every public setting used for refresh, +// including the token endpoint, and 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"` +type OAuthGrant struct { + Metadata OAuthMetadata + AccessToken string + RefreshToken string + ClientSecret string } func validOAuthMetadata(metadata OAuthMetadata) bool { @@ -62,17 +61,10 @@ func validOAuthCreation(command CreateOAuthCredential) bool { 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) { +func applyOAuthUpdate(secret OAuthGrant, update UpdateOAuthCredential) (OAuthGrant, error) { if update.AccessToken != nil { secret.AccessToken = *update.AccessToken secret.Metadata.ExpiresAt = nil @@ -85,15 +77,15 @@ func applyOAuthUpdate(secret oauthSecret, update UpdateOAuthCredential) (oauthSe return secret, nil } if secret.Metadata.Refresh == nil { - return oauthSecret{}, ErrInvalidInput + return OAuthGrant{}, ErrInvalidInput } refresh := *secret.Metadata.Refresh if patch.TokenEndpointAuthType != "" && patch.TokenEndpointAuthType != refresh.TokenEndpointAuth { - return oauthSecret{}, ErrInvalidInput + return OAuthGrant{}, ErrInvalidInput } if patch.ClientSecret != nil { if refresh.TokenEndpointAuth == "none" { - return oauthSecret{}, ErrInvalidInput + return OAuthGrant{}, ErrInvalidInput } secret.ClientSecret = *patch.ClientSecret } @@ -110,7 +102,7 @@ func applyOAuthUpdate(secret oauthSecret, update UpdateOAuthCredential) (oauthSe // 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) { +func currentAccessToken(secret OAuthGrant, 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 { @@ -127,7 +119,7 @@ func currentAccessToken(secret oauthSecret, now time.Time) (token string, expire } // refreshRequest is the exchange that renews an expired grant. -func refreshRequest(secret oauthSecret) (oauthrefresh.Request, error) { +func refreshRequest(secret OAuthGrant) (oauthrefresh.Request, error) { refresh := secret.Metadata.Refresh if refresh == nil || secret.RefreshToken == "" { return oauthrefresh.Request{}, errors.New("expired OAuth credential cannot be refreshed") @@ -141,9 +133,9 @@ func refreshRequest(secret oauthSecret) (oauthrefresh.Request, error) { // 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) { +func applyRefreshedToken(secret OAuthGrant, token oauthrefresh.Token, now time.Time) (OAuthGrant, error) { if token.AccessToken == "" || token.ExpiresAt != nil && !now.Before(*token.ExpiresAt) { - return oauthSecret{}, errors.New("OAuth refresh returned an unusable token") + return OAuthGrant{}, errors.New("OAuth refresh returned an unusable token") } secret.AccessToken = token.AccessToken if token.RefreshToken != "" { diff --git a/services/core/internal/vaults/rules_test.go b/services/core/internal/vaults/rules_test.go index e6f9110f..da4843a5 100644 --- a/services/core/internal/vaults/rules_test.go +++ b/services/core/internal/vaults/rules_test.go @@ -185,7 +185,7 @@ func TestOAuthCreationRules(t *testing.T) { func TestOAuthUpdateRules(t *testing.T) { expiry := "2030-01-02T03:04:05Z" - stored := oauthSecret{Version: 1, AccessToken: "access", RefreshToken: "refresh", ClientSecret: "secret", + stored := OAuthGrant{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) { @@ -208,7 +208,7 @@ func TestOAuthUpdateRules(t *testing.T) { withoutRefresh := stored withoutRefresh.Metadata.Refresh = nil for _, tc := range []struct { - secret oauthSecret + secret OAuthGrant patch OAuthRefreshUpdate }{ {withoutRefresh, OAuthRefreshUpdate{RefreshToken: ptr("cannot-add")}}, @@ -238,19 +238,19 @@ func TestOAuthAccessTokenAndRefreshRules(t *testing.T) { {nil, "", "", false, true}, {ptr("not-a-date"), "access", "", false, true}, } { - token, expired, err := currentAccessToken(oauthSecret{AccessToken: tc.access, Metadata: OAuthMetadata{ExpiresAt: tc.expiresAt}}, now) + token, expired, err := currentAccessToken(OAuthGrant{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{ + secret := OAuthGrant{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}} { + for _, unusable := range []OAuthGrant{{RefreshToken: "refresh"}, {Metadata: secret.Metadata}} { if _, err := refreshRequest(unusable); err == nil { t.Fatal("a grant without refresh configuration or token was refreshed") } diff --git a/services/core/internal/vaults/service.go b/services/core/internal/vaults/service.go index 2c2d1b3d..d314066f 100644 --- a/services/core/internal/vaults/service.go +++ b/services/core/internal/vaults/service.go @@ -2,15 +2,12 @@ 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" ) @@ -20,24 +17,23 @@ import ( // has its own tighter bound. const oauthRefreshTimeout = 20 * time.Second -// Service runs the Vault and Credential operations. +// Service runs the Vault and Credential operations. Storage seals and opens +// the secrets; without a credential key, operations that need a secret return +// credentialcrypto.ErrUnavailable and the others keep working. 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) { +// NewService requires storage and the OAuth refresher. +func NewService(storage Storage, 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 + return &Service{storage: storage, refresher: refresher}, nil } type CreateVault struct { @@ -72,8 +68,8 @@ type CreateStaticCredential struct { Name, MCPServerURL, Token string } -// CreateStaticCredential seals a bearer token to its tenant, Vault, new -// Credential ID and destination, and stores it. +// CreateStaticCredential stores a bearer token under a new Credential ID. +// Storage seals it to its tenant, Vault, that ID and destination. func (s *Service) CreateStaticCredential(ctx context.Context, command CreateStaticCredential) (Credential, error) { tenant, ok := canonicalID(command.TenantID) if !ok { @@ -82,22 +78,11 @@ func (s *Service) CreateStaticCredential(ctx context.Context, command CreateStat 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") - } + // Storage reports a missing credential key before a malformed or missing + // Vault. + key := CredentialKey{TenantID: tenant, VaultID: command.VaultID, CredentialID: uuid.NewString()} return s.storage.CreateCredential(ctx, NewCredential{CredentialKey: key, Name: command.Name, - AuthType: AuthStaticBearer, MCPServerURL: command.MCPServerURL, Ciphertext: ciphertext}) + AuthType: AuthStaticBearer, MCPServerURL: command.MCPServerURL, Token: command.Token}) } type UpdateStaticCredential struct { @@ -106,7 +91,7 @@ type UpdateStaticCredential struct { } // UpdateStaticCredential replaces only the token and the update time. The -// stored metadata supplies the seal's scope, which the write checks again, so +// stored destination is 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) @@ -120,15 +105,7 @@ func (s *Service) UpdateStaticCredential(ctx context.Context, command UpdateStat 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}) + return s.storage.ReplaceStaticToken(ctx, StaticTokenReplacement{CredentialKey: key, MCPServerURL: current.MCPServerURL, Token: command.Token}) } // CreateOAuthCredential keeps the write-only secrets apart from the metadata. @@ -139,8 +116,9 @@ type CreateOAuthCredential struct { 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. +// CreateOAuthCredential stores a grant, with the metadata it is refreshed by, +// under a new Credential ID. Storage seals it to its tenant, Vault, that ID +// and destination. func (s *Service) CreateOAuthCredential(ctx context.Context, command CreateOAuthCredential) (Credential, error) { tenant, ok := canonicalID(command.TenantID) if !ok { @@ -149,22 +127,11 @@ func (s *Service) CreateOAuthCredential(ctx context.Context, command CreateOAuth 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}) + // Storage reports a missing credential key before a malformed or missing + // Vault. + key := CredentialKey{TenantID: tenant, VaultID: command.VaultID, CredentialID: uuid.NewString()} + return s.storage.CreateCredential(ctx, NewCredential{CredentialKey: key, Name: command.Name, AuthType: AuthMCPOAuth, MCPServerURL: command.MCPServerURL, + OAuth: OAuthGrant{Metadata: command.OAuth, AccessToken: command.AccessToken, RefreshToken: command.RefreshToken, ClientSecret: command.ClientSecret}}) } // UpdateOAuthCredential patches an OAuth grant. ExpiresAtSet keeps omitted @@ -192,16 +159,15 @@ func (s *Service) UpdateOAuthCredential(ctx context.Context, command UpdateOAuth 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) + err = s.withOAuth(ctx, key, "", "credential update failed", func(tx OAuthTx, grant OAuthGrant) error { + grant, err := applyOAuthUpdate(grant, command) if err != nil { return err } - sealed, err := s.sealOAuth(key.TenantID, credential, secret) - if err != nil { - return err + if !validOAuthMetadata(grant.Metadata) { + return ErrInvalidInput } - updated, err = tx.ApplyOAuthReplacement(ctx, sealed) + updated, err = tx.ApplyOAuthReplacement(ctx, grant) return err }) if err != nil { @@ -283,33 +249,21 @@ func (s *Service) MCPBearerToken(ctx context.Context, command MCPBearerToken) (s } 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, + return s.storage.StaticToken(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()) + err := s.withOAuth(ctx, key, destination, "OAuth credential refresh commit failed", func(tx OAuthTx, grant OAuthGrant) error { + token, expired, err := currentAccessToken(grant, time.Now()) if err != nil || !expired { bearer = token return err } - request, err := refreshRequest(secret) + request, err := refreshRequest(grant) if err != nil { return err } @@ -317,18 +271,17 @@ func (s *Service) oauthBearerToken(ctx context.Context, key CredentialKey, desti if err != nil { return errors.New("OAuth credential refresh failed") } - secret, err = applyRefreshedToken(secret, refreshed, time.Now()) + grant, err = applyRefreshedToken(grant, refreshed, time.Now()) if err != nil { return err } - sealed, err := s.sealOAuth(key.TenantID, credential, secret) - if err != nil { - return err + if !validOAuthMetadata(grant.Metadata) { + return ErrInvalidInput } - if err := tx.ApplyOAuthRefresh(ctx, sealed); err != nil { + if err := tx.ApplyOAuthRefresh(ctx, grant); err != nil { return err } - bearer = secret.AccessToken + bearer = grant.AccessToken return nil }) if err != nil { @@ -340,21 +293,19 @@ func (s *Service) oauthBearerToken(ctx context.Context, key CredentialKey, desti // 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 { +// cascade. A non-empty destination must match the stored one. The opened +// grant's metadata must still be valid, so a changed stored setting never +// reaches a provider. 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, OAuthGrant) 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) + grant, err := tx.LoadOAuthGrant(ctx, destination) + if err == nil && !validOAuthMetadata(grant.Metadata) { + err = errors.New("OAuth credential authentication failed") } if err == nil { - err = apply(tx, credential, secret) + err = apply(tx, grant) } applied = err return err @@ -365,46 +316,6 @@ func (s *Service) withOAuth(ctx context.Context, key CredentialKey, destination, 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) diff --git a/services/core/internal/vaults/service_test.go b/services/core/internal/vaults/service_test.go index 6f67e77a..4b0e49d9 100644 --- a/services/core/internal/vaults/service_test.go +++ b/services/core/internal/vaults/service_test.go @@ -1,9 +1,7 @@ package vaults import ( - "bytes" "context" - "encoding/json" "errors" "reflect" "strings" @@ -18,20 +16,20 @@ import ( // 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) + 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) + staticToken func(context.Context, StaticTokenQuery) (string, error) } func (f *fakeStorage) unexpected(method string) { @@ -123,40 +121,40 @@ func (f *fakeStorage) FindMCPCredentials(ctx context.Context, query MCPCredentia return f.findMCPCredentials(ctx, query) } -func (f *fakeStorage) StaticTokenCiphertext(ctx context.Context, query StaticTokenQuery) ([]byte, error) { - if f.staticTokenCiphertxt == nil { - f.unexpected("StaticTokenCiphertext") +func (f *fakeStorage) StaticToken(ctx context.Context, query StaticTokenQuery) (string, error) { + if f.staticToken == nil { + f.unexpected("StaticToken") } - return f.staticTokenCiphertxt(ctx, query) + return f.staticToken(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) + load func(context.Context, string) (OAuthGrant, error) + refresh func(context.Context, OAuthGrant) error + replacement func(context.Context, OAuthGrant) (Credential, error) } -func (f *fakeOAuthTx) LoadOAuthCredential(ctx context.Context) (Credential, []byte, error) { +func (f *fakeOAuthTx) LoadOAuthGrant(ctx context.Context, destination string) (OAuthGrant, error) { if f.load == nil { - f.t.Fatal("unexpected call to LoadOAuthCredential") + f.t.Fatal("unexpected call to LoadOAuthGrant") } - return f.load(ctx) + return f.load(ctx, destination) } -func (f *fakeOAuthTx) ApplyOAuthRefresh(ctx context.Context, sealed SealedOAuth) error { +func (f *fakeOAuthTx) ApplyOAuthRefresh(ctx context.Context, grant OAuthGrant) error { if f.refresh == nil { f.t.Fatal("unexpected call to ApplyOAuthRefresh") } - return f.refresh(ctx, sealed) + return f.refresh(ctx, grant) } -func (f *fakeOAuthTx) ApplyOAuthReplacement(ctx context.Context, sealed SealedOAuth) (Credential, error) { +func (f *fakeOAuthTx) ApplyOAuthReplacement(ctx context.Context, grant OAuthGrant) (Credential, error) { if f.replacement == nil { f.t.Fatal("unexpected call to ApplyOAuthReplacement") } - return f.replacement(ctx, sealed) + return f.replacement(ctx, grant) } type refreshFunc func(context.Context, oauthrefresh.Request) (oauthrefresh.Token, error) @@ -172,19 +170,10 @@ func unexpectedRefresh(t testing.TB) refreshFunc { } } -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 { +func testService(t testing.TB, storage *fakeStorage, refresher oauthrefresh.Refresher) *Service { t.Helper() storage.t = t - service, err := NewService(storage, cipher, refresher) + service, err := NewService(storage, refresher) if err != nil { t.Fatal(err) } @@ -192,19 +181,16 @@ func testService(t testing.TB, storage *fakeStorage, cipher *credentialcrypto.Ci } func TestNewServiceRequiresStorageAndRefresher(t *testing.T) { - if _, err := NewService(nil, nil, unexpectedRefresh(t)); err == nil { + if _, err := NewService(nil, unexpectedRefresh(t)); err == nil { t.Fatal("missing storage accepted") } - if _, err := NewService(&fakeStorage{t: t}, nil, nil); err == nil { + if _, err := NewService(&fakeStorage{t: t}, 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)) + service := testService(t, &fakeStorage{}, unexpectedRefresh(t)) long := strings.Repeat("x", 257) for _, command := range []CreateVault{ {TenantID: uuid.NewString(), Name: &long}, @@ -221,14 +207,15 @@ func TestCreateVaultValidatesBeforeStorage(t *testing.T) { t.Fatalf("unexpected new Vault %+v", vault) } return stored, nil - }}, nil, unexpectedRefresh(t)) + }}, 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. +// Credential creation checks the body before storage, and passes the Vault ID +// through, so storage reports a missing credential key before a malformed or +// missing Vault. func TestCredentialCreationValidationOrder(t *testing.T) { tenant, url := uuid.NewString(), "https://mcp.example/tools" create := map[string]func(*Service, string, string) error{ @@ -242,41 +229,49 @@ func TestCredentialCreationValidationOrder(t *testing.T) { }, } 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) + if err := run(testService(t, &fakeStorage{}, unexpectedRefresh(t)), "", "not-a-vault"); !errors.Is(err, ErrInvalidInput) { + t.Fatalf("%s creation: an invalid body reached storage: %v", kind, err) + } + service := testService(t, &fakeStorage{createCredential: func(_ context.Context, credential NewCredential) (Credential, error) { + if credential.VaultID != "not-a-vault" { + t.Fatalf("%s creation changed the Vault ID %q", kind, credential.VaultID) } + return Credential{}, credentialcrypto.ErrUnavailable + }}, unexpectedRefresh(t)) + if err := run(service, "valid", "not-a-vault"); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatalf("%s creation: got %v", kind, err) } } } -func TestCreateStaticCredentialSealsToCanonicalIDs(t *testing.T) { +func TestCreateCredentialsPassPlaintextUnderANewID(t *testing.T) { tenant, vault, url := uuid.New(), uuid.New(), "https://mcp.example/tools" - cipher := testCipher(t) - var stored NewCredential + var stored []NewCredential service := testService(t, &fakeStorage{createCredential: func(_ context.Context, credential NewCredential) (Credential, error) { - stored = credential + stored = append(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 { + }}, unexpectedRefresh(t)) + created, err := service.CreateStaticCredential(t.Context(), CreateStaticCredential{TenantID: strings.ToUpper(tenant.String()), VaultID: vault.String(), Name: "static", MCPServerURL: url, Token: "private-token"}) + if err != nil || created.ID != stored[0].CredentialID { + t.Fatal(err) + } + metadata := OAuthMetadata{Refresh: &OAuthRefreshMetadata{ClientID: "client", TokenEndpoint: "https://issuer.example/token", TokenEndpointAuth: "client_secret_post"}} + if _, err := service.CreateOAuthCredential(t.Context(), CreateOAuthCredential{TenantID: tenant.String(), VaultID: vault.String(), Name: "oauth", MCPServerURL: url, + AccessToken: "access", OAuth: metadata, RefreshToken: "refresh", ClientSecret: "client-secret"}); err != nil { 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) + want := []NewCredential{ + {CredentialKey: CredentialKey{tenant.String(), vault.String(), stored[0].CredentialID}, Name: "static", AuthType: AuthStaticBearer, MCPServerURL: url, Token: "private-token"}, + {CredentialKey: CredentialKey{tenant.String(), vault.String(), stored[1].CredentialID}, Name: "oauth", AuthType: AuthMCPOAuth, MCPServerURL: url, + OAuth: OAuthGrant{Metadata: metadata, AccessToken: "access", RefreshToken: "refresh", ClientSecret: "client-secret"}}, } - 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) + if !reflect.DeepEqual(stored, want) || stored[0].CredentialID == stored[1].CredentialID { + t.Fatalf("unexpected stored Credentials %+v", stored) + } + for _, credential := range stored { + if id, ok := canonicalID(credential.CredentialID); !ok || id != credential.CredentialID { + t.Fatal("the new Credential ID is not a canonical UUID", credential.CredentialID) + } } } @@ -284,7 +279,7 @@ 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) { + if _, err := testService(t, &fakeStorage{}, unexpectedRefresh(t)).UpdateStaticCredential(t.Context(), malformed); !errors.Is(err, ErrNotFound) { t.Fatal("a malformed ID named a Credential", err) } } @@ -293,22 +288,17 @@ func TestUpdateStaticCredential(t *testing.T) { 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) { + if _, err := testService(t, &fakeStorage{getCredential: current(AuthMCPOAuth)}, 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) + if replacement != (StaticTokenReplacement{CredentialKey: CredentialKey{tenant, vault, id}, MCPServerURL: url, Token: "replacement"}) { + t.Fatalf("unexpected replacement %+v", replacement) } - return Credential{ID: id}, nil - }}, cipher, unexpectedRefresh(t)) - if _, err := service.UpdateStaticCredential(t.Context(), command); err != nil { - t.Fatal(err) + return Credential{}, credentialcrypto.ErrUnavailable + }}, unexpectedRefresh(t)) + if _, err := service.UpdateStaticCredential(t.Context(), command); !errors.Is(err, credentialcrypto.ErrUnavailable) { + t.Fatal("the storage result was not returned", err) } } @@ -325,7 +315,7 @@ func TestResolveMCPCredentials(t *testing.T) { return count, nil } } - if _, err := testService(t, &fakeStorage{countOwnedVaults: owned(0)}, nil, unexpectedRefresh(t)).ResolveMCPCredentials(t.Context(), command); !errors.Is(err, ErrNotFound) { + if _, err := testService(t, &fakeStorage{countOwnedVaults: owned(0)}, 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) { @@ -336,7 +326,7 @@ func TestResolveMCPCredentials(t *testing.T) { return nil, nil } return []MCPCredentialMatch{{VaultID: vault, CredentialID: id, AuthType: AuthStaticBearer, MCPServerURL: url}}, nil - }}, nil, unexpectedRefresh(t)) + }}, 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) { @@ -346,42 +336,38 @@ func TestResolveMCPCredentials(t *testing.T) { 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) { + lookup := func(token string, err error) func(context.Context, StaticTokenQuery) (string, error) { + return func(_ context.Context, query StaticTokenQuery) (string, 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 + return token, err } } - 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{staticToken: lookup("private-token", nil)}, unexpectedRefresh(t)).MCPBearerToken(t.Context(), command); err != nil || token != "private-token" { + t.Fatal("static token was not returned", 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) + if token, err := testService(t, &fakeStorage{staticToken: lookup("", credentialcrypto.ErrUnavailable)}, unexpectedRefresh(t)).MCPBearerToken(t.Context(), command); !errors.Is(err, credentialcrypto.ErrUnavailable) || token != "" { + t.Fatal("a keyless lookup returned 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) + unscoped := command + unscoped.Binding.CredentialID = "not-a-credential" + if token, err := testService(t, &fakeStorage{}, unexpectedRefresh(t)).MCPBearerToken(t.Context(), unscoped); !errors.Is(err, ErrNotFound) || token != "" { + t.Fatal("a malformed frozen binding reached storage", err) } } -// oauthScenario is one mcp_oauth Credential sealed by service, served by a -// fake transaction. +// oauthScenario is one mcp_oauth Credential's grant, served by a fake +// transaction that emulates the destination check. type oauthScenario struct { - service *Service - credential Credential - command MCPBearerToken - ciphertext []byte - refreshes []SealedOAuth - committed error + service *Service + credential Credential + command MCPBearerToken + grant OAuthGrant + destinations []string + refreshes []OAuthGrant + committed error } func newOAuthScenario(t *testing.T, expiresAt time.Time, refresher oauthrefresh.Refresher) *oauthScenario { @@ -389,17 +375,24 @@ func newOAuthScenario(t *testing.T, expiresAt time.Time, refresher oauthrefresh. 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}} + scenario := &oauthScenario{ + credential: Credential{ID: uuid.NewString(), VaultID: uuid.NewString(), Name: "OAuth", AuthType: AuthMCPOAuth, MCPServerURL: url, OAuth: &metadata}, + grant: OAuthGrant{Metadata: metadata, AccessToken: "stored-access", RefreshToken: "stored-refresh", ClientSecret: "private-client"}, + } 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 + load: func(_ context.Context, destination string) (OAuthGrant, error) { + scenario.destinations = append(scenario.destinations, destination) + if destination != "" && destination != scenario.credential.MCPServerURL { + return OAuthGrant{}, ErrNotFound + } + return scenario.grant, nil }, - refresh: func(_ context.Context, sealed SealedOAuth) error { - scenario.refreshes = append(scenario.refreshes, sealed) + refresh: func(_ context.Context, grant OAuthGrant) error { + scenario.refreshes = append(scenario.refreshes, grant) return nil }, }) @@ -408,12 +401,7 @@ func newOAuthScenario(t *testing.T, expiresAt time.Time, refresher oauthrefresh. } 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.service = testService(t, storage, refresher) 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 @@ -430,12 +418,16 @@ func TestOAuthBearerTokenRefreshesExpiredGrantUnderLock(t *testing.T) { 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)) + if scenario.destinations[0] != scenario.credential.MCPServerURL { + t.Fatal("the grant was loaded without its frozen destination", scenario.destinations) } - scenario.credential.OAuth, scenario.ciphertext = &metadata, sealed.Ciphertext + refreshed := scenario.refreshes[0] + expiry := later.UTC().Format(time.RFC3339Nano) + want := OAuthGrant{Metadata: OAuthMetadata{ExpiresAt: &expiry, Refresh: scenario.grant.Metadata.Refresh}, AccessToken: "renewed-access", RefreshToken: "stored-refresh", ClientSecret: "private-client"} + if !reflect.DeepEqual(refreshed, want) { + t.Fatalf("refreshed grant was not stored with its metadata: %+v", refreshed) + } + scenario.grant = refreshed 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) } @@ -464,11 +456,23 @@ func TestOAuthBearerTokenFailuresReturnNoToken(t *testing.T) { t.Fatal("a grant was used for another destination", err) } }) - t.Run("substituted metadata", func(t *testing.T) { + t.Run("invalid metadata", func(t *testing.T) { + scenario := newOAuthScenario(t, past, unexpectedRefresh(t)) + scenario.grant.Metadata.Refresh = &OAuthRefreshMetadata{TokenEndpoint: "https://issuer.example/token", TokenEndpointAuth: "client_secret_basic"} + if token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command); err == nil || err.Error() != "OAuth credential authentication failed" || token != "" { + t.Fatal("invalid opened metadata reached the provider", err) + } + }) + t.Run("load failure", 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) + storage := scenario.service.storage.(*fakeStorage) + storage.withOAuthCredential = func(ctx context.Context, _ CredentialKey, apply func(OAuthTx) error) error { + return apply(&fakeOAuthTx{t: t, load: func(context.Context, string) (OAuthGrant, error) { + return OAuthGrant{}, credentialcrypto.ErrUnavailable + }}) + } + if token, err := scenario.service.MCPBearerToken(t.Context(), scenario.command); !errors.Is(err, credentialcrypto.ErrUnavailable) || token != "" { + t.Fatal("a failed load was not returned unchanged", err) } }) t.Run("provider error", func(t *testing.T) { @@ -496,11 +500,11 @@ func TestUpdateOAuthCredentialReplacesLockedGrant(t *testing.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 + var replaced []OAuthGrant 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 apply(&fakeOAuthTx{t: t, load: tx.LoadOAuthGrant, replacement: func(_ context.Context, grant OAuthGrant) (Credential, error) { + replaced = append(replaced, grant) return scenario.credential, nil }}) }) @@ -509,13 +513,12 @@ func TestUpdateOAuthCredentialReplacesLockedGrant(t *testing.T) { 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) + want := OAuthGrant{Metadata: OAuthMetadata{Refresh: scenario.grant.Metadata.Refresh}, AccessToken: "manual-access", RefreshToken: "stored-refresh", ClientSecret: "private-client"} + if len(replaced) != 1 || !reflect.DeepEqual(replaced[0], want) || !reflect.DeepEqual(scenario.destinations, []string{""}) { + t.Fatalf("replacement was not the patched locked grant: %+v", replaced) } 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 { + if _, err := scenario.service.UpdateOAuthCredential(t.Context(), command); !errors.Is(err, ErrInvalidInput) || len(replaced) != 1 { t.Fatal("an invalid expiry was stored", err) } } diff --git a/services/core/internal/vaults/storage.go b/services/core/internal/vaults/storage.go index 683d0a4e..1598bd12 100644 --- a/services/core/internal/vaults/storage.go +++ b/services/core/internal/vaults/storage.go @@ -27,9 +27,13 @@ type Storage interface { 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 seals the Credential's secret and stores it in a Vault + // of the tenant. A missing credential key is + // credentialcrypto.ErrUnavailable, checked after the new Credential ID and + // before the Vault. CreateCredential(ctx context.Context, credential NewCredential) (Credential, error) - // ReplaceStaticToken replaces a static_bearer Credential's sealed token. + // ReplaceStaticToken seals and replaces a static_bearer Credential's + // token. Without a credential key it is credentialcrypto.ErrUnavailable. 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) @@ -43,23 +47,29 @@ type Storage interface { // 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) + // StaticToken opens a static_bearer Credential's token when the complete + // frozen scope still names it. A scope that names none is ErrNotFound, + // then a missing credential key is credentialcrypto.ErrUnavailable. + StaticToken(ctx context.Context, query StaticTokenQuery) (string, error) } // OAuthTx is one mcp_oauth Credential inside a WithOAuthCredential // transaction. type OAuthTx interface { - // LoadOAuthCredential locks the Credential until the transaction ends, so + // LoadOAuthGrant 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) + // Vault's, wait. A non-empty destination must match the stored one, or + // the Credential is ErrNotFound. Only then is the grant opened: a missing + // credential key is credentialcrypto.ErrUnavailable, and stored metadata + // that differs from the sealed copy fails authentication. + LoadOAuthGrant(ctx context.Context, destination string) (OAuthGrant, error) + // ApplyOAuthRefresh seals and stores a refreshed grant for the loaded + // Credential. Execution refreshes are not caller writes and record no + // audit row. + ApplyOAuthRefresh(ctx context.Context, grant OAuthGrant) error + // ApplyOAuthReplacement seals and stores a caller's replacement for the + // loaded Credential and audits it. + ApplyOAuthReplacement(ctx context.Context, grant OAuthGrant) (Credential, error) } // CredentialKey names one Credential of one Vault of one tenant. @@ -74,28 +84,22 @@ type NewVault struct { Metadata []byte } -// NewCredential is a Credential with its ID, sealed to that ID, ready to store. +// NewCredential is a Credential with its new ID, ready to store. Storage +// seals its secret to the tenant, the Vault, that ID, the authentication type +// and the destination: Token for static_bearer, OAuth for mcp_oauth. type NewCredential struct { CredentialKey Name, AuthType, MCPServerURL string - // OAuthMetadata is the encoded OAuthMetadata of an mcp_oauth Credential. - OAuthMetadata []byte - Ciphertext []byte + Token string + OAuth OAuthGrant } -// StaticTokenReplacement is a static_bearer token sealed to the destination -// the write matches. +// StaticTokenReplacement is a new static_bearer token. The write matches the +// destination the token is sealed to. 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 + Token string } // MCPCredentialQuery selects a named Credential by ID alone, so its diff --git a/services/core/internal/vaults/vault.go b/services/core/internal/vaults/vault.go index 177bbce7..c5b5aefa 100644 --- a/services/core/internal/vaults/vault.go +++ b/services/core/internal/vaults/vault.go @@ -27,8 +27,7 @@ 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. +// canonicalID returns the canonical form of a nonzero UUID. func canonicalID(value string) (string, bool) { id, err := uuid.Parse(value) if err != nil || id == uuid.Nil {