diff --git a/loopx/chat_event_cache.py b/loopx/chat_event_cache.py new file mode 100644 index 0000000000..0e158e2dd2 --- /dev/null +++ b/loopx/chat_event_cache.py @@ -0,0 +1,52 @@ +"""Revision-bound chat replay snapshots; file I/O and locking stay in the store.""" +from __future__ import annotations + +import threading +from typing import Any + +EventKey = tuple[str, str] +EventRevision = tuple[int, int, int] | None +EventRows = list[dict[str, Any]] + +# These bound terminal retention, not active history or peak decoding memory. +TERMINAL_CACHE_MAX_TURNS = 8 +TERMINAL_CACHE_MAX_ROWS = 4096 +TERMINAL_CACHE_MAX_ENCODED_BYTES = 2 * 1024 * 1024 + + +class ChatEventCache: + def __init__(self, lock: threading.RLock) -> None: + self._lock = lock + self._entries: dict[EventKey, tuple[EventRevision, EventRows]] = {} + self._finished: dict[EventKey, tuple[int, int]] = {} + + def get(self, key: EventKey, revision: EventRevision) -> EventRows | None: + with self._lock: + entry = self._entries.get(key) + return entry[1] if entry is not None and entry[0] == revision else None + + def put(self, key: EventKey, revision: EventRevision, rows: EventRows) -> None: + with self._lock: + self._finished.pop(key, None) + self._entries[key] = (revision, rows) + + def drop(self, key: EventKey) -> None: + with self._lock: + self._entries.pop(key, None) + self._finished.pop(key, None) + + def retain_terminal(self, key: EventKey, rows: EventRows) -> None: + with self._lock: + entry = self._entries.get(key) + # An append or compaction can replace a snapshot while it is replayed. + if entry is None or entry[1] is not rows: + return + revision = entry[0] + self._finished.pop(key, None) + self._finished[key] = (len(rows), revision[1] if revision else 0) + while ( + len(self._finished) > TERMINAL_CACHE_MAX_TURNS + or sum(weight[0] for weight in self._finished.values()) > TERMINAL_CACHE_MAX_ROWS + or sum(weight[1] for weight in self._finished.values()) > TERMINAL_CACHE_MAX_ENCODED_BYTES + ): + self.drop(next(iter(self._finished))) diff --git a/loopx/chat_store.py b/loopx/chat_store.py index 7ead472c8b..9a14330f17 100644 --- a/loopx/chat_store.py +++ b/loopx/chat_store.py @@ -14,6 +14,7 @@ from weakref import WeakValueDictionary from .chat import require_matching_replay, resolve_attached_completion_replay +from .chat_event_cache import ChatEventCache from .file_lock import exclusive_file_lock @@ -143,8 +144,7 @@ def __init__(self, runtime_root: Path) -> None: self._session_lock_guard = threading.Lock() self._session_locks: WeakValueDictionary[str, threading.Lock] = WeakValueDictionary() self._event_lock = threading.RLock() - self._event_cache: dict[tuple[str, str], list[dict[str, Any]]] = {} - self._event_cache_revision: dict[tuple[str, str], tuple[int, int, int] | None] = {} + self._event_cache = ChatEventCache(self._event_lock) self._event_pending: dict[tuple[str, str], list[dict[str, Any]]] = {} self._event_flush_locks: WeakValueDictionary[tuple[str, str], threading.Lock] = WeakValueDictionary() self.sessions_root.mkdir(parents=True, exist_ok=True, mode=0o700) @@ -1294,13 +1294,10 @@ def _event_rows_locked(self, session_id: str, turn_id: str) -> list[dict[str, An key = (session_id, turn_id) path = self._event_path(session_id, turn_id) revision = self._event_revision(path) - with self._event_lock: - if self._event_cache_revision.get(key) == revision and key in self._event_cache: - return self._event_cache[key] - rows = _read_jsonl(path) - with self._event_lock: - self._event_cache[key] = rows - self._event_cache_revision[key] = revision + rows = self._event_cache.get(key, revision) + if rows is None: + rows = _read_jsonl(path) + self._event_cache.put(key, revision, rows) return rows @staticmethod @@ -1315,12 +1312,6 @@ def _event_flush_lock(self, key: tuple[str, str]) -> threading.Lock: with self._event_lock: return self._event_flush_locks.setdefault(key, threading.Lock()) - def _drop_event_cache(self, key: tuple[str, str], expected: list[dict[str, Any]] | None = None) -> None: - with self._event_lock: - if expected is None or self._event_cache.get(key) is expected: - self._event_cache.pop(key, None) - self._event_cache_revision.pop(key, None) - def append_event( self, session_id: str, @@ -1368,11 +1359,9 @@ def flush_events(self, session_id: str, turn_id: str) -> int: event["sequence"] = sequence _append_jsonl_rows(path, pending) if any(row["kind"] in TERMINAL_EVENT_KINDS for row in pending): - self._drop_event_cache(key) + self._event_cache.drop(key) else: - with self._event_lock: - self._event_cache[key] = [*rows, *pending] - self._event_cache_revision[key] = self._event_revision(path) + self._event_cache.put(key, self._event_revision(path), [*rows, *pending]) except Exception: with self._event_lock: later = self._event_pending.get(key, []) @@ -1389,9 +1378,7 @@ def events_after(self, session_id: str, turn_id: str, event_id: str | None) -> l key = (session_id, turn_id) path = self._event_path(session_id, turn_id) revision = self._event_revision(path) - with self._event_lock: - cached = self._event_cache.get(key) - rows = cached if cached is not None and self._event_cache_revision.get(key) == revision else None + rows = self._event_cache.get(key, revision) if rows is None: with exclusive_file_lock(path, agent_id="loopx-chat", operation="read_chat_events"): rows = self._event_rows_locked(session_id, turn_id) @@ -1400,7 +1387,7 @@ def events_after(self, session_id: str, turn_id: str, event_id: str | None) -> l start = bisect_right(rows, after, key=lambda row: int(row.get("sequence") or 0)) result = rows[start:] if rows and rows[-1].get("kind") in TERMINAL_EVENT_KINDS: - self._drop_event_cache(key, rows) + self._event_cache.retain_terminal(key, rows) return result def compact_completed_events(self, *, older_than_hours: float = 24.0) -> int: @@ -1432,7 +1419,7 @@ def compact_completed_events(self, *, older_than_hours: float = 24.0) -> int: _replace_jsonl(event_path, retained) compacted += 1 revision = self._event_revision(event_path) - self._drop_event_cache((session_id, turn_id)) + self._event_cache.drop((session_id, turn_id)) with exclusive_file_lock(turn_path, agent_id="loopx-chat", operation="mark_chat_events_compacted"): current = _read_json(turn_path) if current.get("status") in TERMINAL_TURN_STATES: diff --git a/tests/test_chat_event_cursor.py b/tests/test_chat_event_cursor.py index 4d8229f037..b0e61bd9c2 100644 --- a/tests/test_chat_event_cursor.py +++ b/tests/test_chat_event_cursor.py @@ -17,14 +17,13 @@ def __iter__(self): event_path.parent.mkdir(parents=True) event_path.touch() key = (session_id, turn_id) - store._event_cache[key] = NonIterableRows( + store._event_cache.put(key, store._event_revision(event_path), NonIterableRows( [ {"sequence": 1, "event_id": "1"}, {"sequence": 4, "event_id": "4"}, {"sequence": 7, "event_id": "7"}, ] - ) - store._event_cache_revision[key] = store._event_revision(event_path) + )) assert store.events_after(session_id, turn_id, "4") == [ {"sequence": 7, "event_id": "7"} diff --git a/tests/test_chat_event_retention.py b/tests/test_chat_event_retention.py index 811415ac0a..5235d12d0f 100644 --- a/tests/test_chat_event_retention.py +++ b/tests/test_chat_event_retention.py @@ -5,6 +5,8 @@ import json from pathlib import Path +import pytest + import loopx.chat_store as chat_store from loopx.chat_store import CHAT_TURN_SCHEMA_VERSION, ChatSessionStore @@ -39,27 +41,76 @@ def _write_completed_turn(root: Path, *, with_events: bool = True) -> tuple[Path return turn_path, event_path -def test_terminal_event_history_does_not_remain_in_the_hot_cache( - tmp_path: Path, -) -> None: +def test_completed_replay_reuses_log_until_an_external_append(tmp_path: Path, monkeypatch) -> None: store = ChatSessionStore(tmp_path) key = ("session", "turn") - store.append_event( - *key, - kind="assistant.delta", - payload={"text": "visible delta"}, - buffered=True, - ) + store.append_event(*key, kind="assistant.delta", payload={"text": "visible"}, buffered=True) store.append_event(*key, kind="turn.completed", payload={}) + assert key not in store._event_cache._entries + reads = [] + read = chat_store._read_jsonl + + def counted(path): + reads.append(path) + return read(path) - assert key not in store._event_cache - assert key not in store._event_cache_revision - assert [row["kind"] for row in store.events_after(*key, None)] == [ - "assistant.delta", - "turn.completed", - ] - assert key not in store._event_cache - assert key not in store._event_cache_revision + monkeypatch.setattr(chat_store, "_read_jsonl", counted) + for _ in range(5): + assert [row["sequence"] for row in store.events_after(*key, "0")] == [1, 2] + assert len(reads) == 1 + other = ChatSessionStore(tmp_path) + other.append_event(*key, kind="turn.completed", payload={"recovered": True}) + assert [row["sequence"] for row in store.events_after(*key, "2")] == [3] + assert len(reads) == 3 # independent writer plus invalidated reader + + +def test_completed_replay_evicts_least_recently_used_log(tmp_path: Path) -> None: + store = ChatSessionStore(tmp_path) + for i in range(8): + store.append_event("session", str(i), kind="turn.completed", payload={}) + store.events_after("session", str(i), None) + store.events_after("session", "0", None) + store.append_event("session", "8", kind="turn.completed", payload={}) + store.events_after("session", "8", None) + assert ("session", "0") in store._event_cache._entries + assert ("session", "1") not in store._event_cache._entries + assert len(store._event_cache._finished) == 8 + + +@pytest.mark.parametrize("row_count,text_size", [(1, 2 * 1024 * 1024), (4096, 1)]) +def test_oversized_completed_log_is_replayable_but_not_retained( + tmp_path: Path, row_count: int, text_size: int, +) -> None: + store = ChatSessionStore(tmp_path) + key = ("session", "large") + for _ in range(row_count): + store.append_event(*key, kind="assistant.delta", payload={"text": "x" * text_size}, buffered=True) + store.append_event(*key, kind="turn.completed", payload={}) + assert len(store.events_after(*key, None)) == row_count + 1 + assert key not in store._event_cache._entries + assert key not in store._event_cache._finished + + +def test_completed_replay_budget_is_aggregate(tmp_path: Path) -> None: + store = ChatSessionStore(tmp_path) + for turn in ("a", "b", "c"): + store.append_event("session", turn, kind="assistant.delta", payload={"text": "x" * 800_000}, buffered=True) + store.append_event("session", turn, kind="turn.completed", payload={}) + store.events_after("session", turn, None) + assert ("session", "a") not in store._event_cache._entries + assert set(store._event_cache._finished) == {("session", "b"), ("session", "c")} + + +def test_old_replay_cannot_retain_a_replaced_snapshot(tmp_path: Path) -> None: + store = ChatSessionStore(tmp_path) + key = ("session", "turn") + store.append_event(*key, kind="turn.completed", payload={}) + store.events_after(*key, None) + old = store._event_cache._entries[key][1] + store.append_event(*key, kind="assistant.delta", payload={"text": "continued"}) + store._event_cache.retain_terminal(key, old) + assert key not in store._event_cache._finished + assert store.events_after(*key, "1")[0]["kind"] == "assistant.delta" def test_keyed_locks_are_released_after_callers_drop_them(tmp_path: Path) -> None: @@ -88,7 +139,7 @@ def test_completed_event_compaction_is_skipped_until_the_file_changes( assert turn["event_compaction_revision"] == list( first._event_revision(event_path) or () ) - assert ("session", "turn") not in first._event_cache + assert ("session", "turn") not in first._event_cache._entries original_read_jsonl = chat_store._read_jsonl @@ -100,7 +151,7 @@ def reject_redundant_event_read(path: Path): monkeypatch.setattr(chat_store, "_read_jsonl", reject_redundant_event_read) restarted = ChatSessionStore(tmp_path) - assert not restarted._event_cache + assert not restarted._event_cache._entries assert not restarted._event_flush_locks @@ -123,7 +174,7 @@ def test_completed_event_compaction_marker_is_invalidated_by_a_new_event( assert turn["event_compaction_revision"] == list( restarted._event_revision(event_path) or () ) - assert ("session", "turn") not in restarted._event_cache + assert ("session", "turn") not in restarted._event_cache._entries def test_missing_event_stream_is_marked_without_retaining_empty_state( @@ -135,7 +186,7 @@ def test_missing_event_stream_is_marked_without_retaining_empty_state( turn = json.loads(turn_path.read_text(encoding="utf-8")) assert turn["event_compaction_revision"] == [] - assert not store._event_cache + assert not store._event_cache._entries assert not store._event_flush_locks