From 37de643c788577a0ecd358992b576016a4fd39df Mon Sep 17 00:00:00 2001 From: SaladDay <1203511142@qq.com> Date: Wed, 30 Sep 2026 14:43:33 +0000 Subject: [PATCH] Move PostgreSQL transactions and the execution lease into pgunit persistence/postgres/pgunit now owns Core's transaction mechanics: pooled read-write and snapshot transactions, and the execution lease with its dedicated connection, gate, ownership check, cancellation fence, close and five-second execution deadline. store runs every transaction through it. The connection choice is fixed at construction. store.New builds a pooled Store; store.NewExecution takes the lease and builds the execution writer whose Session and execution-only transactions run on the leased connection. Execution-only operations on a pooled Store fail with ErrExecutionAuthority. ExecutionLease and its Store() view are gone. persistence/postgres/pgtest is the shared test database helper: the oac_*_tests guard, migrations, and isolated databases for database-wide state. It replaces the per-package database setup in store, execution and sandbox/providers tests. Also fix the go vet copylocks warnings in the store dispatch fixtures and record the layering in services/core/IMPLEMENTATION.md. --- CONTRIBUTING.md | 2 +- services/core/IMPLEMENTATION.md | 8 +- .../archive_cancellation_cleanup_test.go | 3 +- .../deployment_provider_observations_test.go | 4 +- .../runtime_retirement_failure_test.go | 74 +----- .../sandbox_deployment_drain_test.go | 93 +++----- .../execution/sandbox_generations_test.go | 3 +- .../internal/execution/sandbox_reset_test.go | 64 +---- .../execution/sandbox_snapshot_budget_test.go | 5 +- services/core/internal/execution/worker.go | 19 +- .../internal/execution/worker_metrics_test.go | 7 +- .../persistence/postgres/pgtest/pgtest.go | 102 ++++++++ .../postgres/pgtest/pgtest_test.go | 19 ++ .../persistence/postgres/pgunit/lease.go | 143 +++++++++++ .../postgres/pgunit/lease_cleanup_test.go} | 57 ++--- .../persistence/postgres/pgunit/lease_test.go | 225 ++++++++++++++++++ .../persistence/postgres/pgunit/pgunit.go | 43 ++++ .../providers/configuration_flow_test.go | 49 +--- .../internal/store/admin_session_archive.go | 4 +- .../store/admin_session_archive_race_test.go | 2 +- .../store/admin_session_archive_test.go | 4 +- .../admin_session_archive_worker_http_test.go | 2 +- services/core/internal/store/admin_summary.go | 2 +- services/core/internal/store/agents.go | 2 +- services/core/internal/store/agents_delete.go | 2 +- services/core/internal/store/agents_update.go | 2 +- .../store/archive_cancellation_test.go | 5 +- .../core/internal/store/artifact_capture.go | 61 +++-- services/core/internal/store/core_metrics.go | 2 +- .../creation_stream_settlement_public_test.go | 5 +- .../store/deployment_model_providers.go | 4 +- .../store/deployment_provider_observations.go | 41 ++-- .../store/device_bootstrap_binding_test.go | 2 +- .../store/environment_claim_worker_test.go | 14 +- .../store/environment_connection_recovery.go | 12 +- .../environment_connection_recovery_test.go | 22 +- .../environment_connection_worker_test.go | 8 +- .../internal/store/environment_connections.go | 7 +- .../store/environment_connections_test.go | 9 +- .../environment_executor_credentials_test.go | 4 +- .../store/environment_executor_management.go | 4 +- .../internal/store/environment_file_writes.go | 10 +- .../store/environment_file_writes_test.go | 9 +- .../store/environment_initial_input_test.go | 11 +- .../store/environment_initial_public_test.go | 6 +- .../store/environment_initialization.go | 4 +- .../store/environment_initialization_test.go | 5 +- .../store/environment_input_activity_test.go | 16 +- .../store/environment_input_claim_test.go | 2 +- .../store/environment_input_expiry.go | 9 +- .../store/environment_input_expiry_test.go | 28 +-- .../store/environment_input_migration_test.go | 15 +- .../environment_input_settlement_test.go | 15 +- .../core/internal/store/environment_inputs.go | 4 +- .../internal/store/environment_inputs_test.go | 10 +- .../store/environment_runtime_fixture_test.go | 9 +- .../store/environment_steering_order_test.go | 2 +- .../internal/store/environment_templates.go | 6 +- .../internal/store/environment_work_test.go | 4 +- .../store/environment_write_audit_test.go | 4 +- services/core/internal/store/execution.go | 68 ++++++ .../store/execution_cancellation_test.go | 87 ------- .../core/internal/store/execution_lease.go | 139 ----------- ...cution_lease_test.go => execution_test.go} | 154 ++++++++---- .../store/executor_credential_target.go | 2 +- .../internal/store/list_cursor_public_test.go | 8 +- .../internal/store/managed_test_store_test.go | 42 +--- .../internal/store/mcp_credentials_oauth.go | 95 ++++---- .../internal/store/prepared_dispatch_test.go | 6 +- .../core/internal/store/project_api_keys.go | 4 +- services/core/internal/store/projects.go | 4 +- .../internal/store/runtime_adoption_test.go | 2 +- .../store/runtime_allocation_state.go | 4 +- .../internal/store/runtime_allocations.go | 4 +- .../store/runtime_allocations_test.go | 11 +- .../core/internal/store/runtime_deployment.go | 8 +- .../internal/store/runtime_deployment_test.go | 16 +- .../runtime_environment_terminal_test.go | 4 +- .../internal/store/runtime_lifecycle_nodes.go | 9 +- .../store/runtime_lifecycle_nodes_test.go | 2 +- .../store/runtime_node_generations.go | 3 +- .../internal/store/runtime_node_presence.go | 4 +- services/core/internal/store/runtime_nodes.go | 2 +- .../core/internal/store/runtime_nodes_test.go | 2 +- .../internal/store/runtime_suspension_test.go | 6 +- .../store/runtime_worker_recovery_test.go | 6 +- .../store/sandbox_deployment_mutations.go | 24 +- .../sandbox_deployment_resources_test.go | 2 +- .../store/sandbox_deployment_setup.go | 12 +- .../store/sandbox_deployment_setup_test.go | 11 +- .../store/sandbox_deployment_switch_test.go | 14 +- .../store/sandbox_deployment_view_test.go | 2 +- .../internal/store/sandbox_generations.go | 103 ++++---- services/core/internal/store/sandbox_reset.go | 8 +- .../core/internal/store/sandbox_reset_test.go | 4 +- .../store/sandbox_specification_store_test.go | 2 +- services/core/internal/store/scheduling.go | 4 +- .../core/internal/store/session_artifacts.go | 94 ++++---- .../store/session_creation_identity.go | 2 +- .../session_deletion_lifecycle_public_test.go | 5 +- .../internal/store/session_diagnostics.go | 4 +- .../store/session_diagnostics_test.go | 2 +- .../session_execution_configuration_test.go | 2 +- .../internal/store/session_initial_input.go | 2 +- .../core/internal/store/session_metadata.go | 2 +- .../internal/store/session_transaction.go | 11 +- services/core/internal/store/sessions.go | 19 +- services/core/internal/store/sessions_test.go | 57 +---- .../core/internal/store/skill_versions.go | 4 +- services/core/internal/store/skills.go | 6 +- services/core/internal/store/source_files.go | 140 +++++------ .../internal/store/subagent_dispatch_test.go | 6 +- .../store/subagent_identities_test.go | 30 +-- .../store/subagent_native_outputs_test.go | 4 +- .../internal/store/subagent_resources_test.go | 2 +- .../store/subagent_visibility_public_test.go | 5 +- services/core/internal/store/turn_events.go | 11 +- .../core/internal/store/vault_credentials.go | 2 +- .../store/vault_credentials_delete.go | 2 +- .../internal/store/vault_credentials_oauth.go | 115 +++++---- .../store/vault_credentials_oauth_test.go | 10 +- .../store/vault_credentials_update.go | 2 +- services/core/internal/store/vaults.go | 2 +- services/core/internal/store/vaults_delete.go | 2 +- 124 files changed, 1436 insertions(+), 1294 deletions(-) create mode 100644 services/core/internal/persistence/postgres/pgtest/pgtest.go create mode 100644 services/core/internal/persistence/postgres/pgtest/pgtest_test.go create mode 100644 services/core/internal/persistence/postgres/pgunit/lease.go rename services/core/internal/{store/execution_lease_cleanup_test.go => persistence/postgres/pgunit/lease_cleanup_test.go} (67%) create mode 100644 services/core/internal/persistence/postgres/pgunit/lease_test.go create mode 100644 services/core/internal/persistence/postgres/pgunit/pgunit.go create mode 100644 services/core/internal/store/execution.go delete mode 100644 services/core/internal/store/execution_cancellation_test.go delete mode 100644 services/core/internal/store/execution_lease.go rename services/core/internal/store/{execution_lease_test.go => execution_test.go} (51%) diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index f591ed678..4b18c3f7d 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -110,7 +110,7 @@ Run `make check` before completing code changes. The `check` target in the [Make | `OAC_TEST_DATABASE_URL` | A dedicated test database. The full gate fails when it is missing. | | `OAC_TEST_OFFICIAL_SDK_PYTHON` | The pinned official SDK interpreter | -The role needs `CREATE DATABASE`: managed-provider tests create and drop isolated `oac_*_tests` databases because provider identity is deployment-wide. Tests must not bypass the production provider-switch guard. +The role needs `CREATE DATABASE`: tests of database-wide state, such as the execution lease and the provider identity, create and drop isolated `oac_*_tests` databases. Tests must not bypass the production provider-switch guard. ### Contract and schema rules diff --git a/services/core/IMPLEMENTATION.md b/services/core/IMPLEMENTATION.md index 24ea4bc2e..d0b27e917 100644 --- a/services/core/IMPLEMENTATION.md +++ b/services/core/IMPLEMENTATION.md @@ -2,6 +2,12 @@ These are the code-level rules of `services/core` that no contract states. Contracts own behavior: the [coverage ledger](../../contracts/agents-api/README.md) and its linked contracts own the public and administrator APIs, the [machine connection API](../../contracts/agents-api/machine-api.md) the `/api/v1` routes, the [Core–Runtime protocol](../../docs/runtime-protocol.md) the daemon wire, and the [Sandbox Provider guide](../../docs/sandbox-provider.md#managed-lifecycle) the managed compute lifecycle. When code changes one of these rules, change the rule here in the same branch. +## Layering + +`internal/persistence/postgres/pgunit` owns Core's PostgreSQL transaction and execution-lease mechanics: pooled read-write and snapshot transactions, the lease's dedicated connection and its gate, the ownership check, the cancellation fence, close, and the execution deadline. Persistence code runs every transaction through it, and nothing outside `persistence` and `store` imports it. `internal/persistence/postgres/pgtest` is test support: it opens the dedicated test database under the `oac_*_tests` guard, applies the migrations, and creates isolated databases for database-wide state such as the execution lease. Only test files import it. + +`store` is transitional. `store.New` builds a pooled Store, and `store.NewExecution` takes the lease and builds the execution writer on it. An execution-only operation on a pooled Store fails with `store.ErrExecutionAuthority`. New adapters do not copy that check: their execution repositories require a `*pgunit.Lease` at construction, their public repositories expose no execution operation, and the check goes away with `store`. + ## Request handling Every Agents API JSON route reads its body through `readJSONObject` before decoding, validation or lookup. The gate requires a JSON Content-Type, applies the route's body limit and rejects invalid UTF-8, malformed JSON (including unpaired surrogate escapes), repeated keys and non-object roots with the official messages; an empty body or `null` becomes `{}`. DELETE, multipart, Core extension and internal routes keep their own readers. Member names match exactly: decode request objects with `decodeInputObject`, or check `inexactMember` before another decoder, so `encoding/json` never matches a case variant. @@ -143,7 +149,7 @@ Item merging never mutates the incoming observation or the previous snapshot: pu ## Worker ownership -Enabling the daemon gateway with `OAC_PUBLIC_URL` also starts the execution Worker. One Worker owns an execution database through a dedicated PostgreSQL advisory-lock connection, and its store view uses that connection for every Session transaction: binding, claim and reconciliation, journal, Items and usage, function callbacks and receipts, and terminal state. These short transactions and the lease pings serialize, with a five-second deadline that includes gate and Session-lock waits. Never hold a transaction across daemon or model work, reconnect the writer or fall back to the pool after losing the lease. Public admission and device maintenance use pooled connections, and pooled reads grant no write authority. +Enabling the daemon gateway with `OAC_PUBLIC_URL` also starts the execution Worker. One Worker owns an execution database through a `pgunit.Lease`: the PostgreSQL advisory lock held by one dedicated connection. A second Worker on the same database fails to start. The Worker's execution writer runs every Session transaction on that connection: binding, claim and reconciliation, journal, Items and usage, function callbacks and receipts, and terminal state. The lease gate serializes these short transactions and the ownership pings, and `pgunit.ExecutionTimeout` (five seconds) bounds each one, including its gate and Session-lock waits. Cancelling an in-flight pgx operation can close the connection that owns the lock, so coordinator-owned work is cancelled only through the lease's cancellation fence, between leased operations. Never hold a transaction across daemon or model work, reconnect the writer or fall back to the pool after losing the lease. Public admission and device maintenance use pooled connections, and pooled reads grant no write authority. Transactions state their isolation: read committed for writes, read-only repeatable read for snapshots. Sandbox reset snapshots bind the deployment relation explicitly to its single row before joining resources, so that even on a fresh database without statistics an inflated join estimate cannot trigger JIT compilation inside the lease deadline. diff --git a/services/core/internal/execution/archive_cancellation_cleanup_test.go b/services/core/internal/execution/archive_cancellation_cleanup_test.go index d1412a553..b8759fae3 100644 --- a/services/core/internal/execution/archive_cancellation_cleanup_test.go +++ b/services/core/internal/execution/archive_cancellation_cleanup_test.go @@ -62,8 +62,7 @@ func TestArchiveWaitingCleanupReceiptBarrier(t *testing.T) { }{{"Kill_no_delivery", false, false}, {"KillCompute_no_delivery", true, false}, {"Kill_live_delivery", false, true}, {"KillCompute_live_delivery", true, true}} { t.Run(scenario.name, func(t *testing.T) { checkpoint := scenario.checkpoint - s, lease := resetManagerStore(t) - writer := lease.Store() + s, writer := resetManagerStore(t) installation := uuid.NewString() if err := writer.ClaimWebSandboxDeployment(t.Context(), installation); err != nil { 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 0509f489c..e223208ed 100644 --- a/services/core/internal/execution/deployment_provider_observations_test.go +++ b/services/core/internal/execution/deployment_provider_observations_test.go @@ -28,7 +28,7 @@ type finishObservationFixture struct { func newFinishObservationFixture(t *testing.T, maxConnections int32) finishObservationFixture { t.Helper() var cfg *pgxpool.Config - s, lease := resetManagerStoreConfig(t, func(c *pgxpool.Config) { + s, writer := resetManagerStoreConfig(t, func(c *pgxpool.Config) { if maxConnections > 0 { c.MaxConns = maxConnections } @@ -56,7 +56,7 @@ func newFinishObservationFixture(t *testing.T, maxConnections int32) finishObser if err != nil { t.Fatal(err) } - return finishObservationFixture{s, lease.Store(), pool, tenant, session, Dispatcher{Store: lease.Store()}} + return finishObservationFixture{s, writer, pool, tenant, session, Dispatcher{Store: writer}} } func (f finishObservationFixture) start(t *testing.T) store.InputReceipt { t.Helper() diff --git a/services/core/internal/execution/runtime_retirement_failure_test.go b/services/core/internal/execution/runtime_retirement_failure_test.go index ec25f666c..dd7418f88 100644 --- a/services/core/internal/execution/runtime_retirement_failure_test.go +++ b/services/core/internal/execution/runtime_retirement_failure_test.go @@ -3,82 +3,14 @@ package execution import ( "context" "errors" - "net" - "os" - "strings" "sync" "sync/atomic" "testing" "time" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sandbox" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgxpool" ) -// Each case has its own real advisory-lock namespace. The delayed driver read -// lets the cancellation fence hit its own deadline without shortening production -// timeouts or depending on an unrelated SQL failure to exercise this branch. -func retirementFailureLease(t *testing.T, armed *atomic.Bool, reading chan struct{}, release <-chan struct{}) (*store.ExecutionLease, *pgxpool.Pool) { - t.Helper() - dsn := os.Getenv("OAC_TEST_DATABASE_URL") - if dsn == "" { - t.Skip("dedicated PostgreSQL required") - } - cfg, err := pgxpool.ParseConfig(dsn) - if err != nil { - t.Fatal(err) - } - if !strings.HasPrefix(cfg.ConnConfig.Database, "oac_") || !strings.HasSuffix(cfg.ConnConfig.Database, "_tests") { - t.Fatal("dedicated test database required") - } - admin, err := pgxpool.NewWithConfig(t.Context(), cfg.Copy()) - if err != nil { - t.Fatal(err) - } - t.Cleanup(admin.Close) - name := "oac_retirement_" + uuid.NewString()[:8] + "_tests" - quoted := pgx.Identifier{name}.Sanitize() - if _, err := admin.Exec(t.Context(), "CREATE DATABASE "+quoted); err != nil { - t.Fatal(err) - } - t.Cleanup(func() { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - if _, err := admin.Exec(ctx, "DROP DATABASE "+quoted+" WITH (FORCE)"); err != nil { - t.Error(err) - } - }) - cfg.ConnConfig.Database = name - dial := cfg.ConnConfig.DialFunc - cfg.ConnConfig.DialFunc = func(ctx context.Context, network, address string) (net.Conn, error) { - c, err := dial(ctx, network, address) - if err != nil { - return nil, err - } - return &delayedLeaseRead{Conn: c, armed: armed, reading: reading, release: release}, nil - } - pool, err := pgxpool.NewWithConfig(t.Context(), cfg) - if err != nil { - t.Fatal(err) - } - t.Cleanup(pool.Close) - lease, err := store.New(pool).AcquireExecutionLease(t.Context()) - if err != nil { - t.Fatal(err) - } - t.Cleanup(func() { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - if err := lease.Close(ctx); err != nil { - t.Error(err) - } - }) - return lease, pool -} - func TestFailedInventoryRetirementClosesAdmissionAndRetainsGate(t *testing.T) { for _, mode := range []string{"gate_timeout", "lease_loss"} { t.Run(mode, func(t *testing.T) { @@ -87,9 +19,9 @@ func TestFailedInventoryRetirementClosesAdmissionAndRetainsGate(t *testing.T) { var readOnce sync.Once unblockRead := func() { readOnce.Do(func() { close(releaseRead) }) } defer unblockRead() - lease, pool := retirementFailureLease(t, &armed, reading, releaseRead) + writer, pool := delayedReadWriter(t, &armed, reading, releaseRead) m := testRuntimeManager(t) - m.store = lease.Store() + m.store = writer m.loadDeployment = func(context.Context) (*RuntimeProvider, error) { return nil, nil } m.mutationGate = make(chan struct{}, 1) // This fixture models an already loaded node deployment; its provider is @@ -117,7 +49,7 @@ func TestFailedInventoryRetirementClosesAdmissionAndRetainsGate(t *testing.T) { defer cancel() queryDone = make(chan error, 1) armed.Store(true) - go func() { queryDone <- lease.Ping(queryCtx) }() + go func() { queryDone <- writer.CheckExecutionOwnership(queryCtx) }() select { case <-reading: case <-time.After(2 * time.Second): diff --git a/services/core/internal/execution/sandbox_deployment_drain_test.go b/services/core/internal/execution/sandbox_deployment_drain_test.go index e6492062b..8c7a9477b 100644 --- a/services/core/internal/execution/sandbox_deployment_drain_test.go +++ b/services/core/internal/execution/sandbox_deployment_drain_test.go @@ -4,18 +4,17 @@ import ( "bytes" "context" "net" - "os" "strings" "sync" "sync/atomic" "testing" "time" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimegateway" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sandbox/node" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" "github.com/google/uuid" - "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) @@ -43,6 +42,35 @@ func (c *delayedLeaseRead) Read(p []byte) (int, error) { return c.Conn.Read(p) } +// delayedReadWriter owns an isolated database, so each case has its own +// advisory-lock namespace. The delayed driver read lets a cancellation fence hit +// its own deadline without shortening production timeouts. +func delayedReadWriter(t *testing.T, armed *atomic.Bool, reading chan struct{}, release <-chan struct{}) (*store.Store, *pgxpool.Pool) { + t.Helper() + pool := pgtest.OpenIsolated(t, func(cfg *pgxpool.Config) { + dial := cfg.ConnConfig.DialFunc + cfg.ConnConfig.DialFunc = func(ctx context.Context, network, address string) (net.Conn, error) { + c, err := dial(ctx, network, address) + if err != nil { + return nil, err + } + return &delayedLeaseRead{Conn: c, armed: armed, reading: reading, release: release}, nil + } + }) + writer, err := store.NewExecution(t.Context(), store.New(pool)) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if err := writer.CloseExecution(ctx); err != nil { + t.Error(err) + } + }) + return writer, pool +} + func TestSandboxDeploymentDrainPreservesLeaseInFlightRead(t *testing.T) { for _, mode := range []string{"deployment", "inventory", "manual"} { t.Run(mode, func(t *testing.T) { @@ -52,62 +80,17 @@ func TestSandboxDeploymentDrainPreservesLeaseInFlightRead(t *testing.T) { } func testLifecycleCancellationPreservesLease(t *testing.T, mode string) { - dsn := os.Getenv("OAC_TEST_DATABASE_URL") - if dsn == "" { - t.Skip("dedicated PostgreSQL required") - } - cfg, err := pgxpool.ParseConfig(dsn) - if err != nil { - t.Fatal(err) - } - if !strings.HasPrefix(cfg.ConnConfig.Database, "oac_") || !strings.HasSuffix(cfg.ConnConfig.Database, "_tests") { - t.Fatal("dedicated test database required") - } - admin, err := pgxpool.NewWithConfig(t.Context(), cfg.Copy()) - if err != nil { - t.Fatal(err) - } - defer admin.Close() - database := pgx.Identifier{"oac_drain_" + uuid.NewString()[:8] + "_tests"}.Sanitize() - if _, err := admin.Exec(t.Context(), "CREATE DATABASE "+database); err != nil { - t.Fatal(err) - } - defer func() { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - if _, err := admin.Exec(ctx, "DROP DATABASE "+database+" WITH (FORCE)"); err != nil { - t.Error(err) - } - }() - cfg.ConnConfig.Database = strings.Trim(database, `"`) var armed atomic.Bool reading, release := make(chan struct{}), make(chan struct{}) var releaseOnce sync.Once unblock := func() { releaseOnce.Do(func() { close(release) }) } defer unblock() - dial := cfg.ConnConfig.DialFunc - cfg.ConnConfig.DialFunc = func(ctx context.Context, network, address string) (net.Conn, error) { - c, err := dial(ctx, network, address) - if err != nil { - return nil, err - } - return &delayedLeaseRead{Conn: c, armed: &armed, reading: reading, release: release}, nil - } - pool, err := pgxpool.NewWithConfig(t.Context(), cfg) - if err != nil { - t.Fatal(err) - } - defer pool.Close() - lease, err := store.New(pool).AcquireExecutionLease(t.Context()) - if err != nil { - t.Fatal(err) - } - defer lease.Close(context.Background()) + writer, _ := delayedReadWriter(t, &armed, reading, release) hub := node.NewHub(node.HubOptions{}) defer hub.Close() id := uuid.NewString() configuration := &RuntimeProvider{InstallationID: id, ProviderKind: "docker", Mode: "nodes", Generation: 1, CoreURL: "https://core.example/api/v1", BackendFingerprint: strings.Repeat("a", 64), Provider: hub.Proxy(uuid.NewString(), "docker", 1)} - m, err := newRuntimeManager(lease.Store(), runtimegateway.NewRegistry(), NewDeferredRuntimeProvider(id, func(context.Context) (*RuntimeProvider, error) { return configuration, nil })) + m, err := newRuntimeManager(writer, runtimegateway.NewRegistry(), NewDeferredRuntimeProvider(id, func(context.Context) (*RuntimeProvider, error) { return configuration, nil })) if err != nil { t.Fatal(err) } @@ -141,11 +124,11 @@ func testLifecycleCancellationPreservesLease(t *testing.T, mode string) { return } defer finish() - queryDone <- lease.Store().CheckExecutionOwnership(ctx) + queryDone <- writer.CheckExecutionOwnership(ctx) <-ctx.Done() // Provider settlement remains outside the cancellation fence. A fresh owner // read must proceed even before this tracked lifecycle operation returns. - leaseFree <- lease.Store().CheckExecutionOwnership(t.Context()) + leaseFree <- writer.CheckExecutionOwnership(t.Context()) }() select { case <-reading: @@ -198,7 +181,7 @@ func testLifecycleCancellationPreservesLease(t *testing.T, mode string) { if delayed != nil { delayed.unblock() } - if err := lease.Store().CheckExecutionOwnership(t.Context()); err != nil { + if err := writer.CheckExecutionOwnership(t.Context()); err != nil { t.Fatalf("deployment drain destroyed the owner connection (in-flight query: %v): %v", queryErr, err) } if queryErr != nil { @@ -240,12 +223,12 @@ func (c *delayedCancellationContext) unblock() { } func TestSandboxDeploymentDrainFailureCannotReactivate(t *testing.T) { - _, lease := resetManagerStore(t) + _, writer := resetManagerStore(t) hub := node.NewHub(node.HubOptions{}) defer hub.Close() id := uuid.NewString() configuration := &RuntimeProvider{InstallationID: id, ProviderKind: "docker", Mode: "nodes", Generation: 1, CoreURL: "https://core.example/api/v1", BackendFingerprint: strings.Repeat("a", 64), Provider: hub.Proxy(uuid.NewString(), "docker", 1)} - m, err := newRuntimeManager(lease.Store(), runtimegateway.NewRegistry(), NewDeferredRuntimeProvider(id, func(context.Context) (*RuntimeProvider, error) { return configuration, nil })) + m, err := newRuntimeManager(writer, runtimegateway.NewRegistry(), NewDeferredRuntimeProvider(id, func(context.Context) (*RuntimeProvider, error) { return configuration, nil })) if err != nil { t.Fatal(err) } @@ -257,7 +240,7 @@ func TestSandboxDeploymentDrainFailureCannotReactivate(t *testing.T) { if err != nil { t.Fatal(err) } - if err := lease.Close(t.Context()); err != nil { + if err := writer.CloseExecution(t.Context()); err != nil { t.Fatal(err) } first := m.pauseDeployment(t.Context()) diff --git a/services/core/internal/execution/sandbox_generations_test.go b/services/core/internal/execution/sandbox_generations_test.go index 88f1ad9d8..7bdc2c78f 100644 --- a/services/core/internal/execution/sandbox_generations_test.go +++ b/services/core/internal/execution/sandbox_generations_test.go @@ -15,8 +15,7 @@ import ( ) func TestE2BReplacementVerifiesTwiceAndNeverPublishesFailedCommit(t *testing.T) { - s, lease := resetManagerStore(t) - writer := lease.Store() + s, writer := resetManagerStore(t) id := uuid.NewString() if err := writer.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) diff --git a/services/core/internal/execution/sandbox_reset_test.go b/services/core/internal/execution/sandbox_reset_test.go index b23e8920f..2fdcfc883 100644 --- a/services/core/internal/execution/sandbox_reset_test.go +++ b/services/core/internal/execution/sandbox_reset_test.go @@ -3,8 +3,6 @@ package execution import ( "bytes" "context" - "database/sql" - "os" "testing" "time" @@ -12,81 +10,39 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/adminaudit" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimegateway" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sandbox/node" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" "github.com/google/uuid" - "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" - "github.com/jackc/pgx/v5/stdlib" - "github.com/pressly/goose/v3" ) -// The execution lease is database-scoped, so this manager test owns a database. -func resetManagerStore(t *testing.T) (*store.Store, *store.ExecutionLease) { +// The execution lease is database-scoped, so these manager tests own a database. +// They receive the pooled Store and the execution writer built on it. +func resetManagerStore(t *testing.T) (*store.Store, *store.Store) { t.Helper() return resetManagerStoreConfig(t, nil) } -func resetManagerStoreConfig(t *testing.T, configure func(*pgxpool.Config)) (*store.Store, *store.ExecutionLease) { +func resetManagerStoreConfig(t *testing.T, configure func(*pgxpool.Config)) (*store.Store, *store.Store) { t.Helper() - url := os.Getenv("OAC_TEST_DATABASE_URL") - if url == "" { - t.Skip("OAC_TEST_DATABASE_URL is required") - } - admin, err := pgxpool.New(t.Context(), url) - if err != nil { - t.Fatal(err) - } - t.Cleanup(admin.Close) - name := "oac_reset_" + uuid.NewString()[:8] + "_tests" - quoted := pgx.Identifier{name}.Sanitize() - if _, err := admin.Exec(t.Context(), "CREATE DATABASE "+quoted); err != nil { - t.Fatal(err) - } - t.Cleanup(func() { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - if _, err := admin.Exec(ctx, "DROP DATABASE "+quoted+" WITH (FORCE)"); err != nil { - t.Error(err) - } - }) - cfg := admin.Config().Copy() - cfg.ConnConfig.Database = name - db := sql.OpenDB(stdlib.GetConnector(*cfg.ConnConfig)) - provider, err := goose.NewProvider(goose.DialectPostgres, db, os.DirFS("../../migrations"), goose.WithTableName("agents_api_schema_version")) - if err != nil { - t.Fatal(err) - } - _, err = provider.Up(t.Context()) - _ = db.Close() - if err != nil { - t.Fatal(err) - } - if configure != nil { - configure(cfg) - } - pool, err := pgxpool.NewWithConfig(t.Context(), cfg) - if err != nil { - t.Fatal(err) - } - t.Cleanup(pool.Close) + pool := pgtest.OpenIsolated(t, configure) cipher, err := credentialcrypto.New(bytes.Repeat([]byte{8}, 32)) if err != nil { t.Fatal(err) } s := store.NewWithCredentialCipher(pool, cipher) - lease, err := s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } - t.Cleanup(func() { _ = lease.Close(context.Background()) }) - return s, lease + t.Cleanup(func() { _ = writer.CloseExecution(context.Background()) }) + return s, writer } func TestSandboxResetPageTimeoutRecoversCommittedOwner(t *testing.T) { - s, lease := resetManagerStore(t) - w := lease.Store() + s, w := resetManagerStore(t) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) diff --git a/services/core/internal/execution/sandbox_snapshot_budget_test.go b/services/core/internal/execution/sandbox_snapshot_budget_test.go index 4e52b91c3..596a29134 100644 --- a/services/core/internal/execution/sandbox_snapshot_budget_test.go +++ b/services/core/internal/execution/sandbox_snapshot_budget_test.go @@ -60,11 +60,10 @@ func (d *snapshotBudget) TraceQueryEnd(ctx context.Context, _ *pgx.Conn, data pg func TestSandboxResetSnapshotFitsPageBudget(t *testing.T) { budget := &snapshotBudget{t: t} - s, lease := resetManagerStoreConfig(t, func(cfg *pgxpool.Config) { + s, w := resetManagerStoreConfig(t, func(cfg *pgxpool.Config) { cfg.ConnConfig.RuntimeParams["jit"] = "on" cfg.ConnConfig.Tracer = budget }) - w := lease.Store() id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) @@ -98,7 +97,7 @@ func TestSandboxResetSnapshotFitsPageBudget(t *testing.T) { budget.armed.Store(true) started := time.Now() err = m.resetStep(t.Context()) - ping := lease.Ping(t.Context()) + ping := w.CheckExecutionOwnership(t.Context()) t.Logf("reset_elapsed=%s reset_error=%v lease_ping=%v", time.Since(started), err, ping) if err != nil || ping != nil { t.Fatal("bounded snapshot lost execution ownership", err, ping) diff --git a/services/core/internal/execution/worker.go b/services/core/internal/execution/worker.go index 03e42f4e1..fd4c5efc8 100644 --- a/services/core/internal/execution/worker.go +++ b/services/core/internal/execution/worker.go @@ -20,7 +20,6 @@ type Worker struct { metrics workerMetricsState dispatcher *Dispatcher admission *store.Store - lease *store.ExecutionLease directoryReads chan directoryReadRequest fileWrites chan fileWriteRequest scheduleWake chan struct{} @@ -34,17 +33,17 @@ func StartWorker(ctx context.Context, dispatcher *Dispatcher) (*Worker, error) { if dispatcher.MaxConcurrentExecutions < 0 || dispatcher.MaxConcurrentExecutions > 1024 { return nil, errors.New("execution concurrency must be between 1 and 1024, or zero for the default") } - lease, err := dispatcher.Store.AcquireExecutionLease(ctx) + writer, err := store.NewExecution(ctx, dispatcher.Store) if err != nil { return nil, err } owned := *dispatcher - owned.Store = lease.Store() + owned.Store = writer owned.notifications = &executionNotifications{} - worker := &Worker{concurrency: dispatcher.MaxConcurrentExecutions, dispatcher: &owned, admission: dispatcher.Store, lease: lease, directoryReads: make(chan directoryReadRequest), fileWrites: make(chan fileWriteRequest), stopped: make(chan struct{}), scheduleWake: make(chan struct{}, 1), enrolledConnections: make(map[string]*runtimeConnection)} + worker := &Worker{concurrency: dispatcher.MaxConcurrentExecutions, dispatcher: &owned, admission: dispatcher.Store, directoryReads: make(chan directoryReadRequest), fileWrites: make(chan fileWriteRequest), stopped: make(chan struct{}), scheduleWake: make(chan struct{}, 1), enrolledConnections: make(map[string]*runtimeConnection)} worker.runtimes, err = newRuntimeManager(owned.Store, owned.Registry, owned.ManagedRuntimes) if err != nil { - _ = lease.Close(context.Background()) + _ = writer.CloseExecution(context.Background()) return nil, err } var deployment *store.RuntimeDeployment @@ -64,21 +63,21 @@ func StartWorker(ctx context.Context, dispatcher *Dispatcher) (*Worker, error) { if worker.runtimes != nil { worker.runtimes.stop() } - _ = lease.Close(context.Background()) + _ = writer.CloseExecution(context.Background()) return nil, err } if err := owned.Store.ReconcileEnvironmentConnections(ctx); err != nil { if worker.runtimes != nil { worker.runtimes.stop() } - _ = lease.Close(context.Background()) + _ = writer.CloseExecution(context.Background()) return nil, err } if err := worker.reconcile(ctx); err != nil { if worker.runtimes != nil { worker.runtimes.stop() } - _ = lease.Close(context.Background()) + _ = writer.CloseExecution(context.Background()) return nil, err } worker.observeOwnership(nil) @@ -87,7 +86,7 @@ func StartWorker(ctx context.Context, dispatcher *Dispatcher) (*Worker, error) { // CheckOwnership checks the same database lease used for execution writes. func (w *Worker) CheckOwnership(ctx context.Context) error { - err := w.lease.Ping(ctx) + err := w.dispatcher.Store.CheckExecutionOwnership(ctx) w.observeOwnership(err) return err } @@ -151,7 +150,7 @@ func (w *Worker) Run(ctx context.Context) (runErr error) { } closeCtx, stop := context.WithTimeout(context.Background(), 5*time.Second) defer stop() - w.observeWorkerClosed(w.lease.Close(closeCtx)) + w.observeWorkerClosed(w.dispatcher.Store.CloseExecution(closeCtx)) }() active := make(map[string]bool) w.observeSlots(len(active)) diff --git a/services/core/internal/execution/worker_metrics_test.go b/services/core/internal/execution/worker_metrics_test.go index e3405c764..93b140769 100644 --- a/services/core/internal/execution/worker_metrics_test.go +++ b/services/core/internal/execution/worker_metrics_test.go @@ -40,11 +40,10 @@ func TestWorkerMetricsUnknownAndDetached(t *testing.T) { } func TestWorkerMetricsFailuresAndClosure(t *testing.T) { - worker := &Worker{lease: &store.ExecutionLease{}} + // A pooled Store has no execution lease, so its ownership check fails. + worker := &Worker{dispatcher: &Dispatcher{Store: &store.Store{}}} worker.observeOwnership(nil) - ctx, cancel := context.WithCancel(context.Background()) - cancel() - if err := worker.CheckOwnership(ctx); !errors.Is(err, context.Canceled) { + if err := worker.CheckOwnership(t.Context()); !errors.Is(err, store.ErrExecutionAuthority) { t.Fatalf("ownership error changed: %v", err) } if worker.MetricsSnapshot().ExecutionOwner != nil { diff --git a/services/core/internal/persistence/postgres/pgtest/pgtest.go b/services/core/internal/persistence/postgres/pgtest/pgtest.go new file mode 100644 index 000000000..8a9fcdcb3 --- /dev/null +++ b/services/core/internal/persistence/postgres/pgtest/pgtest.go @@ -0,0 +1,102 @@ +// Package pgtest opens PostgreSQL databases for Core tests. Only test files +// import it. +package pgtest + +import ( + "context" + "errors" + "os" + "strings" + "testing" + "time" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + "github.com/jackc/pgx/v5/stdlib" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/migrations" +) + +// Open returns a pool on the dedicated test database named by +// OAC_TEST_DATABASE_URL, with Core's migrations applied, and skips the test when +// the variable is unset. Tests on this shared database isolate their data with +// fresh tenant and project IDs. +func Open(t testing.TB) *pgxpool.Pool { + t.Helper() + dsn := os.Getenv("OAC_TEST_DATABASE_URL") + if dsn == "" { + t.Skip("OAC_TEST_DATABASE_URL is not set; dedicated PostgreSQL required") + } + cfg, err := databaseConfig(dsn) + if err != nil { + t.Fatal(err) + } + pool, err := pgxpool.NewWithConfig(context.Background(), cfg) + if err != nil { + t.Fatal(err) + } + t.Cleanup(pool.Close) + // A separate database, not product fixtures or migrations, is sufficient. + var database string + var productTable *string + if err := pool.QueryRow(context.Background(), "SELECT current_database(), to_regclass('workspaces')::text").Scan(&database, &productTable); err != nil || productTable != nil || database != cfg.ConnConfig.Database { + t.Fatal("execution tests require a database without product workspace tables") + } + if err := migrations.Apply(context.Background(), dsn); err != nil { + t.Fatal(err) + } + return pool +} + +// OpenIsolated creates a fresh migrated database beside the one Open uses and +// drops it when the test ends. Tests use it for state that belongs to a whole +// database, such as the execution lease or the sandbox deployment identity. +// configure, when not nil, adjusts the returned pool's configuration. +func OpenIsolated(t testing.TB, configure func(*pgxpool.Config)) *pgxpool.Pool { + t.Helper() + admin := Open(t) + name := "oac_isolated_" + strings.ReplaceAll(uuid.NewString(), "-", "")[:12] + "_tests" + quoted := pgx.Identifier{name}.Sanitize() + if _, err := admin.Exec(t.Context(), "CREATE DATABASE "+quoted); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + if _, err := admin.Exec(ctx, "DROP DATABASE "+quoted+" WITH (FORCE)"); err != nil { + t.Error(err) + } + }) + cfg := admin.Config().Copy() + cfg.ConnConfig.Database = name + connection := stdlib.RegisterConnConfig(cfg.ConnConfig) + err := migrations.Apply(t.Context(), connection) + stdlib.UnregisterConnConfig(connection) + if err != nil { + t.Fatal(err) + } + if configure != nil { + configure(cfg) + } + pool, err := pgxpool.NewWithConfig(context.Background(), cfg) + if err != nil { + t.Fatal(err) + } + t.Cleanup(pool.Close) + return pool +} + +// databaseConfig validates the driver's effective database, so query +// parameters and key/value DSNs cannot redirect tests to another database. +func databaseConfig(dsn string) (*pgxpool.Config, error) { + cfg, err := pgxpool.ParseConfig(dsn) + if err != nil { + return nil, errors.New("invalid test database configuration") + } + database := cfg.ConnConfig.Database + if !strings.HasPrefix(database, "oac_") || !strings.HasSuffix(database, "_tests") { + return nil, errors.New("test database must be named oac_*_tests") + } + return cfg, nil +} diff --git a/services/core/internal/persistence/postgres/pgtest/pgtest_test.go b/services/core/internal/persistence/postgres/pgtest/pgtest_test.go new file mode 100644 index 000000000..0a9190930 --- /dev/null +++ b/services/core/internal/persistence/postgres/pgtest/pgtest_test.go @@ -0,0 +1,19 @@ +package pgtest + +import "testing" + +func TestDatabaseGuardUsesEffectiveDatabase(t *testing.T) { + for _, dsn := range []string{ + "postgres://localhost/oac_local_tests?dbname=agents_api", + "host=localhost dbname=agents_api", + "postgres://localhost/agents_api", + } { + if _, err := databaseConfig(dsn); err == nil { + t.Fatalf("unsafe database accepted: %s", dsn) + } + } + cfg, err := databaseConfig("postgres://localhost/oac_local_tests") + if err != nil || cfg.ConnConfig.Database != "oac_local_tests" { + t.Fatalf("valid dedicated database rejected: %v", err) + } +} diff --git a/services/core/internal/persistence/postgres/pgunit/lease.go b/services/core/internal/persistence/postgres/pgunit/lease.go new file mode 100644 index 000000000..392472d02 --- /dev/null +++ b/services/core/internal/persistence/postgres/pgunit/lease.go @@ -0,0 +1,143 @@ +package pgunit + +import ( + "context" + "errors" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" +) + +// ExecutionTimeout bounds every leased operation, including its wait for the +// connection gate and for row locks. Code that groups several statements under +// one execution deadline uses it too. +const ExecutionTimeout = 5 * time.Second + +var ( + ErrLeaseHeld = errors.New("another execution service owns this database") + ErrLeaseClosed = errors.New("execution lease is closed") +) + +// Lease owns the connection used for execution writes, not just election: the +// PostgreSQL advisory lock belongs to that session, so every execution write runs +// on it. Its gate serializes pgx operations on the connection; no daemon or model +// work holds the gate. Losing or closing the lease never falls back to a pooled +// connection. +type Lease struct { + conn *pgxpool.Conn + gate chan struct{} + cleanupDone <-chan struct{} +} + +// AcquireLease takes the database's execution lease on a dedicated connection +// from pool. It enforces single-service ownership per database and fails with +// ErrLeaseHeld when another service owns it. +func AcquireLease(ctx context.Context, pool *pgxpool.Pool) (*Lease, error) { + conn, err := pool.Acquire(ctx) + if err != nil { + return nil, err + } + acquired, err := sqlc.New(conn).TryExecutionLease(ctx) + if err != nil || !acquired { + _ = conn.Hijack().Close(context.Background()) + if err != nil { + return nil, err + } + return nil, ErrLeaseHeld + } + return &Lease{conn: conn, gate: make(chan struct{}, 1)}, nil +} + +// Transaction runs apply in a read committed transaction on the leased +// connection and commits only when apply returns nil. apply receives the +// context carrying the execution deadline. +func (l *Lease) Transaction(ctx context.Context, apply func(context.Context, pgx.Tx) error) error { + return l.withConn(ctx, func(ctx context.Context, conn *pgxpool.Conn) error { + return run(ctx, conn, readWrite, apply) + }) +} + +// CheckOwnership pings the leased connection, confirming that this service +// still owns the database before external work. +func (l *Lease) CheckOwnership(ctx context.Context) error { + return l.withConn(ctx, func(ctx context.Context, conn *pgxpool.Conn) error { return conn.Ping(ctx) }) +} + +// CancelOperations cancels coordinator-owned contexts between leased +// operations. Cancelling an in-flight pgx operation can close the connection +// that owns the advisory lock, so cancel runs only while the gate is held and +// after the connection answers a ping. cancel must only invoke synchronous +// context cancel functions; it must not perform database, provider or wait work. +// Caller cancellation and operation deadlines keep their own semantics. +func (l *Lease) CancelOperations(ctx context.Context, cancel context.CancelFunc) error { + if cancel == nil { + return errors.New("execution lease cancellation requires a cancel function") + } + return l.withConn(ctx, func(ctx context.Context, conn *pgxpool.Conn) error { + if err := conn.Ping(ctx); err != nil { + return err + } + cancel() + return nil + }) +} + +// Close releases the lease by closing its connection and waits, within ctx, +// for the driver's asynchronous cleanup. A later Close resumes that wait. +func (l *Lease) Close(ctx context.Context) error { + if err := l.lock(ctx); err != nil { + return err + } + defer l.unlock() + if l.conn != nil { + conn := l.conn.Hijack() + l.conn = nil + l.cleanupDone = conn.PgConn().CleanupDone() + if err := conn.Close(ctx); err != nil { + return err + } + } + if l.cleanupDone == nil { + return nil + } + // A cancelled pgx connection can be unusable before its asynchronous cleanup ends. + // Retain the channel so a later Close can continue waiting after this deadline. + select { + case <-l.cleanupDone: + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +// withConn applies the execution deadline to the gate wait and the operation. +func (l *Lease) withConn(ctx context.Context, apply func(context.Context, *pgxpool.Conn) error) error { + ctx, cancel := context.WithTimeout(ctx, ExecutionTimeout) + defer cancel() + if err := l.lock(ctx); err != nil { + return err + } + defer l.unlock() + if l.conn == nil { + return ErrLeaseClosed + } + return apply(ctx, l.conn) +} + +func (l *Lease) lock(ctx context.Context) error { + select { + case l.gate <- struct{}{}: + if err := ctx.Err(); err != nil { + l.unlock() + return err + } + return nil + case <-ctx.Done(): + return ctx.Err() + } +} + +func (l *Lease) unlock() { <-l.gate } diff --git a/services/core/internal/store/execution_lease_cleanup_test.go b/services/core/internal/persistence/postgres/pgunit/lease_cleanup_test.go similarity index 67% rename from services/core/internal/store/execution_lease_cleanup_test.go rename to services/core/internal/persistence/postgres/pgunit/lease_cleanup_test.go index 9f317de07..a7a76c175 100644 --- a/services/core/internal/store/execution_lease_cleanup_test.go +++ b/services/core/internal/persistence/postgres/pgunit/lease_cleanup_test.go @@ -1,54 +1,41 @@ -package store +package pgunit import ( "context" "errors" "net" - "os" "sync" "sync/atomic" "testing" "time" + "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" ) -func TestExecutionLeaseCloseWaitsForCancelledConnectionCleanup(t *testing.T) { - dsn := os.Getenv("OAC_TEST_DATABASE_URL") - if dsn == "" { - t.Skip("dedicated PostgreSQL required") - } - cfg, err := testDatabaseConfig(dsn) - if err != nil { - t.Fatal(err) - } - observer, err := pgxpool.NewWithConfig(t.Context(), cfg.Copy()) - if err != nil { - t.Fatal(err) - } - t.Cleanup(observer.Close) +func TestLeaseCloseWaitsForCancelledConnectionCleanup(t *testing.T) { + observer := pgtest.Open(t) var armed atomic.Bool blocked, release := make(chan struct{}), make(chan struct{}) var signal, unblocked sync.Once unblock := func() { unblocked.Do(func() { close(release) }) } - dial := cfg.ConnConfig.DialFunc - cfg.ConnConfig.DialFunc = func(ctx context.Context, network, address string) (net.Conn, error) { - if armed.Load() { - signal.Do(func() { close(blocked) }) - select { - case <-release: - case <-ctx.Done(): - return nil, ctx.Err() + pool := pgtest.OpenIsolated(t, func(cfg *pgxpool.Config) { + dial := cfg.ConnConfig.DialFunc + cfg.ConnConfig.DialFunc = func(ctx context.Context, network, address string) (net.Conn, error) { + if armed.Load() { + signal.Do(func() { close(blocked) }) + select { + case <-release: + case <-ctx.Done(): + return nil, ctx.Err() + } } + return dial(ctx, network, address) } - return dial(ctx, network, address) - } - pool, err := pgxpool.NewWithConfig(t.Context(), cfg) - if err != nil { - t.Fatal(err) - } - t.Cleanup(pool.Close) - lease, err := New(pool).AcquireExecutionLease(t.Context()) + }) + lease, err := AcquireLease(t.Context(), pool) if err != nil { t.Fatal(err) } @@ -66,8 +53,8 @@ func TestExecutionLeaseCloseWaitsForCancelledConnectionCleanup(t *testing.T) { defer cancelQuery() queryDone := make(chan error, 1) go func() { - queryDone <- lease.withConn(queryCtx, func(conn *pgxpool.Conn) error { - _, err := conn.Exec(queryCtx, "SELECT pg_sleep(10)") + queryDone <- lease.Transaction(queryCtx, func(ctx context.Context, tx pgx.Tx) error { + _, err := tx.Exec(ctx, "SELECT pg_sleep(10)") return err }) }() @@ -101,7 +88,7 @@ func TestExecutionLeaseCloseWaitsForCancelledConnectionCleanup(t *testing.T) { if !errors.Is(err, context.DeadlineExceeded) { t.Fatal("Close returned before blocked cleanup or ignored its deadline", err) } - if err := lease.Store().CheckExecutionOwnership(t.Context()); err == nil { + if err := lease.CheckOwnership(t.Context()); err == nil { t.Fatal("timed-out cleanup restored writer authority") } closed := make(chan error, 1) diff --git a/services/core/internal/persistence/postgres/pgunit/lease_test.go b/services/core/internal/persistence/postgres/pgunit/lease_test.go new file mode 100644 index 000000000..c5377bc66 --- /dev/null +++ b/services/core/internal/persistence/postgres/pgunit/lease_test.go @@ -0,0 +1,225 @@ +package pgunit + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" +) + +// The execution lease is database-wide, so every lease test owns a database. +func acquireLease(t *testing.T, pool *pgxpool.Pool) *Lease { + t.Helper() + lease, err := AcquireLease(t.Context(), pool) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = lease.Close(context.Background()) }) + return lease +} + +// awaitLeaseRelease waits for PostgreSQL to drop the advisory lock of a closed +// or terminated owner; local pgx cleanup does not acknowledge the release. +func awaitLeaseRelease(t *testing.T, pool *pgxpool.Pool) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for { + var held bool + // Match the single-bigint key in queries/scheduling.sql, scoped to this database. + err := pool.QueryRow(t.Context(), `SELECT EXISTS (SELECT 1 FROM pg_locks WHERE locktype='advisory' AND granted AND objsubid=1 + AND classid::bigint * 4294967296 + objid::bigint = 706172736172 + AND database=(SELECT oid FROM pg_database WHERE datname=current_database()))`).Scan(&held) + if err != nil { + t.Fatal("observe execution lease release", err) + } + if !held { + return + } + if time.Now().After(deadline) { + t.Fatal("previous execution lease was not released") + } + time.Sleep(10 * time.Millisecond) + } +} + +func TestLeaseExcludesSecondOwnerUntilClosed(t *testing.T) { + pool := pgtest.OpenIsolated(t, nil) + lease := acquireLease(t, pool) + if _, err := AcquireLease(t.Context(), pool); !errors.Is(err, ErrLeaseHeld) { + t.Fatal("second owner acquired a held lease", err) + } + if err := lease.Close(t.Context()); err != nil { + t.Fatal(err) + } + for _, err := range []error{ + lease.Transaction(t.Context(), func(context.Context, pgx.Tx) error { return nil }), + lease.CheckOwnership(t.Context()), + lease.CancelOperations(t.Context(), func() { t.Error("closed lease cancelled operations") }), + } { + if !errors.Is(err, ErrLeaseClosed) { + t.Fatal("closed lease kept writer authority", err) + } + } + awaitLeaseRelease(t, pool) + successor := acquireLease(t, pool) + if err := successor.CheckOwnership(t.Context()); err != nil { + t.Fatal(err) + } +} + +func TestLeaseLossRejectsOperations(t *testing.T) { + pool := pgtest.OpenIsolated(t, nil) + lease := acquireLease(t, pool) + var killed bool + if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1, 1000)", lease.conn.Conn().PgConn().PID()).Scan(&killed); err != nil || !killed { + t.Fatal(killed, err) + } + operation, cancel := context.WithCancel(t.Context()) + defer cancel() + if err := lease.CancelOperations(t.Context(), cancel); err == nil { + t.Fatal("lost lease accepted cancellation fence") + } + if operation.Err() != nil { + t.Fatal("lost lease ran the cancellation callback") + } + if err := lease.CheckOwnership(t.Context()); err == nil { + t.Fatal("lost owner reported ownership") + } + if err := lease.Transaction(t.Context(), func(context.Context, pgx.Tx) error { return nil }); err == nil { + t.Fatal("lost owner committed a transaction") + } + if err := lease.Close(t.Context()); err != nil { + t.Fatal(err) + } + if err := lease.CancelOperations(t.Context(), cancel); !errors.Is(err, ErrLeaseClosed) { + t.Fatal("closed lease accepted cancellation fence", err) + } + awaitLeaseRelease(t, pool) + successor := acquireLease(t, pool) + if err := successor.CheckOwnership(t.Context()); err != nil { + t.Fatal(err) + } +} + +func TestLeaseSerializesOperations(t *testing.T) { + pool := pgtest.OpenIsolated(t, nil) + lease := acquireLease(t, pool) + // A pgx connection rejects concurrent use, so these succeed only through the gate. + var group sync.WaitGroup + results := make(chan error, 16) + for range 8 { + group.Go(func() { + results <- lease.Transaction(t.Context(), func(ctx context.Context, tx pgx.Tx) error { + _, err := tx.Exec(ctx, "SELECT pg_sleep(0.01)") + return err + }) + }) + group.Go(func() { results <- lease.CheckOwnership(t.Context()) }) + } + group.Wait() + close(results) + for err := range results { + if err != nil { + t.Fatal(err) + } + } + // A waiting operation honours its caller's cancellation. + if err := lease.lock(t.Context()); err != nil { + t.Fatal(err) + } + defer lease.unlock() + cancelled, stop := context.WithCancel(t.Context()) + stop() + if err := lease.CheckOwnership(cancelled); !errors.Is(err, context.Canceled) { + t.Fatal("gate wait ignored cancellation", err) + } +} + +func TestLeaseTransactionDeadlineIncludesLockWait(t *testing.T) { + pool := pgtest.OpenIsolated(t, nil) + if _, err := pool.Exec(t.Context(), "CREATE TABLE lease_probe (id int PRIMARY KEY, value text NOT NULL); INSERT INTO lease_probe VALUES (1, 'before')"); err != nil { + t.Fatal(err) + } + lease := acquireLease(t, pool) + blocker, err := pool.Begin(t.Context()) + if err != nil { + t.Fatal(err) + } + defer blocker.Rollback(context.Background()) + if _, err = blocker.Exec(t.Context(), "SELECT id FROM lease_probe WHERE id=1 FOR UPDATE"); err != nil { + t.Fatal(err) + } + ctx, cancel := context.WithTimeout(t.Context(), 8*time.Second) + defer cancel() + start := time.Now() + err = lease.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + _, err := tx.Exec(ctx, "UPDATE lease_probe SET value='after' WHERE id=1") + return err + }) + if !errors.Is(err, context.DeadlineExceeded) || time.Since(start) >= 7*time.Second { + t.Fatal("leased transaction did not enforce its shorter deadline", err) + } + if err = blocker.Rollback(t.Context()); err != nil { + t.Fatal(err) + } + var value string + if err = pool.QueryRow(t.Context(), "SELECT value FROM lease_probe WHERE id=1").Scan(&value); err != nil || value != "before" { + t.Fatal("timed-out transaction changed the row", value, err) + } +} + +func TestLeaseCancellationFenceHonorsBounds(t *testing.T) { + pool := pgtest.OpenIsolated(t, nil) + lease := acquireLease(t, pool) + for _, test := range []struct { + name string + limit, maximum time.Duration + }{ + {"caller", 25 * time.Millisecond, time.Second}, + {"operation", 8 * time.Second, 6 * time.Second}, + } { + t.Run(test.name, func(t *testing.T) { + if err := lease.lock(t.Context()); err != nil { + t.Fatal(err) + } + held := true + defer func() { + if held { + lease.unlock() + } + }() + operation, stopOperation := context.WithCancel(t.Context()) + defer stopOperation() + ctx, cancel := context.WithTimeout(t.Context(), test.limit) + defer cancel() + started := time.Now() + err := lease.CancelOperations(ctx, stopOperation) + if !errors.Is(err, context.DeadlineExceeded) || time.Since(started) > test.maximum { + t.Fatal("cancellation fence did not preserve its deadline", err, time.Since(started)) + } + if operation.Err() != nil { + t.Fatal("timed-out fence canceled operations outside the lease gate") + } + lease.unlock() + held = false + if err := lease.CheckOwnership(t.Context()); err != nil { + t.Fatal("gate timeout damaged the healthy owner", err) + } + if err := lease.CancelOperations(t.Context(), stopOperation); err != nil { + t.Fatal(err) + } + if !errors.Is(operation.Err(), context.Canceled) { + t.Fatal("successful fence did not cancel synchronously") + } + if err := lease.CheckOwnership(t.Context()); err != nil { + t.Fatal("cancellation damaged the healthy owner", err) + } + }) + } +} diff --git a/services/core/internal/persistence/postgres/pgunit/pgunit.go b/services/core/internal/persistence/postgres/pgunit/pgunit.go new file mode 100644 index 000000000..c33c84d0e --- /dev/null +++ b/services/core/internal/persistence/postgres/pgunit/pgunit.go @@ -0,0 +1,43 @@ +// Package pgunit runs Core's PostgreSQL transactions: pooled transactions for +// public and administrative work, and the execution owner's transactions on the +// connection that holds the database's execution lease. +package pgunit + +import ( + "context" + + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" +) + +// Transactions state their isolation instead of inheriting the server default: +// Core's lock-then-read code relies on read committed statement snapshots. +var ( + readWrite = pgx.TxOptions{IsoLevel: pgx.ReadCommitted} + snapshot = pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly} +) + +// Pool runs transactions on pooled connections. It grants no execution authority. +type Pool struct{ pool *pgxpool.Pool } + +func NewPool(pool *pgxpool.Pool) *Pool { return &Pool{pool: pool} } + +// Transaction runs apply in a read committed transaction and commits only when +// apply returns nil. +func (p *Pool) Transaction(ctx context.Context, apply func(context.Context, pgx.Tx) error) error { + return run(ctx, p.pool, readWrite, apply) +} + +// Snapshot runs apply in a read-only repeatable read transaction, so every +// statement reads the same snapshot. +func (p *Pool) Snapshot(ctx context.Context, apply func(context.Context, pgx.Tx) error) error { + return run(ctx, p.pool, snapshot, apply) +} + +// run passes apply the context that bounds the transaction; statements must use +// it so that they share the transaction's deadline. +func run(ctx context.Context, db interface { + BeginTx(context.Context, pgx.TxOptions) (pgx.Tx, error) +}, options pgx.TxOptions, apply func(context.Context, pgx.Tx) error) error { + return pgx.BeginTxFunc(ctx, db, options, func(tx pgx.Tx) error { return apply(ctx, tx) }) +} diff --git a/services/core/internal/sandbox/providers/configuration_flow_test.go b/services/core/internal/sandbox/providers/configuration_flow_test.go index 0642c206e..a828b81a7 100644 --- a/services/core/internal/sandbox/providers/configuration_flow_test.go +++ b/services/core/internal/sandbox/providers/configuration_flow_test.go @@ -3,21 +3,16 @@ package providers_test import ( "bytes" "context" - "database/sql" "encoding/json" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/api" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/providercontract" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimedevice" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sandbox" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sandbox/providers" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" "github.com/google/uuid" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgxpool" - "github.com/jackc/pgx/v5/stdlib" - "github.com/pressly/goose/v3" "net/http/httptest" - "os" "strings" "testing" ) @@ -80,41 +75,8 @@ func (regionalCodec) DiscoverConfiguration(context.Context, sandbox.Configuratio // A registered native configuration reaches the ordinary API and Store without // adding its fields or kind to either Core package. func TestAdditionalConfigurationProviderUsesCommonAPIAndStore(t *testing.T) { - dsn := os.Getenv("OAC_TEST_DATABASE_URL") - if dsn == "" { - t.Skip("dedicated PostgreSQL required") - } - cfg, err := pgxpool.ParseConfig(dsn) - if err != nil || !strings.HasPrefix(cfg.ConnConfig.Database, "oac_") || !strings.HasSuffix(cfg.ConnConfig.Database, "_tests") { - t.Fatal("dedicated test database required") - } - admin, err := pgxpool.NewWithConfig(t.Context(), cfg) - if err != nil { - t.Fatal(err) - } - defer admin.Close() - name := "oac_provider_" + uuid.NewString()[:8] + "_tests" - quoted := pgx.Identifier{name}.Sanitize() - if _, err = admin.Exec(t.Context(), "CREATE DATABASE "+quoted); err != nil { - t.Fatal(err) - } - defer admin.Exec(context.Background(), "DROP DATABASE "+quoted+" WITH (FORCE)") - cfg.ConnConfig.Database = name - db := sql.OpenDB(stdlib.GetConnector(*cfg.ConnConfig)) - migration, err := goose.NewProvider(goose.DialectPostgres, db, os.DirFS("../../../migrations"), goose.WithTableName("agents_api_schema_version")) - if err != nil { - t.Fatal(err) - } - _, err = migration.Up(t.Context()) - db.Close() - if err != nil { - t.Fatal(err) - } - pool, err := pgxpool.NewWithConfig(t.Context(), cfg) - if err != nil { - t.Fatal(err) - } - defer pool.Close() + // The deployment identity and execution lease are database-wide. + pool := pgtest.OpenIsolated(t, nil) kind := "regional-fixture" adapter, err := providers.Lookup("docker") if err != nil { @@ -123,12 +85,11 @@ func TestAdditionalConfigurationProviderUsesCommonAPIAndStore(t *testing.T) { adapter.Configuration = regionalCodec{} providers.RegisterFixture(t, kind, adapter) s := store.New(pool) - lease, err := s.AcquireExecutionLease(t.Context()) + w, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } - defer lease.Close(context.Background()) - w := lease.Store() + defer w.CloseExecution(context.Background()) installation := uuid.NewString() if err = w.ClaimWebSandboxDeployment(t.Context(), installation); err != nil { t.Fatal(err) diff --git a/services/core/internal/store/admin_session_archive.go b/services/core/internal/store/admin_session_archive.go index f0dbc0f8a..e323ec9d6 100644 --- a/services/core/internal/store/admin_session_archive.go +++ b/services/core/internal/store/admin_session_archive.go @@ -35,8 +35,8 @@ func (s *Store) ArchiveSandboxResetSession(ctx context.Context, tenantID, sessio var ErrSandboxResetSessionBusy = errors.New("the hosted Session is busy") func (s *Store) archiveManagedSession(ctx context.Context, tenantID, sessionID string, expectedGeneration uint64, resetRequestedAt *time.Time) (ManagedSessionArchive, error) { - if s.executionLease == nil { - return ManagedSessionArchive{}, ErrInvalidInput + if err := s.checkExecutionAuthority(); err != nil { + return ManagedSessionArchive{}, err } tenant, err := parseID(tenantID) if err != nil { diff --git a/services/core/internal/store/admin_session_archive_race_test.go b/services/core/internal/store/admin_session_archive_race_test.go index bb8915a12..c699217b3 100644 --- a/services/core/internal/store/admin_session_archive_race_test.go +++ b/services/core/internal/store/admin_session_archive_race_test.go @@ -12,7 +12,7 @@ import ( func TestManagedSessionArchiveReleasesPendingNodePlacement(t *testing.T) { s, _ := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) installation := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), installation); err != nil { t.Fatal(err) diff --git a/services/core/internal/store/admin_session_archive_test.go b/services/core/internal/store/admin_session_archive_test.go index 6770beb1a..1b9bdfbb1 100644 --- a/services/core/internal/store/admin_session_archive_test.go +++ b/services/core/internal/store/admin_session_archive_test.go @@ -21,7 +21,7 @@ func managedArchiveFixture(t *testing.T) (*Store, *Store, string) { t.Fatal(err) } s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) installation := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), installation); err != nil { t.Fatal(err) @@ -72,7 +72,7 @@ func TestManagedSessionArchiveUnallocatedAndGuards(t *testing.T) { t.Fatal("archive accepted wrong generation", generation, err) } } - if _, err := s.ArchiveManagedSession(ctx, tenant, session.ID, 1); !errors.Is(err, ErrInvalidInput) { + if _, err := s.ArchiveManagedSession(ctx, tenant, session.ID, 1); !errors.Is(err, ErrExecutionAuthority) { t.Fatal("unleased archive accepted", err) } for _, other := range []string{uuid.NewString(), "malformed"} { diff --git a/services/core/internal/store/admin_session_archive_worker_http_test.go b/services/core/internal/store/admin_session_archive_worker_http_test.go index 2f137749c..ac657b750 100644 --- a/services/core/internal/store/admin_session_archive_worker_http_test.go +++ b/services/core/internal/store/admin_session_archive_worker_http_test.go @@ -87,7 +87,7 @@ func TestAdminSessionArchiveWorkerHTTPPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if _, err := s.ArchiveManagedSession(ctx, project.TenantID, active.ID, 1); !errors.Is(err, store.ErrInvalidInput) { + if _, err := s.ArchiveManagedSession(ctx, project.TenantID, active.ID, 1); !errors.Is(err, store.ErrExecutionAuthority) { t.Fatal("fixture admission Store unexpectedly holds execution ownership", err) } admin, err := api.NewDeploymentAuthenticator([]string{runtimedevice.HashCredential("archive-administrator")}) diff --git a/services/core/internal/store/admin_summary.go b/services/core/internal/store/admin_summary.go index 14a1c092e..2ed0ff00a 100644 --- a/services/core/internal/store/admin_summary.go +++ b/services/core/internal/store/admin_summary.go @@ -33,7 +33,7 @@ func (s *Store) ReadAdminSummary(ctx context.Context, tenantID string, filter Ad if visit == nil || filter.CreatedAfter != nil && filter.CreatedBefore != nil && !filter.CreatedAfter.Before(*filter.CreatedBefore) { return counts, ErrInvalidInput } - err = pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}, func(tx pgx.Tx) error { + err = s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) raw, err := q.AdminAssetCounts(ctx, tenant) if err != nil { diff --git a/services/core/internal/store/agents.go b/services/core/internal/store/agents.go index 1236deb1f..d5e880a3d 100644 --- a/services/core/internal/store/agents.go +++ b/services/core/internal/store/agents.go @@ -55,7 +55,7 @@ func (s *Store) CreateAgent(ctx context.Context, tenantID string, input CreateAg return SavedAgent{}, err } var created SavedAgent - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.CreateAgent(ctx, sqlc.CreateAgentParams{ ID: pgtype.UUID{Bytes: uuid.New(), Valid: true}, TenantID: tenant, diff --git a/services/core/internal/store/agents_delete.go b/services/core/internal/store/agents_delete.go index 5c842764c..081df8359 100644 --- a/services/core/internal/store/agents_delete.go +++ b/services/core/internal/store/agents_delete.go @@ -21,7 +21,7 @@ func (s *Store) DeleteAgent(ctx context.Context, tenantID, agentID string) (stri return "", err } var deletedID string - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) deleted, err := q.DeleteAgent(ctx, sqlc.DeleteAgentParams{TenantID: tenant, ID: id}) if err != nil { diff --git a/services/core/internal/store/agents_update.go b/services/core/internal/store/agents_update.go index 9ee2882e2..49b315cf9 100644 --- a/services/core/internal/store/agents_update.go +++ b/services/core/internal/store/agents_update.go @@ -55,7 +55,7 @@ func (s *Store) UpdateAgent(ctx context.Context, tenantID, agentID string, input } } var updated SavedAgent - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.LockAgent(ctx, sqlc.LockAgentParams{TenantID: tenant, ID: id}) if errors.Is(err, pgx.ErrNoRows) { diff --git a/services/core/internal/store/archive_cancellation_test.go b/services/core/internal/store/archive_cancellation_test.go index 81b368585..f38df21ae 100644 --- a/services/core/internal/store/archive_cancellation_test.go +++ b/services/core/internal/store/archive_cancellation_test.go @@ -32,16 +32,15 @@ func TestArchiveWaitingCancellationReceipts(t *testing.T) { t.Run(scenario, func(t *testing.T) { heartbeat := scenario != "receipt_without_heartbeat" s, pool := store.NewManagedTestStore(t) - lease, err := s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } t.Cleanup(func() { - if err := lease.Close(context.Background()); err != nil { + if err := writer.CloseExecution(context.Background()); err != nil { t.Error(err) } }) - writer := lease.Store() installation := uuid.NewString() if err := writer.ClaimWebSandboxDeployment(t.Context(), installation); err != nil { t.Fatal(err) diff --git a/services/core/internal/store/artifact_capture.go b/services/core/internal/store/artifact_capture.go index d0584a088..54e83d140 100644 --- a/services/core/internal/store/artifact_capture.go +++ b/services/core/internal/store/artifact_capture.go @@ -46,40 +46,37 @@ func (s *Store) StageTurnArtifacts(ctx context.Context, tenantID, sessionID, tur if configuration.Type != "openai_hosted" && configuration.Type != "self_hosted" { return ErrInvalidInput } - tx, err := s.pool.Begin(ctx) - if err != nil { - return err - } - defer tx.Rollback(context.Background()) - rows, err := captureArtifactArchive(ctx, tx, input) - if err != nil { - return err - } - q := s.queries.WithTx(tx) - locked, err := q.LockSession(ctx, sqlc.LockSessionParams{TenantID: lookup.TenantID, ID: lookup.SessionID}) - if errors.Is(err, pgx.ErrNoRows) { - return ErrNotFound - } - if err != nil { - return err - } - if locked.DeletedAt.Valid { - return ErrNotFound - } - turn, err := q.GetTurn(ctx, lookup) - if err != nil { - return err - } - if turn.Status != TurnInProgress || turn.CancelRequestedAt.Valid { - return ErrTurnConflict - } - for _, row := range rows { - row.SessionID, row.TurnID, row.EnvironmentID = lookup.SessionID, lookup.ID, environment - if err := q.StageSessionArtifact(ctx, row); err != nil { + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + rows, err := captureArtifactArchive(ctx, tx, input) + if err != nil { return err } - } - return tx.Commit(ctx) + q := s.queries.WithTx(tx) + locked, err := q.LockSession(ctx, sqlc.LockSessionParams{TenantID: lookup.TenantID, ID: lookup.SessionID}) + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotFound + } + if err != nil { + return err + } + if locked.DeletedAt.Valid { + return ErrNotFound + } + turn, err := q.GetTurn(ctx, lookup) + if err != nil { + return err + } + if turn.Status != TurnInProgress || turn.CancelRequestedAt.Valid { + return ErrTurnConflict + } + for _, row := range rows { + row.SessionID, row.TurnID, row.EnvironmentID = lookup.SessionID, lookup.ID, environment + if err := q.StageSessionArtifact(ctx, row); err != nil { + return err + } + } + return nil + }) } func captureArtifactArchive(ctx context.Context, tx pgx.Tx, input io.Reader) ([]sqlc.StageSessionArtifactParams, error) { diff --git a/services/core/internal/store/core_metrics.go b/services/core/internal/store/core_metrics.go index 6e2de85ce..6af05ffaf 100644 --- a/services/core/internal/store/core_metrics.go +++ b/services/core/internal/store/core_metrics.go @@ -73,7 +73,7 @@ func (s *Store) ReadCoreExecutionHistory(ctx context.Context, start, end time.Ti result.Buckets[i].Start = start.Add(time.Duration(i) * resolution).UTC() } first, last := pgtype.Timestamptz{Time: start, Valid: true}, pgtype.Timestamptz{Time: end, Valid: true} - err := pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}, func(tx pgx.Tx) error { + err := s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) var err error result.Interrupted, err = q.CoreInterruptedTurns(ctx, sqlc.CoreInterruptedTurnsParams{RangeStart: first, RangeEnd: last}) diff --git a/services/core/internal/store/creation_stream_settlement_public_test.go b/services/core/internal/store/creation_stream_settlement_public_test.go index f0d236a50..b02e1e6b5 100644 --- a/services/core/internal/store/creation_stream_settlement_public_test.go +++ b/services/core/internal/store/creation_stream_settlement_public_test.go @@ -135,12 +135,11 @@ func TestCreationStreamPublicLifetimes(t *testing.T) { handler.ServeHTTP(w, r) })) defer server.Close() - lease, err := s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } - defer func() { _ = lease.Close(context.Background()) }() - writer := lease.Store() + defer func() { _ = writer.CloseExecution(context.Background()) }() connect := func(environment string) { t.Helper() generation := uuid.NewString() diff --git a/services/core/internal/store/deployment_model_providers.go b/services/core/internal/store/deployment_model_providers.go index 886f6a635..da8bae0e5 100644 --- a/services/core/internal/store/deployment_model_providers.go +++ b/services/core/internal/store/deployment_model_providers.go @@ -70,7 +70,7 @@ func (s *Store) SetDeploymentModelProvider(ctx context.Context, harness string, return DeploymentModelProvider{}, ErrCredentialStorageUnavailable } var result DeploymentModelProvider - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.UpsertDeploymentModelProvider(ctx, sqlc.UpsertDeploymentModelProviderParams{ Harness: harness, Protocol: provider.Protocol, BaseUrl: provider.BaseURL, Model: configuration.Model, HarnessConfig: v1.ResolvedHarnessConfig(configuration.HarnessConfig), @@ -88,7 +88,7 @@ func (s *Store) SetDeploymentModelProvider(ctx context.Context, harness string, // DeleteDeploymentModelProvider is idempotent; each successful call is audited. // Sessions that already froze the default keep their snapshot. func (s *Store) DeleteDeploymentModelProvider(ctx context.Context, harness string) error { - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) if _, err := q.DeleteDeploymentModelProvider(ctx, harness); err != nil { return err diff --git a/services/core/internal/store/deployment_provider_observations.go b/services/core/internal/store/deployment_provider_observations.go index 555289d62..0aa88138f 100644 --- a/services/core/internal/store/deployment_provider_observations.go +++ b/services/core/internal/store/deployment_provider_observations.go @@ -5,6 +5,8 @@ import ( "strconv" "time" + "github.com/jackc/pgx/v5" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" ) @@ -20,30 +22,27 @@ func (s *Store) ObserveDeploymentModelProvider(ctx context.Context, tenantID, se } ctx, cancel := context.WithTimeout(ctx, time.Second) defer cancel() - tx, err := s.pool.Begin(ctx) - if err != nil { - return 0, err - } - defer tx.Rollback(ctx) - deadline, _ := ctx.Deadline() - // Leave a small part of the overall budget for returning the server error - // and releasing this metadata-only transaction before the client deadline. - timeout := time.Until(deadline).Milliseconds() - 25 - if timeout <= 0 { - return 0, context.DeadlineExceeded - } - setting := strconv.FormatInt(timeout, 10) + "ms" - if _, err = tx.Exec(ctx, "SELECT set_config('statement_timeout', $1, true), set_config('lock_timeout', $1, true)", setting); err != nil { - return 0, err - } - count, err := sqlc.New(tx).ObserveDeploymentModelProvider(ctx, sqlc.ObserveDeploymentModelProviderParams{ - TenantID: lookup.TenantID, SessionID: lookup.SessionID, TurnID: lookup.ID, + var count int64 + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + deadline, _ := ctx.Deadline() + // Leave a small part of the overall budget for returning the server error + // and releasing this metadata-only transaction before the client deadline. + timeout := time.Until(deadline).Milliseconds() - 25 + if timeout <= 0 { + return context.DeadlineExceeded + } + setting := strconv.FormatInt(timeout, 10) + "ms" + if _, err := tx.Exec(ctx, "SELECT set_config('statement_timeout', $1, true), set_config('lock_timeout', $1, true)", setting); err != nil { + return err + } + var err error + count, err = sqlc.New(tx).ObserveDeploymentModelProvider(ctx, sqlc.ObserveDeploymentModelProviderParams{ + TenantID: lookup.TenantID, SessionID: lookup.SessionID, TurnID: lookup.ID, + }) + return err }) if err != nil { return 0, err } - if err = tx.Commit(ctx); err != nil { - return 0, err - } return count, nil } diff --git a/services/core/internal/store/device_bootstrap_binding_test.go b/services/core/internal/store/device_bootstrap_binding_test.go index d91292750..c6b4c5f21 100644 --- a/services/core/internal/store/device_bootstrap_binding_test.go +++ b/services/core/internal/store/device_bootstrap_binding_test.go @@ -64,7 +64,7 @@ func TestDeviceCredentialWithoutManagedNodeRetainsPublicRouteIdentity(t *testing t.Fatal(err) } _, environment := localEnvironment(t, s, tenant) - allocation, err := executionLease(t, s).Store().ReserveRuntimeAllocation(t.Context(), tenant, environment.ID, uuid.NewString(), runtimedevice.HashCredential("allocation-token")) + allocation, err := executionWriter(t, s).ReserveRuntimeAllocation(t.Context(), tenant, environment.ID, uuid.NewString(), runtimedevice.HashCredential("allocation-token")) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/environment_claim_worker_test.go b/services/core/internal/store/environment_claim_worker_test.go index f9ab22c5c..9fd582233 100644 --- a/services/core/internal/store/environment_claim_worker_test.go +++ b/services/core/internal/store/environment_claim_worker_test.go @@ -16,12 +16,12 @@ func TestWorkerReconcilesEnvironmentPromotionBeforeStart(t *testing.T) { t.Run(map[bool]string{false: "unbound", true: "deleted"}[deleted], func(t *testing.T) { s, pool := store.NewTestStore(t) tenant, pending := newEnvironmentExpiryReservation(t, s) - lease, err := s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } - t.Cleanup(func() { _ = lease.Close(context.Background()) }) - got, err := lease.Store().PromoteEnvironmentInput(t.Context(), tenant, pending.SessionID, pending.ID) + t.Cleanup(func() { _ = writer.CloseExecution(context.Background()) }) + got, err := writer.PromoteEnvironmentInput(t.Context(), tenant, pending.SessionID, pending.ID) if err != nil || len(got.Receipts) != 1 || got.Receipts[0].Replayed { t.Fatal(got, err) } @@ -40,7 +40,7 @@ func TestWorkerReconcilesEnvironmentPromotionBeforeStart(t *testing.T) { } // Simulate owner loss after commit, without sending any daemon Start. awaitRelease := observeExecutionLeaseRelease(t, pool) - if err := lease.Close(t.Context()); err != nil { + if err := writer.CloseExecution(t.Context()); err != nil { t.Fatal(err) } awaitRelease() @@ -70,12 +70,12 @@ func TestWorkerReconcilesEnvironmentPromotionBeforeStart(t *testing.T) { if err != nil || turns != 1 || inputs != 1 || queued != 0 { t.Fatal("restart duplicated or requeued prepared work", turns, inputs, queued, err) } - successor, err := s.AcquireExecutionLease(t.Context()) + successor, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } - t.Cleanup(func() { _ = successor.Close(context.Background()) }) - retry, err := successor.Store().PromoteEnvironmentInput(t.Context(), tenant, pending.SessionID, pending.ID) + t.Cleanup(func() { _ = successor.CloseExecution(context.Background()) }) + retry, err := successor.PromoteEnvironmentInput(t.Context(), tenant, pending.SessionID, pending.ID) if deleted { if !errors.Is(err, store.ErrNotFound) { t.Fatal("deleted reservation was exposed", err) diff --git a/services/core/internal/store/environment_connection_recovery.go b/services/core/internal/store/environment_connection_recovery.go index 3730d54f8..c32fa6467 100644 --- a/services/core/internal/store/environment_connection_recovery.go +++ b/services/core/internal/store/environment_connection_recovery.go @@ -6,26 +6,24 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" "github.com/google/uuid" + "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" - "github.com/jackc/pgx/v5/pgxpool" ) // ReconcileEnvironmentConnections runs before the new owner's connection producers start. // It discards previous process generations and records loss of their connected transport. func (s *Store) ReconcileEnvironmentConnections(ctx context.Context) error { - if s.executionLease == nil { - return errors.New("Environment reconciliation requires an execution lease") + if err := s.checkExecutionAuthority(); err != nil { + return err } after := pgtype.UUID{Valid: true} for { var rows []sqlc.ListEnvironmentConnectionsRow - queryCtx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - err := s.executionLease.withConn(queryCtx, func(conn *pgxpool.Conn) error { + err := s.writer.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { var err error - rows, err = sqlc.New(conn).ListEnvironmentConnections(queryCtx, after) + rows, err = s.queries.WithTx(tx).ListEnvironmentConnections(ctx, after) return err }) - cancel() if err != nil { return err } diff --git a/services/core/internal/store/environment_connection_recovery_test.go b/services/core/internal/store/environment_connection_recovery_test.go index cad363ed3..d6caab3f8 100644 --- a/services/core/internal/store/environment_connection_recovery_test.go +++ b/services/core/internal/store/environment_connection_recovery_test.go @@ -8,7 +8,7 @@ import ( func TestEnvironmentConnectionRecoveryFencesLostOwnerAcrossPages(t *testing.T) { s, pool := testStore(t) - old := executionLease(t, s) + old := executionWriter(t, s) type target struct { tenant string session Session @@ -19,27 +19,27 @@ func TestEnvironmentConnectionRecoveryFencesLostOwnerAcrossPages(t *testing.T) { for range 33 { tenant, session, environment := connectionFixture(t, s) generation := uuid.NewString() - if err := old.Store().ReplaceEnvironmentConnection(t.Context(), tenant, environment.ID, generation); err != nil { + if err := old.ReplaceEnvironmentConnection(t.Context(), tenant, environment.ID, generation); err != nil { t.Fatal(err) } - if err := old.Store().ObserveEnvironmentConnection(t.Context(), tenant, environment.ID, generation, 1, true); err != nil { + if err := old.ObserveEnvironmentConnection(t.Context(), tenant, environment.ID, generation, 1, true); err != nil { t.Fatal(err) } targets = append(targets, target{tenant, session, environment, generation}) } var killed bool - if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1,1000)", old.conn.Conn().PgConn().PID()).Scan(&killed); err != nil || !killed { + if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1,1000)", executionOwnerPID(t, pool)).Scan(&killed); err != nil || !killed { t.Fatal(killed, err) } - next := executionLease(t, s) + next := executionWriter(t, s) first := targets[0] - if err := old.Store().ObserveEnvironmentConnection(t.Context(), first.tenant, first.environment.ID, first.generation, 2, false); err == nil { + if err := old.ObserveEnvironmentConnection(t.Context(), first.tenant, first.environment.ID, first.generation, 2, false); err == nil { t.Fatal("lost owner wrote state") } if err := s.ReconcileEnvironmentConnections(t.Context()); err == nil { t.Fatal("unleased reconciliation accepted") } - if err := next.Store().ReconcileEnvironmentConnections(t.Context()); err != nil { + if err := next.ReconcileEnvironmentConnections(t.Context()); err != nil { t.Fatal(err) } for _, target := range targets { @@ -52,24 +52,24 @@ func TestEnvironmentConnectionRecoveryFencesLostOwnerAcrossPages(t *testing.T) { t.Fatal("recovery lost event snapshots", changes) } before := connectionSnapshot(t, pool, target.environment.ID) - if err := next.Store().ObserveEnvironmentConnection(t.Context(), target.tenant, target.environment.ID, target.generation, 100, true); err != nil { + if err := next.ObserveEnvironmentConnection(t.Context(), target.tenant, target.environment.ID, target.generation, 100, true); err != nil { t.Fatal(err) } if after := connectionSnapshot(t, pool, target.environment.ID); after != before { t.Fatal("old generation survived recovery") } } - if err := next.Store().ReconcileEnvironmentConnections(t.Context()); err != nil { + if err := next.ReconcileEnvironmentConnections(t.Context()); err != nil { t.Fatal(err) } if len(connectionChanges(t, s, first.tenant, first.session.ID)) != 2 { t.Fatal("repeated recovery duplicated a disconnect") } generation := uuid.NewString() - if err := next.Store().ReplaceEnvironmentConnection(t.Context(), first.tenant, first.environment.ID, generation); err != nil { + if err := next.ReplaceEnvironmentConnection(t.Context(), first.tenant, first.environment.ID, generation); err != nil { t.Fatal(err) } - if err := next.Store().ObserveEnvironmentConnection(t.Context(), first.tenant, first.environment.ID, generation, 1, true); err != nil { + if err := next.ObserveEnvironmentConnection(t.Context(), first.tenant, first.environment.ID, generation, 1, true); err != nil { t.Fatal(err) } got, err := s.GetEnvironment(t.Context(), first.tenant, first.environment.ID) diff --git a/services/core/internal/store/environment_connection_worker_test.go b/services/core/internal/store/environment_connection_worker_test.go index f2c5f67b8..29991aae0 100644 --- a/services/core/internal/store/environment_connection_worker_test.go +++ b/services/core/internal/store/environment_connection_worker_test.go @@ -23,19 +23,19 @@ func TestEnvironmentConnectionWorkerReconcilesAndReleasesLease(t *testing.T) { if err != nil { t.Fatal(err) } - lease, err := s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } generation := uuid.NewString() - if err := lease.Store().ReplaceEnvironmentConnection(t.Context(), tenant, environment.ID, generation); err != nil { + if err := writer.ReplaceEnvironmentConnection(t.Context(), tenant, environment.ID, generation); err != nil { t.Fatal(err) } - if err := lease.Store().ObserveEnvironmentConnection(t.Context(), tenant, environment.ID, generation, 1, true); err != nil { + if err := writer.ObserveEnvironmentConnection(t.Context(), tenant, environment.ID, generation, 1, true); err != nil { t.Fatal(err) } awaitRelease := observeExecutionLeaseRelease(t, pool) - if err := lease.Close(t.Context()); err != nil { + if err := writer.CloseExecution(t.Context()); err != nil { t.Fatal(err) } awaitRelease() diff --git a/services/core/internal/store/environment_connections.go b/services/core/internal/store/environment_connections.go index f2db611c8..9c3b3d809 100644 --- a/services/core/internal/store/environment_connections.go +++ b/services/core/internal/store/environment_connections.go @@ -7,6 +7,7 @@ import ( v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" "github.com/google/uuid" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" @@ -76,10 +77,10 @@ func (s *Store) ObserveEnvironmentConnection(ctx context.Context, tenant, enviro } func (s *Store) withEnvironmentConnection(ctx context.Context, tenant, environment string, apply func(context.Context, *sqlc.Queries, sqlc.GetSessionEnvironmentRow) error) error { - if s.executionLease == nil { - return errors.New("Environment observations require an execution lease") + if err := s.checkExecutionAuthority(); err != nil { + return err } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) + ctx, cancel := context.WithTimeout(ctx, pgunit.ExecutionTimeout) defer cancel() owned, err := s.GetEnvironment(ctx, tenant, environment) if err != nil { diff --git a/services/core/internal/store/environment_connections_test.go b/services/core/internal/store/environment_connections_test.go index c6a8397be..a7ba28722 100644 --- a/services/core/internal/store/environment_connections_test.go +++ b/services/core/internal/store/environment_connections_test.go @@ -52,7 +52,7 @@ func connectionChanges(t *testing.T, s *Store, tenant, session string) []Session func TestEnvironmentConnectionOrdersGenerationsAndImmutableEvents(t *testing.T) { s, pool := testStore(t) tenant, session, environment := connectionFixture(t, s) - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) first, second := uuid.NewString(), uuid.NewString() replace := func(gen string) { t.Helper() @@ -135,8 +135,7 @@ func TestEnvironmentConnectionRequiresOwnerAndRollsBackWithEvent(t *testing.T) { if err := s.ObserveEnvironmentConnection(t.Context(), tenant, environment.ID, generation, 1, true); err == nil { t.Fatal("unleased observation accepted") } - lease := executionLease(t, s) - writer := lease.Store() + writer := executionWriter(t, s) if err := writer.ReplaceEnvironmentConnection(t.Context(), tenant, environment.ID, generation); err != nil { t.Fatal(err) } @@ -169,7 +168,7 @@ func TestEnvironmentConnectionRequiresOwnerAndRollsBackWithEvent(t *testing.T) { if err := writer.ObserveEnvironmentConnection(t.Context(), tenant, environment.ID, generation, 1, true); err != nil { t.Fatal(err) } - if err := lease.Close(t.Context()); err != nil { + if err := writer.CloseExecution(t.Context()); err != nil { t.Fatal(err) } before = connectionSnapshot(t, pool, environment.ID) @@ -186,7 +185,7 @@ func TestEnvironmentConnectionDoesNotReviveDeletedOrTerminalResources(t *testing t.Run(status, func(t *testing.T) { s, pool := testStore(t) tenant, session, environment := connectionFixture(t, s) - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) generation := uuid.NewString() if err := writer.ReplaceEnvironmentConnection(t.Context(), tenant, environment.ID, generation); err != nil { t.Fatal(err) diff --git a/services/core/internal/store/environment_executor_credentials_test.go b/services/core/internal/store/environment_executor_credentials_test.go index 7ababfbd4..ab67e3720 100644 --- a/services/core/internal/store/environment_executor_credentials_test.go +++ b/services/core/internal/store/environment_executor_credentials_test.go @@ -122,11 +122,11 @@ func TestEnvironmentExecutorConcurrentIssueAndDeletion(t *testing.T) { t.Fatal(err) } // Provisioning remains control-plane work while the execution owner is active. - lease, err := s.AcquireExecutionLease(ctx) + writer, err := NewExecution(ctx, s) if err != nil { t.Fatal(err) } - defer lease.Close(ctx) + defer writer.CloseExecution(ctx) const attempts = 8 var wg sync.WaitGroup tokens := make(chan string, attempts) diff --git a/services/core/internal/store/environment_executor_management.go b/services/core/internal/store/environment_executor_management.go index 3c198f944..278135247 100644 --- a/services/core/internal/store/environment_executor_management.go +++ b/services/core/internal/store/environment_executor_management.go @@ -60,7 +60,7 @@ func (s *Store) ProjectExecutorCredentialState(ctx context.Context, project iden } environmentID := parsePathID(environment) result := ExecutorCredentialState{Credentials: []ExecutorCredential{}} - err = pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}, func(tx pgx.Tx) error { + err = s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.GetEnvironmentExecutorConnection(ctx, sqlc.GetEnvironmentExecutorConnectionParams{EnvironmentID: environmentID, TenantID: tenant}) if errors.Is(err, pgx.ErrNoRows) { @@ -144,7 +144,7 @@ func (s *Store) RevokeProjectExecutorCredential(ctx context.Context, project ide return err } record := executorCredentialAudit(project, "revoke", keyID) - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) n, err := q.RevokeExecutorCredential(ctx, sqlc.RevokeExecutorCredentialParams{KeyID: id, TenantID: tenant, SubjectKind: pgtype.Text{String: project.SubjectKind, Valid: true}, SubjectID: pgtype.Text{String: project.SubjectID, Valid: true}}) if err != nil { diff --git a/services/core/internal/store/environment_file_writes.go b/services/core/internal/store/environment_file_writes.go index 715e1b1c6..4a389219f 100644 --- a/services/core/internal/store/environment_file_writes.go +++ b/services/core/internal/store/environment_file_writes.go @@ -45,7 +45,10 @@ func (k FileWriteIdentity) valid() bool { // ReserveEnvironmentFileWrite persists intent before external dispatch. A retry // observes the earlier operation and never authorizes resending an unknown write. func (s *Store) ReserveEnvironmentFileWrite(ctx context.Context, tenant, environment string, key FileWriteIdentity) (EnvironmentFileWrite, error) { - if s.executionLease == nil || !key.valid() { + if err := s.checkExecutionAuthority(); err != nil { + return EnvironmentFileWrite{}, err + } + if !key.valid() { return EnvironmentFileWrite{}, ErrInvalidInput } owned, err := s.GetEnvironment(ctx, tenant, environment) @@ -147,7 +150,10 @@ func (s *Store) GetEnvironmentFileWrite(ctx context.Context, tenant, environment // SettleEnvironmentFileWrite requires an independently validated exact receipt. // A missing receipt, cancellation or owner retirement is not a rejected upload. func (s *Store) SettleEnvironmentFileWrite(ctx context.Context, tenant, environment string, key FileWriteIdentity, state string) (EnvironmentFileWrite, error) { - if s.executionLease == nil || !key.valid() || (state != "committed" && state != "rejected") { + if err := s.checkExecutionAuthority(); err != nil { + return EnvironmentFileWrite{}, err + } + if !key.valid() || (state != "committed" && state != "rejected") { return EnvironmentFileWrite{}, ErrInvalidInput } previous, err := s.GetEnvironmentFileWrite(ctx, tenant, environment, key.ID) diff --git a/services/core/internal/store/environment_file_writes_test.go b/services/core/internal/store/environment_file_writes_test.go index 531acf625..b7bc2d353 100644 --- a/services/core/internal/store/environment_file_writes_test.go +++ b/services/core/internal/store/environment_file_writes_test.go @@ -12,7 +12,6 @@ import ( type fileWriteFixture struct { s, writer *Store - lease *ExecutionLease tenant string session Session env Environment @@ -22,14 +21,14 @@ type fileWriteFixture struct { func newFileWriteFixture(t *testing.T) fileWriteFixture { t.Helper() s, _ := testStore(t) - lease := executionLease(t, s) + writer := executionWriter(t, s) tenant := uuid.NewString() session, env := localEnvironment(t, s, tenant) host, err := s.CreateEnvironmentDevice(t.Context(), tenant, env.ID, "file owner", runtimedevice.HashCredential(uuid.NewString())) if err != nil { t.Fatal(err) } - return fileWriteFixture{s: s, writer: lease.Store(), lease: lease, tenant: tenant, session: session, env: env, + return fileWriteFixture{s: s, writer: writer, tenant: tenant, session: session, env: env, key: FileWriteIdentity{ID: uuid.NewString(), DeviceID: host.ID, RequestSHA256: strings.Repeat("a", 64)}} } @@ -41,14 +40,14 @@ func TestEnvironmentFileWriteRetainsUnknownAcrossLeaseLoss(t *testing.T) { t.Fatal(first, err) } var killed bool - if err := f.s.pool.QueryRow(ctx, "SELECT pg_terminate_backend($1, 1000)", f.lease.conn.Conn().PgConn().PID()).Scan(&killed); err != nil || !killed { + if err := f.s.pool.QueryRow(ctx, "SELECT pg_terminate_backend($1, 1000)", executionOwnerPID(t, f.s.pool)).Scan(&killed); err != nil || !killed { t.Fatal(killed, err) } if _, err := f.writer.SettleEnvironmentFileWrite(ctx, f.tenant, f.env.ID, f.key, "committed"); err == nil { t.Fatal("lost writer settled an upload") } reopened, _ := testStore(t) - next := executionLease(t, reopened).Store() + next := executionWriter(t, reopened) got, err := next.ReserveEnvironmentFileWrite(ctx, f.tenant, f.env.ID, f.key) if err != nil || !got.Replayed || got.State != "pending" || !got.CreatedAt.Equal(first.CreatedAt) || got.Identity != f.key { t.Fatal("restart lost unknown write identity", got, err) diff --git a/services/core/internal/store/environment_initial_input_test.go b/services/core/internal/store/environment_initial_input_test.go index 0d8f7f815..d2d3d90b1 100644 --- a/services/core/internal/store/environment_initial_input_test.go +++ b/services/core/internal/store/environment_initial_input_test.go @@ -46,7 +46,7 @@ func TestEnvironmentInitialExpiryRollsBackWithFailureEventAndSerializesPromotion t.Cleanup(func() { _, _ = pool.Exec(context.Background(), "ALTER TABLE session_events DROP CONSTRAINT IF EXISTS "+constraint) }) - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) if _, err := writer.ExpireEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID); err == nil { t.Fatal("expiry committed without its failure event") } @@ -140,7 +140,7 @@ func TestEnvironmentInitialInputCreationRetainsCursorIdentityAndPromotion(t *tes if _, err := other.CreateSession(t.Context(), tenant, changed); !errors.Is(err, ErrIdempotencyConflict) { t.Fatal("changed creator accepted", err) } - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) generation := uuid.NewString() if err := writer.ReplaceEnvironmentConnection(t.Context(), tenant, session.Environment.ID, generation); err != nil { t.Fatal(err) @@ -207,8 +207,7 @@ func TestEnvironmentInitialInputExpiryHasNoTurnAndCannotReplay(t *testing.T) { if _, err := pool.Exec(t.Context(), "UPDATE environment_input_reservations SET deadline=clock_timestamp()-interval '1 second' WHERE id=$1", reservation.ID); err != nil { t.Fatal(err) } - lease := executionLease(t, s) - writer := lease.Store() + writer := executionWriter(t, s) for reservation.State == EnvironmentInputPending { count, err := writer.ExpireEnvironmentInputs(t.Context()) if err != nil || count < 1 || count > 32 { @@ -228,12 +227,12 @@ func TestEnvironmentInitialInputExpiryHasNoTurnAndCannotReplay(t *testing.T) { if err != nil || len(events) != expectedEvents || events[len(events)-1].Event.Type != "agent.session.failed" || events[len(events)-1].Turn != nil || events[len(events)-1].EnvironmentInputActivity.Status != "failed" { t.Fatal("missing pre-Turn failure snapshot", events, err) } - if err := lease.Close(t.Context()); err != nil { + if err := writer.CloseExecution(t.Context()); err != nil { t.Fatal(err) } pool.Close() reopened, reopenedPool := testStore(t) - writer = executionLease(t, reopened).Store() + writer = executionWriter(t, reopened) if _, err := reopened.CreateSession(t.Context(), tenant, input); err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/environment_initial_public_test.go b/services/core/internal/store/environment_initial_public_test.go index 34108ba43..15cfe4b73 100644 --- a/services/core/internal/store/environment_initial_public_test.go +++ b/services/core/internal/store/environment_initial_public_test.go @@ -40,14 +40,14 @@ func TestEnvironmentInitialFailureOfficialClient(t *testing.T) { if err != nil { t.Fatal(err) } - lease, err := s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } t.Cleanup(func() { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - if err := lease.Close(ctx); err != nil { + if err := writer.CloseExecution(ctx); err != nil { t.Error(err) } }) @@ -99,7 +99,7 @@ func TestEnvironmentInitialFailureOfficialClient(t *testing.T) { if err := pool.QueryRow(t.Context(), "UPDATE environment_input_reservations SET deadline=clock_timestamp()-interval '1 second' WHERE session_id=$1 AND is_initial RETURNING id", session.ID).Scan(&reservation); err != nil { t.Fatal(err) } - if result, err := lease.Store().ExpireEnvironmentInput(t.Context(), tenant, session.ID, reservation); err != nil || result.State != store.EnvironmentInputExpired { + if result, err := writer.ExpireEnvironmentInput(t.Context(), tenant, session.ID, reservation); err != nil || result.State != store.EnvironmentInputExpired { t.Fatal("initial reservation did not expire", result, err) } observed := <-done diff --git a/services/core/internal/store/environment_initialization.go b/services/core/internal/store/environment_initialization.go index 62adb9d4e..6964e321e 100644 --- a/services/core/internal/store/environment_initialization.go +++ b/services/core/internal/store/environment_initialization.go @@ -39,8 +39,8 @@ func (s *Store) ListEnvironmentInitializations(ctx context.Context, after string } func (s *Store) mutateEnvironmentInitialization(ctx context.Context, owner EnvironmentInitialization, apply func(*sqlc.Queries, sqlc.GetSessionEnvironmentRow) error) error { - if s.executionLease == nil { - return errors.New("Environment initialization requires execution ownership") + if err := s.checkExecutionAuthority(); err != nil { + return err } tenant, err := parseID(owner.TenantID) if err != nil { diff --git a/services/core/internal/store/environment_initialization_test.go b/services/core/internal/store/environment_initialization_test.go index 046a557c1..446445bc6 100644 --- a/services/core/internal/store/environment_initialization_test.go +++ b/services/core/internal/store/environment_initialization_test.go @@ -197,12 +197,11 @@ func TestEnvironmentInitializationRevocationBeforeClaim(t *testing.T) { return store.EnvironmentInitialization{EnvironmentID: environment.ID, SessionID: session.ID, TenantID: principal.TenantID, DeviceID: enrolled.DeviceID, State: "pending", Engine: "codex"} } revoked, other := create(), create() - lease, err := s.AcquireExecutionLease(t.Context()) + owned, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } - defer lease.Close(context.Background()) - owned := lease.Store() + defer owned.CloseExecution(context.Background()) if err := s.RevokeDevice(t.Context(), principal.TenantID, revoked.DeviceID); err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/environment_input_activity_test.go b/services/core/internal/store/environment_input_activity_test.go index 54280f2fa..b1509e9bb 100644 --- a/services/core/internal/store/environment_input_activity_test.go +++ b/services/core/internal/store/environment_input_activity_test.go @@ -50,7 +50,7 @@ func TestEnvironmentInputActivityWaitsBeforeTurnAndClearsOnConnection(t *testing t.Fatal("missing pre-Turn snapshot", first, err) } reserveEnvironmentInput(t, s, tenant, session.ID, "waiting") - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) generation := uuid.NewString() if err := writer.ReplaceEnvironmentConnection(t.Context(), tenant, environment, generation); err != nil { t.Fatal(err) @@ -100,7 +100,7 @@ func TestEnvironmentInputActivitySettlementAndNewerWork(t *testing.T) { t.Run(state, func(t *testing.T) { s, pool := testStore(t) tenant, session := environmentInputSession(t, s) - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) prior, err := s.SubmitInputs(t.Context(), tenant, session.ID, "prior", []Input{messageInput("prior")}) if err != nil { t.Fatal(err) @@ -183,7 +183,7 @@ func TestEnvironmentInputActivityRollsBackReservationAndConnection(t *testing.T) } reserveEnvironmentInput(t, s, tenant, session.ID, "waiting") value := requireEnvironmentInputActivity(t, s, tenant, session.ID, "requires_action", session.Environment.ID) - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) generation := uuid.NewString() if err := writer.ReplaceEnvironmentConnection(t.Context(), tenant, value.Environment.ID, generation); err != nil { t.Fatal(err) @@ -205,20 +205,20 @@ func TestEnvironmentInputActivityRecoversWaitingActionAndHidesDeletion(t *testin s, pool := testStore(t) tenant, session := environmentInputSession(t, s) reservation := reserveEnvironmentInput(t, s, tenant, session.ID, "waiting") - old := executionLease(t, s) + old := executionWriter(t, s) generation := uuid.NewString() environment := session.Environment.ID - if err := old.Store().ReplaceEnvironmentConnection(t.Context(), tenant, environment, generation); err != nil { + if err := old.ReplaceEnvironmentConnection(t.Context(), tenant, environment, generation); err != nil { t.Fatal(err) } - if err := old.Store().ObserveEnvironmentConnection(t.Context(), tenant, environment, generation, 1, true); err != nil { + if err := old.ObserveEnvironmentConnection(t.Context(), tenant, environment, generation, 1, true); err != nil { t.Fatal(err) } requireEnvironmentInputActivity(t, s, tenant, session.ID, "idle", "") - if err := old.Close(t.Context()); err != nil { + if err := old.CloseExecution(t.Context()); err != nil { t.Fatal(err) } - next := executionLease(t, s).Store() + next := executionWriter(t, s) if err := next.ReconcileEnvironmentConnections(t.Context()); err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/environment_input_claim_test.go b/services/core/internal/store/environment_input_claim_test.go index ec009f673..bcc75f0ef 100644 --- a/services/core/internal/store/environment_input_claim_test.go +++ b/services/core/internal/store/environment_input_claim_test.go @@ -7,7 +7,7 @@ import ( func TestEnvironmentInputConcurrentPromotionClaimsOnce(t *testing.T) { s, pool := testStore(t) - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) tenant, session := environmentInputSession(t, s) pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") reservationCursor, err := s.SessionEventCursor(t.Context(), tenant, session.ID) diff --git a/services/core/internal/store/environment_input_expiry.go b/services/core/internal/store/environment_input_expiry.go index 60c01ce25..6e0548ae8 100644 --- a/services/core/internal/store/environment_input_expiry.go +++ b/services/core/internal/store/environment_input_expiry.go @@ -2,7 +2,6 @@ package store import ( "context" - "errors" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" "github.com/jackc/pgx/v5" @@ -11,13 +10,11 @@ import ( // ExpireEnvironmentInputs settles one bounded batch without creating Turn history. // Only the current execution writer may run this cross-Session maintenance. func (s *Store) ExpireEnvironmentInputs(ctx context.Context) (int64, error) { - if s.executionLease == nil { - return 0, errors.New("Environment input expiry requires an execution lease") + if err := s.checkExecutionAuthority(); err != nil { + return 0, err } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() var expired int64 - err := s.executionLease.transaction(ctx, func(tx pgx.Tx) error { + err := s.writer.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) rows, err := q.ListDueEnvironmentInputs(ctx) if err != nil { diff --git a/services/core/internal/store/environment_input_expiry_test.go b/services/core/internal/store/environment_input_expiry_test.go index c04195403..2704dbf2a 100644 --- a/services/core/internal/store/environment_input_expiry_test.go +++ b/services/core/internal/store/environment_input_expiry_test.go @@ -27,8 +27,8 @@ func TestEnvironmentExpiryBoundsBatchAndRequiresExecutionWriter(t *testing.T) { if err := pool.QueryRow(ctx, "SELECT count(*) FILTER (WHERE state='expired'), count(*) FILTER (WHERE state='pending' AND deadline <= statement_timestamp()) FROM environment_input_reservations").Scan(&before, &due); err != nil { t.Fatal(err) } - lease := executionLease(t, s) - n, err := lease.Store().ExpireEnvironmentInputs(ctx) + writer := executionWriter(t, s) + n, err := writer.ExpireEnvironmentInputs(ctx) if err != nil || n != 32 { t.Fatal("unbounded or incomplete batch", n, err) } @@ -37,7 +37,7 @@ func TestEnvironmentExpiryBoundsBatchAndRequiresExecutionWriter(t *testing.T) { } // Existing test rows may precede this fixture; each pass must make bounded progress. for remaining := due - n; remaining > 0; { - n, err = lease.Store().ExpireEnvironmentInputs(ctx) + n, err = writer.ExpireEnvironmentInputs(ctx) if err != nil || n <= 0 || n > 32 { t.Fatal("expiry backlog did not progress", n, err) } @@ -59,13 +59,13 @@ func TestEnvironmentExpiryFencesLostExecutionOwner(t *testing.T) { if _, err := pool.Exec(t.Context(), "UPDATE environment_input_reservations SET deadline=clock_timestamp()-interval '1 second' WHERE id=$1", pending.ID); err != nil { t.Fatal(err) } - old := executionLease(t, s) + old := executionWriter(t, s) var killed bool - if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1,1000)", old.conn.Conn().PgConn().PID()).Scan(&killed); err != nil || !killed { + if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1,1000)", executionOwnerPID(t, pool)).Scan(&killed); err != nil || !killed { t.Fatal(killed, err) } - successor := executionLease(t, s) - if n, err := old.Store().ExpireEnvironmentInputs(t.Context()); err == nil || n != 0 { + successor := executionWriter(t, s) + if n, err := old.ExpireEnvironmentInputs(t.Context()); err == nil || n != 0 { t.Fatal("lost owner expired input", n, err) } got, err := s.GetEnvironmentInputReservation(t.Context(), tenant, session.ID, pending.ID) @@ -73,7 +73,7 @@ func TestEnvironmentExpiryFencesLostExecutionOwner(t *testing.T) { t.Fatal("lost owner wrote through the pool", got, err) } for got.State == EnvironmentInputPending { - n, err := successor.Store().ExpireEnvironmentInputs(t.Context()) + n, err := successor.ExpireEnvironmentInputs(t.Context()) if err != nil || n == 0 { t.Fatal("successor could not expire input", n, err) } @@ -97,16 +97,16 @@ func TestEnvironmentExpirySerializesWithTargetedSettlement(t *testing.T) { if _, err := pool.Exec(t.Context(), "UPDATE environment_input_reservations SET deadline=clock_timestamp()-interval '1 second' WHERE id=$1", pending.ID); err != nil { t.Fatal(err) } - lease := executionLease(t, s) + writer := executionWriter(t, s) start := make(chan struct{}) results := make(chan error, 2) - go func() { <-start; _, err := lease.Store().ExpireEnvironmentInputs(t.Context()); results <- err }() + go func() { <-start; _, err := writer.ExpireEnvironmentInputs(t.Context()); results <- err }() go func() { <-start var err error switch action { case "promote": - _, err = lease.Store().PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) + _, err = writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) case "cancel": _, err = s.CancelEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) case "delete": @@ -129,19 +129,19 @@ func TestEnvironmentExpirySerializesWithTargetedSettlement(t *testing.T) { } environmentInputHistory(t, pool, session.ID, 0, 0) if action == "delete" { - if _, err := lease.Store().PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID); !errors.Is(err, ErrNotFound) { + if _, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID); !errors.Is(err, ErrNotFound) { t.Fatal("deleted input resurrected", err) } return } later := reserveEnvironmentInput(t, s, tenant, session.ID, uuid.NewString()) - for _, settle := range []func(context.Context, string, string, string) (EnvironmentInputReservation, error){lease.Store().PromoteEnvironmentInput, s.CancelEnvironmentInput, s.ExpireEnvironmentInput} { + for _, settle := range []func(context.Context, string, string, string) (EnvironmentInputReservation, error){writer.PromoteEnvironmentInput, s.CancelEnvironmentInput, s.ExpireEnvironmentInput} { old, err := settle(t.Context(), tenant, session.ID, pending.ID) if err != nil || old.State != state { t.Fatal("old reservation changed", old, err) } } - if _, err := lease.Store().ExpireEnvironmentInputs(t.Context()); err != nil { + if _, err := writer.ExpireEnvironmentInputs(t.Context()); err != nil { t.Fatal(err) } got, err := s.GetEnvironmentInputReservation(t.Context(), tenant, session.ID, later.ID) diff --git a/services/core/internal/store/environment_input_migration_test.go b/services/core/internal/store/environment_input_migration_test.go index c3a82ff5c..e670d627e 100644 --- a/services/core/internal/store/environment_input_migration_test.go +++ b/services/core/internal/store/environment_input_migration_test.go @@ -98,26 +98,25 @@ func TestEnvironmentInputPromotionUsesCurrentExecutionWriter(t *testing.T) { if _, err := s.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID); err == nil { t.Fatal("pooled Store promoted input without execution ownership") } - closed := executionLease(t, s) - if err := closed.Close(t.Context()); err != nil { + closed := executionWriter(t, s) + if err := closed.CloseExecution(t.Context()); err != nil { t.Fatal(err) } - if _, err := closed.Store().PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID); err == nil { + if _, err := closed.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID); err == nil { t.Fatal("closed execution writer promoted pending input") } environmentInputHistory(t, pool, session.ID, 0, 0) - old := executionLease(t, s) - writer := old.Store() + writer := executionWriter(t, s) var killed bool - if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1, 1000)", old.conn.Conn().PgConn().PID()).Scan(&killed); err != nil || !killed { + if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1, 1000)", executionOwnerPID(t, pool)).Scan(&killed); err != nil || !killed { t.Fatal(killed, err) } - successor := executionLease(t, s) + successor := executionWriter(t, s) if _, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID); err == nil { t.Fatal("stale execution writer promoted pending input") } environmentInputHistory(t, pool, session.ID, 0, 0) - got, err := successor.Store().PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) + got, err := successor.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) if err != nil || got.State != EnvironmentInputAdmitted { t.Fatal("successor could not promote", got, err) } diff --git a/services/core/internal/store/environment_input_settlement_test.go b/services/core/internal/store/environment_input_settlement_test.go index 00d176de6..26adc66a6 100644 --- a/services/core/internal/store/environment_input_settlement_test.go +++ b/services/core/internal/store/environment_input_settlement_test.go @@ -14,8 +14,7 @@ func TestEnvironmentInputTerminalReservationsCannotRestart(t *testing.T) { for _, terminal := range []string{EnvironmentInputCancelled, EnvironmentInputExpired} { t.Run(terminal, func(t *testing.T) { s, pool := testStore(t) - lease := executionLease(t, s) - writer := lease.Store() + writer := executionWriter(t, s) tenant, session := environmentInputSession(t, s) ctx := context.Background() pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") @@ -67,8 +66,7 @@ func TestEnvironmentInputPromotionRollsBackHistoryAndSettlement(t *testing.T) { for _, phase := range []string{"input", "settlement", "claim", "claim-event"} { t.Run(phase, func(t *testing.T) { s, pool := testStore(t) - lease := executionLease(t, s) - writer := lease.Store() + writer := executionWriter(t, s) tenant, session := environmentInputSession(t, s) ctx := context.Background() pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") @@ -115,8 +113,7 @@ func TestEnvironmentInputDeadlineIsCheckedAfterSessionLock(t *testing.T) { for _, action := range []string{"promote", "fail"} { t.Run(action, func(t *testing.T) { s, pool := testStore(t) - lease := executionLease(t, s) - writer := lease.Store() + writer := executionWriter(t, s) tenant, session := environmentInputSession(t, s) pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) @@ -186,8 +183,7 @@ func TestEnvironmentInputDeadlineIsCheckedAfterSessionLock(t *testing.T) { func TestEnvironmentInputCancelAndPromotionShareOneOutcome(t *testing.T) { s, pool := testStore(t) - lease := executionLease(t, s) - writer := lease.Store() + writer := executionWriter(t, s) other, _ := testStore(t) tenant, session := environmentInputSession(t, s) ctx := context.Background() @@ -226,8 +222,7 @@ func TestEnvironmentInputDeletionSettlesPendingAndFencesPromotion(t *testing.T) for _, concurrent := range []bool{false, true} { t.Run(map[bool]string{false: "pending", true: "racing-promotion"}[concurrent], func(t *testing.T) { s, pool := testStore(t) - lease := executionLease(t, s) - writer := lease.Store() + writer := executionWriter(t, s) tenant, session := environmentInputSession(t, s) ctx := context.Background() pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") diff --git a/services/core/internal/store/environment_inputs.go b/services/core/internal/store/environment_inputs.go index dc702c19b..7d2edd454 100644 --- a/services/core/internal/store/environment_inputs.go +++ b/services/core/internal/store/environment_inputs.go @@ -164,8 +164,8 @@ func (s *Store) GetEnvironmentInputReservation(ctx context.Context, tenantID, se // PromoteEnvironmentInput admits and claims work for the retained native preparation. func (s *Store) PromoteEnvironmentInput(ctx context.Context, tenantID, sessionID, reservationID string) (EnvironmentInputReservation, error) { - if s.executionLease == nil { - return EnvironmentInputReservation{}, errors.New("Environment input promotion requires an execution lease") + if err := s.checkExecutionAuthority(); err != nil { + return EnvironmentInputReservation{}, err } return s.settleEnvironmentInput(ctx, tenantID, sessionID, reservationID, EnvironmentInputAdmitted) } diff --git a/services/core/internal/store/environment_inputs_test.go b/services/core/internal/store/environment_inputs_test.go index 12e2b30be..55cd01b17 100644 --- a/services/core/internal/store/environment_inputs_test.go +++ b/services/core/internal/store/environment_inputs_test.go @@ -111,8 +111,7 @@ func TestEnvironmentInputReservationConcurrentIdentity(t *testing.T) { func TestEnvironmentInputReservationPromotionAndDirectRetries(t *testing.T) { s, pool := testStore(t) - lease := executionLease(t, s) - writer := lease.Store() + writer := executionWriter(t, s) tenant, session := environmentInputSession(t, s) ctx := context.Background() first := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") @@ -166,12 +165,12 @@ func TestEnvironmentInputReservationPromotionAndDirectRetries(t *testing.T) { if err != nil || len(retry) != 2 || !retry[0].Replayed || retry[0].Sequence != promoted.Receipts[0].Sequence { t.Fatal("direct retry after promotion", retry, err) } - if err := lease.Close(ctx); err != nil { + if err := writer.CloseExecution(ctx); err != nil { t.Fatal(err) } pool.Close() restarted, pool := testStore(t) - after, err := executionLease(t, restarted).Store().PromoteEnvironmentInput(ctx, tenant, session.ID, first.ID) + after, err := executionWriter(t, restarted).PromoteEnvironmentInput(ctx, tenant, session.ID, first.ID) if err != nil || after.State != EnvironmentInputAdmitted || after.Receipts[0].Sequence != promoted.Receipts[0].Sequence { t.Fatal("restart repeated promotion", after, err) } @@ -209,8 +208,7 @@ func TestEnvironmentInputReservationKeepsEarlierDirectIdentity(t *testing.T) { func TestEnvironmentInputReservationRejectsUnsupportedOrForeignState(t *testing.T) { s, _ := testStore(t) - lease := executionLease(t, s) - writer := lease.Store() + writer := executionWriter(t, s) tenant, session := environmentInputSession(t, s) ctx := context.Background() pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") diff --git a/services/core/internal/store/environment_runtime_fixture_test.go b/services/core/internal/store/environment_runtime_fixture_test.go index 0d7df2277..5d2ac39c8 100644 --- a/services/core/internal/store/environment_runtime_fixture_test.go +++ b/services/core/internal/store/environment_runtime_fixture_test.go @@ -37,8 +37,9 @@ func enrollFixtureSession(t *testing.T, s *store.Store, tenant string, session s func connectFixtureRuntime(t *testing.T, h *dispatchHarness, session store.Session) *dispatchHarness { t.Helper() - other := *h - other.session = session + // The Runtime shares the harness's Core, not its connection or write lock. + other := &dispatchHarness{t: h.t, s: h.s, d: h.d, tenant: h.tenant, session: session, registry: h.registry, url: h.url, + admissions: h.admissions, environments: h.environments} other.device, other.credential = enrollFixtureSession(t, h.s, h.tenant, session) u, err := url.Parse(h.url) if err != nil { @@ -51,8 +52,8 @@ func connectFixtureRuntime(t *testing.T, h *dispatchHarness, session store.Sessi t.Fatal("enrolled Runtime connection failed") } t.Cleanup(func() { _ = other.conn.Close() }) - enableWorkerEnvironment(t, &other) - return &other + enableWorkerEnvironment(t, other) + return other } func assertPreparationReleased(t *testing.T, h *dispatchHarness, request, handle string) { diff --git a/services/core/internal/store/environment_steering_order_test.go b/services/core/internal/store/environment_steering_order_test.go index 537dbf78f..f275b431c 100644 --- a/services/core/internal/store/environment_steering_order_test.go +++ b/services/core/internal/store/environment_steering_order_test.go @@ -15,7 +15,7 @@ func TestEnvironmentActiveInputSerializesWithCompletion(t *testing.T) { } t.Run(name, func(t *testing.T) { s, pool := testStore(t) - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) tenant, session := environmentInputSession(t, s) original := submitMessage(t, s, tenant, session.ID, "original") transition(t, s, tenant, session.ID, original.TurnID, TurnQueued, TurnInProgress) diff --git a/services/core/internal/store/environment_templates.go b/services/core/internal/store/environment_templates.go index f0faa5920..0301d8236 100644 --- a/services/core/internal/store/environment_templates.go +++ b/services/core/internal/store/environment_templates.go @@ -105,7 +105,7 @@ func (s *Store) CreateEnvironmentTemplate(ctx context.Context, tenantID string, return EnvironmentTemplate{}, err } var result EnvironmentTemplate - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.CreateEnvironmentTemplate(ctx, sqlc.CreateEnvironmentTemplateParams{ID: pgtype.UUID{Bytes: id, Valid: true}, TenantID: tenant, Name: name, NetworkAccess: in.NetworkAccess, NetworkAllowedDomains: append([]string{}, in.AllowedDomains...), Files: metadata, FileContents: encrypted, Packages: packages, EnvContents: envContents, SetupContents: setupContents, Skills: skills, SkillContents: skillContents, Plugins: plugins, PluginContents: pluginContents, CapabilityDirectories: append([]string{}, in.Initialization.CapabilityDirectories...)}) result, err = templateFromRow(templateMetadataRow(row), err) @@ -162,7 +162,7 @@ func (s *Store) UpdateEnvironmentTemplate(ctx context.Context, tenantID, templat return EnvironmentTemplate{}, err } var result EnvironmentTemplate - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.UpdateEnvironmentTemplate(ctx, sqlc.UpdateEnvironmentTemplateParams{TenantID: tenant, ID: id, Name: name, SetName: in.SetName, NetworkAccess: in.NetworkAccess, NetworkAllowedDomains: append([]string{}, in.AllowedDomains...), SetNetwork: in.SetNetwork, SetFiles: in.SetFiles, Files: metadata, FileContents: encrypted, Packages: packages, EnvContents: envContents, SetupContents: setupContents, SetPackages: in.SetPackages, SetEnv: in.SetEnv, SetSetup: in.SetSetup, SetSkills: in.SetSkills, SetPlugins: in.SetPlugins, SetDirectories: in.SetDirectories, Skills: skills, SkillContents: skillContents, Plugins: plugins, PluginContents: pluginContents, CapabilityDirectories: append([]string{}, in.Initialization.CapabilityDirectories...)}) result, err = templateFromRow(templateMetadataRow(row), err) @@ -184,7 +184,7 @@ func (s *Store) DeleteEnvironmentTemplate(ctx context.Context, tenantID, templat return "", ErrNotFound } var deletedID string - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) result, err := q.DeleteEnvironmentTemplate(ctx, sqlc.DeleteEnvironmentTemplateParams{TenantID: tenant, ID: id}) if err != nil { diff --git a/services/core/internal/store/environment_work_test.go b/services/core/internal/store/environment_work_test.go index b3024194e..cb4010ce3 100644 --- a/services/core/internal/store/environment_work_test.go +++ b/services/core/internal/store/environment_work_test.go @@ -59,9 +59,7 @@ func TestEnvironmentInputWorkFiltersAndPagesDevices(t *testing.T) { t.Fatal("unconnected work selected", work, err) } } - foreign := *h - foreign.tenant = uuid.NewString() - unboundWorkerEnvironmentReservation(t, &foreign) + unboundWorkerEnvironmentReservation(t, &dispatchHarness{s: h.s, tenant: uuid.NewString()}) seen, cursor := 0, "" for _, count := range []int{100, 4, 0} { devices := []string{h.device.ID} diff --git a/services/core/internal/store/environment_write_audit_test.go b/services/core/internal/store/environment_write_audit_test.go index db078cd02..7de751b8a 100644 --- a/services/core/internal/store/environment_write_audit_test.go +++ b/services/core/internal/store/environment_write_audit_test.go @@ -21,10 +21,10 @@ func TestEnvironmentUploadWriteAuditSurvivesRequestAndLeaseContext(t *testing.T) cancel() sessionAuditCount(t, f.s, f.tenant, 0) // Neither the settlement caller nor the newly acquired execution lease owns the request context. - if err := f.lease.Close(context.Background()); err != nil { + if err := f.writer.CloseExecution(context.Background()); err != nil { t.Fatal(err) } - next := executionLease(t, f.s).Store() + next := executionWriter(t, f.s) var group sync.WaitGroup for range 4 { group.Go(func() { diff --git a/services/core/internal/store/execution.go b/services/core/internal/store/execution.go new file mode 100644 index 000000000..bf56baf97 --- /dev/null +++ b/services/core/internal/store/execution.go @@ -0,0 +1,68 @@ +package store + +import ( + "context" + "errors" + + "github.com/jackc/pgx/v5" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" +) + +// ErrExecutionAuthority rejects an execution-only operation on a pooled Store. +var ErrExecutionAuthority = errors.New("operation requires the execution writer") + +// transactor runs one transaction: pgunit's Pool or its execution Lease. +type transactor interface { + Transaction(context.Context, func(context.Context, pgx.Tx) error) error +} + +// NewExecution takes the database's execution lease on a dedicated connection +// and returns the execution writer built on it. The writer's Session and +// execution-only transactions run on the leased connection; reads keep the pool. +// Keep public admission on the pooled Store. Losing or closing the lease never +// falls back to a pooled writer. It fails when another service owns the database. +func NewExecution(ctx context.Context, s *Store) (*Store, error) { + lease, err := pgunit.AcquireLease(ctx, s.pool) + if err != nil { + return nil, err + } + writer := *s + writer.writer, writer.lease = lease, lease + return &writer, nil +} + +// checkExecutionAuthority only validates. The connection was fixed when the +// Store was constructed. +func (s *Store) checkExecutionAuthority() error { + if s.lease == nil { + return ErrExecutionAuthority + } + return nil +} + +// CheckExecutionOwnership validates the current writer before external preparation. +func (s *Store) CheckExecutionOwnership(ctx context.Context) error { + if err := s.checkExecutionAuthority(); err != nil { + return err + } + return s.lease.CheckOwnership(ctx) +} + +// CancelExecutionOperations cancels coordinator-owned contexts between leased +// operations; see pgunit.Lease.CancelOperations for the constraints on cancel. +func (s *Store) CancelExecutionOperations(ctx context.Context, cancel context.CancelFunc) error { + if err := s.checkExecutionAuthority(); err != nil { + return err + } + return s.lease.CancelOperations(ctx, cancel) +} + +// CloseExecution releases the execution lease and waits for connection cleanup +// within ctx. A later call resumes that wait. The writer cannot write afterwards. +func (s *Store) CloseExecution(ctx context.Context) error { + if err := s.checkExecutionAuthority(); err != nil { + return err + } + return s.lease.Close(ctx) +} diff --git a/services/core/internal/store/execution_cancellation_test.go b/services/core/internal/store/execution_cancellation_test.go deleted file mode 100644 index 695bfb59c..000000000 --- a/services/core/internal/store/execution_cancellation_test.go +++ /dev/null @@ -1,87 +0,0 @@ -package store - -import ( - "context" - "errors" - "testing" - "time" -) - -func TestExecutionCancellationFenceHonorsBounds(t *testing.T) { - s, _ := testStore(t) - lease := executionLease(t, s) - for _, test := range []struct { - name string - limit, maximum time.Duration - }{ - {"caller", 25 * time.Millisecond, time.Second}, - {"operation", 8 * time.Second, 6 * time.Second}, - } { - t.Run(test.name, func(t *testing.T) { - if err := lease.lock(t.Context()); err != nil { - t.Fatal(err) - } - held := true - defer func() { - if held { - lease.unlock() - } - }() - operation, stopOperation := context.WithCancel(t.Context()) - defer stopOperation() - ctx, cancel := context.WithTimeout(t.Context(), test.limit) - defer cancel() - started := time.Now() - err := lease.Store().CancelExecutionOperations(ctx, stopOperation) - if !errors.Is(err, context.DeadlineExceeded) || time.Since(started) > test.maximum { - t.Fatal("cancellation fence did not preserve its deadline", err, time.Since(started)) - } - if operation.Err() != nil { - t.Fatal("timed-out fence canceled operations outside the lease gate") - } - lease.unlock() - held = false - if err := lease.Store().CheckExecutionOwnership(t.Context()); err != nil { - t.Fatal("gate timeout damaged the healthy owner", err) - } - if err := lease.Store().CancelExecutionOperations(t.Context(), stopOperation); err != nil { - t.Fatal(err) - } - if !errors.Is(operation.Err(), context.Canceled) { - t.Fatal("successful fence did not cancel synchronously") - } - if err := lease.Store().CheckExecutionOwnership(t.Context()); err != nil { - t.Fatal("cancellation damaged the healthy owner", err) - } - }) - } -} - -func TestExecutionCancellationFenceRejectsLostLease(t *testing.T) { - s, pool := testStore(t) - lease := executionLease(t, s) - var killed bool - if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1, 1000)", lease.conn.Conn().PgConn().PID()).Scan(&killed); err != nil || !killed { - t.Fatal(killed, err) - } - operation, cancel := context.WithCancel(t.Context()) - defer cancel() - if err := lease.Store().CancelExecutionOperations(t.Context(), cancel); err == nil { - t.Fatal("lost lease accepted cancellation fence") - } - if operation.Err() != nil { - t.Fatal("lost lease ran the cancellation callback") - } - if err := lease.Store().CheckExecutionOwnership(t.Context()); err == nil { - t.Fatal("lost owner became writable") - } - if err := lease.Close(t.Context()); err != nil { - t.Fatal(err) - } - if err := lease.Store().CancelExecutionOperations(t.Context(), cancel); err == nil { - t.Fatal("closed lease accepted cancellation fence") - } - if err := s.CancelExecutionOperations(t.Context(), cancel); err == nil { - t.Fatal("pooled Store accepted execution cancellation") - } -} diff --git a/services/core/internal/store/execution_lease.go b/services/core/internal/store/execution_lease.go deleted file mode 100644 index 31a669e6b..000000000 --- a/services/core/internal/store/execution_lease.go +++ /dev/null @@ -1,139 +0,0 @@ -package store - -import ( - "context" - "errors" - "time" - - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgxpool" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" -) - -const executionTransactionTimeout = 5 * time.Second - -// ExecutionLease owns the connection used for execution writes, not just election. -// Its gate serializes pgx operations; no daemon or model work holds this gate. -type ExecutionLease struct { - conn *pgxpool.Conn - gate chan struct{} - writer Store - cleanupDone <-chan struct{} -} - -// AcquireExecutionLease enforces the gateway's single-service ownership per database. -func (s *Store) AcquireExecutionLease(ctx context.Context) (*ExecutionLease, error) { - conn, err := s.pool.Acquire(ctx) - if err != nil { - return nil, err - } - acquired, err := sqlc.New(conn).TryExecutionLease(ctx) - if err != nil || !acquired { - _ = conn.Hijack().Close(context.Background()) - if err != nil { - return nil, err - } - return nil, errors.New("another execution service owns this database") - } - lease := &ExecutionLease{conn: conn, gate: make(chan struct{}, 1), writer: *s} - lease.writer.executionLease = lease - return lease, nil -} - -// Store returns the execution writer view. Session transactions use the leased -// connection; reads retain the pool. Keep public admission on the original Store. -// Losing or closing the lease never falls back to a pooled writer connection. -func (l *ExecutionLease) Store() *Store { return &l.writer } - -// CheckExecutionOwnership validates the current writer before external preparation. -func (s *Store) CheckExecutionOwnership(ctx context.Context) error { - if s.executionLease == nil { - return errors.New("execution operation requires a leased Store") - } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() - return s.executionLease.Ping(ctx) -} - -// CancelExecutionOperations cancels coordinator-owned contexts between leased -// operations. Canceling an in-flight pgx operation can close the connection that -// owns the execution advisory lock. The callback must only invoke synchronous -// context cancel functions; it must not perform database, provider or wait work. -// Caller cancellation and operation deadlines keep their existing semantics. -func (s *Store) CancelExecutionOperations(ctx context.Context, cancelOperations context.CancelFunc) error { - if s.executionLease == nil || cancelOperations == nil { - return ErrInvalidInput - } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() - return s.executionLease.withConn(ctx, func(conn *pgxpool.Conn) error { - if err := conn.Ping(ctx); err != nil { - return err - } - cancelOperations() - return nil - }) -} - -func (l *ExecutionLease) lock(ctx context.Context) error { - select { - case l.gate <- struct{}{}: - if err := ctx.Err(); err != nil { - l.unlock() - return err - } - return nil - case <-ctx.Done(): - return ctx.Err() - } -} - -func (l *ExecutionLease) unlock() { <-l.gate } - -func (l *ExecutionLease) withConn(ctx context.Context, apply func(*pgxpool.Conn) error) error { - if err := l.lock(ctx); err != nil { - return err - } - defer l.unlock() - if l.conn == nil { - return errors.New("execution lease is closed") - } - return apply(l.conn) -} - -func (l *ExecutionLease) transaction(ctx context.Context, apply func(pgx.Tx) error) error { - return l.withConn(ctx, func(conn *pgxpool.Conn) error { - return pgx.BeginFunc(ctx, conn, apply) - }) -} - -func (l *ExecutionLease) Ping(ctx context.Context) error { - return l.withConn(ctx, func(conn *pgxpool.Conn) error { return conn.Ping(ctx) }) -} - -func (l *ExecutionLease) Close(ctx context.Context) error { - if err := l.lock(ctx); err != nil { - return err - } - defer l.unlock() - if l.conn != nil { - conn := l.conn.Hijack() - l.conn = nil - l.cleanupDone = conn.PgConn().CleanupDone() - if err := conn.Close(ctx); err != nil { - return err - } - } - if l.cleanupDone == nil { - return nil - } - // A cancelled pgx connection can be unusable before its asynchronous cleanup ends. - // Retain the channel so a later Close can continue waiting after this deadline. - select { - case <-l.cleanupDone: - return nil - case <-ctx.Done(): - return ctx.Err() - } -} diff --git a/services/core/internal/store/execution_lease_test.go b/services/core/internal/store/execution_test.go similarity index 51% rename from services/core/internal/store/execution_lease_test.go rename to services/core/internal/store/execution_test.go index 3cae83f10..32ff3260f 100644 --- a/services/core/internal/store/execution_lease_test.go +++ b/services/core/internal/store/execution_test.go @@ -9,24 +9,42 @@ import ( "testing" "time" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimedevice" "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/runtimedevice" ) -func executionLease(t *testing.T, s *Store) *ExecutionLease { +// executionWriter builds the execution writer on the shared test database. +// Tests in this package run sequentially, so one writer at a time owns it. +func executionWriter(t *testing.T, s *Store) *Store { t.Helper() - lease, err := s.AcquireExecutionLease(t.Context()) + writer, err := NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } - t.Cleanup(func() { _ = lease.Close(context.Background()) }) - return lease + t.Cleanup(func() { _ = writer.CloseExecution(context.Background()) }) + return writer +} + +// executionOwnerPID finds the backend holding this database's execution lease, +// matching the single-bigint key in queries/scheduling.sql. +func executionOwnerPID(t *testing.T, pool *pgxpool.Pool) int32 { + t.Helper() + var pid int32 + err := pool.QueryRow(t.Context(), `SELECT pid FROM pg_locks WHERE locktype='advisory' AND granted AND objsubid=1 + AND classid::bigint * 4294967296 + objid::bigint = 706172736172 + AND database=(SELECT oid FROM pg_database WHERE datname=current_database())`).Scan(&pid) + if err != nil { + t.Fatal("observe execution lease owner", err) + } + return pid } func TestExecutionLeaseLossFencesAllLifecycleWrites(t *testing.T) { s, pool := testStore(t) - old := executionLease(t, s) - writer := old.Store() + writer := executionWriter(t, s) tenant, active := newTurnSession(t, s) input := submitMessage(t, s, tenant, active.ID, "active") transition(t, writer, tenant, active.ID, input.TurnID, TurnQueued, TurnInProgress) @@ -63,12 +81,12 @@ func TestExecutionLeaseLossFencesAllLifecycleWrites(t *testing.T) { if err != nil { t.Fatal(err) } - // Kill only this test's owner connection. Do not notify the old writer by Ping. + // Kill only this test's owner connection. Do not notify the old writer by a check. var killed bool - if err = pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1, 1000)", old.conn.Conn().PgConn().PID()).Scan(&killed); err != nil || !killed { + if err = pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1, 1000)", executionOwnerPID(t, pool)).Scan(&killed); err != nil || !killed { t.Fatal(killed, err) } - successor := executionLease(t, s) + successor := executionWriter(t, s) mustReject := func(name string, err error) { t.Helper() if err == nil { @@ -85,6 +103,9 @@ func TestExecutionLeaseLossFencesAllLifecycleWrites(t *testing.T) { mustReject("completion", err) _, err = writer.TransitionTurn(t.Context(), tenant, active.ID, input.TurnID, TurnTransition{ExpectedStatus: TurnInProgress, Status: TurnFailed}) mustReject("reconciliation", err) + _, err = writer.ExpireEnvironmentInputs(t.Context()) + mustReject("input expiry", err) + mustReject("ownership check", writer.CheckExecutionOwnership(t.Context())) after, err := s.GetSession(t.Context(), tenant, active.ID) if err != nil || !reflect.DeepEqual(before, after) { t.Fatal("stale state persisted", after, err) @@ -110,29 +131,43 @@ func TestExecutionLeaseLossFencesAllLifecycleWrites(t *testing.T) { } // Public admission remains usable with a dead owner connection. submitMessage(t, s, tenant, queued.ID, "additional") - if err = successor.Ping(t.Context()); err != nil { + if err = successor.CheckExecutionOwnership(t.Context()); err != nil { t.Fatal(err) } - if _, err = successor.Store().CompleteExecution(t.Context(), tenant, active.ID, input.TurnID, TurnCompleted, json.RawMessage(`{"done":{"content":"accepted"}}`), "successor-native", input.Sequence); err != nil { + if _, err = successor.CompleteExecution(t.Context(), tenant, active.ID, input.TurnID, TurnCompleted, json.RawMessage(`{"done":{"content":"accepted"}}`), "successor-native", input.Sequence); err != nil { t.Fatal(err) } bound, err := s.GetSessionExecutionBinding(t.Context(), tenant, active.ID) if err != nil || bound.NativeSessionID != "successor-native" { t.Fatal(bound, err) } - _, err = successor.Store().TransitionTurn(t.Context(), tenant, active.ID, input.TurnID, TurnTransition{ExpectedStatus: TurnInProgress, Status: TurnFailed}) + _, err = successor.TransitionTurn(t.Context(), tenant, active.ID, input.TurnID, TurnTransition{ExpectedStatus: TurnInProgress, Status: TurnFailed}) if !errors.Is(err, ErrTurnConflict) { t.Fatal("terminal CAS changed", err) } - if err = successor.Close(t.Context()); err != nil { + if err = successor.CloseExecution(t.Context()); err != nil { t.Fatal(err) } - mustReject("closed writer", successor.Store().BindSessionDevice(t.Context(), tenant, queued.ID, host.ID)) + mustReject("closed writer", successor.BindSessionDevice(t.Context(), tenant, queued.ID, host.ID)) } -func TestExecutionLeaseSerializesWritesAndPings(t *testing.T) { - s, _ := testStore(t) - lease := executionLease(t, s) +func TestExecutionWriterSerializesWritesOnItsLease(t *testing.T) { + s, pool := testStore(t) + writer := executionWriter(t, s) + owner := executionOwnerPID(t, pool) + backend := func(store *Store) int32 { + t.Helper() + var pid int32 + if err := store.writer.Transaction(t.Context(), func(ctx context.Context, tx pgx.Tx) error { + return tx.QueryRow(ctx, "SELECT pg_backend_pid()").Scan(&pid) + }); err != nil { + t.Fatal(err) + } + return pid + } + if backend(writer) != owner || backend(s) == owner { + t.Fatal("execution transactions do not run on the leased connection") + } type work struct{ tenant, session, turn string } tasks := make([]work, 4) for i := range tasks { @@ -143,13 +178,13 @@ func TestExecutionLeaseSerializesWritesAndPings(t *testing.T) { results := make(chan error, len(tasks)*2) for _, task := range tasks { group.Go(func() { - _, err := lease.Store().TransitionTurn(t.Context(), task.tenant, task.session, task.turn, TurnTransition{ExpectedStatus: TurnQueued, Status: TurnInProgress}) + _, err := writer.TransitionTurn(t.Context(), task.tenant, task.session, task.turn, TurnTransition{ExpectedStatus: TurnQueued, Status: TurnInProgress}) if err == nil { - err = lease.Store().AppendTurnEvents(t.Context(), task.tenant, task.session, task.turn, 1, []ExecutionEvent{{Kind: "delta", Payload: json.RawMessage(`{"delta":"accepted"}`)}}) + err = writer.AppendTurnEvents(t.Context(), task.tenant, task.session, task.turn, 1, []ExecutionEvent{{Kind: "delta", Payload: json.RawMessage(`{"delta":"accepted"}`)}}) } results <- err }) - group.Go(func() { results <- lease.Ping(t.Context()) }) + group.Go(func() { results <- writer.CheckExecutionOwnership(t.Context()) }) } group.Wait() close(results) @@ -159,49 +194,74 @@ func TestExecutionLeaseSerializesWritesAndPings(t *testing.T) { } } // While the owner connection is in use, public admission on another Session - // does not need that connection. A waiting owner operation can be cancelled. - if err := lease.lock(t.Context()); err != nil { - t.Fatal(err) - } - defer lease.unlock() + // does not need that connection. + entered, release := make(chan struct{}), make(chan struct{}) + held := make(chan error, 1) + go func() { + held <- writer.writer.Transaction(t.Context(), func(context.Context, pgx.Tx) error { + close(entered) + <-release + return nil + }) + }() + <-entered ctx, cancel := context.WithTimeout(t.Context(), time.Second) defer cancel() task := tasks[0] - if _, err := s.SubmitMessage(ctx, task.tenant, task.session, "public", json.RawMessage(`{"text":"additional"}`)); err != nil { + _, err := s.SubmitMessage(ctx, task.tenant, task.session, "public", json.RawMessage(`{"text":"additional"}`)) + close(release) + if err != nil { t.Fatal("public admission used owner gate", err) } - cancelled, stop := context.WithCancel(t.Context()) - stop() - if err := lease.Ping(cancelled); !errors.Is(err, context.Canceled) { - t.Fatal("gate wait ignored cancellation", err) + if err := <-held; err != nil { + t.Fatal(err) } } -func TestExecutionLeaseBoundsSessionLockWait(t *testing.T) { - s, pool := testStore(t) - lease := executionLease(t, s) +func TestPooledStoreHasNoExecutionAuthority(t *testing.T) { + s, _ := testStore(t) tenant, session := newTurnSession(t, s) input := submitMessage(t, s, tenant, session.ID, "start") - blocker, err := pool.Begin(t.Context()) + before, err := s.GetSession(t.Context(), tenant, session.ID) if err != nil { t.Fatal(err) } - defer blocker.Rollback(context.Background()) - if _, err = blocker.Exec(t.Context(), "SELECT id FROM sessions WHERE id=$1 FOR UPDATE", session.ID); err != nil { + cursor, err := s.SessionEventCursor(t.Context(), tenant, session.ID) + if err != nil { t.Fatal(err) } - ctx, cancel := context.WithTimeout(t.Context(), 8*time.Second) + operation, cancel := context.WithCancel(t.Context()) defer cancel() - start := time.Now() - _, err = lease.Store().TransitionTurn(ctx, tenant, session.ID, input.TurnID, TurnTransition{ExpectedStatus: TurnQueued, Status: TurnInProgress}) - if !errors.Is(err, context.DeadlineExceeded) || time.Since(start) >= 7*time.Second { - t.Fatal("owner transaction did not enforce its shorter deadline", err) + _, archiveErr := s.ArchiveManagedSession(t.Context(), tenant, session.ID, 0) + _, expiryErr := s.ExpireEnvironmentInputs(t.Context()) + subagent := []ExecutionEvent{{Kind: "delta", Payload: json.RawMessage(`{"delta":"visible"}`)}, subagentIdentityEvent("child", "root", 100)} + for name, err := range map[string]error{ + "ownership check": s.CheckExecutionOwnership(t.Context()), + "cancellation": s.CancelExecutionOperations(t.Context(), cancel), + "close": s.CloseExecution(t.Context()), + "archive": archiveErr, + "input expiry": expiryErr, + "deployment": s.ConfigureRuntimeDeployment(t.Context(), nil), + "reconciliation": s.ReconcileEnvironmentConnections(t.Context()), + "mixed batch": s.AppendTurnEvents(t.Context(), tenant, session.ID, input.TurnID, 1, subagent), + "subagent only": s.AppendTurnEvents(t.Context(), tenant, session.ID, input.TurnID, 1, subagent[1:]), + "connection change": s.ReplaceEnvironmentConnection(t.Context(), tenant, uuid.NewString(), uuid.NewString()), + } { + if !errors.Is(err, ErrExecutionAuthority) { + t.Fatalf("pooled Store ran %s: %v", name, err) + } } - if err = blocker.Rollback(t.Context()); err != nil { - t.Fatal(err) + if operation.Err() != nil { + t.Fatal("pooled Store ran the cancellation callback") + } + after, err := s.GetSession(t.Context(), tenant, session.ID) + if err != nil || !reflect.DeepEqual(before, after) { + t.Fatal("rejected execution operation changed the Session", after, err) + } + if next, err := s.SessionEventCursor(t.Context(), tenant, session.ID); err != nil || next != cursor { + t.Fatal("rejected execution operation published events", next, err) } - turn, err := s.GetTurn(t.Context(), tenant, session.ID, input.TurnID) - if err != nil || turn.Status != TurnQueued { - t.Fatal("timed-out claim changed queued work", turn, err) + if events, err := s.ListTurnEvents(t.Context(), tenant, session.ID, input.TurnID, 0, 100); err != nil || len(events) != 0 { + t.Fatal("rejected batch reached the journal", events, err) } } diff --git a/services/core/internal/store/executor_credential_target.go b/services/core/internal/store/executor_credential_target.go index e7241f686..1ca267618 100644 --- a/services/core/internal/store/executor_credential_target.go +++ b/services/core/internal/store/executor_credential_target.go @@ -47,7 +47,7 @@ func issuedExecutorCredential(id, environment pgtype.UUID, token string) IssuedE func (s *Store) withExecutorCredentialTarget(ctx context.Context, principal identity.Principal, environment pgtype.UUID, apply func(context.Context, *sqlc.Queries) error) error { if !environment.Valid { - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { return apply(ctx, s.queries.WithTx(tx)) }) + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { return apply(ctx, s.queries.WithTx(tx)) }) } owned, err := s.GetEnvironment(ctx, principal.TenantID, uuid.UUID(environment.Bytes).String()) if err != nil { diff --git a/services/core/internal/store/list_cursor_public_test.go b/services/core/internal/store/list_cursor_public_test.go index cc4e7fd04..a8a1ef6d3 100644 --- a/services/core/internal/store/list_cursor_public_test.go +++ b/services/core/internal/store/list_cursor_public_test.go @@ -239,13 +239,13 @@ func TestListCursorErrorsPostgres(t *testing.T) { server := httptest.NewServer(h) defer server.Close() client := pathIDClient{t: t, server: server} - lease, err := s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } - defer func() { _ = lease.Close(t.Context()) }() - a := seedCursorFixture(t, s, lease.Store(), client, owner, ownerTenant, "a") - b := seedCursorFixture(t, s, lease.Store(), client, foreign, foreignTenant, "b") + defer func() { _ = writer.CloseExecution(t.Context()) }() + a := seedCursorFixture(t, s, writer, client, owner, ownerTenant, "a") + b := seedCursorFixture(t, s, writer, client, foreign, foreignTenant, "b") text := func(value string) *string { return &value } var ( diff --git a/services/core/internal/store/managed_test_store_test.go b/services/core/internal/store/managed_test_store_test.go index 31789268e..b6c2db27f 100644 --- a/services/core/internal/store/managed_test_store_test.go +++ b/services/core/internal/store/managed_test_store_test.go @@ -1,54 +1,18 @@ package store import ( - "context" - "database/sql" - "os" "testing" - "time" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" - "github.com/jackc/pgx/v5/stdlib" - "github.com/pressly/goose/v3" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" ) // Deployment identity belongs to a whole database, so managed fixtures cannot // share the ordinary Store fixture database or bypass production startup checks. func newManagedTestStore(t *testing.T) (*Store, *pgxpool.Pool) { t.Helper() - _, admin := testStore(t) - name := "oac_m_" + uuid.NewString()[:8] + "_tests" - quoted := pgx.Identifier{name}.Sanitize() - if _, err := admin.Exec(t.Context(), "CREATE DATABASE "+quoted); err != nil { - t.Fatal(err) - } - t.Cleanup(func() { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - if _, err := admin.Exec(ctx, "DROP DATABASE "+quoted+" WITH (FORCE)"); err != nil { - t.Error(err) - } - }) - cfg := admin.Config().Copy() - cfg.ConnConfig.Database = name - db := sql.OpenDB(stdlib.GetConnector(*cfg.ConnConfig)) - provider, err := goose.NewProvider(goose.DialectPostgres, db, os.DirFS("../../migrations"), goose.WithTableName("agents_api_schema_version")) - if err != nil { - _ = db.Close() - t.Fatal(err) - } - _, err = provider.Up(t.Context()) - _ = db.Close() - if err != nil { - t.Fatal(err) - } - pool, err := pgxpool.NewWithConfig(t.Context(), cfg) - if err != nil { - t.Fatal(err) - } - t.Cleanup(pool.Close) + pool := pgtest.OpenIsolated(t, nil) // Hosted Sessions freeze a model provider, which needs a credential key. return NewWithCredentialCipher(pool, fixtureCipher), pool } diff --git a/services/core/internal/store/mcp_credentials_oauth.go b/services/core/internal/store/mcp_credentials_oauth.go index ac55182b7..4abdb853f 100644 --- a/services/core/internal/store/mcp_credentials_oauth.go +++ b/services/core/internal/store/mcp_credentials_oauth.go @@ -5,6 +5,8 @@ import ( "errors" "time" + "github.com/jackc/pgx/v5" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" ) @@ -15,55 +17,58 @@ func (s *Store) oauthBearerToken(ctx context.Context, tenantID string, binding M // hold a credential indefinitely. The HTTP client supplies its own tighter bound. ctx, cancel := context.WithTimeout(ctx, oauthRefreshTimeout) defer cancel() - tx, credential, secret, err := s.lockOAuth(ctx, tenantID, binding.VaultID, binding.CredentialID, binding.ServerURL) - if err != nil { - return "", err - } - defer tx.Rollback(context.Background()) - if credential.MCPServerURL != binding.ServerURL { - return "", ErrNotFound - } - if secret.Metadata.ExpiresAt == nil { - return oauthAccessToken(secret.AccessToken) - } - expiry, err := time.Parse(time.RFC3339Nano, *secret.Metadata.ExpiresAt) - if err != nil { - return "", errors.New("invalid OAuth token expiry") - } - if time.Now().Before(expiry) { - return oauthAccessToken(secret.AccessToken) - } - refresh := secret.Metadata.Refresh - if refresh == nil || secret.RefreshToken == "" || s.oauthRefresher == nil { - return "", errors.New("expired OAuth credential cannot be refreshed") - } - token, err := s.oauthRefresher.Refresh(ctx, oauthrefresh.Request{ - TokenEndpoint: refresh.TokenEndpoint, ClientID: refresh.ClientID, - AuthMethod: refresh.TokenEndpointAuth, ClientSecret: secret.ClientSecret, - RefreshToken: secret.RefreshToken, Resource: refresh.Resource, Scope: refresh.Scope, + var bearer string + err := s.withOAuth(ctx, tenantID, binding.VaultID, binding.CredentialID, binding.ServerURL, "OAuth credential refresh commit failed", func(ctx context.Context, tx pgx.Tx, credential Credential, secret oauthSecret) error { + if credential.MCPServerURL != binding.ServerURL { + return ErrNotFound + } + if secret.Metadata.ExpiresAt == nil { + var err error + bearer, err = oauthAccessToken(secret.AccessToken) + return err + } + expiry, err := time.Parse(time.RFC3339Nano, *secret.Metadata.ExpiresAt) + if err != nil { + return errors.New("invalid OAuth token expiry") + } + if time.Now().Before(expiry) { + bearer, err = oauthAccessToken(secret.AccessToken) + return err + } + refresh := secret.Metadata.Refresh + if refresh == nil || secret.RefreshToken == "" || s.oauthRefresher == nil { + return errors.New("expired OAuth credential cannot be refreshed") + } + token, err := s.oauthRefresher.Refresh(ctx, oauthrefresh.Request{ + TokenEndpoint: refresh.TokenEndpoint, ClientID: refresh.ClientID, + AuthMethod: refresh.TokenEndpointAuth, ClientSecret: secret.ClientSecret, + RefreshToken: secret.RefreshToken, Resource: refresh.Resource, Scope: refresh.Scope, + }) + if err != nil { + return errors.New("OAuth credential refresh failed") + } + if token.AccessToken == "" || token.ExpiresAt != nil && !time.Now().Before(*token.ExpiresAt) { + return errors.New("OAuth refresh returned an unusable token") + } + secret.AccessToken = token.AccessToken + if token.RefreshToken != "" { + secret.RefreshToken = token.RefreshToken + } + secret.Metadata.ExpiresAt = nil + if token.ExpiresAt != nil { + value := token.ExpiresAt.UTC().Format(time.RFC3339Nano) + secret.Metadata.ExpiresAt = &value + } + if _, err := s.saveOAuth(ctx, tx, tenantID, credential, secret); err != nil { + return err + } + bearer = token.AccessToken + return nil }) if err != nil { - return "", errors.New("OAuth credential refresh failed") - } - if token.AccessToken == "" || token.ExpiresAt != nil && !time.Now().Before(*token.ExpiresAt) { - return "", errors.New("OAuth refresh returned an unusable token") - } - secret.AccessToken = token.AccessToken - if token.RefreshToken != "" { - secret.RefreshToken = token.RefreshToken - } - secret.Metadata.ExpiresAt = nil - if token.ExpiresAt != nil { - value := token.ExpiresAt.UTC().Format(time.RFC3339Nano) - secret.Metadata.ExpiresAt = &value - } - if _, err := s.saveOAuth(ctx, tx, tenantID, credential, secret); err != nil { return "", err } - if err := tx.Commit(ctx); err != nil { - return "", errors.New("OAuth credential refresh commit failed") - } - return token.AccessToken, nil + return bearer, nil } func oauthAccessToken(token string) (string, error) { diff --git a/services/core/internal/store/prepared_dispatch_test.go b/services/core/internal/store/prepared_dispatch_test.go index 7b9efaca5..b9ab4573d 100644 --- a/services/core/internal/store/prepared_dispatch_test.go +++ b/services/core/internal/store/prepared_dispatch_test.go @@ -22,12 +22,12 @@ func preparedDispatchHarness(t *testing.T) (*dispatchHarness, store.EnvironmentI t.Helper() h := newDispatchHarnessForSession(t, []byte(`{"agent":{"model":"test-model","instructions":"Keep this instruction."},"environment":{"type":"self_hosted","workspace_directory":"/workspace"}}`), true) assertNoRuntimeAllocation(t, h) - lease, err := h.s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), h.s) if err != nil { t.Fatal(err) } - t.Cleanup(func() { _ = lease.Close(context.Background()) }) - h.d.Store = lease.Store() + t.Cleanup(func() { _ = writer.CloseExecution(context.Background()) }) + h.d.Store = writer enableWorkerEnvironment(t, h) pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "pending", []store.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}, {Kind: "message", Payload: json.RawMessage(`{"text":"second"}`)}}) if err != nil { diff --git a/services/core/internal/store/project_api_keys.go b/services/core/internal/store/project_api_keys.go index c9ef26654..c4224b526 100644 --- a/services/core/internal/store/project_api_keys.go +++ b/services/core/internal/store/project_api_keys.go @@ -77,7 +77,7 @@ func (s *Store) CreateProjectAPIKey(ctx context.Context, project, id, name strin return IssuedProjectAPIKey{}, err } var result IssuedProjectAPIKey - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) p, err := q.LockProject(ctx, projectID) if errors.Is(err, pgx.ErrNoRows) { @@ -147,7 +147,7 @@ func (s *Store) RevokeProjectAPIKey(ctx context.Context, project, id string) err if err != nil { return ErrNotFound } - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) p, err := q.LockProject(ctx, projectID) if errors.Is(err, pgx.ErrNoRows) { diff --git a/services/core/internal/store/projects.go b/services/core/internal/store/projects.go index bbedcb381..37a7041a6 100644 --- a/services/core/internal/store/projects.go +++ b/services/core/internal/store/projects.go @@ -65,7 +65,7 @@ func (s *Store) CreateProject(ctx context.Context, id, name string) (Project, er } tenantID, _ := parseID(uuid.NewString()) var result Project - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) _, err := q.EnsureProjectScope(ctx, sqlc.EnsureProjectScopeParams{TenantID: tenantID, OrganizationID: "core", ProjectID: "proj_" + id}) if err != nil { @@ -144,7 +144,7 @@ func (s *Store) mutateProject(ctx context.Context, id, name string, archive bool return Project{}, ErrNotFound } var result Project - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) _, err := q.LockProjectForUpdate(ctx, projectID) if errors.Is(err, pgx.ErrNoRows) { diff --git a/services/core/internal/store/runtime_adoption_test.go b/services/core/internal/store/runtime_adoption_test.go index b6629fa1a..de32a0121 100644 --- a/services/core/internal/store/runtime_adoption_test.go +++ b/services/core/internal/store/runtime_adoption_test.go @@ -12,7 +12,7 @@ import ( func legacyAdoptionFixture(t *testing.T) (*Store, *Store, RuntimeDeployment, RuntimeAllocation) { t.Helper() s, _ := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) d := deploymentSelection() deploymentConfigure(t, w, &d) tenant := uuid.NewString() diff --git a/services/core/internal/store/runtime_allocation_state.go b/services/core/internal/store/runtime_allocation_state.go index 0ec4ef4dc..d4b637f36 100644 --- a/services/core/internal/store/runtime_allocation_state.go +++ b/services/core/internal/store/runtime_allocation_state.go @@ -110,8 +110,8 @@ func (s *Store) ReleaseRuntimeAllocation(ctx context.Context, owner RuntimeAlloc } func (s *Store) mutateRuntimeAllocation(ctx context.Context, owner RuntimeAllocation, live bool, apply func(context.Context, *sqlc.Queries, sqlc.RuntimeAllocation) (sqlc.RuntimeAllocation, error)) (RuntimeAllocation, error) { - if s.executionLease == nil { - return RuntimeAllocation{}, ErrInvalidInput + if err := s.checkExecutionAuthority(); err != nil { + return RuntimeAllocation{}, err } previous, err := s.GetRuntimeAllocation(ctx, owner.TenantID, owner.EnvironmentID) if err != nil { diff --git a/services/core/internal/store/runtime_allocations.go b/services/core/internal/store/runtime_allocations.go index 18bccfcdf..0c433b193 100644 --- a/services/core/internal/store/runtime_allocations.go +++ b/services/core/internal/store/runtime_allocations.go @@ -46,8 +46,8 @@ type RuntimeObservationSessionPage struct { // ReserveRuntimeAllocation commits the allocation and dedicated device together // before external Create. Only a fresh receipt authorizes that one Create call. func (s *Store) ReserveRuntimeAllocation(ctx context.Context, tenant, environment, providerKey, credentialHash string) (RuntimeAllocation, error) { - if s.executionLease == nil { - return RuntimeAllocation{}, ErrInvalidInput + if err := s.checkExecutionAuthority(); err != nil { + return RuntimeAllocation{}, err } provider, err := parseConnectionGeneration(providerKey) if err != nil { diff --git a/services/core/internal/store/runtime_allocations_test.go b/services/core/internal/store/runtime_allocations_test.go index eecc62bf6..575e72bc1 100644 --- a/services/core/internal/store/runtime_allocations_test.go +++ b/services/core/internal/store/runtime_allocations_test.go @@ -13,8 +13,7 @@ import ( func TestRuntimeAllocationAtomicOwnershipAndRecovery(t *testing.T) { s, pool := testStore(t) - lease := executionLease(t, s) - w := lease.Store() + w := executionWriter(t, s) tenant, provider := uuid.NewString(), uuid.NewString() session, environment := localEnvironment(t, s, tenant) secret := uuid.NewString() @@ -32,14 +31,14 @@ func TestRuntimeAllocationAtomicOwnershipAndRecovery(t *testing.T) { if _, err := w.ReserveRuntimeAllocation(t.Context(), tenant, environment.ID, uuid.NewString(), runtimedevice.HashCredential(secret)); !errors.Is(err, ErrIdempotencyConflict) { t.Fatalf("provider target changed: %v", err) } - if err := lease.Close(t.Context()); err != nil { + if err := w.CloseExecution(t.Context()); err != nil { t.Fatal(err) } if _, err := w.ObserveRuntimeRunning(t.Context(), owner); err == nil { t.Fatal("lost writer changed allocation") } reopened, _ := testStore(t) - next := executionLease(t, reopened).Store() + next := executionWriter(t, reopened) retry, err := next.ReserveRuntimeAllocation(t.Context(), tenant, environment.ID, provider, runtimedevice.HashCredential(uuid.NewString())) if err != nil || !retry.Replayed || retry.ID != owner.ID || retry.DeviceID != owner.DeviceID { t.Fatalf("restart replaced unknown allocation: %+v %v", retry, err) @@ -96,7 +95,7 @@ func TestRuntimeAllocationAtomicOwnershipAndRecovery(t *testing.T) { func TestRuntimeAllocationOneWinnerAndRollback(t *testing.T) { s, pool := testStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) tenant, provider := uuid.NewString(), uuid.NewString() _, environment := localEnvironment(t, s, tenant) var wg sync.WaitGroup @@ -153,7 +152,7 @@ func TestRuntimeAllocationOneWinnerAndRollback(t *testing.T) { func TestRuntimeAllocationExpiryAndRevocation(t *testing.T) { s, pool := testStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) tenant := uuid.NewString() _, environment := localEnvironment(t, s, tenant) owner, err := w.ReserveRuntimeAllocation(t.Context(), tenant, environment.ID, uuid.NewString(), runtimedevice.HashCredential(uuid.NewString())) diff --git a/services/core/internal/store/runtime_deployment.go b/services/core/internal/store/runtime_deployment.go index 156a49108..90cced855 100644 --- a/services/core/internal/store/runtime_deployment.go +++ b/services/core/internal/store/runtime_deployment.go @@ -27,8 +27,8 @@ type RuntimeDeployment struct { // AdmissionPaused must be committed for the old installation before any switch. // A nil selection never forgets the previous identity or unresolved resources. func (s *Store) ConfigureRuntimeDeployment(ctx context.Context, selected *RuntimeDeployment) error { - if s.executionLease == nil { - return ErrInvalidInput + if err := s.checkExecutionAuthority(); err != nil { + return err } if selected != nil { copy := *selected @@ -46,9 +46,7 @@ func (s *Store) ConfigureRuntimeDeployment(ctx context.Context, selected *Runtim } update = sqlc.SetRuntimeDeploymentParams{InstallationID: id, BackendFingerprint: selected.BackendFingerprint, AdmissionPaused: selected.AdmissionPaused} } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() - return s.executionLease.transaction(ctx, func(tx pgx.Tx) error { + return s.writer.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) previous, err := q.LockRuntimeDeployment(ctx) if err != nil { diff --git a/services/core/internal/store/runtime_deployment_test.go b/services/core/internal/store/runtime_deployment_test.go index 152bf4bb5..788b3646d 100644 --- a/services/core/internal/store/runtime_deployment_test.go +++ b/services/core/internal/store/runtime_deployment_test.go @@ -43,10 +43,10 @@ func legacyRuntimeSpecification(t *testing.T, w *Store, provider string) { func TestRuntimeDeploymentRequiresMaintenanceBeforeIdentityChange(t *testing.T) { s, pool := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) old := deploymentSelection() deploymentConfigure(t, w, &old) - if err := s.ConfigureRuntimeDeployment(t.Context(), &old); !errors.Is(err, ErrInvalidInput) { + if err := s.ConfigureRuntimeDeployment(t.Context(), &old); !errors.Is(err, ErrExecutionAuthority) { t.Fatal("unleased configuration accepted", err) } for _, changeID := range []bool{false, true} { @@ -89,7 +89,7 @@ func TestRuntimeDeploymentRequiresMaintenanceBeforeIdentityChange(t *testing.T) func TestRuntimeDeploymentPendingSessionsCannotMigrate(t *testing.T) { s, _ := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) tenant := uuid.NewString() session, _ := localEnvironment(t, s, tenant) old := deploymentSelection() @@ -113,7 +113,7 @@ func TestRuntimeDeploymentPendingSessionsCannotMigrate(t *testing.T) { func TestRuntimeDeploymentUnknownAllocationsBlockAdoptionAndSwitch(t *testing.T) { s, _ := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) old := deploymentSelection() tenant := uuid.NewString() session, environment := localEnvironment(t, s, tenant) @@ -169,7 +169,7 @@ func TestRuntimeDeploymentUnknownAllocationsBlockAdoptionAndSwitch(t *testing.T) func TestRuntimeDeploymentMaintenancePreservesCreationRetriesAndOtherPlacements(t *testing.T) { s, pool := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) old := deploymentSelection() deploymentConfigure(t, w, &old) tenant := uuid.NewString() @@ -214,7 +214,7 @@ func TestRuntimeDeploymentMaintenancePreservesCreationRetriesAndOtherPlacements( func TestRuntimeDeploymentMaintenanceSerializesHostedCreation(t *testing.T) { s, pool := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) config := deploymentSelection() deploymentConfigure(t, w, &config) ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) @@ -254,7 +254,7 @@ func TestRuntimeDeploymentRetainedResourcesBlockSwitchWithoutMutation(t *testing for _, state := range []string{"running", "suspended", "cleanup_pending"} { t.Run(state, func(t *testing.T) { s, pool := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) old := deploymentSelection() deploymentConfigure(t, w, &old) tenant := uuid.NewString() @@ -300,7 +300,7 @@ func TestRuntimeDeploymentRetainedResourcesBlockSwitchWithoutMutation(t *testing func TestRuntimeDeploymentAllocationBeforeMaintenanceRetainsOwnership(t *testing.T) { s, pool := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) config := deploymentSelection() deploymentConfigure(t, w, &config) tenant := uuid.NewString() diff --git a/services/core/internal/store/runtime_environment_terminal_test.go b/services/core/internal/store/runtime_environment_terminal_test.go index 8a1b6baaf..18e170a18 100644 --- a/services/core/internal/store/runtime_environment_terminal_test.go +++ b/services/core/internal/store/runtime_environment_terminal_test.go @@ -23,7 +23,7 @@ func TestManagedEnvironmentTerminationSettlesInputAndPreservesIdentity(t *testin t.Fatal(err) } reservation := initialEnvironmentReservation(t, s, pool, tenant, session.ID) - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) owner, err := writer.ReserveRuntimeAllocation(t.Context(), tenant, session.Environment.ID, uuid.NewString(), runtimedevice.HashCredential(uuid.NewString())) if err != nil { t.Fatal(err) @@ -118,7 +118,7 @@ func TestManagedEnvironmentFailureRollsBackWithSessionEvent(t *testing.T) { if err != nil { t.Fatal(err) } - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) owner, err := writer.ReserveRuntimeAllocation(t.Context(), tenant, session.Environment.ID, uuid.NewString(), runtimedevice.HashCredential(uuid.NewString())) if err != nil { t.Fatal(err) diff --git a/services/core/internal/store/runtime_lifecycle_nodes.go b/services/core/internal/store/runtime_lifecycle_nodes.go index 6f6bfaa73..90b7fb4ba 100644 --- a/services/core/internal/store/runtime_lifecycle_nodes.go +++ b/services/core/internal/store/runtime_lifecycle_nodes.go @@ -5,6 +5,7 @@ import ( "errors" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgtype" ) @@ -12,7 +13,7 @@ import ( // ListRuntimeLifecycleNodes includes offline nodes: loss of connectivity never // releases their resources. The empty identity is the single legacy lifecycle. func (s *Store) ListRuntimeLifecycleNodes(ctx context.Context) ([]string, error) { - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) + ctx, cancel := context.WithTimeout(ctx, pgunit.ExecutionTimeout) defer cancel() if err := s.CheckExecutionOwnership(ctx); err != nil { return nil, err @@ -50,7 +51,7 @@ func (s *Store) ListRuntimeAllocationsForNode(ctx context.Context, node, after s if err != nil { return nil, err } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) + ctx, cancel := context.WithTimeout(ctx, pgunit.ExecutionTimeout) defer cancel() if err := s.CheckExecutionOwnership(ctx); err != nil { return nil, err @@ -73,7 +74,7 @@ func (s *Store) ListUnallocatedHostedEnvironmentsForNode(ctx context.Context, no if err != nil { return nil, err } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) + ctx, cancel := context.WithTimeout(ctx, pgunit.ExecutionTimeout) defer cancel() if err := s.CheckExecutionOwnership(ctx); err != nil { return nil, err @@ -96,7 +97,7 @@ func (s *Store) ResolveRuntimeLifecycleNode(ctx context.Context, tenant, environ if err != nil { return "", err } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) + ctx, cancel := context.WithTimeout(ctx, pgunit.ExecutionTimeout) defer cancel() if err := s.CheckExecutionOwnership(ctx); err != nil { return "", err diff --git a/services/core/internal/store/runtime_lifecycle_nodes_test.go b/services/core/internal/store/runtime_lifecycle_nodes_test.go index eee1bde1a..8d05b176a 100644 --- a/services/core/internal/store/runtime_lifecycle_nodes_test.go +++ b/services/core/internal/store/runtime_lifecycle_nodes_test.go @@ -230,7 +230,7 @@ func TestRuntimeLifecycleLegacyLaneAndOwnerLoss(t *testing.T) { if _, err := w.ListUnallocatedHostedEnvironmentsForNode(t.Context(), "", "bad"); err == nil { t.Fatal("invalid cursor accepted") } - if err := w.executionLease.Close(t.Context()); err != nil { + if err := w.CloseExecution(t.Context()); err != nil { t.Fatal(err) } if _, err := w.ListRuntimeLifecycleNodes(t.Context()); err == nil { diff --git a/services/core/internal/store/runtime_node_generations.go b/services/core/internal/store/runtime_node_generations.go index 633e4d77f..e69a8f8b7 100644 --- a/services/core/internal/store/runtime_node_generations.go +++ b/services/core/internal/store/runtime_node_generations.go @@ -7,6 +7,7 @@ import ( "math" "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/sandbox" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sandbox/providers" "github.com/jackc/pgx/v5" @@ -51,7 +52,7 @@ func (s *Store) RuntimeNodeRetention(ctx context.Context, nodeID, connectionID s if len(refs) > 8 { return deployment, nil, ErrInvalidInput } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) + ctx, cancel := context.WithTimeout(ctx, pgunit.ExecutionTimeout) defer cancel() err := s.runtimeDeploymentTransaction(ctx, func(q *sqlc.Queries, d sqlc.RuntimeDeployment) error { n, err := nodeConnection(ctx, q, d, nodeID, connectionID, epoch) diff --git a/services/core/internal/store/runtime_node_presence.go b/services/core/internal/store/runtime_node_presence.go index bf6bc7670..eb54fd2e9 100644 --- a/services/core/internal/store/runtime_node_presence.go +++ b/services/core/internal/store/runtime_node_presence.go @@ -21,7 +21,7 @@ func (s *Store) ConnectRuntimeNode(ctx context.Context, nodeID, connectionID str } // A canceled autocommit UPDATE may still finish on PostgreSQL after pgx // returns. An explicit transaction cannot publish that late write. - return pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}, func(tx pgx.Tx) error { + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { changed, err := s.queries.WithTx(tx).ConnectRuntimeNode(ctx, sqlc.ConnectRuntimeNodeParams{ID: id, ConnectionID: connection, OwnerEpoch: int64(epoch)}) if err != nil { return err @@ -93,7 +93,7 @@ func (s *Store) DisconnectRuntimeNode(ctx context.Context, nodeID, connectionID if err != nil { return err } - return pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}, func(tx pgx.Tx) error { + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) // A Connect COMMIT may be uncertain. Wait on the node regardless of // the visible connection, then fence cleanup in a fresh statement snapshot. diff --git a/services/core/internal/store/runtime_nodes.go b/services/core/internal/store/runtime_nodes.go index 961acc311..a1e43cc93 100644 --- a/services/core/internal/store/runtime_nodes.go +++ b/services/core/internal/store/runtime_nodes.go @@ -53,7 +53,7 @@ func (s *Store) runtimeManagerTransaction(ctx context.Context, apply func(*sqlc. // initialized. Node machine routes use it to authenticate their credential // before reporting any deployment state, including an uninitialized one. func (s *Store) runtimeDeploymentTransaction(ctx context.Context, apply func(*sqlc.Queries, sqlc.RuntimeDeployment) error) error { - return pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) deployment, err := q.LockRuntimeDeployment(ctx) if err != nil { diff --git a/services/core/internal/store/runtime_nodes_test.go b/services/core/internal/store/runtime_nodes_test.go index 093651f2f..fbf73b66f 100644 --- a/services/core/internal/store/runtime_nodes_test.go +++ b/services/core/internal/store/runtime_nodes_test.go @@ -17,7 +17,7 @@ import ( func managerFixture(t *testing.T, active, retained int) (*Store, *Store, RuntimeDeployment) { t.Helper() s, _ := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) d := deploymentSelection() d.ProviderKind = "docker" d.LocalNodeID = uuid.NewString() diff --git a/services/core/internal/store/runtime_suspension_test.go b/services/core/internal/store/runtime_suspension_test.go index 53fc609c1..38b7b8766 100644 --- a/services/core/internal/store/runtime_suspension_test.go +++ b/services/core/internal/store/runtime_suspension_test.go @@ -16,7 +16,7 @@ import ( func runtimeSuspensionFixture(t *testing.T) (*Store, *Store, *pgxpool.Pool, RuntimeAllocation) { t.Helper() s, pool := testStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) tenant := uuid.NewString() _, environment := localEnvironment(t, s, tenant) owner, err := w.ReserveRuntimeAllocation(t.Context(), tenant, environment.ID, uuid.NewString(), runtimedevice.HashCredential(uuid.NewString())) @@ -261,7 +261,7 @@ func TestRuntimeSuspensionRetentionAndDeletedSession(t *testing.T) { func TestRuntimeSuspensionCountsUncertainCapacityUntilReleased(t *testing.T) { s, pool := testStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) provider := uuid.NewString() cases := []struct { state, phase string @@ -359,7 +359,7 @@ func TestRuntimeSuspensionExpiredRunningAndLostWriterAreFenced(t *testing.T) { if _, err := w.SetRuntimeCompute(t.Context(), owner, "quiescing", json.RawMessage(`{}`), &until, time.Nanosecond); !errors.Is(err, ErrTurnConflict) { t.Fatal("expired running allocation entered checkpoint", err) } - if err := w.executionLease.Close(t.Context()); err != nil { + if err := w.CloseExecution(t.Context()); err != nil { t.Fatal(err) } if _, err := w.RuntimeActivity(t.Context(), owner); err == nil { diff --git a/services/core/internal/store/runtime_worker_recovery_test.go b/services/core/internal/store/runtime_worker_recovery_test.go index 41676132a..5d257b61d 100644 --- a/services/core/internal/store/runtime_worker_recovery_test.go +++ b/services/core/internal/store/runtime_worker_recovery_test.go @@ -34,12 +34,12 @@ func runtimeWorkerHarness(t *testing.T) (*dispatchHarness, *pgxpool.Pool) { func TestPreparedDispatchKeepsPendingReservationAfterComputeConflict(t *testing.T) { h, _ := runtimeWorkerHarness(t) - lease, err := h.s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), h.s) if err != nil { t.Fatal(err) } - t.Cleanup(func() { _ = lease.Close(context.Background()) }) - h.d.Store = lease.Store() + t.Cleanup(func() { _ = writer.CloseExecution(context.Background()) }) + h.d.Store = writer pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "pending", []store.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) if err != nil { t.Fatal(err) diff --git a/services/core/internal/store/sandbox_deployment_mutations.go b/services/core/internal/store/sandbox_deployment_mutations.go index c53def61a..c6a4eac92 100644 --- a/services/core/internal/store/sandbox_deployment_mutations.go +++ b/services/core/internal/store/sandbox_deployment_mutations.go @@ -93,8 +93,8 @@ func recordConfigurationMetadata(ctx context.Context, q *sqlc.Queries, input San } func (s *Store) InitializeSandboxDeployment(ctx context.Context, installationID string, input SandboxDeploymentSetupRequest) (RuntimeDeploymentView, error) { - if s.executionLease == nil { - return RuntimeDeploymentView{}, ErrInvalidInput + if err := s.checkExecutionAuthority(); err != nil { + return RuntimeDeploymentView{}, err } if err := validateSandboxSelection(input); err != nil { return RuntimeDeploymentView{}, err @@ -103,10 +103,8 @@ func (s *Store) InitializeSandboxDeployment(ctx context.Context, installationID if err != nil { return RuntimeDeploymentView{}, err } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() var result RuntimeDeploymentView - err = s.executionLease.transaction(ctx, func(tx pgx.Tx) error { + err = s.writer.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) d, err := q.LockRuntimeDeployment(ctx) if err != nil { @@ -144,15 +142,13 @@ func (s *Store) InitializeSandboxDeployment(ctx context.Context, installationID // CheckSandboxDeploymentSwitch is a preliminary check only. The mutation repeats // it in the committing transaction; no database lock spans provider work. func (s *Store) CheckSandboxDeploymentSwitch(ctx context.Context, installation string, input SandboxDeploymentUpdateRequest) error { - if s.executionLease == nil { - return ErrInvalidInput + if err := s.checkExecutionAuthority(); err != nil { + return err } if err := validateSandboxSelection(input.SandboxDeploymentSetupRequest); err != nil { return err } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() - return s.executionLease.transaction(ctx, func(tx pgx.Tx) error { + return s.writer.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) d, err := q.LockRuntimeDeployment(ctx) if err != nil { @@ -182,16 +178,14 @@ func checkSandboxSwitch(ctx context.Context, q *sqlc.Queries, d sqlc.RuntimeDepl } func (s *Store) UpdateSandboxDeployment(ctx context.Context, installation string, input SandboxDeploymentUpdateRequest) (RuntimeDeploymentView, error) { - if s.executionLease == nil { - return RuntimeDeploymentView{}, ErrInvalidInput + if err := s.checkExecutionAuthority(); err != nil { + return RuntimeDeploymentView{}, err } if err := validateSandboxSelection(input.SandboxDeploymentSetupRequest); err != nil { return RuntimeDeploymentView{}, err } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() var result RuntimeDeploymentView - err := s.executionLease.transaction(ctx, func(tx pgx.Tx) error { + err := s.writer.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) d, err := q.LockRuntimeDeployment(ctx) if err != nil { diff --git a/services/core/internal/store/sandbox_deployment_resources_test.go b/services/core/internal/store/sandbox_deployment_resources_test.go index b4c9800b1..e8cc68700 100644 --- a/services/core/internal/store/sandbox_deployment_resources_test.go +++ b/services/core/internal/store/sandbox_deployment_resources_test.go @@ -16,7 +16,7 @@ func TestSandboxDeploymentMutationViewsIncludeActualResources(t *testing.T) { t.Fatal(err) } s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) installation := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), installation); err != nil { t.Fatal(err) diff --git a/services/core/internal/store/sandbox_deployment_setup.go b/services/core/internal/store/sandbox_deployment_setup.go index fc636b06b..8191cb5e1 100644 --- a/services/core/internal/store/sandbox_deployment_setup.go +++ b/services/core/internal/store/sandbox_deployment_setup.go @@ -10,6 +10,7 @@ import ( "strings" "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/sandbox" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sandbox/providers" "github.com/jackc/pgx/v5" @@ -32,7 +33,7 @@ type SandboxSetup struct { } func (s *Store) GetSandboxSetup(ctx context.Context) (SandboxSetup, error) { - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) + ctx, cancel := context.WithTimeout(ctx, pgunit.ExecutionTimeout) defer cancel() d, err := s.queries.GetRuntimeDeployment(ctx) if err != nil { @@ -76,13 +77,14 @@ func (s *Store) sandboxSetup(d sqlc.RuntimeDeployment) (SandboxSetup, error) { // ClaimWebSandboxDeployment runs exactly once per execution-owner startup. It // reserves the installation before selection and fences previous node presence. func (s *Store) ClaimWebSandboxDeployment(ctx context.Context, installationID string) error { + if err := s.checkExecutionAuthority(); err != nil { + return err + } id, err := parseConnectionGeneration(installationID) - if err != nil || s.executionLease == nil { + if err != nil { return ErrInvalidInput } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() - return s.executionLease.transaction(ctx, func(tx pgx.Tx) error { + return s.writer.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) d, err := q.LockRuntimeDeployment(ctx) if err != nil { diff --git a/services/core/internal/store/sandbox_deployment_setup_test.go b/services/core/internal/store/sandbox_deployment_setup_test.go index 96a7f858e..a1b67d014 100644 --- a/services/core/internal/store/sandbox_deployment_setup_test.go +++ b/services/core/internal/store/sandbox_deployment_setup_test.go @@ -35,8 +35,7 @@ func TestSandboxCoreURLValidation(t *testing.T) { func TestSandboxDeploymentSetupPersistsWithoutExecution(t *testing.T) { s, pool := newManagedTestStore(t) s.SetPublicURL("https://core.example") - lease := executionLease(t, s) - w := lease.Store() + w := executionWriter(t, s) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) @@ -71,10 +70,10 @@ func TestSandboxDeploymentSetupPersistsWithoutExecution(t *testing.T) { if err := pool.QueryRow(t.Context(), "SELECT (SELECT count(*) FROM sessions)+(SELECT count(*) FROM runtime_nodes)+(SELECT count(*) FROM runtime_allocations)+(SELECT count(*) FROM runtime_placements)").Scan(&sideEffects); err != nil || sideEffects != 0 { t.Fatal("setup or rejected admission created execution state", sideEffects, err) } - if err := lease.Close(context.Background()); err != nil { + if err := w.CloseExecution(context.Background()); err != nil { t.Fatal(err) } - restarted := executionLease(t, s).Store() + restarted := executionWriter(t, s) if err := restarted.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) } @@ -93,7 +92,7 @@ func TestSandboxDeploymentSetupPersistsWithoutExecution(t *testing.T) { func TestSandboxDeploymentSetupConcurrentSelection(t *testing.T) { s, _ := newManagedTestStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) @@ -142,7 +141,7 @@ func TestSandboxDeploymentSetupConcurrentSelection(t *testing.T) { func TestSandboxDeploymentSetupRejectsFileManagedAndUnleasedWrites(t *testing.T) { s, w, selection := managerFixture(t, 4, 16) input := SandboxDeploymentSetupRequest{DeploymentSpec: SandboxDeploymentTestSpec("docker"), Provider: "docker"} - if _, err := s.InitializeSandboxDeployment(t.Context(), selection.InstallationID, input); !errors.Is(err, ErrInvalidInput) { + if _, err := s.InitializeSandboxDeployment(t.Context(), selection.InstallationID, input); !errors.Is(err, ErrExecutionAuthority) { t.Fatal("unleased setup accepted", err) } if _, err := w.InitializeSandboxDeployment(t.Context(), selection.InstallationID, input); !errors.Is(err, ErrSandboxDeploymentConflict) { diff --git a/services/core/internal/store/sandbox_deployment_switch_test.go b/services/core/internal/store/sandbox_deployment_switch_test.go index b64a643d1..ed5f5fc8a 100644 --- a/services/core/internal/store/sandbox_deployment_switch_test.go +++ b/services/core/internal/store/sandbox_deployment_switch_test.go @@ -26,7 +26,7 @@ func TestSandboxE2BEndpointPersistenceAndOnlineSwitch(t *testing.T) { t.Fatal(err) } s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) @@ -57,7 +57,7 @@ func TestSandboxResetClearsCustomE2BEndpoint(t *testing.T) { t.Fatal(err) } s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) installation := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), installation); err != nil { t.Fatal(err) @@ -90,7 +90,7 @@ func TestSandboxDirectDeploymentOwnershipAndCleanSwitch(t *testing.T) { t.Fatal(err) } s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) @@ -175,7 +175,7 @@ func TestSandboxSwitchRetiresNodesAndEnrollment(t *testing.T) { _, pool := newManagedTestStore(t) cipher, _ := credentialcrypto.New(bytes.Repeat([]byte{5}, 32)) s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) @@ -248,7 +248,7 @@ func TestSandboxResetSerializesFreshDirectSessions(t *testing.T) { _, pool := newManagedTestStore(t) cipher, _ := credentialcrypto.New(bytes.Repeat([]byte{6}, 32)) s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) @@ -291,7 +291,7 @@ func TestSandboxSwitchPreservesReleasedAllocationAndItemHistory(t *testing.T) { _, pool := newManagedTestStore(t) cipher, _ := credentialcrypto.New(bytes.Repeat([]byte{8}, 32)) s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) installation := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), installation); err != nil { t.Fatal(err) @@ -375,7 +375,7 @@ func TestUnspecifiedNodeDeploymentRejectedWithoutMutation(t *testing.T) { _, pool := newManagedTestStore(t) cipher, _ := credentialcrypto.New(bytes.Repeat([]byte{7}, 32)) s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) diff --git a/services/core/internal/store/sandbox_deployment_view_test.go b/services/core/internal/store/sandbox_deployment_view_test.go index 21da06627..7626d2a0e 100644 --- a/services/core/internal/store/sandbox_deployment_view_test.go +++ b/services/core/internal/store/sandbox_deployment_view_test.go @@ -20,7 +20,7 @@ func TestSandboxDeploymentViewRecordsTemplateBuildAndSuspension(t *testing.T) { t.Fatal(err) } s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) diff --git a/services/core/internal/store/sandbox_generations.go b/services/core/internal/store/sandbox_generations.go index 8107adedf..fd5742fb8 100644 --- a/services/core/internal/store/sandbox_generations.go +++ b/services/core/internal/store/sandbox_generations.go @@ -42,69 +42,70 @@ func (s *Store) ClassifySandboxDeploymentChange(ctx context.Context, installatio // GetSandboxAllocationSetup reads immutable ownership and the current credential // in one snapshot. A released receipt remains historical but is never rebound. func (s *Store) GetSandboxAllocationSetup(ctx context.Context, ref sandbox.Reference) (SandboxSetup, error) { - tx, err := s.pool.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}) - if err != nil { - return SandboxSetup{}, err - } - defer tx.Rollback(ctx) - q := s.queries.WithTx(tx) - lookup, err := deviceLookup(ref.TenantID, ref.EnvironmentID) - if err != nil { - return SandboxSetup{}, err - } - a, err := q.GetRuntimeAllocation(ctx, sqlc.GetRuntimeAllocationParams{TenantID: lookup.TenantID, EnvironmentID: lookup.ID}) - if err != nil { - return SandboxSetup{}, err - } - if runtimeUUID(a.RuntimeAllocation.ID) != ref.AllocationID || !a.RuntimeAllocation.DeploymentGeneration.Valid || a.RuntimeAllocation.State == "released" { - return SandboxSetup{}, ErrInvalidInput - } - d, err := q.GetRuntimeDeployment(ctx) - if err != nil { - return SandboxSetup{}, err - } - if a.RuntimeAllocation.ProviderKey != d.InstallationID { - return SandboxSetup{}, sandbox.ErrOwnership - } - result, err := s.sandboxSetup(d) - if err != nil { - return SandboxSetup{}, err - } - generation := a.RuntimeAllocation.DeploymentGeneration.Int64 - if generation != d.Generation { - g, err := q.GetSandboxGeneration(ctx, generation) - if errors.Is(err, pgx.ErrNoRows) { - return SandboxSetup{}, ErrSandboxDeploymentConflict - } + var result SandboxSetup + err := s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { + q := s.queries.WithTx(tx) + lookup, err := deviceLookup(ref.TenantID, ref.EnvironmentID) if err != nil { - return SandboxSetup{}, err + return err } - if g.ProviderKind != d.ProviderKind { - return SandboxSetup{}, ErrSandboxDeploymentConflict + a, err := q.GetRuntimeAllocation(ctx, sqlc.GetRuntimeAllocationParams{TenantID: lookup.TenantID, EnvironmentID: lookup.ID}) + if err != nil { + return err } - result.Generation = uint64(generation) - if err = json.Unmarshal(g.Specification, &result.Specification); err != nil { - return SandboxSetup{}, err + if runtimeUUID(a.RuntimeAllocation.ID) != ref.AllocationID || !a.RuntimeAllocation.DeploymentGeneration.Valid || a.RuntimeAllocation.State == "released" { + return ErrInvalidInput } - retained, err := providers.Decode(g.ProviderKind, sandbox.ConfigurationRecord{Public: g.ProviderConfig, Metadata: g.ProviderMetadata}) + d, err := q.GetRuntimeDeployment(ctx) if err != nil { - return SandboxSetup{}, ErrSandboxDeploymentConflict + return err } - needsCredential, err := providers.UsesCredential(g.ProviderKind) + if a.RuntimeAllocation.ProviderKey != d.InstallationID { + return sandbox.ErrOwnership + } + result, err = s.sandboxSetup(d) if err != nil { - return SandboxSetup{}, err + return err } - if needsCredential { - composed, err := providers.WithCredential(sandbox.Selection{Provider: g.ProviderKind, Configuration: retained}, sandbox.Selection{Provider: d.ProviderKind, Configuration: result.Configuration}) + generation := a.RuntimeAllocation.DeploymentGeneration.Int64 + if generation != d.Generation { + g, err := q.GetSandboxGeneration(ctx, generation) + if errors.Is(err, pgx.ErrNoRows) { + return ErrSandboxDeploymentConflict + } + if err != nil { + return err + } + if g.ProviderKind != d.ProviderKind { + return ErrSandboxDeploymentConflict + } + result.Generation = uint64(generation) + if err = json.Unmarshal(g.Specification, &result.Specification); err != nil { + return err + } + retained, err := providers.Decode(g.ProviderKind, sandbox.ConfigurationRecord{Public: g.ProviderConfig, Metadata: g.ProviderMetadata}) + if err != nil { + return ErrSandboxDeploymentConflict + } + needsCredential, err := providers.UsesCredential(g.ProviderKind) if err != nil { - return SandboxSetup{}, ErrSandboxDeploymentConflict + return err } - retained = composed.Configuration + if needsCredential { + composed, err := providers.WithCredential(sandbox.Selection{Provider: g.ProviderKind, Configuration: retained}, sandbox.Selection{Provider: d.ProviderKind, Configuration: result.Configuration}) + if err != nil { + return ErrSandboxDeploymentConflict + } + retained = composed.Configuration + } + result.Configuration = retained } - result.Configuration = retained - + return nil + }) + if err != nil { + return SandboxSetup{}, err } - return result, tx.Commit(ctx) + return result, nil } // SandboxGenerationPage is bounded; callers retain a deadline over the full scan. diff --git a/services/core/internal/store/sandbox_reset.go b/services/core/internal/store/sandbox_reset.go index 318f22bb7..010567c38 100644 --- a/services/core/internal/store/sandbox_reset.go +++ b/services/core/internal/store/sandbox_reset.go @@ -55,12 +55,10 @@ func checkSandboxGeneration(d sqlc.RuntimeDeployment, installation string, gener } func (s *Store) resetTransaction(ctx context.Context, apply func(context.Context, *sqlc.Queries, sqlc.RuntimeDeployment) error) error { - if s.executionLease == nil { - return ErrInvalidInput + if err := s.checkExecutionAuthority(); err != nil { + return err } - ctx, cancel := context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() - return s.executionLease.transaction(ctx, func(tx pgx.Tx) error { + return s.writer.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) d, err := q.LockRuntimeDeployment(ctx) if err != nil { diff --git a/services/core/internal/store/sandbox_reset_test.go b/services/core/internal/store/sandbox_reset_test.go index 607c0b333..8ba70531d 100644 --- a/services/core/internal/store/sandbox_reset_test.go +++ b/services/core/internal/store/sandbox_reset_test.go @@ -352,10 +352,10 @@ func TestSandboxResetOwnerRestartRetainsDeadlineAndProvenance(t *testing.T) { // Model an abrupt owner loss and wait for PostgreSQL to terminate that backend, // rather than assuming local TCP cleanup acknowledges advisory-lock release. var stopped bool - if err := s.pool.QueryRow(t.Context(), `SELECT pg_terminate_backend($1,1000)`, w.executionLease.conn.Conn().PgConn().PID()).Scan(&stopped); err != nil || !stopped { + if err := s.pool.QueryRow(t.Context(), `SELECT pg_terminate_backend($1,1000)`, executionOwnerPID(t, s.pool)).Scan(&stopped); err != nil || !stopped { t.Fatal(stopped, err) } - successor := executionLease(t, s).Store() + successor := executionWriter(t, s) current, err := s.GetRuntimeDeployment(t.Context()) if err != nil || !current.Reset.RequestedAt.Equal(reset.Reset.RequestedAt) || !current.Reset.DeadlineAt.Equal(*reset.Reset.DeadlineAt) { t.Fatal("restart moved reset deadline", current, err) diff --git a/services/core/internal/store/sandbox_specification_store_test.go b/services/core/internal/store/sandbox_specification_store_test.go index 135412bcc..95e9a6f3c 100644 --- a/services/core/internal/store/sandbox_specification_store_test.go +++ b/services/core/internal/store/sandbox_specification_store_test.go @@ -23,7 +23,7 @@ func webSpecificationFixture(t *testing.T, provider string) (*Store, *Store, Run t.Fatal(err) } s := NewWithCredentialCipher(pool, cipher) - w := executionLease(t, s).Store() + w := executionWriter(t, s) id := uuid.NewString() if err := w.ClaimWebSandboxDeployment(t.Context(), id); err != nil { t.Fatal(err) diff --git a/services/core/internal/store/scheduling.go b/services/core/internal/store/scheduling.go index dc9300fa6..e1f632beb 100644 --- a/services/core/internal/store/scheduling.go +++ b/services/core/internal/store/scheduling.go @@ -87,7 +87,7 @@ func (s *Store) sessionActivity(ctx context.Context, session Session, err error) if err != nil { return Session{}, err } - err = pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}, func(tx pgx.Tx) error { + err = s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { var err error session, err = readSessionActivity(ctx, s.queries.WithTx(tx), session) return err @@ -108,7 +108,7 @@ func (s *Store) SessionStreamSnapshot(ctx context.Context, tenantID, sessionID s } var session Session var cursor int64 - err = pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}, func(tx pgx.Tx) error { + err = s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.GetSession(ctx, sqlc.GetSessionParams{TenantID: tenant, ID: id}) if err != nil { diff --git a/services/core/internal/store/session_artifacts.go b/services/core/internal/store/session_artifacts.go index 3be7a5b53..a7b28c86f 100644 --- a/services/core/internal/store/session_artifacts.go +++ b/services/core/internal/store/session_artifacts.go @@ -96,30 +96,24 @@ func (s *Store) ReadSessionArtifact(ctx context.Context, tenantID, sessionID, ar if consume == nil { return ErrInvalidInput } - tx, err := s.pool.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}) - if err != nil { - return err - } - defer tx.Rollback(context.Background()) - row, err := s.queries.WithTx(tx).GetSessionArtifact(ctx, lookup) - if errors.Is(err, pgx.ErrNoRows) { - return ErrNotFound - } - if err != nil { - return err - } - objects := tx.LargeObjects() - body, err := objects.Open(ctx, row.BodyOid.Uint32, pgx.LargeObjectModeRead) - if err != nil { - return err - } - if err := consume(artifactFromRow(row), body); err != nil { - return err - } - if err := body.Close(); err != nil { - return err - } - return tx.Commit(ctx) + return s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { + row, err := s.queries.WithTx(tx).GetSessionArtifact(ctx, lookup) + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotFound + } + if err != nil { + return err + } + objects := tx.LargeObjects() + body, err := objects.Open(ctx, row.BodyOid.Uint32, pgx.LargeObjectModeRead) + if err != nil { + return err + } + if err := consume(artifactFromRow(row), body); err != nil { + return err + } + return body.Close() + }) } func (s *Store) DeleteSessionArtifact(ctx context.Context, tenantID, sessionID, artifactID string) error { @@ -127,35 +121,29 @@ func (s *Store) DeleteSessionArtifact(ctx context.Context, tenantID, sessionID, if err != nil { return err } - tx, err := s.pool.Begin(ctx) - if err != nil { - return err - } - defer tx.Rollback(context.Background()) - q := s.queries.WithTx(tx) - // Use the same lock order as whole-Session deletion and Turn publication. - locked, err := q.LockSession(ctx, sqlc.LockSessionParams{TenantID: lookup.TenantID, ID: lookup.SessionID}) - if errors.Is(err, pgx.ErrNoRows) || err == nil && locked.DeletedAt.Valid { - return ErrNotFound - } - if err != nil { - return err - } - oid, err := q.DeleteSessionArtifact(ctx, sqlc.DeleteSessionArtifactParams(lookup)) - if errors.Is(err, pgx.ErrNoRows) { - return ErrNotFound - } - if err != nil { - return err - } - objects := tx.LargeObjects() - if err := objects.Unlink(ctx, oid.Uint32); err != nil { - return err - } - if err := recordWriteAudit(ctx, q, tenantID, "delete", "artifact", uuid.UUID(lookup.ID.Bytes).String(), uuid.UUID(lookup.SessionID.Bytes).String()); err != nil { - return err - } - return tx.Commit(ctx) + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + q := s.queries.WithTx(tx) + // Use the same lock order as whole-Session deletion and Turn publication. + locked, err := q.LockSession(ctx, sqlc.LockSessionParams{TenantID: lookup.TenantID, ID: lookup.SessionID}) + if errors.Is(err, pgx.ErrNoRows) || err == nil && locked.DeletedAt.Valid { + return ErrNotFound + } + if err != nil { + return err + } + oid, err := q.DeleteSessionArtifact(ctx, sqlc.DeleteSessionArtifactParams(lookup)) + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotFound + } + if err != nil { + return err + } + objects := tx.LargeObjects() + if err := objects.Unlink(ctx, oid.Uint32); err != nil { + return err + } + return recordWriteAudit(ctx, q, tenantID, "delete", "artifact", uuid.UUID(lookup.ID.Bytes).String(), uuid.UUID(lookup.SessionID.Bytes).String()) + }) } func artifactLookup(tenantID, sessionID, artifactID string) (sqlc.GetSessionArtifactParams, error) { diff --git a/services/core/internal/store/session_creation_identity.go b/services/core/internal/store/session_creation_identity.go index 34bd6ccc6..25e6f5229 100644 --- a/services/core/internal/store/session_creation_identity.go +++ b/services/core/internal/store/session_creation_identity.go @@ -124,7 +124,7 @@ func (s *Store) FindSessionCreation(ctx context.Context, tenantID, key string, r } var row sqlc.Session var environment *Environment - err = pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}, func(tx pgx.Tx) error { + err = s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) var err error row, err = q.FindSessionCreation(ctx, sqlc.FindSessionCreationParams{TenantID: tenant, IdempotencyKey: key}) diff --git a/services/core/internal/store/session_deletion_lifecycle_public_test.go b/services/core/internal/store/session_deletion_lifecycle_public_test.go index ee49e9214..2b9931134 100644 --- a/services/core/internal/store/session_deletion_lifecycle_public_test.go +++ b/services/core/internal/store/session_deletion_lifecycle_public_test.go @@ -40,16 +40,15 @@ func TestSessionDeletionLifecyclePostgres(t *testing.T) { server := httptest.NewServer(h) defer server.Close() client := pathIDClient{t: t, server: server} - lease, err := s.AcquireExecutionLease(ctx) + writer, err := store.NewExecution(ctx, s) if err != nil { t.Fatal(err) } t.Cleanup(func() { closing, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() - _ = lease.Close(closing) + _ = writer.CloseExecution(closing) }) - writer := lease.Store() create := func(environment string, initial bool) store.Session { t.Helper() diff --git a/services/core/internal/store/session_diagnostics.go b/services/core/internal/store/session_diagnostics.go index 78738ae4e..b956fc9f5 100644 --- a/services/core/internal/store/session_diagnostics.go +++ b/services/core/internal/store/session_diagnostics.go @@ -33,7 +33,7 @@ func (s *Store) GetSessionDiagnosticsSnapshot(ctx context.Context, tenantID, ses return Session{}, err } var session Session - err = pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}, func(tx pgx.Tx) error { + err = s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.GetSession(ctx, sqlc.GetSessionParams{TenantID: tenant, ID: parsePathID(sessionID)}) if err != nil { @@ -60,7 +60,7 @@ func (s *Store) GetTurnDiagnosticsSnapshot(ctx context.Context, tenantID, sessio return TurnDiagnosticsSnapshot{}, err } result := TurnDiagnosticsSnapshot{Items: []ItemDiagnosticTiming{}} - err = pgx.BeginTxFunc(ctx, s.pool, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}, func(tx pgx.Tx) error { + err = s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) turn, err := q.GetTurn(ctx, params) if err != nil { diff --git a/services/core/internal/store/session_diagnostics_test.go b/services/core/internal/store/session_diagnostics_test.go index 4c80d667f..5ca9cfc5e 100644 --- a/services/core/internal/store/session_diagnostics_test.go +++ b/services/core/internal/store/session_diagnostics_test.go @@ -203,7 +203,7 @@ func TestDiagnosticProvisioningDetailAtomicAndPrivate(t *testing.T) { if err != nil { t.Fatal(err) } - writer := executionLease(t, s).Store() + writer := executionWriter(t, s) owner, err := writer.ReserveRuntimeAllocation(t.Context(), tenant, session.Environment.ID, uuid.NewString(), runtimedevice.HashCredential(uuid.NewString())) if err != nil { t.Fatal(err) diff --git a/services/core/internal/store/session_execution_configuration_test.go b/services/core/internal/store/session_execution_configuration_test.go index 7310abe46..abd8eb0cd 100644 --- a/services/core/internal/store/session_execution_configuration_test.go +++ b/services/core/internal/store/session_execution_configuration_test.go @@ -244,7 +244,7 @@ func TestSessionExecutionConfigurationConcurrentRetryKeepsWinner(t *testing.T) { func TestSessionExecutionConfigurationSurvivesSuspendResume(t *testing.T) { s, pool := testStore(t) - w := executionLease(t, s).Store() + w := executionWriter(t, s) tenant := uuid.NewString() session, err := s.CreateSession(t.Context(), tenant, executionProjectionInput("agent")) if err != nil { diff --git a/services/core/internal/store/session_initial_input.go b/services/core/internal/store/session_initial_input.go index ac98f7f56..95e1a9d9f 100644 --- a/services/core/internal/store/session_initial_input.go +++ b/services/core/internal/store/session_initial_input.go @@ -26,7 +26,7 @@ func validateInitialInputs(inputs []Input) ([]Input, json.RawMessage, error) { func (s *Store) createSessionResources(ctx context.Context, tenant string, params sqlc.CreateSessionParams, inputs []Input, encodedInput json.RawMessage, files []InitialFile, setup EnvironmentSetup, provider *v1.ModelProviderInput, executionConfiguration *v1.SessionExecutionConfiguration, providerSource string, deploymentRevision uuid.UUID) (sqlc.Session, *Environment, error) { var row sqlc.Session var environment *Environment - err := pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err := s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) var err error row, err = q.CreateSession(ctx, params) diff --git a/services/core/internal/store/session_metadata.go b/services/core/internal/store/session_metadata.go index 85cad57d8..12ac7d0f4 100644 --- a/services/core/internal/store/session_metadata.go +++ b/services/core/internal/store/session_metadata.go @@ -22,7 +22,7 @@ func (s *Store) UpdateSessionMetadata(ctx context.Context, tenantID, sessionID s return Session{}, err } var row sqlc.Session - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) var err error row, err = q.UpdateSessionMetadata(ctx, sqlc.UpdateSessionMetadataParams{TenantID: tenant, ID: id, Metadata: encoded}) diff --git a/services/core/internal/store/session_transaction.go b/services/core/internal/store/session_transaction.go index 0b0e8e407..8145d04fd 100644 --- a/services/core/internal/store/session_transaction.go +++ b/services/core/internal/store/session_transaction.go @@ -43,16 +43,7 @@ func (s *Store) withLockedSession(ctx context.Context, tenantID, sessionID strin return err } } - begin := func(ctx context.Context, apply func(pgx.Tx) error) error { - return pgx.BeginFunc(ctx, s.pool, apply) - } - if s.executionLease != nil { - var cancel context.CancelFunc - ctx, cancel = context.WithTimeout(ctx, executionTransactionTimeout) - defer cancel() - begin = s.executionLease.transaction - } - return begin(ctx, func(tx pgx.Tx) error { + return s.writer.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) session, err := q.LockSession(ctx, sqlc.LockSessionParams{TenantID: tenant, ID: id}) if errors.Is(err, pgx.ErrNoRows) { diff --git a/services/core/internal/store/sessions.go b/services/core/internal/store/sessions.go index 88eba9fa7..19f885327 100644 --- a/services/core/internal/store/sessions.go +++ b/services/core/internal/store/sessions.go @@ -22,6 +22,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/identity" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" ) var ( @@ -79,9 +80,16 @@ type SessionPage struct { } type Store struct { - queries *sqlc.Queries - pool *pgxpool.Pool - executionLease *ExecutionLease + queries *sqlc.Queries + // pool supplies the execution lease's dedicated connection; pooled runs + // every other transaction. + pool *pgxpool.Pool + pooled *pgunit.Pool + // writer runs Session and execution-only transactions, and lease grants + // execution authority. New sets writer to pooled and leaves lease nil; + // NewExecution sets both to the same execution lease. Neither changes later. + writer transactor + lease *pgunit.Lease credentialCipher *credentialcrypto.Cipher oauthRefresher oauthrefresh.Refresher // publicURL is OAC_PUBLIC_URL. Core derives every address it gives @@ -89,7 +97,10 @@ type Store struct { publicURL string } -func New(pool *pgxpool.Pool) *Store { return &Store{queries: sqlc.New(pool), pool: pool} } +func New(pool *pgxpool.Pool) *Store { + pooled := pgunit.NewPool(pool) + return &Store{queries: sqlc.New(pool), pool: pool, pooled: pooled, writer: pooled} +} func ValidEngine(engine string) bool { return enginePattern.MatchString(engine) } diff --git a/services/core/internal/store/sessions_test.go b/services/core/internal/store/sessions_test.go index 3bca4bf0e..6dd3ca83b 100644 --- a/services/core/internal/store/sessions_test.go +++ b/services/core/internal/store/sessions_test.go @@ -3,75 +3,22 @@ package store import ( "context" "errors" - "os" "reflect" - "strings" "sync" "testing" "github.com/google/uuid" "github.com/jackc/pgx/v5/pgxpool" - "github.com/MiniMax-AI/OpenAgentCore/services/core/migrations" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" ) func testStore(t *testing.T) (*Store, *pgxpool.Pool) { t.Helper() - dsn := os.Getenv("OAC_TEST_DATABASE_URL") - if dsn == "" { - t.Skip("OAC_TEST_DATABASE_URL is not set; dedicated PostgreSQL required") - } - cfg, err := testDatabaseConfig(dsn) - if err != nil { - t.Fatal(err) - } - ctx := context.Background() - pool, err := pgxpool.NewWithConfig(ctx, cfg) - if err != nil { - t.Fatal(err) - } - t.Cleanup(pool.Close) - // A separate database, not product fixtures or migrations, is sufficient. - var database string - var productTable *string - if err := pool.QueryRow(ctx, "SELECT current_database(), to_regclass('workspaces')::text").Scan(&database, &productTable); err != nil || productTable != nil || database != cfg.ConnConfig.Database { - t.Fatal("execution tests require a database without product workspace tables") - } - if err := migrations.Apply(ctx, dsn); err != nil { - t.Fatal(err) - } + pool := pgtest.Open(t) return New(pool), pool } -// Validate the driver's effective database, including query parameters and DSNs. -func testDatabaseConfig(dsn string) (*pgxpool.Config, error) { - cfg, err := pgxpool.ParseConfig(dsn) - if err != nil { - return nil, errors.New("invalid test database configuration") - } - database := cfg.ConnConfig.Database - if !strings.HasPrefix(database, "oac_") || !strings.HasSuffix(database, "_tests") { - return nil, errors.New("test database must be named oac_*_tests") - } - return cfg, nil -} - -func TestDatabaseGuardUsesEffectiveDatabase(t *testing.T) { - for _, dsn := range []string{ - "postgres://localhost/oac_local_tests?dbname=agents_api", - "host=localhost dbname=agents_api", - "postgres://localhost/agents_api", - } { - if _, err := testDatabaseConfig(dsn); err == nil { - t.Fatalf("unsafe database accepted: %s", dsn) - } - } - cfg, err := testDatabaseConfig("postgres://localhost/oac_local_tests") - if err != nil || cfg.ConnConfig.Database != "oac_local_tests" { - t.Fatalf("valid dedicated database rejected: %v", err) - } -} - func TestSessionsPersistAndStayTenantScoped(t *testing.T) { s, pool := testStore(t) ctx := context.Background() diff --git a/services/core/internal/store/skill_versions.go b/services/core/internal/store/skill_versions.go index 12152b98b..bc7f314ee 100644 --- a/services/core/internal/store/skill_versions.go +++ b/services/core/internal/store/skill_versions.go @@ -23,7 +23,7 @@ func (s *Store) CreateSkillVersion(ctx context.Context, tenantID, skillID string return SkillVersion{}, ErrInvalidInput } var result SkillVersion - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) owner, err := q.LockSkill(ctx, sqlc.LockSkillParams{TenantID: tenant, ID: id}) if err != nil { @@ -98,7 +98,7 @@ func (s *Store) DeleteSkillVersion(ctx context.Context, tenantID, skillID, versi } number := skillPathVersion(version) var result SkillVersion - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) owner, err := q.LockSkill(ctx, sqlc.LockSkillParams{TenantID: tenant, ID: id}) if err != nil { diff --git a/services/core/internal/store/skills.go b/services/core/internal/store/skills.go index 930ac7af8..d10d46a27 100644 --- a/services/core/internal/store/skills.go +++ b/services/core/internal/store/skills.go @@ -43,7 +43,7 @@ func (s *Store) CreateSkill(ctx context.Context, tenantID string, archive []byte return Skill{}, ErrInvalidInput } var result Skill - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) id := pgtype.UUID{Bytes: uuid.New(), Valid: true} row, err := q.CreateSkill(ctx, sqlc.CreateSkillParams{ID: id, TenantID: tenant, Name: metadata.Name, Description: metadata.Description}) @@ -84,7 +84,7 @@ func (s *Store) UpdateSkillDefault(ctx context.Context, tenantID, skillID, versi return Skill{}, err } var result Skill - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) if _, err := q.LockSkill(ctx, sqlc.LockSkillParams{TenantID: tenant, ID: id}); err != nil { return err @@ -111,7 +111,7 @@ func (s *Store) DeleteSkill(ctx context.Context, tenantID, skillID string) error if err != nil { return err } - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) if _, err := q.DeleteSkill(ctx, sqlc.DeleteSkillParams{TenantID: tenant, ID: id}); err != nil { return err diff --git a/services/core/internal/store/source_files.go b/services/core/internal/store/source_files.go index 036ce6ff9..f77a4ec1b 100644 --- a/services/core/internal/store/source_files.go +++ b/services/core/internal/store/source_files.go @@ -45,50 +45,50 @@ func (s *Store) CreateSourceFile(ctx context.Context, tenantID string, upload fu if err != nil || upload == nil { return SourceFile{}, ErrInvalidInput } - tx, err := s.pool.Begin(ctx) - if err != nil { - return SourceFile{}, err - } - defer tx.Rollback(context.Background()) - objects := tx.LargeObjects() - oid, err := objects.Create(ctx, 0) - if err != nil { - return SourceFile{}, err - } - body, err := objects.Open(ctx, oid, pgx.LargeObjectModeWrite) - if err != nil { - return SourceFile{}, err - } - writer := newSourceFileWriter(body) - input, err := upload(writer) - if writer.err != nil { - return SourceFile{}, writer.err - } - if err != nil { - return SourceFile{}, err - } - if !validSourceFilename(input.Filename) || input.Purpose != "user_data" { - return SourceFile{}, ErrInvalidInput - } - if err := body.Close(); err != nil { - return SourceFile{}, err - } - row, err := s.queries.WithTx(tx).CreateSourceFile(ctx, sqlc.CreateSourceFileParams{ - ID: pgtype.UUID{Bytes: uuid.New(), Valid: true}, TenantID: tenant, - Filename: input.Filename, Purpose: input.Purpose, BodyOid: pgtype.Uint32{Uint32: oid, Valid: true}, - SizeBytes: writer.size, Sha256: hex.EncodeToString(writer.hash.Sum(nil)), + var created SourceFile + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + objects := tx.LargeObjects() + oid, err := objects.Create(ctx, 0) + if err != nil { + return err + } + body, err := objects.Open(ctx, oid, pgx.LargeObjectModeWrite) + if err != nil { + return err + } + writer := newSourceFileWriter(body) + input, err := upload(writer) + if writer.err != nil { + return writer.err + } + if err != nil { + return err + } + if !validSourceFilename(input.Filename) || input.Purpose != "user_data" { + return ErrInvalidInput + } + if err := body.Close(); err != nil { + return err + } + row, err := s.queries.WithTx(tx).CreateSourceFile(ctx, sqlc.CreateSourceFileParams{ + ID: pgtype.UUID{Bytes: uuid.New(), Valid: true}, TenantID: tenant, + Filename: input.Filename, Purpose: input.Purpose, BodyOid: pgtype.Uint32{Uint32: oid, Valid: true}, + SizeBytes: writer.size, Sha256: hex.EncodeToString(writer.hash.Sum(nil)), + }) + if err != nil { + return fmt.Errorf("create source file: %w", err) + } + resource := sourceFileFromRow(row) + if err := recordWriteAudit(ctx, s.queries.WithTx(tx), tenantID, "create", "file", resource.ID, "", AuditResource{Type: "file", ID: resource.ID}); err != nil { + return err + } + created = resource + return nil }) if err != nil { - return SourceFile{}, fmt.Errorf("create source file: %w", err) - } - resource := sourceFileFromRow(row) - if err := recordWriteAudit(ctx, s.queries.WithTx(tx), tenantID, "create", "file", resource.ID, "", AuditResource{Type: "file", ID: resource.ID}); err != nil { - return SourceFile{}, err - } - if err := tx.Commit(ctx); err != nil { return SourceFile{}, err } - return sourceFileFromRow(row), nil + return created, nil } func (s *Store) GetSourceFile(ctx context.Context, tenantID, fileID string) (SourceFile, error) { @@ -156,22 +156,16 @@ func (s *Store) ReadSourceFile(ctx context.Context, tenantID, fileID string, con if consume == nil { return ErrInvalidInput } - tx, err := s.pool.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}) - if err != nil { - return err - } - defer tx.Rollback(context.Background()) - row, err := s.queries.WithTx(tx).GetSourceFile(ctx, sqlc.GetSourceFileParams{TenantID: tenant, ID: id}) - if errors.Is(err, pgx.ErrNoRows) { - return ErrNotFound - } - if err != nil { - return err - } - if err := consumeSourceFile(ctx, tx, row, consume); err != nil { - return err - } - return tx.Commit(ctx) + return s.pooled.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { + row, err := s.queries.WithTx(tx).GetSourceFile(ctx, sqlc.GetSourceFileParams{TenantID: tenant, ID: id}) + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotFound + } + if err != nil { + return err + } + return consumeSourceFile(ctx, tx, row, consume) + }) } func consumeSourceFile(ctx context.Context, tx pgx.Tx, row sqlc.SourceFile, consume func(SourceFile, io.Reader) error) error { @@ -191,26 +185,20 @@ func (s *Store) DeleteSourceFile(ctx context.Context, tenantID, fileID string) e if err != nil { return err } - tx, err := s.pool.Begin(ctx) - if err != nil { - return err - } - defer tx.Rollback(context.Background()) - oid, err := s.queries.WithTx(tx).DeleteSourceFile(ctx, sqlc.DeleteSourceFileParams{TenantID: tenant, ID: id}) - if errors.Is(err, pgx.ErrNoRows) { - return ErrNotFound - } - if err != nil { - return err - } - objects := tx.LargeObjects() - if err := objects.Unlink(ctx, oid.Uint32); err != nil { - return err - } - if err := recordWriteAudit(ctx, s.queries.WithTx(tx), tenantID, "delete", "file", fileID, ""); err != nil { - return err - } - return tx.Commit(ctx) + return s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + oid, err := s.queries.WithTx(tx).DeleteSourceFile(ctx, sqlc.DeleteSourceFileParams{TenantID: tenant, ID: id}) + if errors.Is(err, pgx.ErrNoRows) { + return ErrNotFound + } + if err != nil { + return err + } + objects := tx.LargeObjects() + if err := objects.Unlink(ctx, oid.Uint32); err != nil { + return err + } + return recordWriteAudit(ctx, s.queries.WithTx(tx), tenantID, "delete", "file", fileID, "") + }) } func sourceFileIDs(tenantID, fileID string) (pgtype.UUID, pgtype.UUID, error) { diff --git a/services/core/internal/store/subagent_dispatch_test.go b/services/core/internal/store/subagent_dispatch_test.go index ebc8df539..0b955f7e8 100644 --- a/services/core/internal/store/subagent_dispatch_test.go +++ b/services/core/internal/store/subagent_dispatch_test.go @@ -27,12 +27,12 @@ func TestSubagentIdentityUsesLeasedDispatchJournal(t *testing.T) { if err = h.s.BindSessionDevice(ctx, h.tenant, h.session.ID, h.device.ID); err != nil { t.Fatal(err) } - lease, err := h.s.AcquireExecutionLease(ctx) + writer, err := store.NewExecution(ctx, h.s) if err != nil { t.Fatal(err) } - t.Cleanup(func() { _ = lease.Close(context.Background()) }) - h.d.Store = lease.Store() + t.Cleanup(func() { _ = writer.CloseExecution(context.Background()) }) + h.d.Store = writer input := h.message("first", "root message") running := h.run(ctx, input.TurnID) var request proto.PromptRequestPayload diff --git a/services/core/internal/store/subagent_identities_test.go b/services/core/internal/store/subagent_identities_test.go index 245db7fdc..06a046f32 100644 --- a/services/core/internal/store/subagent_identities_test.go +++ b/services/core/internal/store/subagent_identities_test.go @@ -20,8 +20,7 @@ func subagentIdentityEvent(child, parent string, created int64) ExecutionEvent { func TestSubagentIdentityIsAtomicScopedAndImmutable(t *testing.T) { s, pool := testStore(t) - lease := executionLease(t, s) - w := lease.Store() + w := executionWriter(t, s) ctx := t.Context() tenant, session := newSubagentSession(t, s) host, err := s.CreateDevice(ctx, tenant, "identity test", runtimedevice.HashCredential(uuid.NewString())) @@ -35,14 +34,17 @@ func TestSubagentIdentityIsAtomicScopedAndImmutable(t *testing.T) { transition(t, w, tenant, session.ID, input.TurnID, TurnQueued, TurnInProgress) a, b := subagentIdentityEvent("child-a", "root", 102), subagentIdentityEvent("child-b", "root", 101) batch := []ExecutionEvent{a, b, a} - if err = s.AppendTurnEvents(ctx, tenant, session.ID, input.TurnID, 1, batch); err == nil { - t.Fatal("unleased discovery accepted") + if err = s.AppendTurnEvents(ctx, tenant, session.ID, input.TurnID, 1, batch); !errors.Is(err, ErrExecutionAuthority) { + t.Fatal("unleased discovery accepted", err) } for range 2 { if err = w.AppendTurnEvents(ctx, tenant, session.ID, input.TurnID, 1, batch); err != nil { t.Fatal(err) } } + if err = s.AppendTurnEvents(ctx, tenant, session.ID, input.TurnID, 1, batch); !errors.Is(err, ErrExecutionAuthority) { + t.Fatal("unleased replay of a committed discovery accepted", err) + } saved, err := s.GetSubagentIdentity(ctx, tenant, session.ID, "child-a") if err != nil || saved.ID == "" || saved.ID == saved.NativeID || saved.SessionID != session.ID || saved.FirstTurnID != input.TurnID || saved.FirstEventOrdinal != 1 || saved.NativeCreatedAt != 102 || saved.FirstObservedAt.IsZero() { t.Fatal(saved, err) @@ -107,23 +109,23 @@ func TestSubagentIdentityIsAtomicScopedAndImmutable(t *testing.T) { if _, err = w.CompleteExecution(ctx, tenant, session.ID, input.TurnID, TurnCompleted, json.RawMessage(`{}`), "root", input.Sequence); err != nil { t.Fatal(err) } - if err = lease.Close(ctx); err != nil { + if err = w.CloseExecution(ctx); err != nil { t.Fatal(err) } reopened, _ := testStore(t) - nextOwner := executionLease(t, reopened) + nextOwner := executionWriter(t, reopened) again, err := reopened.GetSubagentIdentity(ctx, tenant, session.ID, "child-a") if err != nil || !reflect.DeepEqual(again, saved) { t.Fatal("restart changed identity", again, err) } second := submitMessage(t, reopened, tenant, session.ID, "second") - transition(t, nextOwner.Store(), tenant, session.ID, second.TurnID, TurnQueued, TurnInProgress) + transition(t, nextOwner, tenant, session.ID, second.TurnID, TurnQueued, TurnInProgress) if err = w.AppendTurnEvents(ctx, tenant, session.ID, second.TurnID, 1, []ExecutionEvent{a}); err == nil { t.Fatal("closed owner wrote identity") } continued := proto.SubagentIdentityPayload{NativeID: "child-a", ParentNativeID: "root", NativeCreatedAt: 102, ParentTurnID: "later-native-turn", SourceItemID: "resume-item"} raw, _ := json.Marshal(continued) - if err = nextOwner.Store().AppendTurnEvents(ctx, tenant, session.ID, second.TurnID, 1, []ExecutionEvent{{Kind: proto.TypeSubagentIdentity, Payload: raw}}); err != nil { + if err = nextOwner.AppendTurnEvents(ctx, tenant, session.ID, second.TurnID, 1, []ExecutionEvent{{Kind: proto.TypeSubagentIdentity, Payload: raw}}); err != nil { t.Fatal(err) } again, err = reopened.GetSubagentIdentity(ctx, tenant, session.ID, "child-a") @@ -143,19 +145,19 @@ func TestSubagentIdentityIsAtomicScopedAndImmutable(t *testing.T) { func TestSubagentIdentityRejectsLostLease(t *testing.T) { s, pool := testStore(t) - old := executionLease(t, s) + old := executionWriter(t, s) tenant, session := newSubagentSession(t, s) input := submitMessage(t, s, tenant, session.ID, "first") - transition(t, old.Store(), tenant, session.ID, input.TurnID, TurnQueued, TurnInProgress) + transition(t, old, tenant, session.ID, input.TurnID, TurnQueued, TurnInProgress) var killed bool - if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1, 1000)", old.conn.Conn().PgConn().PID()).Scan(&killed); err != nil || !killed { + if err := pool.QueryRow(t.Context(), "SELECT pg_terminate_backend($1, 1000)", executionOwnerPID(t, pool)).Scan(&killed); err != nil || !killed { t.Fatal(killed, err) } - successor := executionLease(t, s) - if err := old.Store().AppendTurnEvents(t.Context(), tenant, session.ID, input.TurnID, 1, []ExecutionEvent{subagentIdentityEvent("child", "root", 100)}); err == nil { + successor := executionWriter(t, s) + if err := old.AppendTurnEvents(t.Context(), tenant, session.ID, input.TurnID, 1, []ExecutionEvent{subagentIdentityEvent("child", "root", 100)}); err == nil { t.Fatal("lost owner committed identity") } - if err := successor.Ping(context.Background()); err != nil { + if err := successor.CheckExecutionOwnership(context.Background()); err != nil { t.Fatal(err) } events, err := s.ListTurnEvents(t.Context(), tenant, session.ID, input.TurnID, 0, 100) diff --git a/services/core/internal/store/subagent_native_outputs_test.go b/services/core/internal/store/subagent_native_outputs_test.go index dbfc3b7a6..f2c0cb6c0 100644 --- a/services/core/internal/store/subagent_native_outputs_test.go +++ b/services/core/internal/store/subagent_native_outputs_test.go @@ -13,7 +13,7 @@ import ( func TestSubagentNativeFunctionResultDoesNotConsumeOutputIndex(t *testing.T) { s, pool := testStore(t) - owner := executionLease(t, s).Store() + owner := executionWriter(t, s) tenant, session := newSubagentSession(t, s) host, err := s.CreateDevice(t.Context(), tenant, "child outputs", runtimedevice.HashCredential(uuid.NewString())) if err != nil { @@ -87,7 +87,7 @@ func TestSubagentNativeFunctionResultDoesNotConsumeOutputIndex(t *testing.T) { func TestSubagentCancelledPartialMessageSurvivesHistoryReplay(t *testing.T) { s, _ := testStore(t) - owner := executionLease(t, s).Store() + owner := executionWriter(t, s) tenant, session := newSubagentSession(t, s) host, err := s.CreateDevice(t.Context(), tenant, "cancelled child", runtimedevice.HashCredential(uuid.NewString())) if err != nil { diff --git a/services/core/internal/store/subagent_resources_test.go b/services/core/internal/store/subagent_resources_test.go index a651c577b..e1e5aaf62 100644 --- a/services/core/internal/store/subagent_resources_test.go +++ b/services/core/internal/store/subagent_resources_test.go @@ -28,7 +28,7 @@ func subagentFact(kind string, value any) ExecutionEvent { } func TestSubagentResourcesNativeOwnershipLifecycleAndRecovery(t *testing.T) { s, pool := testStore(t) - owner := executionLease(t, s).Store() + owner := executionWriter(t, s) ctx := t.Context() tenant, session := newSubagentSession(t, s) host, err := s.CreateDevice(ctx, tenant, "child resources", runtimedevice.HashCredential(uuid.NewString())) diff --git a/services/core/internal/store/subagent_visibility_public_test.go b/services/core/internal/store/subagent_visibility_public_test.go index 85a5cbcc8..e1e70d522 100644 --- a/services/core/internal/store/subagent_visibility_public_test.go +++ b/services/core/internal/store/subagent_visibility_public_test.go @@ -86,12 +86,11 @@ func TestSubagentVisibilityPublic(t *testing.T) { server := httptest.NewServer(handler) defer server.Close() client := pathIDClient{t: t, server: server} - lease, err := s.AcquireExecutionLease(t.Context()) + writer, err := store.NewExecution(t.Context(), s) if err != nil { t.Fatal(err) } - defer func() { _ = lease.Close(t.Context()) }() - writer := lease.Store() + defer func() { _ = writer.CloseExecution(t.Context()) }() ctx := t.Context() created := openStream(t, server, token, http.MethodPost, "/v1/agents/sessions", diff --git a/services/core/internal/store/turn_events.go b/services/core/internal/store/turn_events.go index 53bdca2e2..d419a26ed 100644 --- a/services/core/internal/store/turn_events.go +++ b/services/core/internal/store/turn_events.go @@ -4,6 +4,7 @@ import ( "context" "encoding/json" "errors" + "slices" "time" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" @@ -34,12 +35,16 @@ func (s *Store) AppendTurnEvents(ctx context.Context, tenantID, sessionID, turnI if first < 1 || len(events) == 0 || len(events) > 64 { return ErrInvalidInput } + // Subagent observations are projected under the execution journal, so only + // the execution writer records a batch that contains one, replays included. + if slices.ContainsFunc(events, func(event ExecutionEvent) bool { return isSubagentObservation(event.Kind) }) { + if err := s.checkExecutionAuthority(); err != nil { + return err + } + } normalized := make([]ExecutionEvent, len(events)) payloadBytes := 0 for i, event := range events { - if isSubagentObservation(event.Kind) && s.executionLease == nil { - return errors.New("subagent discovery requires a leased Store") - } if len(event.Payload) > 512*1024 || !enginePattern.MatchString(event.Kind) { return ErrInvalidInput } diff --git a/services/core/internal/store/vault_credentials.go b/services/core/internal/store/vault_credentials.go index fc94bda14..63e583d43 100644 --- a/services/core/internal/store/vault_credentials.go +++ b/services/core/internal/store/vault_credentials.go @@ -59,7 +59,7 @@ func (s *Store) CreateStaticCredential(ctx context.Context, tenantID, vaultID st return Credential{}, errors.New("credential encryption failed") } var created Credential - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.CreateStaticCredential(ctx, sqlc.CreateStaticCredentialParams{ ID: pgtype.UUID{Bytes: id, Valid: true}, TenantID: tenant, VaultID: vault, diff --git a/services/core/internal/store/vault_credentials_delete.go b/services/core/internal/store/vault_credentials_delete.go index 4bcccb479..432330906 100644 --- a/services/core/internal/store/vault_credentials_delete.go +++ b/services/core/internal/store/vault_credentials_delete.go @@ -24,7 +24,7 @@ func (s *Store) DeleteCredential(ctx context.Context, tenantID, vaultID, credent return "", ErrNotFound } var deletedID string - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) deleted, err := q.DeleteCredential(ctx, sqlc.DeleteCredentialParams{TenantID: tenant, VaultID: vault, ID: id}) if err != nil { diff --git a/services/core/internal/store/vault_credentials_oauth.go b/services/core/internal/store/vault_credentials_oauth.go index c56b0f6d9..d1a720607 100644 --- a/services/core/internal/store/vault_credentials_oauth.go +++ b/services/core/internal/store/vault_credentials_oauth.go @@ -35,7 +35,7 @@ func (s *Store) CreateOAuthCredential(ctx context.Context, tenantID, vaultID str return Credential{}, err } var created Credential - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.CreateOAuthCredential(ctx, sqlc.CreateOAuthCredentialParams{ ID: pgtype.UUID{Bytes: uuid.MustParse(credential.ID), Valid: true}, TenantID: tenant, VaultID: vault, @@ -67,92 +67,103 @@ func (s *Store) UpdateOAuthCredential(ctx context.Context, tenantID, vaultID, cr if current.AuthType != "mcp_oauth" { return Credential{}, ErrInvalidInput } - tx, credential, secret, err := s.lockOAuth(ctx, tenantID, vaultID, credentialID, "") - if err != nil { - return Credential{}, err - } - defer tx.Rollback(context.Background()) - if input.AccessToken != nil { - secret.AccessToken = *input.AccessToken - secret.Metadata.ExpiresAt = nil - } - if input.ExpiresAtSet { - secret.Metadata.ExpiresAt = input.ExpiresAt - } - if update := input.Refresh; update != nil { - refresh := secret.Metadata.Refresh - if refresh == nil { - return Credential{}, ErrInvalidInput + var updated Credential + err = s.withOAuth(ctx, tenantID, vaultID, credentialID, "", "credential update failed", func(ctx context.Context, tx pgx.Tx, credential Credential, secret oauthSecret) error { + if input.AccessToken != nil { + secret.AccessToken = *input.AccessToken + secret.Metadata.ExpiresAt = nil } - if update.TokenEndpointAuthType != "" && update.TokenEndpointAuthType != refresh.TokenEndpointAuth { - return Credential{}, ErrInvalidInput + if input.ExpiresAtSet { + secret.Metadata.ExpiresAt = input.ExpiresAt } - if update.ClientSecret != nil { - if refresh.TokenEndpointAuth == "none" { - return Credential{}, ErrInvalidInput + if update := input.Refresh; update != nil { + refresh := secret.Metadata.Refresh + if refresh == nil { + return ErrInvalidInput + } + if update.TokenEndpointAuthType != "" && update.TokenEndpointAuthType != refresh.TokenEndpointAuth { + return ErrInvalidInput + } + if update.ClientSecret != nil { + if refresh.TokenEndpointAuth == "none" { + return ErrInvalidInput + } + secret.ClientSecret = *update.ClientSecret + } + if update.RefreshToken != nil { + secret.RefreshToken = *update.RefreshToken + } + if update.ScopeSet { + refresh.Scope = update.Scope } - secret.ClientSecret = *update.ClientSecret } - if update.RefreshToken != nil { - secret.RefreshToken = *update.RefreshToken + var err error + updated, err = s.saveOAuth(ctx, tx, tenantID, credential, secret) + if err != nil { + return err } - if update.ScopeSet { - refresh.Scope = update.Scope + if err := recordWriteAudit(ctx, s.queries.WithTx(tx), tenantID, "update", "credential", updated.ID, updated.VaultID); err != nil { + return errors.New("credential update failed") } - } - updated, err := s.saveOAuth(ctx, tx, tenantID, credential, secret) + return nil + }) if err != nil { return Credential{}, err } - if err := recordWriteAudit(ctx, s.queries.WithTx(tx), tenantID, "update", "credential", updated.ID, updated.VaultID); err != nil { - return Credential{}, errors.New("credential update failed") - } - if err := tx.Commit(ctx); err != nil { - return Credential{}, errors.New("credential update failed") - } return updated, nil } -// Row ownership lasts through refresh or manual replacement. PostgreSQL serializes -// competing updates and deletes, including a parent Vault's cascading deletion. -func (s *Store) lockOAuth(ctx context.Context, tenantID, vaultID, credentialID, destination string) (pgx.Tx, Credential, oauthSecret, error) { +// withOAuth locks the credential row for the whole of apply, including an +// external refresh, and commits only when apply succeeds. Row ownership lasts +// through refresh or manual replacement. PostgreSQL serializes competing updates +// and deletes, including a parent Vault's cascading deletion. A failure to begin +// or commit is reported as failure, never with database error text. +func (s *Store) withOAuth(ctx context.Context, tenantID, vaultID, credentialID, destination, failure string, apply func(context.Context, pgx.Tx, Credential, oauthSecret) error) error { tenant, e1 := parseID(tenantID) vault, e2 := parseID(vaultID) id, e3 := parseID(credentialID) if e1 != nil || e2 != nil || e3 != nil { - return nil, Credential{}, oauthSecret{}, ErrNotFound + return ErrNotFound } - tx, err := s.pool.Begin(ctx) - if err != nil { - return nil, Credential{}, oauthSecret{}, errors.New("credential transaction failed") - } - fail := func(err error) (pgx.Tx, Credential, oauthSecret, error) { - _ = tx.Rollback(context.Background()) - return nil, Credential{}, oauthSecret{}, err + var applied error + err := s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + credential, secret, err := s.lockOAuth(ctx, tx, tenant, vault, id, destination) + if err == nil { + err = apply(ctx, tx, credential, secret) + } + applied = err + return err + }) + if err != nil && applied == nil { + return errors.New(failure) } + return err +} + +func (s *Store) lockOAuth(ctx context.Context, tx pgx.Tx, tenant, vault, id pgtype.UUID, destination string) (Credential, oauthSecret, error) { row, err := sqlc.New(tx).GetOAuthCredentialForUpdate(ctx, sqlc.GetOAuthCredentialForUpdateParams{ TenantID: tenant, VaultID: vault, ID: id, }) if errors.Is(err, pgx.ErrNoRows) { - return fail(ErrNotFound) + return Credential{}, oauthSecret{}, ErrNotFound } if err != nil { - return fail(errors.New("credential lookup failed")) + return Credential{}, oauthSecret{}, errors.New("credential lookup failed") } credential, err := credentialFromRow(sqlc.GetCredentialRow{ID: row.ID, VaultID: row.VaultID, Name: row.Name, AuthType: row.AuthType, McpServerUrl: row.McpServerUrl, CreatedAt: row.CreatedAt, UpdatedAt: row.UpdatedAt, OauthMetadata: row.OauthMetadata}) if err != nil { - return fail(err) + return Credential{}, oauthSecret{}, err } if destination != "" && credential.MCPServerURL != destination { - return fail(ErrNotFound) + return Credential{}, oauthSecret{}, ErrNotFound } secret, err := s.openOAuth(uuid.UUID(tenant.Bytes).String(), credential, row.TokenCiphertext) if err != nil { - return fail(err) + return Credential{}, oauthSecret{}, err } - return tx, credential, secret, nil + return credential, secret, nil } func (s *Store) saveOAuth(ctx context.Context, tx pgx.Tx, tenantID string, credential Credential, secret oauthSecret) (Credential, error) { diff --git a/services/core/internal/store/vault_credentials_oauth_test.go b/services/core/internal/store/vault_credentials_oauth_test.go index 529336a67..3a9454732 100644 --- a/services/core/internal/store/vault_credentials_oauth_test.go +++ b/services/core/internal/store/vault_credentials_oauth_test.go @@ -14,6 +14,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/credentialcrypto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/oauthrefresh" "github.com/google/uuid" + "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) @@ -61,13 +62,14 @@ func oauthFixtureBinding(credential Credential) MCPCredentialBinding { func storedOAuthSecret(t *testing.T, s *Store, tenant string, credential Credential) oauthSecret { t.Helper() - tx, _, secret, err := s.lockOAuth(t.Context(), tenant, credential.VaultID, credential.ID, "") + var secret oauthSecret + err := s.withOAuth(t.Context(), tenant, credential.VaultID, credential.ID, "", "credential read failed", func(_ context.Context, _ pgx.Tx, _ Credential, stored oauthSecret) error { + secret = stored + return nil + }) if err != nil { t.Fatal("read private test grant", err) } - if err := tx.Rollback(t.Context()); err != nil { - t.Fatal(err) - } return secret } diff --git a/services/core/internal/store/vault_credentials_update.go b/services/core/internal/store/vault_credentials_update.go index 86d4b087e..c86bcd265 100644 --- a/services/core/internal/store/vault_credentials_update.go +++ b/services/core/internal/store/vault_credentials_update.go @@ -48,7 +48,7 @@ func (s *Store) UpdateStaticCredential(ctx context.Context, tenantID, vaultID, c return Credential{}, errors.New("credential encryption failed") } var updated Credential - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.UpdateStaticCredential(ctx, sqlc.UpdateStaticCredentialParams{ TenantID: tenant, VaultID: vault, ID: id, McpServerUrl: current.MCPServerURL, TokenCiphertext: ciphertext, diff --git a/services/core/internal/store/vaults.go b/services/core/internal/store/vaults.go index 0e5e2937b..ead0e9ccb 100644 --- a/services/core/internal/store/vaults.go +++ b/services/core/internal/store/vaults.go @@ -48,7 +48,7 @@ func (s *Store) CreateVault(ctx context.Context, tenantID string, input CreateVa return Vault{}, err } var created Vault - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) row, err := q.CreateVault(ctx, sqlc.CreateVaultParams{ ID: pgtype.UUID{Bytes: uuid.New(), Valid: true}, TenantID: tenant, diff --git a/services/core/internal/store/vaults_delete.go b/services/core/internal/store/vaults_delete.go index a00d41ba0..451bd8b13 100644 --- a/services/core/internal/store/vaults_delete.go +++ b/services/core/internal/store/vaults_delete.go @@ -20,7 +20,7 @@ func (s *Store) DeleteVault(ctx context.Context, tenantID, vaultID string) (stri return "", ErrNotFound } var deletedID string - err = pgx.BeginFunc(ctx, s.pool, func(tx pgx.Tx) error { + err = s.pooled.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { q := s.queries.WithTx(tx) deleted, err := q.DeleteVault(ctx, sqlc.DeleteVaultParams{TenantID: tenant, ID: id}) if err != nil {