diff --git a/src/bub/builtin/agent.py b/src/bub/builtin/agent.py index 8406d111..da92c822 100644 --- a/src/bub/builtin/agent.py +++ b/src/bub/builtin/agent.py @@ -22,12 +22,12 @@ is_context_length_error, ) from bub.builtin.settings import load_settings -from bub.builtin.tape import Tape from bub.envelope import field_of from bub.framework import BubFramework from bub.skills import discover_skills, render_skills_prompt +from bub.store import AsyncTapeStoreAdapter, InMemoryTapeStore, is_async_tape_store from bub.streaming import AsyncStreamEvents, StreamEvent, StreamState -from bub.tape import AsyncTapeStoreAdapter, InMemoryTapeStore, is_async_tape_store +from bub.tape import Tape from bub.tools import ( REGISTRY, Tool, diff --git a/src/bub/builtin/hook_impl.py b/src/bub/builtin/hook_impl.py index a4d5a29e..4df649dd 100644 --- a/src/bub/builtin/hook_impl.py +++ b/src/bub/builtin/hook_impl.py @@ -21,8 +21,9 @@ from bub.hooks import hookimpl from bub.hooks.interception import ToolCall, ToolCallDecision from bub.model_selection import ModelChoice, ModelOptions +from bub.store import TapeStore from bub.streaming import AsyncStreamEvents -from bub.tape import TapeContext, TapeStore +from bub.tape import TapeContext from bub.turn import TurnState AGENTS_FILE_NAME = "AGENTS.md" @@ -357,7 +358,7 @@ def render_outbound( @hookimpl def provide_tape_store(self) -> TapeStore: import bub - from bub.builtin.store import FileTapeStore + from bub.store import FileTapeStore return FileTapeStore(directory=bub.home / "tapes") diff --git a/src/bub/builtin/model_runner.py b/src/bub/builtin/model_runner.py index aae8e1c0..099f3f68 100644 --- a/src/bub/builtin/model_runner.py +++ b/src/bub/builtin/model_runner.py @@ -28,7 +28,6 @@ from bub.builtin.codex_provider import OpenaiCodexProvider, should_use_openai_codex_provider from bub.builtin.settings import AgentSettings, ModelCandidate -from bub.builtin.tape import Tape from bub.errors import BubError, ErrorKind from bub.hooks.interception import ( AgentHooks, @@ -37,6 +36,7 @@ LlmCallResult, ) from bub.streaming import AsyncStreamEvents, StreamEvent, StreamState +from bub.tape import Tape from bub.tools import Tool, ToolContext, ToolExecutor CONTEXT_LENGTH_PATTERNS = re.compile( diff --git a/src/bub/builtin/store.py b/src/bub/builtin/store.py deleted file mode 100644 index c5f33adf..00000000 --- a/src/bub/builtin/store.py +++ /dev/null @@ -1,284 +0,0 @@ -from __future__ import annotations - -import itertools -import json -import re -import threading -from collections.abc import Iterable -from dataclasses import asdict, replace -from datetime import UTC, datetime -from pathlib import Path -from typing import Any - -from loguru import logger - -from bub.tape import ( - AsyncTapeStore, - InMemoryQueryMixin, - InMemoryTapeStore, - TapeEntry, - TapeQuery, -) -from bub.utils import get_entry_text - -WORD_PATTERN = re.compile(r"[a-z0-9_/-]+") -MIN_FUZZY_QUERY_LENGTH = 3 -MIN_FUZZY_SCORE = 80 -MAX_FUZZY_CANDIDATES = 128 - - -class ForkTapeStore: - def __init__(self, parent: AsyncTapeStore, tape: str) -> None: - self._parent = parent - self._store = InMemoryTapeStore() - self._tape = tape - self._tape_was_reset = False - - async def list_tapes(self) -> list[str]: - return await self._parent.list_tapes() - - async def reset(self, tape: str) -> None: - if tape != self._tape: - await self._parent.reset(tape) - return - self._store.reset(tape) - self._tape_was_reset = True - - async def fetch_all(self, query: TapeQuery[AsyncTapeStore]) -> Iterable[TapeEntry]: - parent_entries: Iterable[TapeEntry] = [] - if not (query.tape == self._tape and self._tape_was_reset): - try: - parent_entries = await self._parent.fetch_all(query) - except Exception: - parent_entries = [] - this_entries: list[TapeEntry] = [] - for entry in self._store.read(query.tape) or []: - if query._kinds and entry.kind not in query._kinds: - continue - if entry.kind == "anchor": # noqa: SIM102 - if query._after_last or (query._after_anchor and entry.payload.get("name") == query._after_anchor): - this_entries.clear() - parent_entries = [] - continue - this_entries.append(entry) - return itertools.chain(parent_entries, this_entries) - - @staticmethod - def _redact_prompt(prompt: list[dict]) -> Any: - if not isinstance(prompt, list): - return prompt - new_prompt = [] - for part in prompt: - if part.get("type") == "text": - new_prompt.append(part) - return new_prompt - - @staticmethod - def _redact_payload(payload: dict) -> None: - if "content" in payload: - payload["content"] = ForkTapeStore._redact_prompt(payload["content"]) - elif "prompt" in payload: - payload["prompt"] = ForkTapeStore._redact_prompt(payload["prompt"]) - - async def append(self, tape: str, entry: TapeEntry) -> None: - self._redact_payload(entry.payload) - self._store.append(tape, entry) - - async def merge_back(self) -> None: - if self._tape_was_reset: - await self._parent.reset(self._tape) - entries = self._store.read(self._tape) - if not entries: - return - count = len(entries) - for entry in entries: - await self._parent.append(self._tape, entry) - logger.info(f'Merged {count} entries into tape "{self._tape}"') - - -class FileTapeStore(InMemoryQueryMixin): - """TapeStore implementation that persists tapes as JSONL files under a directory.""" - - def __init__(self, directory: Path) -> None: - self._directory = directory - self._directory.mkdir(parents=True, exist_ok=True) - self._tape_files: dict[str, TapeFile] = {} - - def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: - if not query._query: - result: Iterable[TapeEntry] = super().fetch_all(query) - return result - unlimited_query = replace(query, _limit=None) - entries: Iterable[TapeEntry] = super().fetch_all(unlimited_query) - return self._filter_entries(list(entries), query._query, query._limit or 20) - - def _filter_entries(self, entries: list[TapeEntry], query: str, limit: int) -> list[TapeEntry]: - normalized_query = query.strip().lower() - if not normalized_query: - return [] - results: list[TapeEntry] = [] - seen: set[str] = set() - - count = 0 - for entry in reversed(entries): - payload_text = get_entry_text(entry).lower() - if payload_text in seen: - continue - seen.add(payload_text) - - if normalized_query in payload_text or self._is_fuzzy_match(normalized_query, payload_text): - results.append(entry) - count += 1 - if count >= limit: - break - return results - - @staticmethod - def _is_fuzzy_match(normalized_query: str, payload_text: str) -> bool: - from rapidfuzz import fuzz, process - - if len(normalized_query) < MIN_FUZZY_QUERY_LENGTH: - return False - - query_tokens = WORD_PATTERN.findall(normalized_query) - if not query_tokens: - return False - query_phrase = " ".join(query_tokens) - window_size = len(query_tokens) - - source_tokens = WORD_PATTERN.findall(payload_text) - if not source_tokens: - return False - - candidates: list[str] = [] - for token in source_tokens: - candidates.append(token) - if len(candidates) >= MAX_FUZZY_CANDIDATES: - break - - if window_size > 1: - max_window_start = len(source_tokens) - window_size + 1 - for idx in range(max(0, max_window_start)): - candidates.append(" ".join(source_tokens[idx : idx + window_size])) - if len(candidates) >= MAX_FUZZY_CANDIDATES: - break - - best_match = process.extractOne( - query_phrase, - candidates, - scorer=fuzz.WRatio, - score_cutoff=MIN_FUZZY_SCORE, - ) - return best_match is not None - - def _tape_file(self, tape: str) -> TapeFile: - if tape not in self._tape_files: - self._tape_files[tape] = TapeFile(self._directory / f"{tape}.jsonl") - return self._tape_files[tape] - - def list_tapes(self) -> list[str]: - result: list[str] = [] - for file in self._directory.glob("*.jsonl"): - filename = file.stem - if filename.count("__") != 1: - continue - result.append(filename) - return result - - def reset(self, tape: str) -> None: - self._tape_file(tape).reset() - - def append(self, tape: str, entry: TapeEntry) -> None: - self._tape_file(tape).append(entry) - - def read(self, tape: str) -> list[TapeEntry] | None: - return self._tape_file(tape).read() - - -class TapeFile: - """Helper for one tape file.""" - - def __init__(self, path: Path) -> None: - self.path = path - self._lock = threading.Lock() - self._read_entries: list[TapeEntry] = [] - self._read_offset = 0 - - def _next_id(self) -> int: - if self._read_entries: - return self._read_entries[-1].id + 1 - return 1 - - def _reset(self) -> None: - self._read_entries = [] - self._read_offset = 0 - - def reset(self) -> None: - with self._lock: - if self.path.exists(): - self.path.unlink() - self._reset() - - def read(self) -> list[TapeEntry]: - with self._lock: - return self._read_locked() - - def _read_locked(self) -> list[TapeEntry]: - if not self.path.exists(): - self._reset() - return [] - - file_size = self.path.stat().st_size - if file_size < self._read_offset: - # The file was truncated or replaced, so cached entries are stale. - self._reset() - - with self.path.open("r", encoding="utf-8") as handle: - handle.seek(self._read_offset) - for raw_line in handle: - line = raw_line.strip() - if not line: - continue - try: - payload = json.loads(line) - except json.JSONDecodeError: - continue - entry = self.entry_from_payload(payload) - if entry is not None: - self._read_entries.append(entry) - self._read_offset = handle.tell() - - return list(self._read_entries) - - @staticmethod - def entry_from_payload(payload: object) -> TapeEntry | None: - if not isinstance(payload, dict): - return None - entry_id = payload.get("id") - kind = payload.get("kind") - entry_payload = payload.get("payload") - meta = payload.get("meta") - if not isinstance(entry_id, int): - return None - if not isinstance(kind, str): - return None - if not isinstance(entry_payload, dict): - return None - if not isinstance(meta, dict): - meta = {} - if "date" in payload: - date = payload["date"] - else: - date = datetime.fromtimestamp(payload.get("timestamp", 0.0), tz=UTC).isoformat() - return TapeEntry(entry_id, kind, dict(entry_payload), dict(meta), date) - - def append(self, entry: TapeEntry) -> None: - with self._lock: - # Keep cache and offset in sync before allocating new IDs. - self._read_locked() - with self.path.open("a", encoding="utf-8") as handle: - next_id = self._next_id() - stored = TapeEntry(next_id, entry.kind, dict(entry.payload), dict(entry.meta), entry.date) - handle.write(json.dumps(asdict(stored), ensure_ascii=False) + "\n") - self._read_entries.append(stored) - self._read_offset = handle.tell() diff --git a/src/bub/builtin/tape.py b/src/bub/builtin/tape.py deleted file mode 100644 index 76de4d29..00000000 --- a/src/bub/builtin/tape.py +++ /dev/null @@ -1,250 +0,0 @@ -from __future__ import annotations - -import contextlib -import hashlib -import inspect -import json -from collections.abc import AsyncGenerator, Mapping -from dataclasses import asdict, dataclass, field, replace -from datetime import UTC, datetime -from pathlib import Path -from typing import Any - -from pydantic import BaseModel - -from bub.builtin.store import ForkTapeStore -from bub.errors import BubError -from bub.tape import ( - AsyncTapeStore, - TapeContext, - TapeEntry, - TapeQuery, - build_messages, -) - - -@dataclass(frozen=True) -class TapeInfo: - """Runtime tape info summary.""" - - name: str - entries: int - anchors: int - last_anchor: str | None - entries_since_last_anchor: int - last_token_usage: int | None - last_token_cache_hit_rate: float | None - - -@dataclass(frozen=True) -class AnchorSummary: - """Rendered anchor summary.""" - - name: str - state: dict[str, object] - - -@dataclass(frozen=True) -class Tape: - """Tape abstraction for recording agent interactions.""" - - archive_path: Path - store: AsyncTapeStore - context: TapeContext - _name: str | None = field(default=None, repr=False) - - @property - def name(self) -> str: - if self._name is None: - raise ValueError("tape is not scoped") - return self._name - - def with_context(self, context: TapeContext) -> Tape: - return replace(self, context=context) - - def scoped(self, name: str, context: TapeContext | None = None) -> Tape: - return replace(self, context=context or self.context, _name=name) - - def query(self) -> TapeQuery[AsyncTapeStore]: - return TapeQuery(tape=self.name, store=self.store) - - async def info(self) -> TapeInfo: - entries = list(await self.store.fetch_all(self.query())) - anchors = [(i, entry) for i, entry in enumerate(entries) if entry.kind == "anchor"] - if anchors: - last_anchor = anchors[-1][1].payload.get("name") - entries_since_last_anchor = len(entries) - anchors[-1][0] - 1 - else: - last_anchor = None - entries_since_last_anchor = len(entries) - last_token_usage: int | None = None - last_token_cache_hit_rate: float | None = None - for entry in reversed(entries): - if entry.kind == "event" and entry.payload.get("name") == "run": - data = entry.payload.get("data") - usage = data.get("usage") if isinstance(data, Mapping) else None - if not isinstance(usage, Mapping): - continue - token_usage = usage.get("total_tokens") - if not isinstance(token_usage, int) or isinstance(token_usage, bool): - continue - last_token_usage = token_usage - prompt_tokens = usage.get("prompt_tokens") - prompt_details = usage.get("prompt_tokens_details") - cached_tokens = prompt_details.get("cached_tokens") if isinstance(prompt_details, Mapping) else None - if ( - isinstance(prompt_tokens, int) - and not isinstance(prompt_tokens, bool) - and prompt_tokens > 0 - and isinstance(cached_tokens, int) - and not isinstance(cached_tokens, bool) - ): - last_token_cache_hit_rate = cached_tokens / prompt_tokens - break - return TapeInfo( - name=self.name, - entries=len(entries), - anchors=len(anchors), - last_anchor=str(last_anchor) if last_anchor else None, - entries_since_last_anchor=entries_since_last_anchor, - last_token_usage=last_token_usage, - last_token_cache_hit_rate=last_token_cache_hit_rate, - ) - - async def ensure_bootstrap_anchor(self) -> None: - anchors = list(await self.store.fetch_all(self.query().kinds("anchor"))) - if not anchors: - await self.handoff(name="session/start", state={"owner": "human"}) - - async def anchors(self, limit: int = 20) -> list[AnchorSummary]: - entries = list(await self.store.fetch_all(self.query().kinds("anchor"))) - results: list[AnchorSummary] = [] - for entry in entries[-limit:]: - name = str(entry.payload.get("name", "-")) - state = entry.payload.get("state") - state_dict: dict[str, object] = dict(state) if isinstance(state, dict) else {} - results.append(AnchorSummary(name=name, state=state_dict)) - return results - - async def search(self, query: TapeQuery[AsyncTapeStore]) -> list[TapeEntry]: - return list(await self.store.fetch_all(query)) - - async def append_event(self, name: str, payload: dict[str, Any], **meta: Any) -> None: - await self.store.append(self.name, TapeEntry.event(name, payload, **meta)) - - async def read_messages(self) -> list[dict[str, Any]]: - query = self.context.build_query(self.query()) - entries = await self.store.fetch_all(query) - messages = build_messages(entries, self.context) - if inspect.isawaitable(messages): - messages = await messages - return messages - - async def handoff( - self, - *, - name: str, - state: dict[str, Any] | None = None, - **meta: Any, - ) -> list[TapeEntry]: - tape_name = self.name - entry = TapeEntry.anchor(name, state=state, **meta) - event = TapeEntry.event("handoff", {"name": name, "state": state or {}}, **meta) - await self.store.append(tape_name, entry) - await self.store.append(tape_name, event) - return [entry, event] - - async def record_chat( # noqa: C901 - self, - *, - run_id: str, - system_prompt: str | None, - new_messages: list[dict[str, Any]], - response_text: str | None, - context_error: BubError | None = None, - tool_calls: list[dict[str, Any]] | None = None, - tool_results: list[Any] | None = None, - error: BubError | None = None, - response: Any | None = None, - provider: str | None = None, - model: str | None = None, - usage: dict[str, Any] | None = None, - ) -> None: - tape_name = self.name - meta = {"run_id": run_id} - if system_prompt: - await self.store.append(tape_name, TapeEntry.system(system_prompt, **meta)) - if context_error is not None: - await self.store.append(tape_name, TapeEntry.error(context_error, **meta)) - for message in new_messages: - await self.store.append(tape_name, TapeEntry.message(message, **meta)) - if tool_calls: - await self.store.append(tape_name, TapeEntry.tool_call(tool_calls, **meta)) - if tool_results is not None: - await self.store.append(tape_name, TapeEntry.tool_result(tool_results, **meta)) - if error is not None and error is not context_error: - await self.store.append(tape_name, TapeEntry.error(error, **meta)) - if response_text is not None: - await self.store.append( - tape_name, TapeEntry.message({"role": "assistant", "content": response_text}, **meta) - ) - - data: dict[str, Any] = {"status": "error" if error is not None else "ok"} - resolved_usage = usage or self._extract_usage(response) - if resolved_usage is not None: - data["usage"] = resolved_usage - if provider: - data["provider"] = provider - if model: - data["model"] = model - await self.store.append(tape_name, TapeEntry.event("run", data, **meta)) - - @staticmethod - def _extract_usage(response: object) -> dict[str, Any] | None: - usage = getattr(response, "usage", None) - if usage is None: - return None - if isinstance(usage, dict): - return usage - if isinstance(usage, BaseModel): - payload = usage.model_dump(exclude_none=True) - return payload if isinstance(payload, dict) else None - return None - - async def _archive(self) -> Path: - tape_name = self.name - stamp = datetime.now(UTC).strftime("%Y%m%dT%H%M%SZ") - self.archive_path.mkdir(parents=True, exist_ok=True) - archive_path = self.archive_path / f"{tape_name}.jsonl.{stamp}.bak" - with archive_path.open("w", encoding="utf-8") as f: - for entry in await self.store.fetch_all(self.query()): - f.write(json.dumps(asdict(entry), ensure_ascii=False) + "\n") - return archive_path - - async def reset(self, *, archive: bool = False) -> str: - archive_path: Path | None = None - if archive: - archive_path = await self._archive() - await self.store.reset(self.name) - state = {"owner": "human"} - if archive_path is not None: - state["archived"] = str(archive_path) - await self.handoff(name="session/start", state=state) - return f"Archived: {archive_path}" if archive_path else "ok" - - def session_tape(self, session_id: str, workspace: Path, context: TapeContext | None = None) -> Tape: - workspace_hash = hashlib.md5(str(workspace.resolve()).encode("utf-8"), usedforsecurity=False).hexdigest()[:16] - tape_name = ( - workspace_hash + "__" + hashlib.md5(session_id.encode("utf-8"), usedforsecurity=False).hexdigest()[:16] - ) - return self.scoped(tape_name, context=context) - - @contextlib.asynccontextmanager - async def fork_tape(self, merge_back: bool = True) -> AsyncGenerator[Tape, None]: - fork_store = ForkTapeStore(self.store, self.name) - forked = replace(self, store=fork_store) - try: - yield forked - finally: - if merge_back: - await fork_store.merge_back() diff --git a/src/bub/channels/cli/__init__.py b/src/bub/channels/cli/__init__.py index 2416526f..cf559766 100644 --- a/src/bub/channels/cli/__init__.py +++ b/src/bub/channels/cli/__init__.py @@ -25,7 +25,6 @@ import bub from bub.builtin.agent import Agent -from bub.builtin.tape import TapeInfo from bub.channels.admission import AdmitDecision, TurnSnapshot from bub.channels.base import Interface from bub.channels.cli.ansi_bridge import render_to_ansi @@ -39,6 +38,7 @@ from bub.channels.message import ChannelMessage from bub.envelope import Envelope, field_of from bub.streaming import StreamEvent +from bub.tape import TapeInfo from bub.tools import REGISTRY, tool_call_reporter _GENERATION_SPINNER: str = SPINNERS["dots"]["frames"] # type: ignore[assignment] diff --git a/src/bub/framework.py b/src/bub/framework.py index 054df5ea..b11f1cb9 100644 --- a/src/bub/framework.py +++ b/src/bub/framework.py @@ -22,7 +22,8 @@ from bub.hooks.runtime import _SKIP_VALUE, HookRuntime from bub.hooks.specs import BUB_HOOK_NAMESPACE, BubHookSpecs from bub.model_selection import ModelOptions -from bub.tape import AsyncTapeStore, TapeContext, TapeStore +from bub.store import AsyncTapeStore, TapeStore +from bub.tape import TapeContext from bub.turn import TurnResult, TurnState from bub.utils import maybe_context_manager diff --git a/src/bub/hooks/specs.py b/src/bub/hooks/specs.py index ca303c8c..6ddf93f3 100644 --- a/src/bub/hooks/specs.py +++ b/src/bub/hooks/specs.py @@ -19,8 +19,9 @@ ToolCallResult, ) from bub.model_selection import ModelOptions +from bub.store import AsyncTapeStore, TapeStore from bub.streaming import AsyncStreamEvents -from bub.tape import AsyncTapeStore, TapeContext, TapeStore +from bub.tape import TapeContext from bub.turn import TurnState if TYPE_CHECKING: diff --git a/src/bub/store.py b/src/bub/store.py new file mode 100644 index 00000000..addb39e5 --- /dev/null +++ b/src/bub/store.py @@ -0,0 +1,537 @@ +from __future__ import annotations + +import asyncio +import inspect +import itertools +import json +import re +import threading +from collections.abc import Coroutine, Iterable, Sequence +from dataclasses import asdict, dataclass, field, replace +from datetime import UTC, datetime, time +from datetime import date as date_type +from pathlib import Path +from typing import Any, NoReturn, Protocol, Self, overload + +from loguru import logger +from typing_extensions import TypeIs + +from bub.errors import BubError, ErrorKind +from bub.tape import TapeEntry +from bub.utils import get_entry_text + +WORD_PATTERN = re.compile(r"[a-z0-9_/-]+") +MIN_FUZZY_QUERY_LENGTH = 3 +MIN_FUZZY_SCORE = 80 +MAX_FUZZY_CANDIDATES = 128 + + +class TapeStore(Protocol): + """Append-only tape storage interface.""" + + def list_tapes(self) -> list[str]: ... + + def reset(self, tape: str) -> None: ... + + def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: ... + + def append(self, tape: str, entry: TapeEntry) -> None: ... + + +class AsyncTapeStore(Protocol): + """Async append-only tape storage interface.""" + + async def list_tapes(self) -> list[str]: ... + + async def reset(self, tape: str) -> None: ... + + async def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: ... + + async def append(self, tape: str, entry: TapeEntry) -> None: ... + + +def is_async_tape_store(store: TapeStore | AsyncTapeStore) -> TypeIs[AsyncTapeStore]: + return hasattr(store, "append") and inspect.iscoroutinefunction(store.append) + + +@dataclass(frozen=True) +class TapeQuery[T: TapeStore | AsyncTapeStore]: + tape: str + store: T + _query: str | None = None + _after_anchor: str | None = None + _after_last: bool = False + _between_anchors: tuple[str, str] | None = None + _between_dates: tuple[str, str] | None = None + _kinds: tuple[str, ...] = field(default_factory=tuple) + _limit: int | None = None + + def query(self, value: str) -> Self: + return replace(self, _query=value) + + def after_anchor(self, name: str) -> Self: + if not name: + return replace(self, _after_anchor=None, _after_last=False) + return replace(self, _after_anchor=name, _after_last=False) + + def last_anchor(self) -> Self: + return replace(self, _after_anchor=None, _after_last=True) + + def between_anchors(self, start: str, end: str) -> Self: + return replace(self, _between_anchors=(start, end)) + + def between_dates(self, start: str | date_type, end: str | date_type) -> Self: + start_value = start.isoformat() if isinstance(start, date_type) else start + end_value = end.isoformat() if isinstance(end, date_type) else end + return replace(self, _between_dates=(start_value, end_value)) + + def kinds(self, *kinds: str) -> Self: + return replace(self, _kinds=kinds) + + def limit(self, value: int) -> Self: + return replace(self, _limit=value) + + @overload + def all(self: TapeQuery[TapeStore]) -> Iterable[TapeEntry]: ... + + @overload + async def all(self: TapeQuery[AsyncTapeStore]) -> Iterable[TapeEntry]: ... + + def all(self) -> Iterable[TapeEntry] | Coroutine[None, None, Iterable[TapeEntry]]: + return self.store.fetch_all(self) + + +def _anchor_index( + entries: Sequence[TapeEntry], + name: str | None, + *, + default: int, + forward: bool, + start: int = 0, +) -> int: + rng = range(start, len(entries)) if forward else range(len(entries) - 1, start - 1, -1) + for idx in rng: + entry = entries[idx] + if entry.kind != "anchor": + continue + if name is not None and entry.payload.get("name") != name: + continue + return idx + return default + + +def _parse_datetime_boundary(value: str, *, is_end: bool) -> datetime: + if "T" not in value and " " not in value: + try: + parsed_date = date_type.fromisoformat(value) + except ValueError: + pass + else: + boundary_time = time.max if is_end else time.min + return datetime.combine(parsed_date, boundary_time, tzinfo=UTC) + try: + parsed = datetime.fromisoformat(value) + except ValueError: + try: + parsed_date = date_type.fromisoformat(value) + except ValueError as exc: + raise BubError(ErrorKind.INVALID_INPUT, f"Invalid ISO date or datetime: '{value}'.") from exc + boundary_time = time.max if is_end else time.min + parsed = datetime.combine(parsed_date, boundary_time, tzinfo=UTC) + if parsed.tzinfo is None: + return parsed.replace(tzinfo=UTC) + return parsed.astimezone(UTC) + + +def _entry_in_datetime_range(entry: TapeEntry, start_dt: datetime, end_dt: datetime) -> bool: + entry_dt = _parse_datetime_boundary(entry.date, is_end=False) + return start_dt <= entry_dt <= end_dt + + +def _entry_matches_query(entry: TapeEntry, query: str) -> bool: + needle = query.casefold() + haystack = json.dumps( + { + "kind": entry.kind, + "date": entry.date, + "payload": entry.payload, + "meta": entry.meta, + }, + sort_keys=True, + default=str, + ).casefold() + return needle in haystack + + +class InMemoryQueryMixin: + """Mixin to implement in-memory query support for simple stores.""" + + def read(self, tape: str) -> list[TapeEntry] | None: + raise NotImplementedError("InMemoryQueryMixin requires a read() method to be implemented.") + + def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: # noqa: C901 + entries = self.read(query.tape) or [] + start_index = 0 + end_index: int | None = None + + if query._between_anchors is not None: + start_name, end_name = query._between_anchors + start_idx = _anchor_index(entries, start_name, default=-1, forward=False) + if start_idx < 0: + raise BubError(ErrorKind.NOT_FOUND, f"Anchor '{start_name}' was not found.") + end_idx = _anchor_index(entries, end_name, default=-1, forward=True, start=start_idx + 1) + if end_idx < 0: + raise BubError(ErrorKind.NOT_FOUND, f"Anchor '{end_name}' was not found.") + start_index = min(start_idx + 1, len(entries)) + end_index = min(max(start_index, end_idx), len(entries)) + elif query._after_last: + anchor_index = _anchor_index(entries, None, default=-1, forward=False) + if anchor_index < 0: + raise BubError(ErrorKind.NOT_FOUND, "No anchors found in tape.") + start_index = min(anchor_index + 1, len(entries)) + elif query._after_anchor is not None: + anchor_index = _anchor_index(entries, query._after_anchor, default=-1, forward=False) + if anchor_index < 0: + raise BubError(ErrorKind.NOT_FOUND, f"Anchor '{query._after_anchor}' was not found.") + start_index = min(anchor_index + 1, len(entries)) + + sliced = entries[start_index:end_index] + if query._between_dates is not None: + start_date, end_date = query._between_dates + start_dt = _parse_datetime_boundary(start_date, is_end=False) + end_dt = _parse_datetime_boundary(end_date, is_end=True) + if start_dt > end_dt: + raise BubError(ErrorKind.INVALID_INPUT, "Start date must be earlier than or equal to end date.") + sliced = [entry for entry in sliced if _entry_in_datetime_range(entry, start_dt, end_dt)] + if query._query: + sliced = [entry for entry in sliced if _entry_matches_query(entry, query._query)] + if query._kinds: + sliced = [entry for entry in sliced if entry.kind in query._kinds] + if query._limit is not None: + sliced = sliced[: query._limit] + return sliced + + +class InMemoryTapeStore(InMemoryQueryMixin): + """In-memory tape storage.""" + + def __init__(self) -> None: + self._tapes: dict[str, list[TapeEntry]] = {} + self._next_id: dict[str, int] = {} + + def list_tapes(self) -> list[str]: + return sorted(self._tapes.keys()) + + def reset(self, tape: str) -> None: + self._tapes.pop(tape, None) + self._next_id.pop(tape, None) + + def read(self, tape: str) -> list[TapeEntry] | None: + entries = self._tapes.get(tape) + if entries is None: + return None + return [entry.copy() for entry in entries] + + def append(self, tape: str, entry: TapeEntry) -> None: + next_id = self._next_id.get(tape, 1) + self._next_id[tape] = next_id + 1 + stored = TapeEntry(next_id, entry.kind, dict(entry.payload), dict(entry.meta), entry.date) + self._tapes.setdefault(tape, []).append(stored) + + +class AsyncTapeStoreAdapter: + """Adapt a sync TapeStore to AsyncTapeStore.""" + + def __init__(self, store: TapeStore) -> None: + self._store = store + + async def list_tapes(self) -> list[str]: + return await asyncio.to_thread(self._store.list_tapes) + + async def reset(self, tape: str) -> None: + await asyncio.to_thread(self._store.reset, tape) + + async def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: + return await asyncio.to_thread(self._store.fetch_all, query) + + async def append(self, tape: str, entry: TapeEntry) -> None: + await asyncio.to_thread(self._store.append, tape, entry) + + +class UnavailableTapeStore: + """Sync TapeStore sentinel that always fails with a clear message.""" + + def __init__(self, message: str) -> None: + self._message = message + + def _raise(self) -> NoReturn: + raise BubError(ErrorKind.INVALID_INPUT, self._message) + + def list_tapes(self) -> list[str]: + self._raise() + + def reset(self, tape: str) -> None: + self._raise() + + def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: + self._raise() + + def append(self, tape: str, entry: TapeEntry) -> None: + self._raise() + + +class ForkTapeStore: + def __init__(self, parent: AsyncTapeStore, tape: str) -> None: + self._parent = parent + self._store = InMemoryTapeStore() + self._tape = tape + self._tape_was_reset = False + + async def list_tapes(self) -> list[str]: + return await self._parent.list_tapes() + + async def reset(self, tape: str) -> None: + if tape != self._tape: + await self._parent.reset(tape) + return + self._store.reset(tape) + self._tape_was_reset = True + + async def fetch_all(self, query: TapeQuery[AsyncTapeStore]) -> Iterable[TapeEntry]: + parent_entries: Iterable[TapeEntry] = [] + if not (query.tape == self._tape and self._tape_was_reset): + try: + parent_entries = await self._parent.fetch_all(query) + except Exception: + parent_entries = [] + this_entries: list[TapeEntry] = [] + for entry in self._store.read(query.tape) or []: + if query._kinds and entry.kind not in query._kinds: + continue + if entry.kind == "anchor": # noqa: SIM102 + if query._after_last or (query._after_anchor and entry.payload.get("name") == query._after_anchor): + this_entries.clear() + parent_entries = [] + continue + this_entries.append(entry) + return itertools.chain(parent_entries, this_entries) + + @staticmethod + def _redact_prompt(prompt: list[dict]) -> Any: + if not isinstance(prompt, list): + return prompt + new_prompt = [] + for part in prompt: + if part.get("type") == "text": + new_prompt.append(part) + return new_prompt + + @staticmethod + def _redact_payload(payload: dict) -> None: + if "content" in payload: + payload["content"] = ForkTapeStore._redact_prompt(payload["content"]) + elif "prompt" in payload: + payload["prompt"] = ForkTapeStore._redact_prompt(payload["prompt"]) + + async def append(self, tape: str, entry: TapeEntry) -> None: + self._redact_payload(entry.payload) + self._store.append(tape, entry) + + async def merge_back(self) -> None: + if self._tape_was_reset: + await self._parent.reset(self._tape) + entries = self._store.read(self._tape) + if not entries: + return + count = len(entries) + for entry in entries: + await self._parent.append(self._tape, entry) + logger.info(f'Merged {count} entries into tape "{self._tape}"') + + +class FileTapeStore(InMemoryQueryMixin): + """TapeStore implementation that persists tapes as JSONL files under a directory.""" + + def __init__(self, directory: Path) -> None: + self._directory = directory + self._directory.mkdir(parents=True, exist_ok=True) + self._tape_files: dict[str, TapeFile] = {} + + def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: + if not query._query: + result: Iterable[TapeEntry] = super().fetch_all(query) + return result + unlimited_query = replace(query, _limit=None) + entries: Iterable[TapeEntry] = super().fetch_all(unlimited_query) + return self._filter_entries(list(entries), query._query, query._limit or 20) + + def _filter_entries(self, entries: list[TapeEntry], query: str, limit: int) -> list[TapeEntry]: + normalized_query = query.strip().lower() + if not normalized_query: + return [] + results: list[TapeEntry] = [] + seen: set[str] = set() + + count = 0 + for entry in reversed(entries): + payload_text = get_entry_text(entry).lower() + if payload_text in seen: + continue + seen.add(payload_text) + + if normalized_query in payload_text or self._is_fuzzy_match(normalized_query, payload_text): + results.append(entry) + count += 1 + if count >= limit: + break + return results + + @staticmethod + def _is_fuzzy_match(normalized_query: str, payload_text: str) -> bool: + from rapidfuzz import fuzz, process + + if len(normalized_query) < MIN_FUZZY_QUERY_LENGTH: + return False + + query_tokens = WORD_PATTERN.findall(normalized_query) + if not query_tokens: + return False + query_phrase = " ".join(query_tokens) + window_size = len(query_tokens) + + source_tokens = WORD_PATTERN.findall(payload_text) + if not source_tokens: + return False + + candidates: list[str] = [] + for token in source_tokens: + candidates.append(token) + if len(candidates) >= MAX_FUZZY_CANDIDATES: + break + + if window_size > 1: + max_window_start = len(source_tokens) - window_size + 1 + for idx in range(max(0, max_window_start)): + candidates.append(" ".join(source_tokens[idx : idx + window_size])) + if len(candidates) >= MAX_FUZZY_CANDIDATES: + break + + best_match = process.extractOne( + query_phrase, + candidates, + scorer=fuzz.WRatio, + score_cutoff=MIN_FUZZY_SCORE, + ) + return best_match is not None + + def _tape_file(self, tape: str) -> TapeFile: + if tape not in self._tape_files: + self._tape_files[tape] = TapeFile(self._directory / f"{tape}.jsonl") + return self._tape_files[tape] + + def list_tapes(self) -> list[str]: + result: list[str] = [] + for file in self._directory.glob("*.jsonl"): + filename = file.stem + if filename.count("__") != 1: + continue + result.append(filename) + return result + + def reset(self, tape: str) -> None: + self._tape_file(tape).reset() + + def append(self, tape: str, entry: TapeEntry) -> None: + self._tape_file(tape).append(entry) + + def read(self, tape: str) -> list[TapeEntry] | None: + return self._tape_file(tape).read() + + +class TapeFile: + """Helper for one tape file.""" + + def __init__(self, path: Path) -> None: + self.path = path + self._lock = threading.Lock() + self._read_entries: list[TapeEntry] = [] + self._read_offset = 0 + + def _next_id(self) -> int: + if self._read_entries: + return self._read_entries[-1].id + 1 + return 1 + + def _reset(self) -> None: + self._read_entries = [] + self._read_offset = 0 + + def reset(self) -> None: + with self._lock: + if self.path.exists(): + self.path.unlink() + self._reset() + + def read(self) -> list[TapeEntry]: + with self._lock: + return self._read_locked() + + def _read_locked(self) -> list[TapeEntry]: + if not self.path.exists(): + self._reset() + return [] + + file_size = self.path.stat().st_size + if file_size < self._read_offset: + # The file was truncated or replaced, so cached entries are stale. + self._reset() + + with self.path.open("r", encoding="utf-8") as handle: + handle.seek(self._read_offset) + for raw_line in handle: + line = raw_line.strip() + if not line: + continue + try: + payload = json.loads(line) + except json.JSONDecodeError: + continue + entry = self.entry_from_payload(payload) + if entry is not None: + self._read_entries.append(entry) + self._read_offset = handle.tell() + + return list(self._read_entries) + + @staticmethod + def entry_from_payload(payload: object) -> TapeEntry | None: + if not isinstance(payload, dict): + return None + entry_id = payload.get("id") + kind = payload.get("kind") + entry_payload = payload.get("payload") + meta = payload.get("meta") + if not isinstance(entry_id, int): + return None + if not isinstance(kind, str): + return None + if not isinstance(entry_payload, dict): + return None + if not isinstance(meta, dict): + meta = {} + if "date" in payload: + date = payload["date"] + else: + date = datetime.fromtimestamp(payload.get("timestamp", 0.0), tz=UTC).isoformat() + return TapeEntry(entry_id, kind, dict(entry_payload), dict(meta), date) + + def append(self, entry: TapeEntry) -> None: + with self._lock: + # Keep cache and offset in sync before allocating new IDs. + self._read_locked() + with self.path.open("a", encoding="utf-8") as handle: + next_id = self._next_id() + stored = TapeEntry(next_id, entry.kind, dict(entry.payload), dict(entry.meta), entry.date) + handle.write(json.dumps(asdict(stored), ensure_ascii=False) + "\n") + self._read_entries.append(stored) + self._read_offset = handle.tell() diff --git a/src/bub/tape.py b/src/bub/tape.py index b48b3066..b4b3fc54 100644 --- a/src/bub/tape.py +++ b/src/bub/tape.py @@ -2,18 +2,78 @@ from __future__ import annotations -import asyncio +import contextlib +import hashlib import inspect import json -from collections.abc import Callable, Coroutine, Iterable, Sequence -from dataclasses import dataclass, field, replace -from datetime import UTC, datetime, time -from datetime import date as date_type -from typing import Any, NoReturn, Protocol, Self, overload - -from typing_extensions import TypeIs - -from bub.errors import BubError, ErrorKind +from collections.abc import AsyncGenerator, Callable, Coroutine, Iterable, Mapping +from dataclasses import asdict, dataclass, field, replace +from datetime import UTC, datetime +from pathlib import Path +from typing import TYPE_CHECKING, Any + +from pydantic import BaseModel + +from bub.errors import BubError + +__all__ = [ + "LAST_ANCHOR", + "AnchorSelector", + "AnchorSummary", + "AsyncTapeStore", + "AsyncTapeStoreAdapter", + "ContextSelector", + "InMemoryQueryMixin", + "InMemoryTapeStore", + "SelectedMessages", + "Tape", + "TapeContext", + "TapeEntry", + "TapeInfo", + "TapeQuery", + "TapeStore", + "UnavailableTapeStore", + "build_messages", + "is_async_tape_store", + "utc_now", +] + +if TYPE_CHECKING: + from bub.store import ( + AsyncTapeStore, + AsyncTapeStoreAdapter, + InMemoryQueryMixin, + InMemoryTapeStore, + TapeQuery, + TapeStore, + UnavailableTapeStore, + is_async_tape_store, + ) + + +_STORE_EXPORTS = frozenset({ + "AsyncTapeStore", + "AsyncTapeStoreAdapter", + "InMemoryQueryMixin", + "InMemoryTapeStore", + "TapeQuery", + "TapeStore", + "UnavailableTapeStore", + "is_async_tape_store", +}) + + +def __getattr__(name: str) -> Any: + if name not in _STORE_EXPORTS: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + from bub import store + + return getattr(store, name) + + +def __dir__() -> list[str]: + return sorted({*globals(), *_STORE_EXPORTS}) def utc_now() -> str: @@ -68,81 +128,6 @@ def event(cls, name: str, data: dict[str, Any] | None = None, **meta: Any) -> Ta return cls(id=0, kind="event", payload=payload, meta=dict(meta)) -class TapeStore(Protocol): - """Append-only tape storage interface.""" - - def list_tapes(self) -> list[str]: ... - - def reset(self, tape: str) -> None: ... - - def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: ... - - def append(self, tape: str, entry: TapeEntry) -> None: ... - - -class AsyncTapeStore(Protocol): - """Async append-only tape storage interface.""" - - async def list_tapes(self) -> list[str]: ... - - async def reset(self, tape: str) -> None: ... - - async def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: ... - - async def append(self, tape: str, entry: TapeEntry) -> None: ... - - -def is_async_tape_store(store: TapeStore | AsyncTapeStore) -> TypeIs[AsyncTapeStore]: - return hasattr(store, "append") and inspect.iscoroutinefunction(store.append) - - -@dataclass(frozen=True) -class TapeQuery[T: TapeStore | AsyncTapeStore]: - tape: str - store: T - _query: str | None = None - _after_anchor: str | None = None - _after_last: bool = False - _between_anchors: tuple[str, str] | None = None - _between_dates: tuple[str, str] | None = None - _kinds: tuple[str, ...] = field(default_factory=tuple) - _limit: int | None = None - - def query(self, value: str) -> Self: - return replace(self, _query=value) - - def after_anchor(self, name: str) -> Self: - if not name: - return replace(self, _after_anchor=None, _after_last=False) - return replace(self, _after_anchor=name, _after_last=False) - - def last_anchor(self) -> Self: - return replace(self, _after_anchor=None, _after_last=True) - - def between_anchors(self, start: str, end: str) -> Self: - return replace(self, _between_anchors=(start, end)) - - def between_dates(self, start: str | date_type, end: str | date_type) -> Self: - start_value = start.isoformat() if isinstance(start, date_type) else start - end_value = end.isoformat() if isinstance(end, date_type) else end - return replace(self, _between_dates=(start_value, end_value)) - - def kinds(self, *kinds: str) -> Self: - return replace(self, _kinds=kinds) - - def limit(self, value: int) -> Self: - return replace(self, _limit=value) - - @overload - def all(self: TapeQuery[TapeStore]) -> Iterable[TapeEntry]: ... - - @overload - async def all(self: TapeQuery[AsyncTapeStore]) -> Iterable[TapeEntry]: ... - - def all(self) -> Iterable[TapeEntry] | Coroutine[None, None, Iterable[TapeEntry]]: - return self.store.fetch_all(self) - - class _LastAnchor: def __repr__(self) -> str: return "LAST_ANCHOR" @@ -187,180 +172,232 @@ def _default_messages(entries: Iterable[TapeEntry]) -> list[dict[str, Any]]: return messages -def _anchor_index( - entries: Sequence[TapeEntry], - name: str | None, - *, - default: int, - forward: bool, - start: int = 0, -) -> int: - rng = range(start, len(entries)) if forward else range(len(entries) - 1, start - 1, -1) - for idx in rng: - entry = entries[idx] - if entry.kind != "anchor": - continue - if name is not None and entry.payload.get("name") != name: - continue - return idx - return default - - -def _parse_datetime_boundary(value: str, *, is_end: bool) -> datetime: - if "T" not in value and " " not in value: - try: - parsed_date = date_type.fromisoformat(value) - except ValueError: - pass - else: - boundary_time = time.max if is_end else time.min - return datetime.combine(parsed_date, boundary_time, tzinfo=UTC) - try: - parsed = datetime.fromisoformat(value) - except ValueError: - try: - parsed_date = date_type.fromisoformat(value) - except ValueError as exc: - raise BubError(ErrorKind.INVALID_INPUT, f"Invalid ISO date or datetime: '{value}'.") from exc - boundary_time = time.max if is_end else time.min - parsed = datetime.combine(parsed_date, boundary_time, tzinfo=UTC) - if parsed.tzinfo is None: - return parsed.replace(tzinfo=UTC) - return parsed.astimezone(UTC) - - -def _entry_in_datetime_range(entry: TapeEntry, start_dt: datetime, end_dt: datetime) -> bool: - entry_dt = _parse_datetime_boundary(entry.date, is_end=False) - return start_dt <= entry_dt <= end_dt - - -def _entry_matches_query(entry: TapeEntry, query: str) -> bool: - needle = query.casefold() - haystack = json.dumps( - { - "kind": entry.kind, - "date": entry.date, - "payload": entry.payload, - "meta": entry.meta, - }, - sort_keys=True, - default=str, - ).casefold() - return needle in haystack - - -class InMemoryQueryMixin: - """Mixin to implement in-memory query support for simple stores.""" - - def read(self, tape: str) -> list[TapeEntry] | None: - raise NotImplementedError("InMemoryQueryMixin requires a read() method to be implemented.") - - def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: # noqa: C901 - entries = self.read(query.tape) or [] - start_index = 0 - end_index: int | None = None - - if query._between_anchors is not None: - start_name, end_name = query._between_anchors - start_idx = _anchor_index(entries, start_name, default=-1, forward=False) - if start_idx < 0: - raise BubError(ErrorKind.NOT_FOUND, f"Anchor '{start_name}' was not found.") - end_idx = _anchor_index(entries, end_name, default=-1, forward=True, start=start_idx + 1) - if end_idx < 0: - raise BubError(ErrorKind.NOT_FOUND, f"Anchor '{end_name}' was not found.") - start_index = min(start_idx + 1, len(entries)) - end_index = min(max(start_index, end_idx), len(entries)) - elif query._after_last: - anchor_index = _anchor_index(entries, None, default=-1, forward=False) - if anchor_index < 0: - raise BubError(ErrorKind.NOT_FOUND, "No anchors found in tape.") - start_index = min(anchor_index + 1, len(entries)) - elif query._after_anchor is not None: - anchor_index = _anchor_index(entries, query._after_anchor, default=-1, forward=False) - if anchor_index < 0: - raise BubError(ErrorKind.NOT_FOUND, f"Anchor '{query._after_anchor}' was not found.") - start_index = min(anchor_index + 1, len(entries)) - - sliced = entries[start_index:end_index] - if query._between_dates is not None: - start_date, end_date = query._between_dates - start_dt = _parse_datetime_boundary(start_date, is_end=False) - end_dt = _parse_datetime_boundary(end_date, is_end=True) - if start_dt > end_dt: - raise BubError(ErrorKind.INVALID_INPUT, "Start date must be earlier than or equal to end date.") - sliced = [entry for entry in sliced if _entry_in_datetime_range(entry, start_dt, end_dt)] - if query._query: - sliced = [entry for entry in sliced if _entry_matches_query(entry, query._query)] - if query._kinds: - sliced = [entry for entry in sliced if entry.kind in query._kinds] - if query._limit is not None: - sliced = sliced[: query._limit] - return sliced - - -class InMemoryTapeStore(InMemoryQueryMixin): - """In-memory tape storage.""" - - def __init__(self) -> None: - self._tapes: dict[str, list[TapeEntry]] = {} - self._next_id: dict[str, int] = {} - - def list_tapes(self) -> list[str]: - return sorted(self._tapes.keys()) - - def reset(self, tape: str) -> None: - self._tapes.pop(tape, None) - self._next_id.pop(tape, None) - - def read(self, tape: str) -> list[TapeEntry] | None: - entries = self._tapes.get(tape) - if entries is None: - return None - return [entry.copy() for entry in entries] - - def append(self, tape: str, entry: TapeEntry) -> None: - next_id = self._next_id.get(tape, 1) - self._next_id[tape] = next_id + 1 - stored = TapeEntry(next_id, entry.kind, dict(entry.payload), dict(entry.meta), entry.date) - self._tapes.setdefault(tape, []).append(stored) - - -class AsyncTapeStoreAdapter: - """Adapt a sync TapeStore to AsyncTapeStore.""" +@dataclass(frozen=True) +class TapeInfo: + """Runtime tape info summary.""" - def __init__(self, store: TapeStore) -> None: - self._store = store + name: str + entries: int + anchors: int + last_anchor: str | None + entries_since_last_anchor: int + last_token_usage: int | None + last_token_cache_hit_rate: float | None - async def list_tapes(self) -> list[str]: - return await asyncio.to_thread(self._store.list_tapes) - async def reset(self, tape: str) -> None: - await asyncio.to_thread(self._store.reset, tape) +@dataclass(frozen=True) +class AnchorSummary: + """Rendered anchor summary.""" - async def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: - return await asyncio.to_thread(self._store.fetch_all, query) + name: str + state: dict[str, object] - async def append(self, tape: str, entry: TapeEntry) -> None: - await asyncio.to_thread(self._store.append, tape, entry) +@dataclass(frozen=True) +class Tape: + """Tape abstraction for recording agent interactions.""" -class UnavailableTapeStore: - """Sync TapeStore sentinel that always fails with a clear message.""" + archive_path: Path + store: AsyncTapeStore + context: TapeContext + _name: str | None = field(default=None, repr=False) - def __init__(self, message: str) -> None: - self._message = message + @property + def name(self) -> str: + if self._name is None: + raise ValueError("tape is not scoped") + return self._name - def _raise(self) -> NoReturn: - raise BubError(ErrorKind.INVALID_INPUT, self._message) + def with_context(self, context: TapeContext) -> Tape: + return replace(self, context=context) - def list_tapes(self) -> list[str]: - self._raise() + def scoped(self, name: str, context: TapeContext | None = None) -> Tape: + return replace(self, context=context or self.context, _name=name) - def reset(self, tape: str) -> None: - self._raise() + def query(self) -> TapeQuery[AsyncTapeStore]: + from bub.store import TapeQuery - def fetch_all(self, query: TapeQuery) -> Iterable[TapeEntry]: - self._raise() + return TapeQuery(tape=self.name, store=self.store) - def append(self, tape: str, entry: TapeEntry) -> None: - self._raise() + async def info(self) -> TapeInfo: + entries = list(await self.store.fetch_all(self.query())) + anchors = [(i, entry) for i, entry in enumerate(entries) if entry.kind == "anchor"] + if anchors: + last_anchor = anchors[-1][1].payload.get("name") + entries_since_last_anchor = len(entries) - anchors[-1][0] - 1 + else: + last_anchor = None + entries_since_last_anchor = len(entries) + last_token_usage: int | None = None + last_token_cache_hit_rate: float | None = None + for entry in reversed(entries): + if entry.kind == "event" and entry.payload.get("name") == "run": + data = entry.payload.get("data") + usage = data.get("usage") if isinstance(data, Mapping) else None + if not isinstance(usage, Mapping): + continue + token_usage = usage.get("total_tokens") + if not isinstance(token_usage, int) or isinstance(token_usage, bool): + continue + last_token_usage = token_usage + prompt_tokens = usage.get("prompt_tokens") + prompt_details = usage.get("prompt_tokens_details") + cached_tokens = prompt_details.get("cached_tokens") if isinstance(prompt_details, Mapping) else None + if ( + isinstance(prompt_tokens, int) + and not isinstance(prompt_tokens, bool) + and prompt_tokens > 0 + and isinstance(cached_tokens, int) + and not isinstance(cached_tokens, bool) + ): + last_token_cache_hit_rate = cached_tokens / prompt_tokens + break + return TapeInfo( + name=self.name, + entries=len(entries), + anchors=len(anchors), + last_anchor=str(last_anchor) if last_anchor else None, + entries_since_last_anchor=entries_since_last_anchor, + last_token_usage=last_token_usage, + last_token_cache_hit_rate=last_token_cache_hit_rate, + ) + + async def ensure_bootstrap_anchor(self) -> None: + anchors = list(await self.store.fetch_all(self.query().kinds("anchor"))) + if not anchors: + await self.handoff(name="session/start", state={"owner": "human"}) + + async def anchors(self, limit: int = 20) -> list[AnchorSummary]: + entries = list(await self.store.fetch_all(self.query().kinds("anchor"))) + results: list[AnchorSummary] = [] + for entry in entries[-limit:]: + name = str(entry.payload.get("name", "-")) + state = entry.payload.get("state") + state_dict: dict[str, object] = dict(state) if isinstance(state, dict) else {} + results.append(AnchorSummary(name=name, state=state_dict)) + return results + + async def search(self, query: TapeQuery[AsyncTapeStore]) -> list[TapeEntry]: + return list(await self.store.fetch_all(query)) + + async def append_event(self, name: str, payload: dict[str, Any], **meta: Any) -> None: + await self.store.append(self.name, TapeEntry.event(name, payload, **meta)) + + async def read_messages(self) -> list[dict[str, Any]]: + query = self.context.build_query(self.query()) + entries = await self.store.fetch_all(query) + messages = build_messages(entries, self.context) + if inspect.isawaitable(messages): + messages = await messages + return messages + + async def handoff( + self, + *, + name: str, + state: dict[str, Any] | None = None, + **meta: Any, + ) -> list[TapeEntry]: + tape_name = self.name + entry = TapeEntry.anchor(name, state=state, **meta) + event = TapeEntry.event("handoff", {"name": name, "state": state or {}}, **meta) + await self.store.append(tape_name, entry) + await self.store.append(tape_name, event) + return [entry, event] + + async def record_chat( # noqa: C901 + self, + *, + run_id: str, + system_prompt: str | None, + new_messages: list[dict[str, Any]], + response_text: str | None, + context_error: BubError | None = None, + tool_calls: list[dict[str, Any]] | None = None, + tool_results: list[Any] | None = None, + error: BubError | None = None, + response: Any | None = None, + provider: str | None = None, + model: str | None = None, + usage: dict[str, Any] | None = None, + ) -> None: + tape_name = self.name + meta = {"run_id": run_id} + if system_prompt: + await self.store.append(tape_name, TapeEntry.system(system_prompt, **meta)) + if context_error is not None: + await self.store.append(tape_name, TapeEntry.error(context_error, **meta)) + for message in new_messages: + await self.store.append(tape_name, TapeEntry.message(message, **meta)) + if tool_calls: + await self.store.append(tape_name, TapeEntry.tool_call(tool_calls, **meta)) + if tool_results is not None: + await self.store.append(tape_name, TapeEntry.tool_result(tool_results, **meta)) + if error is not None and error is not context_error: + await self.store.append(tape_name, TapeEntry.error(error, **meta)) + if response_text is not None: + await self.store.append( + tape_name, TapeEntry.message({"role": "assistant", "content": response_text}, **meta) + ) + + data: dict[str, Any] = {"status": "error" if error is not None else "ok"} + resolved_usage = usage or self._extract_usage(response) + if resolved_usage is not None: + data["usage"] = resolved_usage + if provider: + data["provider"] = provider + if model: + data["model"] = model + await self.store.append(tape_name, TapeEntry.event("run", data, **meta)) + + @staticmethod + def _extract_usage(response: object) -> dict[str, Any] | None: + usage = getattr(response, "usage", None) + if usage is None: + return None + if isinstance(usage, dict): + return usage + if isinstance(usage, BaseModel): + payload = usage.model_dump(exclude_none=True) + return payload if isinstance(payload, dict) else None + return None + + async def _archive(self) -> Path: + tape_name = self.name + stamp = datetime.now(UTC).strftime("%Y%m%dT%H%M%SZ") + self.archive_path.mkdir(parents=True, exist_ok=True) + archive_path = self.archive_path / f"{tape_name}.jsonl.{stamp}.bak" + with archive_path.open("w", encoding="utf-8") as f: + for entry in await self.store.fetch_all(self.query()): + f.write(json.dumps(asdict(entry), ensure_ascii=False) + "\n") + return archive_path + + async def reset(self, *, archive: bool = False) -> str: + archive_path: Path | None = None + if archive: + archive_path = await self._archive() + await self.store.reset(self.name) + state = {"owner": "human"} + if archive_path is not None: + state["archived"] = str(archive_path) + await self.handoff(name="session/start", state=state) + return f"Archived: {archive_path}" if archive_path else "ok" + + def session_tape(self, session_id: str, workspace: Path, context: TapeContext | None = None) -> Tape: + workspace_hash = hashlib.md5(str(workspace.resolve()).encode("utf-8"), usedforsecurity=False).hexdigest()[:16] + tape_name = ( + workspace_hash + "__" + hashlib.md5(session_id.encode("utf-8"), usedforsecurity=False).hexdigest()[:16] + ) + return self.scoped(tape_name, context=context) + + @contextlib.asynccontextmanager + async def fork_tape(self, merge_back: bool = True) -> AsyncGenerator[Tape, None]: + from bub.store import ForkTapeStore + + fork_store = ForkTapeStore(self.store, self.name) + forked = replace(self, store=fork_store) + try: + yield forked + finally: + if merge_back: + await fork_store.merge_back() diff --git a/src/bub/tools.py b/src/bub/tools.py index 9b3bc01e..057046b4 100644 --- a/src/bub/tools.py +++ b/src/bub/tools.py @@ -13,9 +13,9 @@ from loguru import logger from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError, validate_call -from bub.builtin.tape import Tape from bub.errors import BubError, ErrorKind from bub.hooks.interception import ToolCall, ToolCallResult +from bub.tape import Tape if TYPE_CHECKING: from bub.hooks.interception import AgentHooks diff --git a/tests/test_agent_hooks.py b/tests/test_agent_hooks.py index a4a3ae6c..0a181f90 100644 --- a/tests/test_agent_hooks.py +++ b/tests/test_agent_hooks.py @@ -250,8 +250,8 @@ def _runner_and_tape(self, hooks: AgentHooks, captured: dict): import bub from bub.builtin.model_runner import ModelRunner from bub.builtin.settings import AgentSettings - from bub.builtin.tape import Tape - from bub.tape import AsyncTapeStoreAdapter, InMemoryTapeStore, TapeContext + from bub.store import AsyncTapeStoreAdapter, InMemoryTapeStore + from bub.tape import Tape, TapeContext class FakeRunner(ModelRunner): async def completion_response(self, *, model, messages, tools, max_tokens=None, reasoning_effort=None): diff --git a/tests/test_builtin_hook_impl.py b/tests/test_builtin_hook_impl.py index bda5b92f..2d8d95f4 100644 --- a/tests/test_builtin_hook_impl.py +++ b/tests/test_builtin_hook_impl.py @@ -8,12 +8,11 @@ import pytest from bub.builtin.hook_impl import AGENTS_FILE_NAME, DEFAULT_SYSTEM_PROMPT, BuiltinImpl -from bub.builtin.store import FileTapeStore -from bub.builtin.tape import Tape from bub.channels.message import ChannelMessage from bub.framework import BubFramework +from bub.store import AsyncTapeStoreAdapter, FileTapeStore, InMemoryTapeStore from bub.streaming import AsyncStreamEvents, StreamEvent -from bub.tape import AsyncTapeStoreAdapter, InMemoryTapeStore, TapeContext +from bub.tape import Tape, TapeContext class RecordingLifespan: diff --git a/tests/test_builtin_model_runner.py b/tests/test_builtin_model_runner.py index b2894b3f..374ebdf8 100644 --- a/tests/test_builtin_model_runner.py +++ b/tests/test_builtin_model_runner.py @@ -12,8 +12,7 @@ from bub.builtin.model_runner import ModelRunner, tool_invocation_from_native from bub.builtin.settings import AgentSettings, ModelCandidate -from bub.builtin.tape import Tape -from bub.tape import AsyncTapeStoreAdapter, InMemoryTapeStore, TapeContext +from bub.tape import AsyncTapeStoreAdapter, InMemoryTapeStore, Tape, TapeContext from bub.tools import ToolExecutor diff --git a/tests/test_builtin_tools.py b/tests/test_builtin_tools.py index 873219b0..6296a214 100644 --- a/tests/test_builtin_tools.py +++ b/tests/test_builtin_tools.py @@ -10,7 +10,6 @@ import bub.builtin.tools as builtin_tools from bub.builtin.shell_manager import ShellManager -from bub.builtin.tape import Tape from bub.builtin.tools import ( bash, bash_output, @@ -23,7 +22,8 @@ tape_info, ) from bub.errors import ErrorKind -from bub.tape import AsyncTapeStoreAdapter, InMemoryTapeStore, TapeContext +from bub.store import AsyncTapeStoreAdapter, InMemoryTapeStore +from bub.tape import Tape, TapeContext from bub.tools import REGISTRY, Tool, ToolContext, ToolExecutor, tool diff --git a/tests/test_file_tape_store_entry_ids.py b/tests/test_file_tape_store_entry_ids.py index d891de97..40986f68 100644 --- a/tests/test_file_tape_store_entry_ids.py +++ b/tests/test_file_tape_store_entry_ids.py @@ -2,8 +2,8 @@ import pytest -from bub.builtin.store import FileTapeStore, ForkTapeStore -from bub.tape import AsyncTapeStoreAdapter, TapeEntry +from bub.store import AsyncTapeStoreAdapter, FileTapeStore, ForkTapeStore +from bub.tape import TapeEntry @pytest.mark.asyncio diff --git a/tests/test_fork_store_merge_back.py b/tests/test_fork_store_merge_back.py index c0f3d157..e614900e 100644 --- a/tests/test_fork_store_merge_back.py +++ b/tests/test_fork_store_merge_back.py @@ -2,8 +2,8 @@ import pytest -from bub.builtin.store import ForkTapeStore -from bub.tape import AsyncTapeStoreAdapter, InMemoryTapeStore, TapeEntry, TapeQuery +from bub.store import AsyncTapeStoreAdapter, ForkTapeStore, InMemoryTapeStore, TapeQuery +from bub.tape import TapeEntry @pytest.mark.asyncio diff --git a/tests/test_builtin_tape.py b/tests/test_tape.py similarity index 78% rename from tests/test_builtin_tape.py rename to tests/test_tape.py index 25a0e72a..04e25335 100644 --- a/tests/test_builtin_tape.py +++ b/tests/test_tape.py @@ -4,9 +4,27 @@ import pytest -from bub.builtin.store import ForkTapeStore -from bub.builtin.tape import Tape -from bub.tape import AsyncTapeStoreAdapter, InMemoryTapeStore, TapeContext +from bub.store import AsyncTapeStoreAdapter, ForkTapeStore, InMemoryTapeStore +from bub.tape import Tape, TapeContext + + +def test_tape_reexports_legacy_store_objects() -> None: + from bub import store, tape + + expected_exports = { + "AsyncTapeStore", + "AsyncTapeStoreAdapter", + "InMemoryQueryMixin", + "InMemoryTapeStore", + "TapeQuery", + "TapeStore", + "UnavailableTapeStore", + "is_async_tape_store", + } + + assert expected_exports <= set(dir(tape)) + for name in expected_exports: + assert getattr(tape, name) is getattr(store, name) @pytest.mark.asyncio