mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-07 01:48:53 +00:00
refactor(session): separate persistence behind SessionStore (#5170)
This commit is contained in:
parent
c33c188afb
commit
ad6900e56c
@ -11,7 +11,7 @@ from copy import deepcopy
|
|||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Callable, cast
|
from typing import Any, Callable, Protocol, TypedDict, cast
|
||||||
from weakref import WeakValueDictionary
|
from weakref import WeakValueDictionary
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
@ -427,17 +427,492 @@ class Session:
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
class SessionManager:
|
class SessionPayload(TypedDict):
|
||||||
"""
|
key: str
|
||||||
Manages conversation sessions.
|
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):
|
def __init__(self, workspace: Path):
|
||||||
self.workspace = workspace
|
self.sessions_dir = ensure_dir(workspace / "sessions")
|
||||||
self.sessions_dir = ensure_dir(self.workspace / "sessions")
|
|
||||||
self.legacy_sessions_dir = get_legacy_sessions_dir()
|
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()
|
self._cache: OrderedDict[str, Session] = OrderedDict()
|
||||||
# Preserve identity for sessions held by active callers without retaining idle ones.
|
# Preserve identity for sessions held by active callers without retaining idle ones.
|
||||||
self._overflow_cache: WeakValueDictionary[str, Session] = WeakValueDictionary()
|
self._overflow_cache: WeakValueDictionary[str, Session] = WeakValueDictionary()
|
||||||
@ -475,24 +950,17 @@ class SessionManager:
|
|||||||
@staticmethod
|
@staticmethod
|
||||||
def safe_key(key: str) -> str:
|
def safe_key(key: str) -> str:
|
||||||
"""Public helper used by HTTP handlers to map an arbitrary key to a stable filename stem."""
|
"""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
|
@staticmethod
|
||||||
def _storage_key(key: str) -> str:
|
def _storage_key(key: str) -> str:
|
||||||
"""Collision-resistant encoding for internal session storage filenames."""
|
"""Collision-resistant encoding for internal session storage filenames."""
|
||||||
return base64.urlsafe_b64encode(key.encode()).decode().rstrip("=")
|
return JsonlSessionStore.storage_key(key)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _decode_storage_key(stem: str) -> str | None:
|
def _decode_storage_key(stem: str) -> str | None:
|
||||||
"""Reverse _storage_key(): decode a base64url (no-padding) stem back to the original key."""
|
"""Reverse _storage_key(): decode a base64url (no-padding) stem back to the original key."""
|
||||||
try:
|
return JsonlSessionStore.decode_storage_key(stem)
|
||||||
# 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
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def decode_storage_key(stem: str) -> str | None:
|
def decode_storage_key(stem: str) -> str | None:
|
||||||
@ -502,22 +970,19 @@ class SessionManager:
|
|||||||
@classmethod
|
@classmethod
|
||||||
def _session_key_from_path(cls, path: Path) -> str | None:
|
def _session_key_from_path(cls, path: Path) -> str | None:
|
||||||
"""Decode a session key only from a canonical collision-resistant filename."""
|
"""Decode a session key only from a canonical collision-resistant filename."""
|
||||||
key = cls._decode_storage_key(path.stem)
|
return JsonlSessionStore.session_key_from_path(path)
|
||||||
if key is None or cls._storage_key(key) != path.stem:
|
|
||||||
return None
|
|
||||||
return key
|
|
||||||
|
|
||||||
def _get_session_path(self, key: str) -> Path:
|
def _get_session_path(self, key: str) -> Path:
|
||||||
"""Get the collision-resistant workspace path for a session."""
|
"""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:
|
def _get_legacy_lossy_path(self, key: str) -> Path:
|
||||||
"""Previous workspace session path using lossy ':' to '_' replacement."""
|
"""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:
|
def _get_legacy_session_path(self, key: str) -> Path:
|
||||||
"""Legacy global session path (~/.nanobot/sessions/)."""
|
"""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:
|
def get_or_create(self, key: str) -> Session:
|
||||||
"""
|
"""
|
||||||
@ -541,152 +1006,18 @@ class SessionManager:
|
|||||||
return session
|
return session
|
||||||
|
|
||||||
def _load(self, key: str) -> Session | None:
|
def _load(self, key: str) -> Session | None:
|
||||||
"""Load a session from disk."""
|
return self._store.load(key)
|
||||||
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
|
|
||||||
|
|
||||||
def _repair(self, key: str, *, path: Path | None = None) -> Session | None:
|
def _repair(self, key: str, *, path: Path | None = None) -> Session | None:
|
||||||
"""Attempt to recover a session from a corrupt JSONL file."""
|
"""Attempt to recover a session from a corrupt JSONL file."""
|
||||||
if path is None:
|
return self._jsonl_store.repair(key, path=path)
|
||||||
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
|
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _session_payload(session: Session) -> dict[str, Any]:
|
def _session_payload(session: Session) -> SessionPayload:
|
||||||
return {
|
return JsonlSessionStore.session_payload(session)
|
||||||
"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:
|
def save(self, session: Session, *, fsync: bool = False) -> None:
|
||||||
"""Save a session to disk atomically.
|
"""Persist a session and retain it in the cache."""
|
||||||
|
|
||||||
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.
|
|
||||||
"""
|
|
||||||
archiver = self._file_cap_archiver
|
archiver = self._file_cap_archiver
|
||||||
if archiver is not None:
|
if archiver is not None:
|
||||||
session.enforce_file_cap(
|
session.enforce_file_cap(
|
||||||
@ -696,46 +1027,7 @@ class SessionManager:
|
|||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
path = self._get_session_path(session.key)
|
self._store.save(session, fsync=fsync)
|
||||||
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._remember(session)
|
self._remember(session)
|
||||||
|
|
||||||
def flush_all(self) -> int:
|
def flush_all(self) -> int:
|
||||||
@ -762,26 +1054,9 @@ class SessionManager:
|
|||||||
self._overflow_cache.pop(key, None)
|
self._overflow_cache.pop(key, None)
|
||||||
|
|
||||||
def delete_session(self, key: str) -> bool:
|
def delete_session(self, key: str) -> bool:
|
||||||
"""Remove a session from disk (both workspace and legacy locations) and cache.
|
"""Delete a persisted session and invalidate its cache entry."""
|
||||||
|
|
||||||
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),
|
|
||||||
]
|
|
||||||
self.invalidate(key)
|
self.invalidate(key)
|
||||||
deleted = False
|
return self._store.delete(key)
|
||||||
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 fork_session_before_user_index(
|
def fork_session_before_user_index(
|
||||||
self,
|
self,
|
||||||
@ -840,180 +1115,12 @@ class SessionManager:
|
|||||||
return target
|
return target
|
||||||
|
|
||||||
def read_session_file(self, key: str) -> dict[str, Any] | None:
|
def read_session_file(self, key: str) -> dict[str, Any] | None:
|
||||||
"""Load a session from disk without caching; intended for read-only HTTP endpoints.
|
"""Read a session without populating the cache."""
|
||||||
|
return cast(dict[str, Any] | None, self._store.read(key))
|
||||||
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
|
|
||||||
|
|
||||||
def read_session_metadata(self, key: str) -> dict[str, Any] | None:
|
def read_session_metadata(self, key: str) -> dict[str, Any] | None:
|
||||||
"""Load only the metadata record from a session file.
|
"""Read session metadata without loading the transcript."""
|
||||||
|
return cast(dict[str, Any] | None, self._store.read_metadata(key))
|
||||||
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
|
|
||||||
|
|
||||||
def list_sessions(self) -> list[dict[str, Any]]:
|
def list_sessions(self) -> list[dict[str, Any]]:
|
||||||
"""
|
return cast(list[dict[str, Any]], self._store.list_sessions())
|
||||||
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,
|
|
||||||
)
|
|
||||||
|
|||||||
79
tests/session/test_session_store.py
Normal file
79
tests/session/test_session_store.py
Normal file
@ -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)
|
||||||
Loading…
x
Reference in New Issue
Block a user