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
67 changes: 57 additions & 10 deletions src/agents/extensions/memory/encrypt_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,8 +37,9 @@
from typing_extensions import TypedDict

from ...items import TResponseInputItem
from ...memory.session import SessionABC
from ...memory.session import SessionABC, _call_session_method, _get_session_wrapper
from ...memory.session_settings import SessionSettings, resolve_session_limit
from ...run_context import RunContextWrapper


class EncryptedEnvelope(TypedDict):
Expand Down Expand Up @@ -180,34 +181,80 @@ def _unwrap_valid_items(
valid_items.append(item)
return valid_items

async def get_items(self, limit: int | None = None) -> list[TResponseInputItem]:
async def get_items(
self,
limit: int | None = None,
*,
wrapper: RunContextWrapper[Any] | None = None,
) -> list[TResponseInputItem]:
wrapper = _get_session_wrapper(self.underlying_session, wrapper)
effective_limit = resolve_session_limit(limit, self.session_settings)
if effective_limit is not None and effective_limit > 0:
window = effective_limit
while True:
encrypted_items = await self.underlying_session.get_items(window)
encrypted_items = cast(
list[TResponseInputItem],
await _call_session_method(
self.underlying_session.get_items,
window,
wrapper=wrapper,
),
)
valid_items = self._unwrap_valid_items(encrypted_items)
if len(valid_items) >= effective_limit:
return valid_items[-effective_limit:]
if len(encrypted_items) < window:
return valid_items
window *= 2

encrypted_items = await self.underlying_session.get_items(limit)
encrypted_items = cast(
list[TResponseInputItem],
await _call_session_method(
self.underlying_session.get_items,
limit,
wrapper=wrapper,
),
)
return self._unwrap_valid_items(encrypted_items)

async def add_items(self, items: list[TResponseInputItem]) -> None:
async def add_items(
self,
items: list[TResponseInputItem],
*,
wrapper: RunContextWrapper[Any] | None = None,
) -> None:
wrapper = _get_session_wrapper(self.underlying_session, wrapper)
wrapped: list[EncryptedEnvelope] = [self._wrap(it) for it in items]
await self.underlying_session.add_items(cast(list[TResponseInputItem], wrapped))
await _call_session_method(
self.underlying_session.add_items,
cast(list[TResponseInputItem], wrapped),
wrapper=wrapper,
)

async def pop_item(self) -> TResponseInputItem | None:
async def pop_item(
self,
*,
wrapper: RunContextWrapper[Any] | None = None,
) -> TResponseInputItem | None:
wrapper = _get_session_wrapper(self.underlying_session, wrapper)
while True:
enc = await self.underlying_session.pop_item()
enc = await _call_session_method(
self.underlying_session.pop_item,
wrapper=wrapper,
)
if not enc:
return None
item = self._unwrap(enc)
if item is not None:
return item

async def clear_session(self) -> None:
await self.underlying_session.clear_session()
async def clear_session(
self,
*,
wrapper: RunContextWrapper[Any] | None = None,
) -> None:
wrapper = _get_session_wrapper(self.underlying_session, wrapper)
await _call_session_method(
self.underlying_session.clear_session,
wrapper=wrapper,
)
64 changes: 63 additions & 1 deletion src/agents/memory/session.py
Original file line number Diff line number Diff line change
@@ -1,12 +1,14 @@
from __future__ import annotations

import inspect
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING, Literal, Protocol, TypeGuard, runtime_checkable
from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeGuard, runtime_checkable

from typing_extensions import TypedDict

if TYPE_CHECKING:
from ..items import TResponseInputItem
from ..run_context import RunContextWrapper
from .session_settings import SessionSettings


Expand Down Expand Up @@ -148,3 +150,63 @@ def is_openai_responses_compaction_aware_session(
except Exception:
return False
return callable(run_compaction)


def _session_method_accepts_wrapper(method: Any) -> bool:
Comment thread
seratch marked this conversation as resolved.
"""Return whether a session method opts into receiving ``wrapper``.

The public ``Session`` protocol keeps its released signatures so existing structural
implementations remain type-compatible. Custom sessions can opt in by adding a ``wrapper``
parameter that can be passed by keyword.
"""
try:
parameters = inspect.signature(method).parameters.values()
except Exception:
return False

return any(
parameter.name == "wrapper"
and parameter.kind
in (inspect.Parameter.POSITIONAL_OR_KEYWORD, inspect.Parameter.KEYWORD_ONLY)
for parameter in parameters
)


def _session_accepts_wrapper(session: Any) -> bool:
"""Return whether every history operation accepts ``wrapper``."""
try:
methods = (
session.get_items,
session.add_items,
session.pop_item,
session.clear_session,
)
except Exception:
return False
return all(_session_method_accepts_wrapper(method) for method in methods)


def _get_session_wrapper(
session: Any,
wrapper: RunContextWrapper[Any] | None,
) -> RunContextWrapper[Any] | None:
"""Return ``wrapper`` only for sessions with a complete context-aware contract."""
if wrapper is None or not _session_accepts_wrapper(session):
return None
return wrapper


async def _call_session_method(
method: Any,
/,
*args: Any,
wrapper: RunContextWrapper[Any] | None = None,
**kwargs: Any,
) -> Any:
"""Call a session method with its legacy shape unless it opts into ``wrapper``."""
if wrapper is not None and _session_method_accepts_wrapper(method):
kwargs["wrapper"] = wrapper
result = method(*args, **kwargs)
if inspect.isawaitable(result):
return await result
return result
28 changes: 25 additions & 3 deletions src/agents/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,7 @@
NextStepRunAgain,
)
from .run_internal.session_persistence import (
_session_get_items,
persist_session_items_for_guardrail_trip,
prepare_input_with_session,
reconcile_nested_history_owned_session_item_refs,
Expand Down Expand Up @@ -525,6 +526,9 @@ async def run(
previous_response_id=previous_response_id,
auto_previous_response_id=auto_previous_response_id,
)
context_wrapper = ensure_context_wrapper(context)
context = context_wrapper.context
set_agent_tool_state_scope(context_wrapper, None)

server_manages_conversation = (
conversation_id is not None
Expand All @@ -540,6 +544,7 @@ async def run(
run_config.session_settings,
include_history_in_prepared_input=False,
preserve_dropped_new_items=True,
wrapper=context_wrapper,
)
original_input_for_state = raw_input
session_input_items_for_persistence = []
Expand All @@ -552,6 +557,7 @@ async def run(
session,
run_config.session_input_callback,
run_config.session_settings,
wrapper=context_wrapper,
)
original_input_for_state = prepared_input

Expand Down Expand Up @@ -588,7 +594,10 @@ async def run(
session_input_items: list[TResponseInputItem] | None = None
if session is not None:
try:
session_input_items = await session.get_items()
session_input_items = await _session_get_items(
session,
wrapper=context_wrapper,
)
except Exception:
session_input_items = None
server_conversation_tracker.hydrate_from_state(
Expand Down Expand Up @@ -646,8 +655,6 @@ async def run(
generated_items = []
session_items = []
model_responses = []
context_wrapper = ensure_context_wrapper(context)
set_agent_tool_state_scope(context_wrapper, None)
run_state = RunState(
context=context_wrapper,
original_input=original_input,
Expand Down Expand Up @@ -782,6 +789,7 @@ def _finalize_result(result: RunResult) -> RunResult:
[],
run_state,
store=store_setting,
wrapper=context_wrapper,
)
session_input_items_for_persistence = []
except BaseException:
Expand Down Expand Up @@ -825,6 +833,7 @@ def _finalize_result(result: RunResult) -> RunResult:
original_user_input,
run_state,
store=store_setting,
wrapper=context_wrapper,
)
)
raise
Expand Down Expand Up @@ -875,6 +884,7 @@ def _finalize_result(result: RunResult) -> RunResult:
[],
run_state,
store=store_setting,
wrapper=context_wrapper,
)
session_input_items_for_persistence = []
if run_state is not None and run_state._current_step is not None:
Expand Down Expand Up @@ -944,6 +954,7 @@ def _finalize_result(result: RunResult) -> RunResult:
run_state._reasoning_item_id_policy
),
store=store_setting,
wrapper=context_wrapper,
)
)

Expand Down Expand Up @@ -1057,6 +1068,7 @@ def _finalize_result(result: RunResult) -> RunResult:
run_state,
response_id=turn_result.model_response.response_id,
store=store_setting,
wrapper=context_wrapper,
)
result._original_input = copy_input_items(original_input)
return _finalize_result(result)
Expand Down Expand Up @@ -1191,6 +1203,7 @@ def _finalize_result(result: RunResult) -> RunResult:
response_id=None,
reasoning_item_id_policy=resolved_reasoning_item_id_policy,
store=store_setting,
wrapper=context_wrapper,
)
result._original_input = copy_input_items(original_input)
return _finalize_result(result)
Expand Down Expand Up @@ -1245,6 +1258,7 @@ def _finalize_result(result: RunResult) -> RunResult:
original_user_input,
run_state,
store=store_setting,
wrapper=context_wrapper,
)
)
raise
Expand Down Expand Up @@ -1302,6 +1316,7 @@ def _finalize_result(result: RunResult) -> RunResult:
original_user_input,
run_state,
store=store_setting,
wrapper=context_wrapper,
)
)
raise
Expand Down Expand Up @@ -1431,6 +1446,7 @@ def _finalize_result(result: RunResult) -> RunResult:
run_state._reasoning_item_id_policy
),
store=store_setting,
wrapper=context_wrapper,
)
run_state._current_turn_persisted_item_count += saved_count
else:
Expand All @@ -1441,6 +1457,7 @@ def _finalize_result(result: RunResult) -> RunResult:
run_state,
response_id=turn_result.model_response.response_id,
store=store_setting,
wrapper=context_wrapper,
)

# After the first resumed turn, treat subsequent turns as fresh
Expand All @@ -1467,6 +1484,7 @@ def _finalize_result(result: RunResult) -> RunResult:
items=_retained_items_for_blocked_output(items_to_save_turn),
response_id=turn_result.model_response.response_id,
store=store_setting,
wrapper=context_wrapper,
)
raise
except (Exception, asyncio.CancelledError):
Expand All @@ -1480,6 +1498,7 @@ def _finalize_result(result: RunResult) -> RunResult:
items=items_to_save_turn,
response_id=turn_result.model_response.response_id,
store=store_setting,
wrapper=context_wrapper,
)
raise

Expand All @@ -1491,6 +1510,7 @@ def _finalize_result(result: RunResult) -> RunResult:
items=items_to_save_turn,
response_id=turn_result.model_response.response_id,
store=store_setting,
wrapper=context_wrapper,
)

# Ensure starting_input is not None and not RunState
Expand Down Expand Up @@ -1539,6 +1559,7 @@ def _finalize_result(result: RunResult) -> RunResult:
run_state,
response_id=turn_result.model_response.response_id,
store=store_setting,
wrapper=context_wrapper,
)
append_model_response_if_new(
model_responses, turn_result.model_response
Expand Down Expand Up @@ -1600,6 +1621,7 @@ def _finalize_result(result: RunResult) -> RunResult:
items=session_items_for_turn(turn_result),
response_id=turn_result.model_response.response_id,
store=store_setting,
wrapper=context_wrapper,
)
continue
else:
Expand Down
5 changes: 5 additions & 0 deletions src/agents/run_internal/agent_runner_helpers.py
Original file line number Diff line number Diff line change
Expand Up @@ -474,6 +474,7 @@ async def save_turn_items_if_needed(
items: list[RunItem],
response_id: str | None,
store: bool | None = None,
wrapper: RunContextWrapper[Any] | None = None,
) -> None:
"""Persist turn items when persistence is enabled and guardrails allow it."""
if not session_persistence_enabled:
Expand All @@ -489,6 +490,7 @@ async def save_turn_items_if_needed(
run_state,
response_id=response_id,
store=store,
wrapper=wrapper,
)


Expand All @@ -501,6 +503,7 @@ async def save_final_turn_items_after_guardrails(
items: list[RunItem],
response_id: str | None,
store: bool | None = None,
wrapper: RunContextWrapper[Any] | None = None,
) -> None:
"""Persist deferred final-turn items without skipping a partially persisted resumed turn."""
if not session_persistence_enabled or not items:
Expand All @@ -515,6 +518,7 @@ async def save_final_turn_items_after_guardrails(
response_id=response_id,
reasoning_item_id_policy=run_state._reasoning_item_id_policy,
store=store,
wrapper=wrapper,
)
return
await save_result_to_session(
Expand All @@ -524,6 +528,7 @@ async def save_final_turn_items_after_guardrails(
run_state,
response_id=response_id,
store=store,
wrapper=wrapper,
)


Expand Down
Loading