diff --git a/services/core/IMPLEMENTATION.md b/services/core/IMPLEMENTATION.md index 3c5e0278d..9418835da 100644 --- a/services/core/IMPLEMENTATION.md +++ b/services/core/IMPLEMENTATION.md @@ -16,11 +16,11 @@ Two errors are shared across domains, each with one `api` helper: `textvalue.Err Shared vocabulary has one owner each, and domains use it rather than copy it. `internal/environmentconfig` owns Environment setup, Skills, Plugins and initial files with their validation and public metadata; `Setup.Validate` checks requested configuration, where a Skill may be an unresolved reference, and `Setup.ValidateInstalled` checks frozen, installable configuration. `internal/skills` owns `ParseVersion`, the canonical positive decimal Skill version. `internal/metadata` owns the metadata rules: `Validate` for the pair, key and value limits and U+0000, `ValidateStorable` for U+0000 alone, and `Encode` with its 64 KiB bound. `internal/jsonobject` owns `Normalize`, the stable encoding of stored JSON objects that snapshots and retry identities compare. These packages import no persistence. -`internal/sessions` owns the Session vocabulary: Sessions, Turns, inputs, Environments and their provisioning failures, function calls, Item and Artifact reads, executor credentials, and the errors Session operations return, which `api` maps in `writeSessionsError`. It also owns the Session change vocabulary and its decisions: the public changes that report Turn and Session transitions, what a Turn that ends settles, measured Turn usage and the Session activity a change reports. Sequences that several Session operations share are `sessions` procedures, each over the transaction interface it declares: `CancelTurn` and `CancelWork` cancel work, `TrackInputActivity` reports the input activity a write changes, `FailEnvironment` and `TerminateEnvironment` settle an Environment that failed or whose managed compute ended, `CreateEnvironmentDevice` creates and binds a hosted Environment's dedicated Runtime device, `CheckComputeAdmission`, `CheckFileWriteGate` and `CheckInputStart` gate new work, `AdmitFunctionResult` saves a submitted function result for its Turn's call, `LoadRequiredActions` lists the calls a Turn waits on, and `TransitionTurn` moves a root Turn's status as its current status allows and settles a Turn that ends. `AppendTurnEvents` records a batch of a Turn's execution observations, which `NewJournalBatch` validates, and `AppendTurnEvent` records its terminal outcome; both project the new entries. `ProjectSource` projects one journal entry or admitted input into Items, the root Turn's usage and Subagents, and `ProjectInput` projects an admitted input. `ValidTransition` decides the Turn status transitions, and the Subagent projection decides each Subagent's identity, lifecycle, child Turn and child Item changes. `sessions.ExecutionOperations` runs the Session writes only the execution owner makes; each function-call and Turn operation loads the Turn in one leased Session transaction, decides in `sessions`, then applies in `sessionpg`, and `AppendTurnEvents` records and projects a Turn's journal batch in one leased Session transaction. `internal/items` owns public Items: it projects observations, merges them into stored Items and builds the ordered events that report each Item change. Neither imports persistence. `internal/persistence/postgres/sessionpg` loads the facts those decisions read and applies them inside the caller's Session transaction, under the Session lock: it allocates event sequence positions, event IDs, Item positions, output indexes and Subagent IDs, writes the change journal, each Turn's execution journal and its counters, Items, Turn usage, Subagent bindings, child Turns and child Items, and Artifact settlement, and prunes the change journal. It decides nothing. `WithSession` runs a Session operation's transaction on the pool or the lease: it locks the Session, applies the operation, then prunes the change journal. `BindSession` binds one Session to a transaction's queries as a `SessionTx`, which implements the procedures' interfaces. +`internal/sessions` owns the Session vocabulary: Sessions, Turns, inputs, Environments and their provisioning failures, function calls, Item and Artifact reads, executor credentials, and the errors Session operations return, which `api` maps in `writeSessionsError`. It also owns the Session change vocabulary and its decisions: the public changes that report Turn and Session transitions, what a Turn that ends settles, measured Turn usage and the Session activity a change reports. Sequences that several Session operations share are `sessions` procedures, each over the transaction interface it declares: `CancelTurn` and `CancelWork` cancel work, `TrackInputActivity` reports the input activity a write changes, `FailEnvironment` and `TerminateEnvironment` settle an Environment that failed or whose managed compute ended, `CreateEnvironmentDevice` creates and binds a hosted Environment's dedicated Runtime device, `CheckComputeAdmission`, `CheckFileWriteGate` and `CheckInputStart` gate new work, `AdmitInputs` admits a validated input batch in order, each input steering the active Turn, starting a Turn, requesting a cancellation or joining a function result to its call's Turn, `AdmitFunctionResult` saves a submitted function result for its Turn's call, `LoadRequiredActions` lists the calls a Turn waits on, and `TransitionTurn` moves a root Turn's status as its current status allows and settles a Turn that ends. `AppendTurnEvents` records a batch of a Turn's execution observations, which `NewJournalBatch` validates, and `AppendTurnEvent` records its terminal outcome; both project the new entries. `ProjectSource` projects one journal entry or admitted input into Items, the root Turn's usage and Subagents, and `ProjectInput` projects an admitted input. `ValidTransition` decides the Turn status transitions, and the Subagent projection decides each Subagent's identity, lifecycle, child Turn and child Item changes. `sessions.ExecutionOperations` runs the Session writes only the execution owner makes; each function-call and Turn operation loads the Turn in one leased Session transaction, decides in `sessions`, then applies in `sessionpg`, and `AppendTurnEvents` records and projects a Turn's journal batch in one leased Session transaction. `internal/items` owns public Items: it projects observations, merges them into stored Items and builds the ordered events that report each Item change. Neither imports persistence. `internal/persistence/postgres/sessionpg` loads the facts those decisions read and applies them inside the caller's Session transaction, under the Session lock: it allocates event sequence positions, event IDs, Item positions, output indexes and Subagent IDs, writes the change journal, each Turn's execution journal and its counters, Items, Turn usage, Subagent bindings, child Turns and child Items, and Artifact settlement, and prunes the change journal. It decides nothing. `WithSession` runs a Session operation's transaction on the pool or the lease: it locks the Session, applies the operation, then prunes the change journal. `BindSession` binds one Session to a transaction's queries as a `SessionTx`, which implements the procedures' interfaces. Other adapters call only these `sessionpg` participants, on their own transaction's queries: `LockSession`, which orders a write against the Session's Turns; `BindSession`, which runs a Session procedure, such as `CreateEnvironmentDevice` or `TransitionTurn`, in that transaction; `LoadEnvironment`, which reads the tenant's Environment of a live Session; and `CreateSessionEnvironment`, which stores the Environment of a Session being created. The other `sessionpg` exports serve `store` until the Session operations leave it; among them `LoadSessionActivity`, the one Session activity read, adds the activity projection to the Session that `store`'s Session creation returns. `tests/fixtures` is test tooling that composes procedures over a pooled `WithSession`, which is why `WithSession` and `BindSession` stay exported after store goes; no production code may do this. `deploymentpg` may import `sessionpg`, and `sessionpg` never imports `deploymentpg`; `placementpg`, which loads the placement facts and reserves placements, sits below both, so either may call it and it imports neither. Likewise `deployment` may compose the `sessions` procedures, and `sessions` never imports `deployment`. -`store` is transitional. `store.New` builds a pooled Store, and `store.NewExecution(s, lease)` builds the execution writer on a lease it borrows. 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`. +`store` is transitional. `store.New` builds a pooled Store, and `store.NewExecution(s, lease)` builds the execution writer on a lease it borrows. Adapters' execution repositories require a `*pgunit.Lease` at construction, and their public repositories expose no execution operation. `cmd/server` owns the execution lease. It acquires one `pgunit.Lease`, builds every lease-bound adapter on it, and passes the lease and those adapters together as one `execution.Owner` to `execution.StartWorker`, which binds the Dispatcher to the Owner's execution operations through `Dispatcher.Bind`. If anything fails before that call, `cmd/server` closes the lease. From that call the Worker owns cleanup: a failed start closes the lease before it returns, and a started Worker closes it after `Run` has cancelled and drained its work. Each close runs under its own bounded deadline, independent of the cancelled request or run. Lease-bound adapters and the store writer borrow the lease and never close it, and the Worker uses the lease only through `Owner.Lease`, never through an adapter. Store integration tests start the Worker the same way through `startWorker`, and bind a Dispatcher that runs Turns without a Worker through `Dispatcher.Bind`. @@ -36,7 +36,7 @@ Domain owners, each with its PostgreSQL adapter under `internal/persistence/post - `deployment` and its subpackage `deployment/placement` (`deploymentpg`, `placementpg`): the sandbox deployment and its nodes: provider configuration and the sealed credential, the specification and retained generations, setup, update and switch, node enrollment, identity and authentication, generation configuration, capacity, presence, status, host history, reset, and the counts of nodes and sandboxes bound to the public address; and hosted runtime allocations: reservation with the dedicated device, compute ownership and settlement, observation diagnostics, activity, compute phases, wake receipts and cleanup, with the reads that schedule, discover and authorize them. It interprets Sandbox Provider declarations through the `providers.Registry` it is given, whose lookups return typed errors. `cmd/server` builds that registry and calls it only to build a direct Provider and to discover a Provider's configuration; it takes each setup's mode and declared operations from the `Setup` that `deployment` returns. `deployment/placement` owns hosted admission and placement: `cmd/server` builds one `placement.Rules` from the registry and the public URL, both fixed while Core runs, and gives it to `deployment.Service` and to `store`; its pure decisions admit a hosted Session, choose its node, and admit an allocation's reserved node and a restore on it, and its errors keep one status, code and message through `writeStoreError` and `writeDeploymentError`. `placementpg` loads the facts those decisions read and applies them on the caller's transaction-bound queries; it has no Store or transaction runner. The only other reader is `store`: Session creation decides admission and placement with those rules and `placementpg` inside its own transaction. Deployment changes and allocation writes run through `deployment.ExecutionOperations` on `deploymentpg.NewExecution(lease, …)`, which the Worker receives as `execution.Owner.Deployment`. Each allocation write locks the owning Session, decides in `deployment`, applies in `deploymentpg` and prunes the Session's change journal in the same leased transaction; cleanup settles the Session through the `sessions` procedures on a `sessionpg.SessionTx` bound to that transaction. The Worker reads the deployment, prepares a selection's setup and records live-compute activity through the pooled `deployment.Service` in `execution.Dispatcher.Deployment`, and lists the Sessions a reset still has to archive, schedules node lifecycles and reads allocations through the pooled `deployment.Reader` in `execution.Dispatcher.DeploymentReader`; a pooled read carries no lease, and the leased write that follows it rechecks the allocation's owner. Node management and reads use the pooled `deploymentpg.Store`, which `cmd/server` also reads the owner epoch from and runs the host-history sampler on. Provider calls run outside transactions, and the final transaction rechecks the expected generation. Session archive crosses the Session and the deployment, so it is a `deployment.ExecutionOperations` operation: `ArchiveSession`, which the administrator's archive route calls through `api.Execution.SessionArchive`, and `ArchiveResetSession`, which a reset calls, lock the Session through `deploymentpg`'s `WithSessionArchive`, check the deployment's generation and the running reset in `deployment`, then expire the Environment and cancel its work through the `sessions` procedures, revoke the allocation's device and request its cleanup, and record the audit entry in one leased transaction. `deployment.ObservationResolver` resolves a Session's Runtime observation target for `runtimeobs` from `sessions.SessionReader` and `deployment.Reader`. - `coremetrics` (`coremetricspg`): the Core metrics PostgreSQL holds: the root Turn queue counts, the root Turn history, read from one read-only snapshot, and the database size. `cmd/server`'s Core metrics source adds them to the process, pool, Worker and daemon registry measurements. - `runtimehistory` (`runtimehistory/postgresreader`, outside this directory): Runtime history samples: the periodic export, scoped reads and retention, which also prunes node-host samples. `runtimehistory.Service` scopes a read to the Session's hosted Environment through `sessions.EnvironmentReader`. -- `sessions` (`sessionpg`): Session use cases and reads as they leave `store`. The pooled `sessions.Service` runs the use cases on `sessionpg.Store`, which implements `sessions.Storage`, and plain reads use `sessions.Reader`, which `sessionpg.Store` also implements, directly. So far these cover Sessions (their reads, the change journal and stream snapshot, diagnostics, the frozen execution configuration, measured usage and archive state, metadata updates, deletion and the public write audit), Turns (their reads and the execution work scan), the model provider a Session froze, which `sessionpg.Store` opens with the credential key, root Items and Subagents (their reads), Session Artifacts (their reads, deletion and the staging of a Turn's export), Environments (their reads, the initialization list and the frozen setup and initial files, which `sessionpg.Store` opens with the credential key), devices (their reads, which include a Session's execution binding and the execution device list, their creation, revocation, heartbeats and Runtime enrollment, and the archived cancellation receipt the daemon gateway reads), the administrator's views across Projects (a Project's asset counts and Sessions for the summary, and the Sessions whose Runtime the administrator observes), executor credentials (authentication and a Project's credential state through `sessions.ExecutorCredentialReader`, and issuance, rotation and revocation; the Core-key Project operations record their audit entry in the same transaction) and native installation authorization and claims, whose tokens `sessionpg.Store` signs with the credential key. `cmd/server` builds one `sessionpg.Store` with the credential key and wires the Service and the Store into `execution.Dispatcher.Sessions` and `SessionsReader`, into the api fields, into Runtime enrollment and into the daemon gateway; the Worker stages Artifacts through the Service; `cmd/environment-key` builds one without the key for its credential commands. Turn transitions, execution completion, which alone publishes or discards a Turn's staged Artifacts, the start of Artifact capture, function calls and their application receipts, the Turn execution journal, Environment initialization, connection observations and their reconciliation, Session device binding, and file-write reservation and settlement run through `sessions.ExecutionOperations` on `sessionpg.NewExecution(lease)`, which the Worker receives as `execution.Owner.Sessions`. +- `sessions` (`sessionpg`): Session use cases and reads as they leave `store`. The pooled `sessions.Service` runs the use cases on `sessionpg.Store`, which implements `sessions.Storage`, and plain reads use `sessions.Reader`, which `sessionpg.Store` also implements, directly. So far these cover Sessions (their reads, the change journal and stream snapshot, diagnostics, the frozen execution configuration, measured usage and archive state, metadata updates, deletion and the public write audit), Turns (their reads and the execution work scan), input admission (a public input batch, Environment input reservation and the expiry of one reservation, with the reads of a Turn's admitted inputs, a reservation and the Environment input work scan), the model provider a Session froze, which `sessionpg.Store` opens with the credential key, root Items and Subagents (their reads), Session Artifacts (their reads, deletion and the staging of a Turn's export), Environments (their reads, the initialization list and the frozen setup and initial files, which `sessionpg.Store` opens with the credential key), devices (their reads, which include a Session's execution binding and the execution device list, their creation, revocation, heartbeats and Runtime enrollment, and the archived cancellation receipt the daemon gateway reads), the administrator's views across Projects (a Project's asset counts and Sessions for the summary, and the Sessions whose Runtime the administrator observes), executor credentials (authentication and a Project's credential state through `sessions.ExecutorCredentialReader`, and issuance, rotation and revocation; the Core-key Project operations record their audit entry in the same transaction) and native installation authorization and claims, whose tokens `sessionpg.Store` signs with the credential key. `cmd/server` builds one `sessionpg.Store` with the credential key and wires the Service and the Store into `execution.Dispatcher.Sessions` and `SessionsReader`, into the api fields, into Runtime enrollment and into the daemon gateway; the Worker stages Artifacts through the Service; `cmd/environment-key` builds one without the key for its credential commands. Turn transitions, execution completion, which alone publishes or discards a Turn's staged Artifacts, the start of Artifact capture, function calls and their application receipts, the Turn execution journal, Environment initialization, connection observations and their reconciliation, Session device binding, file-write reservation and settlement, and the promotion, failure and bulk expiry of Environment input reservations run through `sessions.ExecutionOperations` on `sessionpg.NewExecution(lease)`, which the Worker receives as `execution.Owner.Sessions`. ## Request handling diff --git a/services/core/internal/api/native_classification_integration_test.go b/services/core/internal/api/native_classification_integration_test.go index 9ce5ea206..12dc3754f 100644 --- a/services/core/internal/api/native_classification_integration_test.go +++ b/services/core/internal/api/native_classification_integration_test.go @@ -25,10 +25,7 @@ func TestNativeClassificationPostgresRoundTripAndPublicPrivacy(t *testing.T) { if err != nil { t.Fatal(err) } - receipt, err := s.SubmitMessage(t.Context(), tenant, session.ID, "input", json.RawMessage(`{"text":"test"}`)) - if err != nil { - t.Fatal(err) - } + receipt := submitMessage(t, pool, tenant, session.ID, "input", json.RawMessage(`{"text":"test"}`)) transitionTurn(t, pool, tenant, session.ID, receipt.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress}) status := 503 result := execution.Result{ErrorCode: "engine_failed", Error: "Bearer secret-canary https://private.example/key", EngineErrorCode: code, EngineHTTPStatus: &status, Done: proto.DonePayload{Usage: proto.Usage{InputTokens: 7, OutputTokens: 3}, Metadata: map[string]any{proto.DoneMetaAgentSessionID: "native-secret-canary"}}} diff --git a/services/core/internal/api/session_diagnostics_public_compat_test.go b/services/core/internal/api/session_diagnostics_public_compat_test.go index c555210b3..f6b56f881 100644 --- a/services/core/internal/api/session_diagnostics_public_compat_test.go +++ b/services/core/internal/api/session_diagnostics_public_compat_test.go @@ -73,6 +73,20 @@ func databaseSessionReads(pool *pgxpool.Pool) func(*Dependencies, *testFakes) { } } +// submitMessage admits one message input through the Session service on pool. +func submitMessage(t *testing.T, pool *pgxpool.Pool, tenant, session, key string, payload json.RawMessage) sessions.InputReceipt { + t.Helper() + service, err := sessions.NewService(sessionpg.New(pgunit.NewPool(pool), nil)) + if err != nil { + t.Fatal(err) + } + receipts, err := service.SubmitInputs(t.Context(), tenant, session, key, []sessions.Input{{Kind: "message", Payload: payload}}) + if err != nil { + t.Fatal(err) + } + return receipts[0] +} + // transitionTurn moves the Turn as the execution owner does, over a pooled // Session transaction. func transitionTurn(t *testing.T, pool *pgxpool.Pool, tenant, session, turn string, transition sessions.TurnTransition) { diff --git a/services/core/internal/api/session_diagnostics_test.go b/services/core/internal/api/session_diagnostics_test.go index 448a959df..0470488c1 100644 --- a/services/core/internal/api/session_diagnostics_test.go +++ b/services/core/internal/api/session_diagnostics_test.go @@ -19,10 +19,7 @@ func TestDiagnosticsCoreHandlerDatabaseBoundary(t *testing.T) { if err != nil { t.Fatal(err) } - receipt, err := s.SubmitMessage(t.Context(), tenant, session.ID, "input", json.RawMessage(`{"text":"input-secret-canary"}`)) - if err != nil { - t.Fatal(err) - } + receipt := submitMessage(t, pool, tenant, session.ID, "input", json.RawMessage(`{"text":"input-secret-canary"}`)) transitionTurn(t, pool, tenant, session.ID, receipt.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress}) transitionTurn(t, pool, tenant, session.ID, receipt.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnInProgress, Status: sessions.TurnFailed, Outcome: json.RawMessage(`{"error_code":"device_disconnected","error":"Bearer raw-secret-canary https://private.example/key","done":{"native_id":"secret-native-canary"}}`)}) base := adminSessionsPath + session.ID diff --git a/services/core/internal/execution/archive_cancellation_cleanup_test.go b/services/core/internal/execution/archive_cancellation_cleanup_test.go index 822676a52..34ac69f46 100644 --- a/services/core/internal/execution/archive_cancellation_cleanup_test.go +++ b/services/core/internal/execution/archive_cancellation_cleanup_test.go @@ -92,10 +92,12 @@ func TestArchiveWaitingCleanupReceiptBarrier(t *testing.T) { if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(t.Context(), project.TenantID, session.ID, "start", json.RawMessage(`{"text":"run"}`)) + _, service := testSessions(t, pool, nil) + inputs, err := service.SubmitInputs(t.Context(), project.TenantID, session.ID, "start", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"run"}`)}}) if err != nil { t.Fatal(err) } + input := inputs[0] // This fixture isolates lifecycle ordering. Protocol-driven waiting is // independently exercised in TestArchiveWaitingCancellationReceipts. for _, transition := range []sessions.TurnTransition{{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress}, {ExpectedStatus: sessions.TurnInProgress, Status: sessions.TurnWaiting}} { diff --git a/services/core/internal/execution/delivery.go b/services/core/internal/execution/delivery.go index ab20d7d7c..b18f4f86e 100644 --- a/services/core/internal/execution/delivery.go +++ b/services/core/internal/execution/delivery.go @@ -303,7 +303,7 @@ func (d *Dispatcher) deliver(ctx context.Context, tenantID, sessionID string, pe continue } if pending == nil { - inputs, err := d.Store.ListTurnInputs(ctx, tenantID, sessionID, request.RunID, result.AppliedThrough, 1) + inputs, err := d.SessionsReader.ListTurnInputs(ctx, tenantID, sessionID, request.RunID, result.AppliedThrough, 1) if err != nil { result.ErrorCode = "execution_state_unavailable" if ctx.Err() != nil { diff --git a/services/core/internal/execution/deployment_provider_observations_test.go b/services/core/internal/execution/deployment_provider_observations_test.go index fa11d22b2..d8a45d3e4 100644 --- a/services/core/internal/execution/deployment_provider_observations_test.go +++ b/services/core/internal/execution/deployment_provider_observations_test.go @@ -82,10 +82,12 @@ func newFinishObservationFixture(t *testing.T, maxConnections int32) finishObser } func (f finishObservationFixture) start(t *testing.T) sessions.InputReceipt { t.Helper() - receipt, err := f.s.SubmitMessage(t.Context(), f.tenant, f.session.ID, uuid.NewString(), json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"fixture"}]}]}`)) + _, service := testSessions(t, f.pool, nil) + receipts, err := service.SubmitInputs(t.Context(), f.tenant, f.session.ID, uuid.NewString(), []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"fixture"}]}]}`)}}) if err != nil { t.Fatal(err) } + receipt := receipts[0] if _, err = f.execution.TransitionTurn(t.Context(), f.tenant, f.session.ID, receipt.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress}); err != nil { t.Fatal(err) } diff --git a/services/core/internal/execution/environment_admission.go b/services/core/internal/execution/environment_admission.go index 835027d90..d078186e7 100644 --- a/services/core/internal/execution/environment_admission.go +++ b/services/core/internal/execution/environment_admission.go @@ -97,7 +97,7 @@ func (w *Worker) submitEnvironmentInputs(ctx context.Context, session sessions.S changed, unsubscribe := w.dispatcher.notifications.subscribe(session.TenantID, session.ID) defer unsubscribe() reserve, cancel := context.WithTimeout(ctx, 5*time.Second) - reservation, err := w.admission.ReserveEnvironmentInput(reserve, session.TenantID, session.ID, key, inputs) + reservation, err := w.dispatcher.Sessions.ReserveEnvironmentInput(reserve, session.TenantID, session.ID, key, inputs) cancel() if err != nil { return nil, err @@ -140,9 +140,9 @@ func (w *Worker) environmentInputOutcome(ctx context.Context, session sessions.S defer cancel() // The database rechecks its clock under the Session lock before settlement. if !time.Now().Before(reservation.Deadline) { - return w.admission.ExpireEnvironmentInput(read, session.TenantID, session.ID, reservation.ID) + return w.dispatcher.Sessions.ExpireEnvironmentInput(read, session.TenantID, session.ID, reservation.ID) } - return w.admission.GetEnvironmentInputReservation(read, session.TenantID, session.ID, reservation.ID) + return w.dispatcher.SessionsReader.GetEnvironmentInputReservation(read, session.TenantID, session.ID, reservation.ID) } func (w *Worker) checkAdmissionOwnership(ctx context.Context) error { diff --git a/services/core/internal/execution/message_input.go b/services/core/internal/execution/message_input.go index 5a44766d8..328567f6f 100644 --- a/services/core/internal/execution/message_input.go +++ b/services/core/internal/execution/message_input.go @@ -38,7 +38,7 @@ func messageInput(raw json.RawMessage) (proto.MessageInput, error) { } func (d *Dispatcher) initialInput(ctx context.Context, tenant, session, turn string) (proto.MessageInput, int64, error) { - inputs, err := d.Store.ListTurnInputs(ctx, tenant, session, turn, 0, 100) + inputs, err := d.SessionsReader.ListTurnInputs(ctx, tenant, session, turn, 0, 100) if err != nil { return nil, 0, err } diff --git a/services/core/internal/execution/preparation.go b/services/core/internal/execution/preparation.go index 59ea9bfb4..ca815d02e 100644 --- a/services/core/internal/execution/preparation.go +++ b/services/core/internal/execution/preparation.go @@ -97,7 +97,9 @@ func (d *Dispatcher) awaitPreparation(ctx context.Context, tenant, session strin case <-ctx.Done(): return pending, ctx.Err() case <-tick.C: - current, err := d.Store.ExpireEnvironmentInput(ctx, tenant, session, pending.ID) + expire, cancel := context.WithTimeout(ctx, 5*time.Second) + current, err := d.Sessions.ExpireEnvironmentInput(expire, tenant, session, pending.ID) + cancel() if err != nil || current.State != sessions.EnvironmentInputPending { return current, err } diff --git a/services/core/internal/execution/prepared_dispatch.go b/services/core/internal/execution/prepared_dispatch.go index 1584b7442..6cb5817f6 100644 --- a/services/core/internal/execution/prepared_dispatch.go +++ b/services/core/internal/execution/prepared_dispatch.go @@ -5,6 +5,7 @@ import ( "encoding/json" "errors" "strings" + "time" v1 "github.com/MiniMax-AI/OpenAgentCore/contracts/agents-api/v1" "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" @@ -17,12 +18,15 @@ type EnvironmentRun struct { } // RunEnvironmentInput reserves a Turn on the Session-owned Runtime Executor. It -// checks lease, the lease d.Store was built on, before any Runtime preparation. +// checks lease, the lease the Dispatcher's execution operations hold, before +// any Runtime preparation. func (d *Dispatcher) RunEnvironmentInput(ctx context.Context, lease Ownership, tenantID, sessionID, reservationID string) (run EnvironmentRun, err error) { if err = lease.CheckOwnership(ctx); err != nil { return run, err } - run.Reservation, err = d.Store.ExpireEnvironmentInput(ctx, tenantID, sessionID, reservationID) + expire, cancel := context.WithTimeout(ctx, 5*time.Second) + run.Reservation, err = d.Sessions.ExpireEnvironmentInput(expire, tenantID, sessionID, reservationID) + cancel() if err != nil || run.Reservation.State != sessions.EnvironmentInputPending { return run, err } @@ -92,7 +96,7 @@ func (d *Dispatcher) RunEnvironmentInput(ctx context.Context, lease Ownership, t if err := d.messageInputSupport(peer, session.Engine, snapshot, messages); err != nil { return run, err } - promoted, err := d.Store.PromoteEnvironmentInput(owner, tenantID, sessionID, reservationID) + promoted, err := d.sessionExecution.PromoteEnvironmentInput(owner, tenantID, sessionID, reservationID) if errors.Is(err, sessions.ErrTurnConflict) { // A rejected claim leaves the reservation pending for a later attempt. return run, err diff --git a/services/core/internal/execution/worker.go b/services/core/internal/execution/worker.go index 0b2b48a0b..6ad0bcc06 100644 --- a/services/core/internal/execution/worker.go +++ b/services/core/internal/execution/worker.go @@ -303,7 +303,7 @@ func (w *Worker) Run(ctx context.Context) (runErr error) { return err } if maintenance { - if _, err := w.dispatcher.Store.ExpireEnvironmentInputs(ctx); err != nil { + if _, err := w.dispatcher.sessionExecution.ExpireEnvironmentInputs(ctx); err != nil { w.observeSchedulerPoll(0, err) return err } diff --git a/services/core/internal/execution/worker_schedule.go b/services/core/internal/execution/worker_schedule.go index fbc0098d8..de5ac027a 100644 --- a/services/core/internal/execution/worker_schedule.go +++ b/services/core/internal/execution/worker_schedule.go @@ -34,7 +34,7 @@ func (s *workerSchedule) selectWork(ctx context.Context, w *Worker, devices []st } var environments []sessions.EnvironmentInputWork if !time.Now().Before(s.nextEnvironmentScan) { - environments, err = w.dispatcher.Store.ListEnvironmentInputWork(ctx, s.environmentCursor, devices) + environments, err = w.dispatcher.SessionsReader.ListEnvironmentInputWork(ctx, s.environmentCursor, devices) if err != nil { return nil, err } @@ -43,7 +43,7 @@ func (s *workerSchedule) selectWork(ctx context.Context, w *Worker, devices []st s.environmentCursor = "" // Retry the first page now instead of spending a scan interval on EOF. // A single refill preserves the candidate bound and cannot spin when empty. - environments, err = w.dispatcher.Store.ListEnvironmentInputWork(ctx, "", devices) + environments, err = w.dispatcher.SessionsReader.ListEnvironmentInputWork(ctx, "", devices) if err != nil { return nil, err } @@ -94,10 +94,10 @@ func (w *Worker) runEnvironmentInput(ctx context.Context, item scheduledWork) er return nil } if errors.Is(err, ErrModelProviderRequired) { - return w.dispatcher.Store.FailEnvironmentInput(ctx, item.TenantID, item.SessionID, item.reservationID, "model_provider_required") + return w.dispatcher.sessionExecution.FailEnvironmentInput(ctx, item.TenantID, item.SessionID, item.reservationID, "model_provider_required") } if errors.Is(err, errPreparationFailed) && run.Reservation.State == sessions.EnvironmentInputPending { - return w.dispatcher.Store.FailEnvironmentInput(ctx, item.TenantID, item.SessionID, item.reservationID, "runtime_preparation_failed") + return w.dispatcher.sessionExecution.FailEnvironmentInput(ctx, item.TenantID, item.SessionID, item.reservationID, "runtime_preparation_failed") } if run.Reservation.State == sessions.EnvironmentInputAdmitted { return err diff --git a/services/core/internal/execution/worker_wakeup.go b/services/core/internal/execution/worker_wakeup.go index 403e2bead..153a6b962 100644 --- a/services/core/internal/execution/worker_wakeup.go +++ b/services/core/internal/execution/worker_wakeup.go @@ -16,7 +16,7 @@ func (w *Worker) wakeScheduler() { } func (w *Worker) admitInputs(ctx context.Context, tenant, session, key string, inputs []sessions.Input) ([]sessions.InputReceipt, error) { - receipts, err := w.admission.SubmitInputs(ctx, tenant, session, key, inputs) + receipts, err := w.dispatcher.Sessions.SubmitInputs(ctx, tenant, session, key, inputs) if err == nil { w.wakeScheduler() w.dispatcher.notifications.notify(tenant, session) diff --git a/services/core/internal/persistence/postgres/sessionpg/environment_inputs_test.go b/services/core/internal/persistence/postgres/sessionpg/environment_inputs_test.go new file mode 100644 index 000000000..9a000de16 --- /dev/null +++ b/services/core/internal/persistence/postgres/sessionpg/environment_inputs_test.go @@ -0,0 +1,165 @@ +package sessionpg + +import ( + "context" + "encoding/json" + "errors" + "reflect" + "sync" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgtype" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" +) + +var reservedBatch = []sessions.Input{messageInput("first"), messageInput("second")} + +// environmentInputHistory checks the Session's Turns and admitted inputs, one +// Item per input, and that no input left Turn or Item events. +func environmentInputHistory(t *testing.T, pool *pgxpool.Pool, session pgtype.UUID, turns, inputs int) { + t.Helper() + var gotTurns, gotInputs, items, events int + err := pool.QueryRow(t.Context(), ` + SELECT (SELECT count(*) FROM turns WHERE session_id=$1), + (SELECT count(*) FROM turn_inputs WHERE session_id=$1), + (SELECT count(*) FROM session_items WHERE session_id=$1), + (SELECT count(*) FROM session_events WHERE session_id=$1 + AND (payload ? 'turn' OR payload->'event' ? 'item'))`, session).Scan(&gotTurns, &gotInputs, &items, &events) + if err != nil || gotTurns != turns || gotInputs != inputs || items != inputs || (inputs == 0 && events != 0) { + t.Fatal("history", gotTurns, gotInputs, items, events, err) + } +} + +// reserve reserves the two-message batch under key. +func reserve(t *testing.T, service *sessions.Service, tenant, session pgtype.UUID, key string) sessions.EnvironmentInputReservation { + t.Helper() + got, err := service.ReserveEnvironmentInput(t.Context(), text(tenant), text(session), key, reservedBatch) + if err != nil { + t.Fatal(err) + } + return got +} + +// passDeadline moves the reservation's deadline into the past. +func passDeadline(t *testing.T, pool *pgxpool.Pool, reservation string) { + t.Helper() + exec(t, pool, `UPDATE environment_input_reservations SET deadline=clock_timestamp()-interval '1 second' WHERE id=$1`, reservation) +} + +// cancelPending cancels the Session's pending input as Session cancellation +// does, then reads the reservation back. +func cancelPending(ctx context.Context, pool *pgxpool.Pool, store *Store, tenant, session pgtype.UUID, reservation string) (sessions.EnvironmentInputReservation, error) { + err := WithSession(ctx, pgunit.NewPool(pool), tenant, session, func(ctx context.Context, q *sqlc.Queries, locked sessions.LockedSession) error { + if err := locked.Public(); err != nil { + return err + } + bound := BindSession(q, tenant, session) + return sessions.TrackInputActivity(ctx, bound, bound.CancelPendingInput) + }) + if err != nil { + return sessions.EnvironmentInputReservation{}, err + } + return store.GetEnvironmentInputReservation(ctx, text(tenant), text(session), reservation) +} + +// Concurrent equivalent reservations from two pools share one pending +// identity and deadline, which a restart keeps. +func TestEnvironmentInputReservationConcurrentIdentity(t *testing.T) { + pool := pgtest.Open(t) + _, service := stagingService(t, pool) + _, other := stagingService(t, pgtest.Open(t)) + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + ctx := t.Context() + batch := []sessions.Input{ + {Kind: "message", Payload: json.RawMessage(`{"text":"first","detail":{"a":1,"b":2}}`)}, + messageInput("second"), + } + const count = 8 + results := make(chan sessions.EnvironmentInputReservation, count) + var wg sync.WaitGroup + for i := range count { + wg.Go(func() { + admission := service + inputs := append([]sessions.Input(nil), batch...) + if i%2 == 0 { + admission = other + inputs[0].Payload = json.RawMessage(` { "detail": {"b": 2, "a": 1}, "text": "first" } `) + } + got, err := admission.ReserveEnvironmentInput(ctx, text(tenant), text(session), "request", inputs) + if err != nil { + t.Error(err) + return + } + results <- got + }) + } + wg.Wait() + close(results) + var first sessions.EnvironmentInputReservation + received := 0 + for result := range results { + received++ + if first.ID == "" { + first = result + } + if !reflect.DeepEqual(first, result) { + t.Fatal("reservation identity changed", first, result) + } + } + if received != count || first.State != sessions.EnvironmentInputPending || first.ID == "" || first.Deadline.Sub(first.CreatedAt) != 5*time.Minute || first.SettledAt != nil || len(first.Receipts) != 0 { + t.Fatal("invalid pending result", received, first) + } + environmentInputHistory(t, pool, session, 0, 0) + for _, changed := range [][]sessions.Input{batch[:1], {batch[1], batch[0]}, {messageInput("changed"), batch[1]}} { + if _, err := service.ReserveEnvironmentInput(ctx, text(tenant), text(session), "request", changed); !errors.Is(err, sessions.ErrIdempotencyConflict) { + t.Fatal("changed request accepted", err) + } + } + if _, err := other.ReserveEnvironmentInput(ctx, text(tenant), text(session), "other", batch); !errors.Is(err, sessions.ErrTurnConflict) { + t.Fatal("second pending request accepted", err) + } + pool.Close() + _, restarted := stagingService(t, pgtest.Open(t)) + got, err := restarted.ReserveEnvironmentInput(ctx, text(tenant), text(session), "request", batch) + if err != nil || !reflect.DeepEqual(first, got) { + t.Fatal("restart changed deadline or identity", got, err) + } +} + +// A key that direct admission already used keeps its receipts and gains no +// reservation, and its retry leaves newer pending input alone. +func TestEnvironmentInputReservationKeepsEarlierDirectIdentity(t *testing.T) { + pool := pgtest.Open(t) + store, service := stagingService(t, pool) + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + ctx := t.Context() + input := messageInput("already admitted") + receipts, err := service.SubmitInputs(ctx, text(tenant), text(session), "direct", []sessions.Input{input}) + if err != nil { + t.Fatal(err) + } + move(t, pool, tenant, session, receipts[0].TurnID, sessions.TurnQueued, sessions.TurnInProgress) + move(t, pool, tenant, session, receipts[0].TurnID, sessions.TurnInProgress, sessions.TurnCompleted) + pending := reserve(t, service, tenant, session, "new") + got, err := service.ReserveEnvironmentInput(ctx, text(tenant), text(session), "direct", []sessions.Input{input}) + if err != nil || got.State != sessions.EnvironmentInputAdmitted || got.ID != "" || !got.Deadline.IsZero() || len(got.Receipts) != 1 || got.Receipts[0].Sequence != receipts[0].Sequence { + t.Fatal("direct admission gained a reservation", got, err) + } + if _, err := service.ReserveEnvironmentInput(ctx, text(tenant), text(session), "direct", []sessions.Input{messageInput("changed")}); !errors.Is(err, sessions.ErrIdempotencyConflict) { + t.Fatal(err) + } + retry, err := service.SubmitInputs(ctx, text(tenant), text(session), "direct", []sessions.Input{input}) + if err != nil || len(retry) != 1 || !retry[0].Replayed { + t.Fatal(retry, err) + } + retained, err := store.GetEnvironmentInputReservation(ctx, text(tenant), text(session), pending.ID) + if err != nil || !reflect.DeepEqual(retained, pending) { + t.Fatal("old retry affected new pending input", retained, err) + } +} diff --git a/services/core/internal/persistence/postgres/sessionpg/execution_inputs.go b/services/core/internal/persistence/postgres/sessionpg/execution_inputs.go new file mode 100644 index 000000000..77c88f0a9 --- /dev/null +++ b/services/core/internal/persistence/postgres/sessionpg/execution_inputs.go @@ -0,0 +1,37 @@ +package sessionpg + +import ( + "context" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" +) + +func (e *Execution) WithInputs(ctx context.Context, tenant, session string, apply func(context.Context, sessions.InputTx) error) error { + return withInputs(ctx, e.lease, tenant, session, apply) +} + +// WithDueInputReservations locks each due reservation's Session with the scan +// itself, skipping Sessions another transaction holds, and prunes each +// Session's change journal after apply. +func (e *Execution) WithDueInputReservations(ctx context.Context, apply func(context.Context, sessions.InputTx, string) error) error { + return e.lease.Transaction(ctx, func(ctx context.Context, tx pgx.Tx) error { + q := sqlc.New(tx) + rows, err := q.ListDueEnvironmentInputs(ctx) + if err != nil { + return err + } + for _, row := range rows { + if err := apply(ctx, BindSession(q, row.TenantID, row.SessionID), uuid.UUID(row.ID.Bytes).String()); err != nil { + return err + } + if err := PruneChanges(ctx, q, row.SessionID); err != nil { + return err + } + } + return nil + }) +} diff --git a/services/core/internal/persistence/postgres/sessionpg/execution_inputs_test.go b/services/core/internal/persistence/postgres/sessionpg/execution_inputs_test.go new file mode 100644 index 000000000..22807a559 --- /dev/null +++ b/services/core/internal/persistence/postgres/sessionpg/execution_inputs_test.go @@ -0,0 +1,776 @@ +package sessionpg + +import ( + "context" + "errors" + "reflect" + "strings" + "sync" + "testing" + "time" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgtype" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" +) + +// leasedInputs opens a database of its own and returns its pool, the Session +// store and service on the pool, and the execution operations on its lease. +func leasedInputs(t *testing.T) (*pgxpool.Pool, *Store, *sessions.Service, *sessions.ExecutionOperations) { + t.Helper() + pool := pgtest.OpenIsolated(t, nil) + store, service := stagingService(t, pool) + operations, _ := sessionExecution(t, pool) + return pool, store, service, operations +} + +// terminateLeaseOwner ends the backend that holds the execution lease of +// pool's database. +func terminateLeaseOwner(t *testing.T, pool *pgxpool.Pool) { + t.Helper() + var killed bool + err := pool.QueryRow(t.Context(), `SELECT pg_terminate_backend(pid, 1000) 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(&killed) + if err != nil || !killed { + t.Fatal("terminate execution lease owner", killed, err) + } +} + +func reservationState(t *testing.T, store *Store, tenant, session pgtype.UUID, reservation string) string { + t.Helper() + got, err := store.GetEnvironmentInputReservation(t.Context(), text(tenant), text(session), reservation) + if err != nil { + t.Fatal(err) + } + return got.State +} + +// lockSession holds the Session's lock in an open transaction and returns the +// holder's backend PID; the transaction rolls back when the test ends unless +// release commits it first. +func lockSession(ctx context.Context, t *testing.T, pool *pgxpool.Pool, session pgtype.UUID) (int32, func(sql string, args ...any)) { + t.Helper() + tx, err := pool.Begin(ctx) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = tx.Rollback(context.Background()) }) + var pid int32 + if err := tx.QueryRow(ctx, "SELECT pg_backend_pid() FROM sessions WHERE id=$1 FOR UPDATE", session).Scan(&pid); err != nil { + t.Fatal(err) + } + return pid, func(sql string, args ...any) { + t.Helper() + if sql != "" { + if _, err := tx.Exec(ctx, sql, args...); err != nil { + t.Fatal(err) + } + } + if err := tx.Commit(ctx); err != nil { + t.Fatal(err) + } + } +} + +// awaitBlocked waits until a backend waits on holder's lock and returns its +// PID. +func awaitBlocked(ctx context.Context, t *testing.T, pool *pgxpool.Pool, holder int32) int32 { + t.Helper() + for ctx.Err() == nil { + var blocked int32 + if err := pool.QueryRow(ctx, "SELECT pid FROM pg_stat_activity WHERE $1=ANY(pg_blocking_pids(pid)) ORDER BY pid LIMIT 1", holder).Scan(&blocked); err == nil { + return blocked + } + time.Sleep(5 * time.Millisecond) + } + t.Fatal("transaction did not wait for the Session lock") + return 0 +} + +// Direct input waits behind a pending reservation; promotion admits it once, +// and every retry, after a new execution owner and pool too, replays that +// admission. +func TestEnvironmentInputReservationPromotionAndDirectRetries(t *testing.T) { + pool := pgtest.OpenIsolated(t, nil) + store, service := stagingService(t, pool) + operations, lease := sessionExecution(t, pool) + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + ctx := t.Context() + first := reserve(t, service, tenant, session, "pending") + for _, request := range []struct { + key string + inputs []sessions.Input + want error + }{ + {"pending", first.Inputs, sessions.ErrTurnConflict}, + {"pending", []sessions.Input{messageInput("changed")}, sessions.ErrIdempotencyConflict}, + {"later", []sessions.Input{messageInput("later")}, sessions.ErrTurnConflict}, + {"cancel", []sessions.Input{cancelInput}, sessions.ErrTurnConflict}, + } { + if _, err := service.SubmitInputs(ctx, text(tenant), text(session), request.key, request.inputs); !errors.Is(err, request.want) { + t.Fatal("direct path bypassed reservation", request.key, err) + } + } + environmentInputHistory(t, pool, session, 0, 0) + promoted, err := operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), first.ID) + if err != nil || promoted.State != sessions.EnvironmentInputAdmitted || promoted.SettledAt == nil || len(promoted.Receipts) != 2 || !promoted.Deadline.Equal(first.Deadline) { + t.Fatal(promoted, err) + } + for i, receipt := range promoted.Receipts { + if receipt.Replayed || receipt.TurnID == "" || receipt.TurnID != promoted.Receipts[0].TurnID || (i > 0 && receipt.Sequence <= promoted.Receipts[i-1].Sequence) { + t.Fatal("promotion receipts", promoted.Receipts) + } + } + environmentInputHistory(t, pool, session, 1, 2) + for _, read := range []func() (sessions.EnvironmentInputReservation, error){ + func() (sessions.EnvironmentInputReservation, error) { + return operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), first.ID) + }, + func() (sessions.EnvironmentInputReservation, error) { + return store.GetEnvironmentInputReservation(ctx, text(tenant), text(session), first.ID) + }, + func() (sessions.EnvironmentInputReservation, error) { + return service.ReserveEnvironmentInput(ctx, text(tenant), text(session), "pending", first.Inputs) + }, + } { + retry, err := read() + if err != nil || retry.ID != first.ID || !retry.Deadline.Equal(first.Deadline) || retry.State != sessions.EnvironmentInputAdmitted || len(retry.Receipts) != 2 { + t.Fatal(retry, err) + } + for i, receipt := range retry.Receipts { + if !receipt.Replayed || receipt.Sequence != promoted.Receipts[i].Sequence || receipt.TurnID != promoted.Receipts[i].TurnID { + t.Fatal("retry changed admission", receipt) + } + } + } + retry, err := service.SubmitInputs(ctx, text(tenant), text(session), "pending", first.Inputs) + if err != nil || len(retry) != 2 || !retry[0].Replayed || retry[0].Sequence != promoted.Receipts[0].Sequence { + t.Fatal("direct retry after promotion", retry, err) + } + awaitRelease := pgtest.ObserveExecutionLeaseRelease(t, pool) + if err := lease.Close(ctx); err != nil { + t.Fatal(err) + } + awaitRelease() + restarted, err := pgxpool.NewWithConfig(ctx, pool.Config()) + if err != nil { + t.Fatal(err) + } + t.Cleanup(restarted.Close) + pool.Close() + successor, _ := sessionExecution(t, restarted) + after, err := successor.PromoteEnvironmentInput(ctx, text(tenant), text(session), first.ID) + if err != nil || after.State != sessions.EnvironmentInputAdmitted || after.Receipts[0].Sequence != promoted.Receipts[0].Sequence { + t.Fatal("restart repeated promotion", after, err) + } + environmentInputHistory(t, restarted, session, 1, 2) +} + +// Reservations resolve only in their tenant's Session, take only messages, +// need an Environment and join an active Turn at once. +func TestEnvironmentInputReservationRejectsUnsupportedOrForeignState(t *testing.T) { + pool, store, service, operations := leasedInputs(t) + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + ctx := t.Context() + pending := reserve(t, service, tenant, session, "pending") + stranger := pgID(uuid.New()) + for _, read := range []func() error{ + func() error { + _, err := store.GetEnvironmentInputReservation(ctx, text(stranger), text(session), pending.ID) + return err + }, + func() error { + _, err := operations.PromoteEnvironmentInput(ctx, text(stranger), text(session), pending.ID) + return err + }, + func() error { + _, err := cancelPending(ctx, pool, store, stranger, session, pending.ID) + return err + }, + func() error { + _, err := service.ReserveEnvironmentInput(ctx, text(stranger), text(session), "new", pending.Inputs) + return err + }, + } { + if err := read(); !errors.Is(err, sessions.ErrNotFound) { + t.Fatal("foreign access", err) + } + } + other := pgID(uuid.New()) + exec(t, pool, `INSERT INTO sessions(id, tenant_id, engine, idempotency_key, request_hash) VALUES ($1, $2, 'codex', 'other', 'hash')`, other, tenant) + for _, action := range []func(context.Context, pgtype.UUID, pgtype.UUID, string) (sessions.EnvironmentInputReservation, error){ + func(ctx context.Context, tenant, session pgtype.UUID, reservation string) (sessions.EnvironmentInputReservation, error) { + return store.GetEnvironmentInputReservation(ctx, text(tenant), text(session), reservation) + }, + func(ctx context.Context, tenant, session pgtype.UUID, reservation string) (sessions.EnvironmentInputReservation, error) { + return operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), reservation) + }, + func(ctx context.Context, tenant, session pgtype.UUID, reservation string) (sessions.EnvironmentInputReservation, error) { + return cancelPending(ctx, pool, store, tenant, session, reservation) + }, + func(ctx context.Context, tenant, session pgtype.UUID, reservation string) (sessions.EnvironmentInputReservation, error) { + return service.ExpireEnvironmentInput(ctx, text(tenant), text(session), reservation) + }, + } { + if _, err := action(ctx, tenant, other, pending.ID); !errors.Is(err, sessions.ErrNotFound) { + t.Fatal("reservation crossed Session ownership", err) + } + if _, err := action(ctx, stranger, session, pending.ID); !errors.Is(err, sessions.ErrNotFound) { + t.Fatal("reservation crossed tenant ownership", err) + } + } + retained, err := store.GetEnvironmentInputReservation(ctx, text(tenant), text(session), pending.ID) + if err != nil || !reflect.DeepEqual(retained, pending) { + t.Fatal("foreign operations changed reservation", retained, err) + } + for _, invalid := range [][]sessions.Input{nil, {cancelInput}, {{Kind: "tool_result", Payload: []byte(`{}`)}}} { + if _, err := service.ReserveEnvironmentInput(ctx, text(tenant), text(session), "invalid", invalid); !errors.Is(err, sessions.ErrInvalidInput) { + t.Fatal("unsupported reservation", err) + } + } + noneTenant, none := newSession(t, pool) + if _, err := service.ReserveEnvironmentInput(ctx, text(noneTenant), text(none), "none", pending.Inputs); !errors.Is(err, sessions.ErrInvalidInput) { + t.Fatal("none reservation", err) + } + activeTenant, active, _ := newEnvironment(t, pool, "self_hosted", "pending") + activeInput := submitMessage(t, service, activeTenant, active, "active") + steer, err := service.ReserveEnvironmentInput(ctx, text(activeTenant), text(active), "new", pending.Inputs) + if err != nil || steer.State != sessions.EnvironmentInputAdmitted || steer.ID != "" || !steer.Deadline.IsZero() || len(steer.Receipts) != len(pending.Inputs) || steer.Receipts[0].TurnID != activeInput.TurnID { + t.Fatal("active input did not retain the existing Turn", steer, err) + } +} + +// A cancelled or expired reservation keeps its outcome and identity: no retry +// or settlement restarts it, and none touches a newer reservation. +func TestEnvironmentInputTerminalReservationsCannotRestart(t *testing.T) { + pool, store, service, operations := leasedInputs(t) + for _, terminal := range []string{sessions.EnvironmentInputCancelled, sessions.EnvironmentInputExpired} { + t.Run(terminal, func(t *testing.T) { + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + ctx := t.Context() + pending := reserve(t, service, tenant, session, "pending") + early, err := service.ExpireEnvironmentInput(ctx, text(tenant), text(session), pending.ID) + if err != nil || early.State != sessions.EnvironmentInputPending || early.SettledAt != nil || !early.Deadline.Equal(pending.Deadline) { + t.Fatal("early expiry", early, err) + } + var settled sessions.EnvironmentInputReservation + if terminal == sessions.EnvironmentInputExpired { + passDeadline(t, pool, pending.ID) + settled, err = operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), pending.ID) + } else { + settled, err = cancelPending(ctx, pool, store, tenant, session, pending.ID) + } + if err != nil || settled.State != terminal || settled.SettledAt == nil || len(settled.Receipts) != 0 { + t.Fatal("terminal settlement", settled, err) + } + retry, err := service.ReserveEnvironmentInput(ctx, text(tenant), text(session), "pending", pending.Inputs) + if err != nil || retry.State != terminal || retry.ID != pending.ID || !retry.Deadline.Equal(settled.Deadline) || !retry.SettledAt.Equal(*settled.SettledAt) { + t.Fatal("terminal retry changed outcome", retry, err) + } + if _, err := service.SubmitInputs(ctx, text(tenant), text(session), "pending", pending.Inputs); !errors.Is(err, sessions.ErrTurnConflict) { + t.Fatal("terminal request reopened through direct path", err) + } + if _, err := service.SubmitInputs(ctx, text(tenant), text(session), "pending", []sessions.Input{messageInput("changed")}); !errors.Is(err, sessions.ErrIdempotencyConflict) { + t.Fatal("terminal identity changed", err) + } + later := reserve(t, service, tenant, session, "later") + // Cancellation takes whatever input is pending, so only the + // settlements that name the reservation run against it here. + for _, finish := range []func(context.Context, string, string, string) (sessions.EnvironmentInputReservation, error){ + operations.PromoteEnvironmentInput, service.ExpireEnvironmentInput, + } { + got, err := finish(ctx, text(tenant), text(session), pending.ID) + if err != nil || got.State != terminal { + t.Fatal("old settlement changed", got, err) + } + } + got, err := store.GetEnvironmentInputReservation(ctx, text(tenant), text(session), later.ID) + if err != nil || got.State != sessions.EnvironmentInputPending || !got.Deadline.Equal(later.Deadline) { + t.Fatal("old settlement touched successor", got, err) + } + environmentInputHistory(t, pool, session, 0, 0) + }) + } +} + +// A promotion that fails part way leaves no Turn, input, Item, claim or +// settlement, and its retry promotes the reservation whole. +func TestEnvironmentInputPromotionRollsBackHistoryAndSettlement(t *testing.T) { + pool, store, service, operations := leasedInputs(t) + for _, phase := range []string{"input", "settlement", "claim", "claim-event"} { + t.Run(phase, func(t *testing.T) { + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + ctx := t.Context() + pending := reserve(t, service, tenant, session, "pending") + name := "reservation_failure_" + strings.ReplaceAll(uuid.NewString(), "-", "") + table, expression := "turn_inputs", "session_id <> '"+text(session)+"'::uuid OR payload->>'text' <> 'second'" + switch phase { + case "settlement": + table, expression = "environment_input_reservations", "id <> '"+pending.ID+"'::uuid OR state <> 'admitted'" + case "claim": + table, expression = "turns", "session_id <> '"+text(session)+"'::uuid OR status <> 'in_progress'" + case "claim-event": + table, expression = "session_events", "session_id <> '"+text(session)+"'::uuid OR payload->'event'->>'type' <> 'agent.session.turn.in_progress'" + } + exec(t, pool, "ALTER TABLE "+table+" ADD CONSTRAINT "+name+" CHECK ("+expression+") NOT VALID") + t.Cleanup(func() { + _, _ = pool.Exec(context.Background(), "ALTER TABLE "+table+" DROP CONSTRAINT IF EXISTS "+name) + }) + if _, err := operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), pending.ID); err == nil { + t.Fatal("injected failure succeeded") + } + environmentInputHistory(t, pool, session, 0, 0) + got, err := store.GetEnvironmentInputReservation(ctx, text(tenant), text(session), pending.ID) + if err != nil || got.State != sessions.EnvironmentInputPending || got.SettledAt != nil || !got.Deadline.Equal(pending.Deadline) { + t.Fatal("partial settlement survived", got, err) + } + exec(t, pool, "ALTER TABLE "+table+" DROP CONSTRAINT "+name) + got, err = operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), pending.ID) + if err != nil || got.State != sessions.EnvironmentInputAdmitted { + t.Fatal(got, err) + } + environmentInputHistory(t, pool, session, 1, 2) + }) + } +} + +// Promotion and failure read the deadline after they take the Session lock, +// so waiting for the lock never extends the reservation's lifetime. +func TestEnvironmentInputDeadlineIsCheckedAfterSessionLock(t *testing.T) { + pool, store, service, operations := leasedInputs(t) + for _, action := range []string{"promote", "fail"} { + t.Run(action, func(t *testing.T) { + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + pending := reserve(t, service, tenant, session, "pending") + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + defer cancel() + blocker, release := lockSession(ctx, t, pool, session) + type outcome struct { + value sessions.EnvironmentInputReservation + err error + } + done := make(chan outcome, 1) + go func() { + var got sessions.EnvironmentInputReservation + var err error + if action == "promote" { + got, err = operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), pending.ID) + } else if err = operations.FailEnvironmentInput(ctx, text(tenant), text(session), pending.ID, "runtime_preparation_failed"); err == nil { + got, err = store.GetEnvironmentInputReservation(ctx, text(tenant), text(session), pending.ID) + } + done <- outcome{got, err} + }() + awaitBlocked(ctx, t, pool, blocker) + // Transaction-start time is now older than the controlled deadline. + release("UPDATE environment_input_reservations SET deadline=clock_timestamp() WHERE id=$1", pending.ID) + result := <-done + if result.err != nil || result.value.State != sessions.EnvironmentInputExpired || result.value.SettledAt == nil { + t.Fatal("lock wait extended input lifetime", result) + } + environmentInputHistory(t, pool, session, 0, 0) + if state := reservationState(t, store, tenant, session, pending.ID); state != sessions.EnvironmentInputExpired { + t.Fatal("expiry was rolled back", state) + } + }) + } +} + +// A racing promotion and cancellation agree on one outcome. +func TestEnvironmentInputCancelAndPromotionShareOneOutcome(t *testing.T) { + pool, store, service, operations := leasedInputs(t) + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + ctx := t.Context() + pending := reserve(t, service, tenant, session, "pending") + start := make(chan struct{}) + results := make(chan sessions.EnvironmentInputReservation, 2) + errs := make(chan error, 2) + for _, finish := range []func() (sessions.EnvironmentInputReservation, error){ + func() (sessions.EnvironmentInputReservation, error) { + return operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), pending.ID) + }, + func() (sessions.EnvironmentInputReservation, error) { + return cancelPending(ctx, pool, store, tenant, session, pending.ID) + }, + } { + go func() { + <-start + got, err := finish() + results <- got + errs <- err + }() + } + close(start) + first, second := <-results, <-results + for range 2 { + if err := <-errs; err != nil { + t.Fatal(err) + } + } + if first.State != second.State || (first.State != sessions.EnvironmentInputAdmitted && first.State != sessions.EnvironmentInputCancelled) { + t.Fatal("competing settlements diverged", first.State, second.State) + } + turns, inputs := 0, 0 + if first.State == sessions.EnvironmentInputAdmitted { + turns, inputs = 1, 2 + } + environmentInputHistory(t, pool, session, turns, inputs) +} + +// Each expiry pass settles at most 32 due reservations and creates no +// history. +func TestEnvironmentExpiryBoundsBatch(t *testing.T) { + pool, store, service, operations := leasedInputs(t) + ctx := t.Context() + type due struct { + tenant, session pgtype.UUID + reservation string + } + var reservations []due + for range 33 { + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + pending := reserve(t, service, tenant, session, "pending") + passDeadline(t, pool, pending.ID) + reservations = append(reservations, due{tenant, session, pending.ID}) + } + for _, want := range []int64{32, 1, 0} { + n, err := operations.ExpireEnvironmentInputs(ctx) + if err != nil || n != want { + t.Fatal("unbounded or incomplete batch", n, want, err) + } + } + for _, r := range reservations { + if state := reservationState(t, store, r.tenant, r.session, r.reservation); state != sessions.EnvironmentInputExpired { + t.Fatal(state) + } + environmentInputHistory(t, pool, r.session, 0, 0) + } +} + +// An execution owner that lost its lease expires nothing; its successor +// does. +func TestEnvironmentExpiryFencesLostExecutionOwner(t *testing.T) { + pool, store, service, old := leasedInputs(t) + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + pending := reserve(t, service, tenant, session, "pending") + passDeadline(t, pool, pending.ID) + terminateLeaseOwner(t, pool) + successor, _ := sessionExecution(t, pool) + if n, err := old.ExpireEnvironmentInputs(t.Context()); err == nil || n != 0 { + t.Fatal("lost owner expired input", n, err) + } + if state := reservationState(t, store, tenant, session, pending.ID); state != sessions.EnvironmentInputPending { + t.Fatal("lost owner expired input", state) + } + if n, err := successor.ExpireEnvironmentInputs(t.Context()); err != nil || n != 1 { + t.Fatal("successor could not expire input", n, err) + } + if state := reservationState(t, store, tenant, session, pending.ID); state != sessions.EnvironmentInputExpired { + t.Fatal(state) + } + environmentInputHistory(t, pool, session, 0, 0) +} + +// Promotion and expiry of a reservation past its deadline, in either order or +// racing, end with one expired outcome and no history, and leave a newer +// reservation pending. +func TestEnvironmentInputPromotionSerializesWithExpiry(t *testing.T) { + pool, store, service, operations := leasedInputs(t) + promote := func(ctx context.Context, tenant, session pgtype.UUID, reservation string) error { + got, err := operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), reservation) + if err == nil && got.State != sessions.EnvironmentInputExpired { + return errors.New("promotion outran expiry: " + got.State) + } + return err + } + sweep := func(ctx context.Context, _, _ pgtype.UUID, _ string) error { + _, err := operations.ExpireEnvironmentInputs(ctx) + return err + } + expire := func(ctx context.Context, tenant, session pgtype.UUID, reservation string) error { + got, err := service.ExpireEnvironmentInput(ctx, text(tenant), text(session), reservation) + if err == nil && got.State != sessions.EnvironmentInputExpired { + return errors.New("expiry missed its deadline: " + got.State) + } + return err + } + type settle func(context.Context, pgtype.UUID, pgtype.UUID, string) error + for name, actions := range map[string][]settle{ + "sweep-then-promote": {sweep, promote}, + "promote-then-sweep": {promote, sweep}, + "racing-sweep": {sweep, promote}, + "racing-expire": {expire, promote}, + } { + t.Run(name, func(t *testing.T) { + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + ctx := t.Context() + pending := reserve(t, service, tenant, session, "pending") + passDeadline(t, pool, pending.ID) + if strings.HasPrefix(name, "racing") { + start := make(chan struct{}) + errs := make(chan error, len(actions)) + for _, action := range actions { + go func() { <-start; errs <- action(ctx, tenant, session, pending.ID) }() + } + close(start) + for range actions { + if err := <-errs; err != nil { + t.Fatal(err) + } + } + } else { + for _, action := range actions { + if err := action(ctx, tenant, session, pending.ID); err != nil { + t.Fatal(err) + } + } + } + if state := reservationState(t, store, tenant, session, pending.ID); state != sessions.EnvironmentInputExpired { + t.Fatal("invalid competing settlement", state) + } + environmentInputHistory(t, pool, session, 0, 0) + later := reserve(t, service, tenant, session, "later") + for _, action := range []settle{promote, expire, sweep} { + if err := action(ctx, tenant, session, pending.ID); err != nil { + t.Fatal("old reservation changed", err) + } + } + got, err := store.GetEnvironmentInputReservation(ctx, text(tenant), text(session), later.ID) + if err != nil || got.State != sessions.EnvironmentInputPending || !got.Deadline.Equal(later.Deadline) { + t.Fatal("old settlement affected successor", got, err) + } + }) + } +} + +// Concurrent promotions claim one Turn once: one fresh admission, one claim +// event after the admission's, and a retry after the Turn ends publishes +// nothing. +func TestEnvironmentInputConcurrentPromotionClaimsOnce(t *testing.T) { + pool, store, service, operations := leasedInputs(t) + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + pending := reserve(t, service, tenant, session, "pending") + before, _ := journal(t, pool, session) + const count = 8 + results := make(chan sessions.EnvironmentInputReservation, count) + var group sync.WaitGroup + for range count { + group.Go(func() { + got, err := operations.PromoteEnvironmentInput(t.Context(), text(tenant), text(session), pending.ID) + if err != nil { + t.Error(err) + return + } + results <- got + }) + } + group.Wait() + close(results) + var turnID string + fresh, received := 0, 0 + for got := range results { + received++ + if got.State != sessions.EnvironmentInputAdmitted || len(got.Receipts) != 2 { + t.Fatal("promotion lost the original batch", got) + } + if turnID == "" { + turnID = got.Receipts[0].TurnID + } + if !got.Receipts[0].Replayed { + fresh++ + } + for _, receipt := range got.Receipts { + if receipt.TurnID != turnID || receipt.Replayed != got.Receipts[0].Replayed { + t.Fatal("promotion changed execution ownership", got.Receipts) + } + } + } + if received != count || fresh != 1 { + t.Fatal("promotion authorized multiple starts", received, fresh) + } + turn, err := store.GetTurn(t.Context(), text(tenant), text(session), turnID) + if err != nil || turn.Status != sessions.TurnInProgress || turn.StartedAt.IsZero() { + t.Fatal("promotion did not persist its execution claim", turn, err) + } + environmentInputHistory(t, pool, session, 1, 2) + _, changes := journal(t, pool, session) + changes = changes[len(before):] + if len(changes) < 2 { + t.Fatal("missing promotion events", kinds(changes)) + } + created, claimed := changes[0], changes[len(changes)-1] + if created.Event.Type != "agent.session.turn.created" || created.Turn == nil || created.Turn.Status != sessions.TurnQueued || claimed.Event.Type != "agent.session.turn.in_progress" || claimed.Turn == nil || claimed.Turn.ID != turnID || claimed.Turn.Status != sessions.TurnInProgress { + t.Fatal("claim reordered or replaced admission snapshots", kinds(changes)) + } + claims := 0 + for _, change := range changes { + if change.Event.Type == "agent.session.turn.in_progress" { + claims++ + } + } + if claims != 1 { + t.Fatal("retry published another claim", claims) + } + move(t, pool, tenant, session, turnID, sessions.TurnInProgress, sessions.TurnCompleted) + later := reserve(t, service, tenant, session, "later") + cursor, _ := journal(t, pool, session) + retry, err := operations.PromoteEnvironmentInput(t.Context(), text(tenant), text(session), pending.ID) + if err != nil || len(retry.Receipts) != 2 || !retry.Receipts[0].Replayed || retry.Receipts[0].TurnID != turnID { + t.Fatal("terminal retry reclaimed execution", retry, err) + } + if after, _ := journal(t, pool, session); !reflect.DeepEqual(after, cursor) { + t.Fatal("terminal retry published events", after, cursor) + } + retained, err := store.GetEnvironmentInputReservation(t.Context(), text(tenant), text(session), later.ID) + if err != nil || retained.State != sessions.EnvironmentInputPending || !retained.Deadline.Equal(later.Deadline) { + t.Fatal("old promotion affected new preparation", retained, err) + } +} + +// Promotion and preparation failure from a closed lease or a lost execution +// owner write nothing; the successor settles both. +func TestEnvironmentInputSettlementUsesCurrentExecutionOwner(t *testing.T) { + pool := pgtest.OpenIsolated(t, nil) + store, service := stagingService(t, pool) + promoteTenant, promoteSession, _ := newEnvironment(t, pool, "self_hosted", "pending") + failTenant, failSession, _ := newEnvironment(t, pool, "self_hosted", "pending") + promoted := reserve(t, service, promoteTenant, promoteSession, "pending") + failed := reserve(t, service, failTenant, failSession, "pending") + settle := func(operations *sessions.ExecutionOperations) (error, error) { + _, promoteErr := operations.PromoteEnvironmentInput(t.Context(), text(promoteTenant), text(promoteSession), promoted.ID) + return promoteErr, operations.FailEnvironmentInput(t.Context(), text(failTenant), text(failSession), failed.ID, "runtime_preparation_failed") + } + unsettled := func() { + t.Helper() + environmentInputHistory(t, pool, promoteSession, 0, 0) + environmentInputHistory(t, pool, failSession, 0, 0) + if reservationState(t, store, promoteTenant, promoteSession, promoted.ID) != sessions.EnvironmentInputPending || reservationState(t, store, failTenant, failSession, failed.ID) != sessions.EnvironmentInputPending { + t.Fatal("fenced owner settled input") + } + } + closed, lease := sessionExecution(t, pool) + awaitRelease := pgtest.ObserveExecutionLeaseRelease(t, pool) + if err := lease.Close(t.Context()); err != nil { + t.Fatal(err) + } + awaitRelease() + if promoteErr, failErr := settle(closed); !errors.Is(promoteErr, pgunit.ErrLeaseClosed) || !errors.Is(failErr, pgunit.ErrLeaseClosed) { + t.Fatal("closed lease settled input", promoteErr, failErr) + } + unsettled() + stale, _ := sessionExecution(t, pool) + terminateLeaseOwner(t, pool) + successor, _ := sessionExecution(t, pool) + if promoteErr, failErr := settle(stale); promoteErr == nil || failErr == nil { + t.Fatal("stale owner settled input", promoteErr, failErr) + } + unsettled() + if promoteErr, failErr := settle(successor); promoteErr != nil || failErr != nil { + t.Fatal("successor could not settle", promoteErr, failErr) + } + environmentInputHistory(t, pool, promoteSession, 1, 2) + environmentInputHistory(t, pool, failSession, 0, 0) + if reservationState(t, store, promoteTenant, promoteSession, promoted.ID) != sessions.EnvironmentInputAdmitted || reservationState(t, store, failTenant, failSession, failed.ID) != sessions.EnvironmentInputFailed { + t.Fatal("successor settlement") + } +} + +// Input reserved while the active Turn completes serializes with the +// completion: input first joins the Turn and blocks completion until applied; +// completion first leaves the input pending for its own Turn. +func TestEnvironmentActiveInputSerializesWithCompletion(t *testing.T) { + pool, _, service, operations := leasedInputs(t) + for _, completionFirst := range []bool{false, true} { + name := "input-first" + if completionFirst { + name = "completion-first" + } + t.Run(name, func(t *testing.T) { + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + original := submitMessage(t, service, tenant, session, "original") + move(t, pool, tenant, session, original.TurnID, sessions.TurnQueued, sessions.TurnInProgress) + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + defer cancel() + blocker, release := lockSession(ctx, t, pool, session) + type admission struct { + value sessions.EnvironmentInputReservation + err error + } + admitted := make(chan admission, 1) + completed := make(chan error, 1) + input := func() { + value, err := service.ReserveEnvironmentInput(ctx, text(tenant), text(session), "racing-input", reservedBatch) + admitted <- admission{value, err} + } + complete := func() { + _, err := operations.CompleteExecution(ctx, text(tenant), text(session), original.TurnID, sessions.TurnCompleted, nil, "", original.Sequence) + completed <- err + } + first, second := input, complete + if completionFirst { + first, second = complete, input + } + go first() + firstPID := awaitBlocked(ctx, t, pool, blocker) + go second() + awaitBlocked(ctx, t, pool, firstPID) + release("") + got, completionErr := <-admitted, <-completed + if got.err != nil { + t.Fatal(got.err) + } + if completionFirst { + if completionErr != nil || got.value.State != sessions.EnvironmentInputPending || got.value.ID == "" || len(got.value.Receipts) != 0 || got.value.Deadline.Sub(got.value.CreatedAt) != 5*time.Minute { + t.Fatal("completion winner did not leave new input waiting for preparation", completionErr, got.value) + } + environmentInputHistory(t, pool, session, 1, 1) + prepared, err := operations.PromoteEnvironmentInput(ctx, text(tenant), text(session), got.value.ID) + if err != nil || len(prepared.Receipts) != 2 || prepared.Receipts[0].TurnID == original.TurnID { + t.Fatal("prepared successor reused terminal work", err) + } + retry, err := service.ReserveEnvironmentInput(ctx, text(tenant), text(session), "racing-input", reservedBatch) + if err != nil || retry.ID != got.value.ID || !retry.Deadline.Equal(got.value.Deadline) || len(retry.Receipts) != 2 || !retry.Receipts[0].Replayed || retry.Receipts[0].TurnID != prepared.Receipts[0].TurnID { + t.Fatal("active retry replaced its original reservation", retry, err) + } + if _, err := operations.CompleteExecution(ctx, text(tenant), text(session), prepared.Receipts[0].TurnID, sessions.TurnCompleted, nil, "", prepared.Receipts[1].Sequence); err != nil { + t.Fatal(err) + } + } else { + if !errors.Is(completionErr, sessions.ErrUnappliedInputs) || got.value.State != sessions.EnvironmentInputAdmitted || got.value.ID != "" || !got.value.Deadline.IsZero() || len(got.value.Receipts) != 2 || got.value.Receipts[0].TurnID != original.TurnID || got.value.Receipts[0].Replayed { + t.Fatal("admitted input escaped the original Turn or application fence", completionErr, got.value) + } + environmentInputHistory(t, pool, session, 1, 3) + if _, err := operations.CompleteExecution(ctx, text(tenant), text(session), original.TurnID, sessions.TurnCompleted, nil, "", got.value.Receipts[1].Sequence); err != nil { + t.Fatal("completion after controlled application failed", err) + } + } + }) + } +} + +// A late preparation failure leaves a cancelled reservation and a newer one +// as they are, and takes only classified codes. +func TestPreparationFailurePreservesCancelledAndNewerInput(t *testing.T) { + pool, store, service, operations := leasedInputs(t) + tenant, session, _ := newEnvironment(t, pool, "self_hosted", "pending") + first := reserve(t, service, tenant, session, "cancelled") + if _, err := cancelPending(t.Context(), pool, store, tenant, session, first.ID); err != nil { + t.Fatal(err) + } + next := reserve(t, service, tenant, session, "new") + if err := operations.FailEnvironmentInput(t.Context(), text(tenant), text(session), first.ID, "runtime_preparation_failed"); err != nil { + t.Fatal(err) + } + for id, state := range map[string]string{first.ID: sessions.EnvironmentInputCancelled, next.ID: sessions.EnvironmentInputPending} { + if current := reservationState(t, store, tenant, session, id); current != state { + t.Fatal("late failure changed another outcome", id, current) + } + } + if err := operations.FailEnvironmentInput(t.Context(), text(tenant), text(session), next.ID, "secret-canary"); !errors.Is(err, sessions.ErrInvalidInput) { + t.Fatal("unclassified diagnostic accepted", err) + } +} diff --git a/services/core/internal/persistence/postgres/sessionpg/inputs.go b/services/core/internal/persistence/postgres/sessionpg/inputs.go new file mode 100644 index 000000000..37e764773 --- /dev/null +++ b/services/core/internal/persistence/postgres/sessionpg/inputs.go @@ -0,0 +1,247 @@ +package sessionpg + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "time" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5" + "github.com/jackc/pgx/v5/pgtype" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" +) + +var ( + _ sessions.InputTx = (*SessionTx)(nil) + _ sessions.InputReader = (*Store)(nil) +) + +func (s *Store) WithInputs(ctx context.Context, tenant, session string, apply func(context.Context, sessions.InputTx) error) error { + return withInputs(ctx, s.units, tenant, session, apply) +} + +// withInputs runs apply on runner in the Session transaction of the tenant's +// visible Session. A malformed tenant is sessions.ErrInvalidInput; a malformed +// Session ID resolves as a missing Session. +func withInputs(ctx context.Context, runner pgunit.Transactor, tenantID, sessionID string, apply func(context.Context, sessions.InputTx) error) error { + tenant, err := parseID(tenantID) + if err != nil { + return err + } + session := pgunit.PathID(sessionID) + return WithSession(ctx, runner, tenant, session, func(ctx context.Context, q *sqlc.Queries, locked sessions.LockedSession) error { + if err := locked.Public(); err != nil { + return err + } + return apply(ctx, BindSession(q, tenant, session)) + }) +} + +func (t *SessionTx) CreateTurn(ctx context.Context) (sessions.Turn, error) { + row, err := t.q.CreateTurn(ctx, sqlc.CreateTurnParams{ID: pgtype.UUID{Bytes: uuid.New(), Valid: true}, SessionID: t.session}) + if err != nil { + return sessions.Turn{}, err + } + return TurnFromRow(row), nil +} + +func (t *SessionTx) CreateTurnInput(ctx context.Context, turn, key string, position int32, input sessions.Input) (int64, error) { + var id pgtype.UUID + if turn != "" { + var err error + if id, err = parseID(turn); err != nil { + return 0, err + } + } + sequence, err := t.q.CreateTurnInput(ctx, sqlc.CreateTurnInputParams{ + SessionID: t.session, TurnID: id, IdempotencyKey: key, Kind: input.Kind, Payload: input.Payload, BatchPosition: position, + }) + return sequence, storable(err) +} + +func (t *SessionTx) LoadInputBatch(ctx context.Context, key string, batch json.RawMessage) ([]sessions.InputReceipt, bool, error) { + rows, err := t.q.FindInputBatch(ctx, sqlc.FindInputBatchParams{SessionID: t.session, IdempotencyKey: key, Batch: batch}) + if err != nil { + return nil, false, storable(err) + } + receipts := make([]sessions.InputReceipt, 0, len(rows)) + matches := true + for _, row := range rows { + matches = matches && row.Matches + receipts = append(receipts, sessions.InputReceipt{Sequence: row.Sequence, TurnID: uuidString(row.TurnID), Replayed: true}) + } + return receipts, matches, nil +} + +func (t *SessionTx) LoadInputGate(ctx context.Context, key string, batch json.RawMessage) (sessions.InputGate, error) { + gate, err := t.q.CheckEnvironmentInputGate(ctx, sqlc.CheckEnvironmentInputGateParams{SessionID: t.session, IdempotencyKey: key, Batch: batch}) + return sessions.InputGate{Matches: gate.Matches, Blocked: gate.Blocked}, storable(err) +} + +func (t *SessionTx) FindInputReservation(ctx context.Context, key string, batch json.RawMessage) (*sessions.EnvironmentInputReservation, bool, error) { + row, err := t.q.FindEnvironmentInputReservation(ctx, sqlc.FindEnvironmentInputReservationParams{SessionID: t.session, IdempotencyKey: key, Batch: batch}) + if errors.Is(err, pgx.ErrNoRows) { + return nil, false, nil + } + if err != nil { + return nil, false, storable(err) + } + if !row.Matches { + return nil, true, nil + } + reservation, err := t.reservationFromRow(ctx, row.EnvironmentInputReservation) + if err != nil { + return nil, false, err + } + return &reservation, true, nil +} + +func (t *SessionTx) LoadInputReservation(ctx context.Context, reservation string) (sessions.EnvironmentInputReservation, error) { + id, err := parseID(reservation) + if err != nil { + return sessions.EnvironmentInputReservation{}, err + } + row, err := t.q.GetEnvironmentInputReservation(ctx, sqlc.GetEnvironmentInputReservationParams{SessionID: t.session, ID: id}) + if errors.Is(err, pgx.ErrNoRows) { + return sessions.EnvironmentInputReservation{}, sessions.ErrNotFound + } + if err != nil { + return sessions.EnvironmentInputReservation{}, err + } + return t.reservationFromRow(ctx, row) +} + +// reservationFromRow maps a stored Environment input reservation; an +// admitted one carries the receipts of its batch, which it must have. +func (t *SessionTx) reservationFromRow(ctx context.Context, row sqlc.EnvironmentInputReservation) (sessions.EnvironmentInputReservation, error) { + result := sessions.EnvironmentInputReservation{ + ID: uuid.UUID(row.ID.Bytes).String(), SessionID: uuid.UUID(row.SessionID.Bytes).String(), Key: row.IdempotencyKey, + State: row.State, IsInitial: row.IsInitial, CreatedAt: row.CreatedAt.Time, Deadline: row.Deadline.Time, + } + if row.SettledAt.Valid { + result.SettledAt = &row.SettledAt.Time + } + if err := json.Unmarshal(row.Batch, &result.Inputs); err != nil { + return sessions.EnvironmentInputReservation{}, fmt.Errorf("decode Environment input: %w", err) + } + if row.State != sessions.EnvironmentInputAdmitted { + return result, nil + } + receipts, matches, err := t.LoadInputBatch(ctx, row.IdempotencyKey, row.Batch) + if err != nil { + return sessions.EnvironmentInputReservation{}, err + } + if len(receipts) == 0 || !matches { + return sessions.EnvironmentInputReservation{}, errors.New("admitted Environment input has no receipts of its batch") + } + result.Receipts = receipts + return result, nil +} + +func (t *SessionTx) CreateInputReservation(ctx context.Context, key string, batch json.RawMessage) (sessions.EnvironmentInputReservation, error) { + row, err := t.q.CreateEnvironmentInputReservation(ctx, sqlc.CreateEnvironmentInputReservationParams{ + ID: pgtype.UUID{Bytes: uuid.New(), Valid: true}, SessionID: t.session, IdempotencyKey: key, Batch: batch, + }) + if err != nil { + return sessions.EnvironmentInputReservation{}, storable(err) + } + return t.reservationFromRow(ctx, row) +} + +func (t *SessionTx) ExpireInputReservation(ctx context.Context, reservation string) error { + id, err := parseID(reservation) + if err != nil { + return err + } + return t.q.ExpireEnvironmentInputReservation(ctx, sqlc.ExpireEnvironmentInputReservationParams{SessionID: t.session, ID: id}) +} + +func (t *SessionTx) AdmitInputReservation(ctx context.Context, reservation string) (time.Time, error) { + id, err := parseID(reservation) + if err != nil { + return time.Time{}, err + } + row, err := t.q.SettleEnvironmentInputReservation(ctx, sqlc.SettleEnvironmentInputReservationParams{SessionID: t.session, ID: id, State: sessions.EnvironmentInputAdmitted}) + if err != nil { + return time.Time{}, err + } + return row.SettledAt.Time, nil +} + +func (t *SessionTx) FailInputReservation(ctx context.Context, reservation, code string) error { + id, err := parseID(reservation) + if err != nil { + return err + } + _, err = t.q.FailEnvironmentInput(ctx, sqlc.FailEnvironmentInputParams{SessionID: t.session, ID: id, FailureCode: pgtype.Text{String: code, Valid: true}}) + return err +} + +func (t *SessionTx) RecordInputAudit(ctx context.Context) error { + return auditpg.RecordWriteAudit(ctx, t.q, optionalID(t.tenant), "send_events", "session", optionalID(t.session), "") +} + +func (s *Store) ListTurnInputs(ctx context.Context, tenant, session, turn string, after int64, limit int) ([]sessions.TurnInput, error) { + params, err := TurnLookup(tenant, session, turn) + if err != nil { + return nil, err + } + if after < 0 || limit < 1 || limit > 100 { + return nil, fmt.Errorf("%w: nonnegative cursor and page size 1..100 required", sessions.ErrInvalidInput) + } + if _, err := s.GetTurn(ctx, tenant, session, turn); err != nil { + return nil, err + } + rows, err := s.units.Queries().ListTurnInputs(ctx, sqlc.ListTurnInputsParams{ + TenantID: params.TenantID, SessionID: params.SessionID, TurnID: params.ID, Sequence: after, Limit: int32(limit), + }) + if err != nil { + return nil, fmt.Errorf("list turn inputs: %w", err) + } + inputs := make([]sessions.TurnInput, 0, len(rows)) + for _, row := range rows { + inputs = append(inputs, sessions.TurnInput{Sequence: row.Sequence, Kind: row.Kind, Payload: row.Payload, CreatedAt: row.CreatedAt.Time}) + } + return inputs, nil +} + +func (s *Store) GetEnvironmentInputReservation(ctx context.Context, tenantID, sessionID, reservationID string) (sessions.EnvironmentInputReservation, error) { + tenant, err := parseID(tenantID) + if err != nil { + return sessions.EnvironmentInputReservation{}, err + } + session := pgunit.PathID(sessionID) + var reservation sessions.EnvironmentInputReservation + err = s.units.Snapshot(ctx, func(ctx context.Context, tx pgx.Tx) error { + q := sqlc.New(tx) + if err := visibleSession(ctx, q, tenant, session); err != nil { + return err + } + var err error + reservation, err = BindSession(q, tenant, session).LoadInputReservation(ctx, reservationID) + return err + }) + return reservation, err +} + +func (s *Store) ListEnvironmentInputWork(ctx context.Context, after string, connectedDevices []string) ([]sessions.EnvironmentInputWork, error) { + id, devices, err := executionWorkCursor(after, connectedDevices) + if err != nil { + return nil, err + } + rows, err := s.units.Queries().ListEnvironmentInputWork(ctx, sqlc.ListEnvironmentInputWorkParams{AfterID: id, ConnectedDevices: devices}) + if err != nil { + return nil, err + } + work := make([]sessions.EnvironmentInputWork, 0, len(rows)) + for _, row := range rows { + work = append(work, sessions.EnvironmentInputWork{TenantID: uuid.UUID(row.TenantID.Bytes).String(), SessionID: uuid.UUID(row.SessionID.Bytes).String(), ReservationID: uuid.UUID(row.ID.Bytes).String()}) + } + return work, nil +} diff --git a/services/core/internal/persistence/postgres/sessionpg/inputs_test.go b/services/core/internal/persistence/postgres/sessionpg/inputs_test.go new file mode 100644 index 000000000..e2f9e1b2f --- /dev/null +++ b/services/core/internal/persistence/postgres/sessionpg/inputs_test.go @@ -0,0 +1,521 @@ +package sessionpg + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "reflect" + "strings" + "sync" + "testing" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgtype" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" +) + +var messagePayload = json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"hello"}]}]}`) + +var cancelInput = sessions.Input{Kind: "cancel", Payload: json.RawMessage(`{}`)} + +func messageInput(text string) sessions.Input { + payload, _ := json.Marshal(map[string]string{"text": text}) + return sessions.Input{Kind: "message", Payload: payload} +} + +// newSession stores a fresh tenant's idle Session without an Environment and +// returns their IDs. +func newSession(t *testing.T, pool *pgxpool.Pool) (pgtype.UUID, pgtype.UUID) { + t.Helper() + tenant, session := uuid.New(), uuid.New() + exec(t, pool, `INSERT INTO sessions(id, tenant_id, engine, idempotency_key, request_hash) VALUES ($1, $2, 'codex', 'key', 'hash')`, session, tenant) + return pgID(tenant), pgID(session) +} + +// submit admits one input through the Session service. +func submit(ctx context.Context, service *sessions.Service, tenant, session pgtype.UUID, key string, input sessions.Input) (sessions.InputReceipt, error) { + receipts, err := service.SubmitInputs(ctx, text(tenant), text(session), key, []sessions.Input{input}) + if err != nil { + return sessions.InputReceipt{}, err + } + return receipts[0], nil +} + +func submitMessage(t *testing.T, service *sessions.Service, tenant, session pgtype.UUID, key string) sessions.InputReceipt { + t.Helper() + receipt, err := submit(t.Context(), service, tenant, session, key, sessions.Input{Kind: "message", Payload: messagePayload}) + if err != nil { + t.Fatal(err) + } + return receipt +} + +func move(t *testing.T, pool *pgxpool.Pool, tenant, session pgtype.UUID, turn, from, to string) sessions.Turn { + t.Helper() + moved, err := transition(t.Context(), pool, tenant, session, turn, sessions.TurnTransition{ExpectedStatus: from, Status: to}) + if err != nil { + t.Fatal(err) + } + return moved +} + +// Concurrent messages from two pools all join one Turn, and concurrent retries +// of one key commit one input. +func TestConcurrentInputsUseOneTurnAndOneRetryReceipt(t *testing.T) { + pool := pgtest.Open(t) + store, service := stagingService(t, pool) + _, other := stagingService(t, pgtest.Open(t)) + tenant, session := newSession(t, pool) + const count = 8 + for _, repeatedKey := range []bool{true, false} { + t.Run(fmt.Sprintf("repeated-key-%v", repeatedKey), func(t *testing.T) { + var wg sync.WaitGroup + receipts := make(chan sessions.InputReceipt, count) + errs := make(chan error, count) + for i := range count { + wg.Go(func() { + key := "same" + if !repeatedKey { + key = fmt.Sprintf("distinct-%d", i) + } + admission := service + if i%2 == 0 { + admission = other + } + r, err := submit(t.Context(), admission, tenant, session, key, sessions.Input{Kind: "message", Payload: messagePayload}) + receipts <- r + errs <- err + }) + } + wg.Wait() + close(receipts) + close(errs) + for err := range errs { + if err != nil { + t.Fatal(err) + } + } + turns, sequences := map[string]bool{}, map[int64]bool{} + newReceipts := 0 + for r := range receipts { + turns[r.TurnID], sequences[r.Sequence] = true, true + if !r.Replayed { + newReceipts++ + } + } + want := count + if repeatedKey { + want = 1 + } + if len(turns) != 1 || turns[""] || len(sequences) != want || newReceipts != want { + t.Fatalf("turns=%v sequences=%v new=%d", turns, sequences, newReceipts) + } + }) + } + first := submitMessage(t, service, tenant, session, "same") + inputs, err := store.ListTurnInputs(t.Context(), text(tenant), text(session), first.TurnID, 0, 100) + if err != nil || len(inputs) != count+1 { + t.Fatalf("lost or duplicated inputs: %d, %v", len(inputs), err) + } +} + +// An active Turn takes steering input, retries replay their receipts across a +// restart, an idle Session starts a new Turn, and admission changes no stored +// Session field besides its event sequence. +func TestTurnInputRetriesAndRestart(t *testing.T) { + pool := pgtest.Open(t) + _, service := stagingService(t, pool) + tenant, session := newSession(t, pool) + ctx := t.Context() + stored := func(pool *pgxpool.Pool) (row string) { + if err := pool.QueryRow(ctx, `SELECT (to_jsonb(s) - 'event_sequence')::text FROM sessions s WHERE id = $1`, session).Scan(&row); err != nil { + t.Fatal(err) + } + return row + } + created := stored(pool) + first := submitMessage(t, service, tenant, session, "first") + move(t, pool, tenant, session, first.TurnID, sessions.TurnQueued, sessions.TurnInProgress) + steer := submitMessage(t, service, tenant, session, "steer") + if steer.TurnID != first.TurnID || steer.Sequence <= first.Sequence { + t.Fatalf("active message did not steer: %+v", steer) + } + reordered := json.RawMessage(`{ "input": [{"content":[{"text":"hello","type":"input_text"}],"role":"user"}] }`) + retry, err := submit(ctx, service, tenant, session, "first", sessions.Input{Kind: "message", Payload: reordered}) + if err != nil || !retry.Replayed || retry.Sequence != first.Sequence || retry.TurnID != first.TurnID { + t.Fatalf("equivalent retry = %+v, %v", retry, err) + } + if _, err := submit(ctx, service, tenant, session, "first", messageInput("changed")); !errors.Is(err, sessions.ErrIdempotencyConflict) { + t.Fatalf("changed payload accepted: %v", err) + } + if _, err := submit(ctx, service, tenant, session, "first", cancelInput); !errors.Is(err, sessions.ErrIdempotencyConflict) { + t.Fatalf("changed input kind accepted: %v", err) + } + completed := move(t, pool, tenant, session, first.TurnID, sessions.TurnInProgress, sessions.TurnCompleted) + next := submitMessage(t, service, tenant, session, "next") + if next.TurnID == first.TurnID { + t.Fatal("idle message did not start a new Turn") + } + pool.Close() + reopened := pgtest.Open(t) + recovered, restarted := stagingService(t, reopened) + retry = submitMessage(t, restarted, tenant, session, "first") + if !retry.Replayed || retry.TurnID != first.TurnID || retry.Sequence != first.Sequence { + t.Fatalf("restart retry changed target: %+v", retry) + } + got, err := recovered.GetTurn(ctx, text(tenant), text(session), first.TurnID) + if err != nil || !reflect.DeepEqual(got, completed) { + t.Fatalf("restart turn: %+v, %v", got, err) + } + var all []sessions.TurnInput + var cursor int64 + for { + page, err := recovered.ListTurnInputs(ctx, text(tenant), text(session), first.TurnID, cursor, 1) + if err != nil { + t.Fatal(err) + } + if len(page) == 0 { + break + } + if page[0].Sequence <= cursor || page[0].CreatedAt.IsZero() { + t.Fatalf("invalid ordered input: %+v", page) + } + cursor = page[0].Sequence + all = append(all, page...) + } + if len(all) != 2 || all[0].Sequence != first.Sequence || all[1].Sequence != steer.Sequence { + t.Fatalf("recovered inputs = %+v", all) + } + snapshot, err := recovered.GetSession(ctx, text(tenant), text(session)) + if err != nil || snapshot.LastTurn == nil || snapshot.LastTurn.ID != next.TurnID || snapshot.LastTurn.Status != sessions.TurnQueued { + t.Fatal("latest Session activity did not survive restart", err) + } + if stored(reopened) != created { + t.Fatal("input admission mutated the stored Session") + } +} + +// Input admission and Turn reads and writes resolve only the tenant's own +// Session, and a Turn only within its Session. +func TestTurnOperationsAreTenantAndSessionScoped(t *testing.T) { + pool := pgtest.Open(t) + store, service := stagingService(t, pool) + tenant, session := newSession(t, pool) + otherTenant, other := newSession(t, pool) + ctx := t.Context() + first := submitMessage(t, service, tenant, session, "input") + for _, scope := range []struct{ tenant, session pgtype.UUID }{{otherTenant, session}, {tenant, other}, {tenant, pgID(uuid.New())}} { + for name, call := range map[string]func() error{ + "submit": func() error { + _, err := submit(ctx, service, scope.tenant, scope.session, "input", sessions.Input{Kind: "message", Payload: messagePayload}) + return err + }, + "cancel": func() error { + _, err := submit(ctx, service, scope.tenant, scope.session, "cancel", cancelInput) + return err + }, + "read": func() error { + _, err := store.GetTurn(ctx, text(scope.tenant), text(scope.session), first.TurnID) + return err + }, + "inputs": func() error { + _, err := store.ListTurnInputs(ctx, text(scope.tenant), text(scope.session), first.TurnID, 0, 10) + return err + }, + "transition": func() error { + _, err := transition(ctx, pool, scope.tenant, scope.session, first.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnFailed}) + return err + }, + } { + if err := call(); !errors.Is(err, sessions.ErrNotFound) { + t.Fatalf("%s escaped scope: %v", name, err) + } + } + } + // Turn IDs cannot be used with another valid Session in the same tenant either. + second := pgID(uuid.New()) + exec(t, pool, `INSERT INTO sessions(id, tenant_id, engine, idempotency_key, request_hash) VALUES ($1, $2, 'codex', 'second', 'hash')`, second, tenant) + if _, err := store.GetTurn(ctx, text(tenant), text(second), first.TurnID); !errors.Is(err, sessions.ErrNotFound) { + t.Fatalf("cross-session turn read: %v", err) + } + if _, err := transition(ctx, pool, tenant, second, first.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnFailed}); !errors.Is(err, sessions.ErrNotFound) { + t.Fatalf("cross-session turn write: %v", err) + } + otherInput := submitMessage(t, service, otherTenant, other, "input") + if otherInput.TurnID == first.TurnID { + t.Fatal("retry identity leaked across Sessions") + } +} + +// Batches from two pools keep their inputs adjacent and in order, and one +// key's concurrent retries commit one batch. +func TestInputBatchesAreOrderedAndIdempotentAcrossConnections(t *testing.T) { + pool := pgtest.Open(t) + store, service := stagingService(t, pool) + _, other := stagingService(t, pgtest.Open(t)) + tenant, session := newSession(t, pool) + ctx := t.Context() + const count = 8 + batch := []sessions.Input{messageInput("first"), messageInput("second")} + for _, repeated := range []bool{true, false} { + var wg sync.WaitGroup + receipts := make(chan []sessions.InputReceipt, count) + errs := make(chan error, count) + for i := range count { + wg.Go(func() { + admission := service + if i%2 == 0 { + admission = other + } + key := "same-batch" + if !repeated { + key = fmt.Sprintf("batch-%d", i) + } + got, err := admission.SubmitInputs(ctx, text(tenant), text(session), key, batch) + receipts <- got + errs <- err + }) + } + wg.Wait() + close(receipts) + close(errs) + for err := range errs { + if err != nil { + t.Fatal(err) + } + } + newBatches := 0 + for got := range receipts { + if len(got) != 2 || got[0].TurnID == "" || got[0].TurnID != got[1].TurnID || got[0].Sequence >= got[1].Sequence || got[0].Replayed != got[1].Replayed { + t.Fatalf("invalid batch receipt: %+v", got) + } + if !got[0].Replayed { + newBatches++ + } + } + want := count + if repeated { + want = 1 + } + if newBatches != want { + t.Fatalf("accepted %d batches, want %d", newBatches, want) + } + } + got, err := service.SubmitInputs(ctx, text(tenant), text(session), "same-batch", batch) + if err != nil { + t.Fatal(err) + } + inputs, err := store.ListTurnInputs(ctx, text(tenant), text(session), got[0].TurnID, 0, 100) + if err != nil || len(inputs) != 2*(count+1) { + t.Fatalf("inputs=%d err=%v", len(inputs), err) + } + for i, input := range inputs { + var payload map[string]string + if err := json.Unmarshal(input.Payload, &payload); err != nil { + t.Fatal(err) + } + want := "first" + if i%2 == 1 { + want = "second" + } + if payload["text"] != want { + t.Fatalf("batch interleaved at %d: %v", i, payload) + } + } +} + +// A batch retry compares the whole normalized request and replays the +// cancellation targets it committed, after a restart too. +func TestBatchRetriesCompareTheWholeRequestAndRetainTargets(t *testing.T) { + pool := pgtest.Open(t) + store, service := stagingService(t, pool) + tenant, session := newSession(t, pool) + ctx := t.Context() + batch := []sessions.Input{cancelInput, messageInput("one"), cancelInput, messageInput("two")} + first, err := service.SubmitInputs(ctx, text(tenant), text(session), "mixed", batch) + if err != nil { + t.Fatal(err) + } + if first[0].TurnID != "" || first[1].TurnID == "" || first[1].TurnID != first[2].TurnID || first[1].TurnID == first[3].TurnID { + t.Fatalf("cancellation targets: %+v", first) + } + cancelled, err := store.GetTurn(ctx, text(tenant), text(session), first[1].TurnID) + if err != nil || cancelled.Status != sessions.TurnCancelled { + t.Fatalf("cancelled turn=%+v err=%v", cancelled, err) + } + move(t, pool, tenant, session, first[3].TurnID, sessions.TurnQueued, sessions.TurnInProgress) + move(t, pool, tenant, session, first[3].TurnID, sessions.TurnInProgress, sessions.TurnCompleted) + next := submitMessage(t, service, tenant, session, "next") + for _, changed := range [][]sessions.Input{batch[:3], append(append([]sessions.Input{}, batch...), cancelInput), {batch[0], batch[3], batch[2], batch[1]}} { + if _, err := service.SubmitInputs(ctx, text(tenant), text(session), "mixed", changed); !errors.Is(err, sessions.ErrIdempotencyConflict) { + t.Fatalf("changed batch accepted: %v", err) + } + } + batch[1].Payload = json.RawMessage(`{ "text" : "one" }`) + pool.Close() + restartedStore, restarted := stagingService(t, pgtest.Open(t)) + retry, err := restarted.SubmitInputs(ctx, text(tenant), text(session), "mixed", batch) + for i := range first { + first[i].Replayed = true + } + if err != nil || !reflect.DeepEqual(retry, first) { + t.Fatalf("restart changed receipts: %+v, %v", retry, err) + } + current, err := restartedStore.GetTurn(ctx, text(tenant), text(session), next.TurnID) + if err != nil || current.Status != sessions.TurnQueued || !current.CancelRequestedAt.IsZero() { + t.Fatalf("retry cancelled later work: %+v, %v", current, err) + } + if _, err := restarted.SubmitInputs(ctx, uuid.NewString(), text(session), "mixed", batch); !errors.Is(err, sessions.ErrNotFound) { + t.Fatalf("batch retry escaped tenant: %v", err) + } +} + +// A batch that fails part way rolls back its earlier cancellation and inputs, +// and its retry then commits it whole. +func TestFailedBatchRollsBackEarlierCancellationAndInputs(t *testing.T) { + pool := pgtest.Open(t) + store, service := stagingService(t, pool) + tenant, session := newSession(t, pool) + ctx := t.Context() + initial := submitMessage(t, service, tenant, session, "initial") + move(t, pool, tenant, session, initial.TurnID, sessions.TurnQueued, sessions.TurnInProgress) + key := uuid.NewString() + constraint := "batch_failure_" + strings.ReplaceAll(key, "-", "") + // Inject a storage error on the second insert, after cancellation has run. + exec(t, pool, "ALTER TABLE turn_inputs ADD CONSTRAINT "+constraint+" CHECK (idempotency_key <> '"+key+"' OR batch_position = 0)") + t.Cleanup(func() { + _, _ = pool.Exec(context.Background(), "ALTER TABLE turn_inputs DROP CONSTRAINT IF EXISTS "+constraint) + }) + batch := []sessions.Input{cancelInput, messageInput("after cancel")} + if got, err := service.SubmitInputs(ctx, text(tenant), text(session), key, batch); err == nil || got != nil { + t.Fatalf("partial batch succeeded: %+v %v", got, err) + } + turn, err := store.GetTurn(ctx, text(tenant), text(session), initial.TurnID) + if err != nil || turn.Status != sessions.TurnInProgress || !turn.CancelRequestedAt.IsZero() { + t.Fatalf("cancellation escaped rollback: %+v %v", turn, err) + } + inputs, err := store.ListTurnInputs(ctx, text(tenant), text(session), initial.TurnID, 0, 100) + if err != nil || len(inputs) != 1 { + t.Fatalf("partial inputs survived: %v %v", inputs, err) + } + exec(t, pool, "ALTER TABLE turn_inputs DROP CONSTRAINT "+constraint) + got, err := service.SubmitInputs(ctx, text(tenant), text(session), key, batch) + if err != nil || len(got) != 2 || got[0].Replayed || got[1].Replayed { + t.Fatalf("retry after rollback: %+v %v", got, err) + } +} + +// A cancel request targets the Turn active when it was admitted, and its retry +// never retargets a later Turn; cancelling queued work stops it at once. +func TestCancellationStaysBoundToItsOriginalTurn(t *testing.T) { + pool := pgtest.Open(t) + store, service := stagingService(t, pool) + tenant, session := newSession(t, pool) + ctx := t.Context() + idle, err := submit(ctx, service, tenant, session, "idle-cancel", cancelInput) + if err != nil || idle.TurnID != "" || idle.Sequence == 0 { + t.Fatalf("idle cancellation: %+v, %v", idle, err) + } + first := submitMessage(t, service, tenant, session, "first") + started := move(t, pool, tenant, session, first.TurnID, sessions.TurnQueued, sessions.TurnInProgress) + if started.StartedAt.IsZero() || !started.CompletedAt.IsZero() { + t.Fatalf("started timestamps: %+v", started) + } + cancel, err := submit(ctx, service, tenant, session, "cancel", cancelInput) + if err != nil || cancel.TurnID != first.TurnID { + t.Fatalf("cancellation target: %+v, %v", cancel, err) + } + pending, err := store.GetTurn(ctx, text(tenant), text(session), first.TurnID) + if err != nil || pending.Status != sessions.TurnInProgress || pending.CancelRequestedAt.IsZero() || !pending.CompletedAt.IsZero() { + t.Fatalf("running cancellation falsely completed: %+v, %v", pending, err) + } + move(t, pool, tenant, session, first.TurnID, sessions.TurnInProgress, sessions.TurnCancelled) + next := submitMessage(t, service, tenant, session, "next") + for key, original := range map[string]sessions.InputReceipt{"cancel": cancel, "idle-cancel": idle} { + retry, err := submit(ctx, service, tenant, session, key, cancelInput) + if err != nil || !retry.Replayed || retry.TurnID != original.TurnID || retry.Sequence != original.Sequence { + t.Fatalf("cancellation retargeted: %+v, %v", retry, err) + } + } + queued, err := store.GetTurn(ctx, text(tenant), text(session), next.TurnID) + if err != nil || queued.Status != sessions.TurnQueued || !queued.CancelRequestedAt.IsZero() { + t.Fatalf("old cancellation affected later Turn: %+v, %v", queued, err) + } + if _, err := submit(ctx, service, tenant, session, "cancel-queued", cancelInput); err != nil { + t.Fatal(err) + } + stopped, err := store.GetTurn(ctx, text(tenant), text(session), next.TurnID) + if err != nil || stopped.Status != sessions.TurnCancelled || stopped.CompletedAt.IsZero() || !stopped.StartedAt.IsZero() { + t.Fatalf("queued work did not stop: %+v, %v", stopped, err) + } + if _, err := transition(ctx, pool, tenant, session, next.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress}); !errors.Is(err, sessions.ErrTurnConflict) { + t.Fatalf("cancelled queued work was started: %v", err) + } +} + +// A waiting Turn takes steering input and keeps its start time when it +// resumes; a cancel request keeps it from resuming. +func TestWaitingTurnRetainsInputsAndStartTime(t *testing.T) { + pool := pgtest.Open(t) + _, service := stagingService(t, pool) + tenant, session := newSession(t, pool) + first := submitMessage(t, service, tenant, session, "first") + started := move(t, pool, tenant, session, first.TurnID, sessions.TurnQueued, sessions.TurnInProgress) + move(t, pool, tenant, session, first.TurnID, sessions.TurnInProgress, sessions.TurnWaiting) + steer := submitMessage(t, service, tenant, session, "steer") + if steer.TurnID != first.TurnID { + t.Fatal("waiting input started another Turn") + } + resumed := move(t, pool, tenant, session, first.TurnID, sessions.TurnWaiting, sessions.TurnInProgress) + if !started.StartedAt.Equal(resumed.StartedAt) { + t.Fatal("resume reset the start time") + } + move(t, pool, tenant, session, first.TurnID, sessions.TurnInProgress, sessions.TurnWaiting) + if _, err := submit(t.Context(), service, tenant, session, "cancel", cancelInput); err != nil { + t.Fatal(err) + } + if _, err := transition(t.Context(), pool, tenant, session, first.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnWaiting, Status: sessions.TurnInProgress}); !errors.Is(err, sessions.ErrTurnConflict) { + t.Fatalf("cancelling Turn resumed: %v", err) + } + move(t, pool, tenant, session, first.TurnID, sessions.TurnWaiting, sessions.TurnCancelled) +} + +// Rejected input and rejected transitions persist nothing. +func TestTurnInputValidationHasNoSideEffects(t *testing.T) { + pool := pgtest.Open(t) + store, service := stagingService(t, pool) + tenant, session := newSession(t, pool) + ctx := t.Context() + for _, raw := range []json.RawMessage{nil, json.RawMessage(`[]`), json.RawMessage(`null`), json.RawMessage(`{} {}`), json.RawMessage(`{"text":"` + string(make([]byte, 512*1024)) + `"}`)} { + if _, err := submit(ctx, service, tenant, session, "first", sessions.Input{Kind: "message", Payload: raw}); !errors.Is(err, sessions.ErrInvalidInput) { + t.Fatalf("invalid input accepted: %v", err) + } + } + cancelled, stop := context.WithCancel(ctx) + stop() + if _, err := submit(cancelled, service, tenant, session, "first", sessions.Input{Kind: "message", Payload: messagePayload}); !errors.Is(err, context.Canceled) { + t.Fatalf("cancelled context: %v", err) + } + first := submitMessage(t, service, tenant, session, "first") + if first.Replayed { + t.Fatal("failed submission persisted a receipt") + } + for _, input := range []sessions.TurnTransition{ + {ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnCompleted}, + {ExpectedStatus: sessions.TurnInProgress, Status: sessions.TurnQueued}, + {ExpectedStatus: sessions.TurnCompleted, Status: sessions.TurnInProgress}, + {ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress, Outcome: json.RawMessage(`{"premature":true}`)}, + } { + if _, err := transition(ctx, pool, tenant, session, first.TurnID, input); !errors.Is(err, sessions.ErrInvalidInput) { + t.Fatalf("invalid transition accepted: %v", err) + } + } + got, err := store.GetTurn(ctx, text(tenant), text(session), first.TurnID) + if err != nil || got.Status != sessions.TurnQueued || !got.StartedAt.IsZero() { + t.Fatalf("invalid transition changed Turn: %+v, %v", got, err) + } +} diff --git a/services/core/internal/persistence/postgres/sessionpg/turn_reads.go b/services/core/internal/persistence/postgres/sessionpg/turn_reads.go index 9848a68d9..28f96b606 100644 --- a/services/core/internal/persistence/postgres/sessionpg/turn_reads.go +++ b/services/core/internal/persistence/postgres/sessionpg/turn_reads.go @@ -72,7 +72,7 @@ func (s *Store) ListTurns(ctx context.Context, tenant, session, cursor string, l } func (s *Store) ListExecutionWork(ctx context.Context, after string, statuses, connectedDevices []string) ([]sessions.ExecutionWork, error) { - id, devices, err := ExecutionWorkCursor(after, connectedDevices) + id, devices, err := executionWorkCursor(after, connectedDevices) if err != nil { return nil, err } @@ -87,10 +87,10 @@ func (s *Store) ListExecutionWork(ctx context.Context, after string, statuses, c return work, nil } -// ExecutionWorkCursor parses the cursor and connected devices of an execution +// executionWorkCursor parses the cursor and connected devices of an execution // work scan; an empty cursor starts at the first ID. A malformed ID is // sessions.ErrInvalidInput. -func ExecutionWorkCursor(after string, connectedDevices []string) (pgtype.UUID, []pgtype.UUID, error) { +func executionWorkCursor(after string, connectedDevices []string) (pgtype.UUID, []pgtype.UUID, error) { id := pgtype.UUID{Valid: true} var err error if after != "" { diff --git a/services/core/internal/sandbox/providers/configuration_flow_test.go b/services/core/internal/sandbox/providers/configuration_flow_test.go index f400a64af..f043ab26d 100644 --- a/services/core/internal/sandbox/providers/configuration_flow_test.go +++ b/services/core/internal/sandbox/providers/configuration_flow_test.go @@ -149,7 +149,7 @@ func TestAdditionalConfigurationProviderUsesCommonAPIAndStore(t *testing.T) { Execution: &api.Execution{ ExecutorURL: "wss://core.example/api/v1/agent-daemon/ws", SessionAdmission: s, - InputAdmission: s, + InputAdmission: struct{ api.InputAdmission }{}, SessionArchive: struct{ api.SessionArchive }{}, Workspaces: struct{ api.EnvironmentWorkspaces }{}, }, diff --git a/services/core/internal/sessions/artifacts_test.go b/services/core/internal/sessions/artifacts_test.go index a31e5b326..b7a45c855 100644 --- a/services/core/internal/sessions/artifacts_test.go +++ b/services/core/internal/sessions/artifacts_test.go @@ -192,6 +192,11 @@ func (s *fakeArtifactStorage) DeleteSessionArtifact(context.Context, string, str return nil } +func (s *fakeArtifactStorage) WithInputs(context.Context, string, string, func(context.Context, InputTx) error) error { + s.t.Fatal("unexpected call to WithInputs") + return nil +} + func (s *fakeArtifactStorage) CreateDevice(context.Context, string, DeviceRegistration) (ExecutionDevice, error) { s.t.Fatal("unexpected call to CreateDevice") return ExecutionDevice{}, nil diff --git a/services/core/internal/sessions/devices_test.go b/services/core/internal/sessions/devices_test.go index b8297e341..1976c4769 100644 --- a/services/core/internal/sessions/devices_test.go +++ b/services/core/internal/sessions/devices_test.go @@ -45,6 +45,8 @@ type fakeStorage struct { deletion *fakeTx updateMetadata func(encoded string) (Session, error) auditOperation func() error + // inputTx is the transaction WithInputs applies in. + inputTx *fakeInputTx } func (s *fakeStorage) record(name string, set bool, detail ...string) { diff --git a/services/core/internal/sessions/doc.go b/services/core/internal/sessions/doc.go index e86c58aab..069ef6765 100644 --- a/services/core/internal/sessions/doc.go +++ b/services/core/internal/sessions/doc.go @@ -7,12 +7,13 @@ // The Session writes that several operations share are procedures here, over // the transaction interfaces declared beside them: cancelling work, failing // and terminating an Environment, tracking input activity, the admission -// gates, moving a Turn's status, admitting a function result, reading a Turn's -// required actions, appending to a Turn's execution journal and projecting its -// observations and admitted inputs into Items, Turn usage and Subagents. Service runs the pooled -// Session use cases, such as staging a Turn's Artifacts, over Storage, and -// Reader declares the Session reads. ExecutionOperations runs the Session -// writes only the execution owner makes, such as moving and completing Turns, -// recording function calls and their application receipts and journaling a -// Turn's observations, over the lease-bound ExecutionStorage. +// gates, admitting an input batch, moving a Turn's status, admitting a function +// result, reading a Turn's required actions, appending to a Turn's execution +// journal and projecting its observations and admitted inputs into Items, Turn +// usage and Subagents. Service runs the pooled Session use cases, such as +// staging a Turn's Artifacts, over Storage, and Reader declares the Session +// reads. ExecutionOperations runs the Session writes only the execution owner +// makes, such as moving and completing Turns, recording function calls and +// their application receipts and journaling a Turn's observations, over the +// lease-bound ExecutionStorage. package sessions diff --git a/services/core/internal/sessions/environment.go b/services/core/internal/sessions/environment.go index 745ecc791..107aa894c 100644 --- a/services/core/internal/sessions/environment.go +++ b/services/core/internal/sessions/environment.go @@ -98,6 +98,8 @@ const ( type EnvironmentInputReservation struct { ID string SessionID string + // Key is the idempotency key the reservation's batch is admitted under. + Key string State string IsInitial bool Inputs []Input diff --git a/services/core/internal/sessions/environment_inputs.go b/services/core/internal/sessions/environment_inputs.go new file mode 100644 index 000000000..c0b02bee8 --- /dev/null +++ b/services/core/internal/sessions/environment_inputs.go @@ -0,0 +1,159 @@ +package sessions + +import ( + "context" + "errors" +) + +// checkEnvironmentAcceptsInput admits new input only while the Session's +// Environment is live: a failed hosted Environment is +// ErrHostedEnvironmentFailed, and any other failed or expired Environment +// ErrEnvironmentUnavailable. +func checkEnvironmentAcceptsInput(environment Environment) error { + if environment.Status == "failed" { + if kind, err := EnvironmentType(environment.Configuration); err == nil && kind == "openai_hosted" { + return ErrHostedEnvironmentFailed + } + } + if environment.Status == "failed" || environment.Status == "expired" { + return ErrEnvironmentUnavailable + } + return nil +} + +// joinsActiveTurn reports whether reserved Environment input joins the +// Session's active Turn at once instead of waiting for its own Turn: it does +// while the Turn has not started capturing its Artifacts, whose sealed native +// input takes nothing more. +func joinsActiveTurn(turn Turn, active bool) bool { + return active && !turn.ArtifactCaptureStarted +} + +// ReserveEnvironmentInput admits a message batch for a Session whose +// Environment runs it. Under the Session lock it joins the active Turn at once, +// unless that Turn captures its Artifacts, or otherwise reserves the batch +// until a deadline, for the execution owner to promote into a new Turn. A +// retry returns the reservation under its key, expired once its deadline has +// passed, or the receipts the batch was admitted with. The checks of +// SubmitInputs apply, and a failed or expired Environment rejects new input. +func (s *Service) ReserveEnvironmentInput(ctx context.Context, tenant, session, key string, inputs []Input) (EnvironmentInputReservation, error) { + if err := ValidateInputKey(key); err != nil { + return EnvironmentInputReservation{}, err + } + batch, encoded, err := ValidateMessageInputs(inputs) + if err != nil { + return EnvironmentInputReservation{}, err + } + var result EnvironmentInputReservation + err = s.storage.WithInputs(ctx, tenant, session, func(ctx context.Context, tx InputTx) error { + return TrackInputActivity(ctx, tx, func(ctx context.Context) error { + previous, used, err := tx.FindInputReservation(ctx, key, encoded) + if err != nil { + return err + } + replay, err := decideReplay(used, previous != nil) + if err != nil { + return err + } + if replay { + if result, err = expireReservation(ctx, tx, *previous); err != nil { + return err + } + if result.State == EnvironmentInputPending || result.State == EnvironmentInputAdmitted { + return tx.RecordInputAudit(ctx) + } + return nil + } + receipts, matches, err := tx.LoadInputBatch(ctx, key, encoded) + if err != nil { + return err + } + if replay, err = decideReplay(len(receipts) > 0, matches); err != nil { + return err + } + if replay { + // Earlier direct admission has receipts, but never had a reservation or deadline. + result = EnvironmentInputReservation{SessionID: session, State: EnvironmentInputAdmitted, Receipts: receipts} + return tx.RecordInputAudit(ctx) + } + if err := CheckFileWriteGate(ctx, tx); err != nil { + return err + } + environment, err := tx.LoadEnvironment(ctx) + if errors.Is(err, ErrNotFound) { + return ErrInvalidInput + } + if err != nil { + return err + } + if err := checkEnvironmentAcceptsInput(environment); err != nil { + return err + } + gate, err := tx.LoadInputGate(ctx, key, encoded) + if err != nil { + return err + } + if err := checkInputGate(gate); err != nil { + return err + } + turn, active, err := tx.LoadActiveTurn(ctx) + if err != nil { + return err + } + if joinsActiveTurn(turn, active) { + receipts, err := AdmitInputs(ctx, tx, key, batch) + if err != nil { + return err + } + result = EnvironmentInputReservation{SessionID: session, State: EnvironmentInputAdmitted, Receipts: receipts} + return tx.RecordInputAudit(ctx) + } + if result, err = tx.CreateInputReservation(ctx, key, encoded); err != nil { + return err + } + return tx.RecordInputAudit(ctx) + }) + }) + if err != nil { + return EnvironmentInputReservation{}, err + } + return result, nil +} + +// ExpireEnvironmentInput returns the Session's Environment input reservation, +// expired first when it is pending and its deadline has passed by the +// database clock. A malformed reservation ID is ErrInvalidInput, and a missing +// reservation ErrNotFound. +func (s *Service) ExpireEnvironmentInput(ctx context.Context, tenant, session, reservation string) (EnvironmentInputReservation, error) { + if !validID(reservation) { + return EnvironmentInputReservation{}, ErrInvalidInput + } + var result EnvironmentInputReservation + err := s.storage.WithInputs(ctx, tenant, session, func(ctx context.Context, tx InputTx) error { + return TrackInputActivity(ctx, tx, func(ctx context.Context) error { + current, err := tx.LoadInputReservation(ctx, reservation) + if err != nil { + return err + } + result, err = expireReservation(ctx, tx, current) + return err + }) + }) + if err != nil { + return EnvironmentInputReservation{}, err + } + return result, nil +} + +// expireReservation expires a pending reservation whose deadline has passed +// and returns it as the transaction has left it. A settled one is returned as +// it is: its outcome is a successful result, so settling never rolls back. +func expireReservation(ctx context.Context, tx InputTx, reservation EnvironmentInputReservation) (EnvironmentInputReservation, error) { + if reservation.State != EnvironmentInputPending { + return reservation, nil + } + if err := tx.ExpireInputReservation(ctx, reservation.ID); err != nil { + return EnvironmentInputReservation{}, err + } + return tx.LoadInputReservation(ctx, reservation.ID) +} diff --git a/services/core/internal/sessions/environment_inputs_test.go b/services/core/internal/sessions/environment_inputs_test.go new file mode 100644 index 000000000..f49fd2aaf --- /dev/null +++ b/services/core/internal/sessions/environment_inputs_test.go @@ -0,0 +1,165 @@ +package sessions + +import ( + "encoding/json" + "errors" + "reflect" + "testing" +) + +const testReservation = "6a1c9e2b-3d4f-4a5b-8c6d-7e8f9a0b1c2d" + +// reservations is a fake LoadInputReservation that reads states one call at +// a time, then the last again. +func reservations(states ...EnvironmentInputReservation) func() (EnvironmentInputReservation, error) { + return func() (EnvironmentInputReservation, error) { + state := states[0] + if len(states) > 1 { + states = states[1:] + } + return state, nil + } +} + +func reservationIn(state string) EnvironmentInputReservation { + return EnvironmentInputReservation{ID: testReservation, SessionID: testSession, Key: inputKey, State: state, Inputs: []Input{messageInput("hi")}} +} + +func TestCheckEnvironmentAcceptsInput(t *testing.T) { + hosted := json.RawMessage(`{"type":"openai_hosted"}`) + for _, test := range []struct { + environment Environment + want error + }{ + {selfHosted, nil}, + {Environment{Status: "failed", Configuration: hosted}, ErrHostedEnvironmentFailed}, + {Environment{Status: "failed", Configuration: selfHosted.Configuration}, ErrEnvironmentUnavailable}, + {Environment{Status: "expired", Configuration: hosted}, ErrEnvironmentUnavailable}, + } { + if err := checkEnvironmentAcceptsInput(test.environment); err != test.want { + t.Errorf("%s %s: %v", test.environment.Status, test.environment.Configuration, err) + } + } +} + +func TestJoinsActiveTurn(t *testing.T) { + capturing := turnWith(TurnInProgress) + capturing.ArtifactCaptureStarted = true + if joinsActiveTurn(Turn{}, false) || !joinsActiveTurn(turnWith(TurnInProgress), true) || joinsActiveTurn(capturing, true) { + t.Fatal("joins") + } +} + +func TestReserveEnvironmentInput(t *testing.T) { + const batch = `[{"kind":"message","payload":{"text":"hi"}}]` + find := "FindInputReservation request " + batch + reserve := func(t *testing.T, tx *fakeInputTx) (EnvironmentInputReservation, error) { + tx.loadEnvironmentInput = inputs() + return deviceService(t, &fakeStorage{t: t, inputTx: tx}).ReserveEnvironmentInput(t.Context(), testTenant, testSession, inputKey, []Input{messageInput("hi")}) + } + // fresh is a transaction where the key holds no batch yet. + fresh := func(t *testing.T) *fakeInputTx { + tx := newInputTx(t) + tx.findInputReservation = func() (*EnvironmentInputReservation, bool, error) { return nil, false, nil } + tx.loadInputBatch = func() ([]InputReceipt, bool, error) { return nil, true, nil } + tx.loadPendingFileWrite, tx.loadEnvironment = returns(false), returns(selfHosted) + return tx + } + checked := []string{"LoadEnvironmentInput", find, "LoadInputBatch request " + batch, "LoadPendingFileWrite", "LoadEnvironment"} + t.Run("a retry returns its reservation", func(t *testing.T) { + for _, test := range []struct { + name string + current EnvironmentInputReservation + audit []string + }{ + {"pending", reservationIn(EnvironmentInputPending), []string{"RecordInputAudit"}}, + {"expired at its deadline", reservationIn(EnvironmentInputExpired), nil}, + } { + t.Run(test.name, func(t *testing.T) { + previous := reservationIn(EnvironmentInputPending) + tx := newInputTx(t) + tx.findInputReservation = func() (*EnvironmentInputReservation, bool, error) { return &previous, true, nil } + tx.expireInputReservation, tx.loadInputReservation, tx.recordInputAudit = done, returns(test.current), done + result, err := reserve(t, tx) + if err != nil || !reflect.DeepEqual(result, test.current) { + t.Fatalf("reservation %+v, %v", result, err) + } + calls := append([]string{"LoadEnvironmentInput", find, "ExpireInputReservation " + testReservation, "LoadInputReservation " + testReservation}, test.audit...) + assertCalls(t, tx.fakeTx, append(calls, "LoadEnvironmentInput")...) + }) + } + }) + t.Run("a retry of a direct admission returns its receipts", func(t *testing.T) { + previous := []InputReceipt{{Sequence: 3, TurnID: testTurn, Replayed: true}} + tx := newInputTx(t) + tx.findInputReservation = func() (*EnvironmentInputReservation, bool, error) { return nil, false, nil } + tx.loadInputBatch = func() ([]InputReceipt, bool, error) { return previous, true, nil } + tx.recordInputAudit = done + result, err := reserve(t, tx) + if err != nil || !reflect.DeepEqual(result, EnvironmentInputReservation{SessionID: testSession, State: EnvironmentInputAdmitted, Receipts: previous}) { + t.Fatalf("reservation %+v, %v", result, err) + } + assertCalls(t, tx.fakeTx, "LoadEnvironmentInput", find, "LoadInputBatch request "+batch, "RecordInputAudit", "LoadEnvironmentInput") + }) + t.Run("an Environment that no longer runs rejects input", func(t *testing.T) { + for _, test := range []struct { + name string + load func() (Environment, error) + want error + }{ + {"failed hosted", returns(Environment{Status: "failed", Configuration: json.RawMessage(`{"type":"openai_hosted"}`)}), ErrHostedEnvironmentFailed}, + {"expired", returns(Environment{Status: "expired", Configuration: selfHosted.Configuration}), ErrEnvironmentUnavailable}, + {"missing", func() (Environment, error) { return Environment{}, ErrNotFound }, ErrInvalidInput}, + } { + t.Run(test.name, func(t *testing.T) { + tx := fresh(t) + tx.loadEnvironment = test.load + if _, err := reserve(t, tx); err != test.want { + t.Fatal(err) + } + assertCalls(t, tx.fakeTx, checked...) + }) + } + }) + t.Run("input joins the active Turn", func(t *testing.T) { + running := turnWith(TurnInProgress) + tx := fresh(t) + tx.loadInputGate, tx.loadActiveTurn, tx.createTurnInput, tx.loadInputSource, tx.recordInputAudit = returns(InputGate{Matches: true}), activeTurn(&running), sequences(4), unprojected, done + result, err := reserve(t, tx) + if err != nil || !reflect.DeepEqual(result, EnvironmentInputReservation{SessionID: testSession, State: EnvironmentInputAdmitted, Receipts: []InputReceipt{{Sequence: 4, TurnID: testTurn}}}) { + t.Fatalf("reservation %+v, %v", result, err) + } + assertCalls(t, tx.fakeTx, append(checked, "LoadInputGate request "+batch, "LoadActiveTurn", + "LoadActiveTurn", "CreateTurnInput "+testTurn+" request 0 message", "LoadInputSource 4", "RecordInputAudit", "LoadEnvironmentInput")...) + }) + t.Run("input waits in a reservation on an idle Session", func(t *testing.T) { + tx := fresh(t) + tx.loadInputGate, tx.loadActiveTurn = returns(InputGate{Matches: true}), activeTurn(nil) + tx.createInputReservation, tx.recordInputAudit = returns(reservationIn(EnvironmentInputPending)), done + result, err := reserve(t, tx) + if err != nil || !reflect.DeepEqual(result, reservationIn(EnvironmentInputPending)) { + t.Fatalf("reservation %+v, %v", result, err) + } + assertCalls(t, tx.fakeTx, append(checked, "LoadInputGate request "+batch, "LoadActiveTurn", "CreateInputReservation request "+batch, "RecordInputAudit", "LoadEnvironmentInput")...) + }) + t.Run("a cancellation touches no storage", func(t *testing.T) { + if _, err := deviceService(t, &fakeStorage{t: t}).ReserveEnvironmentInput(t.Context(), testTenant, testSession, inputKey, []Input{cancelInput}); !errors.Is(err, ErrInvalidInput) { + t.Fatal(err) + } + }) +} + +func TestExpireEnvironmentInput(t *testing.T) { + tx := newInputTx(t) + tx.loadEnvironmentInput, tx.expireInputReservation = inputs(), done + tx.loadInputReservation = reservations(reservationIn(EnvironmentInputPending), reservationIn(EnvironmentInputExpired)) + service := deviceService(t, &fakeStorage{t: t, inputTx: tx}) + result, err := service.ExpireEnvironmentInput(t.Context(), testTenant, testSession, testReservation) + if err != nil || result.State != EnvironmentInputExpired { + t.Fatalf("reservation %+v, %v", result, err) + } + assertCalls(t, tx.fakeTx, "LoadEnvironmentInput", "LoadInputReservation "+testReservation, "ExpireInputReservation "+testReservation, "LoadInputReservation "+testReservation, "LoadEnvironmentInput") + if _, err := service.ExpireEnvironmentInput(t.Context(), testTenant, testSession, "malformed"); !errors.Is(err, ErrInvalidInput) { + t.Fatal(err) + } +} diff --git a/services/core/internal/sessions/execution.go b/services/core/internal/sessions/execution.go index db5eaaf13..c07c25573 100644 --- a/services/core/internal/sessions/execution.go +++ b/services/core/internal/sessions/execution.go @@ -20,5 +20,6 @@ func NewExecutionOperations(storage ExecutionStorage) (*ExecutionOperations, err type ExecutionStorage interface { EnvironmentExecution FunctionExecution + InputExecution TurnExecution } diff --git a/services/core/internal/sessions/execution_functions_test.go b/services/core/internal/sessions/execution_functions_test.go index cc8ca606a..03cdbb09e 100644 --- a/services/core/internal/sessions/execution_functions_test.go +++ b/services/core/internal/sessions/execution_functions_test.go @@ -15,9 +15,11 @@ import ( // their call on tx, then run apply with tx and locked; without tx they fail // the test. type fakeExecutionStorage struct { - t *testing.T - withFunctionTurn func(ctx context.Context, tenant, session, turn string, apply func(context.Context, FunctionTx, Turn) error) error - withTurns func(ctx context.Context, tenant, session string, apply func(context.Context, TurnTx) error) error + t *testing.T + withFunctionTurn func(ctx context.Context, tenant, session, turn string, apply func(context.Context, FunctionTx, Turn) error) error + withTurns func(ctx context.Context, tenant, session string, apply func(context.Context, TurnTx) error) error + withInputs func(ctx context.Context, tenant, session string, apply func(context.Context, InputTx) error) error + withDueInputReservations func(ctx context.Context, apply func(context.Context, InputTx, string) error) error tx *fakeTx locked LockedSession diff --git a/services/core/internal/sessions/execution_inputs.go b/services/core/internal/sessions/execution_inputs.go new file mode 100644 index 000000000..111195033 --- /dev/null +++ b/services/core/internal/sessions/execution_inputs.go @@ -0,0 +1,103 @@ +package sessions + +import "context" + +// InputExecution is the lease-bound storage of Environment input settlement. +type InputExecution interface { + // WithInputs runs apply in one transaction on the execution lease, under + // the lock of the tenant's Session, and commits only when apply succeeds. + // A malformed tenant is ErrInvalidInput; a malformed, missing or deleted + // Session is ErrNotFound. + WithInputs(ctx context.Context, tenant, session string, apply func(context.Context, InputTx) error) error + // WithDueInputReservations runs apply in one transaction on the execution + // lease for each of at most 32 pending Environment input reservations + // whose deadline has passed, earliest deadline first, of live Sessions + // whose lock no other transaction holds, with that Session's transaction + // and the reservation's ID. It commits only when every apply succeeds. + WithDueInputReservations(ctx context.Context, apply func(context.Context, InputTx, string) error) error +} + +// PromoteEnvironmentInput admits a pending reservation's batch into a new Turn +// and claims that Turn as in progress for the execution owner's retained +// native preparation. The reservation settles as admitted with the batch's +// receipts. A reservation that already settled, or expires now because its +// deadline has passed, is returned as it is. The Session must be idle with no +// pending file write; otherwise it is ErrTurnConflict and the reservation +// stays pending. A malformed reservation ID is ErrInvalidInput, and a missing +// reservation ErrNotFound. +func (o *ExecutionOperations) PromoteEnvironmentInput(ctx context.Context, tenant, session, reservation string) (EnvironmentInputReservation, error) { + if !validID(reservation) { + return EnvironmentInputReservation{}, ErrInvalidInput + } + var result EnvironmentInputReservation + err := o.storage.WithInputs(ctx, tenant, session, func(ctx context.Context, tx InputTx) error { + return TrackInputActivity(ctx, tx, func(ctx context.Context) error { + current, err := tx.LoadInputReservation(ctx, reservation) + if err != nil { + return err + } + if result, err = expireReservation(ctx, tx, current); err != nil || result.State != EnvironmentInputPending { + return err + } + if err := CheckInputStart(ctx, tx); err != nil { + return err + } + if result.Receipts, err = AdmitInputs(ctx, tx, result.Key, result.Inputs); err != nil { + return err + } + settled, err := tx.AdmitInputReservation(ctx, reservation) + if err != nil { + return err + } + start := TurnTransition{ExpectedStatus: TurnQueued, Status: TurnInProgress} + if _, err := TransitionTurn(ctx, tx, result.Receipts[0].TurnID, start); err != nil { + return err + } + result.State, result.SettledAt = EnvironmentInputAdmitted, &settled + return nil + }) + }) + if err != nil { + return EnvironmentInputReservation{}, err + } + return result, nil +} + +// FailEnvironmentInput settles a pending reservation as failed before +// admission with code, model_provider_required or runtime_preparation_failed, +// once execution confirmed the failure. A reservation that already settled, +// including one cancelled meanwhile, keeps its outcome, and one whose deadline +// has passed expires instead. Any other code and a malformed reservation ID +// are ErrInvalidInput. +func (o *ExecutionOperations) FailEnvironmentInput(ctx context.Context, tenant, session, reservation, code string) error { + if (code != "model_provider_required" && code != "runtime_preparation_failed") || !validID(reservation) { + return ErrInvalidInput + } + return o.storage.WithInputs(ctx, tenant, session, func(ctx context.Context, tx InputTx) error { + return TrackInputActivity(ctx, tx, func(ctx context.Context) error { + if err := tx.ExpireInputReservation(ctx, reservation); err != nil { + return err + } + return tx.FailInputReservation(ctx, reservation, code) + }) + }) +} + +// ExpireEnvironmentInputs expires one bounded batch of pending reservations +// whose deadline has passed, reports the input activity each changes, creates +// no Turn history and returns how many it expired. It is cross-Session +// maintenance only the execution owner runs. +func (o *ExecutionOperations) ExpireEnvironmentInputs(ctx context.Context) (int64, error) { + var expired int64 + err := o.storage.WithDueInputReservations(ctx, func(ctx context.Context, tx InputTx, reservation string) error { + if err := TrackInputActivity(ctx, tx, func(ctx context.Context) error { return tx.ExpireInputReservation(ctx, reservation) }); err != nil { + return err + } + expired++ + return nil + }) + if err != nil { + return 0, err + } + return expired, nil +} diff --git a/services/core/internal/sessions/execution_inputs_test.go b/services/core/internal/sessions/execution_inputs_test.go new file mode 100644 index 000000000..64a10a1bd --- /dev/null +++ b/services/core/internal/sessions/execution_inputs_test.go @@ -0,0 +1,133 @@ +package sessions + +import ( + "context" + "encoding/json" + "errors" + "reflect" + "testing" + "time" +) + +func (f *fakeExecutionStorage) WithInputs(ctx context.Context, tenant, session string, apply func(context.Context, InputTx) error) error { + f.t.Helper() + if f.withInputs == nil { + f.t.Fatal("unexpected call to WithInputs") + } + return f.withInputs(ctx, tenant, session, apply) +} + +func (f *fakeExecutionStorage) WithDueInputReservations(ctx context.Context, apply func(context.Context, InputTx, string) error) error { + f.t.Helper() + if f.withDueInputReservations == nil { + f.t.Fatal("unexpected call to WithDueInputReservations") + } + return f.withDueInputReservations(ctx, apply) +} + +// inputOperations serves the tenant's Session through tx. +func inputOperations(t *testing.T, tx *fakeInputTx) *ExecutionOperations { + t.Helper() + tx.loadEnvironmentInput = inputs() + operations, err := NewExecutionOperations(&fakeExecutionStorage{t: t, withInputs: func(ctx context.Context, tenant, session string, apply func(context.Context, InputTx) error) error { + if tenant != testTenant || session != testSession { + t.Fatalf("Session %s %s", tenant, session) + } + return apply(ctx, tx) + }}) + if err != nil { + t.Fatal(err) + } + return operations +} + +func TestPromoteEnvironmentInput(t *testing.T) { + expire := []string{"LoadEnvironmentInput", "LoadInputReservation " + testReservation, "ExpireInputReservation " + testReservation, "LoadInputReservation " + testReservation} + promote := func(t *testing.T, tx *fakeInputTx) (EnvironmentInputReservation, error) { + return inputOperations(t, tx).PromoteEnvironmentInput(t.Context(), testTenant, testSession, testReservation) + } + t.Run("a reservation past its deadline expires", func(t *testing.T) { + tx := newInputTx(t) + tx.loadInputReservation, tx.expireInputReservation = reservations(reservationIn(EnvironmentInputPending), reservationIn(EnvironmentInputExpired)), done + result, err := promote(t, tx) + if err != nil || !reflect.DeepEqual(result, reservationIn(EnvironmentInputExpired)) { + t.Fatalf("reservation %+v, %v", result, err) + } + assertCalls(t, tx.fakeTx, append(expire, "LoadEnvironmentInput")...) + }) + t.Run("a pending reservation starts its Turn", func(t *testing.T) { + settled := time.Unix(1700000002, 0) + tx := newInputTx(t) + tx.loadInputReservation, tx.expireInputReservation = returns(reservationIn(EnvironmentInputPending)), done + tx.loadPendingFileWrite, tx.loadActiveTurn = returns(false), activeTurn(nil) + tx.createTurn, tx.createTurnInput, tx.loadInputSource = returns(turnWith(TurnQueued)), sequences(5), unprojected + tx.loadUsage, tx.appendChanges = returns(json.RawMessage(`{}`)), collect(new([]SessionChange)) + tx.admitInputReservation, tx.loadTurn, tx.loadComputeSuspension = returns(settled), returns(turnWith(TurnQueued)), returns(false) + tx.applyTurnStatus = func(TurnStatusChange) (Turn, error) { return turnWith(TurnInProgress), nil } + result, err := promote(t, tx) + want := reservationIn(EnvironmentInputAdmitted) + want.SettledAt, want.Receipts = &settled, []InputReceipt{{Sequence: 5, TurnID: testTurn}} + if err != nil || !reflect.DeepEqual(result, want) { + t.Fatalf("reservation %+v, %v", result, err) + } + assertCalls(t, tx.fakeTx, append(expire, "LoadPendingFileWrite", "LoadActiveTurn", + "LoadActiveTurn", "CreateTurn", "AppendChanges agent.session.turn.created", "CreateTurnInput "+testTurn+" request 0 message", "LoadInputSource 5", "LoadUsage", "AppendChanges agent.session.in_progress", + "AdmitInputReservation "+testReservation, + "LoadTurn "+testTurn, "LoadComputeSuspension", "ApplyTurnStatus "+testTurn+" queued in_progress {}", "AppendChanges agent.session.turn.in_progress", + "LoadEnvironmentInput")...) + }) + t.Run("a busy Session keeps the reservation pending", func(t *testing.T) { + running := turnWith(TurnInProgress) + tx := newInputTx(t) + tx.loadInputReservation, tx.expireInputReservation = returns(reservationIn(EnvironmentInputPending)), done + tx.loadPendingFileWrite, tx.loadActiveTurn = returns(false), activeTurn(&running) + if _, err := promote(t, tx); !errors.Is(err, ErrTurnConflict) { + t.Fatal(err) + } + assertCalls(t, tx.fakeTx, append(expire, "LoadPendingFileWrite", "LoadActiveTurn")...) + }) + if _, err := unusedStorage(t).PromoteEnvironmentInput(t.Context(), testTenant, testSession, "malformed"); !errors.Is(err, ErrInvalidInput) { + t.Fatal(err) + } +} + +func TestFailEnvironmentInput(t *testing.T) { + tx := newInputTx(t) + tx.expireInputReservation, tx.failInputReservation = done, done + if err := inputOperations(t, tx).FailEnvironmentInput(t.Context(), testTenant, testSession, testReservation, "runtime_preparation_failed"); err != nil { + t.Fatal(err) + } + assertCalls(t, tx.fakeTx, "LoadEnvironmentInput", "ExpireInputReservation "+testReservation, "FailInputReservation "+testReservation+" runtime_preparation_failed", "LoadEnvironmentInput") + for reservation, code := range map[string]string{testReservation: "other", "malformed": "model_provider_required"} { + if err := unusedStorage(t).FailEnvironmentInput(t.Context(), testTenant, testSession, reservation, code); !errors.Is(err, ErrInvalidInput) { + t.Fatalf("%s %s: %v", reservation, code, err) + } + } +} + +func TestExpireEnvironmentInputs(t *testing.T) { + due := []string{"6a1c9e2b-0000-4a5b-8c6d-7e8f9a0b1c2d", "6a1c9e2b-1111-4a5b-8c6d-7e8f9a0b1c2d"} + expire := func(t *testing.T, tx *fakeInputTx, failure error) (int64, error) { + tx.loadEnvironmentInput, tx.expireInputReservation = inputs(), done + operations, err := NewExecutionOperations(&fakeExecutionStorage{t: t, withDueInputReservations: func(ctx context.Context, apply func(context.Context, InputTx, string) error) error { + for _, reservation := range due { + if err := apply(ctx, tx, reservation); err != nil { + return err + } + } + return failure + }}) + if err != nil { + t.Fatal(err) + } + return operations.ExpireEnvironmentInputs(t.Context()) + } + tx := newInputTx(t) + if expired, err := expire(t, tx, nil); err != nil || expired != 2 { + t.Fatalf("expired %d, %v", expired, err) + } + assertCalls(t, tx.fakeTx, "LoadEnvironmentInput", "ExpireInputReservation "+due[0], "LoadEnvironmentInput", "LoadEnvironmentInput", "ExpireInputReservation "+due[1], "LoadEnvironmentInput") + if expired, err := expire(t, newInputTx(t), errStorage); !errors.Is(err, errStorage) || expired != 0 { + t.Fatalf("expired %d, %v", expired, err) + } +} diff --git a/services/core/internal/sessions/inputs.go b/services/core/internal/sessions/inputs.go index e041795b5..264cbb106 100644 --- a/services/core/internal/sessions/inputs.go +++ b/services/core/internal/sessions/inputs.go @@ -1,10 +1,14 @@ package sessions import ( + "context" "encoding/json" "fmt" + "slices" "strings" "time" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/jsonobject" ) // Input is a validated execution command, not an upstream wire type. @@ -35,6 +39,334 @@ func ValidateInputKey(key string) error { return nil } +// ValidateInputs checks and normalizes an input batch: 1..64 messages, +// cancellations and function results whose nonempty JSON object payloads total +// at most 512 KiB, an empty cancellation payload and a well-formed function +// result. It returns the normalized batch and its encoding, the batch's retry +// identity. +func ValidateInputs(inputs []Input) ([]Input, json.RawMessage, error) { + if len(inputs) == 0 || len(inputs) > 64 { + return nil, nil, fmt.Errorf("%w: input batch must contain 1..64 events", ErrInvalidInput) + } + batch := make([]Input, len(inputs)) + size := 0 + for i, input := range inputs { + size += len(input.Payload) + if size > 512*1024 || len(input.Payload) == 0 || (input.Kind != "message" && input.Kind != "cancel" && input.Kind != "tool_result") { + return nil, nil, fmt.Errorf("%w: input payloads must be nonempty and total at most 512 KiB", ErrInvalidInput) + } + payload, err := jsonobject.Normalize(input.Payload) + if err != nil { + return nil, nil, fmt.Errorf("%w: %w", ErrInvalidInput, err) + } + if input.Kind == "cancel" && string(payload) != "{}" { + return nil, nil, fmt.Errorf("%w: cancel payload must be empty", ErrInvalidInput) + } + if input.Kind == "tool_result" { + if _, err := ParseFunctionResultInput(payload); err != nil { + return nil, nil, err + } + } + batch[i] = Input{Kind: input.Kind, Payload: payload} + } + encoded, err := json.Marshal(batch) + return batch, encoded, err +} + +// ValidateMessageInputs validates a batch as ValidateInputs does and also +// requires that it holds only messages, as Environment input reservations and +// Session initial input do. +func ValidateMessageInputs(inputs []Input) ([]Input, json.RawMessage, error) { + for _, input := range inputs { + if input.Kind != "message" { + return nil, nil, ErrInvalidInput + } + } + return ValidateInputs(inputs) +} + +// decideReplay decides a request whose idempotency key may already hold a +// batch: an unused key admits the request, a key holding the same batch +// replays what it admitted, and a key holding a different batch is +// ErrIdempotencyConflict. The whole batch is the retry identity, so a replay +// never re-evaluates its inputs, such as a cancellation's target. +func decideReplay(used, matches bool) (bool, error) { + if !used { + return false, nil + } + if !matches { + return false, ErrIdempotencyConflict + } + return true, nil +} + +// InputGate is what the Session's Environment input reservations show about a +// new input batch. +type InputGate struct { + // Matches reports that the Session has no reservation under the batch's + // idempotency key, or one with the same batch. + Matches bool + // Blocked reports that a reservation is pending or one holds the batch's + // idempotency key. + Blocked bool +} + +// checkInputGate admits a new input batch only while no Environment input +// reservation is pending or holds its key: a reservation holding the key with +// a different batch is ErrIdempotencyConflict, and any other blocking +// reservation ErrInputPending. +func checkInputGate(gate InputGate) error { + if !gate.Matches { + return ErrIdempotencyConflict + } + if gate.Blocked { + return ErrInputPending + } + return nil +} + +// inputPlacement is where admission puts a message or cancellation. +type inputPlacement int + +const ( + // steersTurn adds a message to the active Turn. + steersTurn inputPlacement = iota + // cancelsTurn adds a cancellation to the active Turn and requests it. + cancelsTurn + // startsTurn starts a new Turn with a message on an idle Session. + startsTurn + // keepsIdentity records a cancellation on an idle Session without a Turn, + // which keeps only its retry identity. + keepsIdentity +) + +// placeInput decides where a message or cancellation goes, given whether the +// Session has an active Turn. +func placeInput(kind string, active bool) inputPlacement { + switch { + case active && kind == "cancel": + return cancelsTurn + case active: + return steersTurn + case kind == "message": + return startsTurn + } + return keepsIdentity +} + +// InputAdmissionTx is the Session transaction AdmitInputs runs in. +type InputAdmissionTx interface { + FunctionResultTx + TurnCancellationTx + InputProjectionTx + ActiveTurnTx + // CreateTurn creates a queued Turn in the Session and returns it as + // stored. + CreateTurn(ctx context.Context) (Turn, error) + // CreateTurnInput records input at position of the batch admitted under + // key, joined to turn, or to no Turn when turn is empty, and returns the + // Session sequence it allocates. + CreateTurnInput(ctx context.Context, turn, key string, position int32, input Input) (int64, error) +} + +// AdmitInputs admits the validated batch under key, input by input in batch +// order, and returns their receipts. A function result joins the Turn whose +// call it answers. A message steers the Session's active Turn or, on an idle +// Session, starts a new Turn, which publishes turn.created, then the message's +// Items, then the Session activity. A cancellation requests the cancellation +// of the active Turn; on an idle Session it joins no Turn and keeps only its +// retry identity. +func AdmitInputs(ctx context.Context, tx InputAdmissionTx, key string, batch []Input) ([]InputReceipt, error) { + receipts := make([]InputReceipt, 0, len(batch)) + for position, input := range batch { + receipt, err := admitInput(ctx, tx, key, int32(position), input) + if err != nil { + return nil, err + } + receipts = append(receipts, receipt) + } + return receipts, nil +} + +// admitInput admits the input at position of the batch under key. +func admitInput(ctx context.Context, tx InputAdmissionTx, key string, position int32, input Input) (InputReceipt, error) { + if input.Kind == "tool_result" { + result, err := ParseFunctionResultInput(input.Payload) + if err != nil { + return InputReceipt{}, err + } + turn, err := AdmitFunctionResult(ctx, tx, result) + if err != nil { + return InputReceipt{}, err + } + sequence, err := tx.CreateTurnInput(ctx, turn.ID, key, position, input) + if err != nil { + return InputReceipt{}, err + } + return InputReceipt{Sequence: sequence, TurnID: turn.ID}, nil + } + turn, active, err := tx.LoadActiveTurn(ctx) + if err != nil { + return InputReceipt{}, err + } + placement := placeInput(input.Kind, active) + if placement == startsTurn { + if turn, err = tx.CreateTurn(ctx); err != nil { + return InputReceipt{}, err + } + if err := tx.AppendChanges(ctx, TurnChanges(turn, true)...); err != nil { + return InputReceipt{}, err + } + } + sequence, err := tx.CreateTurnInput(ctx, turn.ID, key, position, input) + if err != nil { + return InputReceipt{}, err + } + if placement == cancelsTurn { + if err := CancelTurn(ctx, tx, turn); err != nil { + return InputReceipt{}, err + } + } + if err := ProjectInput(ctx, tx, sequence); err != nil { + return InputReceipt{}, err + } + if placement == startsTurn { + usage, err := tx.LoadUsage(ctx) + if err != nil { + return InputReceipt{}, err + } + if err := tx.AppendChanges(ctx, ActivityChange(turn, usage, nil)); err != nil { + return InputReceipt{}, err + } + } + return InputReceipt{Sequence: sequence, TurnID: turn.ID}, nil +} + +// InputTx is the Session transaction of the input use cases. +type InputTx interface { + InputAdmissionTx + InputActivityTx + InputStartTx + TurnTransitionTx + // LoadInputBatch reads the receipts of the inputs the Session admitted + // under key, in batch order and marked replayed, none when it admitted + // none, and reports whether they are batch. + LoadInputBatch(ctx context.Context, key string, batch json.RawMessage) ([]InputReceipt, bool, error) + // LoadInputGate reads what the Session's Environment input reservations + // show about batch under key. + LoadInputGate(ctx context.Context, key string, batch json.RawMessage) (InputGate, error) + // LoadEnvironment reads the Session's Environment; a Session without one + // is ErrNotFound. + LoadEnvironment(ctx context.Context) (Environment, error) + // FindInputReservation reads the Session's Environment input reservation + // under key when its batch is batch, nil otherwise, and reports whether + // the Session reserved key at all. + FindInputReservation(ctx context.Context, key string, batch json.RawMessage) (*EnvironmentInputReservation, bool, error) + // LoadInputReservation reads one of the Session's Environment input + // reservations; a missing one is ErrNotFound. An admitted reservation + // carries the receipts of its batch. + LoadInputReservation(ctx context.Context, reservation string) (EnvironmentInputReservation, error) + // CreateInputReservation reserves batch under key until the deadline the + // database clock sets, and returns the pending reservation as stored. + CreateInputReservation(ctx context.Context, key string, batch json.RawMessage) (EnvironmentInputReservation, error) + // ExpireInputReservation expires the reservation while it is pending and + // its deadline has passed by the database clock; otherwise it changes + // nothing. + ExpireInputReservation(ctx context.Context, reservation string) error + // AdmitInputReservation settles the pending reservation as admitted and + // returns when the database clock settled it. + AdmitInputReservation(ctx context.Context, reservation string) (time.Time, error) + // FailInputReservation settles the reservation as failed with code while + // it is pending; otherwise it changes nothing. + FailInputReservation(ctx context.Context, reservation, code string) error + // RecordInputAudit records the send_events write audit of the Session. + RecordInputAudit(ctx context.Context) error +} + +// InputStorage is the pooled storage of the input use cases. +type InputStorage interface { + // WithInputs runs apply in one pooled transaction under the lock of the + // tenant's Session and commits only when apply succeeds. A malformed + // tenant is ErrInvalidInput; a malformed, missing or deleted Session is + // ErrNotFound. + WithInputs(ctx context.Context, tenant, session string, apply func(context.Context, InputTx) error) error +} + +// InputReader reads a Session's admitted inputs, its Environment input +// reservations and the reservations execution may admit. +type InputReader interface { + // ListTurnInputs pages a Turn's admitted inputs in sequence order after + // sequence after, at most limit of 1..100. It is an internal recovery + // read, not the public event stream. A malformed ID, a negative after and + // a limit out of range are ErrInvalidInput; a missing Turn is ErrNotFound. + ListTurnInputs(ctx context.Context, tenant, session, turn string, after int64, limit int) ([]TurnInput, error) + // GetEnvironmentInputReservation reads one of the tenant's Session's + // Environment input reservations; an admitted one carries the receipts of + // its batch. A malformed tenant or reservation ID is ErrInvalidInput; a + // missing or deleted Session and a missing reservation are ErrNotFound. + GetEnvironmentInputReservation(ctx context.Context, tenant, session, reservation string) (EnvironmentInputReservation, error) + // ListEnvironmentInputWork lists, in ID order after the reservation ID + // after or from the first when it is empty, at most 100 pending + // reservations before their deadline of live Sessions whose tenant has an + // unrevoked device among connectedDevices, the Session's bound device when + // it has one. A malformed ID is ErrInvalidInput. + ListEnvironmentInputWork(ctx context.Context, after string, connectedDevices []string) ([]EnvironmentInputWork, error) +} + +// SubmitInputs admits a request's input batch in order under one Session +// lock and returns a receipt per input. The whole batch under its key is the +// retry identity: a retry replays the receipts of the batch it admitted, +// whatever has changed since, and the same key with a different batch is +// ErrIdempotencyConflict. A batch with a message waits for a pending file +// write to the Session's Environment (ErrTurnConflict), and no batch is +// admitted while an Environment input reservation is pending or holds the key +// (ErrInputPending). Internal receipts are not the response body of the +// public events endpoint. +func (s *Service) SubmitInputs(ctx context.Context, tenant, session, key string, inputs []Input) ([]InputReceipt, error) { + if err := ValidateInputKey(key); err != nil { + return nil, err + } + batch, encoded, err := ValidateInputs(inputs) + if err != nil { + return nil, err + } + var receipts []InputReceipt + err = s.storage.WithInputs(ctx, tenant, session, func(ctx context.Context, tx InputTx) error { + previous, matches, err := tx.LoadInputBatch(ctx, key, encoded) + if err != nil { + return err + } + replay, err := decideReplay(len(previous) > 0, matches) + if err != nil { + return err + } + if replay { + receipts = previous + return tx.RecordInputAudit(ctx) + } + if slices.ContainsFunc(batch, func(input Input) bool { return input.Kind == "message" }) { + if err := CheckFileWriteGate(ctx, tx); err != nil { + return err + } + } + gate, err := tx.LoadInputGate(ctx, key, encoded) + if err != nil { + return err + } + if err := checkInputGate(gate); err != nil { + return err + } + if receipts, err = AdmitInputs(ctx, tx, key, batch); err != nil { + return err + } + return tx.RecordInputAudit(ctx) + }) + if err != nil { + return nil, fmt.Errorf("submit turn inputs: %w", err) + } + return receipts, nil +} + // FunctionCall retains public identity and its opaque execution-adapter reference. type FunctionCall struct { CallID, ExecutorCallID, Name string diff --git a/services/core/internal/sessions/inputs_test.go b/services/core/internal/sessions/inputs_test.go new file mode 100644 index 000000000..6120dfb46 --- /dev/null +++ b/services/core/internal/sessions/inputs_test.go @@ -0,0 +1,322 @@ +package sessions + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "reflect" + "strings" + "testing" + "time" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/items" +) + +// fakeInputTx is fakeFunctionTx with the input methods, equally strict: it +// shares fakeTx's call log, and a method whose func is unset fails the test. +type fakeInputTx struct { + *fakeFunctionTx + + createTurn func() (Turn, error) + createTurnInput func() (int64, error) + loadInputBatch func() ([]InputReceipt, bool, error) + loadInputGate func() (InputGate, error) + findInputReservation func() (*EnvironmentInputReservation, bool, error) + loadInputReservation func() (EnvironmentInputReservation, error) + createInputReservation func() (EnvironmentInputReservation, error) + expireInputReservation func() error + admitInputReservation func() (time.Time, error) + failInputReservation func() error + recordInputAudit func() error +} + +var _ InputTx = (*fakeInputTx)(nil) + +func newInputTx(t *testing.T) *fakeInputTx { + return &fakeInputTx{fakeFunctionTx: &fakeFunctionTx{fakeTx: &fakeTx{t: t}}} +} + +func (f *fakeInputTx) CreateTurn(context.Context) (Turn, error) { + f.record("CreateTurn", f.createTurn != nil) + return f.createTurn() +} + +func (f *fakeInputTx) CreateTurnInput(_ context.Context, turn, key string, position int32, input Input) (int64, error) { + f.record("CreateTurnInput", f.createTurnInput != nil, turn, key, fmt.Sprint(position), input.Kind) + return f.createTurnInput() +} + +func (f *fakeInputTx) LoadInputBatch(_ context.Context, key string, batch json.RawMessage) ([]InputReceipt, bool, error) { + f.record("LoadInputBatch", f.loadInputBatch != nil, key, string(batch)) + return f.loadInputBatch() +} + +func (f *fakeInputTx) LoadInputGate(_ context.Context, key string, batch json.RawMessage) (InputGate, error) { + f.record("LoadInputGate", f.loadInputGate != nil, key, string(batch)) + return f.loadInputGate() +} + +func (f *fakeInputTx) FindInputReservation(_ context.Context, key string, batch json.RawMessage) (*EnvironmentInputReservation, bool, error) { + f.record("FindInputReservation", f.findInputReservation != nil, key, string(batch)) + return f.findInputReservation() +} + +func (f *fakeInputTx) LoadInputReservation(_ context.Context, reservation string) (EnvironmentInputReservation, error) { + f.record("LoadInputReservation", f.loadInputReservation != nil, reservation) + return f.loadInputReservation() +} + +func (f *fakeInputTx) CreateInputReservation(_ context.Context, key string, batch json.RawMessage) (EnvironmentInputReservation, error) { + f.record("CreateInputReservation", f.createInputReservation != nil, key, string(batch)) + return f.createInputReservation() +} + +func (f *fakeInputTx) ExpireInputReservation(_ context.Context, reservation string) error { + f.record("ExpireInputReservation", f.expireInputReservation != nil, reservation) + return f.expireInputReservation() +} + +func (f *fakeInputTx) AdmitInputReservation(_ context.Context, reservation string) (time.Time, error) { + f.record("AdmitInputReservation", f.admitInputReservation != nil, reservation) + return f.admitInputReservation() +} + +func (f *fakeInputTx) FailInputReservation(_ context.Context, reservation, code string) error { + f.record("FailInputReservation", f.failInputReservation != nil, reservation, code) + return f.failInputReservation() +} + +func (f *fakeInputTx) RecordInputAudit(context.Context) error { + f.record("RecordInputAudit", f.recordInputAudit != nil) + return f.recordInputAudit() +} + +func (s *fakeStorage) WithInputs(ctx context.Context, tenant, session string, apply func(context.Context, InputTx) error) error { + s.record("WithInputs", s.inputTx != nil, tenant, session) + return apply(ctx, s.inputTx) +} + +// inputKey is the idempotency key of the input batches under test. +const inputKey = "request" + +func messageInput(text string) Input { + payload, _ := json.Marshal(map[string]string{"text": text}) + return Input{Kind: "message", Payload: payload} +} + +var cancelInput = Input{Kind: "cancel", Payload: json.RawMessage(`{}`)} + +// sequences is a fake CreateTurnInput that allocates sequences from first. +func sequences(first int64) func() (int64, error) { + return func() (int64, error) { + first++ + return first - 1, nil + } +} + +// unprojected is a fake LoadInputSource: a message without text, which +// projects no Item. +var unprojected = returns(Source{Turn: testTurn, Kind: "message", Payload: json.RawMessage(`{}`)}) + +func TestValidateInputs(t *testing.T) { + batch, encoded, err := ValidateInputs([]Input{{Kind: "message", Payload: json.RawMessage(` {"text":"hi","a":1} `)}, {Kind: "cancel", Payload: json.RawMessage(` { } `)}}) + if err != nil || string(batch[0].Payload) != `{"a":1,"text":"hi"}` || string(encoded) != `[{"kind":"message","payload":{"a":1,"text":"hi"}},{"kind":"cancel","payload":{}}]` { + t.Fatalf("batch %s, %v", encoded, err) + } + for name, inputs := range map[string][]Input{ + "no inputs": nil, + "65 inputs": make([]Input, 65), + "unsupported kind": {messageInput("ok"), {Kind: "unsupported", Payload: json.RawMessage(`{}`)}}, + "missing payload": {{Kind: "message"}}, + "array payload": {{Kind: "message", Payload: json.RawMessage(`[]`)}}, + "cancel with a target": {{Kind: "cancel", Payload: json.RawMessage(`{"target":"other"}`)}}, + "over 512 KiB": {messageInput(strings.Repeat("x", 300*1024)), messageInput(strings.Repeat("y", 300*1024))}, + } { + if _, _, err := ValidateInputs(inputs); !errors.Is(err, ErrInvalidInput) { + t.Errorf("%s: %v", name, err) + } + } +} + +func TestValidateMessageInputs(t *testing.T) { + if _, encoded, err := ValidateMessageInputs([]Input{messageInput("hi")}); err != nil || string(encoded) != `[{"kind":"message","payload":{"text":"hi"}}]` { + t.Fatalf("batch %s, %v", encoded, err) + } + for name, inputs := range map[string][]Input{"cancel": {messageInput("hi"), cancelInput}, "no inputs": nil} { + if _, _, err := ValidateMessageInputs(inputs); !errors.Is(err, ErrInvalidInput) { + t.Errorf("%s: %v", name, err) + } + } +} + +func TestDecideReplay(t *testing.T) { + for _, test := range []struct { + used, matches, replay bool + err error + }{ + {false, false, false, nil}, + {true, true, true, nil}, + {true, false, false, ErrIdempotencyConflict}, + } { + if replay, err := decideReplay(test.used, test.matches); replay != test.replay || err != test.err { + t.Errorf("used %v matches %v: %v, %v", test.used, test.matches, replay, err) + } + } +} + +func TestCheckInputGate(t *testing.T) { + for gate, want := range map[InputGate]error{ + {Matches: true}: nil, + {Matches: true, Blocked: true}: ErrInputPending, + {Matches: false, Blocked: true}: ErrIdempotencyConflict, + } { + if err := checkInputGate(gate); err != want { + t.Errorf("%+v: %v", gate, err) + } + } +} + +func TestPlaceInput(t *testing.T) { + for _, test := range []struct { + kind string + active bool + want inputPlacement + }{ + {"message", true, steersTurn}, + {"cancel", true, cancelsTurn}, + {"message", false, startsTurn}, + {"cancel", false, keepsIdentity}, + } { + if got := placeInput(test.kind, test.active); got != test.want { + t.Errorf("%s active %v: %v", test.kind, test.active, got) + } + } +} + +func TestAdmitInput(t *testing.T) { + running := turnWith(TurnInProgress) + t.Run("an idle message starts a Turn", func(t *testing.T) { + var changes []SessionChange + tx := newInputTx(t) + tx.loadActiveTurn, tx.createTurn, tx.appendChanges = activeTurn(nil), returns(turnWith(TurnQueued)), collect(&changes) + tx.createTurnInput, tx.loadUsage = sequences(7), returns(json.RawMessage(`{}`)) + tx.loadInputSource = returns(Source{Turn: testTurn, Kind: "message", Sequence: 7, Payload: json.RawMessage(`{"text":"hi"}`)}) + tx.loadItem, tx.putItem = returns(items.Stored{}), func(items.Change) (*int32, error) { return nil, nil } + receipt, err := admitInput(t.Context(), tx, inputKey, 0, messageInput("hi")) + if err != nil || receipt != (InputReceipt{Sequence: 7, TurnID: testTurn}) { + t.Fatalf("receipt %+v, %v", receipt, err) + } + item := items.Identity(testTurn, "input:7") + assertCalls(t, tx.fakeTx, "LoadActiveTurn", "CreateTurn", "AppendChanges agent.session.turn.created", "CreateTurnInput "+testTurn+" request 0 message", + "LoadInputSource 7", "LoadItem "+testTurn+" "+item, "PutItem "+testTurn+" "+item, "AppendChanges agent.session.turn.item.added", + "LoadUsage", "AppendChanges agent.session.in_progress") + }) + for _, test := range []struct { + name string + input Input + active *Turn + want InputReceipt + calls []string + }{ + {"a message steers the active Turn", messageInput("more"), &running, InputReceipt{Sequence: 8, TurnID: testTurn}, + []string{"LoadActiveTurn", "CreateTurnInput " + testTurn + " request 1 message", "LoadInputSource 8"}}, + {"a cancellation requests the active Turn's cancellation", cancelInput, &running, InputReceipt{Sequence: 8, TurnID: testTurn}, + []string{"LoadActiveTurn", "CreateTurnInput " + testTurn + " request 1 cancel", "RequestTurnCancel " + testTurn, "LoadInputSource 8"}}, + {"an idle cancellation keeps only its retry identity", cancelInput, nil, InputReceipt{Sequence: 8}, + []string{"LoadActiveTurn", "CreateTurnInput request 1 cancel", "LoadInputSource 8"}}, + } { + t.Run(test.name, func(t *testing.T) { + tx := newInputTx(t) + tx.loadActiveTurn, tx.createTurnInput, tx.loadInputSource, tx.requestTurnCancel = activeTurn(test.active), sequences(8), unprojected, done + receipt, err := admitInput(t.Context(), tx, inputKey, 1, test.input) + if err != nil || receipt != test.want { + t.Fatalf("receipt %+v, %v", receipt, err) + } + assertCalls(t, tx.fakeTx, test.calls...) + }) + } + t.Run("a function result joins the Turn of its call", func(t *testing.T) { + tx := newInputTx(t) + tx.findResultTurn = func() (Turn, bool, error) { return turnWith(TurnWaiting), true, nil } + tx.matchFunctionResult, tx.submitFunctionResult, tx.createTurnInput = returns(FunctionResultMatch{Recorded: true}), done, sequences(9) + payload := json.RawMessage(`{"turn_id":"` + testTurn + `","call_id":"call","result":{"ok":true}}`) + receipt, err := admitInput(t.Context(), tx, inputKey, 0, Input{Kind: "tool_result", Payload: payload}) + if err != nil || receipt != (InputReceipt{Sequence: 9, TurnID: testTurn}) { + t.Fatalf("receipt %+v, %v", receipt, err) + } + assertCalls(t, tx.fakeTx, "FindResultTurn "+testTurn, "MatchFunctionResult "+testTurn+` call {"ok":true}`, "SubmitFunctionResult "+testTurn+` call {"ok":true}`, + "CreateTurnInput "+testTurn+" request 0 tool_result") + }) +} + +func TestSubmitInputs(t *testing.T) { + const batch = `[{"kind":"message","payload":{"text":"hi"}},{"kind":"cancel","payload":{}}]` + inputs := []Input{messageInput("hi"), cancelInput} + running := turnWith(TurnInProgress) + submit := func(t *testing.T, tx *fakeInputTx, inputs []Input) ([]InputReceipt, error) { + storage := &fakeStorage{t: t, inputTx: tx} + receipts, err := deviceService(t, storage).SubmitInputs(t.Context(), testTenant, testSession, inputKey, inputs) + assertStorageCalls(t, storage, "WithInputs "+testTenant+" "+testSession) + return receipts, err + } + t.Run("admits the batch in order", func(t *testing.T) { + tx := newInputTx(t) + tx.loadInputBatch = func() ([]InputReceipt, bool, error) { return nil, true, nil } + tx.loadPendingFileWrite, tx.loadInputGate = returns(false), returns(InputGate{Matches: true}) + tx.loadActiveTurn, tx.createTurnInput, tx.loadInputSource, tx.requestTurnCancel = activeTurn(&running), sequences(1), unprojected, done + tx.recordInputAudit = done + receipts, err := submit(t, tx, inputs) + if err != nil || !reflect.DeepEqual(receipts, []InputReceipt{{Sequence: 1, TurnID: testTurn}, {Sequence: 2, TurnID: testTurn}}) { + t.Fatalf("receipts %+v, %v", receipts, err) + } + assertCalls(t, tx.fakeTx, "LoadInputBatch request "+batch, "LoadPendingFileWrite", "LoadInputGate request "+batch, + "LoadActiveTurn", "CreateTurnInput "+testTurn+" request 0 message", "LoadInputSource 1", + "LoadActiveTurn", "CreateTurnInput "+testTurn+" request 1 cancel", "RequestTurnCancel "+testTurn, "LoadInputSource 2", + "RecordInputAudit") + }) + t.Run("a retry replays its receipts", func(t *testing.T) { + previous := []InputReceipt{{Sequence: 1, TurnID: testTurn, Replayed: true}, {Sequence: 2, TurnID: testTurn, Replayed: true}} + tx := newInputTx(t) + tx.loadInputBatch = func() ([]InputReceipt, bool, error) { return previous, true, nil } + tx.recordInputAudit = done + receipts, err := submit(t, tx, inputs) + if err != nil || !reflect.DeepEqual(receipts, previous) { + t.Fatalf("receipts %+v, %v", receipts, err) + } + assertCalls(t, tx.fakeTx, "LoadInputBatch request "+batch, "RecordInputAudit") + }) + t.Run("a different batch under the key conflicts", func(t *testing.T) { + tx := newInputTx(t) + tx.loadInputBatch = func() ([]InputReceipt, bool, error) { return []InputReceipt{{Sequence: 1}}, false, nil } + if _, err := submit(t, tx, inputs); !errors.Is(err, ErrIdempotencyConflict) { + t.Fatal(err) + } + assertCalls(t, tx.fakeTx, "LoadInputBatch request "+batch) + }) + t.Run("a pending file write holds a message", func(t *testing.T) { + tx := newInputTx(t) + tx.loadInputBatch, tx.loadPendingFileWrite = func() ([]InputReceipt, bool, error) { return nil, true, nil }, returns(true) + if _, err := submit(t, tx, inputs); !errors.Is(err, ErrTurnConflict) { + t.Fatal(err) + } + assertCalls(t, tx.fakeTx, "LoadInputBatch request "+batch, "LoadPendingFileWrite") + }) + t.Run("a pending reservation holds a cancellation", func(t *testing.T) { + tx := newInputTx(t) + tx.loadInputBatch, tx.loadInputGate = func() ([]InputReceipt, bool, error) { return nil, true, nil }, returns(InputGate{Matches: true, Blocked: true}) + if _, err := submit(t, tx, []Input{cancelInput}); !errors.Is(err, ErrInputPending) { + t.Fatal(err) + } + assertCalls(t, tx.fakeTx, `LoadInputBatch request [{"kind":"cancel","payload":{}}]`, `LoadInputGate request [{"kind":"cancel","payload":{}}]`) + }) + t.Run("an invalid request touches no storage", func(t *testing.T) { + service := deviceService(t, &fakeStorage{t: t}) + if _, err := service.SubmitInputs(t.Context(), testTenant, testSession, " ", inputs); !errors.Is(err, ErrInvalidInput) { + t.Fatal(err) + } + if _, err := service.SubmitInputs(t.Context(), testTenant, testSession, inputKey, nil); !errors.Is(err, ErrInvalidInput) { + t.Fatal(err) + } + }) +} diff --git a/services/core/internal/sessions/reader.go b/services/core/internal/sessions/reader.go index 6da1497a6..59e24ba32 100644 --- a/services/core/internal/sessions/reader.go +++ b/services/core/internal/sessions/reader.go @@ -9,6 +9,7 @@ type Reader interface { DeviceReader EnvironmentReader ExecutorCredentialReader + InputReader ItemReader ModelExecutionReader SessionReader diff --git a/services/core/internal/sessions/service.go b/services/core/internal/sessions/service.go index 8612a15e7..df891e54c 100644 --- a/services/core/internal/sessions/service.go +++ b/services/core/internal/sessions/service.go @@ -20,4 +20,5 @@ type Storage interface { DeviceStorage ExecutorCredentialStorage SessionStorage + InputStorage } 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 6903443af..b5670acc4 100644 --- a/services/core/internal/store/admin_session_archive_race_test.go +++ b/services/core/internal/store/admin_session_archive_race_test.go @@ -59,7 +59,7 @@ func TestManagedSessionArchiveOrdersConcurrentInput(t *testing.T) { go func() { defer wg.Done() <-start - _, err := s.ReserveEnvironmentInput(t.Context(), tenant, session.ID, "racing-input", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"racing"}`)}}) + _, err := sessionService(t, s).ReserveEnvironmentInput(t.Context(), tenant, session.ID, "racing-input", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"racing"}`)}}) if err != nil && !errors.Is(err, sessions.ErrEnvironmentUnavailable) { t.Error(err) } diff --git a/services/core/internal/store/admin_session_archive_test.go b/services/core/internal/store/admin_session_archive_test.go index d35a6477a..cda2e569b 100644 --- a/services/core/internal/store/admin_session_archive_test.go +++ b/services/core/internal/store/admin_session_archive_test.go @@ -114,7 +114,7 @@ func TestManagedSessionArchiveUnallocatedAndGuards(t *testing.T) { if _, err := deploymentExecution(t, w).ReserveAllocation(t.Context(), deployment.AllocationKey{TenantID: tenant, EnvironmentID: session.Environment.ID}, installation, runtimedevice.HashCredential(uuid.NewString())); !errors.Is(err, deployment.ErrInvalidInput) { t.Fatal("archived Environment allocated after archive", err) } - if _, err := s.ReserveEnvironmentInput(t.Context(), tenant, session.ID, "later", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"later"}`)}}); !errors.Is(err, sessions.ErrEnvironmentUnavailable) { + if _, err := sessionService(t, s).ReserveEnvironmentInput(t.Context(), tenant, session.ID, "later", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"later"}`)}}); !errors.Is(err, sessions.ErrEnvironmentUnavailable) { t.Fatal("archived Environment accepted new input", err) } view, err := deploymentService(t, s).View(t.Context()) 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 3c86b135c..7085b64fa 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 @@ -96,7 +96,7 @@ func TestAdminSessionArchiveWorkerHTTPPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(t.Context(), project.TenantID, active.ID, "pending-turn", json.RawMessage(`{"text":"pending"}`)) + input, err := store.SendMessage(t.Context(), s, project.TenantID, active.ID, "pending-turn", json.RawMessage(`{"text":"pending"}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/archive_cancellation_test.go b/services/core/internal/store/archive_cancellation_test.go index 9c56f82eb..6e3da3dce 100644 --- a/services/core/internal/store/archive_cancellation_test.go +++ b/services/core/internal/store/archive_cancellation_test.go @@ -123,7 +123,7 @@ func TestArchiveWaitingCancellationReceipts(t *testing.T) { } time.Sleep(time.Millisecond) } - pending, err := s.ReserveEnvironmentInput(t.Context(), h.tenant, session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}, {Kind: "message", Payload: json.RawMessage(`{"text":"second"}`)}}) + pending, err := store.SessionService(t, s).ReserveEnvironmentInput(t.Context(), h.tenant, session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}, {Kind: "message", Payload: json.RawMessage(`{"text":"second"}`)}}) if err != nil { t.Fatal(err) } @@ -142,7 +142,7 @@ func TestArchiveWaitingCancellationReceipts(t *testing.T) { } if scenario == "cancel_revoke_archive" { - if _, err := s.RequestCancel(t.Context(), h.tenant, session.ID, "ordinary-cancel"); err != nil { + if _, err := store.RequestCancel(t.Context(), s, h.tenant, session.ID, "ordinary-cancel"); err != nil { t.Fatal(err) } } diff --git a/services/core/internal/store/claude_execution_test.go b/services/core/internal/store/claude_execution_test.go index b76101ae7..cf3b54681 100644 --- a/services/core/internal/store/claude_execution_test.go +++ b/services/core/internal/store/claude_execution_test.go @@ -152,7 +152,7 @@ func TestClaudeInvalidImageResultRejectsWholeBatchBeforePersistence(t *testing.T if err != nil || turn.Status != sessions.TurnWaiting || !turn.CancelRequestedAt.IsZero() { t.Fatal(turn, err) } - history, err := h.s.ListTurnInputs(t.Context(), h.tenant, h.session.ID, input.TurnID, 0, 100) + history, err := store.SessionAdapter(h.s).ListTurnInputs(t.Context(), h.tenant, h.session.ID, input.TurnID, 0, 100) if err != nil || len(history) != 1 { t.Fatal(history, err) } diff --git a/services/core/internal/store/command_output_test.go b/services/core/internal/store/command_output_test.go index e2e1cb94f..a87885fa2 100644 --- a/services/core/internal/store/command_output_test.go +++ b/services/core/internal/store/command_output_test.go @@ -22,7 +22,7 @@ func TestCommandOutputCommitsFragmentsSnapshotsAndRecovery(t *testing.T) { if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(ctx, tenant, session.ID, "start", json.RawMessage(`{"text":"run commands"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "start", json.RawMessage(`{"text":"run commands"}`)) if err != nil { t.Fatal(err) } @@ -135,7 +135,7 @@ func TestExecutionJournalsCommandOutputBeforeCancellation(t *testing.T) { h.read(testExecutionRequest) h.write(input.TurnID, proto.TypeToolCall, proto.ToolCallPayload{ID: "cmd", Stage: "before", Observation: &proto.ToolObservation{Kind: "command", Command: "wait", Status: "in_progress"}}) h.write(input.TurnID, proto.TypeCommandOutput, proto.CommandOutputPayload{ID: "cmd", Delta: "partial"}) - if _, err := h.s.RequestCancel(ctx, h.tenant, h.session.ID, "cancel"); err != nil { + if _, err := store.RequestCancel(ctx, h.s, h.tenant, h.session.ID, "cancel"); err != nil { t.Fatal(err) } env := h.read(proto.TypePromptCancel) 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 7de1cccaf..2dced38ca 100644 --- a/services/core/internal/store/creation_stream_settlement_public_test.go +++ b/services/core/internal/store/creation_stream_settlement_public_test.go @@ -171,7 +171,7 @@ func TestCreationStreamPublicLifetimes(t *testing.T) { created.ended(t, 5*time.Second) connect(first.Session.Environment.ID) - if _, err := s.ReserveEnvironmentInput(t.Context(), tenant, first.Session.ID, "later", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"later"}`)}}); err != nil { + if _, err := store.SessionService(t, s).ReserveEnvironmentInput(t.Context(), tenant, first.Session.ID, "later", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"later"}`)}}); err != nil { t.Fatal(err) } if current, err := store.SessionAdapter(s).GetSession(t.Context(), tenant, first.Session.ID); err != nil || !current.PendingInput { @@ -211,7 +211,7 @@ func TestCreationStreamPublicLifetimes(t *testing.T) { if err != nil { t.Fatal(err) } - if settled, err := s.CancelEnvironmentInput(t.Context(), tenant, session, reservation); err != nil || settled.State != sessions.EnvironmentInputCancelled { + if settled, err := store.CancelEnvironmentInput(t.Context(), s, tenant, session, reservation); err != nil || settled.State != sessions.EnvironmentInputCancelled { t.Fatal(settled.State, err) } if after, err := store.SessionAdapter(s).SessionEventCursor(t.Context(), tenant, session); err != nil || after != cursor { diff --git a/services/core/internal/store/deployment_model_providers_http_test.go b/services/core/internal/store/deployment_model_providers_http_test.go index 0b0848aa9..774319df1 100644 --- a/services/core/internal/store/deployment_model_providers_http_test.go +++ b/services/core/internal/store/deployment_model_providers_http_test.go @@ -244,7 +244,7 @@ func TestLegacySessionWithoutProviderCannotStartWork(t *testing.T) { } executor := connectFixtureRuntime(t, h, legacy) // Reserved directly, as a pre-upgrade Core did. - pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, legacy.ID, "before-upgrade", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"old"}`)}}) + pending, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, legacy.ID, "before-upgrade", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"old"}`)}}) if err != nil { t.Fatal(err) } @@ -268,7 +268,7 @@ func TestLegacySessionWithoutProviderCannotStartWork(t *testing.T) { t.Fatal("rejected work was queued", before, after) } awaitDaemonRemoteCondition(t, t.Context(), 5*time.Second, "legacy reservation settled", func() bool { - got, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, legacy.ID, pending.ID) + got, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, legacy.ID, pending.ID) return err == nil && got.State == sessions.EnvironmentInputFailed }) session, err := store.SessionAdapter(h.s).GetSession(t.Context(), h.tenant, legacy.ID) diff --git a/services/core/internal/store/dispatch_test.go b/services/core/internal/store/dispatch_test.go index e5ca9514c..90d97092a 100644 --- a/services/core/internal/store/dispatch_test.go +++ b/services/core/internal/store/dispatch_test.go @@ -127,7 +127,7 @@ func newDispatchHarnessForSession(t *testing.T, configuration []byte, local bool func (h *dispatchHarness) message(key, text string) sessions.InputReceipt { h.t.Helper() body, _ := json.Marshal(map[string]string{"text": text}) - r, err := h.s.SubmitMessage(context.Background(), h.tenant, h.session.ID, key, body) + r, err := store.SendMessage(context.Background(), h.s, h.tenant, h.session.ID, key, body) if err != nil { h.t.Fatal(err) } @@ -290,7 +290,7 @@ func TestExecutionCancellationRequiresReceiptAndSurvivesContextEnd(t *testing.T) first := h.message("first", "Run") result := h.run(context.Background(), first.TurnID) h.read(testExecutionRequest) - _, err := h.s.RequestCancel(context.Background(), h.tenant, h.session.ID, "cancel") + _, err := store.RequestCancel(context.Background(), h.s, h.tenant, h.session.ID, "cancel") if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/environment_admission_test.go b/services/core/internal/store/environment_admission_test.go index 48f5a5d89..052666fdb 100644 --- a/services/core/internal/store/environment_admission_test.go +++ b/services/core/internal/store/environment_admission_test.go @@ -73,7 +73,7 @@ func environmentAdmissionPending(t *testing.T, h *dispatchHarness, key string) s awaitDaemonRemoteCondition(t, t.Context(), 3*time.Second, "input reservation", func() bool { return pool.QueryRow(t.Context(), "SELECT id::text FROM environment_input_reservations WHERE session_id=$1 AND idempotency_key=$2", h.session.ID, key).Scan(&id) == nil }) - pending, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, id) + pending, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, id) if err != nil { t.Fatal(err) } @@ -181,7 +181,7 @@ func TestEnvironmentAdmissionSettlementDoesNotCreateTurn(t *testing.T) { } case "cancelled": expected = execution.ErrEnvironmentInputCancelled - if _, err := h.s.CancelEnvironmentInput(t.Context(), h.tenant, h.session.ID, pending.ID); err != nil { + if _, err := store.CancelEnvironmentInput(t.Context(), h.s, h.tenant, h.session.ID, pending.ID); err != nil { t.Fatal(err) } case "deleted": @@ -220,7 +220,7 @@ func TestEnvironmentAdmissionSettlementDoesNotCreateTurn(t *testing.T) { if name == "deleted" { return } - retained, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) + retained, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) 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 a3b57b5b2..62221cd64 100644 --- a/services/core/internal/store/environment_claim_worker_test.go +++ b/services/core/internal/store/environment_claim_worker_test.go @@ -19,8 +19,7 @@ func TestWorkerReconcilesEnvironmentPromotionBeforeStart(t *testing.T) { s, db := newTestStoreDB(t) tenant, pending := newEnvironmentExpiryReservation(t, s) owner := executionOwner(t, db, s) - writer := owner.Store - got, err := writer.PromoteEnvironmentInput(t.Context(), tenant, pending.SessionID, pending.ID) + got, err := owner.Sessions.PromoteEnvironmentInput(t.Context(), tenant, pending.SessionID, pending.ID) if err != nil || len(got.Receipts) != 1 || got.Receipts[0].Replayed { t.Fatal(got, err) } @@ -66,7 +65,7 @@ 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 := executionOwner(t, db, s).Store + successor := executionOwner(t, db, s).Sessions retry, err := successor.PromoteEnvironmentInput(t.Context(), tenant, pending.SessionID, pending.ID) if deleted { if !errors.Is(err, sessions.ErrNotFound) { diff --git a/services/core/internal/store/environment_device_test.go b/services/core/internal/store/environment_device_test.go index 46801cb71..e56cadc9e 100644 --- a/services/core/internal/store/environment_device_test.go +++ b/services/core/internal/store/environment_device_test.go @@ -61,7 +61,7 @@ func TestWorkerEnvironmentSelectsCapableDeviceWithoutMovingBinding(t *testing.T) } stop() for _, value := range []sessions.EnvironmentInputReservation{pending, bound} { - stored, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, value.SessionID, value.ID) + stored, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, value.SessionID, value.ID) if err != nil || stored.State != sessions.EnvironmentInputPending || !stored.Deadline.Equal(value.Deadline) { t.Fatal("device readiness changed pending input", err) } diff --git a/services/core/internal/store/environment_directory_active_test.go b/services/core/internal/store/environment_directory_active_test.go index 161e59146..2f6db1cf1 100644 --- a/services/core/internal/store/environment_directory_active_test.go +++ b/services/core/internal/store/environment_directory_active_test.go @@ -5,12 +5,13 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" ) func TestEnvironmentDirectoryActiveRunUsesExistingOwner(t *testing.T) { h, w, environment := directoryWorker(t) awaitFixtureCapabilities(t, h, workerEnvironmentCapabilities()) - pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "execute", []sessions.Input{{Kind: "message", Payload: []byte(`{"text":"work"}`)}}) + pending, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "execute", []sessions.Input{{Kind: "message", Payload: []byte(`{"text":"work"}`)}}) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/environment_expiry_worker_test.go b/services/core/internal/store/environment_expiry_worker_test.go index 9ce19dd15..93be26e40 100644 --- a/services/core/internal/store/environment_expiry_worker_test.go +++ b/services/core/internal/store/environment_expiry_worker_test.go @@ -26,7 +26,7 @@ func newEnvironmentExpiryReservation(t *testing.T, s *store.Store) (string, sess if err != nil { t.Fatal(err) } - pending, err := s.ReserveEnvironmentInput(t.Context(), tenant, session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"wait for the environment"}`)}}) + pending, err := store.SessionService(t, s).ReserveEnvironmentInput(t.Context(), tenant, session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"wait for the environment"}`)}}) if err != nil { t.Fatal(err) } @@ -68,7 +68,7 @@ func waitEnvironmentExpiry(t *testing.T, s *store.Store, tenant string, pending t.Helper() deadline := time.Now().Add(5 * time.Second) for { - got, err := s.GetEnvironmentInputReservation(t.Context(), tenant, pending.SessionID, pending.ID) + got, err := store.SessionAdapter(s).GetEnvironmentInputReservation(t.Context(), tenant, pending.SessionID, pending.ID) if err != nil { t.Fatal(err) } @@ -107,7 +107,7 @@ func TestWorkerEnvironmentExpiryWithoutDevicesAndAfterRestart(t *testing.T) { d := &execution.Dispatcher{Store: s, Registry: runtimegateway.NewRegistry()} _, stop := startEnvironmentExpiryWorker(t, db, d) waitEnvironmentExpiry(t, s, dueTenant, due) - got, err := s.GetEnvironmentInputReservation(t.Context(), futureTenant, future.SessionID, future.ID) + got, err := store.SessionAdapter(s).GetEnvironmentInputReservation(t.Context(), futureTenant, future.SessionID, future.ID) if err != nil || got.State != sessions.EnvironmentInputPending || !got.Deadline.Equal(future.Deadline) { t.Fatal("future input changed", got, err) } diff --git a/services/core/internal/store/environment_file_writes_test.go b/services/core/internal/store/environment_file_writes_test.go index 68cc9c7c6..73c342b23 100644 --- a/services/core/internal/store/environment_file_writes_test.go +++ b/services/core/internal/store/environment_file_writes_test.go @@ -61,10 +61,10 @@ func TestEnvironmentFileWriteRetainsUnknownAcrossLeaseLoss(t *testing.T) { if _, err := next.ReserveEnvironmentFileWrite(ctx, f.tenant, f.env.ID, another); !errors.Is(err, sessions.ErrTurnConflict) { t.Fatal("restart admitted successor", err) } - if _, err := reopened.ReserveEnvironmentInput(ctx, f.tenant, f.session.ID, "new-input", []sessions.Input{messageInput("new")}); !errors.Is(err, sessions.ErrTurnConflict) { + if _, err := sessionService(t, reopened).ReserveEnvironmentInput(ctx, f.tenant, f.session.ID, "new-input", []sessions.Input{messageInput("new")}); !errors.Is(err, sessions.ErrTurnConflict) { t.Fatal("unknown write admitted input", err) } - if _, err := reopened.SubmitMessage(ctx, f.tenant, f.session.ID, "direct", messageInput("new").Payload); !errors.Is(err, sessions.ErrTurnConflict) { + if _, err := sendMessage(ctx, reopened, f.tenant, f.session.ID, "direct", messageInput("new").Payload); !errors.Is(err, sessions.ErrTurnConflict) { t.Fatal("direct admission bypassed write", err) } if _, err := sessionAdapter(reopened).GetEnvironment(ctx, f.tenant, f.env.ID); err != nil { @@ -73,7 +73,7 @@ func TestEnvironmentFileWriteRetainsUnknownAcrossLeaseLoss(t *testing.T) { if _, err := sessionAdapter(reopened).GetSession(ctx, f.tenant, f.session.ID); err != nil { t.Fatal("write gate prevented recovery read", err) } - if receipt, err := reopened.RequestCancel(ctx, f.tenant, f.session.ID, "idle-cancel"); err != nil || receipt.TurnID != "" { + if receipt, err := requestCancel(ctx, reopened, f.tenant, f.session.ID, "idle-cancel"); err != nil || receipt.TurnID != "" { t.Fatal("write gate imposed mutation admission on idle cancellation", receipt, err) } if got, err := FixtureFileWrite(ctx, reopened.pool, f.tenant, f.env.ID, f.key.ID); err != nil || got.State != "pending" { @@ -178,7 +178,7 @@ func TestEnvironmentFileWriteSerializesWithInputAndRetry(t *testing.T) { group.Go(func() { <-start var err error - pending, err = f.s.ReserveEnvironmentInput(ctx, f.tenant, f.session.ID, inputKey, []sessions.Input{messageInput("race")}) + pending, err = sessionService(t, f.s).ReserveEnvironmentInput(ctx, f.tenant, f.session.ID, inputKey, []sessions.Input{messageInput("race")}) inputs <- err }) close(start) @@ -189,7 +189,7 @@ func TestEnvironmentFileWriteSerializesWithInputAndRetry(t *testing.T) { t.Fatal(err) } } else if inputErr == nil && errors.Is(writeErr, sessions.ErrTurnConflict) { - if _, err := f.s.CancelEnvironmentInput(ctx, f.tenant, f.session.ID, pending.ID); err != nil { + if _, err := cancelEnvironmentInput(ctx, f.s, f.tenant, f.session.ID, pending.ID); err != nil { t.Fatal(err) } } else { diff --git a/services/core/internal/store/environment_initial_input_test.go b/services/core/internal/store/environment_initial_input_test.go index ae58bba5b..2cb963652 100644 --- a/services/core/internal/store/environment_initial_input_test.go +++ b/services/core/internal/store/environment_initial_input_test.go @@ -21,7 +21,7 @@ func initialEnvironmentReservation(t *testing.T, s *Store, pool *pgxpool.Pool, t if err := pool.QueryRow(t.Context(), "SELECT id FROM environment_input_reservations WHERE session_id=$1 AND is_initial", session).Scan(&id); err != nil { t.Fatal(err) } - reservation, err := s.GetEnvironmentInputReservation(t.Context(), tenant, session, id) + reservation, err := sessionAdapter(s).GetEnvironmentInputReservation(t.Context(), tenant, session, id) if err != nil || !reservation.IsInitial { t.Fatal("missing initial origin", reservation, err) } @@ -49,7 +49,7 @@ func TestEnvironmentInitialExpiryRollsBackWithFailureEventAndSerializesPromotion _, _ = pool.Exec(context.Background(), "ALTER TABLE session_events DROP CONSTRAINT IF EXISTS "+constraint) }) writer := executionWriter(t, s) - if _, err := writer.ExpireEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID); err == nil { + if _, err := sessionService(t, writer).ExpireEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID); err == nil { t.Fatal("expiry committed without its failure event") } retained := initialEnvironmentReservation(t, s, pool, tenant, session.ID) @@ -65,7 +65,7 @@ func TestEnvironmentInitialExpiryRollsBackWithFailureEventAndSerializesPromotion err error } results := make(chan result, 2) - for _, settle := range []func(context.Context, string, string, string) (sessions.EnvironmentInputReservation, error){writer.PromoteEnvironmentInput, s.ExpireEnvironmentInput} { + for _, settle := range []func(context.Context, string, string, string) (sessions.EnvironmentInputReservation, error){sessionExecution(t, writer.lease).PromoteEnvironmentInput, sessionService(t, s).ExpireEnvironmentInput} { go func() { r, err := settle(t.Context(), tenant, session.ID, reservation.ID); results <- result{r, err} }() } for i := 0; i < 2; i++ { @@ -156,7 +156,7 @@ func TestEnvironmentInitialInputCreationRetainsCursorIdentityAndPromotion(t *tes } requireEnvironmentInputActivity(t, s, tenant, session.ID, connectedStatus, "") environmentInputHistory(t, pool, session.ID, 0, 0) - promoted, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID) + promoted, err := sessionExecution(t, writer.lease).PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID) if err != nil || promoted.State != sessions.EnvironmentInputAdmitted || !promoted.IsInitial || len(promoted.Receipts) != 2 { t.Fatal("initial batch did not promote", promoted, err) } @@ -164,7 +164,7 @@ func TestEnvironmentInitialInputCreationRetainsCursorIdentityAndPromotion(t *tes if active.LastTurn == nil || active.LastTurn.Status != sessions.TurnInProgress || active.PendingInput { t.Fatal("promotion did not claim its Turn", active.LastTurn) } - replay, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID) + replay, err := sessionExecution(t, writer.lease).PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID) if err != nil || len(replay.Receipts) != 2 || !replay.Receipts[0].Replayed || !replay.Receipts[1].Replayed { t.Fatal("promotion retry granted fresh receipts", replay, err) } @@ -211,7 +211,7 @@ func TestEnvironmentInitialInputExpiryHasNoTurnAndCannotReplay(t *testing.T) { } writer := executionWriter(t, s) for reservation.State == sessions.EnvironmentInputPending { - count, err := writer.ExpireEnvironmentInputs(t.Context()) + count, err := sessionExecution(t, writer.lease).ExpireEnvironmentInputs(t.Context()) if err != nil || count < 1 || count > 32 { t.Fatal("expiry made no bounded progress", count, err) } @@ -250,7 +250,7 @@ func TestEnvironmentInitialInputExpiryHasNoTurnAndCannotReplay(t *testing.T) { if err := sessionExecution(t, writer.lease).ObserveEnvironmentConnection(t.Context(), tenant, session.Environment.ID, generation, 1, true); err != nil { t.Fatal(err) } - late, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID) + late, err := sessionExecution(t, writer.lease).PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID) if err != nil || late.State != sessions.EnvironmentInputExpired || len(late.Receipts) != 0 { t.Fatal("late connection resurrected initial input", late, err) } @@ -269,7 +269,7 @@ func TestEnvironmentInitialInputExpiryHasNoTurnAndCannotReplay(t *testing.T) { if err := sessionService(t, reopened).DeleteSession(t.Context(), sessions.DeleteSessionCommand{TenantID: tenant, SessionID: session.ID}); !errors.Is(err, sessions.ErrNotIdle) { t.Fatal("pending later input deleted", err) } - if _, err := reopened.CancelEnvironmentInput(t.Context(), tenant, session.ID, later.ID); err != nil { + if _, err := cancelEnvironmentInput(t.Context(), reopened, tenant, session.ID, later.ID); err != nil { t.Fatal(err) } if err := sessionService(t, reopened).DeleteSession(t.Context(), sessions.DeleteSessionCommand{TenantID: tenant, SessionID: session.ID}); err != nil { diff --git a/services/core/internal/store/environment_initial_public_test.go b/services/core/internal/store/environment_initial_public_test.go index 82deb5708..dd490c24a 100644 --- a/services/core/internal/store/environment_initial_public_test.go +++ b/services/core/internal/store/environment_initial_public_test.go @@ -38,7 +38,6 @@ func TestEnvironmentInitialFailureOfficialClient(t *testing.T) { t.Fatal(err) } owner := executionOwner(t, db, s) - writer := owner.Store t.Cleanup(func() { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() @@ -94,7 +93,7 @@ func TestEnvironmentInitialFailureOfficialClient(t *testing.T) { if err := db.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 := writer.ExpireEnvironmentInput(t.Context(), tenant, session.ID, reservation); err != nil || result.State != sessions.EnvironmentInputExpired { + if result, err := store.SessionService(t, s).ExpireEnvironmentInput(t.Context(), tenant, session.ID, reservation); err != nil || result.State != sessions.EnvironmentInputExpired { t.Fatal("initial reservation did not expire", result, err) } observed := <-done diff --git a/services/core/internal/store/environment_input_activity.go b/services/core/internal/store/environment_input_activity.go deleted file mode 100644 index 5e32b7e0e..000000000 --- a/services/core/internal/store/environment_input_activity.go +++ /dev/null @@ -1,23 +0,0 @@ -package store - -import ( - "context" - - "github.com/jackc/pgx/v5/pgtype" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/sessionpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" -) - -// withEnvironmentInputSession runs apply in the public Session's transaction -// and reports the input activity it changes. -func (s *Store) withEnvironmentInputSession(ctx context.Context, tenantID, sessionID string, apply func(context.Context, *sqlc.Queries, pgtype.UUID) error) error { - tenant, err := parseID(tenantID) - if err != nil { - return err - } - return s.withPublicSession(ctx, tenantID, sessionID, func(ctx context.Context, q *sqlc.Queries, id pgtype.UUID) error { - return sessions.TrackInputActivity(ctx, sessionpg.BindSession(q, tenant, id), func(ctx context.Context) error { return apply(ctx, q, id) }) - }) -} diff --git a/services/core/internal/store/environment_input_activity_test.go b/services/core/internal/store/environment_input_activity_test.go index 2dd6a8cd9..b874348f9 100644 --- a/services/core/internal/store/environment_input_activity_test.go +++ b/services/core/internal/store/environment_input_activity_test.go @@ -85,7 +85,7 @@ func TestEnvironmentInputActivityWaitsBeforeTurnAndClearsOnConnection(t *testing if err := sessionExecution(t, writer.lease).ObserveEnvironmentConnection(t.Context(), tenant, environment, generation, 3, true); err != nil { t.Fatal(err) } - if _, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID); err != nil { + if _, err := sessionExecution(t, writer.lease).PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID); err != nil { t.Fatal(err) } active := requireEnvironmentInputActivity(t, s, tenant, session.ID, "", "") @@ -103,7 +103,7 @@ func TestEnvironmentInputActivitySettlementAndNewerWork(t *testing.T) { s, pool := testStore(t) tenant, session := environmentInputSession(t, s) writer := executionWriter(t, s) - prior, err := s.SubmitInputs(t.Context(), tenant, session.ID, "prior", []sessions.Input{messageInput("prior")}) + prior, err := submitInputs(t.Context(), s, tenant, session.ID, "prior", []sessions.Input{messageInput("prior")}) if err != nil { t.Fatal(err) } @@ -120,18 +120,18 @@ func TestEnvironmentInputActivitySettlementAndNewerWork(t *testing.T) { t.Fatal(err) } if state == sessions.EnvironmentInputCancelled { - _, err = s.CancelEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID) + _, err = cancelEnvironmentInput(t.Context(), s, tenant, session.ID, reservation.ID) } else { 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) } // Other retained test rows may precede this reservation in bounded batches. for reservation.State == sessions.EnvironmentInputPending { - count, sweepErr := writer.ExpireEnvironmentInputs(t.Context()) + count, sweepErr := sessionExecution(t, writer.lease).ExpireEnvironmentInputs(t.Context()) if sweepErr != nil || count < 1 || count > 32 { t.Fatal("expiry made no bounded progress", count, sweepErr) } - reservation, err = s.GetEnvironmentInputReservation(t.Context(), tenant, session.ID, reservation.ID) + reservation, err = sessionAdapter(s).GetEnvironmentInputReservation(t.Context(), tenant, session.ID, reservation.ID) if err != nil { t.Fatal(err) } @@ -155,7 +155,7 @@ func TestEnvironmentInputActivitySettlementAndNewerWork(t *testing.T) { if after, err := sessionAdapter(s).SessionEventCursor(t.Context(), tenant, session.ID); err != nil || after != cursor { t.Fatal("settled retry repeated activity", after, cursor, err) } - if _, err := s.SubmitInputs(t.Context(), tenant, session.ID, "newer", []sessions.Input{messageInput("newer")}); err != nil { + if _, err := submitInputs(t.Context(), s, tenant, session.ID, "newer", []sessions.Input{messageInput("newer")}); err != nil { t.Fatal(err) } requireEnvironmentInputActivity(t, s, tenant, session.ID, "", "") @@ -173,7 +173,7 @@ func TestEnvironmentInputActivityRollsBackReservationAndConnection(t *testing.T) t.Cleanup(func() { _, _ = pool.Exec(context.Background(), "ALTER TABLE session_events DROP CONSTRAINT IF EXISTS "+constraint) }) - if _, err := s.ReserveEnvironmentInput(t.Context(), tenant, session.ID, "rollback", []sessions.Input{messageInput("pending")}); err == nil { + if _, err := sessionService(t, s).ReserveEnvironmentInput(t.Context(), tenant, session.ID, "rollback", []sessions.Input{messageInput("pending")}); err == nil { t.Fatal("activity failure retained reservation") } var count int @@ -227,7 +227,7 @@ func TestEnvironmentInputActivityRecoversWaitingActionAndHidesDeletion(t *testin t.Fatal(err) } requireEnvironmentInputActivity(t, s, tenant, session.ID, "requires_action", environment) - got, err := s.GetEnvironmentInputReservation(t.Context(), tenant, session.ID, reservation.ID) + got, err := sessionAdapter(s).GetEnvironmentInputReservation(t.Context(), tenant, session.ID, reservation.ID) if err != nil || got.State != sessions.EnvironmentInputPending || !got.Deadline.Equal(reservation.Deadline) { t.Fatal("recovery changed waiting input or its deadline", got, err) } @@ -255,25 +255,3 @@ func TestEnvironmentInputActivityRecoversWaitingActionAndHidesDeletion(t *testin t.Fatal("deleted activity events remained visible", err) } } - -func TestPreparationFailurePreservesCancelledAndNewerInput(t *testing.T) { - s, _ := testStore(t) - tenant, session := environmentInputSession(t, s) - first := reserveEnvironmentInput(t, s, tenant, session.ID, "cancelled") - if _, err := s.CancelEnvironmentInput(t.Context(), tenant, session.ID, first.ID); err != nil { - t.Fatal(err) - } - next := reserveEnvironmentInput(t, s, tenant, session.ID, "new") - if err := s.FailEnvironmentInput(t.Context(), tenant, session.ID, first.ID, "runtime_preparation_failed"); err != nil { - t.Fatal(err) - } - for id, state := range map[string]string{first.ID: sessions.EnvironmentInputCancelled, next.ID: sessions.EnvironmentInputPending} { - current, err := s.GetEnvironmentInputReservation(t.Context(), tenant, session.ID, id) - if err != nil || current.State != state { - t.Fatal("late failure changed another outcome", err) - } - } - if err := s.FailEnvironmentInput(t.Context(), tenant, session.ID, next.ID, "secret-canary"); !errors.Is(err, sessions.ErrInvalidInput) { - t.Fatal("unclassified diagnostic accepted", err) - } -} diff --git a/services/core/internal/store/environment_input_claim_test.go b/services/core/internal/store/environment_input_claim_test.go deleted file mode 100644 index 8ce51ee8d..000000000 --- a/services/core/internal/store/environment_input_claim_test.go +++ /dev/null @@ -1,96 +0,0 @@ -package store - -import ( - "sync" - "testing" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" -) - -func TestEnvironmentInputConcurrentPromotionClaimsOnce(t *testing.T) { - s, pool := testStore(t) - writer := executionWriter(t, s) - tenant, session := environmentInputSession(t, s) - pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") - reservationCursor, err := sessionAdapter(s).SessionEventCursor(t.Context(), tenant, session.ID) - if err != nil { - t.Fatal(err) - } - const count = 8 - results := make(chan sessions.EnvironmentInputReservation, count) - var group sync.WaitGroup - for range count { - group.Go(func() { - got, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) - if err != nil { - t.Error(err) - return - } - results <- got - }) - } - group.Wait() - close(results) - var turnID string - fresh, received := 0, 0 - for got := range results { - received++ - if got.State != sessions.EnvironmentInputAdmitted || len(got.Receipts) != 2 { - t.Fatal("promotion lost the original batch", got) - } - if turnID == "" { - turnID = got.Receipts[0].TurnID - } - if !got.Receipts[0].Replayed { - fresh++ - } - for _, receipt := range got.Receipts { - if receipt.TurnID != turnID || receipt.Replayed != got.Receipts[0].Replayed { - t.Fatal("promotion changed execution ownership", got.Receipts) - } - } - } - if received != count || fresh != 1 { - t.Fatal("promotion authorized multiple starts", received, fresh) - } - turn, err := sessionAdapter(s).GetTurn(t.Context(), tenant, session.ID, turnID) - if err != nil || turn.Status != sessions.TurnInProgress || turn.StartedAt.IsZero() { - t.Fatal("promotion did not persist its execution claim", turn, err) - } - environmentInputHistory(t, pool, session.ID, 1, 2) - changes, err := sessionAdapter(s).ListSessionEvents(t.Context(), tenant, session.ID, reservationCursor) - if err != nil || len(changes) < 2 { - t.Fatal("missing promotion events", changes, err) - } - created, claimed := changes[0], changes[len(changes)-1] - if created.Event.Type != "agent.session.turn.created" || created.Turn == nil || created.Turn.Status != sessions.TurnQueued || claimed.Event.Type != "agent.session.turn.in_progress" || claimed.Turn == nil || claimed.Turn.ID != turnID || claimed.Turn.Status != sessions.TurnInProgress { - t.Fatal("claim reordered or replaced admission snapshots", created, claimed) - } - claimEvents := 0 - for _, change := range changes { - if change.Event.Type == "agent.session.turn.in_progress" { - claimEvents++ - } - } - if claimEvents != 1 { - t.Fatal("retry published another claim", claimEvents) - } - transition(t, writer, tenant, session.ID, turnID, sessions.TurnInProgress, sessions.TurnCompleted) - later := reserveEnvironmentInput(t, s, tenant, session.ID, "later") - cursor, err := sessionAdapter(s).SessionEventCursor(t.Context(), tenant, session.ID) - if err != nil { - t.Fatal(err) - } - retry, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) - if err != nil || len(retry.Receipts) != 2 || !retry.Receipts[0].Replayed || retry.Receipts[0].TurnID != turnID { - t.Fatal("terminal retry reclaimed execution", retry, err) - } - after, err := sessionAdapter(s).SessionEventCursor(t.Context(), tenant, session.ID) - if err != nil || after != cursor { - t.Fatal("terminal retry published events", after, cursor, err) - } - retained, err := s.GetEnvironmentInputReservation(t.Context(), tenant, session.ID, later.ID) - if err != nil || retained.State != sessions.EnvironmentInputPending || !retained.Deadline.Equal(later.Deadline) { - t.Fatal("old promotion affected new preparation", retained, err) - } -} diff --git a/services/core/internal/store/environment_input_expiry.go b/services/core/internal/store/environment_input_expiry.go deleted file mode 100644 index 3724937da..000000000 --- a/services/core/internal/store/environment_input_expiry.go +++ /dev/null @@ -1,43 +0,0 @@ -package store - -import ( - "context" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/sessionpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" - "github.com/jackc/pgx/v5" -) - -// 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 err := s.checkExecutionAuthority(); err != nil { - return 0, err - } - var expired int64 - 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 { - return err - } - for _, row := range rows { - err := sessions.TrackInputActivity(ctx, sessionpg.BindSession(q, row.TenantID, row.SessionID), func(ctx context.Context) error { - return q.ExpireEnvironmentInputReservation(ctx, sqlc.ExpireEnvironmentInputReservationParams{SessionID: row.SessionID, ID: row.ID}) - }) - if err != nil { - return err - } - if err := sessionpg.PruneChanges(ctx, q, row.SessionID); err != nil { - return err - } - expired++ - } - return nil - }) - if err != nil { - return 0, err - } - return expired, nil -} diff --git a/services/core/internal/store/environment_input_expiry_test.go b/services/core/internal/store/environment_input_expiry_test.go index f733964f8..4c003e0e6 100644 --- a/services/core/internal/store/environment_input_expiry_test.go +++ b/services/core/internal/store/environment_input_expiry_test.go @@ -10,86 +10,6 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" ) -func TestEnvironmentExpiryBoundsBatchAndRequiresExecutionWriter(t *testing.T) { - s, pool := testStore(t) - ctx := t.Context() - var reservations []sessions.EnvironmentInputReservation - for range 33 { - tenant, session := environmentInputSession(t, s) - pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") - reservations = append(reservations, pending) - if _, err := pool.Exec(ctx, "UPDATE environment_input_reservations SET deadline=clock_timestamp()-interval '1 second' WHERE id=$1", pending.ID); err != nil { - t.Fatal(err) - } - } - if n, err := s.ExpireEnvironmentInputs(ctx); err == nil || n != 0 { - t.Fatal("unleased maintenance", n, err) - } - var before, after, due int64 - 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) - } - writer := executionWriter(t, s) - n, err := writer.ExpireEnvironmentInputs(ctx) - if err != nil || n != 32 { - t.Fatal("unbounded or incomplete batch", n, err) - } - if err := pool.QueryRow(ctx, "SELECT count(*) FROM environment_input_reservations WHERE state='expired'").Scan(&after); err != nil || after-before != n { - t.Fatal("reported expiry did not match persisted settlement", before, after, n, err) - } - // Existing test rows may precede this fixture; each pass must make bounded progress. - for remaining := due - n; remaining > 0; { - n, err = writer.ExpireEnvironmentInputs(ctx) - if err != nil || n <= 0 || n > 32 { - t.Fatal("expiry backlog did not progress", n, err) - } - remaining -= n - } - for _, pending := range reservations { - var state string - if err := pool.QueryRow(ctx, "SELECT state FROM environment_input_reservations WHERE id=$1", pending.ID).Scan(&state); err != nil || state != sessions.EnvironmentInputExpired { - t.Fatal(state, err) - } - environmentInputHistory(t, pool, pending.SessionID, 0, 0) - } -} - -func TestEnvironmentExpiryFencesLostExecutionOwner(t *testing.T) { - s, pool := testStore(t) - tenant, session := environmentInputSession(t, s) - pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") - 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 := executionWriter(t, s) - var killed bool - 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 := 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) - if err != nil || got.State != sessions.EnvironmentInputPending { - t.Fatal("lost owner wrote through the pool", got, err) - } - for got.State == sessions.EnvironmentInputPending { - n, err := successor.ExpireEnvironmentInputs(t.Context()) - if err != nil || n == 0 { - t.Fatal("successor could not expire input", n, err) - } - got, err = s.GetEnvironmentInputReservation(t.Context(), tenant, session.ID, pending.ID) - if err != nil { - t.Fatal(err) - } - } - if got.State != sessions.EnvironmentInputExpired { - t.Fatal(got) - } - environmentInputHistory(t, pool, session.ID, 0, 0) -} - func TestEnvironmentExpirySerializesWithTargetedSettlement(t *testing.T) { for _, action := range []string{"promote", "cancel", "delete"} { t.Run(action, func(t *testing.T) { @@ -99,18 +19,22 @@ 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) } - writer := executionWriter(t, s) + operations := sessionExecution(t, executionWriter(t, s).lease) start := make(chan struct{}) results := make(chan error, 2) - go func() { <-start; _, err := writer.ExpireEnvironmentInputs(t.Context()); results <- err }() + go func() { + <-start + _, err := operations.ExpireEnvironmentInputs(t.Context()) + results <- err + }() go func() { <-start var err error switch action { case "promote": - _, err = writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) + _, err = operations.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) case "cancel": - _, err = s.CancelEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) + _, err = cancelEnvironmentInput(t.Context(), s, tenant, session.ID, pending.ID) case "delete": err = sessionService(t, s).DeleteSession(t.Context(), sessions.DeleteSessionCommand{TenantID: tenant, SessionID: session.ID}) } @@ -131,22 +55,25 @@ func TestEnvironmentExpirySerializesWithTargetedSettlement(t *testing.T) { } environmentInputHistory(t, pool, session.ID, 0, 0) if action == "delete" { - if _, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID); !errors.Is(err, sessions.ErrNotFound) { + if _, err := operations.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID); !errors.Is(err, sessions.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) (sessions.EnvironmentInputReservation, error){writer.PromoteEnvironmentInput, s.CancelEnvironmentInput, s.ExpireEnvironmentInput} { + for _, settle := range []func(context.Context, string, string, string) (sessions.EnvironmentInputReservation, error){ + operations.PromoteEnvironmentInput, + sessionService(t, 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 := writer.ExpireEnvironmentInputs(t.Context()); err != nil { + if _, err := operations.ExpireEnvironmentInputs(t.Context()); err != nil { t.Fatal(err) } - got, err := s.GetEnvironmentInputReservation(t.Context(), tenant, session.ID, later.ID) + got, err := sessionAdapter(s).GetEnvironmentInputReservation(t.Context(), tenant, session.ID, later.ID) if err != nil || got.State != sessions.EnvironmentInputPending || !got.Deadline.Equal(later.Deadline) { t.Fatal("old settlement affected successor", got, err) } diff --git a/services/core/internal/store/environment_input_migration_test.go b/services/core/internal/store/environment_input_migration_test.go index c04de8f7a..a792de72c 100644 --- a/services/core/internal/store/environment_input_migration_test.go +++ b/services/core/internal/store/environment_input_migration_test.go @@ -7,8 +7,6 @@ import ( "strings" "testing" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" "github.com/google/uuid" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/stdlib" @@ -92,37 +90,3 @@ func TestEnvironmentInputMigrationRetainsHistoryAndRetryIdentity(t *testing.T) { t.Fatal("migration manufactured a Turn", count, err) } } - -func TestEnvironmentInputPromotionUsesCurrentExecutionWriter(t *testing.T) { - s, pool := testStore(t) - tenant, session := environmentInputSession(t, s) - pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") - if _, err := s.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID); err == nil { - t.Fatal("pooled Store promoted input without execution ownership") - } - closed := executionWriter(t, s) - awaitRelease := pgtest.ObserveExecutionLeaseRelease(t, closed.pool) - if err := closed.lease.Close(t.Context()); err != nil { - t.Fatal(err) - } - awaitRelease() - 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) - writer := executionWriter(t, s) - var killed bool - 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 := 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.PromoteEnvironmentInput(t.Context(), tenant, session.ID, pending.ID) - if err != nil || got.State != sessions.EnvironmentInputAdmitted { - t.Fatal("successor could not promote", got, err) - } - environmentInputHistory(t, pool, session.ID, 1, 2) -} diff --git a/services/core/internal/store/environment_input_settlement_test.go b/services/core/internal/store/environment_input_settlement_test.go index 874d8958a..ac8ee7b74 100644 --- a/services/core/internal/store/environment_input_settlement_test.go +++ b/services/core/internal/store/environment_input_settlement_test.go @@ -3,223 +3,11 @@ package store import ( "context" "errors" - "strings" "testing" - "time" - - "github.com/google/uuid" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" ) -func TestEnvironmentInputTerminalReservationsCannotRestart(t *testing.T) { - for _, terminal := range []string{sessions.EnvironmentInputCancelled, sessions.EnvironmentInputExpired} { - t.Run(terminal, func(t *testing.T) { - s, pool := testStore(t) - writer := executionWriter(t, s) - tenant, session := environmentInputSession(t, s) - ctx := context.Background() - pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") - early, err := s.ExpireEnvironmentInput(ctx, tenant, session.ID, pending.ID) - if err != nil || early.State != sessions.EnvironmentInputPending || early.SettledAt != nil || !early.Deadline.Equal(pending.Deadline) { - t.Fatal("early expiry", early, err) - } - var settled sessions.EnvironmentInputReservation - if terminal == sessions.EnvironmentInputExpired { - if _, err := pool.Exec(ctx, "UPDATE environment_input_reservations SET deadline=clock_timestamp()-interval '1 second' WHERE id=$1", pending.ID); err != nil { - t.Fatal(err) - } - settled, err = writer.PromoteEnvironmentInput(ctx, tenant, session.ID, pending.ID) - } else { - settled, err = s.CancelEnvironmentInput(ctx, tenant, session.ID, pending.ID) - } - if err != nil || settled.State != terminal || settled.SettledAt == nil || len(settled.Receipts) != 0 { - t.Fatal("terminal settlement", settled, err) - } - retry, err := s.ReserveEnvironmentInput(ctx, tenant, session.ID, "pending", pending.Inputs) - if err != nil || retry.State != terminal || retry.ID != pending.ID || !retry.Deadline.Equal(settled.Deadline) || !retry.SettledAt.Equal(*settled.SettledAt) { - t.Fatal("terminal retry changed outcome", retry, err) - } - if _, err := s.SubmitInputs(ctx, tenant, session.ID, "pending", pending.Inputs); !errors.Is(err, sessions.ErrTurnConflict) { - t.Fatal("terminal request reopened through direct path", err) - } - if _, err := s.SubmitInputs(ctx, tenant, session.ID, "pending", []sessions.Input{messageInput("changed")}); !errors.Is(err, sessions.ErrIdempotencyConflict) { - t.Fatal("terminal identity changed", err) - } - later := reserveEnvironmentInput(t, s, tenant, session.ID, "later") - for _, finish := range []func(context.Context, string, string, string) (sessions.EnvironmentInputReservation, error){ - writer.PromoteEnvironmentInput, s.CancelEnvironmentInput, s.ExpireEnvironmentInput, - } { - got, err := finish(ctx, tenant, session.ID, pending.ID) - if err != nil || got.State != terminal { - t.Fatal("old settlement changed", got, err) - } - } - got, err := s.GetEnvironmentInputReservation(ctx, tenant, session.ID, later.ID) - if err != nil || got.State != sessions.EnvironmentInputPending || !got.Deadline.Equal(later.Deadline) { - t.Fatal("old settlement touched successor", got, err) - } - environmentInputHistory(t, pool, session.ID, 0, 0) - }) - } -} - -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) - writer := executionWriter(t, s) - tenant, session := environmentInputSession(t, s) - ctx := context.Background() - pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") - name := "reservation_failure_" + strings.ReplaceAll(uuid.NewString(), "-", "") - table := "turn_inputs" - expression := "session_id <> '" + session.ID + "'::uuid OR payload->>'text' <> 'second'" - if phase == "settlement" { - table = "environment_input_reservations" - expression = "id <> '" + pending.ID + "'::uuid OR state <> 'admitted'" - } - if phase == "claim" { - table = "turns" - expression = "session_id <> '" + session.ID + "'::uuid OR status <> 'in_progress'" - } - if phase == "claim-event" { - table = "session_events" - expression = "session_id <> '" + session.ID + "'::uuid OR payload->'event'->>'type' <> 'agent.session.turn.in_progress'" - } - if _, err := pool.Exec(ctx, "ALTER TABLE "+table+" ADD CONSTRAINT "+name+" CHECK ("+expression+") NOT VALID"); err != nil { - t.Fatal(err) - } - t.Cleanup(func() { _, _ = pool.Exec(ctx, "ALTER TABLE "+table+" DROP CONSTRAINT IF EXISTS "+name) }) - if _, err := writer.PromoteEnvironmentInput(ctx, tenant, session.ID, pending.ID); err == nil { - t.Fatal("injected failure succeeded") - } - environmentInputHistory(t, pool, session.ID, 0, 0) - got, err := s.GetEnvironmentInputReservation(ctx, tenant, session.ID, pending.ID) - if err != nil || got.State != sessions.EnvironmentInputPending || got.SettledAt != nil || !got.Deadline.Equal(pending.Deadline) { - t.Fatal("partial settlement survived", got, err) - } - if _, err := pool.Exec(ctx, "ALTER TABLE "+table+" DROP CONSTRAINT "+name); err != nil { - t.Fatal(err) - } - got, err = writer.PromoteEnvironmentInput(ctx, tenant, session.ID, pending.ID) - if err != nil || got.State != sessions.EnvironmentInputAdmitted { - t.Fatal(got, err) - } - environmentInputHistory(t, pool, session.ID, 1, 2) - }) - } -} - -func TestEnvironmentInputDeadlineIsCheckedAfterSessionLock(t *testing.T) { - for _, action := range []string{"promote", "fail"} { - t.Run(action, func(t *testing.T) { - s, pool := testStore(t) - 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) - defer cancel() - tx, err := pool.Begin(ctx) - if err != nil { - t.Fatal(err) - } - defer func() { _ = tx.Rollback(context.Background()) }() - var blocker int32 - if err := tx.QueryRow(ctx, "SELECT pg_backend_pid() FROM sessions WHERE id=$1 FOR UPDATE", session.ID).Scan(&blocker); err != nil { - t.Fatal(err) - } - type outcome struct { - value sessions.EnvironmentInputReservation - err error - } - done := make(chan outcome, 1) - go func() { - var got sessions.EnvironmentInputReservation - var err error - if action == "promote" { - got, err = writer.PromoteEnvironmentInput(ctx, tenant, session.ID, pending.ID) - } else { - err = writer.FailEnvironmentInput(ctx, tenant, session.ID, pending.ID, "runtime_preparation_failed") - if err == nil { - got, err = s.GetEnvironmentInputReservation(ctx, tenant, session.ID, pending.ID) - } - } - done <- outcome{got, err} - }() - for { - var blocked bool - if err := pool.QueryRow(ctx, "SELECT EXISTS (SELECT 1 FROM pg_stat_activity WHERE $1=ANY(pg_blocking_pids(pid)))", blocker).Scan(&blocked); err != nil { - t.Fatal(err) - } - if blocked { - break - } - select { - case result := <-done: - t.Fatal("settlement bypassed Session lock", result) - case <-ctx.Done(): - t.Fatal("settlement lock wait not observed") - case <-time.After(5 * time.Millisecond): - } - } - // Transaction-start time is now older than the controlled deadline. - if _, err := tx.Exec(ctx, "UPDATE environment_input_reservations SET deadline=clock_timestamp() WHERE id=$1", pending.ID); err != nil { - t.Fatal(err) - } - if err := tx.Commit(ctx); err != nil { - t.Fatal(err) - } - result := <-done - if result.err != nil || result.value.State != sessions.EnvironmentInputExpired || result.value.SettledAt == nil { - t.Fatal("lock wait extended input lifetime", result) - } - environmentInputHistory(t, pool, session.ID, 0, 0) - stored, err := s.GetEnvironmentInputReservation(ctx, tenant, session.ID, pending.ID) - if err != nil || stored.State != sessions.EnvironmentInputExpired { - t.Fatal("expiry was rolled back", stored, err) - } - }) - } -} - -func TestEnvironmentInputCancelAndPromotionShareOneOutcome(t *testing.T) { - s, pool := testStore(t) - writer := executionWriter(t, s) - other, _ := testStore(t) - tenant, session := environmentInputSession(t, s) - ctx := context.Background() - pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") - start := make(chan struct{}) - results := make(chan sessions.EnvironmentInputReservation, 2) - errs := make(chan error, 2) - for _, finish := range []func(context.Context, string, string, string) (sessions.EnvironmentInputReservation, error){ - writer.PromoteEnvironmentInput, other.CancelEnvironmentInput, - } { - go func() { - <-start - got, err := finish(ctx, tenant, session.ID, pending.ID) - results <- got - errs <- err - }() - } - close(start) - first, second := <-results, <-results - for range 2 { - if err := <-errs; err != nil { - t.Fatal(err) - } - } - if first.State != second.State || (first.State != sessions.EnvironmentInputAdmitted && first.State != sessions.EnvironmentInputCancelled) { - t.Fatal("competing settlements diverged", first.State, second.State) - } - turns, inputs := 0, 0 - if first.State == sessions.EnvironmentInputAdmitted { - turns, inputs = 1, 2 - } - environmentInputHistory(t, pool, session.ID, turns, inputs) -} - 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) { @@ -231,7 +19,7 @@ func TestEnvironmentInputDeletionSettlesPendingAndFencesPromotion(t *testing.T) done := make(chan error, 1) if concurrent { go func() { - _, err := writer.PromoteEnvironmentInput(ctx, tenant, session.ID, pending.ID) + _, err := sessionExecution(t, writer.lease).PromoteEnvironmentInput(ctx, tenant, session.ID, pending.ID) done <- err }() } @@ -250,13 +38,17 @@ func TestEnvironmentInputDeletionSettlesPendingAndFencesPromotion(t *testing.T) } } for _, action := range []func(context.Context, string, string, string) (sessions.EnvironmentInputReservation, error){ - s.GetEnvironmentInputReservation, writer.PromoteEnvironmentInput, s.CancelEnvironmentInput, + sessionAdapter(s).GetEnvironmentInputReservation, + sessionExecution(t, writer.lease).PromoteEnvironmentInput, + func(ctx context.Context, tenant, session, id string) (sessions.EnvironmentInputReservation, error) { + return cancelEnvironmentInput(ctx, s, tenant, session, id) + }, } { if _, err := action(ctx, tenant, session.ID, pending.ID); !errors.Is(err, sessions.ErrNotFound) { t.Fatal("deleted reservation remained accessible", err) } } - if _, err := s.ReserveEnvironmentInput(ctx, tenant, session.ID, "late", pending.Inputs); !errors.Is(err, sessions.ErrNotFound) { + if _, err := sessionService(t, s).ReserveEnvironmentInput(ctx, tenant, session.ID, "late", pending.Inputs); !errors.Is(err, sessions.ErrNotFound) { t.Fatal("deleted Session accepted reservation", err) } var state string diff --git a/services/core/internal/store/environment_inputs.go b/services/core/internal/store/environment_inputs.go deleted file mode 100644 index c2d2d4b7f..000000000 --- a/services/core/internal/store/environment_inputs.go +++ /dev/null @@ -1,287 +0,0 @@ -package store - -import ( - "context" - "encoding/json" - "errors" - "fmt" - - "github.com/google/uuid" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgtype" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/sessionpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" -) - -// ReserveEnvironmentInput appends to active work or reserves an idle message batch. -// The Session lock decides both paths; only promotion can create a new Turn. -func (s *Store) ReserveEnvironmentInput(ctx context.Context, tenantID, sessionID, key string, inputs []sessions.Input) (sessions.EnvironmentInputReservation, error) { - if err := sessions.ValidateInputKey(key); err != nil { - return sessions.EnvironmentInputReservation{}, err - } - batch, encoded, err := validateInitialInputs(inputs) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - tenant, err := parseID(tenantID) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - var result sessions.EnvironmentInputReservation - err = s.withEnvironmentInputSession(ctx, tenantID, sessionID, func(ctx context.Context, q *sqlc.Queries, session pgtype.UUID) error { - audit := func() error { - return auditpg.RecordWriteAudit(ctx, q, tenantID, "send_events", "session", uuid.UUID(session.Bytes).String(), "") - } - previous, err := q.FindEnvironmentInputReservation(ctx, sqlc.FindEnvironmentInputReservationParams{ - SessionID: session, IdempotencyKey: key, Batch: encoded, - }) - if err == nil { - if !previous.Matches { - return sessions.ErrIdempotencyConflict - } - result, err = settleEnvironmentInput(ctx, q, tenantID, previous.EnvironmentInputReservation, sessions.EnvironmentInputExpired) - if err == nil && (result.State == sessions.EnvironmentInputPending || result.State == sessions.EnvironmentInputAdmitted) { - return audit() - } - return err - } - if !errors.Is(err, pgx.ErrNoRows) { - return err - } - receipts, err := inputBatchReceipts(ctx, q, session, key, encoded) - if err != nil { - return err - } - if len(receipts) > 0 { - // Earlier direct admission has receipts, but never had a reservation or deadline. - result = sessions.EnvironmentInputReservation{SessionID: sessionID, State: sessions.EnvironmentInputAdmitted, Receipts: receipts} - return audit() - } - if err := sessions.CheckFileWriteGate(ctx, sessionpg.BindSession(q, tenant, session)); err != nil { - return err - } - environment, err := q.GetSessionEnvironment(ctx, sqlc.GetSessionEnvironmentParams{TenantID: tenant, ID: session}) - if errors.Is(err, pgx.ErrNoRows) { - return sessions.ErrInvalidInput - } else if err != nil { - return err - } - if environment.Environment.Status == "failed" { - if kind, err := sessions.EnvironmentType(environment.Configuration); err == nil && kind == "openai_hosted" { - return sessions.ErrHostedEnvironmentFailed - } - } - if environment.Environment.Status == "failed" || environment.Environment.Status == "expired" { - return sessions.ErrEnvironmentUnavailable - } - if err := checkEnvironmentInputGate(ctx, q, session, key, encoded); err != nil { - return err - } - if active, err := q.GetActiveTurn(ctx, session); err == nil && !active.ArtifactCaptureStarted { - result = sessions.EnvironmentInputReservation{SessionID: sessionID, State: sessions.EnvironmentInputAdmitted} - for position, input := range batch { - receipt, err := admitInput(ctx, q, tenantID, session, key, int32(position), input) - if err != nil { - return err - } - result.Receipts = append(result.Receipts, receipt) - } - return audit() - } else if err != nil && !errors.Is(err, pgx.ErrNoRows) { - return err - } - row, err := q.CreateEnvironmentInputReservation(ctx, sqlc.CreateEnvironmentInputReservationParams{ - ID: pgtype.UUID{Bytes: uuid.New(), Valid: true}, SessionID: session, IdempotencyKey: key, Batch: encoded, - }) - if err != nil { - return err - } - result, err = environmentInputFromRow(row) - if err != nil { - return err - } - return audit() - }) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - return result, nil -} - -func (s *Store) GetEnvironmentInputReservation(ctx context.Context, tenantID, sessionID, reservationID string) (sessions.EnvironmentInputReservation, error) { - id, err := parseID(reservationID) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - var result sessions.EnvironmentInputReservation - err = s.withPublicSession(ctx, tenantID, sessionID, func(ctx context.Context, q *sqlc.Queries, session pgtype.UUID) error { - row, err := q.GetEnvironmentInputReservation(ctx, sqlc.GetEnvironmentInputReservationParams{SessionID: session, ID: id}) - if errors.Is(err, pgx.ErrNoRows) { - return sessions.ErrNotFound - } - if err != nil { - return err - } - result, err = environmentInputOutcome(ctx, q, row) - return err - }) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - return result, nil -} - -// PromoteEnvironmentInput admits and claims work for the retained native preparation. -func (s *Store) PromoteEnvironmentInput(ctx context.Context, tenantID, sessionID, reservationID string) (sessions.EnvironmentInputReservation, error) { - if err := s.checkExecutionAuthority(); err != nil { - return sessions.EnvironmentInputReservation{}, err - } - return s.settleEnvironmentInput(ctx, tenantID, sessionID, reservationID, sessions.EnvironmentInputAdmitted) -} - -func (s *Store) CancelEnvironmentInput(ctx context.Context, tenantID, sessionID, reservationID string) (sessions.EnvironmentInputReservation, error) { - return s.settleEnvironmentInput(ctx, tenantID, sessionID, reservationID, sessions.EnvironmentInputCancelled) -} - -// FailEnvironmentInput settles a confirmed pre-admission failure. The Session -// lock and pending-state predicate preserve cancellation and newer input. -func (s *Store) FailEnvironmentInput(ctx context.Context, tenantID, sessionID, reservationID, code string) error { - if code != "model_provider_required" && code != "runtime_preparation_failed" { - return sessions.ErrInvalidInput - } - id, err := parseID(reservationID) - if err != nil { - return err - } - return s.withEnvironmentInputSession(ctx, tenantID, sessionID, func(ctx context.Context, q *sqlc.Queries, session pgtype.UUID) error { - if err := q.ExpireEnvironmentInputReservation(ctx, sqlc.ExpireEnvironmentInputReservationParams{SessionID: session, ID: id}); err != nil { - return err - } - _, err := q.FailEnvironmentInput(ctx, sqlc.FailEnvironmentInputParams{SessionID: session, ID: id, FailureCode: pgtype.Text{String: code, Valid: true}}) - return err - }) -} - -func (s *Store) ExpireEnvironmentInput(ctx context.Context, tenantID, sessionID, reservationID string) (sessions.EnvironmentInputReservation, error) { - return s.settleEnvironmentInput(ctx, tenantID, sessionID, reservationID, sessions.EnvironmentInputExpired) -} - -func (s *Store) settleEnvironmentInput(ctx context.Context, tenantID, sessionID, reservationID, state string) (sessions.EnvironmentInputReservation, error) { - id, err := parseID(reservationID) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - var result sessions.EnvironmentInputReservation - err = s.withEnvironmentInputSession(ctx, tenantID, sessionID, func(ctx context.Context, q *sqlc.Queries, session pgtype.UUID) error { - row, err := q.GetEnvironmentInputReservation(ctx, sqlc.GetEnvironmentInputReservationParams{SessionID: session, ID: id}) - if errors.Is(err, pgx.ErrNoRows) { - return sessions.ErrNotFound - } - if err != nil { - return err - } - result, err = settleEnvironmentInput(ctx, q, tenantID, row, state) - return err - }) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - return result, nil -} - -func settleEnvironmentInput(ctx context.Context, q *sqlc.Queries, tenantID string, row sqlc.EnvironmentInputReservation, state string) (sessions.EnvironmentInputReservation, error) { - if row.State != sessions.EnvironmentInputPending { - return environmentInputOutcome(ctx, q, row) - } - if err := q.ExpireEnvironmentInputReservation(ctx, sqlc.ExpireEnvironmentInputReservationParams{SessionID: row.SessionID, ID: row.ID}); err != nil { - return sessions.EnvironmentInputReservation{}, err - } - row, err := q.GetEnvironmentInputReservation(ctx, sqlc.GetEnvironmentInputReservationParams{SessionID: row.SessionID, ID: row.ID}) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - // Terminal outcomes are successful storage results so settlement is not rolled back. - if row.State != sessions.EnvironmentInputPending || state == sessions.EnvironmentInputExpired { - return environmentInputOutcome(ctx, q, row) - } - result, err := environmentInputFromRow(row) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - if state == sessions.EnvironmentInputAdmitted { - tenant, err := parseID(tenantID) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - if err := sessions.CheckInputStart(ctx, sessionpg.BindSession(q, tenant, row.SessionID)); err != nil { - return sessions.EnvironmentInputReservation{}, err - } - for position, input := range result.Inputs { - receipt, err := admitInput(ctx, q, tenantID, row.SessionID, row.IdempotencyKey, int32(position), input) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - result.Receipts = append(result.Receipts, receipt) - } - } - row, err = q.SettleEnvironmentInputReservation(ctx, sqlc.SettleEnvironmentInputReservationParams{SessionID: row.SessionID, ID: row.ID, State: state}) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - if state == sessions.EnvironmentInputAdmitted { - tenant, err := parseID(tenantID) - if err != nil { - return sessions.EnvironmentInputReservation{}, err - } - start := sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress} - if _, err := sessions.TransitionTurn(ctx, sessionpg.BindSession(q, tenant, row.SessionID), result.Receipts[0].TurnID, start); err != nil { - return sessions.EnvironmentInputReservation{}, err - } - } - result.State = row.State - result.SettledAt = &row.SettledAt.Time - return result, nil -} - -func environmentInputOutcome(ctx context.Context, q *sqlc.Queries, row sqlc.EnvironmentInputReservation) (sessions.EnvironmentInputReservation, error) { - result, err := environmentInputFromRow(row) - if err != nil || row.State != sessions.EnvironmentInputAdmitted { - return result, err - } - result.Receipts, err = inputBatchReceipts(ctx, q, row.SessionID, row.IdempotencyKey, row.Batch) - if err == nil && len(result.Receipts) == 0 { - err = errors.New("admitted Environment input has no receipts") - } - return result, err -} - -func environmentInputFromRow(row sqlc.EnvironmentInputReservation) (sessions.EnvironmentInputReservation, error) { - result := sessions.EnvironmentInputReservation{ - ID: uuid.UUID(row.ID.Bytes).String(), SessionID: uuid.UUID(row.SessionID.Bytes).String(), - State: row.State, IsInitial: row.IsInitial, CreatedAt: row.CreatedAt.Time, Deadline: row.Deadline.Time, - } - if row.SettledAt.Valid { - result.SettledAt = &row.SettledAt.Time - } - if err := json.Unmarshal(row.Batch, &result.Inputs); err != nil { - return sessions.EnvironmentInputReservation{}, fmt.Errorf("decode Environment input: %w", err) - } - return result, nil -} - -func checkEnvironmentInputGate(ctx context.Context, q *sqlc.Queries, session pgtype.UUID, key string, batch json.RawMessage) error { - gate, err := q.CheckEnvironmentInputGate(ctx, sqlc.CheckEnvironmentInputGateParams{SessionID: session, IdempotencyKey: key, Batch: batch}) - if err != nil { - return err - } - if !gate.Matches { - return sessions.ErrIdempotencyConflict - } - if gate.Blocked { - return sessions.ErrInputPending - } - return nil -} diff --git a/services/core/internal/store/environment_inputs_test.go b/services/core/internal/store/environment_inputs_test.go deleted file mode 100644 index 2b2074e70..000000000 --- a/services/core/internal/store/environment_inputs_test.go +++ /dev/null @@ -1,274 +0,0 @@ -package store - -import ( - "context" - "encoding/json" - "errors" - "reflect" - "sync" - "testing" - "time" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgtest" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" - "github.com/google/uuid" - "github.com/jackc/pgx/v5/pgxpool" -) - -func environmentInputSession(t *testing.T, s *Store) (string, sessions.Session) { - t.Helper() - tenant := uuid.NewString() - session, err := s.CreateSession(context.Background(), tenant, environmentInput("session", "self_hosted", "/workspace")) - if err != nil { - t.Fatal(err) - } - return tenant, session -} - -func reserveEnvironmentInput(t *testing.T, s *Store, tenant, session, key string) sessions.EnvironmentInputReservation { - t.Helper() - got, err := s.ReserveEnvironmentInput(context.Background(), tenant, session, key, []sessions.Input{messageInput("first"), messageInput("second")}) - if err != nil { - t.Fatal(err) - } - return got -} - -func environmentInputHistory(t *testing.T, pool *pgxpool.Pool, session string, turns, inputs int) { - t.Helper() - var gotTurns, gotInputs, items, events int - err := pool.QueryRow(context.Background(), ` - SELECT (SELECT count(*) FROM turns WHERE session_id=$1), - (SELECT count(*) FROM turn_inputs WHERE session_id=$1), - (SELECT count(*) FROM session_items WHERE session_id=$1), - (SELECT count(*) FROM session_events WHERE session_id=$1 - AND (payload ? 'turn' OR payload->'event' ? 'item'))`, session).Scan(&gotTurns, &gotInputs, &items, &events) - if err != nil || gotTurns != turns || gotInputs != inputs || items != inputs || (inputs == 0 && events != 0) { - t.Fatal("history", gotTurns, gotInputs, items, events, err) - } -} - -func TestEnvironmentInputReservationConcurrentIdentity(t *testing.T) { - s, pool := testStore(t) - other, _ := testStore(t) - tenant, session := environmentInputSession(t, s) - ctx := context.Background() - batch := []sessions.Input{ - {Kind: "message", Payload: json.RawMessage(`{"text":"first","detail":{"a":1,"b":2}}`)}, - messageInput("second"), - } - const count = 8 - results := make(chan sessions.EnvironmentInputReservation, count) - var wg sync.WaitGroup - for i := range count { - wg.Add(1) - go func() { - defer wg.Done() - st := s - inputs := append([]sessions.Input(nil), batch...) - if i%2 == 0 { - st = other - inputs[0].Payload = json.RawMessage(` { "detail": {"b": 2, "a": 1}, "text": "first" } `) - } - got, err := st.ReserveEnvironmentInput(ctx, tenant, session.ID, "request", inputs) - if err != nil { - t.Error(err) - return - } - results <- got - }() - } - wg.Wait() - close(results) - var first sessions.EnvironmentInputReservation - received := 0 - for result := range results { - received++ - if first.ID == "" { - first = result - } - if !reflect.DeepEqual(first, result) { - t.Fatal("reservation identity changed", first, result) - } - } - if received != count || first.State != sessions.EnvironmentInputPending || first.ID == "" || first.Deadline.Sub(first.CreatedAt) != 5*time.Minute || first.SettledAt != nil || len(first.Receipts) != 0 { - t.Fatal("invalid pending result", received, first) - } - environmentInputHistory(t, pool, session.ID, 0, 0) - for _, changed := range [][]sessions.Input{batch[:1], {batch[1], batch[0]}, {messageInput("changed"), batch[1]}} { - if _, err := s.ReserveEnvironmentInput(ctx, tenant, session.ID, "request", changed); !errors.Is(err, sessions.ErrIdempotencyConflict) { - t.Fatal("changed request accepted", err) - } - } - if _, err := other.ReserveEnvironmentInput(ctx, tenant, session.ID, "other", batch); !errors.Is(err, sessions.ErrTurnConflict) { - t.Fatal("second pending request accepted", err) - } - pool.Close() - restarted, _ := testStore(t) - got, err := restarted.ReserveEnvironmentInput(ctx, tenant, session.ID, "request", batch) - if err != nil || !reflect.DeepEqual(first, got) { - t.Fatal("restart changed deadline or identity", got, err) - } -} - -func TestEnvironmentInputReservationPromotionAndDirectRetries(t *testing.T) { - s, pool := testStore(t) - writer := executionWriter(t, s) - tenant, session := environmentInputSession(t, s) - ctx := context.Background() - first := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") - for _, request := range []struct { - key string - inputs []sessions.Input - want error - }{ - {"pending", first.Inputs, sessions.ErrTurnConflict}, - {"pending", []sessions.Input{messageInput("changed")}, sessions.ErrIdempotencyConflict}, - {"later", []sessions.Input{messageInput("later")}, sessions.ErrTurnConflict}, - {"cancel", []sessions.Input{{Kind: "cancel", Payload: json.RawMessage(`{}`)}}, sessions.ErrTurnConflict}, - } { - if _, err := s.SubmitInputs(ctx, tenant, session.ID, request.key, request.inputs); !errors.Is(err, request.want) { - t.Fatal("direct path bypassed reservation", request.key, err) - } - } - environmentInputHistory(t, pool, session.ID, 0, 0) - promoted, err := writer.PromoteEnvironmentInput(ctx, tenant, session.ID, first.ID) - if err != nil || promoted.State != sessions.EnvironmentInputAdmitted || promoted.SettledAt == nil || len(promoted.Receipts) != 2 || !promoted.Deadline.Equal(first.Deadline) { - t.Fatal(promoted, err) - } - for i, receipt := range promoted.Receipts { - if receipt.Replayed || receipt.TurnID == "" || receipt.TurnID != promoted.Receipts[0].TurnID || (i > 0 && receipt.Sequence <= promoted.Receipts[i-1].Sequence) { - t.Fatal("promotion receipts", promoted.Receipts) - } - } - environmentInputHistory(t, pool, session.ID, 1, 2) - for _, read := range []func() (sessions.EnvironmentInputReservation, error){ - func() (sessions.EnvironmentInputReservation, error) { - return writer.PromoteEnvironmentInput(ctx, tenant, session.ID, first.ID) - }, - func() (sessions.EnvironmentInputReservation, error) { - return s.GetEnvironmentInputReservation(ctx, tenant, session.ID, first.ID) - }, - func() (sessions.EnvironmentInputReservation, error) { - return s.ReserveEnvironmentInput(ctx, tenant, session.ID, "pending", first.Inputs) - }, - } { - retry, err := read() - if err != nil || retry.ID != first.ID || !retry.Deadline.Equal(first.Deadline) || retry.State != sessions.EnvironmentInputAdmitted || len(retry.Receipts) != 2 { - t.Fatal(retry, err) - } - for i, receipt := range retry.Receipts { - if !receipt.Replayed || receipt.Sequence != promoted.Receipts[i].Sequence || receipt.TurnID != promoted.Receipts[i].TurnID { - t.Fatal("retry changed admission", receipt) - } - } - } - retry, err := s.SubmitInputs(ctx, tenant, session.ID, "pending", first.Inputs) - if err != nil || len(retry) != 2 || !retry[0].Replayed || retry[0].Sequence != promoted.Receipts[0].Sequence { - t.Fatal("direct retry after promotion", retry, err) - } - awaitRelease := pgtest.ObserveExecutionLeaseRelease(t, writer.pool) - if err := writer.lease.Close(ctx); err != nil { - t.Fatal(err) - } - awaitRelease() - pool.Close() - restarted, pool := testStore(t) - after, err := executionWriter(t, restarted).PromoteEnvironmentInput(ctx, tenant, session.ID, first.ID) - if err != nil || after.State != sessions.EnvironmentInputAdmitted || after.Receipts[0].Sequence != promoted.Receipts[0].Sequence { - t.Fatal("restart repeated promotion", after, err) - } - environmentInputHistory(t, pool, session.ID, 1, 2) -} - -func TestEnvironmentInputReservationKeepsEarlierDirectIdentity(t *testing.T) { - s, _ := testStore(t) - tenant, session := environmentInputSession(t, s) - ctx := context.Background() - input := messageInput("already admitted") - receipts, err := s.SubmitInputs(ctx, tenant, session.ID, "direct", []sessions.Input{input}) - if err != nil { - t.Fatal(err) - } - transition(t, s, tenant, session.ID, receipts[0].TurnID, sessions.TurnQueued, sessions.TurnInProgress) - transition(t, s, tenant, session.ID, receipts[0].TurnID, sessions.TurnInProgress, sessions.TurnCompleted) - pending := reserveEnvironmentInput(t, s, tenant, session.ID, "new") - got, err := s.ReserveEnvironmentInput(ctx, tenant, session.ID, "direct", []sessions.Input{input}) - if err != nil || got.State != sessions.EnvironmentInputAdmitted || got.ID != "" || !got.Deadline.IsZero() || len(got.Receipts) != 1 || got.Receipts[0].Sequence != receipts[0].Sequence { - t.Fatal("direct admission gained a reservation", got, err) - } - if _, err := s.ReserveEnvironmentInput(ctx, tenant, session.ID, "direct", []sessions.Input{messageInput("changed")}); !errors.Is(err, sessions.ErrIdempotencyConflict) { - t.Fatal(err) - } - retry, err := s.SubmitInputs(ctx, tenant, session.ID, "direct", []sessions.Input{input}) - if err != nil || len(retry) != 1 || !retry[0].Replayed { - t.Fatal(retry, err) - } - retained, err := s.GetEnvironmentInputReservation(ctx, tenant, session.ID, pending.ID) - if err != nil || !reflect.DeepEqual(retained, pending) { - t.Fatal("old retry affected new pending input", retained, err) - } -} - -func TestEnvironmentInputReservationRejectsUnsupportedOrForeignState(t *testing.T) { - s, _ := testStore(t) - writer := executionWriter(t, s) - tenant, session := environmentInputSession(t, s) - ctx := context.Background() - pending := reserveEnvironmentInput(t, s, tenant, session.ID, "pending") - for _, read := range []func() error{ - func() error { - _, err := s.GetEnvironmentInputReservation(ctx, uuid.NewString(), session.ID, pending.ID) - return err - }, - func() error { - _, err := writer.PromoteEnvironmentInput(ctx, uuid.NewString(), session.ID, pending.ID) - return err - }, - func() error { - _, err := s.CancelEnvironmentInput(ctx, uuid.NewString(), session.ID, pending.ID) - return err - }, - func() error { - _, err := s.ReserveEnvironmentInput(ctx, uuid.NewString(), session.ID, "new", pending.Inputs) - return err - }, - } { - if err := read(); !errors.Is(err, sessions.ErrNotFound) { - t.Fatal("foreign access", err) - } - } - other, err := s.CreateSession(ctx, tenant, environmentInput("other", "self_hosted", "/workspace")) - if err != nil { - t.Fatal(err) - } - for _, action := range []func(context.Context, string, string, string) (sessions.EnvironmentInputReservation, error){ - s.GetEnvironmentInputReservation, writer.PromoteEnvironmentInput, s.CancelEnvironmentInput, s.ExpireEnvironmentInput, - } { - if _, err := action(ctx, tenant, other.ID, pending.ID); !errors.Is(err, sessions.ErrNotFound) { - t.Fatal("reservation crossed Session ownership", err) - } - if _, err := action(ctx, uuid.NewString(), session.ID, pending.ID); !errors.Is(err, sessions.ErrNotFound) { - t.Fatal("reservation crossed tenant ownership", err) - } - } - retained, err := s.GetEnvironmentInputReservation(ctx, tenant, session.ID, pending.ID) - if err != nil || !reflect.DeepEqual(retained, pending) { - t.Fatal("foreign operations changed reservation", retained, err) - } - for _, invalid := range [][]sessions.Input{nil, {{Kind: "cancel", Payload: json.RawMessage(`{}`)}}, {{Kind: "tool_result", Payload: json.RawMessage(`{}`)}}} { - if _, err := s.ReserveEnvironmentInput(ctx, tenant, session.ID, "invalid", invalid); !errors.Is(err, sessions.ErrInvalidInput) { - t.Fatal("unsupported reservation", err) - } - } - noneTenant, none := newTurnSession(t, s) - if _, err := s.ReserveEnvironmentInput(ctx, noneTenant, none.ID, "none", pending.Inputs); !errors.Is(err, sessions.ErrInvalidInput) { - t.Fatal("none reservation", err) - } - activeTenant, active := environmentInputSession(t, s) - activeInput := submitMessage(t, s, activeTenant, active.ID, "active") - steer, err := s.ReserveEnvironmentInput(ctx, activeTenant, active.ID, "new", pending.Inputs) - if err != nil || steer.State != sessions.EnvironmentInputAdmitted || steer.ID != "" || !steer.Deadline.IsZero() || len(steer.Receipts) != len(pending.Inputs) || steer.Receipts[0].TurnID != activeInput.TurnID { - t.Fatal("active input did not retain the existing Turn", steer, err) - } -} diff --git a/services/core/internal/store/environment_steering_order_test.go b/services/core/internal/store/environment_steering_order_test.go deleted file mode 100644 index 94cfcf094..000000000 --- a/services/core/internal/store/environment_steering_order_test.go +++ /dev/null @@ -1,104 +0,0 @@ -package store - -import ( - "context" - "errors" - "testing" - "time" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" -) - -func TestEnvironmentActiveInputSerializesWithCompletion(t *testing.T) { - for _, completionFirst := range []bool{false, true} { - name := "input-first" - if completionFirst { - name = "completion-first" - } - t.Run(name, func(t *testing.T) { - s, pool := testStore(t) - 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, sessions.TurnQueued, sessions.TurnInProgress) - ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) - defer cancel() - blocker, err := pool.Begin(ctx) - if err != nil { - t.Fatal(err) - } - defer func() { _ = blocker.Rollback(context.Background()) }() - var blockerPID int32 - if err := blocker.QueryRow(ctx, "SELECT pg_backend_pid() FROM sessions WHERE id=$1 FOR UPDATE", session.ID).Scan(&blockerPID); err != nil { - t.Fatal(err) - } - waiter := func(pid int32) int32 { - t.Helper() - for ctx.Err() == nil { - var blocked int32 - if err := pool.QueryRow(ctx, "SELECT pid FROM pg_stat_activity WHERE $1=ANY(pg_blocking_pids(pid)) ORDER BY pid LIMIT 1", pid).Scan(&blocked); err == nil { - return blocked - } - time.Sleep(5 * time.Millisecond) - } - t.Fatal("transaction did not wait for the Session lock") - return 0 - } - type admission struct { - value sessions.EnvironmentInputReservation - err error - } - admitted := make(chan admission, 1) - completed := make(chan error, 1) - batch := []sessions.Input{messageInput("first"), messageInput("second")} - input := func() { - value, err := s.ReserveEnvironmentInput(ctx, tenant, session.ID, "racing-input", batch) - admitted <- admission{value, err} - } - complete := func() { - _, err := completeExecution(ctx, t, writer, tenant, session.ID, original.TurnID, sessions.TurnCompleted, nil, "", original.Sequence) - completed <- err - } - first, second := input, complete - if completionFirst { - first, second = complete, input - } - go first() - firstPID := waiter(blockerPID) - go second() - waiter(firstPID) - if err := blocker.Commit(ctx); err != nil { - t.Fatal(err) - } - got, completionErr := <-admitted, <-completed - if got.err != nil { - t.Fatal(got.err) - } - if completionFirst { - if completionErr != nil || got.value.State != sessions.EnvironmentInputPending || got.value.ID == "" || len(got.value.Receipts) != 0 || got.value.Deadline.Sub(got.value.CreatedAt) != 5*time.Minute { - t.Fatal("completion winner did not leave new input waiting for preparation", completionErr, got.value) - } - environmentInputHistory(t, pool, session.ID, 1, 1) - prepared, err := writer.PromoteEnvironmentInput(ctx, tenant, session.ID, got.value.ID) - if err != nil || len(prepared.Receipts) != 2 || prepared.Receipts[0].TurnID == original.TurnID { - t.Fatal("prepared successor reused terminal work", err) - } - retry, err := s.ReserveEnvironmentInput(ctx, tenant, session.ID, "racing-input", batch) - if err != nil || retry.ID != got.value.ID || !retry.Deadline.Equal(got.value.Deadline) || len(retry.Receipts) != 2 || !retry.Receipts[0].Replayed || retry.Receipts[0].TurnID != prepared.Receipts[0].TurnID { - t.Fatal("active retry replaced its original reservation", retry, err) - } - if _, err := completeExecution(ctx, t, writer, tenant, session.ID, prepared.Receipts[0].TurnID, sessions.TurnCompleted, nil, "", prepared.Receipts[1].Sequence); err != nil { - t.Fatal(err) - } - } else { - if !errors.Is(completionErr, sessions.ErrUnappliedInputs) || got.value.State != sessions.EnvironmentInputAdmitted || got.value.ID != "" || !got.value.Deadline.IsZero() || len(got.value.Receipts) != 2 || got.value.Receipts[0].TurnID != original.TurnID || got.value.Receipts[0].Replayed { - t.Fatal("admitted input escaped the original Turn or application fence", completionErr, got.value) - } - environmentInputHistory(t, pool, session.ID, 1, 3) - if _, err := completeExecution(ctx, t, writer, tenant, session.ID, original.TurnID, sessions.TurnCompleted, nil, "", got.value.Receipts[1].Sequence); err != nil { - t.Fatal("completion after controlled application failed", err) - } - } - }) - } -} diff --git a/services/core/internal/store/environment_work_test.go b/services/core/internal/store/environment_work_test.go index 5085d3585..0c07a6321 100644 --- a/services/core/internal/store/environment_work_test.go +++ b/services/core/internal/store/environment_work_test.go @@ -23,7 +23,7 @@ func TestEnvironmentInputWorkFiltersAndPagesDevices(t *testing.T) { case "expired": makeEnvironmentExpiryDue(t, pool, &pending) case "cancelled": - if _, err := h.s.CancelEnvironmentInput(t.Context(), h.tenant, pending.SessionID, pending.ID); err != nil { + if _, err := store.CancelEnvironmentInput(t.Context(), h.s, h.tenant, pending.SessionID, pending.ID); err != nil { t.Fatal(err) } case "deleted": @@ -44,7 +44,7 @@ func TestEnvironmentInputWorkFiltersAndPagesDevices(t *testing.T) { if err := fixtureSessionService(t, h.db).RevokeDevice(t.Context(), h.tenant, runtime.device.ID); err != nil { t.Fatal(err) } - work, err := h.s.ListEnvironmentInputWork(t.Context(), "", []string{runtime.device.ID}) + work, err := store.SessionAdapter(h.s).ListEnvironmentInputWork(t.Context(), "", []string{runtime.device.ID}) if err != nil || len(work) != 0 { t.Fatal("revoked Runtime selected", work, err) } @@ -55,7 +55,7 @@ func TestEnvironmentInputWorkFiltersAndPagesDevices(t *testing.T) { } } for _, devices := range [][]string{nil, {}, {uuid.NewString()}} { - work, err := h.s.ListEnvironmentInputWork(t.Context(), "", devices) + work, err := store.SessionAdapter(h.s).ListEnvironmentInputWork(t.Context(), "", devices) if err != nil || len(work) != 0 { t.Fatal("unconnected work selected", work, err) } @@ -67,7 +67,7 @@ func TestEnvironmentInputWorkFiltersAndPagesDevices(t *testing.T) { for _, runtime := range h.environments { devices = append(devices, runtime.device.ID) } - work, err := h.s.ListEnvironmentInputWork(t.Context(), cursor, devices) + work, err := store.SessionAdapter(h.s).ListEnvironmentInputWork(t.Context(), cursor, devices) if err != nil || len(work) != count { t.Fatal("environment work page", len(work), count, err) } diff --git a/services/core/internal/store/environment_worker_helpers_test.go b/services/core/internal/store/environment_worker_helpers_test.go index d39561867..4ea8d7485 100644 --- a/services/core/internal/store/environment_worker_helpers_test.go +++ b/services/core/internal/store/environment_worker_helpers_test.go @@ -48,7 +48,7 @@ func unboundWorkerEnvironmentReservation(t *testing.T, h *dispatchHarness) sessi if err != nil { t.Fatal(err) } - pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, session.ID, "work", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) + pending, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, session.ID, "work", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) if err != nil { t.Fatal(err) } @@ -115,7 +115,7 @@ func awaitWorkerEnvironmentRun(t *testing.T, ctx context.Context, s *store.Store var run execution.EnvironmentRun awaitDaemonRemoteCondition(t, ctx, 5*time.Minute, "worker terminal Environment Turn", func() bool { var err error - run.Reservation, err = s.GetEnvironmentInputReservation(ctx, tenant, pending.SessionID, pending.ID) + run.Reservation, err = store.SessionAdapter(s).GetEnvironmentInputReservation(ctx, tenant, pending.SessionID, pending.ID) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/environment_worker_scan_test.go b/services/core/internal/store/environment_worker_scan_test.go index f160dbb04..56860b921 100644 --- a/services/core/internal/store/environment_worker_scan_test.go +++ b/services/core/internal/store/environment_worker_scan_test.go @@ -44,7 +44,7 @@ func TestWorkerEnvironmentRetriesNewlyReadyAtNextScan(t *testing.T) { runtime.write(prepare.ID, proto.TypePreparationStatus, proto.PreparationStatusPayload{Handle: handle, Revision: 2, State: "failed"}) nextWorkerFrame(t, frames, proto.TypeExecutionRelease) stop() - stored, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, pending.SessionID, pending.ID) + stored, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, pending.SessionID, pending.ID) if err != nil || stored.State != sessions.EnvironmentInputPending || !stored.Deadline.Equal(pending.Deadline) || len(stored.Receipts) != 0 { t.Fatal("readiness retry changed pending identity or admitted work", stored, err) } diff --git a/services/core/internal/store/environment_worker_test.go b/services/core/internal/store/environment_worker_test.go index 275e1220f..58d851208 100644 --- a/services/core/internal/store/environment_worker_test.go +++ b/services/core/internal/store/environment_worker_test.go @@ -82,7 +82,7 @@ func TestWorkerEnvironmentSharesCapacityThroughClaimAndCleanup(t *testing.T) { tenant, due := newEnvironmentExpiryReservation(t, h.s) makeEnvironmentExpiryDue(t, pool, &due) waitEnvironmentExpiry(t, h.s, tenant, due) - if _, err := h.s.CancelEnvironmentInput(t.Context(), h.tenant, waiting.SessionID, waiting.ID); err != nil { + if _, err := store.CancelEnvironmentInput(t.Context(), h.s, h.tenant, waiting.SessionID, waiting.ID); err != nil { t.Fatal(err) } // Distinct Runtime sockets do not promise cross-socket delivery order. @@ -159,7 +159,7 @@ func TestWorkerEnvironmentRetriesPendingWithoutExtendingDeadline(t *testing.T) { } stop() nextWorkerFrame(t, frames, proto.TypeExecutionRelease) - stored, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, pending.SessionID, pending.ID) + stored, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, pending.SessionID, pending.ID) if err != nil || stored.State != sessions.EnvironmentInputPending || !stored.Deadline.Equal(pending.Deadline) || len(stored.Receipts) != 0 { t.Fatal("retry or shutdown changed the original reservation", stored, err) } diff --git a/services/core/internal/store/execution.go b/services/core/internal/store/execution.go index 0978511c5..92d5f5fc9 100644 --- a/services/core/internal/store/execution.go +++ b/services/core/internal/store/execution.go @@ -1,13 +1,6 @@ package store -import ( - "errors" - - "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") +import "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" // NewExecution returns the execution writer built on lease, which the caller // acquired and closes. The writer's Session and execution-only transactions run @@ -18,12 +11,3 @@ func NewExecution(s *Store, lease *pgunit.Lease) *Store { writer.writer, writer.lease = lease, lease return &writer } - -// 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 -} diff --git a/services/core/internal/store/execution_events_test.go b/services/core/internal/store/execution_events_test.go index 2ba74d5c9..5e8b963e9 100644 --- a/services/core/internal/store/execution_events_test.go +++ b/services/core/internal/store/execution_events_test.go @@ -39,7 +39,7 @@ func TestExecutionPersistsLiveAndCancelledPartialOutput(t *testing.T) { } time.Sleep(20 * time.Millisecond) } - if _, err := h.s.RequestCancel(ctx, h.tenant, h.session.ID, "stop"); err != nil { + if _, err := store.RequestCancel(ctx, h.s, h.tenant, h.session.ID, "stop"); err != nil { t.Fatal(err) } cancelEnv := h.read(proto.TypePromptCancel) diff --git a/services/core/internal/store/execution_messages_test.go b/services/core/internal/store/execution_messages_test.go index ccb8f31eb..3df7562ee 100644 --- a/services/core/internal/store/execution_messages_test.go +++ b/services/core/internal/store/execution_messages_test.go @@ -54,7 +54,7 @@ func TestExecutionNegotiatesAndPersistsMessageObservations(t *testing.T) { h.write(input.TurnID, proto.TypeOutputMessage, proto.OutputMessagePayload{ID: "a", Status: "completed", Phase: "commentary", Text: &text}) h.write(input.TurnID, proto.TypeOutputMessage, proto.OutputMessagePayload{ID: "b", Status: "in_progress", Phase: "final_answer"}) h.write(input.TurnID, proto.TypeDelta, proto.DeltaPayload{ItemID: "b", Delta: "partial", Sequence: 2}) - if _, err := h.s.RequestCancel(ctx, h.tenant, h.session.ID, "cancel"); err != nil { + if _, err := store.RequestCancel(ctx, h.s, h.tenant, h.session.ID, "cancel"); err != nil { t.Fatal(err) } env = h.read(proto.TypePromptCancel) diff --git a/services/core/internal/store/execution_test.go b/services/core/internal/store/execution_test.go index c50428880..8104027e5 100644 --- a/services/core/internal/store/execution_test.go +++ b/services/core/internal/store/execution_test.go @@ -169,7 +169,7 @@ func TestExecutionLeaseLossFencesAllLifecycleWrites(t *testing.T) { mustReject("completion", err) _, err = operations.TransitionTurn(t.Context(), tenant, active.ID, input.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnInProgress, Status: sessions.TurnFailed}) mustReject("reconciliation", err) - _, err = writer.ExpireEnvironmentInputs(t.Context()) + _, err = operations.ExpireEnvironmentInputs(t.Context()) mustReject("input expiry", err) mustReject("ownership check", writer.lease.CheckOwnership(t.Context())) after, err := sessionAdapter(s).GetSession(t.Context(), tenant, active.ID) @@ -276,7 +276,7 @@ func TestExecutionWriterSerializesWritesOnItsLease(t *testing.T) { ctx, cancel := context.WithTimeout(t.Context(), time.Second) defer cancel() task := tasks[0] - _, err := s.SubmitMessage(ctx, task.tenant, task.session, "public", json.RawMessage(`{"text":"additional"}`)) + _, err := sendMessage(ctx, s, task.tenant, task.session, "public", json.RawMessage(`{"text":"additional"}`)) close(release) if err != nil { t.Fatal("public admission used owner gate", err) @@ -285,27 +285,3 @@ func TestExecutionWriterSerializesWritesOnItsLease(t *testing.T) { t.Fatal(err) } } - -func TestPooledStoreHasNoExecutionAuthority(t *testing.T) { - s, _ := testStore(t) - tenant, session := newTurnSession(t, s) - submitMessage(t, s, tenant, session.ID, "start") - before, err := sessionAdapter(s).GetSession(t.Context(), tenant, session.ID) - if err != nil { - t.Fatal(err) - } - cursor, err := sessionAdapter(s).SessionEventCursor(t.Context(), tenant, session.ID) - if err != nil { - t.Fatal(err) - } - if _, err := s.ExpireEnvironmentInputs(t.Context()); !errors.Is(err, ErrExecutionAuthority) { - t.Fatalf("pooled Store ran input expiry: %v", err) - } - after, err := sessionAdapter(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 := sessionAdapter(s).SessionEventCursor(t.Context(), tenant, session.ID); err != nil || next != cursor { - t.Fatal("rejected execution operation published events", next, err) - } -} diff --git a/services/core/internal/store/execution_tools_test.go b/services/core/internal/store/execution_tools_test.go index fe1484b4d..2db63f908 100644 --- a/services/core/internal/store/execution_tools_test.go +++ b/services/core/internal/store/execution_tools_test.go @@ -40,7 +40,7 @@ func TestExecutionNegotiatesAndPersistsToolObservations(t *testing.T) { } h.write(input.TurnID, proto.TypeToolCall, proto.ToolCallPayload{ID: id, Stage: stage, Observation: &observation}) } - if _, err := h.s.RequestCancel(ctx, h.tenant, h.session.ID, "cancel"); err != nil { + if _, err := store.RequestCancel(ctx, h.s, h.tenant, h.session.ID, "cancel"); err != nil { t.Fatal(err) } env = h.read(proto.TypePromptCancel) diff --git a/services/core/internal/store/export_test.go b/services/core/internal/store/export_test.go index 6e3040cc3..50d908034 100644 --- a/services/core/internal/store/export_test.go +++ b/services/core/internal/store/export_test.go @@ -28,6 +28,27 @@ func SessionAdapter(s *Store) *sessionpg.Store { return sessionAdapter(s) } // SessionService is the Session service on SessionAdapter(s). func SessionService(t testing.TB, s *Store) *sessions.Service { return sessionService(t, s) } +// SubmitInputs admits inputs through the Session service on s's database. +func SubmitInputs(ctx context.Context, s *Store, tenant, session, key string, inputs []sessions.Input) ([]sessions.InputReceipt, error) { + return submitInputs(ctx, s, tenant, session, key, inputs) +} + +// SendMessage admits one message input with payload. +func SendMessage(ctx context.Context, s *Store, tenant, session, key string, payload json.RawMessage) (sessions.InputReceipt, error) { + return sendMessage(ctx, s, tenant, session, key, payload) +} + +// RequestCancel admits one cancel input. +func RequestCancel(ctx context.Context, s *Store, tenant, session, key string) (sessions.InputReceipt, error) { + return requestCancel(ctx, s, tenant, session, key) +} + +// CancelEnvironmentInput cancels the Session's pending Environment input as +// Session cancellation does, then reads the reservation back. +func CancelEnvironmentInput(ctx context.Context, s *Store, tenant, session, reservation string) (sessions.EnvironmentInputReservation, error) { + return cancelEnvironmentInput(ctx, s, tenant, session, reservation) +} + // TransitionTurn moves the Turn as the execution owner does, in a Session // transaction on s's writer. func TransitionTurn(ctx context.Context, s *Store, tenant, session, turn string, transition sessions.TurnTransition) (sessions.Turn, error) { @@ -144,6 +165,6 @@ func SubmitFixtureFunctionResult(ctx context.Context, s *Store, tenantID, sessio if err != nil { return err } - _, err = s.SubmitInputs(ctx, tenantID, sessionID, uuid.NewString(), []sessions.Input{{Kind: "tool_result", Payload: payload}}) + _, err = SubmitInputs(ctx, s, tenantID, sessionID, uuid.NewString(), []sessions.Input{{Kind: "tool_result", Payload: payload}}) return err } diff --git a/services/core/internal/store/function_execution_native_test.go b/services/core/internal/store/function_execution_native_test.go index a6ed06c0a..665b98188 100644 --- a/services/core/internal/store/function_execution_native_test.go +++ b/services/core/internal/store/function_execution_native_test.go @@ -32,7 +32,7 @@ func TestNativeFunctionExecutionPersistsCallsResultsAndContinuity(t *testing.T) t.Fatal(action, state.LastTurn) } if index == 2 { - if _, err := h.s.RequestCancel(ctx, h.tenant, h.session.ID, "native-cancel"); err != nil { + if _, err := store.RequestCancel(ctx, h.s, h.tenant, h.session.ID, "native-cancel"); err != nil { t.Fatal(err) } h.finished(running, sessions.TurnCancelled) diff --git a/services/core/internal/store/function_execution_test.go b/services/core/internal/store/function_execution_test.go index 0aca0f12c..95e6c0835 100644 --- a/services/core/internal/store/function_execution_test.go +++ b/services/core/internal/store/function_execution_test.go @@ -137,7 +137,7 @@ func TestExecutionFunctionsCancellationAndUnconfirmedResults(t *testing.T) { } status := sessions.TurnFailed if cancel { - if _, err := h.s.RequestCancel(t.Context(), h.tenant, h.session.ID, "cancel"); err != nil { + if _, err := store.RequestCancel(t.Context(), h.s, h.tenant, h.session.ID, "cancel"); err != nil { t.Fatal(err) } var request proto.PromptCancelPayload diff --git a/services/core/internal/store/function_images_native_test.go b/services/core/internal/store/function_images_native_test.go index 3a00438fa..54c2284eb 100644 --- a/services/core/internal/store/function_images_native_test.go +++ b/services/core/internal/store/function_images_native_test.go @@ -77,7 +77,7 @@ func TestNativeFunctionImagePublicExecution(t *testing.T) { if err != nil || !call.Applied { t.Fatal("function delivery acknowledgement missing", err) } - inputs, err := h.s.ListTurnInputs(ctx, h.tenant, proof.Session, item.Turn, 0, 100) + inputs, err := store.SessionAdapter(h.s).ListTurnInputs(ctx, h.tenant, proof.Session, item.Turn, 0, 100) if err != nil || len(inputs) != 2 || inputs[0].Kind != "message" || inputs[1].Kind != "tool_result" { t.Fatal("function result admission duplicated or mutated", err) } diff --git a/services/core/internal/store/function_input_execution_test.go b/services/core/internal/store/function_input_execution_test.go index fd0964274..5d3141eb0 100644 --- a/services/core/internal/store/function_input_execution_test.go +++ b/services/core/internal/store/function_input_execution_test.go @@ -20,11 +20,11 @@ func TestExecutionFunctionInputBatchStillSteersMessages(t *testing.T) { state := functionState(t, h, 1) raw, _ := json.Marshal(sessions.FunctionResultInput{TurnID: input.TurnID, CallID: state.RequiredActions[0].CallID, Result: json.RawMessage(`{"success":true,"output":"answer"}`)}) batch := []sessions.Input{{Kind: "tool_result", Payload: raw}, {Kind: "message", Payload: json.RawMessage(`{"text":"Follow up"}`)}} - receipts, err := h.s.SubmitInputs(t.Context(), h.tenant, h.session.ID, "mixed", batch) + receipts, err := store.SubmitInputs(t.Context(), h.s, h.tenant, h.session.ID, "mixed", batch) if err != nil { t.Fatal(err) } - if _, err := h.s.SubmitInputs(t.Context(), h.tenant, h.session.ID, "mixed", batch); err != nil { + if _, err := store.SubmitInputs(t.Context(), h.s, h.tenant, h.session.ID, "mixed", batch); err != nil { t.Fatal(err) } resultSeen, messageSeen := false, false diff --git a/services/core/internal/store/function_inputs_public_test.go b/services/core/internal/store/function_inputs_public_test.go index 7570cbe0b..4ad8bade0 100644 --- a/services/core/internal/store/function_inputs_public_test.go +++ b/services/core/internal/store/function_inputs_public_test.go @@ -34,7 +34,7 @@ func TestFunctionInputsOfficialClientAtomicAdmission(t *testing.T) { if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(ctx, tenant, session.ID, "start", json.RawMessage(`{"text":"fixture"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "start", json.RawMessage(`{"text":"fixture"}`)) if err != nil { t.Fatal(err) } @@ -104,14 +104,14 @@ func TestFunctionInputsOfficialClientAtomicAdmission(t *testing.T) { } } } - history, err := s.ListTurnInputs(ctx, tenant, session.ID, input.TurnID, 0, 100) + history, err := store.SessionAdapter(s).ListTurnInputs(ctx, tenant, session.ID, input.TurnID, 0, 100) if err != nil || len(history) != 6 { t.Fatal(history, err) } if _, err := store.TransitionTurn(ctx, s, tenant, session.ID, input.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnWaiting, Status: sessions.TurnFailed}); err != nil { t.Fatal(err) } - next, err := s.SubmitMessage(ctx, tenant, session.ID, "next", json.RawMessage(`{"text":"next"}`)) + next, err := store.SendMessage(ctx, s, tenant, session.ID, "next", json.RawMessage(`{"text":"next"}`)) if err != nil { t.Fatal(err) } @@ -121,7 +121,7 @@ func TestFunctionInputsOfficialClientAtomicAdmission(t *testing.T) { if err := command.Wait(); err != nil { t.Fatalf("SDK retry: %v %s", err, stderr.String()) } - history, err = s.ListTurnInputs(ctx, tenant, session.ID, next.TurnID, 0, 100) + history, err = store.SessionAdapter(s).ListTurnInputs(ctx, tenant, session.ID, next.TurnID, 0, 100) if err != nil || len(history) != 1 { t.Fatal(history, err) } diff --git a/services/core/internal/store/function_inputs_test.go b/services/core/internal/store/function_inputs_test.go index 5feae8e64..3a82d5ac5 100644 --- a/services/core/internal/store/function_inputs_test.go +++ b/services/core/internal/store/function_inputs_test.go @@ -40,7 +40,7 @@ func TestFunctionInputBatchesPersistAndReplayWithoutRetargeting(t *testing.T) { tenant, session, turn := functionInputFixture(t, s, functionExecution(t)) full := `{"success":false,"output":[{"type":"input_text","text":""},{"type":"input_image","image_url":"data:image/png;base64,AA=="},{"type":"input_text","text":"after"}],"error":"failed"}` batch := []sessions.Input{resultInput(t, turn, "a", full), {Kind: "message", Payload: json.RawMessage(`{"text":"Follow up"}`)}, resultInput(t, turn, "b", `{"success":true,"output":null,"error":null}`), {Kind: "cancel", Payload: json.RawMessage(`{}`)}} - receipts, err := s.SubmitInputs(t.Context(), tenant, session.ID, "batch", batch) + receipts, err := submitInputs(t.Context(), s, tenant, session.ID, "batch", batch) if err != nil || len(receipts) != 4 { t.Fatal(receipts, err) } @@ -60,7 +60,7 @@ func TestFunctionInputBatchesPersistAndReplayWithoutRetargeting(t *testing.T) { next := submitMessage(t, s, tenant, session.ID, "next").TurnID pool.Close() s, _ = testStore(t) - retry, err := s.SubmitInputs(t.Context(), tenant, session.ID, "batch", batch) + retry, err := submitInputs(t.Context(), s, tenant, session.ID, "batch", batch) if err != nil || len(retry) != len(receipts) { t.Fatal(retry, err) } @@ -73,11 +73,11 @@ func TestFunctionInputBatchesPersistAndReplayWithoutRetargeting(t *testing.T) { if !reflect.DeepEqual(retry, receipts) { t.Fatal(retry, receipts) } - history, err := s.ListTurnInputs(t.Context(), tenant, session.ID, turn, 0, 100) + history, err := sessionAdapter(s).ListTurnInputs(t.Context(), tenant, session.ID, turn, 0, 100) if err != nil || len(history) != 5 || history[1].Kind != "tool_result" || history[2].Kind != "message" { t.Fatal(history, err) } - future, err := s.ListTurnInputs(t.Context(), tenant, session.ID, next, 0, 100) + future, err := sessionAdapter(s).ListTurnInputs(t.Context(), tenant, session.ID, next, 0, 100) if err != nil || len(future) != 1 { t.Fatal(future, err) } @@ -86,11 +86,11 @@ func TestFunctionInputBatchesPersistAndReplayWithoutRetargeting(t *testing.T) { t.Fatal(current, err) } changed := []sessions.Input{batch[1], batch[0], batch[2], batch[3]} - if _, err := s.SubmitInputs(t.Context(), tenant, session.ID, "batch", changed); !errors.Is(err, sessions.ErrIdempotencyConflict) { + if _, err := submitInputs(t.Context(), s, tenant, session.ID, "batch", changed); !errors.Is(err, sessions.ErrIdempotencyConflict) { t.Fatal(err) } // A new request identity can repeat an identical saved result, without native application. - if _, err := s.SubmitInputs(t.Context(), tenant, session.ID, "same-result", batch[:1]); err != nil { + if _, err := submitInputs(t.Context(), s, tenant, session.ID, "same-result", batch[:1]); err != nil { t.Fatal(err) } call, err = FixtureFunctionCall(t.Context(), s.pool, tenant, session.ID, turn, "a") @@ -136,7 +136,7 @@ func TestFunctionInputBatchFailureRollsBackEveryWrite(t *testing.T) { batch = append(batch, resultInput(t, turn, "a", `{"success":false}`)) expected = sessions.ErrFunctionResultConflict } - if _, err := s.SubmitInputs(t.Context(), tenant, session.ID, "failed-batch", batch); !errors.Is(err, expected) { + if _, err := submitInputs(t.Context(), s, tenant, session.ID, "failed-batch", batch); !errors.Is(err, expected) { t.Fatal(err) } call, err := FixtureFunctionCall(t.Context(), s.pool, tenant, session.ID, turn, "a") @@ -147,11 +147,11 @@ func TestFunctionInputBatchFailureRollsBackEveryWrite(t *testing.T) { if err != nil || !state.CancelRequestedAt.IsZero() || state.Status != sessions.TurnWaiting { t.Fatal(state, err) } - history, err := s.ListTurnInputs(t.Context(), tenant, session.ID, turn, 0, 100) + history, err := sessionAdapter(s).ListTurnInputs(t.Context(), tenant, session.ID, turn, 0, 100) if err != nil || len(history) != 1 { t.Fatal(history, err) } - if _, err := s.SubmitInputs(t.Context(), tenant, session.ID, "failed-batch", []sessions.Input{first, cancel}); err != nil { + if _, err := submitInputs(t.Context(), s, tenant, session.ID, "failed-batch", []sessions.Input{first, cancel}); err != nil { t.Fatal("failed transaction retained retry identity", err) } }) @@ -169,7 +169,7 @@ func TestFunctionInputConcurrentBatchesSelectOneResult(t *testing.T) { wg.Add(1) go func() { defer wg.Done() - _, err := other.SubmitInputs(t.Context(), tenant, session.ID, fmt.Sprint(i), batch) + _, err := submitInputs(t.Context(), other, tenant, session.ID, fmt.Sprint(i), batch) results <- err }() } @@ -188,7 +188,7 @@ func TestFunctionInputConcurrentBatchesSelectOneResult(t *testing.T) { if wins != 1 || conflicts != 1 { t.Fatal(wins, conflicts) } - history, err := s.ListTurnInputs(t.Context(), tenant, session.ID, turn, 0, 100) + history, err := sessionAdapter(s).ListTurnInputs(t.Context(), tenant, session.ID, turn, 0, 100) if err != nil || len(history) != 3 { t.Fatal(history, err) } @@ -198,7 +198,7 @@ func TestFunctionInputsRejectInvalidTargetsAndStorageObjects(t *testing.T) { s, _ := testStore(t) tenant, session, turn := functionInputFixture(t, s, functionExecution(t)) for _, raw := range []string{`{}`, `{"turn_id":"","call_id":"a","result":{}}`, fmt.Sprintf(`{"turn_id":%q,"call_id":" ","result":{}}`, turn), fmt.Sprintf(`{"turn_id":%q,"call_id":"a"}`, turn), fmt.Sprintf(`{"turn_id":%q,"call_id":"a","result":null}`, turn), fmt.Sprintf(`{"turn_id":%q,"call_id":"a","result":[]}`, turn)} { - _, err := s.SubmitInputs(t.Context(), tenant, session.ID, "invalid", []sessions.Input{{Kind: "tool_result", Payload: json.RawMessage(raw)}}) + _, err := submitInputs(t.Context(), s, tenant, session.ID, "invalid", []sessions.Input{{Kind: "tool_result", Payload: json.RawMessage(raw)}}) if !errors.Is(err, sessions.ErrInvalidInput) { t.Fatal(err) } @@ -207,7 +207,7 @@ func TestFunctionInputsRejectInvalidTargetsAndStorageObjects(t *testing.T) { // Missing, foreign and malformed Sessions are not found, whatever the target. for _, scope := range []struct{ tenant, session string }{{uuid.NewString(), session.ID}, {tenant, uuid.NewString()}, {tenant, "sess_malformed"}} { for _, target := range []sessions.Input{input, resultInput(t, "bad", "missing", `{"success":true}`)} { - if _, err := s.SubmitInputs(t.Context(), scope.tenant, scope.session, "foreign", []sessions.Input{target}); !errors.Is(err, sessions.ErrNotFound) { + if _, err := submitInputs(t.Context(), s, scope.tenant, scope.session, "foreign", []sessions.Input{target}); !errors.Is(err, sessions.ErrNotFound) { t.Fatal(err) } } @@ -221,12 +221,12 @@ func TestFunctionInputsRejectInvalidTargetsAndStorageObjects(t *testing.T) { {uuid.NewString(), "a", sessions.ErrFunctionCallTurnMismatch}, {"bad", "a", sessions.ErrFunctionCallTurnMismatch}, {turn, "missing", sessions.ErrUnknownFunctionCall}, {uuid.NewString(), "missing", sessions.ErrUnknownFunctionCall}, {"bad", "missing", sessions.ErrUnknownFunctionCall}, } { - if _, err := s.SubmitInputs(t.Context(), tenant, session.ID, "target", []sessions.Input{resultInput(t, target.turn, target.call, `{"success":true}`)}); !errors.Is(err, target.want) { + if _, err := submitInputs(t.Context(), s, tenant, session.ID, "target", []sessions.Input{resultInput(t, target.turn, target.call, `{"success":true}`)}); !errors.Is(err, target.want) { t.Fatal(target, err) } } transition(t, s, tenant, session.ID, turn, sessions.TurnWaiting, sessions.TurnFailed) - if _, err := s.SubmitInputs(t.Context(), tenant, session.ID, "late", []sessions.Input{input}); !errors.Is(err, sessions.ErrTurnConflict) { + if _, err := submitInputs(t.Context(), s, tenant, session.ID, "late", []sessions.Input{input}); !errors.Is(err, sessions.ErrTurnConflict) { t.Fatal(err) } current, err := sessionAdapter(s).GetSession(t.Context(), tenant, session.ID) diff --git a/services/core/internal/store/function_state_public_test.go b/services/core/internal/store/function_state_public_test.go index c578c5310..69b329cbf 100644 --- a/services/core/internal/store/function_state_public_test.go +++ b/services/core/internal/store/function_state_public_test.go @@ -31,7 +31,7 @@ func TestFunctionStateOfficialClientReadsAndLiveEvents(t *testing.T) { if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(ctx, tenant, session.ID, "start", json.RawMessage(`{"text":"fixture"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "start", json.RawMessage(`{"text":"fixture"}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/function_state_test.go b/services/core/internal/store/function_state_test.go index 320650e7f..8da65ee48 100644 --- a/services/core/internal/store/function_state_test.go +++ b/services/core/internal/store/function_state_test.go @@ -98,7 +98,7 @@ func TestFunctionStateCancellationAndTerminalCleanup(t *testing.T) { if status == sessions.TurnCancelled { before, _ := sessionAdapter(s).SessionEventCursor(t.Context(), tenant, session.ID) for _, key := range []string{"cancel", "cancel", "another-cancel"} { - if _, err := s.RequestCancel(t.Context(), tenant, session.ID, key); err != nil { + if _, err := requestCancel(t.Context(), s, tenant, session.ID, key); err != nil { t.Fatal(err) } } diff --git a/services/core/internal/store/harness_onboarding_test.go b/services/core/internal/store/harness_onboarding_test.go index 3203b07de..c41003575 100644 --- a/services/core/internal/store/harness_onboarding_test.go +++ b/services/core/internal/store/harness_onboarding_test.go @@ -105,7 +105,7 @@ func TestThirdHarnessPublicOnboarding(t *testing.T) { t.Fatal(err) } var result execution.Result - inputs, inputErr := h.s.ListTurnInputs(ctx, h.tenant, created.ID, first.RunID, 0, 100) + inputs, inputErr := store.SessionAdapter(h.s).ListTurnInputs(ctx, h.tenant, created.ID, first.RunID, 0, 100) if inputErr != nil || len(inputs) != 2 { t.Fatal(inputs, inputErr) } diff --git a/services/core/internal/store/hosted_initialization_failure_public_test.go b/services/core/internal/store/hosted_initialization_failure_public_test.go index 90ccc4f5a..805d0ebb3 100644 --- a/services/core/internal/store/hosted_initialization_failure_public_test.go +++ b/services/core/internal/store/hosted_initialization_failure_public_test.go @@ -223,7 +223,7 @@ func TestHostedInitializationFailureRecordsSafeSessionFailure(t *testing.T) { !last.EnvironmentFailure.FailedAt.Equal(read.EnvironmentFailure.FailedAt) || last.EnvironmentInputActivity != nil || !last.Settled { t.Fatal("failed snapshot", last) } - if _, err := s.ReserveEnvironmentInput(t.Context(), tenant, session.ID, "later", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"later"}`)}}); !errors.Is(err, sessions.ErrHostedEnvironmentFailed) { + if _, err := store.SessionService(t, s).ReserveEnvironmentInput(t.Context(), tenant, session.ID, "later", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"later"}`)}}); !errors.Is(err, sessions.ErrHostedEnvironmentFailed) { t.Fatal("failed hosted Environment admitted input", err) } raw, _ := json.Marshal(events) diff --git a/services/core/internal/store/input_batches_test.go b/services/core/internal/store/input_batches_test.go deleted file mode 100644 index 16574d0fd..000000000 --- a/services/core/internal/store/input_batches_test.go +++ /dev/null @@ -1,189 +0,0 @@ -package store - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "reflect" - "strings" - "sync" - "testing" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" - "github.com/google/uuid" -) - -func messageInput(text string) sessions.Input { - payload, _ := json.Marshal(map[string]string{"text": text}) - return sessions.Input{Kind: "message", Payload: payload} -} - -func TestInputBatchesAreOrderedAndIdempotentAcrossConnections(t *testing.T) { - s, _ := testStore(t) - other, _ := testStore(t) - tenant, session := newTurnSession(t, s) - ctx := context.Background() - const count = 8 - batch := []sessions.Input{messageInput("first"), messageInput("second")} - for _, repeated := range []bool{true, false} { - var wg sync.WaitGroup - receipts := make(chan []sessions.InputReceipt, count) - errs := make(chan error, count) - for i := range count { - wg.Add(1) - go func() { - defer wg.Done() - st := s - if i%2 == 0 { - st = other - } - key := "same-batch" - if !repeated { - key = fmt.Sprintf("batch-%d", i) - } - got, err := st.SubmitInputs(ctx, tenant, session.ID, key, batch) - receipts <- got - errs <- err - }() - } - wg.Wait() - close(receipts) - close(errs) - for err := range errs { - if err != nil { - t.Fatal(err) - } - } - newBatches := 0 - for got := range receipts { - if len(got) != 2 || got[0].TurnID == "" || got[0].TurnID != got[1].TurnID || got[0].Sequence >= got[1].Sequence || got[0].Replayed != got[1].Replayed { - t.Fatalf("invalid batch receipt: %+v", got) - } - if !got[0].Replayed { - newBatches++ - } - } - want := count - if repeated { - want = 1 - } - if newBatches != want { - t.Fatalf("accepted %d batches, want %d", newBatches, want) - } - } - got, err := s.SubmitInputs(ctx, tenant, session.ID, "same-batch", batch) - if err != nil { - t.Fatal(err) - } - inputs, err := s.ListTurnInputs(ctx, tenant, session.ID, got[0].TurnID, 0, 100) - if err != nil || len(inputs) != 2*(count+1) { - t.Fatalf("inputs=%d err=%v", len(inputs), err) - } - for i, input := range inputs { - var payload map[string]string - if err := json.Unmarshal(input.Payload, &payload); err != nil { - t.Fatal(err) - } - want := "first" - if i%2 == 1 { - want = "second" - } - if payload["text"] != want { - t.Fatalf("batch interleaved at %d: %v", i, payload) - } - } -} - -func TestBatchRetriesCompareTheWholeRequestAndRetainTargets(t *testing.T) { - s, pool := testStore(t) - tenant, session := newTurnSession(t, s) - ctx := context.Background() - cancel := sessions.Input{Kind: "cancel", Payload: json.RawMessage(`{}`)} - batch := []sessions.Input{cancel, messageInput("one"), cancel, messageInput("two")} - first, err := s.SubmitInputs(ctx, tenant, session.ID, "mixed", batch) - if err != nil { - t.Fatal(err) - } - if first[0].TurnID != "" || first[1].TurnID == "" || first[1].TurnID != first[2].TurnID || first[1].TurnID == first[3].TurnID { - t.Fatalf("cancellation targets: %+v", first) - } - cancelled, err := sessionAdapter(s).GetTurn(ctx, tenant, session.ID, first[1].TurnID) - if err != nil || cancelled.Status != sessions.TurnCancelled { - t.Fatalf("cancelled turn=%+v err=%v", cancelled, err) - } - transition(t, s, tenant, session.ID, first[3].TurnID, sessions.TurnQueued, sessions.TurnInProgress) - transition(t, s, tenant, session.ID, first[3].TurnID, sessions.TurnInProgress, sessions.TurnCompleted) - next := submitMessage(t, s, tenant, session.ID, "next") - for _, changed := range [][]sessions.Input{batch[:3], append(append([]sessions.Input{}, batch...), cancel), {batch[0], batch[3], batch[2], batch[1]}} { - if _, err := s.SubmitInputs(ctx, tenant, session.ID, "mixed", changed); !errors.Is(err, sessions.ErrIdempotencyConflict) { - t.Fatalf("changed batch accepted: %v", err) - } - } - batch[1].Payload = json.RawMessage(`{ "text" : "one" }`) - pool.Close() - restarted, _ := testStore(t) - retry, err := restarted.SubmitInputs(ctx, tenant, session.ID, "mixed", batch) - for i := range first { - first[i].Replayed = true - } - if err != nil || !reflect.DeepEqual(retry, first) { - t.Fatalf("restart changed receipts: %+v, %v", retry, err) - } - current, err := sessionAdapter(restarted).GetTurn(ctx, tenant, session.ID, next.TurnID) - if err != nil || current.Status != sessions.TurnQueued || !current.CancelRequestedAt.IsZero() { - t.Fatalf("retry cancelled later work: %+v, %v", current, err) - } - otherTenant, _ := newTurnSession(t, restarted) - if _, err := restarted.SubmitInputs(ctx, otherTenant, session.ID, "mixed", batch); !errors.Is(err, sessions.ErrNotFound) { - t.Fatalf("batch retry escaped tenant: %v", err) - } -} - -func TestFailedBatchRollsBackEarlierCancellationAndInputs(t *testing.T) { - s, pool := testStore(t) - tenant, session := newTurnSession(t, s) - ctx := context.Background() - initial := submitMessage(t, s, tenant, session.ID, "initial") - transition(t, s, tenant, session.ID, initial.TurnID, sessions.TurnQueued, sessions.TurnInProgress) - key := uuid.NewString() - constraint := "batch_failure_" + strings.ReplaceAll(key, "-", "") - // Inject a storage error on the second insert, after cancellation has run. - _, err := pool.Exec(ctx, "ALTER TABLE turn_inputs ADD CONSTRAINT "+constraint+" CHECK (idempotency_key <> '"+key+"' OR batch_position = 0)") - if err != nil { - t.Fatal(err) - } - t.Cleanup(func() { _, _ = pool.Exec(ctx, "ALTER TABLE turn_inputs DROP CONSTRAINT IF EXISTS "+constraint) }) - batch := []sessions.Input{{Kind: "cancel", Payload: json.RawMessage(`{}`)}, messageInput("after cancel")} - if got, err := s.SubmitInputs(ctx, tenant, session.ID, key, batch); err == nil || got != nil { - t.Fatalf("partial batch succeeded: %+v %v", got, err) - } - turn, err := sessionAdapter(s).GetTurn(ctx, tenant, session.ID, initial.TurnID) - if err != nil || turn.Status != sessions.TurnInProgress || !turn.CancelRequestedAt.IsZero() { - t.Fatalf("cancellation escaped rollback: %+v %v", turn, err) - } - inputs, err := s.ListTurnInputs(ctx, tenant, session.ID, initial.TurnID, 0, 100) - if err != nil || len(inputs) != 1 { - t.Fatalf("partial inputs survived: %v %v", inputs, err) - } - if _, err := pool.Exec(ctx, "ALTER TABLE turn_inputs DROP CONSTRAINT "+constraint); err != nil { - t.Fatal(err) - } - got, err := s.SubmitInputs(ctx, tenant, session.ID, key, batch) - if err != nil || len(got) != 2 || got[0].Replayed || got[1].Replayed { - t.Fatalf("retry after rollback: %+v %v", got, err) - } -} - -func TestInputBatchValidation(t *testing.T) { - for _, input := range [][]sessions.Input{ - nil, make([]sessions.Input, 65), {messageInput("ok"), {Kind: "unsupported", Payload: json.RawMessage(`{}`)}}, - {{Kind: "message"}}, {{Kind: "message", Payload: json.RawMessage(`[]`)}}, - {{Kind: "cancel", Payload: json.RawMessage(`{"target":"other"}`)}}, - {messageInput(strings.Repeat("x", 300*1024)), messageInput(strings.Repeat("y", 300*1024))}, - } { - if _, _, err := validateInputs(input); !errors.Is(err, sessions.ErrInvalidInput) { - t.Fatalf("invalid batch accepted: %v", err) - } - } -} diff --git a/services/core/internal/store/input_conflicts_public_test.go b/services/core/internal/store/input_conflicts_public_test.go index fbe1e422c..7c739e4b1 100644 --- a/services/core/internal/store/input_conflicts_public_test.go +++ b/services/core/internal/store/input_conflicts_public_test.go @@ -70,7 +70,7 @@ func TestSessionInputConflictsAndResultTargetsPostgres(t *testing.T) { // waiting starts a Turn that waits for one function result. waiting := func(session, key, call string) sessions.InputReceipt { t.Helper() - receipt, err := s.SubmitMessage(ctx, tenant, session, key, json.RawMessage(`{"text":"work"}`)) + receipt, err := store.SendMessage(ctx, s, tenant, session, key, json.RawMessage(`{"text":"work"}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/inputs_test.go b/services/core/internal/store/inputs_test.go new file mode 100644 index 000000000..ec00d3d78 --- /dev/null +++ b/services/core/internal/store/inputs_test.go @@ -0,0 +1,128 @@ +package store + +import ( + "context" + "encoding/json" + "testing" + + "github.com/google/uuid" + "github.com/jackc/pgx/v5/pgtype" + "github.com/jackc/pgx/v5/pgxpool" + + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/sessionpg" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" +) + +var messagePayload = json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"hello"}]}]}`) + +func messageInput(text string) sessions.Input { + payload, _ := json.Marshal(map[string]string{"text": text}) + return sessions.Input{Kind: "message", Payload: payload} +} + +func newTurnSession(t *testing.T, s *Store) (string, sessions.Session) { + t.Helper() + tenant := uuid.NewString() + session, err := s.CreateSession(context.Background(), tenant, sessions.CreateSession{Creator: FixtureCreator(), + Engine: "codex", IdempotencyKey: "session", Configuration: json.RawMessage(`{"agent":{"model":"test","instructions":"original"}}`), + }) + if err != nil { + t.Fatal(err) + } + return tenant, session +} + +func environmentInputSession(t *testing.T, s *Store) (string, sessions.Session) { + t.Helper() + tenant := uuid.NewString() + session, err := s.CreateSession(context.Background(), tenant, environmentInput("session", "self_hosted", "/workspace")) + if err != nil { + t.Fatal(err) + } + return tenant, session +} + +// submitInputs admits inputs through the Session service on s's database, as +// the public events route does. +func submitInputs(ctx context.Context, s *Store, tenant, session, key string, inputs []sessions.Input) ([]sessions.InputReceipt, error) { + service, err := sessions.NewService(sessionAdapter(s)) + if err != nil { + return nil, err + } + return service.SubmitInputs(ctx, tenant, session, key, inputs) +} + +func submitInput(ctx context.Context, s *Store, tenant, session, key string, input sessions.Input) (sessions.InputReceipt, error) { + receipts, err := submitInputs(ctx, s, tenant, session, key, []sessions.Input{input}) + if err != nil { + return sessions.InputReceipt{}, err + } + return receipts[0], nil +} + +func sendMessage(ctx context.Context, s *Store, tenant, session, key string, payload json.RawMessage) (sessions.InputReceipt, error) { + return submitInput(ctx, s, tenant, session, key, sessions.Input{Kind: "message", Payload: payload}) +} + +func requestCancel(ctx context.Context, s *Store, tenant, session, key string) (sessions.InputReceipt, error) { + return submitInput(ctx, s, tenant, session, key, sessions.Input{Kind: "cancel", Payload: json.RawMessage(`{}`)}) +} + +func submitMessage(t *testing.T, s *Store, tenant, session, key string) sessions.InputReceipt { + t.Helper() + receipt, err := sendMessage(context.Background(), s, tenant, session, key, messagePayload) + if err != nil { + t.Fatal(err) + } + return receipt +} + +func reserveEnvironmentInput(t *testing.T, s *Store, tenant, session, key string) sessions.EnvironmentInputReservation { + t.Helper() + got, err := sessionService(t, s).ReserveEnvironmentInput(context.Background(), tenant, session, key, []sessions.Input{messageInput("first"), messageInput("second")}) + if err != nil { + t.Fatal(err) + } + return got +} + +// cancelEnvironmentInput cancels the Session's pending Environment input as +// Session cancellation does, then reads the reservation back. +func cancelEnvironmentInput(ctx context.Context, s *Store, tenant, session, reservation string) (sessions.EnvironmentInputReservation, error) { + owner, err := parseID(tenant) + if err != nil { + return sessions.EnvironmentInputReservation{}, err + } + err = s.withPublicSession(ctx, tenant, session, func(ctx context.Context, q *sqlc.Queries, id pgtype.UUID) error { + bound := sessionpg.BindSession(q, owner, id) + return sessions.TrackInputActivity(ctx, bound, func(ctx context.Context) error { return bound.CancelPendingInput(ctx) }) + }) + if err != nil { + return sessions.EnvironmentInputReservation{}, err + } + return sessionAdapter(s).GetEnvironmentInputReservation(ctx, tenant, session, reservation) +} + +func transition(t *testing.T, s *Store, tenant, session, turn, from, to string) sessions.Turn { + t.Helper() + got, err := transitionTurn(context.Background(), s, tenant, session, turn, sessions.TurnTransition{ExpectedStatus: from, Status: to}) + if err != nil { + t.Fatal(err) + } + return got +} + +func environmentInputHistory(t *testing.T, pool *pgxpool.Pool, session string, turns, inputs int) { + t.Helper() + var gotTurns, gotInputs, items, events int + err := pool.QueryRow(context.Background(), ` + SELECT (SELECT count(*) FROM turns WHERE session_id=$1), + (SELECT count(*) FROM turn_inputs WHERE session_id=$1), + (SELECT count(*) FROM session_items WHERE session_id=$1), + (SELECT count(*) FROM session_events WHERE session_id=$1 + AND (payload ? 'turn' OR payload->'event' ? 'item'))`, session).Scan(&gotTurns, &gotInputs, &items, &events) + if err != nil || gotTurns != turns || gotInputs != inputs || items != inputs || (inputs == 0 && events != 0) { + t.Fatal("history", gotTurns, gotInputs, items, events, err) + } +} diff --git a/services/core/internal/store/item_order_test.go b/services/core/internal/store/item_order_test.go index 924da4b50..5514b2e31 100644 --- a/services/core/internal/store/item_order_test.go +++ b/services/core/internal/store/item_order_test.go @@ -23,7 +23,7 @@ func TestItemObservationOrderSurvivesTiesUpdatesRetriesAndRecovery(t *testing.T) if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(ctx, tenant, session.ID, "first", json.RawMessage(`{"text":"question"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "first", json.RawMessage(`{"text":"question"}`)) if err != nil { t.Fatal(err) } @@ -52,7 +52,7 @@ func TestItemObservationOrderSurvivesTiesUpdatesRetriesAndRecovery(t *testing.T) t.Fatal(err) } } - if _, err = s.SubmitMessage(ctx, tenant, session.ID, "steer", json.RawMessage(`{"text":"continue"}`)); err != nil { + if _, err = store.SendMessage(ctx, s, tenant, session.ID, "steer", json.RawMessage(`{"text":"continue"}`)); err != nil { t.Fatal(err) } page, err = sessionReads(pool).ListItems(ctx, tenant, session.ID, "", 100, true) @@ -132,7 +132,7 @@ func TestItemObservationOrderSurvivesTiesUpdatesRetriesAndRecovery(t *testing.T) t.Fatal(err) } checkOrder() - next, err := s.SubmitMessage(ctx, tenant, session.ID, "next-turn", json.RawMessage(`{"text":"new turn"}`)) + next, err := store.SendMessage(ctx, s, tenant, session.ID, "next-turn", json.RawMessage(`{"text":"new turn"}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/item_reads_test.go b/services/core/internal/store/item_reads_test.go index 836fd4333..861f0d7f4 100644 --- a/services/core/internal/store/item_reads_test.go +++ b/services/core/internal/store/item_reads_test.go @@ -23,7 +23,7 @@ func TestItemsRecoverSnapshotsPartialResultsPaginationAndIsolation(t *testing.T) if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(ctx, tenant, session.ID, "first", json.RawMessage(`{"text":"question"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "first", json.RawMessage(`{"text":"question"}`)) if err != nil { t.Fatal(err) } @@ -123,7 +123,7 @@ func TestItemProjectionFailureRollsBackJournalAndAggregateRecovers(t *testing.T) journal := executionOwner(t, fixtureDB{pool: pool}, s).Sessions tenant := uuid.NewString() session, _ := s.CreateSession(ctx, tenant, sessions.CreateSession{Creator: store.FixtureCreator(), Engine: "codex", IdempotencyKey: "legacy"}) - input, err := s.SubmitMessage(ctx, tenant, session.ID, "input", json.RawMessage(`{"text":"test"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "input", json.RawMessage(`{"text":"test"}`)) if err != nil { t.Fatal(err) } @@ -160,7 +160,7 @@ func TestReceiptOnlyTextRecoversWithoutInventingCompletion(t *testing.T) { tenant := uuid.NewString() for _, receiptOnly := range []bool{true, false} { session, _ := s.CreateSession(ctx, tenant, sessions.CreateSession{Creator: store.FixtureCreator(), Engine: "codex", IdempotencyKey: uuid.NewString()}) - input, err := s.SubmitMessage(ctx, tenant, session.ID, "first", json.RawMessage(`{"text":"test"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "first", json.RawMessage(`{"text":"test"}`)) if err != nil { t.Fatal(err) } @@ -194,7 +194,7 @@ func TestLegacyFailureRetainsPartialAnswerAcrossRecovery(t *testing.T) { if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(ctx, tenant, session.ID, "first", json.RawMessage(`{"text":"question"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "first", json.RawMessage(`{"text":"question"}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/list_cursor_public_test.go b/services/core/internal/store/list_cursor_public_test.go index 70d58d2dc..8d050877c 100644 --- a/services/core/internal/store/list_cursor_public_test.go +++ b/services/core/internal/store/list_cursor_public_test.go @@ -69,7 +69,7 @@ func seedCursorFixture(t *testing.T, s *store.Store, leased execution.Owner, ski if _, err := store.TransitionTurn(ctx, s, tenant, f.session, f.turn, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnCancelled}); err != nil { t.Fatal(err) } - if _, err := s.SubmitMessage(ctx, tenant, f.session, label+"-second", json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"second"}]}]}`)); err != nil { + if _, err := store.SendMessage(ctx, s, tenant, f.session, label+"-second", json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"second"}]}]}`)); err != nil { t.Fatal(err) } f.otherSession = client.created(token, "/v1/agents/sessions", newSession) @@ -91,7 +91,7 @@ func seedCursorFixture(t *testing.T, s *store.Store, leased execution.Owner, ski if err != nil { t.Fatal(err) } - receipt, err := s.SubmitMessage(ctx, tenant, created.ID, key+"-input", json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"delegate"}]}]}`)) + receipt, err := store.SendMessage(ctx, s, tenant, created.ID, key+"-input", json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"delegate"}]}]}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/local_artifact_export_test.go b/services/core/internal/store/local_artifact_export_test.go index cc45e09e5..b39f29ded 100644 --- a/services/core/internal/store/local_artifact_export_test.go +++ b/services/core/internal/store/local_artifact_export_test.go @@ -9,6 +9,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/execution" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" ) func completeLocalArtifactExport(t *testing.T, h *dispatchHarness, worker *execution.Worker, environment sessions.Environment) { @@ -46,7 +47,7 @@ func completeLocalArtifactExport(t *testing.T, h *dispatchHarness, worker *execu t.Fatal("capture published before native completion", err) } completeCaptureDirectoryRead(t, h, worker, environment) - pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "during-artifact-capture", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"run after the completed native execution"}`)}}) + pending, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "during-artifact-capture", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"run after the completed native execution"}`)}}) if err != nil || pending.State != sessions.EnvironmentInputPending || len(pending.Receipts) != 0 { t.Fatalf("input during artifact capture was assigned to the finished executor: %+v %v", pending, err) } diff --git a/services/core/internal/store/local_environment_file_write_test.go b/services/core/internal/store/local_environment_file_write_test.go index ed77c1edf..dd795c914 100644 --- a/services/core/internal/store/local_environment_file_write_test.go +++ b/services/core/internal/store/local_environment_file_write_test.go @@ -57,7 +57,7 @@ func TestLocalEnvironmentFileWriteOwnsMutationBeforeDispatch(t *testing.T) { if err != nil || intent.State != "pending" || intent.Identity.DeviceID != h.device.ID { t.Fatal("dispatch preceded durable ownership", intent, err) } - if _, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "concurrent", []sessions.Input{{Kind: "message", Payload: []byte(`{"text":"work"}`)}}); !errors.Is(err, sessions.ErrTurnConflict) { + if _, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "concurrent", []sessions.Input{{Kind: "message", Payload: []byte(`{"text":"work"}`)}}); !errors.Is(err, sessions.ErrTurnConflict) { t.Fatal("upload admitted concurrent execution", err) } cancel() @@ -99,7 +99,7 @@ func TestLocalEnvironmentFileWriteLostReceiptRemainsPending(t *testing.T) { if err != nil || intent.State != "pending" { t.Fatal("disconnect guessed rejection", intent, err) } - if _, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "after-loss", []sessions.Input{{Kind: "message", Payload: []byte(`{"text":"work"}`)}}); !errors.Is(err, sessions.ErrTurnConflict) { + if _, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "after-loss", []sessions.Input{{Kind: "message", Payload: []byte(`{"text":"work"}`)}}); !errors.Is(err, sessions.ErrTurnConflict) { t.Fatal("unknown upload admitted execution", err) } } diff --git a/services/core/internal/store/local_environment_worker_test.go b/services/core/internal/store/local_environment_worker_test.go index cd8e2429a..6f935956d 100644 --- a/services/core/internal/store/local_environment_worker_test.go +++ b/services/core/internal/store/local_environment_worker_test.go @@ -110,7 +110,7 @@ func TestLocalEnvironmentWorkerRejectsGeneralDeviceDespiteCapability(t *testing. func TestLocalEnvironmentWorkerSchedulesPreparationWithoutRemoteResolver(t *testing.T) { h, worker, environment := localWorker(t, true, true) - reservation, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "local-input", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) + reservation, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "local-input", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) if err != nil { t.Fatal(err) } @@ -144,7 +144,7 @@ func TestLocalEnvironmentWorkerSchedulesPreparationWithoutRemoteResolver(t *test turn, err := store.SessionAdapter(h.s).GetTurn(t.Context(), h.tenant, h.session.ID, start.RunID) return err == nil && turn.Status == sessions.TurnCompleted }) - settled, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, reservation.ID) + settled, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, reservation.ID) if err != nil || settled.State != sessions.EnvironmentInputAdmitted || len(settled.Receipts) != 1 { t.Fatal("local reservation did not settle", err) } diff --git a/services/core/internal/store/mcode_public_native_test.go b/services/core/internal/store/mcode_public_native_test.go index 9cd0dcd28..9e1b3ce01 100644 --- a/services/core/internal/store/mcode_public_native_test.go +++ b/services/core/internal/store/mcode_public_native_test.go @@ -89,7 +89,7 @@ func TestNativeMCodePublicExecution(t *testing.T) { if err != nil { t.Fatal(err) } - inputs, err := h.s.ListTurnInputs(ctx, h.tenant, proof.Session, proof.FirstTurn, 0, 100) + inputs, err := store.SessionAdapter(h.s).ListTurnInputs(ctx, h.tenant, proof.Session, proof.FirstTurn, 0, 100) if err != nil || len(inputs) != 2 { t.Fatal("steering input not in same turn", err) } diff --git a/services/core/internal/store/message_images_native_test.go b/services/core/internal/store/message_images_native_test.go index 40019ebc2..79aa40647 100644 --- a/services/core/internal/store/message_images_native_test.go +++ b/services/core/internal/store/message_images_native_test.go @@ -86,7 +86,7 @@ func TestNativeMessageImagePublicExecution(t *testing.T) { if json.Unmarshal(turn.Outcome, &outcome) != nil || outcome.AppliedThrough < 1 { t.Fatal("native input receipt missing") } - inputs, err := h.s.ListTurnInputs(ctx, h.tenant, proof.Session, proof.Turn, 0, 100) + inputs, err := store.SessionAdapter(h.s).ListTurnInputs(ctx, h.tenant, proof.Session, proof.Turn, 0, 100) if err != nil || len(inputs) != 3 || inputs[0].Kind != "message" || inputs[1].Kind != "message" || inputs[2].Kind != "tool_result" || outcome.AppliedThrough != inputs[2].Sequence { t.Fatal("active image batch was not applied exactly once in the same Turn", err) } diff --git a/services/core/internal/store/model_protocol_native_test.go b/services/core/internal/store/model_protocol_native_test.go index d0ee2dbd2..fac467b0e 100644 --- a/services/core/internal/store/model_protocol_native_test.go +++ b/services/core/internal/store/model_protocol_native_test.go @@ -125,7 +125,7 @@ func TestNativeModelProtocolPublicExecution(t *testing.T) { if err != nil || !call.Applied { t.Fatal("public function result lacks native delivery acknowledgement") } - inputs, err := h.s.ListTurnInputs(ctx, h.tenant, proof.Session, item.Turn, 0, 100) + inputs, err := store.SessionAdapter(h.s).ListTurnInputs(ctx, h.tenant, proof.Session, item.Turn, 0, 100) if err != nil || len(inputs) != 2 || inputs[0].Kind != "message" || inputs[1].Kind != "tool_result" { t.Fatal("public function result input was lost or duplicated") } diff --git a/services/core/internal/store/prepared_dispatch_failure_test.go b/services/core/internal/store/prepared_dispatch_failure_test.go index aef19d944..4c62c3518 100644 --- a/services/core/internal/store/prepared_dispatch_failure_test.go +++ b/services/core/internal/store/prepared_dispatch_failure_test.go @@ -24,7 +24,7 @@ func TestPreparedDispatchSettlesOnlyReadyInput(t *testing.T) { handle := acknowledgePreparation(h, frame.ID) switch action { case "cancel": - if _, err := h.s.CancelEnvironmentInput(t.Context(), h.tenant, h.session.ID, pending.ID); err != nil { + if _, err := store.CancelEnvironmentInput(t.Context(), h.s, h.tenant, h.session.ID, pending.ID); err != nil { t.Fatal(err) } case "expire": @@ -66,7 +66,7 @@ func TestPreparedDispatchSettlesOnlyReadyInput(t *testing.T) { } assertEnvironmentExpiryHasNoHistory(t, pool, h.session.ID) if action == "prepare-failure" || action == "disconnect" { - stored, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) + stored, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) if err != nil || stored.State != sessions.EnvironmentInputPending || !stored.Deadline.Equal(pending.Deadline) { t.Fatal("preparation failure changed the pending identity", stored, err) } @@ -87,7 +87,7 @@ func TestPreparedDispatchHandlesStartRejectionAndPendingStartCancellation(t *tes h.write(frame.ID, proto.TypePreparationStatus, proto.PreparationStatusPayload{State: "rejected", Operation: proto.TypeExecutionStart, ErrorCode: "preparation_not_ready"}) } else { h.write(frame.ID, proto.TypePreparationStatus, proto.PreparationStatusPayload{Handle: handle, Revision: 3, State: "starting", RunID: start.RunID}) - if _, err := h.s.RequestCancel(t.Context(), h.tenant, h.session.ID, "cancel-start"); err != nil { + if _, err := store.RequestCancel(t.Context(), h.s, h.tenant, h.session.ID, "cancel-start"); err != nil { t.Fatal(err) } frame := h.read(proto.TypePromptCancel) @@ -147,7 +147,7 @@ func TestPreparedDispatchCancellationReceiptSurvivesStartFailure(t *testing.T) { handle := acknowledgePreparation(h, prepare.ID) start := readyPreparedDispatch(t, h, prepare.ID, handle) h.write(prepare.ID, proto.TypePreparationStatus, proto.PreparationStatusPayload{Handle: handle, Revision: 3, State: "starting", RunID: start.RunID}) - if _, err := h.s.RequestCancel(t.Context(), h.tenant, h.session.ID, "cancel-start"); err != nil { + if _, err := store.RequestCancel(t.Context(), h.s, h.tenant, h.session.ID, "cancel-start"); err != nil { t.Fatal(err) } frame := h.read(proto.TypePromptCancel) diff --git a/services/core/internal/store/prepared_dispatch_test.go b/services/core/internal/store/prepared_dispatch_test.go index 4c30eadc6..5964853c0 100644 --- a/services/core/internal/store/prepared_dispatch_test.go +++ b/services/core/internal/store/prepared_dispatch_test.go @@ -25,7 +25,7 @@ func preparedDispatchHarness(t *testing.T) (*dispatchHarness, sessions.Environme assertNoRuntimeAllocation(t, h) h.d, h.lease = h.bound(), h.owner().Lease enableWorkerEnvironment(t, h) - pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}, {Kind: "message", Payload: json.RawMessage(`{"text":"second"}`)}}) + pending, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}, {Kind: "message", Payload: json.RawMessage(`{"text":"second"}`)}}) if err != nil { t.Fatal(err) } @@ -128,7 +128,7 @@ func TestPreparedDispatchOwnerOutlivesReservationDeadline(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) } - stored, err := h.d.Store.ExpireEnvironmentInput(t.Context(), h.tenant, h.session.ID, pending.ID) + stored, err := h.d.Sessions.ExpireEnvironmentInput(t.Context(), h.tenant, h.session.ID, pending.ID) if err != nil || stored.State != sessions.EnvironmentInputAdmitted { t.Fatal("admitted execution lost its owner to the pending-input deadline", err) } diff --git a/services/core/internal/store/public_execution_test.go b/services/core/internal/store/public_execution_test.go index 6b6d99340..681841b58 100644 --- a/services/core/internal/store/public_execution_test.go +++ b/services/core/internal/store/public_execution_test.go @@ -145,7 +145,7 @@ func TestWorkerRestartReconcilesClaimedButPreservesQueuedWork(t *testing.T) { } checkMeasurement(false) queued := publicSession(t, h, "queued") - if _, err := h.s.SubmitMessage(ctx, h.tenant, queued.ID, "first", json.RawMessage(`{"text":"Not sent"}`)); err != nil { + if _, err := store.SendMessage(ctx, h.s, h.tenant, queued.ID, "first", json.RawMessage(`{"text":"Not sent"}`)); err != nil { t.Fatal(err) } worker := startOwnedWorker(t, ctx, h.db, h.d, h.owner()) @@ -162,7 +162,7 @@ func TestWorkerRestartReconcilesClaimedButPreservesQueuedWork(t *testing.T) { if err != nil || pending.LastTurn.Status != sessions.TurnQueued { t.Fatal(pending, err) } - if _, err := h.s.RequestCancel(ctx, h.tenant, queued.ID, "stop-before-dispatch"); err != nil { + if _, err := store.RequestCancel(ctx, h.s, h.tenant, queued.ID, "stop-before-dispatch"); err != nil { t.Fatal(err) } pending, err = store.SessionAdapter(h.s).GetSession(ctx, h.tenant, queued.ID) diff --git a/services/core/internal/store/public_handler_fixture_test.go b/services/core/internal/store/public_handler_fixture_test.go index 198938cbd..594490d71 100644 --- a/services/core/internal/store/public_handler_fixture_test.go +++ b/services/core/internal/store/public_handler_fixture_test.go @@ -145,14 +145,14 @@ func withPolicy(policy execution.Policy) func(*api.Dependencies) { return func(d *api.Dependencies) { d.Policy = policy } } -// storeExecution admits Sessions and inputs through the Store without a -// Worker, so nothing runs them. +// storeExecution admits Sessions through the Store and inputs through the +// Session service without a Worker, so nothing runs them. func storeExecution(t testing.TB, s *store.Store) func(*api.Dependencies) { return func(d *api.Dependencies) { d.Execution = &api.Execution{ ExecutorURL: testExecutorURL, SessionAdmission: s, - InputAdmission: s, + InputAdmission: store.SessionService(t, s), SessionArchive: strictStandIn{t}, Workspaces: strictStandIn{t}, } diff --git a/services/core/internal/store/runtime_environment_terminal_test.go b/services/core/internal/store/runtime_environment_terminal_test.go index 3a1ed2d7f..d90215a2d 100644 --- a/services/core/internal/store/runtime_environment_terminal_test.go +++ b/services/core/internal/store/runtime_environment_terminal_test.go @@ -53,11 +53,11 @@ func TestManagedEnvironmentTerminationSettlesInputAndPreservesIdentity(t *testin if _, ok, err := sessionAdapter(s).GetDeviceCredential(t.Context(), owner.DeviceID); err != nil || ok { t.Fatal("terminal credential remained usable", err) } - failed, err := writer.PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID) + failed, err := sessionExecution(t, writer.lease).PromoteEnvironmentInput(t.Context(), tenant, session.ID, reservation.ID) if err != nil || failed.State != sessions.EnvironmentInputFailed || len(failed.Receipts) != 0 || failed.SettledAt == nil || !failed.Deadline.Equal(reservation.Deadline) { t.Fatal("late preparation resurrected failed input", failed, err) } - if _, err := s.ReserveEnvironmentInput(t.Context(), tenant, session.ID, "new", []sessions.Input{messageInput("later")}); !errors.Is(err, sessions.ErrEnvironmentUnavailable) || expired == errors.Is(err, sessions.ErrHostedEnvironmentFailed) { + if _, err := sessionService(t, s).ReserveEnvironmentInput(t.Context(), tenant, session.ID, "new", []sessions.Input{messageInput("later")}); !errors.Is(err, sessions.ErrEnvironmentUnavailable) || expired == errors.Is(err, sessions.ErrHostedEnvironmentFailed) { t.Fatal("terminal environment admitted new input", err) } if _, err := s.CreateSession(t.Context(), tenant, input); err != nil { diff --git a/services/core/internal/store/runtime_input_admission_test.go b/services/core/internal/store/runtime_input_admission_test.go index b80696f0a..346d3c3c2 100644 --- a/services/core/internal/store/runtime_input_admission_test.go +++ b/services/core/internal/store/runtime_input_admission_test.go @@ -17,7 +17,7 @@ func TestManagedRuntimeMaintenancePreservesCancelAndRetry(t *testing.T) { s, db := newManagedTestStoreDB(t) tenant, session, _ := managedSession(t, s, db) inputs := []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"accepted work"}`)}} - accepted, err := s.SubmitInputs(t.Context(), tenant, session.ID, "work", inputs) + accepted, err := store.SubmitInputs(t.Context(), s, tenant, session.ID, "work", inputs) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/runtime_suspension_concurrency_test.go b/services/core/internal/store/runtime_suspension_concurrency_test.go index 67aaf940f..0bd57d443 100644 --- a/services/core/internal/store/runtime_suspension_concurrency_test.go +++ b/services/core/internal/store/runtime_suspension_concurrency_test.go @@ -99,16 +99,16 @@ func TestRuntimeSuspensionPromotionRetainsPendingInput(t *testing.T) { s, w, pool, owner := runtimeSuspensionFixture(t) runtimeSuspensionSQL(t, pool, `UPDATE runtime_allocations SET compute_phase='waking',compute_retained_until=clock_timestamp()+interval '1 hour' WHERE id=$1`, owner.ID) pending := reserveEnvironmentInput(t, s, owner.TenantID, owner.SessionID, "during-wake") - if _, err := w.PromoteEnvironmentInput(t.Context(), owner.TenantID, owner.SessionID, pending.ID); !errors.Is(err, sessions.ErrTurnConflict) { + if _, err := sessionExecution(t, w.lease).PromoteEnvironmentInput(t.Context(), owner.TenantID, owner.SessionID, pending.ID); !errors.Is(err, sessions.ErrTurnConflict) { t.Fatal("waking allocation promoted model input", err) } - got, err := s.GetEnvironmentInputReservation(t.Context(), owner.TenantID, owner.SessionID, pending.ID) + got, err := sessionAdapter(s).GetEnvironmentInputReservation(t.Context(), owner.TenantID, owner.SessionID, pending.ID) if err != nil || got.State != sessions.EnvironmentInputPending || got.SettledAt != nil || len(got.Receipts) != 0 { t.Fatal("blocked promotion partially committed", got, err) } environmentInputHistory(t, pool, owner.SessionID, 0, 0) runtimeSuspensionSQL(t, pool, `UPDATE runtime_allocations SET compute_phase='running',compute_retained_until=NULL WHERE id=$1`, owner.ID) - got, err = w.PromoteEnvironmentInput(t.Context(), owner.TenantID, owner.SessionID, pending.ID) + got, err = sessionExecution(t, w.lease).PromoteEnvironmentInput(t.Context(), owner.TenantID, owner.SessionID, pending.ID) if err != nil || got.State != sessions.EnvironmentInputAdmitted || len(got.Receipts) != 2 { t.Fatal("pending request could not resume once running", got, err) } diff --git a/services/core/internal/store/runtime_wake_hint_integration_test.go b/services/core/internal/store/runtime_wake_hint_integration_test.go index d6744f12d..b7cfc0664 100644 --- a/services/core/internal/store/runtime_wake_hint_integration_test.go +++ b/services/core/internal/store/runtime_wake_hint_integration_test.go @@ -13,6 +13,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/execution" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sandbox" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" ) type wakeHintScanProvider struct { @@ -125,7 +126,7 @@ func (f *wakeHintIntegration) pending(t *testing.T, target wakeHintIntegrationTa "SELECT id::text FROM environment_input_reservations WHERE session_id=$1 AND idempotency_key=$2", target.session.ID, key).Scan(&id) == nil }) - pending, err := f.fixture.store.GetEnvironmentInputReservation(t.Context(), target.tenant, target.session.ID, id) + pending, err := store.SessionAdapter(f.fixture.store).GetEnvironmentInputReservation(t.Context(), target.tenant, target.session.ID, id) if err != nil || pending.State != sessions.EnvironmentInputPending { t.Fatal("input was not durably pending", pending.State, err) } @@ -196,7 +197,7 @@ func TestRuntimeWakeHintRejectedSubmitDoesNotAccelerateScan(t *testing.T) { } else { // Persist directly while the sentinel is blocked. Only the failing // Worker submission could emit a hint; Store persistence cannot. - if _, err := f.fixture.store.ReserveEnvironmentInput(t.Context(), f.target.tenant, f.target.session.ID, key, inputs); err != nil { + if _, err := store.SessionService(t, f.fixture.store).ReserveEnvironmentInput(t.Context(), f.target.tenant, f.target.session.ID, key, inputs); err != nil { t.Fatal(err) } if name == "idempotency conflict" { diff --git a/services/core/internal/store/runtime_worker_recovery_test.go b/services/core/internal/store/runtime_worker_recovery_test.go index 50c70e6fe..99e4a3275 100644 --- a/services/core/internal/store/runtime_worker_recovery_test.go +++ b/services/core/internal/store/runtime_worker_recovery_test.go @@ -36,7 +36,7 @@ func runtimeWorkerHarness(t *testing.T) (*dispatchHarness, *pgxpool.Pool) { func TestPreparedDispatchKeepsPendingReservationAfterComputeConflict(t *testing.T) { h, _ := runtimeWorkerHarness(t) h.d, h.lease = h.bound(), h.owner().Lease - pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) + pending, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) if err != nil { t.Fatal(err) } @@ -57,7 +57,7 @@ func TestPreparedDispatchKeepsPendingReservationAfterComputeConflict(t *testing. func TestWorkerWaitsForComputeAndSurvivesPromotionConflict(t *testing.T) { h, pool := runtimeWorkerHarness(t) - pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) + pending, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "pending", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) if err != nil { t.Fatal(err) } @@ -104,7 +104,7 @@ func TestWorkerWaitsForComputeAndSurvivesPromotionConflict(t *testing.T) { if release.ID != prepare.ID { t.Fatal("conflicted preparation was not released") } - stored, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) + stored, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) if err != nil || stored.State != sessions.EnvironmentInputPending || len(stored.Receipts) != 0 { t.Fatal("conflict consumed queued input", stored, err) } diff --git a/services/core/internal/store/sandbox_deployment_switch_test.go b/services/core/internal/store/sandbox_deployment_switch_test.go index 2c26d8739..63602a7ed 100644 --- a/services/core/internal/store/sandbox_deployment_switch_test.go +++ b/services/core/internal/store/sandbox_deployment_switch_test.go @@ -311,7 +311,7 @@ func TestSandboxSwitchPreservesReleasedAllocationAndItemHistory(t *testing.T) { if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(t.Context(), tenant, history.ID, "history", json.RawMessage(`{"text":"retained request"}`)) + input, err := sendMessage(t.Context(), s, tenant, history.ID, "history", json.RawMessage(`{"text":"retained request"}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/scheduling.go b/services/core/internal/store/scheduling.go deleted file mode 100644 index b094ed530..000000000 --- a/services/core/internal/store/scheduling.go +++ /dev/null @@ -1,26 +0,0 @@ -package store - -import ( - "context" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/sessionpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" - "github.com/google/uuid" -) - -func (s *Store) ListEnvironmentInputWork(ctx context.Context, after string, connectedDevices []string) ([]sessions.EnvironmentInputWork, error) { - id, devices, err := sessionpg.ExecutionWorkCursor(after, connectedDevices) - if err != nil { - return nil, err - } - rows, err := s.queries.ListEnvironmentInputWork(ctx, sqlc.ListEnvironmentInputWorkParams{AfterID: id, ConnectedDevices: devices}) - if err != nil { - return nil, err - } - work := make([]sessions.EnvironmentInputWork, 0, len(rows)) - for _, row := range rows { - work = append(work, sessions.EnvironmentInputWork{TenantID: uuid.UUID(row.TenantID.Bytes).String(), SessionID: uuid.UUID(row.SessionID.Bytes).String(), ReservationID: uuid.UUID(row.ID.Bytes).String()}) - } - return work, nil -} diff --git a/services/core/internal/store/self_hosted_cancel_public_test.go b/services/core/internal/store/self_hosted_cancel_public_test.go index 5da4bc866..aeff47ca7 100644 --- a/services/core/internal/store/self_hosted_cancel_public_test.go +++ b/services/core/internal/store/self_hosted_cancel_public_test.go @@ -118,7 +118,7 @@ func TestSelfHostedCancellationOfficialClient(t *testing.T) { return value } idleReceipts := receipts(created.IdleKey, "") - later, err := s.ReserveEnvironmentInput(t.Context(), tenant, created.LaterID, "controlled-later-input", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"Retain pending input."}`)}}) + later, err := store.SessionService(t, s).ReserveEnvironmentInput(t.Context(), tenant, created.LaterID, "controlled-later-input", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"Retain pending input."}`)}}) if err != nil || later.State != sessions.EnvironmentInputPending || later.IsInitial { t.Fatal("could not establish controlled later reservation", err) } @@ -132,7 +132,7 @@ func TestSelfHostedCancellationOfficialClient(t *testing.T) { start := func() string { t.Helper() // Controlled callbacks isolate HTTP admission; no daemon or model runs in this fixture. - input, err := s.SubmitMessage(t.Context(), tenant, created.ID, uuid.NewString(), json.RawMessage(`{"text":"Controlled active work."}`)) + input, err := store.SendMessage(t.Context(), s, tenant, created.ID, uuid.NewString(), json.RawMessage(`{"text":"Controlled active work."}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/self_hosted_initial_public_test.go b/services/core/internal/store/self_hosted_initial_public_test.go index 174f9d687..ef1b5e050 100644 --- a/services/core/internal/store/self_hosted_initial_public_test.go +++ b/services/core/internal/store/self_hosted_initial_public_test.go @@ -107,7 +107,7 @@ func TestSelfHostedInitialCreationOfficialClient(t *testing.T) { if err := pool.QueryRow(t.Context(), "SELECT id FROM environment_input_reservations WHERE session_id=$1 AND is_initial", item.ID).Scan(&id); err != nil { t.Fatal("public creation did not reserve initial input", err) } - reservation, err := s.GetEnvironmentInputReservation(t.Context(), tenant, item.ID, id) + reservation, err := store.SessionAdapter(s).GetEnvironmentInputReservation(t.Context(), tenant, item.ID, id) if err != nil || !reservation.IsInitial || len(reservation.Inputs) != 1 || len(reservation.Receipts) != 0 { t.Fatal("invalid public initial reservation", err) } diff --git a/services/core/internal/store/session_artifacts_public_test.go b/services/core/internal/store/session_artifacts_public_test.go index 8447e86de..d5e6cc722 100644 --- a/services/core/internal/store/session_artifacts_public_test.go +++ b/services/core/internal/store/session_artifacts_public_test.go @@ -37,7 +37,7 @@ func hostedArtifactSession(t *testing.T, s *store.Store, tenant, key string) (se // through artifacts and settled as completed, and returns the Turn ID. func completeArtifactTurn(t *testing.T, s *store.Store, artifacts *sessions.Service, tenant, session, environment, key string, outputs map[string]string) string { t.Helper() - receipt, err := s.SubmitMessage(t.Context(), tenant, session, key, json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"publish outputs"}]}]}`)) + receipt, err := store.SendMessage(t.Context(), s, tenant, session, key, json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"publish outputs"}]}]}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/session_artifacts_test.go b/services/core/internal/store/session_artifacts_test.go index e35b125f7..ba5957f9f 100644 --- a/services/core/internal/store/session_artifacts_test.go +++ b/services/core/internal/store/session_artifacts_test.go @@ -235,7 +235,7 @@ func TestSessionArtifactTransferDoesNotBlockDeletionOrCancellation(t *testing.T) t.Fatalf("transfer blocked deletion: %v", err) } } else { - if _, err := s.RequestCancel(ctx, tenant, session, "cancel-capture"); err != nil { + if _, err := requestCancel(ctx, s, tenant, session, "cancel-capture"); err != nil { t.Fatalf("transfer blocked cancellation: %v", err) } want = sessions.ErrTurnConflict diff --git a/services/core/internal/store/session_creation_stream_test.go b/services/core/internal/store/session_creation_stream_test.go index fd0e70e19..c7aee3ae1 100644 --- a/services/core/internal/store/session_creation_stream_test.go +++ b/services/core/internal/store/session_creation_stream_test.go @@ -121,7 +121,7 @@ func TestCreationStreamStartsBeforeOwnInputsAndRetriesAtUpsertCursor(t *testing. if err != nil || late.Created || late.Cursor != all[len(all)-1].Sequence { t.Fatal(late, err) } - next, err := s.SubmitInputs(ctx, tenant, id, "next", []sessions.Input{messageInput("later")}) + next, err := submitInputs(ctx, s, tenant, id, "next", []sessions.Input{messageInput("later")}) if err != nil { t.Fatal(err) } @@ -156,7 +156,7 @@ func TestCreationStreamIdleAndNonstreamRetry(t *testing.T) { if err != nil || !first.Created || first.Cursor != 0 || first.Session.LastTurn != nil { t.Fatal(first, err) } - if _, err := s.SubmitInputs(ctx, tenant, first.Session.ID, "message", []sessions.Input{messageInput("later")}); err != nil { + if _, err := submitInputs(ctx, s, tenant, first.Session.ID, "message", []sessions.Input{messageInput("later")}); err != nil { t.Fatal(err) } retry, err := s.CreateSessionStream(ctx, tenant, input) diff --git a/services/core/internal/store/session_deletion_execution_test.go b/services/core/internal/store/session_deletion_execution_test.go index fe40df03a..92798a4b2 100644 --- a/services/core/internal/store/session_deletion_execution_test.go +++ b/services/core/internal/store/session_deletion_execution_test.go @@ -120,7 +120,7 @@ func TestWaitingSessionCancelsThenDeletesThroughWorker(t *testing.T) { if again := functionState(t, h, 1); again.LastTurn == nil || again.LastTurn.Status != sessions.TurnWaiting || !again.LastTurn.CancelRequestedAt.IsZero() { t.Fatal("rejected deletion changed required actions", again) } - if _, err := h.s.RequestCancel(ctx, h.tenant, h.session.ID, "cancel-before-delete"); err != nil { + if _, err := store.RequestCancel(ctx, h.s, h.tenant, h.session.ID, "cancel-before-delete"); err != nil { t.Fatal(err) } // The explicit cancellation, not the rejected deletion, reaches the daemon. 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 04fd25a7f..8852cfda3 100644 --- a/services/core/internal/store/session_deletion_lifecycle_public_test.go +++ b/services/core/internal/store/session_deletion_lifecycle_public_test.go @@ -40,7 +40,6 @@ func TestSessionDeletionLifecyclePostgres(t *testing.T) { defer server.Close() client := pathIDClient{t: t, server: server} leased := executionOwner(t, db, s) - writer := leased.Store create := func(environment string, initial bool) sessions.Session { t.Helper() @@ -61,7 +60,7 @@ func TestSessionDeletionLifecyclePostgres(t *testing.T) { turn := func(to ...string) string { t.Helper() session := create(none, false) - receipt, err := s.SubmitMessage(ctx, tenant, session.ID, "input", json.RawMessage(`{"text":"work"}`)) + receipt, err := store.SendMessage(ctx, s, tenant, session.ID, "input", json.RawMessage(`{"text":"work"}`)) if err != nil { t.Fatal(err) } @@ -69,7 +68,7 @@ func TestSessionDeletionLifecyclePostgres(t *testing.T) { for _, status := range to { switch status { case "cancel": - _, err = s.RequestCancel(ctx, tenant, session.ID, "cancel") + _, err = store.RequestCancel(ctx, s, tenant, session.ID, "cancel") case "function": err = leased.Sessions.RecordFunctionCall(ctx, tenant, session.ID, receipt.TurnID, sessions.FunctionCall{CallID: "pending", ExecutorCallID: "native-pending", Name: "lookup", Arguments: json.RawMessage(`{}`)}) case sessions.TurnCompleted, sessions.TurnFailed: @@ -86,7 +85,7 @@ func TestSessionDeletionLifecyclePostgres(t *testing.T) { } reserve := func(session sessions.Session) sessions.EnvironmentInputReservation { t.Helper() - reservation, err := s.ReserveEnvironmentInput(ctx, tenant, session.ID, "later", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"later"}`)}}) + reservation, err := store.SessionService(t, s).ReserveEnvironmentInput(ctx, tenant, session.ID, "later", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"later"}`)}}) if err != nil || reservation.State != sessions.EnvironmentInputPending { t.Fatal(reservation, err) } @@ -139,12 +138,12 @@ func TestSessionDeletionLifecyclePostgres(t *testing.T) { if _, err := db.pool.Exec(ctx, "UPDATE environment_input_reservations SET deadline=clock_timestamp()-interval '1 second' WHERE session_id=$1", expired.ID); err != nil { t.Fatal(err) } - if count, err := writer.ExpireEnvironmentInputs(ctx); err != nil || count != 1 { + if count, err := leased.Sessions.ExpireEnvironmentInputs(ctx); err != nil || count != 1 { t.Fatal("initial input did not expire", count, err) } settled["self_hosted_input_expired"] = expired.ID withdrawn := create(selfHosted, false) - if _, err := s.CancelEnvironmentInput(ctx, tenant, withdrawn.ID, reserve(withdrawn).ID); err != nil { + if _, err := store.CancelEnvironmentInput(ctx, s, tenant, withdrawn.ID, reserve(withdrawn).ID); err != nil { t.Fatal(err) } settled["later_input_cancelled"] = withdrawn.ID diff --git a/services/core/internal/store/session_deletion_test.go b/services/core/internal/store/session_deletion_test.go index 04b1b5623..425307967 100644 --- a/services/core/internal/store/session_deletion_test.go +++ b/services/core/internal/store/session_deletion_test.go @@ -62,7 +62,7 @@ func TestSessionDeletionWaitsForSettledTurnAndRejectsAdmission(t *testing.T) { if err != nil { t.Fatal(err) } - receipt, err := s.SubmitMessage(ctx, tenant, session.ID, "input", json.RawMessage(`{"text":"retained"}`)) + receipt, err := sendMessage(ctx, s, tenant, session.ID, "input", json.RawMessage(`{"text":"retained"}`)) if err != nil { t.Fatal(err) } @@ -101,7 +101,7 @@ func TestSessionDeletionWaitsForSettledTurnAndRejectsAdmission(t *testing.T) { } // Callers cancel first. A queued Turn cancels at once; a running // Turn stays active until execution settles its cancellation. - if _, err := s.RequestCancel(ctx, tenant, session.ID, "cancel"); err != nil { + if _, err := requestCancel(ctx, s, tenant, session.ID, "cancel"); err != nil { t.Fatal(err) } if status == sessions.TurnInProgress { @@ -139,10 +139,10 @@ func TestSessionDeletionWaitsForSettledTurnAndRejectsAdmission(t *testing.T) { if _, err := fresh.CreateSessionStream(ctx, tenant, input); !errors.Is(err, sessions.ErrIdempotencyConflict) { t.Fatal(err) } - if _, err := fresh.SubmitMessage(ctx, tenant, session.ID, "input", json.RawMessage(`{"text":"retained"}`)); !errors.Is(err, sessions.ErrNotFound) { + if _, err := sendMessage(ctx, fresh, tenant, session.ID, "input", json.RawMessage(`{"text":"retained"}`)); !errors.Is(err, sessions.ErrNotFound) { t.Fatal(err) } - if _, err := fresh.RequestCancel(ctx, tenant, session.ID, "late-cancel"); !errors.Is(err, sessions.ErrNotFound) { + if _, err := requestCancel(ctx, fresh, tenant, session.ID, "late-cancel"); !errors.Is(err, sessions.ErrNotFound) { t.Fatal(err) } if _, err := sessionAdapter(fresh).ListItems(ctx, tenant, session.ID, "", 20, true); !errors.Is(err, sessions.ErrNotFound) { @@ -162,7 +162,7 @@ func TestSessionDeletionWaitsForSettledTurnAndRejectsAdmission(t *testing.T) { if _, err := transitionTurn(ctx, fresh, tenant, session.ID, receipt.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress}); !errors.Is(err, sessions.ErrTurnConflict) { t.Fatal(err) } - inputs, err := fresh.ListTurnInputs(ctx, tenant, session.ID, receipt.TurnID, 0, 20) + inputs, err := sessionAdapter(fresh).ListTurnInputs(ctx, tenant, session.ID, receipt.TurnID, 0, 20) if err != nil || len(inputs) == 0 || inputs[0].Sequence != receipt.Sequence { t.Fatal(inputs, err) } @@ -184,7 +184,7 @@ func TestSessionDeletionSerializesAdmissionBeforeRetryLookup(t *testing.T) { if err != nil { t.Fatal(err) } - if _, err := s.RequestCancel(ctx, tenant, session.ID, "existing"); err != nil { + if _, err := requestCancel(ctx, s, tenant, session.ID, "existing"); err != nil { t.Fatal(err) } tx, err := pool.Begin(ctx) @@ -196,7 +196,7 @@ func TestSessionDeletionSerializesAdmissionBeforeRetryLookup(t *testing.T) { t.Fatal(err) } done := make(chan error, 1) - go func() { _, err := s.RequestCancel(ctx, tenant, session.ID, "existing"); done <- err }() + go func() { _, err := requestCancel(ctx, s, tenant, session.ID, "existing"); done <- err }() if err := tx.Commit(ctx); err != nil { t.Fatal(err) } @@ -262,14 +262,14 @@ func TestSessionDeletionRacesAdmissionUnderSessionLock(t *testing.T) { tenant, session := newTurnSession(t, s) return tenant, session.ID }, func(ctx context.Context, s *Store, tenant, session string) error { - _, err := s.SubmitMessage(ctx, tenant, session, "racing", messagePayload) + _, err := sendMessage(ctx, s, tenant, session, "racing", messagePayload) return err }}, {"environment_input", func(t *testing.T, s *Store) (string, string) { tenant, session := environmentInputSession(t, s) return tenant, session.ID }, func(ctx context.Context, s *Store, tenant, session string) error { - _, err := s.ReserveEnvironmentInput(ctx, tenant, session, "racing", []sessions.Input{messageInput("racing")}) + _, err := sessionService(t, s).ReserveEnvironmentInput(ctx, tenant, session, "racing", []sessions.Input{messageInput("racing")}) return err }}, } @@ -423,7 +423,7 @@ func TestSessionDeletionKeepsProvisioningInputPlacementUntilSettled(t *testing.T if _, err := s.pool.Exec(ctx, "UPDATE environment_input_reservations SET deadline=clock_timestamp()-interval '1 second' WHERE session_id=$1", session.ID); err != nil { t.Fatal(err) } - if count, err := w.ExpireEnvironmentInputs(ctx); err != nil || count != 1 { + if count, err := sessionExecution(t, w.lease).ExpireEnvironmentInputs(ctx); err != nil || count != 1 { t.Fatal("initial input did not expire", count, err) } if err := sessionService(t, s).DeleteSession(ctx, sessions.DeleteSessionCommand{TenantID: tenant, SessionID: session.ID}); err != nil { diff --git a/services/core/internal/store/session_environment_snapshot_test.go b/services/core/internal/store/session_environment_snapshot_test.go index 56e9ce389..eda9b4990 100644 --- a/services/core/internal/store/session_environment_snapshot_test.go +++ b/services/core/internal/store/session_environment_snapshot_test.go @@ -25,7 +25,7 @@ func TestSelfHostedCreationSnapshotRetainsEnvironmentAndCursor(t *testing.T) { if environment.ID == "" || environment.SessionID != created.Session.ID || environment.TenantID != tenant || environment.Status != "pending" { t.Fatal("incorrect creation association", environment) } - pending, err := s.ReserveEnvironmentInput(t.Context(), tenant, created.Session.ID, "later", []sessions.Input{messageInput("later")}) + pending, err := sessionService(t, s).ReserveEnvironmentInput(t.Context(), tenant, created.Session.ID, "later", []sessions.Input{messageInput("later")}) if err != nil || pending.State != sessions.EnvironmentInputPending { t.Fatal(pending, err) } diff --git a/services/core/internal/store/session_events_test.go b/services/core/internal/store/session_events_test.go index 3b7748f58..1c761eb11 100644 --- a/services/core/internal/store/session_events_test.go +++ b/services/core/internal/store/session_events_test.go @@ -57,7 +57,7 @@ func TestSessionEventsCommitSnapshotsRetriesAndIsolation(t *testing.T) { if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(ctx, tenant, session.ID, "start", json.RawMessage(`{"text":"question"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "start", json.RawMessage(`{"text":"question"}`)) if err != nil { t.Fatal(err) } @@ -65,7 +65,7 @@ func TestSessionEventsCommitSnapshotsRetriesAndIsolation(t *testing.T) { if err != nil { t.Fatal(err) } - if _, err = s.SubmitMessage(ctx, tenant, session.ID, "start", json.RawMessage(`{"text":"question"}`)); err != nil { + if _, err = store.SendMessage(ctx, s, tenant, session.ID, "start", json.RawMessage(`{"text":"question"}`)); err != nil { t.Fatal(err) } after, _ := store.SessionAdapter(s).SessionEventCursor(ctx, tenant, session.ID) @@ -170,7 +170,7 @@ func TestSessionEventsRetentionAndQueuedCancellation(t *testing.T) { inputs[i] = sessions.Input{Kind: "message", Payload: json.RawMessage(`{"text":"input"}`)} } for range 5 { - if _, err = s.SubmitInputs(ctx, tenant, session.ID, uuid.NewString(), inputs); err != nil { + if _, err = store.SubmitInputs(ctx, s, tenant, session.ID, uuid.NewString(), inputs); err != nil { t.Fatal(err) } } @@ -185,7 +185,7 @@ func TestSessionEventsRetentionAndQueuedCancellation(t *testing.T) { if err != nil { t.Fatal(err) } - if _, err = s.RequestCancel(ctx, tenant, session.ID, "cancel"); err != nil { + if _, err = store.RequestCancel(ctx, s, tenant, session.ID, "cancel"); err != nil { t.Fatal(err) } changes, err := store.SessionAdapter(s).ListSessionEvents(ctx, tenant, session.ID, cursor) diff --git a/services/core/internal/store/session_initial_input.go b/services/core/internal/store/session_initial_input.go index 15c9642ed..77e89638b 100644 --- a/services/core/internal/store/session_initial_input.go +++ b/services/core/internal/store/session_initial_input.go @@ -21,15 +21,6 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/writeaudit" ) -func validateInitialInputs(inputs []sessions.Input) ([]sessions.Input, json.RawMessage, error) { - for _, input := range inputs { - if input.Kind != "message" { - return nil, nil, sessions.ErrInvalidInput - } - } - return validateInputs(inputs) -} - // The Session upsert locks retries. Only the new row reserves or admits work, so a // retry after completion or later Turns cannot submit the original input again. func (s *Store) createSessionResources(ctx context.Context, tenant string, params sqlc.CreateSessionParams, inputs []sessions.Input, encodedInput json.RawMessage, files []environmentconfig.InitialFile, setup environmentconfig.Setup, provider *v1.ModelProviderInput, executionConfiguration *v1.SessionExecutionConfiguration, providerSource string, deploymentRevision uuid.UUID) (sqlc.Session, *sessions.Environment, error) { @@ -170,10 +161,8 @@ func (s *Store) createSessionResources(ctx context.Context, tenant string, param return err } } else { - for position, input := range inputs { - if _, err := admitInput(ctx, q, tenant, row.ID, key, int32(position), input); err != nil { - return err - } + if _, err := sessions.AdmitInputs(ctx, sessionpg.BindSession(q, row.TenantID, row.ID), key, inputs); err != nil { + return err } } if err := sessionpg.PruneChanges(ctx, q, row.ID); err != nil { diff --git a/services/core/internal/store/session_initial_input_test.go b/services/core/internal/store/session_initial_input_test.go index 040701873..d8d568786 100644 --- a/services/core/internal/store/session_initial_input_test.go +++ b/services/core/internal/store/session_initial_input_test.go @@ -50,7 +50,7 @@ func TestInitialInputCreationRetriesAcrossConnectionsAndLaterTurns(t *testing.T) if first.LastTurn == nil { t.Fatal("missing initial Turn") } - inputs, err := s.ListTurnInputs(ctx, tenant, first.ID, first.LastTurn.ID, 0, 100) + inputs, err := sessionAdapter(s).ListTurnInputs(ctx, tenant, first.ID, first.LastTurn.ID, 0, 100) if err != nil || len(inputs) != 2 { t.Fatal(inputs, err) } @@ -67,11 +67,11 @@ func TestInitialInputCreationRetriesAcrossConnectionsAndLaterTurns(t *testing.T) transition(t, s, tenant, first.ID, first.LastTurn.ID, sessions.TurnQueued, sessions.TurnInProgress) transition(t, s, tenant, first.ID, first.LastTurn.ID, sessions.TurnInProgress, sessions.TurnCompleted) // The same caller key at the events endpoint is an independent request. - next, err := s.SubmitInputs(ctx, tenant, first.ID, input.IdempotencyKey, []sessions.Input{messageInput("later")}) + next, err := submitInputs(ctx, s, tenant, first.ID, input.IdempotencyKey, []sessions.Input{messageInput("later")}) if err != nil || len(next) != 1 || next[0].TurnID == first.LastTurn.ID { t.Fatal(next, err) } - if _, err := s.RequestCancel(ctx, tenant, first.ID, "cancel"); err != nil { + if _, err := requestCancel(ctx, s, tenant, first.ID, "cancel"); err != nil { t.Fatal(err) } if _, err := sessionService(t, s).UpdateSessionMetadata(ctx, sessions.UpdateSessionMetadataCommand{TenantID: tenant, SessionID: first.ID, Metadata: map[string]string{"updated": "yes"}}); err != nil { diff --git a/services/core/internal/store/session_metadata_test.go b/services/core/internal/store/session_metadata_test.go index a506b7b2a..d1be257f1 100644 --- a/services/core/internal/store/session_metadata_test.go +++ b/services/core/internal/store/session_metadata_test.go @@ -114,7 +114,7 @@ func TestSessionMetadataPreservesTerminalActivity(t *testing.T) { t.Fatal(err) } for _, status := range []string{sessions.TurnCompleted, sessions.TurnFailed, sessions.TurnCancelled} { - receipt, err := s.SubmitMessage(ctx, tenant, session.ID, uuid.NewString(), []byte(`{"text":"metadata fixture"}`)) + receipt, err := sendMessage(ctx, s, tenant, session.ID, uuid.NewString(), []byte(`{"text":"metadata fixture"}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/session_write_audit_test.go b/services/core/internal/store/session_write_audit_test.go index fef29312b..61249e770 100644 --- a/services/core/internal/store/session_write_audit_test.go +++ b/services/core/internal/store/session_write_audit_test.go @@ -143,11 +143,11 @@ func TestSessionWriteAuditRollback(t *testing.T) { case "delete": err = sessionService(t, s).DeleteSession(ctx, sessions.DeleteSessionCommand{TenantID: tenant, SessionID: created.ID}) case "events": - _, err = s.SubmitInputs(ctx, tenant, created.ID, "events", []sessions.Input{messageInput("private")}) + _, err = submitInputs(ctx, s, tenant, created.ID, "events", []sessions.Input{messageInput("private")}) case "noop": err = sessionService(t, s).AuditSessionOperation(ctx, sessions.AuditSessionOperationCommand{TenantID: tenant, SessionID: created.ID, Action: "send_events"}) case "reserve": - _, err = s.ReserveEnvironmentInput(ctx, tenant, created.ID, "reserve", []sessions.Input{messageInput("private")}) + _, err = sessionService(t, s).ReserveEnvironmentInput(ctx, tenant, created.ID, "reserve", []sessions.Input{messageInput("private")}) } if err == nil { t.Fatal("mutation bypassed audit failure") @@ -177,10 +177,10 @@ func TestEventsWriteAuditAdmissionAndReplay(t *testing.T) { } submit := func(ctx context.Context) error { if prepared { - _, err := s.ReserveEnvironmentInput(ctx, tenant, created.ID, "batch", []sessions.Input{messageInput("private")}) + _, err := sessionService(t, s).ReserveEnvironmentInput(ctx, tenant, created.ID, "batch", []sessions.Input{messageInput("private")}) return err } - _, err := s.SubmitInputs(ctx, tenant, created.ID, "batch", []sessions.Input{messageInput("private")}) + _, err := submitInputs(ctx, s, tenant, created.ID, "batch", []sessions.Input{messageInput("private")}) return err } first := sessionAuditContext(t, tenant, "a") diff --git a/services/core/internal/store/sessions.go b/services/core/internal/store/sessions.go index 95c89ba13..78dcf1f94 100644 --- a/services/core/internal/store/sessions.go +++ b/services/core/internal/store/sessions.go @@ -89,7 +89,7 @@ func (s *Store) createSession(ctx context.Context, tenantID string, input sessio var batch []sessions.Input var encodedInput json.RawMessage if len(input.InitialInputs) > 0 { - batch, encodedInput, err = validateInitialInputs(input.InitialInputs) + batch, encodedInput, err = sessions.ValidateMessageInputs(input.InitialInputs) if err != nil { return sessions.Creation{}, err } diff --git a/services/core/internal/store/steering_receipts_test.go b/services/core/internal/store/steering_receipts_test.go index beff34513..69a7ab13e 100644 --- a/services/core/internal/store/steering_receipts_test.go +++ b/services/core/internal/store/steering_receipts_test.go @@ -9,6 +9,7 @@ import ( "github.com/MiniMax-AI/OpenAgentCore/internal/agentdaemon/proto" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/execution" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" ) func TestExecutionDurableInputReceiptLifetime(t *testing.T) { @@ -42,7 +43,7 @@ func TestExecutionDurableInputReceiptLifetime(t *testing.T) { case "retry-after-write": h.write(first.TurnID, proto.TypePromptSteerAck, proto.PromptSteerAckPayload{InputID: input.InputID, ErrorCode: "not_ready"}) case "cancel-unknown-first": - if _, err := h.s.RequestCancel(ctx, h.tenant, h.session.ID, "cancel"); err != nil { + if _, err := store.RequestCancel(ctx, h.s, h.tenant, h.session.ID, "cancel"); err != nil { t.Fatal(err) } var request proto.PromptCancelPayload diff --git a/services/core/internal/store/stream_authority_http_test.go b/services/core/internal/store/stream_authority_http_test.go index a03a3fc5c..1d68d49eb 100644 --- a/services/core/internal/store/stream_authority_http_test.go +++ b/services/core/internal/store/stream_authority_http_test.go @@ -101,7 +101,7 @@ func TestLiveStreamClosesAfterKeyRevocationOrProjectArchive(t *testing.T) { t.Fatal("revocation affected peer", valid.StatusCode) } // New events remain available to valid callers after the reader has closed. - if _, err := s.ReserveEnvironmentInput(t.Context(), project.TenantID, session.ID, "after-revocation", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"new event"}`)}}); err != nil { + if _, err := store.SessionService(t, s).ReserveEnvironmentInput(t.Context(), project.TenantID, session.ID, "after-revocation", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"new event"}`)}}); err != nil { t.Fatal(err) } } diff --git a/services/core/internal/store/token_usage_integration_test.go b/services/core/internal/store/token_usage_integration_test.go index a042c21c2..dfddb68fd 100644 --- a/services/core/internal/store/token_usage_integration_test.go +++ b/services/core/internal/store/token_usage_integration_test.go @@ -34,7 +34,7 @@ func TestTokenUsageDurableSnapshotsAndSessionTotals(t *testing.T) { } } for n, status := range []string{sessions.TurnFailed, sessions.TurnCancelled} { - admission, err := s.SubmitMessage(ctx, tenant, session.ID, fmt.Sprint(n), json.RawMessage(`{"text":"measure"}`)) + admission, err := store.SendMessage(ctx, s, tenant, session.ID, fmt.Sprint(n), json.RawMessage(`{"text":"measure"}`)) if err != nil { t.Fatal(err) } @@ -114,7 +114,7 @@ func TestCancellationReceiptUsageSurvivesRecovery(t *testing.T) { if err != nil { t.Fatal(err) } - admission, err := s.SubmitMessage(ctx, tenant, session.ID, "start", json.RawMessage(`{"text":"measure"}`)) + admission, err := store.SendMessage(ctx, s, tenant, session.ID, "start", json.RawMessage(`{"text":"measure"}`)) if err != nil { t.Fatal(err) } @@ -188,7 +188,7 @@ func TestSessionUsageRequiresEveryRootTurnEndedAndMeasured(t *testing.T) { } submit := func(key string) sessions.InputReceipt { t.Helper() - admission, err := s.SubmitMessage(ctx, tenant, session.ID, key, json.RawMessage(`{"text":"measure"}`)) + admission, err := store.SendMessage(ctx, s, tenant, session.ID, key, json.RawMessage(`{"text":"measure"}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/tool_policy_native_test.go b/services/core/internal/store/tool_policy_native_test.go index bfa83bfc0..380a0c4f7 100644 --- a/services/core/internal/store/tool_policy_native_test.go +++ b/services/core/internal/store/tool_policy_native_test.go @@ -95,7 +95,7 @@ func TestNativeToolPolicyPublicExecution(t *testing.T) { if err != nil { t.Fatal(err) } - inputs, err := h.s.ListTurnInputs(ctx, h.tenant, item.ID, item.FirstTurn, 0, 100) + inputs, err := store.SessionAdapter(h.s).ListTurnInputs(ctx, h.tenant, item.ID, item.FirstTurn, 0, 100) var outcome execution.Result if err != nil || len(inputs) != 1 || json.Unmarshal(turn.Outcome, &outcome) != nil || outcome.AppliedThrough != inputs[0].Sequence { t.Fatal("native text input receipt missing", err) diff --git a/services/core/internal/store/turn_events_test.go b/services/core/internal/store/turn_events_test.go index bd4c1c9ec..0e8b8ba0c 100644 --- a/services/core/internal/store/turn_events_test.go +++ b/services/core/internal/store/turn_events_test.go @@ -21,7 +21,7 @@ func TestTurnEventBatchesAreOrderedIsolatedAndDurable(t *testing.T) { if err != nil { t.Fatal(err) } - input, err := s.SubmitMessage(ctx, tenant, session.ID, "start", json.RawMessage(`{"text":"test"}`)) + input, err := store.SendMessage(ctx, s, tenant, session.ID, "start", json.RawMessage(`{"text":"test"}`)) if err != nil { t.Fatal(err) } diff --git a/services/core/internal/store/turn_inputs.go b/services/core/internal/store/turn_inputs.go deleted file mode 100644 index 3e896dafd..000000000 --- a/services/core/internal/store/turn_inputs.go +++ /dev/null @@ -1,238 +0,0 @@ -package store - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "slices" - - "github.com/google/uuid" - "github.com/jackc/pgx/v5" - "github.com/jackc/pgx/v5/pgtype" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/db/sqlc" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/jsonobject" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/auditpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/sessionpg" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" -) - -// SubmitMessage and RequestCancel use the same request-level admission as batches. -func (s *Store) SubmitMessage(ctx context.Context, tenantID, sessionID, key string, payload json.RawMessage) (sessions.InputReceipt, error) { - return s.submitOne(ctx, tenantID, sessionID, key, sessions.Input{Kind: "message", Payload: payload}) -} - -func (s *Store) RequestCancel(ctx context.Context, tenantID, sessionID, key string) (sessions.InputReceipt, error) { - return s.submitOne(ctx, tenantID, sessionID, key, sessions.Input{Kind: "cancel", Payload: json.RawMessage(`{}`)}) -} - -func (s *Store) submitOne(ctx context.Context, tenantID, sessionID, key string, input sessions.Input) (sessions.InputReceipt, error) { - receipts, err := s.SubmitInputs(ctx, tenantID, sessionID, key, []sessions.Input{input}) - if err != nil { - return sessions.InputReceipt{}, err - } - return receipts[0], nil -} - -// SubmitInputs commits a request in order under one Session lock. The entire -// batch is the retry identity; replay never re-evaluates a cancellation target. -// Internal receipts are not the response body of the public events endpoint. -func (s *Store) SubmitInputs(ctx context.Context, tenantID, sessionID, key string, inputs []sessions.Input) ([]sessions.InputReceipt, error) { - if err := sessions.ValidateInputKey(key); err != nil { - return nil, err - } - batch, encoded, err := validateInputs(inputs) - if err != nil { - return nil, err - } - tenant, err := parseID(tenantID) - if err != nil { - return nil, err - } - receipts := make([]sessions.InputReceipt, 0, len(batch)) - err = s.withPublicSession(ctx, tenantID, sessionID, func(ctx context.Context, q *sqlc.Queries, session pgtype.UUID) error { - previous, err := inputBatchReceipts(ctx, q, session, key, encoded) - if err != nil { - return err - } - if len(previous) > 0 { - receipts = previous - return auditpg.RecordWriteAudit(ctx, q, tenantID, "send_events", "session", uuid.UUID(session.Bytes).String(), "") - } - if slices.ContainsFunc(batch, func(input sessions.Input) bool { return input.Kind == "message" }) { - if err := sessions.CheckFileWriteGate(ctx, sessionpg.BindSession(q, tenant, session)); err != nil { - return err - } - } - if err := checkEnvironmentInputGate(ctx, q, session, key, encoded); err != nil { - return err - } - for position, input := range batch { - receipt, err := admitInput(ctx, q, tenantID, session, key, int32(position), input) - if err != nil { - return err - } - receipts = append(receipts, receipt) - } - return auditpg.RecordWriteAudit(ctx, q, tenantID, "send_events", "session", uuid.UUID(session.Bytes).String(), "") - }) - if err != nil { - return nil, fmt.Errorf("submit turn inputs: %w", err) - } - return receipts, nil -} - -func inputBatchReceipts(ctx context.Context, q *sqlc.Queries, session pgtype.UUID, key string, batch json.RawMessage) ([]sessions.InputReceipt, error) { - rows, err := q.FindInputBatch(ctx, sqlc.FindInputBatchParams{SessionID: session, IdempotencyKey: key, Batch: batch}) - if err != nil { - return nil, err - } - receipts := make([]sessions.InputReceipt, 0, len(rows)) - for _, row := range rows { - if !row.Matches { - return nil, sessions.ErrIdempotencyConflict - } - receipts = append(receipts, inputReceipt(row.Sequence, row.TurnID, true)) - } - return receipts, nil -} - -func validateInputs(inputs []sessions.Input) ([]sessions.Input, json.RawMessage, error) { - if len(inputs) == 0 || len(inputs) > 64 { - return nil, nil, fmt.Errorf("%w: input batch must contain 1..64 events", sessions.ErrInvalidInput) - } - batch := make([]sessions.Input, len(inputs)) - size := 0 - for i, input := range inputs { - size += len(input.Payload) - if size > 512*1024 || len(input.Payload) == 0 || (input.Kind != "message" && input.Kind != "cancel" && input.Kind != "tool_result") { - return nil, nil, fmt.Errorf("%w: input payloads must be nonempty and total at most 512 KiB", sessions.ErrInvalidInput) - } - payload, err := jsonobject.Normalize(input.Payload) - if err != nil { - return nil, nil, fmt.Errorf("%w: %w", sessions.ErrInvalidInput, err) - } - if input.Kind == "cancel" && string(payload) != "{}" { - return nil, nil, fmt.Errorf("%w: cancel payload must be empty", sessions.ErrInvalidInput) - } - if input.Kind == "tool_result" { - if _, err := sessions.ParseFunctionResultInput(payload); err != nil { - return nil, nil, err - } - } - batch[i] = sessions.Input{Kind: input.Kind, Payload: payload} - } - encoded, err := json.Marshal(batch) - return batch, encoded, err -} - -func admitInput(ctx context.Context, q *sqlc.Queries, tenantID string, session pgtype.UUID, key string, position int32, input sessions.Input) (sessions.InputReceipt, error) { - if input.Kind == "tool_result" { - result, err := sessions.ParseFunctionResultInput(input.Payload) - if err != nil { - return sessions.InputReceipt{}, err - } - tenant, err := parseID(tenantID) - if err != nil { - return sessions.InputReceipt{}, err - } - turn, err := sessions.AdmitFunctionResult(ctx, sessionpg.BindSession(q, tenant, session), result) - if err != nil { - return sessions.InputReceipt{}, err - } - id, err := parseID(turn.ID) - if err != nil { - return sessions.InputReceipt{}, err - } - sequence, err := q.CreateTurnInput(ctx, sqlc.CreateTurnInputParams{ - SessionID: session, TurnID: id, IdempotencyKey: key, Kind: input.Kind, Payload: input.Payload, BatchPosition: position, - }) - if err != nil { - return sessions.InputReceipt{}, err - } - return inputReceipt(sequence, id, false), nil - } - created := false - turn, err := q.GetActiveTurn(ctx, session) - if errors.Is(err, pgx.ErrNoRows) { - if input.Kind == "message" { - turn, err = q.CreateTurn(ctx, sqlc.CreateTurnParams{ID: pgtype.UUID{Bytes: uuid.New(), Valid: true}, SessionID: session}) - if err == nil { - created = true - err = sessionpg.AppendChanges(ctx, q, session, sessions.TurnChanges(sessionpg.TurnFromRow(turn), true)...) - } - } else { - err = nil // Retain even an idle cancellation's retry identity. - } - } - if err != nil { - return sessions.InputReceipt{}, err - } - sequence, err := q.CreateTurnInput(ctx, sqlc.CreateTurnInputParams{ - SessionID: session, TurnID: turn.ID, IdempotencyKey: key, Kind: input.Kind, Payload: input.Payload, BatchPosition: position, - }) - if err != nil { - return sessions.InputReceipt{}, err - } - tenant, err := parseID(tenantID) - if err != nil { - return sessions.InputReceipt{}, err - } - bound := sessionpg.BindSession(q, tenant, session) - if input.Kind == "cancel" && turn.ID.Valid { - if err := sessions.CancelTurn(ctx, bound, sessionpg.TurnFromRow(turn)); err != nil { - return sessions.InputReceipt{}, err - } - } - if err := sessions.ProjectInput(ctx, bound, sequence); err != nil { - return sessions.InputReceipt{}, err - } - if created { - // A new Turn publishes turn.created, then its user input Items, then the - // Session activity, within this transaction. - usage, err := sessionpg.LoadUsage(ctx, q, session) - if err != nil { - return sessions.InputReceipt{}, err - } - if err := sessionpg.AppendChanges(ctx, q, session, sessions.ActivityChange(sessionpg.TurnFromRow(turn), usage, nil)); err != nil { - return sessions.InputReceipt{}, err - } - } - return inputReceipt(sequence, turn.ID, false), nil -} - -// ListTurnInputs is an internal ordered recovery query, not the public SSE stream. -func (s *Store) ListTurnInputs(ctx context.Context, tenantID, sessionID, turnID string, after int64, limit int) ([]sessions.TurnInput, error) { - params, err := sessionpg.TurnLookup(tenantID, sessionID, turnID) - if err != nil { - return nil, err - } - if after < 0 || limit < 1 || limit > 100 { - return nil, fmt.Errorf("%w: nonnegative cursor and page size 1..100 required", sessions.ErrInvalidInput) - } - if _, err := s.queries.GetTurn(ctx, params); errors.Is(err, pgx.ErrNoRows) { - return nil, sessions.ErrNotFound - } else if err != nil { - return nil, fmt.Errorf("get turn: %w", err) - } - rows, err := s.queries.ListTurnInputs(ctx, sqlc.ListTurnInputsParams{ - TenantID: params.TenantID, SessionID: params.SessionID, TurnID: params.ID, Sequence: after, Limit: int32(limit), - }) - if err != nil { - return nil, fmt.Errorf("list turn inputs: %w", err) - } - inputs := make([]sessions.TurnInput, 0, len(rows)) - for _, row := range rows { - inputs = append(inputs, sessions.TurnInput{Sequence: row.Sequence, Kind: row.Kind, Payload: row.Payload, CreatedAt: row.CreatedAt.Time}) - } - return inputs, nil -} - -func inputReceipt(sequence int64, turn pgtype.UUID, replayed bool) sessions.InputReceipt { - receipt := sessions.InputReceipt{Sequence: sequence, Replayed: replayed} - if turn.Valid { - receipt.TurnID = uuid.UUID(turn.Bytes).String() - } - return receipt -} diff --git a/services/core/internal/store/turn_inputs_test.go b/services/core/internal/store/turn_inputs_test.go deleted file mode 100644 index 3ffc28e26..000000000 --- a/services/core/internal/store/turn_inputs_test.go +++ /dev/null @@ -1,219 +0,0 @@ -package store - -import ( - "context" - "encoding/json" - "errors" - "fmt" - "reflect" - "sync" - "testing" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" - "github.com/google/uuid" -) - -var messagePayload = json.RawMessage(`{"input":[{"role":"user","content":[{"type":"input_text","text":"hello"}]}]}`) - -func newTurnSession(t *testing.T, s *Store) (string, sessions.Session) { - t.Helper() - tenant := uuid.NewString() - session, err := s.CreateSession(context.Background(), tenant, sessions.CreateSession{Creator: FixtureCreator(), - Engine: "codex", IdempotencyKey: "session", Configuration: json.RawMessage(`{"agent":{"model":"test","instructions":"original"}}`), - }) - if err != nil { - t.Fatal(err) - } - return tenant, session -} - -func submitMessage(t *testing.T, s *Store, tenant, session, key string) sessions.InputReceipt { - t.Helper() - receipt, err := s.SubmitMessage(context.Background(), tenant, session, key, messagePayload) - if err != nil { - t.Fatal(err) - } - return receipt -} - -func transition(t *testing.T, s *Store, tenant, session, turn, from, to string) sessions.Turn { - t.Helper() - got, err := transitionTurn(context.Background(), s, tenant, session, turn, sessions.TurnTransition{ExpectedStatus: from, Status: to}) - if err != nil { - t.Fatal(err) - } - return got -} - -func TestConcurrentInputsUseOneTurnAndOneRetryReceipt(t *testing.T) { - s, _ := testStore(t) - otherStore, _ := testStore(t) - tenant, session := newTurnSession(t, s) - ctx := context.Background() - const count = 8 - for _, repeatedKey := range []bool{true, false} { - t.Run(fmt.Sprintf("repeated-key-%v", repeatedKey), func(t *testing.T) { - var wg sync.WaitGroup - receipts := make(chan sessions.InputReceipt, count) - errs := make(chan error, count) - for i := range count { - wg.Add(1) - go func() { - defer wg.Done() - key := "same" - if !repeatedKey { - key = fmt.Sprintf("distinct-%d", i) - } - store := s - if i%2 == 0 { - store = otherStore - } - r, err := store.SubmitMessage(ctx, tenant, session.ID, key, messagePayload) - receipts <- r - errs <- err - }() - } - wg.Wait() - close(receipts) - close(errs) - for err := range errs { - if err != nil { - t.Fatal(err) - } - } - turns, sequences := map[string]bool{}, map[int64]bool{} - newReceipts := 0 - for r := range receipts { - turns[r.TurnID], sequences[r.Sequence] = true, true - if !r.Replayed { - newReceipts++ - } - } - want := count - if repeatedKey { - want = 1 - } - if len(turns) != 1 || turns[""] || len(sequences) != want || newReceipts != want { - t.Fatalf("turns=%v sequences=%v new=%d", turns, sequences, newReceipts) - } - }) - } - first := submitMessage(t, s, tenant, session.ID, "same") - inputs, err := s.ListTurnInputs(ctx, tenant, session.ID, first.TurnID, 0, 100) - if err != nil || len(inputs) != count+1 { - t.Fatalf("lost or duplicated inputs: %d, %v", len(inputs), err) - } -} - -func TestTurnInputRetriesAndRestart(t *testing.T) { - s, pool := testStore(t) - tenant, session := newTurnSession(t, s) - ctx := context.Background() - first := submitMessage(t, s, tenant, session.ID, "first") - transition(t, s, tenant, session.ID, first.TurnID, sessions.TurnQueued, sessions.TurnInProgress) - steer := submitMessage(t, s, tenant, session.ID, "steer") - if steer.TurnID != first.TurnID || steer.Sequence <= first.Sequence { - t.Fatalf("active message did not steer: %+v", steer) - } - reordered := json.RawMessage(`{ "input": [{"content":[{"text":"hello","type":"input_text"}],"role":"user"}] }`) - retry, err := s.SubmitMessage(ctx, tenant, session.ID, "first", reordered) - if err != nil || !retry.Replayed || retry.Sequence != first.Sequence || retry.TurnID != first.TurnID { - t.Fatalf("equivalent retry = %+v, %v", retry, err) - } - if _, err := s.SubmitMessage(ctx, tenant, session.ID, "first", json.RawMessage(`{"text":"changed"}`)); !errors.Is(err, sessions.ErrIdempotencyConflict) { - t.Fatalf("changed payload accepted: %v", err) - } - if _, err := s.RequestCancel(ctx, tenant, session.ID, "first"); !errors.Is(err, sessions.ErrIdempotencyConflict) { - t.Fatalf("changed input kind accepted: %v", err) - } - completed := transition(t, s, tenant, session.ID, first.TurnID, sessions.TurnInProgress, sessions.TurnCompleted) - next := submitMessage(t, s, tenant, session.ID, "next") - if next.TurnID == first.TurnID { - t.Fatal("idle message did not start a new Turn") - } - pool.Close() - recovered, _ := testStore(t) - retry = submitMessage(t, recovered, tenant, session.ID, "first") - if !retry.Replayed || retry.TurnID != first.TurnID || retry.Sequence != first.Sequence { - t.Fatalf("restart retry changed target: %+v", retry) - } - got, err := sessionAdapter(recovered).GetTurn(ctx, tenant, session.ID, first.TurnID) - if err != nil || !reflect.DeepEqual(got, completed) { - t.Fatalf("restart turn: %+v, %v", got, err) - } - var all []sessions.TurnInput - var cursor int64 - for { - page, err := recovered.ListTurnInputs(ctx, tenant, session.ID, first.TurnID, cursor, 1) - if err != nil { - t.Fatal(err) - } - if len(page) == 0 { - break - } - if page[0].Sequence <= cursor || page[0].CreatedAt.IsZero() { - t.Fatalf("invalid ordered input: %+v", page) - } - cursor = page[0].Sequence - all = append(all, page...) - } - if len(all) != 2 || all[0].Sequence != first.Sequence || all[1].Sequence != steer.Sequence { - t.Fatalf("recovered inputs = %+v", all) - } - snapshot, err := sessionAdapter(recovered).GetSession(ctx, tenant, session.ID) - if err != nil || snapshot.LastTurn == nil || snapshot.LastTurn.ID != next.TurnID || snapshot.LastTurn.Status != sessions.TurnQueued { - t.Fatal("latest Session activity did not survive restart", err) - } - snapshot.LastTurn = nil - if !reflect.DeepEqual(snapshot, session) { - t.Fatal("turn submission mutated the Session snapshot", err) - } -} - -func TestTurnOperationsAreTenantAndSessionScoped(t *testing.T) { - s, _ := testStore(t) - tenant, session := newTurnSession(t, s) - otherTenant, other := newTurnSession(t, s) - ctx := context.Background() - first := submitMessage(t, s, tenant, session.ID, "input") - for _, scope := range []struct{ tenant, session string }{{otherTenant, session.ID}, {tenant, other.ID}, {tenant, uuid.NewString()}} { - for name, call := range map[string]func() error{ - "submit": func() error { - _, err := s.SubmitMessage(ctx, scope.tenant, scope.session, "input", messagePayload) - return err - }, - "cancel": func() error { _, err := s.RequestCancel(ctx, scope.tenant, scope.session, "cancel"); return err }, - "read": func() error { - _, err := sessionAdapter(s).GetTurn(ctx, scope.tenant, scope.session, first.TurnID) - return err - }, - "inputs": func() error { - _, err := s.ListTurnInputs(ctx, scope.tenant, scope.session, first.TurnID, 0, 10) - return err - }, - "transition": func() error { - _, err := transitionTurn(ctx, s, scope.tenant, scope.session, first.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnFailed}) - return err - }, - } { - if err := call(); !errors.Is(err, sessions.ErrNotFound) { - t.Fatalf("%s escaped scope: %v", name, err) - } - } - } - // Turn IDs cannot be used with another valid Session in the same tenant either. - second, err := s.CreateSession(ctx, tenant, sessions.CreateSession{Creator: FixtureCreator(), Engine: "codex", IdempotencyKey: "second"}) - if err != nil { - t.Fatal(err) - } - if _, err := sessionAdapter(s).GetTurn(ctx, tenant, second.ID, first.TurnID); !errors.Is(err, sessions.ErrNotFound) { - t.Fatalf("cross-session turn read: %v", err) - } - if _, err := transitionTurn(ctx, s, tenant, second.ID, first.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnFailed}); !errors.Is(err, sessions.ErrNotFound) { - t.Fatalf("cross-session turn write: %v", err) - } - otherInput := submitMessage(t, s, otherTenant, other.ID, "input") - if otherInput.TurnID == first.TurnID { - t.Fatal("retry identity leaked across Sessions") - } -} diff --git a/services/core/internal/store/turns_test.go b/services/core/internal/store/turns_test.go deleted file mode 100644 index 4e85732e1..000000000 --- a/services/core/internal/store/turns_test.go +++ /dev/null @@ -1,114 +0,0 @@ -package store - -import ( - "context" - "encoding/json" - "errors" - "testing" - - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" -) - -func TestCancellationStaysBoundToItsOriginalTurn(t *testing.T) { - s, _ := testStore(t) - tenant, session := newTurnSession(t, s) - ctx := context.Background() - idle, err := s.RequestCancel(ctx, tenant, session.ID, "idle-cancel") - if err != nil || idle.TurnID != "" || idle.Sequence == 0 { - t.Fatalf("idle cancellation: %+v, %v", idle, err) - } - first := submitMessage(t, s, tenant, session.ID, "first") - started := transition(t, s, tenant, session.ID, first.TurnID, sessions.TurnQueued, sessions.TurnInProgress) - if started.StartedAt.IsZero() || !started.CompletedAt.IsZero() { - t.Fatalf("started timestamps: %+v", started) - } - cancel, err := s.RequestCancel(ctx, tenant, session.ID, "cancel") - if err != nil || cancel.TurnID != first.TurnID { - t.Fatalf("cancellation target: %+v, %v", cancel, err) - } - pending, err := sessionAdapter(s).GetTurn(ctx, tenant, session.ID, first.TurnID) - if err != nil || pending.Status != sessions.TurnInProgress || pending.CancelRequestedAt.IsZero() || !pending.CompletedAt.IsZero() { - t.Fatalf("running cancellation falsely completed: %+v, %v", pending, err) - } - transition(t, s, tenant, session.ID, first.TurnID, sessions.TurnInProgress, sessions.TurnCancelled) - next := submitMessage(t, s, tenant, session.ID, "next") - for key, original := range map[string]sessions.InputReceipt{"cancel": cancel, "idle-cancel": idle} { - retry, err := s.RequestCancel(ctx, tenant, session.ID, key) - if err != nil || !retry.Replayed || retry.TurnID != original.TurnID || retry.Sequence != original.Sequence { - t.Fatalf("cancellation retargeted: %+v, %v", retry, err) - } - } - queued, err := sessionAdapter(s).GetTurn(ctx, tenant, session.ID, next.TurnID) - if err != nil || queued.Status != sessions.TurnQueued || !queued.CancelRequestedAt.IsZero() { - t.Fatalf("old cancellation affected later Turn: %+v, %v", queued, err) - } - if _, err := s.RequestCancel(ctx, tenant, session.ID, "cancel-queued"); err != nil { - t.Fatal(err) - } - stopped, err := sessionAdapter(s).GetTurn(ctx, tenant, session.ID, next.TurnID) - if err != nil || stopped.Status != sessions.TurnCancelled || stopped.CompletedAt.IsZero() || !stopped.StartedAt.IsZero() { - t.Fatalf("queued work did not stop: %+v, %v", stopped, err) - } - if _, err := transitionTurn(ctx, s, tenant, session.ID, next.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress}); !errors.Is(err, sessions.ErrTurnConflict) { - t.Fatalf("cancelled queued work was started: %v", err) - } -} - -func TestWaitingTurnRetainsInputsAndStartTime(t *testing.T) { - s, _ := testStore(t) - tenant, session := newTurnSession(t, s) - ctx := context.Background() - first := submitMessage(t, s, tenant, session.ID, "first") - started := transition(t, s, tenant, session.ID, first.TurnID, sessions.TurnQueued, sessions.TurnInProgress) - transition(t, s, tenant, session.ID, first.TurnID, sessions.TurnInProgress, sessions.TurnWaiting) - steer := submitMessage(t, s, tenant, session.ID, "steer") - if steer.TurnID != first.TurnID { - t.Fatal("waiting input started another Turn") - } - resumed := transition(t, s, tenant, session.ID, first.TurnID, sessions.TurnWaiting, sessions.TurnInProgress) - if !started.StartedAt.Equal(resumed.StartedAt) { - t.Fatal("resume reset the start time") - } - transition(t, s, tenant, session.ID, first.TurnID, sessions.TurnInProgress, sessions.TurnWaiting) - if _, err := s.RequestCancel(ctx, tenant, session.ID, "cancel"); err != nil { - t.Fatal(err) - } - if _, err := transitionTurn(ctx, s, tenant, session.ID, first.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnWaiting, Status: sessions.TurnInProgress}); !errors.Is(err, sessions.ErrTurnConflict) { - t.Fatalf("cancelling Turn resumed: %v", err) - } - transition(t, s, tenant, session.ID, first.TurnID, sessions.TurnWaiting, sessions.TurnCancelled) -} - -func TestTurnInputValidationHasNoSideEffects(t *testing.T) { - s, _ := testStore(t) - tenant, session := newTurnSession(t, s) - ctx := context.Background() - for _, raw := range []json.RawMessage{nil, json.RawMessage(`[]`), json.RawMessage(`null`), json.RawMessage(`{} {}`), json.RawMessage(`{"text":"` + string(make([]byte, 512*1024)) + `"}`)} { - if _, err := s.SubmitMessage(ctx, tenant, session.ID, "first", raw); !errors.Is(err, sessions.ErrInvalidInput) { - t.Fatalf("invalid input accepted: %v", err) - } - } - cancelled, stop := context.WithCancel(ctx) - stop() - if _, err := s.SubmitMessage(cancelled, tenant, session.ID, "first", messagePayload); !errors.Is(err, context.Canceled) { - t.Fatalf("cancelled context: %v", err) - } - first := submitMessage(t, s, tenant, session.ID, "first") - if first.Replayed { - t.Fatal("failed submission persisted a receipt") - } - for _, input := range []sessions.TurnTransition{ - {ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnCompleted}, - {ExpectedStatus: sessions.TurnInProgress, Status: sessions.TurnQueued}, - {ExpectedStatus: sessions.TurnCompleted, Status: sessions.TurnInProgress}, - {ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress, Outcome: json.RawMessage(`{"premature":true}`)}, - } { - if _, err := transitionTurn(ctx, s, tenant, session.ID, first.TurnID, input); !errors.Is(err, sessions.ErrInvalidInput) { - t.Fatalf("invalid transition accepted: %v", err) - } - } - got, err := sessionAdapter(s).GetTurn(ctx, tenant, session.ID, first.TurnID) - if err != nil || got.Status != sessions.TurnQueued || !got.StartedAt.IsZero() { - t.Fatalf("invalid transition changed Turn: %+v, %v", got, err) - } -} diff --git a/services/core/internal/store/worker_input_race_test.go b/services/core/internal/store/worker_input_race_test.go index c12653dc4..6ca03ded1 100644 --- a/services/core/internal/store/worker_input_race_test.go +++ b/services/core/internal/store/worker_input_race_test.go @@ -50,7 +50,7 @@ func TestWorkerInputReadSkipsConcurrentlyCancelledCandidate(t *testing.T) { mutated <- h.s.CommitLegacyDeletion(t.Context(), h.tenant, candidateSession) return } - _, err := h.s.SubmitInputs(t.Context(), h.tenant, candidateSession, "cancel", []sessions.Input{{Kind: "cancel", Payload: json.RawMessage(`{}`)}}) + _, err := store.SubmitInputs(t.Context(), h.s, h.tenant, candidateSession, "cancel", []sessions.Input{{Kind: "cancel", Payload: json.RawMessage(`{}`)}}) mutated <- err }} instrumented, err := pgxpool.NewWithConfig(t.Context(), cfg) diff --git a/services/core/internal/store/worker_preparation_failure_test.go b/services/core/internal/store/worker_preparation_failure_test.go index 93b293053..b7c82413e 100644 --- a/services/core/internal/store/worker_preparation_failure_test.go +++ b/services/core/internal/store/worker_preparation_failure_test.go @@ -16,7 +16,7 @@ func TestWorkerSettlesConfirmedPreparationFailureAndAcceptsNewInput(t *testing.T h := newDispatchHarnessForSession(t, []byte(`{"agent":{"model":"test-model"},"environment":{"type":"self_hosted","workspace_directory":"/workspace"}}`), false) enableWorkerEnvironment(t, h) frames := workerFrames(t, h) - pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "first", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) + pending, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "first", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) if err != nil { t.Fatal(err) } @@ -31,7 +31,7 @@ func TestWorkerSettlesConfirmedPreparationFailureAndAcceptsNewInput(t *testing.T h.write(prepare.ID, proto.TypePreparationStatus, proto.PreparationStatusPayload{State: "rejected", Operation: proto.TypeExecutionPrepare, ErrorCode: code}) } awaitDaemonRemoteCondition(t, t.Context(), 3*time.Second, "failed reservation settlement", func() bool { - current, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) + current, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) return err == nil && current.State == sessions.EnvironmentInputFailed }) session, err := store.SessionAdapter(h.s).GetSession(t.Context(), h.tenant, h.session.ID) @@ -59,7 +59,7 @@ func TestWorkerSettlesConfirmedPreparationFailureAndAcceptsNewInput(t *testing.T t.Fatal("failed input retried", frame.Type) case <-time.After(1200 * time.Millisecond): } - next, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "next", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"next"}`)}}) + next, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "next", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"next"}`)}}) if err != nil { t.Fatal("new input remained blocked", err) } @@ -71,7 +71,7 @@ func TestWorkerSettlesConfirmedPreparationFailureAndAcceptsNewInput(t *testing.T if startFrame.DecodePayload(&start) != nil || start.RunID == "" || inputTextForTest(t, start.Input) != "next" { t.Fatal("new input was not admitted") } - current, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, next.ID) + current, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, next.ID) if err != nil || current.State != sessions.EnvironmentInputAdmitted { t.Fatal("new input state", err) } @@ -99,7 +99,7 @@ func TestWorkerRetriesUncertainPreparationFailure(t *testing.T) { h := newDispatchHarnessForSession(t, []byte(`{"agent":{"model":"test-model"},"environment":{"type":"self_hosted","workspace_directory":"/workspace"}}`), false) enableWorkerEnvironment(t, h) frames := workerFrames(t, h) - pending, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "retry", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"retry"}`)}}) + pending, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "retry", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"retry"}`)}}) if err != nil { t.Fatal(err) } @@ -116,11 +116,11 @@ func TestWorkerRetriesUncertainPreparationFailure(t *testing.T) { nextWorkerFrame(t, frames, proto.TypeExecutionRelease) } nextWorkerFrame(t, frames, proto.TypeExecutionPrepare) - current, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) + current, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, pending.ID) if err != nil || current.State != sessions.EnvironmentInputPending || !current.Deadline.Equal(pending.Deadline) { t.Fatal("transient failure settled or extended input", err) } - if _, err := h.s.CancelEnvironmentInput(t.Context(), h.tenant, h.session.ID, pending.ID); err != nil { + if _, err := store.CancelEnvironmentInput(t.Context(), h.s, h.tenant, h.session.ID, pending.ID); err != nil { t.Fatal(err) } }) @@ -131,17 +131,17 @@ func TestWorkerPreparationRejectionPreservesCancellationAndNewerInput(t *testing h := newDispatchHarnessForSession(t, []byte(`{"agent":{"model":"test-model"},"environment":{"type":"self_hosted","workspace_directory":"/workspace"}}`), false) enableWorkerEnvironment(t, h) frames := workerFrames(t, h) - first, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "first", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) + first, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "first", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"first"}`)}}) if err != nil { t.Fatal(err) } _, stop := startEnvironmentExpiryWorker(t, h.db, h.d) defer stop() old := nextWorkerFrame(t, frames, proto.TypeExecutionPrepare) - if _, err := h.s.CancelEnvironmentInput(t.Context(), h.tenant, h.session.ID, first.ID); err != nil { + if _, err := store.CancelEnvironmentInput(t.Context(), h.s, h.tenant, h.session.ID, first.ID); err != nil { t.Fatal(err) } - next, err := h.s.ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "next", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"next"}`)}}) + next, err := store.SessionService(t, h.s).ReserveEnvironmentInput(t.Context(), h.tenant, h.session.ID, "next", []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"next"}`)}}) if err != nil { t.Fatal(err) } @@ -149,7 +149,7 @@ func TestWorkerPreparationRejectionPreservesCancellationAndNewerInput(t *testing h.write(old.ID, proto.TypePreparationStatus, rejection) prepare := nextWorkerFrame(t, frames, proto.TypeExecutionPrepare) for id, state := range map[string]string{first.ID: sessions.EnvironmentInputCancelled, next.ID: sessions.EnvironmentInputPending} { - current, err := h.s.GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, id) + current, err := store.SessionAdapter(h.s).GetEnvironmentInputReservation(t.Context(), h.tenant, h.session.ID, id) if err != nil || current.State != state { t.Fatal("late rejection changed cancellation or newer input", err) } diff --git a/services/core/tests/fixtures/main.go b/services/core/tests/fixtures/main.go index 936afddfd..12f1e8159 100644 --- a/services/core/tests/fixtures/main.go +++ b/services/core/tests/fixtures/main.go @@ -8,8 +8,9 @@ import ( "os" "strings" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/pgunit" + "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/persistence/postgres/sessionpg" "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/sessions" - "github.com/MiniMax-AI/OpenAgentCore/services/core/internal/store" "github.com/google/uuid" "github.com/jackc/pgx/v5/pgxpool" ) @@ -52,12 +53,16 @@ func seed() error { return err } defer pool.Close() - s := store.New(pool) + service, err := sessions.NewService(sessionpg.New(pgunit.NewPool(pool), nil)) + if err != nil { + return err + } for _, status := range []string{sessions.TurnCompleted, sessions.TurnFailed, sessions.TurnCancelled, sessions.TurnInProgress} { - receipt, err := s.SubmitMessage(ctx, f.Tenant, f.Session, uuid.NewString(), json.RawMessage(`{"text":"recovery fixture"}`)) + receipts, err := service.SubmitInputs(ctx, f.Tenant, f.Session, uuid.NewString(), []sessions.Input{{Kind: "message", Payload: json.RawMessage(`{"text":"recovery fixture"}`)}}) if err != nil { return err } + receipt := receipts[0] if err = transitionTurn(ctx, pool, f.Tenant, f.Session, receipt.TurnID, sessions.TurnTransition{ExpectedStatus: sessions.TurnQueued, Status: sessions.TurnInProgress}); err != nil { return err }