Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 34 additions & 5 deletions __init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,11 +156,40 @@ def register(ctx):

def _on_post_llm_call(**kwargs):
history = kwargs.get("conversation_history")
if history:
try:
engine.ingest(history)
except Exception as exc:
logger.debug("LCM post_llm_call ingest error: %s", exc)
if not history:
return
active_engine = kwargs.get("context_compressor")
if not (
active_engine is not None
and getattr(active_engine, "name", None) == "lcm"
and hasattr(active_engine, "ingest")
):
active_engine = engine

session_id = str(kwargs.get("session_id") or "")
conversation_id = str(
kwargs.get("conversation_id")
or kwargs.get("gateway_session_key")
or ""
)
platform = str(kwargs.get("platform") or "")

try:
if session_id and (
str(getattr(active_engine, "current_session_id", "") or "") != session_id
or (
conversation_id
and str(getattr(active_engine, "current_conversation_id", "") or "") != conversation_id
)
):
active_engine.on_session_start(
session_id,
platform=platform,
conversation_id=conversation_id or None,
)
active_engine.ingest(history)
except Exception as exc:
logger.debug("LCM post_llm_call ingest error: %s", exc)

_mgr._hooks.setdefault("post_llm_call", []).append(_on_post_llm_call)
logger.debug("LCM registered post_llm_call hook for per-turn ingest")
Expand Down
23 changes: 22 additions & 1 deletion db_bootstrap.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@

logger = logging.getLogger(__name__)

SCHEMA_VERSION = 4
SCHEMA_VERSION = 5
SQLITE_BUSY_TIMEOUT_MS = 30_000
_MIN_DISK_SPACE_BYTES = 50 * 1024 * 1024
REQUIRED_CORE_TABLES = (
Expand Down Expand Up @@ -175,6 +175,22 @@ def ensure_lifecycle_state_columns(conn: sqlite3.Connection) -> None:
conn.execute("ALTER TABLE lcm_lifecycle_state ADD COLUMN last_reset_at REAL")


def ensure_message_origin_columns(conn: sqlite3.Connection) -> None:
table_row = conn.execute(
"SELECT name FROM sqlite_master WHERE type='table' AND name='messages'"
).fetchone()
if not table_row:
return
columns = {
row[1] for row in conn.execute("PRAGMA table_info(messages)").fetchall()
}
if "conversation_id" not in columns:
conn.execute("ALTER TABLE messages ADD COLUMN conversation_id TEXT DEFAULT ''")
conn.execute(
"CREATE INDEX IF NOT EXISTS idx_msg_conversation_session ON messages(conversation_id, session_id, store_id)"
)


def mark_migration_step_complete(conn: sqlite3.Connection, step_name: str) -> None:
ensure_migration_state_table(conn)
conn.execute(
Expand Down Expand Up @@ -598,4 +614,9 @@ def run_versioned_migrations(conn: sqlite3.Connection) -> None:
mark_migration_step_complete(conn, "v4_lifecycle_debt_columns")
current_version = 4

ensure_message_origin_columns(conn)
if current_version < 5:
mark_migration_step_complete(conn, "v5_message_conversation_id")
current_version = 5

set_schema_version(conn, current_version)
1 change: 1 addition & 0 deletions engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -3585,6 +3585,7 @@ def _ingest_messages(self, messages: List[Dict[str, Any]]) -> List[Dict[str, Any
protected_messages,
estimates,
source=self._session_platform,
conversation_id=self._conversation_id,
)
self._ingest_cursor = n
logger.debug("Ingested %d messages into LCM store", len(messages_to_store_with_index))
Expand Down
7 changes: 7 additions & 0 deletions schemas.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,13 @@
"Use 'unknown' for explicit unknown-source content."
),
},
"conversation_id": {
"type": "string",
"description": (
"Optional gateway conversation/session key filter for lane-scoped retrieval. "
"Use this to restrict Discord searches to one channel/thread/forum topic lane when rows carry metadata."
),
},
"role": {
"type": "string",
"enum": ["system", "user", "assistant", "tool", "unknown"],
Expand Down
65 changes: 54 additions & 11 deletions store.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@
_MESSAGE_ROLE_BIAS_SQL = "CASE m.role WHEN 'user' THEN 0 WHEN 'assistant' THEN 1 WHEN 'tool' THEN 2 ELSE 1 END"
_MESSAGE_SELECT_COLUMNS = (
"store_id, session_id, source, role, content, tool_call_id, "
"tool_calls, tool_name, timestamp, token_estimate, pinned"
"tool_calls, tool_name, timestamp, token_estimate, pinned, conversation_id"
)
_UNKNOWN_SOURCE = "unknown"

Expand All @@ -71,6 +71,10 @@ def _normalize_source_value(source: str | None) -> str:
return normalized or _UNKNOWN_SOURCE


def _normalize_conversation_id_value(conversation_id: str | None) -> str:
return (conversation_id or "").strip()


def _source_filter_clause(column: str, source: str | None) -> tuple[str | None, list[str]]:
normalized = _normalize_source_value(source) if source is not None else ""
if not normalized:
Expand All @@ -80,6 +84,13 @@ def _source_filter_clause(column: str, source: str | None) -> tuple[str | None,
return f"{column} = ?", [normalized]


def _conversation_filter_clause(column: str, conversation_id: str | None) -> tuple[str | None, list[str]]:
normalized = _normalize_conversation_id_value(conversation_id)
if not normalized:
return None, []
return f"{column} = ?", [normalized]


def _message_role_bias(role: str | None) -> float:
if role == "user":
return 0.0
Expand Down Expand Up @@ -240,6 +251,7 @@ def _init_db(self):
store_id INTEGER PRIMARY KEY AUTOINCREMENT,
session_id TEXT NOT NULL,
source TEXT DEFAULT '',
conversation_id TEXT DEFAULT '',
role TEXT NOT NULL,
content TEXT,
tool_call_id TEXT,
Expand All @@ -265,6 +277,7 @@ def _init_db(self):
)
run_versioned_migrations(self._conn)
self._ensure_source_column()
self._ensure_conversation_id_column()
self._conn.commit()

def _ensure_source_column(self) -> None:
Expand All @@ -277,10 +290,21 @@ def _ensure_source_column(self) -> None:
"CREATE INDEX IF NOT EXISTS idx_msg_source_session ON messages(source, session_id, store_id)"
)

def _ensure_conversation_id_column(self) -> None:
columns = {
row[1] for row in self._conn.execute("PRAGMA table_info(messages)").fetchall()
}
if "conversation_id" not in columns:
self._conn.execute("ALTER TABLE messages ADD COLUMN conversation_id TEXT DEFAULT ''")
self._conn.execute(
"CREATE INDEX IF NOT EXISTS idx_msg_conversation_session ON messages(conversation_id, session_id, store_id)"
)

# -- Write operations ---------------------------------------------------

def append(self, session_id: str, msg: Dict[str, Any],
token_estimate: int = 0, source: str = "") -> int:
token_estimate: int = 0, source: str = "",
conversation_id: str = "") -> int:
"""Persist a message and return its store_id."""
msg = protect_message_for_ingest(
msg,
Expand All @@ -294,12 +318,13 @@ def append(self, session_id: str, msg: Dict[str, Any],
with self._write_lock:
cur = self._conn.execute(
"""INSERT INTO messages
(session_id, source, role, content, tool_call_id, tool_calls,
(session_id, source, conversation_id, role, content, tool_call_id, tool_calls,
tool_name, timestamp, token_estimate, pinned)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
(
session_id,
_normalize_source_value(source),
_normalize_conversation_id_value(conversation_id),
msg.get("role", "unknown"),
_normalize_content_value(msg.get("content")),
msg.get("tool_call_id"),
Expand All @@ -316,7 +341,8 @@ def append(self, session_id: str, msg: Dict[str, Any],
def append_batch(self, session_id: str,
messages: List[Dict[str, Any]],
token_estimates: List[int] | None = None,
source: str = "") -> List[int]:
source: str = "",
conversation_id: str = "") -> List[int]:
"""Persist multiple messages in one transaction. Returns store_ids."""
protected_messages = protect_messages_for_ingest(
messages,
Expand All @@ -329,12 +355,14 @@ def append_batch(self, session_id: str,
protected_messages,
token_estimates,
source=source,
conversation_id=conversation_id,
)

def _append_protected_batch(self, session_id: str,
messages: List[Dict[str, Any]],
token_estimates: List[int] | None = None,
source: str = "") -> List[int]:
source: str = "",
conversation_id: str = "") -> List[int]:
"""Persist messages that already passed ingest protection.

This is an internal fast path for callers that need the protected form
Expand All @@ -353,12 +381,13 @@ def _append_protected_batch(self, session_id: str,
ts = time.time()
cur = self._conn.execute(
"""INSERT INTO messages
(session_id, source, role, content, tool_call_id, tool_calls,
(session_id, source, conversation_id, role, content, tool_call_id, tool_calls,
tool_name, timestamp, token_estimate, pinned)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""",
(
session_id,
_normalize_source_value(source),
_normalize_conversation_id_value(conversation_id),
msg.get("role", "unknown"),
_normalize_content_value(msg.get("content")),
msg.get("tool_call_id"),
Expand Down Expand Up @@ -700,6 +729,7 @@ def get_time_bounds(self, store_ids: List[int]) -> tuple[float | None, float | N
def search(self, query: str, session_id: str | None = None,
limit: int = 20, sort: str | None = None,
source: str | None = None,
conversation_id: str | None = None,
role: str | None = None,
time_from: float | None = None,
time_to: float | None = None) -> List[Dict[str, Any]]:
Expand All @@ -712,6 +742,7 @@ def search(self, query: str, session_id: str | None = None,
- ``source`` limits which raw rows inside those sessions are eligible
- ``source='unknown'`` means the explicit unknown-source bucket, with
legacy blank-source rows treated as equivalent for back-compat
- ``conversation_id`` limits rows to one gateway conversation/session key
"""
safe_query = sanitize_fts5_query(query)
terms = extract_search_terms(safe_query)
Expand All @@ -723,6 +754,7 @@ def search(self, query: str, session_id: str | None = None,
limit=limit,
sort=sort,
source=source,
conversation_id=conversation_id,
role=role,
time_from=time_from,
time_to=time_to,
Expand All @@ -738,6 +770,7 @@ def search(self, query: str, session_id: str | None = None,
apply_directness_adjustment = should_apply_directness_rank_adjustment(terms, phrases)
max_rank_bonus = compute_directness_rank_bonus_upper_bound(terms, phrases) * 3e-7
source_clause, source_args = _source_filter_clause("m.source", source)
conversation_clause, conversation_args = _conversation_filter_clause("m.conversation_id", conversation_id)
offset = 0
scanned_rows = 0
results: list[Dict[str, Any]] = []
Expand All @@ -751,6 +784,9 @@ def search(self, query: str, session_id: str | None = None,
if source_clause:
where.append(source_clause)
args.extend(source_args)
if conversation_clause:
where.append(conversation_clause)
args.extend(conversation_args)
if role is not None:
where.append("m.role = ?")
args.append(role)
Expand All @@ -763,7 +799,7 @@ def search(self, query: str, session_id: str | None = None,
args.extend([fetch_limit, offset])
rows = self._conn.execute(
f"""SELECT m.store_id, m.session_id, m.source, m.role, m.content, m.tool_call_id,
m.tool_calls, m.tool_name, m.timestamp, m.token_estimate, m.pinned,
m.tool_calls, m.tool_name, m.timestamp, m.token_estimate, m.pinned, m.conversation_id,
rank as search_rank,
snippet(messages_fts, 0, '>>>', '<<<', '...', 40) as snippet
FROM messages_fts fts
Expand All @@ -781,6 +817,7 @@ def search(self, query: str, session_id: str | None = None,
limit=limit,
sort=sort,
source=source,
conversation_id=conversation_id,
role=role,
time_from=time_from,
time_to=time_to,
Expand All @@ -789,7 +826,7 @@ def search(self, query: str, session_id: str | None = None,
raw_primary_values: list[float] = []
for r in rows:
d = self._row_to_dict(r)
base_columns = 11
base_columns = 12
d["search_rank"] = r[base_columns] if len(r) > base_columns else None
d["snippet"] = r[base_columns + 1] if len(r) > (base_columns + 1) else ""
d["_directness_score"] = _message_directness_score(d.get("role"), d.get("content"), terms, phrases)
Expand Down Expand Up @@ -821,6 +858,7 @@ def search(self, query: str, session_id: str | None = None,
def _search_like(self, query: str, session_id: str | None = None,
limit: int = 20, sort: str | None = None,
source: str | None = None,
conversation_id: str | None = None,
role: str | None = None,
time_from: float | None = None,
time_to: float | None = None) -> List[Dict[str, Any]]:
Expand All @@ -840,6 +878,10 @@ def _search_like(self, query: str, session_id: str | None = None,
if source_clause:
where.append(source_clause)
args.extend(source_args)
conversation_clause, conversation_args = _conversation_filter_clause("conversation_id", conversation_id)
if conversation_clause:
where.append(conversation_clause)
args.extend(conversation_args)
if role is not None:
where.append("role = ?")
args.append(role)
Expand Down Expand Up @@ -1017,10 +1059,11 @@ def _row_to_dict(self, row) -> Dict[str, Any]:
return {}
cols = [
"store_id", "session_id", "source", "role", "content", "tool_call_id",
"tool_calls", "tool_name", "timestamp", "token_estimate", "pinned",
"tool_calls", "tool_name", "timestamp", "token_estimate", "pinned", "conversation_id",
]
d = dict(zip(cols, row[:len(cols)]))
d["source"] = _normalize_source_value(d.get("source"))
d["conversation_id"] = _normalize_conversation_id_value(d.get("conversation_id"))
# Deserialize tool_calls JSON
if d.get("tool_calls"):
try:
Expand Down
Loading