Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 6 additions & 11 deletions core/src/sessions/base_session_service.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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;
}
}
Expand All @@ -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;
}
Expand Down
31 changes: 15 additions & 16 deletions core/src/sessions/database_session_service.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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';
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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;
}
}
93 changes: 93 additions & 0 deletions core/test/runner/temp_state_visibility_test.ts
Original file line number Diff line number Diff line change
@@ -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<LlmResponse, void> {
this.requests.push(llmRequest);
yield {content: {role: 'model', parts: [{text: this.text}]}};
}

async connect(_llmRequest: LlmRequest): Promise<BaseLlmConnection> {
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([]);
});
});
164 changes: 159 additions & 5 deletions core/test/sessions/database_session_service_test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand All @@ -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();
}
Expand Down Expand Up @@ -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();
Expand Down Expand Up @@ -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}[];
Expand Down Expand Up @@ -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}[];
Expand All @@ -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',
Expand All @@ -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};
Expand Down
Loading