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
115 changes: 96 additions & 19 deletions src/agents/memory/openai_responses_compaction_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,6 +142,10 @@ def __init__(
# Serialize wrapper mutations against compaction snapshot/replace/restore so a
# cancellation rollback cannot rewrite past a newer concurrent write.
self._mutation_lock = asyncio.Lock()
# Runner persistence can carry this wrapper-local generation across the
# append-to-compaction gap. A later wrapper mutation revokes that one
# pending automatic replacement without inferring ownership from history.
self._mutation_generation = 0

@property
def client(self) -> AsyncOpenAI:
Expand Down Expand Up @@ -178,6 +182,35 @@ async def run_compaction(
When a run context is provided, the billed compaction request contributes to
that run's usage totals.
"""
# Keep one wrapper mutation boundary from the snapshot through replacement.
# A concurrent add, pop, or clear waits here and then runs against the
# compacted state instead of being overwritten by a stale replacement.
async with self._mutation_lock:
has_expected_generation = wrapper is not None and hasattr(
wrapper, "_session_compaction_generation"
)
expected_generation = (
getattr(wrapper, "_session_compaction_generation", None)
if has_expected_generation
else None
)
if has_expected_generation and (
not isinstance(expected_generation, int)
or expected_generation != self._mutation_generation
):
logger.warning(
"Skipped compaction because Session history changed after this "
"run appended its items."
)
return
await self._run_compaction_locked(args, wrapper=wrapper)
Comment thread
seratch marked this conversation as resolved.

async def _run_compaction_locked(
self,
args: OpenAIResponsesCompactionArgs | None,
*,
wrapper: RunContextWrapper[Any] | None,
) -> None:
if args and args.get("response_id"):
self._response_id = args["response_id"]
requested_mode = args.get("compaction_mode") if args else None
Expand Down Expand Up @@ -245,14 +278,18 @@ async def run_compaction(
_normalize_compaction_output_items(compacted.output or [])
)

async with self._mutation_lock:
previous_items = await self._get_all_underlying_session_items()
previous_items = await self._get_all_underlying_session_items()
try:
await self._replace_underlying_session_items(
output_items=output_items,
previous_items=previous_items,
)
self._compaction_candidate_items = select_compaction_candidate_items(output_items)
self._session_items = output_items
except (Exception, asyncio.CancelledError):
self._mutation_generation += 1
raise
self._mutation_generation += 1
self._compaction_candidate_items = select_compaction_candidate_items(output_items)
self._session_items = output_items

logger.debug(
"compact: done for %s (mode=%s, output=%s, candidates=%s)",
Expand All @@ -265,6 +302,14 @@ async def run_compaction(
async def get_items(self, limit: int | None = None) -> list[TResponseInputItem]:
return await self.underlying_session.get_items(limit)

async def _get_items_with_generation(
self, limit: int | None = None
) -> tuple[list[TResponseInputItem], int]:
"""Read one Runner snapshot with its exact wrapper generation."""
async with self._mutation_lock:
items = await self.underlying_session.get_items(limit)
return items, self._mutation_generation

async def _get_all_underlying_session_items(self) -> list[TResponseInputItem]:
return await self.underlying_session.get_items(limit=_ALL_SESSION_ITEMS_LIMIT)

Expand Down Expand Up @@ -410,37 +455,69 @@ def _clear_deferred_compaction(self) -> None:
self._deferred_response_id = None

async def add_items(self, items: list[TResponseInputItem]) -> None:
async with self._mutation_lock:
await self._add_items_locked(items)

async def _add_items_with_generation(
self,
items: list[TResponseInputItem],
*,
expected_generation: int | None,
) -> int | None:
"""Append one Runner batch and retain ownership only when its read stayed current."""
async with self._mutation_lock:
owns_generation = expected_generation == self._mutation_generation
await self._add_items_locked(items)
return self._mutation_generation if owns_generation else None

async def _add_items_locked(self, items: list[TResponseInputItem]) -> None:
try:
await self.underlying_session.add_items(items)
except (Exception, asyncio.CancelledError):
# The backend may have committed before acknowledgement failed. Re-read its
# authoritative history before compaction instead of retaining a stale cache.
self._compaction_candidate_items = None
self._session_items = None
self._mutation_generation += 1
raise
self._mutation_generation += 1
if self._compaction_candidate_items is not None:
new_items = _normalize_compaction_session_items(items)
new_candidates = select_compaction_candidate_items(new_items)
if new_candidates:
self._compaction_candidate_items.extend(new_candidates)
if self._session_items is not None:
self._session_items.extend(_normalize_compaction_session_items(items))

async def pop_item(self) -> TResponseInputItem | None:
async with self._mutation_lock:
try:
await self.underlying_session.add_items(items)
popped = await self.underlying_session.pop_item()
except (Exception, asyncio.CancelledError):
# The backend may have committed before acknowledgement failed. Re-read its
# authoritative history before compaction instead of retaining a stale cache.
self._compaction_candidate_items = None
self._session_items = None
self._mutation_generation += 1
raise
if self._compaction_candidate_items is not None:
new_items = _normalize_compaction_session_items(items)
new_candidates = select_compaction_candidate_items(new_items)
if new_candidates:
self._compaction_candidate_items.extend(new_candidates)
if self._session_items is not None:
self._session_items.extend(_normalize_compaction_session_items(items))

async def pop_item(self) -> TResponseInputItem | None:
async with self._mutation_lock:
popped = await self.underlying_session.pop_item()
if popped:
self._compaction_candidate_items = None
self._session_items = None
self._mutation_generation += 1
return popped

async def clear_session(self) -> None:
async with self._mutation_lock:
await self.underlying_session.clear_session()
try:
await self.underlying_session.clear_session()
except (Exception, asyncio.CancelledError):
self._compaction_candidate_items = None
self._session_items = None
self._deferred_response_id = None
self._mutation_generation += 1
raise
self._compaction_candidate_items = []
self._session_items = []
self._deferred_response_id = None
self._mutation_generation += 1

async def _ensure_compaction_candidates(
self,
Expand Down
46 changes: 42 additions & 4 deletions src/agents/run_internal/session_persistence.py
Original file line number Diff line number Diff line change
Expand Up @@ -199,8 +199,17 @@ async def _session_get_items(
limit: int | None | object = _SESSION_LIMIT_UNSET,
*,
wrapper: RunContextWrapper[Any] | None = None,
capture_compaction_generation: bool = False,
) -> list[TResponseInputItem]:
"""Read session items while preserving the legacy method call shape."""
get_with_generation = getattr(session, "_get_items_with_generation", None)
if capture_compaction_generation and wrapper is not None and callable(get_with_generation):
if limit is _SESSION_LIMIT_UNSET:
result, generation = await _call_session_method(get_with_generation)
else:
result, generation = await _call_session_method(get_with_generation, limit=limit)
wrapper._session_compaction_generation = generation # type: ignore[attr-defined]
return cast(list[TResponseInputItem], result)
wrapper = _get_session_wrapper(session, wrapper)
if limit is _SESSION_LIMIT_UNSET:
result = await _call_session_method(session.get_items, wrapper=wrapper)
Expand All @@ -216,6 +225,16 @@ async def _session_add_items(
wrapper: RunContextWrapper[Any] | None = None,
) -> None:
"""Append session items while preserving the legacy method call shape."""
add_with_generation = getattr(session, "_add_items_with_generation", None)
if wrapper is not None and callable(add_with_generation):
expected_generation = getattr(wrapper, "_session_compaction_generation", None)
generation = await _call_session_method(
add_with_generation,
items,
expected_generation=expected_generation,
)
wrapper._session_compaction_generation = generation # type: ignore[attr-defined]
return
wrapper = _get_session_wrapper(session, wrapper)
await _call_session_method(session.add_items, items, wrapper=wrapper)

Expand Down Expand Up @@ -360,9 +379,14 @@ async def prepare_input_with_session(
session,
limit=resolved_settings.limit,
wrapper=wrapper,
capture_compaction_generation=True,
)
else:
history = await _session_get_items(session, wrapper=wrapper)
history = await _session_get_items(
session,
wrapper=wrapper,
capture_compaction_generation=True,
)
is_openai_conversation_session = isinstance(session, OpenAIConversationsSession)
converted_history = [
strip_internal_input_item_metadata(ensure_input_item_format(item)) for item in history
Expand Down Expand Up @@ -674,9 +698,13 @@ async def save_result_to_session(
resumed_write_state._current_turn_persisted_item_count + saved_run_items_count
),
}
await resume_pending_session_write(resumed_write_state, session, wrapper=wrapper)
await resume_pending_session_write(
resumed_write_state,
session,
wrapper=compaction_wrapper,
)
else:
await _session_add_items(session, items_to_save, wrapper=wrapper)
await _session_add_items(session, items_to_save, wrapper=compaction_wrapper)

if run_state is not None:
run_state._current_turn_persisted_item_count = already_persisted + saved_run_items_count
Expand Down Expand Up @@ -801,7 +829,15 @@ def digests(items: Sequence[TResponseInputItem]) -> list[str]:
append = True
else:
expected = before + digests(pending["items"])
tail = await _session_get_items(session, limit=len(expected), wrapper=wrapper)
committed_generation: int | None = None
get_with_generation = getattr(session, "_get_items_with_generation", None)
if wrapper is not None and callable(get_with_generation):
tail, committed_generation = await _call_session_method(
get_with_generation,
limit=len(expected),
)
else:
tail = await _session_get_items(session, limit=len(expected), wrapper=wrapper)
observed = digests(tail)
committed = observed == expected
unchanged = observed[-len(before) :] == before if before else not observed
Expand All @@ -811,6 +847,8 @@ def digests(items: Sequence[TResponseInputItem]) -> list[str]:
"Repair the original Session before resuming; do not rerun the completed tool."
)
append = unchanged
if committed and committed_generation is not None and wrapper is not None:
wrapper._session_compaction_generation = committed_generation # type: ignore[attr-defined]
if append:
# Backends may retain or transform their input; the durable checkpoint stays detached.
await _session_add_items(session, copy.deepcopy(pending["items"]), wrapper=wrapper)
Expand Down
Loading
Loading