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
139 changes: 101 additions & 38 deletions src/agents/extensions/memory/advanced_sqlite_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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__(
Expand All @@ -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:
Expand All @@ -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.
Expand All @@ -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
Expand Down Expand Up @@ -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.

Expand All @@ -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
Comment thread
seratch marked this conversation as resolved.
)
# 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))
Comment thread
jbeckwith-oai marked this conversation as resolved.
except Exception as exc:
log_model_and_tool_action_error(self._logger, "Failed to add session items", exc)
raise
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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))

Expand All @@ -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
Expand All @@ -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]:
Expand All @@ -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(
"""
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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"""
Expand Down Expand Up @@ -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

Expand Down
Loading
Loading