From 75fdd85316b0a3f79cf32a3100579e62b4c42d09 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 19:36:48 -0700 Subject: [PATCH] =?UTF-8?q?refactor(gateway):=20session=20=E2=80=94=20pack?= =?UTF-8?q?=20exploded=20argument=20lists?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- gateway/session.py | 115 ++++++++++----------------------- gateway/session_persistence.py | 5 +- gateway/session_recovery.py | 71 +++++--------------- gateway/session_transcript.py | 17 ++--- 4 files changed, 55 insertions(+), 153 deletions(-) diff --git a/gateway/session.py b/gateway/session.py index 0aedf130da..52f3098a93 100644 --- a/gateway/session.py +++ b/gateway/session.py @@ -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 diff --git a/gateway/session_persistence.py b/gateway/session_persistence.py index ed2c30d667..28ee78b4f5 100644 --- a/gateway/session_persistence.py +++ b/gateway/session_persistence.py @@ -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 diff --git a/gateway/session_recovery.py b/gateway/session_recovery.py index f39183a564..4afcd28899 100644 --- a/gateway/session_recovery.py +++ b/gateway/session_recovery.py @@ -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``. diff --git a/gateway/session_transcript.py b/gateway/session_transcript.py index ff67af59c0..c1ce28e11c 100644 --- a/gateway/session_transcript.py +++ b/gateway/session_transcript.py @@ -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"), )