Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 17 additions & 26 deletions src/bub/builtin/agent.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,8 +7,7 @@
import re
import shlex
import time
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Collection, Coroutine, Iterable
from contextlib import AsyncExitStack
from collections.abc import AsyncGenerator, AsyncIterator, Collection, Iterable
from dataclasses import dataclass, replace
from datetime import UTC, datetime
from functools import cached_property
Expand Down Expand Up @@ -73,25 +72,11 @@ async def generator() -> AsyncIterator:

return AsyncStreamEvents(generator())

@staticmethod
def _events_with_callback(
events: AsyncStreamEvents, callback: Callable[[], Coroutine[Any, Any, Any]]
) -> AsyncStreamEvents:
async def generator() -> AsyncIterator[StreamEvent]:
try:
async for event in events:
yield event
finally:
await callback()

return AsyncStreamEvents(generator(), state=events._state)

async def run_stream(
self,
*,
session_id: str,
tape: Tape,
prompt: str | list[dict],
state: TurnState,
model: str | None = None,
allowed_skills: Collection[str] | None = None,
allowed_tools: Collection[str] | None = None,
Expand All @@ -102,14 +87,6 @@ async def run_stream(
StreamEvent("final", {"text": "error: empty prompt", "ok": False}),
])

state.setdefault("session_id", session_id)
tape = self.tape.session_tape(
session_id, workspace_from_state(state), context=replace(self.tape.context, state=state)
)
merge_back = not session_id.startswith("temp/")
stack = AsyncExitStack()
# The fork_tape context manager must not be exited until the last chunk of the stream is consumed.
tape = await stack.enter_async_context(tape.fork_tape(merge_back=merge_back))
await tape.ensure_bootstrap_anchor()
if isinstance(prompt, str) and prompt.strip().startswith(","):
result = await self._run_command(tape=tape, line=prompt.strip())
Expand All @@ -126,7 +103,21 @@ async def run_stream(
allowed_tools=allowed_tools,
)

return self._events_with_callback(events, callback=stack.aclose)
return events

def session_tape(
self,
session_id: str,
state: TurnState,
) -> Tape:
"""Return the ordinary tape addressed by a logical session id."""

state.setdefault("session_id", session_id)
return self.tape.session_tape(
session_id,
workspace_from_state(state),
context=replace(self.tape.context, state=state),
)

async def _run_command(self, tape: Tape, *, line: str) -> str:
line = line[1:].strip()
Expand Down
152 changes: 152 additions & 0 deletions src/bub/builtin/forkmerge.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,152 @@
"""Fork and merge tape writes through a builtin sidecar."""

from __future__ import annotations

import itertools
from collections.abc import Iterable
from dataclasses import dataclass, field, replace
from typing import Any, cast

from loguru import logger

from bub.sidecars import sidecar_owns_tape, sidecar_tape_name
from bub.store import AsyncTapeStore, InMemoryTapeStore, TapeQuery
from bub.tape import Tape, TapeEntry

FORK_MERGE_SIDECAR_NAME = "forkmerge"


class _ForkTapeStore:
"""In-memory write overlay owned by the forkmerge sidecar."""

def __init__(self, parent: AsyncTapeStore, tape: str, *, sidecars: Iterable[str] = ()) -> None:
self._parent = parent
self._store = InMemoryTapeStore()
self._tape = tape
self._sidecars = tuple(dict.fromkeys(sidecars))
self._managed_tapes = {tape, *self._sidecars}
self._reset_tapes: set[str] = set()

async def list_tapes(self) -> list[str]:
return await self._parent.list_tapes()

async def reset(self, tape: str) -> None:
if tape not in self._managed_tapes:
await self._parent.reset(tape)
return
self._store.reset(tape)
self._reset_tapes.add(tape)

async def fetch_all(self, query: TapeQuery[AsyncTapeStore]) -> Iterable[TapeEntry]:
if query.tape not in self._managed_tapes:
return await self._parent.fetch_all(query)

parent_entries: Iterable[TapeEntry] = []
if query.tape not in self._reset_tapes:
try:
parent_entries = await self._parent.fetch_all(query)
except Exception:
parent_entries = []
fork_entries: list[TapeEntry] = []
for entry in self._store.read(query.tape) or []:
if entry.kind == "anchor": # noqa: SIM102
if query._after_last or (query._after_anchor and entry.payload.get("name") == query._after_anchor):
fork_entries.clear()
parent_entries = []
continue
if query._kinds and entry.kind not in query._kinds:
continue
fork_entries.append(entry)
entries = itertools.chain(parent_entries, fork_entries)
return itertools.islice(entries, query._limit) if query._limit is not None else entries

@staticmethod
def _redact_prompt(prompt: list[dict]) -> Any:
if not isinstance(prompt, list):
return prompt
return [part for part in prompt if part.get("type") == "text"]

@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)
if tape not in self._managed_tapes:
await self._parent.append(tape, entry)
return
self._store.append(tape, entry)

async def merge_back(self) -> None:
total = 0
for sidecar in self._sidecars:
entries = self._store.read(sidecar) or []
try:
if sidecar in self._reset_tapes:
await self._parent.reset(sidecar)
for entry in entries:
await self._parent.append(sidecar, entry)
except Exception as exc:
logger.warning('Failed to merge sidecar "{}" into tape "{}": {}', sidecar, self._tape, exc)
self._store.append(
self._tape,
TapeEntry.event(
"sidecar.merge",
{"tape": sidecar, "status": "error", "error": str(exc)},
context=False,
),
)
else:
total += len(entries)

if self._tape in self._reset_tapes:
await self._parent.reset(self._tape)
entries = self._store.read(self._tape) or []
for entry in entries:
await self._parent.append(self._tape, entry)
total += len(entries)
if total:
logger.info('Merged {} entries into tape fork "{}"', total, self._tape)


@dataclass
class TapeFork:
"""An isolated tape whose writes can be merged or discarded."""

tape: Tape
_store: _ForkTapeStore = field(repr=False)
_closed: bool = field(default=False, init=False, repr=False)

async def merge(self) -> None:
if self._closed:
return
await self._store.merge_back()
self._closed = True

async def discard(self) -> None:
self._closed = True


@dataclass(frozen=True)
class ForkMergeSidecar:
"""Isolate tape writes and merge them only when requested."""

name: str = field(default=FORK_MERGE_SIDECAR_NAME, init=False)
owns_tape: bool = field(default=False, init=False)

@classmethod
def mounted(cls, tape: Tape) -> ForkMergeSidecar:
sidecar = tape.get_sidecar(FORK_MERGE_SIDECAR_NAME)
if sidecar is None or not callable(getattr(sidecar, "fork", None)):
raise TypeError("forkmerge sidecar is not mounted")
return cast("ForkMergeSidecar", sidecar)

async def fork(self, tape: Tape) -> TapeFork:
managed_sidecars = tuple(
sidecar_tape_name(tape.name, sidecar.name) for sidecar in tape.sidecars if sidecar_owns_tape(sidecar)
)
store = _ForkTapeStore(tape.store, tape.name, sidecars=managed_sidecars)
return TapeFork(tape=replace(tape, store=store), _store=store)
28 changes: 20 additions & 8 deletions src/bub/builtin/hook_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -227,12 +227,18 @@ async def build_prompt(self, message: ChannelMessage, session_id: str, state: Tu

@hookimpl
async def run_model_stream(self, prompt: str | list[dict], session_id: str, state: TurnState) -> AsyncStreamEvents:
return await self._get_agent().run_stream(
session_id=session_id,
prompt=prompt,
state=state,
model=state.get("model"),
)
from bub.builtin.forkmerge import ForkMergeSidecar

agent = self._get_agent()
tape = agent.session_tape(session_id, state)
tape_fork = await ForkMergeSidecar.mounted(tape).fork(tape)
try:
events = await agent.run_stream(tape=tape_fork.tape, prompt=prompt, model=state.get("model"))
except Exception:
await tape_fork.discard()
raise
finish = tape_fork.discard if session_id.startswith("temp/") else tape_fork.merge
return events.attach(finish)

@hookimpl
def continue_prompt(self, prompt: str | list[dict], tape: Tape, state: StreamState) -> str:
Expand Down Expand Up @@ -386,13 +392,19 @@ def provide_tape_store(self) -> TapeStore:

return FileTapeStore(directory=bub.home / "tapes")

@hookimpl
def provide_tape_sidecar(self) -> TapeSidecar:
@hookimpl(specname="provide_tape_sidecar")
def provide_spill_sidecar(self) -> TapeSidecar:
from bub.builtin.spill import SpillSettings, SpillStore
from bub.configure import ensure_config

return SpillStore(ensure_config(SpillSettings))

@hookimpl(specname="provide_tape_sidecar")
def provide_forkmerge_sidecar(self) -> TapeSidecar:
from bub.builtin.forkmerge import ForkMergeSidecar

return ForkMergeSidecar()

@hookimpl
def build_tape_context(self) -> TapeContext:
return default_tape_context()
Expand Down
1 change: 1 addition & 0 deletions src/bub/builtin/spill.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,7 @@ class SpillStore:

settings: SpillSettings
name: str = field(default=SPILL_SIDECAR_NAME, init=False)
owns_tape: bool = field(default=True, init=False)

async def _record_write(self, tape: Tape, data: dict[str, object], *, run_id: str) -> None:
try:
Expand Down
26 changes: 18 additions & 8 deletions src/bub/builtin/tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -328,6 +328,8 @@ async def web_fetch(url: str, headers: dict | None = None, timeout: int | None =
@tool(name="subagent", context=True, model=SubAgentInput)
async def run_subagent(param: SubAgentInput, *, context: ToolContext) -> str:
"""Run a task with sub-agent using specific model and session."""
from bub.builtin.forkmerge import ForkMergeSidecar

agent = _get_agent(context)
session_id = context.state.get("session_id", "temp/unknown")
if param.session == "inherit":
Expand All @@ -338,15 +340,23 @@ async def run_subagent(param: SubAgentInput, *, context: ToolContext) -> str:
subagent_session = param.session
state = {**context.state, "session_id": subagent_session}
allowed_tools = resolve_tool_names(param.allowed_tools or None, exclude={"subagent"})
tape = agent.session_tape(subagent_session, state)
tape_fork = await ForkMergeSidecar.mounted(tape).fork(tape)
try:
events = await agent.run_stream(
tape=tape_fork.tape,
prompt=param.prompt,
model=param.model,
allowed_tools=allowed_tools,
allowed_skills=param.allowed_skills,
)
except Exception:
await tape_fork.discard()
raise

finish = tape_fork.discard if param.session == "temp" else tape_fork.merge
output = ""
async for event in await agent.run_stream(
session_id=subagent_session,
prompt=param.prompt,
state=state,
model=param.model,
allowed_tools=allowed_tools,
allowed_skills=param.allowed_skills,
):
async for event in events.attach(finish):
if event.kind == "error":
output += f"[Error: {event.data.get('message', 'unknown error')}]"
elif event.kind == "text":
Expand Down
12 changes: 11 additions & 1 deletion src/bub/sidecars.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,12 +8,22 @@


class TapeSidecar(Protocol):
"""A named capability backed by a sibling tape."""
"""A named capability mounted beside a session tape."""

@property
def name(self) -> str: ...


def sidecar_owns_tape(sidecar: TapeSidecar) -> bool:
"""Return whether a sidecar owns a persistent sibling tape.

Existing storage sidecars default to owning one. Capability-only sidecars,
such as the builtin forkmerge plugin, opt out explicitly.
"""

return getattr(sidecar, "owns_tape", True)


def sidecar_tape_name(owner: str, sidecar: str) -> str:
"""Return the physical tape name for a mounted sidecar."""

Expand Down
Loading
Loading