diff --git a/core/src/sessions/base_session_service.ts b/core/src/sessions/base_session_service.ts index 40819b6b3..afec84179 100644 --- a/core/src/sessions/base_session_service.ts +++ b/core/src/sessions/base_session_service.ts @@ -193,13 +193,9 @@ export abstract class BaseSessionService { * @param request The request to update the session state. */ private updateSessionState({session, event}: AppendEventRequest): void { - if (!event.actions || !event.actions.stateDelta) { - return; - } - for (const [key, value] of Object.entries(event.actions.stateDelta)) { - if (key.startsWith(State.TEMP_PREFIX)) { - continue; - } + for (const [key, value] of Object.entries( + event.actions?.stateDelta ?? {}, + )) { session.state[key] = value; } } @@ -216,10 +212,9 @@ export abstract class BaseSessionService { * {@link trimTempDeltaState}). */ export function applyTempState(session: Session, event: Event): void { - if (!event.actions || !event.actions.stateDelta) { - return; - } - for (const [key, value] of Object.entries(event.actions.stateDelta)) { + // `actions` is typed as required but is absent on events deserialized from a + // backend; see the VertexAiSessionService `actions: undefined` path. + for (const [key, value] of Object.entries(event.actions?.stateDelta ?? {})) { if (key.startsWith(State.TEMP_PREFIX)) { session.state[key] = value; } diff --git a/core/src/sessions/database_session_service.ts b/core/src/sessions/database_session_service.ts index 23404eabf..b2be41053 100644 --- a/core/src/sessions/database_session_service.ts +++ b/core/src/sessions/database_session_service.ts @@ -10,7 +10,6 @@ import { Options as MikroDBOptions, MikroORM, } from '@mikro-orm/core'; -import {cloneDeep} from 'lodash-es'; import {Event} from '../events/event.js'; import {randomUUID} from '../utils/env_aware_utils.js'; @@ -379,10 +378,16 @@ export class DatabaseSessionService extends BaseSessionService { return event; } - // trimTempDeltaState mutates its argument, so trim a clone for persistence - // while keeping the original event (and its temp: keys) intact for - // applyTempState after session.state is rebuilt below. - const trimmedEvent = trimTempDeltaState(cloneDeep(event)); + applyTempState(session, event); + // Captured before the transaction so it also covers temp: keys written + // earlier in this invocation via `State.set`. `session.state` is rebuilt + // from persisted rows below, which by construction never hold temp: keys. + const tempState = Object.fromEntries( + Object.entries(session.state).filter(([key]) => + key.startsWith(State.TEMP_PREFIX), + ), + ); + const trimmedEvent = trimTempDeltaState(event); await em.transactional(async (txEm) => { const storageSession = await txEm.findOne( @@ -505,23 +510,17 @@ export class DatabaseSessionService extends BaseSessionService { userStateModel.state, storageSession.state, ); - session.state = newMergedState; - - // Apply temp-scoped state to the in-memory session so it is readable for - // the remainder of the current invocation. Never persisted: - // storageSession.state and the persisted event were built from the - // trimmed delta above. - applyTempState(session, event); + session.state = {...newMergedState, ...tempState}; - const index = session.events.findIndex((e) => e.id === trimmedEvent.id); + const index = session.events.findIndex((e) => e.id === event.id); if (index >= 0) { - session.events[index] = trimmedEvent; + session.events[index] = event; } else { - session.events.push(trimmedEvent); + session.events.push(event); } session.lastUpdateTime = storageSession.updateTime.getTime(); }); - return trimmedEvent; + return event; } } diff --git a/core/test/runner/temp_state_visibility_test.ts b/core/test/runner/temp_state_visibility_test.ts new file mode 100644 index 000000000..efc075ed3 --- /dev/null +++ b/core/test/runner/temp_state_visibility_test.ts @@ -0,0 +1,93 @@ +/** + * @license + * Copyright 2026 Google LLC + * SPDX-License-Identifier: Apache-2.0 + */ + +import { + BaseLlm, + BaseLlmConnection, + Event, + InMemoryRunner, + LlmAgent, + LlmRequest, + LlmResponse, + SequentialAgent, + State, +} from '@google/adk'; +import {describe, expect, it} from 'vitest'; + +const USER_ID = 'test_user'; +const AGENT_1_OUTPUT = 'result_from_agent_1'; + +class RecordingLlm extends BaseLlm { + readonly requests: LlmRequest[] = []; + + constructor(private readonly text: string) { + super({model: 'recording-llm'}); + } + + async *generateContentAsync( + llmRequest: LlmRequest, + ): AsyncGenerator { + this.requests.push(llmRequest); + yield {content: {role: 'model', parts: [{text: this.text}]}}; + } + + async connect(_llmRequest: LlmRequest): Promise { + throw new Error('Method not implemented.'); + } +} + +describe('temp: state visibility within an invocation', () => { + it('makes a temp: outputKey readable by a later sub-agent without persisting it', async () => { + const secondModel = new RecordingLlm('rewritten'); + const agent = new SequentialAgent({ + name: 'seq', + subAgents: [ + new LlmAgent({ + name: 'a1', + model: new RecordingLlm(AGENT_1_OUTPUT), + outputKey: `${State.TEMP_PREFIX}out`, + }), + new LlmAgent({ + name: 'a2', + model: secondModel, + instruction: `Rewrite: {${State.TEMP_PREFIX}out}`, + }), + ], + }); + + const runner = new InMemoryRunner({agent}); + const session = await runner.sessionService.createSession({ + appName: runner.appName, + userId: USER_ID, + }); + + const events: Event[] = []; + for await (const event of runner.runAsync({ + userId: USER_ID, + sessionId: session.id, + newMessage: {role: 'user', parts: [{text: 'go'}]}, + })) { + events.push(event); + } + + expect(events.map((event) => event.author)).toContain('a2'); + expect(secondModel.requests).toHaveLength(1); + expect(secondModel.requests[0].config?.systemInstruction).toContain( + `Rewrite: ${AGENT_1_OUTPUT}`, + ); + + const storedSession = await runner.sessionService.getSession({ + appName: runner.appName, + userId: USER_ID, + sessionId: session.id, + }); + expect( + Object.keys(storedSession?.state ?? {}).filter((key) => + key.startsWith(State.TEMP_PREFIX), + ), + ).toEqual([]); + }); +}); diff --git a/core/test/sessions/database_session_service_test.ts b/core/test/sessions/database_session_service_test.ts index cd051a1cc..9f868a7c2 100644 --- a/core/test/sessions/database_session_service_test.ts +++ b/core/test/sessions/database_session_service_test.ts @@ -17,6 +17,11 @@ import {afterEach, beforeEach, describe, expect, it} from 'vitest'; import {isDatabaseConnectionString} from '../../src/sessions/database_session_service.js'; import {validateDatabaseSchemaVersion} from '../../src/sessions/db/operations.js'; +/** Reaches the ORM the service owns privately, to assert on persisted rows. */ +function ormOf(service: DatabaseSessionService): MikroORM { + return (service as unknown as {orm: MikroORM}).orm; +} + describe('DatabaseSessionService', () => { let service: DatabaseSessionService; @@ -31,7 +36,7 @@ describe('DatabaseSessionService', () => { afterEach(async () => { // MikroORM closing - const orm = (service as unknown as {orm: MikroORM}).orm; + const orm = ormOf(service); if (orm) { await orm.close(); } @@ -371,7 +376,7 @@ describe('DatabaseSessionService', () => { allowGlobalContext: true, }); await internalService.init(); - const orm = (internalService as unknown as {orm: MikroORM}).orm as MikroORM; + const orm = ormOf(internalService); // Manually insert bad version const em = orm.em.fork(); @@ -668,7 +673,7 @@ describe('DatabaseSessionService', () => { await service.appendEvent({session, event}); - const em = (service as unknown as {orm: MikroORM}).orm.em.fork(); + const em = ormOf(service).em.fork(); const storedEvents = (await em.find('StorageEvent', { sessionId: 's-temp', })) as {sessionId: string; eventData: Event}[]; @@ -730,7 +735,7 @@ describe('DatabaseSessionService', () => { expect(fetched?.state).not.toHaveProperty(State.TEMP_PREFIX + 'hide'); // The persisted event data is a single, temp-free StorageEvent row. - const em = (service as unknown as {orm: MikroORM}).orm.em.fork(); + const em = ormOf(service).em.fork(); const storedEvents = (await em.find('StorageEvent', { sessionId: 's-temp-apply', })) as {sessionId: string; eventData: Event}[]; @@ -742,6 +747,155 @@ describe('DatabaseSessionService', () => { ).toBeUndefined(); }); + it('trims the caller event in place and returns that same event', async () => { + const session = await service.createSession({ + appName: 'test-app', + userId: 'test-user', + sessionId: 's-temp-in-place', + }); + + const event = createEvent({ + timestamp: Date.now(), + actions: createEventActions({ + stateDelta: { + [State.TEMP_PREFIX + 'hide']: 'me', + 'keep': 'me', + }, + }), + }); + + const returnedEvent = await service.appendEvent({session, event}); + + expect(returnedEvent).toBe(event); + expect(event.actions?.stateDelta).not.toHaveProperty( + State.TEMP_PREFIX + 'hide', + ); + expect(event.actions?.stateDelta).toHaveProperty('keep', 'me'); + }); + + it('keeps temp: state from an earlier event readable when a later event is appended', async () => { + const session = await service.createSession({ + appName: 'test-app', + userId: 'test-user', + sessionId: 's-temp-across-events', + }); + + await service.appendEvent({ + session, + event: createEvent({ + timestamp: Date.now(), + actions: createEventActions({ + stateDelta: {[State.TEMP_PREFIX + 'out']: 'from-agent-1'}, + }), + }), + }); + + await service.appendEvent({ + session, + event: createEvent({ + timestamp: Date.now() + 1, + actions: createEventActions({stateDelta: {count: 1}}), + }), + }); + + expect(session.state[State.TEMP_PREFIX + 'out']).toBe('from-agent-1'); + expect(session.state['count']).toBe(1); + + const fetched = await service.getSession({ + appName: 'test-app', + userId: 'test-user', + sessionId: 's-temp-across-events', + }); + expect(fetched?.state).toHaveProperty('count', 1); + expect(fetched?.state).not.toHaveProperty(State.TEMP_PREFIX + 'out'); + }); + + it('keeps temp: state written straight onto the session state across an append', async () => { + const session = await service.createSession({ + appName: 'test-app', + userId: 'test-user', + sessionId: 's-temp-state-set', + }); + + // What `State.set('temp:x', ...)` does during an invocation. + session.state[State.TEMP_PREFIX + 'x'] = 'written-by-a-tool'; + + await service.appendEvent({ + session, + event: createEvent({ + timestamp: Date.now(), + actions: createEventActions({stateDelta: {count: 1}}), + }), + }); + + expect(session.state[State.TEMP_PREFIX + 'x']).toBe('written-by-a-tool'); + + const fetched = await service.getSession({ + appName: 'test-app', + userId: 'test-user', + sessionId: 's-temp-state-set', + }); + expect(fetched?.state).not.toHaveProperty(State.TEMP_PREFIX + 'x'); + }); + + it('leaves a partial event untouched: no state applied and no delta trimmed', async () => { + const session = await service.createSession({ + appName: 'test-app', + userId: 'test-user', + sessionId: 's-temp-partial', + }); + + const event = createEvent({ + timestamp: Date.now(), + partial: true, + actions: createEventActions({ + stateDelta: {[State.TEMP_PREFIX + 'hide']: 'me', 'keep': 'me'}, + }), + }); + + await service.appendEvent({session, event}); + + expect(session.state).not.toHaveProperty(State.TEMP_PREFIX + 'hide'); + expect(session.state).not.toHaveProperty('keep'); + expect(session.events).toHaveLength(0); + expect(event.actions?.stateDelta).toHaveProperty( + State.TEMP_PREFIX + 'hide', + 'me', + ); + }); + + it('keeps temp: state when the append reloads a stale session', async () => { + const session = await service.createSession({ + appName: 'test-app', + userId: 'test-user', + sessionId: 's-temp-stale', + }); + + await service.appendEvent({ + session, + event: createEvent({ + timestamp: Date.now(), + actions: createEventActions({ + stateDelta: {[State.TEMP_PREFIX + 'out']: 'from-agent-1'}, + }), + }), + }); + + // Force the stale-session branch, which reloads state from storage. + session.lastUpdateTime = 0; + + await service.appendEvent({ + session, + event: createEvent({ + timestamp: Date.now() + 1, + actions: createEventActions({stateDelta: {count: 1}}), + }), + }); + + expect(session.state[State.TEMP_PREFIX + 'out']).toBe('from-agent-1'); + expect(session.state['count']).toBe(1); + }); + it('should align session updateTime with event timestamp', async () => { const session = await service.createSession({ appName: 'test-app', @@ -756,7 +910,7 @@ describe('DatabaseSessionService', () => { expect(session.lastUpdateTime).toBe(timestamp); - const em = (service as unknown as {orm: MikroORM}).orm.em.fork(); + const em = ormOf(service).em.fork(); const storedSession = (await em.findOne('StorageSession', { id: 's-time', })) as {id: string; updateTime: Date}; diff --git a/core/test/sessions/in_memory_session_service_test.ts b/core/test/sessions/in_memory_session_service_test.ts index cccab1723..efddbbf65 100644 --- a/core/test/sessions/in_memory_session_service_test.ts +++ b/core/test/sessions/in_memory_session_service_test.ts @@ -736,6 +736,30 @@ describe('InMemorySessionService', () => { expect(fetched?.state).not.toHaveProperty(`${State.TEMP_PREFIX}output`); expect(fetched?.state).not.toHaveProperty(`${State.TEMP_PREFIX}output2`); }); + + it('leaves a partial event untouched: no state applied and no delta trimmed', async () => { + const session = await service.createSession({ + appName: 'app', + userId: 'user', + }); + const event = createEvent({ + timestamp: Date.now(), + partial: true, + actions: createEventActions({ + stateDelta: {[`${State.TEMP_PREFIX}k1`]: 'v1', sk: 'v2'}, + }), + }); + + await service.appendEvent({session, event}); + + expect(session.state).not.toHaveProperty(`${State.TEMP_PREFIX}k1`); + expect(session.state).not.toHaveProperty('sk'); + expect(session.events).toHaveLength(0); + expect(event.actions?.stateDelta).toHaveProperty( + `${State.TEMP_PREFIX}k1`, + 'v1', + ); + }); }); describe('applyTempState (helper)', () => { @@ -768,16 +792,14 @@ describe('InMemorySessionService', () => { expect(session.state).not.toHaveProperty(`${State.APP_PREFIX}a`); }); - it('is a no-op (no throw, no mutation) when the event has no actions or no stateDelta', () => { - const noActions = createEvent({timestamp: Date.now()}); - delete (noActions as unknown as {actions?: unknown}).actions; - const noStateDelta = createEvent({timestamp: Date.now()}); - (noStateDelta.actions as unknown as {stateDelta?: unknown}).stateDelta = - undefined; + it('is a no-op (no throw, no mutation) when the event has an empty stateDelta', () => { + const event = createEvent({ + timestamp: Date.now(), + actions: createEventActions({stateDelta: {}}), + }); const session = makeSession({existing: 1}); - expect(() => applyTempState(session, noActions)).not.toThrow(); - applyTempState(session, noStateDelta); + applyTempState(session, event); expect(session.state).toEqual({existing: 1}); }); diff --git a/core/test/sessions/vertex_ai_session_service_test.ts b/core/test/sessions/vertex_ai_session_service_test.ts index 3a8d38659..84dde18c7 100644 --- a/core/test/sessions/vertex_ai_session_service_test.ts +++ b/core/test/sessions/vertex_ai_session_service_test.ts @@ -1098,14 +1098,14 @@ describe('VertexAiSessionService', () => { }); it('applies temp: state to the in-memory session and does not send temp keys to the backend', async () => { - const session = { + const session: Session = { id: 'append-session', appName: '12345', userId: 'testUser', state: {}, events: [], lastUpdateTime: Date.now(), - } as unknown as Session; + }; const event = createEvent({ timestamp: 1620000000000,