fix(state): deduplicate session system prompts
This commit is contained in:
@@ -31,6 +31,7 @@ from hermes_state import (
|
||||
ProgressCallback = Callable[[dict[str, Any]], None]
|
||||
|
||||
_CANONICAL_TABLES = (
|
||||
"system_prompts",
|
||||
"sessions",
|
||||
"messages",
|
||||
"session_model_usage",
|
||||
@@ -912,6 +913,8 @@ def _cleanup_partial_orphans(
|
||||
"""
|
||||
|
||||
result: dict[str, Any] = {
|
||||
"session_prompt_refs_cleared": 0,
|
||||
"system_prompts_removed": 0,
|
||||
"sessions_parent_cleared": 0,
|
||||
"sessions_reconstructed": 0,
|
||||
"messages_retained": 0,
|
||||
@@ -947,6 +950,42 @@ def _cleanup_partial_orphans(
|
||||
)
|
||||
result["sessions_parent_cleared"] = parent_count
|
||||
|
||||
prompt_ref_count = int(
|
||||
destination.execute(
|
||||
"SELECT COUNT(*) FROM sessions "
|
||||
"WHERE system_prompt_hash IS NOT NULL "
|
||||
"AND NOT EXISTS ("
|
||||
"SELECT 1 FROM system_prompts "
|
||||
"WHERE system_prompts.hash = sessions.system_prompt_hash)"
|
||||
).fetchone()[0]
|
||||
)
|
||||
if prompt_ref_count:
|
||||
destination.execute(
|
||||
"UPDATE sessions SET system_prompt_hash = NULL "
|
||||
"WHERE system_prompt_hash IS NOT NULL "
|
||||
"AND NOT EXISTS ("
|
||||
"SELECT 1 FROM system_prompts "
|
||||
"WHERE system_prompts.hash = sessions.system_prompt_hash)"
|
||||
)
|
||||
result["session_prompt_refs_cleared"] = prompt_ref_count
|
||||
|
||||
unreferenced_prompt_count = int(
|
||||
destination.execute(
|
||||
"SELECT COUNT(*) FROM system_prompts "
|
||||
"WHERE NOT EXISTS ("
|
||||
"SELECT 1 FROM sessions "
|
||||
"WHERE sessions.system_prompt_hash = system_prompts.hash)"
|
||||
).fetchone()[0]
|
||||
)
|
||||
if unreferenced_prompt_count:
|
||||
destination.execute(
|
||||
"DELETE FROM system_prompts "
|
||||
"WHERE NOT EXISTS ("
|
||||
"SELECT 1 FROM sessions "
|
||||
"WHERE sessions.system_prompt_hash = system_prompts.hash)"
|
||||
)
|
||||
result["system_prompts_removed"] = unreferenced_prompt_count
|
||||
|
||||
dependent_tables = (
|
||||
("messages", "messages_removed"),
|
||||
("session_model_usage", "session_model_usage_removed"),
|
||||
@@ -983,7 +1022,8 @@ def _cleanup_partial_orphans(
|
||||
# reconstruction counters describe data RETAINED, so summing them here
|
||||
# would report saving the user's messages as if it were losing them.
|
||||
result["total_removed_or_relinked"] = (
|
||||
int(result["sessions_parent_cleared"])
|
||||
int(result["session_prompt_refs_cleared"])
|
||||
+ int(result["sessions_parent_cleared"])
|
||||
+ int(result["messages_removed"])
|
||||
+ int(result["session_model_usage_removed"])
|
||||
+ int(result["compression_locks_removed"])
|
||||
|
||||
236
hermes_state.py
236
hermes_state.py
@@ -17,6 +17,7 @@ Key design decisions:
|
||||
import asyncio
|
||||
import atexit
|
||||
import errno
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
@@ -85,6 +86,10 @@ logger = logging.getLogger(__name__)
|
||||
_COMPRESSION_LOCK_HOLDER_PID_RE = re.compile(r"(?:^|:)pid=(\d+)(?::|$)")
|
||||
|
||||
|
||||
def _system_prompt_hash(system_prompt: str) -> str:
|
||||
return hashlib.sha256(system_prompt.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _compression_lock_holder_process_is_dead(holder: str) -> bool:
|
||||
"""Return True only when a structured lock holder's local PID is gone.
|
||||
|
||||
@@ -1868,6 +1873,36 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
_IMPORT_MAX_SESSION_BYTES = 5 * 1024 * 1024
|
||||
_IMPORT_MAX_TOTAL_BYTES = 25 * 1024 * 1024
|
||||
|
||||
@staticmethod
|
||||
def _store_system_prompt(conn, system_prompt: Optional[str]) -> Optional[str]:
|
||||
if system_prompt is None:
|
||||
return None
|
||||
prompt_hash = _system_prompt_hash(system_prompt)
|
||||
conn.execute(
|
||||
"INSERT OR IGNORE INTO system_prompts (hash, prompt) VALUES (?, ?)",
|
||||
(prompt_hash, system_prompt),
|
||||
)
|
||||
return prompt_hash
|
||||
|
||||
@staticmethod
|
||||
def _delete_unreferenced_system_prompts(conn) -> None:
|
||||
conn.execute(
|
||||
"DELETE FROM system_prompts "
|
||||
"WHERE NOT EXISTS ("
|
||||
"SELECT 1 FROM sessions "
|
||||
"WHERE sessions.system_prompt_hash = system_prompts.hash"
|
||||
")"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _session_row_dict(row: sqlite3.Row) -> Dict[str, Any]:
|
||||
data = dict(row)
|
||||
if "_system_prompt_resolved" in data:
|
||||
resolved = data.pop("_system_prompt_resolved")
|
||||
if "system_prompt" in data:
|
||||
data["system_prompt"] = resolved
|
||||
return data
|
||||
|
||||
def __init__(self, db_path: Path = None, read_only: bool = False):
|
||||
self.db_path = db_path or _default_db_path()
|
||||
self.read_only = read_only
|
||||
@@ -2847,17 +2882,27 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
without a recoverable routing mapping (#59527).
|
||||
"""
|
||||
def _do(conn):
|
||||
system_prompt_hash = self._store_system_prompt(conn, system_prompt)
|
||||
conn.execute(
|
||||
"""INSERT INTO sessions (
|
||||
id, source, user_id, session_key, chat_id, chat_type, thread_id,
|
||||
model, model_config, system_prompt, parent_session_id, cwd,
|
||||
profile_name, git_repo_root, started_at
|
||||
model, model_config, system_prompt, system_prompt_hash,
|
||||
parent_session_id, cwd, profile_name, git_repo_root, started_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(id) DO UPDATE SET
|
||||
model = COALESCE(sessions.model, excluded.model),
|
||||
model_config = COALESCE(sessions.model_config, excluded.model_config),
|
||||
system_prompt = COALESCE(sessions.system_prompt, excluded.system_prompt),
|
||||
system_prompt_hash = COALESCE(
|
||||
sessions.system_prompt_hash,
|
||||
excluded.system_prompt_hash
|
||||
),
|
||||
system_prompt = CASE
|
||||
WHEN sessions.system_prompt_hash IS NULL
|
||||
AND excluded.system_prompt_hash IS NOT NULL
|
||||
THEN NULL
|
||||
ELSE sessions.system_prompt
|
||||
END,
|
||||
session_key = COALESCE(sessions.session_key, excluded.session_key),
|
||||
chat_id = COALESCE(sessions.chat_id, excluded.chat_id),
|
||||
chat_type = COALESCE(sessions.chat_type, excluded.chat_type),
|
||||
@@ -2876,7 +2921,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
thread_id,
|
||||
model,
|
||||
json.dumps(model_config) if model_config else None,
|
||||
system_prompt,
|
||||
system_prompt_hash,
|
||||
parent_session_id,
|
||||
cwd,
|
||||
profile_name,
|
||||
@@ -2884,6 +2929,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
time.time(),
|
||||
),
|
||||
)
|
||||
if system_prompt_hash is not None:
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
if parent_session_id:
|
||||
conn.execute(
|
||||
"""UPDATE sessions
|
||||
@@ -3150,8 +3197,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
self.flush_token_counts()
|
||||
query = f"""
|
||||
SELECT sessions.*,
|
||||
COALESCE(sp.prompt, sessions.system_prompt)
|
||||
AS _system_prompt_resolved,
|
||||
{_sql_session_last_active("sessions")} AS last_active
|
||||
FROM sessions
|
||||
LEFT JOIN system_prompts sp
|
||||
ON sp.hash = sessions.system_prompt_hash
|
||||
WHERE session_key IS NOT NULL
|
||||
AND started_at = (
|
||||
SELECT MAX(s2.started_at) FROM sessions s2
|
||||
@@ -3167,7 +3218,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
query += " ORDER BY last_active DESC"
|
||||
with self._lock:
|
||||
rows = self._conn.execute(query, params).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
return [self._session_row_dict(r) for r in rows]
|
||||
|
||||
def find_session_by_origin(
|
||||
self,
|
||||
@@ -3245,20 +3296,24 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
with self._lock:
|
||||
row = self._conn.execute(
|
||||
"""
|
||||
SELECT * FROM sessions
|
||||
WHERE session_key = ?
|
||||
AND source = ?
|
||||
AND (ended_at IS NULL OR end_reason IN ('agent_close', 'ws_orphan_reap'))
|
||||
AND (COALESCE(message_count, 0) > 0 OR EXISTS (
|
||||
SELECT 1 FROM messages WHERE messages.session_id = sessions.id LIMIT 1
|
||||
SELECT s.*,
|
||||
COALESCE(sp.prompt, s.system_prompt)
|
||||
AS _system_prompt_resolved
|
||||
FROM sessions s
|
||||
LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash
|
||||
WHERE s.session_key = ?
|
||||
AND s.source = ?
|
||||
AND (s.ended_at IS NULL OR s.end_reason IN ('agent_close', 'ws_orphan_reap'))
|
||||
AND (COALESCE(s.message_count, 0) > 0 OR EXISTS (
|
||||
SELECT 1 FROM messages WHERE messages.session_id = s.id LIMIT 1
|
||||
))
|
||||
ORDER BY started_at DESC
|
||||
ORDER BY s.started_at DESC
|
||||
LIMIT 1
|
||||
""",
|
||||
(session_key, source),
|
||||
).fetchone()
|
||||
if row is not None:
|
||||
return dict(row)
|
||||
return self._session_row_dict(row)
|
||||
|
||||
# Conservative fallback for rows created by current code but with a
|
||||
# temporarily-missing exact key: still require the complete peer
|
||||
@@ -3267,22 +3322,26 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
return None
|
||||
row = self._conn.execute(
|
||||
"""
|
||||
SELECT * FROM sessions
|
||||
WHERE source = ?
|
||||
AND COALESCE(user_id, '') = COALESCE(?, '')
|
||||
AND COALESCE(chat_id, '') = COALESCE(?, '')
|
||||
AND COALESCE(chat_type, '') = COALESCE(?, '')
|
||||
AND COALESCE(thread_id, '') = COALESCE(?, '')
|
||||
AND (ended_at IS NULL OR end_reason IN ('agent_close', 'ws_orphan_reap'))
|
||||
AND (COALESCE(message_count, 0) > 0 OR EXISTS (
|
||||
SELECT 1 FROM messages WHERE messages.session_id = sessions.id LIMIT 1
|
||||
SELECT s.*,
|
||||
COALESCE(sp.prompt, s.system_prompt)
|
||||
AS _system_prompt_resolved
|
||||
FROM sessions s
|
||||
LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash
|
||||
WHERE s.source = ?
|
||||
AND COALESCE(s.user_id, '') = COALESCE(?, '')
|
||||
AND COALESCE(s.chat_id, '') = COALESCE(?, '')
|
||||
AND COALESCE(s.chat_type, '') = COALESCE(?, '')
|
||||
AND COALESCE(s.thread_id, '') = COALESCE(?, '')
|
||||
AND (s.ended_at IS NULL OR s.end_reason IN ('agent_close', 'ws_orphan_reap'))
|
||||
AND (COALESCE(s.message_count, 0) > 0 OR EXISTS (
|
||||
SELECT 1 FROM messages WHERE messages.session_id = s.id LIMIT 1
|
||||
))
|
||||
ORDER BY started_at DESC
|
||||
ORDER BY s.started_at DESC
|
||||
LIMIT 1
|
||||
""",
|
||||
(source, user_id, chat_id, chat_type, thread_id),
|
||||
).fetchone()
|
||||
return dict(row) if row else None
|
||||
return self._session_row_dict(row) if row else None
|
||||
|
||||
def find_live_compression_child(
|
||||
self, parent_session_id: str
|
||||
@@ -3310,18 +3369,22 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
return None
|
||||
rows = self._conn.execute(
|
||||
"""
|
||||
SELECT * FROM sessions
|
||||
WHERE parent_session_id = ?
|
||||
AND ended_at IS NULL
|
||||
AND json_extract(COALESCE(model_config, '{}'), '$._branched_from') IS NULL
|
||||
AND json_extract(COALESCE(model_config, '{}'), '$._delegate_from') IS NULL
|
||||
AND COALESCE(source, '') != 'tool'
|
||||
ORDER BY started_at ASC
|
||||
SELECT s.*,
|
||||
COALESCE(sp.prompt, s.system_prompt)
|
||||
AS _system_prompt_resolved
|
||||
FROM sessions s
|
||||
LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash
|
||||
WHERE s.parent_session_id = ?
|
||||
AND s.ended_at IS NULL
|
||||
AND json_extract(COALESCE(s.model_config, '{}'), '$._branched_from') IS NULL
|
||||
AND json_extract(COALESCE(s.model_config, '{}'), '$._delegate_from') IS NULL
|
||||
AND COALESCE(s.source, '') != 'tool'
|
||||
ORDER BY s.started_at ASC
|
||||
LIMIT 2
|
||||
""",
|
||||
(parent_session_id,),
|
||||
).fetchall()
|
||||
return dict(rows[0]) if len(rows) == 1 else None
|
||||
return self._session_row_dict(rows[0]) if len(rows) == 1 else None
|
||||
|
||||
def publish_compression_child(
|
||||
self,
|
||||
@@ -3371,20 +3434,22 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
raise RuntimeError(f"Compression parent already ended: {parent_session_id}")
|
||||
if not messages:
|
||||
raise RuntimeError("Compression child handoff must not be empty")
|
||||
system_prompt_hash = self._store_system_prompt(conn, system_prompt)
|
||||
|
||||
conn.execute(
|
||||
"""INSERT INTO sessions (
|
||||
id, source, model, model_config, system_prompt,
|
||||
system_prompt_hash,
|
||||
parent_session_id, cwd, git_branch, git_repo_root,
|
||||
profile_name, user_id, session_key, chat_id, chat_type,
|
||||
thread_id, display_name, origin_json, started_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
||||
) VALUES (?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
||||
(
|
||||
child_session_id,
|
||||
source,
|
||||
model,
|
||||
json.dumps(model_config) if model_config else None,
|
||||
system_prompt,
|
||||
system_prompt_hash,
|
||||
parent_session_id,
|
||||
cwd or parent["cwd"],
|
||||
parent["git_branch"],
|
||||
@@ -4140,13 +4205,18 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
)
|
||||
self._execute_write(_do)
|
||||
|
||||
def update_system_prompt(self, session_id: str, system_prompt: str) -> None:
|
||||
def update_system_prompt(
|
||||
self, session_id: str, system_prompt: Optional[str]
|
||||
) -> None:
|
||||
"""Store the full assembled system prompt snapshot."""
|
||||
def _do(conn):
|
||||
system_prompt_hash = self._store_system_prompt(conn, system_prompt)
|
||||
conn.execute(
|
||||
"UPDATE sessions SET system_prompt = ? WHERE id = ?",
|
||||
(system_prompt, session_id),
|
||||
"UPDATE sessions "
|
||||
"SET system_prompt_hash = ?, system_prompt = NULL WHERE id = ?",
|
||||
(system_prompt_hash, session_id),
|
||||
)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
self._execute_write(_do)
|
||||
|
||||
def update_session_model(self, session_id: str, model: str) -> None:
|
||||
@@ -4178,10 +4248,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
THEN json_remove(model_config, '$.browser_model_lock')
|
||||
ELSE model_config
|
||||
END,
|
||||
system_prompt = NULL
|
||||
system_prompt = NULL,
|
||||
system_prompt_hash = NULL
|
||||
WHERE id = ?""",
|
||||
(model, session_id),
|
||||
)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
self._execute_write(_do)
|
||||
|
||||
def update_session_runtime_lock(
|
||||
@@ -4232,10 +4304,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
"""UPDATE sessions SET
|
||||
model_config = ?,
|
||||
model = COALESCE(?, model),
|
||||
system_prompt = NULL
|
||||
system_prompt = NULL,
|
||||
system_prompt_hash = NULL
|
||||
WHERE id = ?""",
|
||||
(json.dumps(config), model, session_id),
|
||||
)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
self._execute_write(_do)
|
||||
|
||||
def set_session_yolo(self, session_id: str, enabled: bool) -> None:
|
||||
@@ -4323,10 +4397,12 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
billing_provider = ?,
|
||||
billing_base_url = ?,
|
||||
billing_mode = COALESCE(?, billing_mode),
|
||||
system_prompt = NULL
|
||||
system_prompt = NULL,
|
||||
system_prompt_hash = NULL
|
||||
WHERE id = ?""",
|
||||
(provider, base_url, billing_mode, session_id),
|
||||
)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
self._execute_write(_do)
|
||||
|
||||
# ── Async token accounting ──
|
||||
@@ -4948,6 +5024,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
conn.execute(
|
||||
f"DELETE FROM sessions WHERE id IN ({placeholders})", ids
|
||||
)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
return ids
|
||||
|
||||
removed_ids = self._execute_write(_do) or []
|
||||
@@ -5004,10 +5081,15 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
self.flush_token_counts()
|
||||
with self._read_ctx() as conn:
|
||||
cursor = conn.execute(
|
||||
"SELECT * FROM sessions WHERE id = ?", (session_id,)
|
||||
"SELECT s.*, "
|
||||
"COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved "
|
||||
"FROM sessions s "
|
||||
"LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
|
||||
"WHERE s.id = ?",
|
||||
(session_id,),
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
return dict(row) if row else None
|
||||
return self._session_row_dict(row) if row else None
|
||||
|
||||
def resolve_session_id(self, session_id_or_prefix: str) -> Optional[str]:
|
||||
"""Resolve an exact or uniquely prefixed session ID to the full ID.
|
||||
@@ -5317,10 +5399,15 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
"""Look up a session by exact title. Returns session dict or None."""
|
||||
with self._read_ctx() as conn:
|
||||
cursor = conn.execute(
|
||||
"SELECT * FROM sessions WHERE title = ?", (title,)
|
||||
"SELECT s.*, "
|
||||
"COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved "
|
||||
"FROM sessions s "
|
||||
"LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
|
||||
"WHERE s.title = ?",
|
||||
(title,),
|
||||
)
|
||||
row = cursor.fetchone()
|
||||
return dict(row) if row else None
|
||||
return self._session_row_dict(row) if row else None
|
||||
|
||||
def resolve_session_by_title(self, title: str) -> Optional[str]:
|
||||
"""Resolve a title to a session ID, preferring the latest in a lineage.
|
||||
@@ -5452,7 +5539,9 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
# the projection is derived from SCHEMA_SQL so columns added later via
|
||||
# declarative reconciliation are included automatically instead of
|
||||
# silently dropping out of list rows.
|
||||
_SESSION_COMPACT_EXCLUDED = frozenset({"system_prompt"})
|
||||
_SESSION_COMPACT_EXCLUDED = frozenset(
|
||||
{"system_prompt", "system_prompt_hash"}
|
||||
)
|
||||
_session_compact_cols_sql: Optional[str] = None
|
||||
|
||||
def list_sessions_rich(
|
||||
@@ -5579,6 +5668,14 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
# Snapshot the filter params before the query builders below extend
|
||||
# them with LIMIT/OFFSET — the pinned back-fill reuses the same WHERE.
|
||||
base_where_params = list(params)
|
||||
prompt_select = (
|
||||
"" if compact_rows
|
||||
else ", COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved"
|
||||
)
|
||||
prompt_join = (
|
||||
"" if compact_rows
|
||||
else "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash"
|
||||
)
|
||||
|
||||
# Optional session-id filter, pushed into SQL so callers (Desktop
|
||||
# session-id search) don't have to fetch every row and filter in
|
||||
@@ -5675,7 +5772,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
FROM chain
|
||||
GROUP BY root_id
|
||||
)
|
||||
SELECT {_sel},
|
||||
SELECT {_sel}{prompt_select},
|
||||
COALESCE(
|
||||
(SELECT {_PREVIEW_RAW_SELECT}
|
||||
FROM messages m
|
||||
@@ -5687,6 +5784,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
COALESCE(cm.effective_last_active, s.started_at) AS _effective_last_active
|
||||
FROM sessions s
|
||||
LEFT JOIN chain_max cm ON cm.root_id = s.id
|
||||
{prompt_join}
|
||||
{outer_where}
|
||||
ORDER BY _effective_last_active DESC, s.started_at DESC, s.id DESC
|
||||
LIMIT ? OFFSET ?
|
||||
@@ -5697,7 +5795,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
else:
|
||||
_sel = self._compact_session_cols() if compact_rows else "s.*"
|
||||
query = f"""
|
||||
SELECT {_sel},
|
||||
SELECT {_sel}{prompt_select},
|
||||
COALESCE(
|
||||
(SELECT {_PREVIEW_RAW_SELECT}
|
||||
FROM messages m
|
||||
@@ -5707,6 +5805,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
) AS _preview_raw,
|
||||
{_sql_session_last_active("s")} AS last_active
|
||||
FROM sessions s
|
||||
{prompt_join}
|
||||
{where_sql}
|
||||
ORDER BY s.started_at DESC
|
||||
LIMIT ? OFFSET ?
|
||||
@@ -5717,7 +5816,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
rows = cursor.fetchall()
|
||||
sessions = []
|
||||
for row in rows:
|
||||
s = dict(row)
|
||||
s = self._session_row_dict(row)
|
||||
s["preview"] = _shape_preview(s.pop("_preview_raw", ""))
|
||||
# Drop the internal ordering column so callers see a clean dict.
|
||||
s.pop("_effective_last_active", None)
|
||||
@@ -5735,7 +5834,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
)
|
||||
_sel = self._compact_session_cols() if compact_rows else "s.*"
|
||||
pinned_query = f"""
|
||||
SELECT {_sel},
|
||||
SELECT {_sel}{prompt_select},
|
||||
COALESCE(
|
||||
(SELECT {_PREVIEW_RAW_SELECT}
|
||||
FROM messages m
|
||||
@@ -5748,6 +5847,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
s.started_at
|
||||
) AS last_active
|
||||
FROM sessions s
|
||||
{prompt_join}
|
||||
{pinned_where}
|
||||
ORDER BY s.started_at DESC
|
||||
"""
|
||||
@@ -5755,7 +5855,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
pinned_cursor = conn.execute(pinned_query, base_where_params)
|
||||
pinned_rows = pinned_cursor.fetchall()
|
||||
for row in pinned_rows:
|
||||
s = dict(row)
|
||||
s = self._session_row_dict(row)
|
||||
if s["id"] in seen_ids:
|
||||
continue
|
||||
s["preview"] = _shape_preview(s.pop("_preview_raw", ""))
|
||||
@@ -7278,8 +7378,11 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
the *current* workspace, not the global MRU.
|
||||
"""
|
||||
select_with_last_active = (
|
||||
f"SELECT s.*, {_sql_session_last_active('s')} AS last_active "
|
||||
"SELECT s.*, "
|
||||
"COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved, "
|
||||
f"{_sql_session_last_active('s')} AS last_active "
|
||||
"FROM sessions s "
|
||||
"LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
|
||||
)
|
||||
where_clauses = []
|
||||
params: list = []
|
||||
@@ -7299,7 +7402,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
"ORDER BY last_active DESC, s.started_at DESC, s.id DESC LIMIT ? OFFSET ?",
|
||||
params,
|
||||
)
|
||||
return [dict(row) for row in cursor.fetchall()]
|
||||
return [self._session_row_dict(row) for row in cursor.fetchall()]
|
||||
|
||||
# =========================================================================
|
||||
# Utility
|
||||
@@ -7627,6 +7730,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
)
|
||||
conn.execute("DELETE FROM messages WHERE session_id = ?", (session_id,))
|
||||
conn.execute("DELETE FROM sessions WHERE id = ?", (session_id,))
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
return True
|
||||
|
||||
deleted = self._execute_write(_do)
|
||||
@@ -7671,6 +7775,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
""",
|
||||
(session_id,),
|
||||
)
|
||||
if cursor.rowcount > 0:
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
return cursor.rowcount > 0
|
||||
|
||||
deleted = self._execute_write(_do)
|
||||
@@ -7751,6 +7857,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
f"DELETE FROM sessions WHERE id IN ({existing_placeholders})",
|
||||
existing,
|
||||
)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
removed_ids.extend(existing)
|
||||
return len(existing)
|
||||
|
||||
@@ -7846,6 +7953,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
)
|
||||
conn.execute("DELETE FROM sessions WHERE id = ?", (sid,))
|
||||
removed_ids.append(sid)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
return len(session_ids)
|
||||
|
||||
count = self._execute_write(_do)
|
||||
@@ -8181,6 +8289,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
conn.execute("DELETE FROM messages WHERE session_id = ?", (sid,))
|
||||
conn.execute("DELETE FROM sessions WHERE id = ?", (sid,))
|
||||
removed_ids.append(sid)
|
||||
self._delete_unreferenced_system_prompts(conn)
|
||||
return len(session_ids)
|
||||
|
||||
count = self._execute_write(_do)
|
||||
@@ -8716,6 +8825,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
rows = self._conn.execute(
|
||||
f"""
|
||||
SELECT s.*,
|
||||
COALESCE(sp.prompt, s.system_prompt)
|
||||
AS _system_prompt_resolved,
|
||||
COALESCE(
|
||||
(SELECT {_PREVIEW_RAW_SELECT}
|
||||
FROM messages m
|
||||
@@ -8725,6 +8836,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
) AS _preview_raw,
|
||||
{_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 = 'telegram'
|
||||
AND s.user_id = ?
|
||||
AND NOT EXISTS (
|
||||
@@ -8742,6 +8855,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
rows = self._conn.execute(
|
||||
f"""
|
||||
SELECT s.*,
|
||||
COALESCE(sp.prompt, s.system_prompt)
|
||||
AS _system_prompt_resolved,
|
||||
COALESCE(
|
||||
(SELECT {_PREVIEW_RAW_SELECT}
|
||||
FROM messages m
|
||||
@@ -8751,6 +8866,8 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
) AS _preview_raw,
|
||||
{_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 = 'telegram'
|
||||
AND s.user_id = ?
|
||||
ORDER BY last_active DESC, s.started_at DESC
|
||||
@@ -8761,7 +8878,7 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
|
||||
sessions: List[Dict[str, Any]] = []
|
||||
for row in rows:
|
||||
session = dict(row)
|
||||
session = self._session_row_dict(row)
|
||||
session["preview"] = _shape_preview(session.pop("_preview_raw", ""))
|
||||
sessions.append(session)
|
||||
return sessions
|
||||
@@ -9044,11 +9161,14 @@ class SessionDB(SessionSearchMixin, SessionSchemaMixin, SessionPortabilityMixin)
|
||||
"""
|
||||
try:
|
||||
cur = self._conn.execute(
|
||||
"SELECT * FROM sessions "
|
||||
"WHERE handoff_state = 'pending' "
|
||||
"ORDER BY started_at ASC"
|
||||
"SELECT s.*, "
|
||||
"COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved "
|
||||
"FROM sessions s "
|
||||
"LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
|
||||
"WHERE s.handoff_state = 'pending' "
|
||||
"ORDER BY s.started_at ASC"
|
||||
)
|
||||
return [dict(r) for r in cur.fetchall()]
|
||||
return [self._session_row_dict(r) for r in cur.fetchall()]
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
|
||||
@@ -152,7 +152,7 @@ def _sql_session_last_active_by_id(session_id_expr: str) -> str:
|
||||
)
|
||||
|
||||
|
||||
SCHEMA_VERSION = 24
|
||||
SCHEMA_VERSION = 25
|
||||
|
||||
|
||||
# FTS storage-layout version, tracked INDEPENDENTLY of SCHEMA_VERSION in the
|
||||
@@ -187,6 +187,11 @@ CREATE TABLE IF NOT EXISTS schema_version (
|
||||
version INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS system_prompts (
|
||||
hash TEXT PRIMARY KEY,
|
||||
prompt TEXT NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sessions (
|
||||
id TEXT PRIMARY KEY,
|
||||
source TEXT NOT NULL,
|
||||
@@ -201,6 +206,7 @@ CREATE TABLE IF NOT EXISTS sessions (
|
||||
model TEXT,
|
||||
model_config TEXT,
|
||||
system_prompt TEXT,
|
||||
system_prompt_hash TEXT,
|
||||
parent_session_id TEXT,
|
||||
started_at REAL NOT NULL,
|
||||
ended_at REAL,
|
||||
@@ -239,7 +245,8 @@ CREATE TABLE IF NOT EXISTS sessions (
|
||||
rewind_count INTEGER NOT NULL DEFAULT 0,
|
||||
archived INTEGER NOT NULL DEFAULT 0,
|
||||
pinned INTEGER NOT NULL DEFAULT 0,
|
||||
FOREIGN KEY (parent_session_id) REFERENCES sessions(id)
|
||||
FOREIGN KEY (parent_session_id) REFERENCES sessions(id),
|
||||
FOREIGN KEY (system_prompt_hash) REFERENCES system_prompts(hash)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS messages (
|
||||
@@ -368,6 +375,8 @@ CREATE INDEX IF NOT EXISTS idx_sessions_gateway_peer
|
||||
ON sessions(source, user_id, chat_id, chat_type, thread_id, started_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_sessions_handoff_state
|
||||
ON sessions(handoff_state, started_at);
|
||||
CREATE INDEX IF NOT EXISTS idx_sessions_system_prompt_hash
|
||||
ON sessions(system_prompt_hash);
|
||||
"""
|
||||
|
||||
|
||||
|
||||
@@ -32,7 +32,7 @@ class SessionPortabilityMixin:
|
||||
@classmethod
|
||||
def _compact_session_cols(cls) -> str:
|
||||
"""SELECT list for compact_rows: every ``sessions`` column declared in
|
||||
SCHEMA_SQL except the ``system_prompt`` blob, aliased with the ``s``
|
||||
SCHEMA_SQL except prompt storage internals, aliased with the ``s``
|
||||
prefix used by list_sessions_rich/_get_session_rich_row queries."""
|
||||
if cls._session_compact_cols_sql is None:
|
||||
declared = cls._parse_schema_columns(SCHEMA_SQL)["sessions"]
|
||||
@@ -102,6 +102,7 @@ class SessionPortabilityMixin:
|
||||
|
||||
query = f"""
|
||||
SELECT s.*,
|
||||
COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved,
|
||||
COALESCE(
|
||||
(SELECT {_PREVIEW_RAW_SELECT}
|
||||
FROM messages m
|
||||
@@ -111,6 +112,7 @@ class SessionPortabilityMixin:
|
||||
) AS _preview_raw,
|
||||
{_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 ?
|
||||
@@ -121,7 +123,7 @@ class SessionPortabilityMixin:
|
||||
|
||||
runs: List[Dict[str, Any]] = []
|
||||
for row in rows:
|
||||
s = dict(row)
|
||||
s = self._session_row_dict(row)
|
||||
s["preview"] = _shape_preview(s.pop("_preview_raw", ""))
|
||||
runs.append(s)
|
||||
return runs
|
||||
@@ -175,8 +177,16 @@ class SessionPortabilityMixin:
|
||||
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"
|
||||
)
|
||||
prompt_join = (
|
||||
"" if compact_rows
|
||||
else "LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash"
|
||||
)
|
||||
query = f"""
|
||||
SELECT {_sel},
|
||||
SELECT {_sel}{prompt_select},
|
||||
COALESCE(
|
||||
(SELECT {_PREVIEW_RAW_SELECT}
|
||||
FROM messages m
|
||||
@@ -186,6 +196,7 @@ class SessionPortabilityMixin:
|
||||
) AS _preview_raw,
|
||||
{_sql_session_last_active("s")} AS last_active
|
||||
FROM sessions s
|
||||
{prompt_join}
|
||||
WHERE s.id IN ({placeholders})
|
||||
"""
|
||||
with self._lock:
|
||||
@@ -193,7 +204,7 @@ class SessionPortabilityMixin:
|
||||
rows = cursor.fetchall()
|
||||
result: Dict[str, Dict[str, Any]] = {}
|
||||
for row in rows:
|
||||
s = dict(row)
|
||||
s = self._session_row_dict(row)
|
||||
s["preview"] = _shape_preview(s.pop("_preview_raw", ""))
|
||||
result[s["id"]] = s
|
||||
return result
|
||||
@@ -557,10 +568,14 @@ class SessionPortabilityMixin:
|
||||
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,
|
||||
@@ -571,7 +586,7 @@ class SessionPortabilityMixin:
|
||||
)
|
||||
VALUES (
|
||||
:id, :source, :user_id, :model, :model_config,
|
||||
:system_prompt, NULL, :started_at, :ended_at,
|
||||
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,
|
||||
@@ -586,7 +601,7 @@ class SessionPortabilityMixin:
|
||||
"user_id": raw.get("user_id"),
|
||||
"model": raw.get("model"),
|
||||
"model_config": raw.get("model_config"),
|
||||
"system_prompt": raw.get("system_prompt"),
|
||||
"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"),
|
||||
|
||||
@@ -36,6 +36,27 @@ logger = logging.getLogger("hermes_state")
|
||||
class SessionSchemaMixin:
|
||||
"""See module docstring — mixin for SessionDB (Schema cluster)."""
|
||||
|
||||
def _dedupe_legacy_system_prompts(self, cursor: sqlite3.Cursor) -> None:
|
||||
"""Move inline prompt snapshots into the shared content-addressed table."""
|
||||
try:
|
||||
rows = cursor.execute(
|
||||
"SELECT id, system_prompt FROM sessions "
|
||||
"WHERE system_prompt IS NOT NULL"
|
||||
).fetchall()
|
||||
except sqlite3.OperationalError:
|
||||
return
|
||||
|
||||
for row in rows:
|
||||
session_id = row["id"] if isinstance(row, sqlite3.Row) else row[0]
|
||||
prompt = row["system_prompt"] if isinstance(row, sqlite3.Row) else row[1]
|
||||
prompt_hash = self._store_system_prompt(cursor, prompt)
|
||||
cursor.execute(
|
||||
"UPDATE sessions "
|
||||
"SET system_prompt_hash = ?, system_prompt = NULL "
|
||||
"WHERE id = ?",
|
||||
(prompt_hash, session_id),
|
||||
)
|
||||
|
||||
def _sqlite_supports_fts5(self, cursor: sqlite3.Cursor) -> bool:
|
||||
try:
|
||||
cursor.execute("CREATE VIRTUAL TABLE temp._hermes_fts5_probe USING fts5(x)")
|
||||
@@ -862,6 +883,14 @@ class SessionSchemaMixin:
|
||||
if fts5_available and self._db_has_legacy_inline_fts(cursor):
|
||||
self.set_meta("fts_optimize_available", "1", cursor=cursor)
|
||||
|
||||
if current_version < 25:
|
||||
# v25: de-duplicate per-session system prompt snapshots into
|
||||
# a shared content-addressed table. Keep the old column as a
|
||||
# read fallback for partially migrated or externally written
|
||||
# rows, but clear migrated rows so future writes do not keep
|
||||
# one large prompt copy per session.
|
||||
self._dedupe_legacy_system_prompts(cursor)
|
||||
|
||||
# The FTS storage layout is versioned independently of the main
|
||||
# schema (see the v23 note above). Stamp the current layout so the
|
||||
# main version can always advance: a fresh/optimized DB is at
|
||||
|
||||
@@ -112,8 +112,6 @@ def _orphan_fts_schema(path: Path) -> None:
|
||||
conn.execute("PRAGMA writable_schema=OFF")
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
def _make_page_spanning_source(
|
||||
path: Path,
|
||||
message_count: int = 320,
|
||||
@@ -596,7 +594,58 @@ def test_cli_allow_partial_salvages_rows_across_a_corrupt_leaf(
|
||||
}
|
||||
|
||||
|
||||
|
||||
def test_partial_recovery_clears_only_unreadable_system_prompt_refs(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
source = tmp_path / "corrupt-system-prompts.db"
|
||||
output = tmp_path / "partial-system-prompts.db"
|
||||
session_count = 180
|
||||
_make_many_sessions_source(source, session_count)
|
||||
|
||||
conn = sqlite3.connect(str(source), isolation_level=None)
|
||||
try:
|
||||
row = conn.execute(
|
||||
"SELECT rootpage FROM sqlite_master "
|
||||
"WHERE type = 'table' AND name = 'system_prompts'"
|
||||
).fetchone()
|
||||
assert row is not None
|
||||
prompt_root = int(row[0])
|
||||
finally:
|
||||
conn.close()
|
||||
_corrupt_middle_table_leaf(source, prompt_root)
|
||||
|
||||
report = recover_session_database(
|
||||
source,
|
||||
output,
|
||||
work_dir=tmp_path,
|
||||
chunk_size=8,
|
||||
allow_partial=True,
|
||||
)
|
||||
|
||||
assert report["verified"] is True
|
||||
assert report["partial"] is True
|
||||
assert report["copy"]["sessions"]["status"] == "complete"
|
||||
assert report["copy"]["messages"]["status"] == "complete"
|
||||
assert report["copy"]["system_prompts"]["status"] == "partial"
|
||||
cleared = report["orphan_cleanup"]["session_prompt_refs_cleared"]
|
||||
assert 0 < cleared < session_count
|
||||
assert report["verification"]["foreign_key_check"] == []
|
||||
|
||||
conn = sqlite3.connect(str(output))
|
||||
try:
|
||||
assert conn.execute("PRAGMA integrity_check").fetchall() == [("ok",)]
|
||||
assert conn.execute("PRAGMA foreign_key_check").fetchall() == []
|
||||
assert conn.execute("SELECT COUNT(*) FROM sessions").fetchone()[0] == session_count
|
||||
retained = conn.execute(
|
||||
"SELECT COUNT(*) FROM sessions WHERE system_prompt_hash IS NOT NULL"
|
||||
).fetchone()[0]
|
||||
assert retained == session_count - cleared
|
||||
assert (
|
||||
conn.execute("SELECT COUNT(*) FROM system_prompts").fetchone()[0]
|
||||
== retained
|
||||
)
|
||||
finally:
|
||||
conn.close()
|
||||
|
||||
|
||||
|
||||
|
||||
267
tests/test_session_system_prompt_dedup.py
Normal file
267
tests/test_session_system_prompt_dedup.py
Normal file
@@ -0,0 +1,267 @@
|
||||
"""Behavior coverage for content-addressed session system prompts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import sqlite3
|
||||
import time
|
||||
|
||||
import pytest
|
||||
|
||||
from hermes_state import SCHEMA_VERSION, SessionDB
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def db(tmp_path):
|
||||
session_db = SessionDB(db_path=tmp_path / "state.db")
|
||||
yield session_db
|
||||
session_db.close()
|
||||
|
||||
|
||||
def _prompt_count(db: SessionDB) -> int:
|
||||
return int(
|
||||
db._conn.execute("SELECT COUNT(*) FROM system_prompts").fetchone()[0]
|
||||
)
|
||||
|
||||
|
||||
def test_prompt_snapshots_are_deduplicated_and_hydrated_for_readers(db):
|
||||
prompt = "You are Hermes.\n" + ("Follow the profile policy.\n" * 5)
|
||||
db.create_session(
|
||||
"s1",
|
||||
"telegram",
|
||||
session_key="agent:main:telegram:dm:c1",
|
||||
chat_id="c1",
|
||||
chat_type="dm",
|
||||
system_prompt=prompt,
|
||||
)
|
||||
db.create_session("s2", "cli", system_prompt=prompt)
|
||||
db.request_handoff("s1", "telegram")
|
||||
|
||||
stored = db._conn.execute(
|
||||
"SELECT hash, prompt FROM system_prompts"
|
||||
).fetchall()
|
||||
assert len(stored) == 1
|
||||
assert stored[0]["prompt"] == prompt
|
||||
raw_sessions = db._conn.execute(
|
||||
"SELECT system_prompt, system_prompt_hash FROM sessions ORDER BY id"
|
||||
).fetchall()
|
||||
assert [row["system_prompt"] for row in raw_sessions] == [None, None]
|
||||
assert {row["system_prompt_hash"] for row in raw_sessions} == {
|
||||
stored[0]["hash"]
|
||||
}
|
||||
|
||||
assert db.get_session("s1")["system_prompt"] == prompt
|
||||
assert db.list_sessions_rich()[0]["system_prompt"] == prompt
|
||||
assert db.search_sessions()[0]["system_prompt"] == prompt
|
||||
assert db.export_session("s1")["system_prompt"] == prompt
|
||||
assert db.list_gateway_sessions()[0]["system_prompt"] == prompt
|
||||
assert db.list_pending_handoffs()[0]["system_prompt"] == prompt
|
||||
|
||||
|
||||
def test_prompt_replacement_and_route_changes_collect_only_orphans(db):
|
||||
shared_prompt = "Model: x-ai/grok-4.5\nProvider: nous"
|
||||
db.create_session(
|
||||
"s1",
|
||||
"hermes_browser",
|
||||
model="x-ai/grok-4.5",
|
||||
model_config={"_branched_from": "parent"},
|
||||
system_prompt=shared_prompt,
|
||||
)
|
||||
db.create_session("s2", "cli", system_prompt=shared_prompt)
|
||||
|
||||
db.update_session_runtime_lock(
|
||||
"s1",
|
||||
model="anthropic/claude-opus-4.8",
|
||||
provider="anthropic",
|
||||
confirmed=True,
|
||||
)
|
||||
s1 = db.get_session("s1")
|
||||
assert s1["system_prompt"] is None
|
||||
assert json.loads(s1["model_config"])["_branched_from"] == "parent"
|
||||
assert db.get_session("s2")["system_prompt"] == shared_prompt
|
||||
assert _prompt_count(db) == 1
|
||||
|
||||
db.update_session_billing_route(
|
||||
"s2",
|
||||
provider="openrouter",
|
||||
base_url="https://example.test/v1",
|
||||
)
|
||||
assert db.get_session("s2")["system_prompt"] is None
|
||||
assert _prompt_count(db) == 0
|
||||
|
||||
db.update_system_prompt("s2", "replacement")
|
||||
assert db.get_session("s2")["system_prompt"] == "replacement"
|
||||
db.update_system_prompt("s2", None)
|
||||
assert _prompt_count(db) == 0
|
||||
|
||||
|
||||
def test_existing_session_enrichment_does_not_leak_unused_prompt(db):
|
||||
db.create_session("s1", "cli", system_prompt="original prompt")
|
||||
db.create_session("s1", "cli", system_prompt="unused prompt")
|
||||
|
||||
prompts = [
|
||||
row["prompt"]
|
||||
for row in db._conn.execute("SELECT prompt FROM system_prompts")
|
||||
]
|
||||
assert prompts == ["original prompt"]
|
||||
assert db.get_session("s1")["system_prompt"] == "original prompt"
|
||||
|
||||
|
||||
def test_every_session_deletion_path_reclaims_final_prompt_reference(db):
|
||||
def seed(session_id: str, *, source: str = "cli") -> None:
|
||||
db.create_session(
|
||||
session_id,
|
||||
source,
|
||||
system_prompt=f"unique prompt for {session_id}",
|
||||
)
|
||||
assert _prompt_count(db) == 1
|
||||
|
||||
seed("single-empty")
|
||||
assert db.delete_session_if_empty("single-empty") is True
|
||||
assert _prompt_count(db) == 0
|
||||
|
||||
seed("bulk")
|
||||
assert db.delete_sessions(["bulk"]) == 1
|
||||
assert _prompt_count(db) == 0
|
||||
|
||||
seed("ended-empty")
|
||||
db.end_session("ended-empty", "user_exit")
|
||||
assert db.delete_empty_sessions() == 1
|
||||
assert _prompt_count(db) == 0
|
||||
|
||||
seed("pruned")
|
||||
db.end_session("pruned", "user_exit")
|
||||
assert db.prune_sessions(
|
||||
older_than_days=None,
|
||||
started_before=time.time() + 1,
|
||||
) == 1
|
||||
assert _prompt_count(db) == 0
|
||||
|
||||
seed("ghost", source="tui")
|
||||
db.end_session("ghost", "user_exit")
|
||||
db._conn.execute("UPDATE sessions SET started_at = 0 WHERE id = 'ghost'")
|
||||
db._conn.commit()
|
||||
assert db.prune_empty_ghost_sessions() == 1
|
||||
assert _prompt_count(db) == 0
|
||||
|
||||
|
||||
def test_deleting_one_shared_session_preserves_prompt_until_final_reference(db):
|
||||
prompt = "shared deletion prompt"
|
||||
db.create_session("s1", "cli", system_prompt=prompt)
|
||||
db.create_session("s2", "cli", system_prompt=prompt)
|
||||
|
||||
assert db.delete_session("s1") is True
|
||||
assert _prompt_count(db) == 1
|
||||
assert db.get_session("s2")["system_prompt"] == prompt
|
||||
|
||||
assert db.delete_session("s2") is True
|
||||
assert _prompt_count(db) == 0
|
||||
|
||||
|
||||
def test_compression_child_uses_content_addressed_prompt(db):
|
||||
prompt = "compressed child prompt"
|
||||
db.create_session("parent", "webui")
|
||||
db.append_message("parent", "user", "original")
|
||||
assert db.try_acquire_compression_lock("parent", "holder", ttl_seconds=60)
|
||||
|
||||
db.publish_compression_child(
|
||||
parent_session_id="parent",
|
||||
child_session_id="child",
|
||||
source="webui",
|
||||
system_prompt=prompt,
|
||||
messages=[{"role": "user", "content": "summary"}],
|
||||
compression_lock_holder="holder",
|
||||
)
|
||||
|
||||
raw = db._conn.execute(
|
||||
"SELECT system_prompt, system_prompt_hash FROM sessions WHERE id = 'child'"
|
||||
).fetchone()
|
||||
assert raw["system_prompt"] is None
|
||||
assert raw["system_prompt_hash"] is not None
|
||||
assert db.get_session("child")["system_prompt"] == prompt
|
||||
assert _prompt_count(db) == 1
|
||||
|
||||
|
||||
def test_imported_prompts_are_deduplicated(tmp_path):
|
||||
prompt = "shared imported prompt"
|
||||
source = SessionDB(db_path=tmp_path / "source.db")
|
||||
try:
|
||||
source.create_session("s1", "cli", system_prompt=prompt)
|
||||
source.create_session("s2", "telegram", system_prompt=prompt)
|
||||
exported = [source.export_session("s1"), source.export_session("s2")]
|
||||
finally:
|
||||
source.close()
|
||||
|
||||
target = SessionDB(db_path=tmp_path / "target.db")
|
||||
try:
|
||||
result = target.import_sessions(exported)
|
||||
assert result["ok"] is True
|
||||
assert result["imported"] == 2
|
||||
assert _prompt_count(target) == 1
|
||||
raw = target._conn.execute(
|
||||
"SELECT system_prompt, system_prompt_hash FROM sessions ORDER BY id"
|
||||
).fetchall()
|
||||
assert [row["system_prompt"] for row in raw] == [None, None]
|
||||
assert len({row["system_prompt_hash"] for row in raw}) == 1
|
||||
assert target.get_session("s1")["system_prompt"] == prompt
|
||||
assert target.get_session("s2")["system_prompt"] == prompt
|
||||
finally:
|
||||
target.close()
|
||||
|
||||
|
||||
def test_v24_inline_prompts_migrate_once_to_content_addressed_storage(tmp_path):
|
||||
db_path = tmp_path / "legacy-prompts.db"
|
||||
legacy_prompt = "Legacy system prompt\n" + ("same policy\n" * 20)
|
||||
|
||||
db = SessionDB(db_path=db_path)
|
||||
db.create_session("s1", "cli")
|
||||
db.create_session("s2", "telegram")
|
||||
db._conn.execute(
|
||||
"UPDATE sessions SET system_prompt = ?, system_prompt_hash = NULL",
|
||||
(legacy_prompt,),
|
||||
)
|
||||
db._conn.execute("UPDATE schema_version SET version = 24")
|
||||
db._conn.commit()
|
||||
db.close()
|
||||
|
||||
migrated = SessionDB(db_path=db_path)
|
||||
try:
|
||||
assert migrated.get_session("s1")["system_prompt"] == legacy_prompt
|
||||
assert migrated.get_session("s2")["system_prompt"] == legacy_prompt
|
||||
assert _prompt_count(migrated) == 1
|
||||
raw_sessions = migrated._conn.execute(
|
||||
"SELECT system_prompt, system_prompt_hash FROM sessions ORDER BY id"
|
||||
).fetchall()
|
||||
assert [row["system_prompt"] for row in raw_sessions] == [None, None]
|
||||
assert len({row["system_prompt_hash"] for row in raw_sessions}) == 1
|
||||
assert migrated._conn.execute(
|
||||
"SELECT version FROM schema_version LIMIT 1"
|
||||
).fetchone()[0] == SCHEMA_VERSION
|
||||
finally:
|
||||
migrated.close()
|
||||
|
||||
|
||||
def test_compact_rows_omit_hash_and_never_read_prompt_blob(db):
|
||||
db.create_session("s1", "cli", system_prompt="never materialize me")
|
||||
|
||||
def deny_prompt_reads(action, table, column, database, trigger):
|
||||
if action == sqlite3.SQLITE_READ and table == "system_prompts":
|
||||
return sqlite3.SQLITE_DENY
|
||||
return sqlite3.SQLITE_OK
|
||||
|
||||
db._conn.set_authorizer(deny_prompt_reads)
|
||||
try:
|
||||
rows = db.list_sessions_rich(
|
||||
compact_rows=True,
|
||||
order_by_last_active=True,
|
||||
)
|
||||
rich = db._get_session_rich_row("s1", compact_rows=True)
|
||||
finally:
|
||||
db._conn.set_authorizer(None)
|
||||
|
||||
assert rows[0]["id"] == "s1"
|
||||
assert rich["id"] == "s1"
|
||||
assert "system_prompt" not in rows[0]
|
||||
assert "system_prompt_hash" not in rows[0]
|
||||
assert "system_prompt" not in rich
|
||||
assert "system_prompt_hash" not in rich
|
||||
Reference in New Issue
Block a user