fix(gateway): retry transcript appends

Queue failed session DB appends so disk order cannot silently lag memory.\nRebuild corrupt FTS indexes once and surface repeated failures as warnings.
This commit is contained in:
MaartenDMT 2026-07-16 14:24:30 +02:00 committed by kshitij
parent f32191fd52
commit 53d3588389
2 changed files with 141 additions and 25 deletions

View file

@ -1028,6 +1028,10 @@ class SessionStore:
self._persisted_routing_generation = 0
self._inflight_lock = threading.Lock()
self._inflight_sessions: Dict[str, _SessionFlight] = {}
self._transcript_retry_lock = threading.Lock()
self._dirty_transcripts: Dict[str, List[Dict[str, Any]]] = {}
self._transcript_append_failures: Dict[str, int] = {}
self._fts_rebuild_attempted = False
self._has_active_processes_fn = has_active_processes_fn
# Whether to keep writing the legacy sessions.json mirror alongside
# the primary gateway_routing table in state.db. Default True for
@ -2505,31 +2509,90 @@ class SessionStore:
_flush_messages_to_session_db(), preventing the
duplicate-write bug (#860).
"""
if self._db and not skip_db:
try:
self._db.append_message(
session_id=session_id,
role=message.get("role", "unknown"),
content=message.get("content"),
tool_name=message.get("tool_name"),
tool_calls=message.get("tool_calls"),
tool_call_id=message.get("tool_call_id"),
reasoning=message.get("reasoning") if message.get("role") == "assistant" else None,
reasoning_content=message.get("reasoning_content") if message.get("role") == "assistant" else None,
reasoning_details=message.get("reasoning_details") if message.get("role") == "assistant" else None,
codex_reasoning_items=message.get("codex_reasoning_items") if message.get("role") == "assistant" else None,
codex_message_items=message.get("codex_message_items") if message.get("role") == "assistant" else None,
# Platform-side message id (yuanbao msg_id, telegram update_id, …).
# Accept either explicit ``platform_message_id`` or the legacy
# ``message_id`` key the JSONL transcript used.
platform_message_id=(
message.get("platform_message_id") or message.get("message_id")
),
observed=bool(message.get("observed")),
timestamp=message.get("timestamp"),
)
except Exception as e:
logger.debug("Session DB operation failed: %s", e)
if not self._db or skip_db:
return
with self._transcript_retry_lock:
pending = self._dirty_transcripts.setdefault(session_id, [])
pending.append(dict(message))
while pending:
try:
self._append_transcript_message(session_id, pending[0])
except Exception as exc:
if self._is_fts_corruption_error(exc) and self._rebuild_fts_once():
try:
self._append_transcript_message(session_id, pending[0])
except Exception as retry_exc:
exc = retry_exc
else:
pending.pop(0)
continue
failures = self._transcript_append_failures.get(session_id, 0) + 1
self._transcript_append_failures[session_id] = failures
logger.warning(
"Session DB transcript append failed for %s "
"(failure_count=%d, pending=%d); will retry: %s",
session_id, failures, len(pending), exc,
)
break
else:
pending.pop(0)
if not pending:
self._dirty_transcripts.pop(session_id, None)
self._transcript_append_failures.pop(session_id, None)
def _append_transcript_message(self, session_id: str, message: Dict[str, Any]) -> None:
"""Write one transcript row. Caller serializes retries per session store."""
self._db.append_message(
session_id=session_id,
role=message.get("role", "unknown"),
content=message.get("content"),
tool_name=message.get("tool_name"),
tool_calls=message.get("tool_calls"),
tool_call_id=message.get("tool_call_id"),
reasoning=message.get("reasoning") if message.get("role") == "assistant" else None,
reasoning_content=message.get("reasoning_content") if message.get("role") == "assistant" else None,
reasoning_details=message.get("reasoning_details") if message.get("role") == "assistant" else None,
codex_reasoning_items=message.get("codex_reasoning_items") if message.get("role") == "assistant" else None,
codex_message_items=message.get("codex_message_items") if message.get("role") == "assistant" else None,
platform_message_id=(message.get("platform_message_id") or message.get("message_id")),
observed=bool(message.get("observed")),
timestamp=message.get("timestamp"),
)
@staticmethod
def _is_fts_corruption_error(exc: Exception) -> bool:
text = str(exc).lower()
return "database disk image is malformed" in text or "fts" in text
def _rebuild_fts_once(self) -> bool:
"""Attempt SQLite's documented FTS5 rebuild command once per store."""
if self._fts_rebuild_attempted:
return False
self._fts_rebuild_attempted = True
db = self._db
conn = getattr(db, "_conn", None)
lock = getattr(db, "_lock", None)
if conn is None or lock is None:
return False
rebuilt = 0
try:
with lock:
for table in getattr(db, "_FTS_TABLES", ("messages_fts", "messages_fts_trigram")):
try:
if hasattr(db, "_fts_table_exists") and not db._fts_table_exists(table):
continue
conn.execute(f"INSERT INTO {table}({table}) VALUES('rebuild')")
conn.commit()
rebuilt += 1
except Exception as exc:
conn.rollback()
logger.warning("Session DB FTS rebuild failed for %s: %s", table, exc)
except Exception as exc:
logger.warning("Session DB FTS rebuild failed: %s", exc)
return False
if rebuilt:
logger.warning("Rebuilt %d Session DB FTS index(es) after append corruption", rebuilt)
return rebuilt > 0
def has_platform_message_id(
self, session_id: str, platform_message_id: str

View file

@ -1660,6 +1660,59 @@ class TestRewriteTranscriptPreservesReasoning:
class TestGatewaySessionDbRecovery:
def test_transcript_append_rebuilds_fts_and_retries_dirty_rows_in_order(self):
import threading
class FakeConnection:
def __init__(self):
self.sql = []
def execute(self, sql):
self.sql.append(sql)
def commit(self):
pass
def rollback(self):
pass
class FakeDb:
_FTS_TABLES = ("messages_fts",)
def __init__(self):
self._lock = threading.Lock()
self._conn = FakeConnection()
self.attempts = []
self.persisted = []
def _fts_table_exists(self, _table):
return True
def append_message(self, **kwargs):
content = kwargs["content"]
self.attempts.append(content)
if len(self.attempts) <= 2:
raise RuntimeError("database disk image is malformed")
self.persisted.append(content)
store = object.__new__(SessionStore)
store._db = FakeDb()
store._transcript_retry_lock = threading.Lock()
store._dirty_transcripts = {}
store._transcript_append_failures = {}
store._fts_rebuild_attempted = False
store.append_to_transcript("s1", {"role": "user", "content": "first"})
assert [m["content"] for m in store._dirty_transcripts["s1"]] == ["first"]
store.append_to_transcript("s1", {"role": "assistant", "content": "second"})
assert store._db.persisted == ["first", "second"]
assert "s1" not in store._dirty_transcripts
assert store._db._conn.sql == [
"INSERT INTO messages_fts(messages_fts) VALUES('rebuild')"
]
def test_new_session_records_gateway_peer_fields(self, tmp_path):
store = SessionStore(sessions_dir=tmp_path, config=GatewayConfig())
source = SessionSource(