|
1 | 1 | from __future__ import annotations |
2 | 2 |
|
| 3 | +import asyncio |
| 4 | +import json |
3 | 5 | import tempfile |
4 | 6 | from pathlib import Path |
| 7 | +from types import SimpleNamespace |
5 | 8 | from typing import Any, cast |
| 9 | +from unittest.mock import AsyncMock, MagicMock |
6 | 10 |
|
7 | 11 | import pytest |
8 | 12 |
|
|
12 | 16 |
|
13 | 17 | from agents import ( |
14 | 18 | Agent, |
| 19 | + RunConfig, |
15 | 20 | RunContextWrapper, |
16 | 21 | Runner, |
| 22 | + RunState, |
17 | 23 | SessionSettings, |
18 | 24 | SQLiteSession, |
19 | 25 | TResponseInputItem, |
20 | 26 | ) |
| 27 | +from agents.decorators import tool |
21 | 28 | 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 |
24 | 32 |
|
25 | 33 | # Mark all tests in this file as asyncio |
26 | 34 | pytestmark = pytest.mark.asyncio |
@@ -126,6 +134,228 @@ async def test_encrypted_session_with_runner( |
126 | 134 | underlying_session.close() |
127 | 135 |
|
128 | 136 |
|
| 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 | + |
129 | 359 | async def test_encrypted_session_pop_item(encryption_key: str, underlying_session: SQLiteSession): |
130 | 360 | """Test pop_item functionality.""" |
131 | 361 | session = EncryptedSession( |
|
0 commit comments