Merge branch 'simp/r2-state-b' into simp/r2-state
# Conflicts: # hermes_state_messages.py
This commit is contained in:
@@ -9,7 +9,7 @@ import json
|
|||||||
import logging
|
import logging
|
||||||
import sqlite3
|
import sqlite3
|
||||||
import time
|
import time
|
||||||
from typing import Any, Dict, List, Optional
|
from typing import Any, Dict, List, Optional, Tuple
|
||||||
|
|
||||||
from hermes_state_common import _sql_session_last_active, is_automatic_end_reason
|
from hermes_state_common import _sql_session_last_active, is_automatic_end_reason
|
||||||
|
|
||||||
@@ -27,6 +27,30 @@ def _ended_by_compression(row) -> bool:
|
|||||||
return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression"
|
return row is not None and row["ended_at"] is not None and row["end_reason"] == "compression"
|
||||||
|
|
||||||
|
|
||||||
|
def _cooldown_row(exists: bool, cooldown_until, error) -> Dict[str, Any]:
|
||||||
|
return {
|
||||||
|
"session_exists": exists,
|
||||||
|
"cooldown_until": float(cooldown_until) if cooldown_until is not None else None,
|
||||||
|
"error": error}
|
||||||
|
|
||||||
|
|
||||||
|
def _claim_lease_row(conn, table: str, key_col: str, key: str, holder: str, now: float, expires_at: float,
|
||||||
|
stale) -> Tuple[bool, Optional[str]]:
|
||||||
|
"""Single-transaction lease claim: DELETE a stale holder's row (``stale(holder,
|
||||||
|
expires_at)``), INSERT OR IGNORE ours, then SELECT to confirm ownership (INSERT OR
|
||||||
|
IGNORE gives no rowcount signal). Returns ``(acquired, reclaimed_holder)``."""
|
||||||
|
reclaimed_holder = None
|
||||||
|
row = conn.execute(f"SELECT holder, expires_at FROM {table} WHERE {key_col} = ?", (key,)).fetchone()
|
||||||
|
if row is not None and stale(row["holder"], row["expires_at"]):
|
||||||
|
conn.execute(f"DELETE FROM {table} WHERE {key_col} = ? AND holder = ?", (key, row["holder"]))
|
||||||
|
reclaimed_holder = row["holder"]
|
||||||
|
conn.execute(
|
||||||
|
f"INSERT OR IGNORE INTO {table} ({key_col}, holder, acquired_at, expires_at) VALUES (?, ?, ?, ?)",
|
||||||
|
(key, holder, now, expires_at))
|
||||||
|
owner = conn.execute(f"SELECT holder FROM {table} WHERE {key_col} = ?", (key,)).fetchone()
|
||||||
|
return owner is not None and owner["holder"] == holder, reclaimed_holder
|
||||||
|
|
||||||
|
|
||||||
class SessionCompressionMixin:
|
class SessionCompressionMixin:
|
||||||
"""Compression lineage, cooldown/streak counters, locks and turn leases."""
|
"""Compression lineage, cooldown/streak counters, locks and turn leases."""
|
||||||
|
|
||||||
@@ -96,16 +120,13 @@ class SessionCompressionMixin:
|
|||||||
deleted = conn.execute(
|
deleted = conn.execute(
|
||||||
"DELETE FROM compression_locks "
|
"DELETE FROM compression_locks "
|
||||||
"WHERE session_id = ? AND holder = ? AND expires_at = ?",
|
"WHERE session_id = ? AND holder = ? AND expires_at = ?",
|
||||||
(session_id, lock_row["holder"], expires_at),
|
(session_id, lock_row["holder"], expires_at))
|
||||||
)
|
|
||||||
if deleted.rowcount != 1:
|
if deleted.rowcount != 1:
|
||||||
return False
|
return False
|
||||||
updated = conn.execute(
|
updated = conn.execute(
|
||||||
"UPDATE sessions SET ended_at = NULL, end_reason = NULL "
|
"UPDATE sessions SET ended_at = NULL, end_reason = NULL "
|
||||||
"WHERE id = ? AND ended_at IS NOT NULL "
|
"WHERE id = ? AND ended_at IS NOT NULL AND end_reason = 'compression'",
|
||||||
"AND end_reason = 'compression'",
|
(session_id,))
|
||||||
(session_id,),
|
|
||||||
)
|
|
||||||
# rowcount==1 is guaranteed by the parent SELECT in this same txn. A False
|
# rowcount==1 is guaranteed by the parent SELECT in this same txn. A False
|
||||||
# return added past this point must raise instead: the lease DELETE above
|
# return added past this point must raise instead: the lease DELETE above
|
||||||
# commits unless _do raises.
|
# commits unless _do raises.
|
||||||
@@ -115,7 +136,10 @@ class SessionCompressionMixin:
|
|||||||
|
|
||||||
def _publish_child_session_row(self, conn, parent, *, parent_session_id, child_session_id, source,
|
def _publish_child_session_row(self, conn, parent, *, parent_session_id, child_session_id, source,
|
||||||
model, model_config, system_prompt, cwd, profile_name) -> None:
|
model, model_config, system_prompt, cwd, profile_name) -> None:
|
||||||
"""INSERT the compression child's ``sessions`` row copied from *parent*."""
|
"""INSERT the compression child's ``sessions`` row copied from *parent*. Same contract
|
||||||
|
as _insert_session_row's compression-fork backfill: the child stays on the parent's
|
||||||
|
profile and keeps gateway routing/origin columns; no owner on either side -> this
|
||||||
|
store's profile."""
|
||||||
system_prompt_hash = self._store_system_prompt(conn, system_prompt)
|
system_prompt_hash = self._store_system_prompt(conn, system_prompt)
|
||||||
conn.execute(
|
conn.execute(
|
||||||
"""INSERT INTO sessions (
|
"""INSERT INTO sessions (
|
||||||
@@ -126,28 +150,12 @@ class SessionCompressionMixin:
|
|||||||
thread_id, display_name, origin_json, started_at
|
thread_id, display_name, origin_json, started_at
|
||||||
) VALUES (?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
) VALUES (?, ?, ?, ?, NULL, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
|
||||||
(
|
(
|
||||||
child_session_id,
|
child_session_id, source, model, json.dumps(model_config) if model_config else None,
|
||||||
source,
|
system_prompt_hash, parent_session_id, cwd or parent["cwd"], parent["git_branch"],
|
||||||
model,
|
|
||||||
json.dumps(model_config) if model_config else None,
|
|
||||||
system_prompt_hash,
|
|
||||||
parent_session_id,
|
|
||||||
cwd or parent["cwd"],
|
|
||||||
parent["git_branch"],
|
|
||||||
parent["git_repo_root"],
|
parent["git_repo_root"],
|
||||||
# Same contract as _insert_session_row's compression-fork backfill: the
|
|
||||||
# child stays on the parent's profile and keeps gateway routing/origin
|
|
||||||
# columns; no owner on either side -> this store's profile.
|
|
||||||
profile_name or parent["profile_name"] or self._own_profile_name(),
|
profile_name or parent["profile_name"] or self._own_profile_name(),
|
||||||
parent["user_id"],
|
parent["user_id"], parent["session_key"], parent["chat_id"], parent["chat_type"],
|
||||||
parent["session_key"],
|
parent["thread_id"], parent["display_name"], parent["origin_json"], time.time()),
|
||||||
parent["chat_id"],
|
|
||||||
parent["chat_type"],
|
|
||||||
parent["thread_id"],
|
|
||||||
parent["display_name"],
|
|
||||||
parent["origin_json"],
|
|
||||||
time.time(),
|
|
||||||
),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def publish_compression_child(
|
def publish_compression_child(
|
||||||
@@ -180,8 +188,7 @@ class SessionCompressionMixin:
|
|||||||
conn.execute(
|
conn.execute(
|
||||||
"UPDATE compression_locks SET expires_at = ? "
|
"UPDATE compression_locks SET expires_at = ? "
|
||||||
"WHERE session_id = ? AND holder = ?",
|
"WHERE session_id = ? AND holder = ?",
|
||||||
(time.time() + lease_ttl_seconds, parent_session_id, compression_lock_holder),
|
(time.time() + lease_ttl_seconds, parent_session_id, compression_lock_holder))
|
||||||
)
|
|
||||||
lock_row = conn.execute(_LOCK_ROW_SQL, (parent_session_id,)).fetchone()
|
lock_row = conn.execute(_LOCK_ROW_SQL, (parent_session_id,)).fetchone()
|
||||||
if require_compression_lease and (
|
if require_compression_lease and (
|
||||||
lock_row is None
|
lock_row is None
|
||||||
@@ -209,10 +216,8 @@ class SessionCompressionMixin:
|
|||||||
# Deliberate boundaries still fail closed.
|
# Deliberate boundaries still fail closed.
|
||||||
if is_automatic_end_reason(parent["end_reason"]):
|
if is_automatic_end_reason(parent["end_reason"]):
|
||||||
conn.execute(
|
conn.execute(
|
||||||
"UPDATE sessions SET ended_at = NULL, end_reason = NULL "
|
"UPDATE sessions SET ended_at = NULL, end_reason = NULL WHERE id = ?",
|
||||||
"WHERE id = ?",
|
(parent_session_id,))
|
||||||
(parent_session_id,),
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
raise RuntimeError(f"Compression parent already ended: {parent_session_id}")
|
raise RuntimeError(f"Compression parent already ended: {parent_session_id}")
|
||||||
if not messages:
|
if not messages:
|
||||||
@@ -220,23 +225,17 @@ class SessionCompressionMixin:
|
|||||||
self._publish_child_session_row(
|
self._publish_child_session_row(
|
||||||
conn, parent, parent_session_id=parent_session_id, child_session_id=child_session_id,
|
conn, parent, parent_session_id=parent_session_id, child_session_id=child_session_id,
|
||||||
source=source, model=model, model_config=model_config, system_prompt=system_prompt,
|
source=source, model=model, model_config=model_config, system_prompt=system_prompt,
|
||||||
cwd=cwd, profile_name=profile_name,
|
cwd=cwd, profile_name=profile_name)
|
||||||
)
|
|
||||||
total_messages, total_tool_calls = self._insert_message_rows(conn, child_session_id, messages)
|
total_messages, total_tool_calls = self._insert_message_rows(conn, child_session_id, messages)
|
||||||
if watermark is not None:
|
if watermark is not None:
|
||||||
# Clone the parent's concurrent tail into the child after the handoff;
|
# Clone the parent's concurrent tail into the child after the handoff;
|
||||||
# originals stay in the closed parent for lineage recovery.
|
# originals stay in the closed parent for lineage recovery.
|
||||||
_ceiling_clause = ""
|
bounded = watermark_ceiling is not None
|
||||||
_params: list = [parent_session_id, int(watermark)]
|
|
||||||
if watermark_ceiling is not None:
|
|
||||||
_ceiling_clause = " AND id <= ?"
|
|
||||||
_params.append(int(watermark_ceiling))
|
|
||||||
tail_ids, tail_tool_calls = self._tail_rows_after_watermark(
|
tail_ids, tail_tool_calls = self._tail_rows_after_watermark(
|
||||||
conn,
|
conn, "SELECT id, tool_calls FROM messages "
|
||||||
"SELECT id, tool_calls FROM messages "
|
|
||||||
"WHERE session_id = ? AND active = 1 AND id > ?"
|
"WHERE session_id = ? AND active = 1 AND id > ?"
|
||||||
f"{_ceiling_clause} ORDER BY id",
|
f"{' AND id <= ?' if bounded else ''} ORDER BY id",
|
||||||
_params,
|
[parent_session_id, int(watermark), *([int(watermark_ceiling)] if bounded else [])],
|
||||||
)
|
)
|
||||||
if tail_ids:
|
if tail_ids:
|
||||||
self._clone_message_rows(conn, tail_ids, session_id=child_session_id)
|
self._clone_message_rows(conn, tail_ids, session_id=child_session_id)
|
||||||
@@ -244,13 +243,10 @@ class SessionCompressionMixin:
|
|||||||
total_tool_calls += tail_tool_calls
|
total_tool_calls += tail_tool_calls
|
||||||
conn.execute(
|
conn.execute(
|
||||||
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
|
"UPDATE sessions SET message_count = ?, tool_call_count = ? WHERE id = ?",
|
||||||
(total_messages, total_tool_calls, child_session_id),
|
(total_messages, total_tool_calls, child_session_id))
|
||||||
)
|
|
||||||
updated = conn.execute(
|
updated = conn.execute(
|
||||||
"UPDATE sessions SET ended_at = ?, end_reason = 'compression' "
|
"UPDATE sessions SET ended_at = ?, end_reason = 'compression' "
|
||||||
"WHERE id = ? AND ended_at IS NULL",
|
"WHERE id = ? AND ended_at IS NULL", (time.time(), parent_session_id))
|
||||||
(time.time(), parent_session_id),
|
|
||||||
)
|
|
||||||
if updated.rowcount != 1:
|
if updated.rowcount != 1:
|
||||||
raise RuntimeError(f"Compression parent changed during publication: {parent_session_id}")
|
raise RuntimeError(f"Compression parent changed during publication: {parent_session_id}")
|
||||||
|
|
||||||
@@ -276,8 +272,7 @@ class SessionCompressionMixin:
|
|||||||
" AND compression_failure_cooldown_until > ? "
|
" AND compression_failure_cooldown_until > ? "
|
||||||
"THEN compression_failure_cooldown_until ELSE ? END, "
|
"THEN compression_failure_cooldown_until ELSE ? END, "
|
||||||
"compression_failure_error = ? WHERE id = ?",
|
"compression_failure_error = ? WHERE id = ?",
|
||||||
(cooldown_until, cooldown_until, error, session_id),
|
(cooldown_until, cooldown_until, error, session_id))
|
||||||
)
|
|
||||||
|
|
||||||
def get_compression_failure_cooldown(self, session_id: str) -> Optional[Dict[str, Any]]:
|
def get_compression_failure_cooldown(self, session_id: str) -> Optional[Dict[str, Any]]:
|
||||||
"""Return the active (unexpired) compression-failure cooldown, or None."""
|
"""Return the active (unexpired) compression-failure cooldown, or None."""
|
||||||
@@ -285,24 +280,17 @@ class SessionCompressionMixin:
|
|||||||
return None
|
return None
|
||||||
now = time.time()
|
now = time.time()
|
||||||
row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,))
|
row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,))
|
||||||
if row is None or row[0] is None:
|
if row is None or row[0] is None or float(row[0]) <= now:
|
||||||
return None
|
return None
|
||||||
cooldown_until = float(row[0])
|
return {"cooldown_until": float(row[0]), "remaining_seconds": float(row[0]) - now, "error": row[1]}
|
||||||
if cooldown_until <= now:
|
|
||||||
return None
|
|
||||||
return {"cooldown_until": cooldown_until, "remaining_seconds": cooldown_until - now, "error": row[1]}
|
|
||||||
|
|
||||||
def get_compression_failure_cooldown_row(self, session_id: str) -> Dict[str, Any]:
|
def get_compression_failure_cooldown_row(self, session_id: str) -> Dict[str, Any]:
|
||||||
"""Exact stored cooldown columns, no expiry filtering, so compression
|
"""Exact stored cooldown columns, no expiry filtering, so compression
|
||||||
cancellation can roll back an expired, partially-null, or absent row exactly."""
|
cancellation can roll back an expired, partially-null, or absent row exactly."""
|
||||||
row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) if session_id else None
|
row = self._read_one(_COOLDOWN_ROW_SQL, (session_id,)) if session_id else None
|
||||||
if row is None:
|
if row is None:
|
||||||
return {"session_exists": False, "cooldown_until": None, "error": None}
|
return _cooldown_row(False, None, None)
|
||||||
return {
|
return _cooldown_row(True, row[0], row[1])
|
||||||
"session_exists": True,
|
|
||||||
"cooldown_until": float(row[0]) if row[0] is not None else None,
|
|
||||||
"error": row[1],
|
|
||||||
}
|
|
||||||
|
|
||||||
def restore_compression_failure_cooldown_row(self, session_id: str, snapshot: Dict[str, Any]) -> None:
|
def restore_compression_failure_cooldown_row(self, session_id: str, snapshot: Dict[str, Any]) -> None:
|
||||||
"""Restore and verify an exact cooldown-row snapshot. Unlike record/clear this
|
"""Restore and verify an exact cooldown-row snapshot. Unlike record/clear this
|
||||||
@@ -318,19 +306,12 @@ class SessionCompressionMixin:
|
|||||||
def _do(conn):
|
def _do(conn):
|
||||||
cursor = conn.execute(
|
cursor = conn.execute(
|
||||||
"UPDATE sessions SET compression_failure_cooldown_until = ?, "
|
"UPDATE sessions SET compression_failure_cooldown_until = ?, "
|
||||||
"compression_failure_error = ? WHERE id = ?",
|
"compression_failure_error = ? WHERE id = ?", (deadline, error, session_id))
|
||||||
(deadline, error, session_id),
|
|
||||||
)
|
|
||||||
if cursor.rowcount != 1:
|
if cursor.rowcount != 1:
|
||||||
raise RuntimeError(f"compression cooldown rollback session missing: {session_id}")
|
raise RuntimeError(f"compression cooldown rollback session missing: {session_id}")
|
||||||
|
|
||||||
self._execute_write(_do)
|
self._execute_write(_do)
|
||||||
actual = self.get_compression_failure_cooldown_row(session_id)
|
actual = self.get_compression_failure_cooldown_row(session_id)
|
||||||
expected = {
|
expected = _cooldown_row(True, deadline, error)
|
||||||
"session_exists": True,
|
|
||||||
"cooldown_until": float(deadline) if deadline is not None else None,
|
|
||||||
"error": error,
|
|
||||||
}
|
|
||||||
if actual != expected:
|
if actual != expected:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"compression cooldown rollback verification failed: "
|
f"compression cooldown rollback verification failed: "
|
||||||
@@ -344,9 +325,7 @@ class SessionCompressionMixin:
|
|||||||
self._write_sql_logged(
|
self._write_sql_logged(
|
||||||
"clear_compression_failure_cooldown", session_id,
|
"clear_compression_failure_cooldown", session_id,
|
||||||
"UPDATE sessions SET compression_failure_cooldown_until = NULL, "
|
"UPDATE sessions SET compression_failure_cooldown_until = NULL, "
|
||||||
"compression_failure_error = NULL WHERE id = ?",
|
"compression_failure_error = NULL WHERE id = ?", (session_id,))
|
||||||
(session_id,),
|
|
||||||
)
|
|
||||||
|
|
||||||
def _read_session_number(self, column: str, session_id: str, cast: type, zero: Any) -> Any:
|
def _read_session_number(self, column: str, session_id: str, cast: type, zero: Any) -> Any:
|
||||||
"""Read one numeric ``sessions`` column clamped at ``zero``; a missing session,
|
"""Read one numeric ``sessions`` column clamped at ``zero``; a missing session,
|
||||||
@@ -370,8 +349,7 @@ class SessionCompressionMixin:
|
|||||||
if session_id:
|
if session_id:
|
||||||
self._write_sql(
|
self._write_sql(
|
||||||
"UPDATE sessions SET compression_fallback_streak = ? WHERE id = ?",
|
"UPDATE sessions SET compression_fallback_streak = ? WHERE id = ?",
|
||||||
(max(0, int(streak)), session_id),
|
(max(0, int(streak)), session_id))
|
||||||
)
|
|
||||||
|
|
||||||
def get_compression_ineffective_count(self, session_id: str) -> int:
|
def get_compression_ineffective_count(self, session_id: str) -> int:
|
||||||
"""Persisted ineffective-compaction strike count — the durable half of the
|
"""Persisted ineffective-compaction strike count — the durable half of the
|
||||||
@@ -384,8 +362,7 @@ class SessionCompressionMixin:
|
|||||||
if session_id:
|
if session_id:
|
||||||
self._write_sql(
|
self._write_sql(
|
||||||
"UPDATE sessions SET compression_ineffective_count = ? WHERE id = ?",
|
"UPDATE sessions SET compression_ineffective_count = ? WHERE id = ?",
|
||||||
(max(0, int(count)), session_id),
|
(max(0, int(count)), session_id))
|
||||||
)
|
|
||||||
|
|
||||||
def get_compression_recovery_deadline(self, session_id: str) -> float:
|
def get_compression_recovery_deadline(self, session_id: str) -> float:
|
||||||
"""Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). Durable
|
"""Persisted anti-thrash recovery deadline (epoch; ``0.0`` = not armed). Durable
|
||||||
@@ -402,8 +379,7 @@ class SessionCompressionMixin:
|
|||||||
normalized = 0.0
|
normalized = 0.0
|
||||||
self._write_sql(
|
self._write_sql(
|
||||||
"UPDATE sessions SET compression_recovery_deadline = ? WHERE id = ?",
|
"UPDATE sessions SET compression_recovery_deadline = ? WHERE id = ?",
|
||||||
(normalized if normalized > 0.0 else None, session_id),
|
(normalized or None, session_id))
|
||||||
)
|
|
||||||
|
|
||||||
def refresh_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool:
|
def refresh_compression_lock(self, session_id: str, holder: str, ttl_seconds: float = 300.0) -> bool:
|
||||||
"""Extend the compression lock lease if ``holder`` still owns it.
|
"""Extend the compression lock lease if ``holder`` still owns it.
|
||||||
@@ -419,8 +395,7 @@ class SessionCompressionMixin:
|
|||||||
expires_at = time.time() + ttl_seconds
|
expires_at = time.time() + ttl_seconds
|
||||||
try:
|
try:
|
||||||
return self._write_rowcount(
|
return self._write_rowcount(
|
||||||
"UPDATE compression_locks SET expires_at = ? "
|
"UPDATE compression_locks SET expires_at = ? WHERE session_id = ? AND holder = ?",
|
||||||
"WHERE session_id = ? AND holder = ?",
|
|
||||||
(expires_at, session_id, holder),
|
(expires_at, session_id, holder),
|
||||||
) > 0
|
) > 0
|
||||||
except sqlite3.Error as exc:
|
except sqlite3.Error as exc:
|
||||||
@@ -442,32 +417,14 @@ class SessionCompressionMixin:
|
|||||||
expires_at = now + ttl_seconds
|
expires_at = now + ttl_seconds
|
||||||
|
|
||||||
def _do(conn):
|
def _do(conn):
|
||||||
reclaimed_holder = None
|
return _claim_lease_row(
|
||||||
row = conn.execute(_LOCK_ROW_SQL, (session_id,)).fetchone()
|
conn, "compression_locks", "session_id", session_id, holder, now, expires_at,
|
||||||
if row is not None:
|
lambda h, e: e < now or _compression_lock_holder_process_is_dead(h))
|
||||||
current_holder, current_expires_at = row[0], row[1]
|
|
||||||
if current_expires_at < now or _compression_lock_holder_process_is_dead(current_holder):
|
|
||||||
conn.execute(
|
|
||||||
"DELETE FROM compression_locks "
|
|
||||||
"WHERE session_id = ? AND holder = ?",
|
|
||||||
(session_id, current_holder),
|
|
||||||
)
|
|
||||||
reclaimed_holder = current_holder
|
|
||||||
conn.execute(
|
|
||||||
"INSERT OR IGNORE INTO compression_locks "
|
|
||||||
"(session_id, holder, acquired_at, expires_at) "
|
|
||||||
"VALUES (?, ?, ?, ?)",
|
|
||||||
(session_id, holder, now, expires_at),
|
|
||||||
)
|
|
||||||
row = conn.execute("SELECT holder FROM compression_locks WHERE session_id = ?", (session_id,)).fetchone()
|
|
||||||
return row is not None and row[0] == holder, reclaimed_holder
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
acquired, reclaimed_holder = self._execute_write(_do)
|
acquired, reclaimed_holder = self._execute_write(_do)
|
||||||
if reclaimed_holder:
|
if reclaimed_holder:
|
||||||
logger.warning(
|
logger.warning("Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder)
|
||||||
"Reclaimed stale compression lock for session=%s (holder=%s)", session_id, reclaimed_holder,
|
|
||||||
)
|
|
||||||
return bool(acquired)
|
return bool(acquired)
|
||||||
except sqlite3.Error as exc:
|
except sqlite3.Error as exc:
|
||||||
# False makes the caller skip compression — safe when the lock subsystem is broken.
|
# False makes the caller skip compression — safe when the lock subsystem is broken.
|
||||||
@@ -480,10 +437,8 @@ class SessionCompressionMixin:
|
|||||||
return
|
return
|
||||||
self._write_sql_logged(
|
self._write_sql_logged(
|
||||||
"release_compression_lock", session_id,
|
"release_compression_lock", session_id,
|
||||||
"DELETE FROM compression_locks "
|
"DELETE FROM compression_locks WHERE session_id = ? AND holder = ?",
|
||||||
"WHERE session_id = ? AND holder = ?",
|
(session_id, holder))
|
||||||
(session_id, holder),
|
|
||||||
)
|
|
||||||
|
|
||||||
def _session_turn_lease_key_on_conn(self, conn, session_id: str) -> str:
|
def _session_turn_lease_key_on_conn(self, conn, session_id: str) -> str:
|
||||||
"""Walk compression parents on ``conn`` to the conversation lease key.
|
"""Walk compression parents on ``conn`` to the conversation lease key.
|
||||||
@@ -496,9 +451,7 @@ class SessionCompressionMixin:
|
|||||||
|
|
||||||
def _row(sid: str):
|
def _row(sid: str):
|
||||||
row = conn.execute(
|
row = conn.execute(
|
||||||
"SELECT id, parent_session_id, source, model_config, end_reason "
|
"SELECT id, parent_session_id, source, model_config, end_reason FROM sessions WHERE id = ?", (sid,),
|
||||||
"FROM sessions WHERE id = ?",
|
|
||||||
(sid,),
|
|
||||||
).fetchone()
|
).fetchone()
|
||||||
return dict(row) if row else None
|
return dict(row) if row else None
|
||||||
|
|
||||||
@@ -537,29 +490,10 @@ class SessionCompressionMixin:
|
|||||||
|
|
||||||
def _do(conn):
|
def _do(conn):
|
||||||
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
||||||
row = conn.execute(
|
return _claim_lease_row(
|
||||||
"SELECT holder, expires_at FROM session_turn_leases "
|
conn, "session_turn_leases", "conversation_id", conversation_id, holder, now, expires_at,
|
||||||
"WHERE conversation_id = ?",
|
lambda h, e: float(e) <= now or _compression_lock_holder_process_is_dead(h),
|
||||||
(conversation_id,),
|
)[0]
|
||||||
).fetchone()
|
|
||||||
if row is not None:
|
|
||||||
current_holder = row["holder"]
|
|
||||||
if float(row["expires_at"]) <= now or _compression_lock_holder_process_is_dead(current_holder):
|
|
||||||
conn.execute(
|
|
||||||
"DELETE FROM session_turn_leases "
|
|
||||||
"WHERE conversation_id = ? AND holder = ?",
|
|
||||||
(conversation_id, current_holder),
|
|
||||||
)
|
|
||||||
conn.execute(
|
|
||||||
"INSERT OR IGNORE INTO session_turn_leases "
|
|
||||||
"(conversation_id, holder, acquired_at, expires_at) "
|
|
||||||
"VALUES (?, ?, ?, ?)",
|
|
||||||
(conversation_id, holder, now, expires_at),
|
|
||||||
)
|
|
||||||
owner = conn.execute(
|
|
||||||
"SELECT holder FROM session_turn_leases WHERE conversation_id = ?", (conversation_id,),
|
|
||||||
).fetchone()
|
|
||||||
return owner is not None and owner["holder"] == holder
|
|
||||||
|
|
||||||
return bool(self._execute_write(_do, patience_s=patience_s))
|
return bool(self._execute_write(_do, patience_s=patience_s))
|
||||||
|
|
||||||
@@ -588,8 +522,7 @@ class SessionCompressionMixin:
|
|||||||
logger.debug("session turn lease should_abort callback failed", exc_info=True)
|
logger.debug("session turn lease should_abort callback failed", exc_info=True)
|
||||||
try:
|
try:
|
||||||
if self.try_acquire_session_turn_lease(
|
if self.try_acquire_session_turn_lease(
|
||||||
session_id, holder, ttl_seconds=ttl_seconds, patience_s=acquire_patience_s,
|
session_id, holder, ttl_seconds=ttl_seconds, patience_s=acquire_patience_s):
|
||||||
):
|
|
||||||
return True
|
return True
|
||||||
except sqlite3.Error as exc:
|
except sqlite3.Error as exc:
|
||||||
# Long holder transactions can exhaust one write-patience budget; keep
|
# Long holder transactions can exhaust one write-patience budget; keep
|
||||||
@@ -620,12 +553,10 @@ class SessionCompressionMixin:
|
|||||||
|
|
||||||
def _do(conn):
|
def _do(conn):
|
||||||
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
||||||
cursor = conn.execute(
|
return conn.execute(
|
||||||
"UPDATE session_turn_leases SET expires_at = ? "
|
"UPDATE session_turn_leases SET expires_at = ? "
|
||||||
"WHERE conversation_id = ? AND holder = ?",
|
"WHERE conversation_id = ? AND holder = ?", (expires_at, conversation_id, holder),
|
||||||
(expires_at, conversation_id, holder),
|
).rowcount > 0
|
||||||
)
|
|
||||||
return cursor.rowcount > 0
|
|
||||||
|
|
||||||
return bool(self._execute_write(_do))
|
return bool(self._execute_write(_do))
|
||||||
|
|
||||||
@@ -637,10 +568,8 @@ class SessionCompressionMixin:
|
|||||||
def _do(conn):
|
def _do(conn):
|
||||||
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
conversation_id = self._session_turn_lease_key_on_conn(conn, session_id)
|
||||||
conn.execute(
|
conn.execute(
|
||||||
"DELETE FROM session_turn_leases "
|
"DELETE FROM session_turn_leases WHERE conversation_id = ? AND holder = ?",
|
||||||
"WHERE conversation_id = ? AND holder = ?",
|
(conversation_id, holder))
|
||||||
(conversation_id, holder),
|
|
||||||
)
|
|
||||||
|
|
||||||
self._execute_write(_do)
|
self._execute_write(_do)
|
||||||
|
|
||||||
@@ -649,10 +578,8 @@ class SessionCompressionMixin:
|
|||||||
if not session_id:
|
if not session_id:
|
||||||
return None
|
return None
|
||||||
row = self._read_one(
|
row = self._read_one(
|
||||||
"SELECT holder FROM compression_locks "
|
"SELECT holder FROM compression_locks WHERE session_id = ? AND expires_at >= ?",
|
||||||
"WHERE session_id = ? AND expires_at >= ?",
|
(session_id, time.time()))
|
||||||
(session_id, time.time()),
|
|
||||||
)
|
|
||||||
return None if row is None else row[0]
|
return None if row is None else row[0]
|
||||||
|
|
||||||
def finalize_orphaned_compression_sessions(self) -> int:
|
def finalize_orphaned_compression_sessions(self) -> int:
|
||||||
@@ -660,10 +587,7 @@ class SessionCompressionMixin:
|
|||||||
has messages, no end_reason/ended_at, api_call_count=0, older than 7 days) as
|
has messages, no end_reason/ended_at, api_call_count=0, older than 7 days) as
|
||||||
``orphaned_compression``. Non-destructive."""
|
``orphaned_compression``. Non-destructive."""
|
||||||
cutoff = time.time() - 604800 # 7 days
|
cutoff = time.time() - 604800 # 7 days
|
||||||
|
return self._write_rowcount(
|
||||||
def _do(conn):
|
|
||||||
now = time.time()
|
|
||||||
result = conn.execute(
|
|
||||||
"""
|
"""
|
||||||
UPDATE sessions
|
UPDATE sessions
|
||||||
SET ended_at = ?,
|
SET ended_at = ?,
|
||||||
@@ -684,11 +608,8 @@ class SessionCompressionMixin:
|
|||||||
WHERE m.session_id = sessions.id
|
WHERE m.session_id = sessions.id
|
||||||
)
|
)
|
||||||
""",
|
""",
|
||||||
(now, cutoff),
|
(time.time(), cutoff),
|
||||||
)
|
) or 0
|
||||||
return result.rowcount
|
|
||||||
|
|
||||||
return self._execute_write(_do) or 0
|
|
||||||
|
|
||||||
def get_compression_chain(self, session_id: str) -> List[str]:
|
def get_compression_chain(self, session_id: str) -> List[str]:
|
||||||
"""Walk the compression-continuation chain forward: root-first through the tip
|
"""Walk the compression-continuation chain forward: root-first through the tip
|
||||||
@@ -703,7 +624,7 @@ class SessionCompressionMixin:
|
|||||||
are still live over stale closed siblings such as ``ws_orphan_reap``."""
|
are still live over stale closed siblings such as ``ws_orphan_reap``."""
|
||||||
current = session_id
|
current = session_id
|
||||||
chain = [current] if current else []
|
chain = [current] if current else []
|
||||||
seen = {current} if current else set()
|
seen = set(chain)
|
||||||
for _ in range(100): # defensive bound; chains this deep are pathological
|
for _ in range(100): # defensive bound; chains this deep are pathological
|
||||||
with self._read_ctx() as conn:
|
with self._read_ctx() as conn:
|
||||||
row = conn.execute(
|
row = conn.execute(
|
||||||
@@ -729,9 +650,7 @@ class SessionCompressionMixin:
|
|||||||
""",
|
""",
|
||||||
(current,),
|
(current,),
|
||||||
).fetchone()
|
).fetchone()
|
||||||
if row is None:
|
child_id = row["id"] if row is not None else None
|
||||||
return chain
|
|
||||||
child_id = row["id"]
|
|
||||||
if not child_id or child_id in seen:
|
if not child_id or child_id in seen:
|
||||||
return chain
|
return chain
|
||||||
seen.add(child_id)
|
seen.add(child_id)
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -79,8 +79,7 @@ class SessionTitlesMixin:
|
|||||||
nothing overwrites a user name, re-running the titler on an llm row is a no-op).
|
nothing overwrites a user name, re-running the titler on an llm row is a no-op).
|
||||||
No writer may move a hidden canonical Bot Chat off its title. Read and write are
|
No writer may move a hidden canonical Bot Chat off its title. Read and write are
|
||||||
one compare-and-swap transaction, so a manual ``/title`` racing an in-flight
|
one compare-and-swap transaction, so a manual ``/title`` racing an in-flight
|
||||||
generation is not clobbered.
|
generation is not clobbered."""
|
||||||
"""
|
|
||||||
title = self.sanitize_title(title)
|
title = self.sanitize_title(title)
|
||||||
is_user = source == self.TITLE_SOURCE_USER
|
is_user = source == self.TITLE_SOURCE_USER
|
||||||
new_rank = self._title_rank(source) if not is_user else None
|
new_rank = self._title_rank(source) if not is_user else None
|
||||||
@@ -166,21 +165,16 @@ class SessionTitlesMixin:
|
|||||||
if source not in self._TITLE_SOURCE_RANK:
|
if source not in self._TITLE_SOURCE_RANK:
|
||||||
raise ValueError(f"invalid title source: {source!r}")
|
raise ValueError(f"invalid title source: {source!r}")
|
||||||
return self._write_rowcount(
|
return self._write_rowcount(
|
||||||
"UPDATE sessions SET title_source = ? "
|
"UPDATE sessions SET title_source = ? WHERE id = ? AND title IS NOT NULL",
|
||||||
"WHERE id = ? AND title IS NOT NULL",
|
|
||||||
(source, session_id),
|
(source, session_id),
|
||||||
) > 0
|
) > 0
|
||||||
|
|
||||||
def get_session_by_title(self, title: str) -> Optional[Dict[str, Any]]:
|
def get_session_by_title(self, title: str) -> Optional[Dict[str, Any]]:
|
||||||
"""Look up a session by exact title. Returns session dict or None."""
|
"""Look up a session by exact title. Returns session dict or None."""
|
||||||
row = self._read_one(
|
row = self._read_one(
|
||||||
"SELECT s.*, "
|
"SELECT s.*, COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved "
|
||||||
"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 "
|
||||||
"FROM sessions s "
|
"WHERE s.title = ?", (title,))
|
||||||
"LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash "
|
|
||||||
"WHERE s.title = ?",
|
|
||||||
(title,),
|
|
||||||
)
|
|
||||||
return self._session_row_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]:
|
def resolve_session_by_title(self, title: str) -> Optional[str]:
|
||||||
@@ -191,8 +185,7 @@ class SessionTitlesMixin:
|
|||||||
numbered = self._read_all(
|
numbered = self._read_all(
|
||||||
"SELECT id, title, started_at FROM sessions "
|
"SELECT id, title, started_at FROM sessions "
|
||||||
"WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC",
|
"WHERE title LIKE ? ESCAPE '\\' ORDER BY started_at DESC",
|
||||||
(f"{_escape_like(title)} #%",),
|
(f"{_escape_like(title)} #%",))
|
||||||
)
|
|
||||||
if numbered:
|
if numbered:
|
||||||
return numbered[0]["id"]
|
return numbered[0]["id"]
|
||||||
return exact["id"] if exact else None
|
return exact["id"] if exact else None
|
||||||
@@ -204,13 +197,9 @@ class SessionTitlesMixin:
|
|||||||
base = match.group(1) if match else base_title
|
base = match.group(1) if match else base_title
|
||||||
rows = self._read_all(
|
rows = self._read_all(
|
||||||
"SELECT title FROM sessions WHERE title = ? OR title LIKE ? ESCAPE '\\'",
|
"SELECT title FROM sessions WHERE title = ? OR title LIKE ? ESCAPE '\\'",
|
||||||
(base, f"{_escape_like(base)} #%"),
|
(base, f"{_escape_like(base)} #%"))
|
||||||
)
|
|
||||||
if not rows:
|
if not rows:
|
||||||
return base
|
return base
|
||||||
max_num = 1 # the unnumbered original counts as #1
|
# The unnumbered original counts as #1.
|
||||||
for row in rows:
|
numbers = [int(m.group(1)) for m in (re.match(r'^.* #(\d+)$', row["title"]) for row in rows) if m]
|
||||||
m = re.match(r'^.* #(\d+)$', row["title"])
|
return f"{base} #{max([1, *numbers]) + 1}"
|
||||||
if m:
|
|
||||||
max_num = max(max_num, int(m.group(1)))
|
|
||||||
return f"{base} #{max_num + 1}"
|
|
||||||
|
|||||||
@@ -78,6 +78,14 @@ _MODEL_USAGE_UPSERT_SQL = """INSERT INTO session_model_usage (
|
|||||||
last_seen = excluded.last_seen"""
|
last_seen = excluded.last_seen"""
|
||||||
|
|
||||||
|
|
||||||
|
# Kwargs forwarded verbatim from update_token_counts / record_auxiliary_usage into
|
||||||
|
# _record_model_usage (the per-route attribution row).
|
||||||
|
_MODEL_USAGE_FIELDS = frozenset((
|
||||||
|
"model", "billing_provider", "billing_base_url", "billing_mode", "input_tokens", "output_tokens",
|
||||||
|
"cache_read_tokens", "cache_write_tokens", "reasoning_tokens", "estimated_cost_usd",
|
||||||
|
"actual_cost_usd", "cost_status", "cost_source", "api_call_count"))
|
||||||
|
|
||||||
|
|
||||||
class SessionUsageMixin:
|
class SessionUsageMixin:
|
||||||
"""Coalesced token writer, per-model usage rows, billing route."""
|
"""Coalesced token writer, per-model usage rows, billing route."""
|
||||||
|
|
||||||
@@ -110,16 +118,16 @@ class SessionUsageMixin:
|
|||||||
to the synchronous path and may raise."""
|
to the synchronous path and may raise."""
|
||||||
with self._token_queue_cond:
|
with self._token_queue_cond:
|
||||||
thread = self._token_writer_thread
|
thread = self._token_writer_thread
|
||||||
writer_stopped = self._token_writer_stop and (thread is None or not thread.is_alive())
|
writer_alive = thread is not None and thread.is_alive()
|
||||||
|
writer_stopped = self._token_writer_stop and not writer_alive
|
||||||
if not writer_stopped:
|
if not writer_stopped:
|
||||||
self._token_queue.append((session_id, kwargs))
|
self._token_queue.append((session_id, kwargs))
|
||||||
if thread is None or not thread.is_alive():
|
if not writer_alive:
|
||||||
# Daemon so exit never hangs on accounting; the atexit hook drains
|
# Daemon so exit never hangs on accounting; the atexit hook drains
|
||||||
# leftovers. ``not is_alive()`` (not ``is None``) respawns a writer
|
# leftovers. ``not is_alive()`` (not ``is None``) respawns a writer
|
||||||
# that died from an unexpected escape.
|
# that died from an unexpected escape.
|
||||||
thread = threading.Thread(
|
thread = threading.Thread(
|
||||||
target=self._token_writer_loop, name="session-db-token-writer", daemon=True,
|
target=self._token_writer_loop, name="session-db-token-writer", daemon=True)
|
||||||
)
|
|
||||||
self._token_writer_thread = thread
|
self._token_writer_thread = thread
|
||||||
thread.start()
|
thread.start()
|
||||||
if self._token_atexit_hook is None:
|
if self._token_atexit_hook is None:
|
||||||
@@ -245,9 +253,7 @@ class SessionUsageMixin:
|
|||||||
# Writer stuck mid-apply: leave deltas unapplied rather than race it.
|
# Writer stuck mid-apply: leave deltas unapplied rather than race it.
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"async token accounting: writer did not stop within %.0fs; "
|
"async token accounting: writer did not stop within %.0fs; "
|
||||||
"%d queued delta(s) not persisted",
|
"%d queued delta(s) not persisted", join_timeout, len(self._token_queue))
|
||||||
join_timeout, len(self._token_queue),
|
|
||||||
)
|
|
||||||
return
|
return
|
||||||
# Writer gone: apply leftovers synchronously under the same busy protocol. Wait
|
# Writer gone: apply leftovers synchronously under the same busy protocol. Wait
|
||||||
# out a flush caller-drain that already claimed busy — close() nulls the
|
# out a flush caller-drain that already claimed busy — close() nulls the
|
||||||
@@ -260,8 +266,7 @@ class SessionUsageMixin:
|
|||||||
logger.warning(
|
logger.warning(
|
||||||
"async token accounting: concurrent drain did not "
|
"async token accounting: concurrent drain did not "
|
||||||
"finish within %.0fs; %d queued delta(s) not persisted",
|
"finish within %.0fs; %d queued delta(s) not persisted",
|
||||||
join_timeout, len(self._token_queue),
|
join_timeout, len(self._token_queue))
|
||||||
)
|
|
||||||
return
|
return
|
||||||
self._token_queue_cond.wait(remaining)
|
self._token_queue_cond.wait(remaining)
|
||||||
# busy BEFORE clearing the queue (same ordering as the writer loop).
|
# busy BEFORE clearing the queue (same ordering as the writer loop).
|
||||||
@@ -290,6 +295,7 @@ class SessionUsageMixin:
|
|||||||
"""Update token counters and backfill model if unset. *absolute*=False
|
"""Update token counters and backfill model if unset. *absolute*=False
|
||||||
increments (per-API-call deltas, CLI path); *absolute*=True sets directly
|
increments (per-API-call deltas, CLI path); *absolute*=True sets directly
|
||||||
(gateway path, where the cached agent holds cumulative totals)."""
|
(gateway path, where the cached agent holds cumulative totals)."""
|
||||||
|
usage = {k: v for k, v in locals().items() if k in _MODEL_USAGE_FIELDS}
|
||||||
# Ensure the row exists: under concurrent load create_session() may have failed
|
# Ensure the row exists: under concurrent load create_session() may have failed
|
||||||
# on locking, and the UPDATE would silently affect 0 rows.
|
# on locking, and the UPDATE would silently affect 0 rows.
|
||||||
self._insert_session_row(session_id, "unknown", model=model)
|
self._insert_session_row(session_id, "unknown", model=model)
|
||||||
@@ -305,8 +311,7 @@ class SessionUsageMixin:
|
|||||||
billing_provider if has_accounted_usage else None,
|
billing_provider if has_accounted_usage else None,
|
||||||
billing_base_url if has_accounted_usage else None,
|
billing_base_url if has_accounted_usage else None,
|
||||||
billing_mode if has_accounted_usage else None, model if has_accounted_usage else None,
|
billing_mode if has_accounted_usage else None, model if has_accounted_usage else None,
|
||||||
api_call_count, session_id,
|
api_call_count, session_id)
|
||||||
)
|
|
||||||
# Per-model attribution: the sessions row keeps one (model, provider) pair, so a
|
# Per-model attribution: the sessions row keeps one (model, provider) pair, so a
|
||||||
# mid-session /model switch would attribute every token to the initial model.
|
# mid-session /model switch would attribute every token to the initial model.
|
||||||
# Only the incremental path records here — absolute cumulative updates cannot be
|
# Only the incremental path records here — absolute cumulative updates cannot be
|
||||||
@@ -317,18 +322,14 @@ class SessionUsageMixin:
|
|||||||
row = conn.execute(
|
row = conn.execute(
|
||||||
"SELECT model, billing_provider, api_call_count FROM sessions WHERE id = ?", (session_id,),
|
"SELECT model, billing_provider, api_call_count FROM sessions WHERE id = ?", (session_id,),
|
||||||
).fetchone()
|
).fetchone()
|
||||||
existing_model = row["model"] if row is not None else None
|
existing = dict(row) if row is not None else {}
|
||||||
existing_provider = row["billing_provider"] if row is not None else None
|
|
||||||
existing_api_calls = int((row["api_call_count"] if row is not None else 0) or 0)
|
|
||||||
# create_session records the requested route before any API call. If that
|
# create_session records the requested route before any API call. If that
|
||||||
# fails and fallback succeeds, the first accounted usage is the authoritative
|
# fails and fallback succeeds, the first accounted usage is the authoritative
|
||||||
# route; after that keep the row as is (one row cannot represent mixed usage).
|
# route; after that keep the row as is (one row cannot represent mixed usage).
|
||||||
first_accounted_route = (
|
first_accounted_route = (
|
||||||
existing_api_calls == 0
|
int(existing.get("api_call_count") or 0) == 0 and has_accounted_usage and bool(model)
|
||||||
and has_accounted_usage
|
|
||||||
and bool(model)
|
|
||||||
and bool(billing_provider)
|
and bool(billing_provider)
|
||||||
and (existing_model != model or existing_provider != billing_provider)
|
and (existing.get("model") != model or existing.get("billing_provider") != billing_provider)
|
||||||
)
|
)
|
||||||
if first_accounted_route:
|
if first_accounted_route:
|
||||||
conn.execute(
|
conn.execute(
|
||||||
@@ -340,51 +341,39 @@ class SessionUsageMixin:
|
|||||||
)
|
)
|
||||||
conn.execute(sql, params)
|
conn.execute(sql, params)
|
||||||
if record_model_usage:
|
if record_model_usage:
|
||||||
self._record_model_usage(
|
self._record_model_usage(conn, session_id, **usage)
|
||||||
conn, session_id, model=model, billing_provider=billing_provider,
|
|
||||||
billing_base_url=billing_base_url, billing_mode=billing_mode,
|
|
||||||
input_tokens=input_tokens, output_tokens=output_tokens,
|
|
||||||
cache_read_tokens=cache_read_tokens, cache_write_tokens=cache_write_tokens,
|
|
||||||
reasoning_tokens=reasoning_tokens, estimated_cost_usd=estimated_cost_usd,
|
|
||||||
actual_cost_usd=actual_cost_usd, cost_status=cost_status, cost_source=cost_source,
|
|
||||||
api_call_count=api_call_count,
|
|
||||||
)
|
|
||||||
self._execute_write(_do)
|
self._execute_write(_do)
|
||||||
|
|
||||||
def _record_model_usage(
|
def _record_model_usage(
|
||||||
self, conn, session_id: str, *, model: Optional[str], billing_provider: Optional[str],
|
self, conn, session_id: str, *, model: Optional[str] = None, billing_provider: Optional[str] = None,
|
||||||
billing_base_url: Optional[str], billing_mode: Optional[str], input_tokens: int,
|
billing_base_url: Optional[str] = None, billing_mode: Optional[str] = None, input_tokens: int = 0,
|
||||||
output_tokens: int, cache_read_tokens: int, cache_write_tokens: int, reasoning_tokens: int,
|
output_tokens: int = 0, cache_read_tokens: int = 0, cache_write_tokens: int = 0,
|
||||||
estimated_cost_usd: Optional[float], actual_cost_usd: Optional[float],
|
reasoning_tokens: int = 0, estimated_cost_usd: Optional[float] = None,
|
||||||
cost_status: Optional[str], cost_source: Optional[str], api_call_count: int, task: str = "",
|
actual_cost_usd: Optional[float] = None, cost_status: Optional[str] = None,
|
||||||
|
cost_source: Optional[str] = None, api_call_count: int = 0, task: str = "",
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Accumulate a per-API-call usage delta into session_model_usage, inside the
|
"""Accumulate a per-API-call usage delta into session_model_usage, inside the
|
||||||
caller's write txn after the ``sessions`` UPDATE. A missing model/provider falls
|
caller's write txn after the ``sessions`` UPDATE. A missing model/provider falls
|
||||||
back to the session row (same COALESCE behaviour as the summary update) — except
|
back to the session row (same COALESCE behaviour as the summary update) — except
|
||||||
for aux rows (``task`` set), which must NOT inherit the main-loop route (vision
|
for aux rows (``task`` set), which must NOT inherit the main-loop route (vision
|
||||||
on gemini while the main loop runs anthropic): missing info stays 'unknown'/empty.
|
on gemini while the main loop runs anthropic): missing info stays 'unknown'/empty."""
|
||||||
"""
|
|
||||||
row = conn.execute(
|
row = conn.execute(
|
||||||
"SELECT model, billing_provider, billing_base_url, billing_mode "
|
"SELECT model, billing_provider, billing_base_url, billing_mode "
|
||||||
"FROM sessions WHERE id = ?",
|
"FROM sessions WHERE id = ?", (session_id,),
|
||||||
(session_id,),
|
|
||||||
).fetchone()
|
).fetchone()
|
||||||
sess = dict(row) if (row is not None and not task) else {}
|
sess = dict(row) if (row is not None and not task) else {}
|
||||||
eff_model = model or sess.get("model") or "unknown"
|
|
||||||
eff_provider = billing_provider or sess.get("billing_provider") or ""
|
|
||||||
eff_base_url = billing_base_url or sess.get("billing_base_url") or ""
|
|
||||||
eff_billing_mode = billing_mode or sess.get("billing_mode") or ""
|
|
||||||
counts = [v or 0 for v in (input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens)]
|
counts = [v or 0 for v in (input_tokens, output_tokens, cache_read_tokens, cache_write_tokens, reasoning_tokens)]
|
||||||
now = time.time()
|
now = time.time()
|
||||||
conn.execute(
|
conn.execute(
|
||||||
_MODEL_USAGE_UPSERT_SQL,
|
_MODEL_USAGE_UPSERT_SQL,
|
||||||
(
|
(
|
||||||
session_id, eff_model, eff_provider, eff_base_url, eff_billing_mode, task or "",
|
session_id, model or sess.get("model") or "unknown",
|
||||||
|
billing_provider or sess.get("billing_provider") or "",
|
||||||
|
billing_base_url or sess.get("billing_base_url") or "",
|
||||||
|
billing_mode or sess.get("billing_mode") or "", task or "",
|
||||||
api_call_count or 0, *counts,
|
api_call_count or 0, *counts,
|
||||||
float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0),
|
float(estimated_cost_usd or 0.0), float(actual_cost_usd or 0.0),
|
||||||
cost_status, cost_source, now, now,
|
cost_status, cost_source, now, now))
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
def record_auxiliary_usage(
|
def record_auxiliary_usage(
|
||||||
self, session_id: str, task: str, *, model: Optional[str] = None,
|
self, session_id: str, task: str, *, model: Optional[str] = None,
|
||||||
@@ -398,22 +387,13 @@ class SessionUsageMixin:
|
|||||||
touching the ``sessions`` summary row (the gateway overwrites those counters with
|
touching the ``sessions`` summary row (the gateway overwrites those counters with
|
||||||
absolute main-loop totals). ``api_call_count`` may aggregate N calls. Best-effort:
|
absolute main-loop totals). ``api_call_count`` may aggregate N calls. Best-effort:
|
||||||
callers must never fail an aux call over accounting."""
|
callers must never fail an aux call over accounting."""
|
||||||
|
usage = {k: v for k, v in locals().items() if k in _MODEL_USAGE_FIELDS}
|
||||||
if not session_id or not task:
|
if not session_id or not task:
|
||||||
return
|
return
|
||||||
|
usage["api_call_count"] = 1 if api_call_count is None else int(api_call_count)
|
||||||
# FK to sessions.id: same INSERT OR IGNORE guard as update_token_counts.
|
# FK to sessions.id: same INSERT OR IGNORE guard as update_token_counts.
|
||||||
self._insert_session_row(session_id, "unknown")
|
self._insert_session_row(session_id, "unknown")
|
||||||
|
self._execute_write(lambda conn: self._record_model_usage(conn, session_id, task=task, **usage))
|
||||||
def _do(conn):
|
|
||||||
self._record_model_usage(
|
|
||||||
conn, session_id, model=model, billing_provider=billing_provider,
|
|
||||||
billing_base_url=billing_base_url, billing_mode=None,
|
|
||||||
input_tokens=input_tokens or 0, output_tokens=output_tokens or 0,
|
|
||||||
cache_read_tokens=cache_read_tokens or 0, cache_write_tokens=cache_write_tokens or 0,
|
|
||||||
reasoning_tokens=reasoning_tokens or 0, estimated_cost_usd=estimated_cost_usd,
|
|
||||||
actual_cost_usd=None, cost_status=None, cost_source=None,
|
|
||||||
api_call_count=1 if api_call_count is None else int(api_call_count), task=task,
|
|
||||||
)
|
|
||||||
self._execute_write(_do)
|
|
||||||
|
|
||||||
def usage_totals(self, *, min_message_count: int = 1, include_archived: bool = False) -> Dict[str, float]:
|
def usage_totals(self, *, min_message_count: int = 1, include_archived: bool = False) -> Dict[str, float]:
|
||||||
"""Tokens and spend across the whole store (one scan), so the sidebar total does
|
"""Tokens and spend across the whole store (one scan), so the sidebar total does
|
||||||
|
|||||||
Reference in New Issue
Block a user