class EncryptedSession(SessionABC):
"""Encrypted wrapper for Session implementations with TTL-based expiration.
This class wraps any SessionABC implementation to provide transparent
encryption/decryption of stored items using Fernet encryption with
per-session key derivation and automatic expiration of old data.
When items expire (exceed TTL), they are silently skipped during retrieval.
Note: Expired tokens are rejected based on the system clock of the application server.
To avoid valid tokens being rejected due to clock drift, ensure all servers in
your environment are synchronized using NTP.
"""
def __init__(
self,
session_id: str,
underlying_session: SessionABC,
encryption_key: str,
ttl: int = 600,
):
"""
Args:
session_id: ID for this session
underlying_session: The real session store (e.g. SQLiteSession, SQLAlchemySession)
encryption_key: Master key (Fernet key or raw secret)
ttl: Token time-to-live in seconds (default 10 min)
"""
self.session_id = session_id
self.underlying_session = underlying_session
self.ttl = ttl
master = _ensure_fernet_key_bytes(encryption_key)
self.cipher = _derive_session_fernet_key(master, session_id)
self._kid = "hkdf-v1"
self._ver = 1
def __getattr__(self, name: str) -> Any:
# Expose compaction only when the underlying session actually supports it.
if isinstance(self.underlying_session, OpenAIResponsesCompactionSession):
if name == "run_compaction":
return self._run_compaction
if name == "_defer_compaction":
return self._defer_encrypted_compaction
return getattr(self.underlying_session, name)
async def _run_compaction(
self,
args: OpenAIResponsesCompactionArgs | None = None,
*,
wrapper: RunContextWrapper[Any] | None = None,
) -> None:
session = cast(OpenAIResponsesCompactionSession, self.underlying_session)
await session._run_compaction(
args,
wrapper=wrapper,
read_items=lambda: _call_session_method(
self.get_items, wrapper=_get_session_wrapper(self, wrapper)
),
prepare_items=self._encrypt_items,
)
async def _defer_encrypted_compaction(
self,
response_id: str,
store: bool | None = None,
*,
wrapper: RunContextWrapper[Any] | None = None,
) -> None:
session = cast(OpenAIResponsesCompactionSession, self.underlying_session)
await session._defer_compaction(
response_id,
store,
read_items=lambda: _call_session_method(
self.get_items, wrapper=_get_session_wrapper(self, wrapper)
),
)
@property
def session_settings(self) -> SessionSettings | None:
"""Get session settings from the underlying session."""
return self.underlying_session.session_settings
@session_settings.setter
def session_settings(self, value: SessionSettings | None) -> None:
"""Set session settings on the underlying session."""
self.underlying_session.session_settings = value
def _wrap(self, item: TResponseInputItem) -> EncryptedEnvelope:
if isinstance(item, dict):
payload = item
elif hasattr(item, "model_dump"):
payload = item.model_dump()
elif hasattr(item, "__dict__"):
payload = item.__dict__
else:
payload = dict(item)
token = self.cipher.encrypt(_to_json_bytes(payload)).decode("utf-8")
return {"__enc__": 1, "v": self._ver, "kid": self._kid, "payload": token}
def _unwrap(self, item: TResponseInputItem | EncryptedEnvelope) -> TResponseInputItem | None:
if not _is_encrypted_envelope(item):
return cast(TResponseInputItem, item)
try:
token = item["payload"].encode("utf-8")
plaintext = self.cipher.decrypt(token, ttl=self.ttl)
return cast(TResponseInputItem, _from_json_bytes(plaintext))
except (InvalidToken, KeyError):
return None
def _unwrap_valid_items(
self, encrypted_items: list[TResponseInputItem]
) -> list[TResponseInputItem]:
valid_items: list[TResponseInputItem] = []
for enc in encrypted_items:
item = self._unwrap(enc)
if item is not None:
valid_items.append(item)
return valid_items
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 = 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 = cast(
list[TResponseInputItem],
await _call_session_method(
self.underlying_session.get_items,
limit,
wrapper=wrapper,
),
)
return self._unwrap_valid_items(encrypted_items)
def _encrypt_items(self, items: list[TResponseInputItem]) -> list[TResponseInputItem]:
return cast(list[TResponseInputItem], [self._wrap(item) for item in items])
async def add_items(
self,
items: list[TResponseInputItem],
*,
wrapper: RunContextWrapper[Any] | None = None,
) -> None:
wrapper = _get_session_wrapper(self.underlying_session, wrapper)
await _call_session_method(
self.underlying_session.add_items,
self._encrypt_items(items),
wrapper=wrapper,
)
async def pop_item(
self,
*,
wrapper: RunContextWrapper[Any] | None = None,
) -> TResponseInputItem | None:
wrapper = _get_session_wrapper(self.underlying_session, wrapper)
while True:
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,
*,
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,
)