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
52 changes: 52 additions & 0 deletions loopx/chat_event_cache.py
Original file line number Diff line number Diff line change
@@ -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)))
35 changes: 11 additions & 24 deletions loopx/chat_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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, [])
Expand All @@ -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)
Expand All @@ -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:
Expand Down Expand Up @@ -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:
Expand Down
5 changes: 2 additions & 3 deletions tests/test_chat_event_cursor.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"}
Expand Down
93 changes: 72 additions & 21 deletions tests/test_chat_event_retention.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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

Expand All @@ -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


Expand All @@ -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(
Expand All @@ -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


Expand Down