From ad6900e56c7605629b61998414439d8965c1c29d Mon Sep 17 00:00:00 2001 From: chengyongru <61816729+chengyongru@users.noreply.github.com> Date: Thu, 30 Jul 2026 11:51:13 +0800 Subject: [PATCH] refactor(session): separate persistence behind SessionStore (#5170) --- nanobot/session/manager.py | 899 ++++++++++++++++------------ tests/session/test_session_store.py | 79 +++ 2 files changed, 582 insertions(+), 396 deletions(-) create mode 100644 tests/session/test_session_store.py diff --git a/nanobot/session/manager.py b/nanobot/session/manager.py index 1d8c57dac..b26b2380c 100644 --- a/nanobot/session/manager.py +++ b/nanobot/session/manager.py @@ -11,7 +11,7 @@ from copy import deepcopy from dataclasses import dataclass, field from datetime import datetime from pathlib import Path -from typing import Any, Callable, cast +from typing import Any, Callable, Protocol, TypedDict, cast from weakref import WeakValueDictionary from loguru import logger @@ -427,17 +427,492 @@ class Session: ) -class SessionManager: - """ - Manages conversation sessions. +class SessionPayload(TypedDict): + key: str + created_at: str | None + updated_at: str | None + metadata: dict[str, Any] + messages: list[dict[str, Any]] - Sessions are stored as JSONL files in the sessions directory. - """ + +class SessionMetadataPayload(TypedDict): + key: str + created_at: str | None + updated_at: str | None + metadata: dict[str, Any] + + +class SessionInfo(TypedDict): + key: str + created_at: str + updated_at: str + title: str + preview: str + path: str + + +class SessionStore(Protocol): + def load(self, key: str) -> Session | None: ... + + def save(self, session: Session, *, fsync: bool = False) -> None: ... + + def delete(self, key: str) -> bool: ... + + def read(self, key: str) -> SessionPayload | None: ... + + def read_metadata(self, key: str) -> SessionMetadataPayload | None: ... + + def list_sessions(self) -> list[SessionInfo]: ... + + +class JsonlSessionStore: + """JSONL implementation of session persistence.""" def __init__(self, workspace: Path): - self.workspace = workspace - self.sessions_dir = ensure_dir(self.workspace / "sessions") + self.sessions_dir = ensure_dir(workspace / "sessions") self.legacy_sessions_dir = get_legacy_sessions_dir() + + @staticmethod + def safe_key(key: str) -> str: + return safe_filename(key.replace(":", "_")) + + @staticmethod + def storage_key(key: str) -> str: + return base64.urlsafe_b64encode(key.encode()).decode().rstrip("=") + + @staticmethod + def decode_storage_key(stem: str) -> str | None: + try: + padding = 4 - len(stem) % 4 + if padding != 4: + stem += "=" * padding + return base64.urlsafe_b64decode(stem).decode("utf-8") + except _SESSION_DATA_ERRORS: + return None + + @classmethod + def session_key_from_path(cls, path: Path) -> str | None: + key = cls.decode_storage_key(path.stem) + if key is None or cls.storage_key(key) != path.stem: + return None + return key + + def get_session_path(self, key: str) -> Path: + return self.sessions_dir / f"{self.storage_key(key)}.jsonl" + + def get_legacy_lossy_path(self, key: str) -> Path: + return self.sessions_dir / f"{safe_filename(key.replace(':', '_'))}.jsonl" + + def get_legacy_session_path(self, key: str) -> Path: + return self.legacy_sessions_dir / f"{self.safe_key(key)}.jsonl" + + def load(self, key: str) -> Session | None: + path = self.get_session_path(key) + if not path.exists(): + return None + + try: + messages: list[dict[str, Any]] = [] + metadata: dict[str, Any] = {} + created_at: datetime | None = None + updated_at: datetime | None = None + last_consolidated = 0 + + with open(path, encoding="utf-8") as f: + for line in f: + line = line.strip() + if not line: + continue + + raw_data: object = json.loads(line) + data = _json_object(raw_data) + + if data.get("_type") == "metadata": + metadata_value = cast(object, data.get("metadata", {})) + metadata = ( + cast(dict[str, Any], metadata_value) + if isinstance(metadata_value, dict) + else {} + ) + created_at_value = cast(object, data.get("created_at")) + updated_at_value = cast(object, data.get("updated_at")) + created_at = ( + datetime.fromisoformat(created_at_value) + if isinstance(created_at_value, str) and created_at_value + else None + ) + updated_at = ( + datetime.fromisoformat(updated_at_value) + if isinstance(updated_at_value, str) and updated_at_value + else None + ) + offset = cast(object, data.get("last_consolidated", 0)) + last_consolidated = ( + offset + if isinstance(offset, int) and not isinstance(offset, bool) + else 0 + ) + else: + messages.append(data) + + return Session( + key=key, + messages=messages, + created_at=created_at or datetime.now(), + updated_at=updated_at or datetime.now(), + metadata=metadata, + last_consolidated=last_consolidated, + ) + except _SESSION_DATA_ERRORS as e: + logger.warning("Failed to load session {}: {}", key, e) + repaired = self.repair(key) + if repaired is not None: + logger.info( + "Recovered session {} from corrupt file ({} messages)", + key, + len(repaired.messages), + ) + return repaired + + def repair(self, key: str, *, path: Path | None = None) -> Session | None: + if path is None: + path = self.get_session_path(key) + if not path.exists(): + return None + + try: + messages: list[dict[str, Any]] = [] + metadata: dict[str, Any] = {} + created_at: datetime | None = None + updated_at: datetime | None = None + last_consolidated = 0 + skipped = 0 + + with open(path, encoding="utf-8") as f: + for line in f: + line = line.strip() + if not line: + continue + try: + raw_data: object = json.loads(line) + except json.JSONDecodeError: + skipped += 1 + continue + if not isinstance(raw_data, dict): + skipped += 1 + continue + data = cast(dict[str, Any], raw_data) + + if data.get("_type") == "metadata": + metadata_value = cast(object, data.get("metadata", {})) + metadata = ( + cast(dict[str, Any], metadata_value) + if isinstance(metadata_value, dict) + else {} + ) + created_at_value = cast(object, data.get("created_at")) + if isinstance(created_at_value, str) and created_at_value: + with suppress(ValueError): + created_at = datetime.fromisoformat(created_at_value) + updated_at_value = cast(object, data.get("updated_at")) + if isinstance(updated_at_value, str) and updated_at_value: + with suppress(ValueError): + updated_at = datetime.fromisoformat(updated_at_value) + offset = cast(object, data.get("last_consolidated", 0)) + last_consolidated = ( + offset + if isinstance(offset, int) and not isinstance(offset, bool) + else 0 + ) + else: + messages.append(data) + + if skipped: + logger.warning("Skipped {} corrupt lines in session {}", skipped, key) + + if not messages and not metadata: + return None + + return Session( + key=key, + messages=messages, + created_at=created_at or datetime.now(), + updated_at=updated_at or datetime.now(), + metadata=metadata, + last_consolidated=last_consolidated, + ) + except _SESSION_DATA_ERRORS as e: + logger.warning("Repair failed for session {}: {}", key, e) + return None + + @staticmethod + def session_payload(session: Session) -> SessionPayload: + return { + "key": session.key, + "created_at": session.created_at.isoformat(), + "updated_at": session.updated_at.isoformat(), + "metadata": session.metadata, + "messages": session.messages, + } + + def save(self, session: Session, *, fsync: bool = False) -> None: + path = self.get_session_path(session.key) + tmp_path = path.with_suffix(".jsonl.tmp") + + try: + with open(tmp_path, "w", encoding="utf-8") as f: + metadata_line = { + "_type": "metadata", + "key": session.key, + "created_at": session.created_at.isoformat(), + "updated_at": session.updated_at.isoformat(), + "metadata": session.metadata, + "last_consolidated": session.last_consolidated, + } + f.write(json.dumps(metadata_line, ensure_ascii=False) + "\n") + for msg in session.messages: + f.write(json.dumps(msg, ensure_ascii=False) + "\n") + if fsync: + f.flush() + os.fsync(f.fileno()) + + os.replace(tmp_path, path) + + if fsync: + with suppress(PermissionError): + fd = os.open(str(path.parent), os.O_RDONLY) + try: + os.fsync(fd) + except OSError as exc: + if exc.errno != errno.EINVAL: + raise + finally: + os.close(fd) + except BaseException: + tmp_path.unlink(missing_ok=True) + raise + + def delete(self, key: str) -> bool: + paths = [ + self.get_session_path(key), + self.get_legacy_lossy_path(key), + self.get_legacy_session_path(key), + ] + deleted = False + for path in paths: + if not path.exists(): + continue + try: + path.unlink() + deleted = True + except OSError as e: + logger.warning("Failed to delete session file {}: {}", path, e) + return deleted + + def read(self, key: str) -> SessionPayload | None: + path = self.get_session_path(key) + if not path.exists(): + return None + try: + messages: list[dict[str, Any]] = [] + metadata: dict[str, Any] = {} + created_at: str | None = None + updated_at: str | None = None + stored_key: str | None = None + with open(path, encoding="utf-8") as f: + for line in f: + line = line.strip() + if not line: + continue + raw_data: object = json.loads(line) + data = _json_object(raw_data) + if data.get("_type") == "metadata": + metadata_value = cast(object, data.get("metadata", {})) + metadata = ( + cast(dict[str, Any], metadata_value) + if isinstance(metadata_value, dict) + else {} + ) + created_at_value = cast(object, data.get("created_at")) + updated_at_value = cast(object, data.get("updated_at")) + stored_key_value = cast(object, data.get("key")) + created_at = ( + created_at_value if isinstance(created_at_value, str) else None + ) + updated_at = ( + updated_at_value if isinstance(updated_at_value, str) else None + ) + stored_key = ( + stored_key_value if isinstance(stored_key_value, str) else None + ) + else: + messages.append(data) + return { + "key": stored_key or key, + "created_at": created_at, + "updated_at": updated_at, + "metadata": metadata, + "messages": messages, + } + except _SESSION_DATA_ERRORS as e: + logger.warning("Failed to read session {}: {}", key, e) + repaired = self.repair(key, path=path) + if repaired is not None: + logger.info("Recovered read-only session view {} from corrupt file", key) + return self.session_payload(repaired) + return None + + def read_metadata(self, key: str) -> SessionMetadataPayload | None: + path = self.get_session_path(key) + if not path.exists(): + return None + try: + with open(path, encoding="utf-8") as f: + for line in f: + line = line.strip() + if not line: + continue + raw_data: object = json.loads(line) + data = _json_object(raw_data) + if data.get("_type") != "metadata": + return None + metadata_value = cast(object, data.get("metadata", {})) + key_value = cast(object, data.get("key")) + created_at_value = cast(object, data.get("created_at")) + updated_at_value = cast(object, data.get("updated_at")) + return { + "key": key_value if isinstance(key_value, str) and key_value else key, + "created_at": ( + created_at_value if isinstance(created_at_value, str) else None + ), + "updated_at": ( + updated_at_value if isinstance(updated_at_value, str) else None + ), + "metadata": ( + cast(dict[str, Any], metadata_value) + if isinstance(metadata_value, dict) + else {} + ), + } + return None + except _SESSION_DATA_ERRORS as e: + logger.warning("Failed to read session metadata {}: {}", key, e) + repaired = self.repair(key, path=path) + if repaired is not None: + logger.info("Recovered read-only session metadata {} from corrupt file", key) + return { + "key": repaired.key, + "created_at": repaired.created_at.isoformat(), + "updated_at": repaired.updated_at.isoformat(), + "metadata": repaired.metadata, + } + return None + + def list_sessions(self) -> list[SessionInfo]: + sessions: list[SessionInfo] = [] + + for path in self.sessions_dir.glob("*.jsonl"): + storage_key = self.session_key_from_path(path) + if storage_key is None: + continue + try: + with open(path, encoding="utf-8") as f: + first_line = f.readline().strip() + if first_line: + raw_data: object = json.loads(first_line) + data = _json_object(raw_data) + if data.get("_type") == "metadata": + key_value = cast(object, data.get("key")) + key = ( + key_value + if isinstance(key_value, str) and key_value + else storage_key + ) + metadata = cast(object, data.get("metadata", {})) + title = _metadata_title(metadata) + preview = "" + fallback_preview = "" + scanned_records = 0 + scanned_chars = 0 + for line in f: + if not line.strip(): + continue + scanned_records += 1 + scanned_chars += len(line) + if ( + scanned_records > _SESSION_LIST_PREVIEW_MAX_RECORDS + or scanned_chars > _SESSION_LIST_PREVIEW_MAX_CHARS + ): + break + raw_item: object = json.loads(line) + item = _json_object(raw_item) + if item.get("_type") == "metadata": + continue + text = _message_preview_text(item) + if not text: + continue + if item.get("role") == "user": + preview = text + break + if not fallback_preview and item.get("role") == "assistant": + fallback_preview = text + preview = preview or fallback_preview + fallback_time = datetime.fromtimestamp(path.stat().st_mtime).isoformat() + created_at = cast(object, data.get("created_at")) + updated_at = cast(object, data.get("updated_at")) + sessions.append( + { + "key": key, + "created_at": ( + created_at + if isinstance(created_at, str) and created_at + else fallback_time + ), + "updated_at": ( + updated_at + if isinstance(updated_at, str) and updated_at + else fallback_time + ), + "title": title, + "preview": preview, + "path": str(path), + } + ) + except FileNotFoundError: + continue + except _SESSION_DATA_ERRORS: + repaired = self.repair(storage_key, path=path) + if repaired is not None: + sessions.append( + { + "key": repaired.key, + "created_at": repaired.created_at.isoformat(), + "updated_at": repaired.updated_at.isoformat(), + "title": _metadata_title(repaired.metadata), + "preview": next( + ( + text + for msg in repaired.messages + if (text := _message_preview_text(msg)) + ), + "", + ), + "path": str(path), + } + ) + continue + return sorted(sessions, key=lambda item: item["updated_at"], reverse=True) + + +class SessionManager: + """Manage session identity, caching, retention, and persistence.""" + + def __init__(self, workspace: Path, *, store: SessionStore | None = None): + self.workspace = workspace + self._jsonl_store = JsonlSessionStore(workspace) + self._store: SessionStore = store if store is not None else self._jsonl_store + self.sessions_dir = self._jsonl_store.sessions_dir + self.legacy_sessions_dir = self._jsonl_store.legacy_sessions_dir self._cache: OrderedDict[str, Session] = OrderedDict() # Preserve identity for sessions held by active callers without retaining idle ones. self._overflow_cache: WeakValueDictionary[str, Session] = WeakValueDictionary() @@ -475,24 +950,17 @@ class SessionManager: @staticmethod def safe_key(key: str) -> str: """Public helper used by HTTP handlers to map an arbitrary key to a stable filename stem.""" - return safe_filename(key.replace(":", "_")) + return JsonlSessionStore.safe_key(key) @staticmethod def _storage_key(key: str) -> str: """Collision-resistant encoding for internal session storage filenames.""" - return base64.urlsafe_b64encode(key.encode()).decode().rstrip("=") + return JsonlSessionStore.storage_key(key) @staticmethod def _decode_storage_key(stem: str) -> str | None: """Reverse _storage_key(): decode a base64url (no-padding) stem back to the original key.""" - try: - # Restore padding stripped by rstrip("=") - padding = 4 - len(stem) % 4 - if padding != 4: - stem += "=" * padding - return base64.urlsafe_b64decode(stem).decode("utf-8") - except _SESSION_DATA_ERRORS: - return None + return JsonlSessionStore.decode_storage_key(stem) @staticmethod def decode_storage_key(stem: str) -> str | None: @@ -502,22 +970,19 @@ class SessionManager: @classmethod def _session_key_from_path(cls, path: Path) -> str | None: """Decode a session key only from a canonical collision-resistant filename.""" - key = cls._decode_storage_key(path.stem) - if key is None or cls._storage_key(key) != path.stem: - return None - return key + return JsonlSessionStore.session_key_from_path(path) def _get_session_path(self, key: str) -> Path: """Get the collision-resistant workspace path for a session.""" - return self.sessions_dir / f"{self._storage_key(key)}.jsonl" + return self._jsonl_store.get_session_path(key) def _get_legacy_lossy_path(self, key: str) -> Path: """Previous workspace session path using lossy ':' to '_' replacement.""" - return self.sessions_dir / f"{safe_filename(key.replace(':', '_'))}.jsonl" + return self._jsonl_store.get_legacy_lossy_path(key) def _get_legacy_session_path(self, key: str) -> Path: """Legacy global session path (~/.nanobot/sessions/).""" - return self.legacy_sessions_dir / f"{self.safe_key(key)}.jsonl" + return self._jsonl_store.get_legacy_session_path(key) def get_or_create(self, key: str) -> Session: """ @@ -541,152 +1006,18 @@ class SessionManager: return session def _load(self, key: str) -> Session | None: - """Load a session from disk.""" - path = self._get_session_path(key) - if not path.exists(): - return None - - try: - messages: list[dict[str, Any]] = [] - metadata: object = {} - created_at: datetime | None = None - updated_at: datetime | None = None - last_consolidated: object = 0 - - with open(path, encoding="utf-8") as f: - for line in f: - line = line.strip() - if not line: - continue - - raw_data: object = json.loads(line) - data = _json_object(raw_data) - - if data.get("_type") == "metadata": - metadata = cast(object, data.get("metadata", {})) - created_at_value = cast(object, data.get("created_at")) - updated_at_value = cast(object, data.get("updated_at")) - created_at = ( - datetime.fromisoformat(cast(str, created_at_value)) - if created_at_value - else None - ) - updated_at = ( - datetime.fromisoformat(cast(str, updated_at_value)) - if updated_at_value - else None - ) - last_consolidated = cast( - object, - data.get("last_consolidated", 0), - ) - else: - messages.append(data) - - return Session( - key=key, - messages=messages, - created_at=created_at or datetime.now(), - updated_at=updated_at or datetime.now(), - metadata=cast(dict[str, Any], metadata), - last_consolidated=cast(int, last_consolidated), - ) - except _SESSION_DATA_ERRORS as e: - logger.warning("Failed to load session {}: {}", key, e) - repaired = self._repair(key) - if repaired is not None: - logger.info("Recovered session {} from corrupt file ({} messages)", key, len(repaired.messages)) - return repaired + return self._store.load(key) def _repair(self, key: str, *, path: Path | None = None) -> Session | None: """Attempt to recover a session from a corrupt JSONL file.""" - if path is None: - path = self._get_session_path(key) - if not path.exists(): - return None - - try: - messages: list[dict[str, Any]] = [] - metadata: object = {} - created_at: datetime | None = None - updated_at: datetime | None = None - last_consolidated: object = 0 - skipped = 0 - - with open(path, encoding="utf-8") as f: - for line in f: - line = line.strip() - if not line: - continue - try: - raw_data: object = json.loads(line) - except json.JSONDecodeError: - skipped += 1 - continue - if not isinstance(raw_data, dict): - skipped += 1 - continue - data = cast(dict[str, Any], raw_data) - - if data.get("_type") == "metadata": - metadata = cast(object, data.get("metadata", {})) - created_at_value = cast(object, data.get("created_at")) - if created_at_value: - with suppress(ValueError, TypeError): - created_at = datetime.fromisoformat( - cast(str, created_at_value) - ) - updated_at_value = cast(object, data.get("updated_at")) - if updated_at_value: - with suppress(ValueError, TypeError): - updated_at = datetime.fromisoformat( - cast(str, updated_at_value) - ) - last_consolidated = cast( - object, - data.get("last_consolidated", 0), - ) - else: - messages.append(data) - - if skipped: - logger.warning("Skipped {} corrupt lines in session {}", skipped, key) - - if not messages and not metadata: - return None - - return Session( - key=key, - messages=messages, - created_at=created_at or datetime.now(), - updated_at=updated_at or datetime.now(), - metadata=cast(dict[str, Any], metadata), - last_consolidated=cast(int, last_consolidated), - ) - except _SESSION_DATA_ERRORS as e: - logger.warning("Repair failed for session {}: {}", key, e) - return None + return self._jsonl_store.repair(key, path=path) @staticmethod - def _session_payload(session: Session) -> dict[str, Any]: - return { - "key": session.key, - "created_at": session.created_at.isoformat(), - "updated_at": session.updated_at.isoformat(), - "metadata": session.metadata, - "messages": session.messages, - } + def _session_payload(session: Session) -> SessionPayload: + return JsonlSessionStore.session_payload(session) def save(self, session: Session, *, fsync: bool = False) -> None: - """Save a session to disk atomically. - - When *fsync* is ``True`` the final file and its parent directory are - explicitly flushed to durable storage. This is intentionally off by - default (the OS page-cache is sufficient for normal operation) but - should be enabled during graceful shutdown so that filesystems with - write-back caching (e.g. rclone VFS, NFS, FUSE mounts) do not lose - the most recent writes. - """ + """Persist a session and retain it in the cache.""" archiver = self._file_cap_archiver if archiver is not None: session.enforce_file_cap( @@ -696,46 +1027,7 @@ class SessionManager: ) ) - path = self._get_session_path(session.key) - tmp_path = path.with_suffix(".jsonl.tmp") - - try: - with open(tmp_path, "w", encoding="utf-8") as f: - metadata_line = { - "_type": "metadata", - "key": session.key, - "created_at": session.created_at.isoformat(), - "updated_at": session.updated_at.isoformat(), - "metadata": session.metadata, - "last_consolidated": session.last_consolidated - } - f.write(json.dumps(metadata_line, ensure_ascii=False) + "\n") - for msg in session.messages: - f.write(json.dumps(msg, ensure_ascii=False) + "\n") - if fsync: - f.flush() - os.fsync(f.fileno()) - - os.replace(tmp_path, path) - - if fsync: - # fsync the directory so the rename is durable. - # On Windows, opening a directory with O_RDONLY raises - # PermissionError; some shared filesystems allow the open but - # reject directory fsync with EINVAL. - with suppress(PermissionError): - fd = os.open(str(path.parent), os.O_RDONLY) - try: - os.fsync(fd) - except OSError as exc: - if exc.errno != errno.EINVAL: - raise - finally: - os.close(fd) - except BaseException: - tmp_path.unlink(missing_ok=True) - raise - + self._store.save(session, fsync=fsync) self._remember(session) def flush_all(self) -> int: @@ -762,26 +1054,9 @@ class SessionManager: self._overflow_cache.pop(key, None) def delete_session(self, key: str) -> bool: - """Remove a session from disk (both workspace and legacy locations) and cache. - - Returns True if at least one JSONL file was found and unlinked. - """ - paths = [ - self._get_session_path(key), - self._get_legacy_lossy_path(key), - self._get_legacy_session_path(key), - ] + """Delete a persisted session and invalidate its cache entry.""" self.invalidate(key) - deleted = False - for path in paths: - if not path.exists(): - continue - try: - path.unlink() - deleted = True - except OSError as e: - logger.warning("Failed to delete session file {}: {}", path, e) - return deleted + return self._store.delete(key) def fork_session_before_user_index( self, @@ -840,180 +1115,12 @@ class SessionManager: return target def read_session_file(self, key: str) -> dict[str, Any] | None: - """Load a session from disk without caching; intended for read-only HTTP endpoints. - - Returns ``{"key", "created_at", "updated_at", "metadata", "messages"}`` or - ``None`` when the session file does not exist or fails to parse. - """ - path = self._get_session_path(key) - if not path.exists(): - return None - try: - messages: list[dict[str, Any]] = [] - metadata: object = {} - created_at: object = None - updated_at: object = None - stored_key: object = None - with open(path, encoding="utf-8") as f: - for line in f: - line = line.strip() - if not line: - continue - raw_data: object = json.loads(line) - data = _json_object(raw_data) - if data.get("_type") == "metadata": - metadata = cast(object, data.get("metadata", {})) - created_at = cast(object, data.get("created_at")) - updated_at = cast(object, data.get("updated_at")) - stored_key = cast(object, data.get("key")) - else: - messages.append(data) - return { - "key": stored_key or key, - "created_at": created_at, - "updated_at": updated_at, - "metadata": metadata, - "messages": messages, - } - except _SESSION_DATA_ERRORS as e: - logger.warning("Failed to read session {}: {}", key, e) - repaired = self._repair(key, path=path) - if repaired is not None: - logger.info("Recovered read-only session view {} from corrupt file", key) - return self._session_payload(repaired) - return None + """Read a session without populating the cache.""" + return cast(dict[str, Any] | None, self._store.read(key)) def read_session_metadata(self, key: str) -> dict[str, Any] | None: - """Load only the metadata record from a session file. - - This is used by WebUI routes that need session-level metadata but not the - full conversation transcript. - """ - path = self._get_session_path(key) - if not path.exists(): - return None - try: - with open(path, encoding="utf-8") as f: - for line in f: - line = line.strip() - if not line: - continue - raw_data: object = json.loads(line) - data = _json_object(raw_data) - if data.get("_type") != "metadata": - return None - metadata = cast(object, data.get("metadata", {})) - return { - "key": data.get("key") or key, - "created_at": data.get("created_at"), - "updated_at": data.get("updated_at"), - "metadata": ( - cast(dict[str, Any], metadata) - if isinstance(metadata, dict) - else {} - ), - } - return None - except _SESSION_DATA_ERRORS as e: - logger.warning("Failed to read session metadata {}: {}", key, e) - repaired = self._repair(key, path=path) - if repaired is not None: - logger.info("Recovered read-only session metadata {} from corrupt file", key) - return { - "key": repaired.key, - "created_at": repaired.created_at.isoformat(), - "updated_at": repaired.updated_at.isoformat(), - "metadata": repaired.metadata, - } - return None + """Read session metadata without loading the transcript.""" + return cast(dict[str, Any] | None, self._store.read_metadata(key)) def list_sessions(self) -> list[dict[str, Any]]: - """ - List all sessions. - - Returns: - List of session info dicts. - """ - sessions: list[dict[str, Any]] = [] - - for path in self.sessions_dir.glob("*.jsonl"): - storage_key = self._session_key_from_path(path) - if storage_key is None: - continue - try: - # Read the metadata line and a small preview for session lists. - with open(path, encoding="utf-8") as f: - first_line = f.readline().strip() - if first_line: - raw_data: object = json.loads(first_line) - data = _json_object(raw_data) - if data.get("_type") == "metadata": - key = cast(object, data.get("key")) or storage_key - metadata = cast(object, data.get("metadata", {})) - title = _metadata_title(metadata) - preview = "" - fallback_preview = "" - scanned_records = 0 - scanned_chars = 0 - for line in f: - if not line.strip(): - continue - scanned_records += 1 - scanned_chars += len(line) - if ( - scanned_records > _SESSION_LIST_PREVIEW_MAX_RECORDS - or scanned_chars > _SESSION_LIST_PREVIEW_MAX_CHARS - ): - break - raw_item: object = json.loads(line) - item = _json_object(raw_item) - if item.get("_type") == "metadata": - continue - text = _message_preview_text(item) - if not text: - continue - if item.get("role") == "user": - preview = text - break - if not fallback_preview and item.get("role") == "assistant": - fallback_preview = text - preview = preview or fallback_preview - fallback_time = datetime.fromtimestamp(path.stat().st_mtime).isoformat() - sessions.append( - { - "key": key, - "created_at": data.get("created_at") or fallback_time, - "updated_at": data.get("updated_at") or fallback_time, - "title": title, - "preview": preview, - "path": str(path), - } - ) - except FileNotFoundError: - continue - except _SESSION_DATA_ERRORS: - repaired = self._repair(storage_key, path=path) - if repaired is not None: - sessions.append( - { - "key": repaired.key, - "created_at": repaired.created_at.isoformat(), - "updated_at": repaired.updated_at.isoformat(), - "title": _metadata_title(repaired.metadata), - "preview": next( - ( - text - for msg in repaired.messages - if (text := _message_preview_text(msg)) - ), - "", - ), - "path": str(path), - } - ) - continue - return sorted( - sessions, - key=lambda item: cast(str, item.get("updated_at", "")), - reverse=True, - ) + return cast(list[dict[str, Any]], self._store.list_sessions()) diff --git a/tests/session/test_session_store.py b/tests/session/test_session_store.py new file mode 100644 index 000000000..dc8e8c0b7 --- /dev/null +++ b/tests/session/test_session_store.py @@ -0,0 +1,79 @@ +from unittest.mock import MagicMock + +import nanobot.session as session_api +from nanobot.session import Session, SessionManager +from nanobot.session.manager import FILE_MAX_MESSAGES, SessionStore + + +def test_store_types_are_not_public_session_api() -> None: + assert not hasattr(session_api, "SessionStore") + assert not hasattr(session_api, "JsonlSessionStore") + + +def test_manager_delegates_persistence_to_store(tmp_path) -> None: + stored = Session(key="cli:test") + stored.add_message("user", "hello") + payload = { + "key": stored.key, + "created_at": stored.created_at.isoformat(), + "updated_at": stored.updated_at.isoformat(), + "metadata": {}, + "messages": stored.messages, + } + metadata = { + "key": stored.key, + "created_at": stored.created_at.isoformat(), + "updated_at": stored.updated_at.isoformat(), + "metadata": {}, + } + listing = [ + { + "key": stored.key, + "created_at": stored.created_at.isoformat(), + "updated_at": stored.updated_at.isoformat(), + "title": "", + "preview": "hello", + "path": "session.db", + } + ] + store = MagicMock(spec=SessionStore) + store.load.return_value = stored + store.read.return_value = payload + store.read_metadata.return_value = metadata + store.list_sessions.return_value = listing + store.delete.return_value = True + manager = SessionManager(tmp_path, store=store) + + assert manager.get_or_create(stored.key) is stored + assert manager.get_or_create(stored.key) is stored + store.load.assert_called_once_with(stored.key) + + manager.save(stored, fsync=True) + store.save.assert_called_once_with(stored, fsync=True) + assert manager.read_session_file(stored.key) == payload + assert manager.read_session_metadata(stored.key) == metadata + assert manager.list_sessions() == listing + + assert manager.delete_session(stored.key) is True + store.delete.assert_called_once_with(stored.key) + assert manager.get_cached(stored.key) is None + + +def test_manager_applies_file_cap_before_store_save(tmp_path) -> None: + store = MagicMock(spec=SessionStore) + archiver = MagicMock() + manager = SessionManager(tmp_path, store=store) + manager.set_file_cap_archiver(archiver) + session = Session( + key="cli:large", + messages=[ + {"role": "user", "content": str(index)} + for index in range(FILE_MAX_MESSAGES + 1) + ], + ) + + manager.save(session) + + assert len(session.messages) == FILE_MAX_MESSAGES + archiver.assert_called_once() + store.save.assert_called_once_with(session, fsync=False)