From 7366a7f3a8505996f08a90d242f187098759695c Mon Sep 17 00:00:00 2001 From: Harshit Sharma <66710144+harshitethic@users.noreply.github.com> Date: Mon, 14 Sep 2026 18:32:07 +0530 Subject: [PATCH] fix(sessions): preserve ciphertext on wrong-key pop --- .../extensions/memory/encrypt_session.py | 48 ++++++++++++++++--- .../memory/test_encrypt_session_wrong_key.py | 36 ++++++++++++++ 2 files changed, 77 insertions(+), 7 deletions(-) create mode 100644 tests/extensions/memory/test_encrypt_session_wrong_key.py diff --git a/src/agents/extensions/memory/encrypt_session.py b/src/agents/extensions/memory/encrypt_session.py index 9b2d05884f..b897e3f7d4 100644 --- a/src/agents/extensions/memory/encrypt_session.py +++ b/src/agents/extensions/memory/encrypt_session.py @@ -212,9 +212,34 @@ def _unwrap(self, item: TResponseInputItem | EncryptedEnvelope) -> TResponseInpu token = item["payload"].encode("utf-8") plaintext = self.cipher.decrypt(token, ttl=self.ttl) return cast(TResponseInputItem, _from_json_bytes(plaintext)) - except (InvalidToken, KeyError): - return None - + except (InvalidToken, KeyError): + return None + + def _unwrap_for_pop( + self, item: TResponseInputItem | EncryptedEnvelope + ) -> tuple[TResponseInputItem | None, bool]: + """Unwrap a popped item and report whether authentication failed. + + Fernet raises ``InvalidToken`` for both an expired token and a token that + cannot be authenticated with the configured key. ``pop_item`` needs to + distinguish those cases so a wrong key cannot drain recoverable history. + """ + if not _is_encrypted_envelope(item): + return cast(TResponseInputItem, item), False + + try: + token = item["payload"].encode("utf-8") + plaintext = self.cipher.decrypt(token, ttl=self.ttl) + return cast(TResponseInputItem, _from_json_bytes(plaintext)), False + except KeyError: + return None, False + except InvalidToken: + try: + self.cipher.decrypt(token) + except InvalidToken: + return None, True + return None, False + def _unwrap_valid_items( self, encrypted_items: list[TResponseInputItem] ) -> list[TResponseInputItem]: @@ -290,10 +315,19 @@ async def pop_item( ) if not enc: return None - item = self._unwrap(enc) - if item is not None: - return item - + item, authentication_failed = self._unwrap_for_pop(enc) + if item is not None: + return item + if authentication_failed: + # Put back the exact item returned by the atomic backend pop. + # This avoids a get-then-pop race while refusing to drain history. + await _call_session_method( + self.underlying_session.add_items, + [cast(TResponseInputItem, enc)], + wrapper=wrapper, + ) + return None + async def clear_session( self, *, diff --git a/tests/extensions/memory/test_encrypt_session_wrong_key.py b/tests/extensions/memory/test_encrypt_session_wrong_key.py new file mode 100644 index 0000000000..cba1a72ff5 --- /dev/null +++ b/tests/extensions/memory/test_encrypt_session_wrong_key.py @@ -0,0 +1,36 @@ +from __future__ import annotations + +import pytest + +pytest.importorskip("cryptography") + +from agents import SQLiteSession +from agents.extensions.memory.encrypt_session import EncryptedSession + +pytestmark = pytest.mark.asyncio + + +async def test_wrong_key_pop_preserves_recoverable_ciphertext(tmp_path) -> None: + store = SQLiteSession("conversation", tmp_path / "history.db") + try: + correct = EncryptedSession( + session_id="conversation", underlying_session=store, + encryption_key="example-correct-key", ttl=3600, + ) + wrong = EncryptedSession( + session_id="conversation", underlying_session=store, + encryption_key="example-wrong-key", ttl=3600, + ) + messages = [ + {"role": "user", "content": "hello"}, + {"role": "assistant", "content": "hi"}, + {"role": "user", "content": "follow-up"}, + ] + await correct.add_items(messages) + stored_before = await store.get_items() + + assert await wrong.pop_item() is None + assert await store.get_items() == stored_before + assert await correct.get_items() == messages + finally: + store.close()