From 873712fa0e31e6f808c6fc55d8977960aa2e8dab Mon Sep 17 00:00:00 2001 From: Max Khor Date: Mon, 3 Aug 2026 20:18:11 +0100 Subject: [PATCH 01/15] Phase 1: collapse the loop and the agent into one implementation each There were four near-copies of the ReAct loop (sync/async x result/stream) and two independent agent classes. They had already drifted: - The streaming loop caught provider errors and returned without building a RunResult, so token usage from every completed turn was lost. A metered deployment silently under-billed itself. - AsyncAgent was missing from_dataframe, from_csv, enable_subagents, add_mcp_server, enable_cache, close, explain, and the on_code/code_only approval gate. _build_tools_for used getattr defaults to paper over it. Now: AsyncHarness._turns is the only loop, parameterised by whether the turn response comes from stream_events or chat. Harness is a facade that bridges a sync adapter onto it and drives it to completion, including from inside a running event loop. _AgentBase holds all configuration and tool wiring; Agent and AsyncAgent differ only in adapter kind and coroutine-ness. Also: - AsyncHarness.last_result exposes the RunResult for streamed runs. - Harness/AsyncHarness gain public inspection properties (system, tools, max_turns, cache, messages, reminders); tests no longer poke privates. - resolve_async_adapter mirrors resolve_adapter through shared routing. - make_subagent_spec accepts either adapter kind. 578 passing (was 547), 31 new regression tests. --- data_harness/__init__.py | 9 +- data_harness/agent.py | 927 +++++++++++++++---------------- data_harness/loop.py | 848 ++++++++++++++-------------- data_harness/quickstart.py | 114 ++-- data_harness/tools/subagent.py | 23 +- tests/smoke_tests.py | 2 +- tests/test_agent.py | 46 +- tests/test_approval_and_cache.py | 2 +- tests/test_async_loop.py | 2 +- tests/test_integration.py | 4 +- tests/test_loop.py | 4 +- tests/test_session_inspection.py | 2 +- tests/test_tool_interpreter.py | 2 +- tests/test_unified_loop.py | 491 ++++++++++++++++ 14 files changed, 1487 insertions(+), 989 deletions(-) create mode 100644 tests/test_unified_loop.py diff --git a/data_harness/__init__.py b/data_harness/__init__.py index 8f69551..72c2fd4 100644 --- a/data_harness/__init__.py +++ b/data_harness/__init__.py @@ -15,7 +15,13 @@ ProviderAdapter, StopReason, ) -from data_harness.quickstart import Chat, SmartFrame, ask, resolve_adapter +from data_harness.quickstart import ( + Chat, + SmartFrame, + ask, + resolve_adapter, + resolve_async_adapter, +) from data_harness.result import CacheStorageInfo, RunResult, Usage from data_harness.streaming import ( ContentBlockDeltaEvent, @@ -84,4 +90,5 @@ "load_dataframe", "mcp_tool_specs", "resolve_adapter", + "resolve_async_adapter", ] diff --git a/data_harness/agent.py b/data_harness/agent.py index d575487..6747b8d 100644 --- a/data_harness/agent.py +++ b/data_harness/agent.py @@ -1,10 +1,16 @@ """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 @@ -16,7 +22,7 @@ from typing import Any from data_harness.cache import SessionCache -from data_harness.loop import AsyncHarness, Harness +from data_harness.loop import AsyncHarness, Harness, run_coroutine_blocking from data_harness.providers.base import AsyncProviderAdapter, ProviderAdapter from data_harness.result import CacheStorageInfo, RunResult, Usage from data_harness.schema import infer_input_schema @@ -50,7 +56,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 +94,16 @@ 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:: - - 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, @@ -192,48 +114,52 @@ def __init__( on_code: Callable[[str], Any] | None = None, code_only: bool = False, ) -> 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] = {} + # ── construction helpers ──────────────────────────────────────────────── + + @classmethod + def _default_adapter(cls, model: str | None) -> Any: + raise NotImplementedError + @classmethod def from_dataframe( cls, 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. + ): + """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.quickstart import _DEFAULT_SYSTEM 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 +169,27 @@ 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, path: str | Path, **kwargs: Any): + """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 exec_cache(self) -> Any: + """The `ExecutionCache`, or ``None`` if caching is disabled.""" + return self._exec_cache + + # ── feature toggles ───────────────────────────────────────────────────── + def connector(self, name: str, *, description: str) -> ConnectorBuilder: """Register a named connector and return a builder for attaching tools. @@ -277,7 +208,7 @@ def connector(self, name: str, *, description: str) -> ConnectorBuilder: ) return ConnectorBuilder(self, name) - def enable_planner(self) -> Agent: + def enable_planner(self): """Enable the planning tool and suffix-based nag reminders. The planner escalates reminders at turns 4, 8, and 12 when no progress @@ -289,9 +220,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, *, adapter_factory: Callable[[], Any]): """Enable the subagent tool, using ``adapter_factory`` for spawned agents. Each spawned subagent gets a fresh adapter, fresh message history, and @@ -299,8 +228,9 @@ 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`; sync ones are bridged automatically. Returns: ``self``, for method chaining. @@ -308,6 +238,46 @@ def enable_subagents( self._subagent_factory = adapter_factory return self + def enable_sql(self, *, engine_url: str | None = None): + """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, path: Any = None): + """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.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, name: str, @@ -316,7 +286,7 @@ def add_mcp_server( args: list[str] | None = None, env: dict[str, str] | None = None, client: Any = None, - ) -> Agent: + ): """Connect an MCP server and expose its tools (via progressive disclosure). The server's tools become a connector named ``name`` — hidden until the @@ -352,71 +322,159 @@ 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.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.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) + ) - def _replay(self, key: str, cached, user_message: str) -> RunResult: - from data_harness.exec_cache import make_key # noqa: F401 (re-export anchor) + if self._connectors or self._mcp_clients: + tools.extend(self._build_connector_tools(target_cache)) - tools = self._build_tools(cache=self._cache) - tool_map = {t.name: t for t in tools} - 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"] + return tools + + def _build_connector_tools(self, cache: SessionCache) -> list[ToolSpec]: + from data_harness.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), ) - for name, meta in self._cache.storage_metadata().items() + 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, } + if self._run_dir is not None: + kwargs["run_dir"] = str(self._run_dir) + return kwargs, effective_planner + + # ── replay cache ──────────────────────────────────────────────────────── + + def _replay_key(self, user_message: str) -> str | None: + if self._exec_cache is None: + return None + from data_harness.exec_cache import make_key + + return make_key(user_message, self._cache, self._system) + + async def _replay(self, cached: Any) -> RunResult: + """Re-execute a cached run's recorded tool steps without calling the model.""" + tool_map = {t.name: t for t in self._build_tools(cache=self._cache)} + for step in cached.steps: + spec = tool_map.get(step["tool"]) + if spec is None or spec.handler is None: + continue + try: + await _call_handler(spec.handler, step["input"]) + except Exception: # noqa: BLE001 + # A recorded step may fail against fresh data; skip it rather + # than aborting the whole replay. + continue return RunResult( text=cached.text, status="success", @@ -425,234 +483,126 @@ def _replay(self, key: str, cached, user_message: str) -> RunResult: 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. - - Returns: - A new `AgentSession` backed by a copy of this agent's cache. - """ - return AgentSession(self) + def _record_replay(self, key: str | None, harness: Any, result: RunResult) -> None: + if key is None or result.status != "success": + return + from data_harness.exec_cache import CachedRun, extract_steps - def run_result(self, user_message: str) -> RunResult: - """Run the agent and return the full `RunResult`. + self._exec_cache.put( + key, CachedRun(steps=extract_steps(harness.messages), text=result.text) + ) - Builds a fresh `Harness` with a fresh message history for each call. - Args: - user_message: The user prompt to send. +async def _call_handler(handler: Callable[..., Any], tool_input: dict) -> Any: + import asyncio + import functools - 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 + if asyncio.iscoroutinefunction(handler): + return await handler(**tool_input) + return await asyncio.to_thread(functools.partial(handler, **tool_input)) - 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) - 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 +class Agent(_AgentBase): + """High-level synchronous agent. - if key is not None and result.status == "success": - from data_harness.exec_cache import CachedRun, extract_steps + `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. - steps = extract_steps(harness._messages) - self._exec_cache.put(key, CachedRun(steps=steps, text=result.text)) - return result + Example:: - def run(self, user_message: str) -> str: - """Run the agent and return the final text response. + from data_harness import Agent + from data_harness.providers.anthropic import AnthropicAdapter - Args: - user_message: The user prompt to send. + 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].")) - 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. - - `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.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 put(self, name: str, value: Any, *, overwrite: bool = False) -> str: - """Store a value in the session cache and return the handle used. + def last_harness(self) -> Harness | None: + return self._last_harness - 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 run_coroutine_blocking(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_result = result - self._turns += result.turns - self._agent._last_harness = self._harness - self._agent._last_run_file = self._harness.run_file + self._last_run_file = 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 +611,86 @@ 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 + from data_harness.loop import _unwrap + + return _unwrap(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.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 - 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._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. + """ + from data_harness.loop import _unwrap + + return _unwrap(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,6 +698,9 @@ 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_harness.last_result`` + holds the `RunResult` for the run, including token usage. + Usage:: async for event in agent.run_stream("hello"): @@ -771,51 +715,28 @@ async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non 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) - def _make_harness( self, *, 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`. + + 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: AsyncAgent) -> None: + def __init__(self, agent: _AgentBase) -> None: self._agent = agent self._cache = SessionCache( sample_size=agent.cache.sample_size, @@ -828,7 +749,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,46 +770,120 @@ 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 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. + """ + from data_harness.loop import _unwrap + + return _unwrap(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. + """ + from data_harness.loop import _unwrap + + return _unwrap(await self.ask_result(user_message)) async def ask_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: - """Stream events for a follow-up turn.""" + """Stream events for a follow-up turn. + + After the generator is exhausted, ``session.harness.last_result`` holds + the `RunResult` for the turn, including token usage. + """ 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 + result = self._harness.last_result + if result is not None: + self._record(result) + 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 _EXPLAIN_TEMPLATE = """\ diff --git a/data_harness/loop.py b/data_harness/loop.py index 6db5ee4..fdb8847 100644 --- a/data_harness/loop.py +++ b/data_harness/loop.py @@ -1,10 +1,29 @@ +"""The ReAct loop. + +There is exactly one loop implementation: `AsyncHarness._turns`. Everything +else in this module is a facade over it. + +- `AsyncHarness.run_result` / `ask_result` drain the loop and return a + `RunResult`. +- `AsyncHarness.run_stream` / `ask_stream` yield the loop's `StreamEvent`s and + leave the `RunResult` on `last_result`. +- `Harness` is the synchronous facade: it bridges a sync `ProviderAdapter` onto + the async loop and drives it to completion. + +Before this collapse there were four near-copies of the loop (sync/async x +result/stream) that had already drifted apart: the streaming path silently +dropped token usage on provider errors, and `AsyncAgent` was missing features +`Agent` had. One implementation means one place for behaviour to live. +""" + from __future__ import annotations import asyncio +import concurrent.futures import dataclasses import functools -from collections.abc import AsyncGenerator, Callable -from typing import Literal, cast +from collections.abc import AsyncGenerator, Callable, Coroutine +from typing import Any, Literal, TypeVar, cast from data_harness.cache import SessionCache from data_harness.exceptions import MaxTurnsExceeded @@ -36,6 +55,8 @@ "Do not use any more tools. Respond with your answer directly." ) +_T = TypeVar("_T") + def _evaluate_code_gate( on_code: Callable[[str], object] | None, @@ -58,26 +79,79 @@ def _evaluate_code_gate( return None -class Harness: - """The core synchronous ReAct loop. +def run_coroutine_blocking(coro: Coroutine[Any, Any, _T]) -> _T: + """Run ``coro`` to completion from synchronous code. + + Uses `asyncio.run` when no event loop is running. When one already is (a + Jupyter kernel, an async web handler calling into sync code), `asyncio.run` + would raise, so the coroutine is handed to a dedicated worker thread with + its own loop instead. + """ + try: + asyncio.get_running_loop() + except RuntimeError: + return asyncio.run(coro) + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: + return pool.submit(asyncio.run, coro).result() + + +class _SyncAdapterBridge(AsyncProviderAdapter): + """Presents a synchronous `ProviderAdapter` as an `AsyncProviderAdapter`. + + Used only by the synchronous `Harness` facade, which owns its own event + loop. The provider call is made inline rather than on a worker thread: + nothing else is scheduled on that loop, so blocking it is harmless and it + keeps provider calls on the caller's thread exactly as before the collapse. + """ + + def __init__(self, adapter: ProviderAdapter) -> None: + self._adapter = adapter - `Harness` owns the message list, dispatches tools, applies suffix-only + async def chat( + self, + system: str, + messages: list[Message], + tools: list[ToolSpec], + ) -> NormalizedResponse: + return self._adapter.chat(system=system, messages=messages, tools=tools) + + def format_cache_control(self, obj: dict) -> dict: + return self._adapter.format_cache_control(obj) + + +def as_async_adapter( + adapter: ProviderAdapter | AsyncProviderAdapter, +) -> AsyncProviderAdapter: + """Return ``adapter`` as an async adapter, bridging a synchronous one. + + Lets callers accept either adapter kind without branching. The bridge is a + pass-through for adapters that are already async. + """ + if isinstance(adapter, AsyncProviderAdapter): + return adapter + return _SyncAdapterBridge(adapter) + + +class AsyncHarness: + """The core ReAct loop. + + `AsyncHarness` 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. + (`AsyncProviderAdapter`, `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``. + For most use cases, prefer `AsyncAgent` over constructing `AsyncHarness` + directly. Use `AsyncHarness` 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. + adapter: Async 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. @@ -85,11 +159,16 @@ class Harness: 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``. + on_code: Optional approval gate called with interpreter code before it + runs. Return ``False`` to block, or a string to return that string + to the model instead of executing. + code_only: When ``True``, interpreter code is echoed back as a dry run + and never executed. """ def __init__( self, - adapter: ProviderAdapter, + adapter: AsyncProviderAdapter, system: str, tools: list[ToolSpec], max_turns: int = 25, @@ -111,6 +190,7 @@ def __init__( self._run_file: str | None = None self._on_code = on_code self._code_only = code_only + self._last_result: RunResult | None = None def register_reminder(self, hook: Callable[[int, int], str | None]) -> None: """Register a suffix reminder hook called before each provider turn. @@ -123,7 +203,61 @@ def register_reminder(self, hook: Callable[[int, int], str | None]) -> None: """ self._reminders.append(hook) - def run_result( + @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 itself carries no run-level + summary. + """ + return self._last_result + + # ── inspection ────────────────────────────────────────────────────────── + # + # The loop's state is readable so callers can assert on what the model was + # actually shown. `tools`, `messages`, and `reminders` return the live + # lists: mutating them mutates the harness, which is how tools are added + # to an already-constructed session. + + @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 cache(self) -> SessionCache: + """The `SessionCache` backing tool results and handles.""" + return self._cache + + @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 + + # ── entry points ──────────────────────────────────────────────────────── + + async def run_result( self, user_message: str, *, @@ -143,12 +277,11 @@ def run_result( 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) + self._begin_run(user_message) + result = await self._drain(stream=False) + return self._stamp(result, run_id, session_id) - def ask_result( + async def ask_result( self, user_message: str, *, @@ -158,7 +291,6 @@ def ask_result( """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. @@ -168,19 +300,13 @@ def ask_result( 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) + self._begin_ask(user_message) + result = await self._drain(stream=False) + return self._stamp(result, run_id, session_id) - def run(self, user_message: str) -> str: + async 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. @@ -191,14 +317,9 @@ def run(self, user_message: str) -> str: 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 + return _unwrap(await self.run_result(user_message)) - def ask(self, user_message: str) -> str: + async def ask(self, user_message: str) -> str: """Append a follow-up message and return the final text response. Args: @@ -211,56 +332,114 @@ def ask(self, user_message: str) -> str: 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 + return _unwrap(await self.ask_result(user_message)) - @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 + 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. + + When the generator is exhausted, `last_result` holds the `RunResult` + for the run, including token usage and any error. + """ + self._begin_run(user_message) + async for event in self._turns(stream=True): + yield event - def _run_loop_result(self) -> RunResult: + async def ask_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: + """Stream events for a follow-up turn in a session. + + When the generator is exhausted, `last_result` holds the `RunResult`. + """ + self._begin_ask(user_message) + async for event in self._turns(stream=True): + yield event + + # ── loop plumbing ─────────────────────────────────────────────────────── + + def _begin_run(self, user_message: str) -> None: + self._run_file = setup_logger(self._run_dir) + self._messages = [Message(role="user", content=[TextBlock(text=user_message)])] + + def _begin_ask(self, user_message: str) -> None: + 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 def _drain(self, *, stream: bool) -> RunResult: + async for _ in self._turns(stream=stream): + pass + result = self._last_result + if result is None: # pragma: no cover - _turns always sets a result + raise RuntimeError("loop finished without producing a result") + return result + + def _stamp( + self, result: RunResult, run_id: str | None, session_id: str | None + ) -> RunResult: + stamped = dataclasses.replace(result, run_id=run_id, session_id=session_id) + self._last_result = stamped + return stamped + + async def _turns(self, *, stream: bool) -> AsyncGenerator[StreamEvent, None]: + """The loop. The only one. + + ``stream`` selects how a turn's response is obtained: real provider + stream events, or a single assembled `chat` call. Every other concern + (reminders, tool dispatch, logging, result construction, error + handling) is shared. + """ if self._run_file is None: raise RuntimeError("run_file must be initialised before running the loop") + self._last_result = None 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( + events_this_turn: list[StreamEvent] = [] + with time_block() as tb: + try: + if stream: + async for evt in self._adapter.stream_events( + system=self._system, + messages=self._messages, + tools=visible_tools, + ): + events_this_turn.append(evt) + yield evt + response = accumulate_stream_events(events_this_turn) + else: + 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, - tools=visible_tools, + error=repr(exc), + run_file=self._run_file, ) - 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), - ) + self._last_result = self._build_result( + text="", + status="error", + turns=turn, + stop_reason=None, + usage=total_usage, + error=repr(exc), + ) + return latency = tb.elapsed_ms @@ -276,11 +455,23 @@ def _run_loop_result(self) -> RunResult: 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_results = await self._dispatch_tools(response.content) - tool_error_count = sum(1 for r in tool_results if r.is_error) + if stream: + tool_name_map = { + b.tool_use_id: b.tool_name + for b in response.content + if isinstance(b, ToolUseBlock) + } + for result_block in tool_results: + yield ToolResultEvent( + tool_use_id=result_block.tool_use_id, + tool_name=tool_name_map.get(result_block.tool_use_id, ""), + content=result_block.content, + is_error=result_block.is_error, + ) + + self._messages.append(Message(role="user", content=list(tool_results))) log_turn( turn=turn, @@ -292,37 +483,53 @@ def _run_loop_result(self) -> RunResult: 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, + tool_error_count=sum(1 for r in tool_results if r.is_error), all_tools=self._tools, ) if response.stop_reason != StopReason.TOOL_USE: - return RunResult( - text=self._extract_text(response), + self._last_result = self._build_result( + text=_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(), ) + return if turn == self._max_turns: - return RunResult( - text=self._extract_text(response), + self._last_result = self._build_result( + text=_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(), ) + return + + 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, + ) -> RunResult: + return RunResult( + text=text, + status=status, + turns=turns, + run_file=self._run_file, + stop_reason=stop_reason, + usage=usage, + cache_snapshots=self._cache.list_handles(), + cache_storage=self._build_cache_storage(), + value=self._cache.get_answer(), + charts=self._cache.list_charts(), + error=error, + ) def _build_cache_storage(self) -> dict[str, CacheStorageInfo]: raw = self._cache.storage_metadata() @@ -349,8 +556,7 @@ def _apply_reminders(self, turn: int) -> None: if not reminder_texts: return - combined = "\n\n".join(reminder_texts) - reminder_block = TextBlock(text=combined) + 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": @@ -358,9 +564,9 @@ def _apply_reminders(self, turn: int) -> None: else: self._messages.append(Message(role="user", content=[reminder_block])) - def _dispatch_tools(self, content: list) -> list[ToolResultBlock]: + async def _dispatch_tools(self, content: list) -> list[ToolResultBlock]: tool_uses = [b for b in content if isinstance(b, ToolUseBlock)] - results = [] + results: list[ToolResultBlock] = [] tool_map = {t.name: t for t in self._tools} for tub in tool_uses: @@ -388,14 +594,18 @@ def _dispatch_tools(self, content: list) -> list[ToolResultBlock]: ) continue try: - raw = spec.handler(**tub.tool_input) + 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, + content=repr(exc), is_error=True, ) ) @@ -410,21 +620,32 @@ def _dispatch_tools(self, content: list) -> list[ToolResultBlock]: 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 Harness: + """Synchronous facade over `AsyncHarness`. -class AsyncHarness: - """Async variant of Harness. Requires an AsyncProviderAdapter. + Holds no loop logic of its own: it wraps the supplied synchronous + `ProviderAdapter` in an async bridge and drives `AsyncHarness` to + completion. Behaviour matches `AsyncHarness` exactly, because it *is* + `AsyncHarness`. + + Safe to call from inside a running event loop (a Jupyter kernel, for + instance); the loop is then driven on a dedicated worker thread. - Exposes the same run_result / ask_result / run / ask surface as Harness, - plus run_stream / ask_stream for token-level streaming. + 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 cache is created if ``None``. + on_code: Optional approval gate for interpreter code. + code_only: When ``True``, interpreter code is never executed. """ def __init__( self, - adapter: AsyncProviderAdapter, + adapter: ProviderAdapter, system: str, tools: list[ToolSpec], max_turns: int = 25, @@ -433,368 +654,113 @@ def __init__( 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._inner = AsyncHarness( + adapter=_SyncAdapterBridge(adapter), + system=system, + tools=tools, + max_turns=max_turns, + run_dir=run_dir, + cache=cache, + on_code=on_code, + code_only=code_only, + ) 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) + """Register a suffix reminder hook called before each provider turn.""" + self._inner.register_reminder(hook) @property def run_file(self) -> str | None: - return self._run_file + """Path to the JSONL log for this run, or ``None`` before the first run.""" + return self._inner.run_file - async def run_result( + @property + def last_result(self) -> RunResult | None: + """`RunResult` for the most recently completed run.""" + return self._inner.last_result + + @property + def system(self) -> str: + """The system prompt. Byte-identical across every turn of a run.""" + return self._inner.system + + @property + def tools(self) -> list[ToolSpec]: + """The full tool list, including invisible tools. Live and mutable.""" + return self._inner.tools + + @property + def max_turns(self) -> int: + """Hard cap on provider turns per run.""" + return self._inner.max_turns + + @property + def cache(self) -> SessionCache: + """The `SessionCache` backing tool results and handles.""" + return self._inner.cache + + @property + def messages(self) -> list[Message]: + """The conversation history the model sees. Live and mutable.""" + return self._inner.messages + + @property + def reminders(self) -> list[Callable[[int, int], str | None]]: + """Registered suffix reminder hooks, in registration order.""" + return self._inner.reminders + + 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) + """Start a fresh run and return the full `RunResult`.""" + return run_coroutine_blocking( + self._inner.run_result(user_message, run_id=run_id, session_id=session_id) + ) - async def ask_result( + 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)]) + """Append a follow-up message and continue the existing run.""" + return run_coroutine_blocking( + self._inner.ask_result(user_message, run_id=run_id, session_id=session_id) ) - 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. + def run(self, user_message: str) -> str: + """Start a fresh run and return the final text response. - 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. + Raises: + MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. + RuntimeError: If the provider raises an exception during the run. """ - 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, - ) + return _unwrap(self.run_result(user_message)) - 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) + def ask(self, user_message: str) -> str: + """Append a follow-up message and return the final text response. - 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])) + Raises: + MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. + RuntimeError: If the provider raises an exception during the run. + """ + return _unwrap(self.ask_result(user_message)) - 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, - ) - ) +def _unwrap(result: RunResult) -> str: + """Return the text of a successful run, or raise the matching exception.""" + if result.status == "max_turns_exceeded": + raise MaxTurnsExceeded(result.turns) + if result.status == "error": + raise RuntimeError(result.error or "unknown error") + return result.text - 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) +def _extract_text(response: NormalizedResponse) -> str: + return "\n".join(b.text for b in response.content if isinstance(b, TextBlock)) diff --git a/data_harness/quickstart.py b/data_harness/quickstart.py index 7bc3571..d14d791 100644 --- a/data_harness/quickstart.py +++ b/data_harness/quickstart.py @@ -12,6 +12,7 @@ from __future__ import annotations import dataclasses +import importlib import importlib.util import os from collections.abc import Callable @@ -22,7 +23,7 @@ from data_harness.result import RunResult if TYPE_CHECKING: - from data_harness.providers.base import ProviderAdapter + from data_harness.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,50 @@ 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.providers.anthropic" + if provider == "anthropic" + else "data_harness.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 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/tools/subagent.py b/data_harness/tools/subagent.py index f0b449b..eb4a9ea 100644 --- a/data_harness/tools/subagent.py +++ b/data_harness/tools/subagent.py @@ -5,7 +5,7 @@ from typing import Callable from data_harness.cache import SessionCache -from data_harness.providers.base import ProviderAdapter +from data_harness.providers.base import AsyncProviderAdapter, ProviderAdapter from data_harness.tools.interpreter import PythonInterpreter from data_harness.tools.variables import make_list_variables_spec from data_harness.types import ToolSpec @@ -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,11 @@ def make_subagent_spec( ) -> ToolSpec: """Create a subagent tool with an explicit cache boundary. + ``adapter_factory`` may return either a `ProviderAdapter` or an + `AsyncProviderAdapter`; a synchronous one is bridged onto the async 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 +50,11 @@ def subagent( input_handles: list[str] | None = None, output_policy: str = "text_only", ) -> str: - from data_harness.loop import Harness + from data_harness.loop import ( + AsyncHarness, + as_async_adapter, + run_coroutine_blocking, + ) # Validate input_handles against parent cache if input_handles: @@ -95,10 +104,8 @@ def subagent( system = _WORKER_SYSTEM_TEMPLATE.format(task=task, input_handles=handles_str) # Spawn fresh adapter - sub_adapter = adapter_factory() - - sub_harness = Harness( - adapter=sub_adapter, + sub_harness = AsyncHarness( + adapter=as_async_adapter(adapter_factory()), system=system, tools=sub_tools, run_dir=run_dir, @@ -106,7 +113,7 @@ def subagent( ) try: - final_text = sub_harness.run(task) + final_text = run_coroutine_blocking(sub_harness.run(task)) except Exception as exc: return f"Error: subagent failed: {type(exc).__name__}: {exc}" diff --git a/tests/smoke_tests.py b/tests/smoke_tests.py index c99bce8..79b4b3e 100644 --- a/tests/smoke_tests.py +++ b/tests/smoke_tests.py @@ -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..d6a2905 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -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")]) @@ -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,20 @@ 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. + + `Harness` is a facade over `AsyncHarness`, so a parent run constructs one + too. The subagent is identified by its worker system prompt rather than by + construction order, which is an implementation detail. + """ + 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")]) @@ -505,7 +519,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.loop import AsyncHarness as RealHarness captured = [] @@ -514,7 +528,7 @@ def recording_harness(*args, **kwargs): captured.append(harness) return harness - monkeypatch.setattr("data_harness.loop.Harness", recording_harness) + monkeypatch.setattr("data_harness.loop.AsyncHarness", recording_harness) adapter = FakeAdapter( [ FakeAdapter.tool_use("tu_1", "subagent", {"task": "work"}), @@ -530,11 +544,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.loop import AsyncHarness as RealHarness captured = [] @@ -543,7 +556,7 @@ def recording_harness(*args, **kwargs): captured.append(harness) return harness - monkeypatch.setattr("data_harness.loop.Harness", recording_harness) + monkeypatch.setattr("data_harness.loop.AsyncHarness", recording_harness) adapter = FakeAdapter( [ FakeAdapter.tool_use( @@ -569,10 +582,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 +594,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..2879875 100644 --- a/tests/test_approval_and_cache.py +++ b/tests/test_approval_and_cache.py @@ -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 diff --git a/tests/test_async_loop.py b/tests/test_async_loop.py index 4555c50..a97a63f 100644 --- a/tests/test_async_loop.py +++ b/tests/test_async_loop.py @@ -199,7 +199,7 @@ def exploding_tool() -> str: # Inspect the tool result in the message history from data_harness.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_integration.py b/tests/test_integration.py index 5d879b8..f2d00c2 100644 --- a/tests/test_integration.py +++ b/tests/test_integration.py @@ -179,10 +179,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_loop.py b/tests/test_loop.py index f93aa90..81fe3b1 100644 --- a/tests/test_loop.py +++ b/tests/test_loop.py @@ -325,8 +325,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_session_inspection.py b/tests/test_session_inspection.py index a348b02..ec4f1a0 100644 --- a/tests/test_session_inspection.py +++ b/tests/test_session_inspection.py @@ -246,6 +246,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_tool_interpreter.py b/tests/test_tool_interpreter.py index edadfc9..21f4fc2 100644 --- a/tests/test_tool_interpreter.py +++ b/tests/test_tool_interpreter.py @@ -75,7 +75,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_unified_loop.py b/tests/test_unified_loop.py new file mode 100644 index 0000000..9a371f7 --- /dev/null +++ b/tests/test_unified_loop.py @@ -0,0 +1,491 @@ +"""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 pytest + +from data_harness.agent import Agent, AsyncAgent +from data_harness.loop import ( + AsyncHarness, + Harness, + as_async_adapter, + run_coroutine_blocking, +) +from data_harness.providers.base import ( + AsyncProviderAdapter, + NormalizedResponse, + ProviderAdapter, +) +from data_harness.streaming import MessageDeltaEvent, ToolResultEvent +from data_harness.testing import FakeAdapter, FakeAsyncAdapter +from data_harness.types import Message, ToolSpec + +# ── 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 sync facade really is the async loop ──────────────────────────────── + + +def test_sync_and_async_harness_produce_identical_results(tmp_path): + """`Harness` holds no loop logic; it drives `AsyncHarness`.""" + + 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): + """The bridge must forward to the wrapped adapter, not around it.""" + 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): + """Notebook kernels and async web handlers already have a loop running. + + `asyncio.run` would raise there, so the facade falls back to a worker + thread with its own loop. + """ + + 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 + + +# ── adapter bridging ──────────────────────────────────────────────────────── + + +def test_as_async_adapter_passes_async_adapters_through(): + adapter = FakeAsyncAdapter([]) + assert as_async_adapter(adapter) is adapter + + +def test_as_async_adapter_wraps_sync_adapters(): + bridged = as_async_adapter(FakeAdapter([])) + assert isinstance(bridged, AsyncProviderAdapter) + assert not isinstance(bridged, ProviderAdapter) + + +def test_bridge_forwards_cache_control(): + bridged = as_async_adapter(FakeAdapter([])) + assert bridged.format_cache_control({"type": "text"}) == { + "type": "text", + "cache_control": {"type": "ephemeral"}, + } + + +# ── 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_bridges_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 From 41e258f53060eb5e7b55161f0ff60f343d311ea6 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Mon, 3 Aug 2026 20:40:56 +0100 Subject: [PATCH 02/15] Phase 1 review fixes: make the sync driver genuinely synchronous An independent review found the sync facade was not behaviour-preserving. Driving the loop through asyncio for sync callers broke three things: - Tool handlers moved onto an asyncio.to_thread worker, so a connector holding a sqlite3 connection built at setup time started returning 'SQLite objects created in a thread can only be used in that same thread' to the model instead of data. - Ctrl-C stopped landing promptly: neither to_thread nor the executor shutdown is cancellable, so a 3s tool call absorbed the interrupt for its full duration (measured 0.41s before, 3.01s after). - asyncio.run cleared the calling thread's event loop, so a program that installed one and then used the *synchronous* API found it gone. Restructured so the loop is a driver-agnostic generator instead. _plan yields CallProvider / CallTool / ToolFinished effects and performs no I/O; failures are sent back in as Failed(exc) and the loop decides what they mean. Harness performs those effects inline on the calling thread; AsyncHarness awaits them and offloads blocking handlers to to_thread. Both inherit _HarnessBase, so loop state and the _call_tool seam exist on each and a subclass override is actually called (a facade silently ignored it). Also from the review: - run_coroutine_blocking restores the ambient event loop. - _SyncAdapterBridge documents that it blocks its loop; as_async_adapter says to prefer a real async adapter where concurrency matters. - The subagent tool gives a sync adapter the sync driver, so it keeps the parent's threading semantics. - Replay dispatches inline for Agent, offloaded for AsyncAgent. - A failed anthropic import no longer becomes "requires the None extra"; the real ImportError propagates. - Restored self-returning annotations on from_dataframe/from_csv/enable_*. Tests: 611 passing (was 547 at baseline). New coverage pins each defect above, plus the gaps the review named: resolve_async_adapter routing, session parity, streaming max_turns_exceeded, abandoned streams, tool events for failures and missing tools, nested driving, and the test_streaming case that should have caught the usage loss and did not. --- data_harness/agent.py | 84 ++-- data_harness/loop.py | 890 ++++++++++++++++++++------------- data_harness/quickstart.py | 4 + data_harness/tools/subagent.py | 31 +- tests/test_agent.py | 13 +- tests/test_streaming.py | 67 +++ tests/test_unified_loop.py | 465 ++++++++++++++++- 7 files changed, 1142 insertions(+), 412 deletions(-) diff --git a/data_harness/agent.py b/data_harness/agent.py index 6747b8d..03f16b9 100644 --- a/data_harness/agent.py +++ b/data_harness/agent.py @@ -15,11 +15,13 @@ from __future__ import annotations +import asyncio +import functools import uuid from collections.abc import AsyncGenerator, Callable from dataclasses import dataclass from pathlib import Path -from typing import Any +from typing import Any, TypeVar from data_harness.cache import SessionCache from data_harness.loop import AsyncHarness, Harness, run_coroutine_blocking @@ -94,6 +96,9 @@ def tool( return fn +_SelfT = TypeVar("_SelfT", bound="_AgentBase") + + class _AgentBase: """Configuration, tool wiring, and feature toggles shared by both agents. @@ -140,7 +145,7 @@ def _default_adapter(cls, model: str | None) -> Any: @classmethod def from_dataframe( - cls, + cls: type[_SelfT], data: Any, *, adapter: Any = None, @@ -148,7 +153,7 @@ def from_dataframe( system: str | None = None, semantics: dict[str, dict] | None = None, **kwargs: Any, - ): + ) -> _SelfT: """Build an agent with ``data`` preloaded as cache handles. Accepts a DataFrame, a ``{name: value}`` mapping, a file path, or a list @@ -169,7 +174,7 @@ def from_dataframe( return agent @classmethod - def from_csv(cls, path: str | Path, **kwargs: Any): + 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) @@ -208,7 +213,7 @@ def connector(self, name: str, *, description: str) -> ConnectorBuilder: ) return ConnectorBuilder(self, name) - def enable_planner(self): + 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 @@ -220,7 +225,7 @@ def enable_planner(self): self._planner_enabled = True return self - def enable_subagents(self, *, adapter_factory: Callable[[], Any]): + 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 @@ -238,7 +243,7 @@ def enable_subagents(self, *, adapter_factory: Callable[[], Any]): self._subagent_factory = adapter_factory return self - def enable_sql(self, *, engine_url: str | None = None): + 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 @@ -255,7 +260,7 @@ def enable_sql(self, *, engine_url: str | None = None): self._sql_engine_url = engine_url return self - def enable_cache(self, path: Any = None): + 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 @@ -279,14 +284,14 @@ def enable_cache(self, path: Any = None): 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, - ): + ) -> _SelfT: """Connect an MCP server and expose its tools (via progressive disclosure). The server's tools become a connector named ``name`` — hidden until the @@ -462,21 +467,19 @@ def _replay_key(self, user_message: str) -> str | None: return make_key(user_message, self._cache, self._system) - async def _replay(self, cached: Any) -> RunResult: - """Re-execute a cached run's recorded tool steps without calling the model.""" + 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 None or spec.handler is None: - continue - try: - await _call_handler(spec.handler, step["input"]) - except Exception: # noqa: BLE001 - # A recorded step may fail against fresh data; skip it rather - # than aborting the whole replay. - continue + if spec is not None and spec.handler is not None: + steps.append((spec.handler, step["input"])) + return steps + + def _replay_result(self, text: str) -> RunResult: return RunResult( - text=cached.text, + text=text, status="success", turns=0, run_file=None, @@ -504,15 +507,6 @@ def _record_replay(self, key: str | None, harness: Any, result: RunResult) -> No ) -async def _call_handler(handler: Callable[..., Any], tool_input: dict) -> Any: - import asyncio - import functools - - if asyncio.iscoroutinefunction(handler): - return await handler(**tool_input) - return await asyncio.to_thread(functools.partial(handler, **tool_input)) - - class Agent(_AgentBase): """High-level synchronous agent. @@ -564,6 +558,20 @@ def _default_adapter(cls, model: str | None) -> ProviderAdapter: def last_harness(self) -> Harness | None: return self._last_harness + 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) + def session(self) -> AgentSession: """Create a stateful `AgentSession` for multi-turn conversations. @@ -587,7 +595,7 @@ def run_result(self, user_message: str) -> RunResult: if key is not None: cached = self._exec_cache.get(key) if cached is not None: - return run_coroutine_blocking(self._replay(cached)) + return self._replay(cached) harness = self._make_harness() self._last_harness = harness @@ -660,6 +668,20 @@ def _default_adapter(cls, model: str | None) -> AsyncProviderAdapter: def last_harness(self) -> AsyncHarness | None: return self._last_harness + 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) diff --git a/data_harness/loop.py b/data_harness/loop.py index fdb8847..d3dac4d 100644 --- a/data_harness/loop.py +++ b/data_harness/loop.py @@ -1,19 +1,27 @@ """The ReAct loop. -There is exactly one loop implementation: `AsyncHarness._turns`. Everything -else in this module is a facade over it. - -- `AsyncHarness.run_result` / `ask_result` drain the loop and return a - `RunResult`. -- `AsyncHarness.run_stream` / `ask_stream` yield the loop's `StreamEvent`s and - leave the `RunResult` on `last_result`. -- `Harness` is the synchronous facade: it bridges a sync `ProviderAdapter` onto - the async loop and drives it to completion. - -Before this collapse there were four near-copies of the loop (sync/async x -result/stream) that had already drifted apart: the streaming path silently -dropped token usage on provider errors, and `AsyncAgent` was missing features -`Agent` had. One implementation means one place for behaviour to live. +There is exactly one loop implementation: `_HarnessBase._plan`. It is a +generator that owns every decision in a turn (reminders, message assembly, +tool gating, logging, when to stop, what the `RunResult` says) and performs no +I/O itself. Instead it *asks* for I/O by yielding an effect: + +- `CallProvider` — run one provider turn, send back a `NormalizedResponse` +- `CallTool` — invoke one tool handler, send back its return value +- `ToolFinished` — a tool result block is ready, emit an event if you want one + +Either failure is sent back in as `Failed(exc)`; the loop decides what that +means. Two drivers perform those 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 @@ -22,7 +30,9 @@ import concurrent.futures import dataclasses import functools -from collections.abc import AsyncGenerator, Callable, Coroutine +import warnings +from collections.abc import AsyncGenerator, Callable, Coroutine, Generator +from dataclasses import dataclass from typing import Any, Literal, TypeVar, cast from data_harness.cache import SessionCache @@ -58,6 +68,53 @@ _T = TypeVar("_T") +# ── effects ───────────────────────────────────────────────────────────────── + + +@dataclass(frozen=True) +class CallProvider: + """Run one provider turn. Send back a `NormalizedResponse` or `Failed`.""" + + system: str + messages: list[Message] + tools: list[ToolSpec] + + +@dataclass(frozen=True) +class CallTool: + """Invoke one tool handler. Send back its 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. Send back ``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 Failed: + """An effect raised. The loop, not the driver, decides what that means.""" + + error: BaseException + + +Effect = CallProvider | CallTool | ToolFinished + + +# ── helpers ───────────────────────────────────────────────────────────────── + + def _evaluate_code_gate( on_code: Callable[[str], object] | None, code_only: bool, @@ -79,18 +136,43 @@ def _evaluate_code_gate( return None +def _ambient_event_loop() -> asyncio.AbstractEventLoop | None: + """The loop installed on this thread, without installing one as a side effect.""" + with warnings.catch_warnings(): + warnings.simplefilter("ignore", DeprecationWarning) + try: + return asyncio.get_event_loop_policy().get_event_loop() + except RuntimeError: + return None + + def run_coroutine_blocking(coro: Coroutine[Any, Any, _T]) -> _T: """Run ``coro`` to completion from synchronous code. - Uses `asyncio.run` when no event loop is running. When one already is (a - Jupyter kernel, an async web handler calling into sync code), `asyncio.run` - would raise, so the coroutine is handed to a dedicated worker thread with - its own loop instead. + 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: - return asyncio.run(coro) + 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() + asyncio.set_event_loop(previous) + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: return pool.submit(asyncio.run, coro).result() @@ -98,10 +180,10 @@ def run_coroutine_blocking(coro: Coroutine[Any, Any, _T]) -> _T: class _SyncAdapterBridge(AsyncProviderAdapter): """Presents a synchronous `ProviderAdapter` as an `AsyncProviderAdapter`. - Used only by the synchronous `Harness` facade, which owns its own event - loop. The provider call is made inline rather than on a worker thread: - nothing else is scheduled on that loop, so blocking it is harmless and it - keeps provider calls on the caller's thread exactly as before the collapse. + The wrapped call is made inline, so it *blocks the event loop it runs on* + for the duration of the provider request. That is fine on a loop created + solely to drive one run, and wrong on a shared application loop. Prefer a + real async adapter anywhere concurrency matters. """ def __init__(self, adapter: ProviderAdapter) -> None: @@ -124,51 +206,27 @@ def as_async_adapter( ) -> AsyncProviderAdapter: """Return ``adapter`` as an async adapter, bridging a synchronous one. - Lets callers accept either adapter kind without branching. The bridge is a - pass-through for adapters that are already async. + The bridge blocks its event loop during provider calls; see + `_SyncAdapterBridge`. Pass an adapter that is already async when the loop + is shared with anything else. """ if isinstance(adapter, AsyncProviderAdapter): return adapter return _SyncAdapterBridge(adapter) -class AsyncHarness: - """The core ReAct loop. +# ── the loop ──────────────────────────────────────────────────────────────── - `AsyncHarness` 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 - (`AsyncProviderAdapter`, `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 `AsyncAgent` over constructing `AsyncHarness` - directly. Use `AsyncHarness` when you need full control over tool wiring, - as shown in ``examples/advanced_wiring.py``. +class _HarnessBase: + """Loop state, pure turn logic, and the effect generator. - Args: - adapter: Async 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``. - on_code: Optional approval gate called with interpreter code before it - runs. Return ``False`` to block, or a string to return that string - to the model instead of executing. - code_only: When ``True``, interpreter code is echoed back as a dry run - and never executed. + 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, - adapter: AsyncProviderAdapter, system: str, tools: list[ToolSpec], max_turns: int = 25, @@ -179,7 +237,6 @@ def __init__( ) -> 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 @@ -203,6 +260,11 @@ def register_reminder(self, hook: Callable[[int, int], str | None]) -> None: """ self._reminders.append(hook) + # ── 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.""" @@ -213,18 +275,11 @@ 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 itself carries no run-level - summary. + 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 - # ── inspection ────────────────────────────────────────────────────────── - # - # The loop's state is readable so callers can assert on what the model was - # actually shown. `tools`, `messages`, and `reminders` return the live - # lists: mutating them mutates the harness, which is how tools are added - # to an already-constructed session. - @property def system(self) -> str: """The system prompt. Byte-identical across every turn of a run.""" @@ -255,111 +310,7 @@ def reminders(self) -> list[Callable[[int, int], str | None]]: """Registered suffix reminder hooks, in registration order.""" return self._reminders - # ── entry points ──────────────────────────────────────────────────────── - - 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`. - - 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) - result = await self._drain(stream=False) - return self._stamp(result, 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. - - Appends ``user_message`` to the current history without resetting it. - - 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) - result = await self._drain(stream=False) - return self._stamp(result, run_id, session_id) - - async 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(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. - - 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(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. - - When the generator is exhausted, `last_result` holds the `RunResult` - for the run, including token usage and any error. - """ - self._begin_run(user_message) - async for event in self._turns(stream=True): - yield event - - async def ask_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: - """Stream events for a follow-up turn in a session. - - When the generator is exhausted, `last_result` holds the `RunResult`. - """ - self._begin_ask(user_message) - async for event in self._turns(stream=True): - yield event - - # ── loop plumbing ─────────────────────────────────────────────────────── + # ── run setup ─────────────────────────────────────────────────────────── def _begin_run(self, user_message: str) -> None: self._run_file = setup_logger(self._run_dir) @@ -372,14 +323,6 @@ def _begin_ask(self, user_message: str) -> None: Message(role="user", content=[TextBlock(text=user_message)]) ) - async def _drain(self, *, stream: bool) -> RunResult: - async for _ in self._turns(stream=stream): - pass - result = self._last_result - if result is None: # pragma: no cover - _turns always sets a result - raise RuntimeError("loop finished without producing a result") - return result - def _stamp( self, result: RunResult, run_id: str | None, session_id: str | None ) -> RunResult: @@ -387,13 +330,12 @@ def _stamp( self._last_result = stamped return stamped - async def _turns(self, *, stream: bool) -> AsyncGenerator[StreamEvent, None]: - """The loop. The only one. + # ── the loop ──────────────────────────────────────────────────────────── - ``stream`` selects how a turn's response is obtained: real provider - stream events, or a single assembled `chat` call. Every other concern - (reminders, tool dispatch, logging, result construction, error - handling) is shared. + 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. """ if self._run_file is None: raise RuntimeError("run_file must be initialised before running the loop") @@ -405,42 +347,32 @@ async def _turns(self, *, stream: bool) -> AsyncGenerator[StreamEvent, None]: 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: - if stream: - async for evt in self._adapter.stream_events( - system=self._system, - messages=self._messages, - tools=visible_tools, - ): - events_this_turn.append(evt) - yield evt - response = accumulate_stream_events(events_this_turn) - else: - 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, - ) - self._last_result = self._build_result( - text="", - status="error", - turns=turn, - stop_reason=None, - usage=total_usage, - error=repr(exc), - ) - return + outcome = yield CallProvider( + system=self._system, + messages=self._messages, + tools=visible_tools, + ) + 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 latency = tb.elapsed_ms total_usage = total_usage + Usage( @@ -455,22 +387,12 @@ async def _turns(self, *, stream: bool) -> AsyncGenerator[StreamEvent, None]: tool_results: list[ToolResultBlock] = [] if response.stop_reason == StopReason.TOOL_USE: - tool_results = await self._dispatch_tools(response.content) - - if stream: - tool_name_map = { - b.tool_use_id: b.tool_name - for b in response.content - if isinstance(b, ToolUseBlock) - } - for result_block in tool_results: - yield ToolResultEvent( - tool_use_id=result_block.tool_use_id, - tool_name=tool_name_map.get(result_block.tool_use_id, ""), - content=result_block.content, - is_error=result_block.is_error, - ) - + # 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) + tool_results = [block for _, block in finished] + for tool_name, block in finished: + yield ToolFinished(tool_name=tool_name, block=block) self._messages.append(Message(role="user", content=list(tool_results))) log_turn( @@ -507,6 +429,95 @@ async def _turns(self, *, stream: bool) -> AsyncGenerator[StreamEvent, None]: ) return + def _dispatch( + self, content: list + ) -> 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]] = [] + + 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: + finished.append( + ( + tub.tool_name, + 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: + finished.append( + ( + tub.tool_name, + ToolResultBlock( + tool_use_id=tub.tool_use_id, + content=gate, + is_error=False, + ), + ) + ) + continue + + raw = yield CallTool( + tool_use_id=tub.tool_use_id, + tool_name=tub.tool_name, + handler=spec.handler, + tool_input=tub.tool_input, + ) + + if isinstance(raw, Failed): + finished.append( + ( + tub.tool_name, + ToolResultBlock( + tool_use_id=tub.tool_use_id, + content=repr(raw.error), + is_error=True, + ), + ) + ) + continue + + try: + output = format_tool_output(raw, cache=self._cache) + except Exception as exc: # noqa: BLE001 - surfaced to the model + finished.append( + ( + tub.tool_name, + ToolResultBlock( + tool_use_id=tub.tool_use_id, + content=repr(exc), + is_error=True, + ), + ) + ) + continue + + finished.append( + ( + tub.tool_name, + ToolResultBlock( + tool_use_id=tub.tool_use_id, + content=output, + is_error=False, + ), + ) + ) + + return finished + + # ── pure helpers ──────────────────────────────────────────────────────── + def _build_result( self, *, @@ -564,83 +575,42 @@ def _apply_reminders(self, turn: int) -> None: 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: list[ToolResultBlock] = [] - 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: - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=repr(exc), - is_error=True, - ) - ) - continue - results.append( - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=output, - is_error=False, - ) - ) - return results +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. -class Harness: - """Synchronous facade over `AsyncHarness`. + `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. - Holds no loop logic of its own: it wraps the supplied synchronous - `ProviderAdapter` in an async bridge and drives `AsyncHarness` to - completion. Behaviour matches `AsyncHarness` exactly, because it *is* - `AsyncHarness`. + 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. - Safe to call from inside a running event loop (a Jupyter kernel, for - instance); the loop is then driven on a dedicated worker thread. + 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. + 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. - max_turns: Hard cap on provider turns. - run_dir: Directory where JSONL logs are written. + 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``. - on_code: Optional approval gate for interpreter code. - code_only: When ``True``, interpreter code is never executed. + on_code: Approval gate called with interpreter code before it runs. + code_only: When ``True``, interpreter code is echoed, never executed. """ def __init__( @@ -654,8 +624,7 @@ def __init__( on_code: Callable[[str], object] | None = None, code_only: bool = False, ) -> None: - self._inner = AsyncHarness( - adapter=_SyncAdapterBridge(adapter), + super().__init__( system=system, tools=tools, max_turns=max_turns, @@ -666,51 +635,157 @@ def __init__( ) self._adapter = adapter - def register_reminder(self, hook: Callable[[int, int], str | None]) -> None: - """Register a suffix reminder hook called before each provider turn.""" - self._inner.register_reminder(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`. - @property - def run_file(self) -> str | None: - """Path to the JSONL log for this run, or ``None`` before the first run.""" - return self._inner.run_file + Resets message history. Use `ask_result` for follow-up turns on the + same history. - @property - def last_result(self) -> RunResult | None: - """`RunResult` for the most recently completed run.""" - return self._inner.last_result + Args: + user_message: The initial user prompt. + run_id: Optional identifier stamped into the `RunResult`. + session_id: Optional session identifier stamped into the `RunResult`. - @property - def system(self) -> str: - """The system prompt. Byte-identical across every turn of a run.""" - return self._inner.system + 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) - @property - def tools(self) -> list[ToolSpec]: - """The full tool list, including invisible tools. Live and mutable.""" - return self._inner.tools + 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. - @property - def max_turns(self) -> int: - """Hard cap on provider turns per run.""" - return self._inner.max_turns + 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`. - @property - def cache(self) -> SessionCache: - """The `SessionCache` backing tool results and handles.""" - return self._inner.cache + Returns: + A `RunResult` describing the outcome of this turn sequence. + """ + self._begin_ask(user_message) + return self._stamp(self._drive(), run_id, session_id) - @property - def messages(self) -> list[Message]: - """The conversation history the model sees. Live and mutable.""" - return self._inner.messages + def run(self, user_message: str) -> str: + """Start a fresh run and return the final text response. - @property - def reminders(self) -> list[Callable[[int, int], str | None]]: - """Registered suffix reminder hooks, in registration order.""" - return self._inner.reminders + Args: + user_message: The initial user prompt. - def run_result( + 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(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(self.ask_result(user_message)) + + # ── driver ────────────────────────────────────────────────────────────── + + def _drive(self) -> RunResult: + plan = self._plan() + sent: Any = 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) -> Any: + if isinstance(effect, CallProvider): + try: + return 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 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", + cache: SessionCache | None = None, + on_code: Callable[[str], object] | None = None, + code_only: bool = False, + ) -> None: + super().__init__( + system=system, + tools=tools, + max_turns=max_turns, + run_dir=run_dir, + cache=cache, + on_code=on_code, + code_only=code_only, + ) + self._adapter = adapter + + async def run_result( self, user_message: str, *, @@ -718,11 +793,10 @@ def run_result( session_id: str | None = None, ) -> RunResult: """Start a fresh run and return the full `RunResult`.""" - return run_coroutine_blocking( - self._inner.run_result(user_message, run_id=run_id, session_id=session_id) - ) + self._begin_run(user_message) + return self._stamp(await self._drain(stream=False), run_id, session_id) - def ask_result( + async def ask_result( self, user_message: str, *, @@ -730,27 +804,126 @@ def ask_result( session_id: str | None = None, ) -> RunResult: """Append a follow-up message and continue the existing run.""" - return run_coroutine_blocking( - self._inner.ask_result(user_message, run_id=run_id, session_id=session_id) - ) + self._begin_ask(user_message) + return self._stamp(await self._drain(stream=False), run_id, session_id) - def run(self, user_message: str) -> str: + 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(self.run_result(user_message)) + return _unwrap(await self.run_result(user_message)) - def ask(self, user_message: str) -> str: + 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(self.ask_result(user_message)) + return _unwrap(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) + async for event in self._drive(stream=True): + 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 for event in self._drive(stream=True): + yield event + + # ── driver ────────────────────────────────────────────────────────────── + + async def _drain(self, *, stream: bool) -> RunResult: + async for _ in self._drive(stream=stream): + 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]: + plan = self._plan() + sent: Any = None + while True: + try: + effect = plan.send(sent) + except StopIteration: + return + sent = None + + if isinstance(effect, CallProvider): + if stream: + events: list[StreamEvent] = [] + try: + async for evt in self._adapter.stream_events( + system=effect.system, + messages=effect.messages, + tools=effect.tools, + ): + events.append(evt) + yield evt + # Inside the try on purpose: a malformed event stream is + # a provider failure, and should land in the RunResult + # rather than escape the loop. + sent = accumulate_stream_events(events) + except Exception as exc: # noqa: BLE001 - reported as a RunResult + sent = Failed(exc) + else: + try: + sent = 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 = 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, + ) + + 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 _unwrap(result: RunResult) -> str: @@ -764,3 +937,16 @@ def _unwrap(result: RunResult) -> str: def _extract_text(response: NormalizedResponse) -> str: return "\n".join(b.text for b in response.content if isinstance(b, TextBlock)) + + +__all__ = [ + "AsyncHarness", + "CallProvider", + "CallTool", + "Effect", + "Failed", + "Harness", + "ToolFinished", + "as_async_adapter", + "run_coroutine_blocking", +] diff --git a/data_harness/quickstart.py b/data_harness/quickstart.py index d14d791..11adaf8 100644 --- a/data_harness/quickstart.py +++ b/data_harness/quickstart.py @@ -111,6 +111,10 @@ def _build_adapter(model: str | None, *, is_async: bool) -> Any: try: 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( f"{provider} support requires the {extra!r} extra: pip install " f"'data-harness[{extra}]'." diff --git a/data_harness/tools/subagent.py b/data_harness/tools/subagent.py index eb4a9ea..f328c6e 100644 --- a/data_harness/tools/subagent.py +++ b/data_harness/tools/subagent.py @@ -50,11 +50,7 @@ def subagent( input_handles: list[str] | None = None, output_policy: str = "text_only", ) -> str: - from data_harness.loop import ( - AsyncHarness, - as_async_adapter, - run_coroutine_blocking, - ) + from data_harness.loop import AsyncHarness, Harness, run_coroutine_blocking # Validate input_handles against parent cache if input_handles: @@ -103,17 +99,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 - sub_harness = AsyncHarness( - adapter=as_async_adapter(adapter_factory()), - system=system, - tools=sub_tools, - run_dir=run_dir, - cache=sub_cache, - ) + # 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() + harness_kwargs = { + "system": system, + "tools": sub_tools, + "run_dir": run_dir, + "cache": sub_cache, + } try: - final_text = run_coroutine_blocking(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/tests/test_agent.py b/tests/test_agent.py index d6a2905..eccd06c 100644 --- a/tests/test_agent.py +++ b/tests/test_agent.py @@ -434,9 +434,8 @@ def test_planner_state_does_not_leak_across_runs(self, tmp_path): def _subagent_harness(captured): """Pick the spawned subagent's harness out of every harness constructed. - `Harness` is a facade over `AsyncHarness`, so a parent run constructs one - too. The subagent is identified by its worker system prompt rather than by - construction order, which is an implementation detail. + 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") @@ -519,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 AsyncHarness as RealHarness + from data_harness.loop import Harness as RealHarness captured = [] @@ -528,7 +527,7 @@ def recording_harness(*args, **kwargs): captured.append(harness) return harness - monkeypatch.setattr("data_harness.loop.AsyncHarness", recording_harness) + monkeypatch.setattr("data_harness.loop.Harness", recording_harness) adapter = FakeAdapter( [ FakeAdapter.tool_use("tu_1", "subagent", {"task": "work"}), @@ -547,7 +546,7 @@ def recording_harness(*args, **kwargs): assert _subagent_harness(captured).reminders == [] def test_subagent_connector_tools_are_fresh_and_hidden(self, monkeypatch, tmp_path): - from data_harness.loop import AsyncHarness as RealHarness + from data_harness.loop import Harness as RealHarness captured = [] @@ -556,7 +555,7 @@ def recording_harness(*args, **kwargs): captured.append(harness) return harness - monkeypatch.setattr("data_harness.loop.AsyncHarness", recording_harness) + monkeypatch.setattr("data_harness.loop.Harness", recording_harness) adapter = FakeAdapter( [ FakeAdapter.tool_use( diff --git a/tests/test_streaming.py b/tests/test_streaming.py index 94ce795..7697f0b 100644 --- a/tests/test_streaming.py +++ b/tests/test_streaming.py @@ -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.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 diff --git a/tests/test_unified_loop.py b/tests/test_unified_loop.py index 9a371f7..8191ac0 100644 --- a/tests/test_unified_loop.py +++ b/tests/test_unified_loop.py @@ -8,10 +8,12 @@ from __future__ import annotations import asyncio +import sqlite3 +import threading import pytest -from data_harness.agent import Agent, AsyncAgent +from data_harness.agent import Agent, AgentSession, AsyncAgent, AsyncAgentSession from data_harness.loop import ( AsyncHarness, Harness, @@ -22,10 +24,12 @@ AsyncProviderAdapter, NormalizedResponse, ProviderAdapter, + StopReason, ) +from data_harness.quickstart import resolve_adapter, resolve_async_adapter from data_harness.streaming import MessageDeltaEvent, ToolResultEvent from data_harness.testing import FakeAdapter, FakeAsyncAdapter -from data_harness.types import Message, ToolSpec +from data_harness.types import Message, ToolSpec, ToolUseBlock # ── helpers ───────────────────────────────────────────────────────────────── @@ -281,7 +285,6 @@ def script(cls): def test_sync_harness_still_calls_the_sync_adapter(tmp_path): - """The bridge must forward to the wrapped adapter, not around it.""" adapter = FakeAdapter([FakeAdapter.text("hi")]) harness = Harness(adapter=adapter, system="sys", tools=[], run_dir=str(tmp_path)) @@ -292,11 +295,7 @@ def test_sync_harness_still_calls_the_sync_adapter(tmp_path): def test_sync_harness_works_inside_a_running_event_loop(tmp_path): - """Notebook kernels and async web handlers already have a loop running. - - `asyncio.run` would raise there, so the facade falls back to a worker - thread with its own loop. - """ + """A notebook kernel or async web handler already has a loop running.""" async def outer(): harness = Harness( @@ -317,6 +316,253 @@ async def coro(): 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", + "_cache", + "_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) + + # ── adapter bridging ──────────────────────────────────────────────────────── @@ -489,3 +735,206 @@ def test_harness_inspection_properties_are_live(tmp_path): 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_leaves_no_result(tmp_path): + """`last_result` is documented as set once the generator is exhausted.""" + harness = AsyncHarness( + adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + + async for _ in harness.run_stream("go"): + break + + assert harness.last_result is None + + +@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" From 47248150320ed75b8a860ad433495649a29750cb Mon Sep 17 00:00:00 2001 From: Max Khor Date: Mon, 3 Aug 2026 22:26:41 +0100 Subject: [PATCH 03/15] Phase 1 review round 2: effect protocol and lifecycle fixes Second independent review confirmed the sans-io restructure is correct (byte-identical RunResults, messages, dispatch order, stream event order and JSONL records vs main across seven exit paths) and found seven smaller problems: - Successes were sent back into the loop raw, so a handler that legitimately returned a Failed was reported to the model as having raised. Answers are now Ok(value) | Failed(exc); no return value can impersonate a failure. - _ambient_event_loop used get_event_loop(), which on 3.10/3.11 CREATES a loop when none is set. The helper that exists to protect the ambient loop was leaking one never-run, never-closed loop plus its selector fd per call. - Abandoning a stream never closed the provider's generator: a bare 'async for' does not close its iterator, and each generator can only clean up its own frame. Collapsed the driver back to a single generator and wrapped the outer 'async for' in aclosing, so one unwind releases the provider's HTTP stream instead of waiting for the GC. - AsyncAgentSession.ask_stream recorded an unstamped RunResult, so anything correlating turns by session_id lost every streamed turn. - AsyncAgent.run_stream ignored the replay cache in both directions: it neither served a hit nor recorded a success, so whether a caller paid for a repeated question depended on the entry point. A hit now yields the cached answer as synthetic text events. - _AgentBase grew last_result, since streamed and replayed runs return nothing to the caller. Replaces a test that asserted its deliberate absence under an older plan; the rewrite says why that constraint was lifted. - Dropped _SyncAdapterBridge/as_async_adapter, dead since the subagent stopped needing them. Corrected the module docstring: _plan does no *network* I/O, but it does write the JSONL log, which Phase 3 moves behind the store. Tests: 624 passing (was 547 at baseline). New tests/test_loop_protocol.py pins the seven invariants the reviewer's mutation testing showed the suite did not notice, and drives _plan by hand to fix the effect/answer sequence. --- data_harness/agent.py | 78 ++++- data_harness/loop.py | 247 ++++++++------- tests/test_loop_protocol.py | 523 +++++++++++++++++++++++++++++++ tests/test_session_inspection.py | 17 +- tests/test_unified_loop.py | 33 +- 5 files changed, 750 insertions(+), 148 deletions(-) create mode 100644 tests/test_loop_protocol.py diff --git a/data_harness/agent.py b/data_harness/agent.py index 03f16b9..241f081 100644 --- a/data_harness/agent.py +++ b/data_harness/agent.py @@ -136,6 +136,7 @@ def __init__( self._sql_engine_url: str | None = None self._exec_cache: Any = None self._mcp_clients: dict[str, Any] = {} + self._last_result: RunResult | None = None # ── construction helpers ──────────────────────────────────────────────── @@ -188,6 +189,15 @@ def cache(self) -> SessionCache: 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.""" @@ -478,6 +488,11 @@ def _replay_steps(self, cached: Any) -> list[tuple[Callable[..., Any], dict]]: 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=text, status="success", @@ -603,6 +618,7 @@ def run_result(self, user_message: str) -> RunResult: 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 @@ -700,6 +716,7 @@ async def run_result(self, user_message: str) -> RunResult: 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 @@ -721,7 +738,10 @@ async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non message_stop, tool_result) following the Claude Agent SDK protocol. After the generator is exhausted, ``agent.last_harness.last_result`` - holds the `RunResult` for the run, including token usage. + holds the `RunResult` for the run, including token usage. A replay-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:: @@ -731,11 +751,23 @@ async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non 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: + result = await self._replay(cached) + for event in _synthetic_text_events(result.text): + yield event + return + harness = self._make_harness() self._last_harness = harness async for event in harness.run_stream(user_message): yield event self._last_run_file = harness.run_file + if harness.last_result is not None: + self._last_result = harness.last_result + self._record_replay(key, harness, harness.last_result) def _make_harness( self, @@ -818,6 +850,7 @@ def _record(self, result: RunResult) -> RunResult: self._turns += result.turns 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 @@ -895,19 +928,56 @@ async def ask(self, user_message: str) -> str: async def ask_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, None]: """Stream events for a follow-up turn. - After the generator is exhausted, ``session.harness.last_result`` holds - the `RunResult` for the turn, including token usage. + 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 for event in self._harness.ask_stream(user_message): yield event result = self._harness.last_result if result is not None: - self._record(result) + 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.providers.base import StopReason + from data_harness.streaming import ( + ContentBlockDeltaEvent, + ContentBlockStartEvent, + ContentBlockStopEvent, + MessageDeltaEvent, + MessageStartEvent, + MessageStopEvent, + TextDelta, + ) + from data_harness.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: diff --git a/data_harness/loop.py b/data_harness/loop.py index d3dac4d..819ac35 100644 --- a/data_harness/loop.py +++ b/data_harness/loop.py @@ -2,15 +2,24 @@ There is exactly one loop implementation: `_HarnessBase._plan`. It is a generator that owns every decision in a turn (reminders, message assembly, -tool gating, logging, when to stop, what the `RunResult` says) and performs no -I/O itself. Instead it *asks* for I/O by yielding an effect: +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, send back a `NormalizedResponse` -- `CallTool` — invoke one tool handler, send back its return value +- `CallProvider` — run one provider turn +- `CallTool` — invoke one tool handler - `ToolFinished` — a tool result block is ready, emit an event if you want one -Either failure is sent back in as `Failed(exc)`; the loop decides what that -means. Two drivers perform those effects: +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 @@ -32,6 +41,7 @@ import functools import warnings from collections.abc import AsyncGenerator, Callable, Coroutine, Generator +from contextlib import aclosing from dataclasses import dataclass from typing import Any, Literal, TypeVar, cast @@ -73,7 +83,7 @@ @dataclass(frozen=True) class CallProvider: - """Run one provider turn. Send back a `NormalizedResponse` or `Failed`.""" + """Run one provider turn. Answer with `Ok(NormalizedResponse)` or `Failed`.""" system: str messages: list[Message] @@ -82,7 +92,7 @@ class CallProvider: @dataclass(frozen=True) class CallTool: - """Invoke one tool handler. Send back its return value or `Failed`.""" + """Invoke one tool handler. Answer with `Ok(return_value)` or `Failed`.""" tool_use_id: str tool_name: str @@ -92,7 +102,7 @@ class CallTool: @dataclass(frozen=True) class ToolFinished: - """A tool result block is ready. Send back ``None``. + """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. @@ -102,6 +112,17 @@ class ToolFinished: 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.""" @@ -110,6 +131,7 @@ class Failed: Effect = CallProvider | CallTool | ToolFinished +Answer = Ok | Failed | None # ── helpers ───────────────────────────────────────────────────────────────── @@ -137,8 +159,19 @@ def _evaluate_code_gate( def _ambient_event_loop() -> asyncio.AbstractEventLoop | None: - """The loop installed on this thread, without installing one as a side effect.""" - with warnings.catch_warnings(): + """The loop already installed on this thread, or ``None``. + + Must not install one as a side effect: on Python 3.10 and 3.11 + ``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. The policy's private slot is the + only way to look without touching. + """ + policy = asyncio.get_event_loop_policy() + watcher = getattr(policy, "_local", None) + if watcher is not None: + return getattr(watcher, "_loop", None) + with warnings.catch_warnings(): # pragma: no cover - non-default policy warnings.simplefilter("ignore", DeprecationWarning) try: return asyncio.get_event_loop_policy().get_event_loop() @@ -177,44 +210,6 @@ def run_coroutine_blocking(coro: Coroutine[Any, Any, _T]) -> _T: return pool.submit(asyncio.run, coro).result() -class _SyncAdapterBridge(AsyncProviderAdapter): - """Presents a synchronous `ProviderAdapter` as an `AsyncProviderAdapter`. - - The wrapped call is made inline, so it *blocks the event loop it runs on* - for the duration of the provider request. That is fine on a loop created - solely to drive one run, and wrong on a shared application loop. Prefer a - real async adapter anywhere concurrency matters. - """ - - def __init__(self, adapter: ProviderAdapter) -> None: - self._adapter = adapter - - async def chat( - self, - system: str, - messages: list[Message], - tools: list[ToolSpec], - ) -> NormalizedResponse: - return self._adapter.chat(system=system, messages=messages, tools=tools) - - def format_cache_control(self, obj: dict) -> dict: - return self._adapter.format_cache_control(obj) - - -def as_async_adapter( - adapter: ProviderAdapter | AsyncProviderAdapter, -) -> AsyncProviderAdapter: - """Return ``adapter`` as an async adapter, bridging a synchronous one. - - The bridge blocks its event loop during provider calls; see - `_SyncAdapterBridge`. Pass an adapter that is already async when the loop - is shared with anything else. - """ - if isinstance(adapter, AsyncProviderAdapter): - return adapter - return _SyncAdapterBridge(adapter) - - # ── the loop ──────────────────────────────────────────────────────────────── @@ -372,7 +367,7 @@ def _plan(self) -> Generator[Effect, Any, None]: ) return - response: NormalizedResponse = outcome + response: NormalizedResponse = outcome.value latency = tb.elapsed_ms total_usage = total_usage + Usage( @@ -468,20 +463,20 @@ def _dispatch( ) continue - raw = yield CallTool( + answer = yield CallTool( tool_use_id=tub.tool_use_id, tool_name=tub.tool_name, handler=spec.handler, tool_input=tub.tool_input, ) - if isinstance(raw, Failed): + if isinstance(answer, Failed): finished.append( ( tub.tool_name, ToolResultBlock( tool_use_id=tub.tool_use_id, - content=repr(raw.error), + content=repr(answer.error), is_error=True, ), ) @@ -489,7 +484,7 @@ def _dispatch( continue try: - output = format_tool_output(raw, cache=self._cache) + output = format_tool_output(answer.value, cache=self._cache) except Exception as exc: # noqa: BLE001 - surfaced to the model finished.append( ( @@ -712,7 +707,7 @@ def ask(self, user_message: str) -> str: def _drive(self) -> RunResult: plan = self._plan() - sent: Any = None + sent: Answer = None while True: try: effect = plan.send(sent) @@ -724,19 +719,21 @@ def _drive(self) -> RunResult: raise RuntimeError("loop finished without producing a result") return result - def _perform(self, effect: Effect) -> Any: + def _perform(self, effect: Effect) -> Answer: if isinstance(effect, CallProvider): try: - return self._adapter.chat( - system=effect.system, - messages=effect.messages, - tools=effect.tools, + 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 self._call_tool(effect) + return Ok(self._call_tool(effect)) except Exception as exc: # noqa: BLE001 - surfaced to the model return Failed(exc) return None @@ -839,8 +836,12 @@ async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non generator early leaves `last_result` unset. """ self._begin_run(user_message) - async for event in self._drive(stream=True): - yield event + # `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. @@ -848,69 +849,94 @@ async def ask_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non Once the generator is exhausted, `last_result` holds the `RunResult`. """ self._begin_ask(user_message) - async for event in self._drive(stream=True): - yield event + 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 for _ in self._drive(stream=stream): - pass + 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]: - plan = self._plan() - sent: Any = None - while True: - try: - effect = plan.send(sent) - except StopIteration: - return - sent = None + """Perform the loop's effects. - if isinstance(effect, CallProvider): - if stream: - events: list[StreamEvent] = [] - try: - async for evt in self._adapter.stream_events( - system=effect.system, - messages=effect.messages, - tools=effect.tools, - ): - events.append(evt) - yield evt - # Inside the try on purpose: a malformed event stream is - # a provider failure, and should land in the RunResult - # rather than escape the loop. - sent = accumulate_stream_events(events) - except Exception as exc: # noqa: BLE001 - reported as a RunResult - sent = Failed(exc) - else: - try: - sent = await self._adapter.chat( + 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, ) - except Exception as exc: # noqa: BLE001 - reported as a RunResult + 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, CallTool): - try: - sent = 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, - ) + 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. @@ -940,13 +966,14 @@ def _extract_text(response: NormalizedResponse) -> str: __all__ = [ + "Answer", "AsyncHarness", "CallProvider", "CallTool", "Effect", "Failed", "Harness", + "Ok", "ToolFinished", - "as_async_adapter", "run_coroutine_blocking", ] diff --git a/tests/test_loop_protocol.py b/tests/test_loop_protocol.py new file mode 100644 index 0000000..a1cb677 --- /dev/null +++ b/tests/test_loop_protocol.py @@ -0,0 +1,523 @@ +"""The effect protocol, and the invariants mutation testing showed were loose. + +An independent review mutated `loop.py` and `agent.py` and found seven changes +the suite did not notice: the approval gate's scope, whether a blocked call is +an error, `_stamp` writing through to `last_result`, the `ToolFinished` +send-`None` contract, latency actually being measured, the async replay's +offload, and an errored run's text. Each of those has a test here. +""" + +from __future__ import annotations + +import asyncio +import json +import time + +import pytest + +from data_harness.agent import Agent, AsyncAgent +from data_harness.loop import ( + AsyncHarness, + CallProvider, + CallTool, + Failed, + Harness, + Ok, + ToolFinished, + run_coroutine_blocking, +) +from data_harness.streaming import ContentBlockDeltaEvent, TextDelta +from data_harness.testing import FakeAdapter, FakeAsyncAdapter +from data_harness.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 protocol. + + `ToolFinished` is an announcement answered with ``None``. Answering it + with a value would shift every later send by one and feed a tool result + into the wrong `yield`, which no end-to-end test would localise. + """ + 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 `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. + """ + from data_harness.loop import _ambient_event_loop + + asyncio.set_event_loop(None) + assert _ambient_event_loop() is None + + async def coro(): + return 1 + + assert run_coroutine_blocking(coro()) == 1 + assert _ambient_event_loop() is None + + +@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] diff --git a/tests/test_session_inspection.py b/tests/test_session_inspection.py index ec4f1a0..b1cd88a 100644 --- a/tests/test_session_inspection.py +++ b/tests/test_session_inspection.py @@ -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 # --------------------------------------------------------------------------- diff --git a/tests/test_unified_loop.py b/tests/test_unified_loop.py index 8191ac0..862b7ff 100644 --- a/tests/test_unified_loop.py +++ b/tests/test_unified_loop.py @@ -14,12 +14,7 @@ import pytest from data_harness.agent import Agent, AgentSession, AsyncAgent, AsyncAgentSession -from data_harness.loop import ( - AsyncHarness, - Harness, - as_async_adapter, - run_coroutine_blocking, -) +from data_harness.loop import AsyncHarness, Harness, run_coroutine_blocking from data_harness.providers.base import ( AsyncProviderAdapter, NormalizedResponse, @@ -251,11 +246,11 @@ def script(): assert [m.role for m in streamed.messages] == [m.role for m in direct.messages] -# ── the sync facade really is the async loop ──────────────────────────────── +# ── the two drivers agree ─────────────────────────────────────────────────── def test_sync_and_async_harness_produce_identical_results(tmp_path): - """`Harness` holds no loop logic; it drives `AsyncHarness`.""" + """Both drivers run the same `_plan` generator, so they must agree.""" def script(cls): return [ @@ -563,28 +558,6 @@ def test_both_drivers_carry_the_loop_state(attribute, tmp_path): assert hasattr(asynchronous, attribute) -# ── adapter bridging ──────────────────────────────────────────────────────── - - -def test_as_async_adapter_passes_async_adapters_through(): - adapter = FakeAsyncAdapter([]) - assert as_async_adapter(adapter) is adapter - - -def test_as_async_adapter_wraps_sync_adapters(): - bridged = as_async_adapter(FakeAdapter([])) - assert isinstance(bridged, AsyncProviderAdapter) - assert not isinstance(bridged, ProviderAdapter) - - -def test_bridge_forwards_cache_control(): - bridged = as_async_adapter(FakeAdapter([])) - assert bridged.format_cache_control({"type": "text"}) == { - "type": "text", - "cache_control": {"type": "ephemeral"}, - } - - # ── features that only Agent used to have, now on AsyncAgent ──────────────── From 64848981ed39a781e5b078923f1ffadbea38f63c Mon Sep 17 00:00:00 2001 From: Max Khor Date: Mon, 3 Aug 2026 22:45:31 +0100 Subject: [PATCH 04/15] Phase 1 review round 3: close the stream at the layer callers use Third review found the previous round's stream-close fix stopped one layer short of the public API, and that one of its tests was vacuous. - AsyncAgent.run_stream and AsyncAgentSession.ask_stream still iterated the harness generator with a bare 'async for'. Closing a generator does not close the one it iterates, so every agent-level caller still leaked the provider's connection on an early break. The harness-level test could not see it. Both now use aclosing, and the new tests exercise the agent and session entry points rather than the harness. - The ambient-event-loop test was vacuous: it called set_event_loop(None) first, which flips the policy's _set_called flag, and the buggy implementation raises rather than creating once that is set. Verified by re-injecting the old implementation: the test passed. It now runs the probe in a subprocess, and re-injecting the bug fails it. - _ambient_event_loop's fallback for a non-default policy called get_event_loop(), reintroducing the leak it exists to prevent. It now reports _UNREADABLE and run_coroutine_blocking leaves the thread alone rather than guessing. - CallProvider/CallTool answered with None gave AttributeError several frames from the driver at fault; _require_answer names the mistake instead. - AsyncAgent.run_stream did not stamp a run id, so run_id was present on session streams and replay hits but absent here. - run_stream's docstring pointed at last_harness.last_result, which is None on a replay hit (no harness is built) and otherwise the previous question's result. Points at last_result now, which is correct on both paths. - Dropped two docstrings still describing the deleted sync-adapter bridge. - Corrected this file's own header: answering ToolFinished with a value is harmless, not a desync, so that was never the invariant worth pinning. Tests: 631 passing (was 547 at baseline). --- data_harness/agent.py | 32 +++-- data_harness/loop.py | 64 ++++++--- data_harness/tools/subagent.py | 7 +- tests/test_loop_protocol.py | 245 ++++++++++++++++++++++++++++++--- 4 files changed, 296 insertions(+), 52 deletions(-) diff --git a/data_harness/agent.py b/data_harness/agent.py index 241f081..da946a5 100644 --- a/data_harness/agent.py +++ b/data_harness/agent.py @@ -19,6 +19,7 @@ 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, TypeVar @@ -245,7 +246,9 @@ def enable_subagents(self: _SelfT, *, adapter_factory: Callable[[], Any]) -> _Se Args: adapter_factory: Zero-argument callable returning a fresh adapter for each subagent. Either a `ProviderAdapter` or an - `AsyncProviderAdapter`; sync ones are bridged automatically. + `AsyncProviderAdapter`; whichever it returns picks the matching + driver, so a sync adapter keeps the parent's threading + semantics. Returns: ``self``, for method chaining. @@ -737,10 +740,13 @@ 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_harness.last_result`` - holds the `RunResult` for the run, including token usage. A replay-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 + 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:: @@ -751,6 +757,7 @@ async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non if isinstance(event.delta, TextDelta): print(event.delta.text, end="", flush=True) """ + run_id = str(uuid.uuid4()) key = self._replay_key(user_message) if key is not None: cached = self._exec_cache.get(key) @@ -762,11 +769,15 @@ async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non 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 if harness.last_result is not None: - self._last_result = harness.last_result + self._last_result = harness._stamp(harness.last_result, run_id, None) self._record_replay(key, harness, harness.last_result) def _make_harness( @@ -933,8 +944,9 @@ async def ask_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non as a non-streamed turn would be. """ run_id = str(uuid.uuid4()) - async for event in self._harness.ask_stream(user_message): - yield event + 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)) diff --git a/data_harness/loop.py b/data_harness/loop.py index 819ac35..4d1d3b4 100644 --- a/data_harness/loop.py +++ b/data_harness/loop.py @@ -39,7 +39,6 @@ import concurrent.futures import dataclasses import functools -import warnings from collections.abc import AsyncGenerator, Callable, Coroutine, Generator from contextlib import aclosing from dataclasses import dataclass @@ -131,9 +130,27 @@ class Failed: 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 ───────────────────────────────────────────────────────────────── @@ -158,25 +175,31 @@ def _evaluate_code_gate( return None -def _ambient_event_loop() -> asyncio.AbstractEventLoop | None: - """The loop already installed on this thread, or ``None``. +#: 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 = object() + + +def _ambient_event_loop() -> Any: + """The loop already installed on this thread, ``None``, or `_UNREADABLE`. Must not install one as a side effect: on Python 3.10 and 3.11 - ``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. The policy's private slot is the - only way to look without touching. + ``asyncio.get_event_loop()`` *creates* one 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. The stdlib policy keeps + the answer in a private slot and offers no public read-only accessor, so + that slot is what we read. + + Under a policy that does not expose it, we decline to guess: creating a + loop to find out whether one exists is the exact bug this guards against, + so the caller is told it cannot know and leaves the thread alone. """ policy = asyncio.get_event_loop_policy() - watcher = getattr(policy, "_local", None) - if watcher is not None: - return getattr(watcher, "_loop", None) - with warnings.catch_warnings(): # pragma: no cover - non-default policy - warnings.simplefilter("ignore", DeprecationWarning) - try: - return asyncio.get_event_loop_policy().get_event_loop() - except RuntimeError: - return None + local = getattr(policy, "_local", None) + if local is None: # pragma: no cover - non-default policy + return _UNREADABLE + return getattr(local, "_loop", None) def run_coroutine_blocking(coro: Coroutine[Any, Any, _T]) -> _T: @@ -204,7 +227,8 @@ def run_coroutine_blocking(coro: Coroutine[Any, Any, _T]) -> _T: loop.run_until_complete(loop.shutdown_asyncgens()) finally: loop.close() - asyncio.set_event_loop(previous) + if previous is not _UNREADABLE: + asyncio.set_event_loop(previous) with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool: return pool.submit(asyncio.run, coro).result() @@ -343,11 +367,12 @@ def _plan(self) -> Generator[Effect, Any, None]: visible_tools = [t for t in self._tools if t.visible] with time_block() as tb: - outcome = yield CallProvider( + request = CallProvider( system=self._system, messages=self._messages, tools=visible_tools, ) + outcome = _require_answer(request, (yield request)) if isinstance(outcome, Failed): log_error_turn( @@ -463,12 +488,13 @@ def _dispatch( ) continue - answer = yield CallTool( + 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): finished.append( diff --git a/data_harness/tools/subagent.py b/data_harness/tools/subagent.py index f328c6e..acbf969 100644 --- a/data_harness/tools/subagent.py +++ b/data_harness/tools/subagent.py @@ -35,9 +35,10 @@ def make_subagent_spec( """Create a subagent tool with an explicit cache boundary. ``adapter_factory`` may return either a `ProviderAdapter` or an - `AsyncProviderAdapter`; a synchronous one is bridged onto the async loop. - The tool handler itself stays synchronous, so it works identically whether - the parent is an `Agent` or an `AsyncAgent`. + `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 diff --git a/tests/test_loop_protocol.py b/tests/test_loop_protocol.py index a1cb677..c82f938 100644 --- a/tests/test_loop_protocol.py +++ b/tests/test_loop_protocol.py @@ -1,17 +1,31 @@ """The effect protocol, and the invariants mutation testing showed were loose. -An independent review mutated `loop.py` and `agent.py` and found seven changes -the suite did not notice: the approval gate's scope, whether a blocked call is -an error, `_stamp` writing through to `last_result`, the `ToolFinished` -send-`None` contract, latency actually being measured, the async replay's -offload, and an errored run's text. Each of those has a test here. +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 @@ -24,7 +38,6 @@ Harness, Ok, ToolFinished, - run_coroutine_blocking, ) from data_harness.streaming import ContentBlockDeltaEvent, TextDelta from data_harness.testing import FakeAdapter, FakeAsyncAdapter @@ -275,11 +288,11 @@ def test_a_handler_returning_failed_is_not_mistaken_for_a_failure(tmp_path): def test_the_effect_and_answer_sequence_lines_up(tmp_path): - """Drive `_plan` by hand and pin the protocol. + """Drive `_plan` by hand and pin the effect order and answer shapes. - `ToolFinished` is an announcement answered with ``None``. Answering it - with a value would shift every later send by one and feed a tool result - into the wrong `yield`, which no end-to-end test would localise. + 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([]), @@ -403,20 +416,49 @@ async def handler(value: str) -> str: 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 `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. + 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. """ - from data_harness.loop import _ambient_event_loop + probe = textwrap.dedent( + """ + import asyncio, sys + policy = asyncio.get_event_loop_policy() + assert policy._local._set_called is False, "process was not pristine" + + from data_harness.loop import _ambient_event_loop, run_coroutine_blocking - asyncio.set_event_loop(None) - assert _ambient_event_loop() is None + before = _ambient_event_loop() + if before is not None: + sys.exit(f"FAIL: looking created a loop: {before!r}") - async def coro(): - return 1 + async def coro(): + return 1 - assert run_coroutine_blocking(coro()) == 1 - assert _ambient_event_loop() is None + 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 @@ -521,3 +563,166 @@ class 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" From 76b9e5c336192a5440f8c905fa48d0127f1405b6 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Mon, 3 Aug 2026 23:06:35 +0100 Subject: [PATCH 05/15] Phase 1 review round 4: close the remaining accounting gap Fourth review confirmed all seven round-3 fixes and found one real defect plus the accounting hole below. - The _UNREADABLE branch skipped restoring the ambient loop, but set_event_loop had already run, so it left the thread pointing at the loop it had just CLOSED: every later get_event_loop user got 'Event loop is closed'. Worse than the leak the branch exists to avoid. It now sets None, the only honest answer when the previous loop could not be read. The branch was also completely untested, which is how this got in; there is now a test with a policy that has no readable _local. - Abandoning a stream after the last event of a completed turn left last_result None, so a caller that read the answer and stopped lost every token the provider had already billed. That is the same defect this phase set out to fix, one case over. _plan now catches GeneratorExit and records what the finished turns cost. The test that asserted the old behaviour is rewritten to say why the premise changed. - _UNREADABLE is an enum member rather than object(), so _ambient_event_loop's return type is a union again instead of Any. - Dead run_id binding on the replay-hit path (replay mints its own). - _stamp documents that it is reached from data_harness.agent and mutates _last_result, since a caller reading it afterwards sees the stamped object. - Renamed a test whose name still described the deleted adapter bridge. Tests: 633 passing (was 547 at baseline). --- data_harness/agent.py | 3 +- data_harness/loop.py | 62 ++++++++++++++++++++++++++++++----- tests/test_loop_protocol.py | 65 +++++++++++++++++++++++++++++++++++++ tests/test_unified_loop.py | 54 ++++++++++++++++++++++++++---- 4 files changed, 169 insertions(+), 15 deletions(-) diff --git a/data_harness/agent.py b/data_harness/agent.py index da946a5..7a02a24 100644 --- a/data_harness/agent.py +++ b/data_harness/agent.py @@ -757,16 +757,17 @@ async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non if isinstance(event.delta, TextDelta): print(event.delta.text, end="", flush=True) """ - run_id = str(uuid.uuid4()) 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 # `aclosing`, not a bare `async for`: closing this generator does not diff --git a/data_harness/loop.py b/data_harness/loop.py index 4d1d3b4..c36e6c5 100644 --- a/data_harness/loop.py +++ b/data_harness/loop.py @@ -38,6 +38,7 @@ import asyncio import concurrent.futures import dataclasses +import enum import functools from collections.abc import AsyncGenerator, Callable, Coroutine, Generator from contextlib import aclosing @@ -175,13 +176,19 @@ def _evaluate_code_gate( return None +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 = object() +_UNREADABLE = _Unreadable.TOKEN -def _ambient_event_loop() -> Any: +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: on Python 3.10 and 3.11 @@ -192,12 +199,11 @@ def _ambient_event_loop() -> Any: that slot is what we read. Under a policy that does not expose it, we decline to guess: creating a - loop to find out whether one exists is the exact bug this guards against, - so the caller is told it cannot know and leaves the thread alone. + loop to find out whether one exists is the exact bug this guards against. """ policy = asyncio.get_event_loop_policy() local = getattr(policy, "_local", None) - if local is None: # pragma: no cover - non-default policy + if local is None: return _UNREADABLE return getattr(local, "_loop", None) @@ -227,8 +233,12 @@ def run_coroutine_blocking(coro: Coroutine[Any, Any, _T]) -> _T: loop.run_until_complete(loop.shutdown_asyncgens()) finally: loop.close() - if previous is not _UNREADABLE: - asyncio.set_event_loop(previous) + # `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() @@ -267,6 +277,8 @@ def __init__( 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 hook called before each provider turn. @@ -345,6 +357,13 @@ def _begin_ask(self, user_message: str) -> None: 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.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 @@ -354,7 +373,8 @@ def _stamp( 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. + 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") @@ -362,7 +382,30 @@ def _plan(self) -> Generator[Effect, Any, None]: self._last_result = None total_usage = Usage() + try: + yield from self._turns(total_usage) + 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._apply_reminders(turn) visible_tools = [t for t in self._tools if t.visible] @@ -401,6 +444,9 @@ def _plan(self) -> Generator[Effect, Any, None]: 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._messages.append(Message(role="assistant", content=response.content)) diff --git a/tests/test_loop_protocol.py b/tests/test_loop_protocol.py index c82f938..3333028 100644 --- a/tests/test_loop_protocol.py +++ b/tests/test_loop_protocol.py @@ -38,6 +38,7 @@ Harness, Ok, ToolFinished, + run_coroutine_blocking, ) from data_harness.streaming import ContentBlockDeltaEvent, TextDelta from data_harness.testing import FakeAdapter, FakeAsyncAdapter @@ -726,3 +727,67 @@ def build(responses): assert fresh.last_result is not None assert fresh.last_result.text == "The answer is 6" + + +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.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_unified_loop.py b/tests/test_unified_loop.py index 862b7ff..2c96c03 100644 --- a/tests/test_unified_loop.py +++ b/tests/test_unified_loop.py @@ -600,7 +600,7 @@ async def test_async_agent_code_only_blocks_execution(tmp_path): @pytest.mark.asyncio -async def test_async_agent_subagent_bridges_a_sync_adapter_factory(tmp_path): +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( @@ -824,8 +824,14 @@ async def test_stream_reports_max_turns_exceeded(tmp_path): @pytest.mark.asyncio -async def test_abandoning_a_stream_leaves_no_result(tmp_path): - """`last_result` is documented as set once the generator is exhausted.""" +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", @@ -833,10 +839,46 @@ async def test_abandoning_a_stream_leaves_no_result(tmp_path): run_dir=str(tmp_path), ) - async for _ in harness.run_stream("go"): - break + 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 "") - assert harness.last_result is None + +@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 From df80e85075e7954161ec40507deba579673efcb5 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Mon, 3 Aug 2026 23:19:36 +0100 Subject: [PATCH 06/15] Phase 2: give the package layers, and a test that keeps them data_harness was one flat 7.8k-LOC namespace where the core loop constructed a SessionCache and called format_tool_output directly. It is now four layers, bottom up: llm (provider adapters and the wire types they speak), core (the loop, RunResult, run logging), data (session cache, interpreter, SQL, connectors), app (Agent, ask, CLI). Each may import the layers below it and no others. The structural change is core no longer knowing what a DataFrame is: - core.loop took cache=SessionCache and called format_tool_output. It now takes a RunEnvironment, which answers two questions the loop cannot: how to render a tool's return value, and what the run's final state was. data.environment.CacheEnvironment is the session-cache implementation; NullEnvironment is the domain-free default. - data.harness.Harness/AsyncHarness are the cache-flavoured pair Agent builds and data_harness.loop resolves to, so cache= still works. - core.serialize hardcoded DataFrame and ndarray snapshotting for the run log. It now has a snapshotter registry that data populates on import. - RunResult._repr_html_ checked isinstance(value, pd.DataFrame). It ducks on _repr_html_ instead, which is both layer-clean and works for polars. - _unwrap became core.result.unwrap_text: the one place deciding which failure becomes which exception, reachable from every layer that needs it. Every pre-layering import path still resolves, via a meta-path finder rather than re-export shims, so data_harness.loop.Harness IS data.harness.Harness rather than a same-named copy that would fail isinstance across the two. Two subtleties are documented where they bite: the loader must return the already-imported module (not the target's spec, which re-executes the file), and the finder must go at the FRONT of sys.meta_path (or PathFinder resolves aliased submodules like data_harness.providers.base itself and loads a copy). Tests: 769 passing (was 633). tests/test_layers.py enforces the boundary statically over the AST, because a runtime import check would pass vacuously here: nearly every heavy dependency is already behind a function-local import. tests/test_core_standalone.py runs a full tool loop with no data domain at all, plugs in a third-party environment, and pins the legacy paths to module identity rather than mere importability. --- data_harness/__init__.py | 59 ++++-- data_harness/_legacy_paths.py | 112 ++++++++++++ data_harness/app/__init__.py | 8 + data_harness/{ => app}/agent.py | 80 ++++---- data_harness/{ => app}/cli.py | 2 +- data_harness/{ => app}/notebook.py | 6 +- data_harness/{ => app}/pandas.py | 4 +- data_harness/{ => app}/quickstart.py | 12 +- data_harness/core/__init__.py | 12 ++ data_harness/{ => core}/artifacts.py | 0 data_harness/core/environment.py | 83 +++++++++ data_harness/{ => core}/exceptions.py | 2 +- data_harness/{ => core}/logger.py | 6 +- data_harness/{ => core}/loop.py | 88 ++++----- data_harness/{ => core}/observe.py | 0 data_harness/{ => core}/result.py | 44 +++-- data_harness/{ => core}/schema.py | 0 data_harness/{ => core}/serialize.py | 52 +++--- data_harness/data/__init__.py | 46 +++++ data_harness/{ => data}/_sandbox_runner.py | 6 +- data_harness/{ => data}/cache.py | 2 +- data_harness/data/environment.py | 44 +++++ data_harness/{ => data}/exec_cache.py | 2 +- data_harness/{ => data}/format.py | 2 +- data_harness/data/harness.py | 129 +++++++++++++ data_harness/{ => data}/io.py | 0 data_harness/{ => data}/mcp.py | 2 +- data_harness/data/tools/__init__.py | 1 + data_harness/{ => data}/tools/connectors.py | 6 +- data_harness/{ => data}/tools/interpreter.py | 6 +- data_harness/{ => data}/tools/planner.py | 2 +- data_harness/{ => data}/tools/sandbox.py | 10 +- data_harness/{ => data}/tools/sql.py | 6 +- data_harness/{ => data}/tools/subagent.py | 16 +- data_harness/{ => data}/tools/variables.py | 4 +- data_harness/eval/case.py | 2 +- data_harness/eval/graders.py | 2 +- data_harness/eval/runner.py | 4 +- data_harness/llm/__init__.py | 9 + data_harness/{ => llm}/providers/__init__.py | 0 data_harness/{ => llm}/providers/anthropic.py | 8 +- data_harness/{ => llm}/providers/base.py | 14 +- data_harness/{ => llm}/providers/openai.py | 4 +- data_harness/{ => llm}/streaming.py | 4 +- data_harness/{ => llm}/testing.py | 4 +- data_harness/{ => llm}/types.py | 0 data_harness/tools/__init__.py | 0 tests/smoke_tests.py | 16 +- tests/test_agent.py | 32 ++-- tests/test_approval_and_cache.py | 6 +- tests/test_async_loop.py | 14 +- tests/test_cache.py | 2 +- tests/test_cache_extras.py | 2 +- tests/test_charts.py | 8 +- tests/test_cli.py | 2 +- tests/test_core_standalone.py | 171 +++++++++++++++++ tests/test_docs.py | 8 +- tests/test_eval.py | 6 +- tests/test_format.py | 2 +- tests/test_gaps.py | 24 ++- tests/test_integration.py | 20 +- tests/test_layers.py | 172 ++++++++++++++++++ tests/test_logger.py | 6 +- tests/test_loop.py | 14 +- tests/test_loop_protocol.py | 14 +- tests/test_loop_reminders.py | 10 +- tests/test_mcp.py | 4 +- tests/test_observe.py | 2 +- tests/test_providers.py | 12 +- tests/test_providers_openai.py | 8 +- tests/test_quickstart.py | 30 +-- tests/test_result.py | 36 ++-- tests/test_review_fixes.py | 20 +- tests/test_sandbox.py | 8 +- tests/test_schema.py | 2 +- tests/test_serialize.py | 4 +- tests/test_session_inspection.py | 12 +- tests/test_sql.py | 6 +- tests/test_streaming.py | 18 +- tests/test_testing.py | 12 +- tests/test_tool_annotations.py | 36 ++-- tests/test_tool_connectors.py | 6 +- tests/test_tool_interpreter.py | 14 +- tests/test_tool_planner.py | 2 +- tests/test_tool_subagent.py | 16 +- tests/test_tool_variables.py | 4 +- tests/test_turn_summary.py | 10 +- tests/test_types.py | 6 +- tests/test_unified_loop.py | 16 +- 89 files changed, 1279 insertions(+), 439 deletions(-) create mode 100644 data_harness/_legacy_paths.py create mode 100644 data_harness/app/__init__.py rename data_harness/{ => app}/agent.py (93%) rename data_harness/{ => app}/cli.py (98%) rename data_harness/{ => app}/notebook.py (88%) rename data_harness/{ => app}/pandas.py (83%) rename data_harness/{ => app}/quickstart.py (97%) create mode 100644 data_harness/core/__init__.py rename data_harness/{ => core}/artifacts.py (100%) create mode 100644 data_harness/core/environment.py rename data_harness/{ => core}/exceptions.py (92%) rename data_harness/{ => core}/logger.py (94%) rename data_harness/{ => core}/loop.py (94%) rename data_harness/{ => core}/observe.py (100%) rename data_harness/{ => core}/result.py (76%) rename data_harness/{ => core}/schema.py (100%) rename data_harness/{ => core}/serialize.py (71%) create mode 100644 data_harness/data/__init__.py rename data_harness/{ => data}/_sandbox_runner.py (94%) rename data_harness/{ => data}/cache.py (99%) create mode 100644 data_harness/data/environment.py rename data_harness/{ => data}/exec_cache.py (98%) rename data_harness/{ => data}/format.py (98%) create mode 100644 data_harness/data/harness.py rename data_harness/{ => data}/io.py (100%) rename data_harness/{ => data}/mcp.py (98%) create mode 100644 data_harness/data/tools/__init__.py rename data_harness/{ => data}/tools/connectors.py (96%) rename data_harness/{ => data}/tools/interpreter.py (98%) rename data_harness/{ => data}/tools/planner.py (98%) rename data_harness/{ => data}/tools/sandbox.py (95%) rename data_harness/{ => data}/tools/sql.py (95%) rename data_harness/{ => data}/tools/subagent.py (94%) rename data_harness/{ => data}/tools/variables.py (89%) create mode 100644 data_harness/llm/__init__.py rename data_harness/{ => llm}/providers/__init__.py (100%) rename data_harness/{ => llm}/providers/anthropic.py (97%) rename data_harness/{ => llm}/providers/base.py (96%) rename data_harness/{ => llm}/providers/openai.py (99%) rename data_harness/{ => llm}/streaming.py (97%) rename data_harness/{ => llm}/testing.py (96%) rename data_harness/{ => llm}/types.py (100%) delete mode 100644 data_harness/tools/__init__.py create mode 100644 tests/test_core_standalone.py create mode 100644 tests/test_layers.py diff --git a/data_harness/__init__.py b/data_harness/__init__.py index 72c2fd4..484252b 100644 --- a/data_harness/__init__.py +++ b/data_harness/__init__.py @@ -1,29 +1,52 @@ -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.exceptions import ( # noqa: E402 MaxTurnsExceeded, 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.result import CacheStorageInfo, RunResult, Usage # noqa: E402 +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, - resolve_async_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, @@ -36,7 +59,7 @@ TextDelta, ToolResultEvent, ) -from data_harness.types import ( +from data_harness.llm.types import ( # noqa: E402 ContentBlock, Message, TextBlock, 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 93% rename from data_harness/agent.py rename to data_harness/app/agent.py index 7a02a24..549b46d 100644 --- a/data_harness/agent.py +++ b/data_harness/app/agent.py @@ -24,18 +24,18 @@ from pathlib import Path from typing import Any, TypeVar -from data_harness.cache import SessionCache -from data_harness.loop import AsyncHarness, Harness, run_coroutine_blocking -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 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) @@ -162,8 +162,8 @@ def from_dataframe( 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 + 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 cls._default_adapter(model), @@ -288,7 +288,7 @@ def enable_cache(self: _SelfT, path: Any = None) -> _SelfT: Returns: ``self``, for method chaining. """ - from data_harness.exec_cache import ExecutionCache + from data_harness.data.exec_cache import ExecutionCache if isinstance(path, ExecutionCache): self._exec_cache = path @@ -323,7 +323,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) @@ -365,7 +365,7 @@ def _build_tools( artifacts_dir = str(Path(self._effective_run_dir) / "charts") if self._execution == "subprocess": - from data_harness.tools.sandbox import SubprocessPythonInterpreter + from data_harness.data.tools.sandbox import SubprocessPythonInterpreter interpreter_spec = SubprocessPythonInterpreter.make_tool_spec( target_cache, @@ -383,7 +383,7 @@ def _build_tools( tools.extend(planner.make_tool_specs()) if self._sql_enabled: - from data_harness.tools.sql import make_sql_query_spec + from data_harness.data.tools.sql import make_sql_query_spec tools.append( make_sql_query_spec(target_cache, engine_url=self._sql_engine_url) @@ -395,7 +395,7 @@ def _build_tools( return tools def _build_connector_tools(self, cache: SessionCache) -> list[ToolSpec]: - from data_harness.mcp import mcp_tool_specs + from data_harness.data.mcp import mcp_tool_specs registry = ConnectorRegistry() for connector_name, connector in self._connectors.items(): @@ -476,7 +476,7 @@ def _harness_kwargs( def _replay_key(self, user_message: str) -> str | None: if self._exec_cache is None: return None - from data_harness.exec_cache import make_key + from data_harness.data.exec_cache import make_key return make_key(user_message, self._cache, self._system) @@ -518,7 +518,7 @@ def _build_replay_result(self, text: str) -> RunResult: def _record_replay(self, key: str | None, harness: Any, result: RunResult) -> None: if key is None or result.status != "success": return - from data_harness.exec_cache import CachedRun, extract_steps + from data_harness.data.exec_cache import CachedRun, extract_steps self._exec_cache.put( key, CachedRun(steps=extract_steps(harness.messages), text=result.text) @@ -536,7 +536,7 @@ class Agent(_AgentBase): Example:: from data_harness import Agent - from data_harness.providers.anthropic import AnthropicAdapter + from data_harness.llm.providers.anthropic import AnthropicAdapter agent = Agent( adapter=AnthropicAdapter(model="claude-sonnet-4-6"), @@ -568,7 +568,7 @@ def __init__( @classmethod def _default_adapter(cls, model: str | None) -> ProviderAdapter: - from data_harness.quickstart import resolve_adapter + from data_harness.app.quickstart import resolve_adapter return resolve_adapter(model) @@ -638,9 +638,7 @@ def run(self, user_message: str) -> str: MaxTurnsExceeded: If the loop reaches ``max_turns``. RuntimeError: If the provider raises an exception. """ - from data_harness.loop import _unwrap - - return _unwrap(self.run_result(user_message)) + return unwrap_text(self.run_result(user_message)) def _make_harness( self, @@ -679,7 +677,7 @@ def __init__( @classmethod def _default_adapter(cls, model: str | None) -> AsyncProviderAdapter: - from data_harness.quickstart import resolve_async_adapter + from data_harness.app.quickstart import resolve_async_adapter return resolve_async_adapter(model) @@ -730,9 +728,7 @@ async def run(self, user_message: str) -> str: MaxTurnsExceeded: If the loop reaches ``max_turns``. RuntimeError: If the provider raises an exception. """ - from data_harness.loop import _unwrap - - return _unwrap(await self.run_result(user_message)) + 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. @@ -753,7 +749,7 @@ async def run_stream(self, user_message: str) -> AsyncGenerator[StreamEvent, Non 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) """ @@ -906,9 +902,7 @@ def ask(self, user_message: str) -> str: MaxTurnsExceeded: If the loop reaches ``max_turns``. RuntimeError: If the provider raises an exception. """ - from data_harness.loop import _unwrap - - return _unwrap(self.ask_result(user_message)) + return unwrap_text(self.ask_result(user_message)) class AsyncAgentSession(_SessionBase): @@ -933,9 +927,7 @@ async def ask(self, user_message: str) -> str: MaxTurnsExceeded: If the loop reaches ``max_turns``. RuntimeError: If the provider raises an exception. """ - from data_harness.loop import _unwrap - - return _unwrap(await self.ask_result(user_message)) + 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. @@ -963,8 +955,8 @@ def _synthetic_text_events(text: str) -> list[StreamEvent]: provider stream to relay, but a streaming caller still renders from deltas, so it must receive the answer the same way. """ - 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, ContentBlockStopEvent, @@ -973,7 +965,7 @@ def _synthetic_text_events(text: str) -> list[StreamEvent]: MessageStopEvent, TextDelta, ) - from data_harness.types import TextBlock + from data_harness.llm.types import TextBlock return [ MessageStartEvent(), @@ -994,10 +986,10 @@ def _synthetic_text_events(text: str) -> list[StreamEvent]: _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 97% rename from data_harness/quickstart.py rename to data_harness/app/quickstart.py index 11adaf8..39b0a2d 100644 --- a/data_harness/quickstart.py +++ b/data_harness/app/quickstart.py @@ -18,12 +18,12 @@ 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 AsyncProviderAdapter, ProviderAdapter + from data_harness.llm.providers.base import AsyncProviderAdapter, ProviderAdapter DEFAULT_ANTHROPIC_MODEL = "claude-sonnet-4-6" DEFAULT_OPENAI_MODEL = "gpt-4o-mini" @@ -104,9 +104,9 @@ 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.providers.anthropic" + "data_harness.llm.providers.anthropic" if provider == "anthropic" - else "data_harness.providers.openai" + else "data_harness.llm.providers.openai" ) try: module = importlib.import_module(module_name) 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/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/exceptions.py b/data_harness/core/exceptions.py similarity index 92% rename from data_harness/exceptions.py rename to data_harness/core/exceptions.py index 6cdaa3f..3306907 100644 --- a/data_harness/exceptions.py +++ b/data_harness/core/exceptions.py @@ -3,7 +3,7 @@ from typing import TYPE_CHECKING if TYPE_CHECKING: - from data_harness.providers.base import NormalizedResponse + from data_harness.llm.providers.base import NormalizedResponse class MaxTurnsExceeded(RuntimeError): 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/loop.py b/data_harness/core/loop.py similarity index 94% rename from data_harness/loop.py rename to data_harness/core/loop.py index c36e6c5..4c8eccc 100644 --- a/data_harness/loop.py +++ b/data_harness/core/loop.py @@ -43,26 +43,24 @@ from collections.abc import AsyncGenerator, Callable, Coroutine, Generator from contextlib import aclosing from dataclasses import dataclass -from typing import Any, Literal, TypeVar, 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 ( +from typing import Any, Literal, TypeVar + +from data_harness.core.environment import NullEnvironment, RunEnvironment +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.llm.providers.base import ( AsyncProviderAdapter, NormalizedResponse, ProviderAdapter, StopReason, ) -from data_harness.result import CacheStorageInfo, RunResult, Usage -from data_harness.streaming import ( +from data_harness.llm.streaming import ( StreamEvent, ToolResultEvent, accumulate_stream_events, ) -from data_harness.types import ( +from data_harness.llm.types import ( Message, TextBlock, ToolResultBlock, @@ -260,7 +258,7 @@ def __init__( tools: list[ToolSpec], max_turns: int = 25, run_dir: str = "./runs", - cache: SessionCache | None = None, + environment: RunEnvironment | None = None, on_code: Callable[[str], object] | None = None, code_only: bool = False, ) -> None: @@ -270,7 +268,9 @@ def __init__( 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._environment = ( + environment if environment is not None else NullEnvironment() + ) self._messages: list[Message] = [] self._reminders: list[Callable[[int, int], str | None]] = [] self._run_file: str | None = None @@ -327,9 +327,9 @@ def max_turns(self) -> int: return self._max_turns @property - def cache(self) -> SessionCache: - """The `SessionCache` backing tool results and handles.""" - return self._cache + def environment(self) -> RunEnvironment: + """The domain services this run uses. See `RunEnvironment`.""" + return self._environment @property def messages(self) -> list[Message]: @@ -359,7 +359,7 @@ def _stamp( ) -> RunResult: """Attach run/session ids and make the result the harness's latest. - Internal, but reached from `data_harness.agent` so that a streamed run + 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. @@ -469,7 +469,7 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: tool_results=tool_results, latency_ms=latency, run_file=self._run_file, - cache_storage=self._cache.storage_metadata(), + 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, @@ -556,7 +556,7 @@ def _dispatch( continue try: - output = format_tool_output(answer.value, cache=self._cache) + output = self._environment.render_tool_output(answer.value) except Exception as exc: # noqa: BLE001 - surfaced to the model finished.append( ( @@ -595,6 +595,7 @@ def _build_result( usage: Usage, error: str | None = None, ) -> RunResult: + state = self._environment.capture() return RunResult( text=text, status=status, @@ -602,23 +603,13 @@ def _build_result( run_file=self._run_file, stop_reason=stop_reason, usage=usage, - cache_snapshots=self._cache.list_handles(), - cache_storage=self._build_cache_storage(), - value=self._cache.get_answer(), - charts=self._cache.list_charts(), + cache_snapshots=state.snapshots, + cache_storage=state.storage, + value=state.value, + charts=state.artifacts, error=error, ) - 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] = [] @@ -656,7 +647,7 @@ class Harness(_HarnessBase): 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. + (`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 @@ -675,7 +666,9 @@ class Harness(_HarnessBase): 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``. + 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. """ @@ -687,7 +680,7 @@ def __init__( tools: list[ToolSpec], max_turns: int = 25, run_dir: str = "./runs", - cache: SessionCache | None = None, + environment: RunEnvironment | None = None, on_code: Callable[[str], object] | None = None, code_only: bool = False, ) -> None: @@ -696,7 +689,7 @@ def __init__( tools=tools, max_turns=max_turns, run_dir=run_dir, - cache=cache, + environment=environment, on_code=on_code, code_only=code_only, ) @@ -758,7 +751,7 @@ def run(self, user_message: str) -> str: MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. RuntimeError: If the provider raises an exception during the run. """ - return _unwrap(self.run_result(user_message)) + 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. @@ -773,7 +766,7 @@ def ask(self, user_message: str) -> str: MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. RuntimeError: If the provider raises an exception during the run. """ - return _unwrap(self.ask_result(user_message)) + return unwrap_text(self.ask_result(user_message)) # ── driver ────────────────────────────────────────────────────────────── @@ -839,7 +832,7 @@ def __init__( tools: list[ToolSpec], max_turns: int = 25, run_dir: str = "./runs", - cache: SessionCache | None = None, + environment: RunEnvironment | None = None, on_code: Callable[[str], object] | None = None, code_only: bool = False, ) -> None: @@ -848,7 +841,7 @@ def __init__( tools=tools, max_turns=max_turns, run_dir=run_dir, - cache=cache, + environment=environment, on_code=on_code, code_only=code_only, ) @@ -883,7 +876,7 @@ async def run(self, user_message: str) -> str: MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. RuntimeError: If the provider raises an exception during the run. """ - return _unwrap(await self.run_result(user_message)) + 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. @@ -892,7 +885,7 @@ async def ask(self, user_message: str) -> str: MaxTurnsExceeded: If the loop reaches ``max_turns`` without stopping. RuntimeError: If the provider raises an exception during the run. """ - return _unwrap(await self.ask_result(user_message)) + 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. @@ -1024,15 +1017,6 @@ async def _call_tool(self, call: CallTool) -> Any: ) -def _unwrap(result: RunResult) -> str: - """Return the text of a successful run, or raise the matching exception.""" - 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 _extract_text(response: NormalizedResponse) -> str: return "\n".join(b.text for b in response.content if isinstance(b, TextBlock)) 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 76% rename from data_harness/result.py rename to data_harness/core/result.py index fbdf61b..e691266 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 @@ -101,7 +101,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 +109,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 +123,30 @@ 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. - return isinstance(value, pd.DataFrame) - except ImportError: - return False + 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. + RuntimeError: If the provider failed. + """ + from data_harness.core.exceptions import MaxTurnsExceeded + + if result.status == "max_turns_exceeded": + raise MaxTurnsExceeded(result.turns) + if result.status == "error": + raise RuntimeError(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/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..527d89d --- /dev/null +++ b/data_harness/data/harness.py @@ -0,0 +1,129 @@ +"""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.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.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``. + 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, + ) -> 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, + ) + + @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, + ) -> 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, + ) + + @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 98% rename from data_harness/tools/interpreter.py rename to data_harness/data/tools/interpreter.py index c9553a1..0b17755 100644 --- a/data_harness/tools/interpreter.py +++ b/data_harness/data/tools/interpreter.py @@ -12,9 +12,9 @@ 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.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. 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 95% rename from data_harness/tools/sandbox.py rename to data_harness/data/tools/sandbox.py index 4d79d9f..02f3150 100644 --- a/data_harness/tools/sandbox.py +++ b/data_harness/data/tools/sandbox.py @@ -21,15 +21,15 @@ 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.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 +101,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), ], 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 94% rename from data_harness/tools/subagent.py rename to data_harness/data/tools/subagent.py index acbf969..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 AsyncProviderAdapter, 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" @@ -51,7 +51,11 @@ def subagent( input_handles: list[str] | None = None, output_policy: str = "text_only", ) -> str: - from data_harness.loop import AsyncHarness, Harness, run_coroutine_blocking + from data_harness.data.harness import ( + AsyncHarness, + Harness, + run_coroutine_blocking, + ) # Validate input_handles against parent cache if input_handles: 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/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 100% rename from data_harness/types.py rename to data_harness/llm/types.py diff --git a/data_harness/tools/__init__.py b/data_harness/tools/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/tests/smoke_tests.py b/tests/smoke_tests.py index 79b4b3e..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 diff --git a/tests/test_agent.py b/tests/test_agent.py index eccd06c..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 @@ -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( @@ -507,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)) @@ -518,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 = [] @@ -527,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"}), @@ -546,7 +546,7 @@ def recording_harness(*args, **kwargs): 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 = [] @@ -555,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( diff --git a/tests/test_approval_and_cache.py b/tests/test_approval_and_cache.py index 2879875..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: @@ -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 a97a63f..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,7 +197,7 @@ 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): if msg.role == "user": 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_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_integration.py b/tests/test_integration.py index f2d00c2..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, 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 81fe3b1..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, diff --git a/tests/test_loop_protocol.py b/tests/test_loop_protocol.py index 3333028..4ab8d4a 100644 --- a/tests/test_loop_protocol.py +++ b/tests/test_loop_protocol.py @@ -29,8 +29,8 @@ import pytest -from data_harness.agent import Agent, AsyncAgent -from data_harness.loop import ( +from data_harness.app.agent import Agent, AsyncAgent +from data_harness.data.harness import ( AsyncHarness, CallProvider, CallTool, @@ -40,9 +40,9 @@ ToolFinished, run_coroutine_blocking, ) -from data_harness.streaming import ContentBlockDeltaEvent, TextDelta -from data_harness.testing import FakeAdapter, FakeAsyncAdapter -from data_harness.types import ToolSpec +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: @@ -433,7 +433,7 @@ def test_looking_up_the_ambient_loop_does_not_create_one(): policy = asyncio.get_event_loop_policy() assert policy._local._set_called is False, "process was not pristine" - from data_harness.loop import _ambient_event_loop, run_coroutine_blocking + from data_harness.core.loop import _ambient_event_loop, run_coroutine_blocking before = _ambient_event_loop() if before is not None: @@ -739,7 +739,7 @@ def test_an_unreadable_ambient_loop_leaves_no_closed_loop_behind(): """ import threading - from data_harness.loop import _UNREADABLE, _ambient_event_loop + 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. 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..67a7712 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: 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 b1cd88a..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( @@ -228,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, 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 7697f0b..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, @@ -1017,7 +1017,7 @@ async def test_a_failure_while_assembling_the_turn_is_reported( def explode(_events): raise ValueError("cannot assemble this turn") - monkeypatch.setattr("data_harness.loop.accumulate_stream_events", explode) + monkeypatch.setattr("data_harness.core.loop.accumulate_stream_events", explode) harness = AsyncHarness( adapter=FakeAsyncAdapter([FakeAsyncAdapter.text("hi")]), system="s", @@ -1090,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 21f4fc2..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 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 index 2c96c03..b8da44a 100644 --- a/tests/test_unified_loop.py +++ b/tests/test_unified_loop.py @@ -13,18 +13,18 @@ import pytest -from data_harness.agent import Agent, AgentSession, AsyncAgent, AsyncAgentSession -from data_harness.loop import AsyncHarness, Harness, run_coroutine_blocking -from data_harness.providers.base import ( +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.quickstart import resolve_adapter, resolve_async_adapter -from data_harness.streaming import MessageDeltaEvent, ToolResultEvent -from data_harness.testing import FakeAdapter, FakeAsyncAdapter -from data_harness.types import Message, ToolSpec, ToolUseBlock +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 ───────────────────────────────────────────────────────────────── @@ -539,7 +539,7 @@ async def _call_tool(self, call): "_tools", "_system", "_max_turns", - "_cache", + "_environment", "_reminders", "_run_file", "_on_code", From 3900d8c4f1c0b98f8b8549931b247e9504f429ba Mon Sep 17 00:00:00 2001 From: Max Khor Date: Tue, 4 Aug 2026 08:06:08 +0100 Subject: [PATCH 07/15] Phase 2 review: read the ambient loop the supported way on 3.14+ Verifying the layering against a built wheel in a clean Python 3.14 venv turned up a forward-compatibility bug in the loop. _ambient_event_loop read asyncio's private policy slot unconditionally, because on 3.10-3.13 get_event_loop() CREATES a loop when none is set and would leave a never-run, never-closed one behind. On 3.14 that reasoning is obsolete: get_event_loop() raises instead, and get_event_loop_policy() is deprecated there and removed in 3.16. Under -W error the probe blew up, so any 3.14 user running warnings-as-errors would have hit it. Now version-gated: the supported API where one exists, the private slot only where it is the only option, and the docstring says which is which and why. Verified this round (the Phase 2 verification agent hit its session limit before reporting, so this was done directly): - Wheel builds and ships all four layer subpackages plus _legacy_paths. - 12 attacks on the legacy-path finder, run against the INSTALLED wheel with no source tree on the path: legacy-first import in a fresh interpreter, both submodule import forms, parent attribute access, identity across all 34 aliases, isinstance across paths, install() idempotence, typos still raising, inspect.getsource, pickling, reload, and an end-to-end run. - Snapshot output byte-identical to pre-registry for both the handle snapshot the model sees and to_jsonable's dataframe/ndarray records. - 8 mutants against the layer guards, all killed: top-level and function-local core->data imports, a lazy pandas import in core, naming SessionCache in core code, the alias loader returning a fresh module, the finder appended rather than inserted, NullEnvironment rendering with repr, and a stripped layer docstring. - Suite green on 3.10 (769) and on 3.14 from the wheel (761, the 8 failures being optional deps absent from that venv). Also fixed the import sorting my rewrite left in examples/. --- data_harness/core/loop.py | 28 ++++++++++++++++++++-------- examples/advanced_wiring.py | 18 +++++++++--------- examples/cache_benchmark.py | 5 ++--- examples/demo.ipynb | 4 +++- examples/inspect_run.py | 2 +- examples/mcp_demo.py | 2 +- examples/quickstart.py | 2 +- 7 files changed, 37 insertions(+), 24 deletions(-) diff --git a/data_harness/core/loop.py b/data_harness/core/loop.py index 4c8eccc..bb534d1 100644 --- a/data_harness/core/loop.py +++ b/data_harness/core/loop.py @@ -40,6 +40,7 @@ import dataclasses import enum import functools +import sys from collections.abc import AsyncGenerator, Callable, Coroutine, Generator from contextlib import aclosing from dataclasses import dataclass @@ -189,16 +190,27 @@ class _Unreadable(enum.Enum): 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: on Python 3.10 and 3.11 - ``asyncio.get_event_loop()`` *creates* one 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. The stdlib policy keeps - the answer in a private slot and offers no public read-only accessor, so - that slot is what we read. + Must not install one as a side effect. Reading it needs two strategies, + because the safe way to ask changed: - Under a policy that does not expose it, we decline to guess: creating a - loop to find out whether one exists is the exact bug this guards against. + - 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: 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.") From 3fc052f0ce3f0cd888ced3dcda328fcfeddb4e28 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Tue, 4 Aug 2026 08:13:36 +0100 Subject: [PATCH 08/15] Phase 3: the session tree The harness kept its state in a list of messages and wrote a separate runs/*.jsonl log that re-serialised the entire history every turn. That log was quadratic to write and could not be read back into anything runnable: no resume, no forking, no way to ask what the agent actually saw at turn 7. A session is now an append-only tree of typed entries, each naming its parent, with a movable leaf. The conversation 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 one inversion: - Resume: reopen a JsonlSessionStore in a different process, pass it to a harness, and ask() continues the conversation. Tested end to end. - Forking: move the leaf and append; both branches survive and either can be replayed. This is how a turn gets retried with a different model without destroying the first attempt. - Compaction is an entry that changes where the walk starts, not an edit. The compacted entries stay in the tree, so moving the leaf back before it restores the full history. Reversible and auditable. - Writing is one line per entry. 40 messages is 40 lines, where the old log wrote 800 message-copies. Layering held: the tree lives in core and knows nothing about the data domain. CustomEntry plus a projector is the extension point, so the data layer can record a cache write or a chart without core learning what either is. Projected entries are invisible to the model unless a projector opts them in, because a cache write is recorded for a human to trace, not for the model to re-read. Note the session store gets its own message encoder rather than reusing to_jsonable. That one is the *log* shape: it renames fields for readability and snapshots large values away. Lossy is right for a debugging log and wrong for a store meant to reconstruct a runnable conversation, and a round-trip test pins the difference on tool blocks specifically. Both stores (memory, jsonl) run against the same test suite via a fixture, so neither drifts from the protocol. Corrupt files, unknown entry types, future format versions, duplicate ids, unknown parents, and cycles all raise a typed SessionStoreError with a stable code rather than a KeyError somewhere downstream. Not yet wired: the data layer does not record cache_put/chart entries. The mechanism and its tests are in place; the plumbing is a separate change. Tests: 822 passing (was 769). --- data_harness/core/loop.py | 60 ++- data_harness/core/session/__init__.py | 55 +++ data_harness/core/session/entries.py | 141 +++++++ data_harness/core/session/jsonl.py | 263 +++++++++++++ data_harness/core/session/session.py | 231 ++++++++++++ data_harness/core/session/store.py | 139 +++++++ data_harness/data/harness.py | 7 + tests/test_session_tree.py | 522 ++++++++++++++++++++++++++ 8 files changed, 1410 insertions(+), 8 deletions(-) create mode 100644 data_harness/core/session/__init__.py create mode 100644 data_harness/core/session/entries.py create mode 100644 data_harness/core/session/jsonl.py create mode 100644 data_harness/core/session/session.py create mode 100644 data_harness/core/session/store.py create mode 100644 tests/test_session_tree.py diff --git a/data_harness/core/loop.py b/data_harness/core/loop.py index bb534d1..c5012c6 100644 --- a/data_harness/core/loop.py +++ b/data_harness/core/loop.py @@ -50,6 +50,7 @@ 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, @@ -273,6 +274,7 @@ def __init__( environment: RunEnvironment | None = None, on_code: Callable[[str], object] | None = None, code_only: bool = False, + session: Session | None = None, ) -> None: if max_turns < 1: raise ValueError(f"max_turns must be at least 1, got {max_turns!r}") @@ -283,7 +285,11 @@ def __init__( self._environment = ( environment if environment is not None else NullEnvironment() ) - self._messages: list[Message] = [] + self._session = session if session is not None else Session() + # 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 self._on_code = on_code @@ -343,6 +349,16 @@ 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 it is derived from. They agree at every turn + boundary, which `tests/test_session_tree.py` pins. + """ + return self._session + @property def messages(self) -> list[Message]: """The conversation history the model sees. Live and mutable.""" @@ -355,16 +371,28 @@ def reminders(self) -> list[Callable[[int, int], str | None]]: # ── 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._messages = [Message(role="user", content=[TextBlock(text=user_message)])] + self._messages = [] + self._record(Message(role="user", content=[TextBlock(text=user_message)])) def _begin_ask(self, user_message: str) -> None: 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)]) - ) + self._record(Message(role="user", content=[TextBlock(text=user_message)])) def _stamp( self, result: RunResult, run_id: str | None, session_id: str | None @@ -460,7 +488,7 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: # what it spent; the generator's local is gone by then. self._usage_so_far = total_usage - self._messages.append(Message(role="assistant", content=response.content)) + self._record(Message(role="assistant", content=response.content)) tool_results: list[ToolResultBlock] = [] @@ -471,7 +499,19 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: tool_results = [block for _, block in finished] for tool_name, block in finished: yield ToolFinished(tool_name=tool_name, block=block) - self._messages.append(Message(role="user", content=list(tool_results))) + 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, @@ -643,7 +683,7 @@ def _apply_reminders(self, turn: int) -> None: 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])) + self._record(Message(role="user", content=[reminder_block])) class Harness(_HarnessBase): @@ -695,6 +735,7 @@ def __init__( environment: RunEnvironment | None = None, on_code: Callable[[str], object] | None = None, code_only: bool = False, + session: Session | None = None, ) -> None: super().__init__( system=system, @@ -704,6 +745,7 @@ def __init__( environment=environment, on_code=on_code, code_only=code_only, + session=session, ) self._adapter = adapter @@ -847,6 +889,7 @@ def __init__( environment: RunEnvironment | None = None, on_code: Callable[[str], object] | None = None, code_only: bool = False, + session: Session | None = None, ) -> None: super().__init__( system=system, @@ -856,6 +899,7 @@ def __init__( environment=environment, on_code=on_code, code_only=code_only, + session=session, ) self._adapter = adapter 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..754c483 --- /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. `data_harness.data.session` defines those. + """ + + 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..7a8a63a --- /dev/null +++ b/data_harness/core/session/jsonl.py @@ -0,0 +1,263 @@ +"""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 turn log it replaces re-serialised the whole message history +every turn, which is quadratic and cannot be read back into a runnable state. + +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 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 + + # ── 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"]) + for number, line in enumerate(lines[1:], start=2): + try: + raw = json.loads(line) + except json.JSONDecodeError as exc: + 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 + + @property + def path(self) -> Path: + return self._path + + # ── 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: + 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}", + ) + with self._path.open("a") as handle: + handle.write(json.dumps(encode_entry(entry)) + "\n") + 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..b290e36 --- /dev/null +++ b/data_harness/core/session/session.py @@ -0,0 +1,231 @@ +"""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 + +#: 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: + 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 keeping: + 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.""" + 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 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 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..10b0a93 --- /dev/null +++ b/data_harness/core/session/store.py @@ -0,0 +1,139 @@ +"""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 + +from typing import Protocol, runtime_checkable + +from data_harness.core.session.entries import Entry, leaf_after + + +class SessionStoreError(Exception): + """Raised for a session that cannot be read or is internally inconsistent. + + Carries a stable ``code`` so a caller can tell a missing entry from a + corrupt file without matching on message text. + """ + + def __init__(self, code: str, message: str) -> None: + super().__init__(message) + self.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: + 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/harness.py b/data_harness/data/harness.py index 527d89d..d09930f 100644 --- a/data_harness/data/harness.py +++ b/data_harness/data/harness.py @@ -29,6 +29,7 @@ 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 @@ -49,6 +50,8 @@ class Harness(_CoreHarness): 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. on_code: Approval gate called with interpreter code before it runs. code_only: When ``True``, interpreter code is echoed, never executed. """ @@ -63,6 +66,7 @@ def __init__( cache: SessionCache | None = None, on_code: Callable[[str], object] | None = None, code_only: bool = False, + session: Session | None = None, ) -> None: super().__init__( adapter=adapter, @@ -73,6 +77,7 @@ def __init__( environment=_environment_for(cache), on_code=on_code, code_only=code_only, + session=session, ) @property @@ -97,6 +102,7 @@ def __init__( cache: SessionCache | None = None, on_code: Callable[[str], object] | None = None, code_only: bool = False, + session: Session | None = None, ) -> None: super().__init__( adapter=adapter, @@ -107,6 +113,7 @@ def __init__( environment=_environment_for(cache), on_code=on_code, code_only=code_only, + session=session, ) @property diff --git a/tests/test_session_tree.py b/tests/test_session_tree.py new file mode 100644 index 0000000..008fe56 --- /dev/null +++ b/tests/test_session_tree.py @@ -0,0 +1,522 @@ +"""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_satisfy_the_protocol(store): + assert isinstance(store, SessionStore) + + +# ── 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_line_names_the_line(tmp_path): + path = tmp_path / "corrupt.jsonl" + path.write_text( + json.dumps({"type": "session", "version": 1, "id": "x"}) + "\nnot json\n" + ) + with pytest.raises(SessionStoreError) as excinfo: + JsonlSessionStore.open(path) + assert excinfo.value.code == "invalid_entry" + assert ":2" in str(excinfo.value) + + +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 From 802b9f9e94b42c0ddc736085441b23b09e095a36 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Tue, 4 Aug 2026 08:40:24 +0100 Subject: [PATCH 09/15] Phase 3 review: fix ten defects in the session tree An independent review found the tree structure sound but the derivation layer and loop wiring broken in three serious ways and seven smaller ones. Serious: - A crash or Ctrl-C during a tool call persists an assistant tool_use with no matching tool_result. Resuming that session sent the provider a transcript it rejects with a 400, so resume failed precisely for the sessions most worth resuming. build_context now drops unpaired tool calls and unpaired tool results, so a derived context is always something a provider will accept. - _begin_run reset the working copy but not the session leaf, so after a second run() messages and build_context disagreed: the model was sent two messages and the log claimed four, linked as one conversation. A fresh run now starts a fresh root. The docstring claiming they 'agree at every turn boundary' was false and unpinned; it is now true and pinned for runs, repeated runs, asks, and streams. - context_entries silently discarded all pre-compaction history when first_kept_entry_id was off-path or pointed forward, and stacked compactions emitted two summaries in reverse chronological order while resurrecting entries the older one had dropped. append_compaction now rejects a cut that is not on the path, and an older compaction inside the kept tail is skipped rather than replayed. Smaller: - Reminders mutate an already-recorded message, so the JSONL store (which serialises on write) showed a prompt the model never saw. Both stores now snapshot on write, and the reminder is recorded as its own entry rather than retroactively editing an immutable one. - A torn final line, the normal result of a crash mid-append, made the whole file unreadable. The tail is now dropped and flagged via the truncated flag; corruption anywhere else is still fatal. - Added open_or_create, since create() truncates and the obvious restart mistake destroyed the session the feature exists to preserve. - A non-serialisable custom payload raised a bare TypeError from one store and was accepted by the other; it is a typed SessionStoreError now. - Corrected two false claims: the quadratic runs/*.jsonl log is still written alongside, and entries.py referenced a module that does not exist. Tests: 840 passing (was 822). Three mutants had survived the previous round and now die: an encoder writing every role as 'user', is_error always False (the old assertion checked the dataclass default), and removed cycle detection. test_both_stores_satisfy_the_protocol was near-vacuous, since isinstance against a runtime_checkable Protocol matches method names only; it now exercises the behaviour. --- data_harness/core/loop.py | 22 +- data_harness/core/session/entries.py | 2 +- data_harness/core/session/jsonl.py | 42 +++- data_harness/core/session/session.py | 87 +++++++- data_harness/core/session/store.py | 9 + tests/test_session_tree.py | 323 ++++++++++++++++++++++++++- 6 files changed, 471 insertions(+), 14 deletions(-) diff --git a/data_harness/core/loop.py b/data_harness/core/loop.py index c5012c6..3499777 100644 --- a/data_harness/core/loop.py +++ b/data_harness/core/loop.py @@ -354,8 +354,13 @@ 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 it is derived from. They agree at every turn - boundary, which `tests/test_session_tree.py` pins. + 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 @@ -387,6 +392,12 @@ def _record(self, message: Message) -> Message: def _begin_run(self, user_message: str) -> None: self._run_file = setup_logger(self._run_dir) 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: @@ -682,6 +693,13 @@ def _apply_reminders(self, turn: int) -> None: # 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])) diff --git a/data_harness/core/session/entries.py b/data_harness/core/session/entries.py index 754c483..39eaaed 100644 --- a/data_harness/core/session/entries.py +++ b/data_harness/core/session/entries.py @@ -109,7 +109,7 @@ class CustomEntry(BaseEntry): 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. `data_harness.data.session` defines those. + being produced, a reminder injected into a prompt. """ type: Literal["custom"] = "custom" diff --git a/data_harness/core/session/jsonl.py b/data_harness/core/session/jsonl.py index 7a8a63a..a8e68d7 100644 --- a/data_harness/core/session/jsonl.py +++ b/data_harness/core/session/jsonl.py @@ -1,8 +1,10 @@ """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 turn log it replaces re-serialised the whole message history -every turn, which is quadratic and cannot be read back into a runnable state. +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 @@ -128,6 +130,8 @@ def __init__(self, path: str | Path, session_id: str) -> None: 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 ─────────────────────────────────────────────────────────── @@ -175,10 +179,18 @@ def open(cls, path: str | Path) -> JsonlSessionStore: ) 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 @@ -188,10 +200,27 @@ def open(cls, path: str | Path) -> JsonlSessionStore: 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 @@ -212,8 +241,15 @@ def append(self, entry: Entry) -> None: "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(json.dumps(encode_entry(entry)) + "\n") + handle.write(line + "\n") self._entries.append(entry) self._by_id[entry.id] = entry self._leaf_id = leaf_after(entry) diff --git a/data_harness/core/session/session.py b/data_harness/core/session/session.py index b290e36..5a8aaf1 100644 --- a/data_harness/core/session/session.py +++ b/data_harness/core/session/session.py @@ -26,7 +26,12 @@ utc_now, ) from data_harness.core.session.store import MemorySessionStore, SessionStore -from data_harness.llm.types import Message, TextBlock +from data_harness.llm.types import ( + Message, + TextBlock, + ToolResultBlock, + ToolUseBlock, +) #: Turns a custom entry into messages for the model, or nothing. #: @@ -90,6 +95,24 @@ def append_custom(self, custom_type: str, data: dict[str, Any]) -> str: 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, @@ -170,13 +193,28 @@ def context_entries(self, from_id: str | None = None) -> list[Entry]: for entry in path[:compaction_index]: if entry.id == newest_compaction.first_kept_entry_id: keeping = True - if keeping: - kept.append(entry) + 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.""" + """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: @@ -199,7 +237,7 @@ def build_context(self, from_id: str | None = None) -> list[Message]: projector = self.projectors.get(entry.custom_type) if projector is not None: messages.extend(projector(entry)) - return messages + return _drop_unpaired_tool_blocks(messages) # ── reporting ─────────────────────────────────────────────────────────── @@ -226,6 +264,45 @@ def custom_entries(self, custom_type: str) -> list[CustomEntry]: ] +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 index 10b0a93..bd5dd92 100644 --- a/data_harness/core/session/store.py +++ b/data_harness/core/session/store.py @@ -8,6 +8,7 @@ from __future__ import annotations +import copy from typing import Protocol, runtime_checkable from data_harness.core.session.entries import Entry, leaf_after @@ -79,6 +80,14 @@ 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" diff --git a/tests/test_session_tree.py b/tests/test_session_tree.py index 008fe56..7849b67 100644 --- a/tests/test_session_tree.py +++ b/tests/test_session_tree.py @@ -257,9 +257,30 @@ def test_an_entry_naming_an_unknown_parent_is_rejected(store): assert excinfo.value.code == "missing_parent" -def test_both_stores_satisfy_the_protocol(store): +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 ───────────────────────────────────────────────────────────── @@ -352,10 +373,17 @@ def test_a_future_format_version_is_rejected_not_guessed(tmp_path): assert excinfo.value.code == "invalid_session" -def test_a_corrupt_entry_line_names_the_line(tmp_path): +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" + json.dumps({"type": "session", "version": 1, "id": "x"}) + + "\nnot json\n" + + good + + "\n" ) with pytest.raises(SessionStoreError) as excinfo: JsonlSessionStore.open(path) @@ -363,6 +391,58 @@ def test_a_corrupt_entry_line_names_the_line(tmp_path): 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( @@ -520,3 +600,240 @@ async def test_a_streamed_run_records_the_same_way(tmp_path): 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. + + Doing so emitted the older summary *after* the newer one and resurrected + the very entries the older one had dropped. + """ + session = Session(store) + session.append_message(say("q1")) + session.append_message(say("a1", "assistant")) + session.append_compaction( + "first summary", first_kept_entry_id=None, tokens_before=1 + ) + keep = 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 + assert "second summary" in context[0] + assert "q1" not in context and "a1" not in context + assert context[1:] == ["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")) + + assert texts(memory.build_context()) == ["original"] + assert texts(on_disk.build_context()) == ["original"] From 9138b0a448d8466b7cc095fd36cd72e1323fbbfc Mon Sep 17 00:00:00 2001 From: Max Khor Date: Tue, 4 Aug 2026 08:42:50 +0100 Subject: [PATCH 10/15] Phase 3 review round 2: finish the snapshot fix, correct two weak tests Mutation-testing the previous round's fixes found two survivors, both of which turned out to be my tests being wrong rather than the code. - The JSONL store serialised a snapshot to the file but kept the caller's object in memory, so a later mutation made its live view disagree with its own file: reading the session back showed something different from reading it now. The previous round only fixed the memory store. Both snapshot now. - test_both_stores_agree_after_a_message_is_mutated compared messages with a helper that reads content[0], so appending a second block was invisible to it. It compares every block now, and that is what exposed the store bug above. - test_stacked_compactions_leave_one_summary_in_order put the older compaction outside the newer one's kept tail, where skipping it changes nothing. The kept tail now starts before the first compaction, which is the case the skip exists for. All ten mutations against this phase's fixes are killed: unpaired tool blocks, stacked compaction, unvalidated compaction cut, both stores' copying, torn-tail tolerance, open_or_create, run() branching, reminder recording, message roles, and is_error. 840 passing. --- data_harness/core/session/jsonl.py | 10 ++++++++++ tests/test_session_tree.py | 25 ++++++++++++++++--------- 2 files changed, 26 insertions(+), 9 deletions(-) diff --git a/data_harness/core/session/jsonl.py b/data_harness/core/session/jsonl.py index a8e68d7..90fa488 100644 --- a/data_harness/core/session/jsonl.py +++ b/data_harness/core/session/jsonl.py @@ -13,6 +13,7 @@ from __future__ import annotations +import copy import dataclasses import json from pathlib import Path @@ -232,6 +233,14 @@ 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" @@ -250,6 +259,7 @@ def append(self, entry: Entry) -> None: ) 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) diff --git a/tests/test_session_tree.py b/tests/test_session_tree.py index 7849b67..9614de9 100644 --- a/tests/test_session_tree.py +++ b/tests/test_session_tree.py @@ -740,26 +740,27 @@ def test_compacting_from_an_entry_off_the_path_is_rejected(store): def test_stacked_compactions_leave_one_summary_in_order(store): """An older compaction inside the kept tail must not be replayed. - Doing so emitted the older summary *after* the newer one and resurrected - the very entries the older one had dropped. + 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")) - session.append_message(say("a1", "assistant")) + keep = session.append_message(say("a1", "assistant")) session.append_compaction( "first summary", first_kept_entry_id=None, tokens_before=1 ) - keep = session.append_message(say("q2")) + 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 + assert sum("summary" in c for c in context) == 1, context assert "second summary" in context[0] - assert "q1" not in context and "a1" not in context - assert context[1:] == ["q2", "q3"] + assert "first summary" not in " ".join(context) + assert context[1:] == ["a1", "q2", "q3"] # ── the working copy and the log agree ────────────────────────────────────── @@ -835,5 +836,11 @@ def test_both_stores_agree_after_a_message_is_mutated(tmp_path): message.content.append(TextBlock(text="added later")) - assert texts(memory.build_context()) == ["original"] - assert texts(on_disk.build_context()) == ["original"] + # 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"]] From f181890c38e29ab7fc4e531fe2e6a92ef838c0e9 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Tue, 4 Aug 2026 08:46:12 +0100 Subject: [PATCH 11/15] Phase 4: hooks, and the approval gate rebuilt on top of them The loop had one extension point, 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 tool dispatcher and was keyed on the literal string 'python_interpreter', so the one general thing the loop could do about a tool call was reachable by exactly one feature. 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 proof that the mechanism is sufficient rather than decorative is that the approval gate is now built from it: make_code_gate returns an ordinary BeforeToolCall hook, registered like any other when on_code or code_only is set, and a test asserts it is there rather than special-cased in the loop. The loop no longer mentions a tool by name. Design points worth keeping: - Every hook sees an event even after another has decided. A hook recording spend must not be skipped because an unrelated one asked to stop; precedence between conflicting decisions belongs to the caller. - Block defaults to is_error=False. A refusal is a decision, not a malfunction, and telling the model its code was broken makes it rewrite and retry rather than stop. - AfterTurn is where a spend cap belongs, because the tokens are already counted and the decision is made on real numbers. BeforeTurn stopping spends nothing at all. - Hooks must not raise. One that does is reported as HookError naming the hook and the event, rather than losing the run mid-turn. - register_reminder still works, as a BeforeTurn hook. 861 passing (was 841). Layer check still clean. --- data_harness/core/hooks.py | 183 +++++++++++++++++ data_harness/core/loop.py | 194 +++++++++++++----- data_harness/data/harness.py | 7 + tests/test_hooks.py | 381 +++++++++++++++++++++++++++++++++++ 4 files changed, 719 insertions(+), 46 deletions(-) create mode 100644 data_harness/core/hooks.py create mode 100644 tests/test_hooks.py diff --git a/data_harness/core/hooks.py b/data_harness/core/hooks.py new file mode 100644 index 0000000..2c87306 --- /dev/null +++ b/data_harness/core/hooks.py @@ -0,0 +1,183 @@ +"""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.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) + + +class HookError(Exception): + """A hook raised. Names the hook and the event so the culprit is obvious.""" + + 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[_E]) -> Any: + """The first decision of type ``kind``, or ``None``.""" + for decision in self.emit(event): + if isinstance(decision, kind): + return decision + return None diff --git a/data_harness/core/loop.py b/data_harness/core/loop.py index 3499777..9300427 100644 --- a/data_harness/core/loop.py +++ b/data_harness/core/loop.py @@ -47,6 +47,18 @@ from typing import Any, Literal, TypeVar from data_harness.core.environment import NullEnvironment, RunEnvironment +from data_harness.core.hooks import ( + AfterToolCall, + AfterTurn, + BeforeToolCall, + BeforeTurn, + Block, + Event, + 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 @@ -155,25 +167,37 @@ def _require_answer(effect: Effect, answer: Answer) -> Ok | Failed: # ── helpers ───────────────────────────────────────────────────────────────── -def _evaluate_code_gate( - on_code: Callable[[str], object] | None, - code_only: bool, - code: str, -) -> str | None: - """Decide whether interpreter ``code`` may run. +#: 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. - Returns a string to short-circuit execution (a dry-run echo or a denial - message returned to the model), or ``None`` to proceed. + 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. """ - 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 + + 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): @@ -275,6 +299,7 @@ def __init__( on_code: Callable[[str], object] | None = None, code_only: bool = False, session: Session | None = None, + hooks: HookRegistry | None = None, ) -> None: if max_turns < 1: raise ValueError(f"max_turns must be at least 1, got {max_turns!r}") @@ -286,6 +311,10 @@ def __init__( environment if environment is not None else NullEnvironment() ) self._session = session if session is not None else Session() + self._hooks = hooks 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 # 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. @@ -299,16 +328,26 @@ def __init__( self._usage_so_far = Usage() def register_reminder(self, hook: Callable[[int, int], str | None]) -> None: - """Register a suffix reminder hook called before each provider turn. + """Register a suffix reminder 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. + 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 @@ -391,6 +430,7 @@ def _record(self, message: Message) -> 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 @@ -401,6 +441,7 @@ def _begin_run(self, user_message: str) -> 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)])) @@ -458,6 +499,18 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: for turn in range(1, self._max_turns + 1): self._turns_completed = turn 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, + ) + return visible_tools = [t for t in self._tools if t.visible] with time_block() as tb: @@ -506,7 +559,7 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: 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) + 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) @@ -538,6 +591,32 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: 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, + ) + return + if response.stop_reason != StopReason.TOOL_USE: self._last_result = self._build_result( text=_extract_text(response), @@ -559,7 +638,7 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: return def _dispatch( - self, content: list + 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} @@ -580,22 +659,24 @@ def _dispatch( ) 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: - finished.append( - ( - tub.tool_name, - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=gate, - is_error=False, - ), - ) + blocked = self._hooks.first( + BeforeToolCall( + turn=turn, tool_name=tub.tool_name, tool_input=tub.tool_input + ), + Block, + ) + if blocked is not None: + finished.append( + ( + tub.tool_name, + ToolResultBlock( + tool_use_id=tub.tool_use_id, + content=blocked.reason, + is_error=blocked.is_error, + ), ) - continue + ) + continue request = CallTool( tool_use_id=tub.tool_use_id, @@ -633,16 +714,25 @@ def _dispatch( ) continue - finished.append( - ( - tub.tool_name, - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=output, - is_error=False, - ), - ) + block = ToolResultBlock( + tool_use_id=tub.tool_use_id, content=output, is_error=False + ) + 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)) return finished @@ -681,6 +771,14 @@ def _apply_reminders(self, turn: int) -> None: 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) @@ -754,6 +852,7 @@ def __init__( on_code: Callable[[str], object] | None = None, code_only: bool = False, session: Session | None = None, + hooks: HookRegistry | None = None, ) -> None: super().__init__( system=system, @@ -764,6 +863,7 @@ def __init__( on_code=on_code, code_only=code_only, session=session, + hooks=hooks, ) self._adapter = adapter @@ -908,6 +1008,7 @@ def __init__( on_code: Callable[[str], object] | None = None, code_only: bool = False, session: Session | None = None, + hooks: HookRegistry | None = None, ) -> None: super().__init__( system=system, @@ -918,6 +1019,7 @@ def __init__( on_code=on_code, code_only=code_only, session=session, + hooks=hooks, ) self._adapter = adapter diff --git a/data_harness/data/harness.py b/data_harness/data/harness.py index d09930f..edd28e8 100644 --- a/data_harness/data/harness.py +++ b/data_harness/data/harness.py @@ -13,6 +13,7 @@ from collections.abc import Callable +from data_harness.core.hooks import HookRegistry from data_harness.core.loop import ( Answer, CallProvider, @@ -52,6 +53,8 @@ class Harness(_CoreHarness): 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. """ @@ -67,6 +70,7 @@ def __init__( on_code: Callable[[str], object] | None = None, code_only: bool = False, session: Session | None = None, + hooks: HookRegistry | None = None, ) -> None: super().__init__( adapter=adapter, @@ -78,6 +82,7 @@ def __init__( on_code=on_code, code_only=code_only, session=session, + hooks=hooks, ) @property @@ -103,6 +108,7 @@ def __init__( on_code: Callable[[str], object] | None = None, code_only: bool = False, session: Session | None = None, + hooks: HookRegistry | None = None, ) -> None: super().__init__( adapter=adapter, @@ -114,6 +120,7 @@ def __init__( on_code=on_code, code_only=code_only, session=session, + hooks=hooks, ) @property diff --git a/tests/test_hooks.py b/tests/test_hooks.py new file mode 100644 index 0000000..774eb1f --- /dev/null +++ b/tests/test_hooks.py @@ -0,0 +1,381 @@ +"""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 + + +def test_after_turn_sees_the_tokens_already_spent(tmp_path): + seen: list[tuple[int, int]] = [] + harness = Harness( + adapter=FakeAdapter([FakeAdapter.text("done")]), + system="sys", + tools=[], + run_dir=str(tmp_path), + ) + harness.on(AfterTurn, lambda e: seen.append((e.input_tokens, e.output_tokens))) + + harness.run("go") + + assert seen == [(0, 0)] # FakeAdapter.text reports no usage + + +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" From 59fb810bba9c3a43e9d24652cd53e37801a57762 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Tue, 4 Aug 2026 09:00:23 +0100 Subject: [PATCH 12/15] Phase 4 review: fix nine defects in the hook system An independent review found the mechanism well shaped but three of the commit's own stated invariants false in the code that shipped. - A HookRegistry passed to the constructor was mutated in place: the harness added its approval gate to the CALLER's object, so a gate configured on one harness governed every other harness sharing the registry. That is exactly the reuse the hooks= parameter is documented to enable. Registries are copied on ingest now. - HookError escaped _plan entirely: no RunResult, last_result left None, and the tokens already billed discarded. Three docstrings and the commit message all claimed it became a failed run. Now it does, carrying the usage spent up to the failure. - Stop(reason) was written to a private field and never read, so a capped run was byte-identical to a model that answered with nothing: same status, same empty text, no error. RunResult.stopped_by records it. Also: - AfterToolCall fired only for successful calls, so a redaction hook could not see the case most likely to need it: an exception repr can carry a connection string. Every result now goes through one settle() path, including tool-not-found, handler failures, and blocked calls. - BeforeToolCall did not fire for an unknown tool name, so a policy hook could not observe attempts on tools it does not know about. - Agent had no way to register hooks at all, and since it builds a fresh Harness per run, one registered on a harness would be gone by the next call. Agent.on() and Agent.hooks now exist and are forwarded. - HookRegistry.first was annotated with the event TypeVar rather than a decision one, so its return was silently Any. - _on_code/_code_only on the harness are dead as switches (the gate closes over them at construction); documented rather than left as a trap. Tests: 873 passing (was 861). Four mutations had survived the previous round and now die: an AfterTurn stop discarding the model's answer, a BeforeTurn stop reporting error, a BeforeTurn stop dropping accumulated usage, and AfterTurn always reporting zero tokens. The token test was vacuous, asserting (0, 0) against an adapter that reports (0, 0); it now uses two turns and a scoring adapter so cumulative semantics are actually pinned. --- data_harness/app/agent.py | 20 +++ data_harness/core/hooks.py | 20 ++- data_harness/core/loop.py | 123 ++++++++------- data_harness/core/result.py | 4 + tests/test_hooks.py | 298 +++++++++++++++++++++++++++++++++++- 5 files changed, 398 insertions(+), 67 deletions(-) diff --git a/data_harness/app/agent.py b/data_harness/app/agent.py index 549b46d..dc3d21f 100644 --- a/data_harness/app/agent.py +++ b/data_harness/app/agent.py @@ -24,6 +24,7 @@ from pathlib import Path 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 @@ -119,6 +120,7 @@ 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._system = system self._max_turns = max_turns @@ -138,6 +140,7 @@ def __init__( 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 ──────────────────────────────────────────────── @@ -204,6 +207,22 @@ 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: @@ -466,6 +485,7 @@ def _harness_kwargs( "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) diff --git a/data_harness/core/hooks.py b/data_harness/core/hooks.py index 2c87306..3466635 100644 --- a/data_harness/core/hooks.py +++ b/data_harness/core/hooks.py @@ -128,6 +128,7 @@ class Stop: Hook = Callable[[Any], Decision | None] _E = TypeVar("_E", bound=Event) +_D = TypeVar("_D", bound=Decision) class HookError(Exception): @@ -175,9 +176,24 @@ def emit(self, event: Event) -> list[Decision]: decisions.append(decision) return decisions - def first(self, event: Event, kind: type[_E]) -> Any: - """The first decision of type ``kind``, or ``None``.""" + 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/core/loop.py b/data_harness/core/loop.py index 9300427..7ef6224 100644 --- a/data_harness/core/loop.py +++ b/data_harness/core/loop.py @@ -54,6 +54,7 @@ BeforeTurn, Block, Event, + HookError, HookRegistry, Reminder, Replace, @@ -311,7 +312,7 @@ def __init__( environment if environment is not None else NullEnvironment() ) self._session = session if session is not None else Session() - self._hooks = hooks if hooks is not None else HookRegistry() + 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 @@ -321,6 +322,9 @@ def __init__( 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 @@ -476,6 +480,19 @@ def _plan(self) -> Generator[Effect, Any, None]: 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 @@ -509,6 +526,7 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: 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] @@ -614,6 +632,7 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: turns=turn, stop_reason=response.stop_reason, usage=total_usage, + stopped_by=self._stop_reason, ) return @@ -644,19 +663,48 @@ def _dispatch( 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: - finished.append( - ( - tub.tool_name, - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=f"Tool not found: {tub.tool_name!r}", - is_error=True, - ), - ) + # 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( @@ -666,16 +714,7 @@ def _dispatch( Block, ) if blocked is not None: - finished.append( - ( - tub.tool_name, - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=blocked.reason, - is_error=blocked.is_error, - ), - ) - ) + settle(tub, blocked.reason, blocked.is_error) continue request = CallTool( @@ -687,52 +726,16 @@ def _dispatch( answer = _require_answer(request, (yield request)) if isinstance(answer, Failed): - finished.append( - ( - tub.tool_name, - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=repr(answer.error), - is_error=True, - ), - ) - ) + 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 - finished.append( - ( - tub.tool_name, - ToolResultBlock( - tool_use_id=tub.tool_use_id, - content=repr(exc), - is_error=True, - ), - ) - ) + settle(tub, repr(exc), True) continue - block = ToolResultBlock( - tool_use_id=tub.tool_use_id, content=output, is_error=False - ) - 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)) + settle(tub, output, False) return finished @@ -747,6 +750,7 @@ def _build_result( stop_reason: StopReason | None, usage: Usage, error: str | None = None, + stopped_by: str | None = None, ) -> RunResult: state = self._environment.capture() return RunResult( @@ -761,6 +765,7 @@ def _build_result( value=state.value, charts=state.artifacts, error=error, + stopped_by=stopped_by, ) def _apply_reminders(self, turn: int) -> None: diff --git a/data_harness/core/result.py b/data_harness/core/result.py index e691266..d0b3e9a 100644 --- a/data_harness/core/result.py +++ b/data_harness/core/result.py @@ -73,6 +73,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 +96,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) diff --git a/tests/test_hooks.py b/tests/test_hooks.py index 774eb1f..fcac354 100644 --- a/tests/test_hooks.py +++ b/tests/test_hooks.py @@ -281,19 +281,32 @@ def test_an_after_turn_hook_can_stop_the_run(tmp_path): assert len(harness.messages) == 3 # user, assistant, tool results -def test_after_turn_sees_the_tokens_already_spent(tmp_path): +@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 = Harness( - adapter=FakeAdapter([FakeAdapter.text("done")]), + harness = AsyncHarness( + adapter=FakeAsyncAdapter( + [ + FakeAsyncAdapter.tool_use("tu_1", "echo", {"value": "hi"}), + FakeAsyncAdapter.text("done"), + ] + ), system="sys", - tools=[], + tools=[echo_spec()], run_dir=str(tmp_path), ) harness.on(AfterTurn, lambda e: seen.append((e.input_tokens, e.output_tokens))) - harness.run("go") + await harness.run("go") - assert seen == [(0, 0)] # FakeAdapter.text reports no usage + # 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): @@ -379,3 +392,276 @@ def test_a_registry_can_be_passed_to_the_constructor(tmp_path): 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] From 1709883d6c605b6e90b1545db41b8a2f61932dc3 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Tue, 4 Aug 2026 09:04:52 +0100 Subject: [PATCH 13/15] Phase 5: compaction, and a real error taxonomy Two things the harness had no answer for. CONTEXT. 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. Compaction summarises older turns and replays the summary in their place. Two properties make it safe: - The cut lands on a turn boundary. An assistant tool call is never separated from the result answering it, and the conversation never opens with an assistant message. Both produce transcripts providers reject outright, and the split pair is the first bug everyone writes here. When there is nowhere safe to cut, nothing is cut. - 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, and a test asserts the undo. The summariser is injected, because summarising means calling a model and which model at what cost is the application's decision. Token estimation is deliberately crude: the decision it feeds is 'are we near the limit', and being wrong by 20% moves the trigger slightly rather than breaking anything, while a real tokeniser would tie core to a provider's vocabulary. Worth recording: compaction costs this library less than it costs a coding agent. 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. Writing the test proved it: a 4000-character tool result never reached the transcript at all, because the cache had already turned it into a handle and a snapshot. ERRORS. Failures were bare RuntimeError/KeyError, or repr(exc) stuffed into a string, so a caller could not tell a rate-limited provider from a typo in the model's pandas from a sandbox timeout. Those are three different decisions: retry, show the user, bill the attempt. Every error now derives from DataHarnessError and carries a stable code; codes are API, messages are not. Added ProviderError, ExecutionError, and ConfigurationError; folded the existing SessionStoreError and HookError into the same base. MaxTurnsExceeded stays a RuntimeError and ToolNotFoundError a KeyError, since callers catch them that way. 903 passing (was 873). --- data_harness/__init__.py | 48 ++++ data_harness/core/compaction.py | 206 +++++++++++++++ data_harness/core/exceptions.py | 82 +++++- data_harness/core/hooks.py | 5 +- data_harness/core/loop.py | 20 ++ data_harness/core/session/store.py | 13 +- data_harness/data/harness.py | 5 + tests/test_compaction.py | 393 +++++++++++++++++++++++++++++ 8 files changed, 757 insertions(+), 15 deletions(-) create mode 100644 data_harness/core/compaction.py create mode 100644 tests/test_compaction.py diff --git a/data_harness/__init__.py b/data_harness/__init__.py index 484252b..69edf28 100644 --- a/data_harness/__init__.py +++ b/data_harness/__init__.py @@ -30,12 +30,39 @@ 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.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 @@ -70,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", @@ -110,6 +156,8 @@ "ToolUseBlock", "Usage", "ask", + "estimate_tokens", + "make_compactor", "load_dataframe", "mcp_tool_specs", "resolve_adapter", diff --git a/data_harness/core/compaction.py b/data_harness/core/compaction.py new file mode 100644 index 0000000..1ecb6f7 --- /dev/null +++ b/data_harness/core/compaction.py @@ -0,0 +1,206 @@ +"""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.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 ValueError( + "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/exceptions.py b/data_harness/core/exceptions.py index 3306907..df1179c 100644 --- a/data_harness/core/exceptions.py +++ b/data_harness/core/exceptions.py @@ -1,3 +1,14 @@ +"""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 @@ -6,23 +17,78 @@ from data_harness.llm.providers.base import NormalizedResponse -class MaxTurnsExceeded(RuntimeError): - """Raised when the ReAct loop reaches ``max_turns`` without an end-turn stop. +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: The number of turns that were executed before the limit was hit. + turns: How many turns ran before the limit was hit. last_response: The final provider response, if available. """ - def __init__(self, turns: int, last_response: "NormalizedResponse | None" = None): + 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(KeyError): - """Raised when a tool invocation names a tool that is not registered.""" +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): + """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. + """ + + code = "provider_error" + + +class ExecutionError(DataHarnessError): + """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 to fix. + """ + + code = "execution_error" + + +class ConfigurationError(DataHarnessError): + """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. + """ -class SubagentRecursionError(RuntimeError): - """Raised when a subagent attempts to spawn another subagent.""" + code = "configuration_error" diff --git a/data_harness/core/hooks.py b/data_harness/core/hooks.py index 3466635..2b6f5a1 100644 --- a/data_harness/core/hooks.py +++ b/data_harness/core/hooks.py @@ -23,6 +23,7 @@ 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 ────────────────────────────────────────────────────────────────── @@ -131,9 +132,11 @@ class Stop: _D = TypeVar("_D", bound=Decision) -class HookError(Exception): +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__( diff --git a/data_harness/core/loop.py b/data_harness/core/loop.py index 7ef6224..d8dc6b2 100644 --- a/data_harness/core/loop.py +++ b/data_harness/core/loop.py @@ -46,6 +46,7 @@ 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.hooks import ( AfterToolCall, @@ -301,6 +302,7 @@ def __init__( code_only: bool = False, session: Session | None = None, hooks: HookRegistry | None = None, + compactor: Compactor | None = None, ) -> None: if max_turns < 1: raise ValueError(f"max_turns must be at least 1, got {max_turns!r}") @@ -316,6 +318,7 @@ def __init__( 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. @@ -515,6 +518,7 @@ def _turns(self, total_usage: Usage) -> Generator[Effect, Any, None]: 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. @@ -768,6 +772,18 @@ def _build_result( 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. + """ + if self._compactor is None: + return + if self._compactor(self._session) is not None: + self._messages[:] = self._session.build_context() + def _apply_reminders(self, turn: int) -> None: reminder_texts: list[str] = [] @@ -858,6 +874,7 @@ def __init__( code_only: bool = False, session: Session | None = None, hooks: HookRegistry | None = None, + compactor: Compactor | None = None, ) -> None: super().__init__( system=system, @@ -869,6 +886,7 @@ def __init__( code_only=code_only, session=session, hooks=hooks, + compactor=compactor, ) self._adapter = adapter @@ -1014,6 +1032,7 @@ def __init__( code_only: bool = False, session: Session | None = None, hooks: HookRegistry | None = None, + compactor: Compactor | None = None, ) -> None: super().__init__( system=system, @@ -1025,6 +1044,7 @@ def __init__( code_only=code_only, session=session, hooks=hooks, + compactor=compactor, ) self._adapter = adapter diff --git a/data_harness/core/session/store.py b/data_harness/core/session/store.py index bd5dd92..19fa7f1 100644 --- a/data_harness/core/session/store.py +++ b/data_harness/core/session/store.py @@ -11,19 +11,20 @@ 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(Exception): - """Raised for a session that cannot be read or is internally inconsistent. +class SessionStoreError(DataHarnessError): + """A session cannot be read, or is internally inconsistent. - Carries a stable ``code`` so a caller can tell a missing entry from a - corrupt file without matching on message text. + 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) - self.code = code + super().__init__(message, code=code) @runtime_checkable diff --git a/data_harness/data/harness.py b/data_harness/data/harness.py index edd28e8..040a878 100644 --- a/data_harness/data/harness.py +++ b/data_harness/data/harness.py @@ -13,6 +13,7 @@ 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, @@ -71,6 +72,7 @@ def __init__( code_only: bool = False, session: Session | None = None, hooks: HookRegistry | None = None, + compactor: Compactor | None = None, ) -> None: super().__init__( adapter=adapter, @@ -83,6 +85,7 @@ def __init__( code_only=code_only, session=session, hooks=hooks, + compactor=compactor, ) @property @@ -109,6 +112,7 @@ def __init__( code_only: bool = False, session: Session | None = None, hooks: HookRegistry | None = None, + compactor: Compactor | None = None, ) -> None: super().__init__( adapter=adapter, @@ -121,6 +125,7 @@ def __init__( code_only=code_only, session=session, hooks=hooks, + compactor=compactor, ) @property diff --git a/tests/test_compaction.py b/tests/test_compaction.py new file mode 100644 index 0000000..633b342 --- /dev/null +++ b/tests/test_compaction.py @@ -0,0 +1,393 @@ +"""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) From 3b62d80862d7c3628ba8418a6f2586f99ab2df83 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Tue, 4 Aug 2026 17:57:44 +0100 Subject: [PATCH 14/15] Phase 5 review: make the taxonomy real, and stop losing runs to compaction Verifying Phase 5 directly (the verification agent hit its session limit before reporting) found the error-taxonomy claim false and one real defect. ProviderError, ExecutionError and ConfigurationError were declared and never raised anywhere. The commit claiming a caller could now tell a rate-limited provider from a typo in the model's pandas from a sandbox timeout was simply untrue: nothing raised any of them. They are wired up now. - unwrap_text raises ProviderError for a failed run, not a bare RuntimeError. - The sandbox raises ExecutionError when code never got to run (wall-clock timeout, killed process) and PythonInterpreterError when the model's code ran and failed. That is the distinction the docstrings claimed; it did not exist. PythonInterpreterError joins the taxonomy too. - max_turns and CompactionSettings validation raise ConfigurationError. - ProviderError and ExecutionError also derive from RuntimeError, and ConfigurationError from ValueError, so callers written before the taxonomy keep working. A compactor that raised killed the run: exception escaped, no RunResult, and usage discarded. Summarising means calling a model, so a rate limit there is ordinary, and compaction is an optimisation: carrying on with the full context is strictly better than losing a run that was otherwise fine. If the context really was too big, the next provider call fails and reports it where it belongs. The attempt is recorded as a compaction_failed entry either way. The layer test caught me trying to import the taxonomy into llm/types.py. llm is the bottom layer and may not import core, so that check stays a plain ValueError, with a comment saying why. Verified directly this round: - Seven adversarial cut-point cases produce no unpaired tool blocks, no assistant-first conversation, and no empty messages: parallel tool calls, a user message mixing text and tool results, an assistant message mixing text and a call, a conversation opening with an assistant message, consecutive user messages, and repeated compaction. - Thirty turns compacting every turn stays bounded (254 tokens against a 300 trigger) rather than growing. - context_entries() is a subset of branch(), so a cut id found in one is never rejected by the other's validation. - Ten mutations, all killed. Three initially survived and each was a bad test rather than good code: the empty-replacement guard was never reached because the fixture was under the compaction trigger, ProviderError was only asserted as a RuntimeError, and the sandbox timeout test exercises the CPU-kill branch rather than the wall-clock one, which needed driving directly since model code cannot sleep. 913 passing (was 903). --- data_harness/core/compaction.py | 3 +- data_harness/core/exceptions.py | 14 +- data_harness/core/loop.py | 21 ++- data_harness/core/result.py | 11 +- data_harness/data/tools/interpreter.py | 12 +- data_harness/data/tools/sandbox.py | 7 +- data_harness/llm/types.py | 4 + tests/test_compaction.py | 201 +++++++++++++++++++++++++ tests/test_sandbox.py | 49 +++++- 9 files changed, 305 insertions(+), 17 deletions(-) diff --git a/data_harness/core/compaction.py b/data_harness/core/compaction.py index 1ecb6f7..4e7a317 100644 --- a/data_harness/core/compaction.py +++ b/data_harness/core/compaction.py @@ -27,6 +27,7 @@ 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, @@ -61,7 +62,7 @@ class CompactionSettings: def __post_init__(self) -> None: if self.keep_recent_tokens >= self.context_window - self.reserve_tokens: - raise ValueError( + 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}" diff --git a/data_harness/core/exceptions.py b/data_harness/core/exceptions.py index df1179c..9940025 100644 --- a/data_harness/core/exceptions.py +++ b/data_harness/core/exceptions.py @@ -64,31 +64,35 @@ class SubagentRecursionError(DataHarnessError, RuntimeError): code = "subagent_recursion" -class ProviderError(DataHarnessError): +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): +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 to fix. + is handed back to it as a tool result to fix. """ code = "execution_error" -class ConfigurationError(DataHarnessError): +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. + 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/loop.py b/data_harness/core/loop.py index d8dc6b2..dc36adf 100644 --- a/data_harness/core/loop.py +++ b/data_harness/core/loop.py @@ -48,6 +48,7 @@ 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, @@ -305,7 +306,7 @@ def __init__( compactor: Compactor | None = None, ) -> None: if max_turns < 1: - raise ValueError(f"max_turns must be at least 1, got {max_turns!r}") + 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 @@ -778,10 +779,26 @@ def _compact_if_needed(self) -> None: 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 - if self._compactor(self._session) is not None: + 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: diff --git a/data_harness/core/result.py b/data_harness/core/result.py index d0b3e9a..e67234f 100644 --- a/data_harness/core/result.py +++ b/data_harness/core/result.py @@ -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'." ) @@ -145,12 +147,13 @@ def unwrap_text(result: RunResult) -> str: Raises: MaxTurnsExceeded: If the run hit its turn cap. - RuntimeError: If the provider failed. + 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 + from data_harness.core.exceptions import MaxTurnsExceeded, ProviderError if result.status == "max_turns_exceeded": raise MaxTurnsExceeded(result.turns) if result.status == "error": - raise RuntimeError(result.error or "unknown error") + raise ProviderError(result.error or "unknown error") return result.text diff --git a/data_harness/data/tools/interpreter.py b/data_harness/data/tools/interpreter.py index 0b17755..43263df 100644 --- a/data_harness/data/tools/interpreter.py +++ b/data_harness/data/tools/interpreter.py @@ -13,6 +13,7 @@ from typing import Any 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 @@ -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/data/tools/sandbox.py b/data_harness/data/tools/sandbox.py index 02f3150..0e237f3 100644 --- a/data_harness/data/tools/sandbox.py +++ b/data_harness/data/tools/sandbox.py @@ -22,6 +22,7 @@ from typing import Any 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, @@ -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/llm/types.py b/data_harness/llm/types.py index af8ba15..70d4c8f 100644 --- a/data_harness/llm/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/tests/test_compaction.py b/tests/test_compaction.py index 633b342..96abebf 100644 --- a/tests/test_compaction.py +++ b/tests/test_compaction.py @@ -391,3 +391,204 @@ def test_the_legacy_base_classes_still_hold(): 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_sandbox.py b/tests/test_sandbox.py index 67a7712..b3c2d0f 100644 --- a/tests/test_sandbox.py +++ b/tests/test_sandbox.py @@ -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) From 400185862039c20dd0669d0108a8a07d941c1459 Mon Sep 17 00:00:00 2001 From: Max Khor Date: Tue, 4 Aug 2026 18:00:20 +0100 Subject: [PATCH 15/15] Document the refactor, and make two tests version-aware - docs/guide/design.md gained sections for the four layers, the one-loop two-driver split, the session tree, hooks, compaction, and the error taxonomy. The page described an architecture that no longer existed. - CHANGELOG has an Unreleased entry covering all five phases. - The two ambient-event-loop tests asserted a 3.10-3.13 precondition (policy._local._set_called) that does not exist on 3.14, where the implementation takes the get_event_loop() branch instead. The pristine check is now version-gated and the custom-policy test skipped above 3.14, with the reason stated. Verified: 913 passing on 3.10; on 3.14 from the built wheel only six tests fail, all for optional dependencies absent from that venv (matplotlib, sqlalchemy, IPython). Zero layer violations. data-harness-ui's usage pattern still produces identical token accounting. --- CHANGELOG.md | 39 ++++++++++++ docs/guide/design.md | 117 ++++++++++++++++++++++++++++++++++++ tests/test_loop_protocol.py | 17 +++++- 3 files changed, 171 insertions(+), 2 deletions(-) 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/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/tests/test_loop_protocol.py b/tests/test_loop_protocol.py index 4ab8d4a..cdecb91 100644 --- a/tests/test_loop_protocol.py +++ b/tests/test_loop_protocol.py @@ -430,8 +430,13 @@ def test_looking_up_the_ambient_loop_does_not_create_one(): probe = textwrap.dedent( """ import asyncio, sys - policy = asyncio.get_event_loop_policy() - assert policy._local._set_called is False, "process was not pristine" + + 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 @@ -729,6 +734,14 @@ def build(responses): 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.