diff --git a/.changeset/turn-owner-write-fence.md b/.changeset/turn-owner-write-fence.md new file mode 100644 index 000000000..be685fa6d --- /dev/null +++ b/.changeset/turn-owner-write-fence.md @@ -0,0 +1,6 @@ +--- +"@truefoundry/trueforge-core": patch +"@truefoundry/trueforge": patch +--- + +Fence turn progress writes with `expected_active_executor_id` so only the owning replica can mutate a running turn; cancel/freeze stays unfenced. Wrong owner raises `TurnExecutorMismatchError`. diff --git a/packages/trueforge-core/src/agent-session/TurnHandle.ts b/packages/trueforge-core/src/agent-session/TurnHandle.ts index f0bb2be03..5544554a7 100644 --- a/packages/trueforge-core/src/agent-session/TurnHandle.ts +++ b/packages/trueforge-core/src/agent-session/TurnHandle.ts @@ -137,6 +137,19 @@ export class TurnHandle> { return this.turn.turn_id; } + /** Keys for progress store writes: session, turn, and owning executor fence. */ + private turnWriteScope(): { + session_id: string; + turn_id: string; + expected_active_executor_id: string; + } { + return { + session_id: this.turn.session_id, + turn_id: this.turn.turn_id, + expected_active_executor_id: this.turn.active_executor_id, + }; + } + get session_id(): string { return this.turn.session_id; } @@ -224,8 +237,7 @@ export class TurnHandle> { thread_id: null, }; await this.store.appendToEvents({ - session_id: this.turn.session_id, - turn_id: this.turn.turn_id, + ...this.turnWriteScope(), events: [turnCreated], }); yield turnCreated; @@ -336,8 +348,7 @@ export class TurnHandle> { if (!frozenByStore) { try { await this.store.updateTurnState({ - session_id: this.turn.session_id, - turn_id: this.turn.turn_id, + ...this.turnWriteScope(), state: terminalState, turn_done_event: turnDone, }); @@ -401,10 +412,7 @@ export class TurnHandle> { * the event should be emitted to the consumer (null = side-effect only / skip). */ private async persistExecutionEvent(event: AgentThreadExecutionEvent): Promise { - const scope = { - session_id: this.turn.session_id, - turn_id: this.turn.turn_id, - }; + const scope = this.turnWriteScope(); switch (event.type) { case HarnessEventType.MODEL_MESSAGE: diff --git a/packages/trueforge-core/src/agent-session/index.ts b/packages/trueforge-core/src/agent-session/index.ts index 3b9b5d66c..84fc5d17b 100644 --- a/packages/trueforge-core/src/agent-session/index.ts +++ b/packages/trueforge-core/src/agent-session/index.ts @@ -105,6 +105,7 @@ export { SessionStoreInvariantError, SessionStoreNotFoundError, TurnAlreadyExistsError, + TurnExecutorMismatchError, TurnNotFoundError, TurnNotRunningError, } from './store/SessionStoreErrors'; diff --git a/packages/trueforge-core/src/agent-session/store/ISessionStore.ts b/packages/trueforge-core/src/agent-session/store/ISessionStore.ts index b11a7ca1d..407ee8a39 100644 --- a/packages/trueforge-core/src/agent-session/store/ISessionStore.ts +++ b/packages/trueforge-core/src/agent-session/store/ISessionStore.ts @@ -159,6 +159,7 @@ export interface ListTurnsInput { export interface UpdateTurnStateInput { session_id: string; turn_id: string; + expected_active_executor_id: string; state: TerminalTurnState; /** Caller-built turn.done; written atomically with the state flip in the same tx. */ turn_done_event: PersistedTurnEvent; @@ -167,24 +168,28 @@ export interface UpdateTurnStateInput { export interface AppendToEventsInput { session_id: string; turn_id: string; + expected_active_executor_id: string; events: PersistedTurnEvent[]; } export interface AddThreadsInput { session_id: string; turn_id: string; + expected_active_executor_id: string; threads: AgentThreadSnapshot[]; } export interface RemoveThreadsInput { session_id: string; turn_id: string; + expected_active_executor_id: string; thread_ids: string[]; } export interface AppendToThreadContextInput { session_id: string; turn_id: string; + expected_active_executor_id: string; thread_id: string; context: ContextMessage[]; current_context_usage: CurrentContextUsage | null; @@ -194,24 +199,28 @@ export interface AppendToThreadContextInput { export interface OverwriteThreadContextInput { session_id: string; turn_id: string; + expected_active_executor_id: string; event: ThreadOverwriteContextEvent; } export interface PatchMCPServersInput { session_id: string; turn_id: string; + expected_active_executor_id: string; mcp_servers: MCPServerInitInfo[]; } export interface PatchSandboxInfoInput { session_id: string; turn_id: string; + expected_active_executor_id: string; sandbox_info: SandboxInfo; } export interface PatchThreadCapabilityStateInput { session_id: string; turn_id: string; + expected_active_executor_id: string; thread_id: string; key: string; state: JsonValue; diff --git a/packages/trueforge-core/src/agent-session/store/InMemorySessionStore.ts b/packages/trueforge-core/src/agent-session/store/InMemorySessionStore.ts index a2e114b61..4f00c010a 100644 --- a/packages/trueforge-core/src/agent-session/store/InMemorySessionStore.ts +++ b/packages/trueforge-core/src/agent-session/store/InMemorySessionStore.ts @@ -48,6 +48,7 @@ import { SessionNotFoundError, SessionStoreInvariantError, TurnAlreadyExistsError, + TurnExecutorMismatchError, TurnNotFoundError, TurnNotRunningError, } from './SessionStoreErrors'; @@ -465,6 +466,13 @@ export class InMemorySessionStore< if (turn.state.status !== 'running') { throw new TurnNotRunningError(input.turn_id, turn.state); } + if (turn.active_executor_id !== input.expected_active_executor_id) { + throw new TurnExecutorMismatchError({ + turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, + active_executor_id: turn.active_executor_id, + }); + } turn.state = deepCopy(input.state); turn.updated_at = new Date(); const list = this.events.get(tKey); @@ -475,7 +483,7 @@ export class InMemorySessionStore< } async appendToEvents(input: AppendToEventsInput): Promise { - this.requireRunningTurn(input.session_id, input.turn_id); + this.requireTurnProgressAllowed(input); const tKey = turnKey(input); const list = this.events.get(tKey); if (!list) { @@ -507,16 +515,27 @@ export class InMemorySessionStore< return turn; } - private requireRunningTurn(sessionId: string, turnId: string): TurnRecord { - const turn = this.requireTurn(sessionId, turnId); + private requireTurnProgressAllowed(input: { + session_id: string; + turn_id: string; + expected_active_executor_id: string; + }): TurnRecord { + const turn = this.requireTurn(input.session_id, input.turn_id); if (turn.state.status !== 'running') { - throw new TurnNotRunningError(turnId, turn.state); + throw new TurnNotRunningError(input.turn_id, turn.state); + } + if (turn.active_executor_id !== input.expected_active_executor_id) { + throw new TurnExecutorMismatchError({ + turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, + active_executor_id: turn.active_executor_id, + }); } return turn; } async addThreads(input: AddThreadsInput): Promise { - const turn = this.requireRunningTurn(input.session_id, input.turn_id); + const turn = this.requireTurnProgressAllowed(input); for (const thread of input.threads) { turn.snapshot.threads[thread.thread_id] = deepCopy(thread); } @@ -528,7 +547,7 @@ export class InMemorySessionStore< if (input.thread_ids.length === 0) { return; } - const turn = this.requireRunningTurn(input.session_id, input.turn_id); + const turn = this.requireTurnProgressAllowed(input); for (const id of input.thread_ids) { Reflect.deleteProperty(turn.snapshot.threads, id); } @@ -537,7 +556,7 @@ export class InMemorySessionStore< } async appendToThreadContext(input: AppendToThreadContextInput): Promise { - const turn = this.requireRunningTurn(input.session_id, input.turn_id); + const turn = this.requireTurnProgressAllowed(input); const thread = turn.snapshot.threads[input.thread_id]; if (!thread) { throw new SessionStoreInvariantError(`Thread not found: ${input.thread_id}`); @@ -554,7 +573,7 @@ export class InMemorySessionStore< } async overwriteThreadContext(input: OverwriteThreadContextInput): Promise { - const turn = this.requireRunningTurn(input.session_id, input.turn_id); + const turn = this.requireTurnProgressAllowed(input); const threadId = input.event.thread_id; const thread = turn.snapshot.threads[threadId]; if (!thread) { @@ -567,7 +586,7 @@ export class InMemorySessionStore< } async patchMCPServers(input: PatchMCPServersInput): Promise { - const turn = this.requireRunningTurn(input.session_id, input.turn_id); + const turn = this.requireTurnProgressAllowed(input); turn.snapshot.mcp_servers ??= {}; for (const server of input.mcp_servers) { turn.snapshot.mcp_servers[server.id] = deepCopy(server); @@ -577,14 +596,14 @@ export class InMemorySessionStore< } async patchSandboxInfo(input: PatchSandboxInfoInput): Promise { - const turn = this.requireRunningTurn(input.session_id, input.turn_id); + const turn = this.requireTurnProgressAllowed(input); turn.snapshot.sandbox_info = deepCopy(input.sandbox_info); turn.updated_at = new Date(); return; } async patchThreadCapabilityState(input: PatchThreadCapabilityStateInput): Promise { - const turn = this.requireRunningTurn(input.session_id, input.turn_id); + const turn = this.requireTurnProgressAllowed(input); const thread = turn.snapshot.threads[input.thread_id]; if (!thread) { throw new SessionStoreInvariantError(`Thread not found: ${input.thread_id}`); diff --git a/packages/trueforge-core/src/agent-session/store/SessionStoreErrors.ts b/packages/trueforge-core/src/agent-session/store/SessionStoreErrors.ts index 89b31ee90..06def4714 100644 --- a/packages/trueforge-core/src/agent-session/store/SessionStoreErrors.ts +++ b/packages/trueforge-core/src/agent-session/store/SessionStoreErrors.ts @@ -98,6 +98,23 @@ export class TurnNotRunningError extends SessionStoreConflictError { } } +/** Progress write rejected: turn is still running but the caller is not the active executor. */ +export class TurnExecutorMismatchError extends SessionStoreConflictError { + readonly turn_id: string; + readonly expected_active_executor_id: string; + readonly active_executor_id: string; + + constructor(input: { turn_id: string; expected_active_executor_id: string; active_executor_id: string }) { + super( + `Turn ${input.turn_id} is owned by executor ${input.active_executor_id}, not ${input.expected_active_executor_id}`, + ); + this.name = 'TurnExecutorMismatchError'; + this.turn_id = input.turn_id; + this.expected_active_executor_id = input.expected_active_executor_id; + this.active_executor_id = input.active_executor_id; + } +} + export class InvalidPageTokenError extends SessionStoreConflictError { readonly token: string; diff --git a/packages/trueforge-core/tests/agent-session/store/storeContractSuite.ts b/packages/trueforge-core/tests/agent-session/store/storeContractSuite.ts index 28a85e43f..aedacbf00 100644 --- a/packages/trueforge-core/tests/agent-session/store/storeContractSuite.ts +++ b/packages/trueforge-core/tests/agent-session/store/storeContractSuite.ts @@ -13,6 +13,7 @@ import { SessionStoreInvariantError, SessionStoreNotFoundError, TurnAlreadyExistsError, + TurnExecutorMismatchError, TurnNotFoundError, TurnNotRunningError, } from '../../../src/agent-session/store/SessionStoreErrors'; @@ -75,6 +76,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.updateTurnState({ session_id: sessionId, turn_id: turnId, + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state, turn_done_event: makeTurnDoneEvent(state), }); @@ -95,13 +97,14 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { function turnScopedWrites( store: ISessionStore, - keys: { session_id: string; turn_id: string }, + keys: { session_id: string; turn_id: string; expected_active_executor_id: string }, ): (() => Promise)[] { const doneState = makeDoneTurnState(); + const turnKeys = { session_id: keys.session_id, turn_id: keys.turn_id }; return [ () => store.freezeAndGetTurn({ - ...keys, + ...turnKeys, reason: CancellationReason.ClientCancelled, turn_done_event: makeTurnDoneEvent(makeCancelledTurnState(CancellationReason.ClientCancelled)), }), @@ -539,11 +542,13 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent('turn-1'), makeModelMessageEvent()], }); await store.patchThreadCapabilityState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, thread_id: MAIN_THREAD_ID, key: 'tfy.plan', state: { step: 'secret' }, @@ -551,11 +556,13 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.patchMCPServers({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, mcp_servers: [{ id: 'svc', name: 'svc', session_id: 'mcp-1', transport_type: 'streamable-http' }], }); await store.patchSandboxInfo({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, sandbox_info: { sandbox_id: 'sbx-1' }, }); @@ -661,7 +668,11 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.createTurn(makeCreateTurnInput({ sessionId, turnId: 'turn-1' })); await store.deleteSession({ tenant_id: tenant, session_id: sessionId }); - for (const write of turnScopedWrites(store, { session_id: sessionId, turn_id: 'turn-1' })) { + for (const write of turnScopedWrites(store, { + session_id: sessionId, + turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, + })) { await expect(write()).rejects.toBeInstanceOf(TurnNotFoundError); } }); @@ -672,7 +683,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.createTurn(makeCreateTurnInput({ sessionId, turnId: 'turn-1' })); await store.deleteSession({ tenant_id: tenant, session_id: sessionId }); - const keys = { session_id: sessionId, turn_id: 'turn-1' }; + const keys = { session_id: sessionId, turn_id: 'turn-1', expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID }; await expect( store.appendToEvents({ ...keys, @@ -712,7 +723,11 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { source: null, }); await store.createTurn(makeCreateTurnInput({ sessionId: nested, turnId: 'turn-1' })); - const nestedKeys = { session_id: nested, turn_id: 'turn-1' }; + const nestedKeys = { + session_id: nested, + turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, + }; await store.appendToEvents({ ...nestedKeys, events: [makeTurnCreatedEvent('turn-1')] }); await store.deleteSession({ tenant_id: tenant, session_id: sessionId }); @@ -728,7 +743,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { const store = createStore(); await seedSession(store); await store.createTurn(makeCreateTurnInput({ sessionId, turnId: 'turn-1' })); - const keys = { session_id: sessionId, turn_id: 'turn-1' }; + const keys = { session_id: sessionId, turn_id: 'turn-1', expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID }; const results = await Promise.allSettled([ store.deleteSession({ tenant_id: tenant, session_id: sessionId }), @@ -795,7 +810,11 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { const store = createStore(); await seedSession(store); - for (const write of turnScopedWrites(store, { session_id: sessionId, turn_id: missingTurnId })) { + for (const write of turnScopedWrites(store, { + session_id: sessionId, + turn_id: missingTurnId, + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, + })) { await expect(write()).rejects.toBeInstanceOf(TurnNotFoundError); } }); @@ -803,7 +822,11 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { it('rejects event mutations for a missing session or turn', async () => { const store = createStore(); await seedSession(store); - const keys = { session_id: sessionId, turn_id: missingTurnId }; + const keys = { + session_id: sessionId, + turn_id: missingTurnId, + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, + }; await expect( store.appendToEvents({ @@ -1497,6 +1520,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.patchThreadCapabilityState({ session_id: sessionId, turn_id: 't1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, thread_id: MAIN_THREAD_ID, key: 'tfy.plan', state: { step: 2 }, @@ -1601,6 +1625,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.overwriteThreadContext({ session_id: sessionId, turn_id: 't2', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, event: { type: EventType.AGENT_CONTEXT_OVERWRITE, id: newEventId(), @@ -1732,7 +1757,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { }); expect(data.some(e => e.type === EventType.TURN_DONE)).toBe(true); - const keys = { session_id: sessionId, turn_id: 'turn-1' }; + const keys = { session_id: sessionId, turn_id: 'turn-1', expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID }; const fencedWrites: (() => Promise)[] = [ () => store.appendToEvents({ @@ -1797,6 +1822,66 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { } }); + it('rejects progress writes with the wrong expected_active_executor_id while still running', async () => { + const store = createStore(); + await seedSession(store); + await store.createTurn(makeCreateTurnInput({ sessionId, turnId: 'turn-1' })); + + const wrongOwner = { + session_id: sessionId, + turn_id: 'turn-1', + expected_active_executor_id: 'other-executor', + }; + + await expect( + store.appendToEvents({ + ...wrongOwner, + events: [makeTurnCreatedEvent('turn-1')], + }), + ).rejects.toBeInstanceOf(TurnExecutorMismatchError); + + await expect( + store.updateTurnState({ + ...wrongOwner, + state: makeDoneTurnState(), + turn_done_event: makeTurnDoneEvent(makeDoneTurnState()), + }), + ).rejects.toBeInstanceOf(TurnExecutorMismatchError); + + await expect( + store.patchSandboxInfo({ + ...wrongOwner, + sandbox_info: { sandbox_id: 'sbx-wrong' }, + }), + ).rejects.toBeInstanceOf(TurnExecutorMismatchError); + + const stillRunning = mustGet(await store.getTurn({ session_id: sessionId, turn_id: 'turn-1' })); + expect(stillRunning.state.status).toBe('running'); + expect(stillRunning.active_executor_id).toBe(TEST_ACTIVE_EXECUTOR_ID); + }); + + it('accepts progress writes with the matching expected_active_executor_id', async () => { + const store = createStore(); + await seedSession(store); + await store.createTurn(makeCreateTurnInput({ sessionId, turnId: 'turn-1' })); + + await store.appendToEvents({ + session_id: sessionId, + turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, + events: [makeTurnCreatedEvent('turn-1')], + }); + await store.patchSandboxInfo({ + session_id: sessionId, + turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, + sandbox_info: { sandbox_id: 'sbx-ok' }, + }); + + const turn = mustGet(await store.getTurn({ session_id: sessionId, turn_id: 'turn-1' })); + expect(turn.snapshot.sandbox_info).toEqual({ sandbox_id: 'sbx-ok' }); + }); + it('cancels a running turn with the caller-supplied reason', async () => { const store = createStore(); await seedSession(store); @@ -1915,6 +2000,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { store.updateTurnState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state: doneState, turn_done_event: makeTurnDoneEvent(doneState), }), @@ -1950,6 +2036,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { store.updateTurnState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state: cancelledState, turn_done_event: makeTurnDoneEvent(cancelledState), }), @@ -1965,6 +2052,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.updateTurnState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state, turn_done_event: turnDone, }); @@ -1998,6 +2086,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.updateTurnState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state, turn_done_event: makeTurnDoneEvent(state), }); @@ -2024,6 +2113,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.updateTurnState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state, turn_done_event: makeTurnDoneEvent(state), }); @@ -2049,6 +2139,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.updateTurnState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state: doneState, turn_done_event: makeTurnDoneEvent(doneState), }); @@ -2063,6 +2154,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { store.updateTurnState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state: losingState, turn_done_event: makeTurnDoneEvent(losingState), }), @@ -2090,6 +2182,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.updateTurnState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state: turn1Done, turn_done_event: makeTurnDoneEvent(turn1Done), }); @@ -2106,6 +2199,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.updateTurnState({ session_id: sessionId, turn_id: 'turn-2', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state: turn2Done, turn_done_event: makeTurnDoneEvent(turn2Done), }); @@ -2132,6 +2226,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.updateTurnState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state: turn1Done, turn_done_event: makeTurnDoneEvent(turn1Done), }); @@ -2174,6 +2269,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { store.updateTurnState({ session_id: sessionId, turn_id: missingTurnId, + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state, turn_done_event: makeTurnDoneEvent(state), }), @@ -2194,6 +2290,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, // Deliberately reversed: durable ordering comes from event.id. events: [model, created], }); @@ -2215,6 +2312,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.addThreads({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, threads: [ { thread_id: 'child', @@ -2230,6 +2328,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToThreadContext({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, thread_id: 'child', context: [{ role: 'user', content: 'hello' }], current_context_usage: null, @@ -2241,6 +2340,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.overwriteThreadContext({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, event: { type: EventType.AGENT_CONTEXT_OVERWRITE, id: newEventId(), @@ -2258,6 +2358,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.removeThreads({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, thread_ids: ['child'], }); turn = await store.getTurn({ session_id: sessionId, turn_id: 'turn-1' }); @@ -2287,6 +2388,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToThreadContext({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, thread_id: MAIN_THREAD_ID, context: secondBatch, current_context_usage: null, @@ -2315,6 +2417,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.patchThreadCapabilityState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, thread_id: MAIN_THREAD_ID, key: 'tfy.plan', state: { v: 1 }, @@ -2322,6 +2425,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.patchThreadCapabilityState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, thread_id: MAIN_THREAD_ID, key: 'tfy.plan', state: { v: 2 }, @@ -2341,6 +2445,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { store.patchThreadCapabilityState({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, thread_id: 'missing-thread', key: 'tfy.plan', state: { v: 1 }, @@ -2377,6 +2482,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.patchThreadCapabilityState({ session_id: sessionId, turn_id: 't2', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, thread_id: MAIN_THREAD_ID, key: 'plan', state: { step: 2 }, @@ -2463,11 +2569,13 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.patchMCPServers({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, mcp_servers: [{ id: 'svc', name: 'svc', session_id: 'mcp-1', transport_type: 'streamable-http' }], }); await store.patchSandboxInfo({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, sandbox_info: { sandbox_id: 'sbx-1' }, }); const turn = await store.getTurn({ session_id: sessionId, turn_id: 'turn-1' }); @@ -2482,11 +2590,13 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.patchMCPServers({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, mcp_servers: [{ id: 'svc', name: 'svc', session_id: 'mcp-1', transport_type: 'streamable-http' }], }); await store.patchMCPServers({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, mcp_servers: [{ id: 'svc', name: 'svc', transport_type: 'sse' }], }); const turn = await store.getTurn({ session_id: sessionId, turn_id: 'turn-1' }); @@ -2531,6 +2641,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 'turn-1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent('turn-1'), makeModelMessageEvent()], }); const { data } = await store.listTurnEvents({ @@ -2579,6 +2690,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 't1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent('t1')], }); await finishTurn(store, 't1'); @@ -2586,6 +2698,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 't2', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent('t2')], }); // t2 still running — must appear in the feed. @@ -2619,6 +2732,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 't1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [passthrough], }); @@ -2638,6 +2752,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 't1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent('t1')], }); // Two forks off t1; the sibling is created BEFORE the anchor, so a @@ -2653,6 +2768,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: turnId, + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent(turnId)], }); } @@ -2684,6 +2800,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 't1', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent('t1')], }); await finishTurn(store, 't1'); @@ -2691,6 +2808,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 't2', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent('t2')], }); @@ -2712,6 +2830,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: 't2-new-active', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent('t2-new-active')], }); @@ -2753,6 +2872,7 @@ export function runStoreContractSuite(createStore: () => ISessionStore) { await store.appendToEvents({ session_id: sessionId, turn_id: turnId, + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, events: [makeTurnCreatedEvent(turnId)], }); } diff --git a/packages/trueforge/src/db/postgres/session-store/queries/capabilities.ts b/packages/trueforge/src/db/postgres/session-store/queries/capabilities.ts index 362d34ffc..cd526670e 100644 --- a/packages/trueforge/src/db/postgres/session-store/queries/capabilities.ts +++ b/packages/trueforge/src/db/postgres/session-store/queries/capabilities.ts @@ -3,7 +3,7 @@ import { sql, type Kysely } from 'kysely'; import { json } from '../../sqlExpressions'; import type { Database } from '../../types'; import { values } from '../sqlExpressions'; -import { classifyTurnFenceWriteFailure, turnRunningFence } from './turns'; +import { classifyTurnProgressFenceFailure, turnProgressFence } from './turns'; /** * patchThreadCapabilityState — single-statement fenced upsert on the PER-TURN PK; @@ -16,10 +16,11 @@ export async function patchThreadCapabilityState( const keys = { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }; const rows = await db - .with('turn_fence', qb => turnRunningFence(qb, keys)) + .with('turn_fence', qb => turnProgressFence(qb, keys)) .insertInto('thread_capability_state') .columns(['session_id', 'turn_id', 'thread_id', 'key', 'state', 'updated_at']) .expression(eb => @@ -45,6 +46,6 @@ export async function patchThreadCapabilityState( .execute(); if (rows.length === 0) { - await classifyTurnFenceWriteFailure(db, keys); + await classifyTurnProgressFenceFailure(db, keys); } } diff --git a/packages/trueforge/src/db/postgres/session-store/queries/events.ts b/packages/trueforge/src/db/postgres/session-store/queries/events.ts index 29e9da909..0787413bc 100644 --- a/packages/trueforge/src/db/postgres/session-store/queries/events.ts +++ b/packages/trueforge/src/db/postgres/session-store/queries/events.ts @@ -23,7 +23,7 @@ import { sql } from 'kysely'; import { json } from '../../sqlExpressions'; import type { Database } from '../../types'; import { unnestWithOrdinality, values } from '../sqlExpressions'; -import { classifyTurnFenceWriteFailure, turnRunningFence } from './turns'; +import { classifyTurnProgressFenceFailure, turnProgressFence } from './turns'; export async function appendToEvents(db: Kysely, input: AppendToEventsInput): Promise { if (input.events.length === 0) { @@ -33,6 +33,7 @@ export async function appendToEvents(db: Kysely, input: AppendToEvents const keys = { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }; const eventRows = input.events.map(event => ({ @@ -43,7 +44,7 @@ export async function appendToEvents(db: Kysely, input: AppendToEvents })); const inserted = await db - .with('turn_fence', qb => turnRunningFence(qb, keys)) + .with('turn_fence', qb => turnProgressFence(qb, keys)) .insertInto('session_event') .columns(['session_id', 'turn_id', 'event_id', 'event', 'created_at']) .expression(eb => @@ -62,7 +63,7 @@ export async function appendToEvents(db: Kysely, input: AppendToEvents .execute(); if (inserted.length === 0) { - await classifyTurnFenceWriteFailure(db, keys); + await classifyTurnProgressFenceFailure(db, keys); } } diff --git a/packages/trueforge/src/db/postgres/session-store/queries/threads.ts b/packages/trueforge/src/db/postgres/session-store/queries/threads.ts index 1b247a400..e618d3f32 100644 --- a/packages/trueforge/src/db/postgres/session-store/queries/threads.ts +++ b/packages/trueforge/src/db/postgres/session-store/queries/threads.ts @@ -18,10 +18,10 @@ import { json, jsonbSet } from '../../sqlExpressions'; import type { Database, TurnThreadCheckpoint } from '../../types'; import { values } from '../sqlExpressions'; import { - assertTurnRunning, - classifyTurnFenceWriteFailure, - classifyTurnThreadWriteFailure, - turnRunningFence, + assertTurnProgressAllowed, + classifyTurnProgressFenceFailure, + classifyTurnThreadProgressFailure, + turnProgressFence, type TurnKeys, } from './turns'; @@ -42,9 +42,10 @@ interface CapabilityStateInsertRow { */ export async function addThreads(db: Kysely, input: AddThreadsInput): Promise { await db.transaction().execute(async trx => { - await assertTurnRunning(trx, { + await assertTurnProgressAllowed(trx, { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }); const now = new Date(); @@ -156,11 +157,12 @@ export async function removeThreads(db: Kysely, input: RemoveThreadsIn const keys: TurnKeys = { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }; const onFence = sql`EXISTS (SELECT 1 FROM turn_fence)`; const fence = await db - .with('turn_fence', qb => turnRunningFence(qb, keys)) + .with('turn_fence', qb => turnProgressFence(qb, keys)) .with('del_cap', qb => qb .deleteFrom('thread_capability_state') @@ -182,7 +184,7 @@ export async function removeThreads(db: Kysely, input: RemoveThreadsIn .executeTakeFirst(); if (fence === undefined) { - await classifyTurnFenceWriteFailure(db, keys); + await classifyTurnProgressFenceFailure(db, keys); } } @@ -219,7 +221,7 @@ async function fencedTurnThreadContextUpdate( if (context.length === 0) { // No log INSERT — still fence + patch usage/completion / clear-or-keep array. const emptyResult = await db - .with('turn_fence', qb => turnRunningFence(qb, keys)) + .with('turn_fence', qb => turnProgressFence(qb, keys)) .updateTable('turn_thread') .set({ context_ids: replace_array ? sql`'{}'::bigint[]` : sql`context_ids`, @@ -234,7 +236,7 @@ async function fencedTurnThreadContextUpdate( .executeTakeFirst(); if (Number(emptyResult.numUpdatedRows) === 0) { - await classifyTurnThreadWriteFailure(db, keys, thread_id); + await classifyTurnThreadProgressFailure(db, keys, thread_id); } return; } @@ -251,7 +253,7 @@ async function fencedTurnThreadContextUpdate( >`context_ids || coalesce((SELECT array_agg(append_id ORDER BY append_id) FROM new_rows), '{}'::bigint[])`; const result = await db - .with('turn_fence', qb => turnRunningFence(qb, keys)) + .with('turn_fence', qb => turnProgressFence(qb, keys)) .with('new_rows', qb => qb .insertInto('thread_context_log') @@ -285,7 +287,7 @@ async function fencedTurnThreadContextUpdate( .executeTakeFirst(); if (Number(result.numUpdatedRows) === 0) { - await classifyTurnThreadWriteFailure(db, keys, thread_id); + await classifyTurnThreadProgressFailure(db, keys, thread_id); } } @@ -299,6 +301,7 @@ export async function appendToThreadContext(db: Kysely, input: AppendT keys: { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }, thread_id: input.thread_id, context: input.context, @@ -318,6 +321,7 @@ export async function overwriteThreadContext(db: Kysely, input: Overwr keys: { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }, thread_id: input.event.thread_id, context: input.event.context, @@ -330,7 +334,7 @@ export async function overwriteThreadContext(db: Kysely, input: Overwr /** * patchMCPServers — one-shot conditional UPDATE on the fence row itself - * (`state->>'status' = 'running'`). No separate FOR SHARE fence CTE. + * (running + matching active_executor_id). No separate FOR SHARE fence CTE. * Subscript LHS + expression RHS for shallow merge by server id. */ export async function patchMCPServers(db: Kysely, input: PatchMCPServersInput): Promise { @@ -342,6 +346,7 @@ export async function patchMCPServers(db: Kysely, input: PatchMCPServe const keys: TurnKeys = { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }; const result = await db @@ -356,21 +361,23 @@ export async function patchMCPServers(db: Kysely, input: PatchMCPServe .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .where(sql`state->>'status' = 'running'`) + .where('active_executor_id', '=', keys.expected_active_executor_id) .executeTakeFirst(); if (Number(result.numUpdatedRows) === 0) { - await classifyTurnFenceWriteFailure(db, keys); + await classifyTurnProgressFenceFailure(db, keys); } } /** * patchSandboxInfo — one-shot conditional UPDATE on the fence row itself - * (`state->>'status' = 'running'`). LWW replace via subscript assignment. + * (running + matching active_executor_id). LWW replace via subscript assignment. */ export async function patchSandboxInfo(db: Kysely, input: PatchSandboxInfoInput): Promise { const keys: TurnKeys = { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }; const result = await db @@ -380,10 +387,11 @@ export async function patchSandboxInfo(db: Kysely, input: PatchSandbox .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .where(sql`state->>'status' = 'running'`) + .where('active_executor_id', '=', keys.expected_active_executor_id) .executeTakeFirst(); if (Number(result.numUpdatedRows) === 0) { - await classifyTurnFenceWriteFailure(db, keys); + await classifyTurnProgressFenceFailure(db, keys); } } diff --git a/packages/trueforge/src/db/postgres/session-store/queries/turns.ts b/packages/trueforge/src/db/postgres/session-store/queries/turns.ts index 06bcda570..679080c24 100644 --- a/packages/trueforge/src/db/postgres/session-store/queries/turns.ts +++ b/packages/trueforge/src/db/postgres/session-store/queries/turns.ts @@ -18,6 +18,7 @@ import { SessionStoreInvariantError, SessionStoreNotFoundError, TurnAlreadyExistsError, + TurnExecutorMismatchError, TurnNotFoundError, TurnNotRunningError, } from '@truefoundry/trueforge-core/agent-session/store/SessionStoreErrors'; @@ -59,6 +60,8 @@ export interface NewThreadRegistration { export interface TurnKeys { session_id: string; turn_id: string; + /** Must equal the turn row's active_executor_id or the write is rejected. */ + expected_active_executor_id: string; } export interface NewContextAppend { @@ -162,27 +165,27 @@ function terminalTurnState(state: TurnState, turn_id: string): TerminalTurnState } /** - * Locking CTE body for single-statement turn-scoped writes: fence + write in one - * network call. Under READ COMMITTED, FOR SHARE re-checks the predicate after a - * lock wait, so a committed freeze empties the fence and the write inserts 0 rows; - * error classification happens on that rare 0-row path. - * Multi-statement turn-scoped writes use {@link assertTurnRunning} instead. + * Locking CTE for single-statement progress writes: turn must be running and owned + * by expected_active_executor_id. Under READ COMMITTED, FOR SHARE re-checks after a + * lock wait, so freeze or owner change empties the fence (0-row write → classify). + * Multi-statement progress writes use {@link assertTurnProgressAllowed} instead. */ -export function turnRunningFence(db: TurnFenceDb, keys: TurnKeys) { +export function turnProgressFence(db: TurnFenceDb, keys: TurnKeys) { return db .selectFrom('turn') .select(sql`1`.as('one')) .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .where(sql`state->>'status'`, '=', 'running') + .where('active_executor_id', '=', keys.expected_active_executor_id) .forShare(); } -/** Classify a 0-row fenced write: missing turn vs frozen/non-running turn. */ -export async function classifyTurnFenceWriteFailure(db: Kysely, keys: TurnKeys): Promise { +/** Classify a 0-row progress-fenced write: missing, wrong owner, or not running. */ +export async function classifyTurnProgressFenceFailure(db: DbOrTrx, keys: TurnKeys): Promise { const row = await db .selectFrom('turn') - .select('state') + .select(['state', 'active_executor_id']) .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .executeTakeFirst(); @@ -190,20 +193,27 @@ export async function classifyTurnFenceWriteFailure(db: Kysely, keys: if (!row) { throw new TurnNotFoundError(keys.turn_id); } + if (row.state.status === 'running' && row.active_executor_id !== keys.expected_active_executor_id) { + throw new TurnExecutorMismatchError({ + turn_id: keys.turn_id, + expected_active_executor_id: keys.expected_active_executor_id, + active_executor_id: row.active_executor_id, + }); + } throw new TurnNotRunningError(keys.turn_id, terminalTurnState(row.state, keys.turn_id)); } /** - * Classify a 0-row fenced turn_thread UPDATE: turn missing/terminal vs thread row missing. + * Classify a 0-row progress-fenced turn_thread UPDATE: missing/terminal/wrong owner vs thread missing. */ -export async function classifyTurnThreadWriteFailure( - db: Kysely, +export async function classifyTurnThreadProgressFailure( + db: DbOrTrx, keys: TurnKeys, thread_id: string, ): Promise { const row = await db .selectFrom('turn') - .select('state') + .select(['state', 'active_executor_id']) .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .executeTakeFirst(); @@ -214,14 +224,21 @@ export async function classifyTurnThreadWriteFailure( if (row.state.status !== 'running') { throw new TurnNotRunningError(keys.turn_id, terminalTurnState(row.state, keys.turn_id)); } + if (row.active_executor_id !== keys.expected_active_executor_id) { + throw new TurnExecutorMismatchError({ + turn_id: keys.turn_id, + expected_active_executor_id: keys.expected_active_executor_id, + active_executor_id: row.active_executor_id, + }); + } throw new SessionStoreInvariantError(`thread ${thread_id} not found in turn ${keys.turn_id}`); } -export async function assertTurnRunning(db: DbOrTrx, keys: TurnKeys): Promise { +export async function assertTurnProgressAllowed(db: DbOrTrx, keys: TurnKeys): Promise { // SELECT ... FOR SHARE serializes against freezeAndGetTurn's state UPDATE. const row = await db .selectFrom('turn') - .select('state') + .select(['state', 'active_executor_id']) .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .forShare() @@ -233,6 +250,13 @@ export async function assertTurnRunning(db: DbOrTrx, keys: TurnKeys): Promise, input: ListTurnsInput): Pr } /** - * updateTurnState — conditional on state->>'status'='running'. - * 0 rows → SELECT by PK → missing NotFound, present Conflict (first terminal write wins). + * updateTurnState — conditional on running + matching active_executor_id. + * 0 rows → classify: missing / wrong owner / already terminal (first terminal write wins). */ export async function updateTurnState(db: Kysely, input: UpdateTurnStateInput): Promise { await db.transaction().execute(async trx => { @@ -759,22 +783,17 @@ export async function updateTurnState(db: Kysely, input: UpdateTurnSta .where('session_id', '=', input.session_id) .where('turn_id', '=', input.turn_id) .where(sql`state->>'status' = 'running'`) + .where('active_executor_id', '=', input.expected_active_executor_id) .returning(['created_at']) .executeTakeFirst(); - // No RETURNING row: UPDATE matched 0 running turns. + // No RETURNING row: UPDATE matched 0 owned running turns. if (result === undefined) { - const existing = await trx - .selectFrom('turn') - .select('state') - .where('session_id', '=', input.session_id) - .where('turn_id', '=', input.turn_id) - .executeTakeFirst(); - - if (!existing) { - throw new TurnNotFoundError(input.turn_id); - } - throw new TurnNotRunningError(input.turn_id, terminalTurnState(existing.state, input.turn_id)); + return await classifyTurnProgressFenceFailure(trx, { + session_id: input.session_id, + turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, + }); } await addSessionCostAndDuration(trx, { diff --git a/packages/trueforge/src/db/sqlite/session-store/queries/capabilities.ts b/packages/trueforge/src/db/sqlite/session-store/queries/capabilities.ts index 791726bf9..56992e407 100644 --- a/packages/trueforge/src/db/sqlite/session-store/queries/capabilities.ts +++ b/packages/trueforge/src/db/sqlite/session-store/queries/capabilities.ts @@ -2,7 +2,7 @@ import type { PatchThreadCapabilityStateInput } from '@truefoundry/trueforge-cor import { sql, type Kysely } from 'kysely'; import { jsonbBind, nowIso } from '../../sqlExpressions'; import type { Database } from '../../types'; -import { classifyTurnFenceWriteFailure } from './turns'; +import { classifyTurnProgressFenceFailure } from './turns'; /** * patchThreadCapabilityState — single-statement fenced upsert on the PER-TURN PK. @@ -14,17 +14,18 @@ export async function patchThreadCapabilityState( input: PatchThreadCapabilityStateInput, ): Promise { await db.transaction().execute(async trx => { - // Fence inside IMMEDIATE transaction: verify turn is still running. + // Fence inside IMMEDIATE transaction: verify turn is still running and owned. const fenceRow = await trx .selectFrom('turn') .select(sql`1`.as('one')) .where('session_id', '=', input.session_id) .where('turn_id', '=', input.turn_id) .where(sql`state->>'status' = 'running'`) + .where('active_executor_id', '=', input.expected_active_executor_id) .executeTakeFirst(); if (!fenceRow) { - await classifyTurnFenceWriteFailure(trx, input); + await classifyTurnProgressFenceFailure(trx, input); } const now = nowIso(); diff --git a/packages/trueforge/src/db/sqlite/session-store/queries/events.ts b/packages/trueforge/src/db/sqlite/session-store/queries/events.ts index 4255746a0..df262e535 100644 --- a/packages/trueforge/src/db/sqlite/session-store/queries/events.ts +++ b/packages/trueforge/src/db/sqlite/session-store/queries/events.ts @@ -21,7 +21,7 @@ import { import { sql, type Kysely } from 'kysely'; import { jsonbBind, jsonText } from '../../sqlExpressions'; import type { Database } from '../../types'; -import { classifyTurnFenceWriteFailure, type TurnKeys } from './turns'; +import { classifyTurnProgressFenceFailure, type TurnKeys } from './turns'; export async function appendToEvents(db: Kysely, input: AppendToEventsInput): Promise { if (input.events.length === 0) { @@ -31,6 +31,7 @@ export async function appendToEvents(db: Kysely, input: AppendToEvents const keys: TurnKeys = { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }; // Fence check first inside BEGIN IMMEDIATE; then batched insert. @@ -41,10 +42,11 @@ export async function appendToEvents(db: Kysely, input: AppendToEvents .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .where(sql`state->>'status' = 'running'`) + .where('active_executor_id', '=', keys.expected_active_executor_id) .executeTakeFirst(); if (!fenceRow) { - await classifyTurnFenceWriteFailure(trx, keys); + await classifyTurnProgressFenceFailure(trx, keys); } const eventRows = input.events.map(event => ({ diff --git a/packages/trueforge/src/db/sqlite/session-store/queries/threads.ts b/packages/trueforge/src/db/sqlite/session-store/queries/threads.ts index 811af4b1d..df712ce78 100644 --- a/packages/trueforge/src/db/sqlite/session-store/queries/threads.ts +++ b/packages/trueforge/src/db/sqlite/session-store/queries/threads.ts @@ -16,9 +16,9 @@ import { jsonbBind, jsonbSet, nowIso } from '../../sqlExpressions'; import type { Database, TurnThreadCheckpoint } from '../../types'; import { sortedByAppendId } from '../sqlExpressions'; import { - assertTurnRunning, - classifyTurnFenceWriteFailure, - classifyTurnThreadWriteFailure, + assertTurnProgressAllowed, + classifyTurnProgressFenceFailure, + classifyTurnThreadProgressFailure, type TurnKeys, } from './turns'; @@ -30,9 +30,10 @@ type DbOrTrx = Kysely | Transaction; */ export async function addThreads(db: Kysely, input: AddThreadsInput): Promise { await db.transaction().execute(async trx => { - await assertTurnRunning(trx, { + await assertTurnProgressAllowed(trx, { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }); const now = nowIso(); @@ -178,9 +179,10 @@ export async function removeThreads(db: Kysely, input: RemoveThreadsIn } await db.transaction().execute(async trx => { - await assertTurnRunning(trx, { + await assertTurnProgressAllowed(trx, { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }); await trx @@ -247,7 +249,7 @@ async function fencedTurnThreadContextUpdate( const { keys, thread_id, context, replace_array } = args; await db.transaction().execute(async trx => { - await assertTurnRunning(trx, keys); + await assertTurnProgressAllowed(trx, keys); const now = nowIso(); @@ -315,7 +317,7 @@ async function fencedTurnThreadContextUpdate( .executeTakeFirst(); if (Number(updateResult.numUpdatedRows) === 0) { - await classifyTurnThreadWriteFailure(trx, keys, thread_id); + await classifyTurnThreadProgressFailure(trx, keys, thread_id); } }); } @@ -329,6 +331,7 @@ export async function appendToThreadContext(db: Kysely, input: AppendT keys: { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }, thread_id: input.thread_id, context: input.context, @@ -348,6 +351,7 @@ export async function overwriteThreadContext(db: Kysely, input: Overwr keys: { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }, thread_id: input.event.thread_id, context: input.event.context, @@ -359,7 +363,7 @@ export async function overwriteThreadContext(db: Kysely, input: Overwr } /** - * patchMCPServers — conditional UPDATE fenced on state->>'status'='running'. + * patchMCPServers — conditional UPDATE fenced on running + matching active_executor_id. * Shallow merge by server id (Postgres `||`): patched ids replace wholesale. */ export async function patchMCPServers(db: Kysely, input: PatchMCPServersInput): Promise { @@ -371,6 +375,7 @@ export async function patchMCPServers(db: Kysely, input: PatchMCPServe const keys: TurnKeys = { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }; // jsonb_patch is RFC 7396 (deep); rebuild via json_each so each id's value is replaced. @@ -402,10 +407,11 @@ export async function patchMCPServers(db: Kysely, input: PatchMCPServe .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .where(sql`state->>'status' = 'running'`) + .where('active_executor_id', '=', keys.expected_active_executor_id) .executeTakeFirst(); if (Number(result.numUpdatedRows) === 0) { - await classifyTurnFenceWriteFailure(db, keys); + await classifyTurnProgressFenceFailure(db, keys); } } @@ -416,6 +422,7 @@ export async function patchSandboxInfo(db: Kysely, input: PatchSandbox const keys: TurnKeys = { session_id: input.session_id, turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, }; const result = await db @@ -427,9 +434,10 @@ export async function patchSandboxInfo(db: Kysely, input: PatchSandbox .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .where(sql`state->>'status' = 'running'`) + .where('active_executor_id', '=', keys.expected_active_executor_id) .executeTakeFirst(); if (Number(result.numUpdatedRows) === 0) { - await classifyTurnFenceWriteFailure(db, keys); + await classifyTurnProgressFenceFailure(db, keys); } } diff --git a/packages/trueforge/src/db/sqlite/session-store/queries/turns.ts b/packages/trueforge/src/db/sqlite/session-store/queries/turns.ts index db230fd20..44c76d4bc 100644 --- a/packages/trueforge/src/db/sqlite/session-store/queries/turns.ts +++ b/packages/trueforge/src/db/sqlite/session-store/queries/turns.ts @@ -18,6 +18,7 @@ import { SessionStoreInvariantError, SessionStoreNotFoundError, TurnAlreadyExistsError, + TurnExecutorMismatchError, TurnNotFoundError, TurnNotRunningError, } from '@truefoundry/trueforge-core/agent-session/store/SessionStoreErrors'; @@ -59,6 +60,7 @@ export interface NewThreadRegistration { export interface TurnKeys { session_id: string; turn_id: string; + expected_active_executor_id: string; } export interface NewContextAppend { @@ -151,46 +153,77 @@ function terminalTurnState(state: TurnState, turn_id: string): TerminalTurnState } } -async function loadTurnState(db: DbOrTrx, keys: TurnKeys): Promise { +async function loadTurnFenceRow( + db: DbOrTrx, + keys: TurnKeys, +): Promise<{ state: TurnState; active_executor_id: string } | undefined> { const row = await db .selectFrom('turn') - .select([jsonText(sql.ref('state')).as('state')]) + .select([jsonText(sql.ref('state')).as('state'), 'active_executor_id']) .where('session_id', '=', keys.session_id) .where('turn_id', '=', keys.turn_id) .executeTakeFirst(); - return row?.state; + if (!row) { + return undefined; + } + return { state: row.state, active_executor_id: row.active_executor_id }; } -/** Classify a 0-row fenced write: missing turn vs frozen/non-running turn. */ -export async function classifyTurnFenceWriteFailure(db: DbOrTrx, keys: TurnKeys): Promise { - const state = await loadTurnState(db, keys); - if (!state) { +/** Classify a 0-row progress-fenced write: missing, wrong owner, or not running. */ +export async function classifyTurnProgressFenceFailure(db: DbOrTrx, keys: TurnKeys): Promise { + const row = await loadTurnFenceRow(db, keys); + if (!row) { throw new TurnNotFoundError(keys.turn_id); } - throw new TurnNotRunningError(keys.turn_id, terminalTurnState(state, keys.turn_id)); + if (row.state.status === 'running' && row.active_executor_id !== keys.expected_active_executor_id) { + throw new TurnExecutorMismatchError({ + turn_id: keys.turn_id, + expected_active_executor_id: keys.expected_active_executor_id, + active_executor_id: row.active_executor_id, + }); + } + throw new TurnNotRunningError(keys.turn_id, terminalTurnState(row.state, keys.turn_id)); } /** - * Classify a 0-row fenced turn_thread UPDATE: turn missing/terminal vs thread row missing. + * Classify a 0-row progress-fenced turn_thread UPDATE: missing/terminal/wrong owner vs thread missing. */ -export async function classifyTurnThreadWriteFailure(db: DbOrTrx, keys: TurnKeys, thread_id: string): Promise { - const state = await loadTurnState(db, keys); - if (!state) { +export async function classifyTurnThreadProgressFailure( + db: DbOrTrx, + keys: TurnKeys, + thread_id: string, +): Promise { + const row = await loadTurnFenceRow(db, keys); + if (!row) { throw new TurnNotFoundError(keys.turn_id); } - if (state.status !== 'running') { - throw new TurnNotRunningError(keys.turn_id, terminalTurnState(state, keys.turn_id)); + if (row.state.status !== 'running') { + throw new TurnNotRunningError(keys.turn_id, terminalTurnState(row.state, keys.turn_id)); + } + if (row.active_executor_id !== keys.expected_active_executor_id) { + throw new TurnExecutorMismatchError({ + turn_id: keys.turn_id, + expected_active_executor_id: keys.expected_active_executor_id, + active_executor_id: row.active_executor_id, + }); } throw new SessionStoreInvariantError(`thread ${thread_id} not found in turn ${keys.turn_id}`); } -export async function assertTurnRunning(db: DbOrTrx, keys: TurnKeys): Promise { - const state = await loadTurnState(db, keys); - if (!state) { +export async function assertTurnProgressAllowed(db: DbOrTrx, keys: TurnKeys): Promise { + const row = await loadTurnFenceRow(db, keys); + if (!row) { throw new TurnNotFoundError(keys.turn_id); } - if (state.status !== 'running') { - throw new TurnNotRunningError(keys.turn_id, terminalTurnState(state, keys.turn_id)); + if (row.state.status !== 'running') { + throw new TurnNotRunningError(keys.turn_id, terminalTurnState(row.state, keys.turn_id)); + } + if (row.active_executor_id !== keys.expected_active_executor_id) { + throw new TurnExecutorMismatchError({ + turn_id: keys.turn_id, + expected_active_executor_id: keys.expected_active_executor_id, + active_executor_id: row.active_executor_id, + }); } } @@ -797,8 +830,8 @@ export async function listTurns(db: Kysely, input: ListTurnsInput): Pr } /** - * updateTurnState — conditional on state->>'status'='running'. - * 0 rows → SELECT by PK → missing NotFound, present Conflict (first terminal write wins). + * updateTurnState — conditional on running + matching active_executor_id. + * 0 rows → classify: missing / wrong owner / already terminal (first terminal write wins). */ export async function updateTurnState(db: Kysely, input: UpdateTurnStateInput): Promise { await db.transaction().execute(async trx => { @@ -811,22 +844,17 @@ export async function updateTurnState(db: Kysely, input: UpdateTurnSta .where('session_id', '=', input.session_id) .where('turn_id', '=', input.turn_id) .where(sql`state->>'status' = 'running'`) + .where('active_executor_id', '=', input.expected_active_executor_id) .returning(['created_at']) .executeTakeFirst(); - // No RETURNING row: UPDATE matched 0 running turns. + // No RETURNING row: UPDATE matched 0 owned running turns. if (result === undefined) { - const existing = await trx - .selectFrom('turn') - .select([jsonText(sql.ref('state')).as('state')]) - .where('session_id', '=', input.session_id) - .where('turn_id', '=', input.turn_id) - .executeTakeFirst(); - - if (!existing) { - throw new TurnNotFoundError(input.turn_id); - } - throw new TurnNotRunningError(input.turn_id, terminalTurnState(existing.state, input.turn_id)); + return await classifyTurnProgressFenceFailure(trx, { + session_id: input.session_id, + turn_id: input.turn_id, + expected_active_executor_id: input.expected_active_executor_id, + }); } await addSessionCostAndDuration(trx, { diff --git a/packages/trueforge/tests/db/session-metrics/metricsContractSuite.ts b/packages/trueforge/tests/db/session-metrics/metricsContractSuite.ts index a3d1fdab1..c3153dca2 100644 --- a/packages/trueforge/tests/db/session-metrics/metricsContractSuite.ts +++ b/packages/trueforge/tests/db/session-metrics/metricsContractSuite.ts @@ -3,6 +3,7 @@ import { makeCreateTurnInput, makeDoneTurnState, makeTurnDoneEvent, + TEST_ACTIVE_EXECUTOR_ID, } from '../../../../trueforge-core/tests/agent-session/testHelpers'; import type { ISessionMetricsStore } from '../../../src/db/sessionMetricsStore'; @@ -56,6 +57,7 @@ export function runSessionMetricsStoreContractSuite( await sessionStore.updateTurnState({ session_id: 'metrics-session', turn_id: 'metrics-turn', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state, turn_done_event: makeTurnDoneEvent(state), }); @@ -147,6 +149,7 @@ export function runSessionMetricsStoreContractSuite( await sessionStore.updateTurnState({ session_id: definition.id, turn_id: turnId, + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state, turn_done_event: makeTurnDoneEvent(state), }); @@ -207,6 +210,7 @@ export function runSessionMetricsStoreContractSuite( await sessionStore.updateTurnState({ session_id: 'completed-session', turn_id: 'completed-turn', + expected_active_executor_id: TEST_ACTIVE_EXECUTOR_ID, state, turn_done_event: makeTurnDoneEvent(state), });