Skip to content
Open
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
16 changes: 11 additions & 5 deletions src/agents/memory/openai_conversations_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@
from agents.models._openai_shared import get_default_openai_client

from ..items import TResponseInputItem
from .session import SessionABC
from .session import SessionABC, _await_mutation
from .session_settings import SessionSettings, coerce_session_settings, resolve_session_limit


Expand Down Expand Up @@ -137,7 +137,13 @@ async def clear_session(self) -> None:
if self._session_id is None:
return

await self._openai_client.conversations.delete(
conversation_id=self._session_id,
)
self._session_id = None
session_id = self._session_id

async def delete_and_clear_session_id() -> None:
await self._openai_client.conversations.delete(
conversation_id=session_id,
)
if self._session_id == session_id:
self._session_id = None

await _await_mutation(delete_and_clear_session_id())
29 changes: 28 additions & 1 deletion src/agents/memory/session.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,10 @@
from __future__ import annotations

import asyncio
import inspect
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeGuard, runtime_checkable
from collections.abc import Awaitable
from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeGuard, TypeVar, runtime_checkable

from typing_extensions import TypedDict

Expand All @@ -12,6 +14,31 @@
from .session_settings import SessionSettings


_T = TypeVar("_T")


async def _await_mutation(awaitable: Awaitable[_T]) -> _T:
"""Wait for a mutation outcome despite repeated caller cancellation."""
task = asyncio.ensure_future(awaitable)
cancellation: asyncio.CancelledError | None = None
while not task.done():
try:
await asyncio.wait({task})
except asyncio.CancelledError as exc:
if cancellation is None:
cancellation = exc

try:
result = task.result()
except BaseException:
if cancellation is not None:
raise cancellation from None
raise
if cancellation is not None:
raise cancellation from None
return result


@runtime_checkable
class Session(Protocol):
"""Protocol for session implementations.
Expand Down
30 changes: 3 additions & 27 deletions src/agents/memory/sqlite_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,39 +5,15 @@
import sqlite3
import threading
import time
from collections.abc import Awaitable, Iterator
from collections.abc import Iterator
from contextlib import closing, contextmanager
from pathlib import Path
from typing import Any, ClassVar, TypeVar
from typing import Any, ClassVar

from ..items import TResponseInputItem
from .session import SessionABC
from .session import SessionABC, _await_mutation as _await_mutation
from .session_settings import SessionSettings, coerce_session_settings, resolve_session_limit

_T = TypeVar("_T")


async def _await_mutation(awaitable: Awaitable[_T]) -> _T:
"""Wait for a mutation outcome despite repeated caller cancellation."""
task = asyncio.ensure_future(awaitable)
cancellation: asyncio.CancelledError | None = None
while not task.done():
try:
await asyncio.wait({task})
except asyncio.CancelledError as exc:
if cancellation is None:
cancellation = exc

try:
result = task.result()
except BaseException:
if cancellation is not None:
raise cancellation from None
raise
if cancellation is not None:
raise cancellation from None
return result


class SQLiteSession(SessionABC):
"""SQLite-based implementation of session storage.
Expand Down
91 changes: 91 additions & 0 deletions tests/memory/test_openai_conversations_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -336,6 +336,97 @@ async def test_clear_session(self, mock_openai_client):
mock_openai_client.conversations.delete.assert_called_once_with(conversation_id="test_id")
assert session._session_id is None

@pytest.mark.asyncio
async def test_clear_session_cancellation_settles_delete_before_reinitializing(
self, mock_openai_client
):
"""A cancelled clear must settle deletion before the session can be reused."""
delete_started = asyncio.Event()
allow_delete_finish = asyncio.Event()
delete_finished = False

async def slow_delete(*, conversation_id: str) -> None:
nonlocal delete_finished
assert conversation_id == "old_id"
delete_started.set()
await allow_delete_finish.wait()
delete_finished = True

mock_openai_client.conversations.delete.side_effect = slow_delete
session = OpenAIConversationsSession(
conversation_id="old_id", openai_client=mock_openai_client
)
clear_task = asyncio.create_task(session.clear_session())

try:
await delete_started.wait()
clear_task.cancel("caller-cancelled")
await asyncio.sleep(0)

allow_delete_finish.set()
with pytest.raises(asyncio.CancelledError, match="caller-cancelled"):
await clear_task

items: list[Any] = [{"role": "user", "content": "Next turn"}]
await session.add_items(items)
finally:
allow_delete_finish.set()
if not clear_task.done():
clear_task.cancel()
await asyncio.gather(clear_task, return_exceptions=True)

assert session.session_id == "test_conversation_id"
assert delete_finished is True
mock_openai_client.conversations.delete.assert_awaited_once_with(conversation_id="old_id")
mock_openai_client.conversations.create.assert_awaited_once_with(items=[])
mock_openai_client.conversations.items.create.assert_awaited_once_with(
conversation_id="test_conversation_id", items=items
)

@pytest.mark.asyncio
async def test_clear_session_cancellation_preserves_replacement_session_id(
self, mock_openai_client
):
"""A settled delete must not clear a replacement conversation ID."""
delete_started = asyncio.Event()
allow_delete_finish = asyncio.Event()

async def slow_delete(*, conversation_id: str) -> None:
assert conversation_id == "old_id"
delete_started.set()
await allow_delete_finish.wait()

mock_openai_client.conversations.delete.side_effect = slow_delete
session = OpenAIConversationsSession(
conversation_id="old_id", openai_client=mock_openai_client
)
clear_task = asyncio.create_task(session.clear_session())

try:
await delete_started.wait()
clear_task.cancel("caller-cancelled")
await asyncio.sleep(0)

session.session_id = "replacement_id"
allow_delete_finish.set()
with pytest.raises(asyncio.CancelledError, match="caller-cancelled"):
await clear_task

items: list[Any] = [{"role": "user", "content": "Next turn"}]
await session.add_items(items)
finally:
allow_delete_finish.set()
if not clear_task.done():
clear_task.cancel()
await asyncio.gather(clear_task, return_exceptions=True)

assert session.session_id == "replacement_id"
mock_openai_client.conversations.delete.assert_awaited_once_with(conversation_id="old_id")
mock_openai_client.conversations.create.assert_not_awaited()
mock_openai_client.conversations.items.create.assert_awaited_once_with(
conversation_id="replacement_id", items=items
)

@pytest.mark.asyncio
async def test_clear_session_uninitialized_does_not_create_session(self, mock_openai_client):
"""Test that clear_session on an uninitialized session does not call create or delete."""
Expand Down