Skip to content
Merged
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
27 changes: 21 additions & 6 deletions src/agents/extensions/memory/async_sqlite_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -211,12 +211,27 @@ async def pop_item(self) -> TResponseInputItem | None:
await cursor.close()
await conn.commit()

if result:
message_data = result[0]
try:
return cast(TResponseInputItem, json.loads(message_data))
except json.JSONDecodeError:
return None
while result:
message_data = result[0]
try:
return cast(TResponseInputItem, json.loads(message_data))
except (json.JSONDecodeError, TypeError):
cursor = await conn.execute(
f"""
DELETE FROM {self.messages_table}
WHERE id = (
SELECT id FROM {self.messages_table}
WHERE session_id = ?
ORDER BY id DESC
LIMIT 1
)
RETURNING message_data
""",
(self.session_id,),
)
result = await cursor.fetchone()
await cursor.close()
await conn.commit()

return None

Expand Down
64 changes: 32 additions & 32 deletions src/agents/extensions/memory/dapr_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -344,42 +344,42 @@ async def pop_item(self) -> TResponseInputItem | None:
The most recent item if it exists, None if the session is empty
"""
async with self._lock:
attempt = 0
while True:
attempt += 1
response = await self._dapr_client.get_state(
store_name=self._state_store_name,
key=self._messages_key,
state_metadata=self._get_read_metadata(),
)
messages = self._decode_messages(response.data)
if not messages:
return None
last_item = messages.pop()
messages_json = json.dumps(messages, separators=(",", ":"))
etag = getattr(response, "etag", None) or None
etag = getattr(response, "etag", None) or None
try:
await self._dapr_client.save_state(
attempt = 0
while True:
attempt += 1
response = await self._dapr_client.get_state(
store_name=self._state_store_name,
key=self._messages_key,
value=messages_json,
etag=etag,
state_metadata=self._get_metadata(),
options=self._get_state_options(concurrency=Concurrency.first_write),
state_metadata=self._get_read_metadata(),
)
break
except Exception as error:
should_retry = await self._handle_concurrency_conflict(error, attempt)
if should_retry:
continue
raise
try:
if isinstance(last_item, str):
return await self._deserialize_item(last_item)
return last_item # type: ignore[no-any-return]
except (json.JSONDecodeError, TypeError):
return None
messages = self._decode_messages(response.data)
if not messages:
return None
last_item = messages.pop()
messages_json = json.dumps(messages, separators=(",", ":"))
etag = getattr(response, "etag", None) or None
try:
await self._dapr_client.save_state(
store_name=self._state_store_name,
key=self._messages_key,
value=messages_json,
etag=etag,
state_metadata=self._get_metadata(),
options=self._get_state_options(concurrency=Concurrency.first_write),
)
break
except Exception as error:
should_retry = await self._handle_concurrency_conflict(error, attempt)
if should_retry:
continue
raise
try:
if isinstance(last_item, str):
return await self._deserialize_item(last_item)
return last_item # type: ignore[no-any-return]
except (json.JSONDecodeError, TypeError):
continue

async def clear_session(self) -> None:
"""Clear all items for this session."""
Expand Down
33 changes: 17 additions & 16 deletions src/agents/extensions/memory/redis_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,22 +223,23 @@ async def pop_item(self) -> TResponseInputItem | None:
The most recent item if it exists, None if the session is empty
"""
async with self._lock:
# Use RPOP to atomically remove and return the rightmost (most recent) item
raw_msg = await self._redis.rpop(self._messages_key) # type: ignore[misc] # Redis library returns Union[Awaitable[T], T] in async context

if raw_msg is None:
return None

try:
# Handle both bytes (default) and str (decode_responses=True) Redis clients
if isinstance(raw_msg, bytes):
msg_str = raw_msg.decode("utf-8")
else:
msg_str = raw_msg # Already a string
return await self._deserialize_item(msg_str)
except (json.JSONDecodeError, UnicodeDecodeError):
# Return None for corrupted messages (already removed)
return None
while True:
# Use RPOP to atomically remove and return the rightmost (most recent) item
raw_msg = await self._redis.rpop(self._messages_key) # type: ignore[misc] # Redis library returns Union[Awaitable[T], T] in async context

if raw_msg is None:
return None

try:
# Handle both bytes (default) and str (decode_responses=True) Redis clients
if isinstance(raw_msg, bytes):
msg_str = raw_msg.decode("utf-8")
else:
msg_str = raw_msg # Already a string
return await self._deserialize_item(msg_str)
except (json.JSONDecodeError, UnicodeDecodeError):
# Drop corrupted messages and keep looking for a valid item.
continue

async def clear_session(self) -> None:
"""Clear all items for this session."""
Expand Down
53 changes: 27 additions & 26 deletions src/agents/extensions/memory/sqlalchemy_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -385,33 +385,34 @@ async def pop_item(self) -> TResponseInputItem | None:
await self._ensure_tables()
async with self._session_factory() as sess:
async with sess.begin():
# Fallback for all dialects - get ID first, then delete
subq = (
select(self._messages.c.id)
.where(self._messages.c.session_id == self.session_id)
.order_by(
self._messages.c.created_at.desc(),
self._messages.c.id.desc(),
while True:
# Fallback for all dialects - get ID first, then delete
subq = (
select(self._messages.c.id)
.where(self._messages.c.session_id == self.session_id)
.order_by(
self._messages.c.created_at.desc(),
self._messages.c.id.desc(),
)
.limit(1)
)
.limit(1)
)
res = await sess.execute(subq)
row_id = res.scalar_one_or_none()
if row_id is None:
return None
# Fetch data before deleting
res_data = await sess.execute(
select(self._messages.c.message_data).where(self._messages.c.id == row_id)
)
row = res_data.scalar_one_or_none()
await sess.execute(delete(self._messages).where(self._messages.c.id == row_id))

if row is None:
return None
try:
return await self._deserialize_item(row)
except json.JSONDecodeError:
return None
res = await sess.execute(subq)
row_id = res.scalar_one_or_none()
if row_id is None:
return None
# Fetch data before deleting
res_data = await sess.execute(
select(self._messages.c.message_data).where(self._messages.c.id == row_id)
)
row = res_data.scalar_one_or_none()
await sess.execute(delete(self._messages).where(self._messages.c.id == row_id))

if row is None:
continue
try:
return await self._deserialize_item(row)
except (json.JSONDecodeError, TypeError):
continue

async def clear_session(self) -> None:
"""Clear all items for this session."""
Expand Down
24 changes: 19 additions & 5 deletions src/agents/memory/sqlite_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -246,7 +246,7 @@ def _get_items_sync():
try:
item = json.loads(message_data)
items.append(item)
except json.JSONDecodeError:
except (json.JSONDecodeError, TypeError):
# Skip invalid JSON entries
continue

Expand Down Expand Up @@ -297,14 +297,28 @@ def _pop_item_sync():
result = cursor.fetchone()
conn.commit()

if result:
while result:
message_data = result[0]
try:
item = json.loads(message_data)
return item
except json.JSONDecodeError:
# Return None for corrupted JSON entries (already deleted)
return None
except (json.JSONDecodeError, TypeError):
# Drop corrupted JSON entries and keep looking for a valid item.
cursor = conn.execute(
f"""
DELETE FROM {self.messages_table}
WHERE id = (
SELECT id FROM {self.messages_table}
WHERE session_id = ?
ORDER BY id DESC
LIMIT 1
)
RETURNING message_data
""",
(self.session_id,),
)
result = cursor.fetchone()
conn.commit()

return None

Expand Down
41 changes: 41 additions & 0 deletions tests/extensions/memory/test_async_sqlite_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,47 @@ async def test_async_sqlite_session_pop_item():
await session.close()


async def test_async_sqlite_session_pop_item_skips_corrupt_most_recent():
"""pop_item skips corrupt newest rows and returns the next valid item."""
with tempfile.TemporaryDirectory() as temp_dir:
db_path = Path(temp_dir) / "async_pop_corrupt.db"
session = AsyncSQLiteSession("async_pop_corrupt", db_path)

valid_item: TResponseInputItem = {"role": "user", "content": "valid"}
await session.add_items([valid_item])

conn = await session._get_connection()
await conn.execute(
f"INSERT INTO {session.messages_table} (session_id, message_data) VALUES (?, ?)",
(session.session_id, "not valid json {{{"),
)
await conn.commit()

assert await session.pop_item() == valid_item
assert await session.get_items() == []

await session.close()


async def test_async_sqlite_session_pop_item_returns_none_after_dropping_only_corrupt_rows():
"""pop_item removes corrupt rows and returns None when no valid items remain."""
with tempfile.TemporaryDirectory() as temp_dir:
db_path = Path(temp_dir) / "async_pop_only_corrupt.db"
session = AsyncSQLiteSession("async_pop_only_corrupt", db_path)

conn = await session._get_connection()
await conn.execute(
f"INSERT INTO {session.messages_table} (session_id, message_data) VALUES (?, ?)",
(session.session_id, "not valid json {{{"),
)
await conn.commit()

assert await session.pop_item() is None
assert await session.get_items() == []

await session.close()


async def test_async_sqlite_session_get_items_limit():
"""Test AsyncSQLiteSession get_items limit handling."""
with tempfile.TemporaryDirectory() as temp_dir:
Expand Down
35 changes: 35 additions & 0 deletions tests/extensions/memory/test_dapr_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -396,6 +396,41 @@ async def test_pop_from_empty_session(fake_dapr_client: FakeDaprClient):
await session.close()


async def test_pop_item_skips_corrupt_most_recent(fake_dapr_client: FakeDaprClient):
"""pop_item skips corrupt newest entries and returns the next valid item."""
session = await _create_test_session(fake_dapr_client, "pop_corrupt")

try:
valid_item: TResponseInputItem = {"role": "user", "content": "valid"}
fake_dapr_client._state[session._messages_key] = json.dumps(
[await session._serialize_item(valid_item), "not valid json {{{"],
separators=(",", ":"),
).encode("utf-8")

assert await session.pop_item() == valid_item
assert await session.get_items() == []
finally:
await session.close()


async def test_pop_item_returns_none_after_dropping_only_corrupt_entries(
fake_dapr_client: FakeDaprClient,
):
"""pop_item removes corrupt entries and returns None when no valid items remain."""
session = await _create_test_session(fake_dapr_client, "pop_only_corrupt")

try:
fake_dapr_client._state[session._messages_key] = json.dumps(
["not valid json {{{"],
separators=(",", ":"),
).encode("utf-8")

assert await session.pop_item() is None
assert await session.get_items() == []
finally:
await session.close()


async def test_add_empty_items_list(fake_dapr_client: FakeDaprClient):
"""Test that adding an empty list of items is a no-op."""
session = await _create_test_session(fake_dapr_client)
Expand Down
10 changes: 4 additions & 6 deletions tests/extensions/memory/test_redis_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -740,12 +740,11 @@ async def test_corrupted_data_handling():
assert items[0].get("content") == "valid message"
assert items[1].get("content") == "valid after corruption"

# Test pop_item with corrupted data at the end
# Test pop_item with corrupted data at the end.
await _safe_rpush(fake_redis, messages_key, "corrupted at end")

# The corrupted item should be handled gracefully
# Since it's at the end, pop_item will encounter it first and return None
# But first, let's pop the valid items to get to the corrupted one
# The corrupted item should be dropped and pop_item should keep looking
# for the next valid item.
popped1 = await session.pop_item()
assert popped1 is not None
assert popped1.get("content") == "valid after corruption"
Expand All @@ -754,8 +753,7 @@ async def test_corrupted_data_handling():
assert popped2 is not None
assert popped2.get("content") == "valid message"

# Now we should hit the corrupted data - this should gracefully handle it
# by returning None (and removing the corrupted item)
# All corrupt items were removed while looking for valid messages.
popped_corrupted = await session.pop_item()
assert popped_corrupted is None

Expand Down
Loading
Loading