diff --git a/CHANGELOG.md b/CHANGELOG.md index 72a2b9c..0772299 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -1,5 +1,44 @@ # Changelog +### Unreleased + +Structural refactor, taking its shape from the pi agent harness. Every +previous import path still works, and resolves to the same object rather than +a copy. + +- **Layers.** The package was one flat namespace where the loop constructed a + `SessionCache` and called `format_tool_output`. It is now `llm` -> `core` -> + `data` -> `app`, each importing only the layers below it, enforced + statically by `tests/test_layers.py`. The loop takes a `RunEnvironment` + instead of a cache, so `data_harness.core` runs an agent with no pandas + anywhere. +- **One loop.** Four near-copies (sync/async x result/stream) collapsed into + one generator driven by two drivers. They had drifted: the streaming copy + discarded token usage on provider errors, and `AsyncAgent` was missing + `from_dataframe`, subagents, MCP, the replay cache, and the approval gate. + The sync driver runs inline, so Ctrl-C lands promptly and a handler holding + a `sqlite3` connection still works. +- **Session tree.** A session is an append-only tree of typed entries and the + conversation is *derived* from it, never stored. Brings resume across + processes, forking a conversation without losing the original, reversible + compaction, and linear rather than quadratic writes. +- **Hooks.** `BeforeTurn`, `BeforeToolCall`, `AfterToolCall`, `AfterTurn`, + returning `Reminder`, `Block`, `Replace`, or `Stop`. The interpreter + approval gate is now built from this mechanism rather than hardcoded inside + the dispatcher. `AfterTurn` is where a spend cap belongs. +- **Compaction.** `max_turns` was a wall, not a strategy. Cuts land on turn + boundaries so a tool call is never separated from its result, and nothing is + deleted: the summary is an entry, so stepping the leaf back undoes it. +- **Typed errors.** Every failure carries a stable `code` under a common + `DataHarnessError`, so a caller can tell a rate-limited provider from a bug + in the model's code from a sandbox timeout. Existing base classes are kept, + so code catching `RuntimeError`/`ValueError` still works. +- `RunResult` gains `stopped_by`; `Agent` gains `last_result` and `hooks`; + `resolve_async_adapter` mirrors `resolve_adapter`. + +913 tests, up from 547. + + ### 0.13.0 - **MCP bridge** (`[mcp]` extra): data-harness is now an MCP client. `Agent.add_mcp_server(name, command, args=...)` connects any stdio [MCP](https://modelcontextprotocol.io) server and exposes its tools as a connector — hidden until `load_connectors`, with large results routed through the `SessionCache`. Swap the command for a Postgres/SQLite/filesystem server; no per-source code. - `MCPClient` / `MCPServer` / `mcp_tool_specs` exported for wiring MCP into a `Harness` directly; `Agent.close()` shuts servers down. `examples/mcp_demo.py` shows it live against `mcp-server-time`. diff --git a/data_harness/__init__.py b/data_harness/__init__.py index 8f69551..69edf28 100644 --- a/data_harness/__init__.py +++ b/data_harness/__init__.py @@ -1,23 +1,79 @@ -from data_harness.agent import Agent, AgentSession, AsyncAgent, AsyncAgentSession -from data_harness.artifacts import ChartArtifact -from data_harness.exceptions import ( +"""data-harness: a constrained agent harness for data analysis. + +The package is layered, bottom up: + +- `data_harness.llm` — provider adapters and the wire types they speak +- `data_harness.core` — the ReAct loop, `RunResult`, run logging +- `data_harness.data` — the session cache, interpreter, SQL, connectors +- `data_harness.app` — `Agent`, `ask`, the CLI + +Each layer may import the ones below it and no others; `tests/test_layers.py` +enforces that. The names below are the stable public surface and are unaffected +by where a class physically lives. +""" + +from data_harness._legacy_paths import install as _install_legacy_paths + +_install_legacy_paths() + +from data_harness.app.agent import ( # noqa: E402 + Agent, + AgentSession, + AsyncAgent, + AsyncAgentSession, +) +from data_harness.app.quickstart import ( # noqa: E402 + Chat, + SmartFrame, + ask, + resolve_adapter, + resolve_async_adapter, +) +from data_harness.core.artifacts import ChartArtifact # noqa: E402 +from data_harness.core.compaction import ( # noqa: E402 + CompactionSettings, + estimate_tokens, + make_compactor, +) +from data_harness.core.exceptions import ( # noqa: E402 + ConfigurationError, + DataHarnessError, + ExecutionError, MaxTurnsExceeded, + ProviderError, SubagentRecursionError, ToolNotFoundError, ) -from data_harness.exec_cache import ExecutionCache -from data_harness.io import load_dataframe -from data_harness.loop import AsyncHarness -from data_harness.mcp import MCPClient, MCPServer, mcp_tool_specs -from data_harness.providers.base import ( +from data_harness.core.hooks import ( # noqa: E402 + AfterToolCall, + AfterTurn, + BeforeToolCall, + BeforeTurn, + Block, + HookError, + HookRegistry, + Reminder, + Replace, + Stop, +) +from data_harness.core.result import CacheStorageInfo, RunResult, Usage # noqa: E402 +from data_harness.core.session import ( # noqa: E402 + JsonlSessionStore, + MemorySessionStore, + Session, + SessionStoreError, +) +from data_harness.data.exec_cache import ExecutionCache # noqa: E402 +from data_harness.data.harness import AsyncHarness # noqa: E402 +from data_harness.data.io import load_dataframe # noqa: E402 +from data_harness.data.mcp import MCPClient, MCPServer, mcp_tool_specs # noqa: E402 +from data_harness.llm.providers.base import ( # noqa: E402 AsyncProviderAdapter, NormalizedResponse, ProviderAdapter, StopReason, ) -from data_harness.quickstart import Chat, SmartFrame, ask, resolve_adapter -from data_harness.result import CacheStorageInfo, RunResult, Usage -from data_harness.streaming import ( +from data_harness.llm.streaming import ( # noqa: E402 ContentBlockDeltaEvent, ContentBlockStartEvent, ContentBlockStopEvent, @@ -30,7 +86,7 @@ TextDelta, ToolResultEvent, ) -from data_harness.types import ( +from data_harness.llm.types import ( # noqa: E402 ContentBlock, Message, TextBlock, @@ -41,6 +97,25 @@ ) __all__ = [ + "Stop", + "SessionStoreError", + "Session", + "Replace", + "Reminder", + "ProviderError", + "MemorySessionStore", + "JsonlSessionStore", + "HookRegistry", + "HookError", + "ExecutionError", + "DataHarnessError", + "ConfigurationError", + "CompactionSettings", + "Block", + "BeforeTurn", + "BeforeToolCall", + "AfterTurn", + "AfterToolCall", "Agent", "AgentSession", "AsyncAgent", @@ -81,7 +156,10 @@ "ToolUseBlock", "Usage", "ask", + "estimate_tokens", + "make_compactor", "load_dataframe", "mcp_tool_specs", "resolve_adapter", + "resolve_async_adapter", ] diff --git a/data_harness/_legacy_paths.py b/data_harness/_legacy_paths.py new file mode 100644 index 0000000..115efcb --- /dev/null +++ b/data_harness/_legacy_paths.py @@ -0,0 +1,112 @@ +"""Keep the pre-layering module paths importable. + +Modules moved into `llm`, `core`, `data`, and `app` when the package grew +layers. `data_harness.loop` and friends were documented and in use, so they +stay resolvable. + +These are aliases, not copies: `sys.modules["data_harness.loop"]` *is* +`data_harness.core.loop`, so ``isinstance`` checks and identity comparisons +across the two paths agree. Re-exporting names into a shim module would give +the same names but a second module object, and a class imported through one +path would not be the class imported through the other. + +New code should use the layered path. The alias exists so a rename does not +become an upgrade barrier. +""" + +from __future__ import annotations + +import importlib +import importlib.abc +import importlib.machinery +import importlib.util +import sys + +#: legacy dotted path -> the module that now implements it +LEGACY_PATHS: dict[str, str] = { + "data_harness.types": "data_harness.llm.types", + "data_harness.streaming": "data_harness.llm.streaming", + "data_harness.testing": "data_harness.llm.testing", + "data_harness.providers": "data_harness.llm.providers", + "data_harness.providers.base": "data_harness.llm.providers.base", + "data_harness.providers.openai": "data_harness.llm.providers.openai", + "data_harness.providers.anthropic": "data_harness.llm.providers.anthropic", + "data_harness.loop": "data_harness.data.harness", + "data_harness.result": "data_harness.core.result", + "data_harness.artifacts": "data_harness.core.artifacts", + "data_harness.exceptions": "data_harness.core.exceptions", + "data_harness.observe": "data_harness.core.observe", + "data_harness.logger": "data_harness.core.logger", + "data_harness.schema": "data_harness.core.schema", + "data_harness.serialize": "data_harness.core.serialize", + "data_harness.cache": "data_harness.data.cache", + "data_harness.format": "data_harness.data.format", + "data_harness.io": "data_harness.data.io", + "data_harness.exec_cache": "data_harness.data.exec_cache", + "data_harness.mcp": "data_harness.data.mcp", + "data_harness.tools": "data_harness.data.tools", + "data_harness.tools.interpreter": "data_harness.data.tools.interpreter", + "data_harness.tools.sandbox": "data_harness.data.tools.sandbox", + "data_harness.tools.sql": "data_harness.data.tools.sql", + "data_harness.tools.variables": "data_harness.data.tools.variables", + "data_harness.tools.connectors": "data_harness.data.tools.connectors", + "data_harness.tools.planner": "data_harness.data.tools.planner", + "data_harness.tools.subagent": "data_harness.data.tools.subagent", + "data_harness.agent": "data_harness.app.agent", + "data_harness.quickstart": "data_harness.app.quickstart", + "data_harness.cli": "data_harness.app.cli", + "data_harness.notebook": "data_harness.app.notebook", + "data_harness.pandas": "data_harness.app.pandas", +} + + +class _AliasLoader(importlib.abc.Loader): + """Bind a legacy name to an already-imported module. + + `create_module` hands back the real module and `exec_module` does nothing, + because the source has already run under its own name. Returning the + target's own spec instead would make the import machinery build a second + module object and execute the file again: the two names would then hold + separate classes, and `isinstance` across them would quietly be false. + """ + + def __init__(self, target: str) -> None: + self._target = target + + def create_module(self, spec: importlib.machinery.ModuleSpec): + return importlib.import_module(self._target) + + def exec_module(self, module: object) -> None: + return None + + +class _LegacyPathFinder(importlib.abc.MetaPathFinder): + """Resolve a legacy path to the module that replaced it, on first import. + + A meta-path finder rather than eager aliasing: importing every module at + package import time would drag the whole library, and every optional + dependency sitting behind a lazy import, into memory just to register + names nobody may use. + """ + + def find_spec( + self, fullname: str, path: object = None, target: object = None + ) -> importlib.machinery.ModuleSpec | None: + real = LEGACY_PATHS.get(fullname) + if real is None: + return None + return importlib.util.spec_from_loader(fullname, _AliasLoader(real)) + + +def install() -> None: + """Make the legacy paths importable. Idempotent. + + Inserted at the front of `sys.meta_path`, not appended. An aliased package + such as `data_harness.providers` resolves to a real package whose + ``__path__`` still contains `base.py`, so the default `PathFinder` can + resolve `data_harness.providers.base` by itself and would load a second + copy of it. Going first means the alias wins for the names it owns; every + other name falls through on a dict miss. + """ + if not any(isinstance(f, _LegacyPathFinder) for f in sys.meta_path): + sys.meta_path.insert(0, _LegacyPathFinder()) diff --git a/data_harness/app/__init__.py b/data_harness/app/__init__.py new file mode 100644 index 0000000..3ecdca6 --- /dev/null +++ b/data_harness/app/__init__.py @@ -0,0 +1,8 @@ +"""Facades: the API most users touch. + +`Agent`/`AsyncAgent`, the zero-config `ask`/`Chat`/`SmartFrame` entry points, +the CLI, and notebook helpers. Composition only; no behaviour that belongs to +a lower layer. + +May import everything below it. +""" diff --git a/data_harness/agent.py b/data_harness/app/agent.py similarity index 51% rename from data_harness/agent.py rename to data_harness/app/agent.py index d575487..dc3d21f 100644 --- a/data_harness/agent.py +++ b/data_harness/app/agent.py @@ -1,32 +1,42 @@ """High-level `Agent` and `AsyncAgent` convenience layers. -`Agent` wraps `Harness` for sync workflows. -`AsyncAgent` wraps `AsyncHarness` for async and streaming workflows. +`Agent` wraps `Harness` for sync workflows. `AsyncAgent` wraps `AsyncHarness` +for async and streaming workflows. Both are one-shot per `run()` call. Use `session()` / `async_session()` for multi-turn conversations over a shared message history and cache. + +Configuration, tool wiring, connectors, MCP, subagents, and the replay cache +all live once in `_AgentBase`. The two public classes differ only in which +adapter kind they take and whether their run methods are coroutines. They used +to be independent copies, which is how `AsyncAgent` silently ended up missing +`from_dataframe`, subagents, MCP, the replay cache, and the approval gate. """ from __future__ import annotations +import asyncio +import functools import uuid from collections.abc import AsyncGenerator, Callable +from contextlib import aclosing from dataclasses import dataclass from pathlib import Path -from typing import Any - -from data_harness.cache import SessionCache -from data_harness.loop import AsyncHarness, Harness -from data_harness.providers.base import AsyncProviderAdapter, ProviderAdapter -from data_harness.result import CacheStorageInfo, RunResult, Usage -from data_harness.schema import infer_input_schema -from data_harness.streaming import StreamEvent -from data_harness.tools.connectors import ConnectorRegistry -from data_harness.tools.interpreter import PythonInterpreter -from data_harness.tools.planner import Planner -from data_harness.tools.subagent import _copy_cache_value, make_subagent_spec -from data_harness.tools.variables import make_list_variables_spec -from data_harness.types import ToolAnnotations, ToolSpec +from typing import Any, TypeVar + +from data_harness.core.hooks import Decision, Event, HookRegistry +from data_harness.core.result import CacheStorageInfo, RunResult, Usage, unwrap_text +from data_harness.core.schema import infer_input_schema +from data_harness.data.cache import SessionCache +from data_harness.data.harness import AsyncHarness, Harness, run_coroutine_blocking +from data_harness.data.tools.connectors import ConnectorRegistry +from data_harness.data.tools.interpreter import PythonInterpreter +from data_harness.data.tools.planner import Planner +from data_harness.data.tools.subagent import _copy_cache_value, make_subagent_spec +from data_harness.data.tools.variables import make_list_variables_spec +from data_harness.llm.providers.base import AsyncProviderAdapter, ProviderAdapter +from data_harness.llm.streaming import StreamEvent +from data_harness.llm.types import ToolAnnotations, ToolSpec @dataclass(frozen=True) @@ -50,7 +60,7 @@ class ConnectorBuilder: Obtain an instance via `Agent.connector` rather than constructing directly. """ - def __init__(self, agent: Agent | AsyncAgent, name: str) -> None: + def __init__(self, agent: _AgentBase, name: str) -> None: self._agent = agent self._name = name @@ -88,100 +98,19 @@ def tool( return fn -def _build_tools_for( - agent: Agent | AsyncAgent, - *, - planner: Planner | None, - cache: SessionCache, -) -> list[ToolSpec]: - """Shared tool-building logic for Agent and AsyncAgent.""" - run_dir = str(agent._run_dir) if agent._run_dir is not None else "./runs" - artifacts_dir = str(Path(run_dir) / "charts") - if getattr(agent, "_execution", "inprocess") == "subprocess": - from data_harness.tools.sandbox import SubprocessPythonInterpreter - - interpreter_spec = SubprocessPythonInterpreter.make_tool_spec( - cache, artifacts_dir=artifacts_dir, **(agent._sandbox_options or {}) - ) - else: - interpreter_spec = PythonInterpreter.make_tool_spec( - cache, artifacts_dir=artifacts_dir - ) - tools = [ - interpreter_spec, - make_list_variables_spec(cache), - ] - if planner is not None: - tools.extend(planner.make_tool_specs()) - if getattr(agent, "_sql_enabled", False): - from data_harness.tools.sql import make_sql_query_spec - - tools.append(make_sql_query_spec(cache, engine_url=agent._sql_engine_url)) - mcp_clients = getattr(agent, "_mcp_clients", {}) - if agent._connectors or mcp_clients: - registry = ConnectorRegistry() - for connector_name, connector in agent._connectors.items(): - registry.register( - name=connector_name, - description=connector.description, - tools=[ - ToolSpec( - name=f"{connector_name}__{definition.fn.__name__}", - description=definition.description, - input_schema=definition.input_schema, - handler=definition.fn, - visible=False, - annotations=definition.annotations, - ) - for definition in agent._connector_tools - if definition.connector_name == connector_name - ], - ) - for server_name, mcp_client in mcp_clients.items(): - from data_harness.mcp import mcp_tool_specs - - registry.register( - name=server_name, - description=( - f"MCP server '{server_name}' ({len(mcp_client.tools)} tools)" - ), - tools=mcp_tool_specs(mcp_client, prefix=server_name), - ) - tools.append(registry.get_load_connectors_spec()) - tools.extend(registry.make_wrapped_specs(cache)) - return tools - - -class Agent: - """High-level synchronous agent. - - `Agent` composes a `Harness`, a `SessionCache`, and optional tools from a - single configuration. Each call to `run` builds a fresh `Harness` with a - fresh message history. Use `session` when you need multi-turn conversation - state to persist across questions. - - Example:: +_SelfT = TypeVar("_SelfT", bound="_AgentBase") - from data_harness import Agent - from data_harness.providers.anthropic import AnthropicAdapter - agent = Agent( - adapter=AnthropicAdapter(model="claude-sonnet-4-6"), - system="You are a data analyst.", - ) - print(agent.run("Compute the mean of [1, 2, 3].")) +class _AgentBase: + """Configuration, tool wiring, and feature toggles shared by both agents. - Args: - adapter: Synchronous provider adapter. - system: System prompt passed unchanged to every `Harness` run. - max_turns: Hard cap on provider turns per `run` call. - cache: Shared `SessionCache`. A fresh cache is created when ``None``. - run_dir: Directory for JSONL logs. Defaults to ``./runs``. + Holds no provider I/O: subclasses own the adapter and the run methods. + Every ``enable_*`` toggle, connector, MCP server, and the replay cache is + defined here exactly once, so the sync and async agents cannot drift. """ def __init__( self, - adapter: ProviderAdapter, system: str, *, max_turns: int = 25, @@ -191,49 +120,56 @@ def __init__( sandbox_options: dict[str, Any] | None = None, on_code: Callable[[str], Any] | None = None, code_only: bool = False, + hooks: HookRegistry | None = None, ) -> None: - self._adapter = adapter self._system = system self._max_turns = max_turns self._cache = cache if cache is not None else SessionCache() self._run_dir = run_dir - self._last_harness: Harness | None = None + self._execution = execution + self._sandbox_options = sandbox_options + self._on_code = on_code + self._code_only = code_only self._last_run_file: str | None = None self._connectors: dict[str, _ConnectorDefinition] = {} self._connector_tools: list[_ConnectorToolDefinition] = [] self._planner_enabled = False - self._subagent_factory: Callable[[], ProviderAdapter] | None = None + self._subagent_factory: Callable[[], Any] | None = None self._sql_enabled = False self._sql_engine_url: str | None = None - self._execution = execution - self._sandbox_options = sandbox_options - self._on_code = on_code - self._code_only = code_only self._exec_cache: Any = None self._mcp_clients: dict[str, Any] = {} + self._last_result: RunResult | None = None + self._hooks = hooks if hooks is not None else HookRegistry() + + # ── construction helpers ──────────────────────────────────────────────── + + @classmethod + def _default_adapter(cls, model: str | None) -> Any: + raise NotImplementedError @classmethod def from_dataframe( - cls, + cls: type[_SelfT], data: Any, *, - adapter: ProviderAdapter | None = None, + adapter: Any = None, model: str | None = None, system: str | None = None, semantics: dict[str, dict] | None = None, **kwargs: Any, - ) -> Agent: - """Build an `Agent` with ``data`` preloaded as cache handles. + ) -> _SelfT: + """Build an agent with ``data`` preloaded as cache handles. Accepts a DataFrame, a ``{name: value}`` mapping, a file path, or a list of paths. Resolves an adapter from ``model``/the environment and applies the default analyst system prompt unless overridden. """ - from data_harness.io import to_handles - from data_harness.quickstart import _DEFAULT_SYSTEM, resolve_adapter + from data_harness.app.quickstart import _DEFAULT_SYSTEM + from data_harness.data.io import to_handles agent = cls( - adapter=adapter if adapter is not None else resolve_adapter(model), + adapter=adapter if adapter is not None else cls._default_adapter(model), system=system if system is not None else _DEFAULT_SYSTEM, **kwargs, ) @@ -243,22 +179,52 @@ def from_dataframe( return agent @classmethod - def from_csv(cls, path: str | Path, **kwargs: Any) -> Agent: - """Build an `Agent` from a CSV (or other supported file) path.""" + def from_csv(cls: type[_SelfT], path: str | Path, **kwargs: Any) -> _SelfT: + """Build an agent from a CSV (or other supported file) path.""" return cls.from_dataframe(str(path), **kwargs) + # ── inspection ────────────────────────────────────────────────────────── + @property def cache(self) -> SessionCache: return self._cache - @property - def last_harness(self) -> Harness | None: - return self._last_harness - @property def last_run_file(self) -> str | None: return self._last_run_file + @property + def last_result(self) -> RunResult | None: + """`RunResult` for the most recent run, however it was executed. + + Set by streamed runs and replay-cache hits too, neither of which + returns a result to the caller directly. + """ + return self._last_result + + @property + def exec_cache(self) -> Any: + """The `ExecutionCache`, or ``None`` if caching is disabled.""" + return self._exec_cache + + @property + def hooks(self) -> HookRegistry: + """Hooks applied to every harness this agent builds. + + The agent constructs a fresh `Harness` per run, so hooks have to live + here rather than on the harness: one registered on a harness would be + gone by the next call. + """ + return self._hooks + + def on( + self, event_type: type[Event], hook: Callable[[Any], Decision | None] + ) -> None: + """Register ``hook`` for ``event_type``. See `data_harness.core.hooks`.""" + self._hooks.add(event_type, hook) + + # ── feature toggles ───────────────────────────────────────────────────── + def connector(self, name: str, *, description: str) -> ConnectorBuilder: """Register a named connector and return a builder for attaching tools. @@ -277,7 +243,7 @@ def connector(self, name: str, *, description: str) -> ConnectorBuilder: ) return ConnectorBuilder(self, name) - def enable_planner(self) -> Agent: + def enable_planner(self: _SelfT) -> _SelfT: """Enable the planning tool and suffix-based nag reminders. The planner escalates reminders at turns 4, 8, and 12 when no progress @@ -289,9 +255,7 @@ def enable_planner(self) -> Agent: self._planner_enabled = True return self - def enable_subagents( - self, *, adapter_factory: Callable[[], ProviderAdapter] - ) -> Agent: + def enable_subagents(self: _SelfT, *, adapter_factory: Callable[[], Any]) -> _SelfT: """Enable the subagent tool, using ``adapter_factory`` for spawned agents. Each spawned subagent gets a fresh adapter, fresh message history, and @@ -299,8 +263,11 @@ def enable_subagents( ``input_handles``. Args: - adapter_factory: Zero-argument callable that returns a fresh - `ProviderAdapter` for each subagent. + adapter_factory: Zero-argument callable returning a fresh adapter + for each subagent. Either a `ProviderAdapter` or an + `AsyncProviderAdapter`; whichever it returns picks the matching + driver, so a sync adapter keeps the parent's threading + semantics. Returns: ``self``, for method chaining. @@ -308,15 +275,55 @@ def enable_subagents( self._subagent_factory = adapter_factory return self + def enable_sql(self: _SelfT, *, engine_url: str | None = None) -> _SelfT: + """Enable the ``sql_query`` tool. + + With no ``engine_url``, queries run via DuckDB in-process over the + DataFrame handles in the cache. With a SQLAlchemy URL, queries run + against that database instead. + + Args: + engine_url: Optional SQLAlchemy connection URL. + + Returns: + ``self``, for method chaining. + """ + self._sql_enabled = True + self._sql_engine_url = engine_url + return self + + def enable_cache(self: _SelfT, path: Any = None) -> _SelfT: + """Enable the code-replay cache. + + On a repeat ``run``/``run_result`` with the same question and data + schema, the previously recorded interpreter/SQL code is replayed against + the current cache without calling the model (zero tokens, no turns). + + Args: + path: A JSON file path to persist the cache across processes, an + existing `ExecutionCache` to share in-process, or ``None`` for an + in-memory cache. + + Returns: + ``self``, for method chaining. + """ + from data_harness.data.exec_cache import ExecutionCache + + if isinstance(path, ExecutionCache): + self._exec_cache = path + else: + self._exec_cache = ExecutionCache(path) + return self + def add_mcp_server( - self, + self: _SelfT, name: str, command: str | None = None, *, args: list[str] | None = None, env: dict[str, str] | None = None, client: Any = None, - ) -> Agent: + ) -> _SelfT: """Connect an MCP server and expose its tools (via progressive disclosure). The server's tools become a connector named ``name`` — hidden until the @@ -335,7 +342,7 @@ def add_mcp_server( ``self``, for chaining. """ if client is None: - from data_harness.mcp import MCPClient, MCPServer + from data_harness.data.mcp import MCPClient, MCPServer client = MCPClient( MCPServer(command=command or "", args=args or [], env=env) @@ -352,307 +359,297 @@ def close(self) -> None: pass self._mcp_clients.clear() - def enable_sql(self, *, engine_url: str | None = None) -> Agent: - """Enable the ``sql_query`` tool. + def explain(self) -> str: + """Return a string showing the equivalent explicit `Harness` wiring.""" + return _EXPLAIN_TEMPLATE.format( + system=_truncate(self._system), + max_turns=self._max_turns, + run_dir=self._run_dir if self._run_dir is not None else "./runs", + ) - With no ``engine_url``, queries run via DuckDB in-process over the - DataFrame handles in the cache. With a SQLAlchemy URL, queries run - against that database instead. + # ── tool wiring ───────────────────────────────────────────────────────── - Args: - engine_url: Optional SQLAlchemy connection URL. + @property + def _effective_run_dir(self) -> str: + return str(self._run_dir) if self._run_dir is not None else "./runs" - Returns: - ``self``, for method chaining. - """ - self._sql_enabled = True - self._sql_engine_url = engine_url - return self + def _build_tools( + self, + *, + planner: Planner | None = None, + cache: SessionCache | None = None, + ) -> list[ToolSpec]: + """Assemble the full tool list for one harness.""" + target_cache = cache if cache is not None else self._cache + artifacts_dir = str(Path(self._effective_run_dir) / "charts") - def enable_cache(self, path: Any = None) -> Agent: - """Enable the code-replay cache. + if self._execution == "subprocess": + from data_harness.data.tools.sandbox import SubprocessPythonInterpreter - On a repeat ``run``/``run_result`` with the same question and data - schema, the previously recorded interpreter/SQL code is replayed against - the current cache without calling the model (zero tokens, no turns). + interpreter_spec = SubprocessPythonInterpreter.make_tool_spec( + target_cache, + artifacts_dir=artifacts_dir, + **(self._sandbox_options or {}), + ) + else: + interpreter_spec = PythonInterpreter.make_tool_spec( + target_cache, artifacts_dir=artifacts_dir + ) - Args: - path: A JSON file path to persist the cache across processes, an - existing `ExecutionCache` to share in-process, or ``None`` for an - in-memory cache. + tools = [interpreter_spec, make_list_variables_spec(target_cache)] - Returns: - ``self``, for method chaining. - """ - from data_harness.exec_cache import ExecutionCache + if planner is not None: + tools.extend(planner.make_tool_specs()) - if isinstance(path, ExecutionCache): - self._exec_cache = path - else: - self._exec_cache = ExecutionCache(path) - return self + if self._sql_enabled: + from data_harness.data.tools.sql import make_sql_query_spec - @property - def exec_cache(self) -> Any: - """The `ExecutionCache`, or ``None`` if caching is disabled.""" - return self._exec_cache + tools.append( + make_sql_query_spec(target_cache, engine_url=self._sql_engine_url) + ) + + if self._connectors or self._mcp_clients: + tools.extend(self._build_connector_tools(target_cache)) + + return tools + + def _build_connector_tools(self, cache: SessionCache) -> list[ToolSpec]: + from data_harness.data.mcp import mcp_tool_specs + + registry = ConnectorRegistry() + for connector_name, connector in self._connectors.items(): + registry.register( + name=connector_name, + description=connector.description, + tools=[ + ToolSpec( + name=f"{connector_name}__{definition.fn.__name__}", + description=definition.description, + input_schema=definition.input_schema, + handler=definition.fn, + visible=False, + annotations=definition.annotations, + ) + for definition in self._connector_tools + if definition.connector_name == connector_name + ], + ) + for server_name, mcp_client in self._mcp_clients.items(): + tool_count = len(mcp_client.tools) + registry.register( + name=server_name, + description=f"MCP server '{server_name}' ({tool_count} tools)", + tools=mcp_tool_specs(mcp_client, prefix=server_name), + ) + return [ + registry.get_load_connectors_spec(), + *registry.make_wrapped_specs(cache), + ] + + def _resolve_planner(self, planner: Planner | None) -> Planner | None: + if planner is not None: + return planner + return Planner() if self._planner_enabled else None + + def _tools_for_harness( + self, *, cache: SessionCache, planner: Planner | None + ) -> list[ToolSpec]: + """Tool list plus the subagent tool, when subagents are enabled.""" + tools = self._build_tools(planner=planner, cache=cache) + if self._subagent_factory is None: + return tools + tools.append( + make_subagent_spec( + adapter_factory=self._subagent_factory, + parent_tools=self._build_tools(planner=None, cache=cache), + parent_cache=cache, + run_dir=self._effective_run_dir, + make_sub_tools=lambda sub_cache: self._build_tools( + planner=None, cache=sub_cache + ), + ) + ) + return tools + + def _harness_kwargs( + self, *, cache: SessionCache | None, planner: Planner | None + ) -> tuple[dict[str, Any], Planner | None]: + effective_cache = cache if cache is not None else self._cache + effective_planner = self._resolve_planner(planner) + kwargs: dict[str, Any] = { + "system": self._system, + "tools": self._tools_for_harness( + cache=effective_cache, planner=effective_planner + ), + "max_turns": self._max_turns, + "cache": effective_cache, + "on_code": self._on_code, + "code_only": self._code_only, + "hooks": self._hooks, + } + if self._run_dir is not None: + kwargs["run_dir"] = str(self._run_dir) + return kwargs, effective_planner + + # ── replay cache ──────────────────────────────────────────────────────── - def _replay(self, key: str, cached, user_message: str) -> RunResult: - from data_harness.exec_cache import make_key # noqa: F401 (re-export anchor) + def _replay_key(self, user_message: str) -> str | None: + if self._exec_cache is None: + return None + from data_harness.data.exec_cache import make_key - tools = self._build_tools(cache=self._cache) - tool_map = {t.name: t for t in tools} + return make_key(user_message, self._cache, self._system) + + def _replay_steps(self, cached: Any) -> list[tuple[Callable[..., Any], dict]]: + """The dispatchable ``(handler, input)`` pairs recorded by a cached run.""" + tool_map = {t.name: t for t in self._build_tools(cache=self._cache)} + steps = [] for step in cached.steps: spec = tool_map.get(step["tool"]) if spec is not None and spec.handler is not None: - try: - spec.handler(**step["input"]) - except Exception: - # A recorded step may fail against fresh data; skip it - # rather than aborting the whole replay. - continue - storage = { - name: CacheStorageInfo( - location=meta["location"], storage_type=meta["storage_type"] - ) - for name, meta in self._cache.storage_metadata().items() - } + steps.append((spec.handler, step["input"])) + return steps + + def _replay_result(self, text: str) -> RunResult: + result = self._build_replay_result(text) + self._last_result = result + return result + + def _build_replay_result(self, text: str) -> RunResult: return RunResult( - text=cached.text, + text=text, status="success", turns=0, run_file=None, stop_reason=None, usage=Usage(), cache_snapshots=self._cache.list_handles(), - cache_storage=storage, + cache_storage={ + name: CacheStorageInfo( + location=meta["location"], storage_type=meta["storage_type"] + ) + for name, meta in self._cache.storage_metadata().items() + }, value=self._cache.get_answer(), charts=self._cache.list_charts(), run_id=str(uuid.uuid4()), ) - def session(self) -> AgentSession: - """Create a stateful `AgentSession` for multi-turn conversations. + def _record_replay(self, key: str | None, harness: Any, result: RunResult) -> None: + if key is None or result.status != "success": + return + from data_harness.data.exec_cache import CachedRun, extract_steps - Returns: - A new `AgentSession` backed by a copy of this agent's cache. - """ - return AgentSession(self) + self._exec_cache.put( + key, CachedRun(steps=extract_steps(harness.messages), text=result.text) + ) - def run_result(self, user_message: str) -> RunResult: - """Run the agent and return the full `RunResult`. - Builds a fresh `Harness` with a fresh message history for each call. +class Agent(_AgentBase): + """High-level synchronous agent. - Args: - user_message: The user prompt to send. + `Agent` composes a `Harness`, a `SessionCache`, and optional tools from a + single configuration. Each call to `run` builds a fresh `Harness` with a + fresh message history. Use `session` when you need multi-turn conversation + state to persist across questions. - Returns: - A `RunResult` with the text response, token usage, and cache state. - """ - key = None - if self._exec_cache is not None: - from data_harness.exec_cache import make_key + Example:: - key = make_key(user_message, self._cache, self._system) - cached = self._exec_cache.get(key) - if cached is not None: - return self._replay(key, cached, user_message) + from data_harness import Agent + from data_harness.llm.providers.anthropic import AnthropicAdapter - harness = self._make_harness() - self._last_harness = harness - result = harness.run_result( - user_message, run_id=str(uuid.uuid4()), session_id=None + agent = Agent( + adapter=AnthropicAdapter(model="claude-sonnet-4-6"), + system="You are a data analyst.", ) - self._last_run_file = harness.run_file - - if key is not None and result.status == "success": - from data_harness.exec_cache import CachedRun, extract_steps - - steps = extract_steps(harness._messages) - self._exec_cache.put(key, CachedRun(steps=steps, text=result.text)) - return result - - def run(self, user_message: str) -> str: - """Run the agent and return the final text response. - - Args: - user_message: The user prompt to send. - - Returns: - The model's final text response. - - Raises: - MaxTurnsExceeded: If the loop reaches ``max_turns``. - RuntimeError: If the provider raises an exception. - """ - result = self.run_result(user_message) - if result.status == "max_turns_exceeded": - from data_harness.exceptions import MaxTurnsExceeded - - raise MaxTurnsExceeded(result.turns) - if result.status == "error": - raise RuntimeError(result.error or "unknown error") - return result.text - - def explain(self) -> str: - """Return a string showing the equivalent explicit `Harness` wiring.""" - return _EXPLAIN_TEMPLATE.format( - system=_truncate(self._system), - max_turns=self._max_turns, - run_dir=self._run_dir if self._run_dir is not None else "./runs", - ) - - def _build_tools( - self, - *, - planner: Planner | None = None, - cache: SessionCache | None = None, - ) -> list[ToolSpec]: - target_cache = cache if cache is not None else self._cache - return _build_tools_for(self, planner=planner, cache=target_cache) - - def _make_harness( - self, - *, - cache: SessionCache | None = None, - planner: Planner | None = None, - ) -> Harness: - effective_cache = cache if cache is not None else self._cache - effective_planner = ( - planner - if planner is not None - else Planner() - if self._planner_enabled - else None - ) - tools = self._build_tools(planner=effective_planner, cache=effective_cache) - if self._subagent_factory is not None: - subagent_parent_tools = self._build_tools( - planner=None, cache=effective_cache - ) - effective_run_dir = ( - str(self._run_dir) if self._run_dir is not None else "./runs" - ) - tools.append( - make_subagent_spec( - adapter_factory=self._subagent_factory, - parent_tools=subagent_parent_tools, - parent_cache=effective_cache, - run_dir=effective_run_dir, - make_sub_tools=lambda sub_cache: self._build_tools( - planner=None, cache=sub_cache - ), - ) - ) - harness_kwargs: dict = { - "adapter": self._adapter, - "system": self._system, - "tools": tools, - "max_turns": self._max_turns, - "cache": effective_cache, - "on_code": self._on_code, - "code_only": self._code_only, - } - if self._run_dir is not None: - harness_kwargs["run_dir"] = str(self._run_dir) - - harness = Harness(**harness_kwargs) - if effective_planner is not None: - harness.register_reminder(effective_planner.reminder_hook) - return harness - - -class AgentSession: - """Stateful chat session built from an `Agent` definition. + print(agent.run("Compute the mean of [1, 2, 3].")) - `Agent.run()` intentionally stays one-shot for examples and tests. Use - `Agent.session()` when an application needs follow-up questions over the - same message history and cache handles. + Args: + adapter: Synchronous provider adapter. + system: System prompt passed unchanged to every `Harness` run. + max_turns: Hard cap on provider turns per `run` call. + cache: Shared `SessionCache`. A fresh cache is created when ``None``. + run_dir: Directory for JSONL logs. Defaults to ``./runs``. + execution: ``"inprocess"`` or ``"subprocess"`` interpreter isolation. + sandbox_options: Extra options for the subprocess interpreter. + on_code: Approval gate called with interpreter code before it runs. + code_only: When ``True``, interpreter code is echoed, never executed. """ - def __init__(self, agent: Agent) -> None: - self._agent = agent - self._cache = SessionCache( - sample_size=agent.cache.sample_size, - storage_dir=None, - hot_limit=agent.cache.hot_limit, - ) - for name, value in agent.cache.items(): - self._cache.put( - name, - _copy_cache_value(value), - semantics=agent.cache.get_semantics(name), - ) - self._harness = agent._make_harness(cache=self._cache) - self._id: str = str(uuid.uuid4()) - self._last_result: RunResult | None = None - self._turns: int = 0 - - @property - def id(self) -> str: - return self._id - - @property - def last_result(self) -> RunResult | None: - return self._last_result - - @property - def turns(self) -> int: - return self._turns + def __init__( + self, + adapter: ProviderAdapter, + system: str, + **kwargs: Any, + ) -> None: + super().__init__(system, **kwargs) + self._adapter = adapter + self._last_harness: Harness | None = None - @property - def cache(self) -> SessionCache: - return self._cache + @classmethod + def _default_adapter(cls, model: str | None) -> ProviderAdapter: + from data_harness.app.quickstart import resolve_adapter - @property - def harness(self) -> Harness: - return self._harness + return resolve_adapter(model) @property - def run_file(self) -> str | None: - return self._harness.run_file + def last_harness(self) -> Harness | None: + return self._last_harness - def put(self, name: str, value: Any, *, overwrite: bool = False) -> str: - """Store a value in the session cache and return the handle used. + def _replay(self, cached: Any) -> RunResult: + """Re-execute a cached run's tool steps inline, without calling the model.""" + for handler, tool_input in self._replay_steps(cached): + try: + if asyncio.iscoroutinefunction(handler): + run_coroutine_blocking(handler(**tool_input)) + else: + handler(**tool_input) + except Exception: # noqa: BLE001 + # A recorded step may fail against fresh data; skip it rather + # than aborting the whole replay. + continue + return self._replay_result(cached.text) - Args: - name: Desired handle name. Must be a valid Python identifier. - value: Any Python object to store. - overwrite: Replace the existing handle if ``True``. + def session(self) -> AgentSession: + """Create a stateful `AgentSession` for multi-turn conversations. Returns: - The handle name under which the value was stored. + A new `AgentSession` backed by a copy of this agent's cache. """ - return self._cache.put(name, value, overwrite=overwrite) + return AgentSession(self) - def list_handles(self) -> dict[str, str]: - """Return a mapping of all cache handle names to their snapshot strings.""" - return self._cache.list_handles() + def run_result(self, user_message: str) -> RunResult: + """Run the agent and return the full `RunResult`. - def ask_result(self, user_message: str) -> RunResult: - """Send a follow-up message and return the full `RunResult`. + Builds a fresh `Harness` with a fresh message history for each call. Args: - user_message: The follow-up user prompt. + user_message: The user prompt to send. Returns: - A `RunResult` for this turn sequence. + A `RunResult` with the text response, token usage, and cache state. """ - result = self._harness.ask_result( - user_message, run_id=str(uuid.uuid4()), session_id=self._id + key = self._replay_key(user_message) + if key is not None: + cached = self._exec_cache.get(key) + if cached is not None: + return self._replay(cached) + + harness = self._make_harness() + self._last_harness = harness + result = harness.run_result( + user_message, run_id=str(uuid.uuid4()), session_id=None ) + self._last_run_file = harness.run_file self._last_result = result - self._turns += result.turns - self._agent._last_harness = self._harness - self._agent._last_run_file = self._harness.run_file + self._record_replay(key, harness, result) return result - def ask(self, user_message: str) -> str: - """Send a follow-up message and return the final text response. + def run(self, user_message: str) -> str: + """Run the agent and return the final text response. Args: - user_message: The follow-up user prompt. + user_message: The user prompt to send. Returns: The model's final text response. @@ -661,95 +658,97 @@ def ask(self, user_message: str) -> str: MaxTurnsExceeded: If the loop reaches ``max_turns``. RuntimeError: If the provider raises an exception. """ - result = self.ask_result(user_message) - if result.status == "max_turns_exceeded": - from data_harness.exceptions import MaxTurnsExceeded + return unwrap_text(self.run_result(user_message)) - raise MaxTurnsExceeded(result.turns) - if result.status == "error": - raise RuntimeError(result.error or "unknown error") - return result.text + def _make_harness( + self, + *, + cache: SessionCache | None = None, + planner: Planner | None = None, + ) -> Harness: + kwargs, effective_planner = self._harness_kwargs(cache=cache, planner=planner) + harness = Harness(adapter=self._adapter, **kwargs) + if effective_planner is not None: + harness.register_reminder(effective_planner.reminder_hook) + return harness -class AsyncAgent: +class AsyncAgent(_AgentBase): """Async agent for use with `AsyncProviderAdapter`. `run()` and `run_result()` are coroutines. `run_stream()` is an async - generator that yields text tokens as they arrive from the provider. + generator that yields `StreamEvent`s as they arrive from the provider. Use `async_session()` for multi-turn streaming conversations. + + Takes the same configuration and supports the same features as `Agent`: + connectors, MCP servers, SQL, the planner, subagents, the replay cache, + and the interpreter approval gate. """ def __init__( self, adapter: AsyncProviderAdapter, system: str, - *, - max_turns: int = 25, - cache: SessionCache | None = None, - run_dir: str | Path | None = None, + **kwargs: Any, ) -> None: + super().__init__(system, **kwargs) self._adapter = adapter - self._system = system - self._max_turns = max_turns - self._cache = cache if cache is not None else SessionCache() - self._run_dir = run_dir self._last_harness: AsyncHarness | None = None - self._last_run_file: str | None = None - self._connectors: dict[str, _ConnectorDefinition] = {} - self._connector_tools: list[_ConnectorToolDefinition] = [] - self._planner_enabled = False - self._sql_enabled = False - self._sql_engine_url: str | None = None - @property - def cache(self) -> SessionCache: - return self._cache + @classmethod + def _default_adapter(cls, model: str | None) -> AsyncProviderAdapter: + from data_harness.app.quickstart import resolve_async_adapter + + return resolve_async_adapter(model) @property def last_harness(self) -> AsyncHarness | None: return self._last_harness - @property - def last_run_file(self) -> str | None: - return self._last_run_file - - def connector(self, name: str, *, description: str) -> ConnectorBuilder: - self._connectors[name] = _ConnectorDefinition( - name=name, description=description - ) - return ConnectorBuilder(self, name) - - def enable_planner(self) -> AsyncAgent: - self._planner_enabled = True - return self - - def enable_sql(self, *, engine_url: str | None = None) -> AsyncAgent: - """Enable the ``sql_query`` tool (DuckDB in-process, or SQLAlchemy URL).""" - self._sql_enabled = True - self._sql_engine_url = engine_url - return self + async def _replay(self, cached: Any) -> RunResult: + """Re-execute a cached run's tool steps without stalling the event loop.""" + for handler, tool_input in self._replay_steps(cached): + try: + if asyncio.iscoroutinefunction(handler): + await handler(**tool_input) + else: + await asyncio.to_thread(functools.partial(handler, **tool_input)) + except Exception: # noqa: BLE001 + # A recorded step may fail against fresh data; skip it rather + # than aborting the whole replay. + continue + return self._replay_result(cached.text) def async_session(self) -> AsyncAgentSession: + """Create a stateful `AsyncAgentSession` for multi-turn conversations.""" return AsyncAgentSession(self) async def run_result(self, user_message: str) -> RunResult: + """Run the agent and return the full `RunResult`.""" + key = self._replay_key(user_message) + if key is not None: + cached = self._exec_cache.get(key) + if cached is not None: + return await self._replay(cached) + harness = self._make_harness() self._last_harness = harness result = await harness.run_result( user_message, run_id=str(uuid.uuid4()), session_id=None ) self._last_run_file = harness.run_file + self._last_result = result + self._record_replay(key, harness, result) return result async def run(self, user_message: str) -> str: - result = await self.run_result(user_message) - from data_harness.exceptions import MaxTurnsExceeded + """Run the agent and return the final text response. - if result.status == "max_turns_exceeded": - raise MaxTurnsExceeded(result.turns) - if result.status == "error": - raise RuntimeError(result.error or "unknown error") - return result.text + Raises: + MaxTurnsExceeded: If the loop reaches ``max_turns``. + RuntimeError: If the provider raises an exception. + """ + return unwrap_text(await self.run_result(user_message)) async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: """Stream events for a one-shot run. @@ -757,28 +756,46 @@ async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non Yields StreamEvent objects (message_start, content_block_*, message_delta, message_stop, tool_result) following the Claude Agent SDK protocol. + After the generator is exhausted, ``agent.last_result`` holds the + `RunResult` for the run, including token usage. Read it there rather + than through ``last_harness``: a replay-cache hit builds no harness, so + ``last_harness`` is either ``None`` or left over from an earlier run. + + A cache hit yields the cached answer as a single synthetic text block + and reports zero usage, so a streaming caller sees the same answer a + non-streaming one would. + Usage:: async for event in agent.run_stream("hello"): if event.type == "content_block_delta": - from data_harness.streaming import TextDelta + from data_harness.llm.streaming import TextDelta if isinstance(event.delta, TextDelta): print(event.delta.text, end="", flush=True) """ + key = self._replay_key(user_message) + if key is not None: + cached = self._exec_cache.get(key) + if cached is not None: + # `_replay` mints its own run id; nothing here to stamp. + result = await self._replay(cached) + for event in _synthetic_text_events(result.text): + yield event + return + + run_id = str(uuid.uuid4()) harness = self._make_harness() self._last_harness = harness - async for event in harness.run_stream(user_message): - yield event + # `aclosing`, not a bare `async for`: closing this generator does not + # close the one it iterates, so abandoning the stream here would leave + # the harness (and the provider's HTTP stream) suspended. + async with aclosing(harness.run_stream(user_message)) as events: + async for event in events: + yield event self._last_run_file = harness.run_file - - def _build_tools( - self, - *, - planner: Planner | None = None, - cache: SessionCache | None = None, - ) -> list[ToolSpec]: - target_cache = cache if cache is not None else self._cache - return _build_tools_for(self, planner=planner, cache=target_cache) + if harness.last_result is not None: + self._last_result = harness._stamp(harness.last_result, run_id, None) + self._record_replay(key, harness, harness.last_result) def _make_harness( self, @@ -786,36 +803,22 @@ def _make_harness( cache: SessionCache | None = None, planner: Planner | None = None, ) -> AsyncHarness: - effective_cache = cache if cache is not None else self._cache - effective_planner = ( - planner - if planner is not None - else Planner() - if self._planner_enabled - else None - ) - tools = self._build_tools(planner=effective_planner, cache=effective_cache) - - harness_kwargs: dict = { - "adapter": self._adapter, - "system": self._system, - "tools": tools, - "max_turns": self._max_turns, - "cache": effective_cache, - } - if self._run_dir is not None: - harness_kwargs["run_dir"] = str(self._run_dir) - - harness = AsyncHarness(**harness_kwargs) + kwargs, effective_planner = self._harness_kwargs(cache=cache, planner=planner) + harness = AsyncHarness(adapter=self._adapter, **kwargs) if effective_planner is not None: harness.register_reminder(effective_planner.reminder_hook) return harness -class AsyncAgentSession: - """Stateful async chat session built from an `AsyncAgent` definition.""" +class _SessionBase: + """Shared state for `AgentSession` and `AsyncAgentSession`. - def __init__(self, agent: AsyncAgent) -> None: + A session owns its own `SessionCache`, seeded with a deep copy of the + agent's handles so a session cannot mutate the agent definition it was + built from. + """ + + def __init__(self, agent: _AgentBase) -> None: self._agent = agent self._cache = SessionCache( sample_size=agent.cache.sample_size, @@ -828,7 +831,7 @@ def __init__(self, agent: AsyncAgent) -> None: _copy_cache_value(value), semantics=agent.cache.get_semantics(name), ) - self._harness = agent._make_harness(cache=self._cache) + self._harness = agent._make_harness(cache=self._cache) # type: ignore[attr-defined] self._id: str = str(uuid.uuid4()) self._last_result: RunResult | None = None self._turns: int = 0 @@ -849,55 +852,164 @@ def turns(self) -> int: def cache(self) -> SessionCache: return self._cache - @property - def harness(self) -> AsyncHarness: - return self._harness - @property def run_file(self) -> str | None: return self._harness.run_file def put(self, name: str, value: Any, *, overwrite: bool = False) -> str: + """Store a value in the session cache and return the handle used. + + Args: + name: Desired handle name. Must be a valid Python identifier. + value: Any Python object to store. + overwrite: Replace the existing handle if ``True``. + + Returns: + The handle name under which the value was stored. + """ return self._cache.put(name, value, overwrite=overwrite) def list_handles(self) -> dict[str, str]: + """Return a mapping of all cache handle names to their snapshot strings.""" return self._cache.list_handles() - async def ask_result(self, user_message: str) -> RunResult: - result = await self._harness.ask_result( - user_message, run_id=str(uuid.uuid4()), session_id=self._id - ) + def _record(self, result: RunResult) -> RunResult: self._last_result = result self._turns += result.turns - self._agent._last_harness = self._harness + self._agent._last_harness = self._harness # type: ignore[attr-defined] self._agent._last_run_file = self._harness.run_file + self._agent._last_result = result return result + +class AgentSession(_SessionBase): + """Stateful chat session built from an `Agent` definition. + + `Agent.run()` intentionally stays one-shot for examples and tests. Use + `Agent.session()` when an application needs follow-up questions over the + same message history and cache handles. + """ + + @property + def harness(self) -> Harness: + return self._harness + + def ask_result(self, user_message: str) -> RunResult: + """Send a follow-up message and return the full `RunResult`. + + Args: + user_message: The follow-up user prompt. + + Returns: + A `RunResult` for this turn sequence. + """ + return self._record( + self._harness.ask_result( + user_message, run_id=str(uuid.uuid4()), session_id=self._id + ) + ) + + def ask(self, user_message: str) -> str: + """Send a follow-up message and return the final text response. + + Args: + user_message: The follow-up user prompt. + + Returns: + The model's final text response. + + Raises: + MaxTurnsExceeded: If the loop reaches ``max_turns``. + RuntimeError: If the provider raises an exception. + """ + return unwrap_text(self.ask_result(user_message)) + + +class AsyncAgentSession(_SessionBase): + """Stateful async chat session built from an `AsyncAgent` definition.""" + + @property + def harness(self) -> AsyncHarness: + return self._harness + + async def ask_result(self, user_message: str) -> RunResult: + """Send a follow-up message and return the full `RunResult`.""" + return self._record( + await self._harness.ask_result( + user_message, run_id=str(uuid.uuid4()), session_id=self._id + ) + ) + async def ask(self, user_message: str) -> str: - result = await self.ask_result(user_message) - if result.status == "max_turns_exceeded": - from data_harness.exceptions import MaxTurnsExceeded + """Send a follow-up message and return the final text response. - raise MaxTurnsExceeded(result.turns) - if result.status == "error": - raise RuntimeError(result.error or "unknown error") - return result.text + Raises: + MaxTurnsExceeded: If the loop reaches ``max_turns``. + RuntimeError: If the provider raises an exception. + """ + return unwrap_text(await self.ask_result(user_message)) async def ask_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: - """Stream events for a follow-up turn.""" - async for event in self._harness.ask_stream(user_message): - yield event - self._agent._last_harness = self._harness - self._agent._last_run_file = self._harness.run_file + """Stream events for a follow-up turn. + + After the generator is exhausted, ``session.last_result`` holds the + `RunResult` for the turn, stamped with the session and run ids exactly + as a non-streamed turn would be. + """ + run_id = str(uuid.uuid4()) + async with aclosing(self._harness.ask_stream(user_message)) as events: + async for event in events: + yield event + result = self._harness.last_result + if result is not None: + self._record(self._harness._stamp(result, run_id, self._id)) + else: # pragma: no cover - the loop always sets a result + self._agent._last_harness = self._harness + self._agent._last_run_file = self._harness.run_file + + +def _synthetic_text_events(text: str) -> list[StreamEvent]: + """A minimal, well-formed event sequence carrying ``text`` and no usage. + + Used when a streamed run is served from the replay cache: there is no + provider stream to relay, but a streaming caller still renders from + deltas, so it must receive the answer the same way. + """ + from data_harness.llm.providers.base import StopReason + from data_harness.llm.streaming import ( + ContentBlockDeltaEvent, + ContentBlockStartEvent, + ContentBlockStopEvent, + MessageDeltaEvent, + MessageStartEvent, + MessageStopEvent, + TextDelta, + ) + from data_harness.llm.types import TextBlock + + return [ + MessageStartEvent(), + ContentBlockStartEvent(index=0, content_block=TextBlock(text="")), + ContentBlockDeltaEvent(index=0, delta=TextDelta(text=text)), + ContentBlockStopEvent(index=0), + MessageDeltaEvent( + stop_reason=StopReason.END_TURN, + input_tokens=0, + output_tokens=0, + cache_read_tokens=0, + cache_write_tokens=0, + ), + MessageStopEvent(), + ] _EXPLAIN_TEMPLATE = """\ Agent is a thin composition layer. The equivalent explicit wiring is: - from data_harness.cache import SessionCache - from data_harness.loop import Harness - from data_harness.tools.interpreter import PythonInterpreter - from data_harness.tools.variables import make_list_variables_spec + from data_harness.data.cache import SessionCache + from data_harness.data.harness import Harness + from data_harness.data.tools.interpreter import PythonInterpreter + from data_harness.data.tools.variables import make_list_variables_spec cache = SessionCache() tools = [ diff --git a/data_harness/cli.py b/data_harness/app/cli.py similarity index 98% rename from data_harness/cli.py rename to data_harness/app/cli.py index 1ba93a8..7af1b2f 100644 --- a/data_harness/cli.py +++ b/data_harness/app/cli.py @@ -16,7 +16,7 @@ import sys from typing import Any -from data_harness.quickstart import ask +from data_harness.app.quickstart import ask def build_parser() -> argparse.ArgumentParser: diff --git a/data_harness/notebook.py b/data_harness/app/notebook.py similarity index 88% rename from data_harness/notebook.py rename to data_harness/app/notebook.py index f6f077e..f834bb9 100644 --- a/data_harness/notebook.py +++ b/data_harness/app/notebook.py @@ -2,7 +2,7 @@ Load it in a notebook with:: - %load_ext data_harness.notebook + %load_ext data_harness.app.notebook Then ask questions about a DataFrame in the user namespace:: @@ -25,7 +25,7 @@ def _load_magics_class() -> Any: class DataHarnessMagics(Magics): @cell_magic def ask(self, line: str, cell: str) -> Any: - from data_harness.quickstart import ask as _ask + from data_harness.app.quickstart import ask as _ask var = line.strip() if not var: @@ -46,5 +46,5 @@ class UsageError(Exception): # type: ignore[no-redef] def load_ipython_extension(ipython: Any) -> None: - """Entry point for ``%load_ext data_harness.notebook``.""" + """Entry point for ``%load_ext data_harness.app.notebook``.""" ipython.register_magics(_load_magics_class()) diff --git a/data_harness/pandas.py b/data_harness/app/pandas.py similarity index 83% rename from data_harness/pandas.py rename to data_harness/app/pandas.py index b49d2dd..94b6d13 100644 --- a/data_harness/pandas.py +++ b/data_harness/app/pandas.py @@ -4,7 +4,7 @@ namespace, so it is **not** the documented headline (use ``ask`` / ``Chat``). Enable it explicitly:: - import data_harness.pandas # registers .chat on every DataFrame + import data_harness.app.pandas # registers .chat on every DataFrame df.chat("plot revenue by month") """ @@ -21,6 +21,6 @@ def __init__(self, df: pd.DataFrame) -> None: self._df = df def __call__(self, question: str, **kwargs: Any) -> Any: - from data_harness.quickstart import ask + from data_harness.app.quickstart import ask return ask(self._df, question, **kwargs) diff --git a/data_harness/quickstart.py b/data_harness/app/quickstart.py similarity index 77% rename from data_harness/quickstart.py rename to data_harness/app/quickstart.py index 7bc3571..39b0a2d 100644 --- a/data_harness/quickstart.py +++ b/data_harness/app/quickstart.py @@ -12,17 +12,18 @@ from __future__ import annotations import dataclasses +import importlib import importlib.util import os from collections.abc import Callable from typing import TYPE_CHECKING, Any -from data_harness.agent import Agent -from data_harness.io import to_handles -from data_harness.result import RunResult +from data_harness.app.agent import Agent +from data_harness.core.result import RunResult +from data_harness.data.io import to_handles if TYPE_CHECKING: - from data_harness.providers.base import ProviderAdapter + from data_harness.llm.providers.base import AsyncProviderAdapter, ProviderAdapter DEFAULT_ANTHROPIC_MODEL = "claude-sonnet-4-6" DEFAULT_OPENAI_MODEL = "gpt-4o-mini" @@ -52,39 +53,45 @@ def _is_openai_model(model: str) -> bool: return model.lower().startswith(_OPENAI_PREFIXES) -def resolve_adapter(model: str | None = None) -> ProviderAdapter: - """Resolve a `ProviderAdapter` from an explicit model or the environment. +#: provider key -> (sync class, async class, name of the extra that supplies it) +_ADAPTER_CLASSES: dict[str, tuple[str, str, str | None]] = { + "anthropic": ("AnthropicAdapter", "AsyncAnthropicAdapter", None), + "openai": ("OpenAIAdapter", "AsyncOpenAIAdapter", "openai"), + "openrouter": ("OpenRouterAdapter", "AsyncOpenRouterAdapter", "openai"), + "deepseek": ("DeepSeekAdapter", "AsyncDeepSeekAdapter", "openai"), +} + +#: provider key -> environment variable consulted when no model is given, +#: in preference order. +_ENV_PREFERENCE: tuple[tuple[str, str, str], ...] = ( + ("anthropic", "ANTHROPIC_API_KEY", DEFAULT_ANTHROPIC_MODEL), + ("openai", "OPENAI_API_KEY", DEFAULT_OPENAI_MODEL), + ("openrouter", "OPENROUTER_API_KEY", DEFAULT_OPENROUTER_MODEL), + ("deepseek", "DEEPSEEK_API_KEY", DEFAULT_DEEPSEEK_MODEL), +) - With no ``model``, prefers ``ANTHROPIC_API_KEY``, then ``OPENAI_API_KEY``, - then ``OPENROUTER_API_KEY``, then ``DEEPSEEK_API_KEY``. With a ``model``, - routes by name: a ``provider/model`` id (containing ``/``) goes to OpenRouter, - ``deepseek*`` to DeepSeek's direct API, ``gpt*``/``o*`` to OpenAI, otherwise - Anthropic. + +def _route(model: str | None) -> tuple[str, str]: + """Map a model name (or the environment) onto ``(provider_key, model_id)``. + + Sync and async resolution share this routing so the two can never disagree + about which provider a model name belongs to. Raises: RuntimeError: If no provider can be resolved (no key, no model). """ if model is not None: if "/" in model: - return _make_openrouter(model) + return "openrouter", model if model.startswith("deepseek"): - return _make_deepseek(model) + return "deepseek", model if _is_openai_model(model): - return _make_openai(model) - from data_harness.providers.anthropic import AnthropicAdapter + return "openai", model + return "anthropic", model - return AnthropicAdapter(model=model) - - if os.environ.get("ANTHROPIC_API_KEY"): - from data_harness.providers.anthropic import AnthropicAdapter - - return AnthropicAdapter(model=DEFAULT_ANTHROPIC_MODEL) - if os.environ.get("OPENAI_API_KEY"): - return _make_openai(DEFAULT_OPENAI_MODEL) - if os.environ.get("OPENROUTER_API_KEY"): - return _make_openrouter(DEFAULT_OPENROUTER_MODEL) - if os.environ.get("DEEPSEEK_API_KEY"): - return _make_deepseek(DEFAULT_DEEPSEEK_MODEL) + for provider, env_var, default_model in _ENV_PREFERENCE: + if os.environ.get(env_var): + return provider, default_model raise RuntimeError( "No provider configured. Set ANTHROPIC_API_KEY, OPENAI_API_KEY, " @@ -93,37 +100,54 @@ def resolve_adapter(model: str | None = None) -> ProviderAdapter: ) -def _make_openai(model: str) -> ProviderAdapter: +def _build_adapter(model: str | None, *, is_async: bool) -> Any: + provider, model_id = _route(model) + sync_name, async_name, extra = _ADAPTER_CLASSES[provider] + module_name = ( + "data_harness.llm.providers.anthropic" + if provider == "anthropic" + else "data_harness.llm.providers.openai" + ) try: - from data_harness.providers.openai import OpenAIAdapter + module = importlib.import_module(module_name) except ImportError as exc: # pragma: no cover - exercised via install matrix + if extra is None: + # A required dependency failed to import. Say what actually broke + # rather than inventing an extra that does not exist. + raise raise RuntimeError( - "OpenAI support requires the 'openai' extra: pip install " - "'data-harness[openai]'." + f"{provider} support requires the {extra!r} extra: pip install " + f"'data-harness[{extra}]'." ) from exc - return OpenAIAdapter(model=model) + cls = getattr(module, async_name if is_async else sync_name) + return cls(model=model_id) -def _make_openrouter(model: str) -> ProviderAdapter: - try: - from data_harness.providers.openai import OpenRouterAdapter - except ImportError as exc: # pragma: no cover - exercised via install matrix - raise RuntimeError( - "OpenRouter support requires the 'openai' extra: pip install " - "'data-harness[openai]'." - ) from exc - return OpenRouterAdapter(model=model) +def resolve_adapter(model: str | None = None) -> ProviderAdapter: + """Resolve a `ProviderAdapter` from an explicit model or the environment. + With no ``model``, prefers ``ANTHROPIC_API_KEY``, then ``OPENAI_API_KEY``, + then ``OPENROUTER_API_KEY``, then ``DEEPSEEK_API_KEY``. With a ``model``, + routes by name: a ``provider/model`` id (containing ``/``) goes to OpenRouter, + ``deepseek*`` to DeepSeek's direct API, ``gpt*``/``o*`` to OpenAI, otherwise + Anthropic. -def _make_deepseek(model: str) -> ProviderAdapter: - try: - from data_harness.providers.openai import DeepSeekAdapter - except ImportError as exc: # pragma: no cover - exercised via install matrix - raise RuntimeError( - "DeepSeek support requires the 'openai' extra: pip install " - "'data-harness[openai]'." - ) from exc - return DeepSeekAdapter(model=model) + Raises: + RuntimeError: If no provider can be resolved (no key, no model). + """ + return _build_adapter(model, is_async=False) + + +def resolve_async_adapter(model: str | None = None) -> AsyncProviderAdapter: + """Resolve an `AsyncProviderAdapter` using the same routing as `resolve_adapter`. + + Args: + model: Explicit model id, or ``None`` to pick from the environment. + + Raises: + RuntimeError: If no provider can be resolved (no key, no model). + """ + return _build_adapter(model, is_async=True) def _build_agent( diff --git a/data_harness/core/__init__.py b/data_harness/core/__init__.py new file mode 100644 index 0000000..813961c --- /dev/null +++ b/data_harness/core/__init__.py @@ -0,0 +1,12 @@ +"""The harness: the ReAct loop and the types describing a run. + +Owns `Harness`/`AsyncHarness`, the effect protocol they are driven by, +`RunResult`, run logging, and the exception taxonomy. + +Deliberately knows nothing about the data domain: no pandas, no `SessionCache`, +no interpreter. What a tool *returns* and what a run's final state *is* are +supplied by whoever builds the harness, through `RunEnvironment`. That is what +makes the loop reusable for a domain other than data analysis. + +May import `llm`. May not import `data` or `app`. +""" diff --git a/data_harness/artifacts.py b/data_harness/core/artifacts.py similarity index 100% rename from data_harness/artifacts.py rename to data_harness/core/artifacts.py diff --git a/data_harness/core/compaction.py b/data_harness/core/compaction.py new file mode 100644 index 0000000..4e7a317 --- /dev/null +++ b/data_harness/core/compaction.py @@ -0,0 +1,207 @@ +"""Compaction: keeping a long run inside the context window. + +Before this the only context management was `max_turns`, and hitting it raised +`MaxTurnsExceeded`. That is a wall, not a strategy: a run that needs thirty +turns fails at twenty-five having spent the money for all of them. + +Compaction summarises the older part of the conversation and replays the +summary in its place. Two things make it safe rather than lossy: + +**The cut lands on a turn boundary.** An assistant message asking for a tool +is never separated from the result answering it. Split them and the provider +rejects the request outright, which is the first bug everyone writes here. + +**Nothing is deleted.** The summary is a session entry, so the compacted turns +stay in the tree and moving the leaf back restores them. Compaction is a view, +not an edit. + +There is also a reason compaction costs this library less than it costs a +coding agent: the data lives in the session cache under handles, not in the +transcript. Compacting away the turn that loaded a DataFrame does not lose the +DataFrame, only the discussion of it. What is being summarised is reasoning, +not state. +""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass + +from data_harness.core.exceptions import ConfigurationError +from data_harness.core.session.entries import Entry, MessageEntry +from data_harness.llm.types import ( + Message, + TextBlock, + ToolResultBlock, + ToolUseBlock, +) + +#: Rough characters per token. Deliberately crude: the decision this feeds is +#: "are we near the limit", and being wrong by 20% moves the trigger point +#: slightly rather than breaking anything. A real tokeniser would tie the core +#: to a provider's vocabulary for no benefit at this precision. +CHARS_PER_TOKEN = 4 + + +@dataclass(frozen=True) +class CompactionSettings: + """When to compact and how much to keep. + + Attributes: + context_window: The model's limit, in tokens. + reserve_tokens: Headroom left for the next response and its tools. + Compaction triggers once the context exceeds the window minus + this. + keep_recent_tokens: Roughly how much recent conversation to carry + through uncompacted, so the model keeps its immediate footing. + """ + + context_window: int = 128_000 + reserve_tokens: int = 24_000 + keep_recent_tokens: int = 16_000 + + def __post_init__(self) -> None: + if self.keep_recent_tokens >= self.context_window - self.reserve_tokens: + raise ConfigurationError( + "keep_recent_tokens must leave room under the trigger point, " + f"got keep={self.keep_recent_tokens} with window=" + f"{self.context_window} and reserve={self.reserve_tokens}" + ) + + +def estimate_tokens(messages: list[Message]) -> int: + """A crude token estimate for a conversation. See `CHARS_PER_TOKEN`.""" + characters = 0 + for message in messages: + for block in message.content: + if isinstance(block, TextBlock): + characters += len(block.text) + elif isinstance(block, ToolUseBlock): + characters += len(block.tool_name) + len(str(block.tool_input)) + elif isinstance(block, ToolResultBlock): + characters += len(block.content) + return characters // CHARS_PER_TOKEN + + +def should_compact(context_tokens: int, settings: CompactionSettings) -> bool: + """Whether the conversation has grown past its safe size.""" + return context_tokens > settings.context_window - settings.reserve_tokens + + +def starts_a_turn(entries: list[Entry], index: int) -> bool: + """Whether ``index`` is a safe place to cut. + + Safe means a user message that is not carrying tool results. Cutting + before a tool-result message would orphan the assistant tool call that + asked for it, and cutting before an assistant message would open the + conversation with one, which providers reject. + """ + entry = entries[index] + if not isinstance(entry, MessageEntry) or entry.message is None: + return False + if entry.message.role != "user": + return False + return not any( + isinstance(block, ToolResultBlock) for block in entry.message.content + ) + + +def find_cut_point(entries: list[Entry], settings: CompactionSettings) -> str | None: + """The id of the first entry to keep, or ``None`` if there is nowhere safe. + + Walks back from the newest entry accumulating tokens, then keeps walking + until it reaches a turn boundary. Returning ``None`` means the whole + conversation is one indivisible turn, which is not compactable: better to + do nothing than to produce a transcript the provider will refuse. + """ + kept = 0 + boundary: int | None = None + for index in range(len(entries) - 1, -1, -1): + entry = entries[index] + if isinstance(entry, MessageEntry) and entry.message is not None: + kept += estimate_tokens([entry.message]) + if kept >= settings.keep_recent_tokens and starts_a_turn(entries, index): + boundary = index + break + + if boundary is None or boundary == 0: + # Nothing before the boundary to compact. + return None + return entries[boundary].id + + +def messages_before(entries: list[Entry], cut_id: str) -> list[Message]: + """The messages a compaction would summarise: everything before the cut.""" + summarised: list[Message] = [] + for entry in entries: + if entry.id == cut_id: + break + if isinstance(entry, MessageEntry) and entry.message is not None: + summarised.append(entry.message) + return summarised + + +SUMMARY_INSTRUCTIONS = """\ +Summarise the conversation so far for an agent that will continue the work \ +without seeing it. + +Keep: the user's goal, decisions taken and why, findings that later steps \ +depend on, and the names of any cached data handles that were created. + +Do not restate the data itself. It is still available under its handles; only \ +the discussion of it is being replaced.""" + + +#: Given the messages being replaced, return prose standing in for them. +Summarizer = Callable[[list[Message]], str] + +#: Called before each turn with the run's session. Returns the id of the +#: compaction entry it appended, or ``None`` if it did nothing. +Compactor = Callable[[object], "str | None"] + + +def make_compactor( + summarize: Summarizer, + settings: CompactionSettings | None = None, +) -> Compactor: + """Build a `Compactor` the harness consults before each turn. + + ``summarize`` is injected rather than built in because summarising means + calling a model, and which model at what cost is the application's + decision, not the loop's. + """ + resolved = settings if settings is not None else CompactionSettings() + + def compact(session: object) -> str | None: + return maybe_compact(session, summarize, resolved) + + return compact + + +def maybe_compact( + session, summarize: Summarizer, settings: CompactionSettings +) -> str | None: + """Compact ``session`` if it has grown too large. Returns the new entry id. + + Does nothing, and says so by returning ``None``, when the context is small + enough or when there is no safe place to cut. + """ + context = session.build_context() + tokens = estimate_tokens(context) + if not should_compact(tokens, settings): + return None + + entries = session.context_entries() + cut_id = find_cut_point(entries, settings) + if cut_id is None: + return None + + replaced = messages_before(entries, cut_id) + if not replaced: + return None + + return session.append_compaction( + summary=summarize(replaced), + first_kept_entry_id=cut_id, + tokens_before=tokens, + ) diff --git a/data_harness/core/environment.py b/data_harness/core/environment.py new file mode 100644 index 0000000..107b146 --- /dev/null +++ b/data_harness/core/environment.py @@ -0,0 +1,83 @@ +"""What the loop needs from a domain, and nothing more. + +The loop has to do two things it cannot decide by itself: turn a tool's return +value into text for the transcript, and describe the run's final state for the +`RunResult`. For a data harness both answers come from the session cache: a +DataFrame becomes a handle plus a snapshot, and the final state is the set of +handles, the recorded answer, and any charts. + +None of that is true of an agent in another domain, so the loop asks a +`RunEnvironment` instead of reaching for a `SessionCache`. That is the whole +reason `data_harness.core` can be read, tested, and reused without pandas. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Protocol, runtime_checkable + +from data_harness.core.artifacts import ChartArtifact +from data_harness.core.result import CacheStorageInfo + + +@dataclass(frozen=True) +class RunState: + """The domain's contribution to a `RunResult`, captured when a run ends. + + Attributes: + snapshots: Handle name -> compact, payload-free description. + storage: Handle name -> where the value physically lives. + value: The structured final answer the run recorded, if any. + artifacts: Charts or other artefacts produced during the run. + """ + + snapshots: dict[str, str] = field(default_factory=dict) + storage: dict[str, CacheStorageInfo] = field(default_factory=dict) + value: Any = None + artifacts: list[ChartArtifact] = field(default_factory=list) + + +@runtime_checkable +class RunEnvironment(Protocol): + """The domain services the loop depends on. + + Implementations must not raise from these methods: they run while the loop + is assembling a turn or a result, where an exception would lose the run + rather than report it. + """ + + def render_tool_output(self, value: Any) -> str: + """Render a tool's return value as the text the model will read. + + Deliberately not named after `data.format.format_tool_output`, which + is only one possible implementation of it. + """ + ... + + def capture(self) -> RunState: + """Describe the current domain state for a `RunResult`.""" + ... + + def storage_metadata(self) -> dict[str, dict[str, str]]: + """Per-handle storage detail for the turn log. Empty when there is none.""" + ... + + +class NullEnvironment: + """A domain-free environment: no cache, no handles, no artefacts. + + The default, so `data_harness.core` stands alone. Tool output is rendered + with `str`, which is right for text-returning tools and lossy for anything + large: that is exactly the problem the data layer's cache solves. + """ + + def render_tool_output(self, value: Any) -> str: + if isinstance(value, Exception): + return f"Error: {type(value).__name__}: {value}" + return str(value) + + def capture(self) -> RunState: + return RunState() + + def storage_metadata(self) -> dict[str, dict[str, str]]: + return {} diff --git a/data_harness/core/exceptions.py b/data_harness/core/exceptions.py new file mode 100644 index 0000000..9940025 --- /dev/null +++ b/data_harness/core/exceptions.py @@ -0,0 +1,98 @@ +"""The exception taxonomy. + +Every failure this library raises carries a stable ``code``. That is the point +of the taxonomy: a caller deciding between retrying, showing the user an +error, and billing for the attempt needs to tell a rate-limited provider from +a typo in the model's pandas apart from a sandbox timeout. Matching on message +text is how that decision silently rots. + +Codes are part of the API. Messages are not, and may be reworded. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from data_harness.llm.providers.base import NormalizedResponse + + +class DataHarnessError(Exception): + """Base for everything this library raises deliberately. + + ``except DataHarnessError`` catches the library's own failures without + also swallowing a bug in the caller's tool handler. + """ + + #: Stable, machine-readable identifier. Subclasses set it. + code: str = "unknown" + + def __init__(self, message: str, *, code: str | None = None) -> None: + super().__init__(message) + if code is not None: + self.code = code + + +class MaxTurnsExceeded(DataHarnessError, RuntimeError): + """The loop reached ``max_turns`` without the model finishing. + + Still a `RuntimeError` because it was one before the taxonomy existed and + callers catch it that way. + + Attributes: + turns: How many turns ran before the limit was hit. + last_response: The final provider response, if available. + """ + + code = "max_turns_exceeded" + + def __init__(self, turns: int, last_response: NormalizedResponse | None = None): + self.turns = turns + self.last_response = last_response + super().__init__(f"Max turns exceeded: {turns}") + + +class ToolNotFoundError(DataHarnessError, KeyError): + """A tool invocation named a tool that is not registered.""" + + code = "tool_not_found" + + +class SubagentRecursionError(DataHarnessError, RuntimeError): + """A subagent tried to spawn another subagent.""" + + code = "subagent_recursion" + + +class ProviderError(DataHarnessError, RuntimeError): + """The model provider failed: rate limit, auth, timeout, malformed reply. + + Distinct from a tool failing, which is reported to the model as a tool + result rather than raised. A provider failure ends the run. + + Also a `RuntimeError`, because that is what a failed run raised before the + taxonomy existed and callers catch it that way. + """ + + code = "provider_error" + + +class ExecutionError(DataHarnessError, RuntimeError): + """Sandboxed code could not be run: timeout, killed, environment missing. + + Distinct from code that ran and raised, which is the model's problem and + is handed back to it as a tool result to fix. + """ + + code = "execution_error" + + +class ConfigurationError(DataHarnessError, ValueError): + """The harness was wired up in a way that cannot work. + + Raised at construction where possible, so a misconfiguration fails before + any tokens are spent rather than on turn nine. Also a `ValueError`, which + is what these checks raised before the taxonomy. + """ + + code = "configuration_error" diff --git a/data_harness/core/hooks.py b/data_harness/core/hooks.py new file mode 100644 index 0000000..2b6f5a1 --- /dev/null +++ b/data_harness/core/hooks.py @@ -0,0 +1,202 @@ +"""Hooks: the loop's extension points. + +Before this, the only way to influence a run without forking the loop was +`register_reminder`, which could append text to the prompt suffix and nothing +else. Everything more interesting was hardcoded. The interpreter approval gate +lived inside the dispatcher and was keyed on the literal tool name +``"python_interpreter"``; a budget check had to be reimplemented by every +application on top. + +A hook observes an event and may return a decision. Returning ``None`` means +"no opinion", which is what most hooks do most of the time. + +Contract, and it matters: **a hook must not raise.** It runs while the loop is +assembling a turn, where an exception loses the run instead of reporting it. +Hooks that raise anyway are caught and reported through `HookError` rather +than corrupting the run, but a hook that needs to signal a problem should +return a decision saying so. +""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import dataclass, field +from typing import Any, TypeVar + +from data_harness.core.exceptions import DataHarnessError +from data_harness.llm.types import Message, ToolResultBlock + +# ── events ────────────────────────────────────────────────────────────────── + + +@dataclass(frozen=True) +class BeforeTurn: + """A provider turn is about to start. + + Return a `Reminder` to append text to the conversation suffix, or a `Stop` + to end the run before spending anything. + """ + + turn: int + max_turns: int + messages: list[Message] + + +@dataclass(frozen=True) +class BeforeToolCall: + """A tool is about to run, after the model chose it. + + Return a `Block` to refuse the call. The model is told what it was told, + so the reason is part of the prompt, not just a log line. + """ + + turn: int + tool_name: str + tool_input: dict + + +@dataclass(frozen=True) +class AfterToolCall: + """A tool returned. Return a `Replace` to rewrite what the model sees.""" + + turn: int + tool_name: str + tool_input: dict + result: ToolResultBlock + + +@dataclass(frozen=True) +class AfterTurn: + """A turn finished and has been accounted for. + + Return a `Stop` to end the run. This is where a spend cap belongs: the + tokens are already counted, so the decision is made on real numbers. + """ + + turn: int + max_turns: int + input_tokens: int + output_tokens: int + tool_results: list[ToolResultBlock] + + +Event = BeforeTurn | BeforeToolCall | AfterToolCall | AfterTurn + + +# ── decisions ─────────────────────────────────────────────────────────────── + + +@dataclass(frozen=True) +class Reminder: + """Append ``text`` to the conversation suffix before the turn. + + Suffix-only on purpose: the system prompt stays byte-identical across + turns so the provider's cache of it is not invalidated. + """ + + text: str + + +@dataclass(frozen=True) +class Block: + """Refuse the tool call and return ``reason`` to the model instead. + + ``is_error`` defaults to False because a refusal is a decision, not a + malfunction: telling the model its code was broken makes it rewrite and + retry rather than stop. + """ + + reason: str + is_error: bool = False + + +@dataclass(frozen=True) +class Replace: + """Rewrite a tool result before the model sees it.""" + + content: str + is_error: bool = False + + +@dataclass(frozen=True) +class Stop: + """End the run cleanly. ``reason`` is recorded on the `RunResult`.""" + + reason: str + + +Decision = Reminder | Block | Replace | Stop +Hook = Callable[[Any], Decision | None] + +_E = TypeVar("_E", bound=Event) +_D = TypeVar("_D", bound=Decision) + + +class HookError(DataHarnessError): + """A hook raised. Names the hook and the event so the culprit is obvious.""" + + code = "hook_error" + + def __init__(self, event: Event, hook: Hook, cause: BaseException) -> None: + name = getattr(hook, "__qualname__", repr(hook)) + super().__init__( + f"Hook {name} raised on {type(event).__name__}: {cause!r}. " + "Hooks must return a decision rather than raise." + ) + self.event = event + self.hook = hook + self.__cause__ = cause + + +@dataclass +class HookRegistry: + """Hooks grouped by the event type they care about, in registration order.""" + + hooks: dict[type, list[Hook]] = field(default_factory=dict) + + def add(self, event_type: type[_E], hook: Callable[[_E], Decision | None]) -> None: + self.hooks.setdefault(event_type, []).append(hook) + + def emit(self, event: Event) -> list[Decision]: + """Run every hook for ``event`` and collect the decisions. + + Every hook sees the event, even after one has decided: a hook that + records spend should not be skipped because an earlier one asked to + stop. Precedence between conflicting decisions is the caller's to + resolve. + + Raises: + HookError: If a hook raises. The loop turns this into a failed + run rather than letting it escape mid-turn. + """ + decisions: list[Decision] = [] + for hook in self.hooks.get(type(event), ()): + try: + decision = hook(event) + except Exception as exc: # noqa: BLE001 - re-raised as HookError + raise HookError(event, hook, exc) from exc + if decision is not None: + decisions.append(decision) + return decisions + + def first(self, event: Event, kind: type[_D]) -> _D | None: + """The first decision of type ``kind``, or ``None``. + + Later decisions of the same kind are discarded. Conflicting policies + are the caller's problem to avoid, not this class's to arbitrate. + """ + for decision in self.emit(event): + if isinstance(decision, kind): + return decision + return None + + def copy(self) -> HookRegistry: + """An independent registry with the same hooks. + + Taken whenever a registry is handed to a harness. The harness adds its + own hooks (the approval gate, for one), and writing those into the + caller's object would leak one harness's policy into every other + harness sharing it, which is exactly the reuse the parameter exists + to support. + """ + return HookRegistry({k: list(v) for k, v in self.hooks.items()}) diff --git a/data_harness/logger.py b/data_harness/core/logger.py similarity index 94% rename from data_harness/logger.py rename to data_harness/core/logger.py index 1976d97..a31c652 100644 --- a/data_harness/logger.py +++ b/data_harness/core/logger.py @@ -8,9 +8,9 @@ from loguru import logger -from data_harness.providers.base import NormalizedResponse -from data_harness.serialize import to_jsonable -from data_harness.types import Message, ToolAnnotations, ToolResultBlock, ToolSpec +from data_harness.core.serialize import to_jsonable +from data_harness.llm.providers.base import NormalizedResponse +from data_harness.llm.types import Message, ToolAnnotations, ToolResultBlock, ToolSpec def log_error_turn( diff --git a/data_harness/core/loop.py b/data_harness/core/loop.py new file mode 100644 index 0000000..dc36adf --- /dev/null +++ b/data_harness/core/loop.py @@ -0,0 +1,1253 @@ +"""The ReAct loop. + +There is exactly one loop implementation: `_HarnessBase._plan`. It is a +generator that owns every decision in a turn (reminders, message assembly, +tool gating, when to stop, what the `RunResult` says) and does no *network* +I/O itself. Instead it asks for it by yielding an effect: + +- `CallProvider` — run one provider turn +- `CallTool` — invoke one tool handler +- `ToolFinished` — a tool result block is ready, emit an event if you want one + +The driver answers `CallProvider` and `CallTool` with `Ok(value)` or +`Failed(exc)`; the loop, not the driver, decides what a failure means. The +wrapper matters: without it a handler that legitimately returned a `Failed` +would be mistaken for one that raised. + +Turn logging is the one effect not modelled as an effect: `_plan` writes JSONL +to disk directly. That is a deliberate leftover, and it is why the "no I/O" +claim above is scoped to the network. Phase 3 replaces the log with a session +store and the write moves behind an interface. + +Two drivers perform the effects: + +- `Harness` runs everything inline on the calling thread. No event loop is + created, so `KeyboardInterrupt` interrupts promptly, handlers keep their + thread affinity (a `sqlite3` connection built at setup still works), and the + caller's ambient event loop is untouched. +- `AsyncHarness` awaits the provider, offloads blocking handlers to + `asyncio.to_thread` so they cannot stall the event loop, and can stream. + +This replaced four near-copies of the loop (sync/async x result/stream) that +had already drifted: the streaming copy swallowed provider errors without +recording the tokens it had already spent. +""" + +from __future__ import annotations + +import asyncio +import concurrent.futures +import dataclasses +import enum +import functools +import sys +from collections.abc import AsyncGenerator, Callable, Coroutine, Generator +from contextlib import aclosing +from dataclasses import dataclass +from typing import Any, Literal, TypeVar + +from data_harness.core.compaction import Compactor +from data_harness.core.environment import NullEnvironment, RunEnvironment +from data_harness.core.exceptions import ConfigurationError +from data_harness.core.hooks import ( + AfterToolCall, + AfterTurn, + BeforeToolCall, + BeforeTurn, + Block, + Event, + HookError, + HookRegistry, + Reminder, + Replace, + Stop, +) +from data_harness.core.logger import log_error_turn, log_turn, setup_logger +from data_harness.core.observe import time_block +from data_harness.core.result import RunResult, Usage, unwrap_text +from data_harness.core.session import Session +from data_harness.llm.providers.base import ( + AsyncProviderAdapter, + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.streaming import ( + StreamEvent, + ToolResultEvent, + accumulate_stream_events, +) +from data_harness.llm.types import ( + Message, + TextBlock, + ToolResultBlock, + ToolSpec, + ToolUseBlock, +) + +_MAX_TURN_REMINDER = ( + "This is the final turn. You MUST produce your complete final output now. " + "Do not use any more tools. Respond with your answer directly." +) + +_T = TypeVar("_T") + + +# ── effects ───────────────────────────────────────────────────────────────── + + +@dataclass(frozen=True) +class CallProvider: + """Run one provider turn. Answer with `Ok(NormalizedResponse)` or `Failed`.""" + + system: str + messages: list[Message] + tools: list[ToolSpec] + + +@dataclass(frozen=True) +class CallTool: + """Invoke one tool handler. Answer with `Ok(return_value)` or `Failed`.""" + + tool_use_id: str + tool_name: str + handler: Callable[..., Any] + tool_input: dict + + +@dataclass(frozen=True) +class ToolFinished: + """A tool result block is ready. Answer with ``None``. + + Emitted for every result, including tool-not-found and gate-blocked calls, + so a streaming driver can surface them all. + """ + + tool_name: str + block: ToolResultBlock + + +@dataclass(frozen=True) +class Ok: + """An effect succeeded, carrying whatever it produced. + + Successes are wrapped rather than sent raw so that no possible return + value can be confused with a `Failed`. + """ + + value: Any + + +@dataclass(frozen=True) +class Failed: + """An effect raised. The loop, not the driver, decides what that means.""" + + error: BaseException + + +Effect = CallProvider | CallTool | ToolFinished + +#: What a driver may send back. `CallProvider` and `CallTool` require an +#: `Ok` or a `Failed`; `ToolFinished` is an announcement and takes ``None``. +Answer = Ok | Failed | None + + +def _require_answer(effect: Effect, answer: Answer) -> Ok | Failed: + """Reject a driver that answered a request with nothing. + + Without this the loop would do ``answer.value`` and raise + ``AttributeError: 'NoneType' object has no attribute 'value'`` several + frames from the driver that actually got it wrong. + """ + if isinstance(answer, (Ok, Failed)): + return answer + raise TypeError( + f"{type(effect).__name__} must be answered with Ok(...) or Failed(...), " + f"got {answer!r}" + ) + + +# ── helpers ───────────────────────────────────────────────────────────────── + + +#: The tool the approval gate applies to. Only this one takes free-form code, +#: so only this one needs approving. +GATED_TOOL = "python_interpreter" + + +def make_code_gate( + on_code: Callable[[str], object] | None, code_only: bool +) -> Callable[[BeforeToolCall], Block | None]: + """Build the interpreter approval gate as a `BeforeToolCall` hook. + + This used to live inside the tool dispatcher, keyed on the literal tool + name, which meant the one general-purpose thing the loop could do about a + tool call was available only to this one feature. It is a hook now, so the + same mechanism serves budget caps, audit trails, and anything else. + """ + + def gate(event: BeforeToolCall) -> Block | None: + if event.tool_name != GATED_TOOL: + return None + code = event.tool_input.get("code", "") + if code_only: + return Block(f"DRY RUN — code not executed:\n{code}") + if on_code is not None: + decision = on_code(code) + if decision is False: + return Block("Execution blocked by the approval gate.") + if isinstance(decision, str): + return Block(decision) + return None + + return gate + + +class _Unreadable(enum.Enum): + """Sentinel type, so `_ambient_event_loop` keeps a real return union.""" + + TOKEN = "unreadable" + + +#: Returned by `_ambient_event_loop` when the thread's loop cannot be read +#: without risking a side effect. Distinct from ``None``, which means "read +#: successfully, and there is no loop". +_UNREADABLE = _Unreadable.TOKEN + + +def _ambient_event_loop() -> asyncio.AbstractEventLoop | None | _Unreadable: + """The loop already installed on this thread, ``None``, or `_UNREADABLE`. + + Must not install one as a side effect. Reading it needs two strategies, + because the safe way to ask changed: + + - 3.14 and later: `asyncio.get_event_loop` raises when no loop is set + instead of creating one, so asking is safe. The policy API it used to + need is deprecated there and goes away in 3.16. + - 3.10 to 3.13: `get_event_loop` *creates* a loop when none is set, which + would leave a never-run, never-closed loop and its selector fd behind on + a thread that only ever used the synchronous API. There is no public + read-only accessor, so the policy's private slot is what we read. + + Under an exotic policy that does not expose that slot we decline to guess: + creating a loop to find out whether one exists is the exact bug this + guards against. + """ + if sys.version_info >= (3, 14): + try: + return asyncio.get_event_loop() + except RuntimeError: + return None + + policy = asyncio.get_event_loop_policy() + local = getattr(policy, "_local", None) + if local is None: + return _UNREADABLE + return getattr(local, "_loop", None) + + +def run_coroutine_blocking(coro: Coroutine[Any, Any, _T]) -> _T: + """Run ``coro`` to completion from synchronous code. + + Only needed when synchronous code must drive something inherently async + (an `AsyncProviderAdapter`, an async tool handler). The synchronous + `Harness` does not use this for its own loop. + + Restores whatever event loop was installed on the calling thread, because + `asyncio.run` otherwise clears it. When a loop is already *running*, the + coroutine goes to a worker thread with its own loop, since `asyncio.run` + cannot be nested. + """ + try: + asyncio.get_running_loop() + except RuntimeError: + previous = _ambient_event_loop() + loop = asyncio.new_event_loop() + try: + asyncio.set_event_loop(loop) + return loop.run_until_complete(coro) + finally: + try: + loop.run_until_complete(loop.shutdown_asyncgens()) + finally: + loop.close() + # `set_event_loop(loop)` already ran, so simply skipping the + # restore would leave the thread pointing at a *closed* loop, + # which is worse than the leak this branch exists to avoid. + # When the previous loop was unreadable, `None` is the only + # honest answer. + asyncio.set_event_loop(None if previous is _UNREADABLE else previous) + + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: + return pool.submit(asyncio.run, coro).result() + + +# ── the loop ──────────────────────────────────────────────────────────────── + + +class _HarnessBase: + """Loop state, pure turn logic, and the effect generator. + + Both `Harness` and `AsyncHarness` inherit this, so the loop exists once and + subclasses of either see the same state and the same overridable seams. + """ + + def __init__( + self, + system: str, + tools: list[ToolSpec], + max_turns: int = 25, + run_dir: str = "./runs", + environment: RunEnvironment | None = None, + on_code: Callable[[str], object] | None = None, + code_only: bool = False, + session: Session | None = None, + hooks: HookRegistry | None = None, + compactor: Compactor | None = None, + ) -> None: + if max_turns < 1: + raise ConfigurationError(f"max_turns must be at least 1, got {max_turns!r}") + self._system = system + self._tools = list(tools) + self._max_turns = max_turns + self._run_dir = run_dir + self._environment = ( + environment if environment is not None else NullEnvironment() + ) + self._session = session if session is not None else Session() + self._hooks = hooks.copy() if hooks is not None else HookRegistry() + if on_code is not None or code_only: + self._hooks.add(BeforeToolCall, make_code_gate(on_code, code_only)) + self._stop_reason: str | None = None + self._compactor = compactor + # Resuming: a session handed in with history already on it seeds the + # working copy, so `ask` continues the conversation rather than + # starting a second one that shares a log with the first. + self._messages: list[Message] = self._session.build_context() + self._reminders: list[Callable[[int, int], str | None]] = [] + self._run_file: str | None = None + # Kept for introspection only. The gate hook closes over these at + # construction, so assigning them afterwards changes nothing; build a + # new harness, or register your own BeforeToolCall hook. + self._on_code = on_code + self._code_only = code_only + self._last_result: RunResult | None = None + self._turns_completed = 0 + self._usage_so_far = Usage() + + def register_reminder(self, hook: Callable[[int, int], str | None]) -> None: + """Register a suffix reminder called before each provider turn. + + Kept for the callers that predate hooks. It is a `BeforeTurn` hook + returning a `Reminder`, which is the general form and can also refuse + a tool call, rewrite a result, or stop the run. + + Args: + hook: Callable with signature ``(turn: int, max_turns: int) -> str | None``. + """ + self._reminders.append(hook) + + def on(self, event_type: type[Event], hook: Callable[[Any], Any]) -> None: + """Register ``hook`` for ``event_type``. See `data_harness.core.hooks`.""" + self._hooks.add(event_type, hook) + + @property + def hooks(self) -> HookRegistry: + """The hooks this harness will consult.""" + return self._hooks + + # ── inspection ────────────────────────────────────────────────────────── + # + # `tools`, `messages`, and `reminders` return the live lists: mutating them + # mutates the harness, which is how tools get added to a live session. + + @property + def run_file(self) -> str | None: + """Path to the JSONL log for this run, or ``None`` before the first run.""" + return self._run_file + + @property + def last_result(self) -> RunResult | None: + """`RunResult` for the most recently completed loop, including streamed ones. + + Streaming callers get their token usage, cache snapshots, charts, and + error status from here: the event protocol carries no run-level summary. + Stays ``None`` if a stream is abandoned before it is exhausted. + """ + return self._last_result + + @property + def system(self) -> str: + """The system prompt. Byte-identical across every turn of a run.""" + return self._system + + @property + def tools(self) -> list[ToolSpec]: + """The full tool list, including invisible tools. Live and mutable.""" + return self._tools + + @property + def max_turns(self) -> int: + """Hard cap on provider turns per run.""" + return self._max_turns + + @property + def environment(self) -> RunEnvironment: + """The domain services this run uses. See `RunEnvironment`.""" + return self._environment + + @property + def session(self) -> Session: + """The durable record of this harness's runs. + + `messages` is the working copy the loop mutates in flight; the session + is the append-only log. `session.build_context()` equals `messages` + after any run, which `tests/test_session_tree.py` pins for single runs, + repeated runs, asks, and streamed runs. + + The one thing the log holds separately is a reminder appended to an + already-recorded message: entries are immutable, so it becomes its own + ``reminder`` entry rather than retroactively editing history. + """ + return self._session + + @property + def messages(self) -> list[Message]: + """The conversation history the model sees. Live and mutable.""" + return self._messages + + @property + def reminders(self) -> list[Callable[[int, int], str | None]]: + """Registered suffix reminder hooks, in registration order.""" + return self._reminders + + # ── run setup ─────────────────────────────────────────────────────────── + + def _record(self, message: Message) -> Message: + """Append ``message`` to the working copy and to the session log. + + The session holds the same object, not a copy. That is deliberate: a + reminder appended to the last user message before a turn is part of + what the model was actually sent, so the log should show it. It also + means a caller that mutates `messages` rewrites history, which is why + `messages` is documented as the loop's own working copy. + """ + self._messages.append(message) + self._session.append_message(message) + return message + + def _begin_run(self, user_message: str) -> None: + self._run_file = setup_logger(self._run_dir) + self._stop_reason = None + self._messages = [] + # A fresh run is a fresh conversation, so the session starts a new root + # rather than hanging it off the previous run's leaf. Without this the + # log claims a continuity the model was never shown, and + # `build_context()` stops matching `messages`. + if self._session.leaf_id is not None: + self._session.move_to(None) + self._record(Message(role="user", content=[TextBlock(text=user_message)])) + + def _begin_ask(self, user_message: str) -> None: + self._stop_reason = None + if self._run_file is None: + self._run_file = setup_logger(self._run_dir) + self._record(Message(role="user", content=[TextBlock(text=user_message)])) + + def _stamp( + self, result: RunResult, run_id: str | None, session_id: str | None + ) -> RunResult: + """Attach run/session ids and make the result the harness's latest. + + Internal, but reached from `data_harness.app.agent` so that a streamed run + is identified exactly like a non-streamed one. Mutates + `self._last_result`; anything reading it afterwards sees the stamped + object. + """ + stamped = dataclasses.replace(result, run_id=run_id, session_id=session_id) + self._last_result = stamped + return stamped + + # ── the loop ──────────────────────────────────────────────────────────── + + def _plan(self) -> Generator[Effect, Any, None]: + """Drive one run, yielding the I/O it needs. The only loop in the library. + + Sets `last_result` before returning, on every exit path, including + abandonment. + """ + if self._run_file is None: + raise RuntimeError("run_file must be initialised before running the loop") + + self._last_result = None + total_usage = Usage() + + try: + yield from self._turns(total_usage) + except HookError as exc: + # A hook broke its contract. The run is over either way, but the + # tokens spent up to here were still billed, so report them rather + # than letting the exception escape with no RunResult at all. + self._last_result = self._build_result( + text="", + status="error", + turns=self._turns_completed, + stop_reason=None, + usage=self._usage_so_far, + error=repr(exc), + ) + raise + except GeneratorExit: + # A consumer that stops mid-stream still spent whatever the + # completed turns cost. Recording it here is the same rule that + # made provider errors keep their usage: never drop tokens the + # provider has already billed just because the caller walked away. + if self._last_result is None: + self._last_result = self._build_result( + text="", + status="error", + turns=self._turns_completed, + stop_reason=None, + usage=self._usage_so_far, + error="RunAbandoned('stream closed before the run finished')", + ) + raise + + def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: + self._turns_completed = 0 + self._usage_so_far = total_usage + + for turn in range(1, self._max_turns + 1): + self._turns_completed = turn + self._compact_if_needed() + self._apply_reminders(turn) + + # A BeforeTurn hook asked to stop, so nothing is spent on this turn. + if self._stop_reason is not None: + self._last_result = self._build_result( + text="", + status="success", + turns=turn - 1, + stop_reason=None, + usage=total_usage, + error=None, + stopped_by=self._stop_reason, + ) + return + visible_tools = [t for t in self._tools if t.visible] + + with time_block() as tb: + request = CallProvider( + system=self._system, + messages=self._messages, + tools=visible_tools, + ) + outcome = _require_answer(request, (yield request)) + + if isinstance(outcome, Failed): + log_error_turn( + turn=turn, + system=self._system, + messages=self._messages, + error=repr(outcome.error), + run_file=self._run_file, + ) + self._last_result = self._build_result( + text="", + status="error", + turns=turn, + stop_reason=None, + usage=total_usage, + error=repr(outcome.error), + ) + return + + response: NormalizedResponse = outcome.value + latency = tb.elapsed_ms + + total_usage = total_usage + Usage( + input_tokens=response.input_tokens, + output_tokens=response.output_tokens, + cache_read_tokens=response.cache_read_tokens, + cache_write_tokens=response.cache_write_tokens, + ) + # Mirrored onto the instance so an abandoned run can still report + # what it spent; the generator's local is gone by then. + self._usage_so_far = total_usage + + self._record(Message(role="assistant", content=response.content)) + + tool_results: list[ToolResultBlock] = [] + + if response.stop_reason == StopReason.TOOL_USE: + # Every tool is dispatched before any result is announced, so a + # streaming consumer sees the same ordering it always has. + finished = yield from self._dispatch(response.content, turn) + tool_results = [block for _, block in finished] + for tool_name, block in finished: + yield ToolFinished(tool_name=tool_name, block=block) + self._record(Message(role="user", content=list(tool_results))) + + self._session.append_turn( + turn=turn, + input_tokens=response.input_tokens, + output_tokens=response.output_tokens, + cache_read_tokens=response.cache_read_tokens, + cache_write_tokens=response.cache_write_tokens, + latency_ms=latency, + stop_reason=response.stop_reason.value, + tool_error_count=sum(1 for r in tool_results if r.is_error), + visible_tools=[t.name for t in visible_tools], + ) + + log_turn( + turn=turn, + system=self._system, + messages=self._messages, + response=response, + tool_results=tool_results, + latency_ms=latency, + run_file=self._run_file, + cache_storage=self._environment.storage_metadata(), + visible_tools=[t.name for t in visible_tools], + tool_error_count=sum(1 for r in tool_results if r.is_error), + all_tools=self._tools, + ) + + for decision in self._hooks.emit( + AfterTurn( + turn=turn, + max_turns=self._max_turns, + input_tokens=total_usage.input_tokens, + output_tokens=total_usage.output_tokens, + tool_results=tool_results, + ) + ): + if isinstance(decision, Stop): + self._stop_reason = decision.reason + + if self._stop_reason is not None: + # An AfterTurn hook stopped the run: a spend cap, a policy + # check, anything deciding on numbers that only exist once the + # turn is accounted for. Whatever the model just said still + # stands as the answer. + self._last_result = self._build_result( + text=_extract_text(response), + status="success", + turns=turn, + stop_reason=response.stop_reason, + usage=total_usage, + stopped_by=self._stop_reason, + ) + return + + if response.stop_reason != StopReason.TOOL_USE: + self._last_result = self._build_result( + text=_extract_text(response), + status="success", + turns=turn, + stop_reason=response.stop_reason, + usage=total_usage, + ) + return + + if turn == self._max_turns: + self._last_result = self._build_result( + text=_extract_text(response), + status="max_turns_exceeded", + turns=turn, + stop_reason=None, + usage=total_usage, + ) + return + + def _dispatch( + self, content: list, turn: int + ) -> Generator[Effect, Any, list[tuple[str, ToolResultBlock]]]: + """Turn the assistant's tool-use blocks into result blocks, in order.""" + tool_map = {t.name: t for t in self._tools} + finished: list[tuple[str, ToolResultBlock]] = [] + + def settle(tub: ToolUseBlock, content: str, is_error: bool) -> None: + """Record one result, giving `AfterToolCall` a chance at it first. + + Every result goes through here, including failures. A redaction + hook that only saw successes would miss the case most likely to + need it: an exception repr can carry a connection string. + """ + block = ToolResultBlock( + tool_use_id=tub.tool_use_id, content=content, is_error=is_error + ) + replacement = self._hooks.first( + AfterToolCall( + turn=turn, + tool_name=tub.tool_name, + tool_input=tub.tool_input, + result=block, + ), + Replace, + ) + if replacement is not None: + block = ToolResultBlock( + tool_use_id=tub.tool_use_id, + content=replacement.content, + is_error=replacement.is_error, + ) + finished.append((tub.tool_name, block)) + + for tub in (b for b in content if isinstance(b, ToolUseBlock)): + spec = tool_map.get(tub.tool_name) + if spec is None or spec.handler is None: + # The hook still sees the attempt: a policy hook needs to know + # about calls to tools that do not exist. + blocked = self._hooks.first( + BeforeToolCall( + turn=turn, tool_name=tub.tool_name, tool_input=tub.tool_input + ), + Block, + ) + if blocked is not None: + settle(tub, blocked.reason, blocked.is_error) + else: + settle(tub, f"Tool not found: {tub.tool_name!r}", True) + continue + + blocked = self._hooks.first( + BeforeToolCall( + turn=turn, tool_name=tub.tool_name, tool_input=tub.tool_input + ), + Block, + ) + if blocked is not None: + settle(tub, blocked.reason, blocked.is_error) + continue + + request = CallTool( + tool_use_id=tub.tool_use_id, + tool_name=tub.tool_name, + handler=spec.handler, + tool_input=tub.tool_input, + ) + answer = _require_answer(request, (yield request)) + + if isinstance(answer, Failed): + settle(tub, repr(answer.error), True) + continue + + try: + output = self._environment.render_tool_output(answer.value) + except Exception as exc: # noqa: BLE001 - surfaced to the model + settle(tub, repr(exc), True) + continue + + settle(tub, output, False) + + return finished + + # ── pure helpers ──────────────────────────────────────────────────────── + + def _build_result( + self, + *, + text: str, + status: Literal["success", "max_turns_exceeded", "error"], + turns: int, + stop_reason: StopReason | None, + usage: Usage, + error: str | None = None, + stopped_by: str | None = None, + ) -> RunResult: + state = self._environment.capture() + return RunResult( + text=text, + status=status, + turns=turns, + run_file=self._run_file, + stop_reason=stop_reason, + usage=usage, + cache_snapshots=state.snapshots, + cache_storage=state.storage, + value=state.value, + charts=state.artifacts, + error=error, + stopped_by=stopped_by, + ) + + def _compact_if_needed(self) -> None: + """Let the compactor shrink the context before the next turn. + + The working copy is rebuilt from the session rather than edited, so + the conversation the model sees stays derived from the log rather than + becoming a second, divergent thing. + + A compactor that fails does not fail the run. Summarising means + calling a model, so a rate limit or a timeout here is ordinary, and + compaction is an optimisation: carrying on with the full context is + strictly better than losing the run. If the context really was too + big, the provider call fails next and that failure is reported where + it belongs. The attempt is recorded either way, so a run that quietly + stopped compacting can be explained afterwards. + """ + if self._compactor is None: + return + try: + compacted = self._compactor(self._session) + except Exception as exc: # noqa: BLE001 - reported, not fatal + self._session.append_custom( + "compaction_failed", + {"turn": self._turns_completed, "error": repr(exc)}, + ) + return + if compacted is not None: + self._messages[:] = self._session.build_context() + + def _apply_reminders(self, turn: int) -> None: + reminder_texts: list[str] = [] + + for hook in self._reminders: + text = hook(turn, self._max_turns) + if text: + reminder_texts.append(text) + + for decision in self._hooks.emit( + BeforeTurn(turn=turn, max_turns=self._max_turns, messages=self._messages) + ): + if isinstance(decision, Reminder): + reminder_texts.append(decision.text) + elif isinstance(decision, Stop): + self._stop_reason = decision.reason + + # Built-in max-turn reminder + if turn == self._max_turns - 1: + reminder_texts.append(_MAX_TURN_REMINDER) + + if not reminder_texts: + return + + reminder_block = TextBlock(text="\n\n".join(reminder_texts)) + + # Append to existing user message or create a new one + if self._messages and self._messages[-1].role == "user": + self._messages[-1].content.append(reminder_block) + # The message it was appended to is already in the log, and entries + # are immutable, so record the reminder separately. A store that + # snapshots on write would otherwise show a prompt the model never + # actually saw. + self._session.append_custom( + "reminder", {"turn": turn, "text": reminder_block.text} + ) + else: + self._record(Message(role="user", content=[reminder_block])) + + +class Harness(_HarnessBase): + """The synchronous driver for the loop. + + Runs the provider call and every tool handler inline on the calling + thread. No event loop is created, so `KeyboardInterrupt` lands promptly and + handlers keep whatever thread affinity they were built with (a `sqlite3` + connection opened at setup keeps working). The only exception is an + ``async def`` tool handler, which necessarily needs a loop. + + `Harness` owns the message list, dispatches tools, applies suffix-only + reminder hooks, and logs every turn to a JSONL file. It is the central + implementation boundary in data-harness: everything above it (`Agent`, + `AgentSession`) is a convenience layer; everything below it + (`ProviderAdapter`, `RunEnvironment`, `ToolSpec`) is a pure dependency. + + The system prompt is never mutated between turns. Reminders, nags, and + dynamic state are always appended to the conversation suffix so the + provider's KV cache is not invalidated. + + For most use cases, prefer `Agent` over constructing `Harness` directly. + Use `Harness` when you need full control over tool wiring, as shown in + ``examples/advanced_wiring.py``. + + Args: + adapter: Synchronous provider adapter that translates provider SDK + objects into harness types. + system: System prompt. Kept byte-identical across all turns. + tools: Full tool list. Invisible tools (``visible=False``) are excluded + from the provider call but can still be dispatched. + max_turns: Hard cap on provider turns before the loop stops and returns + a ``"max_turns_exceeded"`` result. + run_dir: Directory where JSONL logs are written. Created on first run. + environment: Domain services (see `RunEnvironment`). Defaults to + `NullEnvironment`, which renders tool output with `str` and + contributes no run state. + on_code: Approval gate called with interpreter code before it runs. + code_only: When ``True``, interpreter code is echoed, never executed. + """ + + def __init__( + self, + adapter: ProviderAdapter, + system: str, + tools: list[ToolSpec], + max_turns: int = 25, + run_dir: str = "./runs", + environment: RunEnvironment | None = None, + on_code: Callable[[str], object] | None = None, + code_only: bool = False, + session: Session | None = None, + hooks: HookRegistry | None = None, + compactor: Compactor | None = None, + ) -> None: + super().__init__( + system=system, + tools=tools, + max_turns=max_turns, + run_dir=run_dir, + environment=environment, + on_code=on_code, + code_only=code_only, + session=session, + hooks=hooks, + compactor=compactor, + ) + self._adapter = adapter + + def run_result( + self, + user_message: str, + *, + run_id: str | None = None, + session_id: str | None = None, + ) -> RunResult: + """Start a fresh run and return the full `RunResult`. + + Resets message history. Use `ask_result` for follow-up turns on the + same history. + + Args: + user_message: The initial user prompt. + run_id: Optional identifier stamped into the `RunResult`. + session_id: Optional session identifier stamped into the `RunResult`. + + Returns: + A `RunResult` describing the outcome, token usage, and cache state. + """ + self._begin_run(user_message) + return self._stamp(self._drive(), run_id, session_id) + + def ask_result( + self, + user_message: str, + *, + run_id: str | None = None, + session_id: str | None = None, + ) -> RunResult: + """Append a follow-up message and continue the existing run. + + Args: + user_message: The follow-up user prompt. + run_id: Optional identifier stamped into the `RunResult`. + session_id: Optional session identifier stamped into the `RunResult`. + + Returns: + A `RunResult` describing the outcome of this turn sequence. + """ + self._begin_ask(user_message) + return self._stamp(self._drive(), run_id, session_id) + + def run(self, user_message: str) -> str: + """Start a fresh run and return the final text response. + + Args: + user_message: The initial user prompt. + + Returns: + The model's final text response. + + Raises: + MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. + RuntimeError: If the provider raises an exception during the run. + """ + return unwrap_text(self.run_result(user_message)) + + def ask(self, user_message: str) -> str: + """Append a follow-up message and return the final text response. + + Args: + user_message: The follow-up user prompt. + + Returns: + The model's final text response. + + Raises: + MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. + RuntimeError: If the provider raises an exception during the run. + """ + return unwrap_text(self.ask_result(user_message)) + + # ── driver ────────────────────────────────────────────────────────────── + + def _drive(self) -> RunResult: + plan = self._plan() + sent: Answer = None + while True: + try: + effect = plan.send(sent) + except StopIteration: + break + sent = self._perform(effect) + result = self._last_result + if result is None: # pragma: no cover - _plan always sets a result + raise RuntimeError("loop finished without producing a result") + return result + + def _perform(self, effect: Effect) -> Answer: + if isinstance(effect, CallProvider): + try: + return Ok( + self._adapter.chat( + system=effect.system, + messages=effect.messages, + tools=effect.tools, + ) + ) + except Exception as exc: # noqa: BLE001 - reported as a RunResult + return Failed(exc) + if isinstance(effect, CallTool): + try: + return Ok(self._call_tool(effect)) + except Exception as exc: # noqa: BLE001 - surfaced to the model + return Failed(exc) + return None + + def _call_tool(self, call: CallTool) -> Any: + """Invoke one tool handler on the calling thread. + + The overridable seam for changing dispatch (sandboxing, timeouts, + instrumentation). An ``async def`` handler is the one case that needs + an event loop, and gets a short-lived one. + """ + if asyncio.iscoroutinefunction(call.handler): + return run_coroutine_blocking(call.handler(**call.tool_input)) + return call.handler(**call.tool_input) + + +class AsyncHarness(_HarnessBase): + """The asynchronous driver for the loop. + + Awaits the provider and offloads blocking tool handlers to + `asyncio.to_thread`, so a long pandas operation cannot stall the event + loop it shares with a web server. Adds token-level streaming. + + Same arguments as `Harness`, but takes an `AsyncProviderAdapter`. + """ + + def __init__( + self, + adapter: AsyncProviderAdapter, + system: str, + tools: list[ToolSpec], + max_turns: int = 25, + run_dir: str = "./runs", + environment: RunEnvironment | None = None, + on_code: Callable[[str], object] | None = None, + code_only: bool = False, + session: Session | None = None, + hooks: HookRegistry | None = None, + compactor: Compactor | None = None, + ) -> None: + super().__init__( + system=system, + tools=tools, + max_turns=max_turns, + run_dir=run_dir, + environment=environment, + on_code=on_code, + code_only=code_only, + session=session, + hooks=hooks, + compactor=compactor, + ) + self._adapter = adapter + + async def run_result( + self, + user_message: str, + *, + run_id: str | None = None, + session_id: str | None = None, + ) -> RunResult: + """Start a fresh run and return the full `RunResult`.""" + self._begin_run(user_message) + return self._stamp(await self._drain(stream=False), run_id, session_id) + + async def ask_result( + self, + user_message: str, + *, + run_id: str | None = None, + session_id: str | None = None, + ) -> RunResult: + """Append a follow-up message and continue the existing run.""" + self._begin_ask(user_message) + return self._stamp(await self._drain(stream=False), run_id, session_id) + + async def run(self, user_message: str) -> str: + """Start a fresh run and return the final text response. + + Raises: + MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. + RuntimeError: If the provider raises an exception during the run. + """ + return unwrap_text(await self.run_result(user_message)) + + async def ask(self, user_message: str) -> str: + """Append a follow-up message and return the final text response. + + Raises: + MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. + RuntimeError: If the provider raises an exception during the run. + """ + return unwrap_text(await self.ask_result(user_message)) + + async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: + """Stream events for a one-shot run. + + Yields StreamEvent objects following the same protocol as the Claude + Agent SDK. Each provider turn emits message_start, + content_block_start/delta/stop, message_delta, and message_stop events. + After the harness dispatches a tool call a ToolResultEvent is emitted. + The JSONL logger records fully assembled messages, not individual events. + + Once the generator is exhausted, `last_result` holds the `RunResult` + for the run, including token usage and any error. Abandoning the + generator early leaves `last_result` unset. + """ + self._begin_run(user_message) + # `aclosing` matters: a bare `async for` does not close its iterator, + # so a consumer that stops early would leave the driver (and through + # it the provider's HTTP stream) suspended until the GC got to it. + async with aclosing(self._drive(stream=True)) as events: + async for event in events: + yield event + + async def ask_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: + """Stream events for a follow-up turn in a session. + + Once the generator is exhausted, `last_result` holds the `RunResult`. + """ + self._begin_ask(user_message) + async with aclosing(self._drive(stream=True)) as events: + async for event in events: + yield event + + # ── driver ────────────────────────────────────────────────────────────── + + async def _drain(self, *, stream: bool) -> RunResult: + async with aclosing(self._drive(stream=stream)) as events: + async for _ in events: + pass + result = self._last_result + if result is None: # pragma: no cover - _plan always sets a result + raise RuntimeError("loop finished without producing a result") + return result + + async def _drive(self, *, stream: bool) -> AsyncGenerator[StreamEvent, None]: + """Perform the loop's effects. + + Deliberately one generator rather than a stack of them: a generator can + only clean up the resources held in its *own* frame when + `GeneratorExit` arrives, and an intermediate layer would have to be + closed explicitly by the layer above it. Keeping the provider stream, + the plan, and the yields together means one unwind closes everything. + """ + plan = self._plan() + sent: Answer = None + try: + while True: + try: + effect = plan.send(sent) + except StopIteration: + return + sent = None + + if isinstance(effect, CallProvider): + if stream: + events: list[StreamEvent] = [] + provider = self._adapter.stream_events( + system=effect.system, + messages=effect.messages, + tools=effect.tools, + ) + try: + async for evt in provider: + events.append(evt) + yield evt + # Inside the try on purpose: a failure while + # assembling the turn is a provider failure, and + # should land in the RunResult rather than escape + # into the caller's loop. + sent = Ok(accumulate_stream_events(events)) + except Exception as exc: # noqa: BLE001 - reported as a RunResult + sent = Failed(exc) + finally: + # A consumer that stops early unwinds through here + # via GeneratorExit. Close the provider's generator + # now so a real HTTP stream releases its connection + # rather than waiting for the GC. + aclose = getattr(provider, "aclose", None) + if aclose is not None: + await aclose() + else: + try: + sent = Ok( + await self._adapter.chat( + system=effect.system, + messages=effect.messages, + tools=effect.tools, + ) + ) + except Exception as exc: # noqa: BLE001 - reported as a RunResult + sent = Failed(exc) + + elif isinstance(effect, CallTool): + try: + sent = Ok(await self._call_tool(effect)) + except Exception as exc: # noqa: BLE001 - surfaced to the model + sent = Failed(exc) + + elif isinstance(effect, ToolFinished) and stream: + yield ToolResultEvent( + tool_use_id=effect.block.tool_use_id, + tool_name=effect.tool_name, + content=effect.block.content, + is_error=effect.block.is_error, + ) + finally: + plan.close() + + async def _call_tool(self, call: CallTool) -> Any: + """Invoke one tool handler without stalling the event loop. + + The overridable seam for changing dispatch. Blocking handlers go to a + worker thread, so a handler that needs thread affinity must either be + ``async def`` or manage its own resources per call. + """ + if asyncio.iscoroutinefunction(call.handler): + return await call.handler(**call.tool_input) + return await asyncio.to_thread( + functools.partial(call.handler, **call.tool_input) + ) + + +def _extract_text(response: NormalizedResponse) -> str: + return "\n".join(b.text for b in response.content if isinstance(b, TextBlock)) + + +__all__ = [ + "Answer", + "AsyncHarness", + "CallProvider", + "CallTool", + "Effect", + "Failed", + "Harness", + "Ok", + "ToolFinished", + "run_coroutine_blocking", +] diff --git a/data_harness/observe.py b/data_harness/core/observe.py similarity index 100% rename from data_harness/observe.py rename to data_harness/core/observe.py diff --git a/data_harness/result.py b/data_harness/core/result.py similarity index 71% rename from data_harness/result.py rename to data_harness/core/result.py index fbdf61b..e67234f 100644 --- a/data_harness/result.py +++ b/data_harness/core/result.py @@ -6,8 +6,8 @@ from dataclasses import dataclass, field from typing import Any, Literal -from data_harness.artifacts import ChartArtifact -from data_harness.providers.base import StopReason +from data_harness.core.artifacts import ChartArtifact +from data_harness.llm.providers.base import StopReason @dataclass @@ -50,7 +50,9 @@ class CacheStorageInfo: def __post_init__(self) -> None: if self.location not in ("memory", "disk"): - raise ValueError( + from data_harness.core.exceptions import ConfigurationError + + raise ConfigurationError( f"Invalid location: {self.location!r}. Must be 'memory' or 'disk'." ) @@ -73,6 +75,9 @@ class RunResult: cache_storage: Mapping of handle name → `CacheStorageInfo` describing where each handle is stored. error: Exception repr when ``status == "error"``, otherwise ``None``. + stopped_by: Why a hook ended the run early, if one did. A capped run + is otherwise indistinguishable from a model that answered with + nothing. run_id: Optional UUID assigned by `Agent`; ``None`` when using `Harness` directly. session_id: Optional session UUID when the run is part of an @@ -93,6 +98,7 @@ class RunResult: cache_snapshots: dict[str, str] = field(default_factory=dict) cache_storage: dict[str, CacheStorageInfo] = field(default_factory=dict) error: str | None = None + stopped_by: str | None = None run_id: str | None = None session_id: str | None = None value: Any = field(default=None, repr=False) @@ -101,7 +107,7 @@ class RunResult: # --- Rich display for Jupyter / IPython -------------------------------- def _repr_markdown_(self) -> str: parts = [self.text] if self.text else [] - if self.value is not None and not _is_dataframe(self.value): + if self.value is not None and not _renders_itself(self.value): parts.append(f"\n**answer:** `{self.value!r}`") return "\n".join(parts) if parts else f"_(status: {self.status})_" @@ -109,13 +115,13 @@ def _repr_html_(self) -> str: parts: list[str] = [] if self.text: parts.append(f"

{_html.escape(self.text)}

") - if self.value is not None and not _is_dataframe(self.value): + if self.value is not None and not _renders_itself(self.value): parts.append( f"

answer: " f"{_html.escape(repr(self.value))}

" ) - if _is_dataframe(self.value): - parts.append(self.value.to_html()) + if _renders_itself(self.value): + parts.append(self.value._repr_html_()) for chart in self.charts: parts.append(chart._repr_html_()) if not parts: @@ -123,10 +129,31 @@ def _repr_html_(self) -> str: return "\n".join(parts) -def _is_dataframe(value: Any) -> bool: - try: - import pandas as pd +def _renders_itself(value: Any) -> bool: + """Whether ``value`` can render its own HTML, as a DataFrame can. + + Duck-typed on purpose: core must not know what a DataFrame is, and this + way it also works for a polars frame, a styled frame, or anything else + that follows the IPython display protocol. + """ + return hasattr(value, "_repr_html_") + + +def unwrap_text(result: RunResult) -> str: + """Return a successful run's text, or raise the exception its status means. + + The single place that decides which failure becomes which exception, so + `run`, `ask`, and their async twins cannot disagree about it. + + Raises: + MaxTurnsExceeded: If the run hit its turn cap. + ProviderError: If the run failed. `ProviderError` is a `RuntimeError`, + so callers written before the taxonomy still catch it. + """ + from data_harness.core.exceptions import MaxTurnsExceeded, ProviderError - return isinstance(value, pd.DataFrame) - except ImportError: - return False + if result.status == "max_turns_exceeded": + raise MaxTurnsExceeded(result.turns) + if result.status == "error": + raise ProviderError(result.error or "unknown error") + return result.text diff --git a/data_harness/schema.py b/data_harness/core/schema.py similarity index 100% rename from data_harness/schema.py rename to data_harness/core/schema.py diff --git a/data_harness/serialize.py b/data_harness/core/serialize.py similarity index 71% rename from data_harness/serialize.py rename to data_harness/core/serialize.py index 75da671..5bea33d 100644 --- a/data_harness/serialize.py +++ b/data_harness/core/serialize.py @@ -1,6 +1,7 @@ from __future__ import annotations import dataclasses +from collections.abc import Callable from datetime import datetime from enum import Enum from typing import Any @@ -53,40 +54,31 @@ def _convert(obj: Any) -> Any: if isinstance(obj, (list, tuple)): return [_convert(item) for item in obj] - # Try pandas DataFrame - try: - import pandas as pd - - if isinstance(obj, pd.DataFrame): - return { - "type": "dataframe_snapshot", - "shape": list(obj.shape), - "columns": list(obj.columns), - "sample": obj.head(5).to_dict(orient="records"), - } - except ImportError: - pass - - # Try numpy ndarray - try: - import numpy as np - - if isinstance(obj, np.ndarray): - return { - "type": "ndarray_snapshot", - "shape": list(obj.shape), - "dtype": str(obj.dtype), - "sample": obj.flat[:5].tolist(), - } - except ImportError: - pass + for snapshotter in _SNAPSHOTTERS: + snapshot = snapshotter(obj) + if snapshot is not None: + return snapshot return repr(obj) +#: Type-specific snapshotters, tried in registration order. Each returns a +#: payload-free dict for a value it recognises, or ``None`` to pass. +#: +#: The log has to describe a DataFrame without embedding one, but knowing what +#: a DataFrame *is* belongs to the data layer. `data_harness.data` registers +#: its own on import; core on its own falls back to `repr`. +_SNAPSHOTTERS: list[Callable[[Any], dict | None]] = [] + + +def register_snapshotter(snapshotter: Callable[[Any], dict | None]) -> None: + """Teach `to_jsonable` how to describe one more kind of value.""" + _SNAPSHOTTERS.append(snapshotter) + + def _get_text_block_type(): try: - from data_harness.types import TextBlock + from data_harness.llm.types import TextBlock return TextBlock except ImportError: @@ -95,7 +87,7 @@ def _get_text_block_type(): def _get_tool_use_block_type(): try: - from data_harness.types import ToolUseBlock + from data_harness.llm.types import ToolUseBlock return ToolUseBlock except ImportError: @@ -104,7 +96,7 @@ def _get_tool_use_block_type(): def _get_tool_result_block_type(): try: - from data_harness.types import ToolResultBlock + from data_harness.llm.types import ToolResultBlock return ToolResultBlock except ImportError: diff --git a/data_harness/core/session/__init__.py b/data_harness/core/session/__init__.py new file mode 100644 index 0000000..169f370 --- /dev/null +++ b/data_harness/core/session/__init__.py @@ -0,0 +1,55 @@ +"""The session tree: an append-only log of typed entries, shaped like a tree. + +The conversation the model sees is *derived* from the tree by walking root to +leaf. It is never stored, so there is no second copy to disagree with the log. +Compaction is an entry that changes where the walk starts; retrying a turn is +a second child of the same parent; both branches survive. +""" + +from data_harness.core.session.entries import ( + ENTRY_TYPES, + BaseEntry, + CompactionEntry, + CustomEntry, + Entry, + LabelEntry, + LeafEntry, + MessageEntry, + TurnEntry, + new_entry_id, + utc_now, +) +from data_harness.core.session.jsonl import JsonlSessionStore +from data_harness.core.session.session import ( + CustomProjector, + Session, + SessionStats, + leaf_entries, +) +from data_harness.core.session.store import ( + MemorySessionStore, + SessionStore, + SessionStoreError, +) + +__all__ = [ + "ENTRY_TYPES", + "BaseEntry", + "CompactionEntry", + "CustomEntry", + "CustomProjector", + "Entry", + "JsonlSessionStore", + "LabelEntry", + "LeafEntry", + "MemorySessionStore", + "MessageEntry", + "Session", + "SessionStats", + "SessionStore", + "SessionStoreError", + "TurnEntry", + "leaf_entries", + "new_entry_id", + "utc_now", +] diff --git a/data_harness/core/session/entries.py b/data_harness/core/session/entries.py new file mode 100644 index 0000000..39eaaed --- /dev/null +++ b/data_harness/core/session/entries.py @@ -0,0 +1,141 @@ +"""The typed entries a session is made of. + +A session is an append-only log of these, each naming its parent. That makes +the log a *tree*, not a list, and the tree is the session's actual state: the +conversation the model sees is derived by walking root to leaf, never stored. + +Deriving rather than storing is what makes the rest work. Compaction becomes +an entry that changes where the walk starts, not a destructive edit. Retrying +a turn with a different model is a second child of the same parent, and both +survive. Nothing is ever overwritten, so "what did the agent actually see at +turn 7" is a question with an answer. + +The previous design was a write-only JSONL log that re-serialised the entire +message history on every turn: quadratic to write, and impossible to read back +into a runnable state. +""" + +from __future__ import annotations + +import uuid +from dataclasses import dataclass, field +from datetime import datetime, timezone +from typing import Any, Literal + +from data_harness.llm.types import Message + + +def new_entry_id() -> str: + """A short, sortable-enough id. Collisions are checked by the store.""" + return uuid.uuid4().hex[:12] + + +def utc_now() -> str: + return datetime.now(tz=timezone.utc).isoformat() + + +@dataclass(frozen=True) +class BaseEntry: + """Common shape. ``parent_id`` is ``None`` only for the first entry. + + Frozen: an entry is a fact about something that happened. Correcting the + record means appending, not editing. + """ + + id: str = field(default_factory=new_entry_id) + parent_id: str | None = None + timestamp: str = field(default_factory=utc_now) + + +@dataclass(frozen=True) +class MessageEntry(BaseEntry): + """One message in the conversation.""" + + type: Literal["message"] = "message" + message: Message | None = None + + +@dataclass(frozen=True) +class TurnEntry(BaseEntry): + """What one provider turn cost, for accounting and replay analysis.""" + + type: Literal["turn"] = "turn" + turn: int = 0 + input_tokens: int = 0 + output_tokens: int = 0 + cache_read_tokens: int = 0 + cache_write_tokens: int = 0 + latency_ms: float = 0.0 + stop_reason: str | None = None + tool_error_count: int = 0 + visible_tools: list[str] = field(default_factory=list) + + +@dataclass(frozen=True) +class CompactionEntry(BaseEntry): + """History before ``first_kept_entry_id`` is replaced by ``summary``. + + The compacted entries stay in the tree. Only the derived context shrinks, + so a compaction is auditable and reversible: move the leaf back before it + and the full history returns. + """ + + type: Literal["compaction"] = "compaction" + summary: str = "" + first_kept_entry_id: str | None = None + tokens_before: int = 0 + + +@dataclass(frozen=True) +class LabelEntry(BaseEntry): + """A human-readable name for another entry. ``None`` clears it.""" + + type: Literal["label"] = "label" + target_id: str = "" + label: str | None = None + + +@dataclass(frozen=True) +class LeafEntry(BaseEntry): + """Records that the active leaf moved, so branching is itself in the log.""" + + type: Literal["leaf"] = "leaf" + target_id: str | None = None + + +@dataclass(frozen=True) +class CustomEntry(BaseEntry): + """An entry type the core does not know about. + + The extension point that keeps `core` domain-free while letting the data + layer record what matters to it: a cache handle being written, a chart + being produced, a reminder injected into a prompt. + """ + + type: Literal["custom"] = "custom" + custom_type: str = "" + data: dict[str, Any] = field(default_factory=dict) + + +Entry = ( + MessageEntry | TurnEntry | CompactionEntry | LabelEntry | LeafEntry | CustomEntry +) + +#: Discriminator value -> class, for decoding a stored entry. +ENTRY_TYPES: dict[str, type] = { + "message": MessageEntry, + "turn": TurnEntry, + "compaction": CompactionEntry, + "label": LabelEntry, + "leaf": LeafEntry, + "custom": CustomEntry, +} + + +def leaf_after(entry: Entry) -> str | None: + """Where the leaf sits once ``entry`` is appended. + + Every entry advances the leaf to itself, except a `LeafEntry`, whose whole + purpose is to move it somewhere else. + """ + return entry.target_id if isinstance(entry, LeafEntry) else entry.id diff --git a/data_harness/core/session/jsonl.py b/data_harness/core/session/jsonl.py new file mode 100644 index 0000000..90fa488 --- /dev/null +++ b/data_harness/core/session/jsonl.py @@ -0,0 +1,309 @@ +"""A session persisted as one JSONL file: a header, then one line per entry. + +Appending is one line, so writing a session is linear in the number of +entries. The older `runs/*.jsonl` turn log re-serialises the whole message +history every turn, which is quadratic and cannot be read back into anything +runnable. This store is what replaces it; that log is still written alongside +for now, so per-turn write cost is unchanged until it is removed. + +The format is deliberately boring. A session file is greppable, diffable, and +recoverable by hand, which matters more for a debugging artefact than +compactness does. +""" + +from __future__ import annotations + +import copy +import dataclasses +import json +from pathlib import Path +from typing import Any + +from data_harness.core.session.entries import ( + ENTRY_TYPES, + Entry, + LeafEntry, + MessageEntry, + leaf_after, + new_entry_id, + utc_now, +) +from data_harness.core.session.store import SessionStoreError +from data_harness.llm.types import ( + Message, + TextBlock, + ToolResultBlock, + ToolUseBlock, +) + +FORMAT_VERSION = 1 + + +def _encode_block(block: Any) -> dict[str, Any]: + """Encode one content block so it decodes back to an equal object. + + Deliberately not `to_jsonable`: that produces the *log* shape, which + renames fields for readability and snapshots large values away. Lossy is + right for a debugging log and wrong for a store whose whole purpose is to + reconstruct a runnable conversation. + """ + if isinstance(block, ToolUseBlock): + return { + "type": "tool_use", + "tool_use_id": block.tool_use_id, + "tool_name": block.tool_name, + "tool_input": block.tool_input, + } + if isinstance(block, ToolResultBlock): + return { + "type": "tool_result", + "tool_use_id": block.tool_use_id, + "content": block.content, + "is_error": block.is_error, + } + return {"type": "text", "text": getattr(block, "text", str(block))} + + +def _encode_message(message: Message) -> dict[str, Any]: + return { + "role": message.role, + "content": [_encode_block(b) for b in message.content], + } + + +def _decode_message(raw: dict[str, Any]) -> Message: + blocks: list[Any] = [] + for block in raw.get("content", []): + kind = block.get("type") + if kind == "tool_use": + blocks.append( + ToolUseBlock( + tool_use_id=block["tool_use_id"], + tool_name=block["tool_name"], + tool_input=block.get("tool_input", {}), + ) + ) + elif kind == "tool_result": + blocks.append( + ToolResultBlock( + tool_use_id=block["tool_use_id"], + content=block.get("content", ""), + is_error=block.get("is_error", False), + ) + ) + else: + blocks.append(TextBlock(text=block.get("text", ""))) + return Message(role=raw["role"], content=blocks) + + +def encode_entry(entry: Entry) -> dict[str, Any]: + raw = dataclasses.asdict(entry) + if isinstance(entry, MessageEntry) and entry.message is not None: + raw["message"] = _encode_message(entry.message) + return raw + + +def decode_entry(raw: dict[str, Any], *, source: str, line: int) -> Entry: + kind = raw.get("type") + cls = ENTRY_TYPES.get(kind or "") + if cls is None: + raise SessionStoreError( + "invalid_entry", f"{source}:{line} unknown entry type {kind!r}" + ) + fields = {f.name for f in dataclasses.fields(cls)} + kwargs = {k: v for k, v in raw.items() if k in fields and k != "type"} + if cls is MessageEntry and isinstance(kwargs.get("message"), dict): + kwargs["message"] = _decode_message(kwargs["message"]) + try: + return cls(**kwargs) + except TypeError as exc: + raise SessionStoreError( + "invalid_entry", f"{source}:{line} could not be decoded: {exc}" + ) from exc + + +class JsonlSessionStore: + """A `SessionStore` backed by an append-only JSONL file.""" + + def __init__(self, path: str | Path, session_id: str) -> None: + self._path = Path(path) + self._session_id = session_id + self._entries: list[Entry] = [] + self._by_id: dict[str, Entry] = {} + self._leaf_id: str | None = None + #: Set when opening had to drop a torn final entry. + self._truncated = False + + # ── lifecycle ─────────────────────────────────────────────────────────── + + @classmethod + def create(cls, path: str | Path, session_id: str) -> JsonlSessionStore: + """Start a new session file, overwriting any file already there.""" + store = cls(path, session_id) + store._path.parent.mkdir(parents=True, exist_ok=True) + header = { + "type": "session", + "version": FORMAT_VERSION, + "id": session_id, + "timestamp": utc_now(), + } + store._path.write_text(json.dumps(header) + "\n") + return store + + @classmethod + def open(cls, path: str | Path) -> JsonlSessionStore: + """Load an existing session file, replaying its entries. + + Raises: + SessionStoreError: If the file is missing, has no header, or + contains an entry that cannot be decoded. + """ + p = Path(path) + if not p.exists(): + raise SessionStoreError("not_found", f"No session file at {p}") + lines = [line for line in p.read_text().splitlines() if line.strip()] + if not lines: + raise SessionStoreError("invalid_session", f"{p} is empty") + + try: + header = json.loads(lines[0]) + except json.JSONDecodeError as exc: + raise SessionStoreError( + "invalid_session", f"{p}:1 header is not valid JSON" + ) from exc + if header.get("type") != "session": + raise SessionStoreError("invalid_session", f"{p}:1 is not a session header") + if header.get("version") != FORMAT_VERSION: + raise SessionStoreError( + "invalid_session", + f"{p}:1 unsupported session version {header.get('version')!r}", + ) + + store = cls(p, header["id"]) + last_number = len(lines) + for number, line in enumerate(lines[1:], start=2): + try: + raw = json.loads(line) + except json.JSONDecodeError as exc: + if number == last_number: + # A crash mid-append leaves a torn final line. This is an + # append-only recovery log: losing the last entry is the + # cost of the crash, but refusing the file would lose every + # entry before it, which is the opposite of the point. + store._truncated = True + break + raise SessionStoreError( + "invalid_entry", f"{p}:{number} is not valid JSON" + ) from exc + entry = decode_entry(raw, source=str(p), line=number) + store._entries.append(entry) + store._by_id[entry.id] = entry + store._leaf_id = leaf_after(entry) + return store + + @classmethod + def open_or_create(cls, path: str | Path, session_id: str) -> JsonlSessionStore: + """Open an existing session, or start one if the file is not there. + + What a restarting process almost always wants. `create` truncates, + which on a restart path destroys the very session this store exists to + preserve. + """ + if Path(path).exists(): + return cls.open(path) + return cls.create(path, session_id) + + @property + def path(self) -> Path: + return self._path + + @property + def truncated(self) -> bool: + """Whether opening the file had to drop a torn final entry.""" + return self._truncated + + # ── SessionStore ──────────────────────────────────────────────────────── + + @property + def session_id(self) -> str: + return self._session_id + + @property + def leaf_id(self) -> str | None: + return self._leaf_id + + def append(self, entry: Entry) -> None: + """Write ``entry`` to the file and keep a snapshot of it in memory. + + The in-memory copy has to be a copy for the same reason the file does: + the line on disk is frozen at write time, so retaining the caller's + object would let a later mutation make this store's live view disagree + with its own file. Reading the session back would then show something + different from reading it now. + """ + if entry.id in self._by_id: + raise SessionStoreError( + "duplicate_entry", f"Entry {entry.id} already exists" + ) + if entry.parent_id is not None and entry.parent_id not in self._by_id: + raise SessionStoreError( + "missing_parent", + f"Entry {entry.id} names unknown parent {entry.parent_id}", + ) + try: + line = json.dumps(encode_entry(entry)) + except (TypeError, ValueError) as exc: + raise SessionStoreError( + "unserializable_entry", + f"Entry {entry.id} cannot be written as JSON: {exc}", + ) from exc + with self._path.open("a") as handle: + handle.write(line + "\n") + entry = copy.deepcopy(entry) + self._entries.append(entry) + self._by_id[entry.id] = entry + self._leaf_id = leaf_after(entry) + + def get(self, entry_id: str) -> Entry | None: + return self._by_id.get(entry_id) + + def entries(self) -> list[Entry]: + return list(self._entries) + + def set_leaf(self, entry_id: str | None) -> None: + if entry_id is not None and entry_id not in self._by_id: + raise SessionStoreError("not_found", f"Entry {entry_id} not found") + self.append( + LeafEntry( + id=new_entry_id(), + parent_id=self._leaf_id, + timestamp=utc_now(), + target_id=entry_id, + ) + ) + + def path_to_root(self, entry_id: str | None) -> list[Entry]: + if entry_id is None: + return [] + path: list[Entry] = [] + seen: set[str] = set() + current = self._by_id.get(entry_id) + if current is None: + raise SessionStoreError("not_found", f"Entry {entry_id} not found") + while current is not None: + if current.id in seen: + raise SessionStoreError( + "cycle", f"Entry {current.id} is its own ancestor" + ) + seen.add(current.id) + path.append(current) + if current.parent_id is None: + break + parent = self._by_id.get(current.parent_id) + if parent is None: + raise SessionStoreError( + "missing_parent", f"Entry {current.parent_id} not found" + ) + current = parent + path.reverse() + return path diff --git a/data_harness/core/session/session.py b/data_harness/core/session/session.py new file mode 100644 index 0000000..5a8aaf1 --- /dev/null +++ b/data_harness/core/session/session.py @@ -0,0 +1,308 @@ +"""The session tree, and the context derived from it. + +`Session` is the only thing that gives entries meaning. It appends, it moves +the leaf, and it answers the one question the loop actually needs: what +messages should the model see right now. + +That answer is computed on every call from the path root-to-leaf. Nothing +caches it, because a cache is what turns "the history" into a second source of +truth that can disagree with the log. +""" + +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Any, Callable + +from data_harness.core.session.entries import ( + CompactionEntry, + CustomEntry, + Entry, + LabelEntry, + LeafEntry, + MessageEntry, + TurnEntry, + new_entry_id, + utc_now, +) +from data_harness.core.session.store import MemorySessionStore, SessionStore +from data_harness.llm.types import ( + Message, + TextBlock, + ToolResultBlock, + ToolUseBlock, +) + +#: Turns a custom entry into messages for the model, or nothing. +#: +#: Custom entries are invisible to the model by default: the data layer +#: records a cache write so a human can trace it, not so the model re-reads it. +#: A projector opts one in. +CustomProjector = Callable[[CustomEntry], list[Message]] + + +@dataclass +class SessionStats: + """Totals across every entry, including branches no longer on the path.""" + + entries: int = 0 + messages: int = 0 + turns: int = 0 + input_tokens: int = 0 + output_tokens: int = 0 + cache_read_tokens: int = 0 + cache_write_tokens: int = 0 + + +@dataclass +class Session: + """An append-only tree of entries, plus the context it derives. + + Args: + store: Where entries live. Defaults to memory. + projectors: Custom entry type -> how to render it for the model. + """ + + store: SessionStore = field(default_factory=MemorySessionStore) + projectors: dict[str, CustomProjector] = field(default_factory=dict) + + # ── appending ─────────────────────────────────────────────────────────── + + def _append(self, make: Callable[[str, str | None, str], Entry]) -> str: + entry = make(new_entry_id(), self.store.leaf_id, utc_now()) + self.store.append(entry) + return entry.id + + def append_message(self, message: Message) -> str: + return self._append( + lambda i, p, t: MessageEntry( + id=i, parent_id=p, timestamp=t, message=message + ) + ) + + def append_turn(self, **fields: Any) -> str: + return self._append( + lambda i, p, t: TurnEntry(id=i, parent_id=p, timestamp=t, **fields) + ) + + def append_custom(self, custom_type: str, data: dict[str, Any]) -> str: + return self._append( + lambda i, p, t: CustomEntry( + id=i, parent_id=p, timestamp=t, custom_type=custom_type, data=data + ) + ) + + def append_compaction( + self, summary: str, first_kept_entry_id: str | None, tokens_before: int + ) -> str: + """Summarise history before ``first_kept_entry_id``. + + The kept id must be on the current path. Anything else is a bug in the + caller: an id from an abandoned branch, or one recorded before a fork, + would silently amnesia the agent rather than fail. + + Raises: + SessionStoreError: If the kept id is not on the current path. + """ + if first_kept_entry_id is not None: + from data_harness.core.session.store import SessionStoreError + + if first_kept_entry_id not in {e.id for e in self.branch()}: + raise SessionStoreError( + "invalid_cut", + f"Compaction keeps from {first_kept_entry_id}, which is not " + "on the current path", + ) + return self._append( + lambda i, p, t: CompactionEntry( + id=i, + parent_id=p, + timestamp=t, + summary=summary, + first_kept_entry_id=first_kept_entry_id, + tokens_before=tokens_before, + ) + ) + + def label(self, target_id: str, label: str | None) -> str: + """Name an entry, so a branch point can be found again by a human.""" + if self.store.get(target_id) is None: + from data_harness.core.session.store import SessionStoreError + + raise SessionStoreError("not_found", f"Entry {target_id} not found") + return self._append( + lambda i, p, t: LabelEntry( + id=i, parent_id=p, timestamp=t, target_id=target_id, label=label + ) + ) + + # ── navigating ────────────────────────────────────────────────────────── + + @property + def leaf_id(self) -> str | None: + return self.store.leaf_id + + def move_to(self, entry_id: str | None) -> None: + """Make ``entry_id`` the active leaf. The next append branches from it. + + The old branch is untouched and still reachable: this is how a turn is + retried with a different model without losing what the first attempt + did. + """ + self.store.set_leaf(entry_id) + + def branch(self, from_id: str | None = None) -> list[Entry]: + """Root-first path to ``from_id``, defaulting to the active leaf.""" + return self.store.path_to_root( + from_id if from_id is not None else self.store.leaf_id + ) + + def labels(self) -> dict[str, str]: + """Entry id -> current label. Later labels win; ``None`` clears.""" + found: dict[str, str] = {} + for entry in self.store.entries(): + if isinstance(entry, LabelEntry): + if entry.label: + found[entry.target_id] = entry.label + else: + found.pop(entry.target_id, None) + return found + + # ── deriving the context ──────────────────────────────────────────────── + + def context_entries(self, from_id: str | None = None) -> list[Entry]: + """The entries a run should replay, after applying compaction. + + Everything before the newest compaction on the path is dropped, except + the tail it chose to keep. The dropped entries stay in the tree. + """ + path = self.branch(from_id) + + newest_compaction = None + compaction_index = -1 + for index, entry in enumerate(path): + if isinstance(entry, CompactionEntry): + newest_compaction = entry + compaction_index = index + if newest_compaction is None: + return path + + kept: list[Entry] = [newest_compaction] + if newest_compaction.first_kept_entry_id is not None: + keeping = False + for entry in path[:compaction_index]: + if entry.id == newest_compaction.first_kept_entry_id: + keeping = True + if not keeping: + continue + # An older compaction inside the kept tail has already been + # subsumed by this one. Replaying it would emit a second, + # older summary *after* the newer one and resurrect the very + # entries it had dropped. + if isinstance(entry, CompactionEntry): + continue + kept.append(entry) + kept.extend(path[compaction_index + 1 :]) + return kept + + def build_context(self, from_id: str | None = None) -> list[Message]: + """The messages the model should see. Derived, never stored. + + Always something a provider will accept: a tool call with no result, + or a result with no call, is dropped. Both are reachable without any + bug here. A run killed between issuing a tool call and recording its + result leaves the call orphaned on disk, and resuming would otherwise + send a transcript the provider rejects outright, breaking resume for + exactly the sessions most worth resuming. + """ + messages: list[Message] = [] + for entry in self.context_entries(from_id): + if isinstance(entry, MessageEntry) and entry.message is not None: + messages.append(entry.message) + elif isinstance(entry, CompactionEntry): + messages.append( + Message( + role="user", + content=[ + TextBlock( + text=( + "Summary of the earlier conversation:\n" + f"{entry.summary}" + ) + ) + ], + ) + ) + elif isinstance(entry, CustomEntry): + projector = self.projectors.get(entry.custom_type) + if projector is not None: + messages.extend(projector(entry)) + return _drop_unpaired_tool_blocks(messages) + + # ── reporting ─────────────────────────────────────────────────────────── + + def stats(self) -> SessionStats: + stats = SessionStats() + for entry in self.store.entries(): + stats.entries += 1 + if isinstance(entry, MessageEntry): + stats.messages += 1 + elif isinstance(entry, TurnEntry): + stats.turns += 1 + stats.input_tokens += entry.input_tokens + stats.output_tokens += entry.output_tokens + stats.cache_read_tokens += entry.cache_read_tokens + stats.cache_write_tokens += entry.cache_write_tokens + return stats + + def custom_entries(self, custom_type: str) -> list[CustomEntry]: + """Every custom entry of one type, across all branches.""" + return [ + entry + for entry in self.store.entries() + if isinstance(entry, CustomEntry) and entry.custom_type == custom_type + ] + + +def _drop_unpaired_tool_blocks(messages: list[Message]) -> list[Message]: + """Remove tool calls with no result, and results with no call. + + Providers reject either. Both arise legitimately: a run killed between + issuing a call and recording its result orphans the call, and a compaction + cut between the two orphans the result. A message left with no content is + dropped too, since an empty message is itself invalid. + """ + answered = { + block.tool_use_id + for message in messages + for block in message.content + if isinstance(block, ToolResultBlock) + } + requested = { + block.tool_use_id + for message in messages + for block in message.content + if isinstance(block, ToolUseBlock) + } + + repaired: list[Message] = [] + for message in messages: + content = [ + block + for block in message.content + if not ( + isinstance(block, ToolUseBlock) and block.tool_use_id not in answered + ) + and not ( + isinstance(block, ToolResultBlock) + and block.tool_use_id not in requested + ) + ] + if content: + repaired.append(Message(role=message.role, content=content)) + return repaired + + +def leaf_entries(session: Session) -> list[LeafEntry]: + """Every recorded leaf move, oldest first. The session's branch history.""" + return [e for e in session.store.entries() if isinstance(e, LeafEntry)] diff --git a/data_harness/core/session/store.py b/data_harness/core/session/store.py new file mode 100644 index 0000000..19fa7f1 --- /dev/null +++ b/data_harness/core/session/store.py @@ -0,0 +1,149 @@ +"""Where a session's entries live. + +The store is deliberately dumb: append, look up, walk parents. All the meaning +is in `Session`. That split is what lets an application swap Postgres in +without the session logic knowing, and it is why the loop no longer opens +files itself. +""" + +from __future__ import annotations + +import copy +from typing import Protocol, runtime_checkable + +from data_harness.core.exceptions import DataHarnessError +from data_harness.core.session.entries import Entry, leaf_after + + +class SessionStoreError(DataHarnessError): + """A session cannot be read, or is internally inconsistent. + + Codes: ``not_found``, ``invalid_session``, ``invalid_entry``, + ``duplicate_entry``, ``missing_parent``, ``cycle``, ``invalid_cut``, + ``unserializable_entry``. + """ + + def __init__(self, code: str, message: str) -> None: + super().__init__(message, code=code) + + +@runtime_checkable +class SessionStore(Protocol): + """Append-only storage for one session's entries.""" + + @property + def session_id(self) -> str: ... + + @property + def leaf_id(self) -> str | None: + """The entry the next append will hang from. ``None`` for an empty session.""" + ... + + def append(self, entry: Entry) -> None: + """Persist ``entry`` and advance the leaf. Ids must be unique.""" + ... + + def get(self, entry_id: str) -> Entry | None: ... + + def entries(self) -> list[Entry]: + """Every entry, in the order appended. Includes abandoned branches.""" + ... + + def set_leaf(self, entry_id: str | None) -> None: + """Move the active leaf, recording the move as an entry.""" + ... + + def path_to_root(self, entry_id: str | None) -> list[Entry]: + """Root-first path to ``entry_id``, following ``parent_id``.""" + ... + + +class MemorySessionStore: + """In-process store. The default, and what tests use. + + Also the honest default for a server that keeps sessions in memory: it + makes the durability boundary explicit rather than implying a session + survives a restart when it does not. + """ + + def __init__(self, session_id: str = "session") -> None: + self._session_id = session_id + self._entries: list[Entry] = [] + self._by_id: dict[str, Entry] = {} + self._leaf_id: str | None = None + + @property + def session_id(self) -> str: + return self._session_id + + @property + def leaf_id(self) -> str | None: + return self._leaf_id + + def append(self, entry: Entry) -> None: + """Store a snapshot of ``entry``. + + Copied, not referenced. A caller holding the same `Message` object can + otherwise keep editing it and silently rewrite history, and the JSONL + store (which serialises on write) would then disagree with this one + about what was sent. + """ + entry = copy.deepcopy(entry) + if entry.id in self._by_id: + raise SessionStoreError( + "duplicate_entry", f"Entry {entry.id} already exists" + ) + if entry.parent_id is not None and entry.parent_id not in self._by_id: + raise SessionStoreError( + "missing_parent", + f"Entry {entry.id} names unknown parent {entry.parent_id}", + ) + self._entries.append(entry) + self._by_id[entry.id] = entry + self._leaf_id = leaf_after(entry) + + def get(self, entry_id: str) -> Entry | None: + return self._by_id.get(entry_id) + + def entries(self) -> list[Entry]: + return list(self._entries) + + def set_leaf(self, entry_id: str | None) -> None: + from data_harness.core.session.entries import LeafEntry, new_entry_id, utc_now + + if entry_id is not None and entry_id not in self._by_id: + raise SessionStoreError("not_found", f"Entry {entry_id} not found") + self.append( + LeafEntry( + id=new_entry_id(), + parent_id=self._leaf_id, + timestamp=utc_now(), + target_id=entry_id, + ) + ) + + def path_to_root(self, entry_id: str | None) -> list[Entry]: + if entry_id is None: + return [] + path: list[Entry] = [] + seen: set[str] = set() + current = self._by_id.get(entry_id) + if current is None: + raise SessionStoreError("not_found", f"Entry {entry_id} not found") + while current is not None: + if current.id in seen: + raise SessionStoreError( + "cycle", f"Entry {current.id} is its own ancestor" + ) + seen.add(current.id) + path.append(current) + if current.parent_id is None: + break + parent = self._by_id.get(current.parent_id) + if parent is None: + raise SessionStoreError( + "missing_parent", f"Entry {current.parent_id} not found" + ) + current = parent + path.reverse() + return path diff --git a/data_harness/data/__init__.py b/data_harness/data/__init__.py new file mode 100644 index 0000000..8347797 --- /dev/null +++ b/data_harness/data/__init__.py @@ -0,0 +1,46 @@ +"""The data domain: what makes this a *data* harness rather than a generic one. + +Owns the `SessionCache` and its handle/snapshot discipline, the sandboxed +Python interpreter, SQL and connector tools, MCP, and the DataFrame-aware +formatting that keeps large payloads out of the transcript. + +May import `llm` and `core`. May not import `app`. +""" + +from data_harness.core.serialize import register_snapshotter + + +def _dataframe_snapshot(obj: object) -> dict | None: + try: + import pandas as pd + except ImportError: # pragma: no cover - exercised via the install matrix + return None + if not isinstance(obj, pd.DataFrame): + return None + return { + "type": "dataframe_snapshot", + "shape": list(obj.shape), + "columns": list(obj.columns), + "sample": obj.head(5).to_dict(orient="records"), + } + + +def _ndarray_snapshot(obj: object) -> dict | None: + try: + import numpy as np + except ImportError: # pragma: no cover - exercised via the install matrix + return None + if not isinstance(obj, np.ndarray): + return None + return { + "type": "ndarray_snapshot", + "shape": list(obj.shape), + "dtype": str(obj.dtype), + "sample": obj.flat[:5].tolist(), + } + + +# Registered on import: the run log has to describe a frame without embedding +# one, but only this layer knows what a frame is. +register_snapshotter(_dataframe_snapshot) +register_snapshotter(_ndarray_snapshot) diff --git a/data_harness/_sandbox_runner.py b/data_harness/data/_sandbox_runner.py similarity index 94% rename from data_harness/_sandbox_runner.py rename to data_harness/data/_sandbox_runner.py index 00849e4..794cf50 100644 --- a/data_harness/_sandbox_runner.py +++ b/data_harness/data/_sandbox_runner.py @@ -1,10 +1,10 @@ """Child-process entry point for the subprocess sandbox. -Invoked as ``python -m data_harness._sandbox_runner ``. Reads a +Invoked as ``python -m data_harness.data._sandbox_runner ``. Reads a pickled payload (code, handle values, allowlist, limits), executes the code under resource limits with networking disabled, and writes a pickled result. -It deliberately reuses :func:`data_harness.tools.interpreter.execute_namespace` +It deliberately reuses :func:`data_harness.data.tools.interpreter.execute_namespace` so the sandbox shares the exact same security boundary as the in-process path. """ @@ -68,7 +68,7 @@ def _capture_charts(artifacts_dir: str) -> list[tuple[str, str | None]]: def _run(payload: dict) -> dict: - from data_harness.tools.interpreter import ( + from data_harness.data.tools.interpreter import ( _DEFAULT_ALLOWLIST, _SENTINEL, PythonInterpreterError, diff --git a/data_harness/cache.py b/data_harness/data/cache.py similarity index 99% rename from data_harness/cache.py rename to data_harness/data/cache.py index 36e05c5..274e74b 100644 --- a/data_harness/cache.py +++ b/data_harness/data/cache.py @@ -10,7 +10,7 @@ from pathlib import Path from typing import Any -from data_harness.artifacts import ChartArtifact +from data_harness.core.artifacts import ChartArtifact _VALID_IDENTIFIER = re.compile(r"^[a-zA-Z_][a-zA-Z0-9_]*$") diff --git a/data_harness/data/environment.py b/data_harness/data/environment.py new file mode 100644 index 0000000..1521d3f --- /dev/null +++ b/data_harness/data/environment.py @@ -0,0 +1,44 @@ +"""The data domain's answer to `RunEnvironment`, backed by a `SessionCache`. + +This is the seam where the handle/snapshot discipline plugs into the loop. A +tool that returns a DataFrame gets it stored under a handle and reports a +compact snapshot; the model works against the handle name and the raw frame +never enters the transcript. +""" + +from __future__ import annotations + +from typing import Any, Literal, cast + +from data_harness.core.environment import RunState +from data_harness.core.result import CacheStorageInfo +from data_harness.data.cache import SessionCache +from data_harness.data.format import format_tool_output + + +class CacheEnvironment: + """A `RunEnvironment` whose state is a `SessionCache`.""" + + def __init__(self, cache: SessionCache | None = None) -> None: + self.cache = cache if cache is not None else SessionCache() + + def render_tool_output(self, value: Any) -> str: + """Inline small values; spill large ones into the cache as handles.""" + return format_tool_output(value, cache=self.cache) + + def capture(self) -> RunState: + return RunState( + snapshots=self.cache.list_handles(), + storage={ + name: CacheStorageInfo( + location=cast(Literal["memory", "disk"], meta["location"]), + storage_type=meta["storage_type"], + ) + for name, meta in self.cache.storage_metadata().items() + }, + value=self.cache.get_answer(), + artifacts=self.cache.list_charts(), + ) + + def storage_metadata(self) -> dict[str, dict[str, str]]: + return self.cache.storage_metadata() diff --git a/data_harness/exec_cache.py b/data_harness/data/exec_cache.py similarity index 98% rename from data_harness/exec_cache.py rename to data_harness/data/exec_cache.py index cdf8c8d..273443d 100644 --- a/data_harness/exec_cache.py +++ b/data_harness/data/exec_cache.py @@ -15,7 +15,7 @@ from dataclasses import dataclass, field from pathlib import Path -from data_harness.types import Message, ToolResultBlock, ToolUseBlock +from data_harness.llm.types import Message, ToolResultBlock, ToolUseBlock _REPLAYABLE_TOOLS = ("python_interpreter", "sql_query") diff --git a/data_harness/format.py b/data_harness/data/format.py similarity index 98% rename from data_harness/format.py rename to data_harness/data/format.py index 9b13675..285bed9 100644 --- a/data_harness/format.py +++ b/data_harness/data/format.py @@ -4,7 +4,7 @@ from typing import TYPE_CHECKING, Any if TYPE_CHECKING: - from data_harness.cache import SessionCache + from data_harness.data.cache import SessionCache _INLINE_STR_MAX = 500 _INLINE_JSON_MAX = 2000 diff --git a/data_harness/data/harness.py b/data_harness/data/harness.py new file mode 100644 index 0000000..040a878 --- /dev/null +++ b/data_harness/data/harness.py @@ -0,0 +1,148 @@ +"""The data-flavoured harness: `core`'s loop wired to a `SessionCache`. + +`data_harness.core.loop` is domain-free — it takes a `RunEnvironment` and has +no idea what a DataFrame is. Almost nobody wants that directly. These two +classes are the same loop with the data domain plugged in, and they are what +`Agent` builds and what `data_harness.loop` resolves to. + +The only difference from the core classes is the constructor: ``cache=`` in +place of ``environment=``, defaulting to a fresh `SessionCache`. +""" + +from __future__ import annotations + +from collections.abc import Callable + +from data_harness.core.compaction import Compactor +from data_harness.core.hooks import HookRegistry +from data_harness.core.loop import ( + Answer, + CallProvider, + CallTool, + Effect, + Failed, + Ok, + ToolFinished, + run_coroutine_blocking, +) +from data_harness.core.loop import ( + AsyncHarness as _CoreAsyncHarness, +) +from data_harness.core.loop import ( + Harness as _CoreHarness, +) +from data_harness.core.session import Session +from data_harness.data.cache import SessionCache +from data_harness.data.environment import CacheEnvironment +from data_harness.llm.providers.base import AsyncProviderAdapter, ProviderAdapter +from data_harness.llm.types import ToolSpec + + +def _environment_for(cache: SessionCache | None) -> CacheEnvironment: + return CacheEnvironment(cache) + + +class Harness(_CoreHarness): + """Synchronous harness over a `SessionCache`. See `core.loop.Harness`. + + Args: + adapter: Synchronous provider adapter. + system: System prompt. Kept byte-identical across all turns. + tools: Full tool list. + max_turns: Hard cap on provider turns. + run_dir: Directory where JSONL logs are written. + cache: Shared `SessionCache`. A fresh one is created if ``None``. + session: Durable record of the run. Pass one opened from storage + to resume a conversation across processes. + hooks: Pre-built `HookRegistry`, so an application can define its + policy once and reuse it across harnesses. + on_code: Approval gate called with interpreter code before it runs. + code_only: When ``True``, interpreter code is echoed, never executed. + """ + + def __init__( + self, + adapter: ProviderAdapter, + system: str, + tools: list[ToolSpec], + max_turns: int = 25, + run_dir: str = "./runs", + cache: SessionCache | None = None, + on_code: Callable[[str], object] | None = None, + code_only: bool = False, + session: Session | None = None, + hooks: HookRegistry | None = None, + compactor: Compactor | None = None, + ) -> None: + super().__init__( + adapter=adapter, + system=system, + tools=tools, + max_turns=max_turns, + run_dir=run_dir, + environment=_environment_for(cache), + on_code=on_code, + code_only=code_only, + session=session, + hooks=hooks, + compactor=compactor, + ) + + @property + def cache(self) -> SessionCache: + """The `SessionCache` backing tool results and handles.""" + return self._environment.cache + + +class AsyncHarness(_CoreAsyncHarness): + """Async harness over a `SessionCache`. See `core.loop.AsyncHarness`. + + Same arguments as `Harness`, but takes an `AsyncProviderAdapter`. + """ + + def __init__( + self, + adapter: AsyncProviderAdapter, + system: str, + tools: list[ToolSpec], + max_turns: int = 25, + run_dir: str = "./runs", + cache: SessionCache | None = None, + on_code: Callable[[str], object] | None = None, + code_only: bool = False, + session: Session | None = None, + hooks: HookRegistry | None = None, + compactor: Compactor | None = None, + ) -> None: + super().__init__( + adapter=adapter, + system=system, + tools=tools, + max_turns=max_turns, + run_dir=run_dir, + environment=_environment_for(cache), + on_code=on_code, + code_only=code_only, + session=session, + hooks=hooks, + compactor=compactor, + ) + + @property + def cache(self) -> SessionCache: + """The `SessionCache` backing tool results and handles.""" + return self._environment.cache + + +__all__ = [ + "Answer", + "AsyncHarness", + "CallProvider", + "CallTool", + "Effect", + "Failed", + "Harness", + "Ok", + "ToolFinished", + "run_coroutine_blocking", +] diff --git a/data_harness/io.py b/data_harness/data/io.py similarity index 100% rename from data_harness/io.py rename to data_harness/data/io.py diff --git a/data_harness/mcp.py b/data_harness/data/mcp.py similarity index 98% rename from data_harness/mcp.py rename to data_harness/data/mcp.py index fceda2b..a33f388 100644 --- a/data_harness/mcp.py +++ b/data_harness/data/mcp.py @@ -18,7 +18,7 @@ from dataclasses import dataclass, field from typing import Any -from data_harness.types import ToolAnnotations, ToolSpec +from data_harness.llm.types import ToolAnnotations, ToolSpec @dataclass diff --git a/data_harness/data/tools/__init__.py b/data_harness/data/tools/__init__.py new file mode 100644 index 0000000..c66ef88 --- /dev/null +++ b/data_harness/data/tools/__init__.py @@ -0,0 +1 @@ +"""Tool implementations for the data domain.""" diff --git a/data_harness/tools/connectors.py b/data_harness/data/tools/connectors.py similarity index 96% rename from data_harness/tools/connectors.py rename to data_harness/data/tools/connectors.py index d7c1867..be8cd66 100644 --- a/data_harness/tools/connectors.py +++ b/data_harness/data/tools/connectors.py @@ -2,9 +2,9 @@ from typing import Any, Callable -from data_harness.cache import SessionCache -from data_harness.format import format_tool_output -from data_harness.types import ToolSpec +from data_harness.data.cache import SessionCache +from data_harness.data.format import format_tool_output +from data_harness.llm.types import ToolSpec class ConnectorRegistry: diff --git a/data_harness/tools/interpreter.py b/data_harness/data/tools/interpreter.py similarity index 95% rename from data_harness/tools/interpreter.py rename to data_harness/data/tools/interpreter.py index c9553a1..43263df 100644 --- a/data_harness/tools/interpreter.py +++ b/data_harness/data/tools/interpreter.py @@ -12,9 +12,10 @@ from pathlib import Path from typing import Any -from data_harness.artifacts import ChartArtifact -from data_harness.cache import SessionCache -from data_harness.types import ToolAnnotations, ToolSpec +from data_harness.core.artifacts import ChartArtifact +from data_harness.core.exceptions import DataHarnessError +from data_harness.data.cache import SessionCache +from data_harness.llm.types import ToolAnnotations, ToolSpec # Force a headless backend so plotting works without a display. Set before any # user code imports matplotlib. @@ -63,8 +64,15 @@ _SENTINEL = object() -class PythonInterpreterError(Exception): - """Raised by PythonInterpreter.run() on any execution failure.""" +class PythonInterpreterError(DataHarnessError): + """The model's code ran and failed, or was refused by the allow-list. + + The model's problem to fix, so it is handed back as a tool result rather + than ending the run. Distinct from `ExecutionError`, which means the code + never got to run at all. + """ + + code = "interpreter_error" def __repr__(self) -> str: return str(self) diff --git a/data_harness/tools/planner.py b/data_harness/data/tools/planner.py similarity index 98% rename from data_harness/tools/planner.py rename to data_harness/data/tools/planner.py index 0fc6198..6d08724 100644 --- a/data_harness/tools/planner.py +++ b/data_harness/data/tools/planner.py @@ -3,7 +3,7 @@ import uuid from typing import Any -from data_harness.types import ToolSpec +from data_harness.llm.types import ToolSpec class Planner: diff --git a/data_harness/tools/sandbox.py b/data_harness/data/tools/sandbox.py similarity index 91% rename from data_harness/tools/sandbox.py rename to data_harness/data/tools/sandbox.py index 4d79d9f..0e237f3 100644 --- a/data_harness/tools/sandbox.py +++ b/data_harness/data/tools/sandbox.py @@ -21,15 +21,16 @@ from pathlib import Path from typing import Any -from data_harness.artifacts import ChartArtifact -from data_harness.cache import SessionCache -from data_harness.tools.interpreter import ( +from data_harness.core.artifacts import ChartArtifact +from data_harness.core.exceptions import ExecutionError +from data_harness.data.cache import SessionCache +from data_harness.data.tools.interpreter import ( _DEFAULT_ALLOWLIST, _EMPTY_OUTPUT_GUIDANCE, PythonInterpreterError, _interpreter_description, ) -from data_harness.types import ToolAnnotations, ToolSpec +from data_harness.llm.types import ToolAnnotations, ToolSpec _SANDBOX_ANNOTATIONS = ToolAnnotations( title="Python Interpreter (sandboxed)", @@ -101,7 +102,7 @@ def run(self, code: str) -> str: [ sys.executable, "-m", - "data_harness._sandbox_runner", + "data_harness.data._sandbox_runner", str(payload_path), str(result_path), ], @@ -109,13 +110,15 @@ def run(self, code: str) -> str: timeout=self._timeout, ) except subprocess.TimeoutExpired: - raise PythonInterpreterError( + # The code never finished, so this is the environment failing + # rather than the model's code being wrong. + raise ExecutionError( f"Execution timed out after {self._timeout}s in the sandbox." ) from None if proc.returncode != 0 or not result_path.exists(): stderr = proc.stderr.decode("utf-8", "replace")[-800:] - raise PythonInterpreterError( + raise ExecutionError( f"Sandbox process failed (exit {proc.returncode}). {stderr}" ) diff --git a/data_harness/tools/sql.py b/data_harness/data/tools/sql.py similarity index 95% rename from data_harness/tools/sql.py rename to data_harness/data/tools/sql.py index b7e1233..189e5ec 100644 --- a/data_harness/tools/sql.py +++ b/data_harness/data/tools/sql.py @@ -11,11 +11,11 @@ from typing import TYPE_CHECKING -from data_harness.format import format_tool_output -from data_harness.types import ToolAnnotations, ToolSpec +from data_harness.data.format import format_tool_output +from data_harness.llm.types import ToolAnnotations, ToolSpec if TYPE_CHECKING: - from data_harness.cache import SessionCache + from data_harness.data.cache import SessionCache def _is_dataframe(value: object) -> bool: diff --git a/data_harness/tools/subagent.py b/data_harness/data/tools/subagent.py similarity index 82% rename from data_harness/tools/subagent.py rename to data_harness/data/tools/subagent.py index f0b449b..b7961ad 100644 --- a/data_harness/tools/subagent.py +++ b/data_harness/data/tools/subagent.py @@ -4,11 +4,11 @@ import dataclasses from typing import Callable -from data_harness.cache import SessionCache -from data_harness.providers.base import ProviderAdapter -from data_harness.tools.interpreter import PythonInterpreter -from data_harness.tools.variables import make_list_variables_spec -from data_harness.types import ToolSpec +from data_harness.data.cache import SessionCache +from data_harness.data.tools.interpreter import PythonInterpreter +from data_harness.data.tools.variables import make_list_variables_spec +from data_harness.llm.providers.base import AsyncProviderAdapter, ProviderAdapter +from data_harness.llm.types import ToolSpec _SUBAGENT_TOOL_NAME = "subagent" @@ -25,7 +25,7 @@ def make_subagent_spec( - adapter_factory: Callable[[], ProviderAdapter], + adapter_factory: Callable[[], ProviderAdapter | AsyncProviderAdapter], parent_tools: list[ToolSpec], parent_cache: SessionCache, run_dir: str = "./runs", @@ -34,6 +34,12 @@ def make_subagent_spec( ) -> ToolSpec: """Create a subagent tool with an explicit cache boundary. + ``adapter_factory`` may return either a `ProviderAdapter` or an + `AsyncProviderAdapter`. Whichever it returns picks the matching driver, so + a sync adapter runs on the calling thread rather than being pushed through + an event loop. The tool handler itself stays synchronous, so it works + identically whether the parent is an `Agent` or an `AsyncAgent`. + If parent_tools include cache-bound wrappers such as ConnectorRegistry wrapped specs, pass make_sub_tools so those handlers can be rebuilt against the subagent cache. The fallback path only copies cache-independent tools @@ -45,7 +51,11 @@ def subagent( input_handles: list[str] | None = None, output_policy: str = "text_only", ) -> str: - from data_harness.loop import Harness + from data_harness.data.harness import ( + AsyncHarness, + Harness, + run_coroutine_blocking, + ) # Validate input_handles against parent cache if input_handles: @@ -94,19 +104,24 @@ def subagent( handles_str = str(input_handles) if input_handles else "none" system = _WORKER_SYSTEM_TEMPLATE.format(task=task, input_handles=handles_str) - # Spawn fresh adapter + # Spawn fresh adapter. A sync adapter gets the sync driver so the + # subagent keeps the parent's threading semantics; only an async + # adapter needs a loop spun up to drive it from this sync handler. sub_adapter = adapter_factory() - - sub_harness = Harness( - adapter=sub_adapter, - system=system, - tools=sub_tools, - run_dir=run_dir, - cache=sub_cache, - ) + harness_kwargs = { + "system": system, + "tools": sub_tools, + "run_dir": run_dir, + "cache": sub_cache, + } try: - final_text = sub_harness.run(task) + if isinstance(sub_adapter, AsyncProviderAdapter): + final_text = run_coroutine_blocking( + AsyncHarness(adapter=sub_adapter, **harness_kwargs).run(task) + ) + else: + final_text = Harness(adapter=sub_adapter, **harness_kwargs).run(task) except Exception as exc: return f"Error: subagent failed: {type(exc).__name__}: {exc}" diff --git a/data_harness/tools/variables.py b/data_harness/data/tools/variables.py similarity index 89% rename from data_harness/tools/variables.py rename to data_harness/data/tools/variables.py index 1872fee..41f2217 100644 --- a/data_harness/tools/variables.py +++ b/data_harness/data/tools/variables.py @@ -1,7 +1,7 @@ from __future__ import annotations -from data_harness.cache import SessionCache -from data_harness.types import ToolAnnotations, ToolSpec +from data_harness.data.cache import SessionCache +from data_harness.llm.types import ToolAnnotations, ToolSpec _LIST_VARIABLES_ANNOTATIONS = ToolAnnotations( title="List Variables", diff --git a/data_harness/eval/case.py b/data_harness/eval/case.py index b80e688..53be981 100644 --- a/data_harness/eval/case.py +++ b/data_harness/eval/case.py @@ -6,8 +6,8 @@ from typing import TYPE_CHECKING, Any, Callable if TYPE_CHECKING: + from data_harness.core.result import RunResult from data_harness.eval.graders import Grade - from data_harness.result import RunResult # A grader inspects the run's outcome and decides pass/fail. Grader = Callable[["RunResult", "EvalCase"], "Grade"] diff --git a/data_harness/eval/graders.py b/data_harness/eval/graders.py index 04cd440..9dccd65 100644 --- a/data_harness/eval/graders.py +++ b/data_harness/eval/graders.py @@ -12,8 +12,8 @@ from typing import TYPE_CHECKING, Any if TYPE_CHECKING: + from data_harness.core.result import RunResult from data_harness.eval.case import EvalCase, Grader - from data_harness.result import RunResult _NUMBER_RE = re.compile(r"-?\d[\d,]*\.?\d*") diff --git a/data_harness/eval/runner.py b/data_harness/eval/runner.py index 3269af1..b74c82f 100644 --- a/data_harness/eval/runner.py +++ b/data_harness/eval/runner.py @@ -6,12 +6,12 @@ from time import perf_counter from typing import TYPE_CHECKING, Any +from data_harness.app.quickstart import Chat, ask from data_harness.eval.case import ConversationCase, EvalCase from data_harness.eval.report import CaseResult, EvalReport -from data_harness.quickstart import Chat, ask if TYPE_CHECKING: - from data_harness.providers.base import ProviderAdapter + from data_harness.llm.providers.base import ProviderAdapter def evaluate( diff --git a/data_harness/exceptions.py b/data_harness/exceptions.py deleted file mode 100644 index 6cdaa3f..0000000 --- a/data_harness/exceptions.py +++ /dev/null @@ -1,28 +0,0 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING - -if TYPE_CHECKING: - from data_harness.providers.base import NormalizedResponse - - -class MaxTurnsExceeded(RuntimeError): - """Raised when the ReAct loop reaches ``max_turns`` without an end-turn stop. - - Attributes: - turns: The number of turns that were executed before the limit was hit. - last_response: The final provider response, if available. - """ - - def __init__(self, turns: int, last_response: "NormalizedResponse | None" = None): - self.turns = turns - self.last_response = last_response - super().__init__(f"Max turns exceeded: {turns}") - - -class ToolNotFoundError(KeyError): - """Raised when a tool invocation names a tool that is not registered.""" - - -class SubagentRecursionError(RuntimeError): - """Raised when a subagent attempts to spawn another subagent.""" diff --git a/data_harness/llm/__init__.py b/data_harness/llm/__init__.py new file mode 100644 index 0000000..0648b87 --- /dev/null +++ b/data_harness/llm/__init__.py @@ -0,0 +1,9 @@ +"""Provider layer: the wire format and everything that speaks it. + +Owns the types a model provider exchanges (`Message`, `ToolSpec`, +`ContentBlock`), the streaming event protocol, and the adapters that translate +provider SDKs into them. + +Knows nothing about agents, tools-as-behaviour, caches, or DataFrames. It is +the bottom of the stack: nothing here may import `core`, `data`, or `app`. +""" diff --git a/data_harness/providers/__init__.py b/data_harness/llm/providers/__init__.py similarity index 100% rename from data_harness/providers/__init__.py rename to data_harness/llm/providers/__init__.py diff --git a/data_harness/providers/anthropic.py b/data_harness/llm/providers/anthropic.py similarity index 97% rename from data_harness/providers/anthropic.py rename to data_harness/llm/providers/anthropic.py index bcbd531..9d5ce70 100644 --- a/data_harness/providers/anthropic.py +++ b/data_harness/llm/providers/anthropic.py @@ -6,13 +6,13 @@ import anthropic -from data_harness.providers.base import ( +from data_harness.llm.providers.base import ( AsyncProviderAdapter, NormalizedResponse, ProviderAdapter, StopReason, ) -from data_harness.types import ( +from data_harness.llm.types import ( Message, TextBlock, ToolResultBlock, @@ -21,7 +21,7 @@ ) if TYPE_CHECKING: - from data_harness.streaming import StreamEvent + from data_harness.llm.streaming import StreamEvent _STOP_REASON_MAP = { "end_turn": StopReason.END_TURN, @@ -167,7 +167,7 @@ async def stream_events( tools: list[ToolSpec], ) -> AsyncGenerator[StreamEvent, None]: """Yield StreamEvents by mapping raw Anthropic SSE events.""" - from data_harness.streaming import ( + from data_harness.llm.streaming import ( ContentBlockDeltaEvent, ContentBlockStartEvent, ContentBlockStopEvent, diff --git a/data_harness/providers/base.py b/data_harness/llm/providers/base.py similarity index 96% rename from data_harness/providers/base.py rename to data_harness/llm/providers/base.py index 32c798e..6417e5c 100644 --- a/data_harness/providers/base.py +++ b/data_harness/llm/providers/base.py @@ -7,10 +7,16 @@ from enum import Enum from typing import TYPE_CHECKING -from data_harness.types import ContentBlock, Message, TextBlock, ToolSpec, ToolUseBlock +from data_harness.llm.types import ( + ContentBlock, + Message, + TextBlock, + ToolSpec, + ToolUseBlock, +) if TYPE_CHECKING: - from data_harness.streaming import StreamEvent + from data_harness.llm.streaming import StreamEvent class StopReason(Enum): @@ -130,7 +136,7 @@ async def stream_events( standard event types from the assembled response. Override in provider subclasses to emit real token-level events. """ - from data_harness.streaming import ( + from data_harness.llm.streaming import ( ContentBlockDeltaEvent, ContentBlockStartEvent, ContentBlockStopEvent, @@ -180,7 +186,7 @@ async def stream( on_chunk: Callable[[str], Awaitable[None]], ) -> NormalizedResponse: """Backward-compat text-only streaming; calls stream_events() internally.""" - from data_harness.streaming import ( + from data_harness.llm.streaming import ( ContentBlockDeltaEvent, TextDelta, accumulate_stream_events, diff --git a/data_harness/providers/openai.py b/data_harness/llm/providers/openai.py similarity index 99% rename from data_harness/providers/openai.py rename to data_harness/llm/providers/openai.py index 0a452b5..e72d9c2 100644 --- a/data_harness/providers/openai.py +++ b/data_harness/llm/providers/openai.py @@ -6,13 +6,13 @@ import openai -from data_harness.providers.base import ( +from data_harness.llm.providers.base import ( AsyncProviderAdapter, NormalizedResponse, ProviderAdapter, StopReason, ) -from data_harness.types import ( +from data_harness.llm.types import ( Message, TextBlock, ToolResultBlock, diff --git a/data_harness/streaming.py b/data_harness/llm/streaming.py similarity index 97% rename from data_harness/streaming.py rename to data_harness/llm/streaming.py index bc20251..407e8a7 100644 --- a/data_harness/streaming.py +++ b/data_harness/llm/streaming.py @@ -22,8 +22,8 @@ from dataclasses import dataclass, field from typing import Literal -from data_harness.providers.base import NormalizedResponse, StopReason -from data_harness.types import TextBlock, ToolUseBlock +from data_harness.llm.providers.base import NormalizedResponse, StopReason +from data_harness.llm.types import TextBlock, ToolUseBlock # --------------------------------------------------------------------------- # Delta types (carried by ContentBlockDeltaEvent) diff --git a/data_harness/testing.py b/data_harness/llm/testing.py similarity index 96% rename from data_harness/testing.py rename to data_harness/llm/testing.py index 9e7ed03..f3a437c 100644 --- a/data_harness/testing.py +++ b/data_harness/llm/testing.py @@ -10,13 +10,13 @@ import copy from typing import Any -from data_harness.providers.base import ( +from data_harness.llm.providers.base import ( AsyncProviderAdapter, NormalizedResponse, ProviderAdapter, StopReason, ) -from data_harness.types import Message, TextBlock, ToolSpec, ToolUseBlock +from data_harness.llm.types import Message, TextBlock, ToolSpec, ToolUseBlock class FakeAdapter(ProviderAdapter): diff --git a/data_harness/types.py b/data_harness/llm/types.py similarity index 93% rename from data_harness/types.py rename to data_harness/llm/types.py index af8ba15..70d4c8f 100644 --- a/data_harness/types.py +++ b/data_harness/llm/types.py @@ -89,6 +89,10 @@ class Message: def __post_init__(self) -> None: if self.role not in ("user", "assistant"): + # Plain ValueError, not ConfigurationError: `llm` is the bottom + # layer and may not import `core`, which is where the taxonomy + # lives. A malformed message is a caller error anyway, not a + # harness misconfiguration. raise ValueError( f"Invalid role: {self.role!r}. Must be 'user' or 'assistant'." ) diff --git a/data_harness/loop.py b/data_harness/loop.py deleted file mode 100644 index 6db5ee4..0000000 --- a/data_harness/loop.py +++ /dev/null @@ -1,800 +0,0 @@ -from __future__ import annotations - -import asyncio -import dataclasses -import functools -from collections.abc import AsyncGenerator, Callable -from typing import Literal, cast - -from data_harness.cache import SessionCache -from data_harness.exceptions import MaxTurnsExceeded -from data_harness.format import format_tool_output -from data_harness.logger import log_error_turn, log_turn, setup_logger -from data_harness.observe import time_block -from data_harness.providers.base import ( - AsyncProviderAdapter, - NormalizedResponse, - ProviderAdapter, - StopReason, -) -from data_harness.result import CacheStorageInfo, RunResult, Usage -from data_harness.streaming import ( - StreamEvent, - ToolResultEvent, - accumulate_stream_events, -) -from data_harness.types import ( - Message, - TextBlock, - ToolResultBlock, - ToolSpec, - ToolUseBlock, -) - -_MAX_TURN_REMINDER = ( - "This is the final turn. You MUST produce your complete final output now. " - "Do not use any more tools. Respond with your answer directly." -) - - -def _evaluate_code_gate( - on_code: Callable[[str], object] | None, - code_only: bool, - code: str, -) -> str | None: - """Decide whether interpreter ``code`` may run. - - Returns a string to short-circuit execution (a dry-run echo or a denial - message returned to the model), or ``None`` to proceed. - """ - if code_only: - return f"DRY RUN — code not executed:\n{code}" - if on_code is not None: - decision = on_code(code) - if decision is False: - return "Execution blocked by the approval gate." - if isinstance(decision, str): - return decision - return None - - -class Harness: - """The core synchronous ReAct loop. - - `Harness` owns the message list, dispatches tools, applies suffix-only - reminder hooks, and logs every turn to a JSONL file. It is the central - implementation boundary in data-harness: everything above it (`Agent`, - `AgentSession`) is a convenience layer; everything below it - (`ProviderAdapter`, `SessionCache`, `ToolSpec`) is a pure dependency. - - The system prompt is never mutated between turns. Reminders, nags, and - dynamic state are always appended to the conversation suffix so the - provider's KV cache is not invalidated. - - For most use cases, prefer `Agent` over constructing `Harness` directly. - Use `Harness` when you need full control over tool wiring, as shown in - ``examples/advanced_wiring.py``. - - Args: - adapter: Synchronous provider adapter that translates provider SDK - objects into harness types. - system: System prompt. Kept byte-identical across all turns. - tools: Full tool list. Invisible tools (``visible=False``) are excluded - from the provider call but can still be dispatched. - max_turns: Hard cap on provider turns before the loop stops and returns - a ``"max_turns_exceeded"`` result. - run_dir: Directory where JSONL logs are written. Created on first run. - cache: Shared `SessionCache`. A fresh cache is created if ``None``. - """ - - def __init__( - self, - adapter: ProviderAdapter, - system: str, - tools: list[ToolSpec], - max_turns: int = 25, - run_dir: str = "./runs", - cache: SessionCache | None = None, - on_code: Callable[[str], object] | None = None, - code_only: bool = False, - ) -> None: - if max_turns < 1: - raise ValueError(f"max_turns must be at least 1, got {max_turns!r}") - self._adapter = adapter - self._system = system - self._tools = list(tools) - self._max_turns = max_turns - self._run_dir = run_dir - self._cache = cache if cache is not None else SessionCache() - self._messages: list[Message] = [] - self._reminders: list[Callable[[int, int], str | None]] = [] - self._run_file: str | None = None - self._on_code = on_code - self._code_only = code_only - - def register_reminder(self, hook: Callable[[int, int], str | None]) -> None: - """Register a suffix reminder hook called before each provider turn. - - The hook receives ``(current_turn, max_turns)`` and returns a reminder - string to append to the conversation suffix, or ``None`` to skip. - - Args: - hook: Callable with signature ``(turn: int, max_turns: int) -> str | None``. - """ - self._reminders.append(hook) - - def run_result( - self, - user_message: str, - *, - run_id: str | None = None, - session_id: str | None = None, - ) -> RunResult: - """Start a fresh run and return the full `RunResult`. - - Resets message history. Use `ask_result` for follow-up turns on the - same history. - - Args: - user_message: The initial user prompt. - run_id: Optional identifier stamped into the `RunResult`. - session_id: Optional session identifier stamped into the `RunResult`. - - Returns: - A `RunResult` describing the outcome, token usage, and cache state. - """ - self._run_file = setup_logger(self._run_dir) - self._messages = [Message(role="user", content=[TextBlock(text=user_message)])] - result = self._run_loop_result() - return dataclasses.replace(result, run_id=run_id, session_id=session_id) - - def ask_result( - self, - user_message: str, - *, - run_id: str | None = None, - session_id: str | None = None, - ) -> RunResult: - """Append a follow-up message and continue the existing run. - - Appends ``user_message`` to the current history without resetting it. - Useful for multi-turn sessions when driving `Harness` directly. - - Args: - user_message: The follow-up user prompt. - run_id: Optional identifier stamped into the `RunResult`. - session_id: Optional session identifier stamped into the `RunResult`. - - Returns: - A `RunResult` describing the outcome of this turn sequence. - """ - if self._run_file is None: - self._run_file = setup_logger(self._run_dir) - self._messages.append( - Message(role="user", content=[TextBlock(text=user_message)]) - ) - result = self._run_loop_result() - return dataclasses.replace(result, run_id=run_id, session_id=session_id) - - def run(self, user_message: str) -> str: - """Start a fresh run and return the final text response. - - Raises `MaxTurnsExceeded` if the loop hits ``max_turns``. - - Args: - user_message: The initial user prompt. - - Returns: - The model's final text response. - - Raises: - MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. - RuntimeError: If the provider raises an exception during the run. - """ - result = self.run_result(user_message) - if result.status == "max_turns_exceeded": - raise MaxTurnsExceeded(result.turns) - if result.status == "error": - raise RuntimeError(result.error or "unknown error") - return result.text - - def ask(self, user_message: str) -> str: - """Append a follow-up message and return the final text response. - - Args: - user_message: The follow-up user prompt. - - Returns: - The model's final text response. - - Raises: - MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. - RuntimeError: If the provider raises an exception during the run. - """ - result = self.ask_result(user_message) - if result.status == "max_turns_exceeded": - raise MaxTurnsExceeded(result.turns) - if result.status == "error": - raise RuntimeError(result.error or "unknown error") - return result.text - - @property - def run_file(self) -> str | None: - """Path to the JSONL log for this run, or ``None`` before the first run.""" - return self._run_file - - def _run_loop_result(self) -> RunResult: - if self._run_file is None: - raise RuntimeError("run_file must be initialised before running the loop") - - total_usage = Usage() - - for turn in range(1, self._max_turns + 1): - self._apply_reminders(turn) - visible_tools = [t for t in self._tools if t.visible] - - try: - with time_block() as tb: - response = self._adapter.chat( - system=self._system, - messages=self._messages, - tools=visible_tools, - ) - except Exception as exc: - log_error_turn( - turn=turn, - system=self._system, - messages=self._messages, - error=repr(exc), - run_file=self._run_file, - ) - return RunResult( - text="", - status="error", - turns=turn, - run_file=self._run_file, - stop_reason=None, - usage=total_usage, - cache_snapshots=self._cache.list_handles(), - cache_storage=self._build_cache_storage(), - value=self._cache.get_answer(), - charts=self._cache.list_charts(), - error=repr(exc), - ) - - latency = tb.elapsed_ms - - total_usage = total_usage + Usage( - input_tokens=response.input_tokens, - output_tokens=response.output_tokens, - cache_read_tokens=response.cache_read_tokens, - cache_write_tokens=response.cache_write_tokens, - ) - - self._messages.append(Message(role="assistant", content=response.content)) - - tool_results: list[ToolResultBlock] = [] - - if response.stop_reason == StopReason.TOOL_USE: - tool_results = self._dispatch_tools(response.content) - user_msg = Message(role="user", content=list(tool_results)) - self._messages.append(user_msg) - - tool_error_count = sum(1 for r in tool_results if r.is_error) - - log_turn( - turn=turn, - system=self._system, - messages=self._messages, - response=response, - tool_results=tool_results, - latency_ms=latency, - run_file=self._run_file, - cache_storage=self._cache.storage_metadata(), - visible_tools=[t.name for t in visible_tools], - tool_error_count=tool_error_count, - all_tools=self._tools, - ) - - if response.stop_reason != StopReason.TOOL_USE: - return RunResult( - text=self._extract_text(response), - status="success", - turns=turn, - run_file=self._run_file, - stop_reason=response.stop_reason, - usage=total_usage, - cache_snapshots=self._cache.list_handles(), - cache_storage=self._build_cache_storage(), - value=self._cache.get_answer(), - charts=self._cache.list_charts(), - ) - - if turn == self._max_turns: - return RunResult( - text=self._extract_text(response), - status="max_turns_exceeded", - turns=turn, - run_file=self._run_file, - stop_reason=None, - usage=total_usage, - cache_snapshots=self._cache.list_handles(), - cache_storage=self._build_cache_storage(), - value=self._cache.get_answer(), - charts=self._cache.list_charts(), - ) - - def _build_cache_storage(self) -> dict[str, CacheStorageInfo]: - raw = self._cache.storage_metadata() - return { - name: CacheStorageInfo( - location=cast(Literal["memory", "disk"], meta["location"]), - storage_type=meta["storage_type"], - ) - for name, meta in raw.items() - } - - def _apply_reminders(self, turn: int) -> None: - reminder_texts: list[str] = [] - - for hook in self._reminders: - text = hook(turn, self._max_turns) - if text: - reminder_texts.append(text) - - # Built-in max-turn reminder - if turn == self._max_turns - 1: - reminder_texts.append(_MAX_TURN_REMINDER) - - if not reminder_texts: - return - - combined = "\n\n".join(reminder_texts) - reminder_block = TextBlock(text=combined) - - # Append to existing user message or create a new one - if self._messages and self._messages[-1].role == "user": - self._messages[-1].content.append(reminder_block) - else: - self._messages.append(Message(role="user", content=[reminder_block])) - - def _dispatch_tools(self, content: list) -> list[ToolResultBlock]: - tool_uses = [b for b in content if isinstance(b, ToolUseBlock)] - results = [] - tool_map = {t.name: t for t in self._tools} - - for tub in tool_uses: - spec = tool_map.get(tub.tool_name) - if spec is None or spec.handler is None: - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=f"Tool not found: {tub.tool_name!r}", - is_error=True, - ) - ) - continue - if tub.tool_name == "python_interpreter": - gate = _evaluate_code_gate( - self._on_code, self._code_only, tub.tool_input.get("code", "") - ) - if gate is not None: - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=gate, - is_error=False, - ) - ) - continue - try: - raw = spec.handler(**tub.tool_input) - output = format_tool_output(raw, cache=self._cache) - except Exception as exc: - output = repr(exc) - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=output, - is_error=True, - ) - ) - continue - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=output, - is_error=False, - ) - ) - - return results - - def _extract_text(self, response: NormalizedResponse) -> str: - texts = [b.text for b in response.content if isinstance(b, TextBlock)] - return "\n".join(texts) - - -class AsyncHarness: - """Async variant of Harness. Requires an AsyncProviderAdapter. - - Exposes the same run_result / ask_result / run / ask surface as Harness, - plus run_stream / ask_stream for token-level streaming. - """ - - def __init__( - self, - adapter: AsyncProviderAdapter, - system: str, - tools: list[ToolSpec], - max_turns: int = 25, - run_dir: str = "./runs", - cache: SessionCache | None = None, - on_code: Callable[[str], object] | None = None, - code_only: bool = False, - ) -> None: - if max_turns < 1: - raise ValueError(f"max_turns must be at least 1, got {max_turns!r}") - self._adapter = adapter - self._system = system - self._tools = list(tools) - self._max_turns = max_turns - self._run_dir = run_dir - self._cache = cache if cache is not None else SessionCache() - self._messages: list[Message] = [] - self._reminders: list[Callable[[int, int], str | None]] = [] - self._run_file: str | None = None - self._on_code = on_code - self._code_only = code_only - - def register_reminder(self, hook: Callable[[int, int], str | None]) -> None: - self._reminders.append(hook) - - @property - def run_file(self) -> str | None: - return self._run_file - - async def run_result( - self, - user_message: str, - *, - run_id: str | None = None, - session_id: str | None = None, - ) -> RunResult: - self._run_file = setup_logger(self._run_dir) - self._messages = [Message(role="user", content=[TextBlock(text=user_message)])] - result = await self._run_loop_result() - return dataclasses.replace(result, run_id=run_id, session_id=session_id) - - async def ask_result( - self, - user_message: str, - *, - run_id: str | None = None, - session_id: str | None = None, - ) -> RunResult: - if self._run_file is None: - self._run_file = setup_logger(self._run_dir) - self._messages.append( - Message(role="user", content=[TextBlock(text=user_message)]) - ) - result = await self._run_loop_result() - return dataclasses.replace(result, run_id=run_id, session_id=session_id) - - async def run(self, user_message: str) -> str: - result = await self.run_result(user_message) - if result.status == "max_turns_exceeded": - raise MaxTurnsExceeded(result.turns) - if result.status == "error": - raise RuntimeError(result.error or "unknown error") - return result.text - - async def ask(self, user_message: str) -> str: - result = await self.ask_result(user_message) - if result.status == "max_turns_exceeded": - raise MaxTurnsExceeded(result.turns) - if result.status == "error": - raise RuntimeError(result.error or "unknown error") - return result.text - - async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: - """Stream events for a one-shot run. - - Yields StreamEvent objects following the same protocol as the Claude - Agent SDK. Each provider turn emits message_start, - content_block_start/delta/stop, message_delta, and message_stop events. - After the harness dispatches a tool call a ToolResultEvent is emitted. - The JSONL logger records fully assembled messages, not individual events. - """ - self._run_file = setup_logger(self._run_dir) - self._messages = [Message(role="user", content=[TextBlock(text=user_message)])] - async for event in self._run_loop_stream(): - yield event - - async def ask_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: - """Stream events for a follow-up turn in a session.""" - if self._run_file is None: - self._run_file = setup_logger(self._run_dir) - self._messages.append( - Message(role="user", content=[TextBlock(text=user_message)]) - ) - async for event in self._run_loop_stream(): - yield event - - async def _run_loop_result(self) -> RunResult: - if self._run_file is None: - raise RuntimeError("run_file must be initialised before running the loop") - - total_usage = Usage() - - for turn in range(1, self._max_turns + 1): - self._apply_reminders(turn) - visible_tools = [t for t in self._tools if t.visible] - - try: - with time_block() as tb: - response = await self._adapter.chat( - system=self._system, - messages=self._messages, - tools=visible_tools, - ) - except Exception as exc: - log_error_turn( - turn=turn, - system=self._system, - messages=self._messages, - error=repr(exc), - run_file=self._run_file, - ) - return RunResult( - text="", - status="error", - turns=turn, - run_file=self._run_file, - stop_reason=None, - usage=total_usage, - cache_snapshots=self._cache.list_handles(), - cache_storage=self._build_cache_storage(), - value=self._cache.get_answer(), - charts=self._cache.list_charts(), - error=repr(exc), - ) - - latency = tb.elapsed_ms - - total_usage = total_usage + Usage( - input_tokens=response.input_tokens, - output_tokens=response.output_tokens, - cache_read_tokens=response.cache_read_tokens, - cache_write_tokens=response.cache_write_tokens, - ) - - self._messages.append(Message(role="assistant", content=response.content)) - - tool_results: list[ToolResultBlock] = [] - - if response.stop_reason == StopReason.TOOL_USE: - tool_results = await self._dispatch_tools(response.content) - self._messages.append(Message(role="user", content=list(tool_results))) - - tool_error_count = sum(1 for r in tool_results if r.is_error) - - log_turn( - turn=turn, - system=self._system, - messages=self._messages, - response=response, - tool_results=tool_results, - latency_ms=latency, - run_file=self._run_file, - cache_storage=self._cache.storage_metadata(), - visible_tools=[t.name for t in visible_tools], - tool_error_count=tool_error_count, - all_tools=self._tools, - ) - - if response.stop_reason != StopReason.TOOL_USE: - return RunResult( - text=self._extract_text(response), - status="success", - turns=turn, - run_file=self._run_file, - stop_reason=response.stop_reason, - usage=total_usage, - cache_snapshots=self._cache.list_handles(), - cache_storage=self._build_cache_storage(), - value=self._cache.get_answer(), - charts=self._cache.list_charts(), - ) - - if turn == self._max_turns: - return RunResult( - text=self._extract_text(response), - status="max_turns_exceeded", - turns=turn, - run_file=self._run_file, - stop_reason=None, - usage=total_usage, - cache_snapshots=self._cache.list_handles(), - cache_storage=self._build_cache_storage(), - value=self._cache.get_answer(), - charts=self._cache.list_charts(), - ) - - async def _run_loop_stream(self) -> AsyncGenerator[StreamEvent, None]: - if self._run_file is None: - raise RuntimeError("run_file must be initialised before running the loop") - - total_usage = Usage() - - for turn in range(1, self._max_turns + 1): - self._apply_reminders(turn) - visible_tools = [t for t in self._tools if t.visible] - - events_this_turn: list[StreamEvent] = [] - - with time_block() as tb: - try: - async for evt in self._adapter.stream_events( - system=self._system, - messages=self._messages, - tools=visible_tools, - ): - events_this_turn.append(evt) - yield evt - except Exception as exc: - log_error_turn( - turn=turn, - system=self._system, - messages=self._messages, - error=repr(exc), - run_file=self._run_file, - ) - return - - latency = tb.elapsed_ms - response = accumulate_stream_events(events_this_turn) - - total_usage = total_usage + Usage( - input_tokens=response.input_tokens, - output_tokens=response.output_tokens, - cache_read_tokens=response.cache_read_tokens, - cache_write_tokens=response.cache_write_tokens, - ) - - self._messages.append(Message(role="assistant", content=response.content)) - - tool_results: list[ToolResultBlock] = [] - - if response.stop_reason == StopReason.TOOL_USE: - tool_results = await self._dispatch_tools(response.content) - - tool_name_map = { - b.tool_use_id: b.tool_name - for b in response.content - if isinstance(b, ToolUseBlock) - } - for result in tool_results: - yield ToolResultEvent( - tool_use_id=result.tool_use_id, - tool_name=tool_name_map.get(result.tool_use_id, ""), - content=result.content, - is_error=result.is_error, - ) - - self._messages.append(Message(role="user", content=list(tool_results))) - - tool_error_count = sum(1 for r in tool_results if r.is_error) - - log_turn( - turn=turn, - system=self._system, - messages=self._messages, - response=response, - tool_results=tool_results, - latency_ms=latency, - run_file=self._run_file, - cache_storage=self._cache.storage_metadata(), - visible_tools=[t.name for t in visible_tools], - tool_error_count=tool_error_count, - all_tools=self._tools, - ) - - if response.stop_reason != StopReason.TOOL_USE: - return - - if turn == self._max_turns: - return - - def _build_cache_storage(self) -> dict[str, CacheStorageInfo]: - raw = self._cache.storage_metadata() - return { - name: CacheStorageInfo( - location=cast(Literal["memory", "disk"], meta["location"]), - storage_type=meta["storage_type"], - ) - for name, meta in raw.items() - } - - def _apply_reminders(self, turn: int) -> None: - reminder_texts: list[str] = [] - - for hook in self._reminders: - text = hook(turn, self._max_turns) - if text: - reminder_texts.append(text) - - if turn == self._max_turns - 1: - reminder_texts.append(_MAX_TURN_REMINDER) - - if not reminder_texts: - return - - combined = "\n\n".join(reminder_texts) - reminder_block = TextBlock(text=combined) - - if self._messages and self._messages[-1].role == "user": - self._messages[-1].content.append(reminder_block) - else: - self._messages.append(Message(role="user", content=[reminder_block])) - - async def _dispatch_tools(self, content: list) -> list[ToolResultBlock]: - tool_uses = [b for b in content if isinstance(b, ToolUseBlock)] - results = [] - tool_map = {t.name: t for t in self._tools} - - for tub in tool_uses: - spec = tool_map.get(tub.tool_name) - if spec is None or spec.handler is None: - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=f"Tool not found: {tub.tool_name!r}", - is_error=True, - ) - ) - continue - if tub.tool_name == "python_interpreter": - gate = _evaluate_code_gate( - self._on_code, self._code_only, tub.tool_input.get("code", "") - ) - if gate is not None: - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=gate, - is_error=False, - ) - ) - continue - try: - if asyncio.iscoroutinefunction(spec.handler): - raw = await spec.handler(**tub.tool_input) - else: - raw = await asyncio.to_thread( - functools.partial(spec.handler, **tub.tool_input) - ) - output = format_tool_output(raw, cache=self._cache) - except Exception as exc: - output = repr(exc) - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=output, - is_error=True, - ) - ) - continue - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=output, - is_error=False, - ) - ) - - return results - - def _extract_text(self, response: NormalizedResponse) -> str: - texts = [b.text for b in response.content if isinstance(b, TextBlock)] - return "\n".join(texts) diff --git a/data_harness/tools/__init__.py b/data_harness/tools/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/docs/guide/design.md b/docs/guide/design.md index 7f0be60..e9582bb 100644 --- a/docs/guide/design.md +++ b/docs/guide/design.md @@ -5,6 +5,123 @@ exist. Understanding them makes the API surface predictable. --- +## Layers + +The package is layered, bottom up. A layer may import the ones below it and +none above; `tests/test_layers.py` enforces that statically. + +``` +data_harness/ + llm/ provider adapters and the wire types they speak + core/ the loop, the session tree, RunResult, hooks, compaction + data/ the session cache, interpreter, SQL, connectors, MCP + app/ Agent, ask(), the CLI +``` + +The boundary that matters is `core` not knowing what a DataFrame is. The loop +takes a `RunEnvironment` supplying two things it cannot decide for itself: how +to render a tool's return value, and what the run's final state was. The data +layer's implementation is backed by the `SessionCache`; `NullEnvironment` is +the domain-free default. That is why the harness can be read, tested, and +reused without pandas. + +Every pre-layering import path (`data_harness.loop`, `data_harness.types`, …) +still resolves, to the same module object rather than a copy. + +--- + +## One loop, two drivers + +`_HarnessBase._plan` is the only ReAct loop. It is a generator that owns every +decision in a turn and performs no network I/O: it *asks* for I/O by yielding +`CallProvider`, `CallTool`, or `ToolFinished`, and the driver answers with +`Ok(value)` or `Failed(exc)`. + +- `Harness` performs those inline on the calling thread. No event loop is + created, so Ctrl-C lands promptly and a tool handler keeps its thread + affinity: a `sqlite3` connection opened at setup still works. +- `AsyncHarness` awaits the provider and offloads blocking handlers to a + worker thread, so a long pandas call cannot stall a shared event loop. + +There used to be four near-copies of this loop (sync/async x result/stream). +They had drifted: the streaming copy discarded token usage on provider errors, +and `AsyncAgent` was missing eight features `Agent` had. + +--- + +## The session tree + +A session is an append-only tree of typed entries, each naming its parent, +with a movable leaf. **The conversation the model sees is derived by walking +root to leaf, and is never stored**, so there is no second copy to disagree +with the log. + +Everything else follows from that: + +- **Resume** — reopen a session file in another process and carry on. +- **Forking** — move the leaf and append; both branches survive, so a turn can + be retried with a different model without destroying the first attempt. +- **Compaction is an entry, not an edit.** The compacted turns stay in the + tree; moving the leaf back restores them. +- Writing is one line per entry, where the older `runs/*.jsonl` turn log + re-serialises the whole history every turn. + +A derived context is always something a provider will accept: a tool call with +no result, or a result with no call, is dropped. Both are reachable without a +bug, because a run killed mid-tool leaves an orphaned call on disk. + +--- + +## Hooks + +Four events, each with a decision a hook may return: + +``` +BeforeTurn -> Reminder(text) | Stop(reason) +BeforeToolCall -> Block(reason, is_error) +AfterToolCall -> Replace(content, is_error) +AfterTurn -> Stop(reason) +``` + +The interpreter approval gate is built from exactly this mechanism rather than +being special-cased inside the loop, which is the evidence that the mechanism +is sufficient. `AfterTurn` is where a spend cap belongs: the tokens are already +counted, so the decision is made on real numbers. + +Hooks must not raise. One that does is reported as `HookError`, and the run +fails with its usage intact rather than vanishing. + +--- + +## Compaction + +`max_turns` is a wall, not a strategy: a run needing thirty turns fails at +twenty-five having paid for all of them. Compaction summarises older turns and +replays the summary in their place. + +The cut always lands on a turn boundary, so an assistant tool call is never +separated from the result answering it. When there is nowhere safe to cut, +nothing is cut. + +Compaction costs this library less than it costs a coding agent: the data +lives in the cache under handles, not in the transcript, so compacting away +the turn that loaded a DataFrame loses the discussion of it, not the frame. + +--- + +## Errors + +Every failure carries a stable `code`, because a caller deciding between +retrying, showing the user an error, and billing the attempt needs to tell a +rate-limited provider from a typo in the model's pandas from a sandbox +timeout. Codes are API; messages are not. + +`ExecutionError` means the code never ran (timeout, killed process). +`PythonInterpreterError` means it ran and failed, which is the model's problem +and is handed back to it as a tool result. + +--- + ## No bash Giving an agent shell access is the path of least resistance, but it creates diff --git a/examples/advanced_wiring.py b/examples/advanced_wiring.py index d66a64c..c73ddfb 100644 --- a/examples/advanced_wiring.py +++ b/examples/advanced_wiring.py @@ -13,14 +13,14 @@ import pandas as pd -from data_harness.cache import SessionCache -from data_harness.loop import Harness -from data_harness.tools.connectors import ConnectorRegistry -from data_harness.tools.interpreter import PythonInterpreter -from data_harness.tools.planner import Planner -from data_harness.tools.subagent import make_subagent_spec -from data_harness.tools.variables import make_list_variables_spec -from data_harness.types import ToolSpec +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.data.tools.connectors import ConnectorRegistry +from data_harness.data.tools.interpreter import PythonInterpreter +from data_harness.data.tools.planner import Planner +from data_harness.data.tools.subagent import make_subagent_spec +from data_harness.data.tools.variables import make_list_variables_spec +from data_harness.llm.types import ToolSpec DATA_PATH = Path(__file__).parent / "data" / "fred_unrate_2024.csv" @@ -84,7 +84,7 @@ def main() -> None: print("ANTHROPIC_API_KEY not set. Skipping live demo.") sys.exit(0) - from data_harness.providers.anthropic import AnthropicAdapter + from data_harness.llm.providers.anthropic import AnthropicAdapter session_cache = SessionCache(sample_size=5) diff --git a/examples/cache_benchmark.py b/examples/cache_benchmark.py index 19acb65..6556cf0 100644 --- a/examples/cache_benchmark.py +++ b/examples/cache_benchmark.py @@ -10,14 +10,13 @@ from __future__ import annotations -import time - import dataclasses +import time import pandas as pd from data_harness import Agent, ExecutionCache -from data_harness.testing import FakeAdapter +from data_harness.llm.testing import FakeAdapter SALES = pd.DataFrame( {"month": ["Jan", "Feb", "Mar", "Apr"], "revenue": [120, 150, 90, 200]} diff --git a/examples/demo.ipynb b/examples/demo.ipynb index 47daadb..78817ec 100644 --- a/examples/demo.ipynb +++ b/examples/demo.ipynb @@ -109,9 +109,11 @@ ], "source": [ "from dotenv import load_dotenv\n", + "\n", "load_dotenv()\n", "import pandas as pd\n", - "from data_harness import ask, Chat, Agent, ExecutionCache\n", + "\n", + "from data_harness import Agent, Chat, ExecutionCache, ask\n", "\n", "MODEL = 'gpt-4o-mini'\n", "sales = pd.DataFrame({\n", diff --git a/examples/inspect_run.py b/examples/inspect_run.py index 9f56aab..18520cd 100644 --- a/examples/inspect_run.py +++ b/examples/inspect_run.py @@ -4,7 +4,7 @@ """ from data_harness import Agent, RunResult -from data_harness.testing import FakeAdapter +from data_harness.llm.testing import FakeAdapter adapter = FakeAdapter([FakeAdapter.text("The mean of [1, 2, 3] is 2.0")]) diff --git a/examples/mcp_demo.py b/examples/mcp_demo.py index ccc0772..e198da3 100644 --- a/examples/mcp_demo.py +++ b/examples/mcp_demo.py @@ -15,7 +15,7 @@ from dotenv import load_dotenv from data_harness import Agent -from data_harness.quickstart import resolve_adapter +from data_harness.app.quickstart import resolve_adapter def main() -> None: diff --git a/examples/quickstart.py b/examples/quickstart.py index 79dba54..6fb6bf9 100644 --- a/examples/quickstart.py +++ b/examples/quickstart.py @@ -30,7 +30,7 @@ def build_agent(adapter, system="You are a data analyst."): print("ANTHROPIC_API_KEY not set. Skipping live quick start.") sys.exit(0) - from data_harness.providers.anthropic import AnthropicAdapter + from data_harness.llm.providers.anthropic import AnthropicAdapter agent = build_agent(AnthropicAdapter(model="claude-sonnet-4-6")) result = agent.run("Compute the mean of [1, 2, 3, 4, 5] and print it.") diff --git a/tests/smoke_tests.py b/tests/smoke_tests.py index c99bce8..62725c8 100644 --- a/tests/smoke_tests.py +++ b/tests/smoke_tests.py @@ -28,14 +28,14 @@ import pytest from data_harness import Agent -from data_harness.cache import SessionCache -from data_harness.loop import Harness -from data_harness.providers.base import StopReason -from data_harness.providers.openai import OpenRouterAdapter -from data_harness.result import CacheStorageInfo, RunResult -from data_harness.tools.planner import Planner -from data_harness.tools.subagent import make_subagent_spec -from data_harness.types import ToolAnnotations, ToolSpec +from data_harness.core.result import CacheStorageInfo, RunResult +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.data.tools.planner import Planner +from data_harness.data.tools.subagent import make_subagent_spec +from data_harness.llm.providers.base import StopReason +from data_harness.llm.providers.openai import OpenRouterAdapter +from data_harness.llm.types import ToolAnnotations, ToolSpec from examples.advanced_wiring import build_base_tools, load_unemployment_rate pytestmark = pytest.mark.live @@ -75,7 +75,7 @@ def _latest_jsonl(run_dir: Path) -> list[dict]: def _all_text_from_messages(harness: Harness) -> str: parts: list[str] = [] - for message in harness._messages: + for message in harness.messages: for block in message.content: text = getattr(block, "text", None) or getattr(block, "content", None) if text: diff --git a/tests/test_agent.py b/tests/test_agent.py index 436b63f..c0ccbe2 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -6,19 +6,19 @@ import pytest -from data_harness.agent import Agent -from data_harness.cache import SessionCache -from data_harness.loop import Harness -from data_harness.testing import FakeAdapter -from data_harness.tools.planner import Planner -from data_harness.types import ToolResultBlock, ToolSpec +from data_harness.app.agent import Agent +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.data.tools.planner import Planner +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import ToolResultBlock, ToolSpec def test_agent_is_exported_from_top_level_package(): from data_harness import Agent as TopLevelAgent from data_harness import AgentSession as TopLevelAgentSession - from data_harness.agent import Agent as ModuleAgent - from data_harness.agent import AgentSession as ModuleAgentSession + from data_harness.app.agent import Agent as ModuleAgent + from data_harness.app.agent import AgentSession as ModuleAgentSession assert TopLevelAgent is ModuleAgent assert TopLevelAgentSession is ModuleAgentSession @@ -100,7 +100,7 @@ def test_max_turns_propagates(self, tmp_path): adapter = FakeAdapter([FakeAdapter.text("done")]) agent = Agent(adapter=adapter, system="sys", max_turns=7, run_dir=str(tmp_path)) agent.run("hi") - assert agent.last_harness._max_turns == 7 + assert agent.last_harness.max_turns == 7 def test_explain_returns_readable_sketch(self, tmp_path): adapter = FakeAdapter([FakeAdapter.text("done")]) @@ -128,7 +128,7 @@ def test_second_run_does_not_see_first_run_messages(self, tmp_path): agent.run("first user prompt") agent.run("second user prompt") - from data_harness.types import TextBlock + from data_harness.llm.types import TextBlock second_call_msgs = adapter.calls[1]["messages"] all_text = " ".join( @@ -150,7 +150,7 @@ def test_session_preserves_message_history_across_asks(self, tmp_path): assert session.ask("first question") == "first" assert session.ask("follow-up question") == "second" - from data_harness.types import TextBlock + from data_harness.llm.types import TextBlock second_call_msgs = adapter.calls[1]["messages"] all_text = " ".join( @@ -198,7 +198,7 @@ def test_session_does_not_change_agent_run_one_shot_behaviour(self, tmp_path): agent.session().ask("session question") agent.run("standalone question") - from data_harness.types import TextBlock + from data_harness.llm.types import TextBlock standalone_call_msgs = adapter.calls[1]["messages"] all_text = " ".join( @@ -331,7 +331,7 @@ def fetch_ohlcv(symbol: str) -> list[str]: agent.run("hi") - names = {t.name for t in agent.last_harness._tools} + names = {t.name for t in agent.last_harness.tools} assert "market_data__fetch_ohlcv" in names def test_explicit_input_schema_override_bypasses_inference(self, tmp_path): @@ -354,7 +354,7 @@ def fetch(payload: dict) -> str: agent.run("hi") - specs = {t.name: t for t in agent.last_harness._tools} + specs = {t.name: t for t in agent.last_harness.tools} assert specs["market_data__fetch"].input_schema is schema @@ -376,7 +376,7 @@ def test_enable_planner_registers_reminder_hook(self, tmp_path): agent.enable_planner() agent.run("plan") - reminders = agent.last_harness._reminders + reminders = agent.last_harness.reminders assert len(reminders) == 1 assert isinstance(reminders[0].__self__, Planner) @@ -388,7 +388,7 @@ def test_planner_absent_when_not_enabled(self, tmp_path): names = {t.name for t in adapter.calls[0]["tools"]} assert not any(name.startswith("planner__") for name in names) - assert agent.last_harness._reminders == [] + assert agent.last_harness.reminders == [] def test_enable_planner_twice_does_not_duplicate_specs_or_hooks(self, tmp_path): adapter = FakeAdapter([FakeAdapter.text("done")]) @@ -402,7 +402,7 @@ def test_enable_planner_twice_does_not_duplicate_specs_or_hooks(self, tmp_path): assert names.count("planner__add") == 1 assert names.count("planner__update") == 1 assert names.count("planner__list") == 1 - assert len(agent.last_harness._reminders) == 1 + assert len(agent.last_harness.reminders) == 1 def test_planner_state_does_not_leak_across_runs(self, tmp_path): adapter = FakeAdapter( @@ -431,6 +431,19 @@ def test_planner_state_does_not_leak_across_runs(self, tmp_path): assert "task A" not in tool_results[-1].content +def _subagent_harness(captured): + """Pick the spawned subagent's harness out of every harness constructed. + + Selects by the worker system prompt rather than by construction order, + which is an implementation detail of where the patch happens to bite. + """ + subs = [ + h for h in captured if h.system.startswith("You are a clean-context worker") + ] + assert subs, "no subagent harness was constructed" + return subs[0] + + class TestAgentSubagents: def test_enable_subagents_adds_subagent_tool(self, tmp_path): adapter = FakeAdapter([FakeAdapter.text("done")]) @@ -494,7 +507,7 @@ def recording_make_subagent_spec( ) monkeypatch.setattr( - "data_harness.agent.make_subagent_spec", recording_make_subagent_spec + "data_harness.app.agent.make_subagent_spec", recording_make_subagent_spec ) adapter = FakeAdapter([FakeAdapter.text("done")]) agent = Agent(adapter=adapter, system="sys", run_dir=str(tmp_path)) @@ -505,7 +518,7 @@ def recording_make_subagent_spec( assert "subagent" not in captured_names def test_subagent_does_not_inherit_planner_hooks(self, monkeypatch, tmp_path): - from data_harness.loop import Harness as RealHarness + from data_harness.data.harness import Harness as RealHarness captured = [] @@ -514,7 +527,7 @@ def recording_harness(*args, **kwargs): captured.append(harness) return harness - monkeypatch.setattr("data_harness.loop.Harness", recording_harness) + monkeypatch.setattr("data_harness.data.harness.Harness", recording_harness) adapter = FakeAdapter( [ FakeAdapter.tool_use("tu_1", "subagent", {"task": "work"}), @@ -530,11 +543,10 @@ def recording_harness(*args, **kwargs): agent.run("delegate") - assert captured - assert captured[0]._reminders == [] + assert _subagent_harness(captured).reminders == [] def test_subagent_connector_tools_are_fresh_and_hidden(self, monkeypatch, tmp_path): - from data_harness.loop import Harness as RealHarness + from data_harness.data.harness import Harness as RealHarness captured = [] @@ -543,7 +555,7 @@ def recording_harness(*args, **kwargs): captured.append(harness) return harness - monkeypatch.setattr("data_harness.loop.Harness", recording_harness) + monkeypatch.setattr("data_harness.data.harness.Harness", recording_harness) adapter = FakeAdapter( [ FakeAdapter.tool_use( @@ -569,10 +581,9 @@ def fetch_ohlcv(symbol: str) -> list[str]: agent.run("load then delegate") - assert captured - sub_harness = captured[0] - sub_specs = {t.name: t for t in sub_harness._tools} - parent_specs = {t.name: t for t in agent.last_harness._tools} + sub_harness = _subagent_harness(captured) + sub_specs = {t.name: t for t in sub_harness.tools} + parent_specs = {t.name: t for t in agent.last_harness.tools} assert sub_specs["load_connectors"].visible is True assert sub_specs["market_data__fetch_ohlcv"].visible is False assert parent_specs["market_data__fetch_ohlcv"].visible is True @@ -582,7 +593,7 @@ def fetch_ohlcv(symbol: str) -> list[str]: first_sub_connector_id = id(sub_specs["market_data__fetch_ohlcv"]) agent.run("fresh second run") - second_parent_specs = {t.name: t for t in agent.last_harness._tools} + second_parent_specs = {t.name: t for t in agent.last_harness.tools} assert id(second_parent_specs["market_data__fetch_ohlcv"]) != id( parent_specs["market_data__fetch_ohlcv"] ) diff --git a/tests/test_approval_and_cache.py b/tests/test_approval_and_cache.py index 11bb45d..8fb0b65 100644 --- a/tests/test_approval_and_cache.py +++ b/tests/test_approval_and_cache.py @@ -5,7 +5,7 @@ import pandas as pd from data_harness import Agent, ExecutionCache -from data_harness.testing import FakeAdapter +from data_harness.llm.testing import FakeAdapter def _df() -> pd.DataFrame: @@ -55,7 +55,7 @@ def test_code_only_dry_run(tmp_path): res = agent.run_result("sum") assert res.value is None # the interpreter result echoed the code rather than running it - last_user = agent.last_harness._messages[-2] + last_user = agent.last_harness.messages[-2] contents = " ".join(b.content for b in last_user.content if hasattr(b, "content")) assert "DRY RUN" in contents @@ -109,8 +109,8 @@ def test_cache_persists_to_disk(tmp_path): def test_extract_steps_skips_errored_calls(): - from data_harness.exec_cache import extract_steps - from data_harness.types import ( + from data_harness.data.exec_cache import extract_steps + from data_harness.llm.types import ( Message, TextBlock, ToolResultBlock, diff --git a/tests/test_async_loop.py b/tests/test_async_loop.py index 4555c50..175ebeb 100644 --- a/tests/test_async_loop.py +++ b/tests/test_async_loop.py @@ -4,18 +4,18 @@ import pytest -from data_harness.agent import AsyncAgent -from data_harness.exceptions import MaxTurnsExceeded -from data_harness.loop import AsyncHarness -from data_harness.streaming import ( +from data_harness.app.agent import AsyncAgent +from data_harness.core.exceptions import MaxTurnsExceeded +from data_harness.data.harness import AsyncHarness +from data_harness.llm.streaming import ( ContentBlockDeltaEvent, MessageStopEvent, StreamEvent, TextDelta, ToolResultEvent, ) -from data_harness.testing import FakeAsyncAdapter -from data_harness.types import ToolSpec +from data_harness.llm.testing import FakeAsyncAdapter +from data_harness.llm.types import ToolSpec # --------------------------------------------------------------------------- # AsyncHarness — basic run_result / run @@ -197,9 +197,9 @@ def exploding_tool() -> str: result = await harness.run_result("go") assert result.status == "success" # Inspect the tool result in the message history - from data_harness.types import ToolResultBlock + from data_harness.llm.types import ToolResultBlock - for msg in reversed(harness._messages): + for msg in reversed(harness.messages): if msg.role == "user": tool_results = [b for b in msg.content if isinstance(b, ToolResultBlock)] if tool_results: diff --git a/tests/test_cache.py b/tests/test_cache.py index eb468d1..c5a13de 100644 --- a/tests/test_cache.py +++ b/tests/test_cache.py @@ -1,6 +1,6 @@ import pytest -from data_harness.cache import SessionCache +from data_harness.data.cache import SessionCache class TestPutGet: diff --git a/tests/test_cache_extras.py b/tests/test_cache_extras.py index ab219ff..6eec7d2 100644 --- a/tests/test_cache_extras.py +++ b/tests/test_cache_extras.py @@ -7,7 +7,7 @@ import pandas as pd import pytest -from data_harness.cache import SessionCache +from data_harness.data.cache import SessionCache def test_answer_slot_roundtrip(): diff --git a/tests/test_charts.py b/tests/test_charts.py index 6f79bf3..69d8785 100644 --- a/tests/test_charts.py +++ b/tests/test_charts.py @@ -2,10 +2,10 @@ from __future__ import annotations -from data_harness.artifacts import ChartArtifact -from data_harness.cache import SessionCache -from data_harness.result import RunResult, Usage -from data_harness.tools.interpreter import PythonInterpreter +from data_harness.core.artifacts import ChartArtifact +from data_harness.core.result import RunResult, Usage +from data_harness.data.cache import SessionCache +from data_harness.data.tools.interpreter import PythonInterpreter def _interp(tmp_path) -> PythonInterpreter: diff --git a/tests/test_cli.py b/tests/test_cli.py index 23112f6..d5e7399 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -7,7 +7,7 @@ import pandas as pd from data_harness import cli -from data_harness.result import RunResult, Usage +from data_harness.core.result import RunResult, Usage def _fake_result(text="the answer is 6", value=6, charts=None) -> RunResult: diff --git a/tests/test_compaction.py b/tests/test_compaction.py new file mode 100644 index 0000000..96abebf --- /dev/null +++ b/tests/test_compaction.py @@ -0,0 +1,594 @@ +"""Phase 5: compaction and the error taxonomy. + +Before this the only context management was `max_turns`, which is a wall +rather than a strategy: a run needing thirty turns failed at twenty-five +having paid for all of them. + +The tests that matter most are the two safety properties. A cut that separates +an assistant tool call from its result produces a transcript the provider +rejects outright, and it is the first bug everyone writes here. And nothing +may actually be deleted, because a compaction that cannot be undone is an edit +to history rather than a view of it. +""" + +from __future__ import annotations + +import pytest + +from data_harness.core.compaction import ( + CompactionSettings, + estimate_tokens, + find_cut_point, + make_compactor, + maybe_compact, + messages_before, + should_compact, + starts_a_turn, +) +from data_harness.core.exceptions import ( + ConfigurationError, + DataHarnessError, + ExecutionError, + MaxTurnsExceeded, + ProviderError, + SubagentRecursionError, + ToolNotFoundError, +) +from data_harness.core.hooks import HookError +from data_harness.core.session import MessageEntry, Session, SessionStoreError +from data_harness.data.harness import Harness +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import ( + Message, + TextBlock, + ToolResultBlock, + ToolUseBlock, +) + + +def say(text: str, role: str = "user") -> Message: + return Message(role=role, content=[TextBlock(text=text)]) + + +def texts(messages: list[Message]) -> list[str]: + return [m.content[0].text for m in messages] + + +def small_settings(**overrides) -> CompactionSettings: + base = dict(context_window=1000, reserve_tokens=200, keep_recent_tokens=100) + base.update(overrides) + return CompactionSettings(**base) + + +def bulk(label: str, tokens: int) -> Message: + """A message of roughly ``tokens`` tokens.""" + return say(f"{label}:" + "x" * (tokens * 4)) + + +# ── estimation and the trigger ────────────────────────────────────────────── + + +def test_token_estimate_counts_every_block_kind(): + messages = [ + say("a" * 40), + Message( + role="assistant", + content=[ToolUseBlock(tool_use_id="t1", tool_name="echo", tool_input={})], + ), + Message( + role="user", + content=[ToolResultBlock(tool_use_id="t1", content="b" * 40)], + ), + ] + # Not exact by design; what matters is that nothing is counted as zero. + assert estimate_tokens(messages) >= 20 + + +def test_the_trigger_leaves_room_for_a_reply(): + settings = small_settings() # window 1000, reserve 200 + assert should_compact(801, settings) + assert not should_compact(800, settings) + + +def test_settings_that_could_never_trigger_are_rejected(): + """Keeping more than the trigger point would compact forever.""" + with pytest.raises(ValueError, match="keep_recent_tokens"): + CompactionSettings( + context_window=1000, reserve_tokens=200, keep_recent_tokens=900 + ) + + +# ── the cut lands on a turn boundary ──────────────────────────────────────── + + +def test_a_user_message_starts_a_turn(): + session = Session() + session.append_message(say("q1")) + entries = session.context_entries() + assert starts_a_turn(entries, 0) + + +def test_a_tool_result_message_does_not_start_a_turn(): + """It is a user message, but cutting before it orphans the call above.""" + session = Session() + session.append_message( + Message( + role="user", + content=[ToolResultBlock(tool_use_id="t1", content="r")], + ) + ) + assert not starts_a_turn(session.context_entries(), 0) + + +def test_an_assistant_message_does_not_start_a_turn(): + """Cutting here would open the conversation with an assistant message.""" + session = Session() + session.append_message(say("a", "assistant")) + assert not starts_a_turn(session.context_entries(), 0) + + +def test_the_cut_never_splits_a_tool_call_from_its_result(): + """The bug everyone writes. A split pair is rejected by the provider.""" + session = Session() + session.append_message(bulk("q1", 200)) + session.append_message( + Message( + role="assistant", + content=[ToolUseBlock(tool_use_id="t1", tool_name="echo", tool_input={})], + ) + ) + session.append_message( + Message(role="user", content=[ToolResultBlock(tool_use_id="t1", content="r")]) + ) + session.append_message(bulk("q2", 200)) + + entries = session.context_entries() + cut_id = find_cut_point(entries, small_settings()) + + kept = [e for e in entries if e.id == cut_id] + assert kept, "expected a cut point" + entry = kept[0] + assert isinstance(entry, MessageEntry) + assert entry.message.role == "user" + assert not any(isinstance(b, ToolResultBlock) for b in entry.message.content) + + +def test_no_cut_is_made_when_there_is_nowhere_safe(): + """One indivisible turn is better left alone than cut badly.""" + session = Session() + session.append_message(bulk("only question", 500)) + session.append_message(say("thinking", "assistant")) + + assert find_cut_point(session.context_entries(), small_settings()) is None + + +def test_nothing_is_cut_when_the_boundary_is_the_first_entry(): + session = Session() + session.append_message(bulk("q1", 500)) + + assert find_cut_point(session.context_entries(), small_settings()) is None + + +def test_messages_before_the_cut_are_what_gets_summarised(): + session = Session() + session.append_message(say("old one")) + session.append_message(say("old two", "assistant")) + cut = session.append_message(say("recent")) + + replaced = messages_before(session.context_entries(), cut) + + assert texts(replaced) == ["old one", "old two"] + + +# ── compacting a session ──────────────────────────────────────────────────── + + +def test_compaction_replaces_old_turns_with_a_summary(): + session = Session() + for i in range(6): + session.append_message(bulk(f"q{i}", 150)) + session.append_message(bulk(f"a{i}", 150)) + + entry_id = maybe_compact(session, lambda msgs: "they talked", small_settings()) + + assert entry_id is not None + context = texts(session.build_context()) + assert context[0].startswith("Summary of the earlier conversation:") + assert "they talked" in context[0] + assert len(context) < 12 + + +def test_the_summariser_sees_exactly_what_it_replaces(): + seen: list[list[str]] = [] + session = Session() + for i in range(6): + session.append_message(bulk(f"q{i}", 150)) + session.append_message(bulk(f"a{i}", 150)) + + def summarize(messages): + seen.append([m.content[0].text.split(":")[0] for m in messages]) + return "summary" + + maybe_compact(session, summarize, small_settings()) + + kept = {t.split(":")[0] for t in texts(session.build_context())[1:]} + assert seen, "summariser was not called" + assert not (set(seen[0]) & kept), "summarised and kept overlap" + + +def test_a_small_conversation_is_left_alone(): + session = Session() + session.append_message(say("q1")) + + assert maybe_compact(session, lambda m: "never", small_settings()) is None + assert texts(session.build_context()) == ["q1"] + + +def test_nothing_is_deleted_by_compacting(): + """A compaction that cannot be undone is an edit, not a view.""" + session = Session() + for i in range(6): + session.append_message(bulk(f"q{i}", 150)) + session.append_message(bulk(f"a{i}", 150)) + before = len(session.store.entries()) + + maybe_compact(session, lambda m: "summary", small_settings()) + + stored = [e for e in session.store.entries() if isinstance(e, MessageEntry)] + assert len(stored) == before + assert any(m.message.content[0].text.startswith("q0") for m in stored) + + +def test_a_compaction_can_be_undone_by_moving_the_leaf(): + session = Session() + for i in range(6): + session.append_message(bulk(f"q{i}", 150)) + last_before = session.leaf_id + + maybe_compact(session, lambda m: "summary", small_settings()) + assert "Summary" in texts(session.build_context())[0] + + session.move_to(last_before) + restored = texts(session.build_context()) + assert not restored[0].startswith("Summary") + assert len(restored) == 6 + + +# ── the harness consults a compactor ──────────────────────────────────────── + + +def test_the_harness_compacts_between_turns(tmp_path): + """A long transcript, not a long tool output. + + A big tool result never reaches the transcript: the cache stores it under + a handle and the model sees a snapshot. That is the design working, and it + is why compaction costs this library less than it costs a coding agent. + So the pressure here comes from the conversation itself. + """ + calls: list[int] = [] + + def summarize(messages): + calls.append(len(messages)) + return "earlier work" + + tiny = CompactionSettings( + context_window=400, reserve_tokens=100, keep_recent_tokens=40 + ) + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.text("a" * 900), + FakeAdapter.text("b" * 900), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[], + run_dir=str(tmp_path), + compactor=make_compactor(summarize, tiny), + ) + + harness.run("q" * 900) + harness.ask("second question") + harness.ask("third question") + + assert calls, "compactor never fired" + assert any("earlier work" in str(m.content) for m in harness.messages) + + +def test_the_working_copy_is_rebuilt_from_the_session_after_compacting(tmp_path): + """Not edited in place: the conversation stays derived from the log.""" + tiny = CompactionSettings( + context_window=400, reserve_tokens=100, keep_recent_tokens=40 + ) + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.text("a" * 900), + FakeAdapter.text("b" * 900), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[], + run_dir=str(tmp_path), + compactor=make_compactor(lambda m: "summary", tiny), + ) + + harness.run("q" * 900) + harness.ask("second question") + harness.ask("third question") + + assert harness.messages == harness.session.build_context() + + +def test_a_harness_without_a_compactor_never_compacts(tmp_path): + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.run("go") + + from data_harness.core.session import CompactionEntry + + assert not [ + e for e in harness.session.store.entries() if isinstance(e, CompactionEntry) + ] + + +# ── the error taxonomy ────────────────────────────────────────────────────── + + +@pytest.mark.parametrize( + ("error", "code"), + [ + (MaxTurnsExceeded(3), "max_turns_exceeded"), + (ToolNotFoundError("x"), "tool_not_found"), + (SubagentRecursionError("x"), "subagent_recursion"), + (ProviderError("rate limited"), "provider_error"), + (ExecutionError("timed out"), "execution_error"), + (ConfigurationError("bad wiring"), "configuration_error"), + (SessionStoreError("not_found", "gone"), "not_found"), + ], +) +def test_every_error_carries_a_stable_code(error, code): + """Codes are the API. Messages are not, and may be reworded.""" + assert error.code == code + + +def test_every_library_error_shares_one_base(): + """`except DataHarnessError` catches the library's failures and no others.""" + for error in [ + MaxTurnsExceeded(1), + ToolNotFoundError("x"), + SubagentRecursionError("x"), + ProviderError("x"), + ExecutionError("x"), + ConfigurationError("x"), + SessionStoreError("not_found", "x"), + ]: + assert isinstance(error, DataHarnessError) + + +def test_a_hook_error_is_part_of_the_taxonomy(): + def bad(event): + raise ValueError("boom") + + error = HookError("event", bad, ValueError("boom")) + assert isinstance(error, DataHarnessError) + assert error.code == "hook_error" + + +def test_the_legacy_base_classes_still_hold(): + """Callers caught these as RuntimeError/KeyError before the taxonomy.""" + assert isinstance(MaxTurnsExceeded(1), RuntimeError) + assert isinstance(SubagentRecursionError("x"), RuntimeError) + assert isinstance(ToolNotFoundError("x"), KeyError) + + +def test_a_tool_handler_bug_is_not_caught_by_the_library_base(): + """The base must not be so broad it swallows the caller's own mistakes.""" + assert not isinstance(ValueError("caller bug"), DataHarnessError) + + +# ── review: a failing compactor must not cost the run ─────────────────────── + + +def test_a_compactor_that_fails_does_not_fail_the_run(tmp_path): + """Summarising calls a model, so failing here is ordinary. + + Compaction is an optimisation. Carrying on with the full context is + strictly better than losing a run that was otherwise fine, and if the + context really was too big the provider call fails next and reports it + where it belongs. + """ + + def broken(session): + raise ValueError("summariser API down") + + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.text("a" * 900), + FakeAdapter.text("b" * 900), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[], + run_dir=str(tmp_path), + compactor=broken, + ) + + harness.run("q" * 900) + harness.ask("second question") + result = harness.ask_result("third question") + + assert result.status == "success" + assert result.text == "done" + + +def test_a_failed_compaction_is_recorded(tmp_path): + """A run that quietly stopped compacting must be explainable afterwards.""" + + def broken(session): + raise ValueError("summariser API down") + + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("ok")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + compactor=broken, + ) + harness.run("go") + + failures = harness.session.custom_entries("compaction_failed") + assert len(failures) == 1 + assert "summariser API down" in failures[0].data["error"] + + +def test_a_cut_never_leaves_unpaired_tool_blocks(tmp_path): + """Parallel tool calls, mixed blocks, and assistant-first conversations. + + Each of these produces a transcript a provider rejects if the cut lands + wrong, and none of them is exotic. + """ + session = Session() + session.append_message(bulk("q1", 200)) + session.append_message(bulk("a1", 200)) + session.append_message(bulk("q2", 200)) + session.append_message( + Message( + role="assistant", + content=[ + ToolUseBlock(tool_use_id="t1", tool_name="e", tool_input={}), + ToolUseBlock(tool_use_id="t2", tool_name="e", tool_input={}), + ], + ) + ) + session.append_message( + Message( + role="user", + content=[ + ToolResultBlock(tool_use_id="t1", content="r1"), + ToolResultBlock(tool_use_id="t2", content="r2"), + ], + ) + ) + session.append_message(bulk("q3", 200)) + + maybe_compact(session, lambda m: "summary", small_settings()) + context = session.build_context() + + uses = { + b.tool_use_id for m in context for b in m.content if isinstance(b, ToolUseBlock) + } + results = { + b.tool_use_id + for m in context + for b in m.content + if isinstance(b, ToolResultBlock) + } + assert uses == results + assert context[0].role == "user" + assert all(m.content for m in context) + + +def test_compacting_every_turn_keeps_the_context_bounded(): + """Thirty turns, compacting each one, must not grow without bound.""" + settings = small_settings() + session = Session() + for i in range(30): + session.append_message(bulk(f"q{i}", 60)) + session.append_message(bulk(f"a{i}", 60)) + maybe_compact(session, lambda m: "summary", settings) + + tokens = estimate_tokens(session.build_context()) + assert tokens <= settings.context_window - settings.reserve_tokens + + +def test_repeated_compaction_leaves_one_summary(): + session = Session() + for i in range(4): + session.append_message(bulk(f"q{i}", 200)) + session.append_message(bulk(f"a{i}", 200)) + assert maybe_compact(session, lambda m: "summary one", small_settings()) + + for i in range(4, 8): + session.append_message(bulk(f"q{i}", 200)) + session.append_message(bulk(f"a{i}", 200)) + assert maybe_compact(session, lambda m: "summary two", small_settings()) + + summaries = [t for t in texts(session.build_context()) if t.startswith("Summary")] + assert len(summaries) == 1 + assert "summary two" in summaries[0] + + +def test_nothing_is_compacted_when_the_cut_leaves_no_messages(tmp_path): + """A safe cut can still have no *messages* before it. + + Entries are not all messages: a turn record before the first user message + makes the boundary land at index 1 with nothing summarisable behind it. + Compacting there would append a summary of nothing and count as progress. + """ + # Small enough that this conversation is genuinely over the trigger; + # otherwise the run never reaches the guard and the test proves nothing. + settings = CompactionSettings( + context_window=400, reserve_tokens=100, keep_recent_tokens=40 + ) + session = Session() + session.append_turn(turn=0, input_tokens=1) + session.append_message(say("q" * 900)) + session.append_message(say("a" * 900, "assistant")) + session.append_message(say("hi")) + + assert should_compact(estimate_tokens(session.build_context()), settings) + + entries = session.context_entries() + cut = find_cut_point(entries, settings) + + assert cut is not None + assert messages_before(entries, cut) == [] + assert maybe_compact(session, lambda m: "never called", settings) is None + + +def test_a_failed_run_raises_a_provider_error(tmp_path): + """Specifically `ProviderError`, not a bare RuntimeError. + + Telling a rate limit apart from a bug in the caller's own code is the + entire reason the taxonomy exists. + """ + + class Broken(FakeAdapter): + def chat(self, system, messages, tools): + raise RuntimeError("429 rate limited") + + harness = Harness(adapter=Broken([]), system="sys", tools=[], run_dir=str(tmp_path)) + + with pytest.raises(ProviderError) as excinfo: + harness.run("go") + + assert excinfo.value.code == "provider_error" + assert "rate limited" in str(excinfo.value) + # Still a RuntimeError, which is how callers caught this before. + assert isinstance(excinfo.value, RuntimeError) + + +def test_max_turns_does_not_raise_a_provider_error(tmp_path): + """The two failure kinds must stay distinguishable.""" + harness = Harness( + adapter=FakeAdapter([FakeAdapter.tool_use("t1", "missing", {})]), + system="sys", + tools=[], + run_dir=str(tmp_path), + max_turns=1, + ) + + with pytest.raises(MaxTurnsExceeded) as excinfo: + harness.run("go") + + assert not isinstance(excinfo.value, ProviderError) + assert excinfo.value.code == "max_turns_exceeded" diff --git a/tests/test_core_standalone.py b/tests/test_core_standalone.py new file mode 100644 index 0000000..e7d9126 --- /dev/null +++ b/tests/test_core_standalone.py @@ -0,0 +1,171 @@ +"""The boundary is worth something, not just declared. + +Two claims are made by splitting the package into layers, and each is only +credible if it is exercised: + +1. `data_harness.core` runs an agent with no data domain at all. +2. Every pre-layering import path still resolves, to the same object. + +A layering that no one can use without the layer above it is decoration, and a +rename that breaks every downstream import is not a refactor. +""" + +from __future__ import annotations + +import importlib + +import pytest + +from data_harness._legacy_paths import LEGACY_PATHS +from data_harness.core.environment import NullEnvironment, RunEnvironment, RunState +from data_harness.core.loop import AsyncHarness, Harness +from data_harness.llm.testing import FakeAdapter, FakeAsyncAdapter +from data_harness.llm.types import ToolSpec + + +def upper_spec() -> ToolSpec: + return ToolSpec( + name="shout", + description="Upper-case the input.", + input_schema={ + "type": "object", + "properties": {"text": {"type": "string"}}, + "required": ["text"], + }, + handler=lambda text: text.upper(), + ) + + +# ── core runs on its own ──────────────────────────────────────────────────── + + +def test_core_harness_runs_a_tool_loop_with_no_domain(tmp_path): + """No cache, no pandas, no interpreter: still a working agent.""" + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "shout", {"text": "hello"}), + FakeAdapter.text("I shouted it."), + ] + ), + system="sys", + tools=[upper_spec()], + run_dir=str(tmp_path), + ) + + result = harness.run_result("go") + + assert result.status == "success" + assert result.text == "I shouted it." + assert harness.messages[2].content[0].content == "HELLO" + + +@pytest.mark.asyncio +async def test_core_async_harness_runs_with_no_domain(tmp_path): + harness = AsyncHarness( + adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("fine")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + + assert await harness.run("go") == "fine" + + +def test_the_null_environment_contributes_no_run_state(tmp_path): + """A domain-free run reports no handles, no answer, and no charts.""" + result = Harness( + adapter=FakeAdapter([FakeAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ).run_result("go") + + assert result.cache_snapshots == {} + assert result.cache_storage == {} + assert result.value is None + assert result.charts == [] + + +def test_a_custom_environment_shapes_the_transcript_and_the_result(tmp_path): + """The seam is real: a third domain can plug in without touching core. + + This is the actual test of the boundary. If the loop still reached for a + `SessionCache`, there would be no way to write this. + """ + + class ShoutingEnvironment: + """Renders every tool result in caps and records one fixed handle.""" + + def render_tool_output(self, value: object) -> str: + return f"<<{value}>>" + + def capture(self) -> RunState: + return RunState(snapshots={"answer": "42"}, value=42) + + def storage_metadata(self) -> dict[str, dict[str, str]]: + return {"answer": {"location": "memory", "storage_type": "memory"}} + + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "shout", {"text": "hi"}), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[upper_spec()], + run_dir=str(tmp_path), + environment=ShoutingEnvironment(), + ) + + result = harness.run_result("go") + + assert harness.messages[2].content[0].content == "<>" + assert result.value == 42 + assert result.cache_snapshots == {"answer": "42"} + + +def test_null_environment_satisfies_the_protocol(): + assert isinstance(NullEnvironment(), RunEnvironment) + + +def test_the_cache_environment_also_satisfies_it(): + from data_harness.data.environment import CacheEnvironment + + assert isinstance(CacheEnvironment(), RunEnvironment) + + +# ── the compatibility promise ─────────────────────────────────────────────── + + +@pytest.mark.parametrize("legacy", sorted(LEGACY_PATHS)) +def test_legacy_import_path_still_resolves(legacy: str): + assert importlib.import_module(legacy) is not None + + +@pytest.mark.parametrize("legacy,current", sorted(LEGACY_PATHS.items())) +def test_legacy_path_is_the_same_module_not_a_copy(legacy: str, current: str): + """Identity matters, not just importability. + + Re-exporting names into a shim would give a second module object, and a + class reached through the old path would fail an `isinstance` check + against the same class reached through the new one. + """ + assert importlib.import_module(legacy) is importlib.import_module(current) + + +def test_a_class_is_the_same_class_through_either_path(): + from data_harness.loop import Harness as ViaLegacy + + from data_harness.core.loop import Harness as ViaCore + from data_harness.data.harness import Harness as ViaData + + assert ViaLegacy is ViaData + assert issubclass(ViaData, ViaCore) + + +def test_an_unknown_data_harness_module_still_fails(): + """The finder must not swallow genuine typos.""" + with pytest.raises(ModuleNotFoundError): + importlib.import_module("data_harness.no_such_module") diff --git a/tests/test_docs.py b/tests/test_docs.py index ef31270..bdf57f9 100644 --- a/tests/test_docs.py +++ b/tests/test_docs.py @@ -2,10 +2,10 @@ from __future__ import annotations -from data_harness.cache import SessionCache -from data_harness.loop import Harness -from data_harness.testing import FakeAdapter -from data_harness.types import ToolResultBlock +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import ToolResultBlock from examples.advanced_wiring import build_base_tools, load_unemployment_rate from examples.quickstart import build_agent diff --git a/tests/test_eval.py b/tests/test_eval.py index 5edd3d3..12cb49a 100644 --- a/tests/test_eval.py +++ b/tests/test_eval.py @@ -4,7 +4,8 @@ import pandas as pd -from data_harness.artifacts import ChartArtifact +from data_harness.core.artifacts import ChartArtifact +from data_harness.core.result import RunResult, Usage from data_harness.eval import ( EvalCase, EvalReport, @@ -21,8 +22,7 @@ refuses, wtq_row_to_case, ) -from data_harness.result import RunResult, Usage -from data_harness.testing import FakeAdapter +from data_harness.llm.testing import FakeAdapter def _result(text="", value=None, charts=None) -> RunResult: diff --git a/tests/test_format.py b/tests/test_format.py index aad0eef..c170711 100644 --- a/tests/test_format.py +++ b/tests/test_format.py @@ -1,6 +1,6 @@ import pytest -from data_harness.format import format_tool_output +from data_harness.data.format import format_tool_output # ────────────────────────────────────────────── # Phase 0: inline-path tests (no cache needed) diff --git a/tests/test_gaps.py b/tests/test_gaps.py index c91bd39..1278915 100644 --- a/tests/test_gaps.py +++ b/tests/test_gaps.py @@ -9,11 +9,15 @@ import pytest -from data_harness.cache import SessionCache -from data_harness.loop import Harness -from data_harness.providers.base import NormalizedResponse, ProviderAdapter, StopReason -from data_harness.testing import FakeAdapter -from data_harness.types import TextBlock, ToolAnnotations, ToolSpec, ToolUseBlock +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.llm.providers.base import ( + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import TextBlock, ToolAnnotations, ToolSpec, ToolUseBlock def make_text_response( @@ -212,7 +216,7 @@ def test_cache_snapshots_reflect_post_tool_mutation(self, tmp_path): """cache_snapshots in RunResult should capture the state AFTER any tool calls, not the pre-run state. """ - from data_harness.cache import SessionCache + from data_harness.data.cache import SessionCache cache = SessionCache() initial_data = [1, 2, 3] @@ -334,8 +338,8 @@ def test_annotations_set_not_in_provider_dict(self): class TestConnectorBuilderAnnotations: def test_connector_tool_accepts_annotations(self, tmp_path): - from data_harness.agent import Agent - from data_harness.testing import FakeAdapter + from data_harness.app.agent import Agent + from data_harness.llm.testing import FakeAdapter adapter = FakeAdapter([FakeAdapter.text("done")]) agent = Agent(adapter=adapter, system="s", run_dir=str(tmp_path)) @@ -352,8 +356,8 @@ def query_db(sql: str) -> str: assert result.status == "success" def test_connector_tool_annotations_propagate_to_spec(self, tmp_path): - from data_harness.agent import Agent - from data_harness.testing import FakeAdapter + from data_harness.app.agent import Agent + from data_harness.llm.testing import FakeAdapter adapter = FakeAdapter([FakeAdapter.text("done")]) agent = Agent(adapter=adapter, system="s", run_dir=str(tmp_path)) diff --git a/tests/test_hooks.py b/tests/test_hooks.py new file mode 100644 index 0000000..fcac354 --- /dev/null +++ b/tests/test_hooks.py @@ -0,0 +1,667 @@ +"""Phase 4: hooks. + +The loop used to have exactly one extension point, `register_reminder`, which +could append text to the prompt and nothing else. Everything more interesting +was hardcoded: the interpreter approval gate lived inside the dispatcher and +was keyed on the literal string "python_interpreter". + +The test that matters most here is not that hooks fire. It is that the gate is +now built out of the same public mechanism anyone else has, which is what +makes it evidence the mechanism is sufficient rather than decorative. +""" + +from __future__ import annotations + +import pytest + +from data_harness.core.hooks import ( + AfterToolCall, + AfterTurn, + BeforeToolCall, + BeforeTurn, + Block, + HookError, + HookRegistry, + Reminder, + Replace, + Stop, +) +from data_harness.data.harness import AsyncHarness, Harness +from data_harness.llm.testing import FakeAdapter, FakeAsyncAdapter +from data_harness.llm.types import ToolSpec + + +def echo_spec(calls: list | None = None) -> ToolSpec: + def handler(value: str) -> str: + if calls is not None: + calls.append(value) + return value + + return ToolSpec( + name="echo", + description="Echo the input.", + input_schema={ + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, + handler=handler, + ) + + +def tool_then_text(cls=FakeAdapter): + return [ + cls.tool_use("tu_1", "echo", {"value": "hi"}), + cls.text("done"), + ] + + +# ── the registry ──────────────────────────────────────────────────────────── + + +def test_hooks_run_in_registration_order(): + seen: list[str] = [] + registry = HookRegistry() + registry.add(BeforeTurn, lambda e: seen.append("first") or None) + registry.add(BeforeTurn, lambda e: seen.append("second") or None) + + registry.emit(BeforeTurn(turn=1, max_turns=5, messages=[])) + + assert seen == ["first", "second"] + + +def test_every_hook_sees_the_event_even_after_one_decides(): + """A hook that records spend must not be skipped because another stopped. + + Precedence between conflicting decisions belongs to the caller, not to + whichever hook happened to be registered first. + """ + seen: list[str] = [] + registry = HookRegistry() + registry.add(BeforeTurn, lambda e: Stop("out of budget")) + registry.add(BeforeTurn, lambda e: seen.append("still ran") or None) + + decisions = registry.emit(BeforeTurn(turn=1, max_turns=5, messages=[])) + + assert seen == ["still ran"] + assert isinstance(decisions[0], Stop) + + +def test_a_hook_for_another_event_is_not_called(): + registry = HookRegistry() + registry.add(AfterTurn, lambda e: pytest.fail("wrong event")) + + registry.emit(BeforeTurn(turn=1, max_turns=1, messages=[])) + + +def test_a_raising_hook_is_named_not_swallowed(): + """Hooks must not raise, so when one does the error says which.""" + + def badly_behaved(event): + raise ValueError("I should have returned a decision") + + registry = HookRegistry() + registry.add(BeforeTurn, badly_behaved) + + with pytest.raises(HookError) as excinfo: + registry.emit(BeforeTurn(turn=1, max_turns=1, messages=[])) + + assert "badly_behaved" in str(excinfo.value) + assert "BeforeTurn" in str(excinfo.value) + assert isinstance(excinfo.value.__cause__, ValueError) + + +# ── the gate is built from the public mechanism ───────────────────────────── + + +def test_the_approval_gate_is_an_ordinary_hook(tmp_path): + """`code_only` registers a `BeforeToolCall` hook like any other. + + If the gate still lived inside the dispatcher there would be nothing here + to find, and no reason to believe the hook mechanism was enough for it. + """ + harness = Harness( + adapter=FakeAdapter([]), + system="sys", + tools=[], + run_dir=str(tmp_path), + code_only=True, + ) + + assert len(harness.hooks.hooks[BeforeToolCall]) == 1 + + +def test_no_gate_hook_is_registered_when_the_gate_is_off(tmp_path): + harness = Harness( + adapter=FakeAdapter([]), system="sys", tools=[], run_dir=str(tmp_path) + ) + assert BeforeToolCall not in harness.hooks.hooks + + +# ── blocking a tool call ──────────────────────────────────────────────────── + + +def test_a_hook_can_refuse_a_tool_call(tmp_path): + calls: list[str] = [] + harness = Harness( + adapter=FakeAdapter(tool_then_text()), + system="sys", + tools=[echo_spec(calls)], + run_dir=str(tmp_path), + ) + harness.on(BeforeToolCall, lambda e: Block("not allowed right now")) + + harness.run("go") + + assert calls == [] + block = harness.messages[2].content[0] + assert block.content == "not allowed right now" + assert block.is_error is False + + +def test_a_refusal_can_be_marked_an_error_when_it_really_is_one(tmp_path): + harness = Harness( + adapter=FakeAdapter(tool_then_text()), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.on(BeforeToolCall, lambda e: Block("malformed input", is_error=True)) + + harness.run("go") + + assert harness.messages[2].content[0].is_error is True + + +def test_a_hook_sees_what_the_model_asked_for(tmp_path): + seen: list[tuple[str, dict]] = [] + harness = Harness( + adapter=FakeAdapter(tool_then_text()), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.on(BeforeToolCall, lambda e: seen.append((e.tool_name, e.tool_input))) + + harness.run("go") + + assert seen == [("echo", {"value": "hi"})] + + +# ── rewriting a tool result ───────────────────────────────────────────────── + + +def test_a_hook_can_rewrite_a_tool_result(tmp_path): + harness = Harness( + adapter=FakeAdapter(tool_then_text()), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.on(AfterToolCall, lambda e: Replace(f"[redacted: {len(e.result.content)}]")) + + harness.run("go") + + assert harness.messages[2].content[0].content == "[redacted: 2]" + + +def test_after_tool_call_sees_the_real_result(tmp_path): + seen: list[str] = [] + harness = Harness( + adapter=FakeAdapter(tool_then_text()), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.on(AfterToolCall, lambda e: seen.append(e.result.content)) + + harness.run("go") + + assert seen == ["hi"] + + +# ── reminders ─────────────────────────────────────────────────────────────── + + +def test_a_hook_can_append_a_reminder(tmp_path): + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.on(BeforeTurn, lambda e: Reminder(f"turn {e.turn} of {e.max_turns}")) + + harness.run("go") + + assert "turn 1 of 25" in harness.messages[0].content[-1].text + + +def test_the_legacy_reminder_api_still_works(tmp_path): + """`register_reminder` predates hooks and must keep working.""" + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.register_reminder(lambda turn, max_turns: "legacy reminder") + harness.on(BeforeTurn, lambda e: Reminder("hook reminder")) + + harness.run("go") + + text = harness.messages[0].content[-1].text + assert "legacy reminder" in text + assert "hook reminder" in text + + +# ── stopping a run ────────────────────────────────────────────────────────── + + +def test_an_after_turn_hook_can_stop_the_run(tmp_path): + """Where a spend cap belongs: the tokens are counted, so the decision is + made on real numbers rather than a guess before the call.""" + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAdapter.text("never reached"), + ] + ), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.on(AfterTurn, lambda e: Stop("budget exhausted") if e.turn >= 1 else None) + + result = harness.run_result("go") + + assert result.status == "success" + assert result.turns == 1 + assert len(harness.messages) == 3 # user, assistant, tool results + + +@pytest.mark.asyncio +async def test_after_turn_sees_the_tokens_already_spent(tmp_path): + """Cumulative, and real. + + The old version asserted (0, 0) against an adapter that reports (0, 0), + so it could not tell cumulative from per-turn from hardcoded. A mutation + making the event always report zero survived the whole suite. + """ + seen: list[tuple[int, int]] = [] + harness = AsyncHarness( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAsyncAdapter.text("done"), + ] + ), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.on(AfterTurn, lambda e: seen.append((e.input_tokens, e.output_tokens))) + + await harness.run("go") + + # FakeAsyncAdapter reports 10/5 a turn, and the totals accumulate. + assert seen == [(10, 5), (20, 10)] + + +def test_a_before_turn_stop_spends_nothing(tmp_path): + adapter = FakeAdapter([FakeAdapter.text("should not be called")]) + harness = Harness(adapter=adapter, system="sys", tools=[], run_dir=str(tmp_path)) + harness.on(BeforeTurn, lambda e: Stop("refused before starting")) + + result = harness.run_result("go") + + assert adapter.calls == [] + assert result.turns == 0 + + +def test_a_stop_does_not_leak_into_the_next_run(tmp_path): + stops = iter([True, False]) + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("one"), FakeAdapter.text("two")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.on(BeforeTurn, lambda e: Stop("first only") if next(stops, False) else None) + + harness.run_result("first") + second = harness.run_result("second") + + assert second.text == "one" + + +# ── both drivers ──────────────────────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_hooks_work_the_same_under_the_async_driver(tmp_path): + calls: list[str] = [] + harness = AsyncHarness( + adapter=FakeAsyncAdapter(tool_then_text(FakeAsyncAdapter)), + system="sys", + tools=[echo_spec(calls)], + run_dir=str(tmp_path), + ) + harness.on(BeforeToolCall, lambda e: Block("no")) + + await harness.run("go") + + assert calls == [] + assert harness.messages[2].content[0].content == "no" + + +@pytest.mark.asyncio +async def test_hooks_apply_to_streamed_runs(tmp_path): + harness = AsyncHarness( + adapter=FakeAsyncAdapter(tool_then_text(FakeAsyncAdapter)), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.on(AfterToolCall, lambda e: Replace("rewritten")) + + from data_harness.llm.streaming import ToolResultEvent + + events = [e async for e in harness.run_stream("go")] + tool_events = [e for e in events if isinstance(e, ToolResultEvent)] + + assert [e.content for e in tool_events] == ["rewritten"] + + +# ── a hook registry can be supplied up front ──────────────────────────────── + + +def test_a_registry_can_be_passed_to_the_constructor(tmp_path): + """So an application can build its policy once and reuse it.""" + registry = HookRegistry() + registry.add(BeforeToolCall, lambda e: Block("policy says no")) + + harness = Harness( + adapter=FakeAdapter(tool_then_text()), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + hooks=registry, + ) + harness.run("go") + + assert harness.messages[2].content[0].content == "policy says no" + + +# ── review round 1: what the mutation pass showed was unguarded ───────────── + + +def test_an_after_turn_stop_keeps_the_model_answer(tmp_path): + """Stopping after a turn must not discard what that turn produced. + + A mutation returning empty text here survived the whole suite, and the + claim that the answer still stands was the headline of the feature. + """ + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("the answer"), FakeAdapter.text("more")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.on(AfterTurn, lambda e: Stop("capped")) + + result = harness.run_result("go") + + assert result.text == "the answer" + assert result.status == "success" + + +def test_a_before_turn_stop_is_a_success_not_an_error(tmp_path): + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("never")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.on(BeforeTurn, lambda e: Stop("refused")) + + result = harness.run_result("go") + + assert result.status == "success" + assert result.error is None + + +def test_a_mid_run_stop_keeps_the_usage_already_spent(tmp_path): + """Stopping at turn 3 must still report what turns 1 and 2 cost.""" + harness = AsyncHarness( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "a"}), + FakeAsyncAdapter.tool_use("tu_2", "echo", {"value": "b"}), + FakeAsyncAdapter.text("never reached"), + ] + ), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.on(BeforeTurn, lambda e: Stop("capped") if e.turn == 3 else None) + + result = run_sync(harness.run_result("go")) + + assert result.turns == 2 + assert result.usage.input_tokens == 20 + assert result.usage.output_tokens == 10 + + +def run_sync(coro): + from data_harness.core.loop import run_coroutine_blocking + + return run_coroutine_blocking(coro) + + +# ── why a run stopped is recorded ─────────────────────────────────────────── + + +def test_a_capped_run_says_so(tmp_path): + """A capped run was otherwise indistinguishable from an empty answer.""" + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("never")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.on(BeforeTurn, lambda e: Stop("out of budget")) + + result = harness.run_result("go") + + assert result.stopped_by == "out of budget" + + +def test_an_after_turn_stop_also_says_why(tmp_path): + """Both stop paths record it; a mutation on this one survived without it.""" + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("the answer"), FakeAdapter.text("more")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.on(AfterTurn, lambda e: Stop("spend cap reached")) + + result = harness.run_result("go") + + assert result.stopped_by == "spend cap reached" + assert result.text == "the answer" + + +def test_an_ordinary_run_is_not_marked_as_stopped(tmp_path): + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + assert harness.run_result("go").stopped_by is None + + +# ── a raising hook fails the run instead of vanishing with it ─────────────── + + +def test_a_raising_hook_still_produces_a_result(tmp_path): + """The tokens were billed; losing the RunResult loses the accounting.""" + + def bad(event): + if event.turn == 2: + raise ValueError("boom") + return None + + harness = AsyncHarness( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAsyncAdapter.text("done"), + ] + ), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.on(BeforeTurn, bad) + + with pytest.raises(HookError): + run_sync(harness.run_result("go")) + + result = harness.last_result + assert result is not None + assert result.status == "error" + assert "boom" in (result.error or "") + # Turn 1 completed and was billed before the hook broke. + assert result.usage.input_tokens == 10 + + +# ── a supplied registry is not written to ─────────────────────────────────── + + +def test_a_shared_registry_does_not_leak_one_harness_policy_into_another(tmp_path): + """The gate is added by the harness, so it must not land in the caller's + registry and govern every other harness sharing it.""" + shared = HookRegistry() + + gated = Harness( + adapter=FakeAdapter([]), + system="sys", + tools=[], + run_dir=str(tmp_path), + code_only=True, + hooks=shared, + ) + ungated = Harness( + adapter=FakeAdapter(tool_then_text()), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + hooks=shared, + ) + + assert BeforeToolCall not in shared.hooks + assert BeforeToolCall in gated.hooks.hooks + assert BeforeToolCall not in ungated.hooks.hooks + + calls: list[str] = [] + ungated.tools[:] = [echo_spec(calls)] + ungated.run("go") + assert calls == ["hi"] + + +def test_copy_is_independent(): + original = HookRegistry() + original.add(BeforeTurn, lambda e: None) + + duplicate = original.copy() + duplicate.add(BeforeTurn, lambda e: None) + duplicate.add(AfterTurn, lambda e: None) + + assert len(original.hooks[BeforeTurn]) == 1 + assert AfterTurn not in original.hooks + + +# ── hooks see every tool outcome, not only the happy one ──────────────────── + + +def test_a_failing_tool_result_reaches_after_tool_call(tmp_path): + """A redaction hook has to see failures: an exception repr can carry a + connection string, which is exactly what needs redacting.""" + + def explode(value: str) -> str: + raise RuntimeError("postgres://user:hunter2@host/db") + + spec = ToolSpec( + name="echo", + description="Fail.", + input_schema={ + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, + handler=explode, + ) + harness = Harness( + adapter=FakeAdapter(tool_then_text()), + system="sys", + tools=[spec], + run_dir=str(tmp_path), + ) + harness.on( + AfterToolCall, + lambda e: Replace("[redacted]", is_error=True) if e.result.is_error else None, + ) + + harness.run("go") + + block = harness.messages[2].content[0] + assert block.content == "[redacted]" + assert block.is_error is True + + +def test_a_call_to_an_unknown_tool_reaches_the_hooks(tmp_path): + """A policy hook needs to see attempts on tools that do not exist.""" + seen: list[str] = [] + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "no_such_tool", {}), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.on(BeforeToolCall, lambda e: seen.append(e.tool_name)) + + harness.run("go") + + assert seen == ["no_such_tool"] + assert harness.messages[2].content[0].is_error is True + + +# ── the agent exposes hooks, since it builds a fresh harness per run ──────── + + +def test_an_agent_applies_its_hooks_to_every_run(tmp_path): + from data_harness.app.agent import Agent + + agent = Agent( + adapter=FakeAdapter([FakeAdapter.text("one"), FakeAdapter.text("two")]), + system="sys", + run_dir=str(tmp_path), + ) + seen: list[int] = [] + agent.on(BeforeTurn, lambda e: seen.append(e.turn)) + + agent.run("first") + agent.run("second") + + # A hook put on a harness would be gone by the second run. + assert seen == [1, 1] diff --git a/tests/test_integration.py b/tests/test_integration.py index 5d879b8..54b5e48 100644 --- a/tests/test_integration.py +++ b/tests/test_integration.py @@ -8,14 +8,18 @@ import pandas as pd -from data_harness.cache import SessionCache -from data_harness.loop import Harness -from data_harness.providers.base import NormalizedResponse, ProviderAdapter, StopReason -from data_harness.tools.connectors import ConnectorRegistry -from data_harness.tools.interpreter import PythonInterpreter -from data_harness.tools.subagent import make_subagent_spec -from data_harness.tools.variables import make_list_variables_spec -from data_harness.types import ( +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.data.tools.connectors import ConnectorRegistry +from data_harness.data.tools.interpreter import PythonInterpreter +from data_harness.data.tools.subagent import make_subagent_spec +from data_harness.data.tools.variables import make_list_variables_spec +from data_harness.llm.providers.base import ( + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.types import ( Message, TextBlock, ToolResultBlock, @@ -179,10 +183,10 @@ def test_4_turn_scripted_flow(self, tmp_path): assert len(set(systems)) == 1 # Invariant 5: cache_control only in adapter-bound payloads, not harness objects - for msg in harness._messages: + for msg in harness.messages: for block in msg.content: assert not hasattr(block, "cache_control") - for tool in harness._tools: + for tool in harness.tools: assert not hasattr(tool, "cache_control") def test_tool_use_ordering_invariant(self, tmp_path): diff --git a/tests/test_layers.py b/tests/test_layers.py new file mode 100644 index 0000000..3f3d03d --- /dev/null +++ b/tests/test_layers.py @@ -0,0 +1,172 @@ +"""The layer boundaries, enforced. + +`data_harness` is layered llm -> core -> data -> app. A layer may import the +ones below it and nothing above. Without a test this is a paragraph in a +README that decays on the first deadline, which is how the package became one +flat namespace in the first place. + +The check is static, over the AST, not by importing. Runtime import checks +would pass vacuously here: nearly every heavyweight dependency in this package +is already behind a function-local import, so `import data_harness.core` pulls +in no pandas today even though the source referred to `SessionCache` all over. +Static analysis sees the reference regardless of where it sits. +""" + +from __future__ import annotations + +import ast +import pathlib + +import pytest + +PACKAGE = pathlib.Path(__file__).resolve().parent.parent / "data_harness" + +#: bottom to top. A layer may import itself and anything earlier. +LAYERS = ("llm", "core", "data", "app") +RANK = {name: index for index, name in enumerate(LAYERS)} + +#: Not layered. `eval` is a research harness that sits beside the stack, and +#: the two package-level modules exist to serve the legacy import paths. +UNLAYERED = {"eval", "_legacy_paths", "__init__"} + + +def _layer_of(module: str) -> str | None: + """`data_harness.core.loop` -> `core`. ``None`` for unlayered modules.""" + parts = module.split(".") + if len(parts) < 2 or parts[0] != "data_harness": + return None + return parts[1] if parts[1] in RANK else None + + +def _module_name(path: pathlib.Path) -> str: + relative = path.relative_to(PACKAGE.parent).with_suffix("") + name = ".".join(relative.parts) + return name.removesuffix(".__init__") + + +def _imports(path: pathlib.Path) -> set[str]: + found: set[str] = set() + for node in ast.walk(ast.parse(path.read_text())): + if isinstance(node, ast.ImportFrom) and node.module: + if node.module.startswith("data_harness"): + found.add(node.module) + elif isinstance(node, ast.Import): + for alias in node.names: + if alias.name.startswith("data_harness"): + found.add(alias.name) + return found + + +def _source_files() -> list[pathlib.Path]: + return [p for p in sorted(PACKAGE.rglob("*.py")) if "__pycache__" not in p.parts] + + +def test_the_package_is_actually_layered(): + """Guards the guard: if the directories vanish, the rest is vacuous.""" + present = { + p.name for p in PACKAGE.iterdir() if p.is_dir() and p.name != "__pycache__" + } + assert set(LAYERS) <= present + + +@pytest.mark.parametrize("path", _source_files(), ids=_module_name) +def test_module_only_imports_its_own_layer_or_below(path: pathlib.Path): + """No module may import from a layer above its own. + + `core` importing `data` is the specific violation this was written for: + the loop used to construct a `SessionCache` and call `format_tool_output` + directly, which made the harness untestable and unusable without pandas. + It asks a `RunEnvironment` now. + """ + module = _module_name(path) + layer = _layer_of(module) + if layer is None: + return + + violations = [] + for imported in sorted(_imports(path)): + target = _layer_of(imported) + if target is not None and RANK[target] > RANK[layer]: + violations.append(f"{layer}/{module} -> {target}/{imported}") + + assert not violations, "layer violation:\n " + "\n ".join(violations) + + +#: Domain names the lower layers must not use in code. Prose is fine and +#: often necessary: explaining *why* the boundary exists means naming what is +#: on the other side of it. +BANNED_IN_CODE = ("SessionCache", "format_tool_output") + +#: Third-party packages the lower layers must not require, even lazily. A +#: function-local `import pandas` is still a dependency; it just fails later. +BANNED_DEPENDENCIES = ("pandas", "numpy", "duckdb", "pyarrow") + + +def _code_identifiers(path: pathlib.Path) -> set[str]: + """Identifiers referenced in code. Excludes docstrings and comments.""" + names: set[str] = set() + for node in ast.walk(ast.parse(path.read_text())): + if isinstance(node, ast.Name): + names.add(node.id) + elif isinstance(node, ast.Attribute): + names.add(node.attr) + elif isinstance(node, ast.alias): + names.add(node.asname or node.name.split(".")[0]) + elif isinstance(node, ast.arg) and node.annotation is None: + names.add(node.arg) + return names + + +@pytest.mark.parametrize("layer", ["llm", "core"]) +def test_lower_layers_do_not_name_the_data_domain(layer: str): + """The loop must not reach for a `SessionCache`, however indirectly. + + It used to construct one and call `format_tool_output` on it, which is why + the harness could not be exercised or reused without the data stack. It + asks a `RunEnvironment` now. + """ + offenders = [] + for path in _source_files(): + if _layer_of(_module_name(path)) != layer: + continue + used = _code_identifiers(path) + for name in BANNED_IN_CODE: + if name in used: + offenders.append(f"{_module_name(path)} uses {name!r}") + assert not offenders, "\n".join(offenders) + + +@pytest.mark.parametrize("layer", ["llm", "core"]) +def test_lower_layers_do_not_depend_on_the_data_stack(layer: str): + """No pandas in `llm` or `core`, not even behind a lazy import. + + This is the claim that makes the boundary worth having, so it is checked + statically: at runtime almost every heavy import in this package is + function-local, so nothing would fail until the branch was reached. + """ + offenders = [] + for path in _source_files(): + if _layer_of(_module_name(path)) != layer: + continue + for node in ast.walk(ast.parse(path.read_text())): + modules = [] + if isinstance(node, ast.Import): + modules = [a.name for a in node.names] + elif isinstance(node, ast.ImportFrom) and node.module: + modules = [node.module] + for module in modules: + if module.split(".")[0] in BANNED_DEPENDENCIES: + offenders.append(f"{_module_name(path)} imports {module!r}") + assert not offenders, "\n".join(offenders) + + +def test_every_layer_says_what_it_is_for(): + """Each layer package documents its own boundary, next to the code.""" + for layer in LAYERS: + docstring = ast.get_docstring( + ast.parse((PACKAGE / layer / "__init__.py").read_text()) + ) + assert docstring, f"{layer} has no docstring" + assert "import" in docstring.lower(), ( + f"{layer}'s docstring does not state what it may import" + ) diff --git a/tests/test_logger.py b/tests/test_logger.py index c759aa5..d323b72 100644 --- a/tests/test_logger.py +++ b/tests/test_logger.py @@ -3,9 +3,9 @@ import pytest -from data_harness.logger import log_turn, setup_logger -from data_harness.providers.base import NormalizedResponse, StopReason -from data_harness.types import Message, TextBlock, ToolResultBlock +from data_harness.core.logger import log_turn, setup_logger +from data_harness.llm.providers.base import NormalizedResponse, StopReason +from data_harness.llm.types import Message, TextBlock, ToolResultBlock def make_response(text="OK"): diff --git a/tests/test_loop.py b/tests/test_loop.py index f93aa90..e19a334 100644 --- a/tests/test_loop.py +++ b/tests/test_loop.py @@ -8,11 +8,15 @@ import pytest -from data_harness.cache import SessionCache -from data_harness.exceptions import MaxTurnsExceeded -from data_harness.loop import Harness -from data_harness.providers.base import NormalizedResponse, ProviderAdapter, StopReason -from data_harness.types import ( +from data_harness.core.exceptions import MaxTurnsExceeded +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.llm.providers.base import ( + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.types import ( Message, TextBlock, ToolResultBlock, @@ -325,8 +329,8 @@ def test_adapter_input_immutability(self, tmp_path): harness.run("go") # The harness internal messages should not be mutated by adapter calls # We verify by checking stored messages are structurally sound - assert harness._messages is not None - for m in harness._messages: + assert harness.messages is not None + for m in harness.messages: assert m.role in ("user", "assistant") def test_visible_tool_flip_reflects_in_next_call(self, tmp_path): diff --git a/tests/test_loop_protocol.py b/tests/test_loop_protocol.py new file mode 100644 index 0000000..cdecb91 --- /dev/null +++ b/tests/test_loop_protocol.py @@ -0,0 +1,806 @@ +"""The effect protocol, and the invariants mutation testing showed were loose. + +Independent reviews mutated `loop.py` and `agent.py` and found changes the +suite did not notice: the approval gate's scope, whether a blocked call counts +as an error, `_stamp` writing through to `last_result`, latency actually being +measured, the async replay's offload, an errored run's text, and `last_result` +being cleared when a new run starts. Each has a test here. + +Two of the reviews' findings were about tests in this file rather than the +code, and are worth remembering: + +- A probe for the ambient-event-loop bug must run in a *subprocess*. The bug + only fires on the main thread of a process whose loop has never been set, so + an in-process version passes against the bug it exists to catch. +- Answering `ToolFinished` with a value is harmless, not a desync: each + `yield` receives its own send value. The ordering of effects is the real + invariant, and that is what is pinned. +""" + +from __future__ import annotations + +import asyncio +import json +import subprocess +import sys +import textwrap +import time +from pathlib import Path + +import pytest + +from data_harness.app.agent import Agent, AsyncAgent +from data_harness.data.harness import ( + AsyncHarness, + CallProvider, + CallTool, + Failed, + Harness, + Ok, + ToolFinished, + run_coroutine_blocking, +) +from data_harness.llm.streaming import ContentBlockDeltaEvent, TextDelta +from data_harness.llm.testing import FakeAdapter, FakeAsyncAdapter +from data_harness.llm.types import ToolSpec + + +def echo_spec() -> ToolSpec: + return ToolSpec( + name="echo", + description="Echo the input back.", + input_schema={ + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, + handler=lambda value: value, + ) + + +def _stream_text(events) -> str: + return "".join( + e.delta.text + for e in events + if isinstance(e, ContentBlockDeltaEvent) and isinstance(e.delta, TextDelta) + ) + + +# ── the approval gate ─────────────────────────────────────────────────────── + + +def test_the_approval_gate_only_covers_the_interpreter(tmp_path): + """The gate is scoped to `python_interpreter`, not every tool. + + Widening it would silently turn `code_only` into "disable all tools", + which no caller asks for and nothing else would notice. + """ + calls: list[str] = [] + spec = ToolSpec( + name="not_the_interpreter", + description="An ordinary tool.", + input_schema={"type": "object", "properties": {}}, + handler=lambda: calls.append("ran") or "ran", + ) + + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "not_the_interpreter", {}), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[spec], + run_dir=str(tmp_path), + code_only=True, + ) + harness.run("go") + + assert calls == ["ran"] + assert "DRY RUN" not in harness.messages[2].content[0].content + + +def test_a_blocked_interpreter_call_is_not_an_error(tmp_path): + """A refused call is a normal result, not an error. + + Flagging it as an error tells the model its code was broken rather than + declined, so it rewrites and retries instead of stopping. + """ + agent = Agent(adapter=FakeAdapter([]), system="s", run_dir=str(tmp_path)) + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "python_interpreter", {"code": "x = 1"}), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=agent._build_tools(), + run_dir=str(tmp_path), + on_code=lambda code: False, + ) + harness.run("go") + + block = harness.messages[2].content[0] + assert block.is_error is False + assert "blocked" in block.content.lower() + + +def test_code_only_echoes_without_executing(tmp_path): + agent = Agent(adapter=FakeAdapter([]), system="s", run_dir=str(tmp_path)) + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use( + "tu_1", "python_interpreter", {"code": "save('x', 1)"} + ), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=agent._build_tools(cache=agent.cache), + run_dir=str(tmp_path), + cache=agent.cache, + code_only=True, + ) + harness.run("go") + + block = harness.messages[2].content[0] + assert block.is_error is False + assert "DRY RUN" in block.content + assert "x" not in agent.cache.handle_names() + + +# ── run identity ──────────────────────────────────────────────────────────── + + +def test_run_ids_reach_last_result(tmp_path): + """`_stamp` must write through to `last_result`, not only the return value.""" + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("hi")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + + returned = harness.run_result("go", run_id="run-7", session_id="sess-9") + + assert returned.run_id == "run-7" + assert harness.last_result is not None + assert harness.last_result.run_id == "run-7" + assert harness.last_result.session_id == "sess-9" + + +@pytest.mark.asyncio +async def test_streamed_session_turns_are_stamped(tmp_path): + """A streamed turn carries the same ids a non-streamed one does. + + Without it, anything correlating turns by `session_id` silently drops + every streamed turn. + """ + agent = AsyncAgent( + adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("one")]), + system="sys", + run_dir=str(tmp_path), + ) + session = agent.async_session() + + async for _ in session.ask_stream("first"): + pass + + assert session.last_result is not None + assert session.last_result.session_id == session.id + assert session.last_result.run_id is not None + + +# ── what gets logged ──────────────────────────────────────────────────────── + + +def test_latency_is_measured_not_hardcoded(tmp_path): + """The JSONL log's latency must reflect the real provider call.""" + + class SlowAdapter(FakeAdapter): + def chat(self, system, messages, tools): + time.sleep(0.05) + return super().chat(system, messages, tools) + + harness = Harness( + adapter=SlowAdapter([FakeAdapter.text("hi")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.run("go") + + assert harness.run_file is not None + records = [ + json.loads(line) + for line in open(harness.run_file).read().splitlines() + if line.strip() + ] + assert records[0]["metrics"]["latency_ms"] >= 50 + + +def test_a_provider_error_yields_no_partial_text(tmp_path): + """An errored run reports empty text; the detail lives in `error`.""" + + class Broken(FakeAdapter): + def chat(self, system, messages, tools): + raise RuntimeError("down") + + result = Harness( + adapter=Broken([]), system="sys", tools=[], run_dir=str(tmp_path) + ).run_result("go") + + assert result.status == "error" + assert result.text == "" + assert "down" in (result.error or "") + + +@pytest.mark.asyncio +async def test_provider_error_on_the_first_turn(tmp_path): + """Turn-1 failure: zero usage, still a well-formed result.""" + + class Broken(FakeAsyncAdapter): + async def chat(self, system, messages, tools): + raise RuntimeError("down") + + result = await AsyncHarness( + adapter=Broken([]), system="sys", tools=[], run_dir=str(tmp_path) + ).run_result("go") + + assert result.status == "error" + assert result.turns == 1 + assert result.usage.input_tokens == 0 + assert result.run_file is not None + + +# ── the effect protocol ───────────────────────────────────────────────────── + + +def test_a_handler_returning_failed_is_not_mistaken_for_a_failure(tmp_path): + """Successes are wrapped in `Ok`, so no return value can impersonate one. + + Sending raw handler values back into the loop meant a handler that + legitimately returned a `Failed` was reported to the model as raising. + """ + sentinel = Failed(ValueError("I am a legitimate return value")) + spec = ToolSpec( + name="returns_failed", + description="Return a Failed instance.", + input_schema={"type": "object", "properties": {}}, + handler=lambda: sentinel, + ) + + harness = Harness( + adapter=FakeAdapter( + [FakeAdapter.tool_use("tu_1", "returns_failed", {}), FakeAdapter.text("d")] + ), + system="sys", + tools=[spec], + run_dir=str(tmp_path), + ) + harness.run("go") + + block = harness.messages[2].content[0] + assert block.is_error is False + assert "legitimate return value" in block.content + + +def test_the_effect_and_answer_sequence_lines_up(tmp_path): + """Drive `_plan` by hand and pin the effect order and answer shapes. + + Each `yield` receives its own send value, so answering `ToolFinished` with + something is merely ignored, not a desync. What does matter, and is pinned + here, is which effect arrives when and what each one carries. + """ + harness = Harness( + adapter=FakeAdapter([]), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness._begin_run("go") + plan = harness._plan() + + effect = plan.send(None) + assert isinstance(effect, CallProvider) + assert [t.name for t in effect.tools] == ["echo"] + + effect = plan.send(Ok(FakeAdapter.tool_use("tu_1", "echo", {"value": "hi"}))) + assert isinstance(effect, CallTool) + assert effect.tool_name == "echo" + assert effect.tool_input == {"value": "hi"} + + effect = plan.send(Ok("echoed")) + assert isinstance(effect, ToolFinished) + assert effect.tool_name == "echo" + assert effect.block.content == "echoed" + assert effect.block.is_error is False + + effect = plan.send(None) + assert isinstance(effect, CallProvider) + + with pytest.raises(StopIteration): + plan.send(Ok(FakeAdapter.text("final"))) + + assert harness.last_result is not None + assert harness.last_result.text == "final" + assert harness.last_result.turns == 2 + + +def test_every_tool_is_dispatched_before_any_result_is_announced(tmp_path): + """Dispatch-then-announce ordering, which streaming consumers observe.""" + harness = Harness( + adapter=FakeAdapter([]), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness._begin_run("go") + plan = harness._plan() + plan.send(None) + + two_calls = FakeAdapter.tool_use("tu_1", "echo", {"value": "a"}) + two_calls.content.append( + type(two_calls.content[0])( + tool_use_id="tu_2", tool_name="echo", tool_input={"value": "b"} + ) + ) + + kinds = [] + answer = Ok(two_calls) + while True: + effect = plan.send(answer) + kinds.append(type(effect).__name__) + if isinstance(effect, CallTool): + answer = Ok(effect.tool_input["value"]) + elif isinstance(effect, ToolFinished): + answer = None + else: + break + + assert kinds == [ + "CallTool", + "CallTool", + "ToolFinished", + "ToolFinished", + "CallProvider", + ] + plan.close() + + +def test_sync_driver_awaits_async_tool_handlers(tmp_path): + """An `async def` handler under the sync driver must actually be awaited. + + It used to be called and its coroutine object stringified into the + transcript, so the model was shown ''. + """ + + async def handler(value: str) -> str: + await asyncio.sleep(0) + return f"awaited:{value}" + + spec = ToolSpec( + name="async_echo", + description="Async echo.", + input_schema={ + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, + handler=handler, + ) + + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "async_echo", {"value": "v"}), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[spec], + run_dir=str(tmp_path), + ) + harness.run("go") + + block = harness.messages[2].content[0] + assert block.content == "awaited:v" + assert block.is_error is False + + +# ── event-loop hygiene ────────────────────────────────────────────────────── + + +def test_looking_up_the_ambient_loop_does_not_create_one(): + """The helper protecting the ambient loop must not install one itself. + + On Python 3.10/3.11 `asyncio.get_event_loop()` creates a loop when none is + set, leaving a never-run, never-closed loop and its selector fd behind on + a thread that only ever used the synchronous API. + + Must run in a subprocess. The bug only fires on the *main* thread of a + process whose event loop has never been set: `get_event_loop()` raises on + a worker thread and once `set_event_loop` has been called anywhere, so a + probe run inside the pytest process passes against the bug it exists to + catch. This test was written that way first, and was worthless. + """ + probe = textwrap.dedent( + """ + import asyncio, sys + + if sys.version_info < (3, 14): + # Guard the guard: on these versions get_event_loop() only creates + # a loop while this flag is False, so a process where something + # already set a loop would pass against the bug. + policy = asyncio.get_event_loop_policy() + assert policy._local._set_called is False, "process was not pristine" + + from data_harness.core.loop import _ambient_event_loop, run_coroutine_blocking + + before = _ambient_event_loop() + if before is not None: + sys.exit(f"FAIL: looking created a loop: {before!r}") + + async def coro(): + return 1 + + if run_coroutine_blocking(coro()) != 1: + sys.exit("FAIL: coroutine did not run") + + after = _ambient_event_loop() + if after is not None: + sys.exit(f"FAIL: a loop was left behind: {after!r}") + print("OK") + """ + ) + completed = subprocess.run( + [sys.executable, "-c", probe], + capture_output=True, + text=True, + cwd=str(Path(__file__).resolve().parent.parent), + ) + + assert completed.returncode == 0, completed.stdout + completed.stderr + assert "OK" in completed.stdout + + +@pytest.mark.asyncio +async def test_abandoning_a_stream_closes_the_provider_stream(tmp_path): + """A consumer that stops early must not leave the provider's stream open. + + With a real HTTP adapter that generator is holding the connection. + """ + closed: list[bool] = [] + + class TrackingAdapter(FakeAsyncAdapter): + async def stream_events(self, system, messages, tools): + try: + async for evt in super().stream_events(system, messages, tools): + yield evt + finally: + closed.append(True) + + harness = AsyncHarness( + adapter=TrackingAdapter([FakeAsyncAdapter.text("hello")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + + stream = harness.run_stream("go") + async for _ in stream: + break + await stream.aclose() + + assert closed == [True] + + +# ── replay parity between streamed and non-streamed runs ──────────────────── + + +@pytest.mark.asyncio +async def test_streamed_runs_use_and_fill_the_replay_cache(tmp_path): + """`run_stream` must not quietly opt out of the replay cache. + + It previously neither served a hit nor recorded a success, so whether a + caller paid for a repeated question depended on the entry point it chose. + """ + pd = pytest.importorskip("pandas") + shared = tmp_path / "replay.json" + + def build(responses): + agent = AsyncAgent.from_dataframe( + pd.DataFrame({"a": [1, 2, 3]}), + adapter=FakeAsyncAdapter(responses), + run_dir=str(tmp_path), + ) + return agent.enable_cache(str(shared)) + + warming = build( + [ + FakeAsyncAdapter.tool_use( + "tu_1", "python_interpreter", {"code": "answer(6)"} + ), + FakeAsyncAdapter.text("The answer is 6"), + ] + ) + async for _ in warming.run_stream("what is the total"): + pass + assert warming.last_result is not None + assert warming.last_result.text == "The answer is 6" + + # The streamed run recorded the answer, so an empty adapter is enough. + replayed = build([]) + events = [e async for e in replayed.run_stream("what is the total")] + + assert _stream_text(events) == "The answer is 6" + assert replayed.last_result is not None + assert replayed.last_result.turns == 0 + assert replayed.last_result.value == 6 + + +@pytest.mark.asyncio +async def test_async_replay_does_not_block_the_event_loop(tmp_path): + """Replayed steps are offloaded, matching how the async driver dispatches.""" + import threading + + seen: list[str] = [] + agent = AsyncAgent( + adapter=FakeAsyncAdapter([]), system="sys", run_dir=str(tmp_path) + ) + agent.enable_cache() + + spec = ToolSpec( + name="whereami", + description="Report the executing thread.", + input_schema={"type": "object", "properties": {}}, + handler=lambda: seen.append(threading.current_thread().name) or "ok", + ) + + class Cached: + steps = [{"tool": "whereami", "input": {}}] + text = "cached answer" + + agent._build_tools = lambda **kwargs: [spec] # type: ignore[method-assign] + result = await agent._replay(Cached()) + + assert result.text == "cached answer" + assert seen and seen != [threading.current_thread().name] + + +# ── the answer contract is enforced, not assumed ──────────────────────────── + + +def test_a_request_answered_with_none_fails_loudly(tmp_path): + """A driver that answers a request with nothing must be told so. + + Unchecked, the loop does `answer.value` and raises AttributeError several + frames away from the driver that actually got it wrong. + """ + harness = Harness( + adapter=FakeAdapter([]), system="sys", tools=[], run_dir=str(tmp_path) + ) + harness._begin_run("go") + plan = harness._plan() + plan.send(None) + + with pytest.raises(TypeError, match="CallProvider must be answered"): + plan.send(None) + + +def test_a_tool_request_answered_with_none_fails_loudly(tmp_path): + harness = Harness( + adapter=FakeAdapter([]), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness._begin_run("go") + plan = harness._plan() + plan.send(None) + plan.send(Ok(FakeAdapter.tool_use("tu_1", "echo", {"value": "hi"}))) + + with pytest.raises(TypeError, match="CallTool must be answered"): + plan.send(None) + + +def test_a_new_run_clears_the_previous_result(tmp_path): + """`last_result` must not survive into a run that has not finished. + + Otherwise an abandoned stream appears to have completed, reporting the + previous run's tokens and answer. + """ + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("first")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.run("go") + assert harness.last_result is not None + + harness._begin_run("again") + plan = harness._plan() + plan.send(None) + + assert harness.last_result is None + plan.close() + + +# ── abandoning a stream through the public entry points ───────────────────── + + +class _TrackingAdapter(FakeAsyncAdapter): + """Records whether its stream generator was finalised.""" + + def __init__(self, responses): + super().__init__(responses) + self.closed: list[bool] = [] + + async def stream_events(self, system, messages, tools): + try: + async for evt in super().stream_events(system, messages, tools): + yield evt + finally: + self.closed.append(True) + + +@pytest.mark.asyncio +async def test_abandoning_an_agent_stream_closes_the_provider_stream(tmp_path): + """The fix has to hold at the layer callers actually use. + + Closing a generator does not close the one it iterates, so wrapping only + `AsyncHarness.run_stream` left every `AsyncAgent` caller leaking the + provider's connection. A harness-level test cannot see that. + """ + adapter = _TrackingAdapter([FakeAsyncAdapter.text("hello there")]) + agent = AsyncAgent(adapter=adapter, system="sys", run_dir=str(tmp_path)) + + stream = agent.run_stream("go") + async for _ in stream: + break + await stream.aclose() + + assert adapter.closed == [True] + + +@pytest.mark.asyncio +async def test_abandoning_a_session_stream_closes_the_provider_stream(tmp_path): + adapter = _TrackingAdapter([FakeAsyncAdapter.text("hello there")]) + session = AsyncAgent( + adapter=adapter, system="sys", run_dir=str(tmp_path) + ).async_session() + + stream = session.ask_stream("go") + async for _ in stream: + break + await stream.aclose() + + assert adapter.closed == [True] + + +@pytest.mark.asyncio +async def test_streamed_agent_runs_are_stamped(tmp_path): + """`run_stream` must stamp a run id like every other entry point.""" + agent = AsyncAgent( + adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("done")]), + system="sys", + run_dir=str(tmp_path), + ) + + async for _ in agent.run_stream("go"): + pass + + assert agent.last_result is not None + assert agent.last_result.run_id is not None + + +@pytest.mark.asyncio +async def test_a_replay_hit_is_readable_through_last_result(tmp_path): + """A cache hit builds no harness, so `last_harness` cannot be the answer.""" + pd = pytest.importorskip("pandas") + shared = tmp_path / "replay.json" + + def build(responses): + return AsyncAgent.from_dataframe( + pd.DataFrame({"a": [1, 2, 3]}), + adapter=FakeAsyncAdapter(responses), + run_dir=str(tmp_path), + ).enable_cache(str(shared)) + + warming = build( + [ + FakeAsyncAdapter.tool_use( + "tu_1", "python_interpreter", {"code": "answer(6)"} + ), + FakeAsyncAdapter.text("The answer is 6"), + ] + ) + async for _ in warming.run_stream("total"): + pass + + # A fresh agent has never built a harness, which is exactly the case that + # made the old docstring's `last_harness.last_result` an AttributeError. + fresh = build([]) + assert fresh.last_harness is None + + async for _ in fresh.run_stream("total"): + pass + + assert fresh.last_result is not None + assert fresh.last_result.text == "The answer is 6" + + +@pytest.mark.skipif( + sys.version_info >= (3, 14), + reason=( + "Event loop policies are deprecated from 3.14 and removed in 3.16. " + "_ambient_event_loop takes the get_event_loop() branch there, which " + "cannot create a loop, so there is no unreadable case to handle." + ), +) +def test_an_unreadable_ambient_loop_leaves_no_closed_loop_behind(): + """Under a policy we cannot introspect, don't guess and don't poison. + + Skipping the restore would leave the thread pointing at the loop + `run_coroutine_blocking` just closed, which is worse than the leak the + branch exists to avoid: every later `get_event_loop()` user gets + 'Event loop is closed'. + """ + import threading + + from data_harness.core.loop import _UNREADABLE, _ambient_event_loop + + class OpaquePolicy(asyncio.AbstractEventLoopPolicy): + """A policy that keeps its current loop somewhere we cannot read. + + Third-party policies are not obliged to expose the stdlib's private + `_local` slot, so this is the shape the fallback branch must survive. + """ + + def __init__(self) -> None: + self.current: asyncio.AbstractEventLoop | None = None + # Delegate loop construction; building it via `asyncio.new_event_loop` + # would route back through this policy and recurse. + self._factory = asyncio.DefaultEventLoopPolicy() + + def get_event_loop(self): + if self.current is None: + raise RuntimeError("no current loop") + return self.current + + def set_event_loop(self, loop): + self.current = loop + + def new_event_loop(self): + return self._factory.new_event_loop() + + observed: dict[str, object] = {} + + def on_a_fresh_thread() -> None: + previous_policy = asyncio.get_event_loop_policy() + policy = OpaquePolicy() + asyncio.set_event_loop_policy(policy) + try: + observed["probe"] = _ambient_event_loop() + + async def coro(): + return 42 + + observed["result"] = run_coroutine_blocking(coro()) + observed["left_behind"] = policy.current + finally: + asyncio.set_event_loop_policy(previous_policy) + + worker = threading.Thread(target=on_a_fresh_thread) + worker.start() + worker.join() + + assert observed["probe"] is _UNREADABLE + assert observed["result"] == 42 + # Not a closed loop. `None` is the honest answer when we could not read + # what was there before. + assert observed["left_behind"] is None diff --git a/tests/test_loop_reminders.py b/tests/test_loop_reminders.py index a33bcce..64bc67c 100644 --- a/tests/test_loop_reminders.py +++ b/tests/test_loop_reminders.py @@ -4,9 +4,13 @@ import copy -from data_harness.loop import Harness -from data_harness.providers.base import NormalizedResponse, ProviderAdapter, StopReason -from data_harness.types import ( +from data_harness.data.harness import Harness +from data_harness.llm.providers.base import ( + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.types import ( Message, TextBlock, ToolSpec, diff --git a/tests/test_mcp.py b/tests/test_mcp.py index 581b33f..cec232c 100644 --- a/tests/test_mcp.py +++ b/tests/test_mcp.py @@ -5,8 +5,8 @@ from types import SimpleNamespace from data_harness import Agent -from data_harness.mcp import _result_to_text, mcp_tool_specs -from data_harness.testing import FakeAdapter +from data_harness.data.mcp import _result_to_text, mcp_tool_specs +from data_harness.llm.testing import FakeAdapter class _FakeMCPClient: diff --git a/tests/test_observe.py b/tests/test_observe.py index a08383c..235ee19 100644 --- a/tests/test_observe.py +++ b/tests/test_observe.py @@ -1,6 +1,6 @@ import time -from data_harness.observe import TurnMetrics, time_block +from data_harness.core.observe import TurnMetrics, time_block def test_turn_metrics_fields(): diff --git a/tests/test_providers.py b/tests/test_providers.py index 471cc18..81c9f13 100644 --- a/tests/test_providers.py +++ b/tests/test_providers.py @@ -3,8 +3,8 @@ from contextlib import asynccontextmanager from unittest.mock import MagicMock, patch -from data_harness.providers.base import StopReason -from data_harness.streaming import ( +from data_harness.llm.providers.base import StopReason +from data_harness.llm.streaming import ( ContentBlockDeltaEvent, ContentBlockStartEvent, InputJSONDelta, @@ -13,7 +13,7 @@ MessageStopEvent, TextDelta, ) -from data_harness.types import Message, TextBlock, ToolSpec, ToolUseBlock +from data_harness.llm.types import Message, TextBlock, ToolSpec, ToolUseBlock def make_anthropic_response(stop_reason="end_turn", content_blocks=None, usage=None): @@ -62,7 +62,7 @@ def test_values(self): class TestAnthropicAdapter: def _make_adapter(self): with patch("anthropic.Anthropic"): - from data_harness.providers.anthropic import AnthropicAdapter + from data_harness.llm.providers.anthropic import AnthropicAdapter adapter = AnthropicAdapter(model="claude-3-5-sonnet-20241022") return adapter @@ -329,7 +329,7 @@ def _make_tool_sse_sequence( class TestAsyncAnthropicAdapterStreamEvents: def _make_async_adapter(self): with patch("anthropic.AsyncAnthropic"): - from data_harness.providers.anthropic import AsyncAnthropicAdapter + from data_harness.llm.providers.anthropic import AsyncAnthropicAdapter return AsyncAnthropicAdapter(model="claude-sonnet-4-6") @@ -437,7 +437,7 @@ async def test_tool_use_stream_content_block_start(self): starts = [e for e in events if isinstance(e, ContentBlockStartEvent)] assert len(starts) == 1 - from data_harness.types import ToolUseBlock + from data_harness.llm.types import ToolUseBlock assert isinstance(starts[0].content_block, ToolUseBlock) assert starts[0].content_block.tool_use_id == "tu1" diff --git a/tests/test_providers_openai.py b/tests/test_providers_openai.py index 5f0fddf..4727fac 100644 --- a/tests/test_providers_openai.py +++ b/tests/test_providers_openai.py @@ -9,8 +9,8 @@ pytest.importorskip("openai") -from data_harness.providers.base import StopReason -from data_harness.types import ( +from data_harness.llm.providers.base import StopReason +from data_harness.llm.types import ( Message, TextBlock, ToolResultBlock, @@ -51,7 +51,7 @@ def make_openai_tool_call(id_, name, arguments): class TestOpenAIAdapter: def _make_adapter(self): with patch("openai.OpenAI"): - from data_harness.providers.openai import OpenAIAdapter + from data_harness.llm.providers.openai import OpenAIAdapter adapter = OpenAIAdapter(model="gpt-test") return adapter @@ -310,7 +310,7 @@ def test_openai_live_smoke(): if not os.environ.get("OPENAI_API_KEY"): pytest.skip("OPENAI_API_KEY not set") - from data_harness.providers.openai import OpenAIAdapter + from data_harness.llm.providers.openai import OpenAIAdapter adapter = OpenAIAdapter(model=os.environ.get("OPENAI_MODEL", "gpt-4o-mini")) response = adapter.chat( diff --git a/tests/test_quickstart.py b/tests/test_quickstart.py index d54ece4..81fa9d8 100644 --- a/tests/test_quickstart.py +++ b/tests/test_quickstart.py @@ -6,9 +6,9 @@ import pytest from data_harness import Chat, SmartFrame, ask -from data_harness.io import load_dataframe, sanitise_handle, to_handles -from data_harness.quickstart import resolve_adapter -from data_harness.testing import FakeAdapter +from data_harness.app.quickstart import resolve_adapter +from data_harness.data.io import load_dataframe, sanitise_handle, to_handles +from data_harness.llm.testing import FakeAdapter def _frame() -> pd.DataFrame: @@ -26,8 +26,8 @@ def _answer_adapter(code: str, final: str) -> FakeAdapter: # --- provider resolution --------------------------------------------------- def test_resolve_adapter_routes_by_model_name(monkeypatch): - from data_harness.providers.anthropic import AnthropicAdapter - from data_harness.providers.openai import OpenAIAdapter + from data_harness.llm.providers.anthropic import AnthropicAdapter + from data_harness.llm.providers.openai import OpenAIAdapter monkeypatch.setenv("ANTHROPIC_API_KEY", "test-key") monkeypatch.setenv("OPENAI_API_KEY", "test-key") @@ -37,7 +37,7 @@ def test_resolve_adapter_routes_by_model_name(monkeypatch): def test_resolve_adapter_prefers_anthropic_env(monkeypatch): - from data_harness.providers.anthropic import AnthropicAdapter + from data_harness.llm.providers.anthropic import AnthropicAdapter monkeypatch.setenv("ANTHROPIC_API_KEY", "test-key") monkeypatch.setenv("OPENAI_API_KEY", "test-key") @@ -45,7 +45,7 @@ def test_resolve_adapter_prefers_anthropic_env(monkeypatch): def test_resolve_adapter_falls_back_to_openai(monkeypatch): - from data_harness.providers.openai import OpenAIAdapter + from data_harness.llm.providers.openai import OpenAIAdapter monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) monkeypatch.setenv("OPENAI_API_KEY", "test-key") @@ -61,7 +61,7 @@ def test_resolve_adapter_raises_without_keys(monkeypatch): def test_resolve_adapter_routes_slash_model_to_openrouter(monkeypatch): - from data_harness.providers.openai import OPENROUTER_BASE_URL, OpenRouterAdapter + from data_harness.llm.providers.openai import OPENROUTER_BASE_URL, OpenRouterAdapter monkeypatch.setenv("OPENROUTER_API_KEY", "test-key") adapter = resolve_adapter("anthropic/claude-3.5-sonnet") @@ -70,7 +70,7 @@ def test_resolve_adapter_routes_slash_model_to_openrouter(monkeypatch): def test_resolve_adapter_falls_back_to_openrouter(monkeypatch): - from data_harness.providers.openai import OpenRouterAdapter + from data_harness.llm.providers.openai import OpenRouterAdapter monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) monkeypatch.delenv("OPENAI_API_KEY", raising=False) @@ -79,7 +79,7 @@ def test_resolve_adapter_falls_back_to_openrouter(monkeypatch): def test_openrouter_adapter_uses_env_key(monkeypatch): - from data_harness.providers.openai import OpenRouterAdapter + from data_harness.llm.providers.openai import OpenRouterAdapter monkeypatch.setenv("OPENROUTER_API_KEY", "router-secret") adapter = OpenRouterAdapter(model="openai/gpt-4o-mini") @@ -87,7 +87,7 @@ def test_openrouter_adapter_uses_env_key(monkeypatch): def test_resolve_adapter_routes_deepseek_direct(monkeypatch): - from data_harness.providers.openai import DEEPSEEK_BASE_URL, DeepSeekAdapter + from data_harness.llm.providers.openai import DEEPSEEK_BASE_URL, DeepSeekAdapter monkeypatch.setenv("DEEPSEEK_API_KEY", "ds-key") adapter = resolve_adapter("deepseek-chat") @@ -96,7 +96,7 @@ def test_resolve_adapter_routes_deepseek_direct(monkeypatch): def test_resolve_adapter_deepseek_slash_goes_to_openrouter(monkeypatch): - from data_harness.providers.openai import OpenRouterAdapter + from data_harness.llm.providers.openai import OpenRouterAdapter monkeypatch.setenv("OPENROUTER_API_KEY", "test-key") # a slash means OpenRouter, even for deepseek @@ -104,7 +104,7 @@ def test_resolve_adapter_deepseek_slash_goes_to_openrouter(monkeypatch): def test_resolve_adapter_falls_back_to_deepseek(monkeypatch): - from data_harness.providers.openai import DeepSeekAdapter + from data_harness.llm.providers.openai import DeepSeekAdapter for var in ("ANTHROPIC_API_KEY", "OPENAI_API_KEY", "OPENROUTER_API_KEY"): monkeypatch.delenv(var, raising=False) @@ -213,7 +213,7 @@ def test_load_dataframe_unsupported(tmp_path): # --- pandas accessor + notebook magic -------------------------------------- def test_pandas_chat_accessor(tmp_path): - import data_harness.pandas # noqa: F401 registers the accessor + import data_harness.app.pandas # noqa: F401 registers the accessor df = _frame() adapter = FakeAdapter([FakeAdapter.text("via accessor")]) @@ -222,7 +222,7 @@ def test_pandas_chat_accessor(tmp_path): def test_notebook_magic_class_builds(): - from data_harness.notebook import _load_magics_class + from data_harness.app.notebook import _load_magics_class cls = _load_magics_class() assert hasattr(cls, "ask") diff --git a/tests/test_result.py b/tests/test_result.py index f3ac5ed..344b65d 100644 --- a/tests/test_result.py +++ b/tests/test_result.py @@ -9,13 +9,13 @@ import pytest -from data_harness.cache import SessionCache -from data_harness.exceptions import MaxTurnsExceeded -from data_harness.loop import Harness -from data_harness.providers.base import NormalizedResponse, StopReason -from data_harness.result import CacheStorageInfo, RunResult, Usage -from data_harness.testing import FakeAdapter -from data_harness.types import TextBlock, ToolSpec, ToolUseBlock +from data_harness.core.exceptions import MaxTurnsExceeded +from data_harness.core.result import CacheStorageInfo, RunResult, Usage +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.llm.providers.base import NormalizedResponse, StopReason +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import TextBlock, ToolSpec, ToolUseBlock # --------------------------------------------------------------------------- # Helpers @@ -256,14 +256,14 @@ def test_stop_reason_populated(self, tmp_path): result = make_harness([make_text_response("ok")], tmp_path=tmp_path).run_result( "x" ) - from data_harness.providers.base import StopReason + from data_harness.llm.providers.base import StopReason assert result.stop_reason == StopReason.END_TURN def test_max_turns_returns_result_not_raises(self, tmp_path): # run_result should NOT raise MaxTurnsExceeded - it returns status instead # All responses are TOOL_USE to force max turns - from data_harness.providers.base import NormalizedResponse as NR + from data_harness.llm.providers.base import NormalizedResponse as NR tool_responses = [ NR( @@ -379,7 +379,7 @@ def test_ask_still_returns_string(self, tmp_path): class TestAgentRunResult: def test_returns_run_result(self, tmp_path): - from data_harness.agent import Agent + from data_harness.app.agent import Agent adapter = FakeAdapter([FakeAdapter.text("done")]) agent = Agent(adapter=adapter, system="sys", run_dir=str(tmp_path)) @@ -387,7 +387,7 @@ def test_returns_run_result(self, tmp_path): assert isinstance(result, RunResult) def test_text_matches_run(self, tmp_path): - from data_harness.agent import Agent + from data_harness.app.agent import Agent adapter_a = FakeAdapter([FakeAdapter.text("hello")]) adapter_b = FakeAdapter([FakeAdapter.text("hello")]) @@ -396,7 +396,7 @@ def test_text_matches_run(self, tmp_path): assert agent_a.run_result("hi").text == agent_b.run("hi") def test_status_success(self, tmp_path): - from data_harness.agent import Agent + from data_harness.app.agent import Agent adapter = FakeAdapter([FakeAdapter.text("ok")]) agent = Agent(adapter=adapter, system="sys", run_dir=str(tmp_path)) @@ -404,7 +404,7 @@ def test_status_success(self, tmp_path): assert result.status == "success" def test_run_file_populated(self, tmp_path): - from data_harness.agent import Agent + from data_harness.app.agent import Agent adapter = FakeAdapter([FakeAdapter.text("ok")]) agent = Agent(adapter=adapter, system="sys", run_dir=str(tmp_path)) @@ -413,7 +413,7 @@ def test_run_file_populated(self, tmp_path): assert Path(result.run_file).exists() def test_run_still_returns_string(self, tmp_path): - from data_harness.agent import Agent + from data_harness.app.agent import Agent adapter = FakeAdapter([FakeAdapter.text("hello")]) agent = Agent(adapter=adapter, system="sys", run_dir=str(tmp_path)) @@ -427,7 +427,7 @@ def test_run_still_returns_string(self, tmp_path): class TestAgentSessionAskResult: def test_returns_run_result(self, tmp_path): - from data_harness.agent import Agent + from data_harness.app.agent import Agent adapter = FakeAdapter([FakeAdapter.text("first"), FakeAdapter.text("second")]) agent = Agent(adapter=adapter, system="sys", run_dir=str(tmp_path)) @@ -438,8 +438,8 @@ def test_returns_run_result(self, tmp_path): assert result.text == "second" def test_usage_per_ask(self, tmp_path): - from data_harness.agent import Agent - from data_harness.providers.base import NormalizedResponse as NR + from data_harness.app.agent import Agent + from data_harness.llm.providers.base import NormalizedResponse as NR r1 = NR( stop_reason=StopReason.END_TURN, @@ -466,7 +466,7 @@ def test_usage_per_ask(self, tmp_path): assert result.usage.output_tokens == 3 def test_ask_still_returns_string(self, tmp_path): - from data_harness.agent import Agent + from data_harness.app.agent import Agent adapter = FakeAdapter([FakeAdapter.text("hi")]) agent = Agent(adapter=adapter, system="sys", run_dir=str(tmp_path)) diff --git a/tests/test_review_fixes.py b/tests/test_review_fixes.py index 74efc1e..82b5a80 100644 --- a/tests/test_review_fixes.py +++ b/tests/test_review_fixes.py @@ -10,10 +10,14 @@ import pytest -from data_harness.loop import Harness -from data_harness.providers.base import NormalizedResponse, ProviderAdapter, StopReason -from data_harness.testing import FakeAdapter -from data_harness.types import TextBlock, ToolAnnotations, ToolSpec +from data_harness.data.harness import Harness +from data_harness.llm.providers.base import ( + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import TextBlock, ToolAnnotations, ToolSpec def make_text_response(text: str) -> NormalizedResponse: @@ -35,7 +39,7 @@ def make_text_response(text: str) -> NormalizedResponse: class TestAgentSessionAskRaisesOnError: def test_ask_raises_runtime_error_on_adapter_exception(self, tmp_path): """AgentSession.ask() must raise RuntimeError when adapter raises.""" - from data_harness.agent import Agent + from data_harness.app.agent import Agent class BoomAdapter(ProviderAdapter): def chat(self, system, messages, tools): @@ -51,7 +55,7 @@ def format_cache_control(self, obj): def test_ask_result_still_returns_error_status_not_raises(self, tmp_path): """ask_result() must return RunResult(status='error'), not raise.""" - from data_harness.agent import Agent + from data_harness.app.agent import Agent class BoomAdapter(ProviderAdapter): def chat(self, system, messages, tools): @@ -158,7 +162,7 @@ def test_stop_sequence_terminates_immediately(self, tmp_path): class TestAgentSessionDeepCacheIsolation: def test_session_cache_is_isolated_from_agent_cache(self, tmp_path): """Mutating session cache must not affect the agent-level cache.""" - from data_harness.agent import Agent + from data_harness.app.agent import Agent agent = Agent( adapter=FakeAdapter([make_text_response("ok")]), @@ -179,7 +183,7 @@ def test_session_cache_is_isolated_from_agent_cache(self, tmp_path): def test_agent_cache_mutation_does_not_affect_session(self, tmp_path): """Mutating agent cache after session creation must not affect session cache.""" - from data_harness.agent import Agent + from data_harness.app.agent import Agent shared_list = [10, 20] agent = Agent( diff --git a/tests/test_sandbox.py b/tests/test_sandbox.py index d4ec64f..b3c2d0f 100644 --- a/tests/test_sandbox.py +++ b/tests/test_sandbox.py @@ -6,10 +6,10 @@ import pytest from data_harness import Agent -from data_harness.cache import SessionCache -from data_harness.testing import FakeAdapter -from data_harness.tools.interpreter import PythonInterpreterError -from data_harness.tools.sandbox import SubprocessPythonInterpreter +from data_harness.data.cache import SessionCache +from data_harness.data.tools.interpreter import PythonInterpreterError +from data_harness.data.tools.sandbox import SubprocessPythonInterpreter +from data_harness.llm.testing import FakeAdapter def _interp(tmp_path, *, timeout=30, **kw) -> SubprocessPythonInterpreter: @@ -60,11 +60,58 @@ def test_sandbox_propagates_runtime_error(tmp_path): def test_sandbox_timeout(tmp_path): + """A timeout is the environment failing, not the model's code being wrong. + + It used to raise PythonInterpreterError, the same as a NameError in the + model's code, which left a caller unable to tell "your sandbox is too + small" from "the model wrote a bug". Both remain DataHarnessError, so a + caller that does not care about the difference is unaffected. + """ + from data_harness.core.exceptions import DataHarnessError, ExecutionError + interp = _interp(tmp_path, timeout=2, cpu_seconds=1) - with pytest.raises(PythonInterpreterError): + with pytest.raises(ExecutionError) as excinfo: # busy loop exceeds both the CPU and wall-clock limit interp.run("while True:\n pass") + assert excinfo.value.code == "execution_error" + assert isinstance(excinfo.value, DataHarnessError) + + +def test_sandbox_wall_clock_timeout(tmp_path, monkeypatch): + """The wall-clock branch, driven directly. + + The other timeout test burns CPU, so the kernel kills the child and the + sandbox reports a non-zero exit instead: the `TimeoutExpired` branch never + runs, and a mutation to it survived the whole suite. Model code cannot + sleep either, since `time` is not on the allow-list, so the only honest + way to reach this branch is to make the subprocess call time out. + """ + import subprocess + + from data_harness.core.exceptions import ExecutionError + + def times_out(*args, **kwargs): + raise subprocess.TimeoutExpired(cmd="python", timeout=1) + + monkeypatch.setattr(subprocess, "run", times_out) + interp = _interp(tmp_path, timeout=1) + + with pytest.raises(ExecutionError, match="timed out"): + interp.run("1 + 1") + + +def test_a_bug_in_the_models_code_is_not_an_execution_error(tmp_path): + """The other side of the distinction: this one is the model's to fix.""" + from data_harness.core.exceptions import ExecutionError + + interp = _interp(tmp_path) + with pytest.raises(PythonInterpreterError) as excinfo: + interp.run("this_name_does_not_exist") + + assert not isinstance(excinfo.value, ExecutionError) + assert excinfo.value.code == "interpreter_error" + def test_sandbox_handles_roundtrip_dataframe(tmp_path): interp = _interp(tmp_path) diff --git a/tests/test_schema.py b/tests/test_schema.py index f77ce69..89f30d5 100644 --- a/tests/test_schema.py +++ b/tests/test_schema.py @@ -7,7 +7,7 @@ import pytest -from data_harness.schema import infer_input_schema +from data_harness.core.schema import infer_input_schema OVERRIDE_HINT = "pass input_schema=... to override" diff --git a/tests/test_serialize.py b/tests/test_serialize.py index 6320c73..6fb8e67 100644 --- a/tests/test_serialize.py +++ b/tests/test_serialize.py @@ -4,8 +4,8 @@ import pytest -from data_harness.serialize import to_jsonable -from data_harness.types import TextBlock, ToolResultBlock, ToolUseBlock +from data_harness.core.serialize import to_jsonable +from data_harness.llm.types import TextBlock, ToolResultBlock, ToolUseBlock class Color(Enum): diff --git a/tests/test_session_inspection.py b/tests/test_session_inspection.py index a348b02..40de033 100644 --- a/tests/test_session_inspection.py +++ b/tests/test_session_inspection.py @@ -5,11 +5,11 @@ from __future__ import annotations -from data_harness.agent import Agent -from data_harness.providers.base import NormalizedResponse, StopReason -from data_harness.result import RunResult -from data_harness.testing import FakeAdapter -from data_harness.types import TextBlock +from data_harness.app.agent import Agent +from data_harness.core.result import RunResult +from data_harness.llm.providers.base import NormalizedResponse, StopReason +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import TextBlock def make_text_response( @@ -117,15 +117,24 @@ def test_one_shot_run_result_has_no_session_id(self, tmp_path): result = agent.run_result("q") assert result.session_id is None - def test_agent_has_no_last_result_attribute(self, tmp_path): - """Per plan: Agent should NOT grow last_result in this phase.""" + def test_agent_records_last_result(self, tmp_path): + """`Agent` deliberately grew `last_result`. + + It previously did not, on the grounds that `run_result` already + returns one. That stopped being sufficient once a run could be + streamed or served from the replay cache: neither hands the caller a + result, so there was no way to reach a streamed run's token usage. + """ agent = Agent( adapter=FakeAdapter([make_text_response("hi")]), system="s", run_dir=str(tmp_path), ) - agent.run_result("q") - assert not hasattr(agent, "last_result") + assert agent.last_result is None + + returned = agent.run_result("q") + + assert agent.last_result is returned # --------------------------------------------------------------------------- @@ -219,7 +228,7 @@ def test_turns_counts_model_turns_not_asks(self, tmp_path): A two-turn run (tool-use + end-turn) triggered by one ask() should contribute 2 to session.turns. """ - from data_harness.types import ToolSpec, ToolUseBlock + from data_harness.llm.types import ToolSpec, ToolUseBlock tool_resp = NormalizedResponse( stop_reason=StopReason.TOOL_USE, @@ -246,6 +255,6 @@ def test_turns_counts_model_turns_not_asks(self, tmp_path): # by using the harness directly through session session = agent.session() # Patch: add the echo spec to the harness tools for this test - session.harness._tools.append(echo_spec) + session.harness.tools.append(echo_spec) session.ask("go") assert session.turns == 2 diff --git a/tests/test_session_tree.py b/tests/test_session_tree.py new file mode 100644 index 0000000..9614de9 --- /dev/null +++ b/tests/test_session_tree.py @@ -0,0 +1,846 @@ +"""Phase 3: the session tree. + +The claim is that a session is an append-only tree and the conversation is +*derived* from it. Everything worth having follows from that: resume, forking, +reversible compaction, and a log that is linear rather than quadratic to +write. Each of those is tested here, because each is a claim a README could +make falsely. +""" + +from __future__ import annotations + +import json + +import pytest + +from data_harness.core.session import ( + JsonlSessionStore, + LeafEntry, + MemorySessionStore, + MessageEntry, + Session, + SessionStore, + SessionStoreError, + TurnEntry, +) +from data_harness.data.harness import Harness +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import Message, TextBlock, ToolSpec + + +def say(text: str, role: str = "user") -> Message: + return Message(role=role, content=[TextBlock(text=text)]) + + +def texts(messages: list[Message]) -> list[str]: + return [m.content[0].text for m in messages] + + +def echo_spec() -> ToolSpec: + return ToolSpec( + name="echo", + description="Echo the input.", + input_schema={ + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, + handler=lambda value: value, + ) + + +@pytest.fixture(params=["memory", "jsonl"]) +def store(request, tmp_path) -> SessionStore: + """Both stores, everywhere, so neither drifts from the protocol.""" + if request.param == "memory": + return MemorySessionStore("sess") + return JsonlSessionStore.create(tmp_path / "session.jsonl", "sess") + + +# ── the tree ──────────────────────────────────────────────────────────────── + + +def test_context_is_derived_from_the_path(store): + session = Session(store) + session.append_message(say("q1")) + session.append_message(say("a1", "assistant")) + session.append_message(say("q2")) + + assert texts(session.build_context()) == ["q1", "a1", "q2"] + + +def test_branching_keeps_both_branches(store): + """The point of a tree. Retrying a turn must not destroy the first try.""" + session = Session(store) + session.append_message(say("q1")) + fork_point = session.append_message(say("a1", "assistant")) + original = session.append_message(say("original follow-up")) + + session.move_to(fork_point) + session.append_message(say("different follow-up")) + + assert texts(session.build_context()) == ["q1", "a1", "different follow-up"] + assert texts(session.build_context(original)) == ["q1", "a1", "original follow-up"] + + +def test_moving_the_leaf_is_itself_recorded(store): + session = Session(store) + first = session.append_message(say("q1")) + session.append_message(say("q2")) + session.move_to(first) + + moves = [e for e in store.entries() if isinstance(e, LeafEntry)] + assert [m.target_id for m in moves] == [first] + + +def test_moving_to_none_starts_a_fresh_root(store): + session = Session(store) + session.append_message(say("q1")) + session.move_to(None) + session.append_message(say("a new beginning")) + + assert texts(session.build_context()) == ["a new beginning"] + + +def test_an_unknown_entry_cannot_be_made_the_leaf(store): + session = Session(store) + session.append_message(say("q1")) + with pytest.raises(SessionStoreError) as excinfo: + session.move_to("nope") + assert excinfo.value.code == "not_found" + + +def test_entries_are_never_removed(store): + """History only grows. That is what makes a run auditable after the fact.""" + session = Session(store) + session.append_message(say("q1")) + fork = session.append_message(say("a1", "assistant")) + session.append_message(say("abandoned")) + session.move_to(fork) + session.append_message(say("kept")) + + recorded = [ + e.message.content[0].text + for e in store.entries() + if isinstance(e, MessageEntry) + ] + assert recorded == ["q1", "a1", "abandoned", "kept"] + + +# ── compaction ────────────────────────────────────────────────────────────── + + +def test_compaction_shrinks_the_context_but_not_the_tree(store): + session = Session(store) + session.append_message(say("ancient q")) + session.append_message(say("ancient a", "assistant")) + keep_from = session.append_message(say("recent q")) + session.append_compaction( + summary="They discussed something ancient.", + first_kept_entry_id=keep_from, + tokens_before=1234, + ) + session.append_message(say("newest q")) + + context = texts(session.build_context()) + assert context[0].startswith("Summary of the earlier conversation:") + assert "recent q" in context + assert "newest q" in context + assert "ancient q" not in context + + # Nothing was deleted. + stored = [ + e.message.content[0].text + for e in store.entries() + if isinstance(e, MessageEntry) + ] + assert "ancient q" in stored + + +def test_compaction_is_reversible_by_moving_the_leaf(store): + """A compaction is an entry, not an edit, so stepping back undoes it.""" + session = Session(store) + session.append_message(say("ancient q")) + before_compaction = session.append_message(say("recent q")) + session.append_compaction("a summary", first_kept_entry_id=None, tokens_before=1) + session.append_message(say("after")) + + assert "ancient q" not in texts(session.build_context()) + + session.move_to(before_compaction) + assert texts(session.build_context()) == ["ancient q", "recent q"] + + +def test_compaction_without_a_kept_tail_drops_everything_before_it(store): + session = Session(store) + session.append_message(say("q1")) + session.append_message(say("q2")) + session.append_compaction("summary", first_kept_entry_id=None, tokens_before=9) + + assert len(session.build_context()) == 1 + + +def test_the_newest_compaction_wins(store): + session = Session(store) + session.append_message(say("q1")) + session.append_compaction("first summary", None, 1) + session.append_message(say("q2")) + session.append_compaction("second summary", None, 2) + session.append_message(say("q3")) + + context = texts(session.build_context()) + assert any("second summary" in c for c in context) + assert not any("first summary" in c for c in context) + + +# ── custom entries: the extension point that keeps core domain-free ──────── + + +def test_custom_entries_are_invisible_to_the_model_by_default(store): + session = Session(store) + session.append_message(say("q1")) + session.append_custom("cache_put", {"handle": "sales_df", "snapshot": "3x2"}) + + assert texts(session.build_context()) == ["q1"] + assert [e.data["handle"] for e in session.custom_entries("cache_put")] == [ + "sales_df" + ] + + +def test_a_projector_opts_a_custom_entry_into_the_context(store): + session = Session( + store, + projectors={ + "note": lambda entry: [say(f"note: {entry.data['text']}")], + }, + ) + session.append_message(say("q1")) + session.append_custom("note", {"text": "remember this"}) + + assert texts(session.build_context()) == ["q1", "note: remember this"] + + +# ── labels ────────────────────────────────────────────────────────────────── + + +def test_labels_name_a_branch_point(store): + session = Session(store) + first = session.append_message(say("q1")) + session.label(first, "before the detour") + + assert session.labels() == {first: "before the detour"} + + session.label(first, None) + assert session.labels() == {} + + +def test_labelling_an_unknown_entry_is_an_error(store): + session = Session(store) + with pytest.raises(SessionStoreError): + session.label("nope", "x") + + +# ── store integrity ───────────────────────────────────────────────────────── + + +def test_a_duplicate_entry_id_is_rejected(store): + entry = MessageEntry(id="fixed", parent_id=None, message=say("q")) + store.append(entry) + with pytest.raises(SessionStoreError) as excinfo: + store.append(MessageEntry(id="fixed", parent_id=None, message=say("q"))) + assert excinfo.value.code == "duplicate_entry" + + +def test_an_entry_naming_an_unknown_parent_is_rejected(store): + with pytest.raises(SessionStoreError) as excinfo: + store.append(MessageEntry(id="a", parent_id="ghost", message=say("q"))) + assert excinfo.value.code == "missing_parent" + + +def test_both_stores_behave_the_same_way(store): + """`isinstance` against a Protocol checks names only, so exercise it. + + A class with every method present and every signature wrong passes + `isinstance(..., SessionStore)`, which makes that assertion close to + worthless on its own. + """ + assert isinstance(store, SessionStore) + + first = MessageEntry(id="e1", parent_id=None, message=say("q1")) + second = MessageEntry(id="e2", parent_id="e1", message=say("a1", "assistant")) + store.append(first) + store.append(second) + + assert store.leaf_id == "e2" + assert store.get("e1") == first + assert store.get("missing") is None + assert [e.id for e in store.entries()] == ["e1", "e2"] + assert [e.id for e in store.path_to_root("e2")] == ["e1", "e2"] + assert store.path_to_root(None) == [] + + store.set_leaf("e1") + assert store.leaf_id == "e1" + + +# ── persistence ───────────────────────────────────────────────────────────── + + +def test_a_jsonl_session_round_trips(tmp_path): + path = tmp_path / "s.jsonl" + session = Session(JsonlSessionStore.create(path, "sess-1")) + session.append_message(say("q1")) + session.append_message(say("a1", "assistant")) + session.append_turn( + turn=1, input_tokens=10, output_tokens=5, stop_reason="end_turn" + ) + session.append_custom("chart", {"path": "/tmp/x.png"}) + + reopened = Session(JsonlSessionStore.open(path)) + + assert texts(reopened.build_context()) == ["q1", "a1"] + assert reopened.stats().input_tokens == 10 + assert reopened.custom_entries("chart")[0].data["path"] == "/tmp/x.png" + assert reopened.leaf_id == session.leaf_id + + +def test_tool_messages_survive_a_round_trip(tmp_path): + """Tool blocks are the part a naive text-only encoder would silently lose.""" + from data_harness.llm.types import ToolResultBlock, ToolUseBlock + + path = tmp_path / "s.jsonl" + session = Session(JsonlSessionStore.create(path, "sess")) + session.append_message( + Message( + role="assistant", + content=[ + ToolUseBlock(tool_use_id="t1", tool_name="echo", tool_input={"v": 1}) + ], + ) + ) + session.append_message( + Message( + role="user", + content=[ToolResultBlock(tool_use_id="t1", content="ok", is_error=False)], + ) + ) + + reopened = Session(JsonlSessionStore.open(path)).build_context() + + use = reopened[0].content[0] + result = reopened[1].content[0] + assert isinstance(use, ToolUseBlock) + assert use.tool_name == "echo" and use.tool_input == {"v": 1} + assert isinstance(result, ToolResultBlock) + assert result.tool_use_id == "t1" and result.is_error is False + + +def test_the_file_grows_linearly_not_quadratically(tmp_path): + """The log this replaces re-serialised the whole history every turn. + + At 40 messages that is 800 message-copies on disk; here it is 40 lines. + """ + path = tmp_path / "s.jsonl" + session = Session(JsonlSessionStore.create(path, "sess")) + for i in range(40): + session.append_message(say(f"message {i}")) + + lines = [line for line in path.read_text().splitlines() if line.strip()] + assert len(lines) == 41 # header + one line per entry + + # Each message appears exactly once in the file. + assert path.read_text().count('"message 0"') == 1 + + +def test_a_missing_session_file_is_a_typed_error(tmp_path): + with pytest.raises(SessionStoreError) as excinfo: + JsonlSessionStore.open(tmp_path / "nope.jsonl") + assert excinfo.value.code == "not_found" + + +def test_a_file_without_a_header_is_rejected(tmp_path): + path = tmp_path / "bad.jsonl" + path.write_text('{"type": "message", "id": "a"}\n') + with pytest.raises(SessionStoreError) as excinfo: + JsonlSessionStore.open(path) + assert excinfo.value.code == "invalid_session" + + +def test_a_future_format_version_is_rejected_not_guessed(tmp_path): + path = tmp_path / "future.jsonl" + path.write_text(json.dumps({"type": "session", "version": 99, "id": "x"}) + "\n") + with pytest.raises(SessionStoreError) as excinfo: + JsonlSessionStore.open(path) + assert excinfo.value.code == "invalid_session" + + +def test_a_corrupt_entry_in_the_middle_names_the_line(tmp_path): + """Corruption anywhere but the tail means the file cannot be trusted.""" + path = tmp_path / "corrupt.jsonl" + good = json.dumps( + {"type": "message", "id": "a", "parent_id": None, "timestamp": "t"} + ) + path.write_text( + json.dumps({"type": "session", "version": 1, "id": "x"}) + + "\nnot json\n" + + good + + "\n" + ) + with pytest.raises(SessionStoreError) as excinfo: + JsonlSessionStore.open(path) + assert excinfo.value.code == "invalid_entry" + assert ":2" in str(excinfo.value) + + +def test_a_torn_final_line_costs_one_entry_not_the_file(tmp_path): + """A crash mid-append must not make every earlier entry unreadable. + + This is an append-only recovery log. Refusing the whole file because the + process died halfway through the last write is the opposite of the point. + """ + path = tmp_path / "torn.jsonl" + session = Session(JsonlSessionStore.create(path, "sess")) + session.append_message(say("q1")) + session.append_message(say("a1", "assistant")) + + with path.open("a") as handle: + handle.write('{"type": "message", "id": "hal') # power cut + + store = JsonlSessionStore.open(path) + recovered = Session(store) + + assert store.truncated is True + assert texts(recovered.build_context()) == ["q1", "a1"] + + +def test_open_or_create_does_not_destroy_an_existing_session(tmp_path): + """`create` truncates, which on a restart path is the worst possible move.""" + path = tmp_path / "s.jsonl" + Session(JsonlSessionStore.create(path, "sess")).append_message(say("q1")) + + reopened = Session(JsonlSessionStore.open_or_create(path, "sess")) + assert texts(reopened.build_context()) == ["q1"] + + fresh = Session(JsonlSessionStore.open_or_create(tmp_path / "new.jsonl", "sess2")) + assert fresh.build_context() == [] + + +def test_an_unserializable_custom_payload_is_a_typed_error(tmp_path): + """Every store failure is a SessionStoreError, not a bare TypeError.""" + session = Session(JsonlSessionStore.create(tmp_path / "s.jsonl", "sess")) + with pytest.raises(SessionStoreError) as excinfo: + session.append_custom("bad", {"handle": object()}) + assert excinfo.value.code == "unserializable_entry" + + +def test_a_cycle_in_a_hand_edited_file_is_detected(store): + """Unreachable through the API, reachable by editing a file.""" + from data_harness.core.session.entries import MessageEntry as ME + + store.append(ME(id="a", parent_id=None, message=say("q1"))) + store._by_id["a"] = ME(id="a", parent_id="a", message=say("q1")) + with pytest.raises(SessionStoreError) as excinfo: + store.path_to_root("a") + assert excinfo.value.code == "cycle" + + +def test_an_unknown_entry_type_is_rejected(tmp_path): + path = tmp_path / "unknown.jsonl" + path.write_text( + json.dumps({"type": "session", "version": 1, "id": "x"}) + + "\n" + + json.dumps({"type": "from_the_future", "id": "a", "parent_id": None}) + + "\n" + ) + with pytest.raises(SessionStoreError) as excinfo: + JsonlSessionStore.open(path) + assert excinfo.value.code == "invalid_entry" + + +# ── the harness records into the session ──────────────────────────────────── + + +def test_a_run_records_every_message_it_sent(tmp_path): + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.run("go") + + # The working copy and the durable log agree. They are two representations + # of one thing, and this is what stops them becoming two sources of truth. + assert harness.session.build_context() == harness.messages + assert [m.role for m in harness.session.build_context()] == [ + "user", + "assistant", + "user", + "assistant", + ] + + +def test_a_run_records_what_each_turn_cost(tmp_path): + harness = Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.run("go") + + turns = [e for e in harness.session.store.entries() if isinstance(e, TurnEntry)] + assert [t.turn for t in turns] == [1, 2] + assert turns[0].stop_reason == "tool_use" + assert turns[1].stop_reason == "end_turn" + assert turns[0].visible_tools == ["echo"] + + +def test_a_session_spans_several_runs(tmp_path): + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("one"), FakeAdapter.text("two")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.run("first") + harness.run("second") + + recorded = [ + e.message.content[0].text + for e in harness.session.store.entries() + if isinstance(e, MessageEntry) + ] + assert recorded == ["first", "one", "second", "two"] + + +def test_a_conversation_resumes_across_a_restart(tmp_path): + """The headline capability, and what the old write-only log could not do.""" + path = tmp_path / "session.jsonl" + + first = Harness( + adapter=FakeAdapter([FakeAdapter.text("4")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + session=Session(JsonlSessionStore.create(path, "sess-1")), + ) + first.run("what is 2+2") + + # A different process, holding nothing but the file. + second = Harness( + adapter=FakeAdapter([FakeAdapter.text("6")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + session=Session(JsonlSessionStore.open(path)), + ) + assert texts(second.messages) == ["what is 2+2", "4"] + + second.ask("and 3+3") + + assert texts(second.messages) == ["what is 2+2", "4", "and 3+3", "6"] + assert second.session.stats().turns == 2 + + +def test_a_resumed_session_can_fork_instead_of_continuing(tmp_path): + """Re-ask an earlier question differently, keeping the first attempt.""" + path = tmp_path / "session.jsonl" + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("first answer")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + session=Session(JsonlSessionStore.create(path, "sess-1")), + ) + harness.run("the question") + original_leaf = harness.session.leaf_id + + session = Session(JsonlSessionStore.open(path)) + first_message = next( + e for e in session.store.entries() if isinstance(e, MessageEntry) + ) + session.move_to(first_message.id) + + retried = Harness( + adapter=FakeAdapter([FakeAdapter.text("second answer")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + session=session, + ) + retried.ask("") + + assert texts(retried.session.build_context())[-1] == "second answer" + assert texts(retried.session.build_context(original_leaf))[-1] == "first answer" + + +@pytest.mark.asyncio +async def test_a_streamed_run_records_the_same_way(tmp_path): + from data_harness.data.harness import AsyncHarness + from data_harness.llm.testing import FakeAsyncAdapter + + harness = AsyncHarness( + adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("streamed")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + async for _ in harness.run_stream("go"): + pass + + assert texts(harness.session.build_context()) == ["go", "streamed"] + assert harness.session.stats().turns == 1 + + +# ── the review's surviving mutants ────────────────────────────────────────── + + +def test_message_roles_survive_a_round_trip(tmp_path): + """Nothing pinned this, so an encoder writing every role as `user` passed.""" + path = tmp_path / "s.jsonl" + session = Session(JsonlSessionStore.create(path, "sess")) + session.append_message(say("q", "user")) + session.append_message(say("a", "assistant")) + session.append_message(say("q2", "user")) + + roles = [m.role for m in Session(JsonlSessionStore.open(path)).build_context()] + assert roles == ["user", "assistant", "user"] + + +def test_a_failed_tool_result_round_trips_as_failed(tmp_path): + """`is_error` defaults to False, so asserting False proved nothing.""" + from data_harness.llm.types import ToolResultBlock, ToolUseBlock + + path = tmp_path / "s.jsonl" + session = Session(JsonlSessionStore.create(path, "sess")) + session.append_message( + Message( + role="assistant", + content=[ToolUseBlock(tool_use_id="t1", tool_name="boom", tool_input={})], + ) + ) + session.append_message( + Message( + role="user", + content=[ + ToolResultBlock(tool_use_id="t1", content="it broke", is_error=True) + ], + ) + ) + + result = Session(JsonlSessionStore.open(path)).build_context()[1].content[0] + assert result.is_error is True + assert result.content == "it broke" + + +# ── a derived context is always something a provider will accept ──────────── + + +def test_an_orphaned_tool_call_is_dropped_from_the_context(tmp_path): + """A run killed mid-tool leaves a call with no result on disk. + + Resuming must not hand the provider a transcript it rejects outright, + which is precisely the session a user most wants back. + """ + from data_harness.llm.types import ToolUseBlock + + def explode(value: str) -> str: + raise KeyboardInterrupt + + spec = ToolSpec( + name="boom", + description="Die.", + input_schema={ + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, + handler=explode, + ) + path = tmp_path / "s.jsonl" + harness = Harness( + adapter=FakeAdapter([FakeAdapter.tool_use("t1", "boom", {"value": "x"})]), + system="sys", + tools=[spec], + run_dir=str(tmp_path), + session=Session(JsonlSessionStore.create(path, "sess")), + ) + with pytest.raises(KeyboardInterrupt): + harness.run("go") + + # The orphan really is on disk. + stored = Session(JsonlSessionStore.open(path)) + raw = [ + b + for e in stored.store.entries() + if isinstance(e, MessageEntry) + for b in e.message.content + ] + assert any(isinstance(b, ToolUseBlock) for b in raw) + + # It is not in the context a resumed run would send. + context = stored.build_context() + assert not any(isinstance(b, ToolUseBlock) for m in context for b in m.content) + assert texts(context) == ["go"] + + +def test_a_compaction_that_orphans_a_tool_result_drops_it(store): + from data_harness.llm.types import ToolResultBlock, ToolUseBlock + + session = Session(store) + session.append_message(say("q1")) + session.append_message( + Message( + role="assistant", + content=[ToolUseBlock(tool_use_id="t1", tool_name="echo", tool_input={})], + ) + ) + orphan = session.append_message( + Message( + role="user", + content=[ToolResultBlock(tool_use_id="t1", content="r", is_error=False)], + ) + ) + session.append_compaction("summary", first_kept_entry_id=orphan, tokens_before=1) + session.append_message(say("next")) + + context = session.build_context() + assert not any(isinstance(b, ToolResultBlock) for m in context for b in m.content) + + +# ── compaction is validated and composes ──────────────────────────────────── + + +def test_compacting_from_an_entry_off_the_path_is_rejected(store): + """A stale id would silently amnesia the agent rather than fail.""" + session = Session(store) + root = session.append_message(say("q1")) + session.append_message(say("a1", "assistant")) + session.move_to(root) + abandoned_branch_entry = None + for entry in store.entries(): + if isinstance(entry, MessageEntry) and entry.message.role == "assistant": + abandoned_branch_entry = entry.id + + with pytest.raises(SessionStoreError) as excinfo: + session.append_compaction("s", abandoned_branch_entry, 1) + assert excinfo.value.code == "invalid_cut" + + +def test_stacked_compactions_leave_one_summary_in_order(store): + """An older compaction inside the kept tail must not be replayed. + + The kept tail deliberately starts *before* the first compaction, so the + first compaction sits inside it. Replaying it emitted the older summary + after the newer one and resurrected the entries it had dropped. + """ + session = Session(store) + session.append_message(say("q1")) + keep = session.append_message(say("a1", "assistant")) + session.append_compaction( + "first summary", first_kept_entry_id=None, tokens_before=1 + ) + session.append_message(say("q2")) + session.append_compaction( + "second summary", first_kept_entry_id=keep, tokens_before=2 + ) + session.append_message(say("q3")) + + context = texts(session.build_context()) + assert sum("summary" in c for c in context) == 1, context + assert "second summary" in context[0] + assert "first summary" not in " ".join(context) + assert context[1:] == ["a1", "q2", "q3"] + + +# ── the working copy and the log agree ────────────────────────────────────── + + +def test_a_second_run_starts_a_new_branch_not_a_false_continuation(tmp_path): + """`run` resets the conversation, so the log must not imply otherwise. + + Otherwise the tree claims a continuity the model was never shown, and + resuming replays a context that was never sent. + """ + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("one"), FakeAdapter.text("two")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.run("first") + assert harness.session.build_context() == harness.messages + + harness.run("second") + assert harness.session.build_context() == harness.messages + assert texts(harness.session.build_context()) == ["second", "two"] + + +def test_the_log_and_the_working_copy_agree_on_every_entry_point(tmp_path): + harness = Harness( + adapter=FakeAdapter( + [FakeAdapter.text("a"), FakeAdapter.text("b"), FakeAdapter.text("c")] + ), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.run("q1") + assert harness.session.build_context() == harness.messages + harness.ask("q2") + assert harness.session.build_context() == harness.messages + harness.ask("q3") + assert harness.session.build_context() == harness.messages + + +def test_a_reminder_appended_to_a_recorded_message_is_still_logged(tmp_path): + """Entries are immutable, so the reminder becomes its own entry. + + Without it the JSONL store, which serialises on write, would show a + prompt the model never actually saw. + """ + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + max_turns=2, + ) + harness.register_reminder(lambda turn, max_turns: "stay on task") + harness.run("go") + + reminders = harness.session.custom_entries("reminder") + assert len(reminders) == 1 + # The built-in max-turn nag rides along in the same suffix block. + assert reminders[0].data["text"].startswith("stay on task") + assert reminders[0].data["turn"] == 1 + + +def test_both_stores_agree_after_a_message_is_mutated(tmp_path): + """Stores snapshot on write, so neither can be rewritten after the fact.""" + message = say("original") + memory = Session(MemorySessionStore("m")) + on_disk = Session(JsonlSessionStore.create(tmp_path / "s.jsonl", "d")) + memory.append_message(message) + on_disk.append_message(message) + + message.content.append(TextBlock(text="added later")) + + # Compare every block, not just the first: the earlier version of this + # test used a helper that read `content[0]` and so could not see an + # appended block at all. + def blocks(session): + return [[b.text for b in m.content] for m in session.build_context()] + + assert blocks(memory) == [["original"]] + assert blocks(on_disk) == [["original"]] diff --git a/tests/test_sql.py b/tests/test_sql.py index 12ab663..57fb331 100644 --- a/tests/test_sql.py +++ b/tests/test_sql.py @@ -7,9 +7,9 @@ import pandas as pd from data_harness import Agent -from data_harness.cache import SessionCache -from data_harness.testing import FakeAdapter -from data_harness.tools.sql import make_sql_query_spec +from data_harness.data.cache import SessionCache +from data_harness.data.tools.sql import make_sql_query_spec +from data_harness.llm.testing import FakeAdapter def _sales() -> pd.DataFrame: diff --git a/tests/test_streaming.py b/tests/test_streaming.py index 94ce795..1fdd07d 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -23,14 +23,14 @@ from collections.abc import AsyncGenerator from typing import Any -from data_harness.agent import AsyncAgent -from data_harness.loop import AsyncHarness -from data_harness.providers.base import ( +from data_harness.app.agent import AsyncAgent +from data_harness.data.harness import AsyncHarness +from data_harness.llm.providers.base import ( AsyncProviderAdapter, NormalizedResponse, StopReason, ) -from data_harness.streaming import ( +from data_harness.llm.streaming import ( ContentBlockDeltaEvent, ContentBlockStartEvent, ContentBlockStopEvent, @@ -43,8 +43,8 @@ ToolResultEvent, accumulate_stream_events, ) -from data_harness.testing import FakeAsyncAdapter -from data_harness.types import Message, TextBlock, ToolSpec, ToolUseBlock +from data_harness.llm.testing import FakeAsyncAdapter +from data_harness.llm.types import Message, TextBlock, ToolSpec, ToolUseBlock # --------------------------------------------------------------------------- # Helpers @@ -681,7 +681,7 @@ def tool_b(y: str) -> str: ), ] - from data_harness.providers.base import NormalizedResponse + from data_harness.llm.providers.base import NormalizedResponse two_tool_response = NormalizedResponse( stop_reason=StopReason.TOOL_USE, @@ -965,6 +965,73 @@ async def stream_events(self, system, messages, tools): # Stream should have stopped (no infinite loop) assert len(events) < 100 + # The stream stopping is not enough: the run must still be accounted + # for. This test used to end at the line above, which is how the + # streaming loop got away with discarding its RunResult on provider + # errors, and with it the tokens already spent. + result = harness.last_result + assert result is not None + assert result.status == "error" + assert "provider blew up" in (result.error or "") + assert result.run_file == harness.run_file + + async def test_unknown_event_types_are_ignored_not_fatal(self, tmp_path): + """The accumulator skips events it does not recognise, on purpose. + + Providers add event types over time; an unknown one must not break a + run. It contributes nothing, so the turn assembles from the rest. + """ + + class ChattyAdapter(AsyncProviderAdapter): + async def chat(self, system, messages, tools): + return FakeAsyncAdapter.text("x") + + def format_cache_control(self, obj): + return obj + + async def stream_events(self, system, messages, tools): + yield MessageStartEvent() + yield "some future event type" + yield MessageStopEvent() + + harness = AsyncHarness( + adapter=ChattyAdapter(), system="s", tools=[], run_dir=str(tmp_path) + ) + async for _ in harness.run_stream("q"): + pass + + result = harness.last_result + assert result is not None + assert result.status == "success" + + async def test_a_failure_while_assembling_the_turn_is_reported( + self, tmp_path, monkeypatch + ): + """Accumulation runs inside the loop's provider try-block, deliberately. + + If assembling the turn blows up, that is a provider failure and belongs + in the RunResult. Outside the try it would escape into the caller's + `async for` and bypass all run accounting. + """ + + def explode(_events): + raise ValueError("cannot assemble this turn") + + monkeypatch.setattr("data_harness.core.loop.accumulate_stream_events", explode) + harness = AsyncHarness( + adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("hi")]), + system="s", + tools=[], + run_dir=str(tmp_path), + ) + async for _ in harness.run_stream("q"): + pass + + result = harness.last_result + assert result is not None + assert result.status == "error" + assert "cannot assemble this turn" in (result.error or "") + # --------------------------------------------------------------------------- # 12. AsyncAgent.run_stream() passthrough @@ -1023,7 +1090,7 @@ def raw_tool(v: str) -> str: raw_calls.append(v) return "raw_result" - from data_harness.loop import AsyncHarness + from data_harness.data.harness import AsyncHarness harness = AsyncHarness( adapter=adapter2, diff --git a/tests/test_testing.py b/tests/test_testing.py index e2de393..2f68ed8 100644 --- a/tests/test_testing.py +++ b/tests/test_testing.py @@ -1,10 +1,14 @@ -"""Tests for `data_harness.testing` — the public FakeAdapter for docs and tests.""" +"""Tests for `data_harness.llm.testing` — the public FakeAdapter for docs and tests.""" from __future__ import annotations -from data_harness.providers.base import NormalizedResponse, ProviderAdapter, StopReason -from data_harness.testing import FakeAdapter -from data_harness.types import TextBlock, ToolUseBlock +from data_harness.llm.providers.base import ( + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import TextBlock, ToolUseBlock def _text_response(text: str) -> NormalizedResponse: diff --git a/tests/test_tool_annotations.py b/tests/test_tool_annotations.py index dced6c6..af17dc1 100644 --- a/tests/test_tool_annotations.py +++ b/tests/test_tool_annotations.py @@ -10,14 +10,14 @@ import pytest -from data_harness.loop import Harness -from data_harness.providers.base import StopReason -from data_harness.testing import FakeAdapter -from data_harness.types import TextBlock, ToolSpec +from data_harness.data.harness import Harness +from data_harness.llm.providers.base import StopReason +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import TextBlock, ToolSpec def make_text_response(text: str): - from data_harness.providers.base import NormalizedResponse + from data_harness.llm.providers.base import NormalizedResponse return NormalizedResponse( stop_reason=StopReason.END_TURN, @@ -41,12 +41,12 @@ def read_jsonl(path: str) -> list[dict]: class TestToolAnnotations: def test_import(self): - from data_harness.types import ToolAnnotations + from data_harness.llm.types import ToolAnnotations assert ToolAnnotations is not None def test_all_fields_optional(self): - from data_harness.types import ToolAnnotations + from data_harness.llm.types import ToolAnnotations ann = ToolAnnotations() assert ann.title is None @@ -56,7 +56,7 @@ def test_all_fields_optional(self): assert ann.open_world is None def test_explicit_fields(self): - from data_harness.types import ToolAnnotations + from data_harness.llm.types import ToolAnnotations ann = ToolAnnotations( title="Echo", @@ -70,7 +70,7 @@ def test_explicit_fields(self): assert ann.cache_mutating is False def test_frozen(self): - from data_harness.types import ToolAnnotations + from data_harness.llm.types import ToolAnnotations ann = ToolAnnotations(title="Echo") with pytest.raises((AttributeError, TypeError)): @@ -92,7 +92,7 @@ def test_annotations_field_defaults_to_none(self): assert spec.annotations is None def test_annotations_field_set(self): - from data_harness.types import ToolAnnotations + from data_harness.llm.types import ToolAnnotations ann = ToolAnnotations(read_only=True) spec = ToolSpec( @@ -105,7 +105,7 @@ def test_annotations_field_set(self): assert spec.annotations.read_only is True def test_to_provider_dict_excludes_annotations(self): - from data_harness.types import ToolAnnotations + from data_harness.llm.types import ToolAnnotations ann = ToolAnnotations(read_only=True, destructive=False) spec = ToolSpec( @@ -128,24 +128,24 @@ def test_to_provider_dict_excludes_annotations(self): class TestBuiltinToolAnnotations: def test_list_variables_read_only(self): - from data_harness.cache import SessionCache - from data_harness.tools.variables import make_list_variables_spec + from data_harness.data.cache import SessionCache + from data_harness.data.tools.variables import make_list_variables_spec spec = make_list_variables_spec(SessionCache()) assert spec.annotations is not None assert spec.annotations.read_only is True def test_python_interpreter_cache_mutating(self): - from data_harness.cache import SessionCache - from data_harness.tools.interpreter import PythonInterpreter + from data_harness.data.cache import SessionCache + from data_harness.data.tools.interpreter import PythonInterpreter spec = PythonInterpreter.make_tool_spec(SessionCache()) assert spec.annotations is not None assert spec.annotations.cache_mutating is True def test_python_interpreter_not_open_world(self): - from data_harness.cache import SessionCache - from data_harness.tools.interpreter import PythonInterpreter + from data_harness.data.cache import SessionCache + from data_harness.data.tools.interpreter import PythonInterpreter spec = PythonInterpreter.make_tool_spec(SessionCache()) assert spec.annotations.open_world is False @@ -158,7 +158,7 @@ def test_python_interpreter_not_open_world(self): class TestAnnotationsInLog: def test_annotations_serialised_in_jsonl(self, tmp_path): - from data_harness.types import ToolAnnotations + from data_harness.llm.types import ToolAnnotations ann = ToolAnnotations(title="Echo tool", read_only=True) echo_spec = ToolSpec( diff --git a/tests/test_tool_connectors.py b/tests/test_tool_connectors.py index 510d145..99a8e75 100644 --- a/tests/test_tool_connectors.py +++ b/tests/test_tool_connectors.py @@ -2,9 +2,9 @@ import json -from data_harness.cache import SessionCache -from data_harness.tools.connectors import ConnectorRegistry -from data_harness.types import ToolSpec +from data_harness.data.cache import SessionCache +from data_harness.data.tools.connectors import ConnectorRegistry +from data_harness.llm.types import ToolSpec def make_market_data_connector(): diff --git a/tests/test_tool_interpreter.py b/tests/test_tool_interpreter.py index edadfc9..f437e0f 100644 --- a/tests/test_tool_interpreter.py +++ b/tests/test_tool_interpreter.py @@ -3,16 +3,20 @@ import pandas as pd import pytest -from data_harness.cache import SessionCache -from data_harness.loop import Harness -from data_harness.providers.base import NormalizedResponse, ProviderAdapter, StopReason -from data_harness.tools.interpreter import ( +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.data.tools.interpreter import ( _EMPTY_OUTPUT_GUIDANCE, _LOCALS_ERROR, PythonInterpreter, PythonInterpreterError, ) -from data_harness.types import TextBlock, ToolResultBlock, ToolUseBlock +from data_harness.llm.providers.base import ( + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.types import TextBlock, ToolResultBlock, ToolUseBlock # --------------------------------------------------------------------------- # Helpers @@ -75,7 +79,7 @@ def _run_interpreter_via_harness( harness = Harness(adapter=adapter, system="sys", tools=[spec], cache=cache) harness.run("go") # The tool result is the second-to-last message (user message with tool result) - for msg in reversed(harness._messages): + for msg in reversed(harness.messages): if msg.role == "user": for block in msg.content: if isinstance(block, ToolResultBlock): diff --git a/tests/test_tool_planner.py b/tests/test_tool_planner.py index 5b176d0..e542cae 100644 --- a/tests/test_tool_planner.py +++ b/tests/test_tool_planner.py @@ -1,6 +1,6 @@ """Tests for the Planner tool.""" -from data_harness.tools.planner import Planner +from data_harness.data.tools.planner import Planner class TestPlannerBasic: diff --git a/tests/test_tool_subagent.py b/tests/test_tool_subagent.py index 9678568..5315d79 100644 --- a/tests/test_tool_subagent.py +++ b/tests/test_tool_subagent.py @@ -4,12 +4,16 @@ import copy -from data_harness.cache import SessionCache -from data_harness.providers.base import NormalizedResponse, ProviderAdapter, StopReason -from data_harness.tools.connectors import ConnectorRegistry -from data_harness.tools.interpreter import PythonInterpreter -from data_harness.tools.subagent import make_subagent_spec -from data_harness.types import Message, TextBlock, ToolSpec, ToolUseBlock +from data_harness.data.cache import SessionCache +from data_harness.data.tools.connectors import ConnectorRegistry +from data_harness.data.tools.interpreter import PythonInterpreter +from data_harness.data.tools.subagent import make_subagent_spec +from data_harness.llm.providers.base import ( + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.types import Message, TextBlock, ToolSpec, ToolUseBlock class FakeAdapter(ProviderAdapter): diff --git a/tests/test_tool_variables.py b/tests/test_tool_variables.py index ad71caa..89c41ab 100644 --- a/tests/test_tool_variables.py +++ b/tests/test_tool_variables.py @@ -2,8 +2,8 @@ import pytest -from data_harness.cache import SessionCache -from data_harness.tools.variables import make_list_variables_spec +from data_harness.data.cache import SessionCache +from data_harness.data.tools.variables import make_list_variables_spec class TestListVariables: diff --git a/tests/test_turn_summary.py b/tests/test_turn_summary.py index 8dfb823..afb4ec7 100644 --- a/tests/test_turn_summary.py +++ b/tests/test_turn_summary.py @@ -8,11 +8,11 @@ import json from pathlib import Path -from data_harness.cache import SessionCache -from data_harness.loop import Harness -from data_harness.providers.base import NormalizedResponse, StopReason -from data_harness.testing import FakeAdapter -from data_harness.types import TextBlock, ToolSpec, ToolUseBlock +from data_harness.data.cache import SessionCache +from data_harness.data.harness import Harness +from data_harness.llm.providers.base import NormalizedResponse, StopReason +from data_harness.llm.testing import FakeAdapter +from data_harness.llm.types import TextBlock, ToolSpec, ToolUseBlock def make_text_response( diff --git a/tests/test_types.py b/tests/test_types.py index 8facbce..b10d571 100644 --- a/tests/test_types.py +++ b/tests/test_types.py @@ -1,8 +1,8 @@ import pytest -from data_harness.providers.base import ProviderAdapter -from data_harness.serialize import to_jsonable -from data_harness.types import ( +from data_harness.core.serialize import to_jsonable +from data_harness.llm.providers.base import ProviderAdapter +from data_harness.llm.types import ( Message, TextBlock, ToolResultBlock, diff --git a/tests/test_unified_loop.py b/tests/test_unified_loop.py new file mode 100644 index 0000000..b8da44a --- /dev/null +++ b/tests/test_unified_loop.py @@ -0,0 +1,955 @@ +"""Phase 1: one loop, one agent base. + +These tests exist because the sync/async and stream/non-stream copies had +already drifted apart before they were collapsed. Each test pins one of the +specific defects that drift produced, so a future split would fail loudly. +""" + +from __future__ import annotations + +import asyncio +import sqlite3 +import threading + +import pytest + +from data_harness.app.agent import Agent, AgentSession, AsyncAgent, AsyncAgentSession +from data_harness.app.quickstart import resolve_adapter, resolve_async_adapter +from data_harness.data.harness import AsyncHarness, Harness, run_coroutine_blocking +from data_harness.llm.providers.base import ( + AsyncProviderAdapter, + NormalizedResponse, + ProviderAdapter, + StopReason, +) +from data_harness.llm.streaming import MessageDeltaEvent, ToolResultEvent +from data_harness.llm.testing import FakeAdapter, FakeAsyncAdapter +from data_harness.llm.types import Message, ToolSpec, ToolUseBlock + +# ── helpers ───────────────────────────────────────────────────────────────── + + +class ExplodingAsyncAdapter(FakeAsyncAdapter): + """Serves scripted responses, then raises. Models a provider dying mid-run.""" + + def __init__(self, responses, error: Exception) -> None: + super().__init__(responses) + self._error = error + + async def chat(self, system, messages, tools) -> NormalizedResponse: + if not self._responses: + raise self._error + return await super().chat(system, messages, tools) + + +def echo_spec() -> ToolSpec: + return ToolSpec( + name="echo", + description="Echo the input back.", + input_schema={ + "type": "object", + "properties": {"value": {"type": "string"}}, + "required": ["value"], + }, + handler=lambda value: value, + ) + + +# ── API parity: the two agents must not drift again ───────────────────────── + + +def test_agent_and_async_agent_expose_the_same_features(): + """Both agents share one base, so their public surfaces differ only by design. + + `AsyncAgent` previously lacked from_dataframe, from_csv, subagents, MCP, + the replay cache, close(), and explain(). Nothing flagged it. + """ + sync_api = {n for n in dir(Agent) if not n.startswith("_")} + async_api = {n for n in dir(AsyncAgent) if not n.startswith("_")} + + assert sync_api - async_api == {"session"} + assert async_api - sync_api == {"async_session", "run_stream"} + + +@pytest.mark.parametrize( + "feature", + [ + "from_dataframe", + "from_csv", + "enable_subagents", + "enable_planner", + "enable_sql", + "enable_cache", + "add_mcp_server", + "connector", + "close", + "explain", + "exec_cache", + ], +) +def test_async_agent_has_feature(feature): + assert hasattr(AsyncAgent, feature) + + +def test_async_agent_accepts_the_sandbox_and_gate_options(tmp_path): + """`on_code`, `code_only`, `execution` used to be Agent-only constructor args.""" + agent = AsyncAgent( + adapter=FakeAsyncAdapter([]), + system="sys", + run_dir=str(tmp_path), + execution="inprocess", + on_code=lambda code: True, + code_only=True, + ) + assert agent._code_only is True + + +# ── streaming keeps its accounting ────────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_stream_records_run_result_on_success(tmp_path): + """A streamed run exposes usage and status via `last_result`. + + The event protocol carries no run-level summary, so before this the only + way to get usage out of a stream was to reassemble it from raw events. + """ + harness = AsyncHarness( + adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + + events = [evt async for evt in harness.run_stream("go")] + + assert any(isinstance(e, MessageDeltaEvent) for e in events) + result = harness.last_result + assert result is not None + assert result.status == "success" + assert result.text == "done" + assert result.usage.input_tokens == 10 + assert result.usage.output_tokens == 5 + + +@pytest.mark.asyncio +async def test_stream_keeps_usage_when_the_provider_fails_mid_run(tmp_path): + """Tokens already billed before a provider error must still be accounted for. + + The old streaming loop caught the exception, logged, and returned without + building a RunResult, so usage from every completed turn was lost. A + metered deployment silently under-billed itself. + """ + adapter = ExplodingAsyncAdapter( + [FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "hi"})], + RuntimeError("provider down"), + ) + harness = AsyncHarness( + adapter=adapter, + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + + events = [evt async for evt in harness.ask_stream("go")] + + assert any(isinstance(e, ToolResultEvent) for e in events) + result = harness.last_result + assert result is not None + assert result.status == "error" + assert "provider down" in (result.error or "") + # Turn 1 completed and was billed before turn 2 failed. + assert result.usage.input_tokens == 10 + assert result.usage.output_tokens == 5 + + +@pytest.mark.asyncio +async def test_non_streaming_error_also_reports_usage(tmp_path): + """The non-streaming path always behaved correctly here; it must keep doing so.""" + adapter = ExplodingAsyncAdapter( + [FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "hi"})], + RuntimeError("provider down"), + ) + harness = AsyncHarness( + adapter=adapter, system="sys", tools=[echo_spec()], run_dir=str(tmp_path) + ) + + result = await harness.run_result("go") + + assert result.status == "error" + assert result.usage.input_tokens == 10 + + +@pytest.mark.asyncio +async def test_session_ask_stream_updates_session_totals(tmp_path): + """A streamed session turn counts toward `turns` and `last_result`.""" + agent = AsyncAgent( + adapter=FakeAsyncAdapter( + [FakeAsyncAdapter.text("one"), FakeAsyncAdapter.text("two")] + ), + system="sys", + run_dir=str(tmp_path), + ) + session = agent.async_session() + + async for _ in session.ask_stream("first"): + pass + + assert session.turns == 1 + assert session.last_result is not None + assert session.last_result.text == "one" + assert session.last_result.usage.output_tokens == 5 + + await session.ask_result("second") + assert session.turns == 2 + + +# ── streaming and non-streaming agree ─────────────────────────────────────── + + +@pytest.mark.asyncio +async def test_stream_and_result_paths_agree(tmp_path): + """Same script, same outcome, whichever entry point is used. + + They are one loop now, so this asserts the parameterisation did not change + behaviour rather than that two implementations happen to match. + """ + + def script(): + return [ + FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAsyncAdapter.text("final answer"), + ] + + streamed = AsyncHarness( + adapter=FakeAsyncAdapter(script()), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + async for _ in streamed.run_stream("go"): + pass + + direct = AsyncHarness( + adapter=FakeAsyncAdapter(script()), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + direct_result = await direct.run_result("go") + + stream_result = streamed.last_result + assert stream_result is not None + assert stream_result.text == direct_result.text == "final answer" + assert stream_result.turns == direct_result.turns == 2 + assert stream_result.usage == direct_result.usage + assert [m.role for m in streamed.messages] == [m.role for m in direct.messages] + + +# ── the two drivers agree ─────────────────────────────────────────────────── + + +def test_sync_and_async_harness_produce_identical_results(tmp_path): + """Both drivers run the same `_plan` generator, so they must agree.""" + + def script(cls): + return [ + cls.tool_use("tu_1", "echo", {"value": "hi"}), + cls.text("final answer"), + ] + + sync_result = Harness( + adapter=FakeAdapter(script(FakeAdapter)), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ).run_result("go") + + async_result = asyncio.run( + AsyncHarness( + adapter=FakeAsyncAdapter(script(FakeAsyncAdapter)), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ).run_result("go") + ) + + assert sync_result.text == async_result.text + assert sync_result.turns == async_result.turns + assert sync_result.status == async_result.status + + +def test_sync_harness_still_calls_the_sync_adapter(tmp_path): + adapter = FakeAdapter([FakeAdapter.text("hi")]) + harness = Harness(adapter=adapter, system="sys", tools=[], run_dir=str(tmp_path)) + + harness.run("go") + + assert len(adapter.calls) == 1 + assert adapter.calls[0]["system"] == "sys" + + +def test_sync_harness_works_inside_a_running_event_loop(tmp_path): + """A notebook kernel or async web handler already has a loop running.""" + + async def outer(): + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("from inside a loop")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + return harness.run("go") + + assert asyncio.run(outer()) == "from inside a loop" + + +def test_run_coroutine_blocking_without_a_running_loop(): + async def coro(): + return 7 + + assert run_coroutine_blocking(coro()) == 7 + + +# ── the sync driver must stay genuinely synchronous ───────────────────────── + + +def test_sync_driver_runs_tool_handlers_on_the_calling_thread(tmp_path): + """Handlers must keep thread affinity. + + A connector holding a `sqlite3` connection built at setup time is the most + ordinary pattern this library has, and sqlite3 refuses cross-thread use. + Driving the sync path through an event loop moved handlers onto an + `asyncio.to_thread` worker and turned that into a tool error string. + """ + seen: list[str] = [] + + spec = ToolSpec( + name="whereami", + description="Report the executing thread.", + input_schema={"type": "object", "properties": {}}, + handler=lambda: seen.append(threading.current_thread().name) or "ok", + ) + + Harness( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "whereami", {}), + FakeAdapter.text("done"), + ] + ), + system="sys", + tools=[spec], + run_dir=str(tmp_path), + ).run("go") + + assert seen == [threading.current_thread().name] + + +def test_sync_driver_keeps_a_thread_bound_resource_usable(tmp_path): + """The concrete failure the thread-affinity rule protects against.""" + connection = sqlite3.connect(":memory:") + connection.execute("create table t (v integer)") + connection.execute("insert into t values (42)") + + spec = ToolSpec( + name="query", + description="Read the fixed row.", + input_schema={"type": "object", "properties": {}}, + handler=lambda: str(connection.execute("select v from t").fetchone()), + ) + + harness = Harness( + adapter=FakeAdapter( + [FakeAdapter.tool_use("tu_1", "query", {}), FakeAdapter.text("done")] + ), + system="sys", + tools=[spec], + run_dir=str(tmp_path), + ) + harness.run("go") + + tool_result = harness.messages[2].content[0] + assert tool_result.is_error is False + assert "42" in tool_result.content + + +def test_sync_driver_leaves_the_ambient_event_loop_alone(tmp_path): + """`asyncio.run` clears the thread's loop; the sync driver must not. + + A program that installs its own loop and then calls the *synchronous* API + used to find the loop gone afterwards, failing much later and elsewhere. + """ + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + try: + Harness( + adapter=FakeAdapter([FakeAdapter.text("hi")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ).run("go") + assert asyncio.get_event_loop() is loop + finally: + asyncio.set_event_loop(None) + loop.close() + + +def test_run_coroutine_blocking_restores_the_ambient_event_loop(): + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + + async def coro(): + return 1 + + try: + assert run_coroutine_blocking(coro()) == 1 + assert asyncio.get_event_loop() is loop + finally: + asyncio.set_event_loop(None) + loop.close() + + +def test_keyboard_interrupt_escapes_the_sync_driver(tmp_path): + """Ctrl-C must reach the caller, not be absorbed as a tool error. + + Tool failures are caught as `Exception` and reported to the model. + `KeyboardInterrupt` is a `BaseException` and must pass straight through, + and with inline dispatch it does so without waiting on a worker thread. + """ + + def interrupt(): + raise KeyboardInterrupt + + spec = ToolSpec( + name="boom", + description="Interrupt.", + input_schema={"type": "object", "properties": {}}, + handler=interrupt, + ) + + with pytest.raises(KeyboardInterrupt): + Harness( + adapter=FakeAdapter( + [FakeAdapter.tool_use("tu_1", "boom", {}), FakeAdapter.text("done")] + ), + system="sys", + tools=[spec], + run_dir=str(tmp_path), + ).run("go") + + +@pytest.mark.asyncio +async def test_async_driver_offloads_blocking_handlers(tmp_path): + """The async driver must NOT run blocking handlers on the event loop. + + The mirror image of the sync rule: a long pandas call on the loop thread + would stall every other task sharing it. + """ + seen: list[str] = [] + + spec = ToolSpec( + name="whereami", + description="Report the executing thread.", + input_schema={"type": "object", "properties": {}}, + handler=lambda: seen.append(threading.current_thread().name) or "ok", + ) + + await AsyncHarness( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use("tu_1", "whereami", {}), + FakeAsyncAdapter.text("done"), + ] + ), + system="sys", + tools=[spec], + run_dir=str(tmp_path), + ).run("go") + + assert seen and seen != [threading.current_thread().name] + + +# ── the drivers stay overridable ──────────────────────────────────────────── + + +def test_sync_driver_dispatch_is_overridable(tmp_path): + """`_call_tool` is the documented seam for sandboxing or instrumentation. + + A facade that forwarded only public methods would accept the subclass and + silently never call its override. + """ + + class Instrumented(Harness): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.dispatched: list[str] = [] + + def _call_tool(self, call): + self.dispatched.append(call.tool_name) + return super()._call_tool(call) + + harness = Instrumented( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAdapter.text("d"), + ] + ), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + harness.run("go") + + assert harness.dispatched == ["echo"] + + +@pytest.mark.asyncio +async def test_async_driver_dispatch_is_overridable(tmp_path): + class Instrumented(AsyncHarness): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + self.dispatched: list[str] = [] + + async def _call_tool(self, call): + self.dispatched.append(call.tool_name) + return await super()._call_tool(call) + + harness = Instrumented( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAsyncAdapter.text("d"), + ] + ), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + await harness.run("go") + + assert harness.dispatched == ["echo"] + + +@pytest.mark.parametrize( + "attribute", + [ + "_messages", + "_tools", + "_system", + "_max_turns", + "_environment", + "_reminders", + "_run_file", + "_on_code", + "_code_only", + ], +) +def test_both_drivers_carry_the_loop_state(attribute, tmp_path): + """Loop state lives on the shared base, so neither driver is a hollow shell.""" + sync = Harness( + adapter=FakeAdapter([]), system="sys", tools=[], run_dir=str(tmp_path) + ) + asynchronous = AsyncHarness( + adapter=FakeAsyncAdapter([]), system="sys", tools=[], run_dir=str(tmp_path) + ) + assert hasattr(sync, attribute) + assert hasattr(asynchronous, attribute) + + +# ── features that only Agent used to have, now on AsyncAgent ──────────────── + + +@pytest.mark.asyncio +async def test_async_agent_from_dataframe_preloads_handles(tmp_path): + pd = pytest.importorskip("pandas") + agent = AsyncAgent.from_dataframe( + pd.DataFrame({"a": [1, 2, 3]}), + adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("ok")]), + run_dir=str(tmp_path), + ) + assert agent.cache.list_handles() + assert await agent.run("summarise") == "ok" + + +@pytest.mark.asyncio +async def test_async_agent_code_only_blocks_execution(tmp_path): + """The interpreter approval gate used to be unreachable from AsyncAgent.""" + agent = AsyncAgent( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use( + "tu_1", "python_interpreter", {"code": "save('x', 1)"} + ), + FakeAsyncAdapter.text("done"), + ] + ), + system="sys", + run_dir=str(tmp_path), + code_only=True, + ) + + await agent.run_result("compute") + + assert "x" not in agent.cache.handle_names() + harness = agent.last_harness + assert harness is not None + tool_result = harness.messages[2].content[0] + assert "DRY RUN" in tool_result.content + + +@pytest.mark.asyncio +async def test_async_agent_subagent_accepts_a_sync_adapter_factory(tmp_path): + """A sync adapter factory must still work under an async parent.""" + agent = AsyncAgent( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use("tu_1", "subagent", {"task": "work"}), + FakeAsyncAdapter.text("parent done"), + ] + ), + system="sys", + run_dir=str(tmp_path), + ) + agent.enable_subagents( + adapter_factory=lambda: FakeAdapter([FakeAdapter.text("sub done")]) + ) + + await agent.run_result("delegate") + + harness = agent.last_harness + assert harness is not None + tool_result = harness.messages[2].content[0] + assert "sub done" in tool_result.content + assert tool_result.is_error is False + + +@pytest.mark.asyncio +async def test_async_agent_subagent_accepts_an_async_adapter_factory(tmp_path): + agent = AsyncAgent( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use("tu_1", "subagent", {"task": "work"}), + FakeAsyncAdapter.text("parent done"), + ] + ), + system="sys", + run_dir=str(tmp_path), + ) + agent.enable_subagents( + adapter_factory=lambda: FakeAsyncAdapter([FakeAsyncAdapter.text("async sub")]) + ) + + await agent.run_result("delegate") + + harness = agent.last_harness + assert harness is not None + assert "async sub" in harness.messages[2].content[0].content + + +@pytest.mark.asyncio +async def test_async_agent_replay_cache_skips_the_model(tmp_path): + """`enable_cache` used to be Agent-only, so async callers paid for every repeat.""" + pd = pytest.importorskip("pandas") + + def build(responses): + agent = AsyncAgent.from_dataframe( + pd.DataFrame({"a": [1, 2, 3]}), + adapter=FakeAsyncAdapter(responses), + run_dir=str(tmp_path), + ) + return agent + + shared_path = tmp_path / "replay.json" + + first = build( + [ + FakeAsyncAdapter.tool_use( + "tu_1", "python_interpreter", {"code": "answer(6)"} + ), + FakeAsyncAdapter.text("The answer is 6"), + ] + ) + first.enable_cache(str(shared_path)) + first_result = await first.run_result("what is the total") + assert first_result.text == "The answer is 6" + + second = build([]) + second.enable_cache(str(shared_path)) + second_result = await second.run_result("what is the total") + + assert second_result.text == "The answer is 6" + assert second_result.turns == 0 + assert second_result.value == 6 + + +# ── inspection surface ────────────────────────────────────────────────────── + + +def test_harness_inspection_properties_are_live(tmp_path): + """`tools` and `messages` return the live lists, so sessions can add tools.""" + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("hi")]), + system="the system prompt", + tools=[], + run_dir=str(tmp_path), + max_turns=9, + ) + + assert harness.system == "the system prompt" + assert harness.max_turns == 9 + assert harness.tools == [] + assert harness.reminders == [] + + harness.tools.append(echo_spec()) + assert [t.name for t in harness.tools] == ["echo"] + + harness.run("go") + assert isinstance(harness.messages[0], Message) + assert harness.last_result is not None + + +# ── async adapter resolution ──────────────────────────────────────────────── + + +@pytest.mark.parametrize( + ("model", "expected"), + [ + ("claude-sonnet-4-6", "AsyncAnthropicAdapter"), + ("gpt-4o-mini", "AsyncOpenAIAdapter"), + ("o3-mini", "AsyncOpenAIAdapter"), + ("deepseek-chat", "AsyncDeepSeekAdapter"), + ("openai/gpt-4o-mini", "AsyncOpenRouterAdapter"), + ], +) +def test_resolve_async_adapter_routes_like_the_sync_one(model, expected, monkeypatch): + """The two resolvers share `_route`, so they cannot disagree about a model.""" + monkeypatch.setenv("ANTHROPIC_API_KEY", "x") + monkeypatch.setenv("OPENAI_API_KEY", "x") + monkeypatch.setenv("OPENROUTER_API_KEY", "x") + monkeypatch.setenv("DEEPSEEK_API_KEY", "x") + + assert type(resolve_async_adapter(model)).__name__ == expected + assert type(resolve_adapter(model)).__name__ == expected.replace("Async", "", 1) + + +def test_resolve_async_adapter_reads_the_environment(monkeypatch): + for var in ("ANTHROPIC_API_KEY", "OPENAI_API_KEY", "OPENROUTER_API_KEY"): + monkeypatch.delenv(var, raising=False) + monkeypatch.setenv("DEEPSEEK_API_KEY", "x") + + assert type(resolve_async_adapter()).__name__ == "AsyncDeepSeekAdapter" + + +def test_resolve_async_adapter_without_any_provider(monkeypatch): + for var in ( + "ANTHROPIC_API_KEY", + "OPENAI_API_KEY", + "OPENROUTER_API_KEY", + "DEEPSEEK_API_KEY", + ): + monkeypatch.delenv(var, raising=False) + + with pytest.raises(RuntimeError, match="No provider configured"): + resolve_async_adapter() + + +def test_async_agent_from_dataframe_resolves_an_async_adapter(monkeypatch): + """Without an explicit adapter, `AsyncAgent` must not pick a sync one.""" + pd = pytest.importorskip("pandas") + monkeypatch.setenv("DEEPSEEK_API_KEY", "x") + for var in ("ANTHROPIC_API_KEY", "OPENAI_API_KEY", "OPENROUTER_API_KEY"): + monkeypatch.delenv(var, raising=False) + + agent = AsyncAgent.from_dataframe(pd.DataFrame({"a": [1]})) + + assert isinstance(agent._adapter, AsyncProviderAdapter) + assert isinstance( + Agent.from_dataframe(pd.DataFrame({"a": [1]}))._adapter, ProviderAdapter + ) + + +# ── session parity ────────────────────────────────────────────────────────── + + +def test_sessions_expose_the_same_surface(): + """The sessions are the remaining hand-written pair; guard them too.""" + sync_api = {n for n in dir(AgentSession) if not n.startswith("_")} + async_api = {n for n in dir(AsyncAgentSession) if not n.startswith("_")} + + assert sync_api - async_api == set() + assert async_api - sync_api == {"ask_stream"} + + +def test_sessions_seed_from_a_copy_of_the_agent_cache(tmp_path): + pd = pytest.importorskip("pandas") + frame = pd.DataFrame({"a": [1, 2, 3]}) + + sync_session = Agent.from_dataframe( + frame, adapter=FakeAdapter([]), run_dir=str(tmp_path) + ).session() + async_session = AsyncAgent.from_dataframe( + frame, adapter=FakeAsyncAdapter([]), run_dir=str(tmp_path) + ).async_session() + + assert sync_session.list_handles().keys() == async_session.list_handles().keys() + assert sync_session.turns == async_session.turns == 0 + + +# ── streaming reaches every terminal state ────────────────────────────────── + + +@pytest.mark.asyncio +async def test_stream_reports_max_turns_exceeded(tmp_path): + """The streaming path used to produce no result at all for this outcome.""" + harness = AsyncHarness( + adapter=FakeAsyncAdapter( + [FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "hi"})] + ), + system="sys", + tools=[echo_spec()], + max_turns=1, + run_dir=str(tmp_path), + ) + + async for _ in harness.run_stream("go"): + pass + + result = harness.last_result + assert result is not None + assert result.status == "max_turns_exceeded" + assert result.turns == 1 + assert result.usage.input_tokens == 10 + + +@pytest.mark.asyncio +async def test_abandoning_a_stream_still_accounts_for_it(tmp_path): + """An abandoned run reports what it spent, rather than nothing at all. + + This used to assert `last_result is None`. That was the same defect the + phase set out to fix, one case over: a consumer that reads the final + event and breaks has paid for that turn in full, and reporting nothing + loses those tokens exactly the way the old streaming error path did. + """ + harness = AsyncHarness( + adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + + stream = harness.run_stream("go") + async for _ in stream: + break # consume one event, then walk away + await stream.aclose() + + result = harness.last_result + assert result is not None + assert result.status == "error" + assert "RunAbandoned" in (result.error or "") + + +@pytest.mark.asyncio +async def test_an_abandoned_stream_keeps_the_usage_of_finished_turns(tmp_path): + """Turns that completed before the caller walked away are still billed.""" + harness = AsyncHarness( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAsyncAdapter.text("done"), + ] + ), + system="sys", + tools=[echo_spec()], + run_dir=str(tmp_path), + ) + + stream = harness.run_stream("go") + seen = 0 + async for event in stream: + seen += 1 + if isinstance(event, ToolResultEvent): + break + await stream.aclose() + + result = harness.last_result + assert result is not None + assert result.status == "error" + # Turn 1 completed and was billed before the caller stopped reading. + assert result.usage.input_tokens == 10 + assert result.usage.output_tokens == 5 + + +@pytest.mark.asyncio +async def test_tool_result_events_cover_failures_and_missing_tools(tmp_path): + """Every result block gets an event, not only the ones that succeeded.""" + + def raises(): + raise ValueError("nope") + + boom = ToolSpec( + name="boom", + description="Fail.", + input_schema={"type": "object", "properties": {}}, + handler=raises, + ) + harness = AsyncHarness( + adapter=FakeAsyncAdapter( + [ + NormalizedResponse( + stop_reason=StopReason.TOOL_USE, + content=[ + ToolUseBlock( + tool_use_id="tu_1", tool_name="boom", tool_input={} + ), + ToolUseBlock( + tool_use_id="tu_2", tool_name="ghost", tool_input={} + ), + ], + input_tokens=1, + output_tokens=1, + cache_read_tokens=0, + cache_write_tokens=0, + ), + FakeAsyncAdapter.text("done"), + ] + ), + system="sys", + tools=[boom], + run_dir=str(tmp_path), + ) + + events = [e async for e in harness.run_stream("go")] + tool_events = [e for e in events if isinstance(e, ToolResultEvent)] + + assert [e.tool_name for e in tool_events] == ["boom", "ghost"] + assert all(e.is_error for e in tool_events) + + +# ── the deepest path the drivers support ──────────────────────────────────── + + +def test_subagent_inside_a_sync_harness_inside_a_running_loop(tmp_path): + """Async caller -> sync Agent -> sync subagent. Must not deadlock.""" + + async def outer(): + agent = Agent( + adapter=FakeAdapter( + [ + FakeAdapter.tool_use("tu_1", "subagent", {"task": "work"}), + FakeAdapter.text("parent done"), + ] + ), + system="sys", + run_dir=str(tmp_path), + ) + agent.enable_subagents( + adapter_factory=lambda: FakeAdapter([FakeAdapter.text("sub done")]) + ) + return agent.run_result("delegate") + + result = asyncio.run(outer()) + + assert result.status == "success" + assert result.text == "parent done"