Skip to content

Commit 63cab19

Browse files
committed
fix: Preserve Session transformations during generation tracking
1 parent 02c205f commit 63cab19

4 files changed

Lines changed: 269 additions & 27 deletions

File tree

src/agents/memory/openai_responses_compaction_session.py

Lines changed: 11 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -303,11 +303,12 @@ async def get_items(self, limit: int | None = None) -> list[TResponseInputItem]:
303303
return await self.underlying_session.get_items(limit)
304304

305305
async def _get_items_with_generation(
306-
self, limit: int | None = None
306+
self, read_items: Callable[[], Awaitable[list[TResponseInputItem]]]
307307
) -> tuple[list[TResponseInputItem], int]:
308308
"""Read one Runner snapshot with its exact wrapper generation."""
309309
async with self._mutation_lock:
310-
items = await self.underlying_session.get_items(limit)
310+
# Read through the outer Session so its decryption and filtering still apply.
311+
items = await read_items()
311312
return items, self._mutation_generation
312313

313314
async def _get_all_underlying_session_items(self) -> list[TResponseInputItem]:
@@ -460,15 +461,18 @@ async def add_items(self, items: list[TResponseInputItem]) -> None:
460461

461462
async def _add_items_with_generation(
462463
self,
463-
items: list[TResponseInputItem],
464+
write_items: Callable[[], Awaitable[None]],
464465
*,
465466
expected_generation: int | None,
466467
) -> int | None:
467468
"""Append one Runner batch and retain ownership only when its read stayed current."""
468-
async with self._mutation_lock:
469-
owns_generation = expected_generation == self._mutation_generation
470-
await self._add_items_locked(items)
471-
return self._mutation_generation if owns_generation else None
469+
# The outer Session applies its transformations before our public add_items acquires
470+
# the mutation lock. Any intervening mutation revokes ownership, including a write
471+
# that completes while the outer Session is awaiting acknowledgement after our append.
472+
await write_items()
473+
if expected_generation is not None and self._mutation_generation == expected_generation + 1:
474+
return self._mutation_generation
475+
return None
472476

473477
async def _add_items_locked(self, items: list[TResponseInputItem]) -> None:
474478
try:

src/agents/run_internal/session_persistence.py

Lines changed: 21 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -202,20 +202,23 @@ async def _session_get_items(
202202
capture_compaction_generation: bool = False,
203203
) -> list[TResponseInputItem]:
204204
"""Read session items while preserving the legacy method call shape."""
205-
get_with_generation = getattr(session, "_get_items_with_generation", None)
206-
if capture_compaction_generation and wrapper is not None and callable(get_with_generation):
205+
session_wrapper = _get_session_wrapper(session, wrapper)
206+
207+
async def read_items() -> list[TResponseInputItem]:
207208
if limit is _SESSION_LIMIT_UNSET:
208-
result, generation = await _call_session_method(get_with_generation)
209+
result = await _call_session_method(session.get_items, wrapper=session_wrapper)
209210
else:
210-
result, generation = await _call_session_method(get_with_generation, limit=limit)
211+
result = await _call_session_method(
212+
session.get_items, limit=limit, wrapper=session_wrapper
213+
)
214+
return cast(list[TResponseInputItem], result)
215+
216+
get_with_generation = getattr(session, "_get_items_with_generation", None)
217+
if capture_compaction_generation and wrapper is not None and callable(get_with_generation):
218+
result, generation = await _call_session_method(get_with_generation, read_items)
211219
wrapper._session_compaction_generation = generation # type: ignore[attr-defined]
212220
return cast(list[TResponseInputItem], result)
213-
wrapper = _get_session_wrapper(session, wrapper)
214-
if limit is _SESSION_LIMIT_UNSET:
215-
result = await _call_session_method(session.get_items, wrapper=wrapper)
216-
else:
217-
result = await _call_session_method(session.get_items, limit=limit, wrapper=wrapper)
218-
return cast(list[TResponseInputItem], result)
221+
return await read_items()
219222

220223

221224
async def _session_add_items(
@@ -225,18 +228,22 @@ async def _session_add_items(
225228
wrapper: RunContextWrapper[Any] | None = None,
226229
) -> None:
227230
"""Append session items while preserving the legacy method call shape."""
231+
session_wrapper = _get_session_wrapper(session, wrapper)
232+
233+
async def write_items() -> None:
234+
await _call_session_method(session.add_items, items, wrapper=session_wrapper)
235+
228236
add_with_generation = getattr(session, "_add_items_with_generation", None)
229237
if wrapper is not None and callable(add_with_generation):
230238
expected_generation = getattr(wrapper, "_session_compaction_generation", None)
231239
generation = await _call_session_method(
232240
add_with_generation,
233-
items,
241+
write_items,
234242
expected_generation=expected_generation,
235243
)
236244
wrapper._session_compaction_generation = generation # type: ignore[attr-defined]
237245
return
238-
wrapper = _get_session_wrapper(session, wrapper)
239-
await _call_session_method(session.add_items, items, wrapper=wrapper)
246+
await write_items()
240247

241248

242249
async def _session_pop_item(
@@ -834,7 +841,7 @@ def digests(items: Sequence[TResponseInputItem]) -> list[str]:
834841
if wrapper is not None and callable(get_with_generation):
835842
tail, committed_generation = await _call_session_method(
836843
get_with_generation,
837-
limit=len(expected),
844+
lambda: _session_get_items(session, limit=len(expected), wrapper=wrapper),
838845
)
839846
else:
840847
tail = await _session_get_items(session, limit=len(expected), wrapper=wrapper)

tests/extensions/memory/test_encrypt_session.py

Lines changed: 232 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,8 +1,12 @@
11
from __future__ import annotations
22

3+
import asyncio
4+
import json
35
import tempfile
46
from pathlib import Path
7+
from types import SimpleNamespace
58
from typing import Any, cast
9+
from unittest.mock import AsyncMock, MagicMock
610

711
import pytest
812

@@ -12,15 +16,19 @@
1216

1317
from agents import (
1418
Agent,
19+
RunConfig,
1520
RunContextWrapper,
1621
Runner,
22+
RunState,
1723
SessionSettings,
1824
SQLiteSession,
1925
TResponseInputItem,
2026
)
27+
from agents.decorators import tool
2128
from agents.extensions.memory.encrypt_session import EncryptedSession
22-
from agents.testing import ScriptedModel
23-
from tests.test_responses import get_text_message
29+
from agents.memory import OpenAIResponsesCompactionSession
30+
from agents.testing import ModelStep, ScriptedModel
31+
from tests.test_responses import get_function_tool_call, get_text_message
2432

2533
# Mark all tests in this file as asyncio
2634
pytestmark = pytest.mark.asyncio
@@ -126,6 +134,228 @@ async def test_encrypted_session_with_runner(
126134
underlying_session.close()
127135

128136

137+
async def _run_encrypted_session(
138+
agent: Agent[Any], value: str | RunState[Any], session: EncryptedSession, streamed: bool
139+
):
140+
config = RunConfig(tracing_disabled=True)
141+
if not streamed:
142+
return await Runner.run(agent, value, session=session, run_config=config)
143+
result = Runner.run_streamed(agent, value, session=session, run_config=config)
144+
async for _ in result.stream_events():
145+
pass
146+
return result
147+
148+
149+
def _decrypt_stored_items(
150+
session: EncryptedSession, stored: list[TResponseInputItem]
151+
) -> list[dict[str, Any]]:
152+
# Inspect the storage envelope and decrypt directly, independently of Session.get_items.
153+
envelopes = cast(list[dict[str, Any]], stored)
154+
assert all(set(item) == {"__enc__", "v", "kid", "payload"} for item in envelopes)
155+
assert all(item["__enc__"] == 1 for item in envelopes)
156+
return [json.loads(session.cipher.decrypt(item["payload"].encode())) for item in envelopes]
157+
158+
159+
@pytest.mark.parametrize("streamed", [False, True])
160+
async def test_runner_encrypts_items_around_compaction(
161+
streamed: bool, encryption_key: str, tmp_path: Path
162+
) -> None:
163+
backend = SQLiteSession("encrypted-compaction", tmp_path / "history.db")
164+
session = EncryptedSession(
165+
backend.session_id,
166+
OpenAIResponsesCompactionSession(
167+
backend.session_id, backend, should_trigger_compaction=lambda _: False
168+
),
169+
encryption_key,
170+
)
171+
output = get_text_message("private assistant answer")
172+
model = ScriptedModel([[output], [get_text_message("followup answer")]])
173+
agent = Agent(name="test", model=model)
174+
expected = [
175+
{"role": "user", "content": "private user input"},
176+
output.model_dump(exclude_unset=True),
177+
]
178+
try:
179+
result = await _run_encrypted_session(agent, "private user input", session, streamed)
180+
assert result.final_output == "private assistant answer"
181+
assert _decrypt_stored_items(session, await backend.get_items()) == expected
182+
183+
await _run_encrypted_session(agent, "followup", session, streamed)
184+
assert model.calls[-1].input == expected + [{"role": "user", "content": "followup"}]
185+
finally:
186+
backend.close()
187+
188+
189+
@pytest.mark.parametrize("streamed", [False, True])
190+
async def test_runner_decrypts_existing_compaction_history_with_ttl_and_limit(
191+
streamed: bool, encryption_key: str, tmp_path: Path, set_fernet_time: Any
192+
) -> None:
193+
backend = SQLiteSession("encrypted-history", tmp_path / "history.db")
194+
session = EncryptedSession(
195+
backend.session_id,
196+
OpenAIResponsesCompactionSession(
197+
backend.session_id, backend, should_trigger_compaction=lambda _: False
198+
),
199+
encryption_key,
200+
ttl=10,
201+
)
202+
expected_history = [{"role": "user", "content": "retained history"}]
203+
try:
204+
set_fernet_time(1_000)
205+
await session.add_items([{"role": "assistant", "content": "expired history"}])
206+
expired = (await backend.get_items())[0]
207+
set_fernet_time(1_020)
208+
await session.add_items(cast(list[TResponseInputItem], expected_history))
209+
# An expired tail forces EncryptedSession to expand its one-item retrieval window.
210+
await backend.add_items([expired])
211+
session.session_settings = SessionSettings(limit=1)
212+
model = ScriptedModel([[get_text_message("answer")]])
213+
214+
await _run_encrypted_session(Agent(name="test", model=model), "next", session, streamed)
215+
216+
assert model.calls[0].input == expected_history + [{"role": "user", "content": "next"}]
217+
finally:
218+
backend.close()
219+
220+
221+
@pytest.mark.parametrize("streamed", [False, True])
222+
async def test_runner_recovers_encrypted_resumed_append(
223+
streamed: bool, encryption_key: str, tmp_path: Path
224+
) -> None:
225+
class LostAckSQLiteSession(SQLiteSession):
226+
"""Fail after a real committed batch to exercise public Runner recovery."""
227+
228+
fail_after_commit = False
229+
230+
async def add_items(self, items: list[TResponseInputItem]) -> None:
231+
await super().add_items(items)
232+
if self.fail_after_commit:
233+
self.fail_after_commit = False
234+
raise RuntimeError("append acknowledgement lost")
235+
236+
backend = LostAckSQLiteSession("encrypted-resume", tmp_path / "history.db")
237+
session = EncryptedSession(
238+
backend.session_id,
239+
OpenAIResponsesCompactionSession(
240+
backend.session_id, backend, should_trigger_compaction=lambda _: False
241+
),
242+
encryption_key,
243+
)
244+
effects: list[str] = []
245+
246+
@tool(needs_approval=True)
247+
async def record() -> str:
248+
effects.append("recorded")
249+
return "private tool receipt"
250+
251+
call = get_function_tool_call("record", "{}", call_id="record-1")
252+
output = get_text_message("private final answer")
253+
model = ScriptedModel([[call], [output]])
254+
agent = Agent(name="test", model=model, tools=[record])
255+
try:
256+
paused = await _run_encrypted_session(agent, "private request", session, streamed)
257+
state = paused.to_state()
258+
state.approve(state.get_interruptions()[0])
259+
backend.fail_after_commit = True
260+
with pytest.raises(RuntimeError, match="append acknowledgement lost"):
261+
await _run_encrypted_session(agent, state, session, streamed)
262+
assert effects == ["recorded"]
263+
assert len(model.calls) == 1
264+
state = await RunState.from_json(agent, state.to_json())
265+
266+
result = await _run_encrypted_session(agent, state, session, streamed)
267+
268+
assert result.final_output == "private final answer"
269+
assert effects == ["recorded"]
270+
assert "pending_session_write" not in result.to_state().to_json()
271+
assert _decrypt_stored_items(session, await backend.get_items()) == [
272+
{"role": "user", "content": "private request"},
273+
call.model_dump(exclude_unset=True),
274+
{
275+
"type": "function_call_output",
276+
"call_id": "record-1",
277+
"output": "private tool receipt",
278+
},
279+
output.model_dump(exclude_unset=True),
280+
]
281+
finally:
282+
backend.close()
283+
284+
285+
@pytest.mark.parametrize(
286+
"pause_before_append", [False, True], ids=["after-append", "before-append"]
287+
)
288+
async def test_encrypted_runner_skips_compaction_after_interleaved_outer_write(
289+
pause_before_append: bool, encryption_key: str, tmp_path: Path
290+
) -> None:
291+
paused = asyncio.Event()
292+
release = asyncio.Event()
293+
294+
class PausingEncryptedSession(EncryptedSession):
295+
"""Suspend at the outer public append boundary while another Runner completes."""
296+
297+
async def add_items(
298+
self,
299+
items: list[TResponseInputItem],
300+
*,
301+
wrapper: RunContextWrapper[Any] | None = None,
302+
) -> None:
303+
is_run_a_output = "answer-A" in json.dumps(items)
304+
if is_run_a_output and pause_before_append:
305+
paused.set()
306+
await release.wait()
307+
await super().add_items(items, wrapper=wrapper)
308+
if is_run_a_output and not pause_before_append:
309+
paused.set()
310+
await release.wait()
311+
312+
backend = SQLiteSession("encrypted-concurrent", tmp_path / "history.db")
313+
client = MagicMock()
314+
client.responses.compact = AsyncMock(return_value=SimpleNamespace(output=[]))
315+
session = PausingEncryptedSession(
316+
backend.session_id,
317+
OpenAIResponsesCompactionSession(
318+
backend.session_id,
319+
backend,
320+
client=client,
321+
should_trigger_compaction=lambda context: context["response_id"] == "resp-a",
322+
),
323+
encryption_key,
324+
)
325+
output_a = get_text_message("answer-A")
326+
output_b = get_text_message("answer-B")
327+
agent_a = Agent(
328+
name="A", model=ScriptedModel([ModelStep(output=[output_a], response_id="resp-a")])
329+
)
330+
agent_b = Agent(
331+
name="B", model=ScriptedModel([ModelStep(output=[output_b], response_id="resp-b")])
332+
)
333+
run_a = asyncio.create_task(_run_encrypted_session(agent_a, "input-A", session, False))
334+
try:
335+
await asyncio.wait_for(paused.wait(), timeout=2)
336+
result_b = await asyncio.wait_for(
337+
_run_encrypted_session(agent_b, "input-B", session, False), timeout=2
338+
)
339+
assert result_b.final_output == "answer-B"
340+
release.set()
341+
result_a = await asyncio.wait_for(run_a, timeout=2)
342+
assert result_a.final_output == "answer-A"
343+
client.responses.compact.assert_not_awaited()
344+
expected = [
345+
{"role": "user", "content": "input-A"},
346+
{"role": "user", "content": "input-B"},
347+
output_b.model_dump(exclude_unset=True),
348+
]
349+
expected.insert(3 if pause_before_append else 1, output_a.model_dump(exclude_unset=True))
350+
assert _decrypt_stored_items(session, await backend.get_items()) == expected
351+
finally:
352+
release.set()
353+
if not run_a.done():
354+
run_a.cancel()
355+
await asyncio.gather(run_a, return_exceptions=True)
356+
backend.close()
357+
358+
129359
async def test_encrypted_session_pop_item(encryption_key: str, underlying_session: SQLiteSession):
130360
"""Test pop_item functionality."""
131361
session = EncryptedSession(

0 commit comments

Comments
 (0)