From 3530e4e0242f62e39a73c5888118c1944ac0cde4 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 16:04:59 -0700 Subject: [PATCH] refactor(state): split import_sessions into validate/insert/attach-parents helpers; shared rich SELECT builder --- hermes_state_portability.py | 665 ++++++++++++++++-------------------- 1 file changed, 297 insertions(+), 368 deletions(-) diff --git a/hermes_state_portability.py b/hermes_state_portability.py index 2f740786ff..1c42ed34cf 100644 --- a/hermes_state_portability.py +++ b/hermes_state_portability.py @@ -12,16 +12,74 @@ import time from typing import Any, Dict, List, Optional from agent.skill_commands import SKILL_SCAFFOLD_SQL_LIKE -from hermes_state_common import ( - SCHEMA_SQL, - _PREVIEW_RAW_SUBQUERY_SQL, - _shape_preview, - _sql_session_last_active, -) +from hermes_state_common import SCHEMA_SQL, _PREVIEW_RAW_SUBQUERY_SQL, _shape_preview, _sql_session_last_active # Keep the pre-split logger identity so log filtering/capture is unchanged. logger = logging.getLogger("hermes_state") +_IMPORT_SESSION_TEXT_FIELDS = ( + "source", "user_id", "model", "system_prompt", "end_reason", "cwd", + "git_branch", "git_repo_root", "billing_provider", "billing_base_url", + "billing_mode", "cost_status", "cost_source", "pricing_version", "title", +) +# ``role`` is validated separately (non-empty string). +_IMPORT_MESSAGE_TEXT_FIELDS = ( + "tool_call_id", "tool_name", "effect_disposition", "finish_reason", + "reasoning", "reasoning_content", "platform_message_id", "message_id", +) +_IMPORT_MESSAGE_JSON_FIELDS = ("reasoning_details", "codex_reasoning_items", "codex_message_items") +_IMPORT_SESSION_INSERT_SQL = """INSERT INTO sessions ( + id, source, user_id, model, model_config, system_prompt, + system_prompt_hash, + parent_session_id, started_at, ended_at, end_reason, + message_count, tool_call_count, input_tokens, output_tokens, + cache_read_tokens, cache_write_tokens, reasoning_tokens, + cwd, git_branch, git_repo_root, + billing_provider, billing_base_url, billing_mode, + estimated_cost_usd, actual_cost_usd, cost_status, cost_source, + pricing_version, title, api_call_count, archived + ) + VALUES ( + :id, :source, :user_id, :model, :model_config, + NULL, :system_prompt_hash, NULL, :started_at, :ended_at, + :end_reason, 0, 0, :input_tokens, :output_tokens, + :cache_read_tokens, :cache_write_tokens, + :reasoning_tokens, :cwd, :git_branch, :git_repo_root, + :billing_provider, :billing_base_url, :billing_mode, + :estimated_cost_usd, :actual_cost_usd, :cost_status, + :cost_source, :pricing_version, :title, + :api_call_count, :archived + )""" +# Columns copied verbatim from the payload; typed columns are converted below. +_IMPORT_PASSTHROUGH_COLS = ( + "user_id", "model", "model_config", "end_reason", "cwd", "git_branch", "git_repo_root", + "billing_provider", "billing_base_url", "billing_mode", "cost_status", "cost_source", + "pricing_version", "title", +) +_IMPORT_INT_COLS = ( + "input_tokens", "output_tokens", "cache_read_tokens", "cache_write_tokens", "reasoning_tokens", + "api_call_count", +) +_IMPORT_FLOAT_COLS = ("ended_at", "estimated_cost_usd", "actual_cost_usd") + + +def _rich_select(select_cols: str, where: str, tail: str = "", prompt_select: Optional[str] = "") -> str: + """``list_sessions_rich``-shaped SELECT: resolved prompt (``prompt_select`` + fragment; None omits prompt columns AND the join), preview, last_active. + Whitespace matches the historical inline queries (SQL text is pinned).""" + prompt_join = "" if prompt_select is None else "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash" + return f""" + SELECT {select_cols}{prompt_select or ""}, + {_PREVIEW_RAW_SUBQUERY_SQL}, + {_sql_session_last_active("s")} AS last_active + FROM sessions s + {prompt_join} + WHERE {where}{tail} + """ + + +_PROMPT_RESOLVED_SQL = "COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved" + class SessionPortabilityMixin: """See module docstring — mixin for SessionDB (Port cluster).""" @@ -33,8 +91,7 @@ class SessionPortabilityMixin: if cls._session_compact_cols_sql is None: declared = cls._parse_schema_columns(SCHEMA_SQL)["sessions"] cls._session_compact_cols_sql = ", ".join( - f"s.{name}" for name in declared - if name not in cls._SESSION_COMPACT_EXCLUDED + f"s.{name}" for name in declared if name not in cls._SESSION_COMPACT_EXCLUDED ) return cls._session_compact_cols_sql @@ -45,82 +102,60 @@ class SessionPortabilityMixin: s["preview"] = _shape_preview(s.pop("_preview_raw", "")) return s + def _locked_rows(self, sql: str, params=()) -> list: + with self._lock: + return self._conn.execute(sql, params).fetchall() + def distinct_session_cwds(self, include_archived: bool = False) -> List[Dict[str, Any]]: """Distinct non-empty session cwds with usage stats, for repo discovery. - - Aggregates across ALL history (not one page) so every repo the user - worked in surfaces. Children/branches count: a worktree session is a - real workspace signal. - """ + Aggregates across ALL history so every repo the user worked in + surfaces; children/branches count (a worktree session is a real + workspace signal).""" where = "cwd IS NOT NULL AND TRIM(cwd) != ''" if not include_archived: where += " AND archived = 0" - with self._lock: - rows = self._conn.execute( - "SELECT cwd AS cwd, COUNT(*) AS sessions, " - "MAX(COALESCE(ended_at, started_at, 0)) AS last_active " - f"FROM sessions WHERE {where} GROUP BY cwd" - ).fetchall() + rows = self._locked_rows( + "SELECT cwd AS cwd, COUNT(*) AS sessions, " + "MAX(COALESCE(ended_at, started_at, 0)) AS last_active " + f"FROM sessions WHERE {where} GROUP BY cwd" + ) return [ - { - "cwd": r["cwd"], - "sessions": int(r["sessions"] or 0), - "last_active": float(r["last_active"] or 0), - } + {"cwd": r["cwd"], "sessions": int(r["sessions"] or 0), "last_active": float(r["last_active"] or 0)} for r in rows ] - def list_cron_job_runs( - self, - job_id: str, - limit: int = 20, - offset: int = 0, - ) -> List[Dict[str, Any]]: + def list_cron_job_runs(self, job_id: str, limit: int = 20, offset: int = 0) -> List[Dict[str, Any]]: """List the run sessions produced by a single cron job, newest first. Cron runs are flat sessions with id ``cron_{job_id}_{timestamp}``; they never compress or branch, so this skips ``list_sessions_rich``'s - compression-chain CTE / leading-wildcard ``id_query`` path, which seeds - from EVERY ``source='cron'`` row and scales with the whole cron pile. - Instead: a ``[prefix, prefix_hi)`` index range scan on id, filtered to - ``source='cron'``, so work scales with the requested window. - - Returns the ``list_sessions_rich`` row shape (``preview`` + ``last_active``). + compression-chain CTE / leading-wildcard ``id_query`` path (which + seeds from EVERY ``source='cron'`` row). Instead a ``[prefix, + prefix_hi)`` index range scan on id filtered to ``source='cron'``, so + work scales with the requested window. Returns the + ``list_sessions_rich`` row shape (``preview`` + ``last_active``). """ prefix = f"cron_{job_id}_" # Half-open upper bound: bump the final byte so the range covers exactly # the ids starting with ``prefix``. prefix_hi = prefix[:-1] + chr(ord(prefix[-1]) + 1) - - query = f""" - SELECT s.*, - COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved, - {_PREVIEW_RAW_SUBQUERY_SQL}, - {_sql_session_last_active("s")} AS last_active - FROM sessions s - LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash - WHERE s.source = 'cron' AND s.id >= ? AND s.id < ? - ORDER BY s.started_at DESC, s.id DESC - LIMIT ? OFFSET ? - """ - with self._lock: - rows = self._conn.execute(query, (prefix, prefix_hi, limit, offset)).fetchall() + query = _rich_select( + "s.*", "s.source = 'cron' AND s.id >= ? AND s.id < ?", + "\n ORDER BY s.started_at DESC, s.id DESC\n LIMIT ? OFFSET ?", + prompt_select=f",\n {_PROMPT_RESOLVED_SQL}", + ) + rows = self._locked_rows(query, (prefix, prefix_hi, limit, offset)) return [self._rich_row(row) for row in rows] def _get_session_rich_row(self, session_id: str, compact_rows: bool = False) -> Optional[Dict[str, Any]]: """One session with the ``list_sessions_rich`` enriched columns, or None. ``compact_rows=True`` omits the ``system_prompt`` blob.""" - return self._get_session_rich_rows_batch( - [session_id], compact_rows=compact_rows - ).get(session_id) + return self._get_session_rich_rows_batch([session_id], compact_rows=compact_rows).get(session_id) - def _get_session_rich_rows_batch( - self, session_ids, compact_rows: bool = False - ) -> Dict[str, Dict[str, Any]]: + def _get_session_rich_rows_batch(self, session_ids, compact_rows: bool = False) -> Dict[str, Dict[str, Any]]: """Enriched rows for many sessions in one query, keyed by id; missing ids are simply absent. Resolves a page of compression tips in one - round trip instead of one query per root row. - """ + round trip instead of one query per root row.""" ids = [sid for sid in session_ids if sid] if not ids: return {} @@ -130,34 +165,16 @@ class SessionPortabilityMixin: if len(ids) > _CHUNK: result: Dict[str, Dict[str, Any]] = {} for start in range(0, len(ids), _CHUNK): - result.update( - self._get_session_rich_rows_batch( - ids[start:start + _CHUNK], compact_rows=compact_rows - ) - ) + result.update(self._get_session_rich_rows_batch(ids[start:start + _CHUNK], compact_rows=compact_rows)) return result # Same read-your-writes guarantee as list_sessions_rich. self.flush_token_counts() - _sel = self._compact_session_cols() if compact_rows else "s.*" - placeholders = ",".join("?" for _ in ids) - prompt_select = ( - "" if compact_rows - else ", COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved" + query = _rich_select( + self._compact_session_cols() if compact_rows else "s.*", + f"s.id IN ({','.join('?' for _ in ids)})", + prompt_select=None if compact_rows else f", {_PROMPT_RESOLVED_SQL}", ) - prompt_join = ( - "" if compact_rows - else "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash" - ) - query = f""" - SELECT {_sel}{prompt_select}, - {_PREVIEW_RAW_SUBQUERY_SQL}, - {_sql_session_last_active("s")} AS last_active - FROM sessions s - {prompt_join} - WHERE s.id IN ({placeholders}) - """ - with self._lock: - rows = self._conn.execute(query, ids).fetchall() + rows = self._locked_rows(query, ids) return {s["id"]: s for s in map(self._rich_row, rows)} def get_session_rich_row(self, session_id: str, compact_rows: bool = False) -> Optional[Dict[str, Any]]: @@ -166,14 +183,11 @@ class SessionPortabilityMixin: def list_skill_scaffolded_sessions(self, limit: int = 200) -> List[Dict[str, Any]]: """Titled sessions whose first user turn was a ``/skill`` invocation. - Their titles were generated from the expanded skill body, so they describe the skill, not the request. Returns ``id``, ``title`` and the - first-turn ``content`` so callers can re-derive what was typed. Newest first. - """ - with self._lock: - rows = self._conn.execute( - """ + first-turn ``content`` so callers can re-derive what was typed. Newest first.""" + rows = self._locked_rows( + """ SELECT s.id, s.title, m.content FROM sessions s JOIN messages m ON m.id = ( @@ -186,16 +200,19 @@ class SessionPortabilityMixin: ORDER BY s.started_at DESC LIMIT ? """, - (SKILL_SCAFFOLD_SQL_LIKE, int(limit)), - ).fetchall() + (SKILL_SCAFFOLD_SQL_LIKE, int(limit)), + ) return [dict(row) for row in rows] + # ── Export ───────────────────────────────────────────────────────────── + + def _with_messages(self, session: Dict[str, Any]) -> Dict[str, Any]: + return {**session, "messages": self.get_messages(session["id"])} + def export_session(self, session_id: str) -> Optional[Dict[str, Any]]: """Export a single session with all its messages as a dict.""" session = self.get_session(session_id) - if not session: - return None - return {**session, "messages": self.get_messages(session_id)} + return self._with_messages(session) if session else None def export_session_lineage(self, session_id: str) -> Optional[Dict[str, Any]]: """Export a compression lineage as one logical session dict.""" @@ -216,19 +233,13 @@ class SessionPortabilityMixin: def export_all(self, source: str = None) -> List[Dict[str, Any]]: """Export all sessions (with messages) as dicts, e.g. for JSONL backup.""" - return [ - {**session, "messages": self.get_messages(session["id"])} - for session in self.search_sessions(source=source, limit=100000) - ] + return [self._with_messages(s) for s in self.search_sessions(source=source, limit=100000)] def adopt_session_lineage_from( - self, - donor_db: Any, # a full SessionDB (mixin cannot import it — cycle) - session_id: str, - *, - retire_donor: bool = True, + self, donor_db: Any, session_id: str, *, retire_donor: bool = True, ) -> Dict[str, Any]: - """Adopt *session_id*'s full compression lineage from *donor_db*. + """Adopt *session_id*'s full compression lineage from *donor_db* (a + full SessionDB — the mixin cannot import it). Stranded-bot-session heal: before the desktop routed session RPCs by target session, a profile bot's rows accumulated in the DEFAULT @@ -251,7 +262,6 @@ class SessionPortabilityMixin: "ok": False, "adopted": False, "donor_retired": False, "error": f"session {session_id!r} not found in donor store", } - segments = payload.get("segments") or [payload] # Divergence guard: a segment we will SKIP (already here) may have kept @@ -281,57 +291,52 @@ class SessionPortabilityMixin: logger.warning( "adoption of %s did not complete: imported=%s skipped=%s " "of %s segment(s); errors=%s", - session_id, imported, skipped, len(segments), - result.get("errors"), + session_id, imported, skipped, len(segments), result.get("errors"), ) donor_retired = False if adopted and retire_donor and not donor_ahead: - retire_ok = True - for seg in segments: - seg_id = seg.get("id") - if not seg_id: - continue - try: - # TOCTOU close-out: the guard above used EXPORT-TIME counts; - # re-read both stores right before stamping so donor growth - # never lands behind a non-recoverable archive. (Equal-count - # CONTENT divergence is accepted: bytes stay in the donor - # either way, only reachability differs.) - donor_now = len(donor_db.get_messages(seg_id)) - local_now = len(self.get_messages(seg_id)) - if donor_now > local_now: - retire_ok = False - logger.warning( - "adoption divergence at retire time: donor " - "segment %s grew to %d messages (local %d) — " - "leaving donor unretired", - seg_id, donor_now, local_now, - ) - continue - # First end_reason wins in end_session(); reopen so the - # adoption boundary is stamped even on ended segments. - donor_db.reopen_session(seg_id) - donor_db.end_session(seg_id, "adopted_by_profile") - donor_db.set_session_archived(seg_id, True) - except Exception: - # Best-effort: a retirement failure must not fail the adoption - # (a later resume retries idempotently) — but never claim - # success we didn't have. - retire_ok = False - logger.warning( - "failed to retire donor segment %s after adoption", - seg_id, exc_info=True, - ) - donor_retired = retire_ok - + donor_retired = all( + self._retire_donor_segment(donor_db, seg["id"]) for seg in segments if seg.get("id") + ) return {**result, "adopted": adopted, "donor_retired": donor_retired} + def _retire_donor_segment(self, donor_db: Any, seg_id: str) -> bool: + """Archive one adopted donor segment; False when skipped or failed. + + TOCTOU close-out: the divergence guard used EXPORT-TIME counts; re-read + both stores right before stamping so donor growth never lands behind a + non-recoverable archive (equal-count CONTENT divergence is accepted — + bytes stay in the donor either way, only reachability differs). A + retirement failure must not fail the adoption (a later resume retries + idempotently), but never claims success it didn't have. + """ + try: + donor_now = len(donor_db.get_messages(seg_id)) + local_now = len(self.get_messages(seg_id)) + if donor_now > local_now: + logger.warning( + "adoption divergence at retire time: donor " + "segment %s grew to %d messages (local %d) — " + "leaving donor unretired", + seg_id, donor_now, local_now, + ) + return False + # First end_reason wins in end_session(); reopen so the adoption + # boundary is stamped even on ended segments. + donor_db.reopen_session(seg_id) + donor_db.end_session(seg_id, "adopted_by_profile") + donor_db.set_session_archived(seg_id, True) + return True + except Exception: + logger.warning("failed to retire donor segment %s after adoption", seg_id, exc_info=True) + return False + + # ── Import ───────────────────────────────────────────────────────────── + @staticmethod def _import_text_or_none(value: Any, field: str) -> Optional[str]: - if value is None: - return None - if isinstance(value, str): + if value is None or isinstance(value, str): return value raise ValueError(f"{field} must be a string") @@ -397,6 +402,146 @@ class SessionPortabilityMixin: item["session_id"] = session_id return item + def _normalize_import_session(self, raw: Dict[str, Any], session_id: str, messages: list) -> Dict[str, Any]: + """Type-check one payload session + its messages; raises ValueError.""" + clean_session = dict(raw) + clean_session["id"] = session_id + clean_session["model_config"] = self._import_json_object_or_none(clean_session.get("model_config"), "model_config") + clean_session["parent_session_id"] = self._import_text_or_none( + clean_session.get("parent_session_id"), "parent_session_id" + ) + for field in _IMPORT_SESSION_TEXT_FIELDS: + clean_session[field] = self._import_text_or_none(clean_session.get(field), field) + clean_messages: List[Dict[str, Any]] = [] + for message_index, message in enumerate(messages): + clean_message = dict(message) + role = clean_message.get("role") + if not isinstance(role, str) or not role: + raise ValueError(f"messages[{message_index}].role must be a non-empty string") + for field in _IMPORT_MESSAGE_TEXT_FIELDS: + clean_message[field] = self._import_text_or_none(clean_message.get(field), field) + clean_message["token_count"] = self._import_int_or_none(clean_message.get("token_count"), "token_count") + clean_messages.append(clean_message) + return {"session": clean_session, "messages": clean_messages} + + def _validate_import_payload(self, sessions: List[Dict[str, Any]]) -> tuple: + """Size/shape/type validation of the whole payload; returns + ``(normalized_items, errors)``. Every rejected entry is reported.""" + normalized: List[Dict[str, Any]] = [] + errors: List[Dict[str, Any]] = [] + seen_ids: set[str] = set() + total_messages = 0 + total_bytes = 0 + for index, raw in enumerate(sessions): + if not isinstance(raw, dict): + errors.append(self._import_error(index, "", "session must be an object")) + continue + session_id = str(raw.get("id") or "").strip() + if not session_id: + errors.append(self._import_error(index, "", "session id is required")) + continue + + def _reject(msg: str) -> None: + errors.append(self._import_error(index, session_id, msg)) + + if session_id in seen_ids: + _reject("duplicate session id") + continue + messages = raw.get("messages") or [] + if not isinstance(messages, list): + _reject("messages must be a list") + continue + if len(messages) > self._IMPORT_MAX_MESSAGES_PER_SESSION: + _reject("messages exceeds the per-session import limit") + continue + if any(not isinstance(msg, dict) for msg in messages): + _reject("messages must contain only objects") + continue + try: + session_bytes = len(json.dumps(raw, ensure_ascii=False, separators=(",", ":")).encode("utf-8")) + except (TypeError, ValueError): + _reject("session must be JSON serializable") + continue + if session_bytes > self._IMPORT_MAX_SESSION_BYTES: + _reject("session exceeds the import size limit") + continue + total_bytes += session_bytes + if total_bytes > self._IMPORT_MAX_TOTAL_BYTES: + _reject("import exceeds the total size limit") + continue + try: + item = self._normalize_import_session(raw, session_id, messages) + except ValueError as exc: + _reject(str(exc)) + continue + total_messages += len(item["messages"]) + if total_messages > self._IMPORT_MAX_TOTAL_MESSAGES: + _reject("messages exceeds the total import limit") + continue + seen_ids.add(session_id) + normalized.append({"index": index, **item}) + return normalized, errors + + def _import_session_row(self, conn, raw: Dict[str, Any], messages: List[Dict[str, Any]], session_id: str) -> None: + """INSERT one normalized session + its messages; counts fixed up after.""" + started_at = self._float_or_none(raw.get("started_at")) + params = { + "id": session_id, + "source": str(raw.get("source") or "import"), + "system_prompt_hash": self._store_system_prompt(conn, raw.get("system_prompt")), + "started_at": time.time() if started_at is None else started_at, + "archived": 1 if raw.get("archived") else 0, + **{col: raw.get(col) for col in _IMPORT_PASSTHROUGH_COLS}, + **{col: self._float_or_none(raw.get(col)) for col in _IMPORT_FLOAT_COLS}, + **{col: self._int_or_default(raw.get(col)) for col in _IMPORT_INT_COLS}, + } + conn.execute(_IMPORT_SESSION_INSERT_SQL, params) + sanitized_messages = [ + {**msg, **{key: self._reasoning_json_value(msg.get(key)) for key in _IMPORT_MESSAGE_JSON_FIELDS}} + for msg in messages + ] + total_messages, total_tool_calls = self._insert_message_rows(conn, session_id, sanitized_messages) + conn.execute( + "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", + (total_messages, total_tool_calls, session_id), + ) + + @staticmethod + def _attach_import_parents(conn, parent_updates: List[tuple]) -> int: + """Re-attach imported children whose parent exists (in the store or + the same payload) without creating a cycle; returns the detached count. + Only the closing edge of a cycle is dropped, so later entries can + still attach to the now-root session.""" + parent_by_child = dict(parent_updates) + + def _would_create_cycle(session_id: str, parent_id: str) -> bool: + seen = {session_id} + current = parent_id + while current: + if current in seen: + return True + seen.add(current) + if current in parent_by_child: + current = parent_by_child[current] + continue + row = conn.execute( + "SELECT parent_session_id FROM sessions WHERE id = ? LIMIT 1", (current,), + ).fetchone() + if row is None: + return False + current = row["parent_session_id"] + return False + + detached = 0 + for session_id, parent_id in parent_updates: + parent_exists = conn.execute("SELECT 1 FROM sessions WHERE id = ? LIMIT 1", (parent_id,)).fetchone() + if parent_exists and not _would_create_cycle(session_id, parent_id): + conn.execute("UPDATE sessions SET parent_session_id = ? WHERE id = ?", (parent_id, session_id)) + else: + parent_by_child.pop(session_id, None) + detached += 1 + return detached + def import_sessions(self, sessions: List[Dict[str, Any]]) -> Dict[str, Any]: """Import sessions exported by :meth:`export_session` or ``export_all``. @@ -414,102 +559,8 @@ class SessionPortabilityMixin: if not isinstance(sessions, list): raise ValueError("sessions must be a list") if len(sessions) > self._IMPORT_MAX_SESSIONS: - raise ValueError( - f"sessions must contain at most {self._IMPORT_MAX_SESSIONS} entries" - ) - - normalized: List[Dict[str, Any]] = [] - errors: List[Dict[str, Any]] = [] - seen_ids: set[str] = set() - total_messages = 0 - total_bytes = 0 - session_text_fields = ( - "source", "user_id", "model", "system_prompt", "end_reason", "cwd", - "git_branch", "git_repo_root", "billing_provider", "billing_base_url", - "billing_mode", "cost_status", "cost_source", "pricing_version", "title", - ) - # ``role`` is validated separately below (non-empty string). - message_text_fields = ( - "tool_call_id", "tool_name", "effect_disposition", "finish_reason", - "reasoning", "reasoning_content", "platform_message_id", "message_id", - ) - - for index, raw in enumerate(sessions): - if not isinstance(raw, dict): - errors.append(self._import_error(index, "", "session must be an object")) - continue - session_id = str(raw.get("id") or "").strip() - if not session_id: - errors.append(self._import_error(index, "", "session id is required")) - continue - if session_id in seen_ids: - errors.append(self._import_error(index, session_id, "duplicate session id")) - continue - messages = raw.get("messages") or [] - if not isinstance(messages, list): - errors.append(self._import_error(index, session_id, "messages must be a list")) - continue - if len(messages) > self._IMPORT_MAX_MESSAGES_PER_SESSION: - errors.append(self._import_error(index, session_id, "messages exceeds the per-session import limit")) - continue - if any(not isinstance(msg, dict) for msg in messages): - errors.append(self._import_error(index, session_id, "messages must contain only objects")) - continue - - try: - session_bytes = len(json.dumps(raw, ensure_ascii=False, separators=(",", ":")).encode("utf-8")) - except (TypeError, ValueError): - errors.append(self._import_error(index, session_id, "session must be JSON serializable")) - continue - if session_bytes > self._IMPORT_MAX_SESSION_BYTES: - errors.append(self._import_error(index, session_id, "session exceeds the import size limit")) - continue - total_bytes += session_bytes - if total_bytes > self._IMPORT_MAX_TOTAL_BYTES: - errors.append(self._import_error(index, session_id, "import exceeds the total size limit")) - continue - - try: - clean_session = dict(raw) - clean_session["id"] = session_id - clean_session["model_config"] = self._import_json_object_or_none( - clean_session.get("model_config"), "model_config" - ) - clean_session["parent_session_id"] = self._import_text_or_none( - clean_session.get("parent_session_id"), "parent_session_id" - ) - for field in session_text_fields: - clean_session[field] = self._import_text_or_none( - clean_session.get(field), field - ) - - clean_messages: List[Dict[str, Any]] = [] - for message_index, message in enumerate(messages): - clean_message = dict(message) - role = clean_message.get("role") - if not isinstance(role, str) or not role: - raise ValueError(f"messages[{message_index}].role must be a non-empty string") - for field in message_text_fields: - clean_message[field] = self._import_text_or_none( - clean_message.get(field), field - ) - clean_message["token_count"] = self._import_int_or_none( - clean_message.get("token_count"), "token_count" - ) - clean_messages.append(clean_message) - except ValueError as exc: - errors.append(self._import_error(index, session_id, str(exc))) - continue - - total_messages += len(clean_messages) - if total_messages > self._IMPORT_MAX_TOTAL_MESSAGES: - errors.append(self._import_error(index, session_id, "messages exceeds the total import limit")) - continue - seen_ids.add(session_id) - normalized.append( - {"index": index, "session": clean_session, "messages": clean_messages} - ) - + raise ValueError(f"sessions must contain at most {self._IMPORT_MAX_SESSIONS} entries") + normalized, errors = self._validate_import_payload(sessions) if errors: return {"ok": False, "imported": 0, "skipped": 0, "detached": 0, "errors": errors} @@ -517,140 +568,18 @@ class SessionPortabilityMixin: imported_ids: List[str] = [] skipped_ids: List[str] = [] parent_updates: List[tuple[str, str]] = [] - detached = 0 - for item in normalized: raw = item["session"] - messages = item["messages"] session_id = str(raw.get("id") or "").strip() - exists = conn.execute( - "SELECT 1 FROM sessions WHERE id = ? LIMIT 1", - (session_id,), - ).fetchone() - if exists: + if conn.execute("SELECT 1 FROM sessions WHERE id = ? LIMIT 1", (session_id,)).fetchone(): skipped_ids.append(session_id) continue - - started_at = self._float_or_none(raw.get("started_at")) - if started_at is None: - started_at = time.time() - archived = 1 if raw.get("archived") else 0 - system_prompt_hash = self._store_system_prompt( - conn, raw.get("system_prompt") - ) - - conn.execute( - """INSERT INTO sessions ( - id, source, user_id, model, model_config, system_prompt, - system_prompt_hash, - parent_session_id, started_at, ended_at, end_reason, - message_count, tool_call_count, input_tokens, output_tokens, - cache_read_tokens, cache_write_tokens, reasoning_tokens, - cwd, git_branch, git_repo_root, - billing_provider, billing_base_url, billing_mode, - estimated_cost_usd, actual_cost_usd, cost_status, cost_source, - pricing_version, title, api_call_count, archived - ) - VALUES ( - :id, :source, :user_id, :model, :model_config, - NULL, :system_prompt_hash, NULL, :started_at, :ended_at, - :end_reason, 0, 0, :input_tokens, :output_tokens, - :cache_read_tokens, :cache_write_tokens, - :reasoning_tokens, :cwd, :git_branch, :git_repo_root, - :billing_provider, :billing_base_url, :billing_mode, - :estimated_cost_usd, :actual_cost_usd, :cost_status, - :cost_source, :pricing_version, :title, - :api_call_count, :archived - )""", - { - "id": session_id, - "source": str(raw.get("source") or "import"), - "user_id": raw.get("user_id"), - "model": raw.get("model"), - "model_config": raw.get("model_config"), - "system_prompt_hash": system_prompt_hash, - "started_at": started_at, - "ended_at": self._float_or_none(raw.get("ended_at")), - "end_reason": raw.get("end_reason"), - "input_tokens": self._int_or_default(raw.get("input_tokens")), - "output_tokens": self._int_or_default(raw.get("output_tokens")), - "cache_read_tokens": self._int_or_default(raw.get("cache_read_tokens")), - "cache_write_tokens": self._int_or_default(raw.get("cache_write_tokens")), - "reasoning_tokens": self._int_or_default(raw.get("reasoning_tokens")), - "cwd": raw.get("cwd"), - "git_branch": raw.get("git_branch"), - "git_repo_root": raw.get("git_repo_root"), - "billing_provider": raw.get("billing_provider"), - "billing_base_url": raw.get("billing_base_url"), - "billing_mode": raw.get("billing_mode"), - "estimated_cost_usd": self._float_or_none(raw.get("estimated_cost_usd")), - "actual_cost_usd": self._float_or_none(raw.get("actual_cost_usd")), - "cost_status": raw.get("cost_status"), - "cost_source": raw.get("cost_source"), - "pricing_version": raw.get("pricing_version"), - "title": raw.get("title"), - "api_call_count": self._int_or_default(raw.get("api_call_count")), - "archived": archived, - }, - ) - - sanitized_messages: List[Dict[str, Any]] = [] - for msg in messages: - clean = dict(msg) - for key in ("reasoning_details", "codex_reasoning_items", "codex_message_items"): - clean[key] = self._reasoning_json_value(clean.get(key)) - sanitized_messages.append(clean) - - total_messages, total_tool_calls = self._insert_message_rows( - conn, session_id, sanitized_messages - ) - conn.execute( - "UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?", - (total_messages, total_tool_calls, session_id), - ) - + self._import_session_row(conn, raw, item["messages"], session_id) parent_id = str(raw.get("parent_session_id") or "").strip() if parent_id: parent_updates.append((session_id, parent_id)) imported_ids.append(session_id) - - parent_by_child = dict(parent_updates) - - def _would_create_cycle(session_id: str, parent_id: str) -> bool: - seen = {session_id} - current = parent_id - while current: - if current in seen: - return True - seen.add(current) - if current in parent_by_child: - current = parent_by_child[current] - continue - row = conn.execute( - "SELECT parent_session_id FROM sessions WHERE id = ? LIMIT 1", - (current,), - ).fetchone() - if row is None: - return False - current = row["parent_session_id"] - return False - - for session_id, parent_id in parent_updates: - parent_exists = conn.execute( - "SELECT 1 FROM sessions WHERE id = ? LIMIT 1", - (parent_id,), - ).fetchone() - if parent_exists and not _would_create_cycle(session_id, parent_id): - conn.execute( - "UPDATE sessions SET parent_session_id = ? WHERE id = ?", - (parent_id, session_id), - ) - else: - # Drop only the closing edge; later entries can still attach - # to this now-root session. - parent_by_child.pop(session_id, None) - detached += 1 - + detached = self._attach_import_parents(conn, parent_updates) return { "ok": True, "imported": len(imported_ids),