diff --git a/better_memory/_common.py b/better_memory/_common.py new file mode 100644 index 0000000..bb8359d --- /dev/null +++ b/better_memory/_common.py @@ -0,0 +1,64 @@ +"""Tiny shared helpers used across hooks, services, storage, and the MCP layer. + +Pure stdlib with no intra-package imports, so hooks can import this module +without paying for SQLite / sqlite-vec / boto3 / config at invocation time +(hooks must start fast and never fail). Everything heavier belongs in +:mod:`better_memory.config` or the relevant service module. +""" + +from __future__ import annotations + +import os +from datetime import UTC, datetime +from pathlib import Path +from uuid import uuid4 + +_DEFAULT_HOME = "~/.better-memory" + + +def env_session_id() -> str | None: + """Return the Claude session id from the environment, or ``None``. + + Reads ``CLAUDE_SESSION_ID`` first (kept as the primary name for + back-compat), then ``CLAUDE_CODE_SESSION_ID`` (the name Claude Code + actually exports). The shared resolution order makes hook-written + events and MCP-written observations agree on the same session id. + """ + return ( + os.environ.get("CLAUDE_SESSION_ID") + or os.environ.get("CLAUDE_CODE_SESSION_ID") + ) + + +def get_session_id() -> str: + """Return the current Claude session id, generating one if no env var + is set. Falls back to a fresh ``uuid4().hex`` (32 chars).""" + return env_session_id() or uuid4().hex + + +def default_clock() -> datetime: + """UTC-aware ``now``. The conventional default for injectable ``clock`` + parameters across services.""" + return datetime.now(UTC) + + +def resolve_home() -> Path: + """Return ``BETTER_MEMORY_HOME`` (or its default) with ``~`` expanded.""" + raw = os.environ.get("BETTER_MEMORY_HOME", _DEFAULT_HOME) + return Path(raw).expanduser() + + +def default_spool_dir() -> Path: + """Return ``$BETTER_MEMORY_HOME/spool``, defaulting to ``~/.better-memory``.""" + return resolve_home() / "spool" + + +def safe_timestamp(raw: str | None) -> str: + """Return a filesystem-safe timestamp component. + + Replaces ``:`` (illegal on NTFS) with ``-``. Falls back to current UTC + time if ``raw`` is missing or empty. + """ + if not raw: + raw = datetime.now(UTC).isoformat() + return raw.replace(":", "-") diff --git a/better_memory/config.py b/better_memory/config.py index c007c86..a704aee 100644 --- a/better_memory/config.py +++ b/better_memory/config.py @@ -23,8 +23,8 @@ from typing import Literal from better_memory import _diag +from better_memory._common import resolve_home -_DEFAULT_HOME = "~/.better-memory" _DEFAULT_OLLAMA_HOST = "http://localhost:11434" _DEFAULT_EMBED_MODEL = "nomic-embed-text" _DEFAULT_EMBEDDINGS_BACKEND = "ollama" @@ -34,12 +34,6 @@ _DEFAULT_AGENTCORE_REGION = "eu-west-2" -def resolve_home() -> Path: - """Return ``BETTER_MEMORY_HOME`` (or its default) with ``~`` expanded.""" - raw = os.environ.get("BETTER_MEMORY_HOME", _DEFAULT_HOME) - return Path(raw).expanduser() - - # Maps absolute cwd string → resolved project name (or None when no git tree # is reachable). Both successes and ``None`` are cached for the process # lifetime: the walk is deterministic given filesystem state, so caching the diff --git a/better_memory/hooks/observer.py b/better_memory/hooks/observer.py index aa397e0..b4de183 100644 --- a/better_memory/hooks/observer.py +++ b/better_memory/hooks/observer.py @@ -16,7 +16,8 @@ import sys import time from datetime import UTC, datetime -from pathlib import Path + +from better_memory._common import default_spool_dir, safe_timestamp # Cap stdin reads so a malicious or accidentally huge payload can't starve the # hook process of memory. 1 MiB is far larger than anything Claude Code emits @@ -24,29 +25,6 @@ _MAX_STDIN_BYTES = 1_048_576 -def _default_spool_dir() -> Path: - """Return ``$BETTER_MEMORY_HOME/spool``, defaulting to ``~/.better-memory``. - - Kept separate from :func:`better_memory.config.get_config` so hooks do not - import SQLite / sqlite-vec / anything heavyweight at invocation time. - """ - home = os.environ.get("BETTER_MEMORY_HOME") - if home: - return Path(home).expanduser() / "spool" - return Path.home() / ".better-memory" / "spool" - - -def _safe_timestamp(raw: str | None) -> str: - """Return a filesystem-safe timestamp component. - - Replaces ``:`` (illegal on NTFS) with ``-``. Falls back to current UTC - time if ``raw`` is missing or empty. - """ - if not raw: - raw = datetime.now(UTC).isoformat() - return raw.replace(":", "-") - - def _safe_tool(raw: object) -> str: """Return a filesystem-safe tool component.""" if not raw or not isinstance(raw, str): @@ -75,10 +53,10 @@ def main() -> None: if "timestamp" not in data or not data["timestamp"]: data["timestamp"] = datetime.now(UTC).isoformat() - spool_dir = _default_spool_dir() + spool_dir = default_spool_dir() spool_dir.mkdir(parents=True, exist_ok=True) - ts_component = _safe_timestamp(data.get("timestamp")) + ts_component = safe_timestamp(data.get("timestamp")) tool_component = _safe_tool(data.get("tool")) # SHA-256 prefix of the serialised payload — cheap collision avoidance # for two events that happen in the same second on the same tool. The diff --git a/better_memory/hooks/post_commit.py b/better_memory/hooks/post_commit.py index 1de01f4..736306c 100644 --- a/better_memory/hooks/post_commit.py +++ b/better_memory/hooks/post_commit.py @@ -29,9 +29,13 @@ import time from datetime import UTC, datetime from pathlib import Path -from uuid import uuid4 -from better_memory.config import project_name, resolve_home +from better_memory._common import ( + default_spool_dir, + get_session_id, + safe_timestamp, +) +from better_memory.config import project_name _MAX_STDIN_BYTES = 1_048_576 @@ -42,17 +46,6 @@ _TRAILER_KEY = "closes-episode" -def _default_spool_dir() -> Path: - """Return ``$BETTER_MEMORY_HOME/spool``, defaulting to ``~/.better-memory``.""" - return resolve_home() / "spool" - - -def _safe_timestamp(raw: str | None) -> str: - if not raw: - raw = datetime.now(UTC).isoformat() - return raw.replace(":", "-") - - def _read_head_commit_message() -> tuple[str, str]: """Return ``(subject_plus_body, commit_sha)`` for HEAD. @@ -146,11 +139,7 @@ def main() -> None: sys.exit(0) now_iso = datetime.now(UTC).isoformat() - session_id = ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - or uuid4().hex - ) + session_id = get_session_id() cwd = _resolve_cwd() data: dict[str, object] = { @@ -162,10 +151,10 @@ def main() -> None: "commit_sha": commit_sha, } - spool_dir = _default_spool_dir() + spool_dir = default_spool_dir() spool_dir.mkdir(parents=True, exist_ok=True) - ts_component = _safe_timestamp(now_iso) + ts_component = safe_timestamp(now_iso) serialised = json.dumps(data, sort_keys=True).encode("utf-8") salt = f"{time.time_ns()}:{os.getpid()}".encode() hash_hex = hashlib.sha256(serialised + salt).hexdigest()[:12] diff --git a/better_memory/hooks/session_bootstrap.py b/better_memory/hooks/session_bootstrap.py index d05743d..52f75de 100644 --- a/better_memory/hooks/session_bootstrap.py +++ b/better_memory/hooks/session_bootstrap.py @@ -13,8 +13,8 @@ import sys from contextlib import closing from pathlib import Path -from uuid import uuid4 +from better_memory._common import get_session_id from better_memory.config import get_config from better_memory.db.connection import connect from better_memory.hooks._error_log import record_hook_error @@ -62,11 +62,7 @@ def main() -> None: session_id = ( str(payload.get("session_id")) if payload.get("session_id") - else ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - or uuid4().hex - ) + else get_session_id() ) cwd_str = str(payload.get("cwd")) if payload.get("cwd") else os.getcwd() diff --git a/better_memory/hooks/session_close.py b/better_memory/hooks/session_close.py index 62a18c1..6b4047a 100644 --- a/better_memory/hooks/session_close.py +++ b/better_memory/hooks/session_close.py @@ -16,9 +16,14 @@ import sys import time from datetime import UTC, datetime -from pathlib import Path -from uuid import uuid4 +from better_memory._common import ( + default_spool_dir, + env_session_id, + get_session_id, + resolve_home, + safe_timestamp, +) from better_memory.config import get_config # Mirror the observer cap: reject any stdin payload above 1 MiB without @@ -26,35 +31,13 @@ _MAX_STDIN_BYTES = 1_048_576 -def _default_spool_dir() -> Path: - """Return ``$BETTER_MEMORY_HOME/spool``, defaulting to ``~/.better-memory``. - - Mirrors the observer hook. Kept duplicated to avoid a cross-module import - that would slow hook startup. - """ - home = os.environ.get("BETTER_MEMORY_HOME") - if home: - return Path(home).expanduser() / "spool" - return Path.home() / ".better-memory" / "spool" - - -def _safe_timestamp(raw: str | None) -> str: - if not raw: - raw = datetime.now(UTC).isoformat() - return raw.replace(":", "-") - - def _synthesise_marker() -> dict[str, str]: """Build a minimal ``session_end`` payload from env + clock.""" return { "event_type": "session_end", "timestamp": datetime.now(UTC).isoformat(), "cwd": os.environ.get("PWD") or os.getcwd(), - "session_id": ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - or uuid4().hex - ), + "session_id": get_session_id(), } @@ -103,12 +86,7 @@ def _fire_agentcore_closure(*, session_id: str, project: str) -> bool: resolve_actor_id, ) - home_env = os.environ.get("BETTER_MEMORY_HOME") - home = ( - Path(home_env).expanduser() - if home_env - else Path.home() / ".better-memory" - ) + home = resolve_home() cfg = load_agentcore_config(home) if cfg is None: return False @@ -267,19 +245,12 @@ def main() -> None: if "timestamp" not in data or not data["timestamp"]: data["timestamp"] = datetime.now(UTC).isoformat() if "session_id" not in data or not data["session_id"]: - data["session_id"] = ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - or uuid4().hex - ) + data["session_id"] = get_session_id() if "cwd" not in data or not data["cwd"]: data["cwd"] = os.environ.get("PWD") or os.getcwd() session_id_str = ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - or data.get("session_id") - or "" + env_session_id() or data.get("session_id") or "" ) if session_id_str and _emit_rating_directive_if_unrated( str(session_id_str) @@ -307,10 +278,10 @@ def main() -> None: project=str(project_for_closure), ) - spool_dir = _default_spool_dir() + spool_dir = default_spool_dir() spool_dir.mkdir(parents=True, exist_ok=True) - ts_component = _safe_timestamp(str(data.get("timestamp"))) + ts_component = safe_timestamp(str(data.get("timestamp"))) # Salt the hash with monotonic-nanosecond clock + PID so two # byte-identical payloads in the same second can't collide on # filename. The salt does NOT appear in the written body. diff --git a/better_memory/mcp/_util.py b/better_memory/mcp/_util.py index 3e5d9d3..6baf892 100644 --- a/better_memory/mcp/_util.py +++ b/better_memory/mcp/_util.py @@ -8,13 +8,13 @@ from __future__ import annotations import logging -import os import time from collections.abc import Callable from pathlib import Path from typing import Any from better_memory import _diag +from better_memory._common import env_session_id from better_memory.runtime.session_marker import read_session_id logger = logging.getLogger(__name__) @@ -64,8 +64,4 @@ def resolve_session_id(home: Path) -> str | None: propagate the session id into the spawned stdio MCP server's env, so the marker file is the fallback for every rating call. """ - return ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - or read_session_id(home) - ) + return env_session_id() or read_session_id(home) diff --git a/better_memory/mcp/handlers/sessions.py b/better_memory/mcp/handlers/sessions.py index 1560c4d..97e2d20 100644 --- a/better_memory/mcp/handlers/sessions.py +++ b/better_memory/mcp/handlers/sessions.py @@ -8,12 +8,12 @@ import json import os -import uuid from pathlib import Path from typing import Any from mcp.types import TextContent +from better_memory._common import get_session_id from better_memory.mcp._util import resolve_session_id from better_memory.services import ui_launcher from better_memory.services.memory_rating import MemoryRatingService @@ -45,12 +45,7 @@ def tools(self) -> dict[str, Any]: async def session_bootstrap(self, args: dict[str, Any]) -> list[TextContent]: cwd_arg = args.get("cwd") or os.getcwd() - session_id_arg = ( - args.get("session_id") - or os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - or uuid.uuid4().hex - ) + session_id_arg = args.get("session_id") or get_session_id() result = self._session_bootstrap.bootstrap( source=args.get("source"), session_id=session_id_arg, diff --git a/better_memory/search/hybrid.py b/better_memory/search/hybrid.py index db48440..4135250 100644 --- a/better_memory/search/hybrid.py +++ b/better_memory/search/hybrid.py @@ -30,6 +30,8 @@ import sqlite_vec +from better_memory._common import default_clock + Outcome = Literal["success", "failure", "neutral"] @@ -108,7 +110,7 @@ def hybrid_search( if second_source == "trigram" and (query_text is None or not query_text.strip()): return [] - now = (clock or _default_clock)() + now = (clock or default_clock)() where_sql, where_params = _build_where(filters, now=now) # Gather candidate rowids from each active source. @@ -416,5 +418,3 @@ def _parse_sqlite_datetime(value: str) -> datetime: return datetime.fromisoformat(normalized) -def _default_clock() -> datetime: - return datetime.now(UTC) diff --git a/better_memory/services/episode.py b/better_memory/services/episode.py index 892096e..1e1dc86 100644 --- a/better_memory/services/episode.py +++ b/better_memory/services/episode.py @@ -18,14 +18,11 @@ import sqlite3 from collections.abc import Callable from dataclasses import dataclass -from datetime import UTC, datetime +from datetime import datetime from uuid import uuid4 from better_memory import _diag - - -def _default_clock() -> datetime: - return datetime.now(UTC) +from better_memory._common import default_clock @dataclass(frozen=True) @@ -74,7 +71,7 @@ def __init__( clock: Callable[[], datetime] | None = None, ) -> None: self._conn = conn - self._clock: Callable[[], datetime] = clock or _default_clock + self._clock: Callable[[], datetime] = clock or default_clock def open_background(self, *, session_id: str, project: str) -> str: """Create a background episode (goal=NULL) for ``session_id``. diff --git a/better_memory/services/knowledge.py b/better_memory/services/knowledge.py index b50418a..93f0aad 100644 --- a/better_memory/services/knowledge.py +++ b/better_memory/services/knowledge.py @@ -32,6 +32,7 @@ from datetime import UTC, datetime from pathlib import Path +from better_memory._common import default_clock from better_memory.config import project_name from better_memory.search.query import sanitize_fts5_query @@ -77,10 +78,6 @@ class ReindexReport: removed: int -def _default_clock() -> datetime: - return datetime.now(UTC) - - def _doc_id(relative_path: str) -> str: """Stable 16-char id derived from the POSIX relative path.""" return hashlib.sha256(relative_path.encode("utf-8")).hexdigest()[:16] @@ -181,7 +178,7 @@ def __init__( ) -> None: self._conn = conn self._knowledge_base = Path(knowledge_base) if knowledge_base else None - self._clock: Callable[[], datetime] = clock or _default_clock + self._clock: Callable[[], datetime] = clock or default_clock # ------------------------------------------------------------------ public diff --git a/better_memory/services/memory_rating.py b/better_memory/services/memory_rating.py index 970a90c..afafcf1 100644 --- a/better_memory/services/memory_rating.py +++ b/better_memory/services/memory_rating.py @@ -13,13 +13,10 @@ import sqlite3 from collections.abc import Callable -from datetime import UTC, datetime +from datetime import datetime from typing import Literal, TypedDict - -def _default_clock() -> datetime: - return datetime.now(UTC) - +from better_memory._common import default_clock Kind = Literal["reflection", "semantic"] Classification = Literal["cited", "shaped", "ignored", "misled", "overlooked"] @@ -83,7 +80,7 @@ def __init__( clock: Callable[[], datetime] | None = None, ) -> None: self._conn = conn - self._clock: Callable[[], datetime] = clock or _default_clock + self._clock: Callable[[], datetime] = clock or default_clock # --------------------------------------------------------------- credit_one def credit_one( diff --git a/better_memory/services/observation.py b/better_memory/services/observation.py index 86dc594..cfad6fe 100644 --- a/better_memory/services/observation.py +++ b/better_memory/services/observation.py @@ -52,17 +52,17 @@ from __future__ import annotations -import os import sqlite3 from collections.abc import Callable from dataclasses import dataclass -from datetime import UTC, datetime +from datetime import datetime from typing import Any, Literal from uuid import uuid4 import sqlite_vec from better_memory import _diag +from better_memory._common import default_clock, get_session_id from better_memory.config import get_config, project_name from better_memory.search.hybrid import ( SearchFilters, @@ -86,11 +86,6 @@ class BucketedResults: neutral: list[SearchResult] -def _default_clock() -> datetime: - """UTC-aware ``now``. Kept as a module-level function for clarity.""" - return datetime.now(UTC) - - class ObservationService: """Service for creating observations and recording their reinforcement. @@ -126,26 +121,17 @@ def __init__( ) -> None: self._conn = conn self._embedder = embedder - self._clock: Callable[[], datetime] = clock or _default_clock + self._clock: Callable[[], datetime] = clock or default_clock self._project_resolver: Callable[[], str] = ( project_resolver if project_resolver is not None else project_name ) self._scope_resolver: Callable[[], str | None] = ( scope_resolver if scope_resolver is not None else (lambda: None) ) - # Resolution order: explicit kwarg > CLAUDE_SESSION_ID env var > - # CLAUDE_CODE_SESSION_ID env var > uuid4(). The env vars make - # hook-written events and MCP-written observations share the same - # session id (Claude Code exports CLAUDE_CODE_SESSION_ID; the older - # CLAUDE_SESSION_ID is kept as the primary name for back-compat). + # Resolution order: explicit kwarg > env vars > uuid4() — see + # better_memory._common.get_session_id for the env-var order. self.session_id: str = ( - session_id - if session_id is not None - else ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - or uuid4().hex - ) + session_id if session_id is not None else get_session_id() ) # ``None`` defers to the resolved config value so tests can inject # ``False`` without having to monkeypatch the environment. diff --git a/better_memory/services/reflection.py b/better_memory/services/reflection.py index b26cda6..34442e6 100644 --- a/better_memory/services/reflection.py +++ b/better_memory/services/reflection.py @@ -25,21 +25,17 @@ from __future__ import annotations import json -import os import sqlite3 import sys from collections.abc import Callable from dataclasses import dataclass -from datetime import UTC, datetime, timedelta +from datetime import datetime, timedelta from better_memory import _diag +from better_memory._common import default_clock, env_session_id from better_memory.services.memory_rating import OVERLOOKED_RANKING_WEIGHT -def _default_clock() -> datetime: - return datetime.now(UTC) - - def _later_ts(a: str | None, b: str | None) -> str | None: """Return the later of two ISO-8601 timestamps, treating NULL as -inf. @@ -307,7 +303,7 @@ def __init__( clock: Callable[[], datetime] | None = None, ) -> None: self._conn = conn - self._clock: Callable[[], datetime] = clock or _default_clock + self._clock: Callable[[], datetime] = clock or default_clock @staticmethod def _normalize_tech(tech: str | None) -> str | None: # Mirror EpisodeService.start_foreground / ObservationService.create @@ -1267,10 +1263,7 @@ def retrieve_reflections( # (e.g., test or non-Claude context) — see spec §5.2.1. # track_exposure=False is used by SessionBootstrapService.bootstrap, # which manages its own exposure write via _record_exposure. - sid = ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - ) + sid = env_session_id() if track_exposure: _diag.step(fn, "exposure_track_begin", sid=bool(sid)) if not sid: @@ -1335,7 +1328,7 @@ def __init__( clock: Callable[[], datetime] | None = None, ) -> None: self._conn = conn - self._clock: Callable[[], datetime] = clock or _default_clock + self._clock: Callable[[], datetime] = clock or default_clock def confirm(self, *, reflection_id: str) -> None: """pending_review → confirmed; no-op on confirmed; raise on retired/superseded.""" diff --git a/better_memory/services/retention.py b/better_memory/services/retention.py index ff631f5..64aed6d 100644 --- a/better_memory/services/retention.py +++ b/better_memory/services/retention.py @@ -10,11 +10,9 @@ import sqlite3 from collections.abc import Callable from dataclasses import dataclass -from datetime import UTC, datetime, timedelta +from datetime import datetime, timedelta - -def _default_clock() -> datetime: - return datetime.now(UTC) +from better_memory._common import default_clock @dataclass(frozen=True) @@ -54,7 +52,7 @@ def __init__( clock: Callable[[], datetime] | None = None, ) -> None: self._conn = conn - self._clock: Callable[[], datetime] = clock or _default_clock + self._clock: Callable[[], datetime] = clock or default_clock # --------------------------------------------------------- public diff --git a/better_memory/services/retention_scheduler.py b/better_memory/services/retention_scheduler.py index 82307f8..9fb18bb 100644 --- a/better_memory/services/retention_scheduler.py +++ b/better_memory/services/retention_scheduler.py @@ -10,9 +10,10 @@ import sqlite3 from collections.abc import Callable -from datetime import UTC, datetime, timedelta +from datetime import datetime, timedelta from better_memory import _diag +from better_memory._common import default_clock from better_memory.services.retention import RetentionReport, RetentionService _RETENTION_DAYS = 90 @@ -20,10 +21,6 @@ _GUARD_HOURS = 24 -def _default_clock() -> datetime: - return datetime.now(UTC) - - class RetentionScheduler: """24h-guarded wrapper around RetentionService.""" @@ -36,7 +33,7 @@ def __init__( ) -> None: self._conn = conn self._auto_prune = auto_prune - self._clock: Callable[[], datetime] = clock or _default_clock + self._clock: Callable[[], datetime] = clock or default_clock def maybe_run(self, *, triggered_by: str) -> None: """Run retention IF >24h since last run. Records to retention_runs. diff --git a/better_memory/services/semantic.py b/better_memory/services/semantic.py index 9e7f50e..840010c 100644 --- a/better_memory/services/semantic.py +++ b/better_memory/services/semantic.py @@ -9,20 +9,16 @@ from __future__ import annotations -import os import sqlite3 from collections.abc import Callable from dataclasses import dataclass -from datetime import UTC, datetime +from datetime import datetime from uuid import uuid4 +from better_memory._common import default_clock, env_session_id from better_memory.services.memory_rating import OVERLOOKED_RANKING_WEIGHT -def _default_clock() -> datetime: - return datetime.now(UTC) - - @dataclass(frozen=True) class SemanticMemory: """Read model returned by retrieve.""" @@ -58,7 +54,7 @@ def __init__( clock: Callable[[], datetime] | None = None, ) -> None: self._conn = conn - self._clock: Callable[[], datetime] = clock or _default_clock + self._clock: Callable[[], datetime] = clock or default_clock def create( self, *, content: str, project: str, scope: str = "project" @@ -266,10 +262,7 @@ def list_for_project( # Best-effort exposure tracking — see spec §5.2.1. # track_exposure=False is used by SessionBootstrapService.bootstrap, # which manages its own exposure write via _record_exposure. - sid = ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - ) + sid = env_session_id() if track_exposure: if not sid: # Best-effort: bump diagnostics counter. Swallow any error so diff --git a/better_memory/storage/session.py b/better_memory/storage/session.py index 1aeb61a..b98d0e1 100644 --- a/better_memory/storage/session.py +++ b/better_memory/storage/session.py @@ -10,10 +10,10 @@ from __future__ import annotations -import os import re from typing import Literal -from uuid import uuid4 + +from better_memory._common import get_session_id _NamespaceKind = Literal["reflections", "episodes", "semantic", "retired"] _VALID_NAMESPACE_KINDS: tuple[_NamespaceKind, ...] = ( @@ -34,13 +34,9 @@ def resolve_actor_id(project: str | None) -> str: def resolve_session_id() -> str: """Return the current Claude session id, generating one if no env var - is set. Reads CLAUDE_SESSION_ID first, then CLAUDE_CODE_SESSION_ID, - then generates a uuid4 hex (32 chars).""" - return ( - os.environ.get("CLAUDE_SESSION_ID") - or os.environ.get("CLAUDE_CODE_SESSION_ID") - or uuid4().hex - ) + is set. Delegates to :func:`better_memory._common.get_session_id` + (CLAUDE_SESSION_ID, then CLAUDE_CODE_SESSION_ID, then uuid4 hex).""" + return get_session_id() def resolve_namespace(actor_id: str, kind: _NamespaceKind) -> str: diff --git a/tests/hooks/test_session_close_agentcore.py b/tests/hooks/test_session_close_agentcore.py index 194a3db..313de18 100644 --- a/tests/hooks/test_session_close_agentcore.py +++ b/tests/hooks/test_session_close_agentcore.py @@ -104,7 +104,7 @@ def test_spool_marker_written_even_when_closure_event_raises( ) # Force the hook to read from agentcore_config_present's tmp_path monkeypatch.setattr( - "better_memory.hooks.session_close._default_spool_dir", + "better_memory.hooks.session_close.default_spool_dir", lambda: Path(agentcore_config_present) / "spool", ) # Feed an empty stdin so the hook synthesises the marker diff --git a/tests/test_common.py b/tests/test_common.py new file mode 100644 index 0000000..517e782 --- /dev/null +++ b/tests/test_common.py @@ -0,0 +1,90 @@ +"""Tests for the shared lightweight helpers in better_memory._common.""" + +from __future__ import annotations + +import subprocess +import sys +from datetime import UTC + +from better_memory import _common + + +class TestEnvSessionId: + def test_claude_session_id_wins(self, monkeypatch) -> None: + monkeypatch.setenv("CLAUDE_SESSION_ID", "primary") + monkeypatch.setenv("CLAUDE_CODE_SESSION_ID", "secondary") + assert _common.env_session_id() == "primary" + + def test_falls_back_to_claude_code_session_id(self, monkeypatch) -> None: + monkeypatch.delenv("CLAUDE_SESSION_ID", raising=False) + monkeypatch.setenv("CLAUDE_CODE_SESSION_ID", "secondary") + assert _common.env_session_id() == "secondary" + + def test_none_when_no_env(self, monkeypatch) -> None: + monkeypatch.delenv("CLAUDE_SESSION_ID", raising=False) + monkeypatch.delenv("CLAUDE_CODE_SESSION_ID", raising=False) + assert _common.env_session_id() is None + + +class TestGetSessionId: + def test_uses_env_when_set(self, monkeypatch) -> None: + monkeypatch.setenv("CLAUDE_SESSION_ID", "env-session") + assert _common.get_session_id() == "env-session" + + def test_generates_uuid_hex_when_no_env(self, monkeypatch) -> None: + monkeypatch.delenv("CLAUDE_SESSION_ID", raising=False) + monkeypatch.delenv("CLAUDE_CODE_SESSION_ID", raising=False) + sid = _common.get_session_id() + assert len(sid) == 32 + int(sid, 16) # raises if not hex + + +class TestDefaultClock: + def test_returns_utc_aware_now(self) -> None: + now = _common.default_clock() + assert now.tzinfo is UTC + + +class TestResolveHome: + def test_env_override(self, monkeypatch, tmp_path) -> None: + monkeypatch.setenv("BETTER_MEMORY_HOME", str(tmp_path)) + assert _common.resolve_home() == tmp_path + + def test_default_is_dot_better_memory(self, monkeypatch) -> None: + monkeypatch.delenv("BETTER_MEMORY_HOME", raising=False) + home = _common.resolve_home() + assert home.name == ".better-memory" + + +class TestDefaultSpoolDir: + def test_under_home(self, monkeypatch, tmp_path) -> None: + monkeypatch.setenv("BETTER_MEMORY_HOME", str(tmp_path)) + assert _common.default_spool_dir() == tmp_path / "spool" + + +class TestSafeTimestamp: + def test_replaces_colons(self) -> None: + assert ( + _common.safe_timestamp("2026-06-11T12:34:56+00:00") + == "2026-06-11T12-34-56+00-00" + ) + + def test_falls_back_to_now_when_empty(self) -> None: + out = _common.safe_timestamp(None) + assert ":" not in out + assert out.startswith("20") + + +def test_import_is_lightweight() -> None: + """Hooks import _common at startup — it must not drag in sqlite3, + boto3, or better_memory.config.""" + code = ( + "import sys; import better_memory._common; " + "banned = {'sqlite3', 'boto3', 'better_memory.config'}; " + "loaded = banned & set(sys.modules); " + "sys.exit(1 if loaded else 0)" + ) + result = subprocess.run( + [sys.executable, "-c", code], capture_output=True, text=True + ) + assert result.returncode == 0, result.stderr