refactor(gateway): session — pack exploded argument lists

This commit is contained in:
Teknium
2026-09-02 19:36:48 -07:00
parent 4f943a1cdd
commit 75fdd85316
4 changed files with 55 additions and 153 deletions

View File

@@ -185,12 +185,10 @@ class SessionSource:
if name != "chat_type"
}
return cls(
platform=Platform(data["platform"]),
chat_id=str(data["chat_id"]),
platform=Platform(data["platform"]), chat_id=str(data["chat_id"]),
chat_type=data.get("chat_type", "dm"),
scope_id=data.get("scope_id", data.get("guild_id")),
auto_thread_created=bool(data.get("auto_thread_created", False)),
**plain,
auto_thread_created=bool(data.get("auto_thread_created", False)), **plain,
)
@@ -654,20 +652,14 @@ class SessionEntry:
plain = {name: data.get(name, defaults[name]) for name in cls._PLAIN_FIELDS + cls._RESET_FIELDS}
plain["expiry_finalized"] = data.get("expiry_finalized", data.get("memory_flushed", False))
return cls(
session_key=session_key,
session_id=session_id,
session_key=session_key, session_id=session_id,
created_at=datetime.fromisoformat(data["created_at"]),
updated_at=datetime.fromisoformat(data["updated_at"]),
origin=origin,
display_name=data.get("display_name"),
platform=platform,
chat_type=data.get("chat_type", "dm"),
metadata=dict(data.get("metadata") or {}),
updated_at=datetime.fromisoformat(data["updated_at"]), origin=origin,
display_name=data.get("display_name"), platform=platform,
chat_type=data.get("chat_type", "dm"), metadata=dict(data.get("metadata") or {}),
last_resume_marked_at=_parse_iso(data.get("last_resume_marked_at")),
active_turn_token=active_turn_token,
active_turn_started_at=active_turn_started_at,
model_override=sanitize_model_override(data.get("model_override")),
**plain,
active_turn_token=active_turn_token, active_turn_started_at=active_turn_started_at,
model_override=sanitize_model_override(data.get("model_override")), **plain,
)
@@ -697,9 +689,7 @@ def build_channel_continuity_note(entry: "SessionEntry", source: SessionSource)
def is_shared_multi_user_session(
source: SessionSource,
*,
group_sessions_per_user: bool = True,
source: SessionSource, *, group_sessions_per_user: bool = True,
thread_sessions_per_user: bool = False,
) -> bool:
"""True when a non-DM session is shared across participants (mirrors the
@@ -733,10 +723,8 @@ def _canonical_participant(source: SessionSource) -> Optional[str]:
def build_session_key(
source: SessionSource,
group_sessions_per_user: bool = True,
thread_sessions_per_user: bool = False,
profile: Optional[str] = None,
source: SessionSource, group_sessions_per_user: bool = True,
thread_sessions_per_user: bool = False, profile: Optional[str] = None,
) -> str:
"""Build a deterministic session key from a message source (single source of truth).
@@ -1048,12 +1036,9 @@ class SessionStore(
self._save_entries()
self._finish_route_transition(
session_key,
end_session_id=decision.prev_session_id,
end_reason=decision.reset_reason or "session_reset",
create_kwargs=create_kwargs,
origin=source,
display_name=decision.entry.display_name,
session_key, end_session_id=decision.prev_session_id,
end_reason=decision.reset_reason or "session_reset", create_kwargs=create_kwargs,
origin=source, display_name=decision.entry.display_name,
)
return decision.entry
@@ -1065,12 +1050,8 @@ class SessionStore(
return _RouteChecks(sid, canonical, is_stale, self._route_reset_reason(entry, source, now))
def _apply_route_checks(
self,
session_key: str,
checks: Optional[_RouteChecks],
force_new: bool,
touch_activity: bool,
now: datetime,
self, session_key: str, checks: Optional[_RouteChecks], force_new: bool,
touch_activity: bool, now: datetime,
) -> _RouteDecision:
"""Apply stale/reset decisions to ``_entries`` under ``_lock``.
@@ -1144,30 +1125,18 @@ class SessionStore(
decision.needs_save = True
def _route_create(
self,
decision: _RouteDecision,
session_key: str,
source: SessionSource,
now: datetime,
force_new: bool,
observed: Optional[SessionEntry],
self, decision: _RouteDecision, session_key: str, source: SessionSource, now: datetime,
force_new: bool, observed: Optional[SessionEntry],
) -> Optional[Dict[str, Any]]:
"""Create a candidate outside the lock, publish it only if another worker
has not already populated this routing key; returns ``create_session``
kwargs when the candidate won."""
session_id = _new_session_id(now)
candidate = SessionEntry(
session_key=session_key,
session_id=session_id,
created_at=now,
updated_at=now,
origin=source,
display_name=source.chat_name,
platform=source.platform,
chat_type=source.chat_type,
was_auto_reset=decision.reset_reason is not None,
auto_reset_reason=decision.reset_reason,
reset_had_activity=decision.reset_had_activity,
session_key=session_key, session_id=session_id, created_at=now, updated_at=now,
origin=source, display_name=source.chat_name, platform=source.platform,
chat_type=source.chat_type, was_auto_reset=decision.reset_reason is not None,
auto_reset_reason=decision.reset_reason, reset_had_activity=decision.reset_had_activity,
prev_session_id=decision.prev_session_id,
)
with self._lock:
@@ -1179,11 +1148,8 @@ class SessionStore(
if current is not candidate:
return None
return self._session_create_kwargs(
session_id=session_id,
session_key=session_key,
origin=source,
source_value=source.platform.value,
display_name=source.chat_name,
session_id=session_id, session_key=session_key, origin=source,
source_value=source.platform.value, display_name=source.chat_name,
parent_session_id=decision.prev_session_id,
)
@@ -1261,34 +1227,22 @@ class SessionStore(
is_fresh_reset=True,
)
db_create_kwargs = self._session_create_kwargs(
session_id=session_id,
session_key=session_key,
origin=old_entry.origin,
session_id=session_id, session_key=session_key, origin=old_entry.origin,
source_value=old_entry.platform.value if old_entry.platform else "unknown",
display_name=old_entry.display_name,
parent_session_id=old_entry.session_id,
display_name=old_entry.display_name, parent_session_id=old_entry.session_id,
)
self._finish_route_transition(
session_key,
end_session_id=old_entry.session_id,
end_reason="session_reset",
create_kwargs=db_create_kwargs,
origin=old_entry.origin,
display_name=new_entry.display_name,
during=" during reset",
session_key, end_session_id=old_entry.session_id, end_reason="session_reset",
create_kwargs=db_create_kwargs, origin=old_entry.origin,
display_name=new_entry.display_name, during=" during reset",
)
return new_entry
def _replace_route_locked(self, session_key, old_entry, session_id, now, **fields) -> SessionEntry:
"""Publish a fresh entry (inheriting origin/platform/chat_type) and save. Lock held."""
new_entry = SessionEntry(
session_key=session_key,
session_id=session_id,
created_at=now,
updated_at=now,
origin=old_entry.origin,
platform=old_entry.platform,
chat_type=old_entry.chat_type,
session_key=session_key, session_id=session_id, created_at=now, updated_at=now,
origin=old_entry.origin, platform=old_entry.platform, chat_type=old_entry.chat_type,
**fields,
)
self._entries[session_key] = new_entry
@@ -1317,11 +1271,8 @@ class SessionStore(
if self._db_for_key(session_key):
self._reopen_session_row(session_key, target_session_id, log_prefix="Session DB reopen_session failed")
self._record_gateway_session_peer(
target_session_id,
session_key,
new_entry.origin,
display_name=new_entry.display_name,
include_compression_ancestors=True,
target_session_id, session_key, new_entry.origin,
display_name=new_entry.display_name, include_compression_ancestors=True,
)
return new_entry

View File

@@ -554,10 +554,7 @@ class SessionPersistenceMixin:
self._persist_routing_data(data, generation)
def _save_entry(
self,
session_key: str,
*,
entry_data: Optional[Dict[str, Any]] = None,
self, session_key: str, *, entry_data: Optional[Dict[str, Any]] = None,
lock_held: bool = False,
) -> None:
"""Persist ONE routing entry via UPSERT — the per-turn fast path

View File

@@ -162,23 +162,14 @@ class SessionRecoveryMixin:
if had_activity is None:
had_activity = bool(row.get("message_count") or 0) or last_activity is not None
return SessionEntry(
session_key=session_key,
session_id=str(row["id"]),
created_at=created_at,
updated_at=updated_at,
origin=source,
display_name=source.chat_name,
platform=source.platform,
chat_type=source.chat_type,
session_key=session_key, session_id=str(row["id"]), created_at=created_at,
updated_at=updated_at, origin=source, display_name=source.chat_name,
platform=source.platform, chat_type=source.chat_type,
reset_had_activity=bool(had_activity),
)
def _find_gateway_session_row(
self,
*,
session_key: str,
source: SessionSource,
allow_peer_fallback: bool,
self, *, session_key: str, source: SessionSource, allow_peer_fallback: bool,
raise_on_lookup_error: bool = False,
) -> Optional[Dict[str, Any]]:
"""Query one durable gateway session row.
@@ -194,9 +185,7 @@ class SessionRecoveryMixin:
return None
try:
return finder(
source=source.platform.value,
user_id=source.user_id,
session_key=session_key,
source=source.platform.value, user_id=source.user_id, session_key=session_key,
chat_id=source.chat_id if allow_peer_fallback else None,
chat_type=source.chat_type if allow_peer_fallback else None,
thread_id=source.thread_id,
@@ -208,11 +197,7 @@ class SessionRecoveryMixin:
return None
def _recover_session_from_db(
self,
*,
session_key: str,
source: SessionSource,
now: datetime,
self, *, session_key: str, source: SessionSource, now: datetime,
raise_on_lookup_error: bool = False,
) -> Optional[SessionEntry]:
"""Rebuild a missing session-key mapping from durable state.db data.
@@ -222,9 +207,7 @@ class SessionRecoveryMixin:
durably promoted to a reset boundary instead of resurrected.
"""
entry, migrated_legacy = self._query_recoverable_row(
session_key=session_key,
source=source,
now=now,
session_key=session_key, source=source, now=now,
raise_on_lookup_error=raise_on_lookup_error,
)
if entry is None:
@@ -274,17 +257,13 @@ class SessionRecoveryMixin:
"""
legacy_key = self._legacy_slack_session_key(source)
recovered = self._find_gateway_session_row(
session_key=session_key,
source=source,
allow_peer_fallback=legacy_key is None,
session_key=session_key, source=source, allow_peer_fallback=legacy_key is None,
raise_on_lookup_error=raise_on_lookup_error,
)
migrated_legacy = False
if not recovered and legacy_key and self._claim_legacy_slack_key(legacy_key):
recovered = self._find_gateway_session_row(
session_key=legacy_key,
source=source,
allow_peer_fallback=False,
session_key=legacy_key, source=source, allow_peer_fallback=False,
raise_on_lookup_error=raise_on_lookup_error,
)
migrated_legacy = bool(recovered)
@@ -337,12 +316,8 @@ class SessionRecoveryMixin:
logger.debug("Gateway session DB reopen failed for %s: %s", session_key, exc)
def _record_gateway_session_peer(
self,
session_id: str,
session_key: str,
source: Optional[SessionSource],
display_name: Optional[str] = None,
include_compression_ancestors: bool = False,
self, session_id: str, session_key: str, source: Optional[SessionSource],
display_name: Optional[str] = None, include_compression_ancestors: bool = False,
) -> None:
"""Persist the routing peer for an existing gateway session row."""
db = self._db_for_key(session_key)
@@ -352,18 +327,12 @@ class SessionRecoveryMixin:
if not callable(recorder):
return
peer = dict(
source=source.platform.value,
user_id=source.user_id,
session_key=session_key,
chat_id=source.chat_id,
chat_type=source.chat_type,
thread_id=source.thread_id,
source=source.platform.value, user_id=source.user_id, session_key=session_key,
chat_id=source.chat_id, chat_type=source.chat_type, thread_id=source.thread_id,
)
try:
recorder(
session_id,
**peer,
display_name=display_name or source.chat_name,
session_id, **peer, display_name=display_name or source.chat_name,
origin_json=_origin_json(source),
include_compression_ancestors=include_compression_ancestors,
)
@@ -412,15 +381,9 @@ class SessionRecoveryMixin:
)
def _finish_route_transition(
self,
session_key: str,
*,
end_session_id: Optional[str],
end_reason: str,
create_kwargs: Optional[Dict[str, Any]],
origin: Optional[SessionSource],
display_name: Optional[str],
during: str = "",
self, session_key: str, *, end_session_id: Optional[str], end_reason: str,
create_kwargs: Optional[Dict[str, Any]], origin: Optional[SessionSource],
display_name: Optional[str], during: str = "",
) -> None:
"""SQLite side of a routing transition, outside ``_lock``.

View File

@@ -66,9 +66,7 @@ class SessionTranscriptMixin:
return session_id
def _heal_compression_tip_locked(
self,
entry: "SessionEntry",
original_session_id: Optional[str],
self, entry: "SessionEntry", original_session_id: Optional[str],
canonical_session_id: Optional[str],
) -> bool:
"""Rewrite *entry* to the compression continuation if stale. Lock held."""
@@ -466,10 +464,7 @@ class SessionTranscriptMixin:
return False
def rewrite_transcript(
self,
session_id: str,
messages: List[Dict[str, Any]],
active_only: bool = False,
self, session_id: str, messages: List[Dict[str, Any]], active_only: bool = False,
reject_active_turn_lease: bool = False,
) -> bool:
"""Replace a session's transcript (/retry, /compress).
@@ -487,9 +482,7 @@ class SessionTranscriptMixin:
with self._get_transcript_drain_lock():
try:
db.replace_messages(
session_id,
messages,
active_only=active_only,
session_id, messages, active_only=active_only,
reject_active_turn_lease=reject_active_turn_lease,
)
except Exception as e:
@@ -580,9 +573,7 @@ class SessionTranscriptMixin:
target_text = retryable_user_text(target_view.get("content"))
try:
result = db.rewind_to_message(
session_id,
target_id,
preserve_compaction_handoff=handoff is not None,
session_id, target_id, preserve_compaction_handoff=handoff is not None,
expected_active_ids=expected_active_ids,
expected_target_content=target_view.get("content"),
)