diff --git a/src/agents/extensions/memory/advanced_sqlite_session.py b/src/agents/extensions/memory/advanced_sqlite_session.py index d7c13d845d..2100d2a613 100644 --- a/src/agents/extensions/memory/advanced_sqlite_session.py +++ b/src/agents/extensions/memory/advanced_sqlite_session.py @@ -5,10 +5,13 @@ import logging import re import sqlite3 +import threading import time +from collections.abc import Callable, Coroutine +from concurrent.futures import Future from contextlib import closing from pathlib import Path -from typing import Any, ClassVar, cast +from typing import Any, ClassVar, TypeVar, cast from agents.result import RunResult from agents.usage import Usage @@ -25,6 +28,8 @@ from ...memory.session_settings import SessionSettings, resolve_session_limit from ...memory.sqlite_session import _await_mutation +_T = TypeVar("_T") + def _allow_all_sqlite_actions( _action: int, @@ -73,6 +78,13 @@ def __init__( logger: The logger to use. Defaults to the module logger **kwargs: Additional keyword arguments to pass to the superclass """ # noqa: E501 + # Publish branch and generation together so caller-thread snapshots cannot + # combine a cleared branch with the new generation from a worker update. + self._branch_state: tuple[str, int] = ("main", 0) + # Order history mutations and usage capture before dispatching workers. + # The bookkeeping lock never protects SQLite work or an await. + self._turn_operation_lock = threading.Lock() + self._turn_operation_tail: Future[None] | None = None self._create_structure_tables_on_init = create_tables try: super().__init__( @@ -87,13 +99,16 @@ def __init__( except BaseException: pass raise - self._current_branch_id = "main" - # Synchronized with the durable session_clear_generations row whenever a - # branch pointer is established or a write begins. A mismatch means - # another instance cleared the session, so the local pointer resets to main. - self._generation = 0 self._logger = logger if logger is not None else logging.getLogger(__name__) + @property + def _current_branch_id(self) -> str: + return self._branch_state[0] + + @property + def _generation(self) -> int: + return self._branch_state[1] + def _init_db_for_connection(self, conn: sqlite3.Connection) -> None: """Initialize base tables only after validating advanced-table ownership.""" if self._create_structure_tables_on_init: @@ -104,6 +119,7 @@ def _init_db_for_connection(self, conn: sqlite3.Connection) -> None: self._claim_structure_tables(conn) self._create_schema_for_connection(conn) conn.commit() + self._refresh_branch_after_external_clear(conn, initialize=False) def _commit_branch_pointer(self, branch_id: str, generation: int) -> bool: """Set the current-branch pointer unless a clear has committed meanwhile. @@ -123,11 +139,9 @@ def _commit_branch_pointer(self, branch_id: str, generation: int) -> bool: ).fetchone() durable_generation = row[0] if row is not None else 0 if durable_generation != generation: - self._generation = durable_generation - self._current_branch_id = "main" + self._branch_state = ("main", durable_generation) return False - self._generation = durable_generation - self._current_branch_id = branch_id + self._branch_state = (branch_id, durable_generation) return True # The structure tables that record which base-table pair owns a database file, and the @@ -344,6 +358,33 @@ def _init_structure_tables(self, conn: sqlite3.Connection) -> None: ON turn_usage(session_id, branch_id, user_turn_number) """) + def _run_ordered_turn_operation( + self, worker: Callable[..., _T], *args: Any + ) -> Coroutine[Any, Any, _T]: + """Reserve history/capture order before the caller first yields. + + Completion futures are independent of any event loop, like the session's + thread-local database connections. A failed operation still releases its + successor; each caller owns cancellation settlement via _await_mutation. + """ + completion: Future[None] = Future() + with self._turn_operation_lock: + predecessor = self._turn_operation_tail + self._turn_operation_tail = completion + + async def _run() -> _T: + try: + if predecessor is not None: + await asyncio.wrap_future(predecessor) + return await asyncio.to_thread(worker, *args) + finally: + completion.set_result(None) + with self._turn_operation_lock: + if self._turn_operation_tail is completion: + self._turn_operation_tail = None + + return _run() + async def add_items(self, items: list[TResponseInputItem]) -> None: """Add items to the session. @@ -356,17 +397,24 @@ async def add_items(self, items: list[TResponseInputItem]) -> None: if not items: return + branch_id, generation = self._branch_state + def _add_items_sync(): """Synchronous helper to add items and structure metadata together.""" with self._write_connection() as conn: self._refresh_branch_after_external_clear(conn) + # Match pop_item's reset behavior after a clear, while keeping a + # later branch switch from redirecting an already queued append. + target_branch = ( + self._current_branch_id if self._generation != generation else branch_id + ) # Keep both writes in one transaction so metadata failures do not leave orphans. self._insert_items(conn, items) - self._insert_structure_metadata(conn, items) + self._insert_structure_metadata(conn, items, target_branch) conn.commit() try: - await _await_mutation(asyncio.to_thread(_add_items_sync)) + await _await_mutation(self._run_ordered_turn_operation(_add_items_sync)) except Exception as exc: log_model_and_tool_action_error(self._logger, "Failed to add session items", exc) raise @@ -473,8 +521,7 @@ async def pop_item(self) -> TResponseInputItem | None: # Snapshot the current branch at call time so a concurrent # switch_to_branch() cannot redirect this pop to a different branch once # it has been dispatched to the worker thread. - branch_id = self._current_branch_id - generation = self._generation + branch_id, generation = self._branch_state def _pop_item_sync(): with self._write_connection() as conn: @@ -551,7 +598,7 @@ def _pop_item_sync(): # Drop corrupted JSON entries and keep looking for a valid item. continue - return await _await_mutation(asyncio.to_thread(_pop_item_sync)) + return await _await_mutation(self._run_ordered_turn_operation(_pop_item_sync)) async def clear_session(self) -> None: """Clear all items for this session. @@ -611,8 +658,7 @@ def _clear_session_sync(): # the pointer still references a deleted branch. Bumping the # generation invalidates any in-flight switch/create that # captured the pre-clear generation. - self._generation = generation - self._current_branch_id = "main" + self._branch_state = ("main", generation) await _await_mutation(asyncio.to_thread(_clear_session_sync)) @@ -621,6 +667,8 @@ async def store_run_usage(self, result: RunResult) -> None: This is designed to be called after `Runner.run()` completes. Session-level usage can be aggregated from turn data when needed. + Appends and pops started later on this session instance cannot change + which turn this call captures. Args: result: The result from the run @@ -635,14 +683,24 @@ async def store_run_usage(self, result: RunResult) -> None: # write is skipped. The anchor is scoped to this branch/turn, so # unrelated removals (e.g. delete_branch on another branch) do # not drop this write. - current_turn, branch_id, turn_anchor = self._capture_current_turn() - # Only update turn-level usage - session usage is aggregated on demand - await self._update_turn_usage_internal( - current_turn, - result.context_wrapper.usage, - branch_id=branch_id, - turn_anchor=turn_anchor, + branch_id, generation = self._branch_state + usage = result.context_wrapper.usage + capture = self._run_ordered_turn_operation( + self._capture_current_turn, branch_id, generation ) + + async def _store_usage() -> None: + current_turn, captured_branch, turn_anchor = await capture + # Only update turn-level usage; session usage is aggregated on demand. + await self._update_turn_usage_internal( + current_turn, + usage, + branch_id=captured_branch, + turn_anchor=turn_anchor, + ) + + # Capture and write must both settle before caller cancellation propagates. + await _await_mutation(_store_usage()) except Exception as e: def diagnostic_extra() -> dict[str, object]: @@ -655,17 +713,20 @@ def diagnostic_extra() -> dict[str, object]: diagnostic_extra=diagnostic_extra, ) - def _capture_current_turn(self) -> tuple[int, str, int | None]: + def _capture_current_turn(self, branch_id: str, generation: int) -> tuple[int, str, int | None]: """Return (current_turn, branch_id, turn_anchor) in one locked read. ``turn_anchor`` is the smallest ``message_structure.id`` of the current - turn on the current branch (``None`` if the turn has no rows). Because + turn on the captured branch (``None`` if the turn has no rows). Because ids are monotonic and never reused, it uniquely identifies this turn incarnation, so a later pop+recreate that reuses the numeric turn id yields a different anchor. """ with self._locked_connection() as conn: - branch_id = self._resolve_read_branch(conn, None) + self._refresh_branch_after_external_clear(conn, initialize=False) + if self._generation != generation: + # A clear invalidates usage belonging to the previous session history. + return 0, branch_id, None with closing(conn.cursor()) as cursor: cursor.execute( """ @@ -792,7 +853,9 @@ def _insert_structure_metadata( self, conn: sqlite3.Connection, items: list[TResponseInputItem], + branch_id: str | None = None, ) -> None: + target_branch = self._current_branch_id if branch_id is None else branch_id # Get the IDs of messages we just inserted, in order. with closing(conn.cursor()) as cursor: cursor.execute( @@ -830,7 +893,7 @@ def _insert_structure_metadata( FROM message_structure WHERE session_id = ? AND branch_id = ? """, - (self.session_id, self._current_branch_id), + (self.session_id, target_branch), ) result = cursor.fetchone() current_turn = result[0] if result else 0 @@ -856,7 +919,7 @@ def _insert_structure_metadata( ( self.session_id, msg_id, - self._current_branch_id, + target_branch, msg_type, seq_start + i + 1, item_turn, @@ -1212,7 +1275,9 @@ def _delete_sync(): return usage_deleted, structure_deleted, orphaned_messages_deleted usage_deleted, structure_deleted, orphaned_messages_deleted = await _await_mutation( - asyncio.to_thread(_delete_sync) + # Earlier queued appends must settle before deletion, or their captured + # branch target could recreate rows after this deletion commits. + self._run_ordered_turn_operation(_delete_sync) ) self._logger.info( @@ -1349,8 +1414,7 @@ def _refresh_branch_after_external_clear( ).fetchone() generation = row[0] if row is not None else 0 if generation != self._generation: - self._generation = generation - self._current_branch_id = "main" + self._branch_state = ("main", generation) def _resolve_read_branch( self, @@ -1420,8 +1484,7 @@ def _copy_sync() -> tuple[str, Any, str, int]: conn.execute("BEGIN IMMEDIATE") self._ensure_branch_reservations_table(conn) self._refresh_branch_after_external_clear(conn) - source_branch_id = self._current_branch_id - generation = self._generation + source_branch_id, generation = self._branch_state with closing(conn.cursor()) as cursor: cursor.execute( f""" @@ -1927,15 +1990,15 @@ async def _update_turn_usage_internal( turn reused the same numeric id. Because the check is scoped to this branch/turn, unrelated removals (e.g. delete_branch on another branch) do not drop this write. ``None`` means the branch - had no turn when it was read, so there is nothing to attribute - the usage to and the write is skipped. + had no turn when it was read or a session clear invalidated the + pending capture, so the write is skipped. """ target_branch = branch_id if branch_id is not None else self._current_branch_id if turn_anchor is None: - # ``_capture_current_turn`` returns no anchor only when the branch has no - # turn rows; recording usage would invent a phantom turn 0. + # A missing anchor means there is no captured turn or a session clear + # invalidated the capture; neither case can safely receive usage. self._logger.debug("Skipping usage store: no current turn on branch %r", target_branch) return diff --git a/tests/extensions/memory/test_advanced_sqlite_session.py b/tests/extensions/memory/test_advanced_sqlite_session.py index f8d084adf9..bbf3839ac3 100644 --- a/tests/extensions/memory/test_advanced_sqlite_session.py +++ b/tests/extensions/memory/test_advanced_sqlite_session.py @@ -215,11 +215,12 @@ def _insert_structure_metadata( self, conn: Any, items: list[TResponseInputItem], + branch_id: str | None = None, ) -> None: if self.fail_structure_metadata_once: self.fail_structure_metadata_once = False raise RuntimeError("structure metadata failed") - super()._insert_structure_metadata(conn, items) + super()._insert_structure_metadata(conn, items, branch_id) class PartiallyFailingStructureMetadataSession(AdvancedSQLiteSession): @@ -229,6 +230,7 @@ def _insert_structure_metadata( self, conn: Any, items: list[TResponseInputItem], + branch_id: str | None = None, ) -> None: cursor = conn.execute( f"SELECT id FROM {self.messages_table} WHERE session_id = ? ORDER BY id ASC LIMIT 1", @@ -245,7 +247,16 @@ def _insert_structure_metadata( user_turn_number, branch_turn_number, tool_name) VALUES (?, ?, ?, ?, ?, ?, ?, ?) """, - (self.session_id, row[0], self._current_branch_id, "user", 1, 1, 1, None), + ( + self.session_id, + row[0], + self._current_branch_id if branch_id is None else branch_id, + "user", + 1, + 1, + 1, + None, + ), ) raise RuntimeError("structure metadata failed after partial write") @@ -3654,6 +3665,557 @@ async def test_clear_before_branch_transaction_prevents_stale_reservation(): session.close() +async def test_store_run_usage_waits_for_database_lock_off_event_loop( + tmp_path: Path, + usage_data: Usage, + monkeypatch: pytest.MonkeyPatch, +): + """A contended usage read must let the event loop release the database lock.""" + session = AdvancedSQLiteSession( + session_id="responsive_usage", + db_path=tmp_path / "responsive_usage.db", + create_tables=True, + ) + await session.add_items([{"role": "user", "content": "first turn"}]) + loop = asyncio.get_running_loop() + database_lock = session._lock + lock_held = threading.Event() + release_lock = threading.Event() + event_loop_progress: list[bool] = [] + + def hold_database_lock() -> None: + with database_lock: + lock_held.set() + # Bound a broken implementation so the test cannot deadlock its event loop. + event_loop_progress.append(release_lock.wait(timeout=2)) + + class ObservedLock: + # Instrument the existing lock acquisition without replacing SQLite or usage writes. + def __enter__(self): + loop.call_soon_threadsafe(release_lock.set) + return database_lock.__enter__() + + def __exit__(self, *args: Any): + return database_lock.__exit__(*args) + + holder = asyncio.create_task(asyncio.to_thread(hold_database_lock)) + try: + assert await asyncio.to_thread(lock_held.wait, 2) + with monkeypatch.context() as patcher: + patcher.setattr(session, "_lock", ObservedLock()) + await session.store_run_usage(create_mock_run_result(usage_data)) + await holder + + assert event_loop_progress == [True] + assert await session.get_items() == [{"role": "user", "content": "first turn"}] + recorded_usage = await session.get_turn_usage(1) + assert isinstance(recorded_usage, dict) + assert recorded_usage["total_tokens"] == usage_data.total_tokens + finally: + release_lock.set() + await holder + session.close() + + +@pytest.mark.parametrize("create_tables", [False, True]) +async def test_store_run_usage_after_reopening_cleared_session( + tmp_path: Path, usage_data: Usage, create_tables: bool +): + """Reopening existing history must not discard its first usage update.""" + db_path = tmp_path / "reopened_usage.db" + original = AdvancedSQLiteSession( + session_id="reopened_usage", db_path=db_path, create_tables=True + ) + items: list[TResponseInputItem] = [{"role": "user", "content": "current turn"}] + try: + await original.add_items([{"role": "user", "content": "old turn"}]) + await original.clear_session() + await original.add_items(items) + finally: + original.close() + + reopened = AdvancedSQLiteSession( + session_id="reopened_usage", db_path=db_path, create_tables=create_tables + ) + try: + # Store before any read can refresh the reopened session's local state. + await reopened.store_run_usage(create_mock_run_result(usage_data)) + recorded_usage = await reopened.get_turn_usage(1) + assert isinstance(recorded_usage, dict) + assert recorded_usage["total_tokens"] == usage_data.total_tokens + assert await reopened.get_items() == items + finally: + reopened.close() + + +async def test_store_run_usage_cancellation_during_capture_waits_for_commit(usage_data: Usage): + """Cancellation during capture must settle the complete usage mutation.""" + session = AdvancedSQLiteSession(session_id="usage_capture_cancel", create_tables=True) + task: asyncio.Task[None] | None = None + try: + await session.add_items([{"role": "user", "content": "usage turn"}]) + with _gate_worker("_capture_current_turn") as (started, real_to_thread, release): + try: + task = asyncio.create_task( + session.store_run_usage(create_mock_run_result(usage_data)) + ) + assert await real_to_thread(started.wait, 5) + task.cancel("capture-cancel") + await asyncio.sleep(0) + cancelled_before_release = task.done() + release.set() + with pytest.raises(asyncio.CancelledError) as exc_info: + await task + _assert_cancel_message(exc_info.value, "capture-cancel") + finally: + release.set() + if task is not None: + await asyncio.gather(task, return_exceptions=True) + + recorded_usage = await session.get_turn_usage(1) + assert isinstance(recorded_usage, dict) + assert recorded_usage["total_tokens"] == usage_data.total_tokens + assert not cancelled_before_release + finally: + session.close() + + +@pytest.mark.parametrize("operation", ["switch", "clear"]) +async def test_store_run_usage_keeps_target_while_capture_is_waiting( + usage_data: Usage, operation: str +): + """Usage must not move to another branch or replacement session history.""" + session = AdvancedSQLiteSession(session_id=f"usage_capture_{operation}", create_tables=True) + task: asyncio.Task[None] | None = None + replacement: asyncio.Task[None] | None = None + try: + await session.add_items([{"role": "user", "content": "main turn"}]) + if operation == "switch": + await session.create_branch_from_turn(1, "branch_a") + await session.add_items([{"role": "user", "content": "branch turn"}]) + + with _gate_worker("_capture_current_turn") as (started, real_to_thread, release): + try: + task = asyncio.create_task( + session.store_run_usage(create_mock_run_result(usage_data)) + ) + assert await real_to_thread(started.wait, 5) + if operation == "switch": + await session.switch_to_branch("main") + else: + await session.clear_session() + replacement = asyncio.create_task( + session.add_items([{"role": "user", "content": "replacement turn"}]) + ) + release.set() + await task + if replacement is not None: + await replacement + finally: + release.set() + if task is not None: + await asyncio.gather(task, return_exceptions=True) + if replacement is not None: + await asyncio.gather(replacement, return_exceptions=True) + + assert await session.get_turn_usage(branch_id="main") == [] + if operation == "switch": + branch_usage = await session.get_turn_usage(branch_id="branch_a") + assert isinstance(branch_usage, list) + assert len(branch_usage) == 1 + assert branch_usage[0]["total_tokens"] == usage_data.total_tokens + assert await session.get_items() == [{"role": "user", "content": "main turn"}] + else: + assert await session.get_items() == [{"role": "user", "content": "replacement turn"}] + finally: + session.close() + + +@pytest.mark.parametrize("operation", ["append", "replace"]) +async def test_store_run_usage_keeps_turn_when_history_changes_during_capture( + usage_data: Usage, operation: str +): + """Later history must not receive usage from a capture that is still waiting.""" + session = AdvancedSQLiteSession(session_id=f"capture_turn_{operation}", create_tables=True) + newer_usage = Usage(requests=1, input_tokens=7, output_tokens=4, total_tokens=11) + mutation_entered = asyncio.Event() + mutation_dispatched = asyncio.Event() + tasks: list[asyncio.Task[None]] = [] + try: + await session.add_items([{"role": "user", "content": "original turn"}]) + with _gate_worker("_capture_current_turn") as (started, real_to_thread, release): + gated_to_thread = asyncio.to_thread + + async def observe_dispatch(func, /, *args, **kwargs): + if getattr(func, "__name__", "") in ("_add_items_sync", "_pop_item_sync"): + mutation_dispatched.set() + return await gated_to_thread(func, *args, **kwargs) + + async def change_history() -> None: + mutation_entered.set() + if operation == "replace": + await session.pop_item() + await session.add_items([{"role": "user", "content": "newer turn"}]) + await session.store_run_usage(create_mock_run_result(newer_usage)) + + with patch( + "agents.extensions.memory.advanced_sqlite_session.asyncio.to_thread", + observe_dispatch, + ): + try: + tasks.append( + asyncio.create_task( + session.store_run_usage(create_mock_run_result(usage_data)) + ) + ) + assert await real_to_thread(started.wait, 5) + tasks.append(asyncio.create_task(change_history())) + await mutation_entered.wait() + # Let the completion-owned mutation task reach its first await. + await asyncio.sleep(0) + if mutation_dispatched.is_set(): + # Expose the original race if history can overtake capture. + await tasks[1] + release.set() + await asyncio.gather(*tasks) + finally: + release.set() + await asyncio.gather(*tasks, return_exceptions=True) + + newer_turn = 2 if operation == "append" else 1 + recorded_newer = await session.get_turn_usage(newer_turn) + assert isinstance(recorded_newer, dict) + assert recorded_newer["total_tokens"] == newer_usage.total_tokens + if operation == "append": + recorded_original = await session.get_turn_usage(1) + assert isinstance(recorded_original, dict) + assert recorded_original["total_tokens"] == usage_data.total_tokens + else: + assert await session.get_items() == [{"role": "user", "content": "newer turn"}] + finally: + session.close() + + +@pytest.mark.parametrize("publication", ["clear", "refresh"]) +async def test_add_captures_coherent_branch_after_clear_publication( + tmp_path: Path, publication: str +): + """A clear publication cannot pair a removed branch with its new generation.""" + published = threading.Event() + release = threading.Event() + + class PausedPublicationSession(AdvancedSQLiteSession): + pause_publication = False + + def __setattr__(self, name: str, value: Any) -> None: + super().__setattr__(name, value) + # The legacy path exposes its split-field publication; the current + # path pauses after publishing the complete immutable branch state. + if name in {"_generation", "_branch_state"} and self.pause_publication: + self.pause_publication = False + published.set() + assert release.wait(5), "Branch-state publication was not released" + + db_path = tmp_path / "coherent_branch.db" + session = PausedPublicationSession( + session_id="coherent_branch", db_path=db_path, create_tables=True + ) + clearer = AdvancedSQLiteSession(session_id="coherent_branch", db_path=db_path) + tasks: list[asyncio.Task[Any]] = [] + added: TResponseInputItem = {"role": "user", "content": "after clear"} + try: + await session.add_items([{"role": "user", "content": "main"}]) + await session.create_branch_from_turn(1, "selected") + await session.add_items([{"role": "user", "content": "old selected"}]) + if publication == "refresh": + await clearer.clear_session() + session.pause_publication = True + tasks.append( + asyncio.create_task( + session.clear_session() if publication == "clear" else session.get_items() + ) + ) + assert await asyncio.to_thread(published.wait, 5) + entered = asyncio.Event() + + async def append() -> None: + entered.set() + await session.add_items([added]) + + tasks.append(asyncio.create_task(append())) + await entered.wait() + release.set() + await asyncio.gather(*tasks) + assert await session.get_items() == [added] + assert await session.get_items(branch_id="selected") == [] + assert session._current_branch_id == "main" + finally: + release.set() + await asyncio.gather(*tasks, return_exceptions=True) + session.close() + clearer.close() + + +@pytest.mark.parametrize("branch_turn", [1, 2]) +@pytest.mark.parametrize("operation", ["switch", "clear"]) +async def test_queued_add_keeps_selected_branch( + usage_data: Usage, branch_turn: int, operation: str +): + """A queued append retains branch ownership without reviving cleared history.""" + session = AdvancedSQLiteSession(session_id="queued_add_branch", create_tables=True) + original: list[TResponseInputItem] = [ + {"role": "user", "content": "main first"}, + {"role": "user", "content": "main second"}, + ] + added: TResponseInputItem = {"role": "user", "content": "queued branch item"} + tasks: list[asyncio.Task[None]] = [] + try: + await session.add_items(original) + await session.create_branch_from_turn(branch_turn, "selected") + with _gate_worker("_capture_current_turn") as (started, real_to_thread, release): + entered = asyncio.Event() + + async def append() -> None: + entered.set() + await session.add_items([added]) + + try: + capture = asyncio.create_task( + session.store_run_usage(create_mock_run_result(usage_data)) + ) + tasks.append(capture) + assert await real_to_thread(started.wait, 5) + mutation = asyncio.create_task(append()) + tasks.append(mutation) + # The add has reserved its queue position before this task resumes. + await entered.wait() + assert not mutation.done() + if operation == "switch": + await session.switch_to_branch("main") + else: + await session.clear_session() + release.set() + await asyncio.gather(*tasks) + finally: + release.set() + await asyncio.gather(*tasks, return_exceptions=True) + + assert session._current_branch_id == "main" + if operation == "switch": + assert await session.get_items() == original + assert await session.get_items(branch_id="selected") == original[: branch_turn - 1] + [ + added + ] + turns = await session.get_conversation_by_turns("selected") + assert list(turns) == list(range(1, branch_turn + 1)) + else: + assert await session.get_items() == [added] + assert await session.get_items(branch_id="selected") == [] + assert await session.get_turn_usage() == [] + finally: + session.close() + + +@pytest.mark.parametrize("force", [False, True]) +async def test_queued_append_cannot_resurrect_deleted_branch( + usage_data: Usage, monkeypatch: pytest.MonkeyPatch, force: bool +): + session = AdvancedSQLiteSession(session_id="queued_append_delete", create_tables=True) + main_items: list[TResponseInputItem] = [ + {"role": "user", "content": "main first"}, + {"role": "user", "content": "main second"}, + ] + tasks: list[asyncio.Task[Any]] = [] + append_entered = asyncio.Event() + delete_scheduled = asyncio.Event() + delete_queued = False + order_operation = session._run_ordered_turn_operation + + def observe_order(worker, *args): + nonlocal delete_queued + operation = order_operation(worker, *args) + if worker.__name__ == "_delete_sync": + delete_queued = True + delete_scheduled.set() + return operation + + async def append() -> None: + append_entered.set() + await session.add_items([{"role": "user", "content": "queued branch item"}]) + + try: + await session.add_items(main_items) + await session.create_branch_from_turn(2, "selected") + with _gate_worker("_capture_current_turn") as (started, real_to_thread, release): + gated_to_thread = asyncio.to_thread + + async def observe_dispatch(worker, /, *args, **kwargs): + if worker.__name__ == "_delete_sync": + delete_scheduled.set() + return await gated_to_thread(worker, *args, **kwargs) + + with monkeypatch.context() as patcher: + patcher.setattr(session, "_run_ordered_turn_operation", observe_order) + patcher.setattr(asyncio, "to_thread", observe_dispatch) + try: + tasks.append( + asyncio.create_task( + session.store_run_usage(create_mock_run_result(usage_data)) + ) + ) + assert await real_to_thread(started.wait, 5) + tasks.append(asyncio.create_task(append())) + await append_entered.wait() + if not force: + await session.switch_to_branch("main") + deletion = asyncio.create_task(session.delete_branch("selected", force=force)) + tasks.append(deletion) + await asyncio.wait_for(delete_scheduled.wait(), 5) + if not delete_queued: + # Expose the old ordering: deletion can commit before the + # earlier append is dispatched. No timing-based sleep is needed. + await deletion + release.set() + await asyncio.gather(*tasks) + finally: + release.set() + await asyncio.gather(*tasks, return_exceptions=True) + + assert session._current_branch_id == "main" + assert await session.get_items() == main_items + assert await session.get_items(branch_id="selected") == [] + assert {branch["branch_id"] for branch in await session.list_branches()} == {"main"} + assert await session.get_turn_usage(branch_id="selected") == [] + assert _count_rows(session, session.messages_table) == len(main_items) + assert _count_rows(session, "message_structure") == len(main_items) + assert _count_rows(session, "turn_usage") == 0 + finally: + session.close() + + +@pytest.mark.parametrize("outcome", ["failure", "cancel"]) +async def test_queued_turn_mutation_releases_later_usage_capture(usage_data: Usage, outcome: str): + """A failed or caller-cancelled queued mutation cannot strand later usage.""" + session = FailingOnceStructureMetadataSession( + session_id=f"queued_turn_{outcome}", create_tables=True + ) + session.fail_structure_metadata_once = False + user: TResponseInputItem = {"role": "user", "content": "original turn"} + assistant: TResponseInputItem = {"role": "assistant", "content": "original response"} + replacement: TResponseInputItem = {"role": "user", "content": "next turn"} + tasks: list[asyncio.Task[Any]] = [] + try: + await session.add_items([user, assistant]) + session.fail_structure_metadata_once = outcome == "failure" + with _gate_worker("_pop_item_sync") as (started, real_to_thread, release): + try: + predecessor = asyncio.create_task(session.pop_item()) + tasks.append(predecessor) + assert await real_to_thread(started.wait, 5) + mutation_entered = asyncio.Event() + capture_entered = asyncio.Event() + + async def add_next_turn() -> None: + mutation_entered.set() + await session.add_items([replacement]) + + async def capture_usage() -> None: + capture_entered.set() + await session.store_run_usage(create_mock_run_result(usage_data)) + + mutation = asyncio.create_task(add_next_turn()) + tasks.append(mutation) + # The public add call has run to its first await before this resumes. + await mutation_entered.wait() + if outcome == "cancel": + mutation.cancel("queued-add-cancel") + capture = asyncio.create_task(capture_usage()) + tasks.append(capture) + await capture_entered.wait() + assert not mutation.done() + assert not capture.done() + release.set() + _, pending = await asyncio.wait(tasks, timeout=5) + assert not pending, "A settled predecessor stranded the operation queue" + assert await predecessor == assistant + if outcome == "failure": + with pytest.raises(RuntimeError, match="structure metadata failed"): + await mutation + else: + with pytest.raises(asyncio.CancelledError) as exc_info: + await mutation + _assert_cancel_message(exc_info.value, "queued-add-cancel") + await capture + finally: + release.set() + if tasks: + await asyncio.wait(tasks, timeout=5) + + expected_turn = 1 if outcome == "failure" else 2 + assert await session.get_items() == ( + [user] if outcome == "failure" else [user, replacement] + ) + recorded = await session.get_turn_usage() + assert isinstance(recorded, list) + assert [(row["user_turn_number"], row["total_tokens"]) for row in recorded] == [ + (expected_turn, usage_data.total_tokens) + ] + finally: + session.close() + + +async def test_turn_ordering_survives_reuse_on_another_event_loop(usage_data: Usage): + """Both loops must contend: uncontended reuse would miss asyncio.Lock binding.""" + session = AdvancedSQLiteSession(session_id="usage_two_loops", create_tables=True) + + async def exercise_loop(original_turn: int) -> None: + await session.add_items([{"role": "user", "content": f"usage owner {original_turn}"}]) + tasks: list[asyncio.Task[Any]] = [] + with _gate_worker("_capture_current_turn") as (started, real_to_thread, release): + try: + capture = asyncio.create_task( + session.store_run_usage(create_mock_run_result(usage_data)) + ) + tasks.append(capture) + assert await real_to_thread(started.wait, 5) + append_entered = asyncio.Event() + + async def append_turn() -> None: + append_entered.set() + await session.add_items( + [{"role": "user", "content": f"following {original_turn}"}] + ) + + append = asyncio.create_task(append_turn()) + tasks.append(append) + await append_entered.wait() + assert not append.done() + release.set() + _, pending = await asyncio.wait(tasks, timeout=5) + assert not pending, "Contended operations did not settle on the new loop" + await asyncio.gather(*tasks) + finally: + release.set() + if tasks: + await asyncio.wait(tasks, timeout=5) + + def use_two_loops() -> None: + asyncio.run(exercise_loop(1)) + asyncio.run(exercise_loop(3)) + + try: + # Run outside pytest's active loop; both asyncio.run calls share the session. + await asyncio.to_thread(use_two_loops) + recorded = await session.get_turn_usage() + assert isinstance(recorded, list) + assert [(row["user_turn_number"], row["total_tokens"]) for row in recorded] == [ + (1, usage_data.total_tokens), + (3, usage_data.total_tokens), + ] + finally: + session.close() + + async def test_stale_store_run_usage_skipped_when_turn_removed_by_pop(usage_data: Usage): """A store_run_usage that reads a turn and then races with pop_item removing that turn must not reinsert usage for the now-nonexistent turn.