From 4beb5f5bf47344eb04d46230c985f8b87365a816 Mon Sep 17 00:00:00 2001 From: Henry Su Date: Fri, 21 Aug 2026 12:29:02 -0500 Subject: [PATCH] fix(memory): skip malformed encrypted session envelopes --- .../extensions/memory/encrypt_session.py | 18 +++++-- .../extensions/memory/test_encrypt_session.py | 48 +++++++++++++++++++ 2 files changed, 63 insertions(+), 3 deletions(-) diff --git a/src/agents/extensions/memory/encrypt_session.py b/src/agents/extensions/memory/encrypt_session.py index 8b2eb18226..4f8dd3e12b 100644 --- a/src/agents/extensions/memory/encrypt_session.py +++ b/src/agents/extensions/memory/encrypt_session.py @@ -165,10 +165,22 @@ def _unwrap(self, item: TResponseInputItem | EncryptedEnvelope) -> TResponseInpu return cast(TResponseInputItem, item) try: - token = item["payload"].encode("utf-8") + payload = item["payload"] + if not isinstance(payload, str): + return None + token = payload.encode("utf-8") plaintext = self.cipher.decrypt(token, ttl=self.ttl) - return cast(TResponseInputItem, _from_json_bytes(plaintext)) - except (InvalidToken, KeyError): + decoded = _from_json_bytes(plaintext) + if not isinstance(decoded, dict): + return None + return cast(TResponseInputItem, decoded) + except ( + InvalidToken, + KeyError, + UnicodeDecodeError, + UnicodeEncodeError, + json.JSONDecodeError, + ): return None def _unwrap_valid_items( diff --git a/tests/extensions/memory/test_encrypt_session.py b/tests/extensions/memory/test_encrypt_session.py index 5ccbf59f8f..ff0dacec3c 100644 --- a/tests/extensions/memory/test_encrypt_session.py +++ b/tests/extensions/memory/test_encrypt_session.py @@ -33,6 +33,20 @@ def _invalid_encrypted_envelope() -> TResponseInputItem: ) +def _malformed_encrypted_envelope() -> TResponseInputItem: + return cast( + TResponseInputItem, + {"__enc__": 1, "v": 1, "kid": "hkdf-v1", "payload": None}, + ) + + +def _malformed_unicode_encrypted_envelope() -> TResponseInputItem: + return cast( + TResponseInputItem, + {"__enc__": 1, "v": 1, "kid": "hkdf-v1", "payload": "\ud800"}, + ) + + @pytest.fixture def agent() -> Agent: """Fixture for a basic agent with a scripted model.""" @@ -381,6 +395,40 @@ async def test_encrypted_session_get_items_limit_skips_invalid_latest_envelope( underlying_session.close() +async def test_encrypted_session_skips_malformed_envelopes( + encryption_key: str, underlying_session: SQLiteSession +): + """Malformed persisted envelopes should be skipped like invalid tokens.""" + session = EncryptedSession( + session_id="test_session", + underlying_session=underlying_session, + encryption_key=encryption_key, + ) + + await session.add_items([{"role": "user", "content": "valid"}]) + await underlying_session.add_items( + [_malformed_encrypted_envelope(), _malformed_unicode_encrypted_envelope()] + ) + await underlying_session.add_items( + [ + cast( + TResponseInputItem, + { + "__enc__": 1, + "v": 1, + "kid": "hkdf-v1", + "payload": session.cipher.encrypt(b"[]").decode("utf-8"), + }, + ) + ] + ) + + items = await session.get_items() + assert [item.get("content") for item in items] == ["valid"] + + underlying_session.close() + + async def test_encrypted_session_get_items_limit_returns_latest_valid_items_after_invalids( encryption_key: str, underlying_session: SQLiteSession ):