diff --git a/.env.example b/.env.example index 9a1f7c8..f26ec94 100644 --- a/.env.example +++ b/.env.example @@ -6,3 +6,8 @@ CHAT_MODEL=deepseek-v4-flash LLM_TIMEOUT=30 LLM_TEMPERATURE=0.7 LLM_INCLUDE_USAGE=true + +# Optional Anthropic Messages adapter +# ANTHROPIC_API_KEY=sk-ant-x +# ANTHROPIC_MODEL=your-claude-model +# ANTHROPIC_MAX_TOKENS=4096 diff --git a/.gitmodules b/.gitmodules deleted file mode 100644 index 506b937..0000000 --- a/.gitmodules +++ /dev/null @@ -1,3 +0,0 @@ -[submodule "pi"] - path = pi - url = https://github.com/earendil-works/pi.git diff --git a/README.md b/README.md index 3d91f91..dd46c27 100644 --- a/README.md +++ b/README.md @@ -2,1035 +2,197 @@ [English](README.md) | [简体中文](README_zh-CN.md) -EJAgent Core is an OpenAI-first Agent Harness Core built around eight capability -areas: execution, control, extensibility, safety, context management, durable -history, branching, and observability. It provides the runtime mechanisms for -state, orchestration, context construction, tool dispatch, middleware, MCP, and -skills, while derived agents own concrete tools such as shell, file editing, -Git, or explicit completion. +EJAgent Core is a small, extensible runtime for one logical agent. It separates +one model–tool execution (`RuntimeKernel`) from durable state, lifecycle, and +control (`AgentHarness`). Conversation, Run Audit, and disposable model context +are distinct data domains. Requires Python 3.12 or newer. -## Core capabilities - -- Stateful `BaseAgent` with persistent conversation history and `reset()` -- Provider-neutral `ModelAdapter` boundary with an OpenAI-compatible adapter -- Public `AgentOrchestrator` for the provider-tool loop -- Structured `AgentRunResult`, `RunStatus`, and `StopReason` -- Explicit `RuntimePolicy` for loop and completion behavior -- `AgentContextBuilder` for non-mutating per-turn context projection -- OpenAI-first canonical `ToolDefinition` with legacy dictionary compatibility -- Composable `BaseHandler` and `MethodToolHandler` tool contracts -- `ToolRuntime` lifecycle, routing, middleware, and repeat-call protection -- Generic `ToolMiddleware` interception -- Optional fail-closed Tool Execution Policy middleware with approval and - atomic per-run call limits -- Optional local JSON Schema argument validation before policy and handlers -- Structured, cancellable Tool Progress events -- Provider-neutral token Usage and per-run budget guards -- Context pressure estimates, independent window budgets, and non-mutating - compaction preparation -- Explicit, cancellable compaction through a pluggable `Compactor`, canonical - `SummaryEntry`, and resumable Session snapshots -- Opt-in automatic compaction on context pressure and one safe recovery attempt - for provider-normalized context overflow -- Versioned Session serialization, append-only JSONL journals, and explicit - cross-process restoration -- Optional MCP integration through `McpToolHandler` -- Local skill discovery, metadata projection, and explicit context activation - -The core intentionally does not provide Bash, Git, filesystem, approval UI, or -finish tools. Those belong to a derived agent such as a CodeAgent. - -This core-boundary change removes the former `BashHandler`, `GitDiffHandler`, -`FinishHandler`, `HumanApproval`, and `BashApprovalMiddleware` public exports. -Derived agents should provide equivalent implementations when needed. - -## Installation +## Install ```bash pip install ejagent-core -# or, for a source checkout: uv sync -``` - -MCP support is optional. Install its extra only when the agent uses MCP: - -```bash -uv sync --extra mcp -# or: pip install "ejagent-core[mcp]" -``` - -Version 0.6.0 renames the PyPI distribution to `ejagent-core` and the Python -import package to `ejagent`: - -```python -from ejagent import BaseAgent +# source checkout +uv sync --locked --all-extras --group dev ``` -The former `simagentplg` import is not retained as a compatibility alias. -Existing JSONL Sessions remain readable; legacy internal message metadata is -filtered from provider requests during restoration. +Anthropic support is optional: `pip install 'ejagent-core[anthropic]'`. -## Configuration - -Copy `.env.example` to `.env` and provide model credentials: +Configure an OpenAI-compatible endpoint in `.env`: ```env MODEL_API_KEY=sk-xxxxxxxx -MODEL_URL=https://api.deepseek.com -CHAT_MODEL=deepseek-v4-flash +MODEL_URL=https://api.example.com/v1 +CHAT_MODEL=your-model LLM_TIMEOUT=60 LLM_TEMPERATURE=0.7 LLM_INCLUDE_USAGE=true ``` -`ModelConfig` belongs to `OpenAIModelAdapter`, rather than to `BaseAgent`. -Configuration can also be supplied directly: - -```python -from ejagent import ModelConfig - -config = ModelConfig( - model="deepseek-v4-flash", - api_key="sk-xxxxxxxx", - base_url="https://api.deepseek.com", -) -``` - -Other model providers can integrate with the core by implementing -`ModelAdapter.complete()` and optionally overriding `ModelAdapter.stream()`. -The adapter owns provider client creation, response normalization, streaming, -and optional startup/shutdown resources; `BaseAgent` only consumes -provider-neutral stream events and the normalized `AssistantMessage` contract. - -## Plain agent - -Conversation history is preserved across calls: - -```python -from ejagent import BaseAgent, ModelConfig, OpenAIModelAdapter - -agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="tutor", - system_prompt="You are a concise Python tutor.", -) - -first = await agent.runtime(task="Remember that I prefer Python.") -second = await agent.runtime(task="Which language do I prefer?") - -agent.reset() -await agent.shutdown() -``` - -Calls on the same agent are serialized to protect conversation state. - -## Structured runs - -`run()` exposes the core result protocol: - -```python -result = await agent.run(task="Explain the repository architecture.") - -print(result.status) -print(result.stop_reason) -print(result.turns) -print(result.output) -``` - -`runtime()` remains a compatibility wrapper. It returns `result.output` for a -completed run and raises `AgentRunError` for failed, rejected, or cancelled -runs. - -### Cancelling a run - -Each run owns an independent cancellation token. `abort()` requests -cancellation without waiting, while `wait_for_idle()` settles only after the -terminal event and all awaited event sinks have completed: - -```python -import asyncio - -run = asyncio.create_task(agent.run(task="Perform a long operation.")) - -agent.abort("stopped by user") -await agent.wait_for_idle() -result = await run -``` - -An externally aborted run returns `RunStatus.CANCELLED` with -`StopReason.EXTERNAL_ABORT`. The same agent can be reused for another run. -Model adapters, tool middleware, and tool handlers receive the run's -`CancellationToken`; long-running handlers should also use `try/finally` to -release resources such as subprocesses. - -### Steering an active run - -`steer()` queues guidance for the active Run without interrupting an in-flight -model response or Tool call: - -```python -receipt = await agent.steer( - "Do not modify the database; produce a migration plan instead." -) -if not receipt.accepted: - print(receipt.status) -``` - -Each Run owns a bounded FIFO queue; configure it with -`BaseAgent(..., steering_queue_capacity=16)`. Submission returns a typed -`ControlReceipt` with `ACCEPTED`, `AGENT_IDLE`, `QUEUE_FULL`, or `RUN_CLOSING` -instead of silently dropping input. Accepted guidance is consumed immediately -before the next Provider context is finalized, including a model retry after -Context Overflow recovery. A Tool that has already started always settles -before Steering reaches the model. - -Applied guidance becomes a normal User message in Agent history and emits -`SteeringApplied`; `SessionRecorder` stores it as `steering_applied`. If the Run -ends or is cancelled before another model-call safe point, the pending input is -not persisted and emits `SteeringDiscarded`. An accepted receipt therefore -means “queued”, while these events expose the final applied/discarded outcome. - -### Follow-up runs - -`follow_up()` queues an independent Run behind the active Run chain and returns -a waitable handle: - -```python -import asyncio - -initial = asyncio.create_task(agent.run(task="Inspect the deployment failure.")) -await asyncio.sleep(0) # allow the Run chain to become active -follow_up = await agent.follow_up("Now propose the smallest safe fix.") - -initial_result = await initial -if follow_up.accepted: - follow_up_result = await follow_up.wait() -``` - -Each Follow-up receives its own `run_id`, lifecycle events, Session Run, and -`AgentRunResult`. It starts only after the previous `AgentFinished` and every -awaited Event Sink have settled. A Run-chain gate keeps already-waiting direct -`run()`, `compact()`, and lifecycle operations from interleaving ahead of the -Follow-up FIFO. - -Configure the queue with `follow_up_queue_capacity`. Rejected handles expose an -immediate `ControlReceipt` with `AGENT_IDLE`, `QUEUE_FULL`, or `RUN_CLOSING`, and -`wait()` raises `FollowUpRejectedError`. By default, a failed or cancelled Run -discards remaining accepted handles with `FollowUpDiscardedError`; opt into -`FollowUpFailurePolicy.CONTINUE` to keep processing. Cancelling one caller's -`wait()` does not cancel the queued Run. `shutdown()` discards pending -Follow-ups, while `wait_for_idle()` includes accepted Follow-ups and their -terminal sinks. - -Pending Follow-ups remain process-local. `SessionRecorder` writes `run_started` -only when a Follow-up actually begins, so the queue is not presented as a -durable job system. - -### Continue existing history - -`continue_run()` starts a distinct Run from committed history without appending -another User message: - -```python -if agent.can_continue: - result = await agent.continue_run() -else: - print(agent.continue_rejection_reason) -``` - -Continue allocates a new `run_id`, emits `AgentContinued`, and is persisted as -`run_continued`. Session Runs expose `SessionRunIntent.TASK` or -`SessionRunIntent.CONTINUE`, with `task=None` for Continue. Restoring a finished -Session also restores the latest terminal result, so a new Agent instance can -continue the checked-out history. - -Only explicit safe terminal reasons are continuable. Active Agents, Shutdown, -missing previous Runs, unsupported failure/cancellation reasons, and unmatched -Tool Calls raise `ContinueRejectedError` with a typed -`ContinueRejectedReason`. Continue reuses Cancellation, Steering safe points, -Follow-up chaining, automatic Compaction, terminal Event Sink barriers, and -`wait_for_idle()` semantics. - -### Behavior Hooks - -`BehaviorHook` controls whether a non-terminal full Turn may advance to the -next Provider request. It is deliberately separate from read-only Event Sinks -and Tool Middleware: - -```python -from ejagent import BehaviorDecision - - -class StopAfterFirstToolTurn: - async def after_turn(self, snapshot, *, cancellation): - if snapshot.turn >= 1: - return BehaviorDecision.stop("paused at a safe Turn boundary") - return None - - -agent = BaseAgent( - model, - agent_id="controlled-agent", - behavior_hooks=[StopAfterFirstToolTurn()], -) -``` - -The decision point runs after `TurnCompleted` work, including every committed -Tool Result, and before another `TurnStarted` or Provider request. Hooks receive -a detached `TurnSnapshot` plus the Run's `CancellationToken`; they never receive -mutable `AgentState`. Multiple Hooks are awaited in declaration order, and the -first STOP short-circuits the rest. - -STOP completes the Run with `StopReason.BEHAVIOR_STOP`, normal `AgentFinished` -and Session `run_finished` records. This is a safe Continue boundary. Hook -exceptions become `RUNTIME_ERROR`; abort interrupts a slow Hook, and -`wait_for_idle()` includes Hook backpressure. Hooks are not invoked after a Turn -that already produced a terminal result. - -### Parallel read-only tools - -Parallel Tool Calls are opt-in at both the tool and Agent levels. A handler -must explicitly declare a tool `READ_ONLY`, and the Runtime Policy must enable -parallel calls: - -```python -from ejagent import ( - MethodToolHandler, - RuntimePolicy, - ToolDefinition, - ToolEffect, -) - -LOOKUP_TOOL = ToolDefinition( - name="lookup", - description="Look up one value.", - parameters={ - "type": "object", - "properties": {"key": {"type": "string"}}, - "required": ["key"], - }, - effect=ToolEffect.READ_ONLY, -) - - -class LookupHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((LOOKUP_TOOL,)) - - -agent = BaseAgent( - model, - agent_id="parallel-lookups", - handlers=[LookupHandler()], - runtime_policy=RuntimePolicy( - parallel_tool_calls=True, - max_parallel_tool_calls=4, - ), -) -``` - -The scheduler only parallelizes contiguous read-only calls. Unannotated, -unknown, and `SIDE_EFFECTING` tools remain sequential and form barriers between -read batches. The default policy is fully sequential, so existing handlers do -not change behavior after upgrading. - -Within a parallel batch, `ToolStarted` is emitted in source order and Progress -may interleave by `tool_call_id`. Handlers execute concurrently, but -`ToolCompleted`, Agent history, and Session Tool Messages are committed in the -Assistant Tool Call order. Errors remain ordinary Tool Results. External abort -settles every active and pending call; if a read-only batch returns a terminal -Tool Control, active peers settle before the first terminal decision in source -order stops the Run. Declaring `READ_ONLY` is a concurrency-safety contract for -the handler and its middleware, not an automatic side-effect inference. - -### Streaming responses - -`BaseAgent.run()` still returns one final `AgentRunResult`, while provisional -text and provisional reasoning are observed through typed Delta events: - -```python -from ejagent import AssistantThinkingDelta, AssistantTextDelta - - -class ConsoleSink: - async def emit(self, event): - if isinstance(event.payload, AssistantThinkingDelta): - print("[thinking]", event.payload.delta, end="") - elif isinstance(event.payload, AssistantTextDelta): - print(event.payload.delta, end="", flush=True) -``` - -`OpenAIModelAdapter` uses a real streaming request. Tool-call fragments are -assembled inside the provider adapter and only complete `AssistantMessage` -objects enter Agent state. Thinking Delta remains observation-only and is not -mixed into normal text or persisted to Session. Existing complete-only adapters -remain compatible through the default `ModelAdapter.stream()` implementation. -Session recording ignores provisional deltas and persists only -`MessageCompleted`. - -### Usage and run budgets - -`ModelResponseCompleted` carries optional provider-neutral `ModelUsage`. -Reported Usage is attached to internal agent messages and Session history, but -`AgentContextBuilder` removes it from the final `llm_messages` sent to the -Provider. `AgentRunResult.usage` aggregates all attempted requests while -preserving whether every request actually reported Usage: - -```python -result = await agent.run(task="Inspect the project.") - -print(result.usage.total_tokens) -print(result.usage.request_count) -print(result.usage.complete) -``` - -Unknown Usage is distinct from zero. Complete-only adapters remain compatible -and produce an incomplete `RunUsage` unless they override `stream()` with a -terminal Usage value. - -### Context pressure and compaction preparation - -Context window capacity is independent of cumulative run spend. Configure an -optional `CompactionPolicy` to assess the complete provider request before each -model call: - -```python -from ejagent import CompactionPolicy, ContextBudget - -context_policy = CompactionPolicy( - ContextBudget( - context_window=128_000, - reserve_tokens=16_000, - keep_recent_tokens=20_000, - ) -) - -agent = BaseAgent( - model, - agent_id="context-aware", - compaction_policy=context_policy, -) -``` - -The estimate combines the latest assistant `ModelUsage`, trailing messages, -and a UTF-8-aware heuristic lower bound that includes current tool schemas. -Each configured turn emits `ContextPressureEvaluated`. When the threshold is -reached, its `CompactionPreparation` separates protected messages, complete -old User/Assistant/Tool turns to summarize, and recent turns to keep. Tool -calls and results remain in the same turn. - -`CompactionPolicy` alone remains observation-only. Applications can call -`estimate_context_usage()` and `prepare_compaction()` directly, and can replace -the fallback through `MessageTokenEstimator`. - -### Automatic compaction and overflow recovery - -Automatic behavior is opt-in and reuses the same `CompactionPolicy` and -`Compactor`: - -```python -from ejagent import AutoCompactionPolicy - -agent = BaseAgent( - model, - agent_id="context-aware", - compaction_policy=context_policy, - compactor=my_compactor, - auto_compaction_policy=AutoCompactionPolicy(), -) -``` - -At the configured pressure threshold, Core compacts old complete turns, -rebuilds context, and dispatches the model request in the same Agent Run. If a -provider adapter raises `ContextOverflowError`, Core can compact, rebuild, and -retry once. A second overflow returns `StopReason.CONTEXT_OVERFLOW`; compactor -failure returns `StopReason.COMPACTION_FAILED`. Core never retries after text or -thinking deltas have been exposed, preventing duplicate provisional output. - -`AutoCompactionPolicy(compact_on_pressure=False)` keeps overflow recovery while -disabling proactive compaction. Set `enabled=False` or omit the policy to keep -all automatic behavior off. Provider adapters normalize overflow, rate-limit, -timeout, authentication, and other failures through `ModelProviderError` and -`ModelErrorKind`. - -### Explicit compaction - -A derived agent supplies the summary behavior through the cancellable -`Compactor` protocol, then invokes `compact()` explicitly: - -```python -agent = BaseAgent( - model, - agent_id="context-aware", - compaction_policy=context_policy, - compactor=my_compactor, -) - -compaction = await agent.compact() -print(compaction.status) -print(compaction.summary) -``` - -`ModelCompactor` adapts a borrowed `ModelAdapter` into this protocol while the -application still owns the summary prompt: - -```python -compactor = ModelCompactor( - summary_model, - context_builder=build_summary_context, - source="summary-model:v1", -) -``` - -The injected builder receives `CompactionRequest` and returns the complete -`ContextBuildResult`. The caller owns the borrowed model lifecycle, so Core -does not silently create another provider client or choose a prompt. - -The Core calls the Compactor with `CompactionRequest`, creates trusted range and -token metadata in `SummaryEntry`, then atomically installs protected messages + -Summary + recent turns. Failure or cancellation returns a structured -`CompactionResult` and leaves history unchanged. Repeated compaction passes the -previous Summary to the Compactor for merging and replaces the old Summary -message. - -`CompactionStarted`, `CompactionCompleted`, and `CompactionFailed` expose the -lifecycle. `abort()` and `wait_for_idle()` apply to compaction as well as normal -runs. `SessionRecorder` stores a compacted recovery snapshot while retaining -the original `SessionMessage` audit entries. Each operation exposes a stable -`operation_id` and `CompactionTrigger`. The Core does not choose a summary model -or prompt. - -## Durable Session journals - -`SessionRecorder` can use `JsonlSessionStorage` to append a versioned semantic -record for each accepted lifecycle mutation: - -```python -from ejagent import JsonlSessionStorage, SessionRecorder - -storage = JsonlSessionStorage("./sessions") -recorder = SessionRecorder(session_id="project-42", storage=storage) -agent = BaseAgent(model, agent_id="core-agent", event_sink=recorder) -await agent.run(task="remember this decision") -``` - -A different process can load the completed snapshot and explicitly restore a -new Agent: +## Quick Start ```python -saved = await storage.load("project-42") -if saved is not None: - resumed = BaseAgent(model, agent_id="core-agent", event_sink=recorder) - resumed.restore_session(saved) -``` +from ejagent.contracts import SystemMessage +from ejagent.harness import AgentHarness +from ejagent.providers import ModelConfig, OpenAIModelPort +from ejagent.tools import FunctionToolExecutor -Each JSONL record carries a monotonic `revision`, immutable `record_id`, -`parent_id`, and `branch_id`. File order defines the global revision while -parent links define the logical tree. `SessionRecorder` appends compact mutations such as -`run_started`, `message_appended`, `compaction_applied`, and `run_finished`; -explicit `save()` appends a full Checkpoint for imports and exports. - -Branches retain their source history without copying or rewriting records: - -```python -forked = await storage.fork("project-42", branch_id="experiment") -rolled_back = await storage.rollback( - "project-42", - to_record_id="a-completed-ancestor-record", - branch_id="rollback-before-change", -) -retry = await storage.prepare_retry( - "project-42", - run_id="run-to-repeat", - branch_id="retry-run", +harness = AgentHarness( + agent_id="assistant", + model=OpenAIModelPort(ModelConfig.from_env()), + tools=FunctionToolExecutor(), + initial_messages=(SystemMessage("Answer precisely."),), ) -``` - -`fork()` creates a general branch at a completed projection. `rollback()` -requires the target to be an ancestor of the source head. `prepare_retry()` -branches immediately before a Run and returns its original task; it never -executes that task automatically because Tool calls may have external side -effects. Use `checkout()`, `head()`, and `list_branches()` to inspect the tree. -To continue a branch, restore the checkout and give `SessionRecorder` the same -`branch_id`. - -Session IDs are mapped to hashed filenames. Each complete line is encoded -before one append write and followed by `fsync`; an incomplete final line from -an interrupted write is ignored and repaired before the next append. Invalid -JSON in a completed line and unsupported journal schema versions raise -`SessionSerializationError` instead of looking like a missing Session. - -`JsonlSessionStorage` coordinates every read-validate-append transaction with a -stable `.jsonl.lock` sidecar. The lock combines process-local coordination with -a POSIX advisory file lock, so separate storage instances and Python processes -cannot allocate the same revision or silently replace a branch head. Conditional -appends still report stale heads as `SessionConflictError`; lock acquisition -timeouts report `SessionLockTimeoutError`. Configure the deadline with -`JsonlSessionStorage(root, lock_timeout=...)`, or pass `None` to wait until the -operation is cancelled. Lock waits run outside the event loop. - -`restore_session()` verifies Agent identity and rejects unfinished Runs. Core -does not replay an interrupted Tool call because it may already have produced -an external side effect. The file-backed lock contract targets local POSIX -filesystems; network filesystem deployments must verify their advisory-lock -semantics or provide another `SessionTreeStorage` backend. - -## Runtime policy - -Tool availability and completion policy are independent: - -```python -from ejagent import RuntimePolicy -policy = RuntimePolicy( - max_steps=20, - max_no_tool_responses=3, - max_repeated_tool_calls=3, - max_run_tokens=None, - require_explicit_finish=False, - parallel_tool_calls=False, - max_parallel_tool_calls=None, -) +async with harness: + first = await harness.run("Remember that my project is EJAgent.") + second = await harness.run("What is my project?") + print(first.result.output, second.result.output) ``` -`max_run_tokens` is an optional cumulative model-request budget. It is checked -between turns: the current response and its requested tools settle first, then -the guard prevents another Provider request with -`StopReason.TOKEN_BUDGET_EXCEEDED`. If another request is needed but Usage was -not reported, the run stops with `StopReason.USAGE_UNAVAILABLE` instead of -treating unknown Usage as zero. +The Harness starts resources transactionally, serializes Runs, commits only +accepted outcomes, and shuts resources down in reverse order. The Kernel never +mutates committed state directly. -By default, an agent may call tools and later complete with ordinary text. A -derived autonomous agent can require a completion tool: +## Architecture -```python -policy = RuntimePolicy(require_explicit_finish=True) +```text +AgentHarness + ├─ Conversation snapshot and revision + ├─ lifecycle, cancellation, steering, follow-ups + ├─ SessionStore compare-and-commit + └─ RuntimeKernel + ├─ ContextPipeline → disposable ContextView + ├─ ModelPort → normalized stream + └─ ToolExecutor → normalized Tool result ``` -That agent must register one of its own tools that returns -`ToolControl.COMPLETE`. +- `Conversation` contains typed messages usable by future Runs. +- `RunAudit` records completed, failed, cancelled, and rejected attempts. +- `ContextView` may contain summaries, Skills, or steering without rewriting + Conversation. -## Custom tools - -Tools are grouped into handlers. `MethodToolHandler` maps a tool named `add` to -an async `do_add()` method: +## Function Tools ```python -from collections.abc import Mapping -from typing import Any - -from ejagent import ( - CancellationToken, - MethodToolHandler, - StepOutcome, - ToolDefinition, - ToolEffect, +from ejagent.contracts import ( + CancellationToken, ToolCall, ToolDefinition, + ToolExecutionResult, ToolSemantics, ) +from ejagent.tools import FunctionTool, FunctionToolExecutor -ADD_TOOL = ToolDefinition( +definition = ToolDefinition( name="add", description="Add two numbers.", - parameters={ - "type": "object", - "properties": { - "left": {"type": "number"}, - "right": {"type": "number"}, - }, - "required": ["left", "right"], - "additionalProperties": False, - }, - effect=ToolEffect.READ_ONLY, - strict=True, -) - - -class MathHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((ADD_TOOL,)) - - async def do_add( - self, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - return StepOutcome( - {"value": arguments["left"] + arguments["right"]} - ) -``` - -Register it explicitly: - -```python -agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="calculator", - handlers=[MathHandler()], + input_schema={"type": "object"}, + semantics=ToolSemantics.read_only(), ) -``` - -Duplicate tool names fail during startup instead of being silently -overwritten. `ToolDefinition` is the Core source of truth for OpenAI function -name, description, Parameters, optional `strict`, and execution Effect. -`to_openai_tool()` returns the provider request shape without exposing Core-only -Effect metadata. - -Existing OpenAI function-calling dictionaries remain accepted and are -normalized once when the Handler is created. The compatibility form is not -deprecated and does not require an immediate migration: - -```python -MethodToolHandler((OPENAI_TOOL_DICTIONARY,)) -``` - -### Tool progress - -Long-running tools can optionally accept a scoped `progress` reporter. Existing -`do_*` methods that do not declare this keyword remain compatible: - -```python -from ejagent import ToolProgressReporter, ToolProgressUpdate +async def add(call: ToolCall, cancellation: CancellationToken): + cancellation.raise_if_cancelled() + return ToolExecutionResult({ + "value": call.arguments["left"] + call.arguments["right"] + }) -async def do_index( - self, - arguments, - *, - cancellation, - progress: ToolProgressReporter | None = None, -) -> StepOutcome: - if progress is not None: - await progress.report( - ToolProgressUpdate( - "indexing files", - {"completed": 12, "total": 40}, - ) - ) - return StepOutcome({"indexed": 40}) +tools = FunctionToolExecutor((FunctionTool(definition, add),)) ``` -Each accepted update becomes a `ToolProgressed` event correlated with the -current run, turn, and tool call. Updates are ordered, stop after cancellation, -and are ignored after `ToolCompleted`. They never change `StepOutcome` or -`ToolControl`, and are not persisted to Agent state or Session. +Use `CompositeToolExecutor` to combine independent executors. Duplicate Tool +names fail at the composition boundary. -### Tool control signals +## MCP -Tool payload and runtime control are separate: +MCP is optional: -```python -from ejagent import StepOutcome, ToolControl - -StepOutcome(data) # continue the provider-tool loop -StepOutcome(data, control=ToolControl.COMPLETE) -StepOutcome(data, control=ToolControl.REJECT) -StepOutcome(data, control=ToolControl.CANCEL) -``` - -This lets the runtime distinguish successful completion, policy rejection, -and tool-requested cancellation. `ToolControl.CANCEL` is a tool's business -decision; external `agent.abort()` uses the separate run cancellation -protocol. - -## Tool middleware - -`ToolMiddleware` is the single interception chain around tool execution: - -```python -from ejagent import ToolMiddleware - - -class AuditMiddleware(ToolMiddleware): - async def __call__(self, context, call_next): - print("before", context.tool_name) - result = await call_next(context) - print("after", context.tool_name) - return result +```bash +uv sync --extra mcp ``` -The optional `ToolPolicyMiddleware` is one concrete middleware built on that -same chain. It does not add a second policy path or a special `BaseAgent` -parameter: - ```python -from ejagent import ( - RuleBasedToolPolicy, - ToolApprovalDecision, - ToolEffect, - ToolPolicyAction, - ToolPolicyMiddleware, - ToolPolicyRule, -) - - -class ConsoleApprover: - async def approve(self, request): - # A real application can bridge this request to its UI or RPC layer. - return ToolApprovalDecision(approved=False, reason="operator denied") +from ejagent.tools import McpToolExecutor - -policy = RuleBasedToolPolicy( - ( - ToolPolicyRule( - rule_id="approve-side-effects", - action=ToolPolicyAction.REQUIRE_APPROVAL, - effects=frozenset({ToolEffect.SIDE_EFFECTING}), - max_calls_per_run=5, - reason="this tool can change external state", - ), - ToolPolicyRule( - rule_id="allow-reads", - action=ToolPolicyAction.ALLOW, - effects=frozenset({ToolEffect.READ_ONLY}), - max_calls_per_run=20, - ), - ), - default_action=ToolPolicyAction.DENY, -) - -agent = BaseAgent( - model, - agent_id="policy-agent", - handlers=[handler], - middlewares=[ - ToolPolicyMiddleware(policy, approver=ConsoleApprover()), - AuditMiddleware(), - ], -) +tools = McpToolExecutor("examples/mcp_config.json") ``` -Rules use ordered, first-match semantics and may select exact tool names, -`ToolEffect`, and a synchronous or asynchronous `when(context)` predicate. -`max_calls_per_run` reserves attempts atomically, including parallel read-only -calls, and resets at each Agent Run start. Approval is fail-closed: a missing -approver, an exception, an invalid decision, or a denied request returns -`ToolControl.REJECT` without invoking the handler. Rejection payloads include -the tool, safe reason, and matching `rule_id`, but never echo arguments. +MCP metadata is normalized to Core `ToolDefinition` values after startup. Tool +execution uses the same cancellation and result contract as local functions. -Core supplies the policy/approval protocol, not an approval UI or -shell/filesystem-specific risk rules. Those remain application concerns. - -`ToolSchemaValidationMiddleware` can be placed before policy so structurally -invalid model arguments never consume policy limits, request approval, or -reach a handler: +## Skills and Context ```python -from ejagent import ( - ToolPolicyMiddleware, - ToolSchemaValidationMiddleware, -) +from ejagent.context import SkillsContextPipeline -agent = BaseAgent( - model, - agent_id="validated-agent", - handlers=[handler], - middlewares=[ - ToolSchemaValidationMiddleware(max_errors=8), - ToolPolicyMiddleware(policy, approver=approver), - AuditMiddleware(), - ], -) +context = SkillsContextPipeline("examples/skills") ``` -The middleware compiles each canonical `ToolDefinition.parameters` schema -during Agent startup. An invalid registered schema raises -`ToolSchemaConfigurationError` and rolls back startup. A model argument -failure instead becomes a normal, non-terminal Tool Result: - -```json -{ - "status": "error", - "tool": "transfer", - "code": "invalid_tool_arguments", - "errors": [ - { - "path": "/amount", - "keyword": "type", - "message": "value does not match the required type" - } - ] -} -``` +The pipeline discovers child directories containing `SKILL.md`. It projects a +compact index on every model request and full instructions when the latest user +task names `$skill_name` or `skill:skill_name`. Skill text remains transient. -Error paths use JSON Pointer. Messages describe the failed rule without -echoing argument values, and `max_errors` bounds the payload. Tools without a -parameters schema pass through unchanged. Validation results use -`ToolControl.CONTINUE`, so the model can correct its call; permission denial -remains the distinct terminal `ToolControl.REJECT` path. +For long histories, wrap a `ContextCompactor` with +`DerivedCompactionPipeline`. Summaries are derived views and never overwrite +the committed Conversation. -## MCP tools +## Persistence and Recovery -MCP uses the same handler contract: +Use `MemorySessionStore` for process-local state or `JsonlSessionStore` for an +append-only durable journal: ```python -from ejagent import ( - BaseAgent, - McpToolHandler, - ModelConfig, - OpenAIModelAdapter, -) +from ejagent.storage import JsonlSessionStore -agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="browser", - handlers=[McpToolHandler("examples/mcp_config.json")], -) +store = JsonlSessionStore(".ejagent-sessions") ``` -An MCP-enabled agent can execute MCP tools and then complete with plain text. -It does not need a separate finish tool unless its `RuntimePolicy` explicitly -requires one. - -## Skills +Durable commits use revision and message compare-and-swap, idempotent Run IDs, +cross-process locking, `fsync`, and partial-tail recovery. A Store failure +cannot advance Harness state. -Skills are prompt and resource extensions independent of handler tools: +To import an old JSONL Session once: ```python -from pathlib import Path - -from ejagent import BaseAgent, ModelConfig, OpenAIModelAdapter - -agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="skilled-agent", - skills_dir=Path("examples/skills"), +store = JsonlSessionStore( + ".ejagent-sessions", + legacy_session_id="old-session-id", ) ``` -`SkillManager` discovers child folders containing `SKILL.md` and injects compact -metadata containing each skill's name, description, and file location. Users -can explicitly select a skill with `$skill_name` or `skill:skill_name`, which -injects its full instructions into the current context. The core does not -register a special skill tool; a derived agent with a file-reading tool can use -the advertised location for progressive loading. - -```text -examples/skills/ - release_notes/ - SKILL.md - template.md - examples/ - sample.md -``` - -## Core boundary - -EJAgent Core owns mechanisms: - -```text -Orchestration + State + Context + Runtime Policy + Run Result -+ Model Adapter + Tool Protocol + Middleware + MCP + Skills -+ Lifecycle Events + Session Tree + Runtime Cancellation -+ Provider Streaming + Tool Progress + Usage Accounting + Run Budget -+ Context Pressure + Compaction Preparation -+ Model Compactor + Summary Entry + Durable Session Journal -+ Canonical Tool Definition + OpenAI Schema Compatibility View -``` - -Derived agents own concrete capabilities and policies: - -```text -Shell + Filesystem + Git + Workspace + Approval UI -+ Sandbox + Completion Tool + Product Interface -``` +Migration reads original entries rather than compacted projections. Unsupported +or unfinished legacy data raises `SessionMigrationError` with remediation. -See [the Pi Harness comparison](docs/pi-harness-gap-analysis.md) for the -architecture analysis and future roadmap. +## Runtime Control -## Examples +- `harness.cancel(reason)` cooperatively cancels the active Run. +- `harness.steer(content)` admits transient input for the next model safe point. +- `harness.follow_up(task)` queues an independent FIFO Run and returns a handle. +- `harness.continue_run()` starts a Run without appending a new user task. -```bash -# Provider-backed examples -uv run python examples/01_stateful_chat.py -uv run python examples/02_custom_tool.py -uv run python examples/04_mcp_tools.py -uv run python examples/06_skill.py +Arbitrary mid-Run pause/resume and multi-agent management are intentionally out +of scope. -# Harness examples using the configured real provider -uv run python examples/07_event_observers.py -uv run python examples/08_session_resume.py -uv run python examples/09_runtime_control.py -uv run python examples/10_composed_harness.py -uv run python examples/11_streaming_events.py -uv run python examples/12_tool_progress.py -uv run python examples/13_usage_budget.py -uv run python examples/14_context_pressure.py -uv run python examples/15_explicit_compaction.py -uv run python examples/16_durable_session.py record -uv run python examples/16_durable_session.py resume -``` +## Extension Contracts -See [the examples guide](examples/README.md) for the capability demonstrated by -each file. +Implement narrow protocols from `ejagent.contracts`: -## Tests +- `ModelPort` for another Provider protocol. +- `ToolExecutor` for another Tool backend. +- `ContextPipeline` or `ContextCompactor` for context policy. +- `SessionStore` for another durable backend. +- `RunObserver` for post-decision observation. -```bash -uv run python -m unittest discover -s tests -p 'test*.py' -q -``` +Expected operational failures use typed error contracts. Invalid configuration +and protocol violations raise exceptions. -Run the complete local quality gate before submitting a change: +## Development ```bash -uv sync --locked --all-extras --group dev -uv run ruff check src tests examples -uv run ruff format --check src tests examples +uv run ruff check src tests examples benchmarks +uv run ruff format --check src tests examples benchmarks uv run mypy +uv run python -m unittest discover -s tests -p 'test*.py' -q uv build ``` -## Release - -`.github/workflows/release.yml` builds one verified wheel/sdist pair, attaches -both files to a generated GitHub Release, and publishes the same distributions -to PyPI through Trusted Publishing. No long-lived API token is stored in -GitHub. After configuring the `pypi` environment and PyPI publisher, merge the -release commit into `main`, then push a version-matching tag: - -```text -PyPI project: ejagent-core -GitHub owner: jyh20030112 -Repository: EJAgent -Workflow: release.yml -Environment: pypi -``` - -Because `ejagent-core` is a new PyPI project name, configure its Trusted -Publisher before creating the `v0.6.0` release tag. - -Protect the GitHub `pypi` environment with required reviewers and restrict -creation of `v*` tags to maintainers. Then publish with: - -```bash -git tag v0.6.0 -git push origin v0.6.0 -``` - -The workflow rejects tags whose commit is not on `main` or whose value does not -match `project.version`, reruns the complete quality matrix, builds and smoke -tests the distributions, then publishes them to GitHub Releases and PyPI. The -two publishing jobs consume the same immutable Actions artifact; only the PyPI -job receives a short-lived OIDC identity. - -## Public API - -The package root exports: - -- Agent: `BaseAgent`, `AgentOrchestrator`, `AgentState`, `AgentStatus` -- Providers: `ModelAdapter`, `OpenAIModelAdapter`, `ModelConfig`, `AssistantMessage`, `ModelToolCall`, `ModelUsage`, `ModelStreamEvent`, `ModelTextDelta`, `ModelThinkingDelta`, `ModelResponseCompleted`, `ModelErrorKind`, `ModelProviderError`, `ContextOverflowError`, `ModelRateLimitError`, `ModelTimeoutError`, `ModelAuthenticationError` -- Runtime: `RuntimePolicy`, `AgentRunResult`, `RunUsage`, `AgentRunError`, `RunStatus`, `StopReason` -- Behavior: `BehaviorHook`, `BehaviorAction`, `BehaviorDecision`, `BehaviorHookError`, `TurnSnapshot` -- Cancellation: `CancellationToken`, `CancellationSource`, `AgentCancelledError` -- Events: `AgentEvent`, `AgentEventSink`, `CompositeAgentEventSink`, `AgentContinued`, `AssistantTextDelta`, `AssistantThinkingDelta`, `ToolProgressed`, `SteeringApplied`, `SteeringDiscarded`, `ContextPressureEvaluated`, `CompactionStarted`, `CompactionCompleted`, `CompactionFailed` -- Control: `ControlInputKind`, `ControlStatus`, `ControlInput`, `ControlReceipt`, `ContinueRejectedReason`, `ContinueRejectedError`, `FollowUpFailurePolicy`, `FollowUpDiscardReason`, `FollowUpHandle`, `FollowUpError`, `FollowUpRejectedError`, `FollowUpDiscardedError` -- Session: `AgentSession`, `SessionRun`, `SessionRunIntent`, `SessionRecorder`, `SessionStorage`, `SessionJournalStorage`, `SessionTreeStorage`, `MemorySessionStorage`, `JsonlSessionStorage`, `SessionCompaction`, `SessionRecord`, `SessionRecordDraft`, `SessionRecordKind`, `SessionBranchIntent`, `SessionBranch`, `SessionCheckout`, `SessionRetry`, `DEFAULT_SESSION_BRANCH`, `SESSION_SCHEMA_VERSION`, `SESSION_JOURNAL_SCHEMA_VERSION`, `session_to_dict`, `session_from_dict`, `SessionError`, `SessionSerializationError`, `SessionStorageError`, `SessionConflictError`, `SessionLockTimeoutError` -- Context: `AgentContextBuilder`, `ContextBuildResult`, `ContextBudget`, `ContextUsageEstimate`, `CompactionPolicy`, `AutoCompactionPolicy`, `CompactionDecision`, `CompactionPreparation`, `MessageTokenEstimator`, `estimate_context_usage`, `prepare_compaction` -- Compaction: `CompactionRuntime`, `Compactor`, `ModelCompactor`, `CompactionContextBuilder`, `CompactorOutput`, `CompactionRequest`, `CompactionResult`, `CompactionStatus`, `CompactionTrigger`, `SummaryEntry` -- Tools: `ToolDefinition`, `ToolDefinitionError`, `ToolEffect`, `StepOutcome`, `ToolControl`, `ToolProgressUpdate`, `ToolProgressReporter`, `BaseHandler`, `MethodToolHandler`, `McpToolHandler` -- Middleware: `Middleware`, `ToolMiddleware`, `ToolCallContext`, `ToolNext`, `ToolExecutionPolicy`, `RuleBasedToolPolicy`, `ToolPolicyRule`, `ToolPolicyPredicate`, `ToolPolicyAction`, `ToolPolicyDecision`, `ToolPolicyMiddleware`, `ToolApprover`, `ToolApprovalRequest`, `ToolApprovalDecision`, `ToolSchemaValidationMiddleware`, `ToolSchemaConfigurationError` -- Extensions: `McpServerManager`, `SkillManager` - -## License - -MIT +See [the design document](docs/runtime-kernel-harness-design.md) and +[examples](examples/README.md) for the normative boundaries and runnable usage. diff --git a/README_zh-CN.md b/README_zh-CN.md index ec69dae..3509ff2 100644 --- a/README_zh-CN.md +++ b/README_zh-CN.md @@ -2,860 +2,189 @@ [English](README.md) | [简体中文](README_zh-CN.md) -EJAgent Core 是一个 OpenAI-first Agent Harness Core,围绕八个核心维度构建: -执行内核、运行控制、可组合扩展、安全边界、上下文管理、持久历史、分支演化与 -可观测性。Core 提供状态、编排、上下文、工具调度、Middleware、MCP 和 Skill 等 -运行机制;Shell、文件编辑、Git、审批界面和完成工具由具体派生 Agent 自行实现。 +EJAgent Core 是面向单个逻辑 Agent 的小型可扩展运行时。它将一次模型—工具执行 +(`RuntimeKernel`)与持久状态、生命周期和控制(`AgentHarness`)分离,并明确区分 +Conversation、Run Audit 与一次性的模型 Context。 需要 Python 3.12 或更高版本。 -## 核心能力 - -- 有状态的 `BaseAgent`,支持持久对话历史和 `reset()` -- Provider 无关的 `ModelAdapter` 边界,以及 OpenAI-compatible 适配器 -- 公开的 `AgentOrchestrator`,负责模型—工具运行循环 -- 结构化的 `AgentRunResult`、`RunStatus` 和 `StopReason` -- 显式的 `RuntimePolicy`,控制循环和完成策略 -- `AgentContextBuilder`,构造不修改历史的每轮上下文 -- OpenAI-first 的强类型 `ToolDefinition`,并兼容旧工具字典 -- 可组合的 `BaseHandler` 和 `MethodToolHandler` 工具协议 -- `ToolRuntime` 生命周期、路由、Middleware 和重复调用保护 -- 通用 `ToolMiddleware` 拦截机制 -- 可选的 fail-closed 工具执行策略 Middleware、审批协议和 Run 级原子调用限额 -- 在 Policy 和 Handler 前执行的可选本地 JSON Schema 参数校验 -- 类型化生命周期事件、Text / Thinking Streaming 和 Tool Progress -- Run Cancellation、`abort()`、Steering Safe Point 与 `wait_for_idle()` 终态屏障 -- Follow-up Run Chain、Continue Safe Point 和 `after_turn` Behavior Hooks -- 显式 Side-effect 属性与可选的只读 Tool Call 并行调度 -- Provider-neutral Token Usage 与单次 Run 预算保护 -- 上下文压力估算、独立窗口预算和非变异压缩准备 -- 通过可插拔 `Compactor`、标准 `SummaryEntry` 和 Session 快照提供显式可取消压缩 -- 可选的阈值自动压缩,以及 Provider 上下文溢出后的单次安全恢复 -- 版本化 Session 序列化、追加式 JSONL Journal 和显式跨进程恢复 -- 通过 `McpToolHandler` 提供可选 MCP 集成 -- 通过 `SkillManager` 发现本地 Skill、投影 metadata 并显式激活上下文 - -Core 刻意不再内置 Bash、Git、文件系统、审批 UI 或 Finish 工具。这些能力属于 -CodeAgent 等派生 Agent。 - -本次 Core 边界调整移除了原有的 `BashHandler`、`GitDiffHandler`、 -`FinishHandler`、`HumanApproval` 和 `BashApprovalMiddleware` 公共导出;需要 -这些能力的派生 Agent 应自行提供对应实现。 - ## 安装 ```bash pip install ejagent-core -# 或在源码仓库中执行:uv sync -``` - -MCP 支持是可选能力。只有使用 MCP 的 Agent 才需要安装额外依赖: - -```bash -uv sync --extra mcp -# 或:pip install "ejagent-core[mcp]" -``` - -0.6.0 将 PyPI Distribution 改名为 `ejagent-core`,Python Import Package 改为 -`ejagent`: - -```python -from ejagent import BaseAgent +# 源码仓库 +uv sync --locked --all-extras --group dev ``` -旧的 `simagentplg` Import 不再作为兼容别名保留。已有 JSONL Session 仍可读取; -Restore 时会继续从 Provider 请求中剥离旧版内部消息元数据。 +Anthropic Provider 适配器按需安装:`pip install 'ejagent-core[anthropic]'`。 -## 配置 - -复制 `.env.example` 为 `.env` 并填写模型配置: +在 `.env` 中配置 OpenAI-compatible Endpoint: ```env MODEL_API_KEY=sk-xxxxxxxx -MODEL_URL=https://api.deepseek.com -CHAT_MODEL=deepseek-v4-flash +MODEL_URL=https://api.example.com/v1 +CHAT_MODEL=your-model LLM_TIMEOUT=60 LLM_TEMPERATURE=0.7 LLM_INCLUDE_USAGE=true ``` -`ModelConfig` 属于 `OpenAIModelAdapter`,不再属于 `BaseAgent`。也可以直接构造: - -```python -from ejagent import ModelConfig - -config = ModelConfig( - model="deepseek-v4-flash", - api_key="sk-xxxxxxxx", - base_url="https://api.deepseek.com", -) -``` - -接入其他模型 Provider 时,只需实现 `ModelAdapter.complete()`。适配器负责 Provider -Client 的创建、响应归一化以及可选的启动/关闭资源;`BaseAgent` 只消费归一化后的 -`AssistantMessage` 协议。 - -## 普通 Agent - -多次调用之间会保留对话历史: - -```python -from ejagent import BaseAgent, ModelConfig, OpenAIModelAdapter - -agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="tutor", - system_prompt="你是一名回答简洁的 Python 导师。", -) - -first = await agent.runtime(task="请记住我更喜欢 Python。") -second = await agent.runtime(task="我更喜欢哪种编程语言?") - -agent.reset() -await agent.shutdown() -``` - -同一个 Agent 的调用会串行执行,以保护对话状态。 - -## 结构化运行结果 - -`run()` 暴露 Core 的运行结果协议: - -```python -result = await agent.run(task="解释这个仓库的架构。") - -print(result.status) -print(result.stop_reason) -print(result.turns) -print(result.output) -print(result.usage.total_tokens) -print(result.usage.complete) -``` - -`runtime()` 继续作为兼容接口。任务完成时返回 `result.output`;运行失败、被拒绝或 -取消时抛出 `AgentRunError`。 - -`ModelResponseCompleted` 可以携带标准化 `ModelUsage`。Usage 会保存在 Agent 内部消息 -和 Session 中,但 `AgentContextBuilder` 会在构造 `llm_messages` 时移除,不会发送给 -Provider。`AgentRunResult.usage` 聚合一次 Run 的所有模型请求;`complete=False` 表示 -至少一次请求没有报告 Usage,不能把它当成零消耗。 - -### 取消正在执行的 Run - -`abort()` 会取消当前模型请求、Compactor、Tool、Middleware、Behavior Hook 或 MCP 调用, -而不等待 Agent 的串行操作锁: - -```python -run = asyncio.create_task(agent.run(task="执行一个较长任务")) -agent.abort("用户停止") -result = await run -``` - -取消返回结构化 `AgentRunResult`。`wait_for_idle()` 会一直等待到 `AgentFinished` 和全部同步 -Event Sink 完成,之后同一个 Agent 可以安全复用。 - -### Steering 当前 Run - -`steer()` 可以为正在执行的 Run 提交有界 FIFO 控制输入: - -```python -receipt = agent.steer("优先检查配置文件,不要修改代码") -print(receipt.status) -``` - -Steering 不会中断已经开始的模型响应或 Tool Call。Core 会在下一次 Provider Context 最终 -确定前的安全点按顺序应用输入,并发布 `SteeringApplied`;Run 提前结束时,未消费输入会 -发布 `SteeringDiscarded`。回执明确区分已接受、Agent 空闲、队列已满和 Run 正在收尾。 - -### Follow-up Run Chain - -`follow_up()` 把独立任务加入当前 Run Chain,并返回可等待的 Handle: - -```python -handle = agent.follow_up("根据刚才的结果生成测试计划") -result = await handle.wait() -``` - -每个 Follow-up 都有独立的 `run_id`、事件序列和 Session Run,并等待前一个 Run 的 -`AgentFinished` 及 Event Sink 完成。Queue 是进程内有界 FIFO,不伪装成持久任务系统; -失败后的丢弃或继续由 `FollowUpFailurePolicy` 显式控制。 - -### Continue 已有历史 - -`continue_run()` 从已提交历史启动新的 Run,但不会追加新的 User Message: - -```python -if agent.can_continue: - result = await agent.continue_run() -else: - print(agent.continue_rejection_reason) -``` - -Continue 分配新的 `run_id`、发布 `AgentContinued`,并通过 `SessionRunIntent.CONTINUE` -持久化。只有显式安全的终态和完整 Tool Result 历史可以 Continue;其他情况会抛出带类型化 -原因的 `ContinueRejectedError`。 - -### Behavior Hooks - -`BehaviorHook.after_turn()` 决定一个非终态完整 Turn 是否允许进入下一次 Provider 请求: - -```python -from ejagent import BehaviorDecision - - -class StopAfterFirstToolTurn: - async def after_turn(self, snapshot, *, cancellation): - if snapshot.turn >= 1: - return BehaviorDecision.stop("在完整 Turn 边界暂停") - return None - - -agent = BaseAgent( - model, - agent_id="controlled-agent", - behavior_hooks=[StopAfterFirstToolTurn()], -) -``` - -决策点位于所有 Tool Result 和 `TurnCompleted` 提交之后、下一次 `TurnStarted` 之前。Hook -接收隔离的 `TurnSnapshot`,不能修改 `AgentState`;多个 Hook 按声明顺序执行,首个 STOP -短路后续 Hook,并以 `StopReason.BEHAVIOR_STOP` 完成 Run。该终态可以显式 Continue。 - -### 并行只读工具 - -并行 Tool Call 需要工具和 Agent 两侧同时显式开启:Handler 必须把工具声明为 -`ToolEffect.READ_ONLY`,Runtime Policy 也必须允许并行: - -```python -from ejagent import ( - MethodToolHandler, - RuntimePolicy, - ToolDefinition, - ToolEffect, -) - -LOOKUP_TOOL = ToolDefinition( - name="lookup", - description="查询一个值。", - parameters={ - "type": "object", - "properties": {"key": {"type": "string"}}, - "required": ["key"], - }, - effect=ToolEffect.READ_ONLY, -) - - -class LookupHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((LOOKUP_TOOL,)) - - -agent = BaseAgent( - model, - agent_id="parallel-lookups", - handlers=[LookupHandler()], - runtime_policy=RuntimePolicy( - parallel_tool_calls=True, - max_parallel_tool_calls=4, - ), -) -``` - -调度器只并行连续的只读调用。未标注、未知和 `SIDE_EFFECTING` 工具保持顺序执行,并在只读 -批次之间形成屏障;默认策略仍是完全串行,因此升级不会改变现有 Handler 行为。 - -并行批次的 `ToolStarted` 按源顺序发布,Progress 可以按 `tool_call_id` 交错。Handler 可以 -并发完成,但 `ToolCompleted`、Agent History 和 Session Tool Message 始终按 Assistant -Tool Call 的原始顺序提交。外部取消会补齐活动和待执行调用;出现终态 Tool Control 时, -活动 peer 先安全收尾,再由源顺序中的首个终态结果停止 Run。`READ_ONLY` 是 Handler 及其 -Middleware 作出的并发安全契约,Core 不会自动推断副作用。 - -### 流式模型响应 - -`BaseAgent.run()` 仍只返回最终 `AgentRunResult`;临时文本和推理片段通过 -`AssistantTextDelta`、`AssistantThinkingDelta` 事件观察。Provider Adapter 负责组装完整 -Tool Call,只有完整 `AssistantMessage` 才会进入 Agent State。Delta 不写入 Session, -只实现 `complete()` 的现有 Adapter 仍可通过默认 `stream()` 回退工作。 - -## 上下文压力与压缩准备 - -Context Window 容量和累计 Run 消耗是两个独立概念。可以为 Agent 配置可选的 -`CompactionPolicy`,在每次模型请求前评估完整 Provider 上下文: - -```python -from ejagent import CompactionPolicy, ContextBudget - -context_policy = CompactionPolicy( - ContextBudget( - context_window=128_000, - reserve_tokens=16_000, - keep_recent_tokens=20_000, - ) -) - -agent = BaseAgent( - model, - agent_id="context-aware", - compaction_policy=context_policy, -) -``` - -评估会组合最近一次 Assistant `ModelUsage`、它之后新增的消息,以及包含当前 Tool Schema -的 UTF-8 感知保守估算。配置策略后,每轮都会发布 `ContextPressureEvaluated`;达到阈值 -时,事件中的 `CompactionPreparation` 会分离受保护消息、待摘要的旧完整 -User/Assistant/Tool Turn,以及需要原文保留的最近 Turn。Tool Call 和对应 Tool Result -不会被切开。 - -只配置 `CompactionPolicy` 时,压力评估仍是只读观察。应用也可以直接调用 -`estimate_context_usage()` 和 `prepare_compaction()`,并通过 -`MessageTokenEstimator` 替换默认估算器。 - -## 自动压缩与 Overflow 恢复 - -自动行为默认关闭;启用时复用同一个 `CompactionPolicy` 和 `Compactor`: - -```python -from ejagent import AutoCompactionPolicy - -agent = BaseAgent( - model, - agent_id="context-aware", - compaction_policy=context_policy, - compactor=my_compactor, - auto_compaction_policy=AutoCompactionPolicy(), -) -``` - -达到压力阈值后,Core 会在同一次 Agent Run 内压缩旧完整 Turn、重建上下文,再请求模型。 -Provider Adapter 抛出 `ContextOverflowError` 时,Core 最多执行一次“压缩—重建—重试”。 -第二次溢出返回 `StopReason.CONTEXT_OVERFLOW`;Compactor 失败返回 -`StopReason.COMPACTION_FAILED`。一旦 Text 或 Thinking Delta 已对外发布,Core 不会重试, -从而避免重复的流式输出。 - -`AutoCompactionPolicy(compact_on_pressure=False)` 可以只保留 Overflow 恢复;省略该策略或 -设置 `enabled=False` 会关闭全部自动行为。Provider Adapter 通过 `ModelProviderError` 和 -`ModelErrorKind` 统一区分上下文溢出、限流、超时、认证及普通 Provider 错误。 - -## 显式上下文压缩 - -派生 Agent 通过可取消的 `Compactor` 协议提供摘要行为,然后显式调用: - -```python -agent = BaseAgent( - model, - agent_id="context-aware", - compaction_policy=context_policy, - compactor=my_compactor, -) - -compaction = await agent.compact() -print(compaction.status) -print(compaction.summary) -``` - -`ModelCompactor` 可以把借用的 `ModelAdapter` 接入该协议,同时由应用继续拥有摘要 Prompt: - -```python -compactor = ModelCompactor( - summary_model, - context_builder=build_summary_context, - source="summary-model:v1", -) -``` - -注入的 Builder 接收 `CompactionRequest`,返回完整 `ContextBuildResult`。调用方负责借用模型 -的生命周期,因此 Core 不会静默创建另一个 Provider Client,也不会替应用选择 Prompt。 - -Core 将 `CompactionRequest` 交给 Compactor,由 Core 在 `SummaryEntry` 中写入可信的范围和 -Token metadata,最后原子替换成“受保护消息 + Summary + 最近 Turn”。失败或取消返回 -结构化 `CompactionResult`,历史保持不变。重复压缩时,旧 Summary 会传给 Compactor -合并,并由新 Summary 消息替换。 - -生命周期通过 `CompactionStarted`、`CompactionCompleted` 和 `CompactionFailed` 发布。 -`abort()`、`wait_for_idle()` 同时适用于普通 Run 和压缩。`SessionRecorder` 保存紧凑恢复 -快照,同时保留原始 `SessionMessage` 审计条目。每次压缩都有独立 `operation_id` 和 -`CompactionTrigger`。Core 不会替派生 Agent 选择摘要模型或 Prompt。 - -## 持久化 Session Journal - -`SessionRecorder` 可以使用 `JsonlSessionStorage`,为每个已接受的生命周期变更追加一条 -版本化语义 Record: - -```python -from ejagent import JsonlSessionStorage, SessionRecorder - -storage = JsonlSessionStorage("./sessions") -recorder = SessionRecorder(session_id="project-42", storage=storage) -agent = BaseAgent(model, agent_id="core-agent", event_sink=recorder) -await agent.run(task="remember this decision") -``` - -另一个进程可以读取完成快照并显式恢复新的 Agent: +## 快速开始 ```python -saved = await storage.load("project-42") -if saved is not None: - resumed = BaseAgent(model, agent_id="core-agent", event_sink=recorder) - resumed.restore_session(saved) -``` +from ejagent.contracts import SystemMessage +from ejagent.harness import AgentHarness +from ejagent.providers import ModelConfig, OpenAIModelPort +from ejagent.tools import FunctionToolExecutor -每条 JSONL Record 都包含单调递增的 `revision`、不可变 `record_id`、`parent_id` 和 -`branch_id`。文件顺序定义全局 Revision,父指针定义逻辑树。`SessionRecorder` 追加 `run_started`、`message_appended`、 -`compaction_applied`、`run_finished` 等紧凑 Mutation;显式 `save()` 用完整 Checkpoint -支持导入和导出。 - -Branch 会复用源历史,不复制或改写旧 Record: - -```python -forked = await storage.fork("project-42", branch_id="experiment") -rolled_back = await storage.rollback( - "project-42", - to_record_id="a-completed-ancestor-record", - branch_id="rollback-before-change", +harness = AgentHarness( + agent_id="assistant", + model=OpenAIModelPort(ModelConfig.from_env()), + tools=FunctionToolExecutor(), + initial_messages=(SystemMessage("请准确回答。"),), ) -retry = await storage.prepare_retry( - "project-42", - run_id="run-to-repeat", - branch_id="retry-run", -) -``` - -`fork()` 在已完成的投影上创建通用 Branch;`rollback()` 要求目标是源 Head 的祖先; -`prepare_retry()` 在指定 Run 之前创建 Branch 并返回原始 Task。Core 不会自动执行重试, -因为 Tool Call 可能已经产生外部副作用。可以使用 `checkout()`、`head()` 和 -`list_branches()` 检查 Session 树。继续某个 Branch 时,需要恢复其 Checkout,并把相同 -的 `branch_id` 传给 `SessionRecorder`。 - -Session ID 会映射为哈希文件名。每一行先完整编码,再通过一次追加写入并执行 `fsync`; -中断产生的不完整尾行会在读取时忽略,并在下一次追加前修复。已经换行的损坏 JSON 或 -未知 Journal Schema 会抛出 `SessionSerializationError`,不会被误认为 Session 不存在。 - -`JsonlSessionStorage` 使用稳定的 `.jsonl.lock` Sidecar,把进程内协调和 POSIX Advisory -Lock 组合在同一个 read-validate-append 事务中。不同 Storage 实例和不同进程不会分配重复 -Revision 或静默覆盖 Branch Head;过期 Head 返回 `SessionConflictError`,锁超时返回 -`SessionLockTimeoutError`。锁等待在线程中执行,不阻塞 Event Loop。 - -`restore_session()` 会校验 Agent 身份,并拒绝包含未完成 Run 的 Session。Core 不会重放 -中断的 Tool Call,因为它可能已经产生外部副作用。文件锁契约面向本地 POSIX 文件系统; -网络文件系统需要验证其 Advisory Lock 语义,或提供其他 `SessionTreeStorage` Backend。 - -## RuntimePolicy -工具是否存在和任务是否必须显式完成已经解耦: - -```python -from ejagent import RuntimePolicy - -policy = RuntimePolicy( - max_steps=20, - max_no_tool_responses=3, - max_repeated_tool_calls=3, - max_run_tokens=None, - require_explicit_finish=False, - parallel_tool_calls=False, - max_parallel_tool_calls=None, -) +async with harness: + first = await harness.run("记住我的项目是 EJAgent。") + second = await harness.run("我的项目是什么?") + print(first.result.output, second.result.output) ``` -`parallel_tool_calls` 默认关闭;开启后也只并行显式声明为 `READ_ONLY` 的工具。 -`max_parallel_tool_calls` 可以限制每个只读批次的并发量。 - -可选的 `max_run_tokens` 在轮次边界阻止下一次模型请求。当前响应及其请求的工具会先完整 -收尾;达到预算时返回 `TOKEN_BUDGET_EXCEEDED`,需要继续但 Provider 未报告 Usage 时 -返回 `USAGE_UNAVAILABLE`。 +Harness 以事务方式启动资源、串行化 Run、只提交 Store 接受的结果,并按相反顺序 +关闭资源。Kernel 永远不会直接修改已提交状态。 -默认情况下,Agent 可以调用工具,之后用普通文本完成任务。自主型派生 Agent 可以要求 -必须调用完成工具: +## 架构 -```python -policy = RuntimePolicy(require_explicit_finish=True) +```text +AgentHarness + ├─ Conversation 快照与 Revision + ├─ 生命周期、取消、Steering、Follow-up + ├─ SessionStore Compare-and-Commit + └─ RuntimeKernel + ├─ ContextPipeline → 一次性 ContextView + ├─ ModelPort → 标准化 Stream + └─ ToolExecutor → 标准化 Tool Result ``` -此时派生 Agent 必须自行注册一个返回 `ToolControl.COMPLETE` 的工具。 - -## 自定义工具 +- `Conversation` 只包含可用于未来 Run 的类型化消息。 +- `RunAudit` 记录成功、失败、取消和拒绝的尝试。 +- `ContextView` 可以包含摘要、Skill 或 Steering,但不会改写 Conversation。 -工具通过 Handler 组织。`MethodToolHandler` 会把名为 `add` 的工具映射到异步 -`do_add()` 方法: +## Function Tool ```python -from collections.abc import Mapping -from typing import Any - -from ejagent import ( - MethodToolHandler, - StepOutcome, - ToolDefinition, - ToolEffect, +from ejagent.contracts import ( + CancellationToken, ToolCall, ToolDefinition, + ToolExecutionResult, ToolSemantics, ) +from ejagent.tools import FunctionTool, FunctionToolExecutor -ADD_TOOL = ToolDefinition( +definition = ToolDefinition( name="add", - description="计算两个数的和。", - parameters={ - "type": "object", - "properties": { - "left": {"type": "number"}, - "right": {"type": "number"}, - }, - "required": ["left", "right"], - "additionalProperties": False, - }, - effect=ToolEffect.READ_ONLY, - strict=True, + description="Add two numbers.", + input_schema={"type": "object"}, + semantics=ToolSemantics.read_only(), ) +async def add(call: ToolCall, cancellation: CancellationToken): + cancellation.raise_if_cancelled() + return ToolExecutionResult({ + "value": call.arguments["left"] + call.arguments["right"] + }) -class MathHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((ADD_TOOL,)) - - async def do_add(self, arguments: Mapping[str, Any]) -> StepOutcome: - return StepOutcome( - {"value": arguments["left"] + arguments["right"]} - ) -``` - -显式注册到 Agent: - -```python -agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="calculator", - handlers=[MathHandler()], -) -``` - -重复工具名会在启动阶段报错,不会静默覆盖。`ToolDefinition` 统一保存 OpenAI Function 的 -名称、描述、Parameters、可选 `strict` 和执行 Effect;`to_openai_tool()` 只输出 Provider -请求字段,不会泄漏 Core 专用的 Effect metadata。 - -原有 OpenAI function-calling 字典仍然兼容,并在 Handler 构造时归一化一次。兼容写法不会 -立即弃用,现有项目不需要同步迁移: - -```python -MethodToolHandler((OPENAI_TOOL_DICTIONARY,)) -``` - -### 工具执行进度 - -长时间运行的工具可以选择声明一个作用域限定的 `progress` Reporter。没有声明该参数的 -现有 `do_*` 方法仍然兼容: - -```python -from ejagent import ToolProgressReporter, ToolProgressUpdate - - -async def do_index( - self, - arguments, - *, - cancellation, - progress: ToolProgressReporter | None = None, -) -> StepOutcome: - if progress is not None: - await progress.report( - ToolProgressUpdate( - "正在建立文件索引", - {"completed": 12, "total": 40}, - ) - ) - return StepOutcome({"indexed": 40}) -``` - -每条有效更新都会生成关联当前 run、turn 和 tool call 的 `ToolProgressed` 事件。 -Progress 保持顺序,在取消或 `ToolCompleted` 后停止接收;它不会改变 `StepOutcome`、 -`ToolControl`,也不会写入 Agent State、Session 或模型上下文。 - -### 工具控制信号 - -工具结果数据与运行控制已经分离: - -```python -from ejagent import StepOutcome, ToolControl - -StepOutcome(data) # 继续模型—工具循环 -StepOutcome(data, control=ToolControl.COMPLETE) -StepOutcome(data, control=ToolControl.REJECT) -StepOutcome(data, control=ToolControl.CANCEL) +tools = FunctionToolExecutor((FunctionTool(definition, add),)) ``` -运行时由此可以区分正常完成、策略拒绝和取消。 +通过 `CompositeToolExecutor` 可以组合多个独立 Executor;重复 Tool 名称会在组合 +边界立即失败。 -## Tool Middleware +## MCP -`ToolMiddleware` 是工具执行外围唯一的拦截链: +MCP 是可选依赖: -```python -from ejagent import ToolMiddleware - - -class AuditMiddleware(ToolMiddleware): - async def __call__(self, context, call_next): - print("before", context.tool_name) - result = await call_next(context) - print("after", context.tool_name) - return result +```bash +uv sync --extra mcp ``` -可选的 `ToolPolicyMiddleware` 是建立在同一条链上的官方具体 Middleware;它不会增加第二条 -Policy 通道,也不会给 `BaseAgent` 增加特殊参数: - ```python -from ejagent import ( - RuleBasedToolPolicy, - ToolApprovalDecision, - ToolEffect, - ToolPolicyAction, - ToolPolicyMiddleware, - ToolPolicyRule, -) - +from ejagent.tools import McpToolExecutor -class ConsoleApprover: - async def approve(self, request): - # 实际应用可以把 Request 转发给自己的 UI 或 RPC 层。 - return ToolApprovalDecision(approved=False, reason="操作员拒绝") - - -policy = RuleBasedToolPolicy( - ( - ToolPolicyRule( - rule_id="approve-side-effects", - action=ToolPolicyAction.REQUIRE_APPROVAL, - effects=frozenset({ToolEffect.SIDE_EFFECTING}), - max_calls_per_run=5, - reason="该工具可能修改外部状态", - ), - ToolPolicyRule( - rule_id="allow-reads", - action=ToolPolicyAction.ALLOW, - effects=frozenset({ToolEffect.READ_ONLY}), - max_calls_per_run=20, - ), - ), - default_action=ToolPolicyAction.DENY, -) - -agent = BaseAgent( - model, - agent_id="policy-agent", - handlers=[handler], - middlewares=[ - ToolPolicyMiddleware(policy, approver=ConsoleApprover()), - AuditMiddleware(), - ], -) +tools = McpToolExecutor("examples/mcp_config.json") ``` -规则按声明顺序采用 first-match 语义,可以匹配精确工具名、`ToolEffect`,也可以提供同步或 -异步的 `when(context)` 条件。`max_calls_per_run` 会原子预留调用次数,能覆盖并行只读调用, -并在每个 Agent Run 开始时重置。审批默认 fail-closed:没有 Approver、发生异常、返回非法 -结果或明确拒绝时,都会返回 `ToolControl.REJECT`,Handler 绝不会执行。拒绝载荷包含工具名、 -安全的原因和匹配的 `rule_id`,不会回显参数。 +MCP 元数据会在启动后转换成 Core `ToolDefinition`,执行时复用本地 Tool 相同的取消 +与结果协议。 -Core 只提供策略和审批协议,不提供审批 UI,也不内置 Shell、文件系统等领域风险规则;这些 -仍由具体应用负责。 - -将 `ToolSchemaValidationMiddleware` 放在 Policy 前,可以保证结构非法的模型参数不会消耗 -Policy 限额、请求审批或到达 Handler: +## Skill 与 Context ```python -from ejagent import ( - ToolPolicyMiddleware, - ToolSchemaValidationMiddleware, -) +from ejagent.context import SkillsContextPipeline -agent = BaseAgent( - model, - agent_id="validated-agent", - handlers=[handler], - middlewares=[ - ToolSchemaValidationMiddleware(max_errors=8), - ToolPolicyMiddleware(policy, approver=approver), - AuditMiddleware(), - ], -) +context = SkillsContextPipeline("examples/skills") ``` -Middleware 会在 Agent 启动阶段编译每个 Canonical `ToolDefinition.parameters` Schema。注册的 -Schema 本身非法时,启动会抛出 `ToolSchemaConfigurationError` 并回滚;模型参数不符合 -Schema 时则返回普通、非终态 Tool Result: - -```json -{ - "status": "error", - "tool": "transfer", - "code": "invalid_tool_arguments", - "errors": [ - { - "path": "/amount", - "keyword": "type", - "message": "value does not match the required type" - } - ] -} -``` +Pipeline 会发现包含 `SKILL.md` 的子目录。每次模型请求都会获得精简索引;最新用户 +任务出现 `$skill_name` 或 `skill:skill_name` 时才投影完整说明。Skill 文本不会进入 +Conversation。 -错误路径使用 JSON Pointer;错误消息只说明失败规则,不回显参数值,`max_errors` 负责限制 -载荷大小。没有 Parameters Schema 的工具保持原行为。Validation 使用 -`ToolControl.CONTINUE`,允许模型修正调用;权限拒绝仍使用独立的终态 -`ToolControl.REJECT`。 +长历史可以通过 `DerivedCompactionPipeline` 和自定义 `ContextCompactor` 生成摘要; +摘要始终是派生视图,不会覆盖已提交 Conversation。 -## MCP 工具 +## 持久化与恢复 -MCP 使用相同的 Handler 协议: +`MemorySessionStore` 用于进程内状态,`JsonlSessionStore` 提供追加式持久 Journal: ```python -from ejagent import ( - BaseAgent, - McpToolHandler, - ModelConfig, - OpenAIModelAdapter, -) +from ejagent.storage import JsonlSessionStore -agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="browser", - handlers=[McpToolHandler("examples/mcp_config.json")], -) +store = JsonlSessionStore(".ejagent-sessions") ``` -启用 MCP 的 Agent 可以执行 MCP 工具,然后直接用普通文本完成任务。只有 -`RuntimePolicy` 明确要求时,才需要额外的完成工具。 - -## Skill +持久提交具备 Revision 和消息 Compare-and-Swap、幂等 Run ID、跨进程锁、`fsync` +和不完整尾部恢复。Store 失败不会推进 Harness 状态。 -Skill 是独立于 Handler 工具的提示词和资源扩展: +一次性导入旧 JSONL Session: ```python -from pathlib import Path - -from ejagent import BaseAgent, ModelConfig, OpenAIModelAdapter - -agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="skilled-agent", - skills_dir=Path("examples/skills"), +store = JsonlSessionStore( + ".ejagent-sessions", + legacy_session_id="old-session-id", ) ``` -`SkillManager` 会发现包含 `SKILL.md` 的子目录,并注入包含名称、描述和文件位置的 -紧凑 metadata。用户可以用 `$skill_name` 或 `skill:skill_name` 显式选择 Skill, -将其完整指令注入当前上下文。Core 不注册特殊的 Skill 工具;未来带文件读取工具的 -派生 Agent 可以根据 metadata 中的位置渐进加载 Skill。 +迁移以原始 Entry 为准,不使用旧 Compaction 投影。无法无损表示或尚未完成的数据会 +抛出带修复建议的 `SessionMigrationError`。 -```text -examples/skills/ - release_notes/ - SKILL.md - template.md - examples/ - sample.md -``` +## 运行控制 -## Core 边界 +- `harness.cancel(reason)`:协作式取消当前 Run。 +- `harness.steer(content)`:为下一个模型安全点提交临时输入。 +- `harness.follow_up(task)`:排队一个独立的 FIFO Run,并返回 Handle。 +- `harness.continue_run()`:不追加新用户任务,直接从已提交 Revision 继续。 -EJAgent Core 负责机制: +任意时刻暂停/恢复和多 Agent 管理明确不在 Core 范围内。 -```text -Orchestration + State + Context + Runtime Policy + Run Result -+ Model Adapter + Tool Protocol + Middleware + MCP + Skills -+ Lifecycle Events + Session + Streaming + Tool Progress + Usage Budget -+ Context Pressure + Compaction Preparation -+ Model Compactor + Summary Entry + Durable Session Journal + Session Tree -+ Cancellation + Steering + Follow-up + Continue + Behavior Hooks -+ Parallel Read-only Tool Scheduler + Side-effect Barrier -+ Canonical Tool Definition + OpenAI Schema Compatibility View -``` +## 扩展协议 -派生 Agent 负责具体能力与策略: +通过 `ejagent.contracts` 中的窄协议扩展: -```text -Shell + Filesystem + Git + Workspace + Approval UI -+ Sandbox + Completion Tool + Product Interface -``` - -架构分析和后续路线参见 -[Pi Harness 对照分析](docs/pi-harness-gap-analysis.md)。 - -## 示例 - -```bash -uv run python examples/01_stateful_chat.py -uv run python examples/02_custom_tool.py -uv run python examples/04_mcp_tools.py -uv run python examples/06_skill.py -uv run python examples/13_usage_budget.py -uv run python examples/14_context_pressure.py -uv run python examples/15_explicit_compaction.py -uv run python examples/16_durable_session.py record -uv run python examples/16_durable_session.py resume -``` - -## 测试 +- `ModelPort`:接入其他 Provider 协议。 +- `ToolExecutor`:接入其他 Tool Backend。 +- `ContextPipeline` / `ContextCompactor`:定义 Context 策略。 +- `SessionStore`:接入其他持久化 Backend。 +- `RunObserver`:观察 Store 决策后的 Run。 -```bash -uv run python -m unittest discover -s tests -p 'test*.py' -q -``` +预期的运行故障使用类型化错误;无效配置和协议违规使用异常。 -提交变更前运行完整的本地质量门: +## 开发 ```bash -uv sync --locked --all-extras --group dev -uv run ruff check src tests examples -uv run ruff format --check src tests examples +uv run ruff check src tests examples benchmarks +uv run ruff format --check src tests examples benchmarks uv run mypy +uv run python -m unittest discover -s tests -p 'test*.py' -q uv build ``` -## 发布 - -`.github/workflows/release.yml` 只构建一份经过验证的 wheel/sdist,将同一份产物同时 -附加到自动生成的 GitHub Release,并通过 Trusted Publishing 发布到 PyPI。GitHub 中 -不保存长期 API Token。配置 `pypi` Environment 和 PyPI Publisher 后,先把发布提交 -合并到 `main`,再推送与项目版本一致的 Tag: - -```text -PyPI 项目:ejagent-core -GitHub Owner:jyh20030112 -Repository:EJAgent -Workflow:release.yml -Environment:pypi -``` - -`ejagent-core` 是新的 PyPI 项目名,因此创建 `v0.6.0` Release Tag 前,需要先为它 -配置 Trusted Publisher。 - -建议为 GitHub `pypi` Environment 配置 Required Reviewer,并限制只有 Maintainer 可以 -创建 `v*` Tag。然后执行: - -```bash -git tag v0.6.0 -git push origin v0.6.0 -``` - -工作流会拒绝不在 `main` 上或与 `project.version` 不一致的 Tag,重新执行完整质量矩阵, -构建并 smoke test 发布产物,然后发布到 GitHub Releases 和 PyPI。两个发布 Job 使用同一份 -不可变 Actions Artifact,只有 PyPI Job 会获得短期 OIDC 身份。 - -## 公共 API - -包根目录导出: - -- Agent:`BaseAgent`、`AgentOrchestrator`、`AgentState`、`AgentStatus`、`BehaviorHook`、`BehaviorDecision`、`BehaviorAction`、`TurnSnapshot` -- Provider:`ModelAdapter`、`OpenAIModelAdapter`、`ModelConfig`、`AssistantMessage`、`ModelToolCall`、`ModelUsage`、`ModelErrorKind`、`ModelProviderError`、`ContextOverflowError`、`ModelRateLimitError`、`ModelTimeoutError`、`ModelAuthenticationError` -- Runtime:`RuntimePolicy`、`AgentRunResult`、`RunUsage`、`AgentRunError`、`RunStatus`、`StopReason` -- 控制面:`ControlInput`、`ControlReceipt`、`ControlStatus`、`FollowUpHandle`、`FollowUpFailurePolicy`、`ContinueRejectedError`、`ContinueRejectedReason` -- Session:`AgentSession`、`SessionRecorder`、`SessionStorage`、`SessionJournalStorage`、`SessionTreeStorage`、`MemorySessionStorage`、`JsonlSessionStorage`、`SessionRunIntent`、`SessionCompaction`、`SessionRecord`、`SessionRecordDraft`、`SessionRecordKind`、`SessionBranchIntent`、`SessionBranch`、`SessionCheckout`、`SessionRetry`、`DEFAULT_SESSION_BRANCH`、`SESSION_SCHEMA_VERSION`、`SESSION_JOURNAL_SCHEMA_VERSION`、`session_to_dict`、`session_from_dict`、`SessionError`、`SessionSerializationError`、`SessionStorageError`、`SessionConflictError`、`SessionLockTimeoutError` -- Context:`AgentContextBuilder`、`ContextBuildResult`、`ContextBudget`、`ContextUsageEstimate`、`CompactionPolicy`、`AutoCompactionPolicy`、`CompactionDecision`、`CompactionPreparation`、`MessageTokenEstimator` -- Compaction:`CompactionRuntime`、`Compactor`、`ModelCompactor`、`CompactionContextBuilder`、`CompactorOutput`、`CompactionRequest`、`CompactionResult`、`CompactionStatus`、`CompactionTrigger`、`SummaryEntry` -- Tool:`ToolDefinition`、`ToolDefinitionError`、`ToolEffect`、`StepOutcome`、`ToolControl`、`ToolProgressReporter`、`ToolProgressUpdate`、`BaseHandler`、`MethodToolHandler`、`McpToolHandler` -- Middleware:`Middleware`、`ToolMiddleware`、`ToolCallContext`、`ToolNext`、`ToolExecutionPolicy`、`RuleBasedToolPolicy`、`ToolPolicyRule`、`ToolPolicyPredicate`、`ToolPolicyAction`、`ToolPolicyDecision`、`ToolPolicyMiddleware`、`ToolApprover`、`ToolApprovalRequest`、`ToolApprovalDecision`、`ToolSchemaValidationMiddleware`、`ToolSchemaConfigurationError` -- Event:`AgentEvent`、`AgentEventSink`、`AgentStarted`、`AgentContinued`、`TurnStarted`、`MessageCompleted`、`ToolStarted`、`ToolProgressed`、`ToolCompleted`、`TurnCompleted`、`AgentFinished` -- 扩展:`McpServerManager`、`SkillManager` - -## License - -MIT +规范边界和可运行示例见[设计文档](docs/runtime-kernel-harness-design.md)与 +[示例指南](examples/README.md)。 diff --git a/benchmarks/session_journal.py b/benchmarks/session_journal.py index 5bb9835..f3b4c09 100644 --- a/benchmarks/session_journal.py +++ b/benchmarks/session_journal.py @@ -1,4 +1,4 @@ -"""Measure JSONL Session journal append and replay scaling.""" +"""Measure append and replay scaling for the Core JSONL SessionStore.""" from __future__ import annotations @@ -10,37 +10,66 @@ from pathlib import Path from typing import Any -from ejagent import AgentSession, JsonlSessionStorage +from ejagent.contracts import ( + AssistantMessage, + ConversationSnapshot, + RunDelta, + RunOutcome, + RunResult, + RunStatus, + SessionCommit, + StopReason, + UserMessage, +) +from ejagent.storage import JsonlSessionStore + +AGENT_ID = "benchmark-agent" def journal_size(directory: str) -> int: - return next(Path(directory).glob("*.jsonl")).stat().st_size + return next(Path(directory).glob("*.core.jsonl")).stat().st_size + + +def commit(index: int, base: ConversationSnapshot) -> SessionCommit: + run_id = f"benchmark-{index}" + return SessionCommit( + agent_id=AGENT_ID, + base=base, + outcome=RunOutcome( + result=RunResult( + run_id=run_id, + status=RunStatus.COMPLETED, + stop_reason=StopReason.TEXT_RESPONSE, + turns=1, + output="ok", + ), + delta=RunDelta( + base_revision=base.revision, + messages=(UserMessage(run_id), AssistantMessage("ok")), + ), + ), + ) async def benchmark(record_count: int) -> dict[str, Any]: with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - session = AgentSession(session_id=f"benchmark-{record_count}") - session.bind_agent("benchmark-agent") - + store = JsonlSessionStore(directory) + conversation = ConversationSnapshot() started = time.perf_counter() - for _ in range(record_count): - await storage.save(session) + for index in range(record_count): + snapshot = await store.commit(commit(index, conversation)) + conversation = snapshot.conversation build_seconds = time.perf_counter() - started started = time.perf_counter() - await JsonlSessionStorage(directory).load(session.session_id) + await JsonlSessionStore(directory).load(AGENT_ID) load_seconds = time.perf_counter() - started - byte_count = await asyncio.to_thread(journal_size, directory) return { "records": record_count, "bytes": byte_count, "build_seconds": round(build_seconds, 6), - "average_append_ms": round( - build_seconds * 1000 / record_count, - 6, - ), + "average_append_ms": round(build_seconds * 1000 / record_count, 6), "load_ms": round(load_seconds * 1000, 6), } @@ -52,12 +81,7 @@ async def main(record_counts: list[int]) -> None: if __name__ == "__main__": parser = argparse.ArgumentParser() - parser.add_argument( - "records", - type=int, - nargs="*", - default=[100, 500, 1000], - ) + parser.add_argument("records", type=int, nargs="*", default=[100, 500, 1000]) args = parser.parse_args() if any(record_count <= 0 for record_count in args.records): parser.error("record counts must be greater than zero") diff --git a/docs/pi-harness-gap-analysis.md b/docs/pi-harness-gap-analysis.md deleted file mode 100644 index 610cf3b..0000000 --- a/docs/pi-harness-gap-analysis.md +++ /dev/null @@ -1,617 +0,0 @@ -# EJAgent Core 与 Pi Agent Harness 能力对照 - -> 更新日期:2026-07-22 -> -> EJAgent Core 基线:`58d28f1` 后的 0.6.0 Rename 工作树 -> -> Pi 子模块基线:`f4e9ca74` -> -> 对照范围:`pi/packages/agent`、`pi/packages/coding-agent` 与 EJAgent Core - -## 1. 当前结论 - -EJAgent Core 已经不只是 Agent Loop。当前实现覆盖了通用 Agent Core 的主要执行机制: - -- Provider-neutral 模型边界和 OpenAI-compatible Adapter -- 模型—工具循环、结构化终止结果与 Runtime Policy -- Tool Runtime、Middleware、MCP Adapter 和 Skill Resource -- 只读生命周期事件、Text / Thinking Stream 和 Tool Progress -- Cancellation、`abort()`、`wait_for_idle()` 与终态事件屏障 -- Usage 聚合、Run Token Budget 和 Context Pressure -- 显式与自动 Summary Compaction -- Provider Context Overflow 标准化、最多一次 compact-and-retry -- Durable JSONL Session Journal -- Session Tree、Branch Head、Checkout、Fork、Rollback 和整 Run Retry -- 跨实例 / 跨进程 Session Journal 原子写入与 `SessionTreeStorage` 协议 -- 有界 FIFO Steering Queue、Model-call Safe Point 与 Session 审计 -- Follow-up Run Chain、有界 FIFO、Waitable Handle 与终态 Sink 屏障 -- Continue / Resume Safe Point、独立 Run Intent 与 Session 回放 -- `after_turn` Behavior Hooks、隔离 Turn Snapshot 与结构化策略停止 -- Parallel Read-only Tool Calls、显式 Side-effect 属性与确定性结果提交 -- OpenAI-first Canonical `ToolDefinition`、强类型校验与旧字典兼容视图 -- 基于现有 Middleware 链的 Tool Execution Policy、异步审批与并发安全 Run 限额 -- Policy 前的本地 Tool Argument JSON Schema Validation 与安全错误载荷 -- Python 3.12 / 3.13 CI、构建 smoke test 与 PyPI Trusted Publishing CD - -因此,`BaseAgent` 已经是一个可独立使用的轻量 Harness 组装根,而不再只是 -`AgentOrchestrator` 的薄包装。 - -与 Pi 的主要差距已经从“缺少持久化和自动压缩”转移到以下领域: - -1. **行为控制面**:Steering、Follow-up、Continue 与 `after_turn` Hook 基础版已完成; - `before_model` 等更广义行为注入仍未实现。 -2. **工具调度**:只读 Tool Call、Canonical Definition、Execution Policy 和本地参数校验已 - 完成基础版;工具集合刻意保持构造/启动时冻结,不向模型或 Active Run 暴露动态修改能力。 -3. **Provider 广度**:当前明确采用 OpenAI-first,仅实现 OpenAI-compatible Adapter;在出现 - 真实需求前不建设多 Provider Schema 抽象。 -4. **长 Journal 性能**:写入与读取仍会线性重建索引,需要按实际规模决定索引和归档策略。 -5. **ExecutionEnv / CodeAgent**:文件、Shell、Git、Workspace、Sandbox、Approval 和 UI/RPC - 明确留给派生层,当前仓库尚未实现该层。 - -Session Journal 并发边界、行为控制链、安全并行调度、OpenAI-first Canonical Tool -Definition、Tool Execution Policy 和参数 Schema Validation 已补齐。Dynamic Tool Set 不进入 -近期路线。 - -## 2. 当前架构 - -```text -BaseAgent - ├── Active Operation / CancellationSource - ├── Bounded Steering Queue / Control Receipt - ├── Follow-up FIFO / Run Handle / Run-chain Gate - ├── Continue Safe Point / Run Intent - ├── BehaviorHook / TurnSnapshot / BehaviorDecision - ├── ModelAdapter - │ └── OpenAIModelAdapter - ├── AgentOrchestrator - │ ├── RuntimePolicy - │ ├── AgentContextBuilder - │ ├── UsageAccumulator - │ └── AutoCompactionPolicy - ├── CompactionPolicy / ContextBudget - ├── CompactionRuntime - │ └── Compactor / ModelCompactor - ├── AgentEventEmitter - │ └── AgentEventSink / CompositeAgentEventSink - │ └── SessionRecorder(branch_id) - │ ├── MemorySessionStorage - │ └── JsonlSessionStorage - │ └── Session Tree / Branch Heads - ├── AgentState - ├── ToolRuntime - │ ├── Read-only Batch Scheduler / Side-effect Barrier - │ ├── ToolDefinition / ToolEffect / OpenAI Compatibility View - │ ├── BaseHandler / MethodToolHandler - │ ├── McpToolHandler - │ └── ToolMiddleware - │ ├── ToolSchemaValidationMiddleware - │ └── ToolPolicyMiddleware / RuleBasedToolPolicy - └── SkillManager -``` - -### 2.1 组件职责 - -| 组件 | 当前职责 | -|---|---| -| `BaseAgent` | 依赖组装、Run Chain、Steering、Follow-up、Behavior Hook、取消、空闲等待、资源生命周期、Session Restore | -| `AgentOrchestrator` | 模型—工具循环、after-Turn 决策、自动压缩、Overflow Recovery、结构化终止 | -| `AgentState` | 当前消息、Turn、任务状态和运行结果 | -| `AgentContextBuilder` | 构造 Agent 投影和 Provider-safe 请求,注入 Skill 与临时控制消息 | -| `RuntimePolicy` | 步数、无工具响应、重复调用、显式完成和 Run Token Budget | -| `CompactionPolicy` | 评估 Context Pressure,并按完整 User Turn 准备安全切分 | -| `AutoCompactionPolicy` | 控制压力触发和最多一次 Overflow Recovery | -| `CompactionRuntime` | 取消、摘要编排、原子替换和 Compaction 终态事件 | -| `ModelCompactor` | 将借用的 `ModelAdapter` 适配为 `Compactor`,Prompt 仍由应用提供 | -| `AgentEventEmitter` | 分配 `run_id` 和事件序号,按顺序发布只读事件 | -| `SessionRecorder` | 将生命周期事件转换为指定 Branch 的语义 Journal Record | -| `JsonlSessionStorage` | 原子追加 JSONL、跨进程锁、校验树、回放节点、管理 Branch 和 Head | -| `ToolRuntime` | 工具路由、Middleware、执行策略、Progress、控制信号和重复调用保护 | -| `ModelAdapter` | Provider Client 生命周期、请求和响应归一化 | -| `SkillManager` | Skill 发现、metadata、显式选择和上下文投影 | - -### 2.2 Runtime 主链路 - -```text -BaseAgent.run(task) - → create CancellationToken + run_id - → AgentState.begin_task(task) - → AgentStarted - → drain Steering at model-call safe point - → AgentContextBuilder.build() - → ContextPressureEvaluated - → optional pressure-triggered Compaction - → ModelAdapter.stream() - → Text / Thinking Delta* - → ModelResponseCompleted + ModelUsage - → MessageCompleted - → ToolRuntime.execute_tool_call()* - → ToolProgressed* / ToolCompleted - → AgentRunResult - → AgentFinished - → SessionRecorder locked append + fsync - → optional next Follow-up Run -``` - -Provider 在输出任何 Delta 前报告 Context Overflow 时,Orchestrator 可以执行一次自动压缩、 -重建 Context 并重试。已经开始输出后发生 Overflow 不会重试,以免重复可见输出。 - -## 3. 已完成的 Core 能力 - -### 3.1 执行与终止语义 - -`AgentRunResult`、`RunStatus`、`StopReason` 和 `ToolControl` 已能稳定区分: - -- 文本完成和工具显式完成 -- 工具拒绝、工具取消和外部取消 -- 空响应、步数限制、无工具响应限制和重复工具调用 -- Run Token Budget 超限与 Usage 缺失 -- Context Overflow、Compaction 失败和一般 Runtime 错误 - -工具存在与任务是否必须显式完成已经解耦。`BaseAgent.runtime()` 只作为兼容包装,核心 -终态接口是 `run() -> AgentRunResult`。 - -### 3.2 Provider 与 Tool Runtime - -Core 只依赖 `ModelAdapter`,当前 `OpenAIModelAdapter` 负责 Streaming、Tool Call 组装、 -Usage 和 Provider Error 归一化。只实现 `complete()` 的 Adapter 仍可通过基类回退工作。 - -所有工具进入统一 `ToolRuntime`: - -- Handler 生命周期和确定性路由 -- 重复工具名校验和 JSON 参数解析 -- Tool Middleware -- 标准 Tool Message 和结构化 Tool Control -- Cancellation 与 Progress -- 重复调用保护 -- MCP Tool 适配 -- OpenAI-first `ToolDefinition`:Name、Description、Parameters、`strict` 与 Effect -- 旧 OpenAI 字典在 Handler 边界一次归一化 -- 显式 `ToolEffect.READ_ONLY` / `SIDE_EFFECTING` -- Opt-in Parallel Read-only Batch 与稳定源顺序提交 -- 可选 `ToolPolicyMiddleware`、默认拒绝、异步审批和原子 Run 调用限额 -- 可选 `ToolSchemaValidationMiddleware`、启动时 Schema 编译和安全的参数错误结果 - -默认仍按顺序执行。启用 `RuntimePolicy.parallel_tool_calls` 后,只有连续且显式标注为 -`READ_ONLY` 的调用才并行;未标注、未知和 `SIDE_EFFECTING` 工具形成串行屏障。开始事件按 -源顺序发出,Progress 可以交错,完成事件、Agent History 和 Session Tool Message 按 Assistant -Tool Call 源顺序提交。并行上限由 `max_parallel_tool_calls` 控制。 - -Tool Policy 不新增第二条执行链。`ToolExecutionPolicy` 只产生 ALLOW、DENY 或 -REQUIRE_APPROVAL 决策,`ToolPolicyMiddleware` 通过现有 `ToolMiddleware` 链强制执行。 -Rule-based 基础实现支持精确工具名、`ToolEffect`、同步/异步参数条件和每 Run 调用限额; -并发只读调用使用原子预留,缺少 Approver、策略/审批异常或非法返回值一律 fail closed。 -Core 只提供审批协议,不包含具体审批 UI 或 Shell/Filesystem 风险规则。 - -Schema Validation 也是同一 Middleware 链的具体实现。推荐顺序为 Validation、Policy、其他 -Middleware、Handler。它在启动阶段通过 `configure_tools()` 编译静态 Canonical Parameters; -Schema 配置错误会中止并回滚启动。参数错误使用 `ToolControl.CONTINUE` 返回有界、无参数值的 -JSON Pointer 错误,让模型可以修正;无效调用不会触发 Policy、Approver 或 Handler。 - -### 3.3 事件、Streaming 与取消 - -事件信封包含 `agent_id`、`run_id` 和单调 `sequence`。当前主要事件包括: - -```text -AgentStarted -AgentContinued -TurnStarted -SteeringApplied / SteeringDiscarded -ContextPressureEvaluated -AssistantThinkingDelta* -AssistantTextDelta* -MessageCompleted -ToolStarted / ToolProgressed* / ToolCompleted -TurnCompleted -CompactionStarted / CompactionCompleted / CompactionFailed -AgentFinished -``` - -事件是只读观察协议,不允许 Sink 改写行为。Partial Assistant Message 不进入 -`AgentState` 或 Session,只有 `MessageCompleted` 才是提交点。Thinking 和 Progress 默认 -不进入后续模型上下文。 - -`CancellationToken` 已传递到 Model、Compactor、Tool Runtime、Middleware、Handler、Behavior Hook 和 -MCP。`abort()` 不等待操作锁,`wait_for_idle()` 覆盖终态 Sink 收尾;取消后同一 Agent -可以复用。 - -### 3.4 Steering 控制面 - -`BaseAgent.steer(content)` 为当前 Run 提交有界 FIFO 控制输入,并返回带稳定 `input_id` 的 -`ControlReceipt`。即时状态明确区分 `accepted`、`agent_idle`、`queue_full` 和 -`run_closing`,不会静默丢弃输入。 - -Steering 不会中断正在执行的 Model Response 或 Tool。Core 在下一次 Provider Context 最终 -确定前消费队列;该安全点也覆盖 Context Pressure Compaction 和 Overflow Retry。已经开始的 -Tool Call 必须先产生对应 Tool Result,Steering 才会作为新的 User Message 进入上下文。 - -成功应用会发出 `SteeringApplied` 并写入 `steering_applied` Session Record。Run 在下一个 -Model Call 前结束时,未消费输入发出 `SteeringDiscarded`,但不伪装成 Durable History。 - -### 3.5 Follow-up Run Chain - -`BaseAgent.follow_up(task)` 将独立任务提交到当前 Run Chain 的有界 FIFO,并返回 -`FollowUpHandle`。Handle 的 Receipt 只表示是否入队;`wait()` 最终返回该独立 Run 的 -`AgentRunResult`。取消某个等待者不会取消队列任务。 - -Follow-up 必须等待前一个 `AgentFinished` 和全部 Event Sink 完成后才开始。每项都有独立 -`run_id`、事件序列和 Session Run。额外的 Run-chain Gate 会阻止已经等待锁的直接 `run()`、 -`compact()` 或生命周期操作插入 Follow-up 之前;首个直接 `run()` 仍会在自身结果完成后立即 -返回,不等待整条 Follow-up Chain。 - -默认 `FollowUpFailurePolicy.DISCARD` 会在前一 Run 失败或取消时丢弃余项;显式选择 -`CONTINUE` 才继续。Queue Full、Agent Idle、Shutdown、前序 Run 非成功和等待者取消都有明确 -Receipt 或异常。Pending Queue 只存在于进程内,只有真正开始的 Follow-up 才通过 -`SessionRecorder` 写入 `run_started`。 - -### 3.6 Continue / Resume Safe Point - -`BaseAgent.continue_run()` 从已提交历史启动新的独立 Run,但不会追加 User Message。Continue -分配新的 `run_id`,发出 `AgentContinued`,并通过 `run_continued` Journal Record 持久化; -`SessionRunIntent` 可区分普通 Task 与 Continue,Restore 后仍能恢复最新终态并继续。 - -Continue 只允许显式安全的 Stop Reason。Active Run、Shutdown、无前序 Run、不支持的失败或 -取消原因,以及未匹配 Tool Result 的 Tool Call 都通过 `ContinueRejectedReason` 明确拒绝。 -调用方可以先读取 `can_continue` 和 `continue_rejection_reason`,也可以处理 -`ContinueRejectedError`。 - -Continue 复用普通 Run 的 Cancellation、Steering Safe Point、Auto Compaction、Follow-up -Run Chain、`AgentFinished` Sink Barrier 和 `wait_for_idle()`。因此它不是修改旧 Run,也不是 -Retry Branch;它只是从安全历史边界继续模型循环。 - -### 3.7 Behavior Hooks - -`BehaviorHook.after_turn()` 是独立于只读 Event Sink 和 Tool Middleware 的行为决策点。它只在 -本来还会进入下一 Turn 的完整非终态回合执行,位置固定为:Tool Result 提交、 -`TurnCompleted` 完成之后,下一次 `TurnStarted` 和 Provider Request 之前。 - -Hook 接收隔离的 `TurnSnapshot` 和当前 `CancellationToken`,不接触可变 `AgentState`。多个 -Hook 按声明顺序 await;`None` 或 `BehaviorDecision.CONTINUE` 继续,首个 STOP 短路后续 Hook, -并产生 `RunStatus.COMPLETED + StopReason.BEHAVIOR_STOP`。正常 `AgentFinished` 和 Session -`run_finished` 仍会完成,因此该终态也是允许显式 Continue 的安全边界。 - -Hook 异常或非法返回值转换为 `RUNTIME_ERROR`,慢 Hook 形成有意的有界背压,`abort()` 可以 -中断它,`wait_for_idle()` 会等待 Hook 与终态 Sink。第一批不加入 `before_model` Context 注入, -避免与 `AgentContextBuilder` 职责重叠。 - -### 3.8 Usage、Context 与 Compaction - -当前上下文管理链路已经完整接通: - -```text -ModelUsage - → RunUsage / max_run_tokens - → ContextUsageEstimate - → ContextBudget - → CompactionPolicy.prepare() - → Compactor / ModelCompactor - → SummaryEntry - → atomic AgentState replacement -``` - -关键语义: - -- Run Budget 与单次请求 Context Window 分离。 -- Usage 的 unknown 与明确的 zero 分离。 -- UTF-8 启发式估算和 Tool Schema 作为完整请求下界。 -- 只在完整 User Turn 边界切分,避免拆开 Tool Call 与 Tool Result。 -- 显式 `compact()` 与自动压缩共享取消和原子替换机制。 -- 自动压力压缩是 opt-in。 -- Provider Overflow 最多 compact-and-retry 一次。 -- Compactor 失败安全终止,不安装部分 Summary。 -- Compaction 通过 Session Record 持久化,原始审计 Entry 不删除。 - -### 3.9 Durable JSONL Session Tree - -`SessionRecorder` 对 Journal Storage 追加语义 Record: - -```text -run_started -message_appended -messages_appended -steering_applied -compaction_applied -run_finished -branch_created -checkpoint -``` - -每条 Record 包含: - -```text -record_id + parent_id + branch_id + revision -session_id + agent_id + sequence + type + data -``` - -文件顺序定义全局 `revision`,`parent_id` 定义逻辑树。当前支持: - -- `load()`:兼容地加载 `main` Head -- `checkout()`:加载 Branch Head 或任意 Record -- `head()` / `list_branches()` -- `fork()`:从已完成投影创建通用 Branch -- `rollback()`:只允许回到源 Branch 的祖先并创建新 Branch -- `prepare_retry()`:回到目标 Run 之前并返回原始 Task -- `SessionRecorder(branch_id=...)`:在指定 Branch 继续执行 -- `expected_head_id` 与 `SessionConflictError` - -Fork、Rollback 和 Retry 都不改写旧历史。Checkout 可以审计未完成节点,但正常 Fork 和 -Restore 会拒绝未完成 Run。Retry 是“准备新 Branch + 返回原 Task”,不会自动重放可能 -有外部副作用的 Tool。 - -JSONL 的当前耐久性保证: - -- Session ID 映射为哈希文件名 -- 每条完整 Record 通过一次追加写入并 `fsync` -- 中断产生的不完整尾行可忽略并在下次追加前修复 -- 完整但损坏的 JSON、错误 Schema、父链和 Branch Head 会明确失败 - -尚未保证的是多个进程对同一 Session 的 read-validate-append 原子性。 - -### 3.10 交付与兼容性 - -- Python 3.12 和 3.13 质量矩阵 -- Ruff、Mypy、Unit Test、sdist/wheel 和安装后 Public API smoke test -- PyPI Trusted Publishing CD -- Release Tag 与 `project.version` 一致性检查 -- Release Commit 必须属于 `main` - -CD 属于交付能力,不改变 Core Runtime 语义。 - -## 4. 与 Pi 的能力矩阵 - -| Harness 能力 | Pi 当前实现 | EJAgent Core 当前实现 | 结论 | -|---|---|---|---| -| Agent Loop | `agentLoop` / `Agent` | `AgentOrchestrator` / `BaseAgent` | 已对齐核心能力 | -| 结构化终止 | Assistant stop reason 为主 | `AgentRunResult` + `StopReason` | Sim 更显式 | -| Provider 边界 | 多 Provider / Model Registry | `ModelAdapter`,仅 OpenAI-compatible | 边界已建,广度不足 | -| Context Transform | `transformContext` | `AgentContextBuilder` | 基础对齐 | -| Event Stream | Message Snapshot + Delta | 分型 Delta + 原子 Message Commit | 语义不同,均可用 | -| Event Barrier | `Agent` Subscriber | 顺序 await `AgentEventSink` | 已具备 | -| Tool Middleware / Hook | `beforeToolCall`、`afterToolCall` | `ToolMiddleware` + Execution Policy | 基础具备 | -| Stop-after-turn Hook | `shouldStopAfterTurn` | 顺序 await `BehaviorHook.after_turn` | 基础已具备 | -| Steering | Queue + safe turn boundary | 有界 FIFO + Model-call Safe Point | 基础已具备 | -| Follow-up | Agent outer loop queue | 独立 FIFO Run Chain + Waitable Handle | 基础已具备 | -| Continue | `continue()` | 独立 Run + 无新增 User Message + Safe Point | 基础已具备 | -| Parallel Tool Calls | 默认并行,可强制顺序 | Opt-in 只读并行 + 副作用屏障 | 基础已具备,更保守 | -| Canonical Tool Definition | 内部强类型定义 | OpenAI-first `ToolDefinition` + 字典兼容 | 基础已具备 | -| Tool Argument Validation | Provider / Tool 层处理 | 可选本地 JSON Schema Middleware | 已具备 | -| Dynamic Tool Set | Runtime 可替换 | 构造/启动时冻结 | 刻意延后 | -| Cancellation | AbortSignal | CancellationToken | 已对齐 | -| Tool Progress | `onUpdate` | `ToolProgressReporter` | 已具备 | -| Usage / Budget | Message Usage + Cost | ModelUsage + RunUsage + Run Budget | 基础已具备,无 Cost | -| Context Pressure | Usage + Estimate | Usage + UTF-8 Estimate + Tool 下界 | 已具备 | -| Explicit Compaction | Harness / Extension Hook | `compact()` + `Compactor` | 已具备 | -| Auto Compaction | Coding Agent 已接通 | `AutoCompactionPolicy` | 已具备 | -| Overflow Recovery | Coding Agent Auto Retry | 标准错误 + 最多一次 Retry | 已具备 | -| Durable Session | JSONL | JSONL Semantic Journal | 已具备 | -| Session Tree | Entry parent tree / forked repo | 同文件 Branch Tree | 已具备,模型不同 | -| Fork | Before / at entry | 从完成 Record Fork | 已具备 | -| Rollback | 通过选 Leaf / Fork 表达 | 独立 `rollback()` 意图 | Sim 更显式 | -| Retry | Continue / Fork 组合 | `prepare_retry(run_id)` | 已具备完整 Run Retry | -| Branch Summary | `branch_summary` Entry / Hook | 只有通用 Compaction Summary | 部分具备 | -| Custom Session Entry | Extension Custom Entry | 固定 Session Record Kind | 未实现扩展协议 | -| Labels / Model Change Entry | 已具备 | 无 | 未实现 | -| ExecutionEnv | Coding Agent 集成 | 无 | 未实现 | -| File / Shell / Git Tools | Coding Agent 内置 | 明确不属于 Core | 待派生 Agent | -| Extension / RPC / TUI | Coding Agent 已具备 | 无 | 待产品层 | - -## 5. 关键设计差异 - -### 5.1 只读 Event 与行为型 Hook - -Pi 的 Agent 和 Coding Agent 允许 Hook 阻止 Tool、改写结果、停止 Turn、注入资源或修改 -Context。EJAgent Core 当前坚持: - -- `AgentEventSink` 只观察,不修改行为。 -- Tool 行为改写只通过 `ToolMiddleware`。 -- Context 改写通过显式 `AgentContextBuilder`。 -- 自动压缩是 Orchestrator 的明确依赖,不通过 Event 回调重入 Agent。 - -这个边界避免事件观察者意外改变运行。需要在完整 Turn 后停止时使用独立的 -`BehaviorHook.after_turn`;后续扩展决策点也不应把可变 Hook 塞进 `AgentEventSink`。 - -### 5.2 Session Tree 模型 - -Pi 当前同时存在 Harness Session Repo 和 Coding Agent Session Manager: - -- Entry 使用 `id/parentId` 形成树。 -- Coding Agent 支持 Compaction、Branch Summary、Custom Entry、Label、Model Change 等类型。 -- Repo Fork 可以把选定 Entry 路径复制到新的 JSONL Session,并记录 Parent Session。 - -EJAgent Core 使用一个 Session 一个 JSONL 文件,所有 Branch 共享祖先 Record: - -```text -main: r1 → r2 → r3 → r4 - ↘ branch_created → r5 → r6 -``` - -这种结构避免复制历史,Fork / Rollback / Retry 意图也可直接审计;代价是同一文件的并发 -协调和索引更重要。 - -### 5.3 Retry 与外部副作用 - -EJAgent Core 不从 `ToolCompleted` 中间恢复,也不自动执行 Retry: - -1. 找到目标 `RUN_STARTED`。 -2. 回到该 Run 之前的完成状态。 -3. 创建 Retry Branch。 -4. 返回原 Task,由调用方显式执行。 - -这比无条件 Continue 更保守,因为历史 Tool 可能已经发送邮件、付款或写入外部系统。 - -### 5.4 Compaction 边界 - -Pi Coding Agent 支持更丰富的 Branch Summary、Extension Compaction 和 mid-turn 相关语义。 -EJAgent Core 当前只在完整 User Turn 边界切分,并让应用拥有 Summary Prompt。它更适合作为 -通用 Core,但在超大单 Turn 或复杂分支摘要方面能力较弱。 - -### 5.5 Tool 执行顺序 - -Pi 当前默认并行执行允许并行的 Tool Call,同时保证最终 Tool Result 按 Assistant 源顺序 -写入 transcript。EJAgent Core 采用更保守的双重 opt-in:Handler 必须声明 `READ_ONLY`,Agent -还必须开启并行策略。连续只读调用并发运行,但 `ToolCompleted` 与 Tool Message 在批次完成后 -按源顺序提交;副作用工具不会与相邻只读批次重叠。 - -并行批次共享 Run Cancellation,每个 Progress 保留自己的 `tool_call_id`。普通异常转换为对应 -Tool Result;外部取消会补齐批次和后续未启动调用;并行批次中的 active peer 全部收尾后,按 -源顺序选择首个 COMPLETE / REJECT / CANCEL。`READ_ONLY` 是 Handler 与 Middleware 对并发安全 -作出的契约,Core 不尝试从名称或参数自动推断副作用。 - -### 5.6 Tool Middleware、Schema Validation 与 Execution Policy - -`ToolMiddleware` 仍是唯一的工具行为拦截协议。`ToolSchemaValidationMiddleware` 和 -`ToolPolicyMiddleware` 都是它的标准实现,与 Audit、Metrics 等 Middleware 使用同一条组合链, -不在 Tool Runtime 前后创建额外通道,也不给 `BaseAgent` 增加专用参数。 - -策略规则属于应用配置,不写入 `ToolDefinition`:Definition 描述工具能力和副作用类别,同一 -工具在不同应用中的权限可以不同。Rule-based 策略采用有序 first-match,推荐 `default=DENY`; -审批由应用提供异步 `ToolApprover`。拒绝沿用 `ToolControl.REJECT`、`RunStatus.REJECTED` 和现有 -Session Tool Result,Handler 不会收到被拒绝的调用。 - -Validation 与 Policy 的失败语义不同:结构错误返回普通 Tool Error 并继续模型循环,权限拒绝 -返回 `ToolControl.REJECT` 并终止 Run。Validation 错误只包含 Tool、Code、JSON Pointer、失败 -Keyword 和通用消息;模型提供的参数值不会被复制到结果或 Session。 - -## 6. 当前技术边界与风险 - -### 6.1 JSONL 多进程写入正确性——已加固 - -`JsonlSessionStorage` 现在为每个 Session 使用稳定 `.jsonl.lock` Sidecar,并组合进程内共享锁 -与 POSIX Advisory Lock。读取、修复尾行、校验 `expected_head_id`、分配 Revision / Parent、 -追加和 `fsync` 位于同一临界区。CAS 失败继续返回 `SessionConflictError`,锁等待超时返回 -`SessionLockTimeoutError`;等待在 Worker Thread 中执行,可以取消而不阻塞 Event Loop。 - -该实现面向 macOS / Linux 本地文件系统。NFS 等网络文件系统需要验证 Advisory Lock 语义, -或者提供其他 `SessionTreeStorage` Backend。 - -### 6.2 Tree Storage 协议——已完成 - -公开的 `SessionTreeStorage` 在 `SessionJournalStorage` 之上统一 `records()`、`head()`、 -`list_branches()`、`fork()`、`rollback()` 和 `prepare_retry()`。上层可以仅依赖协议编程, -不必绑定 JSONL 具体实现。 - -### 6.3 JSONL 读取成本随 Journal 线性增长 - -每次操作都会验证完整文件并重建索引;任意节点投影还需要沿父链回放。当前 Checkpoint -可以缩短语义重建,但没有持久 Sidecar Index、Checkpoint Policy 或归档策略。长生命周期 -服务需要明确性能上限。 - -可重复基准位于 `benchmarks/session_journal.py`。本机 1,000 条小型 Checkpoint Journal -约 401 KB,完整 Load 约 14.7 ms;从空文件逐条构建约 7.68 s,平均追加约 7.68 ms。 -这验证了当前规模仍可用,但也确认逐次追加的累计成本呈二次增长。现阶段不增加 Index; -当真实 Session 接近数千至数万 Record 时,再设计可完全从 JSONL 重建的 Sidecar Index。 - -### 6.4 Canonical Tool Definition——已完成 OpenAI-first 基础版 - -`ToolDefinition` 已成为 Core 内部事实来源,统一保存 Name、Description、Parameters、可选 -`strict` 和 `ToolEffect`。`MethodToolHandler` 同时接受强类型定义和旧 OpenAI 字典;旧字典只在 -构造边界归一化一次,`.tools` 继续提供兼容的 OpenAI 请求视图。Tool Runtime 的路由、重复名称 -校验、副作用调度和 Agent 日志不再解析 `tool["function"]["name"]`。 - -该设计明确是 OpenAI-first,不尝试抽取所有 Provider 的最低公分母。Parameters 只做 JSON -可序列化校验,具体 JSON Schema 能力仍由 OpenAI-compatible Provider 决定。未来如有真实多 -Provider 需求,可以在稳定定义外增加 Adapter 转换,而无需现在预建复杂抽象。 - -### 6.5 Provider 广度——按需求延后 - -目前只有 OpenAI-compatible Adapter。这是当前产品边界,不再视为近期优先缺口;只有真实接入 -需求出现后,才验证 Thinking、Cache Usage、Tool Schema 和错误分类差异。 - -### 6.6 Event Backpressure 是同步的 - -事件按顺序 await,保证确定性和终态屏障,但慢 UI 或网络 Sink 会拖慢 Agent。未来应提供 -有界缓冲 Sink Adapter,而不是在 Core 中创建无界后台任务。 - -### 6.7 Tool Schema Validation——已完成基础版 - -`ToolSchemaValidationMiddleware` 使用 `jsonschema` 验证 Canonical `parameters`,支持 Schema -选择的标准 Draft 语义。它是显式可选 Middleware,不改变未注册它的旧 Agent 行为,也不把 -Provider `strict` 当作本地安全边界。 - -当前不启用外部 Format Checker,也不提供领域语义验证;路径范围、金额上限等约束继续属于 -Tool Policy。远程 Schema Registry、动态 Schema 和自定义 Format 暂不进入 Core。 - -## 7. 后续建设顺序 - -### 阶段一:执行内核——已完成 - -- Agent Loop、Runtime Policy、Structured Result -- Provider Adapter、Tool Runtime、Skill Resource -- Cancellation、Streaming、Progress、Usage Budget - -### 阶段二:Context 与 Compaction——已完成基础版 - -- Context Pressure、Budget 和完整 Turn Preparation -- Explicit Compaction、ModelCompactor -- Auto Compaction 和 Context Overflow Recovery - -### 阶段三:Durable Session Tree——已完成基础版 - -- Versioned Codec 和 Semantic JSONL Journal -- Checkpoint、Crash Tail Recovery 和 Replay -- Branch Head、Checkout、Fork、Rollback、Retry -- Branch-aware SessionRecorder 和 Head Conflict - -### 阶段四:Session Journal Hardening——已完成 - -- 跨实例、跨进程 read-validate-append 原子性 -- 后端无关的 `SessionTreeStorage` 协议 -- 多进程竞争和 Crash Recovery 测试 -- 锁超时、取消和异常退出恢复 -- 长 Journal 基准;Sidecar Index 延后到真实规模需要时 - -### 阶段五:Harness 行为控制面——进行中 - -- Steering Queue、Safe Point、回执和 Session 审计——已完成基础版 -- Follow-up Queue、Run Handle、失败策略和终态屏障——已完成基础版 -- Continue / Resume-safe-point、Run Intent 和 Session 回放——已完成基础版 -- 独立于只读 Event 的 `after_turn` Behavior Hook——已完成基础版 -- Queue 与 Session Journal 的进程内 / Durable 边界——已明确 -- `before_model` 等扩展决策点——按真实需求延后 - -### 阶段六:工具调度与 Provider 扩展 - -- Parallel Tool Calls 与 Side-effect Policy——已完成基础版 -- Canonical Tool Definition——已完成 OpenAI-first 基础版 -- Tool Execution Policy Middleware——已完成基础版 -- Tool Argument Schema Validation Middleware——已完成基础版 -- Dynamic Tool Set——按当前静态工具集设计延后 -- 第二个真实 Provider Adapter——按真实需求延后 - -### 阶段七:ExecutionEnv 与派生 CodeAgent - -```text -CodeAgent - ├── ExecutionEnv / Workspace - ├── Read / Write / Edit - ├── Grep / Find / List - ├── Bash / Git - ├── Sandbox / Approval Policy - ├── Completion Tool(可选策略) - └── CLI / RPC / TUI Adapter -``` - -## 8. 下一步任务建议 - -Tool Argument Validation 完成后,下一项建议是 **Tool Policy Decision Audit Events**。当前 -拒绝结果可以通过 `ToolCompleted` 和 Session Tool Message 审计,但被允许的 Rule、审批请求和 -审批结果没有独立的只读事件,长期运行服务难以回答“为什么这个副作用调用被允许”。 - -### 8.1 建议实现范围 - -1. 增加只读 `ToolPolicyEvaluated`、`ToolApprovalRequested` 和 `ToolApprovalResolved` 事件。 -2. 事件只包含 `tool_call_id`、Tool Name、Action、`rule_id`、安全 Reason 和审批状态,不包含参数。 -3. 事件由现有 `AgentEventEmitter` 顺序发布,不允许 Event Sink 改写 Policy 决策。 -4. 审批请求和结果必须成对;异常、取消和缺少 Approver 也有明确 Resolution。 -5. 并行调用允许 Policy Event 交错,但最终 Tool Result 和 Session 顺序继续保持确定性。 -6. 是否持久化独立 Policy Journal Record 需显式决定,不能把只读临时事件自动混入历史。 - -### 8.2 验收标准 - -- 每个 Policy/Approval 决策都有稳定关联的 `run_id`、`tool_call_id` 和事件序号。 -- 允许、拒绝、无 Approver、审批拒绝、审批异常和取消都有确定性事件测试。 -- Event Sink 失败仍被隔离,不得改变工具是否执行。 -- 事件不泄露 Arguments,Session 持久化策略有明确测试和文档。 -- 现有 ToolCompleted、并行提交、Cancellation 和 Policy fail-closed 语义保持不变。 - -完成 Policy 审计事件后,再根据真实使用数据决定是否增加 Tool Output 大小边界或领域级规则库; -Dynamic Tool Set 和第二个 Provider 继续按需求延后。 diff --git a/docs/runtime-kernel-harness-design.md b/docs/runtime-kernel-harness-design.md new file mode 100644 index 0000000..502294a --- /dev/null +++ b/docs/runtime-kernel-harness-design.md @@ -0,0 +1,189 @@ +# Runtime Kernel and Agent Harness Design + +## Status + +This document defines the implemented pre-1.0 EJAgent Core architecture and is +normative for its contracts. The former stateful agent loop, dictionary-message +runtime, Handler/Middleware stack, and event-driven Session writer have been +removed. Only a Store-private decoder remains for one-way legacy migration. + +Implementation progress: + +- Typed message, Run, Model, Tool, cancellation, usage, and JSON contracts are + implemented in `ejagent.contracts`. +- `ejagent.kernel.RuntimeKernel` executes one Run over a private workspace and + returns a deterministic Delta plus Audit records. +- `ejagent.harness.AgentHarness` owns single-agent resource lifecycle, + cancellation, FIFO Run admission, snapshot recovery, and atomic + compare-and-commit through the new `SessionStore` contract. +- `MemorySessionStore` provides an idempotent in-process adapter and rejects + stale revisions and reused Run IDs with different content. +- Conversation recovery now uses an immutable `ConversationSnapshot`; durable + Run facts are read separately as `RunAudit` values. +- Every model call receives a disposable `ContextView` from a + `ContextPipeline`. Identity and derived-compaction implementations preserve + committed history while permitting summaries and transient instructions. +- Steering is admitted by the Harness and consumed by the Kernel only at model + call safe points; unapplied inputs remain auditable without entering + Conversation. +- Follow-ups run as independent FIFO Runs. `RunObserver` delivery happens + asynchronously after the Store decision, so observer latency and failure + cannot alter execution or commit semantics. +- `JsonlSessionStore` provides locked, append-only durable commits with CAS, + idempotent Run IDs, crash-tail recovery, and typed Conversation/Audit codecs. + Legacy JSONL Sessions can be imported once or rejected with an actionable + `SessionMigrationError`; compacted projections are not treated as source + history. +- `OpenAIModelPort` translates typed Context and Tool values at the Provider + seam. Function, composite, and MCP ToolExecutors share the Kernel Tool + contract; `SkillsContextPipeline` contributes disposable skill instructions. +- `AnthropicModelPort` independently translates system instructions, content + blocks, streamed tool input, and cache-aware usage from the Anthropic + Messages protocol. Both Provider adapters pass the same Kernel-facing stream + contract without adding Provider concepts to Core. +- Representative OpenAI, Anthropic, local Tool, MCP, Skills, memory recovery, + and durable recovery examples run through `AgentHarness`. +- `ejagent` now exports only the new composition surface. Legacy execution and + commit packages are absent from built artifacts. Migration stages 1–8 are + complete. + +## Scope + +EJAgent Core contains two layers for one logical agent: + +```text +AgentHarness + ├─ long-lived conversation state + ├─ resource and control lifecycle + ├─ context projection + ├─ durable commit coordination + └─ RuntimeKernel + └─ one deterministic model-tool Run +``` + +Multi-agent management and arbitrary mid-Run pause/resume are out of scope. +The supported controls are cancellation, steering at model-call safe points, +queued follow-ups, and continuation from a committed revision. + +## Ownership Boundaries + +### RuntimeKernel + +The Kernel owns only Run-local state: turn counters, usage, repeat guards, +cancellation, the private message workspace, and temporary context. It receives +an immutable `RunSpec` and returns a `RunOutcome`. It never starts resources, +persists Sessions, or mutates Harness state. + +### AgentHarness + +The Harness owns conversation revisions, the last committed result, Provider +and Tool resource lifecycle, control queues, context policy, and SessionStore +commit coordination. Each Run receives a snapshot of the current revision. +Configuration changes take effect only on the next Run. + +## Data Domains + +Three data domains must remain distinct: + +- **Conversation** contains immutable, typed messages that are valid input to a + future Run. +- **Audit** is an append-only account of what actually happened, including + partial failed or cancelled Runs and external side effects. +- **ContextView** is a disposable projection for one model request. It may use + summaries, windows, Skills, and transient instructions without rewriting the + Conversation or Audit. + +Provider payload dictionaries are produced only by Provider adapters. Core +messages use a closed, Provider-neutral type union. + +## Run Transaction + +```text +committed revision N + │ snapshot + ▼ +private RunWorkspace + │ Kernel execution + ▼ +RunOutcome(result, delta, audit) + │ compare-and-commit N + ▼ +committed revision N+1 +``` + +`RunDelta.base_revision` prevents stale writes. A configured durable Store is a +correctness boundary: the Harness cannot report a committed completion until +the Store accepts the idempotent commit. Observer failures never substitute for +Store failures and do not alter execution. + +Failed and cancelled Runs remain auditable but do not advance active +Conversation history by default. A Harness policy may explicitly promote an +eligible Delta. + +## Extension Boundaries + +New extension points are limited to narrow, ordered protocols: + +1. `ModelPort` normalizes Provider requests and streams. +2. `ContextPipeline` contributes, transforms, and serializes ContextViews. +3. `ToolRuntime` validates, authorizes, schedules, and executes Tools. +4. `RunPolicy` decides budgets, retries, termination, and commit eligibility. +5. `RunObserver` observes events without changing Run results. +6. `SessionStore` commits revisions and durable Run records. + +Resources are started transactionally by the Harness and shut down in reverse +order. Kernel and Tool dispatch methods require ready dependencies. + +## Tool Semantics + +Tool scheduling does not infer safety from implementation names. Definitions +declare effect, idempotency, and an optional concurrency key. Execution may be +concurrent only when semantics and policy permit it. Completion and +Conversation commit order always follow the model's source order; actual timing +is retained in Audit records. Non-idempotent Tools are never retried by Core. + +## Failure Contract + +Expected operational failures return structured outcomes. These include model +timeouts and rate limits, Tool failures, policy rejection, cancellation, +budgets, context overflow, and persistence failure. Exceptions are reserved for +invalid configuration, protocol violations, and broken invariants. + +## Public API + +`AgentHarness` is the primary entry point. `RuntimeKernel` has a narrow advanced +API for direct single-Run embedding. `BaseAgent` is removed rather than kept as +a compatibility execution path. Public contracts are intentionally small; +workspace mechanics, schedulers, commit coordinators, and state-machine phases +remain internal. + +The initial refactor remains in one distribution. Module dependency direction +is enforced before considering separate Provider or backend packages. + +## Migration Strategy + +1. Add typed message and Run contracts with contract tests. +2. Implement the Kernel over a private workspace and deterministic Delta. +3. Implement the Harness, resource lifecycle, controls, and atomic commit. +4. Split Conversation, Audit, and ContextView; make compaction derived. +5. adapt Memory/JSONL Stores and provide legacy Session decoding/migration. +6. Adapt OpenAI, MCP, Skills, examples, and documentation. +7. Replace the public API and delete the old execution path. +8. Validate `ModelPort` with a second, materially different Provider protocol. + +Intermediate steps may coexist on a development branch, but no release should +expose two competing execution or commit semantics. + +## Acceptance Criteria + +- Kernel tests prove no mutation escapes before a Harness commit. +- Identical RunSpec and scripted dependencies produce identical committed + message order. +- Durable commit failure cannot advance the active revision or report success. +- Observer failure cannot change a Run result. +- Context compaction is rebuildable from immutable Conversation history. +- Tool execution timing may vary, but Delta order remains deterministic. +- ModelPort and SessionStore implementations pass shared contract suites. +- Legacy Session fixtures either migrate losslessly or fail with a typed, + actionable migration error. +- Kernel modules do not import concrete Providers, Stores, MCP, or Skills. diff --git a/examples/01_stateful_chat.py b/examples/01_stateful_chat.py index c854ca1..45252fa 100644 --- a/examples/01_stateful_chat.py +++ b/examples/01_stateful_chat.py @@ -1,31 +1,27 @@ -"""Plain chat with stateful conversation memory.""" +"""Run two tasks against one stateful AgentHarness.""" import asyncio -from ejagent import BaseAgent, ModelConfig, OpenAIModelAdapter +from ejagent.contracts import SystemMessage +from ejagent.harness import AgentHarness +from ejagent.providers import ModelConfig, OpenAIModelPort +from ejagent.tools import FunctionToolExecutor async def main() -> None: - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), + harness = AgentHarness( agent_id="tutor", - system_prompt="You are a concise Python tutor.", + model=OpenAIModelPort(ModelConfig.from_env()), + tools=FunctionToolExecutor(), + initial_messages=(SystemMessage("You are a concise Python tutor."),), ) - try: - first = await agent.runtime( - task="Remember that my preferred language is Python." - ) - print(f"First response: {first}") + async with harness: + first = await harness.run("Remember that my preferred language is Python.") + print(f"First response: {first.result.output}") - second = await agent.runtime(task="Which programming language do I prefer?") - print(f"Memory response: {second}") - - agent.reset() - third = await agent.runtime(task="Which programming language do I prefer?") - print(f"After reset: {third}") - finally: - await agent.shutdown() + second = await harness.run("Which programming language do I prefer?") + print(f"Memory response: {second.result.output}") if __name__ == "__main__": diff --git a/examples/02_custom_tool.py b/examples/02_custom_tool.py index 3afc6c9..dccbf95 100644 --- a/examples/02_custom_tool.py +++ b/examples/02_custom_tool.py @@ -1,25 +1,23 @@ -"""Compose an agent with a custom atomic math tool.""" +"""Compose an AgentHarness with a custom atomic math tool.""" import asyncio -from collections.abc import Mapping -from typing import Any -from ejagent import ( - BaseAgent, +from ejagent.contracts import ( CancellationToken, - MethodToolHandler, - ModelConfig, - OpenAIModelAdapter, - StepOutcome, + ToolCall, ToolControl, ToolDefinition, - ToolEffect, + ToolExecutionResult, + ToolSemantics, ) +from ejagent.harness import AgentHarness +from ejagent.providers import ModelConfig, OpenAIModelPort +from ejagent.tools import FunctionTool, FunctionToolExecutor ADD_TOOL = ToolDefinition( name="add", description="Add two numbers.", - parameters={ + input_schema={ "type": "object", "properties": { "left": {"type": "number"}, @@ -27,44 +25,45 @@ }, "required": ["left", "right"], }, - effect=ToolEffect.READ_ONLY, + semantics=ToolSemantics.read_only(), ) -class MathHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((ADD_TOOL,)) - - async def do_add( - self, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - left = arguments.get("left") - right = arguments.get("right") - if not isinstance(left, (int, float)) or not isinstance(right, (int, float)): - return StepOutcome( - {"status": "error", "error": "left and right must be numbers"} - ) - return StepOutcome( - {"status": "success", "value": left + right}, - control=ToolControl.COMPLETE, +async def add( + call: ToolCall, + cancellation: CancellationToken, +) -> ToolExecutionResult: + cancellation.raise_if_cancelled() + left = call.arguments.get("left") + right = call.arguments.get("right") + if ( + isinstance(left, bool) + or not isinstance(left, (int, float)) + or isinstance(right, bool) + or not isinstance(right, (int, float)) + ): + return ToolExecutionResult( + {"status": "error", "error": "left and right must be numbers"}, + error="left and right must be numbers", ) + value = left + right + return ToolExecutionResult( + {"status": "success", "value": value}, + control=ToolControl.COMPLETE, + output=str(value), + ) async def main() -> None: - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), + harness = AgentHarness( agent_id="calculator", - handlers=[MathHandler()], + model=OpenAIModelPort(ModelConfig.from_env()), + tools=FunctionToolExecutor((FunctionTool(ADD_TOOL, add),)), ) - try: - result = await agent.runtime(task="Use the add tool to calculate 19.5 + 22.5.") - print(result) - finally: - await agent.shutdown() + async with harness: + outcome = await harness.run("Use the add tool to calculate 19.5 + 22.5.") + print(outcome.result.output) if __name__ == "__main__": diff --git a/examples/03_anthropic_chat.py b/examples/03_anthropic_chat.py new file mode 100644 index 0000000..5df96ee --- /dev/null +++ b/examples/03_anthropic_chat.py @@ -0,0 +1,25 @@ +"""Run one task through the Anthropic Messages Provider adapter.""" + +import asyncio + +from ejagent.contracts import SystemMessage +from ejagent.harness import AgentHarness +from ejagent.providers import AnthropicConfig, AnthropicModelPort +from ejagent.tools import FunctionToolExecutor + + +async def main() -> None: + harness = AgentHarness( + agent_id="anthropic-assistant", + model=AnthropicModelPort(AnthropicConfig.from_env()), + tools=FunctionToolExecutor(), + initial_messages=(SystemMessage("Answer precisely."),), + ) + + async with harness: + outcome = await harness.run("Explain the ModelPort boundary in one sentence.") + print(outcome.result.output) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/examples/04_mcp_tools.py b/examples/04_mcp_tools.py index d24e9cd..87b866c 100644 --- a/examples/04_mcp_tools.py +++ b/examples/04_mcp_tools.py @@ -1,36 +1,32 @@ -"""Expose tools from an MCP server to an agent.""" +"""Expose tools from an MCP server to an AgentHarness.""" import asyncio from pathlib import Path -from ejagent import ( - BaseAgent, - McpToolHandler, - ModelConfig, - OpenAIModelAdapter, -) +from ejagent.contracts import SystemMessage +from ejagent.harness import AgentHarness +from ejagent.providers import ModelConfig, OpenAIModelPort +from ejagent.tools import McpToolExecutor MCP_CONFIG = Path(__file__).with_name("mcp_config.json") async def main() -> None: - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), + harness = AgentHarness( agent_id="browser", - system_prompt=( - "Use MCP tools to inspect the requested page, then answer with the " - "page title and relevant result summary." + model=OpenAIModelPort(ModelConfig.from_env()), + tools=McpToolExecutor(MCP_CONFIG), + initial_messages=( + SystemMessage( + "Use MCP tools to inspect the requested page, then answer with " + "the page title and relevant result summary." + ), ), - handlers=[McpToolHandler(MCP_CONFIG)], ) - try: - result = await agent.runtime( - task="Open https://baidu.com and report the page title." - ) - print(result) - finally: - await agent.shutdown() + async with harness: + outcome = await harness.run("Open https://baidu.com and report the page title.") + print(outcome.result.output) if __name__ == "__main__": diff --git a/examples/05_skill.py b/examples/05_skill.py new file mode 100644 index 0000000..bd06085 --- /dev/null +++ b/examples/05_skill.py @@ -0,0 +1,38 @@ +"""Project a local release-note skill into one Run's ContextView.""" + +import asyncio +from pathlib import Path + +from ejagent.context import SkillsContextPipeline +from ejagent.contracts import SystemMessage +from ejagent.harness import AgentHarness +from ejagent.providers import ModelConfig, OpenAIModelPort +from ejagent.tools import FunctionToolExecutor + +SKILLS_DIR = Path(__file__).with_name("skills") + + +async def main() -> None: + harness = AgentHarness( + agent_id="release-writer", + model=OpenAIModelPort(ModelConfig.from_env()), + tools=FunctionToolExecutor(), + context=SkillsContextPipeline(SKILLS_DIR), + initial_messages=( + SystemMessage( + "Use the explicitly loaded local skill and return the final " + "deliverable directly." + ), + ), + ) + + async with harness: + outcome = await harness.run( + "$release_notes Write release notes for EJAgent Core 0.6.0. " + "Changes: added a runtime kernel, lifecycle harness, and durable store." + ) + print(outcome.result.output) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/examples/06_session_resume.py b/examples/06_session_resume.py new file mode 100644 index 0000000..d293ffd --- /dev/null +++ b/examples/06_session_resume.py @@ -0,0 +1,45 @@ +"""Resume a committed Conversation in a new AgentHarness instance.""" + +import asyncio + +from ejagent.contracts import SystemMessage +from ejagent.harness import AgentHarness, MemorySessionStore +from ejagent.providers import ModelConfig, OpenAIModelPort +from ejagent.tools import FunctionToolExecutor + +AGENT_ID = "session-demo" +SYSTEM_PROMPT = ( + "Preserve user-provided facts exactly. Answer only from the conversation." +) + + +async def main() -> None: + store = MemorySessionStore() + first = AgentHarness( + agent_id=AGENT_ID, + model=OpenAIModelPort(ModelConfig.from_env()), + tools=FunctionToolExecutor(), + store=store, + initial_messages=(SystemMessage(SYSTEM_PROMPT),), + ) + async with first: + outcome = await first.run("Remember this exact code: CORE-2048") + print(f"first response: {outcome.result.output}") + + resumed = AgentHarness( + agent_id=AGENT_ID, + model=OpenAIModelPort(ModelConfig.from_env()), + tools=FunctionToolExecutor(), + store=store, + initial_messages=(SystemMessage(SYSTEM_PROMPT),), + ) + async with resumed: + outcome = await resumed.run("Return only the exact stored code.") + print(f"resumed response: {outcome.result.output}") + + print(f"committed revision: {resumed.revision}") + print(f"audited runs: {len(await store.load_audit(AGENT_ID))}") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/examples/06_skill.py b/examples/06_skill.py deleted file mode 100644 index ceb400f..0000000 --- a/examples/06_skill.py +++ /dev/null @@ -1,36 +0,0 @@ -"""Load a local release-note skill by explicit name.""" - -import asyncio -from pathlib import Path - -from ejagent import BaseAgent, ModelConfig, OpenAIModelAdapter - -SKILLS_DIR = Path(__file__).with_name("skills") - - -async def main() -> None: - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="release-writer", - system_prompt=( - "Use the explicitly loaded local skill to complete the task. " - "Return the final deliverable directly." - ), - skills_dir=SKILLS_DIR, - ) - - try: - result = await agent.runtime( - task=( - "$release_notes Write release notes for EJAgent Core 0.2.4. Changes: " - "added an orchestrator, structured run results, and runtime " - "policy controls." - ) - ) - print(result) - finally: - await agent.shutdown() - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/07_durable_session.py b/examples/07_durable_session.py new file mode 100644 index 0000000..ebb720a --- /dev/null +++ b/examples/07_durable_session.py @@ -0,0 +1,66 @@ +"""Persist Core Conversation and Audit state across process invocations.""" + +import argparse +import asyncio +import os +from pathlib import Path + +from ejagent.contracts import SystemMessage +from ejagent.harness import AgentHarness +from ejagent.providers import ModelConfig, OpenAIModelPort +from ejagent.storage import JsonlSessionStore +from ejagent.tools import FunctionToolExecutor + +AGENT_ID = "durable-session-agent" +SYSTEM_PROMPT = ( + "Preserve user-provided facts exactly and answer only from durable history." +) + + +def harness(store: JsonlSessionStore) -> AgentHarness: + return AgentHarness( + agent_id=AGENT_ID, + model=OpenAIModelPort(ModelConfig.from_env()), + tools=FunctionToolExecutor(), + store=store, + initial_messages=(SystemMessage(SYSTEM_PROMPT),), + ) + + +async def record(store: JsonlSessionStore) -> None: + agent = harness(store) + async with agent: + outcome = await agent.run( + "Remember this exact project code and reply only ACK: CORE-2048" + ) + print(f"record response: {outcome.result.output}") + print(f"committed revision: {agent.revision}") + + +async def resume(store: JsonlSessionStore) -> None: + if await store.load(AGENT_ID) is None: + raise RuntimeError("run the record command before resume") + agent = harness(store) + async with agent: + outcome = await agent.run("Return only the exact project code stored before.") + print(f"resume response: {outcome.result.output}") + print(f"audited runs: {len(await store.load_audit(AGENT_ID))}") + + +async def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("command", choices=("record", "resume")) + parser.add_argument( + "--session-dir", + default=os.getenv("EJAGENT_SESSION_DIR", ".ejagent-sessions"), + ) + args = parser.parse_args() + store = JsonlSessionStore(Path(args.session_dir)) + if args.command == "record": + await record(store) + else: + await resume(store) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/examples/07_event_observers.py b/examples/07_event_observers.py deleted file mode 100644 index 1a9f2c2..0000000 --- a/examples/07_event_observers.py +++ /dev/null @@ -1,54 +0,0 @@ -"""Observe a real provider-backed agent through multiple event sinks.""" - -import asyncio -from collections import Counter - -from ejagent import ( - AgentEvent, - AgentEventKind, - BaseAgent, - CompositeAgentEventSink, - ModelConfig, - OpenAIModelAdapter, -) - - -class ConsoleEventSink: - async def emit(self, event: AgentEvent) -> None: - print(f"event #{event.sequence}: {event.kind} run={event.run_id[:8]}") - - -class EventMetricsSink: - def __init__(self) -> None: - self.counts: Counter[AgentEventKind] = Counter() - - async def emit(self, event: AgentEvent) -> None: - self.counts[event.kind] += 1 - - -async def main() -> None: - metrics = EventMetricsSink() - event_sink = CompositeAgentEventSink([ConsoleEventSink(), metrics]) - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="event-demo", - system_prompt="Answer concisely in one paragraph.", - event_sink=event_sink, - ) - - try: - result = await agent.run( - task="Explain why lifecycle events are useful in an Agent Harness." - ) - print(f"result: {result.status} / {result.stop_reason}") - print(f"output: {result.output}") - print( - "metrics:", - {kind.value: count for kind, count in metrics.counts.items()}, - ) - finally: - await agent.shutdown() - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/08_session_resume.py b/examples/08_session_resume.py deleted file mode 100644 index 0597936..0000000 --- a/examples/08_session_resume.py +++ /dev/null @@ -1,70 +0,0 @@ -"""Persist a real conversation and resume it in a new agent instance.""" - -import asyncio - -from ejagent import ( - BaseAgent, - MemorySessionStorage, - ModelConfig, - OpenAIModelAdapter, - SessionRecorder, -) - - -async def main() -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="resume-demo", storage=storage) - - first_agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="session-demo", - system_prompt=( - "Preserve user-provided facts exactly. Answer only from the " - "conversation and do not invent implementation details." - ), - event_sink=recorder, - ) - try: - first = await first_agent.run( - task=( - "Store this exact statement and reply only ACK: EJAgent Core " - "uses lifecycle events to build Sessions without coupling " - "persistence to the Agent Loop." - ) - ) - print(f"first response: {first.output}") - finally: - await first_agent.shutdown() - - saved = await recorder.load() - if saved is None: - raise RuntimeError("session was not saved") - - resumed_agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="session-demo", - system_prompt=( - "Preserve user-provided facts exactly. Answer only from the " - "conversation and do not invent implementation details." - ), - event_sink=recorder, - ) - resumed_agent.reset(saved.messages) - try: - resumed = await resumed_agent.run( - task="Repeat the exact stored EJAgent Core statement." - ) - print(f"resumed response: {resumed.output}") - finally: - await resumed_agent.shutdown() - - updated = await recorder.load() - if updated is None: - raise RuntimeError("resumed session was not saved") - - print(f"saved runs: {len(updated.runs)}") - print(f"message roles: {[message['role'] for message in updated.messages]}") - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/09_runtime_control.py b/examples/09_runtime_control.py deleted file mode 100644 index 353dc63..0000000 --- a/examples/09_runtime_control.py +++ /dev/null @@ -1,79 +0,0 @@ -"""Abort a real provider request, wait for idle, and reuse the agent.""" - -import asyncio -import os -from collections.abc import AsyncIterator -from typing import Any - -from ejagent import ( - AgentEvent, - BaseAgent, - CancellationToken, - ModelConfig, - ModelStreamEvent, - OpenAIModelAdapter, -) - - -class ObservableOpenAIModelAdapter(OpenAIModelAdapter): - """Expose when a real provider request is about to start.""" - - def __init__(self, config: ModelConfig) -> None: - super().__init__(config) - self.request_started = asyncio.Event() - - async def stream( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - self.request_started.set() - async for event in super().stream( - context, - cancellation=cancellation, - ): - yield event - - -class ConsoleEventSink: - async def emit(self, event: AgentEvent) -> None: - print(f"event #{event.sequence}: {event.kind}") - - -async def main() -> None: - abort_delay = float(os.getenv("HARNESS_ABORT_DELAY", "0.2")) - model = ObservableOpenAIModelAdapter(ModelConfig.from_env()) - agent = BaseAgent( - model, - agent_id="runtime-control-demo", - event_sink=ConsoleEventSink(), - ) - - try: - run = asyncio.create_task( - agent.run( - task=( - "Write a detailed 2000-word technical essay about Agent " - "Harness architecture, runtime control, and persistence." - ) - ) - ) - await model.request_started.wait() - await asyncio.sleep(abort_delay) - - accepted = agent.abort("cancelled by the runtime-control example") - await agent.wait_for_idle() - first_result = await run - - print(f"abort accepted: {accepted}") - print(f"first result: {first_result.status} / {first_result.stop_reason}") - - reused = await agent.run(task="Reply with exactly: the agent is reusable") - print(f"reused result: {reused.output}") - finally: - await agent.shutdown() - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/10_composed_harness.py b/examples/10_composed_harness.py deleted file mode 100644 index cade140..0000000 --- a/examples/10_composed_harness.py +++ /dev/null @@ -1,161 +0,0 @@ -"""Compose the real provider, tools, events, Session, and runtime policy.""" - -import asyncio -from collections.abc import Mapping -from typing import Any - -from ejagent import ( - AgentEvent, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - MemorySessionStorage, - MethodToolHandler, - ModelConfig, - OpenAIModelAdapter, - RuntimePolicy, - SessionRecorder, - StepOutcome, - ToolCallContext, - ToolControl, - ToolMiddleware, - ToolNext, -) - -INSPECT_TOOL = { - "type": "function", - "function": { - "name": "inspect_project", - "description": "Return current EJAgent Core Harness metadata.", - "parameters": {"type": "object", "properties": {}}, - }, -} - -FINISH_TOOL = { - "type": "function", - "function": { - "name": "finish", - "description": "Explicitly finish after inspecting the project.", - "parameters": { - "type": "object", - "properties": {"summary": {"type": "string"}}, - "required": ["summary"], - }, - }, -} - - -class HarnessTools(MethodToolHandler): - def __init__(self) -> None: - super().__init__((INSPECT_TOOL, FINISH_TOOL)) - self.inspected = False - - async def on_task_start(self) -> None: - self.inspected = False - - async def do_inspect_project( - self, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - self.inspected = True - return StepOutcome( - { - "agent_core": "EJAgent Core", - "orchestrator": "AgentOrchestrator", - "events": "AgentEvent + CompositeAgentEventSink", - "session": "SessionRecorder + MemorySessionStorage", - "runtime_control": "CancellationToken + abort + wait_for_idle", - } - ) - - async def do_finish( - self, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - if not self.inspected: - return StepOutcome( - { - "status": "error", - "error": "inspect_project must be called before finish", - } - ) - return StepOutcome( - {"summary": arguments["summary"]}, - control=ToolControl.COMPLETE, - ) - - -class ToolAuditMiddleware(ToolMiddleware): - def __init__(self) -> None: - super().__init__() - self.calls: list[str] = [] - - async def __call__( - self, - context: ToolCallContext, - call_next: ToolNext, - ) -> StepOutcome: - if context.cancellation is None: - raise RuntimeError("tool call has no cancellation token") - self.calls.append(context.tool_name) - print(f"middleware before: {context.tool_name}") - outcome = await call_next(context) - print(f"middleware after: {context.tool_name}") - return outcome - - -class ConsoleEventSink: - async def emit(self, event: AgentEvent) -> None: - print(f"event #{event.sequence}: {event.kind}") - - -async def main() -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder( - session_id="composed-harness", - storage=storage, - ) - middleware = ToolAuditMiddleware() - event_sink = CompositeAgentEventSink([ConsoleEventSink(), recorder]) - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="composed-harness", - system_prompt=( - "You are demonstrating an Agent Harness. First call " - "inspect_project exactly once. Then call finish with a concise " - "summary of the returned metadata. Never finish with plain text." - ), - handlers=[HarnessTools()], - middlewares=[middleware], - runtime_policy=RuntimePolicy( - max_steps=6, - max_no_tool_responses=2, - require_explicit_finish=True, - ), - event_sink=event_sink, - ) - - try: - result = await agent.run( - task="Inspect the registered Harness metadata and finish." - ) - finally: - await agent.shutdown() - - session = await recorder.load() - if session is None: - raise RuntimeError("session was not recorded") - - print(f"result: {result.status} / {result.stop_reason}") - print(f"output: {result.output}") - print(f"tool calls: {middleware.calls}") - print(f"session roles: {[message['role'] for message in session.messages]}") - print(f"session runs: {len(session.runs)}") - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/11_streaming_events.py b/examples/11_streaming_events.py deleted file mode 100644 index c8dfad1..0000000 --- a/examples/11_streaming_events.py +++ /dev/null @@ -1,59 +0,0 @@ -"""Render real provider Thinking and Text deltas from Harness events.""" - -import asyncio - -from ejagent import ( - AgentEvent, - AgentFinished, - AssistantTextDelta, - AssistantThinkingDelta, - BaseAgent, - MessageCompleted, - ModelConfig, - OpenAIModelAdapter, -) - - -class StreamingConsoleSink: - def __init__(self) -> None: - self.received_delta = False - self.in_thinking = False - - async def emit(self, event: AgentEvent) -> None: - payload = event.payload - if isinstance(payload, AssistantThinkingDelta): - if not self.in_thinking: - self.in_thinking = True - print("[thinking] ", end="", flush=True) - print(payload.delta, end="", flush=True) - elif isinstance(payload, AssistantTextDelta): - if self.in_thinking: - self.in_thinking = False - print("\n[answer] ", end="", flush=True) - self.received_delta = True - print(payload.delta, end="", flush=True) - elif isinstance(payload, MessageCompleted) and self.received_delta: - print() - elif isinstance(payload, AgentFinished): - print(f"\nfinished: {payload.result.status} / {payload.result.stop_reason}") - - -async def main() -> None: - sink = StreamingConsoleSink() - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="streaming-demo", - system_prompt="Answer concisely and stream the response normally.", - event_sink=sink, - ) - - try: - result = await agent.run(task="你有什么能力") - if not result.succeeded: - result.raise_for_status() - finally: - await agent.shutdown() - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/12_tool_progress.py b/examples/12_tool_progress.py deleted file mode 100644 index 87edc3e..0000000 --- a/examples/12_tool_progress.py +++ /dev/null @@ -1,115 +0,0 @@ -"""Stream real tool execution progress through Harness events.""" - -import asyncio -from collections.abc import Mapping -from typing import Any - -from ejagent import ( - AgentEvent, - AgentFinished, - BaseAgent, - CancellationToken, - MethodToolHandler, - ModelConfig, - OpenAIModelAdapter, - RuntimePolicy, - StepOutcome, - ToolCompleted, - ToolControl, - ToolProgressed, - ToolProgressReporter, - ToolProgressUpdate, - ToolStarted, -) - -INDEX_TOOL = { - "type": "function", - "function": { - "name": "index_project", - "description": ( - "Index the demo project and complete the task. This tool must be " - "called exactly once." - ), - "parameters": {"type": "object", "properties": {}}, - }, -} - - -class ProjectTools(MethodToolHandler): - def __init__(self) -> None: - super().__init__((INDEX_TOOL,)) - - async def do_index_project( - self, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - progress: ToolProgressReporter | None = None, - ) -> StepOutcome: - if progress is None: - raise RuntimeError("tool progress reporter is unavailable") - - updates = ( - ToolProgressUpdate( - "discovering files", - {"completed": 1, "total": 3}, - ), - ToolProgressUpdate( - "parsing Python modules", - {"completed": 2, "total": 3}, - ), - ToolProgressUpdate( - "building symbol index", - {"completed": 3, "total": 3}, - ), - ) - for update in updates: - await progress.report(update) - await asyncio.sleep(0.2) - - return StepOutcome( - {"status": "indexed", "files": 24, "symbols": 137}, - control=ToolControl.COMPLETE, - ) - - -class ProgressConsoleSink: - async def emit(self, event: AgentEvent) -> None: - payload = event.payload - if isinstance(payload, ToolStarted): - print(f"tool started: {payload.tool_call.name}") - elif isinstance(payload, ToolProgressed): - print(f" progress: {payload.update.message} {payload.update.data}") - elif isinstance(payload, ToolCompleted): - print(f"tool completed: {payload.tool_call.name}") - elif isinstance(payload, AgentFinished): - print(f"finished: {payload.result.status} / {payload.result.stop_reason}") - - -async def main() -> None: - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="tool-progress-demo", - system_prompt=( - "You demonstrate tool progress. Always call index_project exactly " - "once for the user's request. Never answer with plain text." - ), - handlers=[ProjectTools()], - runtime_policy=RuntimePolicy( - max_steps=3, - max_no_tool_responses=2, - require_explicit_finish=True, - ), - event_sink=ProgressConsoleSink(), - ) - - try: - result = await agent.run(task="Index the demo project now.") - if not result.succeeded: - result.raise_for_status() - finally: - await agent.shutdown() - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/13_usage_budget.py b/examples/13_usage_budget.py deleted file mode 100644 index a111b23..0000000 --- a/examples/13_usage_budget.py +++ /dev/null @@ -1,105 +0,0 @@ -"""Observe real Provider usage and stop before an over-budget follow-up.""" - -import asyncio -import os -from collections.abc import Mapping -from typing import Any - -from ejagent import ( - AgentEvent, - AgentFinished, - BaseAgent, - CancellationToken, - MessageCompleted, - MethodToolHandler, - ModelConfig, - OpenAIModelAdapter, - RuntimePolicy, - StepOutcome, -) - -INSPECT_TOOL = { - "type": "function", - "function": { - "name": "inspect_usage_demo", - "description": ( - "Inspect the usage-budget demo once. Return the observation to " - "the model so it would normally need another turn." - ), - "parameters": {"type": "object", "properties": {}}, - }, -} - - -class UsageTools(MethodToolHandler): - def __init__(self) -> None: - super().__init__((INSPECT_TOOL,)) - - async def do_inspect_usage_demo( - self, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - print("tool executed: inspect_usage_demo") - return StepOutcome( - { - "status": "inspected", - "next_action": "explain the observation in another turn", - } - ) - - -class UsageConsoleSink: - async def emit(self, event: AgentEvent) -> None: - payload = event.payload - if isinstance(payload, MessageCompleted): - if payload.usage is None: - print(f"turn {payload.turn} usage: unavailable") - else: - print( - f"turn {payload.turn} usage: " - f"input={payload.usage.input_tokens}, " - f"output={payload.usage.output_tokens}, " - f"total={payload.usage.total_tokens}" - ) - elif isinstance(payload, AgentFinished): - usage = payload.result.usage - print(f"finished: {payload.result.status} / {payload.result.stop_reason}") - print( - f"run usage: total={usage.total_tokens}, " - f"requests={usage.reported_request_count}/" - f"{usage.request_count}, complete={usage.complete}" - ) - - -async def main() -> None: - max_run_tokens = int(os.getenv("HARNESS_MAX_RUN_TOKENS", "1")) - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="usage-budget-demo", - system_prompt=( - "Call inspect_usage_demo exactly once. Wait for its result before " - "producing any explanation. Do not call it more than once." - ), - handlers=[UsageTools()], - runtime_policy=RuntimePolicy( - max_steps=4, - max_run_tokens=max_run_tokens, - require_explicit_finish=True, - ), - event_sink=UsageConsoleSink(), - ) - - try: - result = await agent.run(task="Inspect the usage budget demo.") - finally: - await agent.shutdown() - - print(f"configured max_run_tokens: {max_run_tokens}") - if result.error: - print(f"guard detail: {result.error}") - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/14_context_pressure.py b/examples/14_context_pressure.py deleted file mode 100644 index abce0d3..0000000 --- a/examples/14_context_pressure.py +++ /dev/null @@ -1,180 +0,0 @@ -"""Evaluate context pressure and prepare compaction with a real Provider.""" - -import asyncio -import os -from collections.abc import Mapping -from typing import Any - -from ejagent import ( - AgentEvent, - BaseAgent, - CancellationToken, - CompactionPolicy, - ContextBudget, - ContextPressureEvaluated, - MethodToolHandler, - ModelConfig, - OpenAIModelAdapter, - StepOutcome, -) - -READ_TOOL = { - "type": "function", - "function": { - "name": "read_demo_file", - "description": "Read one synthetic file used by the context demo.", - "parameters": { - "type": "object", - "properties": {"path": {"type": "string"}}, - "required": ["path"], - }, - }, -} - - -class DemoTools(MethodToolHandler): - def __init__(self) -> None: - super().__init__((READ_TOOL,)) - - async def do_read_demo_file( - self, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - return StepOutcome( - { - "path": arguments["path"], - "content": "A fresh synthetic file has no warnings.", - } - ) - - -OLD_TOOL_OUTPUT = "\n".join( - f"diagnostic line {index}: legacy warning" for index in range(80) -) - -HISTORY = [ - { - "role": "user", - "content": "Inspect the old synthetic report.", - }, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "old-call-1", - "type": "function", - "function": { - "name": "read_demo_file", - "arguments": '{"path":"legacy-report.txt"}', - }, - } - ], - "usage": { - "input_tokens": 180, - "output_tokens": 20, - "total_tokens": 200, - "cache_read_tokens": None, - "cache_write_tokens": None, - "reasoning_tokens": None, - }, - }, - { - "role": "tool", - "tool_call_id": "old-call-1", - "content": OLD_TOOL_OUTPUT, - }, - { - "role": "assistant", - "content": "The legacy report contains repeated warnings.", - "usage": { - "input_tokens": 900, - "output_tokens": 18, - "total_tokens": 918, - "cache_read_tokens": None, - "cache_write_tokens": None, - "reasoning_tokens": None, - }, - }, - { - "role": "user", - "content": "Remember that finding for the next check.", - }, - { - "role": "assistant", - "content": "I will retain the legacy finding as context.", - "usage": { - "input_tokens": 950, - "output_tokens": 16, - "total_tokens": 966, - "cache_read_tokens": None, - "cache_write_tokens": None, - "reasoning_tokens": None, - }, - }, -] - - -class ContextConsoleSink: - async def emit(self, event: AgentEvent) -> None: - payload = event.payload - if not isinstance(payload, ContextPressureEvaluated): - return - - estimate = payload.decision.estimate - print( - f"turn {payload.turn} context: total={estimate.total_tokens}, " - f"reported={estimate.reported_tokens}, " - f"trailing={estimate.trailing_tokens}, " - f"heuristic={estimate.heuristic_tokens}, " - f"source={estimate.source}" - ) - print( - f"threshold={payload.decision.threshold_tokens}, " - f"should_compact={payload.decision.should_compact}" - ) - if payload.preparation is not None: - preparation = payload.preparation - roles = [ - message.get("role") for message in preparation.messages_to_summarize - ] - print( - f"preparation: can_compact={preparation.can_compact}, " - f"first_kept_index={preparation.first_kept_index}, " - f"summarize_roles={roles}" - ) - - -async def main() -> None: - budget = ContextBudget( - context_window=int(os.getenv("HARNESS_CONTEXT_WINDOW", "600")), - reserve_tokens=int(os.getenv("HARNESS_CONTEXT_RESERVE", "100")), - keep_recent_tokens=int(os.getenv("HARNESS_KEEP_RECENT_TOKENS", "40")), - ) - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id="context-pressure-demo", - system_prompt=( - "This is a context-pressure demonstration. Answer the latest " - "question in one short sentence without calling a tool." - ), - handlers=[DemoTools()], - compaction_policy=CompactionPolicy(budget), - event_sink=ContextConsoleSink(), - ) - agent.reset(history=HISTORY) - - try: - result = await agent.run(task="Does the old report contain repeated warnings?") - finally: - await agent.shutdown() - - print(f"finished: {result.status} / {result.stop_reason}") - print(f"output: {result.output}") - print("history was evaluated but was not compacted or mutated") - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/15_explicit_compaction.py b/examples/15_explicit_compaction.py deleted file mode 100644 index 2353b40..0000000 --- a/examples/15_explicit_compaction.py +++ /dev/null @@ -1,169 +0,0 @@ -"""Explicitly compact history with a real Provider-backed Compactor.""" - -import asyncio -import json -import os - -from ejagent import ( - AgentEvent, - BaseAgent, - CompactionCompleted, - CompactionFailed, - CompactionPolicy, - CompactionRequest, - CompactionStarted, - CompositeAgentEventSink, - ContextBudget, - ContextBuildResult, - MemorySessionStorage, - ModelCompactor, - ModelConfig, - OpenAIModelAdapter, - SessionRecorder, -) - - -def build_compaction_context(request: CompactionRequest) -> ContextBuildResult: - """Application-owned prompt policy injected into ModelCompactor.""" - - visible_messages = [ - { - key: value - for key, value in message.items() - if key not in {"usage", "_ejagent_summary", "_simagentplg_summary"} - } - for message in request.preparation.messages_to_summarize - ] - previous = ( - request.previous_summary.content - if request.previous_summary is not None - else "(none)" - ) - prompt = ( - "Create a concise continuation summary. Preserve user goals, " - "tool findings, decisions, and unfinished work. Do not invent " - "facts. Return only the summary.\n\n" - f"Previous summary:\n{previous}\n\n" - "New messages:\n" + json.dumps(visible_messages, ensure_ascii=False, indent=2) - ) - messages = ( - { - "role": "system", - "content": "You produce reliable Agent context summaries.", - }, - {"role": "user", "content": prompt}, - ) - return ContextBuildResult( - agent_messages=messages, - llm_messages=messages, - tools=(), - ) - - -class CompactionConsoleSink: - async def emit(self, event: AgentEvent) -> None: - payload = event.payload - if isinstance(payload, CompactionStarted): - preparation = payload.request.preparation - print( - "compaction started: " - f"summarize={len(preparation.messages_to_summarize)}, " - f"keep={len(preparation.messages_to_keep)}" - ) - elif isinstance(payload, CompactionCompleted): - print(f"compaction completed: {payload.result.status}") - elif isinstance(payload, CompactionFailed): - print( - f"compaction failed: {payload.result.status}, " - f"error={payload.result.error}" - ) - - -OLD_TOOL_OUTPUT = "\n".join( - f"diagnostic line {index}: repeated legacy warning" for index in range(80) -) - -HISTORY = [ - {"role": "user", "content": "Inspect the legacy report."}, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "legacy-call", - "type": "function", - "function": { - "name": "read_report", - "arguments": '{"path":"legacy.txt"}', - }, - } - ], - }, - { - "role": "tool", - "tool_call_id": "legacy-call", - "content": OLD_TOOL_OUTPUT, - }, - { - "role": "assistant", - "content": "The legacy report contains repeated warnings.", - }, - { - "role": "user", - "content": "Keep that finding available for the next question.", - }, - { - "role": "assistant", - "content": "The finding will remain available.", - }, -] - - -async def main() -> None: - model = OpenAIModelAdapter(ModelConfig.from_env()) - storage = MemorySessionStorage() - recorder = SessionRecorder( - session_id="explicit-compaction-demo", - storage=storage, - ) - agent = BaseAgent( - model, - agent_id="explicit-compaction-demo", - system_prompt=( - "Answer using the retained conversation summary and recent history." - ), - compaction_policy=CompactionPolicy( - ContextBudget( - context_window=int(os.getenv("HARNESS_CONTEXT_WINDOW", "4096")), - reserve_tokens=int(os.getenv("HARNESS_CONTEXT_RESERVE", "512")), - keep_recent_tokens=int(os.getenv("HARNESS_KEEP_RECENT_TOKENS", "20")), - ) - ), - compactor=ModelCompactor( - model, - context_builder=build_compaction_context, - source=f"openai-compatible:{model.config.model}", - ), - event_sink=CompositeAgentEventSink([recorder, CompactionConsoleSink()]), - ) - agent.reset(history=HISTORY) - - try: - compacted = await agent.compact() - print(f"summary source: {compacted.summary.source}") - print(f"state roles after compact: {[m['role'] for m in agent.messages]}") - result = await agent.run( - task="What did the legacy report contain? Reply in one sentence." - ) - session = await recorder.load() - finally: - await agent.shutdown() - - print(f"agent answer: {result.output}") - assert session is not None - print(f"session compactions: {len(session.compactions)}") - print(f"session roles: {[message['role'] for message in session.messages]}") - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/16_durable_session.py b/examples/16_durable_session.py deleted file mode 100644 index 4710667..0000000 --- a/examples/16_durable_session.py +++ /dev/null @@ -1,91 +0,0 @@ -"""Persist a Session to disk and resume it in a separate invocation.""" - -import argparse -import asyncio -import os -from pathlib import Path - -from ejagent import ( - BaseAgent, - JsonlSessionStorage, - ModelConfig, - OpenAIModelAdapter, - SessionRecorder, -) - -SESSION_ID = "durable-session-demo" -AGENT_ID = "durable-session-agent" -SYSTEM_PROMPT = ( - "Preserve user-provided facts exactly and answer only from durable " - "conversation history." -) - - -async def record(storage: JsonlSessionStorage) -> None: - recorder = SessionRecorder(session_id=SESSION_ID, storage=storage) - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id=AGENT_ID, - system_prompt=SYSTEM_PROMPT, - event_sink=recorder, - ) - try: - result = await agent.run( - task=( - "Remember this exact project code and reply only ACK: SIM-SESSION-2048" - ) - ) - print(f"record response: {result.output}") - finally: - await agent.shutdown() - - saved = await storage.load(SESSION_ID) - if saved is None: - raise RuntimeError("durable Session was not saved") - print(f"saved runs: {len(saved.runs)}") - - -async def resume(storage: JsonlSessionStorage) -> None: - saved = await storage.load(SESSION_ID) - if saved is None: - raise RuntimeError("run the record command before resume") - - recorder = SessionRecorder(session_id=SESSION_ID, storage=storage) - agent = BaseAgent( - OpenAIModelAdapter(ModelConfig.from_env()), - agent_id=AGENT_ID, - system_prompt=SYSTEM_PROMPT, - event_sink=recorder, - ) - agent.restore_session(saved) - try: - result = await agent.run( - task="Return only the exact project code stored previously." - ) - print(f"resume response: {result.output}") - finally: - await agent.shutdown() - - updated = await storage.load(SESSION_ID) - if updated is None: - raise RuntimeError("resumed Session was not saved") - print(f"saved runs after resume: {len(updated.runs)}") - - -async def main() -> None: - parser = argparse.ArgumentParser() - parser.add_argument("command", choices=("record", "resume")) - parser.add_argument( - "--session-dir", - default=os.getenv("EJAGENT_SESSION_DIR", ".ejagent-sessions"), - ) - args = parser.parse_args() - storage = JsonlSessionStorage(Path(args.session_dir)) - if args.command == "record": - await record(storage) - else: - await resume(storage) - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/17_session_tree.py b/examples/17_session_tree.py deleted file mode 100644 index 6c7ed59..0000000 --- a/examples/17_session_tree.py +++ /dev/null @@ -1,65 +0,0 @@ -"""Inspect and create durable Session branches without replaying a model.""" - -import argparse -import asyncio -import os -from pathlib import Path - -from ejagent import JsonlSessionStorage - - -async def main() -> None: - parser = argparse.ArgumentParser() - parser.add_argument("--session-id", default="durable-session-demo") - parser.add_argument( - "--session-dir", - default=os.getenv("EJAGENT_SESSION_DIR", ".ejagent-sessions"), - ) - commands = parser.add_subparsers(dest="command", required=True) - commands.add_parser("branches") - - fork_parser = commands.add_parser("fork") - fork_parser.add_argument("branch_id") - fork_parser.add_argument("--from-record-id") - - rollback_parser = commands.add_parser("rollback") - rollback_parser.add_argument("to_record_id") - rollback_parser.add_argument("branch_id") - - retry_parser = commands.add_parser("retry") - retry_parser.add_argument("run_id") - retry_parser.add_argument("branch_id") - - args = parser.parse_args() - storage = JsonlSessionStorage(Path(args.session_dir)) - if args.command == "branches": - for branch in await storage.list_branches(args.session_id): - print( - f"{branch.branch_id}: head={branch.head_record_id} " - f"intent={branch.intent or 'main'}" - ) - elif args.command == "fork": - checkout = await storage.fork( - args.session_id, - from_record_id=args.from_record_id, - branch_id=args.branch_id, - ) - print(f"created {checkout.branch.branch_id} at {checkout.head.record_id}") - elif args.command == "rollback": - checkout = await storage.rollback( - args.session_id, - to_record_id=args.to_record_id, - branch_id=args.branch_id, - ) - print(f"created {checkout.branch.branch_id} at {checkout.head.record_id}") - else: - retry = await storage.prepare_retry( - args.session_id, - run_id=args.run_id, - branch_id=args.branch_id, - ) - print(f"prepared {retry.checkout.branch.branch_id}; rerun task: {retry.task}") - - -if __name__ == "__main__": - asyncio.run(main()) diff --git a/examples/README.md b/examples/README.md index 780a070..c8ab18b 100644 --- a/examples/README.md +++ b/examples/README.md @@ -1,97 +1,50 @@ # EJAgent Core Examples -All examples use the environment variables documented in the project README. -Copy `.env.example` to `.env` and fill in credentials for an OpenAI-compatible -provider before running them. Examples `07` through `16` exercise the Harness -against the real `OpenAIModelAdapter`; they do not use scripted model results. - -Run an example from the repository root: +Copy `.env.example` to `.env` and configure a Provider. +Run examples from the repository root, for example: ```bash uv run python examples/01_stateful_chat.py ``` -Every `BaseAgent` declares its own immutable `agent_id`. +All examples use the `AgentHarness`/`RuntimeKernel` path and make real Provider +requests. + +## Core Examples -## Examples +- `01_stateful_chat.py`: typed Conversation state across two committed Runs. +- `02_custom_tool.py`: a provider-neutral Tool registered through + `FunctionToolExecutor`. +- `03_anthropic_chat.py`: the same Harness contract over Anthropic Messages. +- `04_mcp_tools.py`: opt-in MCP integration through `McpToolExecutor`. +- `05_skill.py`: disposable skill indexing and explicit instruction projection. +- `06_session_resume.py`: recovery from `MemorySessionStore` in a new Harness. +- `07_durable_session.py`: cross-process recovery from `JsonlSessionStore`. -- `01_stateful_chat.py`: plain chat, conversation memory, and `reset()` -- `02_custom_tool.py`: an OpenAI-first canonical `ToolDefinition` registered - through `MethodToolHandler` -- `04_mcp_tools.py`: opt-in MCP integration with a custom config file -- `06_skill.py`: local skill discovery, indexing, template, and sample injection -- `07_event_observers.py`: ordered lifecycle events and multiple independent - sinks -- `08_session_resume.py`: linear Session recording and recovery in a new agent -- `09_runtime_control.py`: `abort()`, `wait_for_idle()`, terminal events, and reuse -- `10_composed_harness.py`: Model Adapter, tools, Middleware, events, Session, - runtime policy, and explicit completion composed as one Harness -- `11_streaming_events.py`: real provider Thinking and Text Delta events - rendered as they arrive while the final message remains atomically committed -- `12_tool_progress.py`: a real provider-triggered tool reports structured, - ordered progress before its final result is committed -- `13_usage_budget.py`: real Provider Usage is normalized, aggregated, and used - to stop before an intentionally over-budget follow-up request -- `14_context_pressure.py`: the complete real Provider request is assessed, - old full turns are prepared for compaction, and current history remains - unchanged -- `15_explicit_compaction.py`: a real Provider-backed Compactor atomically - replaces old Tool history with a Summary Entry, then the Agent and Session - continue from the compacted projection -- `16_durable_session.py`: append a versioned JSONL Session Journal and restore - it in a separate process invocation -- `17_session_tree.py`: inspect branches and prepare Fork, Rollback, or Retry - branches without automatically executing a model +Every `AgentHarness` has one immutable `agent_id`. Tool-enabled Harnesses expose +only their configured `ToolExecutor.definitions`; plain assistant text completes +a Run by default. A function Tool may return +`ToolExecutionResult(..., control=ToolControl.COMPLETE)` to finish explicitly. -Run the composed Harness example directly: +The MCP example requires the commands in `mcp_config.json`; the sample starts +Playwright MCP through `npx`. Install optional support first: ```bash -uv run python examples/10_composed_harness.py +uv sync --extra mcp ``` -The Harness examples make real provider requests, so response text and exact -timing depend on the configured model. Their control flow, event ordering, -Session projection, and terminal result protocols remain deterministic. - -Tool-enabled agents expose only the handlers explicitly passed to -`BaseAgent`. Tool availability does not force an explicit completion call; -plain text completes a task by default. A derived agent can require explicit -completion through `RuntimePolicy` and return -`StepOutcome(..., control=ToolControl.COMPLETE)` from its own completion tool. -The repeated-tool-call guard remains configurable through the same policy. -Example `13` defaults `HARNESS_MAX_RUN_TOKENS` to `1` so its first reported -response exhausts the budget after the requested tool settles. Set a larger -value to observe additional turns. -Example `14` uses `HARNESS_CONTEXT_WINDOW`, `HARNESS_CONTEXT_RESERVE`, and -`HARNESS_KEEP_RECENT_TOKENS` only as Harness-side demonstration values. The -default low threshold makes the seeded old tool output produce a compaction -suggestion; it does not change the configured model's real context window. -Example `15` uses the configured Provider twice: once through an example-owned -summary prompt and once for normal Agent continuation. The Core supplies the -protocol and atomic state transition but does not own that prompt. +The Anthropic example requires `uv sync --extra anthropic` plus +`ANTHROPIC_MODEL` and `ANTHROPIC_API_KEY`. -The MCP example requires the commands declared in `mcp_config.json` to be -available locally. Its sample configuration starts the Playwright MCP server -through `npx`. +The Skills example loads `examples/skills/release_notes/`. The catalog index and +explicitly selected `SKILL.md`, template, and sample are projected into a +disposable `ContextView`; they are never committed to Conversation. -The skill example loads `examples/skills/release_notes/`. Skill metadata, -including its file location, is injected into the model context. Naming the -skill explicitly loads its full `SKILL.md`, optional `template.md`, and optional -`examples/sample.md` content. The core does not register a special skill tool. +For durable recovery, run: -The event examples use `CompositeAgentEventSink` to attach observers without -changing the Agent Loop. The Session examples use `SessionRecorder` and -`MemorySessionStorage`; only real user, assistant, and tool messages are saved. -Tool Progress is observation-only and is not written to Agent state, Session, -or the model transcript. -The runtime-control example distinguishes external `abort()` from -`ToolControl.CANCEL` and waits until terminal event sinks have settled before -reporting idle. It waits briefly before aborting the provider request; set -`HARNESS_ABORT_DELAY` to adjust that delay. +```bash +uv run python examples/07_durable_session.py record +uv run python examples/07_durable_session.py resume +``` -`16_durable_session.py` appends a versioned JSONL Session Journal. Run it once -with `record` and again with `resume`; the second invocation creates a new Agent -process and restores the saved conversation through `restore_session()`. -`17_session_tree.py` operates on that same journal. Its `branches`, `fork`, -`rollback`, and `retry` commands expose tree management while leaving model -execution explicit. +Conversation recovery and append-only Run Audit are separate Store domains. diff --git a/pi b/pi deleted file mode 160000 index f4e9ca7..0000000 --- a/pi +++ /dev/null @@ -1 +0,0 @@ -Subproject commit f4e9ca7466b5576090d1093c27fe38d73909f3d2 diff --git a/pyproject.toml b/pyproject.toml index ba2cb2d..ccd8c06 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -6,11 +6,9 @@ readme = "README.md" license = {file = "LICENSE"} requires-python = ">=3.12" dependencies = [ - "jsonschema>=4.26.0", "openai>=2.41.0", "python-dotenv>=1.2.2", "pyyaml>=6.0.3", - "referencing>=0.37.0", ] [project.urls] @@ -19,6 +17,9 @@ Repository = "https://github.com/jyh20030112/EJAgent" Issues = "https://github.com/jyh20030112/EJAgent/issues" [project.optional-dependencies] +anthropic = [ + "anthropic>=0.104.1", +] mcp = [ "fastmcp>=3.4.2", ] @@ -27,7 +28,6 @@ mcp = [ dev = [ "mypy>=1.18.2", "ruff>=0.14.0", - "types-jsonschema>=4.26.0.20260518", "types-pyyaml>=6.0.12.20250915", ] @@ -38,12 +38,6 @@ build-backend = "hatchling.build" [tool.hatch.build.targets.wheel] packages = ["src/ejagent"] -[tool.hatch.build.targets.sdist] -exclude = [ - "/.gitmodules", - "/pi", -] - [tool.ruff] target-version = "py312" line-length = 88 diff --git a/src/ejagent/__init__.py b/src/ejagent/__init__.py index 58711a9..6928aac 100644 --- a/src/ejagent/__init__.py +++ b/src/ejagent/__init__.py @@ -1,341 +1,43 @@ -"""Composable stateful agents with tool handlers and MCP integration. +"""Composable runtime kernel and lifecycle harness for one logical agent.""" -.. note:: - - This package provides the core agent framework — see :class:`BaseAgent` for - the primary entry point. -""" - -from ejagent.agent.base import BaseAgent -from ejagent.agent.behavior import ( - BehaviorAction, - BehaviorDecision, - BehaviorHook, - BehaviorHookError, - TurnSnapshot, -) -from ejagent.agent.cancellation import ( - AgentCancelledError, - CancellationSource, - CancellationToken, -) -from ejagent.agent.compaction import ( - CompactionRequest, - CompactionResult, - CompactionRuntime, - CompactionStatus, - CompactionTrigger, - Compactor, - CompactorOutput, - SummaryEntry, -) -from ejagent.agent.context_builder import AgentContextBuilder, ContextBuildResult -from ejagent.agent.context_management import ( - AutoCompactionPolicy, - CompactionDecision, - CompactionDecisionReason, - CompactionPolicy, - CompactionPreparation, - ContextBudget, - ContextUsageEstimate, - ContextUsageSource, - HeuristicMessageTokenEstimator, - MessageTokenEstimator, - estimate_context_usage, - prepare_compaction, -) -from ejagent.agent.control import ( - ContinueRejectedError, - ContinueRejectedReason, - ControlInput, - ControlInputKind, - ControlReceipt, - ControlStatus, - FollowUpDiscardedError, - FollowUpDiscardReason, - FollowUpError, - FollowUpFailurePolicy, - FollowUpHandle, - FollowUpRejectedError, -) -from ejagent.agent.events import ( - AgentContinued, - AgentEvent, - AgentEventKind, - AgentEventPayload, - AgentEventSink, - AgentEventSinkError, - AgentFinished, - AgentStarted, - AssistantTextDelta, - AssistantThinkingDelta, - CompactionCompleted, - CompactionFailed, - CompactionStarted, - CompositeAgentEventSink, - ContextPressureEvaluated, - MessageCompleted, - SteeringApplied, - SteeringDiscarded, - ToolCompleted, - ToolProgressed, - ToolStarted, - TurnCompleted, - TurnStarted, -) -from ejagent.agent.model_compactor import ( - CompactionContextBuilder, - ModelCompactor, -) -from ejagent.agent.orchestrator import AgentOrchestrator -from ejagent.agent.result import ( - AgentRunError, - AgentRunResult, - RunStatus, - StopReason, -) -from ejagent.agent.runtime_policy import RuntimePolicy -from ejagent.agent.state import AgentState, AgentStatus -from ejagent.agent.types import ( - StepOutcome, - ToolCallResult, - ToolControl, - ToolProgressReporter, - ToolProgressUpdate, -) -from ejagent.agent.usage import RunUsage -from ejagent.handlers import ( - BaseHandler, - McpToolHandler, - MethodToolHandler, - ToolDefinition, - ToolDefinitionError, - ToolEffect, - UnknownToolError, -) -from ejagent.middleware import ( - Middleware, - RuleBasedToolPolicy, - ToolApprovalDecision, - ToolApprovalRequest, - ToolApprover, - ToolCallContext, - ToolExecutionPolicy, - ToolMiddleware, - ToolNext, - ToolPolicyAction, - ToolPolicyDecision, - ToolPolicyMiddleware, - ToolPolicyPredicate, - ToolPolicyRule, - ToolSchemaConfigurationError, - ToolSchemaValidationMiddleware, - compose_tool_middlewares, - format_tool_call_preview, +from ejagent.context import ( + DerivedCompactionPipeline, + IdentityContextPipeline, + SkillsContextPipeline, ) -from ejagent.plugins.mcp.mcp_manager import McpServerManager -from ejagent.plugins.skill.skill_manager import SkillManager +from ejagent.harness import AgentHarness, MemorySessionStore +from ejagent.kernel import RuntimeKernel from ejagent.providers import ( - AssistantMessage, - ContextOverflowError, - ModelAdapter, - ModelAuthenticationError, + AnthropicConfig, + AnthropicModelPort, ModelConfig, - ModelErrorKind, - ModelProviderError, - ModelRateLimitError, - ModelResponseCompleted, - ModelStreamEvent, - ModelTextDelta, - ModelThinkingDelta, - ModelTimeoutError, - ModelToolCall, - ModelUsage, - OpenAIModelAdapter, + OpenAIModelPort, ) -from ejagent.session import ( - DEFAULT_SESSION_BRANCH, - SESSION_JOURNAL_SCHEMA_VERSION, - SESSION_SCHEMA_VERSION, - AgentSession, - JsonlSessionStorage, - MemorySessionStorage, - SessionBranch, - SessionBranchIntent, - SessionCheckout, - SessionCompaction, - SessionConflictError, - SessionError, - SessionJournalStorage, - SessionLockTimeoutError, - SessionMessage, - SessionRecord, - SessionRecordDraft, - SessionRecorder, - SessionRecordKind, - SessionRetry, - SessionRun, - SessionRunIntent, - SessionSerializationError, - SessionStorage, - SessionStorageError, - SessionTreeStorage, - session_from_dict, - session_to_dict, +from ejagent.skills import Skill, SkillCatalog +from ejagent.storage import JsonlSessionStore +from ejagent.tools import ( + CompositeToolExecutor, + FunctionTool, + FunctionToolExecutor, + McpToolExecutor, ) __all__ = [ - "BaseAgent", - "BehaviorAction", - "BehaviorDecision", - "BehaviorHook", - "BehaviorHookError", - "TurnSnapshot", - "AgentCancelledError", - "CancellationSource", - "CancellationToken", - "Compactor", - "CompactorOutput", - "CompactionRequest", - "CompactionResult", - "CompactionRuntime", - "CompactionStatus", - "CompactionTrigger", - "ControlInputKind", - "ContinueRejectedReason", - "ContinueRejectedError", - "ControlStatus", - "ControlInput", - "ControlReceipt", - "FollowUpFailurePolicy", - "FollowUpDiscardReason", - "FollowUpHandle", - "FollowUpError", - "FollowUpRejectedError", - "FollowUpDiscardedError", - "SummaryEntry", - "CompactionContextBuilder", - "ModelCompactor", + "AgentHarness", + "AnthropicConfig", + "AnthropicModelPort", + "CompositeToolExecutor", + "DerivedCompactionPipeline", + "FunctionTool", + "FunctionToolExecutor", + "IdentityContextPipeline", + "JsonlSessionStore", + "McpToolExecutor", + "MemorySessionStore", "ModelConfig", - "ModelAdapter", - "OpenAIModelAdapter", - "AssistantMessage", - "ModelErrorKind", - "ModelProviderError", - "ContextOverflowError", - "ModelRateLimitError", - "ModelTimeoutError", - "ModelAuthenticationError", - "ModelStreamEvent", - "ModelTextDelta", - "ModelThinkingDelta", - "ModelResponseCompleted", - "ModelToolCall", - "ModelUsage", - "AgentState", - "AgentStatus", - "AgentContextBuilder", - "ContextBuildResult", - "ContextBudget", - "ContextUsageEstimate", - "ContextUsageSource", - "MessageTokenEstimator", - "HeuristicMessageTokenEstimator", - "CompactionPolicy", - "AutoCompactionPolicy", - "CompactionDecision", - "CompactionDecisionReason", - "CompactionPreparation", - "estimate_context_usage", - "prepare_compaction", - "AgentEvent", - "AgentEventKind", - "AgentEventPayload", - "AgentEventSink", - "AgentEventSinkError", - "CompositeAgentEventSink", - "CompactionStarted", - "CompactionCompleted", - "CompactionFailed", - "ContextPressureEvaluated", - "AgentStarted", - "AgentContinued", - "AssistantThinkingDelta", - "AssistantTextDelta", - "TurnStarted", - "MessageCompleted", - "SteeringApplied", - "SteeringDiscarded", - "ToolStarted", - "ToolProgressed", - "ToolCompleted", - "TurnCompleted", - "AgentFinished", - "AgentOrchestrator", - "RuntimePolicy", - "AgentRunResult", - "AgentRunError", - "RunStatus", - "StopReason", - "RunUsage", - "StepOutcome", - "ToolCallResult", - "ToolControl", - "ToolProgressReporter", - "ToolProgressUpdate", - "Middleware", - "ToolMiddleware", - "ToolCallContext", - "ToolNext", - "ToolExecutionPolicy", - "RuleBasedToolPolicy", - "ToolPolicyRule", - "ToolPolicyAction", - "ToolPolicyDecision", - "ToolPolicyMiddleware", - "ToolPolicyPredicate", - "ToolApprover", - "ToolApprovalRequest", - "ToolApprovalDecision", - "ToolSchemaConfigurationError", - "ToolSchemaValidationMiddleware", - "compose_tool_middlewares", - "format_tool_call_preview", - "BaseHandler", - "MethodToolHandler", - "McpToolHandler", - "ToolDefinition", - "ToolDefinitionError", - "ToolEffect", - "UnknownToolError", - "McpServerManager", - "SkillManager", - "AgentSession", - "SessionMessage", - "SessionRun", - "SessionRunIntent", - "SessionCompaction", - "SessionStorage", - "SessionJournalStorage", - "SessionTreeStorage", - "JsonlSessionStorage", - "MemorySessionStorage", - "SessionRecorder", - "SESSION_SCHEMA_VERSION", - "session_to_dict", - "session_from_dict", - "SessionError", - "SessionConflictError", - "SessionLockTimeoutError", - "SessionSerializationError", - "SessionStorageError", - "SESSION_JOURNAL_SCHEMA_VERSION", - "DEFAULT_SESSION_BRANCH", - "SessionRecordKind", - "SessionRecordDraft", - "SessionRecord", - "SessionBranchIntent", - "SessionBranch", - "SessionCheckout", - "SessionRetry", + "OpenAIModelPort", + "RuntimeKernel", + "Skill", + "SkillCatalog", + "SkillsContextPipeline", ] diff --git a/src/ejagent/agent/__init__.py b/src/ejagent/agent/__init__.py deleted file mode 100644 index 0b8de8f..0000000 --- a/src/ejagent/agent/__init__.py +++ /dev/null @@ -1,223 +0,0 @@ -"""Agent runtime primitives.""" - -from ejagent.agent.base import BaseAgent -from ejagent.agent.behavior import ( - BehaviorAction, - BehaviorDecision, - BehaviorHook, - BehaviorHookError, - TurnSnapshot, -) -from ejagent.agent.cancellation import ( - AgentCancelledError, - CancellationSource, - CancellationToken, -) -from ejagent.agent.compaction import ( - CompactionRequest, - CompactionResult, - CompactionRuntime, - CompactionStatus, - CompactionTrigger, - Compactor, - CompactorOutput, - SummaryEntry, -) -from ejagent.agent.context_builder import AgentContextBuilder, ContextBuildResult -from ejagent.agent.context_management import ( - AutoCompactionPolicy, - CompactionDecision, - CompactionDecisionReason, - CompactionPolicy, - CompactionPreparation, - ContextBudget, - ContextUsageEstimate, - ContextUsageSource, - HeuristicMessageTokenEstimator, - MessageTokenEstimator, - estimate_context_usage, - prepare_compaction, -) -from ejagent.agent.control import ( - ContinueRejectedError, - ContinueRejectedReason, - ControlInput, - ControlInputKind, - ControlReceipt, - ControlStatus, - FollowUpDiscardedError, - FollowUpDiscardReason, - FollowUpError, - FollowUpFailurePolicy, - FollowUpHandle, - FollowUpRejectedError, -) -from ejagent.agent.events import ( - AgentContinued, - AgentEvent, - AgentEventKind, - AgentEventPayload, - AgentEventSink, - AgentEventSinkError, - AgentFinished, - AgentStarted, - AssistantTextDelta, - AssistantThinkingDelta, - CompactionCompleted, - CompactionFailed, - CompactionStarted, - CompositeAgentEventSink, - ContextPressureEvaluated, - MessageCompleted, - SteeringApplied, - SteeringDiscarded, - ToolCompleted, - ToolProgressed, - ToolStarted, - TurnCompleted, - TurnStarted, -) -from ejagent.agent.model_compactor import ( - CompactionContextBuilder, - ModelCompactor, -) -from ejagent.agent.orchestrator import AgentOrchestrator -from ejagent.agent.result import ( - AgentRunError, - AgentRunResult, - RunStatus, - StopReason, -) -from ejagent.agent.runtime_policy import RuntimePolicy -from ejagent.agent.state import AgentState, AgentStatus -from ejagent.agent.types import ( - StepOutcome, - ToolCallResult, - ToolControl, - ToolProgressReporter, - ToolProgressUpdate, -) -from ejagent.agent.usage import RunUsage -from ejagent.middleware import ( - Middleware, - RuleBasedToolPolicy, - ToolApprovalDecision, - ToolApprovalRequest, - ToolApprover, - ToolCallContext, - ToolExecutionPolicy, - ToolMiddleware, - ToolNext, - ToolPolicyAction, - ToolPolicyDecision, - ToolPolicyMiddleware, - ToolPolicyPredicate, - ToolPolicyRule, - ToolSchemaConfigurationError, - ToolSchemaValidationMiddleware, - compose_tool_middlewares, - format_tool_call_preview, -) - -__all__ = [ - "BaseAgent", - "BehaviorAction", - "BehaviorDecision", - "BehaviorHook", - "BehaviorHookError", - "TurnSnapshot", - "AgentCancelledError", - "CancellationSource", - "CancellationToken", - "Compactor", - "CompactorOutput", - "CompactionRequest", - "CompactionResult", - "CompactionRuntime", - "CompactionStatus", - "CompactionTrigger", - "SummaryEntry", - "CompactionContextBuilder", - "ModelCompactor", - "AgentState", - "AgentStatus", - "AgentContextBuilder", - "ContextBuildResult", - "ContextBudget", - "ContextUsageEstimate", - "ContextUsageSource", - "MessageTokenEstimator", - "HeuristicMessageTokenEstimator", - "CompactionPolicy", - "AutoCompactionPolicy", - "CompactionDecision", - "CompactionDecisionReason", - "CompactionPreparation", - "estimate_context_usage", - "prepare_compaction", - "ControlInputKind", - "ContinueRejectedReason", - "ContinueRejectedError", - "ControlStatus", - "ControlInput", - "ControlReceipt", - "FollowUpFailurePolicy", - "FollowUpDiscardReason", - "FollowUpHandle", - "FollowUpError", - "FollowUpRejectedError", - "FollowUpDiscardedError", - "AgentEvent", - "AgentEventKind", - "AgentEventPayload", - "AgentEventSink", - "AgentEventSinkError", - "CompositeAgentEventSink", - "CompactionStarted", - "CompactionCompleted", - "CompactionFailed", - "ContextPressureEvaluated", - "AgentStarted", - "AgentContinued", - "AssistantThinkingDelta", - "AssistantTextDelta", - "TurnStarted", - "MessageCompleted", - "SteeringApplied", - "SteeringDiscarded", - "ToolStarted", - "ToolProgressed", - "ToolCompleted", - "TurnCompleted", - "AgentFinished", - "AgentOrchestrator", - "RuntimePolicy", - "AgentRunResult", - "AgentRunError", - "RunStatus", - "StopReason", - "RunUsage", - "StepOutcome", - "ToolCallResult", - "ToolControl", - "ToolProgressReporter", - "ToolProgressUpdate", - "Middleware", - "ToolMiddleware", - "ToolCallContext", - "ToolNext", - "ToolExecutionPolicy", - "RuleBasedToolPolicy", - "ToolPolicyRule", - "ToolPolicyAction", - "ToolPolicyDecision", - "ToolPolicyMiddleware", - "ToolPolicyPredicate", - "ToolApprover", - "ToolApprovalRequest", - "ToolApprovalDecision", - "ToolSchemaConfigurationError", - "ToolSchemaValidationMiddleware", - "compose_tool_middlewares", - "format_tool_call_preview", -] diff --git a/src/ejagent/agent/base.py b/src/ejagent/agent/base.py deleted file mode 100644 index 1077302..0000000 --- a/src/ejagent/agent/base.py +++ /dev/null @@ -1,666 +0,0 @@ -from __future__ import annotations - -import asyncio -from collections.abc import Iterable, Mapping, Sequence -from dataclasses import dataclass -from pathlib import Path -from typing import TYPE_CHECKING, Any - -from ejagent.agent.behavior import BehaviorHook -from ejagent.agent.cancellation import ( - CancellationSource, -) -from ejagent.agent.compaction import ( - CompactionResult, - CompactionRuntime, - Compactor, -) -from ejagent.agent.context_builder import AgentContextBuilder -from ejagent.agent.context_management import ( - AutoCompactionPolicy, - CompactionPolicy, - MessageTokenEstimator, -) -from ejagent.agent.control import ( - ContinueRejectedError, - ContinueRejectedReason, - ControlInput, - ControlReceipt, - ControlStatus, - FollowUpDiscardReason, - FollowUpFailurePolicy, - FollowUpHandle, - _continue_history_rejection_reason, - _FollowUpQueue, - _SteeringQueue, -) -from ejagent.agent.events import ( - AgentEventEmitter, - AgentEventSink, -) -from ejagent.agent.orchestrator import AgentOrchestrator -from ejagent.agent.result import AgentRunResult -from ejagent.agent.runtime_policy import RuntimePolicy -from ejagent.agent.state import AgentState -from ejagent.agent.tool_runtime import ToolRuntime -from ejagent.agent.types import StepOutcome -from ejagent.logger import get_logger -from ejagent.middleware import ToolMiddleware -from ejagent.plugins.skill.skill_manager import SkillManager -from ejagent.providers.base import ModelAdapter - -if TYPE_CHECKING: - from ejagent.handlers.base import BaseHandler - from ejagent.handlers.definition import ToolDefinition - from ejagent.session.types import AgentSession - -DEFAULT_SYSTEM_PROMPT = "You are a helpful, concise assistant." - -TOOL_PROTOCOL_PROMPT = """ -You can call external tools when they are available. - -Tool protocol: -- Use tool calls for actions that require a registered tool. -- Wait for tool results before deciding the next action. -- Do not repeat the same ineffective tool call. -""".strip() - -EXPLICIT_FINISH_PROTOCOL_PROMPT = """ -This agent requires explicit tool completion. -Plain text does not finish the task. After completing all work, call a tool -that returns the completion control signal. -""".strip() - - -@dataclass(slots=True) -class _ActiveOperation: - cancellation: CancellationSource - steering: _SteeringQueue | None = None - - -class BaseAgent: - """Stateful agent core composed with a model adapter and tool handlers.""" - - def __init__( - self, - model: ModelAdapter, - *, - agent_id: str, - system_prompt: str = DEFAULT_SYSTEM_PROMPT, - handlers: Iterable[BaseHandler] | None = None, - middlewares: Iterable[ToolMiddleware] | None = None, - behavior_hooks: Iterable[BehaviorHook] | None = None, - skills_dir: str | Path | None = None, - context_builder: AgentContextBuilder | None = None, - compaction_policy: CompactionPolicy | None = None, - compactor: Compactor | None = None, - auto_compaction_policy: AutoCompactionPolicy | None = None, - context_token_estimator: MessageTokenEstimator | None = None, - runtime_policy: RuntimePolicy | None = None, - event_sink: AgentEventSink | None = None, - steering_queue_capacity: int = 16, - follow_up_queue_capacity: int = 16, - follow_up_failure_policy: FollowUpFailurePolicy = ( - FollowUpFailurePolicy.DISCARD - ), - ) -> None: - self._agent_id = agent_id.strip() - if not self._agent_id: - raise ValueError("agent_id must not be empty") - if isinstance(steering_queue_capacity, bool) or not isinstance( - steering_queue_capacity, int - ): - raise TypeError("steering_queue_capacity must be an integer") - if steering_queue_capacity <= 0: - raise ValueError("steering_queue_capacity must be greater than zero") - if isinstance(follow_up_queue_capacity, bool) or not isinstance( - follow_up_queue_capacity, int - ): - raise TypeError("follow_up_queue_capacity must be an integer") - if follow_up_queue_capacity <= 0: - raise ValueError("follow_up_queue_capacity must be greater than zero") - if not isinstance(follow_up_failure_policy, FollowUpFailurePolicy): - raise TypeError("follow_up_failure_policy must be a FollowUpFailurePolicy") - policy = runtime_policy or RuntimePolicy() - self.model = model - self.system_prompt = system_prompt - self.runtime_policy = policy - self.compaction_policy = compaction_policy - self.compactor = compactor - self.auto_compaction_policy = auto_compaction_policy - if auto_compaction_policy is not None and auto_compaction_policy.enabled: - if compaction_policy is None: - raise ValueError("automatic compaction requires a CompactionPolicy") - if compactor is None: - raise ValueError("automatic compaction requires a Compactor") - self.context_token_estimator = context_token_estimator - self.handlers = list(handlers or ()) - self.middlewares = list(middlewares or ()) - self.behavior_hooks = tuple(behavior_hooks or ()) - self.event_sink = event_sink - self.steering_queue_capacity = steering_queue_capacity - self.follow_up_queue_capacity = follow_up_queue_capacity - self.follow_up_failure_policy = follow_up_failure_policy - self._operation_lock = asyncio.Lock() - self._shutdown_lock = asyncio.Lock() - self._run_chain_claim_lock = asyncio.Lock() - self._run_chain_gate = asyncio.Event() - self._run_chain_gate.set() - self._active_operation: _ActiveOperation | None = None - self._follow_ups = _FollowUpQueue(follow_up_queue_capacity) - self._follow_up_chain_open = False - self._follow_up_worker: asyncio.Task[None] | None = None - self._shutting_down = False - self._pending_operations = 0 - self._idle_event = asyncio.Event() - self._idle_event.set() - self._started = False - self._skill_manager = SkillManager(skills_dir) if skills_dir else None - self.state = AgentState() - self._context_builder = context_builder or AgentContextBuilder( - skill_manager=self._skill_manager, - ) - self.logger = get_logger(f"{self.agent_id}") - self._event_emitter = AgentEventEmitter( - agent_id=self.agent_id, - sink=self.event_sink, - logger=self.logger, - ) - self._compaction_runtime = CompactionRuntime( - state=self.state, - policy=self.compaction_policy, - estimator=self.context_token_estimator, - event_emitter=self._event_emitter, - ) - self._tool_runtime = ToolRuntime( - self.handlers, - self.middlewares, - state=self.state, - logger=self.logger, - event_emitter=self._event_emitter, - max_repeated_tool_calls=policy.max_repeated_tool_calls, - ) - self.orchestrator = AgentOrchestrator( - agent_id=self.agent_id, - state=self.state, - context_builder=self._context_builder, - model_stream=self.model.stream, - tool_runtime=self._tool_runtime, - skill_manager=self._skill_manager, - policy=self.runtime_policy, - compaction_policy=self.compaction_policy, - auto_compaction_policy=self.auto_compaction_policy, - compactor=self.compactor, - compaction_runtime=self._compaction_runtime, - context_token_estimator=self.context_token_estimator, - behavior_hooks=self.behavior_hooks, - event_emitter=self._event_emitter, - ) - self.reset() - - @property - def agent_id(self) -> str: - """Return the immutable identity of this agent.""" - - return self._agent_id - - @property - def tools(self) -> list[dict[str, Any]]: - """Return the currently registered function tool definitions.""" - - return self.orchestrator.tools - - @property - def tool_definitions(self) -> tuple[ToolDefinition, ...]: - """Return the canonical definitions registered with this Agent.""" - - return self.orchestrator.tool_definitions - - @property - def messages(self) -> list[dict[str, Any]]: - """Return the agent's persistent conversation history.""" - - return self.state.messages - - @property - def pending_steering_count(self) -> int: - """Return queued Steering inputs for the active Run.""" - - active_operation = self._active_operation - if active_operation is None or active_operation.steering is None: - return 0 - return active_operation.steering.size - - @property - def pending_follow_up_count(self) -> int: - """Return accepted Follow-ups that have not started their Run.""" - - return self._follow_ups.size - - @property - def continue_rejection_reason(self) -> ContinueRejectedReason | None: - """Return why Continue is currently unavailable, if any.""" - - if self._shutting_down: - return ContinueRejectedReason.AGENT_SHUTTING_DOWN - if self._pending_operations or not self._run_chain_gate.is_set(): - return ContinueRejectedReason.AGENT_ACTIVE - return _continue_history_rejection_reason( - self.state.messages, - self.state.last_run_result, - ) - - @property - def can_continue(self) -> bool: - """Return whether a Continue Run could be started now.""" - - return self.continue_rejection_reason is None - - def reset( - self, - history: Sequence[Mapping[str, Any]] | None = None, - ) -> None: - """Reset conversation memory while preserving the agent identity.""" - - messages = [{"role": "system", "content": self.system_prompt}] - if self.handlers: - messages.append({"role": "system", "content": TOOL_PROTOCOL_PROMPT}) - if self.runtime_policy.require_explicit_finish: - messages.append( - { - "role": "system", - "content": EXPLICIT_FINISH_PROTOCOL_PROMPT, - } - ) - if history: - messages.extend(dict(message) for message in history) - self.state.reset(messages) - - def restore_session(self, session: AgentSession) -> None: - """Restore one finished Session projection into this Agent's history.""" - - if self._pending_operations or self._operation_lock.locked(): - raise RuntimeError("cannot restore a Session while the agent is active") - snapshot = session.snapshot() - if snapshot.agent_id is not None and snapshot.agent_id != self.agent_id: - raise ValueError( - f"session {snapshot.session_id!r} belongs to agent " - f"{snapshot.agent_id!r}, not {self.agent_id!r}" - ) - unfinished = [run.run_id for run in snapshot.runs if not run.finished] - if unfinished: - raise ValueError( - "cannot restore a Session with unfinished run(s): " - + ", ".join(unfinished) - ) - self.reset(snapshot.messages) - if snapshot.runs: - self.state.last_run_result = snapshot.runs[-1].result - - async def startup(self) -> None: - """Start the model adapter, handlers, and middleware resources.""" - - await self._claim_run_chain() - try: - async with self._operation_lock: - await self._startup() - finally: - self._release_run_chain() - - async def _startup(self) -> None: - await self._ensure_skills_discovered() - - if self._started: - return - - try: - await self.model.startup() - await self._tool_runtime.startup() - if self._tool_runtime.tool_definitions: - self.logger.info( - "Loaded %d handler(s); registered tools: %s", - len(self.handlers), - ", ".join( - sorted( - tool.name for tool in self._tool_runtime.tool_definitions - ) - ), - ) - except Exception: - try: - await self._tool_runtime.shutdown() - except Exception as shutdown_error: - self.logger.warning( - "Tool runtime rollback shutdown failed: %s", - shutdown_error, - ) - try: - await self.model.shutdown() - except Exception as shutdown_error: - self.logger.warning( - "Model adapter rollback shutdown failed: %s", - shutdown_error, - ) - raise - - self._started = True - - async def shutdown(self) -> None: - """Release all resources owned by this agent.""" - - async with self._shutdown_lock: - self._shutting_down = True - self._discard_follow_ups(FollowUpDiscardReason.AGENT_SHUTDOWN) - chain_claimed = False - try: - await self._claim_run_chain() - chain_claimed = True - async with self._operation_lock: - await self._shutdown() - finally: - self._shutting_down = False - if chain_claimed: - self._release_run_chain() - - async def _shutdown(self) -> None: - if not self._started: - return - - errors: list[Exception] = [] - try: - await self._tool_runtime.shutdown() - except Exception as exc: - errors.append(exc) - try: - await self.model.shutdown() - except Exception as exc: - errors.append(exc) - self._started = False - if errors: - raise RuntimeError( - f"failed to shut down {len(errors)} agent resource(s)" - ) from errors[0] - - async def dispatch( - self, - tool_name: str, - arguments: Mapping[str, Any], - ) -> StepOutcome: - """Dispatch a tool call to its explicitly registered handler.""" - - await self._claim_run_chain() - try: - async with self._operation_lock: - await self._startup() - return await self._tool_runtime.dispatch(tool_name, arguments) - finally: - self._release_run_chain() - - async def run(self, *, task: str) -> AgentRunResult: - """Run one task and return a structured terminal result.""" - - self._track_operation() - chain_claimed = False - result: AgentRunResult | None = None - try: - await self._claim_run_chain() - chain_claimed = True - self._follow_up_chain_open = True - async with self._operation_lock: - result = await self._run_task(task) - return result - finally: - self._settle_operation() - if chain_claimed: - self._finish_run_chain(result) - - async def continue_run(self) -> AgentRunResult: - """Resume existing history in a new Run without adding a user message.""" - - reason = self.continue_rejection_reason - if reason is not None: - raise ContinueRejectedError(reason) - if not await self._try_claim_run_chain(): - raise ContinueRejectedError(ContinueRejectedReason.AGENT_ACTIVE) - - tracked = False - result: AgentRunResult | None = None - try: - reason = _continue_history_rejection_reason( - self.state.messages, - self.state.last_run_result, - ) - if reason is not None: - raise ContinueRejectedError(reason) - self._track_operation() - tracked = True - self._follow_up_chain_open = True - async with self._operation_lock: - result = await self._continue_task() - return result - finally: - if tracked: - self._settle_operation() - self._finish_run_chain(result) - else: - self._release_run_chain() - - async def compact( - self, - *, - compactor: Compactor | None = None, - ) -> CompactionResult: - """Explicitly summarize old turns and atomically replace history.""" - - active_compactor = self.compactor if compactor is None else compactor - if active_compactor is None: - raise RuntimeError("explicit compaction requires a Compactor") - if self.compaction_policy is None: - raise RuntimeError("explicit compaction requires a CompactionPolicy") - self._track_operation() - chain_claimed = False - try: - await self._claim_run_chain() - chain_claimed = True - async with self._operation_lock: - source = CancellationSource() - active_operation = _ActiveOperation(source) - self._active_operation = active_operation - try: - await self._startup() - return await self._compaction_runtime.compact( - active_compactor, - cancellation=source.token, - ) - finally: - if self._active_operation is active_operation: - self._active_operation = None - finally: - self._settle_operation() - if chain_claimed: - self._release_run_chain() - - async def _run_task(self, task: str) -> AgentRunResult: - source = CancellationSource() - steering = _SteeringQueue(self.steering_queue_capacity) - active_operation = _ActiveOperation(source, steering) - self._active_operation = active_operation - try: - await self._startup() - return await self.orchestrator.run( - task=task, - cancellation=source.token, - steering=steering, - ) - finally: - steering.close() - if self._active_operation is active_operation: - self._active_operation = None - - async def _continue_task(self) -> AgentRunResult: - source = CancellationSource() - steering = _SteeringQueue(self.steering_queue_capacity) - active_operation = _ActiveOperation(source, steering) - self._active_operation = active_operation - try: - await self._startup() - return await self.orchestrator.continue_run( - cancellation=source.token, - steering=steering, - ) - finally: - steering.close() - if self._active_operation is active_operation: - self._active_operation = None - - async def _claim_run_chain(self) -> None: - while True: - await self._run_chain_gate.wait() - async with self._run_chain_claim_lock: - if self._run_chain_gate.is_set(): - self._run_chain_gate.clear() - return - - async def _try_claim_run_chain(self) -> bool: - async with self._run_chain_claim_lock: - if not self._run_chain_gate.is_set(): - return False - self._run_chain_gate.clear() - return True - - def _release_run_chain(self) -> None: - self._follow_up_chain_open = False - self._run_chain_gate.set() - - def _finish_run_chain(self, result: AgentRunResult | None) -> None: - if self._shutting_down: - self._discard_follow_ups(FollowUpDiscardReason.AGENT_SHUTDOWN) - self._release_run_chain() - return - can_continue = ( - result is not None and result.succeeded - ) or self.follow_up_failure_policy is FollowUpFailurePolicy.CONTINUE - if not can_continue: - self._discard_follow_ups(FollowUpDiscardReason.PREVIOUS_RUN_NOT_COMPLETED) - self._release_run_chain() - return - if self._follow_ups.size == 0: - self._release_run_chain() - return - self._follow_up_worker = asyncio.create_task( - self._drain_follow_ups(), - name=f"{self.agent_id}-follow-ups", - ) - - async def _drain_follow_ups(self) -> None: - try: - while self._follow_up_chain_open and not self._shutting_down: - handle = self._follow_ups.pop() - if handle is None: - break - result: AgentRunResult | None = None - try: - async with self._operation_lock: - result = await self._run_task(handle.control.content) - except asyncio.CancelledError: - handle._discard(FollowUpDiscardReason.AGENT_SHUTDOWN) - self._discard_follow_ups(FollowUpDiscardReason.AGENT_SHUTDOWN) - raise - except Exception as exc: - handle._set_exception(exc) - else: - handle._set_result(result) - finally: - self._settle_operation() - - can_continue = ( - result is not None and result.succeeded - ) or self.follow_up_failure_policy is FollowUpFailurePolicy.CONTINUE - if not can_continue: - self._discard_follow_ups( - FollowUpDiscardReason.PREVIOUS_RUN_NOT_COMPLETED - ) - break - finally: - if self._shutting_down: - self._discard_follow_ups(FollowUpDiscardReason.AGENT_SHUTDOWN) - self._follow_up_worker = None - self._release_run_chain() - - def _discard_follow_ups(self, reason: FollowUpDiscardReason) -> None: - for handle in self._follow_ups.drain(): - handle._discard(reason) - self._settle_operation() - - def _track_operation(self) -> None: - self._pending_operations += 1 - self._idle_event.clear() - - def _settle_operation(self) -> None: - self._pending_operations -= 1 - if self._pending_operations < 0: - raise RuntimeError("agent operation accounting became negative") - if self._pending_operations == 0: - self._idle_event.set() - - def abort(self, reason: str | None = None) -> bool: - """Cancel the active run or compaction without waiting for it.""" - - active_operation = self._active_operation - if active_operation is None: - return False - return active_operation.cancellation.cancel(reason) - - async def steer(self, content: str) -> ControlReceipt: - """Queue guidance for the active Run's next model-call safe point.""" - - control = ControlInput.steering(content) - active_operation = self._active_operation - if active_operation is None or active_operation.steering is None: - return ControlReceipt( - control=control, - status=ControlStatus.AGENT_IDLE, - queue_size=0, - ) - if active_operation.cancellation.token.cancelled: - return ControlReceipt( - control=control, - status=ControlStatus.RUN_CLOSING, - queue_size=active_operation.steering.size, - ) - return active_operation.steering.submit(control) - - async def follow_up(self, task: str) -> FollowUpHandle: - """Queue one independent Run after the active Run chain.""" - - control = ControlInput.follow_up(task) - if self._shutting_down: - return self._follow_ups.rejected(control, ControlStatus.RUN_CLOSING) - active_operation = self._active_operation - if ( - active_operation is not None - and active_operation.cancellation.token.cancelled - ): - return self._follow_ups.rejected(control, ControlStatus.RUN_CLOSING) - if not self._follow_up_chain_open: - return self._follow_ups.rejected(control, ControlStatus.AGENT_IDLE) - handle = self._follow_ups.submit(control) - if handle.accepted: - self._track_operation() - return handle - - async def wait_for_idle(self) -> None: - """Wait for requested runs or compactions and their sinks to settle.""" - - await self._idle_event.wait() - - async def runtime(self, *, task: str) -> str | None: - """Compatibility wrapper returning completed output as text.""" - - result = await self.run(task=task) - result.raise_for_status() - return result.output - - async def _ensure_skills_discovered(self) -> None: - if self._skill_manager is not None: - await self._skill_manager.discover() diff --git a/src/ejagent/agent/behavior.py b/src/ejagent/agent/behavior.py deleted file mode 100644 index 8a93342..0000000 --- a/src/ejagent/agent/behavior.py +++ /dev/null @@ -1,115 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping, Sequence -from copy import deepcopy -from dataclasses import dataclass -from enum import StrEnum -from typing import Protocol - -from ejagent.agent.cancellation import CancellationToken -from ejagent.agent.types import AgentMessage -from ejagent.agent.usage import RunUsage -from ejagent.providers.base import AssistantMessage - - -class BehaviorAction(StrEnum): - """Action selected by one behavior decision point.""" - - CONTINUE = "continue" - STOP = "stop" - - -@dataclass(frozen=True, slots=True) -class BehaviorDecision: - """Typed instruction returned by a Behavior Hook.""" - - action: BehaviorAction = BehaviorAction.CONTINUE - output: str | None = None - - def __post_init__(self) -> None: - if not isinstance(self.action, BehaviorAction): - raise TypeError("action must be a BehaviorAction") - if self.action is BehaviorAction.CONTINUE and self.output is not None: - raise ValueError("Continue behavior decision must not contain output") - - @classmethod - def continue_run(cls) -> BehaviorDecision: - """Allow the Agent Loop to start another Turn.""" - - return cls() - - @classmethod - def stop(cls, output: str | None = None) -> BehaviorDecision: - """Complete the Run at the current full-Turn boundary.""" - - return cls(action=BehaviorAction.STOP, output=output) - - -@dataclass(frozen=True, slots=True, init=False) -class TurnSnapshot: - """Detached read snapshot exposed at the after-Turn safe point.""" - - agent_id: str - run_id: str - turn: int - task: str | None - response: AssistantMessage - usage: RunUsage - _messages: tuple[AgentMessage, ...] - - def __init__( - self, - *, - agent_id: str, - run_id: str, - turn: int, - task: str | None, - response: AssistantMessage, - usage: RunUsage, - messages: Sequence[Mapping[str, object]], - ) -> None: - if not agent_id: - raise ValueError("agent_id must not be empty") - if not run_id: - raise ValueError("run_id must not be empty") - if turn <= 0: - raise ValueError("turn must be greater than zero") - object.__setattr__(self, "agent_id", agent_id) - object.__setattr__(self, "run_id", run_id) - object.__setattr__(self, "turn", turn) - object.__setattr__(self, "task", task) - object.__setattr__(self, "response", response) - object.__setattr__(self, "usage", usage) - object.__setattr__( - self, - "_messages", - tuple(deepcopy(dict(message)) for message in messages), - ) - - @property - def messages(self) -> tuple[AgentMessage, ...]: - """Return a fresh detached transcript for this completed Turn.""" - - return tuple(deepcopy(message) for message in self._messages) - - -class BehaviorHook(Protocol): - """Behavior extension evaluated after a non-terminal full Turn.""" - - async def after_turn( - self, - snapshot: TurnSnapshot, - *, - cancellation: CancellationToken, - ) -> BehaviorDecision | None: - """Return STOP or allow the next Turn with CONTINUE/None.""" - - -class BehaviorHookError(RuntimeError): - """One Behavior Hook failed or returned an unsupported decision.""" - - def __init__(self, hook: BehaviorHook, error: BaseException) -> None: - self.hook = hook - self.error = error - hook_name = getattr(hook, "name", type(hook).__name__) - super().__init__(f"behavior hook {hook_name!r} failed: {error}") diff --git a/src/ejagent/agent/cancellation.py b/src/ejagent/agent/cancellation.py deleted file mode 100644 index 70e09d1..0000000 --- a/src/ejagent/agent/cancellation.py +++ /dev/null @@ -1,98 +0,0 @@ -from __future__ import annotations - -import asyncio -import inspect -from collections.abc import Awaitable -from typing import TypeVar - -T = TypeVar("T") - - -class AgentCancelledError(RuntimeError): - """Raised cooperatively when the active agent run is aborted.""" - - -class CancellationToken: - """Read-only cancellation signal shared across one agent run.""" - - def __init__(self) -> None: - self._event = asyncio.Event() - self._reason: str | None = None - - @property - def cancelled(self) -> bool: - """Return whether cancellation has been requested.""" - - return self._event.is_set() - - @property - def reason(self) -> str | None: - """Return the cancellation reason when one was supplied.""" - - return self._reason - - async def wait(self) -> None: - """Wait until cancellation is requested.""" - - await self._event.wait() - - def raise_if_cancelled(self) -> None: - """Raise the agent-level cancellation exception when cancelled.""" - - if self.cancelled: - raise AgentCancelledError(self.reason or "agent run was aborted") - - async def run(self, awaitable: Awaitable[T]) -> T: - """Await work while interrupting it when this token is cancelled.""" - - if self.cancelled: - if inspect.iscoroutine(awaitable): - awaitable.close() - self.raise_if_cancelled() - - work = asyncio.ensure_future(awaitable) - cancellation = asyncio.create_task(self.wait()) - try: - done, _ = await asyncio.wait( - (work, cancellation), - return_when=asyncio.FIRST_COMPLETED, - ) - if work in done: - return await work - - work.cancel() - await asyncio.gather(work, return_exceptions=True) - self.raise_if_cancelled() - raise RuntimeError("cancellation wait completed without a signal") - finally: - cancellation.cancel() - if not work.done(): - work.cancel() - await asyncio.gather( - work, - cancellation, - return_exceptions=True, - ) - - def _cancel(self, reason: str | None) -> bool: - if self.cancelled: - return False - self._reason = reason or "agent run was aborted" - self._event.set() - return True - - -class CancellationSource: - """Mutable owner of one public read-only cancellation token.""" - - def __init__(self) -> None: - self._token = CancellationToken() - - @property - def token(self) -> CancellationToken: - return self._token - - def cancel(self, reason: str | None = None) -> bool: - """Request cancellation once and report whether state changed.""" - - return self._token._cancel(reason) diff --git a/src/ejagent/agent/compaction.py b/src/ejagent/agent/compaction.py deleted file mode 100644 index acfeec0..0000000 --- a/src/ejagent/agent/compaction.py +++ /dev/null @@ -1,439 +0,0 @@ -from __future__ import annotations - -import asyncio -from collections.abc import Mapping -from copy import deepcopy -from dataclasses import dataclass, field -from enum import StrEnum -from typing import TYPE_CHECKING, Any, Protocol -from uuid import uuid4 - -from ejagent.agent.cancellation import ( - AgentCancelledError, - CancellationToken, -) -from ejagent.agent.context_management import ( - CompactionPolicy, - CompactionPreparation, - MessageTokenEstimator, -) -from ejagent.agent.state import AgentState -from ejagent.agent.types import INTERNAL_METADATA_PREFIX, AgentMessage - -if TYPE_CHECKING: - from ejagent.agent.events import AgentEventEmitter - - -SUMMARY_METADATA_KEY = f"{INTERNAL_METADATA_PREFIX}summary" -SUMMARY_CONTEXT_HEADER = "Conversation summary from earlier turns:" - - -class CompactionStatus(StrEnum): - """Terminal state of one compaction operation.""" - - COMPLETED = "completed" - SKIPPED = "skipped" - FAILED = "failed" - CANCELLED = "cancelled" - - -class CompactionTrigger(StrEnum): - """Reason one compaction operation was requested.""" - - EXPLICIT = "explicit" - PRESSURE = "pressure" - OVERFLOW = "overflow" - - -@dataclass(frozen=True, slots=True) -class CompactorOutput: - """Provider-neutral summary text returned by a concrete Compactor.""" - - content: str - source: str - - def __post_init__(self) -> None: - if not self.content.strip(): - raise ValueError("compactor output content must not be empty") - if not self.source.strip(): - raise ValueError("compactor output source must not be empty") - - -@dataclass(frozen=True, slots=True) -class SummaryEntry: - """Canonical metadata for one summary projected into model context.""" - - content: str - source: str - history_start_index: int - first_kept_index: int - summarized_message_count: int - tokens_before: int - - def __post_init__(self) -> None: - if not self.content.strip(): - raise ValueError("summary content must not be empty") - if not self.source.strip(): - raise ValueError("summary source must not be empty") - if self.history_start_index < 0: - raise ValueError("history_start_index must not be negative") - if self.first_kept_index <= self.history_start_index: - raise ValueError("first_kept_index must follow history_start_index") - if self.summarized_message_count <= 0: - raise ValueError("summarized_message_count must be greater than zero") - if self.tokens_before < 0: - raise ValueError("tokens_before must not be negative") - - def to_dict(self) -> dict[str, Any]: - """Return a detached JSON-compatible metadata representation.""" - - return { - "content": self.content, - "source": self.source, - "history_start_index": self.history_start_index, - "first_kept_index": self.first_kept_index, - "summarized_message_count": self.summarized_message_count, - "tokens_before": self.tokens_before, - } - - @classmethod - def from_dict(cls, value: Mapping[str, Any]) -> SummaryEntry: - """Restore validated summary metadata from an internal message.""" - - content = value["content"] - source = value["source"] - integer_fields = { - name: value[name] - for name in ( - "history_start_index", - "first_kept_index", - "summarized_message_count", - "tokens_before", - ) - } - if not isinstance(content, str) or not isinstance(source, str): - raise TypeError("summary content and source must be strings") - if any( - not isinstance(item, int) or isinstance(item, bool) - for item in integer_fields.values() - ): - raise TypeError("summary range and token fields must be integers") - return cls( - content=content, - source=source, - history_start_index=integer_fields["history_start_index"], - first_kept_index=integer_fields["first_kept_index"], - summarized_message_count=integer_fields["summarized_message_count"], - tokens_before=integer_fields["tokens_before"], - ) - - def to_agent_message(self) -> AgentMessage: - """Project this entry as a system message with internal metadata.""" - - return { - "role": "system", - "content": f"{SUMMARY_CONTEXT_HEADER}\n\n{self.content}", - SUMMARY_METADATA_KEY: self.to_dict(), - } - - -@dataclass(frozen=True, slots=True) -class CompactionRequest: - """Input supplied to a concrete summary generator.""" - - preparation: CompactionPreparation - previous_summary: SummaryEntry | None = None - trigger: CompactionTrigger = CompactionTrigger.EXPLICIT - operation_id: str = field(default_factory=lambda: uuid4().hex) - - def __post_init__(self) -> None: - if not self.operation_id: - raise ValueError("compaction operation_id must not be empty") - - -class Compactor(Protocol): - """Pluggable, cancellable behavior for generating one summary.""" - - async def compact( - self, - request: CompactionRequest, - *, - cancellation: CancellationToken | None = None, - ) -> CompactorOutput: - """Summarize the prepared history without mutating Agent State.""" - - -@dataclass(frozen=True, slots=True) -class CompactionResult: - """Structured terminal result of one compaction operation.""" - - status: CompactionStatus - preparation: CompactionPreparation - summary: SummaryEntry | None = None - messages: tuple[AgentMessage, ...] = field(default_factory=tuple) - error: str | None = None - trigger: CompactionTrigger = CompactionTrigger.EXPLICIT - operation_id: str = field(default_factory=lambda: uuid4().hex) - - def __post_init__(self) -> None: - object.__setattr__(self, "messages", deepcopy(self.messages)) - if not self.operation_id: - raise ValueError("compaction operation_id must not be empty") - if self.status is CompactionStatus.COMPLETED: - if self.summary is None or not self.messages or self.error is not None: - raise ValueError( - "completed compaction requires summary and messages only" - ) - if not self.preparation.can_compact: - raise ValueError("completed compaction requires old turns") - elif self.status is CompactionStatus.SKIPPED: - if self.summary is not None or self.messages or self.error is not None: - raise ValueError("skipped compaction must not contain output") - if self.preparation.can_compact: - raise ValueError("skipped compaction must not have old turns") - else: - if self.summary is not None or self.messages or not self.error: - raise ValueError( - "failed or cancelled compaction requires only an error" - ) - - @property - def completed(self) -> bool: - return self.status is CompactionStatus.COMPLETED - - -def build_summary_entry( - output: CompactorOutput, - preparation: CompactionPreparation, - *, - previous_summary: SummaryEntry | None = None, -) -> SummaryEntry: - """Combine trusted Core range metadata with concrete summary text.""" - - previous_count = ( - previous_summary.summarized_message_count if previous_summary is not None else 0 - ) - return SummaryEntry( - content=output.content.strip(), - source=output.source.strip(), - history_start_index=preparation.history_start_index, - first_kept_index=preparation.first_kept_index, - summarized_message_count=( - previous_count + len(preparation.messages_to_summarize) - ), - tokens_before=preparation.estimated_history_tokens, - ) - - -def find_previous_summary( - messages: tuple[AgentMessage, ...], -) -> SummaryEntry | None: - """Return the latest valid internal Summary Entry, if present.""" - - for message in reversed(messages): - raw = message.get(SUMMARY_METADATA_KEY) - if not isinstance(raw, Mapping): - continue - try: - return SummaryEntry.from_dict(raw) - except (KeyError, TypeError, ValueError): - continue - return None - - -def build_compacted_state_messages( - preparation: CompactionPreparation, - summary: SummaryEntry, -) -> list[AgentMessage]: - """Build the atomic Agent State replacement without old Summary entries.""" - - protected = [ - deepcopy(message) - for message in preparation.protected_messages - if not _is_valid_summary_message(message) - ] - return [ - *protected, - summary.to_agent_message(), - *(deepcopy(message) for message in preparation.messages_to_keep), - ] - - -def build_compacted_session_messages( - preparation: CompactionPreparation, - summary: SummaryEntry, -) -> tuple[AgentMessage, ...]: - """Build the compacted conversation projection stored by Session.""" - - leading_system_end = 0 - while ( - leading_system_end < len(preparation.protected_messages) - and preparation.protected_messages[leading_system_end].get("role") == "system" - ): - leading_system_end += 1 - preserved = ( - deepcopy(message) - for message in preparation.protected_messages[leading_system_end:] - if not _is_valid_summary_message(message) - ) - return ( - *preserved, - summary.to_agent_message(), - *(deepcopy(message) for message in preparation.messages_to_keep), - ) - - -def _is_valid_summary_message(message: Mapping[str, Any]) -> bool: - raw = message.get(SUMMARY_METADATA_KEY) - if not isinstance(raw, Mapping): - return False - try: - SummaryEntry.from_dict(raw) - except (KeyError, TypeError, ValueError): - return False - return True - - -class CompactionRuntime: - """Execute compaction with atomic state replacement and lifecycle events.""" - - def __init__( - self, - *, - state: AgentState, - policy: CompactionPolicy | None, - estimator: MessageTokenEstimator | None, - event_emitter: AgentEventEmitter, - ) -> None: - self.state = state - self.policy = policy - self.estimator = estimator - self.event_emitter = event_emitter - - async def compact( - self, - compactor: Compactor, - *, - cancellation: CancellationToken, - ) -> CompactionResult: - """Run explicit compaction in a dedicated event operation.""" - - operation_id = self.event_emitter.begin_run() - try: - return await self.compact_in_active_run( - compactor, - cancellation=cancellation, - trigger=CompactionTrigger.EXPLICIT, - ) - finally: - self.event_emitter.end_run(operation_id) - - async def compact_in_active_run( - self, - compactor: Compactor, - *, - cancellation: CancellationToken, - trigger: CompactionTrigger, - preparation: CompactionPreparation | None = None, - ) -> CompactionResult: - """Compact while reusing the caller's active Agent event run.""" - - from ejagent.agent.events import ( - CompactionCompleted, - CompactionFailed, - CompactionStarted, - ) - - if self.policy is None: - raise RuntimeError("explicit compaction requires a CompactionPolicy") - - before = self.state.snapshot().messages - active_preparation = preparation or self.policy.prepare( - before, - estimator=self.estimator, - ) - previous_summary = find_previous_summary(active_preparation.protected_messages) - request = CompactionRequest( - active_preparation, - previous_summary, - trigger=trigger, - ) - await self.event_emitter.emit(CompactionStarted(request)) - if not active_preparation.can_compact: - result = CompactionResult( - status=CompactionStatus.SKIPPED, - preparation=active_preparation, - trigger=trigger, - operation_id=request.operation_id, - ) - await self.event_emitter.emit(CompactionCompleted(result)) - return result - - try: - output = await cancellation.run( - compactor.compact( - request, - cancellation=cancellation, - ) - ) - if not isinstance(output, CompactorOutput): - raise TypeError("Compactor.compact() must return CompactorOutput") - cancellation.raise_if_cancelled() - if self.state.messages != before: - raise RuntimeError("agent history changed during compaction") - - summary = build_summary_entry( - output, - active_preparation, - previous_summary=previous_summary, - ) - state_messages = build_compacted_state_messages( - active_preparation, - summary, - ) - session_messages = build_compacted_session_messages( - active_preparation, - summary, - ) - result = CompactionResult( - status=CompactionStatus.COMPLETED, - preparation=active_preparation, - summary=summary, - messages=session_messages, - trigger=trigger, - operation_id=request.operation_id, - ) - completed_event = CompactionCompleted(result) - self.state.replace_messages(state_messages) - await self.event_emitter.emit(completed_event) - return result - except AgentCancelledError as exc: - result = CompactionResult( - status=CompactionStatus.CANCELLED, - preparation=active_preparation, - error=str(exc), - trigger=trigger, - operation_id=request.operation_id, - ) - await self.event_emitter.emit(CompactionFailed(result)) - return result - except asyncio.CancelledError: - result = CompactionResult( - status=CompactionStatus.CANCELLED, - preparation=active_preparation, - error="compaction coroutine was cancelled", - trigger=trigger, - operation_id=request.operation_id, - ) - await self.event_emitter.emit(CompactionFailed(result)) - raise - except Exception as exc: - result = CompactionResult( - status=CompactionStatus.FAILED, - preparation=active_preparation, - error=str(exc), - trigger=trigger, - operation_id=request.operation_id, - ) - await self.event_emitter.emit(CompactionFailed(result)) - return result diff --git a/src/ejagent/agent/context_builder.py b/src/ejagent/agent/context_builder.py deleted file mode 100644 index 0a8fe60..0000000 --- a/src/ejagent/agent/context_builder.py +++ /dev/null @@ -1,125 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping, Sequence -from dataclasses import dataclass -from typing import Any - -from ejagent.agent.state import AgentState -from ejagent.agent.types import AgentMessage, is_internal_metadata_key -from ejagent.handlers.definition import ( - ToolDefinition, - ToolDefinitionInput, - normalize_tool_definitions, -) -from ejagent.plugins.skill.skill_manager import SkillManager - - -@dataclass(frozen=True, slots=True) -class ContextBuildResult: - """One complete model request built from agent state.""" - - agent_messages: tuple[AgentMessage, ...] - llm_messages: tuple[AgentMessage, ...] - tools: tuple[dict[str, Any], ...] - tool_definitions: tuple[ToolDefinition, ...] = () - - -class AgentContextBuilder: - """Build provider-ready context from agent state without mutating it.""" - - def __init__( - self, - *, - skill_manager: SkillManager | None = None, - ) -> None: - self._skill_manager = skill_manager - - def build( - self, - state: AgentState, - *, - tools: Sequence[ToolDefinitionInput] = (), - transient_messages: Sequence[Mapping[str, Any]] = (), - ) -> ContextBuildResult: - """Build the context for one model request. - - Persistent conversation is copied from ``state``. Skill instructions, - per-turn control messages, and tool definitions are combined into a - complete provider request without mutating the state history. - """ - - context = self._copy_messages(state.messages) - self._insert_skill_context(context, state.active_skill_name) - context.extend(self._copy_messages(transient_messages)) - projected = self.convert_to_llm_messages(context) - llm_messages = self._strip_internal_metadata(projected) - definitions = normalize_tool_definitions(tools) - return ContextBuildResult( - agent_messages=tuple(context), - llm_messages=tuple(llm_messages), - tools=tuple(tool.to_openai_tool() for tool in definitions), - tool_definitions=definitions, - ) - - def convert_to_llm_messages( - self, - messages: Sequence[Mapping[str, Any]], - ) -> list[AgentMessage]: - """Convert internal messages into provider-compatible messages. - - The default is copy-only. Subclasses can inspect internal metadata, - filter records, or adapt custom message types without mutating - ``AgentState``. Core-only metadata is removed after this hook returns. - """ - - return self._copy_messages(messages) - - @staticmethod - def _strip_internal_metadata( - messages: Sequence[Mapping[str, Any]], - ) -> list[AgentMessage]: - return [ - { - key: value - for key, value in message.items() - if key != "usage" and not is_internal_metadata_key(key) - } - for message in messages - ] - - def _insert_skill_context( - self, - context: list[AgentMessage], - active_skill_name: str | None, - ) -> None: - if self._skill_manager is None: - return - - skill_messages: list[AgentMessage] = [] - index_message = self._skill_manager.build_index_message() - if index_message is not None: - skill_messages.append(dict(index_message)) - if active_skill_name is not None: - skill_messages.append( - self._skill_manager.build_skill_context_message(active_skill_name) - ) - if not skill_messages: - return - - insert_at = self._system_message_end(context) - context[insert_at:insert_at] = skill_messages - - @staticmethod - def _copy_messages( - messages: Sequence[Mapping[str, Any]], - ) -> list[AgentMessage]: - return [dict(message) for message in messages] - - @staticmethod - def _system_message_end( - messages: Sequence[Mapping[str, Any]], - ) -> int: - index = 0 - while index < len(messages) and messages[index].get("role") == "system": - index += 1 - return index diff --git a/src/ejagent/agent/context_management.py b/src/ejagent/agent/context_management.py deleted file mode 100644 index 4e62ce9..0000000 --- a/src/ejagent/agent/context_management.py +++ /dev/null @@ -1,409 +0,0 @@ -from __future__ import annotations - -import json -import math -from collections.abc import Mapping, Sequence -from copy import deepcopy -from dataclasses import dataclass -from enum import StrEnum -from typing import Any, Protocol - -from ejagent.agent.types import AgentMessage, is_internal_metadata_key -from ejagent.providers.base import ModelUsage - - -class ContextUsageSource(StrEnum): - """Origin of one context pressure estimate.""" - - ESTIMATED = "estimated" - REPORTED = "reported" - MIXED = "mixed" - - -@dataclass(frozen=True, slots=True) -class ContextUsageEstimate: - """Conservative token estimate for one complete provider request.""" - - reported_tokens: int - trailing_tokens: int - heuristic_tokens: int - total_tokens: int - last_usage_index: int | None - source: ContextUsageSource - - def __post_init__(self) -> None: - for name in ( - "reported_tokens", - "trailing_tokens", - "heuristic_tokens", - "total_tokens", - ): - if getattr(self, name) < 0: - raise ValueError(f"{name} must not be negative") - if self.last_usage_index is not None and self.last_usage_index < 0: - raise ValueError("last_usage_index must not be negative") - if self.total_tokens < self.heuristic_tokens: - raise ValueError("total_tokens must cover heuristic_tokens") - if self.total_tokens < self.reported_tokens + self.trailing_tokens: - raise ValueError("total_tokens must cover reported and trailing tokens") - if ( - self.last_usage_index is None - and self.source is not ContextUsageSource.ESTIMATED - ): - raise ValueError("an estimate without usage must use estimated source") - - @property - def usage_based_tokens(self) -> int: - """Return the reported baseline plus messages added after it.""" - - return self.reported_tokens + self.trailing_tokens - - -class MessageTokenEstimator(Protocol): - """Replaceable token estimator for provider-visible context values.""" - - def estimate_message(self, message: Mapping[str, Any]) -> int: - """Estimate one provider-visible message.""" - - def estimate_tools(self, tools: Sequence[Mapping[str, Any]]) -> int: - """Estimate the complete tool definition collection.""" - - -class HeuristicMessageTokenEstimator: - """UTF-8-aware fallback used when no provider tokenizer is available.""" - - message_overhead_tokens = 4 - tool_overhead_tokens = 8 - - def estimate_message(self, message: Mapping[str, Any]) -> int: - visible = { - key: value - for key, value in message.items() - if key != "usage" and not is_internal_metadata_key(key) - } - return self.message_overhead_tokens + self._estimate_value(visible) - - def estimate_tools(self, tools: Sequence[Mapping[str, Any]]) -> int: - if not tools: - return 0 - return self.tool_overhead_tokens + self._estimate_value(list(tools)) - - @staticmethod - def _estimate_value(value: Any) -> int: - serialized = json.dumps( - value, - ensure_ascii=False, - separators=(",", ":"), - default=str, - ) - ascii_count = sum(ord(character) < 128 for character in serialized) - non_ascii_count = len(serialized) - ascii_count - return math.ceil(ascii_count / 4 + non_ascii_count) - - -def estimate_context_usage( - messages: Sequence[Mapping[str, Any]], - *, - tools: Sequence[Mapping[str, Any]] = (), - estimator: MessageTokenEstimator | None = None, -) -> ContextUsageEstimate: - """Estimate the full request from the latest Usage plus a safe fallback. - - The most recent assistant Usage is the best baseline for the repeated - context. Messages appended after that response are estimated. A complete - heuristic pass, including tool definitions, guards against changed - projections and is used as a lower bound. - """ - - active_estimator = estimator or HeuristicMessageTokenEstimator() - heuristic_tokens = sum( - active_estimator.estimate_message(message) for message in messages - ) + active_estimator.estimate_tools(tools) - - usage_info = _last_usage_info(messages) - if usage_info is None: - return ContextUsageEstimate( - reported_tokens=0, - trailing_tokens=heuristic_tokens, - heuristic_tokens=heuristic_tokens, - total_tokens=heuristic_tokens, - last_usage_index=None, - source=ContextUsageSource.ESTIMATED, - ) - - usage_index, reported_tokens = usage_info - trailing_tokens = sum( - active_estimator.estimate_message(message) - for message in messages[usage_index + 1 :] - ) - usage_based_tokens = reported_tokens + trailing_tokens - total_tokens = max(usage_based_tokens, heuristic_tokens) - source = ( - ContextUsageSource.REPORTED - if trailing_tokens == 0 and total_tokens == reported_tokens - else ContextUsageSource.MIXED - ) - return ContextUsageEstimate( - reported_tokens=reported_tokens, - trailing_tokens=trailing_tokens, - heuristic_tokens=heuristic_tokens, - total_tokens=total_tokens, - last_usage_index=usage_index, - source=source, - ) - - -def _last_usage_info( - messages: Sequence[Mapping[str, Any]], -) -> tuple[int, int] | None: - for index in range(len(messages) - 1, -1, -1): - message = messages[index] - if message.get("role") != "assistant" or "usage" not in message: - continue - total_tokens = _usage_total_tokens(message["usage"]) - if total_tokens is not None: - return index, total_tokens - return None - - -def _usage_total_tokens(value: Any) -> int | None: - if isinstance(value, ModelUsage): - return value.total_tokens - if not isinstance(value, Mapping): - return None - total_tokens = value.get("total_tokens") - if ( - isinstance(total_tokens, int) - and not isinstance(total_tokens, bool) - and total_tokens >= 0 - ): - return total_tokens - return None - - -@dataclass(frozen=True, slots=True) -class ContextBudget: - """Capacity reserved for one model request, independent of run spend.""" - - context_window: int - reserve_tokens: int - keep_recent_tokens: int - - def __post_init__(self) -> None: - if self.context_window <= 0: - raise ValueError("context_window must be greater than zero") - if self.reserve_tokens < 0: - raise ValueError("reserve_tokens must not be negative") - if self.reserve_tokens >= self.context_window: - raise ValueError("reserve_tokens must be less than context_window") - if self.keep_recent_tokens <= 0: - raise ValueError("keep_recent_tokens must be greater than zero") - if self.keep_recent_tokens > self.threshold_tokens: - raise ValueError("keep_recent_tokens must not exceed the context threshold") - - @property - def threshold_tokens(self) -> int: - """Return the pressure threshold before reserved capacity.""" - - return self.context_window - self.reserve_tokens - - -class CompactionDecisionReason(StrEnum): - """Reason behind one compaction policy decision.""" - - DISABLED = "disabled" - BELOW_THRESHOLD = "below_threshold" - THRESHOLD_REACHED = "threshold_reached" - - -@dataclass(frozen=True, slots=True) -class CompactionDecision: - """Pure policy result derived from a context usage estimate.""" - - estimate: ContextUsageEstimate - threshold_tokens: int - should_compact: bool - reason: CompactionDecisionReason - - @property - def pressure_ratio(self) -> float: - """Return context pressure relative to the policy threshold.""" - - return self.estimate.total_tokens / self.threshold_tokens - - -@dataclass(frozen=True, slots=True) -class CompactionPreparation: - """Non-mutating plan for summarizing old persistent conversation turns.""" - - protected_messages: tuple[AgentMessage, ...] - messages_to_summarize: tuple[AgentMessage, ...] - messages_to_keep: tuple[AgentMessage, ...] - history_start_index: int - first_kept_index: int - estimated_history_tokens: int - estimated_summarized_tokens: int - estimated_kept_tokens: int - - @property - def can_compact(self) -> bool: - """Return whether at least one complete old turn can be summarized.""" - - return bool(self.messages_to_summarize) - - -@dataclass(frozen=True, slots=True) -class CompactionPolicy: - """Decide when context pressure warrants preparing compaction.""" - - budget: ContextBudget - enabled: bool = True - - def evaluate( - self, - estimate: ContextUsageEstimate, - ) -> CompactionDecision: - if not self.enabled: - return CompactionDecision( - estimate=estimate, - threshold_tokens=self.budget.threshold_tokens, - should_compact=False, - reason=CompactionDecisionReason.DISABLED, - ) - should_compact = estimate.total_tokens >= self.budget.threshold_tokens - return CompactionDecision( - estimate=estimate, - threshold_tokens=self.budget.threshold_tokens, - should_compact=should_compact, - reason=( - CompactionDecisionReason.THRESHOLD_REACHED - if should_compact - else CompactionDecisionReason.BELOW_THRESHOLD - ), - ) - - def prepare( - self, - messages: Sequence[Mapping[str, Any]], - *, - estimator: MessageTokenEstimator | None = None, - ) -> CompactionPreparation: - return prepare_compaction( - messages, - keep_recent_tokens=self.budget.keep_recent_tokens, - estimator=estimator, - ) - - -@dataclass(frozen=True, slots=True) -class AutoCompactionPolicy: - """Opt-in behavior for proactive compaction and overflow recovery.""" - - enabled: bool = True - compact_on_pressure: bool = True - recover_on_overflow: bool = True - max_overflow_retries: int = 1 - - def __post_init__(self) -> None: - if not isinstance(self.enabled, bool): - raise TypeError("enabled must be a bool") - if not isinstance(self.compact_on_pressure, bool): - raise TypeError("compact_on_pressure must be a bool") - if not isinstance(self.recover_on_overflow, bool): - raise TypeError("recover_on_overflow must be a bool") - if ( - not isinstance(self.max_overflow_retries, int) - or isinstance(self.max_overflow_retries, bool) - or self.max_overflow_retries not in {0, 1} - ): - raise ValueError("max_overflow_retries must be zero or one") - - -def prepare_compaction( - messages: Sequence[Mapping[str, Any]], - *, - keep_recent_tokens: int, - estimator: MessageTokenEstimator | None = None, -) -> CompactionPreparation: - """Prepare a safe full-turn cut without changing Agent State. - - Leading system messages are protected. Only a contiguous suffix containing - user, assistant, and tool messages is eligible. Cuts occur immediately - before a user message, so assistant tool calls and their tool results are - never separated by this first compaction implementation. - """ - - if keep_recent_tokens <= 0: - raise ValueError("keep_recent_tokens must be greater than zero") - active_estimator = estimator or HeuristicMessageTokenEstimator() - copied = [deepcopy(dict(message)) for message in messages] - history_start = _compactable_history_start(copied) - ranges = _turn_ranges(copied, history_start) - - if not ranges: - return CompactionPreparation( - protected_messages=tuple(copied), - messages_to_summarize=(), - messages_to_keep=(), - history_start_index=history_start, - first_kept_index=history_start, - estimated_history_tokens=0, - estimated_summarized_tokens=0, - estimated_kept_tokens=0, - ) - - turn_tokens = [ - sum(active_estimator.estimate_message(message) for message in copied[start:end]) - for start, end in ranges - ] - kept_turn_index = len(ranges) - 1 - kept_tokens = turn_tokens[kept_turn_index] - while kept_turn_index > 0 and kept_tokens < keep_recent_tokens: - kept_turn_index -= 1 - kept_tokens += turn_tokens[kept_turn_index] - - first_kept = ranges[kept_turn_index][0] - summarized_tokens = sum(turn_tokens[:kept_turn_index]) - history_tokens = sum(turn_tokens) - return CompactionPreparation( - protected_messages=tuple(copied[:history_start]), - messages_to_summarize=tuple(copied[history_start:first_kept]), - messages_to_keep=tuple(copied[first_kept:]), - history_start_index=history_start, - first_kept_index=first_kept, - estimated_history_tokens=history_tokens, - estimated_summarized_tokens=summarized_tokens, - estimated_kept_tokens=history_tokens - summarized_tokens, - ) - - -def _compactable_history_start(messages: Sequence[Mapping[str, Any]]) -> int: - history_start = 0 - while ( - history_start < len(messages) - and messages[history_start].get("role") == "system" - ): - history_start += 1 - - compactable_roles = {"user", "assistant", "tool"} - for index in range(history_start, len(messages)): - if messages[index].get("role") not in compactable_roles: - history_start = index + 1 - return history_start - - -def _turn_ranges( - messages: Sequence[Mapping[str, Any]], - history_start: int, -) -> list[tuple[int, int]]: - if history_start >= len(messages): - return [] - - starts = [history_start] - for index in range(history_start + 1, len(messages)): - if messages[index].get("role") == "user": - starts.append(index) - return [ - (start, starts[index + 1] if index + 1 < len(starts) else len(messages)) - for index, start in enumerate(starts) - ] diff --git a/src/ejagent/agent/control.py b/src/ejagent/agent/control.py deleted file mode 100644 index 45cae36..0000000 --- a/src/ejagent/agent/control.py +++ /dev/null @@ -1,334 +0,0 @@ -from __future__ import annotations - -import asyncio -from collections import deque -from collections.abc import Mapping, Sequence -from dataclasses import dataclass -from enum import StrEnum -from uuid import uuid4 - -from ejagent.agent.result import AgentRunResult, StopReason -from ejagent.agent.types import AgentMessage - - -class ControlInputKind(StrEnum): - """Kind of external input submitted to an active Agent operation.""" - - STEERING = "steering" - FOLLOW_UP = "follow_up" - - -class ControlStatus(StrEnum): - """Immediate result of submitting one control input.""" - - ACCEPTED = "accepted" - AGENT_IDLE = "agent_idle" - QUEUE_FULL = "queue_full" - RUN_CLOSING = "run_closing" - - -class FollowUpFailurePolicy(StrEnum): - """Whether a Run chain continues after a non-successful result.""" - - DISCARD = "discard" - CONTINUE = "continue" - - -class FollowUpDiscardReason(StrEnum): - """Reason an accepted Follow-up will not start a Run.""" - - PREVIOUS_RUN_NOT_COMPLETED = "previous_run_not_completed" - AGENT_SHUTDOWN = "agent_shutdown" - - -class ContinueRejectedReason(StrEnum): - """Reason a Continue Run could not start.""" - - AGENT_ACTIVE = "agent_active" - AGENT_SHUTTING_DOWN = "agent_shutting_down" - NO_PREVIOUS_RUN = "no_previous_run" - UNSUPPORTED_STOP_REASON = "unsupported_stop_reason" - INCOMPLETE_TOOL_STATE = "incomplete_tool_state" - - -class ContinueRejectedError(RuntimeError): - """A Continue request was rejected before allocating a Run.""" - - def __init__(self, reason: ContinueRejectedReason) -> None: - self.reason = reason - super().__init__(f"continue run was rejected: {reason.value}") - - -_CONTINUABLE_STOP_REASONS = frozenset( - { - StopReason.TEXT_RESPONSE, - StopReason.TOOL_COMPLETION, - StopReason.BEHAVIOR_STOP, - StopReason.EMPTY_RESPONSE, - StopReason.MAX_STEPS, - StopReason.MAX_NO_TOOL_RESPONSES, - StopReason.TOKEN_BUDGET_EXCEEDED, - StopReason.USAGE_UNAVAILABLE, - } -) - - -def _continue_history_rejection_reason( - messages: Sequence[AgentMessage], - result: AgentRunResult | None, -) -> ContinueRejectedReason | None: - if result is None: - return ContinueRejectedReason.NO_PREVIOUS_RUN - if result.stop_reason not in _CONTINUABLE_STOP_REASONS: - return ContinueRejectedReason.UNSUPPORTED_STOP_REASON - - pending: set[str] = set() - for message in messages: - role = message.get("role") - if role == "assistant": - tool_calls = message.get("tool_calls") - if tool_calls is None: - continue - if not isinstance(tool_calls, list): - return ContinueRejectedReason.INCOMPLETE_TOOL_STATE - for tool_call in tool_calls: - if not isinstance(tool_call, Mapping): - return ContinueRejectedReason.INCOMPLETE_TOOL_STATE - tool_call_id = tool_call.get("id") - if not isinstance(tool_call_id, str) or not tool_call_id: - return ContinueRejectedReason.INCOMPLETE_TOOL_STATE - pending.add(tool_call_id) - elif role == "tool": - tool_call_id = message.get("tool_call_id") - if isinstance(tool_call_id, str): - pending.discard(tool_call_id) - if pending: - return ContinueRejectedReason.INCOMPLETE_TOOL_STATE - return None - - -@dataclass(frozen=True, slots=True) -class ControlInput: - """One immutable external instruction with a stable correlation id.""" - - input_id: str - kind: ControlInputKind - content: str - - def __post_init__(self) -> None: - input_id = self.input_id.strip() - content = self.content.strip() - if not input_id: - raise ValueError("input_id must not be empty") - if not content: - raise ValueError("control content must not be empty") - object.__setattr__(self, "input_id", input_id) - object.__setattr__(self, "content", content) - - @classmethod - def steering(cls, content: str) -> ControlInput: - """Create one Steering instruction for the active Run.""" - - return cls( - input_id=uuid4().hex, - kind=ControlInputKind.STEERING, - content=content, - ) - - @classmethod - def follow_up(cls, task: str) -> ControlInput: - """Create one Follow-up task for the active Run chain.""" - - return cls( - input_id=uuid4().hex, - kind=ControlInputKind.FOLLOW_UP, - content=task, - ) - - -@dataclass(frozen=True, slots=True) -class ControlReceipt: - """Immediate acknowledgement for one submitted control input.""" - - control: ControlInput - status: ControlStatus - queue_size: int - - def __post_init__(self) -> None: - if self.queue_size < 0: - raise ValueError("queue_size must not be negative") - - @property - def accepted(self) -> bool: - """Return whether the input entered the active Run's queue.""" - - return self.status is ControlStatus.ACCEPTED - - -class FollowUpError(RuntimeError): - """Base error raised while waiting for a Follow-up Run.""" - - -class FollowUpRejectedError(FollowUpError): - """The Follow-up never entered the Run chain.""" - - def __init__(self, receipt: ControlReceipt) -> None: - self.receipt = receipt - super().__init__(f"follow-up was not accepted: {receipt.status.value}") - - -class FollowUpDiscardedError(FollowUpError): - """An accepted Follow-up was discarded before its Run started.""" - - def __init__( - self, - control: ControlInput, - reason: FollowUpDiscardReason, - ) -> None: - self.control = control - self.reason = reason - super().__init__(f"follow-up was discarded: {reason.value}") - - -class FollowUpHandle: - """Waitable handle for one accepted or rejected Follow-up submission.""" - - def __init__(self, receipt: ControlReceipt) -> None: - if receipt.control.kind is not ControlInputKind.FOLLOW_UP: - raise ValueError("FollowUpHandle requires a Follow-up control input") - self.receipt = receipt - self._future: asyncio.Future[AgentRunResult] = ( - asyncio.get_running_loop().create_future() - ) - # Retrieving the exception here prevents warnings when a rejected or - # discarded handle is intentionally never awaited. Awaiters still see it. - self._future.add_done_callback(self._consume_exception) - if not receipt.accepted: - self._future.set_exception(FollowUpRejectedError(receipt)) - - @staticmethod - def _consume_exception(future: asyncio.Future[AgentRunResult]) -> None: - if not future.cancelled(): - future.exception() - - @property - def control(self) -> ControlInput: - """Return the immutable task and correlation id.""" - - return self.receipt.control - - @property - def accepted(self) -> bool: - """Return whether this Follow-up entered the pending queue.""" - - return self.receipt.accepted - - @property - def done(self) -> bool: - """Return whether the Run completed or the submission was rejected.""" - - return self._future.done() - - async def wait(self) -> AgentRunResult: - """Wait without allowing caller cancellation to cancel the queued Run.""" - - return await asyncio.shield(self._future) - - def _set_result(self, result: AgentRunResult) -> None: - if not self._future.done(): - self._future.set_result(result) - - def _set_exception(self, error: BaseException) -> None: - if not self._future.done(): - self._future.set_exception(error) - - def _discard(self, reason: FollowUpDiscardReason) -> None: - self._set_exception(FollowUpDiscardedError(self.control, reason)) - - -class _SteeringQueue: - """Run-scoped bounded FIFO consumed only at model-call safe points.""" - - def __init__(self, capacity: int) -> None: - self.capacity = capacity - self._items: deque[ControlInput] = deque() - self._closed = False - - @property - def size(self) -> int: - return len(self._items) - - @property - def has_pending(self) -> bool: - return bool(self._items) - - def submit(self, control: ControlInput) -> ControlReceipt: - if self._closed: - status = ControlStatus.RUN_CLOSING - elif len(self._items) >= self.capacity: - status = ControlStatus.QUEUE_FULL - else: - self._items.append(control) - status = ControlStatus.ACCEPTED - return ControlReceipt( - control=control, - status=status, - queue_size=len(self._items), - ) - - def drain(self) -> tuple[ControlInput, ...]: - items = tuple(self._items) - self._items.clear() - return items - - def close(self) -> tuple[ControlInput, ...]: - self._closed = True - return self.drain() - - -class _FollowUpQueue: - """Agent-scoped bounded FIFO of independently waitable tasks.""" - - def __init__(self, capacity: int) -> None: - self.capacity = capacity - self._items: deque[FollowUpHandle] = deque() - - @property - def size(self) -> int: - return len(self._items) - - def submit(self, control: ControlInput) -> FollowUpHandle: - if len(self._items) >= self.capacity: - return self.rejected(control, ControlStatus.QUEUE_FULL) - handle = FollowUpHandle( - ControlReceipt( - control=control, - status=ControlStatus.ACCEPTED, - queue_size=len(self._items) + 1, - ) - ) - self._items.append(handle) - return handle - - def rejected( - self, - control: ControlInput, - status: ControlStatus, - ) -> FollowUpHandle: - return FollowUpHandle( - ControlReceipt( - control=control, - status=status, - queue_size=len(self._items), - ) - ) - - def pop(self) -> FollowUpHandle | None: - if not self._items: - return None - return self._items.popleft() - - def drain(self) -> tuple[FollowUpHandle, ...]: - items = tuple(self._items) - self._items.clear() - return items diff --git a/src/ejagent/agent/events.py b/src/ejagent/agent/events.py deleted file mode 100644 index 4afac91..0000000 --- a/src/ejagent/agent/events.py +++ /dev/null @@ -1,367 +0,0 @@ -from __future__ import annotations - -import logging -from collections.abc import Iterable -from dataclasses import dataclass -from enum import StrEnum -from typing import ClassVar, Protocol, TypeAlias -from uuid import uuid4 - -from ejagent.agent.compaction import ( - CompactionRequest, - CompactionResult, - CompactionStatus, -) -from ejagent.agent.context_management import ( - CompactionDecision, - CompactionPreparation, -) -from ejagent.agent.control import ControlInput, ControlInputKind -from ejagent.agent.result import AgentRunResult, StopReason -from ejagent.agent.types import ToolCallResult, ToolProgressUpdate -from ejagent.providers.base import ( - AssistantMessage, - ModelToolCall, - ModelUsage, -) - - -class AgentEventKind(StrEnum): - """Stable discriminator for one observable lifecycle event.""" - - AGENT_STARTED = "agent_started" - AGENT_CONTINUED = "agent_continued" - TURN_STARTED = "turn_started" - STEERING_APPLIED = "steering_applied" - STEERING_DISCARDED = "steering_discarded" - CONTEXT_PRESSURE_EVALUATED = "context_pressure_evaluated" - COMPACTION_STARTED = "compaction_started" - COMPACTION_COMPLETED = "compaction_completed" - COMPACTION_FAILED = "compaction_failed" - MESSAGE_COMPLETED = "message_completed" - ASSISTANT_TEXT_DELTA = "assistant_text_delta" - ASSISTANT_THINKING_DELTA = "assistant_thinking_delta" - TOOL_STARTED = "tool_started" - TOOL_PROGRESSED = "tool_progressed" - TOOL_COMPLETED = "tool_completed" - TURN_COMPLETED = "turn_completed" - AGENT_FINISHED = "agent_finished" - - -@dataclass(frozen=True, slots=True) -class AgentStarted: - """A task entered the agent runtime.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.AGENT_STARTED - task: str - - -@dataclass(frozen=True, slots=True) -class AgentContinued: - """A new Run resumed the existing conversation without a user message.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.AGENT_CONTINUED - - -@dataclass(frozen=True, slots=True) -class TurnStarted: - """One provider turn started.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.TURN_STARTED - turn: int - - -@dataclass(frozen=True, slots=True) -class SteeringApplied: - """One queued Steering input entered persistent history at a safe point.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.STEERING_APPLIED - control: ControlInput - target_turn: int - - def __post_init__(self) -> None: - if self.control.kind is not ControlInputKind.STEERING: - raise ValueError("SteeringApplied requires a Steering control input") - if self.target_turn <= 0: - raise ValueError("target_turn must be greater than zero") - - -@dataclass(frozen=True, slots=True) -class SteeringDiscarded: - """One accepted Steering input was not applied before its Run ended.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.STEERING_DISCARDED - control: ControlInput - reason: StopReason - - def __post_init__(self) -> None: - if self.control.kind is not ControlInputKind.STEERING: - raise ValueError("SteeringDiscarded requires a Steering control input") - - -@dataclass(frozen=True, slots=True) -class ContextPressureEvaluated: - """One complete model request was assessed before provider dispatch.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.CONTEXT_PRESSURE_EVALUATED - turn: int - decision: CompactionDecision - preparation: CompactionPreparation | None = None - - def __post_init__(self) -> None: - if self.turn <= 0: - raise ValueError("turn must be greater than zero") - if self.preparation is not None and not self.decision.should_compact: - raise ValueError("compaction preparation requires a positive decision") - - -@dataclass(frozen=True, slots=True) -class CompactionStarted: - """One compaction operation began summary generation.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.COMPACTION_STARTED - request: CompactionRequest - - -@dataclass(frozen=True, slots=True) -class CompactionCompleted: - """One compaction completed or safely skipped.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.COMPACTION_COMPLETED - result: CompactionResult - - def __post_init__(self) -> None: - if self.result.status not in { - CompactionStatus.COMPLETED, - CompactionStatus.SKIPPED, - }: - raise ValueError("completed event requires completed or skipped result") - - -@dataclass(frozen=True, slots=True) -class CompactionFailed: - """One compaction failed or was cancelled before mutation.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.COMPACTION_FAILED - result: CompactionResult - - def __post_init__(self) -> None: - if self.result.status not in { - CompactionStatus.FAILED, - CompactionStatus.CANCELLED, - }: - raise ValueError("failed event requires failed or cancelled result") - - -@dataclass(frozen=True, slots=True) -class MessageCompleted: - """One complete provider-neutral assistant message was accepted.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.MESSAGE_COMPLETED - turn: int - message: AssistantMessage - usage: ModelUsage | None = None - - -@dataclass(frozen=True, slots=True) -class AssistantTextDelta: - """One provisional assistant text fragment for the active turn.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.ASSISTANT_TEXT_DELTA - turn: int - delta: str - - def __post_init__(self) -> None: - if self.turn <= 0: - raise ValueError("turn must be greater than zero") - if not self.delta: - raise ValueError("assistant text delta must not be empty") - - -@dataclass(frozen=True, slots=True) -class AssistantThinkingDelta: - """One provisional reasoning fragment for the active turn.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.ASSISTANT_THINKING_DELTA - turn: int - delta: str - - def __post_init__(self) -> None: - if self.turn <= 0: - raise ValueError("turn must be greater than zero") - if not self.delta: - raise ValueError("assistant thinking delta must not be empty") - - -@dataclass(frozen=True, slots=True) -class ToolStarted: - """Execution of one normalized model tool call started.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.TOOL_STARTED - turn: int - tool_call: ModelToolCall - - -@dataclass(frozen=True, slots=True) -class ToolProgressed: - """One provisional update from the active tool execution.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.TOOL_PROGRESSED - turn: int - tool_call: ModelToolCall - update: ToolProgressUpdate - - def __post_init__(self) -> None: - if self.turn <= 0: - raise ValueError("turn must be greater than zero") - - -@dataclass(frozen=True, slots=True) -class ToolCompleted: - """Execution of one normalized model tool call settled.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.TOOL_COMPLETED - turn: int - tool_call: ModelToolCall - result: ToolCallResult - - -@dataclass(frozen=True, slots=True) -class TurnCompleted: - """One provider turn and its requested tool calls settled.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.TURN_COMPLETED - turn: int - - -@dataclass(frozen=True, slots=True) -class AgentFinished: - """One agent run reached its structured terminal result.""" - - kind: ClassVar[AgentEventKind] = AgentEventKind.AGENT_FINISHED - result: AgentRunResult - - -AgentEventPayload: TypeAlias = ( - AgentStarted - | AgentContinued - | TurnStarted - | SteeringApplied - | SteeringDiscarded - | ContextPressureEvaluated - | CompactionStarted - | CompactionCompleted - | CompactionFailed - | MessageCompleted - | AssistantTextDelta - | AssistantThinkingDelta - | ToolStarted - | ToolProgressed - | ToolCompleted - | TurnCompleted - | AgentFinished -) - - -@dataclass(frozen=True, slots=True) -class AgentEvent: - """Immutable event envelope shared by every lifecycle payload.""" - - agent_id: str - run_id: str - sequence: int - payload: AgentEventPayload - - @property - def kind(self) -> AgentEventKind: - return self.payload.kind - - -class AgentEventSink(Protocol): - """Read-only observer receiving ordered events for an agent run.""" - - async def emit(self, event: AgentEvent) -> None: - """Observe an event without changing agent behavior.""" - - -class AgentEventSinkError(RuntimeError): - """Raised after one or more sinks fail during event fan-out.""" - - -class CompositeAgentEventSink: - """Forward each event to multiple ordered read-only observers.""" - - def __init__(self, sinks: Iterable[AgentEventSink]) -> None: - self.sinks = tuple(sinks) - - async def emit(self, event: AgentEvent) -> None: - errors: list[Exception] = [] - for sink in self.sinks: - try: - await sink.emit(event) - except Exception as exc: - errors.append(exc) - if errors: - raise AgentEventSinkError( - f"{len(errors)} agent event sink(s) failed" - ) from errors[0] - - -class AgentEventEmitter: - """Create ordered event envelopes and isolate observer failures.""" - - def __init__( - self, - *, - agent_id: str, - sink: AgentEventSink | None, - logger: logging.Logger, - ) -> None: - self.agent_id = agent_id - self.sink = sink - self.logger = logger - self._run_id: str | None = None - self._sequence = 0 - - def begin_run(self) -> str: - """Open a new event sequence and return its correlation id.""" - - if self._run_id is not None: - raise RuntimeError("an agent event run is already active") - self._run_id = uuid4().hex - self._sequence = 0 - return self._run_id - - def end_run(self, run_id: str) -> None: - """Close the active event sequence.""" - - if self._run_id != run_id: - raise RuntimeError("agent event run id does not match active run") - self._run_id = None - self._sequence = 0 - - async def emit(self, payload: AgentEventPayload) -> AgentEvent: - """Publish one event without allowing Sink failures to fail the run.""" - - if self._run_id is None: - raise RuntimeError("agent event run is not active") - - self._sequence += 1 - event = AgentEvent( - agent_id=self.agent_id, - run_id=self._run_id, - sequence=self._sequence, - payload=payload, - ) - if self.sink is None: - return event - - try: - await self.sink.emit(event) - except Exception as exc: - self.logger.warning( - "Agent event sink failed for %s: %s", - event.kind, - exc, - ) - return event diff --git a/src/ejagent/agent/model_compactor.py b/src/ejagent/agent/model_compactor.py deleted file mode 100644 index d8a2832..0000000 --- a/src/ejagent/agent/model_compactor.py +++ /dev/null @@ -1,64 +0,0 @@ -from __future__ import annotations - -from typing import Protocol - -from ejagent.agent.cancellation import CancellationToken -from ejagent.agent.compaction import ( - CompactionRequest, - CompactorOutput, -) -from ejagent.agent.context_builder import ContextBuildResult -from ejagent.providers.base import AssistantMessage, ModelAdapter - - -class CompactionContextBuilder(Protocol): - """Build one provider request from a prepared compaction operation.""" - - def __call__(self, request: CompactionRequest) -> ContextBuildResult: - """Return the model context, including the application-owned prompt.""" - - -class ModelCompactor: - """Adapt a borrowed ``ModelAdapter`` into the ``Compactor`` protocol. - - The caller owns the model lifecycle and supplies all prompt construction. - This keeps the Core provider-neutral and avoids embedding a summary policy. - """ - - def __init__( - self, - model: ModelAdapter, - *, - context_builder: CompactionContextBuilder, - source: str, - ) -> None: - source = source.strip() - if not source: - raise ValueError("model compactor source must not be empty") - self.model = model - self.context_builder = context_builder - self.source = source - - async def compact( - self, - request: CompactionRequest, - *, - cancellation: CancellationToken | None = None, - ) -> CompactorOutput: - if cancellation is not None: - cancellation.raise_if_cancelled() - context = self.context_builder(request) - if not isinstance(context, ContextBuildResult): - raise TypeError("CompactionContextBuilder must return ContextBuildResult") - response = await self.model.complete( - context, - cancellation=cancellation, - ) - if not isinstance(response, AssistantMessage): - raise TypeError("ModelAdapter.complete() must return AssistantMessage") - if cancellation is not None: - cancellation.raise_if_cancelled() - content = response.content.strip() if response.content is not None else "" - if not content: - raise RuntimeError("compaction model returned empty content") - return CompactorOutput(content=content, source=self.source) diff --git a/src/ejagent/agent/orchestrator.py b/src/ejagent/agent/orchestrator.py deleted file mode 100644 index 73d3fca..0000000 --- a/src/ejagent/agent/orchestrator.py +++ /dev/null @@ -1,866 +0,0 @@ -from __future__ import annotations - -import asyncio -from collections.abc import AsyncIterator -from contextlib import suppress -from typing import Any, Protocol - -from ejagent.agent.behavior import ( - BehaviorAction, - BehaviorDecision, - BehaviorHook, - BehaviorHookError, - TurnSnapshot, -) -from ejagent.agent.cancellation import ( - AgentCancelledError, - CancellationSource, - CancellationToken, -) -from ejagent.agent.compaction import ( - CompactionRuntime, - CompactionStatus, - CompactionTrigger, - Compactor, -) -from ejagent.agent.context_builder import ( - AgentContextBuilder, - ContextBuildResult, -) -from ejagent.agent.context_management import ( - AutoCompactionPolicy, - CompactionPolicy, - CompactionPreparation, - MessageTokenEstimator, - estimate_context_usage, -) -from ejagent.agent.control import ( - ContinueRejectedError, - _continue_history_rejection_reason, - _SteeringQueue, -) -from ejagent.agent.events import ( - AgentContinued, - AgentEventEmitter, - AgentFinished, - AgentStarted, - AssistantTextDelta, - AssistantThinkingDelta, - ContextPressureEvaluated, - MessageCompleted, - SteeringApplied, - SteeringDiscarded, - TurnCompleted, - TurnStarted, -) -from ejagent.agent.result import AgentRunResult, RunStatus, StopReason -from ejagent.agent.runtime_policy import RuntimePolicy -from ejagent.agent.state import AgentState -from ejagent.agent.tool_runtime import ( - RepeatedToolCallError, - ToolRuntime, -) -from ejagent.agent.types import ToolCallResult, ToolControl -from ejagent.agent.usage import UsageAccumulator -from ejagent.handlers.definition import ToolDefinition, ToolEffect -from ejagent.plugins.skill.skill_manager import SkillManager -from ejagent.providers.base import ( - AssistantMessage, - ContextOverflowError, - ModelResponseCompleted, - ModelStreamEvent, - ModelTextDelta, - ModelThinkingDelta, - ModelToolCall, - serialize_assistant_message, -) - -TOOL_COMPLETION_RETRY_PROMPT = """ -Explicit-finish mode requires a completing tool call to finish the task. -If the work is complete, call a tool that returns completion control now. -Do not end with plain text. -""".strip() - - -class ModelStream(Protocol): - """Provider stream shape consumed by the orchestrator.""" - - def __call__( - self, - context: ContextBuildResult, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: ... - - -class _AutomaticCompactionError(RuntimeError): - """Automatic compaction failed before provider dispatch or retry.""" - - -class AgentOrchestrator: - """Coordinate one agent task across model, state, and tool runtimes.""" - - def __init__( - self, - *, - agent_id: str, - state: AgentState, - context_builder: AgentContextBuilder, - model_stream: ModelStream, - tool_runtime: ToolRuntime, - skill_manager: SkillManager | None, - policy: RuntimePolicy, - compaction_policy: CompactionPolicy | None = None, - auto_compaction_policy: AutoCompactionPolicy | None = None, - compactor: Compactor | None = None, - compaction_runtime: CompactionRuntime | None = None, - context_token_estimator: MessageTokenEstimator | None = None, - behavior_hooks: tuple[BehaviorHook, ...] = (), - event_emitter: AgentEventEmitter, - ) -> None: - self.agent_id = agent_id - self.state = state - self.context_builder = context_builder - self.model_stream = model_stream - self.tool_runtime = tool_runtime - self.skill_manager = skill_manager - self.policy = policy - self.compaction_policy = compaction_policy - self.auto_compaction_policy = auto_compaction_policy - self.compactor = compactor - self.compaction_runtime = compaction_runtime - self.context_token_estimator = context_token_estimator - self.behavior_hooks = behavior_hooks - self.event_emitter = event_emitter - self._usage = UsageAccumulator() - - @property - def tools(self) -> list[dict[str, Any]]: - """Return every tool definition available to the model.""" - - return self.tool_runtime.tools - - @property - def tool_definitions(self) -> tuple[ToolDefinition, ...]: - """Return the canonical tools used by Core routing and scheduling.""" - - return self.tool_runtime.tool_definitions - - async def run( - self, - *, - task: str, - cancellation: CancellationToken | None = None, - steering: _SteeringQueue | None = None, - ) -> AgentRunResult: - """Run one task and return its structured terminal result.""" - - return await self._run( - task=task, - cancellation=cancellation, - steering=steering, - ) - - async def continue_run( - self, - *, - cancellation: CancellationToken | None = None, - steering: _SteeringQueue | None = None, - ) -> AgentRunResult: - """Resume existing history in a distinct Run without a user message.""" - - reason = _continue_history_rejection_reason( - self.state.messages, - self.state.last_run_result, - ) - if reason is not None: - raise ContinueRejectedError(reason) - return await self._run( - task=None, - cancellation=cancellation, - steering=steering, - ) - - async def _run( - self, - *, - task: str | None, - cancellation: CancellationToken | None, - steering: _SteeringQueue | None, - ) -> AgentRunResult: - """Execute one task or Continue Run through the shared lifecycle.""" - - token = cancellation or CancellationSource().token - self._usage = UsageAccumulator() - run_id = self.event_emitter.begin_run() - try: - caller_cancellation: asyncio.CancelledError | None = None - try: - if task is None: - self.state.begin_continue() - await self.event_emitter.emit(AgentContinued()) - else: - self.state.begin_task(task) - await self.event_emitter.emit(AgentStarted(task)) - await token.run(self._prepare_task()) - result = await self._run_loop(token, steering, run_id=run_id) - except AgentCancelledError as exc: - result = self._cancelled(str(exc)) - except RepeatedToolCallError as exc: - result = self._failure( - StopReason.REPEATED_TOOL_CALL, - str(exc), - ) - except ContextOverflowError as exc: - result = self._failure(StopReason.CONTEXT_OVERFLOW, str(exc)) - except _AutomaticCompactionError as exc: - result = self._failure(StopReason.COMPACTION_FAILED, str(exc)) - except asyncio.CancelledError as exc: - caller_cancellation = exc - result = self._cancelled("agent run coroutine was cancelled") - except Exception as exc: - result = self._failure(StopReason.RUNTIME_ERROR, str(exc)) - if steering is not None: - for control in steering.close(): - await self.event_emitter.emit( - SteeringDiscarded(control, result.stop_reason) - ) - self._commit_result(result) - await self.event_emitter.emit(AgentFinished(result)) - if caller_cancellation is not None: - raise caller_cancellation - return result - finally: - self.event_emitter.end_run(run_id) - - async def _prepare_task(self) -> None: - await self.tool_runtime.on_task_start() - self._activate_explicit_skill() - - async def _run_loop( - self, - cancellation: CancellationToken, - steering: _SteeringQueue | None = None, - *, - run_id: str, - ) -> AgentRunResult: - for _ in range(self.policy.max_steps): - cancellation.raise_if_cancelled() - budget_failure = self._budget_failure() - if budget_failure is not None: - return budget_failure - turn = self.state.advance_turn() - await self.event_emitter.emit(TurnStarted(turn)) - try: - response = await self._chat_next_turn(cancellation, steering) - message = response.message - self._usage.record(response.usage) - self.state.add_message( - serialize_assistant_message( - message, - usage=response.usage, - ) - ) - await self.event_emitter.emit( - MessageCompleted(turn, message, response.usage) - ) - cancellation.raise_if_cancelled() - - if not message.tool_calls: - if steering is not None and steering.has_pending: - self.state.no_tool_response_count = 0 - elif not self.policy.require_explicit_finish: - if message.content: - return AgentRunResult( - status=RunStatus.COMPLETED, - stop_reason=StopReason.TEXT_RESPONSE, - turns=self.state.turn, - output=message.content, - usage=self._usage.snapshot(), - ) - return self._failure( - StopReason.EMPTY_RESPONSE, - "chat completion returned empty content", - ) - - else: - self.state.no_tool_response_count += 1 - if ( - self.state.no_tool_response_count - >= self.policy.max_no_tool_responses - ): - return self._failure( - StopReason.MAX_NO_TOOL_RESPONSES, - "explicit-finish mode produced plain text without a " - "completing tool call " - f"{self.state.no_tool_response_count} " - "consecutive times", - ) - else: - self.state.no_tool_response_count = 0 - tool_result = await self._execute_tool_calls( - message, - cancellation, - ) - self.state.add_messages(list(tool_result.messages)) - if tool_result.cancelled: - raise AgentCancelledError( - tool_result.error - or cancellation.reason - or "agent run was aborted" - ) - terminal_result = self._terminal_tool_result(tool_result) - if terminal_result is not None: - return terminal_result - finally: - await self.event_emitter.emit(TurnCompleted(turn)) - - behavior_result = await self._evaluate_after_turn( - run_id=run_id, - turn=turn, - response=message, - cancellation=cancellation, - ) - if behavior_result is not None: - return behavior_result - - return self._failure( - StopReason.MAX_STEPS, - f"agent {self.agent_id!r} did not finish within " - f"{self.policy.max_steps} steps", - ) - - async def _evaluate_after_turn( - self, - *, - run_id: str, - turn: int, - response: AssistantMessage, - cancellation: CancellationToken, - ) -> AgentRunResult | None: - for hook in self.behavior_hooks: - snapshot = TurnSnapshot( - agent_id=self.agent_id, - run_id=run_id, - turn=turn, - task=self.state.task, - response=response, - usage=self._usage.snapshot(), - messages=self.state.messages, - ) - try: - decision = await cancellation.run( - hook.after_turn(snapshot, cancellation=cancellation) - ) - if decision is not None and not isinstance(decision, BehaviorDecision): - raise TypeError("after_turn must return BehaviorDecision or None") - except (AgentCancelledError, asyncio.CancelledError): - raise - except Exception as exc: - raise BehaviorHookError(hook, exc) from exc - - if decision is not None and decision.action is BehaviorAction.STOP: - return AgentRunResult( - status=RunStatus.COMPLETED, - stop_reason=StopReason.BEHAVIOR_STOP, - turns=self.state.turn, - output=( - decision.output - if decision.output is not None - else response.content - ), - usage=self._usage.snapshot(), - ) - return None - - async def _chat_next_turn( - self, - cancellation: CancellationToken, - steering: _SteeringQueue | None = None, - ) -> ModelResponseCompleted: - context, preparation = await self._build_evaluated_context(steering) - auto_policy = self.auto_compaction_policy - if ( - auto_policy is not None - and auto_policy.enabled - and auto_policy.compact_on_pressure - and preparation is not None - ): - result = await self._compact_automatically( - cancellation, - trigger=CompactionTrigger.PRESSURE, - preparation=preparation, - ) - if result: - context, _ = await self._build_evaluated_context(steering) - - overflow_retries = 0 - while True: - try: - return await self._request_model(context, cancellation) - except ContextOverflowError as exc: - if ( - exc.response_started - or auto_policy is None - or not auto_policy.enabled - or not auto_policy.recover_on_overflow - or overflow_retries >= auto_policy.max_overflow_retries - ): - raise - compacted = await self._compact_automatically( - cancellation, - trigger=CompactionTrigger.OVERFLOW, - ) - if not compacted: - raise - overflow_retries += 1 - context, _ = await self._build_evaluated_context(steering) - - async def _build_evaluated_context( - self, - steering: _SteeringQueue | None = None, - ) -> tuple[ContextBuildResult, CompactionPreparation | None]: - while True: - await self._apply_steering(steering) - context = self.context_builder.build( - self.state, - tools=self.tool_definitions, - transient_messages=self._runtime_context_messages(), - ) - preparation = await self._evaluate_context_pressure(context) - if steering is None or not steering.has_pending: - return context, preparation - - async def _apply_steering(self, steering: _SteeringQueue | None) -> None: - if steering is None: - return - controls = steering.drain() - if not controls: - return - self.state.no_tool_response_count = 0 - for control in controls: - self.state.add_message({"role": "user", "content": control.content}) - await self.event_emitter.emit( - SteeringApplied(control=control, target_turn=self.state.turn) - ) - self._activate_explicit_skill() - - async def _request_model( - self, - context: ContextBuildResult, - cancellation: CancellationToken, - ) -> ModelResponseCompleted: - self._usage.begin_request() - stream = self.model_stream( - context, - cancellation=cancellation, - ) - iterator = stream.__aiter__() - completed: ModelResponseCompleted | None = None - response_started = False - try: - while completed is None: - cancellation.raise_if_cancelled() - try: - event = await cancellation.run(anext(iterator)) - except StopAsyncIteration: - break - - if isinstance(event, ModelTextDelta): - response_started = True - await self.event_emitter.emit( - AssistantTextDelta(self.state.turn, event.delta) - ) - elif isinstance(event, ModelThinkingDelta): - response_started = True - await self.event_emitter.emit( - AssistantThinkingDelta( - self.state.turn, - event.delta, - ) - ) - elif isinstance(event, ModelResponseCompleted): - completed = event - else: - raise TypeError( - "model stream returned an unsupported event: " - f"{type(event).__name__}" - ) - except ContextOverflowError as exc: - if response_started and not exc.response_started: - raise ContextOverflowError( - str(exc), - response_started=True, - ) from exc - raise - finally: - close = getattr(iterator, "aclose", None) - if close is not None: - with suppress(Exception): - await close() - - if completed is None: - raise RuntimeError("model stream ended without a completed response") - return completed - - async def _evaluate_context_pressure( - self, - context: ContextBuildResult, - ) -> CompactionPreparation | None: - policy = self.compaction_policy - if policy is None: - return None - - estimate = estimate_context_usage( - context.agent_messages, - tools=context.tools, - estimator=self.context_token_estimator, - ) - decision = policy.evaluate(estimate) - preparation = ( - policy.prepare( - self.state.messages, - estimator=self.context_token_estimator, - ) - if decision.should_compact - else None - ) - await self.event_emitter.emit( - ContextPressureEvaluated( - turn=self.state.turn, - decision=decision, - preparation=preparation, - ) - ) - return preparation - - async def _compact_automatically( - self, - cancellation: CancellationToken, - *, - trigger: CompactionTrigger, - preparation: CompactionPreparation | None = None, - ) -> bool: - runtime = self.compaction_runtime - compactor = self.compactor - if runtime is None or compactor is None: - raise _AutomaticCompactionError( - "automatic compaction is missing its runtime or Compactor" - ) - result = await runtime.compact_in_active_run( - compactor, - cancellation=cancellation, - trigger=trigger, - preparation=preparation, - ) - if result.status is CompactionStatus.CANCELLED: - raise AgentCancelledError( - result.error or cancellation.reason or "agent run was aborted" - ) - if result.status is CompactionStatus.FAILED: - raise _AutomaticCompactionError( - result.error or "automatic compaction failed" - ) - return result.status is CompactionStatus.COMPLETED - - def _budget_failure(self) -> AgentRunResult | None: - limit = self.policy.max_run_tokens - if limit is None: - return None - - usage = self._usage.snapshot() - if usage.request_count == 0: - return None - if not usage.complete: - return self._failure( - StopReason.USAGE_UNAVAILABLE, - "run token budget cannot continue because " - f"{usage.missing_request_count} model request(s) did not " - "report usage", - ) - if usage.total_tokens >= limit: - return self._failure( - StopReason.TOKEN_BUDGET_EXCEEDED, - f"run used {usage.total_tokens} tokens and cannot start " - f"another model request under the {limit}-token budget", - ) - return None - - def _runtime_context_messages(self) -> list[dict[str, str]]: - if self.state.no_tool_response_count == 0: - return [] - return [ - { - "role": "system", - "content": self._tool_completion_retry_prompt( - self.state.no_tool_response_count - ), - } - ] - - def _tool_completion_retry_prompt( - self, - no_tool_response_count: int, - ) -> str: - if no_tool_response_count <= 1: - return TOOL_COMPLETION_RETRY_PROMPT - return ( - TOOL_COMPLETION_RETRY_PROMPT - + "\n\n" - + ( - f"Retry {no_tool_response_count}/" - f"{self.policy.max_no_tool_responses}: " - "the previous response still did not include a tool call." - ) - ) - - def _activate_explicit_skill(self) -> None: - if self.skill_manager is None: - return - - skill_name = self.skill_manager.select_explicit_skill(self.state.messages) - if skill_name is not None: - self.state.active_skill_name = skill_name - - async def _execute_tool_calls( - self, - message: AssistantMessage, - cancellation: CancellationToken, - ) -> ToolCallResult: - result_messages: list[dict[str, Any]] = [] - - tool_calls = message.tool_calls or () - index = 0 - while index < len(tool_calls): - tool_call = tool_calls[index] - next_index = index + 1 - if ( - self.policy.parallel_tool_calls - and self.tool_runtime.tool_effect(tool_call.name) - is ToolEffect.READ_ONLY - ): - while ( - next_index < len(tool_calls) - and self.tool_runtime.tool_effect(tool_calls[next_index].name) - is ToolEffect.READ_ONLY - ): - next_index += 1 - - batch = tool_calls[index:next_index] - if len(batch) > 1: - tool_result = await self._execute_read_only_tool_calls( - batch, - cancellation, - ) - else: - tool_result = await self.tool_runtime.execute_tool_call( - tool_call, - cancellation=cancellation, - ) - result_messages.extend(tool_result.messages) - if tool_result.cancelled: - reason = ( - tool_result.error or cancellation.reason or "agent run was aborted" - ) - for pending_call in tool_calls[next_index:]: - pending_result = await self.tool_runtime.cancel_tool_call( - pending_call, - reason=reason, - ) - result_messages.extend(pending_result.messages) - return ToolCallResult( - tuple(result_messages), - error=reason, - cancelled=True, - ) - if tool_result.control is not ToolControl.CONTINUE: - return ToolCallResult( - tuple(result_messages), - control=tool_result.control, - output=tool_result.output, - ) - index = next_index - - return ToolCallResult(tuple(result_messages)) - - async def _execute_read_only_tool_calls( - self, - tool_calls: tuple[ModelToolCall, ...], - cancellation: CancellationToken, - ) -> ToolCallResult: - """Run safe calls concurrently while committing results in source order.""" - - result_messages: list[dict[str, Any]] = [] - if not self.tool_runtime.can_parallelize_tool_calls(tool_calls): - return await self._execute_sequential_tool_calls( - tool_calls, - cancellation, - ) - limit = self.policy.max_parallel_tool_calls or len(tool_calls) - offset = 0 - while offset < len(tool_calls): - batch = tool_calls[offset : offset + limit] - prepared = await self.tool_runtime.prepare_parallel_tool_calls(batch) - if not prepared: - return await self._execute_sequential_tool_calls( - tool_calls[offset:], - cancellation, - prefix=result_messages, - ) - - results = await asyncio.gather( - *( - self.tool_runtime.execute_tool_call( - tool_call, - cancellation=cancellation, - prepared=True, - defer_completion=True, - ) - for tool_call in batch - ) - ) - for tool_call, result in zip(batch, results, strict=True): - await self.tool_runtime.complete_tool_call(tool_call, result) - result_messages.extend(result.messages) - - cancelled = next((result for result in results if result.cancelled), None) - if cancelled is not None: - reason = ( - cancelled.error or cancellation.reason or "agent run was aborted" - ) - for pending_call in tool_calls[offset + len(batch) :]: - pending_result = await self.tool_runtime.cancel_tool_call( - pending_call, - reason=reason, - ) - result_messages.extend(pending_result.messages) - return ToolCallResult( - tuple(result_messages), - error=reason, - cancelled=True, - ) - - terminal = next( - ( - result - for result in results - if result.control is not ToolControl.CONTINUE - ), - None, - ) - if terminal is not None: - reason = "skipped after a terminal read-only tool result" - for pending_call in tool_calls[offset + len(batch) :]: - pending_result = await self.tool_runtime.cancel_tool_call( - pending_call, - reason=reason, - ) - result_messages.extend(pending_result.messages) - return ToolCallResult( - tuple(result_messages), - control=terminal.control, - output=terminal.output, - ) - offset += len(batch) - - return ToolCallResult(tuple(result_messages)) - - async def _execute_sequential_tool_calls( - self, - tool_calls: tuple[ModelToolCall, ...], - cancellation: CancellationToken, - *, - prefix: list[dict[str, Any]] | None = None, - ) -> ToolCallResult: - result_messages = list(prefix or ()) - for index, tool_call in enumerate(tool_calls): - result = await self.tool_runtime.execute_tool_call( - tool_call, - cancellation=cancellation, - ) - result_messages.extend(result.messages) - if result.cancelled: - reason = result.error or cancellation.reason or "agent run was aborted" - for pending_call in tool_calls[index + 1 :]: - pending_result = await self.tool_runtime.cancel_tool_call( - pending_call, - reason=reason, - ) - result_messages.extend(pending_result.messages) - return ToolCallResult( - tuple(result_messages), - error=reason, - cancelled=True, - ) - if result.control is not ToolControl.CONTINUE: - return ToolCallResult( - tuple(result_messages), - control=result.control, - output=result.output, - ) - return ToolCallResult(tuple(result_messages)) - - def _cancelled(self, error: str) -> AgentRunResult: - return AgentRunResult( - status=RunStatus.CANCELLED, - stop_reason=StopReason.EXTERNAL_ABORT, - turns=self.state.turn, - error=error, - usage=self._usage.snapshot(), - ) - - def _terminal_tool_result( - self, - tool_result: ToolCallResult, - ) -> AgentRunResult | None: - if tool_result.control is ToolControl.CONTINUE: - return None - if tool_result.control is ToolControl.COMPLETE: - return AgentRunResult( - status=RunStatus.COMPLETED, - stop_reason=StopReason.TOOL_COMPLETION, - turns=self.state.turn, - output=tool_result.output, - usage=self._usage.snapshot(), - ) - if tool_result.control is ToolControl.REJECT: - return AgentRunResult( - status=RunStatus.REJECTED, - stop_reason=StopReason.TOOL_REJECTED, - turns=self.state.turn, - output=tool_result.output, - error="tool execution was rejected", - usage=self._usage.snapshot(), - ) - return AgentRunResult( - status=RunStatus.CANCELLED, - stop_reason=StopReason.TOOL_CANCELLED, - turns=self.state.turn, - output=tool_result.output, - error="tool execution was cancelled", - usage=self._usage.snapshot(), - ) - - def _failure( - self, - stop_reason: StopReason, - error: str, - ) -> AgentRunResult: - return AgentRunResult( - status=RunStatus.FAILED, - stop_reason=stop_reason, - turns=self.state.turn, - error=error, - usage=self._usage.snapshot(), - ) - - def _commit_result(self, result: AgentRunResult) -> None: - self.state.last_run_result = result - if result.status is RunStatus.COMPLETED: - self.state.complete(result.output or "") - elif result.status is RunStatus.REJECTED: - self.state.reject(result.output) - elif result.status is RunStatus.CANCELLED: - self.state.cancel(result.output) - else: - self.state.fail(result.error or result.stop_reason.value) diff --git a/src/ejagent/agent/result.py b/src/ejagent/agent/result.py deleted file mode 100644 index dbf5d73..0000000 --- a/src/ejagent/agent/result.py +++ /dev/null @@ -1,65 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass, field -from enum import StrEnum - -from ejagent.agent.usage import RunUsage - - -class RunStatus(StrEnum): - """Terminal status of one orchestrator run.""" - - COMPLETED = "completed" - FAILED = "failed" - CANCELLED = "cancelled" - REJECTED = "rejected" - - -class StopReason(StrEnum): - """Reason why an orchestrator run stopped.""" - - TEXT_RESPONSE = "text_response" - TOOL_COMPLETION = "tool_completion" - TOOL_REJECTED = "tool_rejected" - TOOL_CANCELLED = "tool_cancelled" - BEHAVIOR_STOP = "behavior_stop" - EXTERNAL_ABORT = "external_abort" - EMPTY_RESPONSE = "empty_response" - MAX_STEPS = "max_steps" - MAX_NO_TOOL_RESPONSES = "max_no_tool_responses" - REPEATED_TOOL_CALL = "repeated_tool_call" - TOKEN_BUDGET_EXCEEDED = "token_budget_exceeded" - USAGE_UNAVAILABLE = "usage_unavailable" - CONTEXT_OVERFLOW = "context_overflow" - COMPACTION_FAILED = "compaction_failed" - RUNTIME_ERROR = "runtime_error" - - -@dataclass(frozen=True, slots=True) -class AgentRunResult: - """Structured terminal result produced by ``AgentOrchestrator``.""" - - status: RunStatus - stop_reason: StopReason - turns: int - output: str | None = None - error: str | None = None - usage: RunUsage = field(default_factory=RunUsage) - - @property - def succeeded(self) -> bool: - return self.status is RunStatus.COMPLETED - - def raise_for_status(self) -> None: - """Raise a compatibility error when the run did not complete.""" - - if not self.succeeded: - raise AgentRunError(self) - - -class AgentRunError(RuntimeError): - """Compatibility exception wrapping a structured failed run result.""" - - def __init__(self, result: AgentRunResult) -> None: - self.result = result - super().__init__(result.error or result.stop_reason.value) diff --git a/src/ejagent/agent/runtime_policy.py b/src/ejagent/agent/runtime_policy.py deleted file mode 100644 index a549926..0000000 --- a/src/ejagent/agent/runtime_policy.py +++ /dev/null @@ -1,33 +0,0 @@ -from dataclasses import dataclass - - -@dataclass(frozen=True, slots=True) -class RuntimePolicy: - """Policy controlling the provider-tool loop for one agent.""" - - max_steps: int = 20 - max_no_tool_responses: int = 3 - max_repeated_tool_calls: int = 3 - max_run_tokens: int | None = None - require_explicit_finish: bool = False - parallel_tool_calls: bool = False - max_parallel_tool_calls: int | None = None - - def __post_init__(self) -> None: - if not isinstance(self.parallel_tool_calls, bool): - raise TypeError("parallel_tool_calls must be a boolean") - if self.max_steps <= 0: - raise ValueError("max_steps must be greater than zero") - if self.max_no_tool_responses <= 0: - raise ValueError("max_no_tool_responses must be greater than zero") - if self.max_repeated_tool_calls <= 0: - raise ValueError("max_repeated_tool_calls must be greater than zero") - if self.max_run_tokens is not None and self.max_run_tokens <= 0: - raise ValueError("max_run_tokens must be greater than zero") - if self.max_parallel_tool_calls is not None: - if isinstance(self.max_parallel_tool_calls, bool) or not isinstance( - self.max_parallel_tool_calls, int - ): - raise TypeError("max_parallel_tool_calls must be an integer") - if self.max_parallel_tool_calls <= 0: - raise ValueError("max_parallel_tool_calls must be greater than zero") diff --git a/src/ejagent/agent/state.py b/src/ejagent/agent/state.py deleted file mode 100644 index 841e86e..0000000 --- a/src/ejagent/agent/state.py +++ /dev/null @@ -1,136 +0,0 @@ -from __future__ import annotations - -from copy import deepcopy -from dataclasses import dataclass, field -from enum import StrEnum - -from ejagent.agent.result import AgentRunResult -from ejagent.agent.types import AgentMessage - - -class AgentStatus(StrEnum): - """Lifecycle status of the current agent task.""" - - IDLE = "idle" - RUNNING = "running" - COMPLETED = "completed" - FAILED = "failed" - CANCELLED = "cancelled" - REJECTED = "rejected" - - -@dataclass(slots=True) -class AgentState: - """Persistent conversation and current-task state for one agent.""" - - messages: list[AgentMessage] = field(default_factory=list) - task: str | None = None - status: AgentStatus = AgentStatus.IDLE - turn: int = 0 - no_tool_response_count: int = 0 - active_skill_name: str | None = None - result: str | None = None - error: str | None = None - last_run_result: AgentRunResult | None = None - - def reset(self, messages: list[AgentMessage]) -> None: - """Replace conversation history and clear the current task state.""" - - self.messages = [dict(message) for message in messages] - self.task = None - self.status = AgentStatus.IDLE - self.turn = 0 - self.no_tool_response_count = 0 - self.active_skill_name = None - self.result = None - self.error = None - self.last_run_result = None - - def begin_task(self, task: str) -> None: - """Start a new task while preserving the conversation history.""" - - self.task = task - self.status = AgentStatus.RUNNING - self.turn = 0 - self.no_tool_response_count = 0 - self.active_skill_name = None - self.result = None - self.error = None - self.last_run_result = None - self.messages.append({"role": "user", "content": task}) - - def begin_continue(self) -> None: - """Start a new Run without adding another user task message.""" - - self.task = None - self.status = AgentStatus.RUNNING - self.turn = 0 - self.no_tool_response_count = 0 - self.active_skill_name = None - self.result = None - self.error = None - self.last_run_result = None - - def advance_turn(self) -> int: - """Record and return the next model turn number.""" - - self.turn += 1 - return self.turn - - def add_message(self, message: AgentMessage) -> None: - """Append one conversation message.""" - - self.messages.append(dict(message)) - - def add_messages(self, messages: list[AgentMessage]) -> None: - """Append multiple conversation messages.""" - - self.messages.extend(dict(message) for message in messages) - - def replace_messages(self, messages: list[AgentMessage]) -> None: - """Atomically replace persistent history with detached messages.""" - - self.messages = deepcopy(messages) - - def complete(self, result: str) -> None: - """Mark the current task as completed.""" - - self.status = AgentStatus.COMPLETED - self.result = result - self.error = None - - def fail(self, error: Exception | str) -> None: - """Mark the current task as failed without discarding its history.""" - - self.status = AgentStatus.FAILED - self.result = None - self.error = str(error) - - def cancel(self, result: str | None = None) -> None: - """Mark the current task as cancelled.""" - - self.status = AgentStatus.CANCELLED - self.result = result - self.error = None - - def reject(self, result: str | None = None) -> None: - """Mark the current task as rejected by a tool or policy.""" - - self.status = AgentStatus.REJECTED - self.result = result - self.error = None - - def snapshot(self) -> AgentState: - """Return an independent copy suitable for observation or persistence.""" - - return AgentState( - messages=deepcopy(self.messages), - task=self.task, - status=self.status, - turn=self.turn, - no_tool_response_count=self.no_tool_response_count, - active_skill_name=self.active_skill_name, - result=self.result, - error=self.error, - last_run_result=self.last_run_result, - ) diff --git a/src/ejagent/agent/tool_runtime.py b/src/ejagent/agent/tool_runtime.py deleted file mode 100644 index 9c50e3d..0000000 --- a/src/ejagent/agent/tool_runtime.py +++ /dev/null @@ -1,517 +0,0 @@ -import asyncio -import json -import logging -from collections.abc import Iterable, Mapping, Sequence -from typing import Any - -from ejagent.agent.cancellation import ( - AgentCancelledError, - CancellationSource, - CancellationToken, -) -from ejagent.agent.events import ( - AgentEventEmitter, - ToolCompleted, - ToolProgressed, - ToolStarted, -) -from ejagent.agent.state import AgentState -from ejagent.agent.types import ( - StepOutcome, - ToolCallResult, - ToolControl, - ToolProgressReporter, - ToolProgressUpdate, -) -from ejagent.handlers.base import BaseHandler -from ejagent.handlers.definition import ToolDefinition, ToolEffect -from ejagent.middleware import ( - ToolCallContext, - ToolMiddleware, - ToolNext, - compose_tool_middlewares, -) -from ejagent.providers.base import ModelToolCall - - -class RepeatedToolCallError(RuntimeError): - """Raised when an identical tool call reaches the configured limit.""" - - -class _ScopedToolProgressReporter(ToolProgressReporter): - """Serialize updates and reject them after one tool call settles.""" - - def __init__( - self, - *, - event_emitter: AgentEventEmitter, - turn: int, - tool_call: ModelToolCall, - cancellation: CancellationToken, - ) -> None: - self._event_emitter = event_emitter - self._turn = turn - self._tool_call = tool_call - self._cancellation = cancellation - self._lock = asyncio.Lock() - self._accepting = True - - async def report(self, update: ToolProgressUpdate) -> None: - if not isinstance(update, ToolProgressUpdate): - raise TypeError( - "tool progress reporter requires ToolProgressUpdate, " - f"got {type(update).__name__}" - ) - - async with self._lock: - if not self._accepting or self._cancellation.cancelled: - return - emission = asyncio.create_task( - self._event_emitter.emit( - ToolProgressed( - self._turn, - self._tool_call, - update, - ) - ) - ) - try: - await asyncio.shield(emission) - except asyncio.CancelledError: - await asyncio.gather(emission, return_exceptions=True) - raise - - async def close(self) -> None: - """Wait for an in-flight update, then reject future reports.""" - - async with self._lock: - self._accepting = False - - -class ToolRuntime: - """Lifecycle, routing, middleware, and execution for tool handlers.""" - - def __init__( - self, - handlers: Iterable[BaseHandler], - middlewares: Iterable[ToolMiddleware], - *, - state: AgentState, - logger: logging.Logger, - event_emitter: AgentEventEmitter, - max_repeated_tool_calls: int = 3, - ) -> None: - if max_repeated_tool_calls <= 0: - raise ValueError("max_repeated_tool_calls must be greater than zero") - self.handlers = list(handlers) - self.middlewares = list(middlewares) - self.state = state - self.logger = logger - self.event_emitter = event_emitter - self.max_repeated_tool_calls = max_repeated_tool_calls - self._tool_routes: dict[str, BaseHandler] = {} - self._active_middlewares: list[ToolMiddleware] = [] - self._tool_chain: ToolNext | None = None - self._started = False - self._last_tool_signature: tuple[str, str] | None = None - self._repeated_tool_calls = 0 - - @property - def tools(self) -> list[dict[str, Any]]: - return [tool.to_openai_tool() for tool in self.tool_definitions] - - @property - def tool_definitions(self) -> tuple[ToolDefinition, ...]: - """Return the canonical tool collection registered with this runtime.""" - - definitions: list[ToolDefinition] = [] - for handler in self.handlers: - for tool in handler.tool_definitions: - effect = handler.tool_effect(tool.name) - definitions.append( - tool if effect is tool.effect else tool.with_effect(effect) - ) - return tuple(definitions) - - async def startup(self) -> None: - if self._started: - return - - started_handlers: list[BaseHandler] = [] - started_middlewares: list[ToolMiddleware] = [] - try: - for handler in self.handlers: - await handler.startup() - started_handlers.append(handler) - self._tool_routes = self._build_tool_routes() - self._active_middlewares = ( - [middleware for middleware in self.middlewares if middleware.enabled] - if self._tool_routes - else [] - ) - tool_definitions = self.tool_definitions - for middleware in self._active_middlewares: - middleware.configure_tools(tool_definitions) - for middleware in self._active_middlewares: - await middleware.startup() - started_middlewares.append(middleware) - self._tool_chain = compose_tool_middlewares( - self._active_middlewares, - self._dispatch_handler, - ) - except Exception: - for middleware in reversed(started_middlewares): - try: - await middleware.shutdown() - except Exception as shutdown_error: - self.logger.warning( - "Middleware %s rollback shutdown failed: %s", - type(middleware).__name__, - shutdown_error, - ) - for handler in reversed(started_handlers): - try: - await handler.shutdown() - except Exception as shutdown_error: - self.logger.warning( - "Handler %s rollback shutdown failed: %s", - type(handler).__name__, - shutdown_error, - ) - self._tool_routes.clear() - self._active_middlewares.clear() - self._tool_chain = None - raise - - self._started = True - - async def shutdown(self) -> None: - if not self._started: - return - - errors: list[Exception] = [] - for middleware in reversed(self._active_middlewares): - try: - await middleware.shutdown() - except Exception as exc: - errors.append(exc) - for handler in reversed(self.handlers): - try: - await handler.shutdown() - except Exception as exc: - errors.append(exc) - - self._tool_routes.clear() - self._active_middlewares.clear() - self._tool_chain = None - self._started = False - if errors: - raise RuntimeError( - f"failed to shut down {len(errors)} handler(s)" - ) from errors[0] - - async def on_task_start(self) -> None: - """Reset per-task tool state and notify handlers/middleware.""" - - self._last_tool_signature = None - self._repeated_tool_calls = 0 - for handler in self.handlers: - await handler.on_task_start() - for middleware in self._active_middlewares: - await middleware.on_task_start() - - async def dispatch( - self, - tool_name: str, - arguments: Mapping[str, Any], - *, - tool_call_id: str | None = None, - cancellation: CancellationToken | None = None, - progress: ToolProgressReporter | None = None, - ) -> StepOutcome: - if not self._started: - await self.startup() - - handler = self._get_handler(tool_name) - if self._tool_chain is None: - raise RuntimeError("tool middleware chain is not initialized") - token = cancellation or CancellationSource().token - context = ToolCallContext( - state=self.state, - tool_name=tool_name, - arguments=dict(arguments), - tool_call_id=tool_call_id, - cancellation=token, - progress=progress, - tool_definition=self._get_tool_definition(handler, tool_name), - ) - return await token.run(self._tool_chain(context)) - - async def execute_tool_call( - self, - tool_call: ModelToolCall, - *, - cancellation: CancellationToken, - prepared: bool = False, - defer_completion: bool = False, - ) -> ToolCallResult: - tool_name = tool_call.name - raw_arguments = tool_call.arguments - if not prepared: - await self.event_emitter.emit(ToolStarted(self.state.turn, tool_call)) - progress = _ScopedToolProgressReporter( - event_emitter=self.event_emitter, - turn=self.state.turn, - tool_call=tool_call, - cancellation=cancellation, - ) - error: str | None = None - immediate_result: ToolCallResult | None = None - repeated_error: RepeatedToolCallError | None = None - outcome: StepOutcome | None = None - try: - try: - if not prepared: - self._check_repeated_tool_call(tool_name, raw_arguments) - arguments = json.loads(raw_arguments) - if not isinstance(arguments, dict): - raise TypeError("tool arguments must be a JSON object") - outcome = await self.dispatch( - tool_name, - arguments, - tool_call_id=tool_call.id, - cancellation=cancellation, - progress=progress, - ) - except AgentCancelledError as exc: - immediate_result = self._cancelled_tool_result( - tool_call, - str(exc), - ) - except RepeatedToolCallError as exc: - immediate_result = ToolCallResult((), error=str(exc)) - repeated_error = exc - except Exception as exc: - error = str(exc) - outcome = StepOutcome( - { - "status": "error", - "tool": tool_name, - "error": error, - } - ) - finally: - await progress.close() - - if immediate_result is not None: - if not defer_completion: - await self.complete_tool_call(tool_call, immediate_result) - if repeated_error is not None: - raise repeated_error - return immediate_result - - assert outcome is not None - serialized = serialize_tool_result(outcome.data) - message = { - "role": "tool", - "tool_call_id": tool_call.id, - "content": serialized, - } - if outcome.control is not ToolControl.CONTINUE: - result = ToolCallResult( - (message,), - control=outcome.control, - output=serialized, - error=error, - ) - else: - result = ToolCallResult((message,), error=error) - if not defer_completion: - await self.complete_tool_call(tool_call, result) - return result - - async def prepare_parallel_tool_calls( - self, - tool_calls: Sequence[ModelToolCall], - ) -> bool: - """Reserve repetition state and emit ordered starts for one safe batch.""" - - if not self._reserve_repeated_tool_calls(tool_calls): - return False - for tool_call in tool_calls: - await self.event_emitter.emit(ToolStarted(self.state.turn, tool_call)) - return True - - def can_parallelize_tool_calls( - self, - tool_calls: Sequence[ModelToolCall], - ) -> bool: - """Return whether the full source group can reserve repetition state.""" - - return self._repetition_state_after(tool_calls) is not None - - async def complete_tool_call( - self, - tool_call: ModelToolCall, - result: ToolCallResult, - ) -> None: - """Emit one completion at the scheduler's deterministic commit point.""" - - await self.event_emitter.emit(ToolCompleted(self.state.turn, tool_call, result)) - - def tool_effect(self, tool_name: str) -> ToolEffect: - """Return one route's declared effect, defaulting unknown calls to unsafe.""" - - handler = self._tool_routes.get(tool_name) - if handler is None: - return ToolEffect.SIDE_EFFECTING - return handler.tool_effect(tool_name) - - async def cancel_tool_call( - self, - tool_call: ModelToolCall, - *, - reason: str, - ) -> ToolCallResult: - """Settle an unstarted call skipped after run cancellation.""" - - await self.event_emitter.emit(ToolStarted(self.state.turn, tool_call)) - result = self._cancelled_tool_result(tool_call, reason) - await self.event_emitter.emit(ToolCompleted(self.state.turn, tool_call, result)) - return result - - def _build_tool_routes(self) -> dict[str, BaseHandler]: - routes: dict[str, BaseHandler] = {} - for handler in self.handlers: - for tool_name in handler.tool_names: - if tool_name in routes: - first = type(routes[tool_name]).__name__ - second = type(handler).__name__ - raise ValueError( - f"duplicate tool {tool_name!r} in {first} and {second}" - ) - routes[tool_name] = handler - return routes - - def _get_handler(self, tool_name: str) -> BaseHandler: - try: - return self._tool_routes[tool_name] - except KeyError as exc: - available = ", ".join(sorted(self._tool_routes)) or "none" - raise KeyError( - f"unknown tool {tool_name!r}; available tools: {available}" - ) from exc - - @staticmethod - def _get_tool_definition( - handler: BaseHandler, - tool_name: str, - ) -> ToolDefinition: - for definition in handler.tool_definitions: - if definition.name == tool_name: - effect = handler.tool_effect(tool_name) - return ( - definition - if definition.effect is effect - else definition.with_effect(effect) - ) - raise KeyError(f"handler has no definition for tool {tool_name!r}") - - async def _dispatch_handler(self, context: ToolCallContext) -> StepOutcome: - handler = self._get_handler(context.tool_name) - return await handler.execute(context) - - def _cancelled_tool_result( - self, - tool_call: ModelToolCall, - reason: str, - ) -> ToolCallResult: - message = { - "role": "tool", - "tool_call_id": tool_call.id, - "content": serialize_tool_result( - { - "status": "cancelled", - "tool": tool_call.name, - "error": reason, - } - ), - } - return ToolCallResult( - (message,), - error=reason, - cancelled=True, - ) - - def _check_repeated_tool_call( - self, - tool_name: str, - raw_arguments: str, - ) -> None: - signature = self._tool_signature(tool_name, raw_arguments) - if signature == self._last_tool_signature: - self._repeated_tool_calls += 1 - else: - self._last_tool_signature = signature - self._repeated_tool_calls = 1 - - if self._repeated_tool_calls >= self.max_repeated_tool_calls: - raise RepeatedToolCallError( - f"tool {tool_name!r} was called with the same arguments " - f"{self.max_repeated_tool_calls} consecutive times" - ) - - def _reserve_repeated_tool_calls( - self, - tool_calls: Sequence[ModelToolCall], - ) -> bool: - """Atomically reserve source-ordered guard state or request fallback.""" - - next_state = self._repetition_state_after(tool_calls) - if next_state is None: - return False - self._last_tool_signature, self._repeated_tool_calls = next_state - return True - - def _repetition_state_after( - self, - tool_calls: Sequence[ModelToolCall], - ) -> tuple[tuple[str, str] | None, int] | None: - last_signature = self._last_tool_signature - repeated_calls = self._repeated_tool_calls - for tool_call in tool_calls: - signature = self._tool_signature( - tool_call.name, - tool_call.arguments, - ) - if signature == last_signature: - repeated_calls += 1 - else: - last_signature = signature - repeated_calls = 1 - if repeated_calls >= self.max_repeated_tool_calls: - return None - return last_signature, repeated_calls - - @staticmethod - def _tool_signature( - tool_name: str, - raw_arguments: str, - ) -> tuple[str, str]: - try: - arguments = json.loads(raw_arguments) - normalized_arguments = json.dumps( - arguments, - ensure_ascii=False, - sort_keys=True, - separators=(",", ":"), - ) - except (TypeError, ValueError): - normalized_arguments = raw_arguments - return tool_name, normalized_arguments - - -def serialize_tool_result(data: Any) -> str: - if isinstance(data, str): - return data - return json.dumps(data, ensure_ascii=False, default=str) diff --git a/src/ejagent/agent/types.py b/src/ejagent/agent/types.py deleted file mode 100644 index 41b6f6f..0000000 --- a/src/ejagent/agent/types.py +++ /dev/null @@ -1,64 +0,0 @@ -from dataclasses import dataclass -from enum import StrEnum -from typing import Any, Protocol - -AgentMessage = dict[str, Any] -INTERNAL_METADATA_PREFIX = "_ejagent_" -LEGACY_INTERNAL_METADATA_PREFIXES = ("_simagentplg_",) -INTERNAL_METADATA_PREFIXES = ( - INTERNAL_METADATA_PREFIX, - *LEGACY_INTERNAL_METADATA_PREFIXES, -) - - -def is_internal_metadata_key(key: str) -> bool: - """Return whether a message key belongs to current or legacy Core metadata.""" - - return key.startswith(INTERNAL_METADATA_PREFIXES) - - -class ToolControl(StrEnum): - """Control signal returned by a tool independently of its payload.""" - - CONTINUE = "continue" - COMPLETE = "complete" - REJECT = "reject" - CANCEL = "cancel" - - -@dataclass(frozen=True, slots=True) -class ToolProgressUpdate: - """One provisional, human-readable update from an executing tool.""" - - message: str - data: Any = None - - def __post_init__(self) -> None: - if not self.message: - raise ValueError("tool progress message must not be empty") - - -class ToolProgressReporter(Protocol): - """Scoped interface used by one tool call to publish progress.""" - - async def report(self, update: ToolProgressUpdate) -> None: - """Publish an update while the owning tool call is active.""" - - -@dataclass(slots=True) -class StepOutcome: - """Normalized result returned by every tool handler.""" - - data: Any - control: ToolControl = ToolControl.CONTINUE - - -@dataclass(frozen=True, slots=True) -class ToolCallResult: - """Normalized result of one model-requested tool execution.""" - - messages: tuple[AgentMessage, ...] - control: ToolControl = ToolControl.CONTINUE - output: str | None = None - error: str | None = None - cancelled: bool = False diff --git a/src/ejagent/context/__init__.py b/src/ejagent/context/__init__.py new file mode 100644 index 0000000..e42f210 --- /dev/null +++ b/src/ejagent/context/__init__.py @@ -0,0 +1,13 @@ +"""Disposable ContextView projection implementations.""" + +from ejagent.context.pipeline import ( + DerivedCompactionPipeline, + IdentityContextPipeline, +) +from ejagent.context.skills import SkillsContextPipeline + +__all__ = [ + "DerivedCompactionPipeline", + "IdentityContextPipeline", + "SkillsContextPipeline", +] diff --git a/src/ejagent/context/pipeline.py b/src/ejagent/context/pipeline.py new file mode 100644 index 0000000..03b2efb --- /dev/null +++ b/src/ejagent/context/pipeline.py @@ -0,0 +1,137 @@ +from __future__ import annotations + +from ejagent.contracts.context import ( + ContextBuildError, + ContextCompactionOutput, + ContextCompactionRequest, + ContextCompactor, + ContextCompactorError, + ContextPipeline, + ContextProtocolError, + ContextRequest, + ContextView, +) +from ejagent.contracts.control import CancellationToken, RunCancelledError +from ejagent.contracts.lifecycle import ManagedResource +from ejagent.contracts.messages import ContextSummary, SystemMessage +from ejagent.contracts.runs import FailureCode + + +class IdentityContextPipeline(ContextPipeline): + """Project committed and pending messages without transformation.""" + + async def build( + self, + request: ContextRequest, + *, + cancellation: CancellationToken, + ) -> ContextView: + cancellation.raise_if_cancelled() + return ContextView( + run_id=request.run_id, + source_revision=request.source_revision, + turn=request.turn, + messages=request.messages, + metadata={**request.metadata, "projection": "identity"}, + ) + + +class DerivedCompactionPipeline(ContextPipeline): + """Replace committed non-system history with a disposable summary view.""" + + def __init__( + self, + compactor: ContextCompactor, + *, + minimum_messages: int = 20, + ) -> None: + if isinstance(minimum_messages, bool) or not isinstance(minimum_messages, int): + raise TypeError("minimum_messages must be an integer") + if minimum_messages <= 0: + raise ValueError("minimum_messages must be greater than zero") + self._compactor = compactor + self._minimum_messages = minimum_messages + + async def start(self) -> None: + if isinstance(self._compactor, ManagedResource): + await self._compactor.start() + + async def shutdown(self) -> None: + if isinstance(self._compactor, ManagedResource): + await self._compactor.shutdown() + + async def build( + self, + request: ContextRequest, + *, + cancellation: CancellationToken, + ) -> ContextView: + cancellation.raise_if_cancelled() + system_end = 0 + while system_end < len(request.committed_messages) and isinstance( + request.committed_messages[system_end], SystemMessage + ): + system_end += 1 + stable_instructions = request.committed_messages[:system_end] + history = request.committed_messages[system_end:] + if len(history) < self._minimum_messages: + return ContextView( + run_id=request.run_id, + source_revision=request.source_revision, + turn=request.turn, + messages=request.messages, + metadata={**request.metadata, "projection": "identity"}, + ) + + revision_start = 1 if request.source_revision > 0 else 0 + compaction_request = ContextCompactionRequest( + messages=history, + source_revision_start=revision_start, + source_revision_end=request.source_revision, + ) + try: + output = await cancellation.run( + self._compactor.compact( + compaction_request, + cancellation=cancellation, + ) + ) + except RunCancelledError: + raise + except ContextCompactorError as exc: + raise ContextBuildError( + FailureCode.COMPACTION_FAILED, + str(exc), + retryable=exc.retryable, + ) from exc + except Exception as exc: + raise ContextProtocolError( + f"ContextCompactor raised an undeclared {type(exc).__name__}" + ) from exc + if not isinstance(output, ContextCompactionOutput): + raise ContextProtocolError( + "ContextCompactor.compact() must return ContextCompactionOutput" + ) + + summary = ContextSummary( + source_revision_start=revision_start, + source_revision_end=request.source_revision, + content=output.content.strip(), + compactor_id=output.compactor_id.strip(), + ) + return ContextView( + run_id=request.run_id, + source_revision=request.source_revision, + turn=request.turn, + messages=( + *stable_instructions, + summary, + *request.pending_messages, + *request.transient_instructions, + ), + metadata={ + **request.metadata, + "projection": "derived_compaction", + "summarized_message_count": len(history), + }, + ) diff --git a/src/ejagent/context/skills.py b/src/ejagent/context/skills.py new file mode 100644 index 0000000..0ef1151 --- /dev/null +++ b/src/ejagent/context/skills.py @@ -0,0 +1,112 @@ +from __future__ import annotations + +from pathlib import Path + +from ejagent.context.pipeline import IdentityContextPipeline +from ejagent.contracts.context import ( + ContextPipeline, + ContextProtocolError, + ContextRequest, + ContextView, +) +from ejagent.contracts.control import CancellationToken +from ejagent.contracts.lifecycle import ManagedResource +from ejagent.contracts.messages import TransientInstruction, UserMessage +from ejagent.skills import SkillCatalog + + +class SkillsContextPipeline(ContextPipeline): + """Decorate ContextViews with a local skill index and explicit instructions.""" + + def __init__( + self, + skills_root: str | Path, + *, + base: ContextPipeline | None = None, + ) -> None: + self._catalog = SkillCatalog(skills_root) + self._base = base or IdentityContextPipeline() + self._started = False + + @property + def catalog(self) -> SkillCatalog: + return self._catalog + + async def start(self) -> None: + if self._started: + return + await self._catalog.discover() + if isinstance(self._base, ManagedResource): + await self._base.start() + self._started = True + + async def shutdown(self) -> None: + if not self._started: + return + try: + if isinstance(self._base, ManagedResource): + await self._base.shutdown() + finally: + self._started = False + + async def build( + self, + request: ContextRequest, + *, + cancellation: CancellationToken, + ) -> ContextView: + if not self._started: + raise ContextProtocolError("SkillsContextPipeline is not started") + cancellation.raise_if_cancelled() + instructions: list[TransientInstruction] = [] + index = self._catalog.build_index_content() + if index is not None: + instructions.append(TransientInstruction(index, "skills:index")) + task = self._latest_user_task(request) + selected = self._catalog.select_explicit_skill_from_text(task) + if selected is not None: + instructions.append( + TransientInstruction( + self._catalog.build_skill_context_content(selected), + f"skills:{selected}", + ) + ) + augmented = ContextRequest( + run_id=request.run_id, + source_revision=request.source_revision, + turn=request.turn, + committed_messages=request.committed_messages, + pending_messages=request.pending_messages, + transient_instructions=( + *instructions, + *request.transient_instructions, + ), + metadata=request.metadata, + ) + view = await cancellation.run( + self._base.build(augmented, cancellation=cancellation) + ) + if not isinstance(view, ContextView): + raise ContextProtocolError( + "wrapped ContextPipeline.build() must return ContextView" + ) + return ContextView( + run_id=view.run_id, + source_revision=view.source_revision, + turn=view.turn, + messages=view.messages, + metadata={ + **view.metadata, + "available_skills": tuple(skill.name for skill in self._catalog.skills), + "selected_skill": selected, + }, + ) + + @staticmethod + def _latest_user_task(request: ContextRequest) -> str: + for message in reversed( + (*request.committed_messages, *request.pending_messages) + ): + if isinstance(message, UserMessage): + return message.content + return "" diff --git a/src/ejagent/contracts/__init__.py b/src/ejagent/contracts/__init__.py new file mode 100644 index 0000000..54bb37d --- /dev/null +++ b/src/ejagent/contracts/__init__.py @@ -0,0 +1,178 @@ +"""Stable, implementation-independent EJAgent Core contracts.""" + +from ejagent.contracts.audit import AuditReader, RunAudit +from ejagent.contracts.context import ( + ContextBuildError, + ContextCompactionOutput, + ContextCompactionRequest, + ContextCompactor, + ContextCompactorError, + ContextPipeline, + ContextProtocolError, + ContextRequest, + ContextView, +) +from ejagent.contracts.control import ( + CancellationSource, + CancellationToken, + ControlKind, + ControlProtocolError, + ControlReceipt, + ControlStatus, + RunCancelledError, + RunControlSource, + SteeringInput, +) +from ejagent.contracts.conversation import ConversationSnapshot +from ejagent.contracts.json import ( + JsonObject, + JsonScalar, + JsonValue, + MutableJsonValue, + freeze_json_object, + freeze_json_value, + thaw_json_value, +) +from ejagent.contracts.lifecycle import ManagedResource +from ejagent.contracts.messages import ( + AssistantMessage, + ContextMessage, + ContextSummary, + ConversationMessage, + SystemMessage, + ToolCall, + ToolResultMessage, + TransientInstruction, + UserMessage, + is_context_message, + is_conversation_message, +) +from ejagent.contracts.model import ( + ModelCallError, + ModelPort, + ModelProtocolError, + ModelRequest, + ModelResponseCompleted, + ModelStreamEvent, + ModelTextDelta, + ModelThinkingDelta, + ModelUsage, +) +from ejagent.contracts.observer import RunObserver +from ejagent.contracts.runs import ( + AuditRecord, + FailureCode, + RunDelta, + RunFailure, + RunIntent, + RunLimits, + RunOutcome, + RunPhase, + RunResult, + RunSpec, + RunStatus, + StopReason, +) +from ejagent.contracts.session import ( + SessionCommit, + SessionConflictError, + SessionMigrationError, + SessionSnapshot, + SessionStore, + SessionStoreError, + SessionStoreLockTimeoutError, + SessionStoreSerializationError, +) +from ejagent.contracts.tools import ( + ToolControl, + ToolDefinition, + ToolEffect, + ToolExecutionError, + ToolExecutionResult, + ToolExecutor, + ToolProtocolError, + ToolSemantics, +) +from ejagent.contracts.usage import RunUsage + +__all__ = [ + "AssistantMessage", + "AuditReader", + "AuditRecord", + "CancellationSource", + "CancellationToken", + "ControlKind", + "ControlProtocolError", + "ControlReceipt", + "ControlStatus", + "ContextMessage", + "ContextBuildError", + "ContextCompactionOutput", + "ContextCompactionRequest", + "ContextCompactor", + "ContextCompactorError", + "ContextPipeline", + "ContextProtocolError", + "ContextRequest", + "ContextSummary", + "ContextView", + "ConversationMessage", + "ConversationSnapshot", + "FailureCode", + "JsonObject", + "JsonScalar", + "JsonValue", + "ManagedResource", + "MutableJsonValue", + "ModelCallError", + "ModelPort", + "ModelProtocolError", + "ModelRequest", + "ModelResponseCompleted", + "ModelStreamEvent", + "ModelTextDelta", + "ModelThinkingDelta", + "ModelUsage", + "RunDelta", + "RunAudit", + "RunFailure", + "RunIntent", + "RunLimits", + "RunOutcome", + "RunPhase", + "RunResult", + "RunSpec", + "RunStatus", + "RunUsage", + "RunCancelledError", + "RunControlSource", + "RunObserver", + "SessionCommit", + "SessionConflictError", + "SessionMigrationError", + "SessionSnapshot", + "SessionStore", + "SessionStoreError", + "SessionStoreLockTimeoutError", + "SessionStoreSerializationError", + "StopReason", + "SteeringInput", + "SystemMessage", + "ToolCall", + "ToolControl", + "ToolDefinition", + "ToolEffect", + "ToolExecutionError", + "ToolExecutionResult", + "ToolExecutor", + "ToolProtocolError", + "ToolResultMessage", + "ToolSemantics", + "TransientInstruction", + "UserMessage", + "freeze_json_object", + "freeze_json_value", + "is_context_message", + "is_conversation_message", + "thaw_json_value", +] diff --git a/src/ejagent/contracts/audit.py b/src/ejagent/contracts/audit.py new file mode 100644 index 0000000..63461a5 --- /dev/null +++ b/src/ejagent/contracts/audit.py @@ -0,0 +1,53 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Protocol + +from ejagent.contracts.runs import AuditRecord, RunFailure, RunResult + + +@dataclass(frozen=True, slots=True) +class RunAudit: + """Append-only facts for one attempted Run, separate from Conversation.""" + + result: RunResult + base_revision: int + resulting_revision: int + committed: bool + records: tuple[AuditRecord, ...] = () + failure: RunFailure | None = None + + def __post_init__(self) -> None: + if not isinstance(self.result, RunResult): + raise TypeError("audit result must be a RunResult") + for name, value in ( + ("base_revision", self.base_revision), + ("resulting_revision", self.resulting_revision), + ): + if isinstance(value, bool) or not isinstance(value, int): + raise TypeError(f"audit {name} must be an integer") + if value < 0: + raise ValueError(f"audit {name} must not be negative") + if not isinstance(self.committed, bool): + raise TypeError("audit committed must be a boolean") + expected_revision = self.base_revision + int(self.committed) + if self.resulting_revision != expected_revision: + raise ValueError("audit resulting revision does not match commit status") + object.__setattr__(self, "records", tuple(self.records)) + if not all(isinstance(record, AuditRecord) for record in self.records): + raise TypeError("audit records must contain AuditRecord values") + if any(record.run_id != self.result.run_id for record in self.records): + raise ValueError("audit records must belong to the audited Run") + if self.failure is not None and not isinstance(self.failure, RunFailure): + raise TypeError("audit failure must be a RunFailure or None") + + @property + def run_id(self) -> str: + return self.result.run_id + + +class AuditReader(Protocol): + """Read-only seam for durable Run audit history.""" + + async def load_audit(self, agent_id: str) -> tuple[RunAudit, ...]: + """Return Run facts in durable insertion order.""" diff --git a/src/ejagent/contracts/context.py b/src/ejagent/contracts/context.py new file mode 100644 index 0000000..4eee40b --- /dev/null +++ b/src/ejagent/contracts/context.py @@ -0,0 +1,218 @@ +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import Protocol + +from ejagent.contracts.control import CancellationToken +from ejagent.contracts.json import JsonObject, freeze_json_object +from ejagent.contracts.messages import ( + ContextMessage, + ConversationMessage, + TransientInstruction, + is_context_message, + is_conversation_message, +) +from ejagent.contracts.runs import FailureCode + + +def _non_negative_integer(value: object, field_name: str) -> None: + if isinstance(value, bool) or not isinstance(value, int): + raise TypeError(f"{field_name} must be an integer") + if value < 0: + raise ValueError(f"{field_name} must not be negative") + + +@dataclass(frozen=True, slots=True) +class ContextRequest: + """Immutable inputs for one disposable model ContextView.""" + + run_id: str + source_revision: int + turn: int + committed_messages: tuple[ConversationMessage, ...] + pending_messages: tuple[ConversationMessage, ...] = () + transient_instructions: tuple[TransientInstruction, ...] = () + metadata: JsonObject = field(default_factory=dict) + + def __post_init__(self) -> None: + if not isinstance(self.run_id, str) or not self.run_id.strip(): + raise ValueError("context run_id must not be empty") + _non_negative_integer(self.source_revision, "context source_revision") + _non_negative_integer(self.turn, "context turn") + if self.turn == 0: + raise ValueError("context turn must be greater than zero") + object.__setattr__(self, "committed_messages", tuple(self.committed_messages)) + object.__setattr__(self, "pending_messages", tuple(self.pending_messages)) + object.__setattr__( + self, + "transient_instructions", + tuple(self.transient_instructions), + ) + if not all( + is_conversation_message(message) for message in self.committed_messages + ): + raise TypeError( + "committed_messages must contain ConversationMessage values" + ) + if not all( + is_conversation_message(message) for message in self.pending_messages + ): + raise TypeError("pending_messages must contain ConversationMessage values") + if not all( + isinstance(instruction, TransientInstruction) + for instruction in self.transient_instructions + ): + raise TypeError( + "transient_instructions must contain TransientInstruction values" + ) + if not isinstance(self.metadata, Mapping): + raise TypeError("context metadata must be a JSON object") + object.__setattr__( + self, + "metadata", + freeze_json_object(self.metadata, label="context metadata"), + ) + + @property + def messages(self) -> tuple[ContextMessage, ...]: + return ( + *self.committed_messages, + *self.pending_messages, + *self.transient_instructions, + ) + + +@dataclass(frozen=True, slots=True) +class ContextView: + """Disposable Provider-neutral projection for one model request.""" + + run_id: str + source_revision: int + turn: int + messages: tuple[ContextMessage, ...] + metadata: JsonObject = field(default_factory=dict) + + def __post_init__(self) -> None: + if not isinstance(self.run_id, str) or not self.run_id.strip(): + raise ValueError("ContextView run_id must not be empty") + _non_negative_integer(self.source_revision, "ContextView source_revision") + _non_negative_integer(self.turn, "ContextView turn") + if self.turn == 0: + raise ValueError("ContextView turn must be greater than zero") + object.__setattr__(self, "messages", tuple(self.messages)) + if not all(is_context_message(message) for message in self.messages): + raise TypeError("ContextView messages must contain ContextMessage values") + if not isinstance(self.metadata, Mapping): + raise TypeError("ContextView metadata must be a JSON object") + object.__setattr__( + self, + "metadata", + freeze_json_object(self.metadata, label="ContextView metadata"), + ) + + +@dataclass(frozen=True, slots=True) +class ContextCompactionRequest: + """Committed history selected for one derived summary projection.""" + + messages: tuple[ConversationMessage, ...] + source_revision_start: int + source_revision_end: int + + def __post_init__(self) -> None: + object.__setattr__(self, "messages", tuple(self.messages)) + if not self.messages: + raise ValueError("compaction messages must not be empty") + if not all(is_conversation_message(message) for message in self.messages): + raise TypeError( + "compaction messages must contain ConversationMessage values" + ) + _non_negative_integer( + self.source_revision_start, + "compaction source_revision_start", + ) + _non_negative_integer( + self.source_revision_end, + "compaction source_revision_end", + ) + if self.source_revision_end < self.source_revision_start: + raise ValueError("compaction revision range is reversed") + + +@dataclass(frozen=True, slots=True) +class ContextCompactionOutput: + """Summary text returned by a ContextCompactor adapter.""" + + content: str + compactor_id: str + + def __post_init__(self) -> None: + if not isinstance(self.content, str) or not self.content.strip(): + raise ValueError("compaction output content must not be empty") + if not isinstance(self.compactor_id, str) or not self.compactor_id.strip(): + raise ValueError("compaction output compactor_id must not be empty") + + +class ContextCompactorError(RuntimeError): + """Expected operational failure from a ContextCompactor.""" + + def __init__(self, message: str, *, retryable: bool = False) -> None: + if not message.strip(): + raise ValueError("compactor error message must not be empty") + if not isinstance(retryable, bool): + raise TypeError("compactor error retryable must be a boolean") + self.retryable = retryable + super().__init__(message) + + +class ContextCompactor(Protocol): + """Generate summary content without owning Conversation state.""" + + async def compact( + self, + request: ContextCompactionRequest, + *, + cancellation: CancellationToken, + ) -> ContextCompactionOutput: + """Summarize the selected immutable committed history.""" + + +class ContextBuildError(RuntimeError): + """Expected operational failure while deriving a ContextView.""" + + def __init__( + self, + code: FailureCode, + message: str, + *, + retryable: bool = False, + ) -> None: + if code not in ( + FailureCode.CONTEXT_OVERFLOW, + FailureCode.COMPACTION_FAILED, + ): + raise ValueError("ContextBuildError requires a context failure code") + if not message.strip(): + raise ValueError("context error message must not be empty") + if not isinstance(retryable, bool): + raise TypeError("context error retryable must be a boolean") + self.code = code + self.retryable = retryable + super().__init__(message) + + +class ContextProtocolError(RuntimeError): + """A ContextPipeline violated the stable projection protocol.""" + + +class ContextPipeline(Protocol): + """Build one ContextView without mutating Conversation or Run state.""" + + async def build( + self, + request: ContextRequest, + *, + cancellation: CancellationToken, + ) -> ContextView: + """Return one disposable view for the next model request.""" diff --git a/src/ejagent/contracts/control.py b/src/ejagent/contracts/control.py new file mode 100644 index 0000000..963b679 --- /dev/null +++ b/src/ejagent/contracts/control.py @@ -0,0 +1,163 @@ +from __future__ import annotations + +import asyncio +import inspect +from collections.abc import Awaitable +from dataclasses import dataclass +from enum import StrEnum +from typing import Protocol, TypeVar + +T = TypeVar("T") + + +class ControlKind(StrEnum): + """Kind of input admitted by an AgentHarness control queue.""" + + STEERING = "steering" + FOLLOW_UP = "follow_up" + + +class ControlStatus(StrEnum): + """Immediate admission decision for one control input.""" + + ACCEPTED = "accepted" + NOT_RUNNING = "not_running" + TOO_LATE = "too_late" + QUEUE_FULL = "queue_full" + CLOSED = "closed" + + +@dataclass(frozen=True, slots=True) +class ControlReceipt: + """Stable immediate result of one control admission attempt.""" + + input_id: str + kind: ControlKind + status: ControlStatus + + def __post_init__(self) -> None: + if not isinstance(self.input_id, str) or not self.input_id.strip(): + raise ValueError("control input_id must not be empty") + if not isinstance(self.kind, ControlKind): + raise TypeError("control kind must be a ControlKind") + if not isinstance(self.status, ControlStatus): + raise TypeError("control status must be a ControlStatus") + + @property + def accepted(self) -> bool: + return self.status is ControlStatus.ACCEPTED + + +@dataclass(frozen=True, slots=True) +class SteeringInput: + """One admitted transient instruction awaiting a model-call safe point.""" + + input_id: str + content: str + + def __post_init__(self) -> None: + if not isinstance(self.input_id, str) or not self.input_id.strip(): + raise ValueError("steering input_id must not be empty") + if not isinstance(self.content, str) or not self.content.strip(): + raise ValueError("steering content must not be empty") + + +class RunControlSource(Protocol): + """Run-local control source consumed only at Kernel safe points.""" + + def drain_steering(self) -> tuple[SteeringInput, ...]: + """Return admitted steering in FIFO order exactly once.""" + + +class ControlProtocolError(RuntimeError): + """A RunControlSource violated the stable Kernel control protocol.""" + + +class RunCancelledError(RuntimeError): + """Raised cooperatively when an active Run is cancelled.""" + + +class CancellationToken: + """Read-only cancellation signal shared across one Run.""" + + def __init__(self) -> None: + self._event = asyncio.Event() + self._reason: str | None = None + + @property + def cancelled(self) -> bool: + """Return whether cancellation has been requested.""" + + return self._event.is_set() + + @property + def reason(self) -> str | None: + """Return the cancellation reason when one was supplied.""" + + return self._reason + + async def wait(self) -> None: + """Wait until cancellation is requested.""" + + await self._event.wait() + + def raise_if_cancelled(self) -> None: + """Raise the Run-level cancellation exception when cancelled.""" + + if self.cancelled: + raise RunCancelledError(self.reason or "Run was cancelled") + + async def run(self, awaitable: Awaitable[T]) -> T: + """Await work while interrupting it when this token is cancelled.""" + + if self.cancelled: + if inspect.iscoroutine(awaitable): + awaitable.close() + self.raise_if_cancelled() + + work = asyncio.ensure_future(awaitable) + cancellation = asyncio.create_task(self.wait()) + try: + done, _ = await asyncio.wait( + (work, cancellation), + return_when=asyncio.FIRST_COMPLETED, + ) + if work in done: + return await work + + work.cancel() + await asyncio.gather(work, return_exceptions=True) + self.raise_if_cancelled() + raise RuntimeError("cancellation wait completed without a signal") + finally: + cancellation.cancel() + if not work.done(): + work.cancel() + await asyncio.gather( + work, + cancellation, + return_exceptions=True, + ) + + def _cancel(self, reason: str | None) -> bool: + if self.cancelled: + return False + self._reason = reason or "Run was cancelled" + self._event.set() + return True + + +class CancellationSource: + """Mutable owner of one public read-only cancellation token.""" + + def __init__(self) -> None: + self._token = CancellationToken() + + @property + def token(self) -> CancellationToken: + return self._token + + def cancel(self, reason: str | None = None) -> bool: + """Request cancellation once and report whether state changed.""" + + return self._token._cancel(reason) diff --git a/src/ejagent/contracts/conversation.py b/src/ejagent/contracts/conversation.py new file mode 100644 index 0000000..8b68b6f --- /dev/null +++ b/src/ejagent/contracts/conversation.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +from dataclasses import dataclass + +from ejagent.contracts.messages import ConversationMessage, is_conversation_message + + +@dataclass(frozen=True, slots=True) +class ConversationSnapshot: + """Immutable committed Conversation at one logical revision.""" + + revision: int = 0 + messages: tuple[ConversationMessage, ...] = () + + def __post_init__(self) -> None: + if isinstance(self.revision, bool) or not isinstance(self.revision, int): + raise TypeError("conversation revision must be an integer") + if self.revision < 0: + raise ValueError("conversation revision must not be negative") + object.__setattr__(self, "messages", tuple(self.messages)) + if not all(is_conversation_message(message) for message in self.messages): + raise TypeError( + "conversation messages must contain ConversationMessage values" + ) diff --git a/src/ejagent/contracts/json.py b/src/ejagent/contracts/json.py new file mode 100644 index 0000000..bc1cac3 --- /dev/null +++ b/src/ejagent/contracts/json.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +import math +from collections.abc import Mapping, Sequence +from types import MappingProxyType +from typing import TypeAlias + +JsonScalar: TypeAlias = str | int | float | bool | None +JsonValue: TypeAlias = JsonScalar | tuple["JsonValue", ...] | Mapping[str, "JsonValue"] +JsonObject: TypeAlias = Mapping[str, JsonValue] +MutableJsonValue: TypeAlias = ( + JsonScalar | list["MutableJsonValue"] | dict[str, "MutableJsonValue"] +) + + +def freeze_json_value(value: object, *, label: str = "value") -> JsonValue: + """Validate and recursively freeze one JSON-compatible value.""" + + if value is None or isinstance(value, (str, bool, int)): + return value + if isinstance(value, float): + if not math.isfinite(value): + raise ValueError(f"{label} must not contain non-finite numbers") + return value + if isinstance(value, Mapping): + frozen: dict[str, JsonValue] = {} + for key, item in value.items(): + if not isinstance(key, str): + raise TypeError(f"{label} object keys must be strings") + frozen[key] = freeze_json_value(item, label=f"{label}.{key}") + return MappingProxyType(frozen) + if isinstance(value, Sequence) and not isinstance( + value, + (str, bytes, bytearray), + ): + return tuple( + freeze_json_value(item, label=f"{label}[{index}]") + for index, item in enumerate(value) + ) + raise TypeError(f"{label} must be JSON-compatible") + + +def freeze_json_object( + value: Mapping[str, object], + *, + label: str = "value", +) -> JsonObject: + """Validate and recursively freeze one JSON object.""" + + frozen = freeze_json_value(value, label=label) + if not isinstance(frozen, Mapping): + raise TypeError(f"{label} must be a JSON object") + return frozen + + +def thaw_json_value(value: JsonValue) -> MutableJsonValue: + """Return a detached mutable representation of one frozen JSON value.""" + + if isinstance(value, Mapping): + return {key: thaw_json_value(item) for key, item in value.items()} + if isinstance(value, tuple): + return [thaw_json_value(item) for item in value] + return value diff --git a/src/ejagent/contracts/lifecycle.py b/src/ejagent/contracts/lifecycle.py new file mode 100644 index 0000000..d112eb1 --- /dev/null +++ b/src/ejagent/contracts/lifecycle.py @@ -0,0 +1,14 @@ +from __future__ import annotations + +from typing import Protocol, runtime_checkable + + +@runtime_checkable +class ManagedResource(Protocol): + """Resource whose lifetime is owned by one AgentHarness.""" + + async def start(self) -> None: + """Acquire resources before the Harness accepts Runs.""" + + async def shutdown(self) -> None: + """Release resources after the Harness stops accepting Runs.""" diff --git a/src/ejagent/contracts/messages.py b/src/ejagent/contracts/messages.py new file mode 100644 index 0000000..74c86a5 --- /dev/null +++ b/src/ejagent/contracts/messages.py @@ -0,0 +1,164 @@ +from __future__ import annotations + +import re +from collections.abc import Mapping +from dataclasses import dataclass, field +from typing import TypeAlias, TypeGuard + +from ejagent.contracts.json import ( + JsonObject, + JsonValue, + freeze_json_object, + freeze_json_value, +) + +_TOOL_NAME_PATTERN = re.compile(r"^[A-Za-z0-9_-]{1,64}$") + + +def _required_text(value: object, field_name: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"{field_name} must not be empty") + return value + + +@dataclass(frozen=True, slots=True) +class ToolCall: + """Provider-neutral Tool invocation requested by an assistant.""" + + id: str + name: str + arguments: JsonObject = field(default_factory=dict) + + def __post_init__(self) -> None: + _required_text(self.id, "tool call id") + if not isinstance(self.name, str) or not _TOOL_NAME_PATTERN.fullmatch( + self.name + ): + raise ValueError( + "tool call name must contain 1-64 letters, digits, " + "underscores, or dashes" + ) + if not isinstance(self.arguments, Mapping): + raise TypeError("tool call arguments must be a JSON object") + object.__setattr__( + self, + "arguments", + freeze_json_object(self.arguments, label="tool call arguments"), + ) + + +@dataclass(frozen=True, slots=True) +class SystemMessage: + """Stable instruction belonging to committed conversation state.""" + + content: str + + def __post_init__(self) -> None: + _required_text(self.content, "system message content") + + +@dataclass(frozen=True, slots=True) +class UserMessage: + """One user input committed to conversation state.""" + + content: str + + def __post_init__(self) -> None: + _required_text(self.content, "user message content") + + +@dataclass(frozen=True, slots=True) +class AssistantMessage: + """Provider-neutral assistant output committed as one message.""" + + content: str | None = None + tool_calls: tuple[ToolCall, ...] = () + + def __post_init__(self) -> None: + if self.content is not None and not isinstance(self.content, str): + raise TypeError("assistant message content must be a string or None") + object.__setattr__(self, "tool_calls", tuple(self.tool_calls)) + if not self.content and not self.tool_calls: + raise ValueError("assistant message must contain text or tool calls") + if not all(isinstance(call, ToolCall) for call in self.tool_calls): + raise TypeError("assistant tool_calls must contain ToolCall values") + + +@dataclass(frozen=True, slots=True) +class ToolResultMessage: + """Provider-neutral result paired with one committed Tool call.""" + + tool_call_id: str + tool_name: str + result: JsonValue + is_error: bool = False + + def __post_init__(self) -> None: + _required_text(self.tool_call_id, "tool result call id") + if not isinstance(self.tool_name, str) or not _TOOL_NAME_PATTERN.fullmatch( + self.tool_name + ): + raise ValueError( + "tool result name must contain 1-64 letters, digits, " + "underscores, or dashes" + ) + if not isinstance(self.is_error, bool): + raise TypeError("tool result is_error must be a boolean") + object.__setattr__( + self, + "result", + freeze_json_value(self.result, label="tool result"), + ) + + +@dataclass(frozen=True, slots=True) +class ContextSummary: + """Derived summary usable in a ContextView but not conversation truth.""" + + source_revision_start: int + source_revision_end: int + content: str + compactor_id: str + + def __post_init__(self) -> None: + if self.source_revision_start < 0: + raise ValueError("summary start revision must not be negative") + if self.source_revision_end < self.source_revision_start: + raise ValueError("summary end revision must not precede start revision") + _required_text(self.content, "summary content") + _required_text(self.compactor_id, "summary compactor_id") + + +@dataclass(frozen=True, slots=True) +class TransientInstruction: + """Run-local instruction projected into context but never Conversation.""" + + content: str + source: str + + def __post_init__(self) -> None: + _required_text(self.content, "transient instruction content") + _required_text(self.source, "transient instruction source") + + +ConversationMessage: TypeAlias = ( + SystemMessage | UserMessage | AssistantMessage | ToolResultMessage +) +ContextMessage: TypeAlias = ConversationMessage | ContextSummary | TransientInstruction + + +def is_conversation_message(value: object) -> TypeGuard[ConversationMessage]: + """Return whether a value belongs to the closed Conversation union.""" + + return isinstance( + value, + (SystemMessage, UserMessage, AssistantMessage, ToolResultMessage), + ) + + +def is_context_message(value: object) -> TypeGuard[ContextMessage]: + """Return whether a value belongs to the closed Context union.""" + + return is_conversation_message(value) or isinstance( + value, (ContextSummary, TransientInstruction) + ) diff --git a/src/ejagent/contracts/model.py b/src/ejagent/contracts/model.py new file mode 100644 index 0000000..dbab14f --- /dev/null +++ b/src/ejagent/contracts/model.py @@ -0,0 +1,162 @@ +from __future__ import annotations + +from collections.abc import AsyncIterator +from dataclasses import dataclass +from typing import Protocol, TypeAlias + +from ejagent.contracts.control import CancellationToken +from ejagent.contracts.messages import ( + AssistantMessage, + ContextMessage, + is_context_message, +) +from ejagent.contracts.runs import FailureCode +from ejagent.contracts.tools import ToolDefinition + + +@dataclass(frozen=True, slots=True) +class ModelUsage: + """Provider-neutral token usage for one completed model response.""" + + input_tokens: int + output_tokens: int + total_tokens: int + cache_read_tokens: int | None = None + cache_write_tokens: int | None = None + reasoning_tokens: int | None = None + + def __post_init__(self) -> None: + values = { + "input_tokens": self.input_tokens, + "output_tokens": self.output_tokens, + "total_tokens": self.total_tokens, + "cache_read_tokens": self.cache_read_tokens, + "cache_write_tokens": self.cache_write_tokens, + "reasoning_tokens": self.reasoning_tokens, + } + for name, value in values.items(): + if value is not None and value < 0: + raise ValueError(f"{name} must not be negative") + if self.total_tokens != self.input_tokens + self.output_tokens: + raise ValueError("total_tokens must equal input_tokens + output_tokens") + if ( + self.cache_read_tokens is not None + and self.cache_read_tokens > self.input_tokens + ): + raise ValueError("cache_read_tokens must not exceed input_tokens") + if ( + self.cache_write_tokens is not None + and self.cache_write_tokens > self.input_tokens + ): + raise ValueError("cache_write_tokens must not exceed input_tokens") + if ( + self.reasoning_tokens is not None + and self.reasoning_tokens > self.output_tokens + ): + raise ValueError("reasoning_tokens must not exceed output_tokens") + + def to_dict(self) -> dict[str, int | None]: + """Return a detached JSON-compatible representation.""" + + return { + "input_tokens": self.input_tokens, + "output_tokens": self.output_tokens, + "total_tokens": self.total_tokens, + "cache_read_tokens": self.cache_read_tokens, + "cache_write_tokens": self.cache_write_tokens, + "reasoning_tokens": self.reasoning_tokens, + } + + +@dataclass(frozen=True, slots=True) +class ModelRequest: + """Provider-neutral request produced for one Kernel turn.""" + + messages: tuple[ContextMessage, ...] + tools: tuple[ToolDefinition, ...] = () + + def __post_init__(self) -> None: + object.__setattr__(self, "messages", tuple(self.messages)) + if not all(is_context_message(message) for message in self.messages): + raise TypeError("model messages must contain ContextMessage values") + object.__setattr__(self, "tools", tuple(self.tools)) + if not all(isinstance(tool, ToolDefinition) for tool in self.tools): + raise TypeError("model tools must contain ToolDefinition values") + + +@dataclass(frozen=True, slots=True) +class ModelTextDelta: + """One Provider-neutral piece of assistant text.""" + + delta: str + + def __post_init__(self) -> None: + if not isinstance(self.delta, str) or not self.delta: + raise ValueError("model text delta must not be empty") + + +@dataclass(frozen=True, slots=True) +class ModelThinkingDelta: + """One Provider-neutral piece of provisional model reasoning.""" + + delta: str + + def __post_init__(self) -> None: + if not isinstance(self.delta, str) or not self.delta: + raise ValueError("model thinking delta must not be empty") + + +@dataclass(frozen=True, slots=True) +class ModelResponseCompleted: + """Terminal stream event containing one normalized assistant response.""" + + message: AssistantMessage + usage: ModelUsage | None = None + + def __post_init__(self) -> None: + if not isinstance(self.message, AssistantMessage): + raise TypeError("completed model message must be AssistantMessage") + if self.usage is not None and not isinstance(self.usage, ModelUsage): + raise TypeError("completed model usage must be ModelUsage or None") + + +ModelStreamEvent: TypeAlias = ( + ModelTextDelta | ModelThinkingDelta | ModelResponseCompleted +) + + +class ModelCallError(RuntimeError): + """Expected operational failure reported by a ModelPort.""" + + def __init__( + self, + code: FailureCode, + message: str, + *, + retryable: bool = False, + ) -> None: + if not isinstance(code, FailureCode): + raise TypeError("Model error code must be a FailureCode") + if not message.strip(): + raise ValueError("Model error message must not be empty") + if not isinstance(retryable, bool): + raise TypeError("Model error retryable must be a boolean") + self.code = code + self.retryable = retryable + super().__init__(message) + + +class ModelProtocolError(RuntimeError): + """A ModelPort violated the stable Kernel stream protocol.""" + + +class ModelPort(Protocol): + """Ready Provider boundary consumed by a Runtime Kernel.""" + + def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + """Stream one normalized model response without managing resources.""" diff --git a/src/ejagent/contracts/observer.py b/src/ejagent/contracts/observer.py new file mode 100644 index 0000000..a2e8c1d --- /dev/null +++ b/src/ejagent/contracts/observer.py @@ -0,0 +1,12 @@ +from __future__ import annotations + +from typing import Protocol + +from ejagent.contracts.audit import RunAudit + + +class RunObserver(Protocol): + """Observe a finished Run without participating in its result or commit.""" + + async def observe(self, audit: RunAudit) -> None: + """Consume one immutable RunAudit after the Store decision.""" diff --git a/src/ejagent/contracts/runs.py b/src/ejagent/contracts/runs.py new file mode 100644 index 0000000..1f87e0c --- /dev/null +++ b/src/ejagent/contracts/runs.py @@ -0,0 +1,298 @@ +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass, field +from datetime import datetime +from enum import StrEnum + +from ejagent.contracts.json import JsonObject, freeze_json_object +from ejagent.contracts.messages import ConversationMessage, is_conversation_message +from ejagent.contracts.usage import RunUsage + + +class RunIntent(StrEnum): + """Why a Harness started one Run.""" + + TASK = "task" + CONTINUE = "continue" + + +class RunStatus(StrEnum): + """Terminal status of one Runtime Kernel Run.""" + + COMPLETED = "completed" + FAILED = "failed" + CANCELLED = "cancelled" + REJECTED = "rejected" + + +class StopReason(StrEnum): + """Stable reason why one Run stopped.""" + + TEXT_RESPONSE = "text_response" + TOOL_COMPLETION = "tool_completion" + TOOL_REJECTED = "tool_rejected" + TOOL_CANCELLED = "tool_cancelled" + BEHAVIOR_STOP = "behavior_stop" + EXTERNAL_ABORT = "external_abort" + EMPTY_RESPONSE = "empty_response" + MAX_STEPS = "max_steps" + MAX_NO_TOOL_RESPONSES = "max_no_tool_responses" + REPEATED_TOOL_CALL = "repeated_tool_call" + TOKEN_BUDGET_EXCEEDED = "token_budget_exceeded" + USAGE_UNAVAILABLE = "usage_unavailable" + CONTEXT_OVERFLOW = "context_overflow" + COMPACTION_FAILED = "compaction_failed" + PERSISTENCE_FAILED = "persistence_failed" + RUNTIME_ERROR = "runtime_error" + + +class RunPhase(StrEnum): + """Execution phase associated with one structured Run failure.""" + + PREPARATION = "preparation" + CONTEXT = "context" + MODEL = "model" + TOOL = "tool" + CONTROL = "control" + COMMIT = "commit" + RUNTIME = "runtime" + + +class FailureCode(StrEnum): + """Provider-neutral category for an expected operational failure.""" + + PROVIDER_ERROR = "provider_error" + RATE_LIMIT = "rate_limit" + TIMEOUT = "timeout" + AUTHENTICATION = "authentication" + TOOL_ERROR = "tool_error" + POLICY_REJECTED = "policy_rejected" + CANCELLED = "cancelled" + BUDGET_EXCEEDED = "budget_exceeded" + CONTEXT_OVERFLOW = "context_overflow" + COMPACTION_FAILED = "compaction_failed" + PERSISTENCE_FAILED = "persistence_failed" + RUNTIME_ERROR = "runtime_error" + + +def _positive_integer(value: object, field_name: str) -> None: + if isinstance(value, bool) or not isinstance(value, int): + raise TypeError(f"{field_name} must be an integer") + if value <= 0: + raise ValueError(f"{field_name} must be greater than zero") + + +@dataclass(frozen=True, slots=True) +class RunLimits: + """Immutable limits captured when one Run starts.""" + + max_turns: int = 20 + max_tokens: int | None = None + max_repeated_tool_calls: int = 3 + parallel_tool_calls: bool = False + max_parallel_tool_calls: int | None = None + + def __post_init__(self) -> None: + _positive_integer(self.max_turns, "max_turns") + _positive_integer( + self.max_repeated_tool_calls, + "max_repeated_tool_calls", + ) + if self.max_tokens is not None: + _positive_integer(self.max_tokens, "max_tokens") + if not isinstance(self.parallel_tool_calls, bool): + raise TypeError("parallel_tool_calls must be a boolean") + if self.max_parallel_tool_calls is not None: + _positive_integer( + self.max_parallel_tool_calls, + "max_parallel_tool_calls", + ) + + +@dataclass(frozen=True, slots=True) +class RunSpec: + """Immutable input captured by a Harness for one Kernel Run.""" + + run_id: str + base_revision: int + intent: RunIntent + task: str | None + messages: tuple[ConversationMessage, ...] + limits: RunLimits = field(default_factory=RunLimits) + configuration_revision: str = "default" + metadata: JsonObject = field(default_factory=dict) + + def __post_init__(self) -> None: + if not isinstance(self.run_id, str) or not self.run_id.strip(): + raise ValueError("run_id must not be empty") + if isinstance(self.base_revision, bool) or not isinstance( + self.base_revision, + int, + ): + raise TypeError("base_revision must be an integer") + if self.base_revision < 0: + raise ValueError("base_revision must not be negative") + if not isinstance(self.intent, RunIntent): + raise TypeError("intent must be a RunIntent") + if self.intent is RunIntent.TASK: + if not isinstance(self.task, str) or not self.task.strip(): + raise ValueError("task Run requires a non-empty task") + elif self.task is not None: + raise ValueError("Continue Run must not contain a task") + object.__setattr__(self, "messages", tuple(self.messages)) + if not all(is_conversation_message(message) for message in self.messages): + raise TypeError("messages must contain ConversationMessage values") + if not isinstance(self.limits, RunLimits): + raise TypeError("limits must be RunLimits") + if ( + not isinstance(self.configuration_revision, str) + or not self.configuration_revision.strip() + ): + raise ValueError("configuration_revision must not be empty") + if not isinstance(self.metadata, Mapping): + raise TypeError("metadata must be a JSON object") + object.__setattr__( + self, + "metadata", + freeze_json_object(self.metadata, label="RunSpec metadata"), + ) + + +@dataclass(frozen=True, slots=True) +class RunDelta: + """Conversation messages proposed against one Harness revision.""" + + base_revision: int + messages: tuple[ConversationMessage, ...] = () + + def __post_init__(self) -> None: + if isinstance(self.base_revision, bool) or not isinstance( + self.base_revision, + int, + ): + raise TypeError("base_revision must be an integer") + if self.base_revision < 0: + raise ValueError("base_revision must not be negative") + object.__setattr__(self, "messages", tuple(self.messages)) + if not all(is_conversation_message(message) for message in self.messages): + raise TypeError("messages must contain ConversationMessage values") + + @property + def next_revision(self) -> int: + """Return the revision produced when this Delta is committed.""" + + return self.base_revision + 1 + + +@dataclass(frozen=True, slots=True) +class RunFailure: + """Structured expected failure produced during one Run.""" + + phase: RunPhase + code: FailureCode + message: str + retryable: bool = False + cause: BaseException | None = field(default=None, repr=False, compare=False) + + def __post_init__(self) -> None: + if not isinstance(self.phase, RunPhase): + raise TypeError("failure phase must be a RunPhase") + if not isinstance(self.code, FailureCode): + raise TypeError("failure code must be a FailureCode") + if not isinstance(self.message, str) or not self.message.strip(): + raise ValueError("failure message must not be empty") + if not isinstance(self.retryable, bool): + raise TypeError("failure retryable must be a boolean") + if self.cause is not None and not isinstance(self.cause, BaseException): + raise TypeError("failure cause must be an exception or None") + + +@dataclass(frozen=True, slots=True) +class RunResult: + """Provider-neutral terminal summary of one Kernel Run.""" + + run_id: str + status: RunStatus + stop_reason: StopReason + turns: int + output: str | None = None + usage: RunUsage = field(default_factory=RunUsage) + + def __post_init__(self) -> None: + if not isinstance(self.run_id, str) or not self.run_id.strip(): + raise ValueError("result run_id must not be empty") + if not isinstance(self.status, RunStatus): + raise TypeError("status must be a RunStatus") + if not isinstance(self.stop_reason, StopReason): + raise TypeError("stop_reason must be a StopReason") + if isinstance(self.turns, bool) or not isinstance(self.turns, int): + raise TypeError("turns must be an integer") + if self.turns < 0: + raise ValueError("turns must not be negative") + if self.output is not None and not isinstance(self.output, str): + raise TypeError("output must be a string or None") + if not isinstance(self.usage, RunUsage): + raise TypeError("usage must be RunUsage") + + @property + def succeeded(self) -> bool: + return self.status is RunStatus.COMPLETED + + +@dataclass(frozen=True, slots=True) +class AuditRecord: + """One immutable fact emitted while a Run was executing.""" + + run_id: str + sequence: int + kind: str + occurred_at: datetime + payload: JsonObject = field(default_factory=dict) + + def __post_init__(self) -> None: + if not isinstance(self.run_id, str) or not self.run_id.strip(): + raise ValueError("audit run_id must not be empty") + _positive_integer(self.sequence, "audit sequence") + if not isinstance(self.kind, str) or not self.kind.strip(): + raise ValueError("audit kind must not be empty") + if not isinstance(self.occurred_at, datetime): + raise TypeError("audit occurred_at must be a datetime") + if self.occurred_at.tzinfo is None or self.occurred_at.utcoffset() is None: + raise ValueError("audit occurred_at must be timezone-aware") + if not isinstance(self.payload, Mapping): + raise TypeError("audit payload must be a JSON object") + object.__setattr__( + self, + "payload", + freeze_json_object(self.payload, label="audit payload"), + ) + + +@dataclass(frozen=True, slots=True) +class RunOutcome: + """Complete Kernel output awaiting a Harness commit decision.""" + + result: RunResult + delta: RunDelta + audit_records: tuple[AuditRecord, ...] = () + failure: RunFailure | None = None + + def __post_init__(self) -> None: + if not isinstance(self.result, RunResult): + raise TypeError("result must be a RunResult") + if not isinstance(self.delta, RunDelta): + raise TypeError("delta must be a RunDelta") + object.__setattr__(self, "audit_records", tuple(self.audit_records)) + if not all(isinstance(record, AuditRecord) for record in self.audit_records): + raise TypeError("audit_records must contain AuditRecord values") + if any(record.run_id != self.result.run_id for record in self.audit_records): + raise ValueError("audit records must belong to the outcome Run") + sequences = [record.sequence for record in self.audit_records] + if sequences != sorted(sequences) or len(sequences) != len(set(sequences)): + raise ValueError("audit record sequences must be unique and ordered") + if self.result.status is RunStatus.FAILED: + if self.failure is None: + raise ValueError("failed RunOutcome requires a RunFailure") + elif self.failure is not None: + raise ValueError("non-failed RunOutcome must not contain a RunFailure") diff --git a/src/ejagent/contracts/session.py b/src/ejagent/contracts/session.py new file mode 100644 index 0000000..f2085ef --- /dev/null +++ b/src/ejagent/contracts/session.py @@ -0,0 +1,145 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import Protocol + +from ejagent.contracts.audit import RunAudit +from ejagent.contracts.conversation import ConversationSnapshot +from ejagent.contracts.messages import ConversationMessage +from ejagent.contracts.runs import RunOutcome, RunResult, RunStatus + + +def _required_id(value: object, field_name: str) -> None: + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"{field_name} must not be empty") + + +@dataclass(frozen=True, slots=True) +class SessionSnapshot: + """Harness recovery state with Conversation kept as its own data domain.""" + + agent_id: str + conversation: ConversationSnapshot = field(default_factory=ConversationSnapshot) + last_result: RunResult | None = None + + def __post_init__(self) -> None: + _required_id(self.agent_id, "snapshot agent_id") + if not isinstance(self.conversation, ConversationSnapshot): + raise TypeError("snapshot conversation must be a ConversationSnapshot") + if self.last_result is not None and not isinstance(self.last_result, RunResult): + raise TypeError("snapshot last_result must be a RunResult or None") + + @property + def revision(self) -> int: + return self.conversation.revision + + @property + def messages(self) -> tuple[ConversationMessage, ...]: + return self.conversation.messages + + +@dataclass(frozen=True, slots=True) +class SessionCommit: + """One idempotent Run record proposed against a committed revision.""" + + agent_id: str + base: ConversationSnapshot + outcome: RunOutcome + + def __post_init__(self) -> None: + _required_id(self.agent_id, "commit agent_id") + if not isinstance(self.base, ConversationSnapshot): + raise TypeError("commit base must be a ConversationSnapshot") + if not isinstance(self.outcome, RunOutcome): + raise TypeError("commit outcome must be a RunOutcome") + if self.outcome.delta.base_revision != self.base_revision: + raise ValueError("commit and Delta base revisions must match") + + @property + def base_revision(self) -> int: + return self.base.revision + + @property + def base_messages(self) -> tuple[ConversationMessage, ...]: + return self.base.messages + + @property + def run_id(self) -> str: + return self.outcome.result.run_id + + @property + def advances_revision(self) -> bool: + """Return whether this Run is eligible for Conversation commit.""" + + return self.outcome.result.status is RunStatus.COMPLETED + + @property + def resulting_revision(self) -> int: + return self.base_revision + int(self.advances_revision) + + @property + def resulting_conversation(self) -> ConversationSnapshot: + if not self.advances_revision: + return self.base + return ConversationSnapshot( + revision=self.resulting_revision, + messages=(*self.base.messages, *self.outcome.delta.messages), + ) + + @property + def audit(self) -> RunAudit: + return RunAudit( + result=self.outcome.result, + base_revision=self.base_revision, + resulting_revision=self.resulting_revision, + committed=self.advances_revision, + records=self.outcome.audit_records, + failure=self.outcome.failure, + ) + + +class SessionStoreError(RuntimeError): + """Expected operational failure at the durable Session seam.""" + + +class SessionConflictError(SessionStoreError): + """A Session commit was based on a stale revision or Conversation.""" + + +class SessionStoreSerializationError(SessionStoreError): + """Durable Session data is malformed, unsupported, or not JSON-compatible.""" + + +class SessionStoreLockTimeoutError(SessionStoreError): + """A durable Session lock could not be acquired within its timeout.""" + + +class SessionMigrationError(SessionStoreError): + """Legacy Session data cannot be converted without losing semantics.""" + + def __init__( + self, + message: str, + *, + remediation: str, + ) -> None: + if not message.strip(): + raise ValueError("migration error message must not be empty") + if not remediation.strip(): + raise ValueError("migration remediation must not be empty") + self.remediation = remediation + super().__init__(f"{message} Remediation: {remediation}") + + +class SessionStore(Protocol): + """Durable compare-and-commit seam owned by an AgentHarness.""" + + async def load(self, agent_id: str) -> SessionSnapshot | None: + """Load the latest committed snapshot, if one exists.""" + + async def commit(self, commit: SessionCommit) -> SessionSnapshot: + """Atomically record one Run and return its resulting snapshot. + + Repeating the same commit for the same ``run_id`` must be idempotent. + A different commit using an existing ``run_id`` must fail. + """ diff --git a/src/ejagent/contracts/tools.py b/src/ejagent/contracts/tools.py new file mode 100644 index 0000000..ac4a131 --- /dev/null +++ b/src/ejagent/contracts/tools.py @@ -0,0 +1,154 @@ +from __future__ import annotations + +import re +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from enum import StrEnum +from typing import Protocol + +from ejagent.contracts.control import CancellationToken +from ejagent.contracts.json import ( + JsonObject, + JsonValue, + freeze_json_object, + freeze_json_value, +) +from ejagent.contracts.messages import ToolCall + +_TOOL_NAME_PATTERN = re.compile(r"^[A-Za-z0-9_-]{1,64}$") + + +class ToolEffect(StrEnum): + """Whether a Tool changes externally observable state.""" + + READ_ONLY = "read_only" + SIDE_EFFECTING = "side_effecting" + + +class ToolControl(StrEnum): + """Terminal control requested by one Tool execution.""" + + CONTINUE = "continue" + COMPLETE = "complete" + REJECT = "reject" + CANCEL = "cancel" + + +@dataclass(frozen=True, slots=True) +class ToolSemantics: + """Execution semantics consumed by scheduling and retry policy.""" + + effect: ToolEffect = ToolEffect.SIDE_EFFECTING + idempotent: bool = False + concurrency_key: str | None = None + + def __post_init__(self) -> None: + if not isinstance(self.effect, ToolEffect): + raise TypeError("tool effect must be a ToolEffect") + if not isinstance(self.idempotent, bool): + raise TypeError("tool idempotent must be a boolean") + if self.concurrency_key is not None and ( + not isinstance(self.concurrency_key, str) + or not self.concurrency_key.strip() + ): + raise ValueError("tool concurrency_key must not be empty") + + @classmethod + def read_only(cls, *, concurrency_key: str | None = None) -> ToolSemantics: + """Return conservative semantics for one read-only Tool.""" + + return cls( + effect=ToolEffect.READ_ONLY, + idempotent=True, + concurrency_key=concurrency_key, + ) + + +@dataclass(frozen=True, slots=True) +class ToolDefinition: + """Provider-neutral Tool definition exposed to a Runtime Kernel.""" + + name: str + description: str | None = None + input_schema: JsonObject = field(default_factory=dict) + semantics: ToolSemantics = field(default_factory=ToolSemantics) + + def __post_init__(self) -> None: + if not isinstance(self.name, str) or not _TOOL_NAME_PATTERN.fullmatch( + self.name + ): + raise ValueError( + "tool name must contain 1-64 letters, digits, underscores, or dashes" + ) + if self.description is not None and not isinstance(self.description, str): + raise TypeError("tool description must be a string or None") + if not isinstance(self.input_schema, Mapping): + raise TypeError("tool input_schema must be a JSON object") + object.__setattr__( + self, + "input_schema", + freeze_json_object(self.input_schema, label="tool input_schema"), + ) + if not isinstance(self.semantics, ToolSemantics): + raise TypeError("tool semantics must be ToolSemantics") + + +@dataclass(frozen=True, slots=True) +class ToolExecutionResult: + """Normalized operational result returned by a ToolExecutor.""" + + result: JsonValue + control: ToolControl = ToolControl.CONTINUE + output: str | None = None + error: str | None = None + + def __post_init__(self) -> None: + object.__setattr__( + self, + "result", + freeze_json_value(self.result, label="tool execution result"), + ) + if not isinstance(self.control, ToolControl): + raise TypeError("tool control must be a ToolControl") + if self.output is not None and not isinstance(self.output, str): + raise TypeError("tool output must be a string or None") + if self.error is not None and ( + not isinstance(self.error, str) or not self.error.strip() + ): + raise ValueError("tool error must not be empty") + + @property + def is_error(self) -> bool: + return self.error is not None + + +class ToolExecutionError(RuntimeError): + """Expected Tool infrastructure failure that prevents a result.""" + + def __init__(self, message: str, *, retryable: bool = False) -> None: + if not message.strip(): + raise ValueError("Tool execution error message must not be empty") + if not isinstance(retryable, bool): + raise TypeError("Tool execution retryable must be a boolean") + self.retryable = retryable + super().__init__(message) + + +class ToolProtocolError(RuntimeError): + """A ToolExecutor violated the stable Kernel protocol.""" + + +class ToolExecutor(Protocol): + """Ready Tool execution boundary consumed by a Runtime Kernel.""" + + @property + def definitions(self) -> Sequence[ToolDefinition]: + """Return the immutable Tool collection captured for a Run.""" + + async def execute( + self, + call: ToolCall, + *, + cancellation: CancellationToken, + ) -> ToolExecutionResult: + """Execute one Tool call without acquiring or releasing resources.""" diff --git a/src/ejagent/agent/usage.py b/src/ejagent/contracts/usage.py similarity index 53% rename from src/ejagent/agent/usage.py rename to src/ejagent/contracts/usage.py index 12dcc52..4890649 100644 --- a/src/ejagent/agent/usage.py +++ b/src/ejagent/contracts/usage.py @@ -2,12 +2,10 @@ from dataclasses import dataclass -from ejagent.providers.base import ModelUsage - @dataclass(frozen=True, slots=True) class RunUsage: - """Aggregated reported usage and coverage for one agent run.""" + """Aggregated reported usage and coverage for one Run.""" input_tokens: int = 0 output_tokens: int = 0 @@ -75,69 +73,3 @@ def to_dict(self) -> dict[str, int | None]: "cache_write_tokens": self.cache_write_tokens, "reasoning_tokens": self.reasoning_tokens, } - - -class UsageAccumulator: - """Mutable per-run collector producing immutable usage snapshots.""" - - def __init__(self) -> None: - self._request_count = 0 - self._reported_request_count = 0 - self._input_tokens = 0 - self._output_tokens = 0 - self._cache_read_tokens: int | None = None - self._cache_write_tokens: int | None = None - self._reasoning_tokens: int | None = None - - def begin_request(self) -> None: - self._request_count += 1 - - def record(self, usage: ModelUsage | None) -> None: - if usage is None: - return - if self._reported_request_count >= self._request_count: - raise RuntimeError("usage was recorded without an active request") - - previous_reports = self._reported_request_count - self._reported_request_count += 1 - self._input_tokens += usage.input_tokens - self._output_tokens += usage.output_tokens - self._cache_read_tokens = self._add_optional( - self._cache_read_tokens, - usage.cache_read_tokens, - previous_reports, - ) - self._cache_write_tokens = self._add_optional( - self._cache_write_tokens, - usage.cache_write_tokens, - previous_reports, - ) - self._reasoning_tokens = self._add_optional( - self._reasoning_tokens, - usage.reasoning_tokens, - previous_reports, - ) - - def snapshot(self) -> RunUsage: - return RunUsage( - input_tokens=self._input_tokens, - output_tokens=self._output_tokens, - total_tokens=self._input_tokens + self._output_tokens, - request_count=self._request_count, - reported_request_count=self._reported_request_count, - cache_read_tokens=self._cache_read_tokens, - cache_write_tokens=self._cache_write_tokens, - reasoning_tokens=self._reasoning_tokens, - ) - - @staticmethod - def _add_optional( - current: int | None, - value: int | None, - previous_reports: int, - ) -> int | None: - if previous_reports == 0: - return value - if current is None or value is None: - return None - return current + value diff --git a/src/ejagent/handlers/__init__.py b/src/ejagent/handlers/__init__.py deleted file mode 100644 index 35f6314..0000000 --- a/src/ejagent/handlers/__init__.py +++ /dev/null @@ -1,23 +0,0 @@ -"""Composable local and external tool handlers.""" - -from ejagent.handlers.base import ( - BaseHandler, - MethodToolHandler, - UnknownToolError, -) -from ejagent.handlers.definition import ( - ToolDefinition, - ToolDefinitionError, - ToolEffect, -) -from ejagent.handlers.mcp import McpToolHandler - -__all__ = [ - "BaseHandler", - "MethodToolHandler", - "ToolDefinition", - "ToolDefinitionError", - "ToolEffect", - "UnknownToolError", - "McpToolHandler", -] diff --git a/src/ejagent/handlers/base.py b/src/ejagent/handlers/base.py deleted file mode 100644 index 5232dd1..0000000 --- a/src/ejagent/handlers/base.py +++ /dev/null @@ -1,178 +0,0 @@ -from abc import ABC, abstractmethod -from collections.abc import Mapping, Sequence -from inspect import Parameter, signature -from typing import Any - -from ejagent.agent.cancellation import ( - CancellationSource, - CancellationToken, -) -from ejagent.agent.types import StepOutcome, ToolProgressReporter -from ejagent.handlers.definition import ( - ToolDefinition, - ToolDefinitionError, - ToolDefinitionInput, - ToolEffect, - ToolSchema, - normalize_tool_definition, - normalize_tool_definitions, -) -from ejagent.middleware.base import ToolCallContext - - -class UnknownToolError(KeyError): - """Raised when a handler is asked to execute an unknown tool.""" - - -class BaseHandler(ABC): - """Interface implemented by reusable groups of related tools.""" - - @property - @abstractmethod - def tools(self) -> Sequence[ToolSchema]: - """Return tool definitions in OpenAI function-calling format.""" - - @property - def tool_definitions(self) -> Sequence[ToolDefinition]: - """Return canonical definitions for legacy BaseHandler subclasses.""" - - return normalize_tool_definitions(self.tools) - - @property - def tool_names(self) -> tuple[str, ...]: - return tuple(tool.name for tool in self.tool_definitions) - - async def startup(self) -> None: - """Initialize optional external resources.""" - - async def shutdown(self) -> None: - """Release optional external resources.""" - - async def on_task_start(self) -> None: - """Prepare handler state for one new agent task.""" - - def can_handle(self, tool_name: str) -> bool: - return tool_name in self.tool_names - - def tool_effect(self, tool_name: str) -> ToolEffect: - """Return the conservative execution effect for one registered tool.""" - - for tool in self.tool_definitions: - if tool.name == tool_name: - return tool.effect - raise UnknownToolError(f"unknown tool {tool_name!r}") - - async def execute(self, context: ToolCallContext) -> StepOutcome: - """Execute a runtime context while adapting legacy dispatch methods.""" - - kwargs: dict[str, Any] = { - "cancellation": context.cancellation, - } - if _accepts_keyword(self.dispatch, "progress"): - kwargs["progress"] = context.progress - return await self.dispatch( - context.tool_name, - context.arguments, - **kwargs, - ) - - @abstractmethod - async def dispatch( - self, - tool_name: str, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - """Execute a registered tool with optional run cancellation.""" - - -class MethodToolHandler(BaseHandler): - """Handler that maps tool names to async ``do_`` methods.""" - - def __init__( - self, - tools: Sequence[ToolDefinitionInput], - *, - tool_effects: Mapping[str, ToolEffect] | None = None, - ) -> None: - definitions = list(normalize_tool_definitions(tools)) - names = tuple(tool.name for tool in definitions) - if len(names) != len(set(names)): - raise ToolDefinitionError("handler contains duplicate tool names") - effects = dict(tool_effects or {}) - unknown = effects.keys() - set(names) - if unknown: - names_text = ", ".join(sorted(unknown)) - raise ToolDefinitionError( - f"tool effects reference unknown tool(s): {names_text}" - ) - for name, effect in effects.items(): - if not isinstance(effect, ToolEffect): - raise ToolDefinitionError( - f"tool effect for {name!r} must be a ToolEffect" - ) - self._tool_definitions = tuple( - normalize_tool_definition( - tool, - effect=effects.get(tool.name), - ) - for tool in definitions - ) - - @property - def tools(self) -> Sequence[ToolSchema]: - return tuple(tool.to_openai_tool() for tool in self._tool_definitions) - - @property - def tool_definitions(self) -> Sequence[ToolDefinition]: - return self._tool_definitions - - def tool_effect(self, tool_name: str) -> ToolEffect: - if not self.can_handle(tool_name): - raise UnknownToolError(f"unknown tool {tool_name!r}") - return next( - tool.effect for tool in self._tool_definitions if tool.name == tool_name - ) - - async def dispatch( - self, - tool_name: str, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - progress: ToolProgressReporter | None = None, - ) -> StepOutcome: - if not self.can_handle(tool_name): - raise UnknownToolError(f"unknown tool {tool_name!r}") - - method = getattr(self, f"do_{tool_name}", None) - if method is None: - raise ToolDefinitionError( - f"{type(self).__name__} must define do_{tool_name}()" - ) - - token = cancellation or CancellationSource().token - kwargs: dict[str, Any] = {"cancellation": token} - if _accepts_keyword(method, "progress"): - kwargs["progress"] = progress - outcome = await method(dict(arguments), **kwargs) - if not isinstance(outcome, StepOutcome): - raise TypeError( - f"do_{tool_name}() must return StepOutcome, " - f"got {type(outcome).__name__}" - ) - return outcome - - -def _accepts_keyword(callable_: Any, keyword: str) -> bool: - """Return whether a callable accepts one explicit or variadic keyword.""" - - parameters = signature(callable_).parameters - parameter = parameters.get(keyword) - if parameter is not None and parameter.kind in { - Parameter.POSITIONAL_OR_KEYWORD, - Parameter.KEYWORD_ONLY, - }: - return True - return any(item.kind is Parameter.VAR_KEYWORD for item in parameters.values()) diff --git a/src/ejagent/handlers/definition.py b/src/ejagent/handlers/definition.py deleted file mode 100644 index de7eb8f..0000000 --- a/src/ejagent/handlers/definition.py +++ /dev/null @@ -1,149 +0,0 @@ -from __future__ import annotations - -import json -import re -from collections.abc import Mapping, Sequence -from copy import deepcopy -from dataclasses import dataclass, replace -from enum import StrEnum -from types import MappingProxyType -from typing import Any, TypeAlias - -ToolSchema: TypeAlias = dict[str, Any] -_TOOL_NAME_PATTERN = re.compile(r"^[A-Za-z0-9_-]{1,64}$") - - -def _mutable_json_copy(value: Any) -> Any: - if isinstance(value, Mapping): - return {key: _mutable_json_copy(item) for key, item in value.items()} - if isinstance(value, (list, tuple)): - return [_mutable_json_copy(item) for item in value] - return deepcopy(value) - - -def _freeze_json(value: Any) -> Any: - if isinstance(value, Mapping): - return MappingProxyType( - {key: _freeze_json(item) for key, item in value.items()} - ) - if isinstance(value, list): - return tuple(_freeze_json(item) for item in value) - return value - - -class ToolEffect(StrEnum): - """Declared side-effect class used by the Core tool scheduler.""" - - READ_ONLY = "read_only" - SIDE_EFFECTING = "side_effecting" - - -class ToolDefinitionError(ValueError): - """Raised when a handler exposes an invalid tool definition.""" - - -@dataclass(frozen=True, slots=True) -class ToolDefinition: - """Canonical OpenAI-first function tool definition used inside Core.""" - - name: str - description: str | None = None - parameters: Mapping[str, Any] | None = None - effect: ToolEffect = ToolEffect.SIDE_EFFECTING - strict: bool | None = None - - def __post_init__(self) -> None: - if not isinstance(self.name, str) or not _TOOL_NAME_PATTERN.fullmatch( - self.name - ): - raise ToolDefinitionError( - "tool name must contain 1-64 letters, digits, underscores, or dashes" - ) - if self.description is not None and not isinstance(self.description, str): - raise ToolDefinitionError("tool description must be a string or None") - if self.parameters is not None: - if not isinstance(self.parameters, Mapping): - raise ToolDefinitionError("tool parameters must be a mapping or None") - parameters = _mutable_json_copy(self.parameters) - try: - json.dumps(parameters, ensure_ascii=False, allow_nan=False) - except (TypeError, ValueError) as exc: - raise ToolDefinitionError( - "tool parameters must be JSON serializable" - ) from exc - object.__setattr__(self, "parameters", _freeze_json(parameters)) - if not isinstance(self.effect, ToolEffect): - raise ToolDefinitionError("tool effect must be a ToolEffect") - if self.strict is not None and not isinstance(self.strict, bool): - raise ToolDefinitionError("tool strict must be a boolean or None") - - @classmethod - def from_openai_tool( - cls, - tool: Mapping[str, Any], - *, - effect: ToolEffect = ToolEffect.SIDE_EFFECTING, - ) -> ToolDefinition: - """Normalize one legacy OpenAI function-calling dictionary.""" - - if not isinstance(tool, Mapping): - raise ToolDefinitionError("tool must be a mapping") - if tool.get("type") != "function": - raise ToolDefinitionError("only named function tools are supported") - function = tool.get("function") - if not isinstance(function, Mapping): - raise ToolDefinitionError("tool must contain a function mapping") - if "name" not in function: - raise ToolDefinitionError("tool must contain function.name") - return cls( - name=function["name"], - description=function.get("description"), - parameters=function.get("parameters"), - effect=effect, - strict=function.get("strict"), - ) - - def with_effect(self, effect: ToolEffect) -> ToolDefinition: - """Return a copy with one compatibility effect override applied.""" - - if not isinstance(effect, ToolEffect): - raise ToolDefinitionError("tool effect must be a ToolEffect") - return replace(self, effect=effect) - - def to_openai_tool(self) -> ToolSchema: - """Serialize this definition to the OpenAI function-calling shape.""" - - function: dict[str, Any] = {"name": self.name} - if self.description is not None: - function["description"] = self.description - if self.parameters is not None: - function["parameters"] = _mutable_json_copy(self.parameters) - if self.strict is not None: - function["strict"] = self.strict - return {"type": "function", "function": function} - - -ToolDefinitionInput: TypeAlias = ToolDefinition | Mapping[str, Any] - - -def normalize_tool_definition( - tool: ToolDefinitionInput, - *, - effect: ToolEffect | None = None, -) -> ToolDefinition: - """Return one detached canonical definition from either public input form.""" - - if isinstance(tool, ToolDefinition): - return tool if effect is None else tool.with_effect(effect) - return ToolDefinition.from_openai_tool( - tool, - effect=effect or ToolEffect.SIDE_EFFECTING, - ) - - -def normalize_tool_definitions( - tools: Sequence[ToolDefinitionInput], -) -> tuple[ToolDefinition, ...]: - """Normalize a tool collection without retaining legacy dictionaries.""" - - return tuple(normalize_tool_definition(tool) for tool in tools) diff --git a/src/ejagent/handlers/mcp.py b/src/ejagent/handlers/mcp.py deleted file mode 100644 index 9042c80..0000000 --- a/src/ejagent/handlers/mcp.py +++ /dev/null @@ -1,114 +0,0 @@ -from collections.abc import Mapping, Sequence -from pathlib import Path -from typing import Any - -from ejagent.agent.cancellation import ( - CancellationSource, - CancellationToken, -) -from ejagent.agent.types import StepOutcome, ToolProgressReporter -from ejagent.handlers.base import ( - BaseHandler, - UnknownToolError, -) -from ejagent.handlers.definition import ( - ToolDefinition, - ToolDefinitionError, - ToolEffect, - ToolSchema, - normalize_tool_definition, - normalize_tool_definitions, -) -from ejagent.plugins.mcp.mcp_manager import McpServerManager - - -class McpToolHandler(BaseHandler): - """Expose tools from configured MCP servers through the handler API.""" - - def __init__( - self, - config_path: str | Path | None = None, - *, - manager: Any | None = None, - tool_effects: Mapping[str, ToolEffect] | None = None, - ) -> None: - if config_path is None and manager is None: - raise ValueError("config_path is required when manager is not provided") - if config_path is not None and manager is not None: - raise ValueError("provide either config_path or manager, not both") - if manager is not None: - self.manager = manager - else: - assert config_path is not None - self.manager = McpServerManager(config_path) - self._tool_definitions: tuple[ToolDefinition, ...] = () - effects = dict(tool_effects or {}) - for name, effect in effects.items(): - if not isinstance(effect, ToolEffect): - raise ToolDefinitionError( - f"tool effect for {name!r} must be a ToolEffect" - ) - self._tool_effects = effects - self._started = False - - @property - def tools(self) -> Sequence[ToolSchema]: - return tuple(tool.to_openai_tool() for tool in self._tool_definitions) - - @property - def tool_definitions(self) -> Sequence[ToolDefinition]: - return self._tool_definitions - - async def startup(self) -> None: - if self._started: - return - await self.manager.startup() - try: - definitions = normalize_tool_definitions(self.manager.get_openai_tools()) - self._tool_definitions = definitions - unknown = self._tool_effects.keys() - set(self.tool_names) - if unknown: - names_text = ", ".join(sorted(unknown)) - raise ToolDefinitionError( - f"tool effects reference unknown MCP tool(s): {names_text}" - ) - self._tool_definitions = tuple( - normalize_tool_definition( - tool, - effect=self._tool_effects.get(tool.name), - ) - for tool in definitions - ) - except Exception: - self._tool_definitions = () - await self.manager.shutdown() - raise - self._started = True - - async def shutdown(self) -> None: - if not self._started: - return - await self.manager.shutdown() - self._tool_definitions = () - self._started = False - - async def dispatch( - self, - tool_name: str, - arguments: Mapping[str, Any], - *, - cancellation: CancellationToken | None = None, - progress: ToolProgressReporter | None = None, - ) -> StepOutcome: - if not self.can_handle(tool_name): - raise UnknownToolError(f"unknown MCP tool {tool_name!r}") - token = cancellation or CancellationSource().token - result = await token.run(self.manager.call_tool(tool_name, dict(arguments))) - return StepOutcome(result) - - def tool_effect(self, tool_name: str) -> ToolEffect: - if not self.can_handle(tool_name): - raise UnknownToolError(f"unknown MCP tool {tool_name!r}") - return next( - tool.effect for tool in self._tool_definitions if tool.name == tool_name - ) diff --git a/src/ejagent/harness/__init__.py b/src/ejagent/harness/__init__.py new file mode 100644 index 0000000..493f15b --- /dev/null +++ b/src/ejagent/harness/__init__.py @@ -0,0 +1,25 @@ +"""Long-lived single-agent lifecycle and durable commit coordination.""" + +from ejagent.harness._control import ( + FollowUpDiscardedError, + FollowUpHandle, + FollowUpRejectedError, +) +from ejagent.harness._memory import MemorySessionStore +from ejagent.harness.core import ( + AgentHarness, + HarnessClosedError, + HarnessStatus, + SessionStoreProtocolError, +) + +__all__ = [ + "AgentHarness", + "FollowUpDiscardedError", + "FollowUpHandle", + "FollowUpRejectedError", + "HarnessClosedError", + "HarnessStatus", + "MemorySessionStore", + "SessionStoreProtocolError", +] diff --git a/src/ejagent/harness/_control.py b/src/ejagent/harness/_control.py new file mode 100644 index 0000000..0e89b69 --- /dev/null +++ b/src/ejagent/harness/_control.py @@ -0,0 +1,89 @@ +from __future__ import annotations + +import asyncio +from collections import deque +from dataclasses import dataclass + +from ejagent.contracts.control import ControlReceipt, SteeringInput +from ejagent.contracts.runs import RunOutcome + + +class FollowUpRejectedError(RuntimeError): + """A follow-up was not admitted to an active Run chain.""" + + def __init__(self, receipt: ControlReceipt) -> None: + self.receipt = receipt + super().__init__(f"follow-up was not accepted: {receipt.status.value}") + + +class FollowUpDiscardedError(RuntimeError): + """An accepted follow-up was discarded before its Run started.""" + + +class FollowUpHandle: + """Asynchronous result holder returned immediately by follow_up().""" + + def __init__(self, receipt: ControlReceipt) -> None: + self.receipt = receipt + self._event = asyncio.Event() + self._outcome: RunOutcome | None = None + self._error: BaseException | None = None + + @property + def accepted(self) -> bool: + return self.receipt.accepted + + async def wait(self) -> RunOutcome: + """Wait for the admitted Run or raise its admission/execution error.""" + + if not self.accepted: + raise FollowUpRejectedError(self.receipt) + await self._event.wait() + if self._error is not None: + raise self._error + assert self._outcome is not None + return self._outcome + + def _resolve(self, outcome: RunOutcome) -> None: + if self._event.is_set(): + return + self._outcome = outcome + self._event.set() + + def _fail(self, error: BaseException) -> None: + if self._event.is_set(): + return + self._error = error + self._event.set() + + +@dataclass(frozen=True, slots=True) +class _QueuedFollowUp: + task: str + handle: FollowUpHandle + + +class _RunControls: + def __init__(self, capacity: int) -> None: + self._capacity = capacity + self._steering: deque[SteeringInput] = deque() + self._closed = False + + @property + def closed(self) -> bool: + return self._closed + + def offer(self, item: SteeringInput) -> bool: + if self._closed or len(self._steering) >= self._capacity: + return False + self._steering.append(item) + return True + + def drain_steering(self) -> tuple[SteeringInput, ...]: + items = tuple(self._steering) + self._steering.clear() + return items + + def close(self) -> tuple[SteeringInput, ...]: + self._closed = True + return self.drain_steering() diff --git a/src/ejagent/harness/_memory.py b/src/ejagent/harness/_memory.py new file mode 100644 index 0000000..5a9c1cc --- /dev/null +++ b/src/ejagent/harness/_memory.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +import asyncio + +from ejagent.contracts.audit import AuditReader, RunAudit +from ejagent.contracts.session import ( + SessionCommit, + SessionConflictError, + SessionSnapshot, + SessionStore, +) + + +class MemorySessionStore(SessionStore, AuditReader): + """Process-local atomic SessionStore with idempotent Run commits.""" + + def __init__(self) -> None: + self._snapshots: dict[str, SessionSnapshot] = {} + self._commits: dict[tuple[str, str], tuple[SessionCommit, SessionSnapshot]] = {} + self._order: dict[str, list[str]] = {} + self._lock = asyncio.Lock() + + async def load(self, agent_id: str) -> SessionSnapshot | None: + async with self._lock: + return self._snapshots.get(agent_id) + + async def commit(self, commit: SessionCommit) -> SessionSnapshot: + async with self._lock: + key = (commit.agent_id, commit.run_id) + existing = self._commits.get(key) + if existing is not None: + previous, snapshot = existing + if previous != commit: + raise SessionConflictError( + f"run_id {commit.run_id!r} already identifies " + "a different commit" + ) + return snapshot + + current = self._snapshots.get(commit.agent_id) + if current is None: + if commit.base_revision != 0: + raise SessionConflictError( + f"agent {commit.agent_id!r} has no revision " + f"{commit.base_revision}" + ) + current = SessionSnapshot( + agent_id=commit.agent_id, + conversation=commit.base, + ) + + if current.revision != commit.base_revision: + raise SessionConflictError( + f"agent {commit.agent_id!r} is at revision " + f"{current.revision}, not {commit.base_revision}" + ) + if current.messages != commit.base_messages: + raise SessionConflictError( + f"agent {commit.agent_id!r} Conversation does not match " + f"revision {commit.base_revision}" + ) + + snapshot = SessionSnapshot( + agent_id=commit.agent_id, + conversation=commit.resulting_conversation, + last_result=( + commit.outcome.result + if commit.advances_revision + else current.last_result + ), + ) + self._snapshots[commit.agent_id] = snapshot + self._commits[key] = (commit, snapshot) + self._order.setdefault(commit.agent_id, []).append(commit.run_id) + return snapshot + + async def load_audit(self, agent_id: str) -> tuple[RunAudit, ...]: + """Return Run audit facts without exposing Conversation deltas.""" + + async with self._lock: + return tuple( + self._commits[(agent_id, run_id)][0].audit + for run_id in self._order.get(agent_id, ()) + ) + + async def commits(self, agent_id: str) -> tuple[SessionCommit, ...]: + """Return recorded commits in insertion order for inspection.""" + + async with self._lock: + return tuple( + self._commits[(agent_id, run_id)][0] + for run_id in self._order.get(agent_id, ()) + ) diff --git a/src/ejagent/harness/core.py b/src/ejagent/harness/core.py new file mode 100644 index 0000000..342aba5 --- /dev/null +++ b/src/ejagent/harness/core.py @@ -0,0 +1,611 @@ +from __future__ import annotations + +import asyncio +from collections import deque +from collections.abc import Callable, Iterable, Mapping +from datetime import UTC, datetime +from enum import StrEnum +from types import TracebackType +from typing import Self +from uuid import uuid4 + +from ejagent.contracts.audit import RunAudit +from ejagent.contracts.context import ContextPipeline +from ejagent.contracts.control import ( + CancellationSource, + ControlKind, + ControlReceipt, + ControlStatus, + SteeringInput, +) +from ejagent.contracts.conversation import ConversationSnapshot +from ejagent.contracts.json import JsonValue +from ejagent.contracts.lifecycle import ManagedResource +from ejagent.contracts.messages import ConversationMessage +from ejagent.contracts.model import ModelPort +from ejagent.contracts.observer import RunObserver +from ejagent.contracts.runs import ( + AuditRecord, + FailureCode, + RunFailure, + RunIntent, + RunLimits, + RunOutcome, + RunPhase, + RunResult, + RunSpec, + RunStatus, + StopReason, +) +from ejagent.contracts.session import ( + SessionCommit, + SessionSnapshot, + SessionStore, + SessionStoreError, +) +from ejagent.contracts.tools import ToolExecutor +from ejagent.harness._control import ( + FollowUpDiscardedError, + FollowUpHandle, + _QueuedFollowUp, + _RunControls, +) +from ejagent.harness._memory import MemorySessionStore +from ejagent.kernel import RuntimeKernel + +RunIdFactory = Callable[[], str] +Clock = Callable[[], datetime] + + +def _utc_now() -> datetime: + return datetime.now(UTC) + + +class HarnessStatus(StrEnum): + """Observable lifecycle phase of one AgentHarness.""" + + NEW = "new" + STARTING = "starting" + READY = "ready" + RUNNING = "running" + STOPPING = "stopping" + CLOSED = "closed" + + +class HarnessClosedError(RuntimeError): + """A Run was requested after Harness shutdown began.""" + + +class SessionStoreProtocolError(RuntimeError): + """A SessionStore returned a snapshot that violates its contract.""" + + +class AgentHarness: + """Own one agent's resources, Conversation, and atomic Run commits.""" + + def __init__( + self, + *, + agent_id: str, + model: ModelPort, + tools: ToolExecutor, + context: ContextPipeline | None = None, + initial_messages: Iterable[ConversationMessage] = (), + store: SessionStore | None = None, + observers: Iterable[RunObserver] = (), + resources: Iterable[object] = (), + limits: RunLimits | None = None, + configuration_revision: str = "default", + run_id_factory: RunIdFactory | None = None, + clock: Clock | None = None, + steering_capacity: int = 16, + follow_up_capacity: int = 16, + ) -> None: + if not isinstance(agent_id, str): + raise TypeError("agent_id must be a string") + agent_id = agent_id.strip() + if not agent_id: + raise ValueError("agent_id must not be empty") + if not isinstance(configuration_revision, str): + raise TypeError("configuration_revision must be a string") + if not configuration_revision.strip(): + raise ValueError("configuration_revision must not be empty") + if limits is not None and not isinstance(limits, RunLimits): + raise TypeError("limits must be RunLimits or None") + if run_id_factory is not None and not callable(run_id_factory): + raise TypeError("run_id_factory must be callable or None") + if clock is not None and not callable(clock): + raise TypeError("clock must be callable or None") + self._validate_capacity(steering_capacity, "steering_capacity") + self._validate_capacity(follow_up_capacity, "follow_up_capacity") + + initial = SessionSnapshot( + agent_id=agent_id, + conversation=ConversationSnapshot(messages=tuple(initial_messages)), + ) + self._agent_id = agent_id + self._model = model + self._tools = tools + self._context = context + self._observers = tuple(observers) + self._store = store if store is not None else MemorySessionStore() + self._snapshot = initial + self._limits = limits or RunLimits() + self._configuration_revision = configuration_revision + self._run_id_factory = run_id_factory or (lambda: str(uuid4())) + self._clock = clock or _utc_now + self._steering_capacity = steering_capacity + self._follow_up_capacity = follow_up_capacity + self._kernel = RuntimeKernel( + model=model, + tools=tools, + context=context, + clock=self._clock, + ) + self._resources = self._managed_resources( + ( + self._store, + self._model, + self._tools, + self._context, + *self._observers, + *resources, + ) + ) + self._started_resources: tuple[ManagedResource, ...] = () + self._status = HarnessStatus.NEW + self._closing = False + self._active_cancellation: CancellationSource | None = None + self._active_controls: _RunControls | None = None + self._follow_ups: deque[_QueuedFollowUp] = deque() + self._outstanding_follow_ups = 0 + self._follow_up_worker: asyncio.Task[None] | None = None + self._observer_tasks: set[asyncio.Task[None]] = set() + self._lifecycle_lock = asyncio.Lock() + self._run_lock = asyncio.Lock() + + @property + def agent_id(self) -> str: + return self._agent_id + + @property + def status(self) -> HarnessStatus: + return self._status + + @property + def revision(self) -> int: + return self._snapshot.revision + + @property + def messages(self) -> tuple[ConversationMessage, ...]: + return self._snapshot.messages + + @property + def last_result(self) -> RunResult | None: + return self._snapshot.last_result + + @property + def snapshot(self) -> SessionSnapshot: + return self._snapshot + + @property + def pending_follow_up_count(self) -> int: + return self._outstanding_follow_ups + + async def __aenter__(self) -> Self: + await self.start() + return self + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> None: + await self.shutdown() + + async def start(self) -> None: + """Start owned resources transactionally and restore Conversation state.""" + + async with self._lifecycle_lock: + if self._status in (HarnessStatus.READY, HarnessStatus.RUNNING): + return + if self._closing or self._status is HarnessStatus.CLOSED: + raise HarnessClosedError(f"agent harness {self._agent_id!r} is closed") + + self._status = HarnessStatus.STARTING + started: list[ManagedResource] = [] + try: + for resource in self._resources: + await resource.start() + started.append(resource) + loaded = await self._store.load(self._agent_id) + if loaded is not None: + self._validate_loaded_snapshot(loaded) + self._snapshot = loaded + except BaseException: + await self._rollback_start(started) + self._status = HarnessStatus.NEW + raise + + self._started_resources = tuple(started) + self._status = HarnessStatus.READY + + async def run( + self, + task: str, + *, + limits: RunLimits | None = None, + metadata: Mapping[str, JsonValue] | None = None, + ) -> RunOutcome: + """Run one task after earlier calls, then atomically commit on success.""" + + if not isinstance(task, str) or not task.strip(): + raise ValueError("task must not be empty") + return await self._execute( + intent=RunIntent.TASK, + task=task, + limits=limits, + metadata=metadata, + ) + + async def continue_run( + self, + *, + limits: RunLimits | None = None, + metadata: Mapping[str, JsonValue] | None = None, + ) -> RunOutcome: + """Continue committed Conversation without appending a user message.""" + + return await self._execute( + intent=RunIntent.CONTINUE, + task=None, + limits=limits, + metadata=metadata, + ) + + def cancel(self, reason: str | None = None) -> bool: + """Request cooperative cancellation of the active Run, if any.""" + + source = self._active_cancellation + return source.cancel(reason) if source is not None else False + + def steer(self, content: str) -> ControlReceipt: + """Queue one transient instruction for the next model-call safe point.""" + + if not isinstance(content, str) or not content.strip(): + raise ValueError("steering content must not be empty") + input_id = uuid4().hex + if self._closing or self._status is HarnessStatus.CLOSED: + status = ControlStatus.CLOSED + elif self._status is not HarnessStatus.RUNNING: + status = ControlStatus.NOT_RUNNING + elif self._active_controls is None or self._active_controls.closed: + status = ControlStatus.TOO_LATE + elif self._active_controls.offer( + SteeringInput(input_id=input_id, content=content.strip()) + ): + status = ControlStatus.ACCEPTED + else: + status = ControlStatus.QUEUE_FULL + return ControlReceipt(input_id, ControlKind.STEERING, status) + + def follow_up(self, task: str) -> FollowUpHandle: + """Submit an independent FIFO Run to follow the active Run chain.""" + + if not isinstance(task, str) or not task.strip(): + raise ValueError("follow-up task must not be empty") + input_id = uuid4().hex + if self._closing or self._status is HarnessStatus.CLOSED: + status = ControlStatus.CLOSED + elif self._status is not HarnessStatus.RUNNING: + status = ControlStatus.NOT_RUNNING + elif self._outstanding_follow_ups >= self._follow_up_capacity: + status = ControlStatus.QUEUE_FULL + else: + status = ControlStatus.ACCEPTED + receipt = ControlReceipt(input_id, ControlKind.FOLLOW_UP, status) + handle = FollowUpHandle(receipt) + if receipt.accepted: + self._follow_ups.append(_QueuedFollowUp(task.strip(), handle)) + self._outstanding_follow_ups += 1 + if self._follow_up_worker is None: + self._follow_up_worker = asyncio.create_task(self._run_follow_ups()) + return handle + + async def shutdown(self) -> None: + """Stop accepting Runs, cancel active work, and release resources.""" + + async with self._lifecycle_lock: + if self._status is HarnessStatus.CLOSED: + return + self._closing = True + self._status = HarnessStatus.STOPPING + self._discard_pending_follow_ups("AgentHarness is shutting down") + self.cancel("AgentHarness is shutting down") + + async with self._run_lock: + await self._flush_observers() + failures: list[BaseException] = [] + for resource in reversed(self._started_resources): + try: + await resource.shutdown() + except BaseException as exc: + failures.append(exc) + self._started_resources = () + self._status = HarnessStatus.CLOSED + + worker = self._follow_up_worker + if worker is not None and worker is not asyncio.current_task(): + await asyncio.gather(worker, return_exceptions=True) + + if failures: + raise BaseExceptionGroup( + "one or more AgentHarness resources failed to shut down", + failures, + ) + + async def _execute( + self, + *, + intent: RunIntent, + task: str | None, + limits: RunLimits | None, + metadata: Mapping[str, JsonValue] | None, + ) -> RunOutcome: + await self.start() + async with self._run_lock: + if self._closing or self._status is HarnessStatus.CLOSED: + raise HarnessClosedError(f"agent harness {self._agent_id!r} is closed") + + base = self._snapshot + spec = RunSpec( + run_id=self._run_id_factory(), + base_revision=base.revision, + intent=intent, + task=task, + messages=base.messages, + limits=limits or self._limits, + configuration_revision=self._configuration_revision, + metadata=metadata or {}, + ) + cancellation = CancellationSource() + controls = _RunControls(self._steering_capacity) + self._active_cancellation = cancellation + self._active_controls = controls + self._status = HarnessStatus.RUNNING + try: + outcome = await self._kernel.run( + spec, + cancellation=cancellation.token, + controls=controls, + ) + discarded = controls.close() + self._active_controls = None + outcome = self._append_discarded_steering(outcome, discarded) + outcome = await self._commit(base, outcome) + self._dispatch_observers(base, outcome) + return outcome + finally: + controls.close() + self._active_cancellation = None + if self._active_controls is controls: + self._active_controls = None + if not self._closing: + self._status = HarnessStatus.READY + + async def _run_follow_ups(self) -> None: + try: + while self._follow_ups: + queued = self._follow_ups.popleft() + try: + if self._closing: + queued.handle._fail( + FollowUpDiscardedError("AgentHarness is shutting down") + ) + continue + try: + outcome = await self._execute( + intent=RunIntent.TASK, + task=queued.task, + limits=None, + metadata={ + "control_input_id": queued.handle.receipt.input_id + }, + ) + except HarnessClosedError: + queued.handle._fail( + FollowUpDiscardedError("AgentHarness is shutting down") + ) + except BaseException as exc: + queued.handle._fail(exc) + else: + queued.handle._resolve(outcome) + finally: + self._outstanding_follow_ups -= 1 + finally: + self._follow_up_worker = None + + def _append_discarded_steering( + self, + outcome: RunOutcome, + discarded: tuple[SteeringInput, ...], + ) -> RunOutcome: + if not discarded: + return outcome + records = list(outcome.audit_records) + for item in discarded: + records.append( + AuditRecord( + run_id=outcome.result.run_id, + sequence=len(records) + 1, + kind="steering_discarded", + occurred_at=self._clock(), + payload={ + "input_id": item.input_id, + "content": item.content, + "reason": "run_finished", + }, + ) + ) + return RunOutcome( + result=outcome.result, + delta=outcome.delta, + audit_records=tuple(records), + failure=outcome.failure, + ) + + def _dispatch_observers( + self, + base: SessionSnapshot, + outcome: RunOutcome, + ) -> None: + if not self._observers: + return + audit = SessionCommit( + agent_id=self._agent_id, + base=base.conversation, + outcome=outcome, + ).audit + for observer in self._observers: + task = asyncio.create_task(self._notify_observer(observer, audit)) + self._observer_tasks.add(task) + task.add_done_callback(self._observer_done) + + def _observer_done(self, task: asyncio.Task[None]) -> None: + self._observer_tasks.discard(task) + if not task.cancelled(): + task.exception() + + async def _flush_observers(self) -> None: + if self._observer_tasks: + await asyncio.gather(*tuple(self._observer_tasks), return_exceptions=True) + + @staticmethod + async def _notify_observer(observer: RunObserver, audit: RunAudit) -> None: + await observer.observe(audit) + + def _discard_pending_follow_ups(self, reason: str) -> None: + while self._follow_ups: + queued = self._follow_ups.popleft() + queued.handle._fail(FollowUpDiscardedError(reason)) + self._outstanding_follow_ups -= 1 + + async def _commit( + self, + base: SessionSnapshot, + outcome: RunOutcome, + ) -> RunOutcome: + commit = SessionCommit( + agent_id=self._agent_id, + base=base.conversation, + outcome=outcome, + ) + try: + snapshot = await self._store.commit(commit) + except SessionStoreError as exc: + return self._persistence_failure(outcome, exc) + + self._validate_committed_snapshot(snapshot, commit, base) + self._snapshot = snapshot + return outcome + + def _persistence_failure( + self, + outcome: RunOutcome, + error: SessionStoreError, + ) -> RunOutcome: + failure = RunFailure( + phase=RunPhase.COMMIT, + code=FailureCode.PERSISTENCE_FAILED, + message=str(error) or type(error).__name__, + retryable=True, + cause=error, + ) + result = RunResult( + run_id=outcome.result.run_id, + status=RunStatus.FAILED, + stop_reason=StopReason.PERSISTENCE_FAILED, + turns=outcome.result.turns, + usage=outcome.result.usage, + ) + audit = ( + *outcome.audit_records, + AuditRecord( + run_id=outcome.result.run_id, + sequence=len(outcome.audit_records) + 1, + kind="commit_failed", + occurred_at=self._clock(), + payload={"error": str(error) or type(error).__name__}, + ), + ) + return RunOutcome( + result=result, + delta=outcome.delta, + audit_records=audit, + failure=failure, + ) + + def _validate_loaded_snapshot(self, snapshot: object) -> None: + if not isinstance(snapshot, SessionSnapshot): + raise SessionStoreProtocolError( + "SessionStore.load() must return SessionSnapshot or None" + ) + if snapshot.agent_id != self._agent_id: + raise SessionStoreProtocolError( + f"SessionStore loaded agent {snapshot.agent_id!r} for " + f"{self._agent_id!r}" + ) + + def _validate_committed_snapshot( + self, + snapshot: object, + commit: SessionCommit, + base: SessionSnapshot, + ) -> None: + if not isinstance(snapshot, SessionSnapshot): + raise SessionStoreProtocolError( + "SessionStore.commit() must return SessionSnapshot" + ) + expected_result = ( + commit.outcome.result if commit.advances_revision else base.last_result + ) + if ( + snapshot.agent_id != self._agent_id + or snapshot.revision != commit.resulting_revision + or snapshot.conversation != commit.resulting_conversation + or snapshot.last_result != expected_result + ): + raise SessionStoreProtocolError( + "SessionStore.commit() returned a snapshot inconsistent " + "with the proposed commit" + ) + + async def _rollback_start(self, started: list[ManagedResource]) -> None: + for resource in reversed(started): + try: + await resource.shutdown() + except BaseException: + pass + + @staticmethod + def _managed_resources( + resources: Iterable[object], + ) -> tuple[ManagedResource, ...]: + managed: list[ManagedResource] = [] + seen: set[int] = set() + for resource in resources: + identity = id(resource) + if identity in seen: + continue + seen.add(identity) + if isinstance(resource, ManagedResource): + managed.append(resource) + return tuple(managed) + + @staticmethod + def _validate_capacity(value: object, field_name: str) -> None: + if isinstance(value, bool) or not isinstance(value, int): + raise TypeError(f"{field_name} must be an integer") + if value <= 0: + raise ValueError(f"{field_name} must be greater than zero") diff --git a/src/ejagent/kernel/__init__.py b/src/ejagent/kernel/__init__.py new file mode 100644 index 0000000..22c2267 --- /dev/null +++ b/src/ejagent/kernel/__init__.py @@ -0,0 +1,5 @@ +"""Single-Run Runtime Kernel.""" + +from ejagent.kernel.runtime import RuntimeKernel + +__all__ = ["RuntimeKernel"] diff --git a/src/ejagent/kernel/_workspace.py b/src/ejagent/kernel/_workspace.py new file mode 100644 index 0000000..310b672 --- /dev/null +++ b/src/ejagent/kernel/_workspace.py @@ -0,0 +1,98 @@ +from __future__ import annotations + +import json + +from ejagent.contracts.json import JsonValue, thaw_json_value +from ejagent.contracts.messages import ( + AssistantMessage, + ConversationMessage, + ToolCall, + ToolResultMessage, + UserMessage, +) +from ejagent.contracts.runs import RunDelta, RunIntent, RunSpec + + +class _RunWorkspace: + """Private mutable message workspace owned by one Kernel invocation.""" + + def __init__(self, spec: RunSpec) -> None: + self.spec = spec + self._messages = list(spec.messages) + self._delta: list[ConversationMessage] = [] + self._tool_call_ids: set[str] = set() + self._last_tool_signature: tuple[str, str] | None = None + self._repeated_tool_calls = 0 + self.turn = 0 + if spec.intent is RunIntent.TASK: + assert spec.task is not None + self.append(UserMessage(spec.task)) + + @property + def messages(self) -> tuple[ConversationMessage, ...]: + return tuple(self._messages) + + @property + def committed_messages(self) -> tuple[ConversationMessage, ...]: + return self.spec.messages + + @property + def pending_messages(self) -> tuple[ConversationMessage, ...]: + return tuple(self._delta) + + @property + def delta(self) -> RunDelta: + return RunDelta( + base_revision=self.spec.base_revision, + messages=tuple(self._delta), + ) + + def advance_turn(self) -> int: + self.turn += 1 + return self.turn + + def append(self, message: ConversationMessage) -> None: + self._messages.append(message) + self._delta.append(message) + + def append_assistant(self, message: AssistantMessage) -> None: + for call in message.tool_calls: + if call.id in self._tool_call_ids: + raise ValueError(f"duplicate tool call id {call.id!r}") + self._tool_call_ids.add(call.id) + self.append(message) + + def append_tool_result( + self, + call: ToolCall, + *, + result: JsonValue, + is_error: bool, + ) -> ToolResultMessage: + message = ToolResultMessage( + tool_call_id=call.id, + tool_name=call.name, + result=result, + is_error=is_error, + ) + self.append(message) + return message + + def record_tool_call(self, call: ToolCall) -> bool: + """Update the repeat guard and return whether its limit was reached.""" + + signature = ( + call.name, + json.dumps( + thaw_json_value(call.arguments), + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ), + ) + if signature == self._last_tool_signature: + self._repeated_tool_calls += 1 + else: + self._last_tool_signature = signature + self._repeated_tool_calls = 1 + return self._repeated_tool_calls >= self.spec.limits.max_repeated_tool_calls diff --git a/src/ejagent/kernel/runtime.py b/src/ejagent/kernel/runtime.py new file mode 100644 index 0000000..d0691a4 --- /dev/null +++ b/src/ejagent/kernel/runtime.py @@ -0,0 +1,694 @@ +from __future__ import annotations + +from collections.abc import Callable, Mapping +from contextlib import suppress +from datetime import UTC, datetime + +from ejagent.context import IdentityContextPipeline +from ejagent.contracts.context import ( + ContextBuildError, + ContextPipeline, + ContextProtocolError, + ContextRequest, + ContextView, +) +from ejagent.contracts.control import ( + CancellationSource, + CancellationToken, + ControlProtocolError, + RunCancelledError, + RunControlSource, + SteeringInput, +) +from ejagent.contracts.json import ( + JsonObject, + JsonValue, + freeze_json_object, +) +from ejagent.contracts.messages import ( + AssistantMessage, + ToolCall, + TransientInstruction, +) +from ejagent.contracts.model import ( + ModelCallError, + ModelPort, + ModelProtocolError, + ModelRequest, + ModelResponseCompleted, + ModelTextDelta, + ModelThinkingDelta, + ModelUsage, +) +from ejagent.contracts.runs import ( + AuditRecord, + FailureCode, + RunFailure, + RunOutcome, + RunPhase, + RunResult, + RunSpec, + RunStatus, + StopReason, +) +from ejagent.contracts.tools import ( + ToolControl, + ToolDefinition, + ToolExecutionError, + ToolExecutionResult, + ToolExecutor, + ToolProtocolError, +) +from ejagent.contracts.usage import RunUsage +from ejagent.kernel._workspace import _RunWorkspace + +Clock = Callable[[], datetime] + + +def _utc_now() -> datetime: + return datetime.now(UTC) + + +class _UsageAccumulator: + def __init__(self) -> None: + self.request_count = 0 + self.reported_request_count = 0 + self.input_tokens = 0 + self.output_tokens = 0 + self.cache_read_tokens: int | None = None + self.cache_write_tokens: int | None = None + self.reasoning_tokens: int | None = None + + def begin_request(self) -> None: + self.request_count += 1 + + def record(self, usage: ModelUsage | None) -> None: + if usage is None: + return + previous_reports = self.reported_request_count + self.reported_request_count += 1 + self.input_tokens += usage.input_tokens + self.output_tokens += usage.output_tokens + self.cache_read_tokens = self._add_optional( + self.cache_read_tokens, + usage.cache_read_tokens, + previous_reports, + ) + self.cache_write_tokens = self._add_optional( + self.cache_write_tokens, + usage.cache_write_tokens, + previous_reports, + ) + self.reasoning_tokens = self._add_optional( + self.reasoning_tokens, + usage.reasoning_tokens, + previous_reports, + ) + + def snapshot(self) -> RunUsage: + return RunUsage( + input_tokens=self.input_tokens, + output_tokens=self.output_tokens, + total_tokens=self.input_tokens + self.output_tokens, + request_count=self.request_count, + reported_request_count=self.reported_request_count, + cache_read_tokens=self.cache_read_tokens, + cache_write_tokens=self.cache_write_tokens, + reasoning_tokens=self.reasoning_tokens, + ) + + @staticmethod + def _add_optional( + current: int | None, + value: int | None, + previous_reports: int, + ) -> int | None: + if previous_reports == 0: + return value + if current is None or value is None: + return None + return current + value + + +class _AuditTrail: + def __init__(self, run_id: str, clock: Clock) -> None: + self.run_id = run_id + self.clock = clock + self.records: list[AuditRecord] = [] + + def append(self, kind: str, payload: Mapping[str, JsonValue] | None = None) -> None: + self.records.append( + AuditRecord( + run_id=self.run_id, + sequence=len(self.records) + 1, + kind=kind, + occurred_at=self.clock(), + payload=payload or {}, + ) + ) + + +class RuntimeKernel: + """Execute one deterministic Model-Tool Run over a private workspace.""" + + def __init__( + self, + *, + model: ModelPort, + tools: ToolExecutor, + context: ContextPipeline | None = None, + clock: Clock | None = None, + ) -> None: + self._model = model + self._tools = tools + self._context = context if context is not None else IdentityContextPipeline() + self._clock = clock or _utc_now + + async def run( + self, + spec: RunSpec, + *, + cancellation: CancellationToken | None = None, + controls: RunControlSource | None = None, + ) -> RunOutcome: + """Execute one Run without mutating the supplied RunSpec or Harness state.""" + + if not isinstance(spec, RunSpec): + raise TypeError("spec must be a RunSpec") + token = cancellation or CancellationSource().token + workspace = _RunWorkspace(spec) + usage = _UsageAccumulator() + audit = _AuditTrail(spec.run_id, self._clock) + definitions = self._snapshot_tool_definitions() + tool_names = {definition.name for definition in definitions} + audit.append( + "run_started", + { + "base_revision": spec.base_revision, + "intent": spec.intent.value, + "configuration_revision": spec.configuration_revision, + }, + ) + + try: + for _ in range(spec.limits.max_turns): + token.raise_if_cancelled() + budget_failure = self._budget_failure(spec, usage) + if budget_failure is not None: + return self._failed( + workspace, + usage, + audit, + stop_reason=budget_failure[0], + failure=budget_failure[1], + ) + + turn = workspace.advance_turn() + audit.append("turn_started", {"turn": turn}) + steering = self._drain_steering(controls) + for item in steering: + audit.append( + "steering_applied", + { + "turn": turn, + "input_id": item.input_id, + "content": item.content, + }, + ) + context = await self._build_context( + workspace, + cancellation=token, + turn=turn, + transient_instructions=tuple( + TransientInstruction(item.content, "steering") + for item in steering + ), + ) + audit.append( + "context_built", + { + "turn": turn, + "source_revision": context.source_revision, + "message_count": len(context.messages), + "metadata": context.metadata, + }, + ) + usage.begin_request() + completed = await self._request_model( + ModelRequest( + messages=context.messages, + tools=definitions, + ), + cancellation=token, + audit=audit, + turn=turn, + ) + usage.record(completed.usage) + message = completed.message + try: + workspace.append_assistant(message) + except ValueError as exc: + raise ModelProtocolError(str(exc)) from exc + audit.append( + "assistant_message", + self._assistant_payload(message, turn), + ) + + if not message.tool_calls: + assert message.content is not None + audit.append("turn_completed", {"turn": turn}) + return self._terminal( + workspace, + usage, + audit, + status=RunStatus.COMPLETED, + stop_reason=StopReason.TEXT_RESPONSE, + output=message.content, + ) + + for call in message.tool_calls: + if workspace.record_tool_call(call): + audit.append("turn_completed", {"turn": turn}) + return self._failed( + workspace, + usage, + audit, + stop_reason=StopReason.REPEATED_TOOL_CALL, + failure=RunFailure( + phase=RunPhase.TOOL, + code=FailureCode.TOOL_ERROR, + message=( + f"tool {call.name!r} was called with the same " + "arguments too many consecutive times" + ), + ), + ) + execution = await self._execute_tool( + call, + known=call.name in tool_names, + cancellation=token, + audit=audit, + turn=turn, + ) + workspace.append_tool_result( + call, + result=execution.result, + is_error=execution.is_error, + ) + if execution.control is not ToolControl.CONTINUE: + audit.append("turn_completed", {"turn": turn}) + return self._tool_terminal( + workspace, + usage, + audit, + execution, + ) + audit.append("turn_completed", {"turn": turn}) + + return self._failed( + workspace, + usage, + audit, + stop_reason=StopReason.MAX_STEPS, + failure=RunFailure( + phase=RunPhase.CONTROL, + code=FailureCode.BUDGET_EXCEEDED, + message=( + f"Run {spec.run_id!r} did not finish within " + f"{spec.limits.max_turns} turns" + ), + ), + ) + except RunCancelledError as exc: + return self._terminal( + workspace, + usage, + audit, + status=RunStatus.CANCELLED, + stop_reason=StopReason.EXTERNAL_ABORT, + output=str(exc), + ) + except ContextBuildError as exc: + return self._failed( + workspace, + usage, + audit, + stop_reason=self._context_stop_reason(exc.code), + failure=RunFailure( + phase=RunPhase.CONTEXT, + code=exc.code, + message=str(exc), + retryable=exc.retryable, + cause=exc, + ), + ) + except ModelCallError as exc: + return self._failed( + workspace, + usage, + audit, + stop_reason=self._model_stop_reason(exc.code), + failure=RunFailure( + phase=RunPhase.MODEL, + code=exc.code, + message=str(exc), + retryable=exc.retryable, + cause=exc, + ), + ) + except ToolExecutionError as exc: + return self._failed( + workspace, + usage, + audit, + stop_reason=StopReason.RUNTIME_ERROR, + failure=RunFailure( + phase=RunPhase.TOOL, + code=FailureCode.TOOL_ERROR, + message=str(exc), + retryable=exc.retryable, + cause=exc, + ), + ) + + async def _build_context( + self, + workspace: _RunWorkspace, + *, + cancellation: CancellationToken, + turn: int, + transient_instructions: tuple[TransientInstruction, ...], + ) -> ContextView: + request = ContextRequest( + run_id=workspace.spec.run_id, + source_revision=workspace.spec.base_revision, + turn=turn, + committed_messages=workspace.committed_messages, + pending_messages=workspace.pending_messages, + transient_instructions=transient_instructions, + metadata=workspace.spec.metadata, + ) + try: + view = await cancellation.run( + self._context.build(request, cancellation=cancellation) + ) + except (RunCancelledError, ContextBuildError, ContextProtocolError): + raise + except Exception as exc: + raise ContextProtocolError( + f"ContextPipeline raised an undeclared {type(exc).__name__}" + ) from exc + if not isinstance(view, ContextView): + raise ContextProtocolError( + "ContextPipeline.build() must return ContextView" + ) + if ( + view.run_id != request.run_id + or view.source_revision != request.source_revision + or view.turn != request.turn + ): + raise ContextProtocolError( + "ContextView identity does not match its ContextRequest" + ) + return view + + @staticmethod + def _drain_steering( + controls: RunControlSource | None, + ) -> tuple[SteeringInput, ...]: + if controls is None: + return () + items = controls.drain_steering() + if not isinstance(items, tuple) or not all( + isinstance(item, SteeringInput) for item in items + ): + raise ControlProtocolError( + "RunControlSource.drain_steering() must return " + "tuple[SteeringInput, ...]" + ) + return items + + def _snapshot_tool_definitions(self) -> tuple[ToolDefinition, ...]: + definitions = tuple(self._tools.definitions) + if not all(isinstance(item, ToolDefinition) for item in definitions): + raise ToolProtocolError( + "ToolExecutor definitions must contain ToolDefinition values" + ) + names = [definition.name for definition in definitions] + if len(names) != len(set(names)): + raise ToolProtocolError("ToolExecutor contains duplicate Tool names") + return definitions + + async def _request_model( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + audit: _AuditTrail, + turn: int, + ) -> ModelResponseCompleted: + stream = self._model.stream(request, cancellation=cancellation) + iterator = stream.__aiter__() + try: + while True: + cancellation.raise_if_cancelled() + try: + event = await cancellation.run(anext(iterator)) + except StopAsyncIteration: + break + if isinstance(event, ModelTextDelta): + audit.append( + "model_text_delta", + {"turn": turn, "delta": event.delta}, + ) + elif isinstance(event, ModelThinkingDelta): + audit.append( + "model_thinking_delta", + {"turn": turn, "delta": event.delta}, + ) + elif isinstance(event, ModelResponseCompleted): + return event + else: + raise ModelProtocolError( + "ModelPort returned an unsupported stream event: " + f"{type(event).__name__}" + ) + finally: + close = getattr(iterator, "aclose", None) + if close is not None: + with suppress(Exception): + await close() + raise ModelProtocolError("ModelPort stream ended without completion") + + async def _execute_tool( + self, + call: ToolCall, + *, + known: bool, + cancellation: CancellationToken, + audit: _AuditTrail, + turn: int, + ) -> ToolExecutionResult: + audit.append( + "tool_started", + { + "turn": turn, + "tool_call_id": call.id, + "tool_name": call.name, + "arguments": call.arguments, + }, + ) + if not known: + execution = ToolExecutionResult( + { + "status": "error", + "tool": call.name, + "error": f"unknown tool {call.name!r}", + }, + error=f"unknown tool {call.name!r}", + ) + else: + try: + execution = await cancellation.run( + self._tools.execute(call, cancellation=cancellation) + ) + except (RunCancelledError, ToolExecutionError): + raise + except Exception as exc: + raise ToolProtocolError( + f"ToolExecutor raised an undeclared {type(exc).__name__}" + ) from exc + if not isinstance(execution, ToolExecutionResult): + raise ToolProtocolError( + "ToolExecutor.execute() must return ToolExecutionResult" + ) + audit.append( + "tool_completed", + { + "turn": turn, + "tool_call_id": call.id, + "tool_name": call.name, + "control": execution.control.value, + "is_error": execution.is_error, + "result": execution.result, + }, + ) + return execution + + def _tool_terminal( + self, + workspace: _RunWorkspace, + usage: _UsageAccumulator, + audit: _AuditTrail, + execution: ToolExecutionResult, + ) -> RunOutcome: + if execution.control is ToolControl.COMPLETE: + status = RunStatus.COMPLETED + reason = StopReason.TOOL_COMPLETION + elif execution.control is ToolControl.REJECT: + status = RunStatus.REJECTED + reason = StopReason.TOOL_REJECTED + else: + status = RunStatus.CANCELLED + reason = StopReason.TOOL_CANCELLED + return self._terminal( + workspace, + usage, + audit, + status=status, + stop_reason=reason, + output=execution.output, + ) + + def _terminal( + self, + workspace: _RunWorkspace, + usage: _UsageAccumulator, + audit: _AuditTrail, + *, + status: RunStatus, + stop_reason: StopReason, + output: str | None, + ) -> RunOutcome: + result = RunResult( + run_id=workspace.spec.run_id, + status=status, + stop_reason=stop_reason, + turns=workspace.turn, + output=output, + usage=usage.snapshot(), + ) + audit.append( + "run_finished", + { + "status": status.value, + "stop_reason": stop_reason.value, + "turns": workspace.turn, + }, + ) + return RunOutcome( + result=result, + delta=workspace.delta, + audit_records=tuple(audit.records), + ) + + def _failed( + self, + workspace: _RunWorkspace, + usage: _UsageAccumulator, + audit: _AuditTrail, + *, + stop_reason: StopReason, + failure: RunFailure, + ) -> RunOutcome: + result = RunResult( + run_id=workspace.spec.run_id, + status=RunStatus.FAILED, + stop_reason=stop_reason, + turns=workspace.turn, + usage=usage.snapshot(), + ) + audit.append( + "run_finished", + { + "status": RunStatus.FAILED.value, + "stop_reason": stop_reason.value, + "turns": workspace.turn, + "failure_code": failure.code.value, + }, + ) + return RunOutcome( + result=result, + delta=workspace.delta, + audit_records=tuple(audit.records), + failure=failure, + ) + + @staticmethod + def _budget_failure( + spec: RunSpec, + usage: _UsageAccumulator, + ) -> tuple[StopReason, RunFailure] | None: + limit = spec.limits.max_tokens + if limit is None or usage.request_count == 0: + return None + snapshot = usage.snapshot() + if not snapshot.complete: + return ( + StopReason.USAGE_UNAVAILABLE, + RunFailure( + phase=RunPhase.CONTROL, + code=FailureCode.BUDGET_EXCEEDED, + message=( + "Run token budget cannot continue because model usage " + "was not reported" + ), + ), + ) + if snapshot.total_tokens >= limit: + return ( + StopReason.TOKEN_BUDGET_EXCEEDED, + RunFailure( + phase=RunPhase.CONTROL, + code=FailureCode.BUDGET_EXCEEDED, + message=( + f"Run used {snapshot.total_tokens} tokens and cannot start " + f"another request under the {limit}-token budget" + ), + ), + ) + return None + + @staticmethod + def _model_stop_reason(code: FailureCode) -> StopReason: + if code is FailureCode.CONTEXT_OVERFLOW: + return StopReason.CONTEXT_OVERFLOW + if code is FailureCode.COMPACTION_FAILED: + return StopReason.COMPACTION_FAILED + return StopReason.RUNTIME_ERROR + + @staticmethod + def _context_stop_reason(code: FailureCode) -> StopReason: + if code is FailureCode.CONTEXT_OVERFLOW: + return StopReason.CONTEXT_OVERFLOW + return StopReason.COMPACTION_FAILED + + @staticmethod + def _assistant_payload(message: AssistantMessage, turn: int) -> JsonObject: + calls: list[JsonValue] = [] + for call in message.tool_calls: + calls.append( + { + "id": call.id, + "name": call.name, + "arguments": call.arguments, + } + ) + return freeze_json_object( + { + "turn": turn, + "content": message.content, + "tool_calls": calls, + }, + label="assistant audit payload", + ) diff --git a/src/ejagent/middleware/__init__.py b/src/ejagent/middleware/__init__.py deleted file mode 100644 index 76d1b80..0000000 --- a/src/ejagent/middleware/__init__.py +++ /dev/null @@ -1,47 +0,0 @@ -"""Composable middleware for core agent execution.""" - -from ejagent.middleware.base import ( - Middleware, - ToolCallContext, - ToolMiddleware, - ToolNext, - compose_tool_middlewares, - format_tool_call_preview, -) -from ejagent.middleware.policy import ( - RuleBasedToolPolicy, - ToolApprovalDecision, - ToolApprovalRequest, - ToolApprover, - ToolExecutionPolicy, - ToolPolicyAction, - ToolPolicyDecision, - ToolPolicyMiddleware, - ToolPolicyPredicate, - ToolPolicyRule, -) -from ejagent.middleware.validation import ( - ToolSchemaConfigurationError, - ToolSchemaValidationMiddleware, -) - -__all__ = [ - "Middleware", - "ToolMiddleware", - "ToolCallContext", - "ToolNext", - "compose_tool_middlewares", - "format_tool_call_preview", - "RuleBasedToolPolicy", - "ToolApprovalDecision", - "ToolApprovalRequest", - "ToolApprover", - "ToolExecutionPolicy", - "ToolPolicyAction", - "ToolPolicyDecision", - "ToolPolicyMiddleware", - "ToolPolicyPredicate", - "ToolPolicyRule", - "ToolSchemaConfigurationError", - "ToolSchemaValidationMiddleware", -] diff --git a/src/ejagent/middleware/base.py b/src/ejagent/middleware/base.py deleted file mode 100644 index 4462d7b..0000000 --- a/src/ejagent/middleware/base.py +++ /dev/null @@ -1,113 +0,0 @@ -from __future__ import annotations - -import json -from collections.abc import Awaitable, Callable, Mapping, Sequence -from dataclasses import dataclass -from typing import TYPE_CHECKING, Any - -from ejagent.agent.types import StepOutcome - -if TYPE_CHECKING: - from ejagent.agent.cancellation import CancellationToken - from ejagent.agent.state import AgentState - from ejagent.agent.types import ToolProgressReporter - from ejagent.handlers.definition import ToolDefinition - - -@dataclass(frozen=True, slots=True) -class ToolCallContext: - """Metadata and cancellation signal for one tool execution.""" - - state: AgentState - tool_name: str - arguments: dict[str, Any] - tool_call_id: str | None = None - cancellation: CancellationToken | None = None - progress: ToolProgressReporter | None = None - tool_definition: ToolDefinition | None = None - - -ToolNext = Callable[[ToolCallContext], Awaitable[StepOutcome]] - - -class Middleware: - """Base class for reusable agent middleware.""" - - def __init__(self, *, name: str | None = None, enabled: bool = True) -> None: - self.name = name or type(self).__name__ - self.enabled = enabled - - async def startup(self) -> None: - """Initialize optional middleware resources.""" - - async def shutdown(self) -> None: - """Release optional middleware resources.""" - - async def on_task_start(self) -> None: - """Prepare middleware state for one new agent task.""" - - -class ToolMiddleware(Middleware): - """Decorator around one tool execution.""" - - def configure_tools( - self, - tool_definitions: Sequence[ToolDefinition], - ) -> None: - """Validate or cache the resolved tool set before middleware startup.""" - - async def __call__( - self, - context: ToolCallContext, - call_next: ToolNext, - ) -> StepOutcome: - """Invoke the next decorator or handler in the execution chain.""" - - return await call_next(context) - - -def compose_tool_middlewares( - middlewares: Sequence[ToolMiddleware], - terminal: ToolNext, -) -> ToolNext: - """Wrap a tool terminal with middleware in declaration order.""" - - call_next = terminal - for middleware in reversed(middlewares): - next_in_chain = call_next - - async def wrapped( - context: ToolCallContext, - *, - middleware: ToolMiddleware = middleware, - call_next: ToolNext = next_in_chain, - ) -> StepOutcome: - return await middleware(context, call_next) - - call_next = wrapped - return call_next - - -def format_tool_call_preview( - tool_name: str, - arguments: Mapping[str, Any], - *, - risk: str | None = None, - review: str | None = None, -) -> str: - """Build a readable approval preview from one tool call.""" - - try: - payload = json.dumps( - dict(arguments), - ensure_ascii=False, - sort_keys=True, - indent=2, - default=str, - ) - except TypeError: - payload = repr(dict(arguments)) - review_text = review or risk - if review_text: - return f"Tool: {tool_name}\nReview: {review_text}\nArguments:\n{payload}" - return f"Tool: {tool_name}\nArguments:\n{payload}" diff --git a/src/ejagent/middleware/policy.py b/src/ejagent/middleware/policy.py deleted file mode 100644 index 8e256e8..0000000 --- a/src/ejagent/middleware/policy.py +++ /dev/null @@ -1,374 +0,0 @@ -from __future__ import annotations - -import asyncio -import inspect -import logging -from abc import ABC, abstractmethod -from collections.abc import Awaitable, Callable, Iterable, Mapping -from copy import deepcopy -from dataclasses import dataclass -from enum import StrEnum -from types import MappingProxyType -from typing import Any, Protocol - -from ejagent.agent.cancellation import AgentCancelledError -from ejagent.agent.types import StepOutcome, ToolControl -from ejagent.handlers.definition import ToolEffect -from ejagent.middleware.base import ToolCallContext, ToolMiddleware, ToolNext - -logger = logging.getLogger(__name__) - -ToolPolicyPredicate = Callable[[ToolCallContext], bool | Awaitable[bool]] - - -class ToolPolicyAction(StrEnum): - """One enforceable decision returned by a tool execution policy.""" - - ALLOW = "allow" - DENY = "deny" - REQUIRE_APPROVAL = "require_approval" - - -@dataclass(frozen=True, slots=True) -class ToolPolicyDecision: - """Detached policy decision for one tool execution attempt.""" - - action: ToolPolicyAction - reason: str | None = None - rule_id: str | None = None - - def __post_init__(self) -> None: - if not isinstance(self.action, ToolPolicyAction): - raise TypeError("tool policy action must be a ToolPolicyAction") - _validate_optional_text(self.reason, "tool policy reason") - _validate_optional_text(self.rule_id, "tool policy rule_id") - - -@dataclass(frozen=True, slots=True) -class ToolApprovalRequest: - """Immutable application-facing request for one policy approval.""" - - tool_name: str - arguments: Mapping[str, Any] - tool_call_id: str | None - reason: str | None - rule_id: str | None - - def __post_init__(self) -> None: - if not self.tool_name: - raise ValueError("approval tool_name must not be empty") - _validate_optional_text(self.tool_call_id, "approval tool_call_id") - _validate_optional_text(self.reason, "approval reason") - _validate_optional_text(self.rule_id, "approval rule_id") - object.__setattr__(self, "arguments", _freeze_value(self.arguments)) - - -@dataclass(frozen=True, slots=True) -class ToolApprovalDecision: - """Application decision for one requested tool approval.""" - - approved: bool - reason: str | None = None - - def __post_init__(self) -> None: - if not isinstance(self.approved, bool): - raise TypeError("approval decision must use a boolean approved value") - _validate_optional_text(self.reason, "approval decision reason") - - -class ToolApprover(Protocol): - """Application-owned asynchronous approval boundary.""" - - async def approve(self, request: ToolApprovalRequest) -> ToolApprovalDecision: - """Approve or reject one policy-gated tool execution.""" - - -class ToolExecutionPolicy(ABC): - """Decision source consumed by :class:`ToolPolicyMiddleware`.""" - - async def startup(self) -> None: - """Initialize optional policy resources.""" - - async def shutdown(self) -> None: - """Release optional policy resources.""" - - async def on_task_start(self) -> None: - """Reset policy state for a newly starting Agent Run.""" - - @abstractmethod - async def evaluate(self, context: ToolCallContext) -> ToolPolicyDecision: - """Return the decision for one normalized tool execution attempt.""" - - -@dataclass(frozen=True, slots=True) -class ToolPolicyRule: - """One ordered exact-name/effect rule in a rule-based policy.""" - - rule_id: str - action: ToolPolicyAction - tool_names: frozenset[str] | None = None - effects: frozenset[ToolEffect] | None = None - when: ToolPolicyPredicate | None = None - max_calls_per_run: int | None = None - reason: str | None = None - - def __post_init__(self) -> None: - _validate_required_text(self.rule_id, "tool policy rule_id") - if not isinstance(self.action, ToolPolicyAction): - raise TypeError("tool policy rule action must be a ToolPolicyAction") - if self.tool_names is not None: - names = frozenset(self.tool_names) - if not names: - raise ValueError("tool policy rule tool_names must not be empty") - for name in names: - _validate_required_text(name, "tool policy rule tool name") - object.__setattr__(self, "tool_names", names) - if self.effects is not None: - effects = frozenset(self.effects) - if not effects: - raise ValueError("tool policy rule effects must not be empty") - if any(not isinstance(effect, ToolEffect) for effect in effects): - raise TypeError( - "tool policy rule effects must contain ToolEffect values" - ) - object.__setattr__(self, "effects", effects) - if self.when is not None and not callable(self.when): - raise TypeError("tool policy rule when must be callable") - if self.max_calls_per_run is not None: - if isinstance(self.max_calls_per_run, bool) or not isinstance( - self.max_calls_per_run, int - ): - raise TypeError("max_calls_per_run must be an integer") - if self.max_calls_per_run <= 0: - raise ValueError("max_calls_per_run must be greater than zero") - if self.action is ToolPolicyAction.DENY: - raise ValueError("a deny rule cannot define max_calls_per_run") - _validate_optional_text(self.reason, "tool policy rule reason") - - async def matches(self, context: ToolCallContext) -> bool: - """Return whether all configured selectors match one call.""" - - if self.tool_names is not None and context.tool_name not in self.tool_names: - return False - effect = ( - context.tool_definition.effect - if context.tool_definition is not None - else ToolEffect.SIDE_EFFECTING - ) - if self.effects is not None and effect not in self.effects: - return False - if self.when is None: - return True - matched = self.when(context) - if inspect.isawaitable(matched): - matched = await matched - if not isinstance(matched, bool): - raise TypeError("tool policy rule predicate must return a boolean") - return matched - - -class RuleBasedToolPolicy(ToolExecutionPolicy): - """Evaluate ordered rules with atomic per-Run call reservations.""" - - def __init__( - self, - rules: Iterable[ToolPolicyRule], - *, - default_action: ToolPolicyAction = ToolPolicyAction.DENY, - default_reason: str | None = None, - ) -> None: - self.rules = tuple(rules) - if any(not isinstance(rule, ToolPolicyRule) for rule in self.rules): - raise TypeError("rules must contain ToolPolicyRule values") - rule_ids = [rule.rule_id for rule in self.rules] - if len(rule_ids) != len(set(rule_ids)): - raise ValueError("tool policy rule_id values must be unique") - if not isinstance(default_action, ToolPolicyAction): - raise TypeError("default_action must be a ToolPolicyAction") - _validate_optional_text(default_reason, "default policy reason") - self.default_action = default_action - self.default_reason = default_reason - self._call_counts: dict[str, int] = {} - self._count_lock = asyncio.Lock() - - async def on_task_start(self) -> None: - async with self._count_lock: - self._call_counts.clear() - - async def evaluate(self, context: ToolCallContext) -> ToolPolicyDecision: - for rule in self.rules: - if not await rule.matches(context): - continue - if rule.max_calls_per_run is not None: - async with self._count_lock: - count = self._call_counts.get(rule.rule_id, 0) - if count >= rule.max_calls_per_run: - return ToolPolicyDecision( - ToolPolicyAction.DENY, - reason=( - f"tool call limit reached for policy rule " - f"{rule.rule_id!r}" - ), - rule_id=rule.rule_id, - ) - self._call_counts[rule.rule_id] = count + 1 - return ToolPolicyDecision( - rule.action, - reason=rule.reason, - rule_id=rule.rule_id, - ) - return ToolPolicyDecision( - self.default_action, - reason=self.default_reason, - ) - - -class ToolPolicyMiddleware(ToolMiddleware): - """Enforce one tool execution policy through the standard middleware chain.""" - - def __init__( - self, - policy: ToolExecutionPolicy, - *, - approver: ToolApprover | None = None, - name: str | None = None, - enabled: bool = True, - ) -> None: - if not isinstance(policy, ToolExecutionPolicy): - raise TypeError("policy must be a ToolExecutionPolicy") - super().__init__(name=name, enabled=enabled) - self.policy = policy - self.approver = approver - - async def startup(self) -> None: - await self.policy.startup() - - async def shutdown(self) -> None: - await self.policy.shutdown() - - async def on_task_start(self) -> None: - await self.policy.on_task_start() - - async def __call__( - self, - context: ToolCallContext, - call_next: ToolNext, - ) -> StepOutcome: - try: - decision: object = await self.policy.evaluate(context) - except AgentCancelledError: - raise - except Exception: - logger.exception("Tool policy evaluation failed closed") - return _rejected_outcome( - context, - reason="tool policy evaluation failed", - ) - - if not isinstance(decision, ToolPolicyDecision): - logger.error( - "Tool policy returned %s instead of ToolPolicyDecision", - type(decision).__name__, - ) - return _rejected_outcome( - context, - reason="tool policy returned an invalid decision", - ) - - if decision.action is ToolPolicyAction.ALLOW: - return await call_next(context) - if decision.action is ToolPolicyAction.DENY: - return _rejected_outcome( - context, - reason=decision.reason or "tool execution denied by policy", - rule_id=decision.rule_id, - ) - return await self._request_approval(context, decision, call_next) - - async def _request_approval( - self, - context: ToolCallContext, - decision: ToolPolicyDecision, - call_next: ToolNext, - ) -> StepOutcome: - if self.approver is None: - return _rejected_outcome( - context, - reason="tool execution requires approval but no approver is configured", - rule_id=decision.rule_id, - ) - request = ToolApprovalRequest( - tool_name=context.tool_name, - arguments=context.arguments, - tool_call_id=context.tool_call_id, - reason=decision.reason, - rule_id=decision.rule_id, - ) - try: - approval: object = await self.approver.approve(request) - except AgentCancelledError: - raise - except Exception: - logger.exception("Tool approval failed closed") - return _rejected_outcome( - context, - reason="tool approval failed", - rule_id=decision.rule_id, - ) - if not isinstance(approval, ToolApprovalDecision): - logger.error( - "Tool approver returned %s instead of ToolApprovalDecision", - type(approval).__name__, - ) - return _rejected_outcome( - context, - reason="tool approver returned an invalid decision", - rule_id=decision.rule_id, - ) - if approval.approved: - return await call_next(context) - return _rejected_outcome( - context, - reason=( - approval.reason - or decision.reason - or "tool execution approval was denied" - ), - rule_id=decision.rule_id, - ) - - -def _rejected_outcome( - context: ToolCallContext, - *, - reason: str, - rule_id: str | None = None, -) -> StepOutcome: - data: dict[str, Any] = { - "status": "rejected", - "tool": context.tool_name, - "reason": reason, - } - if rule_id is not None: - data["rule_id"] = rule_id - return StepOutcome(data, control=ToolControl.REJECT) - - -def _freeze_value(value: Any) -> Any: - if isinstance(value, Mapping): - return MappingProxyType( - {key: _freeze_value(item) for key, item in value.items()} - ) - if isinstance(value, list): - return tuple(_freeze_value(item) for item in value) - return deepcopy(value) - - -def _validate_required_text(value: object, field: str) -> None: - if not isinstance(value, str) or not value.strip(): - raise ValueError(f"{field} must be a non-empty string") - - -def _validate_optional_text(value: object, field: str) -> None: - if value is not None: - _validate_required_text(value, field) diff --git a/src/ejagent/middleware/validation.py b/src/ejagent/middleware/validation.py deleted file mode 100644 index 6860de6..0000000 --- a/src/ejagent/middleware/validation.py +++ /dev/null @@ -1,189 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping, Sequence -from itertools import islice -from typing import Any - -from jsonschema.exceptions import SchemaError, ValidationError -from jsonschema.protocols import Validator -from jsonschema.validators import validator_for -from referencing import Registry -from referencing.exceptions import Unresolvable - -from ejagent.agent.types import StepOutcome -from ejagent.handlers.definition import ToolDefinition -from ejagent.middleware.base import ToolCallContext, ToolMiddleware, ToolNext - - -class ToolSchemaConfigurationError(ValueError): - """Raised when a registered tool contains an invalid JSON Schema.""" - - def __init__(self, tool_name: str) -> None: - self.tool_name = tool_name - super().__init__(f"tool {tool_name!r} contains an invalid parameters schema") - - -class ToolSchemaValidationMiddleware(ToolMiddleware): - """Validate normalized tool arguments before policy and handler execution.""" - - def __init__( - self, - *, - max_errors: int = 8, - name: str | None = None, - enabled: bool = True, - ) -> None: - if isinstance(max_errors, bool) or not isinstance(max_errors, int): - raise TypeError("max_errors must be an integer") - if max_errors <= 0: - raise ValueError("max_errors must be greater than zero") - super().__init__(name=name, enabled=enabled) - self.max_errors = max_errors - self._validators: dict[str, Validator] = {} - - def configure_tools( - self, - tool_definitions: Sequence[ToolDefinition], - ) -> None: - """Compile every static schema before middleware resources start.""" - - validators: dict[str, Validator] = {} - for definition in tool_definitions: - if definition.parameters is None: - continue - validators[definition.name] = _compile_validator(definition) - self._validators = validators - - async def __call__( - self, - context: ToolCallContext, - call_next: ToolNext, - ) -> StepOutcome: - definition = context.tool_definition - if definition is None or definition.parameters is None: - return await call_next(context) - validator = self._validators.get(definition.name) - if validator is None: - validator = _compile_validator(definition) - self._validators[definition.name] = validator - - try: - errors = list( - islice( - validator.iter_errors(context.arguments), - self.max_errors + 1, - ) - ) - except Unresolvable as exc: - raise ToolSchemaConfigurationError(definition.name) from exc - if not errors: - return await call_next(context) - - truncated = len(errors) > self.max_errors - visible_errors = errors[: self.max_errors] - visible_errors.sort(key=_error_sort_key) - data: dict[str, Any] = { - "status": "error", - "tool": context.tool_name, - "code": "invalid_tool_arguments", - "errors": [_serialize_error(error) for error in visible_errors], - } - if truncated: - data["truncated"] = True - return StepOutcome(data) - - -def _compile_validator(definition: ToolDefinition) -> Validator: - assert definition.parameters is not None - schema = _mutable_json_value(definition.parameters) - validator_class = validator_for(schema) - try: - validator_class.check_schema(schema) - except SchemaError as exc: - raise ToolSchemaConfigurationError(definition.name) from exc - return validator_class(schema, registry=Registry()) - - -def _serialize_error(error: ValidationError) -> dict[str, str]: - keyword = str(error.validator or "schema") - path = list(error.absolute_path) - if keyword == "required": - missing = _first_missing_property(error) - if missing is not None: - path.append(missing) - return { - "path": _json_pointer(path), - "keyword": keyword, - "message": _safe_error_message(keyword), - } - - -def _first_missing_property(error: ValidationError) -> str | None: - required = error.validator_value - if not isinstance(required, Sequence) or isinstance(required, str): - return None - return next( - ( - name - for name in required - if isinstance(name, str) - and error.message == f"{name!r} is a required property" - ), - None, - ) - - -def _safe_error_message(keyword: str) -> str: - messages = { - "additionalProperties": "unexpected properties are not allowed", - "anyOf": "value does not match any allowed schema", - "const": "value does not match the required constant", - "contains": "array does not contain a required item", - "dependentRequired": "dependent properties are missing", - "enum": "value is not in the allowed set", - "exclusiveMaximum": "number is above the allowed exclusive maximum", - "exclusiveMinimum": "number is below the allowed exclusive minimum", - "format": "value does not match the required format", - "maxItems": "array contains too many items", - "maxLength": "string is longer than allowed", - "maxProperties": "object contains too many properties", - "maximum": "number is above the allowed maximum", - "minItems": "array contains too few items", - "minLength": "string is shorter than allowed", - "minProperties": "object contains too few properties", - "minimum": "number is below the allowed minimum", - "multipleOf": "number is not an allowed multiple", - "not": "value matches a forbidden schema", - "oneOf": "value does not match exactly one allowed schema", - "pattern": "string does not match the required pattern", - "patternProperties": "property does not match the required schema", - "prefixItems": "array item does not match the required schema", - "propertyNames": "property name is not allowed", - "required": "required property is missing", - "type": "value does not match the required type", - "uniqueItems": "array items must be unique", - } - return messages.get(keyword, "value does not match the tool schema") - - -def _error_sort_key(error: ValidationError) -> tuple[tuple[str, ...], str]: - return ( - tuple(str(part) for part in error.absolute_path), - str(error.validator or "schema"), - ) - - -def _json_pointer(path: Sequence[object]) -> str: - if not path: - return "" - return "/" + "/".join( - str(part).replace("~", "~0").replace("/", "~1") for part in path - ) - - -def _mutable_json_value(value: Any) -> Any: - if isinstance(value, Mapping): - return {key: _mutable_json_value(item) for key, item in value.items()} - if isinstance(value, tuple): - return [_mutable_json_value(item) for item in value] - return value diff --git a/src/ejagent/plugins/__init__.py b/src/ejagent/plugins/__init__.py deleted file mode 100644 index afd4a6b..0000000 --- a/src/ejagent/plugins/__init__.py +++ /dev/null @@ -1,6 +0,0 @@ -"""Plugin integrations for MCP services and local skills.""" - -from ejagent.plugins.mcp.mcp_manager import McpServerManager -from ejagent.plugins.skill.skill_manager import SkillManager - -__all__ = ["McpServerManager", "SkillManager"] diff --git a/src/ejagent/providers/__init__.py b/src/ejagent/providers/__init__.py index d0ca850..dc3f024 100644 --- a/src/ejagent/providers/__init__.py +++ b/src/ejagent/providers/__init__.py @@ -1,38 +1,13 @@ -"""Model provider adapters for the agent core.""" +"""Concrete ModelPort adapters.""" -from ejagent.providers.base import ( - AssistantMessage, - ContextOverflowError, - ModelAdapter, - ModelAuthenticationError, - ModelErrorKind, - ModelProviderError, - ModelRateLimitError, - ModelResponseCompleted, - ModelStreamEvent, - ModelTextDelta, - ModelThinkingDelta, - ModelTimeoutError, - ModelToolCall, - ModelUsage, -) -from ejagent.providers.openai import ModelConfig, OpenAIModelAdapter +from ejagent.providers.anthropic_config import AnthropicConfig +from ejagent.providers.anthropic_port import AnthropicModelPort +from ejagent.providers.config import ModelConfig +from ejagent.providers.openai_port import OpenAIModelPort __all__ = [ - "AssistantMessage", - "ModelErrorKind", - "ModelProviderError", - "ContextOverflowError", - "ModelRateLimitError", - "ModelTimeoutError", - "ModelAuthenticationError", - "ModelAdapter", - "ModelStreamEvent", - "ModelTextDelta", - "ModelThinkingDelta", - "ModelResponseCompleted", - "ModelToolCall", - "ModelUsage", + "AnthropicConfig", + "AnthropicModelPort", "ModelConfig", - "OpenAIModelAdapter", + "OpenAIModelPort", ] diff --git a/src/ejagent/providers/anthropic_config.py b/src/ejagent/providers/anthropic_config.py new file mode 100644 index 0000000..0e9303a --- /dev/null +++ b/src/ejagent/providers/anthropic_config.py @@ -0,0 +1,69 @@ +from __future__ import annotations + +import os +from dataclasses import dataclass + +from dotenv import load_dotenv + + +@dataclass(frozen=True, slots=True) +class AnthropicConfig: + """Connection and generation settings for Anthropic Messages.""" + + model: str + api_key: str + max_tokens: int = 4096 + base_url: str | None = None + timeout: float = 60.0 + temperature: float = 0.7 + + def __post_init__(self) -> None: + if not isinstance(self.model, str) or not self.model.strip(): + raise ValueError("model must not be empty") + if not isinstance(self.api_key, str) or not self.api_key.strip(): + raise ValueError("api_key must not be empty") + if self.base_url is not None and not self.base_url.strip(): + raise ValueError("base_url must not be empty") + if isinstance(self.max_tokens, bool) or not isinstance(self.max_tokens, int): + raise TypeError("max_tokens must be an integer") + if self.max_tokens <= 0: + raise ValueError("max_tokens must be greater than zero") + if isinstance(self.timeout, bool) or not isinstance(self.timeout, (int, float)): + raise TypeError("timeout must be numeric") + if self.timeout <= 0: + raise ValueError("timeout must be greater than zero") + if isinstance(self.temperature, bool) or not isinstance( + self.temperature, (int, float) + ): + raise TypeError("temperature must be numeric") + if not 0 <= self.temperature <= 1: + raise ValueError("temperature must be between zero and one") + + @classmethod + def from_env(cls) -> AnthropicConfig: + """Build a config from Anthropic-specific environment variables.""" + + load_dotenv() + model = os.getenv("ANTHROPIC_MODEL") + api_key = os.getenv("ANTHROPIC_API_KEY") + if not model or not api_key: + raise ValueError("ANTHROPIC_MODEL and ANTHROPIC_API_KEY must be defined") + + try: + max_tokens = int(os.getenv("ANTHROPIC_MAX_TOKENS", "4096")) + timeout = float(os.getenv("ANTHROPIC_TIMEOUT", "60")) + temperature = float(os.getenv("ANTHROPIC_TEMPERATURE", "0.7")) + except ValueError as exc: + raise ValueError( + "ANTHROPIC_MAX_TOKENS, ANTHROPIC_TIMEOUT, and " + "ANTHROPIC_TEMPERATURE must be numeric" + ) from exc + + return cls( + model=model, + api_key=api_key, + max_tokens=max_tokens, + base_url=os.getenv("ANTHROPIC_BASE_URL"), + timeout=timeout, + temperature=temperature, + ) diff --git a/src/ejagent/providers/anthropic_port.py b/src/ejagent/providers/anthropic_port.py new file mode 100644 index 0000000..b356da0 --- /dev/null +++ b/src/ejagent/providers/anthropic_port.py @@ -0,0 +1,452 @@ +from __future__ import annotations + +import asyncio +import importlib +import json +from collections.abc import AsyncIterator, Mapping +from contextlib import suppress +from dataclasses import dataclass, field +from typing import Any + +from ejagent.contracts.control import CancellationToken, RunCancelledError +from ejagent.contracts.json import JsonValue, thaw_json_value +from ejagent.contracts.messages import ( + AssistantMessage, + ContextMessage, + ContextSummary, + SystemMessage, + ToolCall, + ToolResultMessage, + TransientInstruction, + UserMessage, +) +from ejagent.contracts.model import ( + ModelCallError, + ModelPort, + ModelProtocolError, + ModelRequest, + ModelResponseCompleted, + ModelStreamEvent, + ModelTextDelta, + ModelThinkingDelta, + ModelUsage, +) +from ejagent.contracts.runs import FailureCode +from ejagent.contracts.tools import ToolDefinition +from ejagent.providers.anthropic_config import AnthropicConfig + +_CONTEXT_OVERFLOW_PHRASES = ( + "prompt is too long", + "request too large", + "too many tokens", + "context window", +) + + +@dataclass(slots=True) +class _StreamingToolUse: + id: str + name: str + initial_input: Any + fragments: list[str] = field(default_factory=list) + + +class AnthropicModelPort(ModelPort): + """Translate typed Core requests to Anthropic's content-block protocol.""" + + def __init__( + self, + config: AnthropicConfig, + *, + client: Any | None = None, + ) -> None: + if not isinstance(config, AnthropicConfig): + raise TypeError("config must be an AnthropicConfig") + self.config = config + self._client = client + self._owns_client = client is None + self._started = False + + async def start(self) -> None: + if self._started: + return + if self._client is None: + try: + anthropic = importlib.import_module("anthropic") + except ModuleNotFoundError as exc: + raise RuntimeError( + "Anthropic support requires `pip install ejagent-core[anthropic]`" + ) from exc + options: dict[str, Any] = { + "api_key": self.config.api_key, + "timeout": self.config.timeout, + "max_retries": 0, + } + if self.config.base_url is not None: + options["base_url"] = self.config.base_url + self._client = anthropic.AsyncAnthropic(**options) + self._started = True + + async def shutdown(self) -> None: + if not self._started: + return + if self._owns_client and self._client is not None: + await self._client.close() + self._client = None + self._started = False + + async def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + if not self._started: + raise RuntimeError("AnthropicModelPort is not started") + client = self._client + if client is None: + raise RuntimeError("Anthropic client is not initialized") + + system, messages = _request_messages(request.messages) + if not messages: + raise ModelProtocolError( + "Anthropic requests require at least one user or assistant message" + ) + options: dict[str, Any] = { + "model": self.config.model, + "max_tokens": self.config.max_tokens, + "temperature": self.config.temperature, + "messages": messages, + } + if system is not None: + options["system"] = system + if request.tools: + options["tools"] = [_tool_to_anthropic(tool) for tool in request.tools] + + manager: Any | None = None + entered = False + content_parts: list[str] = [] + tool_uses: dict[int, _StreamingToolUse] = {} + input_tokens: int | None = None + output_tokens: int | None = None + cache_read_tokens: int | None = None + cache_write_tokens: int | None = None + stop_reason: str | None = None + has_message_stop = False + try: + manager = client.messages.stream(**options) + response = await cancellation.run(manager.__aenter__()) + entered = True + iterator = response.__aiter__() + while True: + cancellation.raise_if_cancelled() + try: + event = await cancellation.run(anext(iterator)) + except StopAsyncIteration: + break + + event_type = _field(event, "type") + if event_type == "message_start": + usage = _field(_field(event, "message"), "usage") + if usage is not None: + input_tokens = _usage_integer(usage, "input_tokens", 0) + cache_read_tokens = _usage_integer( + usage, "cache_read_input_tokens", None + ) + cache_write_tokens = _usage_integer( + usage, "cache_creation_input_tokens", None + ) + initial_output = _usage_integer(usage, "output_tokens", None) + if initial_output is not None: + output_tokens = initial_output + continue + if event_type == "content_block_start": + block = _field(event, "content_block") + block_type = _field(block, "type") + if block_type == "tool_use": + index = _event_index(event) + call_id = _field(block, "id") + name = _field(block, "name") + if not isinstance(call_id, str) or not call_id: + raise ModelProtocolError( + "Anthropic stream returned a tool use without an id" + ) + if not isinstance(name, str) or not name: + raise ModelProtocolError( + "Anthropic stream returned a tool use without a name" + ) + tool_uses[index] = _StreamingToolUse( + call_id, + name, + _field(block, "input"), + ) + continue + if event_type == "content_block_delta": + delta = _field(event, "delta") + delta_type = _field(delta, "type") + if delta_type == "text_delta": + text = _field(delta, "text") + if isinstance(text, str) and text: + content_parts.append(text) + yield ModelTextDelta(text) + elif delta_type == "thinking_delta": + thinking = _field(delta, "thinking") + if isinstance(thinking, str) and thinking: + yield ModelThinkingDelta(thinking) + elif delta_type == "input_json_delta": + index = _event_index(event) + tool_use = tool_uses.get(index) + if tool_use is None: + raise ModelProtocolError( + "Anthropic tool input delta preceded its start event" + ) + fragment = _field(delta, "partial_json") + if not isinstance(fragment, str): + raise ModelProtocolError( + "Anthropic tool input delta was not a string" + ) + tool_use.fragments.append(fragment) + continue + if event_type == "message_delta": + delta = _field(event, "delta") + raw_reason = _field(delta, "stop_reason") + if raw_reason is not None: + stop_reason = str(raw_reason) + usage = _field(event, "usage") + if usage is not None: + updated_output = _usage_integer(usage, "output_tokens", None) + if updated_output is not None: + output_tokens = updated_output + continue + if event_type == "message_stop": + has_message_stop = True + continue + if event_type == "error": + error = _field(event, "error") + raise RuntimeError(str(_field(error, "message") or error)) + except (asyncio.CancelledError, RunCancelledError): + raise + except ModelProtocolError: + raise + except Exception as exc: + raise _model_call_error(exc, operation="Anthropic message stream") from exc + finally: + if entered and manager is not None: + with suppress(Exception): + await manager.__aexit__(None, None, None) + + if not has_message_stop or stop_reason is None: + raise ModelProtocolError( + "Anthropic stream ended without message_stop and a stop_reason" + ) + calls = tuple( + _completed_tool_use(tool_uses[index]) for index in sorted(tool_uses) + ) + try: + message = AssistantMessage( + content="".join(content_parts) or None, + tool_calls=calls, + ) + except (TypeError, ValueError) as exc: + raise ModelProtocolError( + f"invalid Anthropic assistant message: {exc}" + ) from exc + + usage = _normalize_usage( + input_tokens, + output_tokens, + cache_read_tokens, + cache_write_tokens, + ) + yield ModelResponseCompleted(message=message, usage=usage) + + +def _field(value: Any, name: str) -> Any: + if isinstance(value, Mapping): + return value.get(name) + return getattr(value, name, None) + + +def _event_index(event: Any) -> int: + index = _field(event, "index") + if isinstance(index, bool) or not isinstance(index, int) or index < 0: + raise ModelProtocolError("Anthropic stream returned an invalid block index") + return index + + +def _request_messages( + messages: tuple[ContextMessage, ...], +) -> tuple[str | None, list[dict[str, Any]]]: + system_parts: list[str] = [] + provider_messages: list[dict[str, Any]] = [] + for message in messages: + if isinstance(message, SystemMessage): + system_parts.append(message.content) + elif isinstance(message, ContextSummary): + system_parts.append( + f"[Derived summary: revisions {message.source_revision_start}-" + f"{message.source_revision_end}; {message.compactor_id}]\n" + f"{message.content}" + ) + elif isinstance(message, TransientInstruction): + system_parts.append( + f"[Transient instruction: {message.source}]\n{message.content}" + ) + elif isinstance(message, UserMessage): + provider_messages.append( + { + "role": "user", + "content": [{"type": "text", "text": message.content}], + } + ) + elif isinstance(message, AssistantMessage): + provider_messages.append(_assistant_to_anthropic(message)) + elif isinstance(message, ToolResultMessage): + block = _tool_result_to_anthropic(message) + if _can_append_tool_result(provider_messages): + provider_messages[-1]["content"].append(block) + else: + provider_messages.append({"role": "user", "content": [block]}) + else: + raise TypeError(f"unsupported ContextMessage {type(message).__name__}") + return "\n\n".join(system_parts) or None, provider_messages + + +def _assistant_to_anthropic(message: AssistantMessage) -> dict[str, Any]: + content: list[dict[str, Any]] = [] + if message.content: + content.append({"type": "text", "text": message.content}) + content.extend( + { + "type": "tool_use", + "id": call.id, + "name": call.name, + "input": thaw_json_value(call.arguments), + } + for call in message.tool_calls + ) + return {"role": "assistant", "content": content} + + +def _tool_result_to_anthropic(message: ToolResultMessage) -> dict[str, Any]: + result = thaw_json_value(message.result) + content = ( + result + if isinstance(result, str) + else json.dumps(result, ensure_ascii=False, separators=(",", ":")) + ) + return { + "type": "tool_result", + "tool_use_id": message.tool_call_id, + "content": content, + "is_error": message.is_error, + } + + +def _can_append_tool_result(messages: list[dict[str, Any]]) -> bool: + if not messages or messages[-1]["role"] != "user": + return False + content = messages[-1]["content"] + return bool(content) and all( + block.get("type") == "tool_result" for block in content + ) + + +def _tool_to_anthropic(definition: ToolDefinition) -> dict[str, Any]: + value: dict[str, Any] = { + "name": definition.name, + "input_schema": thaw_json_value(definition.input_schema), + } + if definition.description is not None: + value["description"] = definition.description + return value + + +def _completed_tool_use(tool_use: _StreamingToolUse) -> ToolCall: + if tool_use.fragments: + try: + arguments: JsonValue = json.loads("".join(tool_use.fragments) or "{}") + except json.JSONDecodeError as exc: + raise ModelProtocolError( + f"Anthropic tool {tool_use.name!r} returned invalid input JSON" + ) from exc + else: + arguments = tool_use.initial_input + if not isinstance(arguments, Mapping): + raise ModelProtocolError( + f"Anthropic tool {tool_use.name!r} input must be an object" + ) + try: + return ToolCall(tool_use.id, tool_use.name, arguments) + except (TypeError, ValueError) as exc: + raise ModelProtocolError(f"invalid Anthropic tool use: {exc}") from exc + + +def _usage_integer(usage: Any, name: str, default: int | None) -> int | None: + raw = _field(usage, name) + if raw is None: + return default + if isinstance(raw, bool) or not isinstance(raw, int) or raw < 0: + raise ModelProtocolError(f"Anthropic usage {name} must be non-negative") + return raw + + +def _normalize_usage( + uncached_input: int | None, + output: int | None, + cache_read: int | None, + cache_write: int | None, +) -> ModelUsage | None: + if uncached_input is None and output is None: + return None + input_tokens = (uncached_input or 0) + (cache_read or 0) + (cache_write or 0) + output_tokens = output or 0 + return ModelUsage( + input_tokens=input_tokens, + output_tokens=output_tokens, + total_tokens=input_tokens + output_tokens, + cache_read_tokens=cache_read, + cache_write_tokens=cache_write, + ) + + +def _model_call_error(exc: Exception, *, operation: str) -> ModelCallError: + message = f"{operation} failed: {exc}" + name = type(exc).__name__ + status = getattr(exc, "status_code", None) + if name == "AuthenticationError" or status in {401, 403}: + return ModelCallError(FailureCode.AUTHENTICATION, message) + if name == "RateLimitError" or status == 429: + return ModelCallError(FailureCode.RATE_LIMIT, message, retryable=True) + if ( + name == "APITimeoutError" + or isinstance(exc, TimeoutError) + or status + in { + 408, + 504, + } + ): + return ModelCallError(FailureCode.TIMEOUT, message, retryable=True) + if _is_context_overflow(exc): + return ModelCallError(FailureCode.CONTEXT_OVERFLOW, message) + return ModelCallError( + FailureCode.PROVIDER_ERROR, + message, + retryable=status in {500, 502, 503, 529}, + ) + + +def _is_context_overflow(exc: Exception) -> bool: + values = [str(exc)] + body = getattr(exc, "body", None) + if isinstance(body, Mapping): + error = body.get("error", body) + if isinstance(error, Mapping): + values.extend( + str(value) + for key in ("type", "message") + if (value := error.get(key)) is not None + ) + text = " ".join(values).lower() + return any(phrase in text for phrase in _CONTEXT_OVERFLOW_PHRASES) diff --git a/src/ejagent/providers/base.py b/src/ejagent/providers/base.py deleted file mode 100644 index 32149b9..0000000 --- a/src/ejagent/providers/base.py +++ /dev/null @@ -1,232 +0,0 @@ -from __future__ import annotations - -from abc import ABC, abstractmethod -from collections.abc import AsyncIterator -from dataclasses import dataclass -from enum import StrEnum -from typing import TYPE_CHECKING, Any, TypeAlias - -if TYPE_CHECKING: - from ejagent.agent.cancellation import CancellationToken - from ejagent.agent.context_builder import ContextBuildResult - - -class ModelErrorKind(StrEnum): - """Stable provider-neutral category for model request failures.""" - - CONTEXT_OVERFLOW = "context_overflow" - RATE_LIMIT = "rate_limit" - TIMEOUT = "timeout" - AUTHENTICATION = "authentication" - PROVIDER_ERROR = "provider_error" - - -class ModelProviderError(RuntimeError): - """Normalized model provider failure exposed to the Agent Core.""" - - kind = ModelErrorKind.PROVIDER_ERROR - - -class ContextOverflowError(ModelProviderError): - """The provider rejected a request because its context was too large.""" - - kind = ModelErrorKind.CONTEXT_OVERFLOW - - def __init__(self, message: str, *, response_started: bool = False) -> None: - self.response_started = response_started - super().__init__(message) - - -class ModelRateLimitError(ModelProviderError): - """The provider rejected a request because of rate limiting.""" - - kind = ModelErrorKind.RATE_LIMIT - - -class ModelTimeoutError(ModelProviderError): - """The provider request exceeded its configured time limit.""" - - kind = ModelErrorKind.TIMEOUT - - -class ModelAuthenticationError(ModelProviderError): - """The provider rejected the configured credentials.""" - - kind = ModelErrorKind.AUTHENTICATION - - -@dataclass(frozen=True, slots=True) -class ModelUsage: - """Provider-neutral token usage for one completed model response.""" - - input_tokens: int - output_tokens: int - total_tokens: int - cache_read_tokens: int | None = None - cache_write_tokens: int | None = None - reasoning_tokens: int | None = None - - def __post_init__(self) -> None: - values = { - "input_tokens": self.input_tokens, - "output_tokens": self.output_tokens, - "total_tokens": self.total_tokens, - "cache_read_tokens": self.cache_read_tokens, - "cache_write_tokens": self.cache_write_tokens, - "reasoning_tokens": self.reasoning_tokens, - } - for name, value in values.items(): - if value is not None and value < 0: - raise ValueError(f"{name} must not be negative") - if self.total_tokens != self.input_tokens + self.output_tokens: - raise ValueError("total_tokens must equal input_tokens + output_tokens") - if ( - self.cache_read_tokens is not None - and self.cache_read_tokens > self.input_tokens - ): - raise ValueError("cache_read_tokens must not exceed input_tokens") - if ( - self.cache_write_tokens is not None - and self.cache_write_tokens > self.input_tokens - ): - raise ValueError("cache_write_tokens must not exceed input_tokens") - if ( - self.reasoning_tokens is not None - and self.reasoning_tokens > self.output_tokens - ): - raise ValueError("reasoning_tokens must not exceed output_tokens") - - def to_dict(self) -> dict[str, int | None]: - """Return a detached JSON-compatible representation.""" - - return { - "input_tokens": self.input_tokens, - "output_tokens": self.output_tokens, - "total_tokens": self.total_tokens, - "cache_read_tokens": self.cache_read_tokens, - "cache_write_tokens": self.cache_write_tokens, - "reasoning_tokens": self.reasoning_tokens, - } - - -@dataclass(frozen=True, slots=True) -class ModelToolCall: - """Provider-neutral function call requested by an assistant message.""" - - id: str - name: str - arguments: str - - def to_agent_message(self) -> dict[str, Any]: - """Serialize the call into the current conversation message format.""" - - return { - "id": self.id, - "type": "function", - "function": { - "name": self.name, - "arguments": self.arguments, - }, - } - - -@dataclass(frozen=True, slots=True) -class AssistantMessage: - """Provider-neutral assistant response consumed by the orchestrator.""" - - content: str | None = None - tool_calls: tuple[ModelToolCall, ...] = () - - def to_agent_message(self) -> dict[str, Any]: - """Serialize the response for persistent conversation state.""" - - message: dict[str, Any] = { - "role": "assistant", - "content": self.content, - } - if self.tool_calls: - message["tool_calls"] = [ - tool_call.to_agent_message() for tool_call in self.tool_calls - ] - return message - - -def serialize_assistant_message( - message: AssistantMessage, - *, - usage: ModelUsage | None = None, -) -> dict[str, Any]: - """Attach Core metadata without widening legacy message methods.""" - - serialized = dict(message.to_agent_message()) - if usage is not None: - serialized["usage"] = usage.to_dict() - return serialized - - -@dataclass(frozen=True, slots=True) -class ModelTextDelta: - """One provider-neutral piece of assistant text.""" - - delta: str - - def __post_init__(self) -> None: - if not self.delta: - raise ValueError("model text delta must not be empty") - - -@dataclass(frozen=True, slots=True) -class ModelThinkingDelta: - """One provider-neutral piece of provisional model reasoning.""" - - delta: str - - def __post_init__(self) -> None: - if not self.delta: - raise ValueError("model thinking delta must not be empty") - - -@dataclass(frozen=True, slots=True) -class ModelResponseCompleted: - """Terminal stream event containing the normalized assistant message.""" - - message: AssistantMessage - usage: ModelUsage | None = None - - -ModelStreamEvent: TypeAlias = ( - ModelTextDelta | ModelThinkingDelta | ModelResponseCompleted -) - - -class ModelAdapter(ABC): - """Provider boundary used by the agent core.""" - - async def startup(self) -> None: - """Acquire optional provider resources.""" - - async def shutdown(self) -> None: - """Release optional provider resources.""" - - @abstractmethod - async def complete( - self, - context: ContextBuildResult, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - """Return one complete response and honor the per-run cancellation.""" - - async def stream( - self, - context: ContextBuildResult, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - """Adapt a complete-only provider into one terminal stream event.""" - - message = await self.complete( - context, - cancellation=cancellation, - ) - yield ModelResponseCompleted(message) diff --git a/src/ejagent/providers/config.py b/src/ejagent/providers/config.py new file mode 100644 index 0000000..e986808 --- /dev/null +++ b/src/ejagent/providers/config.py @@ -0,0 +1,61 @@ +from __future__ import annotations + +import os +from dataclasses import dataclass + +from dotenv import load_dotenv + + +@dataclass(frozen=True, slots=True) +class ModelConfig: + """Connection and generation settings for an OpenAI-compatible model.""" + + model: str + api_key: str + base_url: str + timeout: int = 60 + temperature: float = 0.7 + include_usage: bool = True + + def __post_init__(self) -> None: + if not self.model: + raise ValueError("model must not be empty") + if not self.api_key: + raise ValueError("api_key must not be empty") + if not self.base_url: + raise ValueError("base_url must not be empty") + if self.timeout <= 0: + raise ValueError("timeout must be greater than zero") + if not isinstance(self.include_usage, bool): + raise TypeError("include_usage must be a bool") + + @classmethod + def from_env(cls) -> ModelConfig: + """Build a config from the configured model environment variables.""" + + load_dotenv() + model = os.getenv("CHAT_MODEL") + api_key = os.getenv("MODEL_API_KEY") + base_url = os.getenv("MODEL_URL") + + if not model or not api_key or not base_url: + raise ValueError("CHAT_MODEL, MODEL_API_KEY and MODEL_URL must be defined") + + try: + timeout = int(os.getenv("LLM_TIMEOUT", "60")) + temperature = float(os.getenv("LLM_TEMPERATURE", "0.7")) + except ValueError as exc: + raise ValueError("LLM_TIMEOUT and LLM_TEMPERATURE must be numeric") from exc + + include_usage_value = os.getenv("LLM_INCLUDE_USAGE", "true").lower() + if include_usage_value not in {"true", "false"}: + raise ValueError("LLM_INCLUDE_USAGE must be true or false") + + return cls( + model=model, + api_key=api_key, + base_url=base_url, + timeout=timeout, + temperature=temperature, + include_usage=include_usage_value == "true", + ) diff --git a/src/ejagent/providers/openai.py b/src/ejagent/providers/openai.py deleted file mode 100644 index 6d92155..0000000 --- a/src/ejagent/providers/openai.py +++ /dev/null @@ -1,429 +0,0 @@ -from __future__ import annotations - -import asyncio -import inspect -import os -from collections.abc import AsyncIterator, Mapping -from contextlib import suppress -from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, cast - -from dotenv import load_dotenv -from openai import ( - APITimeoutError, - AsyncOpenAI, - AuthenticationError, - RateLimitError, -) - -from ejagent.agent.cancellation import AgentCancelledError -from ejagent.providers.base import ( - AssistantMessage, - ContextOverflowError, - ModelAdapter, - ModelAuthenticationError, - ModelProviderError, - ModelRateLimitError, - ModelResponseCompleted, - ModelStreamEvent, - ModelTextDelta, - ModelThinkingDelta, - ModelTimeoutError, - ModelToolCall, - ModelUsage, -) - -if TYPE_CHECKING: - from ejagent.agent.cancellation import CancellationToken - from ejagent.agent.context_builder import ContextBuildResult - - -@dataclass(frozen=True, slots=True) -class ModelConfig: - """Connection and generation settings for an OpenAI-compatible model.""" - - model: str - api_key: str - base_url: str - timeout: int = 60 - temperature: float = 0.7 - include_usage: bool = True - - def __post_init__(self) -> None: - if not self.model: - raise ValueError("model must not be empty") - if not self.api_key: - raise ValueError("api_key must not be empty") - if not self.base_url: - raise ValueError("base_url must not be empty") - if self.timeout <= 0: - raise ValueError("timeout must be greater than zero") - if not isinstance(self.include_usage, bool): - raise TypeError("include_usage must be a bool") - - @classmethod - def from_env(cls) -> ModelConfig: - """Build a config from the configured model environment variables.""" - - load_dotenv() - model = os.getenv("CHAT_MODEL") - api_key = os.getenv("MODEL_API_KEY") - base_url = os.getenv("MODEL_URL") - - if not model or not api_key or not base_url: - raise ValueError("CHAT_MODEL, MODEL_API_KEY and MODEL_URL must be defined") - - try: - timeout = int(os.getenv("LLM_TIMEOUT", "60")) - temperature = float(os.getenv("LLM_TEMPERATURE", "0.7")) - except ValueError as exc: - raise ValueError("LLM_TIMEOUT and LLM_TEMPERATURE must be numeric") from exc - - include_usage_value = os.getenv("LLM_INCLUDE_USAGE", "true").lower() - if include_usage_value not in {"true", "false"}: - raise ValueError("LLM_INCLUDE_USAGE must be true or false") - - return cls( - model=model, - api_key=api_key, - base_url=base_url, - timeout=timeout, - temperature=temperature, - include_usage=include_usage_value == "true", - ) - - -@dataclass(slots=True) -class _StreamingToolCall: - id: str = "" - name: str = "" - arguments: str = "" - - -def _serialize_context_tools( - context: ContextBuildResult, -) -> tuple[dict[str, Any], ...]: - """Prefer canonical definitions while preserving legacy contexts.""" - - if context.tool_definitions: - return tuple(tool.to_openai_tool() for tool in context.tool_definitions) - return context.tools - - -_CONTEXT_OVERFLOW_CODES = { - "context_length_error", - "context_length_exceeded", - "context_window_exceeded", - "input_too_long", - "prompt_too_long", -} -_CONTEXT_OVERFLOW_PHRASES = ( - "context length exceeded", - "context window exceeded", - "maximum context length", - "prompt is too long", - "too many tokens", -) - - -def _provider_error_values(exc: Exception) -> tuple[str, ...]: - values: list[str] = [] - for candidate in ( - getattr(exc, "code", None), - getattr(exc, "type", None), - getattr(exc, "body", None), - ): - if isinstance(candidate, str): - values.append(candidate) - elif isinstance(candidate, Mapping): - for key in ("code", "type", "message"): - value = candidate.get(key) - if isinstance(value, str): - values.append(value) - nested = candidate.get("error") - if isinstance(nested, Mapping): - for key in ("code", "type", "message"): - value = nested.get(key) - if isinstance(value, str): - values.append(value) - values.append(str(exc)) - return tuple(values) - - -def _is_context_overflow(exc: Exception) -> bool: - values = _provider_error_values(exc) - normalized_codes = {value.strip().lower() for value in values} - if normalized_codes & _CONTEXT_OVERFLOW_CODES: - return True - text = " ".join(normalized_codes) - return any(phrase in text for phrase in _CONTEXT_OVERFLOW_PHRASES) - - -def _normalize_provider_error( - exc: Exception, - *, - operation: str, -) -> ModelProviderError: - if isinstance(exc, ModelProviderError): - return exc - message = f"{operation} failed: {exc}" - if isinstance(exc, AuthenticationError): - return ModelAuthenticationError(message) - if isinstance(exc, RateLimitError): - return ModelRateLimitError(message) - if isinstance(exc, (APITimeoutError, TimeoutError)): - return ModelTimeoutError(message) - if _is_context_overflow(exc): - return ContextOverflowError(message) - return ModelProviderError(message) - - -def _usage_field(value: Any, name: str) -> Any: - if isinstance(value, Mapping): - return value.get(name) - return getattr(value, name, None) - - -def _normalize_usage(raw_usage: Any) -> ModelUsage: - input_tokens = int(_usage_field(raw_usage, "prompt_tokens") or 0) - output_tokens = int(_usage_field(raw_usage, "completion_tokens") or 0) - prompt_details = _usage_field(raw_usage, "prompt_tokens_details") - completion_details = _usage_field( - raw_usage, - "completion_tokens_details", - ) - cache_read = _usage_field(prompt_details, "cached_tokens") - cache_write = _usage_field(prompt_details, "cache_write_tokens") - reasoning = _usage_field(completion_details, "reasoning_tokens") - return ModelUsage( - input_tokens=input_tokens, - output_tokens=output_tokens, - total_tokens=input_tokens + output_tokens, - cache_read_tokens=int(cache_read) if cache_read is not None else None, - cache_write_tokens=(int(cache_write) if cache_write is not None else None), - reasoning_tokens=int(reasoning) if reasoning is not None else None, - ) - - -class OpenAIModelAdapter(ModelAdapter): - """OpenAI-compatible provider adapter for the core model contract.""" - - def __init__( - self, - config: ModelConfig, - *, - client: AsyncOpenAI | None = None, - ) -> None: - self.config = config - self._client = client - self._owns_client = client is None - - async def startup(self) -> None: - if self._client is None: - self._client = AsyncOpenAI( - api_key=self.config.api_key, - base_url=self.config.base_url, - timeout=self.config.timeout, - ) - - async def shutdown(self) -> None: - if self._owns_client and self._client is not None: - await self._client.close() - self._client = None - - async def complete( - self, - context: ContextBuildResult, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - await self.startup() - client = self._client - if client is None: - raise RuntimeError("OpenAI model client is not initialized") - - tools = _serialize_context_tools(context) - try: - request = client.chat.completions.create( - model=self.config.model, - messages=cast(Any, context.llm_messages), - temperature=self.config.temperature, - tools=cast(Any, tools or None), - ) - response = ( - await cancellation.run(request) - if cancellation is not None - else await request - ) - except asyncio.CancelledError: - raise - except AgentCancelledError: - raise - except Exception as exc: - raise _normalize_provider_error( - exc, - operation="chat completion", - ) from exc - - if not response.choices: - raise ModelProviderError("chat completion returned no choices") - message = response.choices[0].message - return AssistantMessage( - content=message.content, - tool_calls=tuple( - ModelToolCall( - id=tool_call.id, - name=tool_call.function.name, - arguments=tool_call.function.arguments, - ) - for tool_call in message.tool_calls or () - if tool_call.type == "function" - ), - ) - - async def stream( - self, - context: ContextBuildResult, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - """Stream and normalize one OpenAI-compatible chat completion.""" - - await self.startup() - client = self._client - if client is None: - raise RuntimeError("OpenAI model client is not initialized") - - tools = _serialize_context_tools(context) - response: Any | None = None - content_parts: list[str] = [] - tool_calls: dict[int, _StreamingToolCall] = {} - has_finish_reason = False - usage: ModelUsage | None = None - try: - request_options: dict[str, Any] = { - "model": self.config.model, - "messages": cast(Any, context.llm_messages), - "temperature": self.config.temperature, - "tools": cast(Any, tools) or None, - "stream": True, - } - if self.config.include_usage: - request_options["stream_options"] = {"include_usage": True} - request = client.chat.completions.create(**request_options) - response = ( - await cancellation.run(request) - if cancellation is not None - else await request - ) - assert response is not None - iterator = response.__aiter__() - while True: - if cancellation is not None: - cancellation.raise_if_cancelled() - try: - next_chunk = anext(iterator) - chunk = ( - await cancellation.run(next_chunk) - if cancellation is not None - else await next_chunk - ) - except StopAsyncIteration: - break - - raw_usage = getattr(chunk, "usage", None) - if raw_usage is not None: - usage = _normalize_usage(raw_usage) - - choices = getattr(chunk, "choices", None) or () - if not choices: - continue - choice = choices[0] - if getattr(choice, "finish_reason", None) is not None: - has_finish_reason = True - delta = getattr(choice, "delta", None) - if delta is None: - continue - - content = getattr(delta, "content", None) - if content: - text = str(content) - content_parts.append(text) - yield ModelTextDelta(text) - - for field in ( - "reasoning_content", - "reasoning", - "reasoning_text", - ): - reasoning = getattr(delta, field, None) - if isinstance(reasoning, str) and reasoning: - yield ModelThinkingDelta(reasoning) - break - - for position, partial in enumerate( - getattr(delta, "tool_calls", None) or () - ): - index = getattr(partial, "index", None) - if index is None: - index = position - call = tool_calls.setdefault(index, _StreamingToolCall()) - call_id = getattr(partial, "id", None) - if call_id: - call.id = str(call_id) - function = getattr(partial, "function", None) - if function is None: - continue - name = getattr(function, "name", None) - if name: - call.name += str(name) - arguments = getattr(function, "arguments", None) - if arguments: - call.arguments += str(arguments) - except asyncio.CancelledError: - raise - except AgentCancelledError: - raise - except Exception as exc: - raise _normalize_provider_error( - exc, - operation="chat completion stream", - ) from exc - finally: - if response is not None: - close = getattr(response, "close", None) - if close is None: - close = getattr(response, "aclose", None) - if close is not None: - with suppress(Exception): - close_result = close() - if inspect.isawaitable(close_result): - await close_result - - if not has_finish_reason: - raise ModelProviderError( - "chat completion stream ended without finish_reason" - ) - - normalized_tool_calls: list[ModelToolCall] = [] - for index in sorted(tool_calls): - call = tool_calls[index] - if not call.id or not call.name: - raise ModelProviderError( - "chat completion stream returned an incomplete tool call" - ) - normalized_tool_calls.append( - ModelToolCall( - id=call.id, - name=call.name, - arguments=call.arguments, - ) - ) - - yield ModelResponseCompleted( - AssistantMessage( - content="".join(content_parts) or None, - tool_calls=tuple(normalized_tool_calls), - ), - usage=usage, - ) diff --git a/src/ejagent/providers/openai_port.py b/src/ejagent/providers/openai_port.py new file mode 100644 index 0000000..c717f0f --- /dev/null +++ b/src/ejagent/providers/openai_port.py @@ -0,0 +1,382 @@ +from __future__ import annotations + +import asyncio +import inspect +import json +from collections.abc import AsyncIterator, Mapping +from contextlib import suppress +from dataclasses import dataclass +from typing import Any, cast + +from openai import ( + APITimeoutError, + AsyncOpenAI, + AuthenticationError, + RateLimitError, +) + +from ejagent.contracts.control import CancellationToken, RunCancelledError +from ejagent.contracts.json import thaw_json_value +from ejagent.contracts.messages import ( + AssistantMessage, + ContextMessage, + ContextSummary, + SystemMessage, + ToolCall, + ToolResultMessage, + TransientInstruction, + UserMessage, +) +from ejagent.contracts.model import ( + ModelCallError, + ModelPort, + ModelProtocolError, + ModelRequest, + ModelResponseCompleted, + ModelStreamEvent, + ModelTextDelta, + ModelThinkingDelta, + ModelUsage, +) +from ejagent.contracts.runs import FailureCode +from ejagent.contracts.tools import ToolDefinition +from ejagent.providers.config import ModelConfig + +_CONTEXT_OVERFLOW_CODES = { + "context_length_error", + "context_length_exceeded", + "context_window_exceeded", + "input_too_long", + "prompt_too_long", +} +_CONTEXT_OVERFLOW_PHRASES = ( + "context length exceeded", + "context window exceeded", + "maximum context length", + "prompt is too long", + "too many tokens", +) + + +@dataclass(slots=True) +class _StreamingToolCall: + id: str = "" + name: str = "" + arguments: str = "" + + +class OpenAIModelPort(ModelPort): + """Translate typed Core requests to OpenAI-compatible chat completions.""" + + def __init__( + self, + config: ModelConfig, + *, + client: AsyncOpenAI | None = None, + ) -> None: + if not isinstance(config, ModelConfig): + raise TypeError("config must be a ModelConfig") + self.config = config + self._client = client + self._owns_client = client is None + self._started = False + + async def start(self) -> None: + if self._started: + return + if self._client is None: + self._client = AsyncOpenAI( + api_key=self.config.api_key, + base_url=self.config.base_url, + timeout=self.config.timeout, + max_retries=0, + ) + self._started = True + + async def shutdown(self) -> None: + if not self._started: + return + if self._owns_client and self._client is not None: + await self._client.close() + self._client = None + self._started = False + + async def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + if not self._started: + raise RuntimeError("OpenAIModelPort is not started") + client = self._client + if client is None: + raise RuntimeError("OpenAI client is not initialized") + + options: dict[str, Any] = { + "model": self.config.model, + "messages": cast( + Any, [_message_to_openai(item) for item in request.messages] + ), + "temperature": self.config.temperature, + "tools": cast(Any, [_tool_to_openai(item) for item in request.tools]) + or None, + "stream": True, + } + if self.config.include_usage: + options["stream_options"] = {"include_usage": True} + + response: Any | None = None + content_parts: list[str] = [] + tool_calls: dict[int, _StreamingToolCall] = {} + usage: ModelUsage | None = None + has_finish_reason = False + try: + response = await cancellation.run(client.chat.completions.create(**options)) + if response is None: + raise ModelProtocolError("OpenAI stream request returned no response") + iterator = response.__aiter__() + while True: + cancellation.raise_if_cancelled() + try: + chunk = await cancellation.run(anext(iterator)) + except StopAsyncIteration: + break + + raw_usage = getattr(chunk, "usage", None) + if raw_usage is not None: + try: + usage = _normalize_usage(raw_usage) + except (TypeError, ValueError) as exc: + raise ModelProtocolError( + f"OpenAI stream returned invalid usage: {exc}" + ) from exc + choices = getattr(chunk, "choices", None) or () + if not choices: + continue + choice = choices[0] + if getattr(choice, "finish_reason", None) is not None: + has_finish_reason = True + delta = getattr(choice, "delta", None) + if delta is None: + continue + + content = getattr(delta, "content", None) + if content: + text = str(content) + content_parts.append(text) + yield ModelTextDelta(text) + for field in ( + "reasoning_content", + "reasoning", + "reasoning_text", + ): + reasoning = getattr(delta, field, None) + if isinstance(reasoning, str) and reasoning: + yield ModelThinkingDelta(reasoning) + break + for position, partial in enumerate( + getattr(delta, "tool_calls", None) or () + ): + index = getattr(partial, "index", None) + if index is None: + index = position + if ( + isinstance(index, bool) + or not isinstance(index, int) + or index < 0 + ): + raise ModelProtocolError( + "OpenAI stream returned an invalid tool call index" + ) + call = tool_calls.setdefault(index, _StreamingToolCall()) + call_id = getattr(partial, "id", None) + if call_id: + call.id = str(call_id) + function = getattr(partial, "function", None) + if function is None: + continue + name = getattr(function, "name", None) + if name: + call.name += str(name) + arguments = getattr(function, "arguments", None) + if arguments: + call.arguments += str(arguments) + except (asyncio.CancelledError, RunCancelledError): + raise + except ModelProtocolError: + raise + except Exception as exc: + raise _model_call_error(exc, operation="chat completion stream") from exc + finally: + if response is not None: + close = getattr(response, "close", None) or getattr( + response, + "aclose", + None, + ) + if close is not None: + with suppress(Exception): + result = close() + if inspect.isawaitable(result): + await result + + if not has_finish_reason: + raise ModelProtocolError("OpenAI stream ended without a finish_reason") + normalized_calls = tuple( + _completed_tool_call(tool_calls[index]) for index in sorted(tool_calls) + ) + try: + message = AssistantMessage( + content="".join(content_parts) or None, + tool_calls=normalized_calls, + ) + except (TypeError, ValueError) as exc: + raise ModelProtocolError( + f"invalid OpenAI assistant message: {exc}" + ) from exc + yield ModelResponseCompleted(message=message, usage=usage) + + +def _message_to_openai(message: ContextMessage) -> dict[str, Any]: + if isinstance(message, SystemMessage): + return {"role": "system", "content": message.content} + if isinstance(message, UserMessage): + return {"role": "user", "content": message.content} + if isinstance(message, AssistantMessage): + value: dict[str, Any] = { + "role": "assistant", + "content": message.content, + } + if message.tool_calls: + value["tool_calls"] = [ + { + "id": call.id, + "type": "function", + "function": { + "name": call.name, + "arguments": json.dumps( + thaw_json_value(call.arguments), + ensure_ascii=False, + separators=(",", ":"), + ), + }, + } + for call in message.tool_calls + ] + return value + if isinstance(message, ToolResultMessage): + result = thaw_json_value(message.result) + content = ( + result + if isinstance(result, str) + else json.dumps(result, ensure_ascii=False, separators=(",", ":")) + ) + return { + "role": "tool", + "tool_call_id": message.tool_call_id, + "content": content, + } + if isinstance(message, ContextSummary): + return { + "role": "system", + "content": ( + f"[Derived summary: revisions {message.source_revision_start}-" + f"{message.source_revision_end}; {message.compactor_id}]\n" + f"{message.content}" + ), + } + if isinstance(message, TransientInstruction): + return { + "role": "system", + "content": f"[Transient instruction: {message.source}]\n{message.content}", + } + raise TypeError(f"unsupported ContextMessage {type(message).__name__}") + + +def _tool_to_openai(definition: ToolDefinition) -> dict[str, Any]: + function: dict[str, Any] = { + "name": definition.name, + "parameters": thaw_json_value(definition.input_schema), + } + if definition.description is not None: + function["description"] = definition.description + return {"type": "function", "function": function} + + +def _completed_tool_call(call: _StreamingToolCall) -> ToolCall: + if not call.id or not call.name: + raise ModelProtocolError("OpenAI stream returned an incomplete tool call") + try: + arguments = json.loads(call.arguments or "{}") + except json.JSONDecodeError as exc: + raise ModelProtocolError( + f"OpenAI tool {call.name!r} returned invalid argument JSON" + ) from exc + if not isinstance(arguments, Mapping): + raise ModelProtocolError( + f"OpenAI tool {call.name!r} arguments must decode to an object" + ) + try: + return ToolCall(id=call.id, name=call.name, arguments=arguments) + except (TypeError, ValueError) as exc: + raise ModelProtocolError(f"invalid OpenAI tool call: {exc}") from exc + + +def _usage_field(value: Any, name: str) -> Any: + if isinstance(value, Mapping): + return value.get(name) + return getattr(value, name, None) + + +def _normalize_usage(raw_usage: Any) -> ModelUsage: + input_tokens = int(_usage_field(raw_usage, "prompt_tokens") or 0) + output_tokens = int(_usage_field(raw_usage, "completion_tokens") or 0) + prompt_details = _usage_field(raw_usage, "prompt_tokens_details") + completion_details = _usage_field(raw_usage, "completion_tokens_details") + cache_read = _usage_field(prompt_details, "cached_tokens") + cache_write = _usage_field(prompt_details, "cache_write_tokens") + reasoning = _usage_field(completion_details, "reasoning_tokens") + return ModelUsage( + input_tokens=input_tokens, + output_tokens=output_tokens, + total_tokens=input_tokens + output_tokens, + cache_read_tokens=int(cache_read) if cache_read is not None else None, + cache_write_tokens=int(cache_write) if cache_write is not None else None, + reasoning_tokens=int(reasoning) if reasoning is not None else None, + ) + + +def _model_call_error(exc: Exception, *, operation: str) -> ModelCallError: + message = f"{operation} failed: {exc}" + if isinstance(exc, AuthenticationError): + return ModelCallError(FailureCode.AUTHENTICATION, message) + if isinstance(exc, RateLimitError): + return ModelCallError(FailureCode.RATE_LIMIT, message, retryable=True) + if isinstance(exc, (APITimeoutError, TimeoutError)): + return ModelCallError(FailureCode.TIMEOUT, message, retryable=True) + if _is_context_overflow(exc): + return ModelCallError(FailureCode.CONTEXT_OVERFLOW, message) + return ModelCallError(FailureCode.PROVIDER_ERROR, message) + + +def _is_context_overflow(exc: Exception) -> bool: + values: list[str] = [] + for candidate in ( + getattr(exc, "code", None), + getattr(exc, "type", None), + getattr(exc, "body", None), + ): + if isinstance(candidate, str): + values.append(candidate) + elif isinstance(candidate, Mapping): + for key in ("code", "type", "message"): + value = candidate.get(key) + if isinstance(value, str): + values.append(value) + values.append(str(exc)) + normalized = {value.strip().lower() for value in values} + if normalized & _CONTEXT_OVERFLOW_CODES: + return True + text = " ".join(normalized) + return any(phrase in text for phrase in _CONTEXT_OVERFLOW_PHRASES) diff --git a/src/ejagent/session/__init__.py b/src/ejagent/session/__init__.py deleted file mode 100644 index d6ba0dc..0000000 --- a/src/ejagent/session/__init__.py +++ /dev/null @@ -1,73 +0,0 @@ -"""Durable Agent Session projections built from lifecycle event trees.""" - -from ejagent.session.codec import ( - SESSION_SCHEMA_VERSION, - session_from_dict, - session_to_dict, -) -from ejagent.session.errors import ( - SessionConflictError, - SessionError, - SessionLockTimeoutError, - SessionSerializationError, - SessionStorageError, -) -from ejagent.session.journal import ( - DEFAULT_SESSION_BRANCH, - SESSION_JOURNAL_SCHEMA_VERSION, - SessionRecord, - SessionRecordDraft, - SessionRecordKind, -) -from ejagent.session.jsonl import JsonlSessionStorage -from ejagent.session.memory import MemorySessionStorage -from ejagent.session.recorder import SessionRecorder -from ejagent.session.storage import ( - SessionJournalStorage, - SessionStorage, - SessionTreeStorage, -) -from ejagent.session.tree import ( - SessionBranch, - SessionBranchIntent, - SessionCheckout, - SessionRetry, -) -from ejagent.session.types import ( - AgentSession, - SessionCompaction, - SessionMessage, - SessionRun, - SessionRunIntent, -) - -__all__ = [ - "AgentSession", - "SessionMessage", - "SessionRun", - "SessionRunIntent", - "SessionCompaction", - "SessionStorage", - "SessionJournalStorage", - "SessionTreeStorage", - "JsonlSessionStorage", - "MemorySessionStorage", - "SessionRecorder", - "SESSION_SCHEMA_VERSION", - "session_to_dict", - "session_from_dict", - "SessionError", - "SessionConflictError", - "SessionLockTimeoutError", - "SessionSerializationError", - "SessionStorageError", - "SESSION_JOURNAL_SCHEMA_VERSION", - "DEFAULT_SESSION_BRANCH", - "SessionRecordKind", - "SessionRecordDraft", - "SessionRecord", - "SessionBranchIntent", - "SessionBranch", - "SessionCheckout", - "SessionRetry", -] diff --git a/src/ejagent/session/codec.py b/src/ejagent/session/codec.py deleted file mode 100644 index 9a6c0d2..0000000 --- a/src/ejagent/session/codec.py +++ /dev/null @@ -1,278 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from copy import deepcopy -from typing import Any - -from ejagent.agent.compaction import SummaryEntry -from ejagent.agent.result import AgentRunResult, RunStatus, StopReason -from ejagent.agent.usage import RunUsage -from ejagent.session.errors import SessionSerializationError -from ejagent.session.types import ( - AgentSession, - SessionCompaction, - SessionMessage, - SessionRun, - SessionRunIntent, -) - -SESSION_SCHEMA_VERSION = 1 - - -def session_to_dict(session: AgentSession) -> dict[str, Any]: - """Encode one detached Session into the stable versioned data shape.""" - - snapshot = session.snapshot() - return { - "schema_version": SESSION_SCHEMA_VERSION, - "session": { - "session_id": snapshot.session_id, - "agent_id": snapshot.agent_id, - "entries": [ - { - "run_id": entry.run_id, - "sequence": entry.sequence, - "message": deepcopy(entry.message), - } - for entry in snapshot.entries - ], - "runs": [_run_to_dict(run) for run in snapshot.runs], - "compactions": [ - { - "operation_id": compaction.operation_id, - "sequence": compaction.sequence, - "summary": compaction.summary.to_dict(), - "messages": deepcopy(list(compaction.messages)), - "covered_entry_count": compaction.covered_entry_count, - } - for compaction in snapshot.compactions - ], - }, - } - - -def session_from_dict(value: Mapping[str, Any]) -> AgentSession: - """Decode and validate one versioned Session data structure.""" - - try: - root = _mapping(value, "session document") - version = _integer(root.get("schema_version"), "schema_version") - if version != SESSION_SCHEMA_VERSION: - raise SessionSerializationError( - f"unsupported session schema_version {version}; " - f"expected {SESSION_SCHEMA_VERSION}" - ) - payload = _mapping(root.get("session"), "session") - session_id = _string(payload.get("session_id"), "session.session_id") - agent_id = _optional_string( - payload.get("agent_id"), - "session.agent_id", - ) - entries = [ - _entry_from_dict(item, index) - for index, item in enumerate( - _list(payload.get("entries"), "session.entries") - ) - ] - runs = [ - _run_from_dict(item, index) - for index, item in enumerate(_list(payload.get("runs"), "session.runs")) - ] - compactions = [ - _compaction_from_dict(item, index) - for index, item in enumerate( - _list(payload.get("compactions"), "session.compactions") - ) - ] - return AgentSession( - session_id=session_id, - agent_id=agent_id, - entries=entries, - runs=runs, - compactions=compactions, - ) - except SessionSerializationError: - raise - except (KeyError, TypeError, ValueError) as exc: - raise SessionSerializationError(f"invalid session payload: {exc}") from exc - - -def _run_to_dict(run: SessionRun) -> dict[str, Any]: - return { - "run_id": run.run_id, - "task": run.task, - "intent": run.intent.value, - "start_sequence": run.start_sequence, - "finish_sequence": run.finish_sequence, - "result": ( - agent_run_result_to_dict(run.result) if run.result is not None else None - ), - } - - -def agent_run_result_to_dict(result: AgentRunResult) -> dict[str, Any]: - """Encode a structured Run result for Session journal records.""" - - return { - "status": result.status.value, - "stop_reason": result.stop_reason.value, - "turns": result.turns, - "output": result.output, - "error": result.error, - "usage": result.usage.to_dict(), - } - - -def _entry_from_dict(value: Any, index: int) -> SessionMessage: - label = f"session.entries[{index}]" - item = _mapping(value, label) - return SessionMessage( - run_id=_string(item.get("run_id"), f"{label}.run_id"), - sequence=_integer(item.get("sequence"), f"{label}.sequence"), - message=_message(item.get("message"), f"{label}.message"), - ) - - -def _run_from_dict(value: Any, index: int) -> SessionRun: - label = f"session.runs[{index}]" - item = _mapping(value, label) - raw_result = item.get("result") - result = ( - None - if raw_result is None - else agent_run_result_from_dict(raw_result, label=label) - ) - finish_sequence = _optional_integer( - item.get("finish_sequence"), - f"{label}.finish_sequence", - ) - return SessionRun( - run_id=_string(item.get("run_id"), f"{label}.run_id"), - task=_optional_string(item.get("task"), f"{label}.task"), - start_sequence=_integer( - item.get("start_sequence"), - f"{label}.start_sequence", - ), - intent=SessionRunIntent(_string(item.get("intent", "task"), f"{label}.intent")), - finish_sequence=finish_sequence, - result=result, - ) - - -def agent_run_result_from_dict( - value: Any, - *, - label: str = "result", -) -> AgentRunResult: - """Decode a structured Run result from a Session journal record.""" - - label = f"{label}.result" if not label.endswith(".result") else label - item = _mapping(value, label) - return AgentRunResult( - status=RunStatus(_string(item.get("status"), f"{label}.status")), - stop_reason=StopReason( - _string(item.get("stop_reason"), f"{label}.stop_reason") - ), - turns=_integer(item.get("turns"), f"{label}.turns"), - output=_optional_string(item.get("output"), f"{label}.output"), - error=_optional_string(item.get("error"), f"{label}.error"), - usage=_usage_from_dict(item.get("usage"), label), - ) - - -def _usage_from_dict(value: Any, parent_label: str) -> RunUsage: - label = f"{parent_label}.usage" - item = _mapping(value, label) - return RunUsage( - input_tokens=_integer(item.get("input_tokens"), f"{label}.input_tokens"), - output_tokens=_integer( - item.get("output_tokens"), - f"{label}.output_tokens", - ), - total_tokens=_integer(item.get("total_tokens"), f"{label}.total_tokens"), - request_count=_integer( - item.get("request_count"), - f"{label}.request_count", - ), - reported_request_count=_integer( - item.get("reported_request_count"), - f"{label}.reported_request_count", - ), - cache_read_tokens=_optional_integer( - item.get("cache_read_tokens"), - f"{label}.cache_read_tokens", - ), - cache_write_tokens=_optional_integer( - item.get("cache_write_tokens"), - f"{label}.cache_write_tokens", - ), - reasoning_tokens=_optional_integer( - item.get("reasoning_tokens"), - f"{label}.reasoning_tokens", - ), - ) - - -def _compaction_from_dict(value: Any, index: int) -> SessionCompaction: - label = f"session.compactions[{index}]" - item = _mapping(value, label) - raw_messages = _list(item.get("messages"), f"{label}.messages") - summary_value = _mapping(item.get("summary"), f"{label}.summary") - return SessionCompaction( - operation_id=_string( - item.get("operation_id"), - f"{label}.operation_id", - ), - sequence=_integer(item.get("sequence"), f"{label}.sequence"), - summary=SummaryEntry.from_dict(summary_value), - messages=tuple( - _message(message, f"{label}.messages[{message_index}]") - for message_index, message in enumerate(raw_messages) - ), - covered_entry_count=_integer( - item.get("covered_entry_count"), - f"{label}.covered_entry_count", - ), - ) - - -def _mapping(value: Any, label: str) -> Mapping[str, Any]: - if not isinstance(value, Mapping): - raise SessionSerializationError(f"{label} must be an object") - if any(not isinstance(key, str) for key in value): - raise SessionSerializationError(f"{label} keys must be strings") - return value - - -def _list(value: Any, label: str) -> list[Any]: - if not isinstance(value, list): - raise SessionSerializationError(f"{label} must be an array") - return value - - -def _message(value: Any, label: str) -> dict[str, Any]: - return deepcopy(dict(_mapping(value, label))) - - -def _string(value: Any, label: str) -> str: - if not isinstance(value, str): - raise SessionSerializationError(f"{label} must be a string") - return value - - -def _optional_string(value: Any, label: str) -> str | None: - if value is None: - return None - return _string(value, label) - - -def _integer(value: Any, label: str) -> int: - if not isinstance(value, int) or isinstance(value, bool): - raise SessionSerializationError(f"{label} must be an integer") - return value - - -def _optional_integer(value: Any, label: str) -> int | None: - if value is None: - return None - return _integer(value, label) diff --git a/src/ejagent/session/errors.py b/src/ejagent/session/errors.py deleted file mode 100644 index bd74aa8..0000000 --- a/src/ejagent/session/errors.py +++ /dev/null @@ -1,18 +0,0 @@ -class SessionError(RuntimeError): - """Base error for durable Session encoding and storage.""" - - -class SessionSerializationError(SessionError): - """A Session payload is invalid, unsupported, or not JSON-compatible.""" - - -class SessionStorageError(SessionError): - """A Session could not be read or atomically persisted.""" - - -class SessionConflictError(SessionStorageError): - """A branch head changed before a conditional append could commit.""" - - -class SessionLockTimeoutError(SessionStorageError): - """A Session journal lock could not be acquired before its deadline.""" diff --git a/src/ejagent/session/journal.py b/src/ejagent/session/journal.py deleted file mode 100644 index 676d465..0000000 --- a/src/ejagent/session/journal.py +++ /dev/null @@ -1,504 +0,0 @@ -from __future__ import annotations - -from collections.abc import Mapping -from copy import deepcopy -from dataclasses import dataclass -from enum import StrEnum -from typing import Any - -from ejagent.agent.compaction import CompactionResult, SummaryEntry -from ejagent.agent.result import AgentRunResult -from ejagent.agent.types import AgentMessage -from ejagent.session.codec import ( - agent_run_result_from_dict, - agent_run_result_to_dict, - session_from_dict, - session_to_dict, -) -from ejagent.session.errors import SessionSerializationError -from ejagent.session.types import AgentSession - -SESSION_JOURNAL_SCHEMA_VERSION = 1 -DEFAULT_SESSION_BRANCH = "main" - - -class SessionRecordKind(StrEnum): - """Mutation represented by one immutable Session journal record.""" - - CHECKPOINT = "checkpoint" - BRANCH_CREATED = "branch_created" - RUN_STARTED = "run_started" - RUN_CONTINUED = "run_continued" - MESSAGE_APPENDED = "message_appended" - MESSAGES_APPENDED = "messages_appended" - STEERING_APPLIED = "steering_applied" - COMPACTION_APPLIED = "compaction_applied" - RUN_FINISHED = "run_finished" - - -@dataclass(frozen=True, slots=True) -class SessionRecordDraft: - """Semantic Session mutation before journal identity is assigned.""" - - session_id: str - agent_id: str | None - sequence: int - kind: SessionRecordKind - data: dict[str, Any] - branch_id: str = DEFAULT_SESSION_BRANCH - - def __post_init__(self) -> None: - _validate_envelope_fields( - session_id=self.session_id, - agent_id=self.agent_id, - sequence=self.sequence, - branch_id=self.branch_id, - ) - object.__setattr__(self, "data", deepcopy(self.data)) - - @classmethod - def checkpoint(cls, session: AgentSession) -> SessionRecordDraft: - return cls( - session_id=session.session_id, - agent_id=session.agent_id, - sequence=0, - kind=SessionRecordKind.CHECKPOINT, - data={"document": session_to_dict(session)}, - ) - - @classmethod - def branch_created( - cls, - *, - session_id: str, - agent_id: str | None, - branch_id: str, - base_record_id: str | None, - source_branch_id: str, - source_head_id: str, - intent: str, - retried_run_id: str | None = None, - ) -> SessionRecordDraft: - data: dict[str, Any] = { - "base_record_id": base_record_id, - "source_branch_id": source_branch_id, - "source_head_id": source_head_id, - "intent": intent, - } - if retried_run_id is not None: - data["retried_run_id"] = retried_run_id - return cls( - session_id=session_id, - agent_id=agent_id, - sequence=0, - kind=SessionRecordKind.BRANCH_CREATED, - data=data, - branch_id=branch_id, - ) - - @classmethod - def run_started( - cls, - *, - session_id: str, - agent_id: str, - sequence: int, - run_id: str, - task: str, - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> SessionRecordDraft: - return cls( - session_id=session_id, - agent_id=agent_id, - sequence=sequence, - kind=SessionRecordKind.RUN_STARTED, - data={"run_id": run_id, "task": task}, - branch_id=branch_id, - ) - - @classmethod - def run_continued( - cls, - *, - session_id: str, - agent_id: str, - sequence: int, - run_id: str, - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> SessionRecordDraft: - return cls( - session_id=session_id, - agent_id=agent_id, - sequence=sequence, - kind=SessionRecordKind.RUN_CONTINUED, - data={"run_id": run_id}, - branch_id=branch_id, - ) - - @classmethod - def message_appended( - cls, - *, - session_id: str, - agent_id: str, - sequence: int, - run_id: str, - message: AgentMessage, - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> SessionRecordDraft: - return cls( - session_id=session_id, - agent_id=agent_id, - sequence=sequence, - kind=SessionRecordKind.MESSAGE_APPENDED, - data={"run_id": run_id, "message": deepcopy(message)}, - branch_id=branch_id, - ) - - @classmethod - def messages_appended( - cls, - *, - session_id: str, - agent_id: str, - sequence: int, - run_id: str, - messages: tuple[AgentMessage, ...], - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> SessionRecordDraft: - return cls( - session_id=session_id, - agent_id=agent_id, - sequence=sequence, - kind=SessionRecordKind.MESSAGES_APPENDED, - data={"run_id": run_id, "messages": deepcopy(list(messages))}, - branch_id=branch_id, - ) - - @classmethod - def steering_applied( - cls, - *, - session_id: str, - agent_id: str, - sequence: int, - run_id: str, - input_id: str, - content: str, - target_turn: int, - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> SessionRecordDraft: - if target_turn <= 0: - raise ValueError("target_turn must be greater than zero") - return cls( - session_id=session_id, - agent_id=agent_id, - sequence=sequence, - kind=SessionRecordKind.STEERING_APPLIED, - data={ - "run_id": run_id, - "input_id": input_id, - "content": content, - "target_turn": target_turn, - }, - branch_id=branch_id, - ) - - @classmethod - def compaction_applied( - cls, - *, - session_id: str, - agent_id: str, - sequence: int, - result: CompactionResult, - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> SessionRecordDraft: - if result.summary is None: - raise ValueError("completed compaction requires a SummaryEntry") - return cls( - session_id=session_id, - agent_id=agent_id, - sequence=sequence, - kind=SessionRecordKind.COMPACTION_APPLIED, - data={ - "operation_id": result.operation_id, - "summary": result.summary.to_dict(), - "messages": deepcopy(list(result.messages)), - }, - branch_id=branch_id, - ) - - @classmethod - def run_finished( - cls, - *, - session_id: str, - agent_id: str, - sequence: int, - run_id: str, - result: AgentRunResult, - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> SessionRecordDraft: - return cls( - session_id=session_id, - agent_id=agent_id, - sequence=sequence, - kind=SessionRecordKind.RUN_FINISHED, - data={ - "run_id": run_id, - "result": agent_run_result_to_dict(result), - }, - branch_id=branch_id, - ) - - -@dataclass(frozen=True, slots=True) -class SessionRecord: - """One immutable, tree-addressable JSONL Session journal record.""" - - record_id: str - parent_id: str | None - branch_id: str - revision: int - session_id: str - agent_id: str | None - sequence: int - kind: SessionRecordKind - data: dict[str, Any] - - def __post_init__(self) -> None: - if not self.record_id: - raise ValueError("record_id must not be empty") - if self.parent_id is not None and not self.parent_id: - raise ValueError("parent_id must not be empty") - if self.revision <= 0: - raise ValueError("revision must be greater than zero") - _validate_envelope_fields( - session_id=self.session_id, - agent_id=self.agent_id, - sequence=self.sequence, - branch_id=self.branch_id, - ) - object.__setattr__(self, "data", deepcopy(self.data)) - - def to_dict(self) -> dict[str, Any]: - return { - "journal_schema_version": SESSION_JOURNAL_SCHEMA_VERSION, - "record_id": self.record_id, - "parent_id": self.parent_id, - "branch_id": self.branch_id, - "revision": self.revision, - "session_id": self.session_id, - "agent_id": self.agent_id, - "sequence": self.sequence, - "type": self.kind.value, - "data": deepcopy(self.data), - } - - @classmethod - def from_dict(cls, value: Mapping[str, Any]) -> SessionRecord: - try: - version = _integer( - value.get("journal_schema_version"), - "journal_schema_version", - ) - if version != SESSION_JOURNAL_SCHEMA_VERSION: - raise SessionSerializationError( - f"unsupported journal_schema_version {version}; " - f"expected {SESSION_JOURNAL_SCHEMA_VERSION}" - ) - raw_parent = value.get("parent_id") - parent_id = None if raw_parent is None else _string(raw_parent, "parent_id") - raw_agent = value.get("agent_id") - agent_id = None if raw_agent is None else _string(raw_agent, "agent_id") - data = value.get("data") - if not isinstance(data, Mapping): - raise SessionSerializationError("data must be an object") - return cls( - record_id=_string(value.get("record_id"), "record_id"), - parent_id=parent_id, - branch_id=_string(value.get("branch_id"), "branch_id"), - revision=_integer(value.get("revision"), "revision"), - session_id=_string(value.get("session_id"), "session_id"), - agent_id=agent_id, - sequence=_integer(value.get("sequence"), "sequence"), - kind=SessionRecordKind(_string(value.get("type"), "type")), - data=deepcopy(dict(data)), - ) - except SessionSerializationError: - raise - except (TypeError, ValueError) as exc: - raise SessionSerializationError( - f"invalid Session journal record: {exc}" - ) from exc - - -def apply_session_record( - session: AgentSession | None, - record: SessionRecord | SessionRecordDraft, -) -> AgentSession: - """Apply one validated mutation to a detached Session projection.""" - - if record.kind is SessionRecordKind.CHECKPOINT: - document = record.data.get("document") - if not isinstance(document, Mapping): - raise SessionSerializationError( - "checkpoint data.document must be an object" - ) - restored = session_from_dict(document) - if restored.session_id != record.session_id: - raise SessionSerializationError( - "checkpoint Session id does not match its journal envelope" - ) - if restored.agent_id != record.agent_id: - raise SessionSerializationError( - "checkpoint Agent id does not match its journal envelope" - ) - return restored - - if record.kind is SessionRecordKind.BRANCH_CREATED: - active = session or AgentSession(session_id=record.session_id) - if active.session_id != record.session_id: - raise SessionSerializationError("journal contains multiple Session ids") - if record.agent_id is not None: - active.bind_agent(record.agent_id) - return active - - active = session or AgentSession(session_id=record.session_id) - if active.session_id != record.session_id: - raise SessionSerializationError("journal contains multiple Session ids") - if record.agent_id is None: - raise SessionSerializationError(f"{record.kind.value} record requires agent_id") - active.bind_agent(record.agent_id) - - if record.kind is SessionRecordKind.RUN_STARTED: - active.begin_run( - _data_string(record, "run_id"), - _data_string(record, "task"), - record.sequence, - ) - elif record.kind is SessionRecordKind.RUN_CONTINUED: - active.begin_continue( - _data_string(record, "run_id"), - record.sequence, - ) - elif record.kind is SessionRecordKind.MESSAGE_APPENDED: - active.append_message( - _data_string(record, "run_id"), - record.sequence, - _data_message(record, "message"), - ) - elif record.kind is SessionRecordKind.MESSAGES_APPENDED: - messages = record.data.get("messages") - if not isinstance(messages, list): - raise SessionSerializationError( - "messages_appended data.messages must be an array" - ) - run_id = _data_string(record, "run_id") - for index, message in enumerate(messages): - if not isinstance(message, Mapping): - raise SessionSerializationError( - f"messages_appended data.messages[{index}] must be an object" - ) - active.append_message( - run_id, - record.sequence, - deepcopy(dict(message)), - ) - elif record.kind is SessionRecordKind.STEERING_APPLIED: - _data_string(record, "input_id") - target_turn = _integer(record.data.get("target_turn"), "target_turn") - if target_turn <= 0: - raise SessionSerializationError( - "steering_applied data.target_turn must be greater than zero" - ) - active.append_message( - _data_string(record, "run_id"), - record.sequence, - { - "role": "user", - "content": _data_string(record, "content"), - }, - ) - elif record.kind is SessionRecordKind.COMPACTION_APPLIED: - summary = record.data.get("summary") - messages = record.data.get("messages") - if not isinstance(summary, Mapping): - raise SessionSerializationError( - "compaction_applied data.summary must be an object" - ) - if not isinstance(messages, list): - raise SessionSerializationError( - "compaction_applied data.messages must be an array" - ) - if any(not isinstance(message, Mapping) for message in messages): - raise SessionSerializationError( - "compaction_applied messages must contain only objects" - ) - normalized_messages = tuple(deepcopy(dict(message)) for message in messages) - active.apply_compaction( - _data_string(record, "operation_id"), - record.sequence, - SummaryEntry.from_dict(summary), - normalized_messages, - ) - elif record.kind is SessionRecordKind.RUN_FINISHED: - active.finish_run( - _data_string(record, "run_id"), - record.sequence, - agent_run_result_from_dict(record.data.get("result")), - ) - else: - raise SessionSerializationError( - f"unsupported Session record type {record.kind.value!r}" - ) - return active - - -def _validate_envelope_fields( - *, - session_id: str, - agent_id: str | None, - sequence: int, - branch_id: str, -) -> None: - if not session_id.strip(): - raise ValueError("session_id must not be empty") - if agent_id is not None and not agent_id.strip(): - raise ValueError("agent_id must not be empty") - if sequence < 0: - raise ValueError("sequence must not be negative") - if not branch_id.strip(): - raise ValueError("branch_id must not be empty") - - -def _data_string( - record: SessionRecord | SessionRecordDraft, - key: str, -) -> str: - return _string(record.data.get(key), f"{record.kind.value} data.{key}") - - -def _data_message( - record: SessionRecord | SessionRecordDraft, - key: str, -) -> AgentMessage: - value = record.data.get(key) - if not isinstance(value, Mapping): - raise SessionSerializationError( - f"{record.kind.value} data.{key} must be an object" - ) - return deepcopy(dict(value)) - - -def _string(value: Any, label: str) -> str: - if not isinstance(value, str): - raise SessionSerializationError(f"{label} must be a string") - return value - - -def _integer(value: Any, label: str) -> int: - if not isinstance(value, int) or isinstance(value, bool): - raise SessionSerializationError(f"{label} must be an integer") - return value diff --git a/src/ejagent/session/jsonl.py b/src/ejagent/session/jsonl.py deleted file mode 100644 index 53eb704..0000000 --- a/src/ejagent/session/jsonl.py +++ /dev/null @@ -1,963 +0,0 @@ -from __future__ import annotations - -import asyncio -import json -import math -import os -import threading -from collections.abc import Callable -from dataclasses import dataclass, replace -from hashlib import sha256 -from pathlib import Path -from typing import Any, TypeVar -from uuid import uuid4 - -from ejagent.session._file_lock import SessionFileLock -from ejagent.session.errors import ( - SessionConflictError, - SessionSerializationError, - SessionStorageError, -) -from ejagent.session.journal import ( - DEFAULT_SESSION_BRANCH, - SessionRecord, - SessionRecordDraft, - SessionRecordKind, - apply_session_record, -) -from ejagent.session.tree import ( - SessionBranch, - SessionBranchIntent, - SessionCheckout, - SessionRetry, -) -from ejagent.session.types import AgentSession - -_ResultT = TypeVar("_ResultT") - - -@dataclass(slots=True) -class _JournalIndex: - records: list[SessionRecord] - records_by_id: dict[str, SessionRecord] - branches: dict[str, SessionBranch] - valid_length: int = 0 - - @classmethod - def empty(cls) -> _JournalIndex: - return cls(records=[], records_by_id={}, branches={}) - - -class JsonlSessionStorage: - """Append and project immutable, tree-addressable Session records. - - One JSONL file contains every branch for a Session. File order assigns a - global revision while ``parent_id`` defines the logical tree. A stable - sidecar lock makes read-validate-append atomic across instances and - processes on POSIX filesystems. - """ - - def __init__( - self, - root: str | Path, - *, - lock_timeout: float | None = 10.0, - ) -> None: - if lock_timeout is not None: - if isinstance(lock_timeout, bool) or not isinstance( - lock_timeout, (int, float) - ): - raise TypeError("lock_timeout must be a number or None") - if lock_timeout < 0 or not math.isfinite(lock_timeout): - raise ValueError("lock_timeout must be finite and non-negative") - lock_timeout = float(lock_timeout) - self.root = Path(root).expanduser() - self.lock_timeout = lock_timeout - self._lock = asyncio.Lock() - - async def load(self, session_id: str) -> AgentSession | None: - """Load the detached ``main`` branch projection for compatibility.""" - - checkout = await self.checkout(session_id) - return checkout.session if checkout is not None else None - - async def checkout( - self, - session_id: str, - *, - branch_id: str = DEFAULT_SESSION_BRANCH, - record_id: str | None = None, - ) -> SessionCheckout | None: - """Project one branch head or an exact record without changing history.""" - - normalized_id = self._normalize_session_id(session_id) - normalized_branch = self._normalize_branch_id(branch_id) - path = self._path_for(normalized_id) - async with self._lock: - index = await self._read_index(path, normalized_id) - return self._checkout_sync( - index, - normalized_id, - branch_id=normalized_branch, - record_id=record_id, - ) - - async def head( - self, - session_id: str, - *, - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> SessionRecord | None: - """Return the immutable record currently heading one branch.""" - - normalized_id = self._normalize_session_id(session_id) - normalized_branch = self._normalize_branch_id(branch_id) - path = self._path_for(normalized_id) - async with self._lock: - index = await self._read_index(path, normalized_id) - branch = index.branches.get(normalized_branch) - if branch is None: - return None - return index.records_by_id[branch.head_record_id] - - async def list_branches(self, session_id: str) -> tuple[SessionBranch, ...]: - """Return branches ordered by their creation revision.""" - - normalized_id = self._normalize_session_id(session_id) - path = self._path_for(normalized_id) - async with self._lock: - index = await self._read_index(path, normalized_id) - return tuple( - sorted( - index.branches.values(), key=lambda branch: branch.created_revision - ) - ) - - async def save(self, session: AgentSession) -> None: - """Append a full logical Checkpoint to the ``main`` branch.""" - - await self.append(SessionRecordDraft.checkpoint(session.snapshot())) - - async def append( - self, - draft: SessionRecordDraft, - *, - expected_head_id: str | None = None, - check_head: bool = False, - ) -> SessionRecord: - """Append one branch mutation, optionally comparing its current head.""" - - if draft.kind is SessionRecordKind.BRANCH_CREATED: - raise ValueError("use fork(), rollback(), or prepare_retry()") - self._validate_draft_json(draft) - path = self._path_for(draft.session_id) - async with self._lock: - return await self._run_locked( - path, - self._append_sync, - path, - draft, - expected_head_id, - check_head, - ) - - async def records(self, session_id: str) -> tuple[SessionRecord, ...]: - """Return detached immutable records in physical append order.""" - - normalized_id = self._normalize_session_id(session_id) - path = self._path_for(normalized_id) - async with self._lock: - index = await self._read_index(path, normalized_id) - return tuple(index.records) - - async def fork( - self, - session_id: str, - *, - source_branch: str = DEFAULT_SESSION_BRANCH, - from_record_id: str | None = None, - branch_id: str | None = None, - ) -> SessionCheckout: - """Create a general-purpose branch at a completed Session projection.""" - - return await self._create_branch( - session_id, - source_branch=source_branch, - base_record_id=from_record_id, - branch_id=branch_id, - intent=SessionBranchIntent.FORK, - ) - - async def rollback( - self, - session_id: str, - *, - to_record_id: str, - source_branch: str = DEFAULT_SESSION_BRANCH, - branch_id: str | None = None, - ) -> SessionCheckout: - """Create a branch at an ancestor without rewriting the source branch.""" - - if not to_record_id: - raise ValueError("to_record_id must not be empty") - return await self._create_branch( - session_id, - source_branch=source_branch, - base_record_id=to_record_id, - branch_id=branch_id, - intent=SessionBranchIntent.ROLLBACK, - ) - - async def prepare_retry( - self, - session_id: str, - *, - run_id: str, - source_branch: str = DEFAULT_SESSION_BRANCH, - branch_id: str | None = None, - ) -> SessionRetry: - """Branch before one Run and return its original task for explicit retry.""" - - run_id = run_id.strip() - if not run_id: - raise ValueError("run_id must not be empty") - normalized_id = self._normalize_session_id(session_id) - normalized_source = self._normalize_branch_id(source_branch) - normalized_branch = ( - self._normalize_branch_id(branch_id) - if branch_id is not None - else self._generated_branch_id(SessionBranchIntent.RETRY) - ) - path = self._path_for(normalized_id) - async with self._lock: - return await self._run_locked( - path, - self._prepare_retry_sync, - path, - normalized_id, - run_id, - normalized_source, - normalized_branch, - ) - - async def _create_branch( - self, - session_id: str, - *, - source_branch: str, - base_record_id: str | None, - branch_id: str | None, - intent: SessionBranchIntent, - ) -> SessionCheckout: - normalized_id = self._normalize_session_id(session_id) - normalized_source = self._normalize_branch_id(source_branch) - normalized_base = ( - self._normalize_record_id(base_record_id) - if base_record_id is not None - else None - ) - normalized_branch = ( - self._normalize_branch_id(branch_id) - if branch_id is not None - else self._generated_branch_id(intent) - ) - path = self._path_for(normalized_id) - async with self._lock: - return await self._run_locked( - path, - self._create_branch_sync, - path, - normalized_id, - normalized_source, - normalized_base, - normalized_branch, - intent, - None, - ) - - @staticmethod - def _normalize_session_id(session_id: str) -> str: - if not isinstance(session_id, str): - raise TypeError("session_id must be a string") - normalized = session_id.strip() - if not normalized: - raise ValueError("session_id must not be empty") - return normalized - - @staticmethod - def _normalize_branch_id(branch_id: str) -> str: - if not isinstance(branch_id, str): - raise TypeError("branch_id must be a string") - normalized = branch_id.strip() - if not normalized: - raise ValueError("branch_id must not be empty") - return normalized - - @staticmethod - def _normalize_record_id(record_id: str) -> str: - if not isinstance(record_id, str): - raise TypeError("record_id must be a string") - normalized = record_id.strip() - if not normalized: - raise ValueError("record_id must not be empty") - return normalized - - @staticmethod - def _generated_branch_id(intent: SessionBranchIntent) -> str: - return f"{intent.value}-{uuid4().hex[:12]}" - - def _path_for(self, session_id: str) -> Path: - normalized = self._normalize_session_id(session_id) - digest = sha256(normalized.encode("utf-8")).hexdigest() - return self.root / f"{digest}.jsonl" - - @staticmethod - def _lock_path_for(path: Path) -> Path: - return path.with_name(f"{path.name}.lock") - - async def _read_index( - self, - path: Path, - expected_session_id: str, - ) -> _JournalIndex: - return await self._run_worker( - self._read_locked_sync, - path, - expected_session_id, - ) - - async def _run_locked( - self, - path: Path, - operation: Callable[..., _ResultT], - *args: Any, - ) -> _ResultT: - return await self._run_worker( - self._call_locked_sync, - path, - operation, - args, - ) - - async def _run_worker( - self, - operation: Callable[..., _ResultT], - *args: Any, - ) -> _ResultT: - cancelled = threading.Event() - worker = asyncio.create_task(asyncio.to_thread(operation, *args, cancelled)) - try: - return await asyncio.shield(worker) - except asyncio.CancelledError: - cancelled.set() - while True: - try: - await asyncio.shield(worker) - break - except asyncio.CancelledError: - continue - except Exception: - break - raise - - def _read_locked_sync( - self, - path: Path, - expected_session_id: str, - cancelled: threading.Event, - ) -> _JournalIndex: - lock_path = self._lock_path_for(path) - if not path.exists() and not lock_path.exists(): - return _JournalIndex.empty() - return self._call_locked_sync( - path, - self._read_sync, - (path, expected_session_id), - cancelled, - ) - - def _call_locked_sync( - self, - path: Path, - operation: Callable[..., _ResultT], - args: tuple[Any, ...], - cancelled: threading.Event, - ) -> _ResultT: - lock = SessionFileLock( - self._lock_path_for(path), - timeout=self.lock_timeout, - ) - with lock.acquire(cancelled): - return operation(*args) - - @staticmethod - def _validate_draft_json(draft: SessionRecordDraft) -> None: - try: - json.dumps(draft.data, allow_nan=False, sort_keys=True) - except (TypeError, ValueError) as exc: - raise SessionSerializationError( - f"Session record is not JSON-compatible: {exc}" - ) from exc - - def _append_sync( - self, - path: Path, - draft: SessionRecordDraft, - expected_head_id: str | None, - check_head: bool, - ) -> SessionRecord: - index = self._read_sync(path, draft.session_id) - branch = index.branches.get(draft.branch_id) - if branch is None and draft.branch_id != DEFAULT_SESSION_BRANCH: - raise ValueError(f"unknown Session branch {draft.branch_id!r}") - current_head = branch.head_record_id if branch is not None else None - if check_head and current_head != expected_head_id: - raise SessionConflictError( - f"Session branch {draft.branch_id!r} head changed from " - f"{expected_head_id!r} to {current_head!r}" - ) - record = SessionRecord( - record_id=uuid4().hex, - parent_id=current_head, - branch_id=draft.branch_id, - revision=len(index.records) + 1, - session_id=draft.session_id, - agent_id=draft.agent_id, - sequence=draft.sequence, - kind=draft.kind, - data=draft.data, - ) - self._validate_candidate(index, record, path) - self._write_record_sync(path, record, index.valid_length) - return record - - def _create_branch_sync( - self, - path: Path, - session_id: str, - source_branch_id: str, - base_record_id: str | None, - branch_id: str, - intent: SessionBranchIntent, - retried_run_id: str | None, - ) -> SessionCheckout: - index = self._read_sync(path, session_id) - source = self._require_branch(index, source_branch_id) - if branch_id in index.branches: - raise ValueError(f"Session branch {branch_id!r} already exists") - base_id = base_record_id or source.head_record_id - if base_id not in index.records_by_id: - raise ValueError(f"unknown Session record {base_id!r}") - if not self._is_ancestor(index, base_id, source.head_record_id): - raise ValueError( - f"record {base_id!r} is not an ancestor of branch {source_branch_id!r}" - ) - base_session = self._project(index, base_id, session_id) - self._require_finished(base_session, label=f"record {base_id!r}") - draft = SessionRecordDraft.branch_created( - session_id=session_id, - agent_id=base_session.agent_id, - branch_id=branch_id, - base_record_id=base_id, - source_branch_id=source_branch_id, - source_head_id=source.head_record_id, - intent=intent.value, - retried_run_id=retried_run_id, - ) - record = SessionRecord( - record_id=uuid4().hex, - parent_id=base_id, - branch_id=branch_id, - revision=len(index.records) + 1, - session_id=session_id, - agent_id=draft.agent_id, - sequence=0, - kind=SessionRecordKind.BRANCH_CREATED, - data=draft.data, - ) - self._validate_candidate(index, record, path) - self._write_record_sync(path, record, index.valid_length) - checkout = self._checkout_sync(index, session_id, branch_id=branch_id) - if checkout is None: - raise RuntimeError("created Session branch is unavailable") - return checkout - - def _prepare_retry_sync( - self, - path: Path, - session_id: str, - run_id: str, - source_branch_id: str, - branch_id: str, - ) -> SessionRetry: - index = self._read_sync(path, session_id) - source = self._require_branch(index, source_branch_id) - if branch_id in index.branches: - raise ValueError(f"Session branch {branch_id!r} already exists") - run_record = self._find_run_started(index, source.head_record_id, run_id) - if run_record is None: - raise ValueError( - f"unknown run {run_id!r} on Session branch {source_branch_id!r}" - ) - task = run_record.data.get("task") - if not isinstance(task, str) or not task: - raise SessionSerializationError( - f"run_started record {run_record.record_id!r} has no valid task" - ) - if run_record.parent_id is None: - base_session = AgentSession(session_id=session_id) - if run_record.agent_id is not None: - base_session.bind_agent(run_record.agent_id) - checkout = self._create_root_retry_sync( - path, - index, - base_session, - source, - branch_id, - run_id, - ) - else: - checkout = self._create_branch_sync( - path, - session_id, - source_branch_id, - run_record.parent_id, - branch_id, - SessionBranchIntent.RETRY, - run_id, - ) - return SessionRetry(checkout=checkout, task=task, retried_run_id=run_id) - - def _create_root_retry_sync( - self, - path: Path, - index: _JournalIndex, - base_session: AgentSession, - source: SessionBranch, - branch_id: str, - run_id: str, - ) -> SessionCheckout: - draft = SessionRecordDraft.branch_created( - session_id=base_session.session_id, - agent_id=base_session.agent_id, - branch_id=branch_id, - base_record_id=None, - source_branch_id=source.branch_id, - source_head_id=source.head_record_id, - intent=SessionBranchIntent.RETRY.value, - retried_run_id=run_id, - ) - record = SessionRecord( - record_id=uuid4().hex, - parent_id=None, - branch_id=branch_id, - revision=len(index.records) + 1, - session_id=base_session.session_id, - agent_id=base_session.agent_id, - sequence=0, - kind=SessionRecordKind.BRANCH_CREATED, - data=draft.data, - ) - self._validate_candidate(index, record, path) - self._write_record_sync(path, record, index.valid_length) - checkout = self._checkout_sync( - index, - base_session.session_id, - branch_id=branch_id, - ) - if checkout is None: - raise RuntimeError("created retry branch is unavailable") - return checkout - - def _validate_candidate( - self, - index: _JournalIndex, - record: SessionRecord, - path: Path, - ) -> None: - self._add_record(index, record, path=path, line_number=record.revision) - base = ( - self._project(index, record.parent_id, record.session_id) - if record.parent_id is not None - else None - ) - self._apply_checked(base, record, path=path, line_number=record.revision) - - def _write_record_sync( - self, - path: Path, - record: SessionRecord, - valid_length: int, - ) -> None: - try: - self.root.mkdir(parents=True, exist_ok=True) - except OSError as exc: - raise SessionStorageError( - f"failed to create Session journal directory {self.root}" - ) from exc - try: - line = ( - json.dumps( - record.to_dict(), - ensure_ascii=False, - separators=(",", ":"), - sort_keys=True, - allow_nan=False, - ) - + "\n" - ).encode("utf-8") - except (TypeError, ValueError) as exc: - raise SessionSerializationError( - f"Session record is not JSON-compatible: {exc}" - ) from exc - - descriptor: int | None = None - try: - descriptor = os.open(path, os.O_RDWR | os.O_CREAT | os.O_APPEND, 0o600) - current_size = os.fstat(descriptor).st_size - if valid_length < current_size: - os.ftruncate(descriptor, valid_length) - os.lseek(descriptor, 0, os.SEEK_END) - written = os.write(descriptor, line) - if written != len(line): - raise OSError( - f"partial Session journal write: {written}/{len(line)} bytes" - ) - os.fsync(descriptor) - except OSError as exc: - raise SessionStorageError( - f"failed to append Session journal {path}" - ) from exc - finally: - if descriptor is not None: - os.close(descriptor) - self._fsync_directory() - - def _read_sync(self, path: Path, expected_session_id: str) -> _JournalIndex: - try: - content = path.read_bytes() - except FileNotFoundError: - return _JournalIndex.empty() - except OSError as exc: - raise SessionStorageError(f"failed to read Session journal {path}") from exc - - index = _JournalIndex.empty() - branch_sessions: dict[str, AgentSession] = {} - lines = content.splitlines(keepends=True) - for line_index, encoded_line in enumerate(lines): - line_number = line_index + 1 - if not encoded_line.endswith(b"\n"): - if index.records and line_index == len(lines) - 1: - break - raise SessionSerializationError( - f"Session journal {path} has an incomplete first record" - ) - try: - raw: Any = json.loads(encoded_line) - except (UnicodeDecodeError, json.JSONDecodeError) as exc: - raise SessionSerializationError( - f"Session journal {path} has invalid JSON at line {line_number}" - ) from exc - if not isinstance(raw, dict): - raise SessionSerializationError( - f"Session journal {path} line {line_number} must be an object" - ) - record = SessionRecord.from_dict(raw) - if record.session_id != expected_session_id: - raise SessionSerializationError( - f"Session journal {path} line {line_number} contains id " - f"{record.session_id!r}, expected {expected_session_id!r}" - ) - previous_session = branch_sessions.get(record.branch_id) - self._add_record(index, record, path=path, line_number=line_number) - if record.kind is SessionRecordKind.BRANCH_CREATED: - base = ( - self._project(index, record.parent_id, expected_session_id) - if record.parent_id is not None - else None - ) - else: - base = previous_session - branch_sessions[record.branch_id] = self._apply_checked( - base, - record, - path=path, - line_number=line_number, - ) - index.valid_length += len(encoded_line) - return index - - def _add_record( - self, - index: _JournalIndex, - record: SessionRecord, - *, - path: Path, - line_number: int, - ) -> None: - if record.record_id in index.records_by_id: - raise SessionSerializationError( - f"Session journal {path} repeats record_id " - f"{record.record_id!r} at line {line_number}" - ) - if record.revision != len(index.records) + 1: - raise SessionSerializationError( - f"Session journal revision jumped at {path}:{line_number}" - ) - if record.parent_id is not None and record.parent_id not in index.records_by_id: - raise SessionSerializationError( - f"Session journal parent changed or is missing at {path}:{line_number}" - ) - - existing = index.branches.get(record.branch_id) - if record.kind is SessionRecordKind.BRANCH_CREATED: - branch = self._branch_from_record(index, record, path, line_number) - elif existing is None: - if record.branch_id != DEFAULT_SESSION_BRANCH: - raise SessionSerializationError( - f"Session branch {record.branch_id!r} was not created" - ) - if record.parent_id is not None: - raise SessionSerializationError( - f"initial main record has a parent at {path}:{line_number}" - ) - branch = SessionBranch( - branch_id=DEFAULT_SESSION_BRANCH, - head_record_id=record.record_id, - base_record_id=None, - source_branch_id=None, - intent=None, - created_revision=record.revision, - ) - else: - if record.parent_id != existing.head_record_id: - raise SessionSerializationError( - f"Session branch head changed at {path}:{line_number}" - ) - branch = replace(existing, head_record_id=record.record_id) - - index.records.append(record) - index.records_by_id[record.record_id] = record - index.branches[record.branch_id] = branch - - def _branch_from_record( - self, - index: _JournalIndex, - record: SessionRecord, - path: Path, - line_number: int, - ) -> SessionBranch: - if record.branch_id == DEFAULT_SESSION_BRANCH: - raise SessionSerializationError( - "main branch must not be created explicitly" - ) - if record.branch_id in index.branches: - raise SessionSerializationError( - f"Session branch {record.branch_id!r} already exists" - ) - base_record_id = self._optional_data_string(record, "base_record_id") - source_branch_id = self._required_data_string(record, "source_branch_id") - source_head_id = self._required_data_string(record, "source_head_id") - if record.parent_id != base_record_id: - raise SessionSerializationError( - f"branch base does not match parent at {path}:{line_number}" - ) - source = index.branches.get(source_branch_id) - if source is None: - raise SessionSerializationError( - f"unknown source branch {source_branch_id!r} at {path}:{line_number}" - ) - if source.head_record_id != source_head_id: - raise SessionSerializationError( - f"source branch head does not match at {path}:{line_number}" - ) - if base_record_id is not None and not self._is_ancestor( - index, - base_record_id, - source_head_id, - ): - raise SessionSerializationError( - f"branch base is not a source ancestor at {path}:{line_number}" - ) - try: - intent = SessionBranchIntent(self._required_data_string(record, "intent")) - except ValueError as exc: - raise SessionSerializationError( - f"invalid branch intent at {path}:{line_number}" - ) from exc - retried_run_id = record.data.get("retried_run_id") - if intent is SessionBranchIntent.RETRY: - if not isinstance(retried_run_id, str) or not retried_run_id: - raise SessionSerializationError( - f"retry branch has no retried_run_id at {path}:{line_number}" - ) - elif retried_run_id is not None: - raise SessionSerializationError( - f"non-retry branch has retried_run_id at {path}:{line_number}" - ) - if base_record_id is None and intent is not SessionBranchIntent.RETRY: - raise SessionSerializationError( - f"only a first-Run retry may branch from the root at " - f"{path}:{line_number}" - ) - return SessionBranch( - branch_id=record.branch_id, - head_record_id=record.record_id, - base_record_id=base_record_id, - source_branch_id=source_branch_id, - intent=intent, - created_revision=record.revision, - ) - - @staticmethod - def _apply_checked( - session: AgentSession | None, - record: SessionRecord, - *, - path: Path, - line_number: int, - ) -> AgentSession: - if ( - session is not None - and session.agent_id is not None - and record.agent_id != session.agent_id - ): - raise SessionSerializationError( - f"Session journal {path} changes agent_id at line {line_number}" - ) - try: - return apply_session_record(session, record) - except (KeyError, TypeError, ValueError) as exc: - if isinstance(exc, SessionSerializationError): - raise - raise SessionSerializationError( - f"invalid Session mutation at {path}:{line_number}: {exc}" - ) from exc - - def _checkout_sync( - self, - index: _JournalIndex, - session_id: str, - *, - branch_id: str, - record_id: str | None = None, - ) -> SessionCheckout | None: - if record_id is None: - branch = index.branches.get(branch_id) - if branch is None: - return None - target_id = branch.head_record_id - else: - target = index.records_by_id.get(record_id) - if target is None: - raise ValueError(f"unknown Session record {record_id!r}") - stored_branch = index.branches[target.branch_id] - branch = replace(stored_branch, head_record_id=target.record_id) - target_id = target.record_id - session = self._project(index, target_id, session_id) - return SessionCheckout( - session=session, - branch=branch, - head=index.records_by_id[target_id], - ) - - def _project( - self, - index: _JournalIndex, - record_id: str, - session_id: str, - ) -> AgentSession: - path: list[SessionRecord] = [] - current_id: str | None = record_id - while current_id is not None: - record = index.records_by_id.get(current_id) - if record is None: - raise SessionSerializationError( - f"Session tree references missing record {current_id!r}" - ) - path.append(record) - current_id = record.parent_id - session: AgentSession | None = None - for record in reversed(path): - session = apply_session_record(session, record) - if session is None: - return AgentSession(session_id=session_id) - return session.snapshot() - - @staticmethod - def _require_branch(index: _JournalIndex, branch_id: str) -> SessionBranch: - branch = index.branches.get(branch_id) - if branch is None: - raise ValueError(f"unknown Session branch {branch_id!r}") - return branch - - @staticmethod - def _require_finished(session: AgentSession, *, label: str) -> None: - unfinished = [run.run_id for run in session.runs if not run.finished] - if unfinished: - raise ValueError( - f"cannot branch from {label} with unfinished run(s): " - + ", ".join(unfinished) - ) - - @staticmethod - def _is_ancestor( - index: _JournalIndex, - ancestor_id: str, - descendant_id: str, - ) -> bool: - current_id: str | None = descendant_id - while current_id is not None: - if current_id == ancestor_id: - return True - current_id = index.records_by_id[current_id].parent_id - return False - - @staticmethod - def _find_run_started( - index: _JournalIndex, - head_record_id: str, - run_id: str, - ) -> SessionRecord | None: - current_id: str | None = head_record_id - while current_id is not None: - record = index.records_by_id[current_id] - if ( - record.kind is SessionRecordKind.RUN_STARTED - and record.data.get("run_id") == run_id - ): - return record - current_id = record.parent_id - return None - - @staticmethod - def _required_data_string(record: SessionRecord, key: str) -> str: - value = record.data.get(key) - if not isinstance(value, str) or not value: - raise SessionSerializationError( - f"branch_created data.{key} must be a non-empty string" - ) - return value - - @staticmethod - def _optional_data_string(record: SessionRecord, key: str) -> str | None: - value = record.data.get(key) - if value is None: - return None - if not isinstance(value, str) or not value: - raise SessionSerializationError( - f"branch_created data.{key} must be null or a non-empty string" - ) - return value - - def _fsync_directory(self) -> None: - try: - descriptor = os.open(self.root, os.O_RDONLY) - except OSError: - return - try: - os.fsync(descriptor) - except OSError: - pass - finally: - os.close(descriptor) diff --git a/src/ejagent/session/memory.py b/src/ejagent/session/memory.py deleted file mode 100644 index 90eebdb..0000000 --- a/src/ejagent/session/memory.py +++ /dev/null @@ -1,20 +0,0 @@ -import asyncio - -from ejagent.session.types import AgentSession - - -class MemorySessionStorage: - """Process-local Session storage with copy isolation.""" - - def __init__(self) -> None: - self._sessions: dict[str, AgentSession] = {} - self._lock = asyncio.Lock() - - async def load(self, session_id: str) -> AgentSession | None: - async with self._lock: - session = self._sessions.get(session_id) - return session.snapshot() if session is not None else None - - async def save(self, session: AgentSession) -> None: - async with self._lock: - self._sessions[session.session_id] = session.snapshot() diff --git a/src/ejagent/session/recorder.py b/src/ejagent/session/recorder.py deleted file mode 100644 index 726f9f0..0000000 --- a/src/ejagent/session/recorder.py +++ /dev/null @@ -1,196 +0,0 @@ -import asyncio - -from ejagent.agent.events import ( - AgentContinued, - AgentEvent, - AgentFinished, - AgentStarted, - CompactionCompleted, - MessageCompleted, - SteeringApplied, - ToolCompleted, -) -from ejagent.providers.base import serialize_assistant_message -from ejagent.session.journal import DEFAULT_SESSION_BRANCH, SessionRecordDraft -from ejagent.session.storage import SessionJournalStorage, SessionStorage -from ejagent.session.types import AgentSession - -_RECORDED_PAYLOADS = ( - AgentStarted, - AgentContinued, - MessageCompleted, - SteeringApplied, - ToolCompleted, - AgentFinished, - CompactionCompleted, -) - - -class SessionRecorder: - """Build one selected Session branch from read-only lifecycle events.""" - - def __init__( - self, - *, - session_id: str, - storage: SessionStorage, - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> None: - session_id = session_id.strip() - if not session_id: - raise ValueError("session_id must not be empty") - self.session_id = session_id - self.storage = storage - self.branch_id = branch_id.strip() - if not self.branch_id: - raise ValueError("branch_id must not be empty") - self._lock = asyncio.Lock() - - async def emit(self, event: AgentEvent) -> None: - payload = event.payload - if not isinstance(payload, _RECORDED_PAYLOADS): - return - if isinstance(payload, CompactionCompleted) and not payload.result.completed: - return - - async with self._lock: - expected_head_id: str | None = None - journal_storage = ( - self.storage - if isinstance(self.storage, SessionJournalStorage) - else None - ) - if journal_storage is not None: - checkout = await journal_storage.checkout( - self.session_id, - branch_id=self.branch_id, - ) - session = checkout.session if checkout is not None else None - expected_head_id = ( - checkout.head.record_id if checkout is not None else None - ) - else: - session = await self.storage.load(self.session_id) - if session is None: - if self.branch_id != DEFAULT_SESSION_BRANCH: - raise ValueError(f"unknown Session branch {self.branch_id!r}") - session = AgentSession(session_id=self.session_id) - session.bind_agent(event.agent_id) - - if isinstance(payload, AgentStarted): - session.begin_run(event.run_id, payload.task, event.sequence) - draft = SessionRecordDraft.run_started( - session_id=self.session_id, - agent_id=event.agent_id, - sequence=event.sequence, - run_id=event.run_id, - task=payload.task, - branch_id=self.branch_id, - ) - elif isinstance(payload, AgentContinued): - session.begin_continue(event.run_id, event.sequence) - draft = SessionRecordDraft.run_continued( - session_id=self.session_id, - agent_id=event.agent_id, - sequence=event.sequence, - run_id=event.run_id, - branch_id=self.branch_id, - ) - elif isinstance(payload, SteeringApplied): - session.append_message( - event.run_id, - event.sequence, - {"role": "user", "content": payload.control.content}, - ) - draft = SessionRecordDraft.steering_applied( - session_id=self.session_id, - agent_id=event.agent_id, - sequence=event.sequence, - run_id=event.run_id, - input_id=payload.control.input_id, - content=payload.control.content, - target_turn=payload.target_turn, - branch_id=self.branch_id, - ) - elif isinstance(payload, CompactionCompleted): - assert payload.result.summary is not None - session.apply_compaction( - payload.result.operation_id, - event.sequence, - payload.result.summary, - payload.result.messages, - ) - draft = SessionRecordDraft.compaction_applied( - session_id=self.session_id, - agent_id=event.agent_id, - sequence=event.sequence, - result=payload.result, - branch_id=self.branch_id, - ) - elif isinstance(payload, MessageCompleted): - message = serialize_assistant_message( - payload.message, - usage=payload.usage, - ) - session.append_message( - event.run_id, - event.sequence, - message, - ) - draft = SessionRecordDraft.message_appended( - session_id=self.session_id, - agent_id=event.agent_id, - sequence=event.sequence, - run_id=event.run_id, - message=message, - branch_id=self.branch_id, - ) - elif isinstance(payload, ToolCompleted): - for message in payload.result.messages: - session.append_message( - event.run_id, - event.sequence, - message, - ) - draft = SessionRecordDraft.messages_appended( - session_id=self.session_id, - agent_id=event.agent_id, - sequence=event.sequence, - run_id=event.run_id, - messages=payload.result.messages, - branch_id=self.branch_id, - ) - else: - session.finish_run( - event.run_id, - event.sequence, - payload.result, - ) - draft = SessionRecordDraft.run_finished( - session_id=self.session_id, - agent_id=event.agent_id, - sequence=event.sequence, - run_id=event.run_id, - result=payload.result, - branch_id=self.branch_id, - ) - - if journal_storage is not None: - await journal_storage.append( - draft, - expected_head_id=expected_head_id, - check_head=True, - ) - else: - await self.storage.save(session) - - async def load(self) -> AgentSession | None: - """Load the currently persisted detached Session snapshot.""" - - if isinstance(self.storage, SessionJournalStorage): - checkout = await self.storage.checkout( - self.session_id, - branch_id=self.branch_id, - ) - return checkout.session if checkout is not None else None - return await self.storage.load(self.session_id) diff --git a/src/ejagent/session/storage.py b/src/ejagent/session/storage.py deleted file mode 100644 index 089767a..0000000 --- a/src/ejagent/session/storage.py +++ /dev/null @@ -1,91 +0,0 @@ -from typing import Protocol, runtime_checkable - -from ejagent.session.journal import ( - DEFAULT_SESSION_BRANCH, - SessionRecord, - SessionRecordDraft, -) -from ejagent.session.tree import SessionBranch, SessionCheckout, SessionRetry -from ejagent.session.types import AgentSession - - -class SessionStorage(Protocol): - """Persistence boundary for detached Agent Session snapshots.""" - - async def load(self, session_id: str) -> AgentSession | None: - """Load a detached Session or return ``None`` when it does not exist.""" - - async def save(self, session: AgentSession) -> None: - """Create or replace one Session snapshot.""" - - -@runtime_checkable -class SessionJournalStorage(SessionStorage, Protocol): - """Storage capable of appending semantic Session journal records.""" - - async def checkout( - self, - session_id: str, - *, - branch_id: str = DEFAULT_SESSION_BRANCH, - record_id: str | None = None, - ) -> SessionCheckout | None: - """Project one branch head or an exact record.""" - - async def append( - self, - draft: SessionRecordDraft, - *, - expected_head_id: str | None = None, - check_head: bool = False, - ) -> SessionRecord: - """Atomically append one mutation and return its assigned envelope.""" - - -@runtime_checkable -class SessionTreeStorage(SessionJournalStorage, Protocol): - """Backend-neutral persistence boundary for addressable Session trees.""" - - async def head( - self, - session_id: str, - *, - branch_id: str = DEFAULT_SESSION_BRANCH, - ) -> SessionRecord | None: - """Return the immutable record currently heading one branch.""" - - async def list_branches(self, session_id: str) -> tuple[SessionBranch, ...]: - """Return the Session's branches in backend-defined stable order.""" - - async def records(self, session_id: str) -> tuple[SessionRecord, ...]: - """Return immutable records in global revision order.""" - - async def fork( - self, - session_id: str, - *, - source_branch: str = DEFAULT_SESSION_BRANCH, - from_record_id: str | None = None, - branch_id: str | None = None, - ) -> SessionCheckout: - """Create a branch at a completed source projection.""" - - async def rollback( - self, - session_id: str, - *, - to_record_id: str, - source_branch: str = DEFAULT_SESSION_BRANCH, - branch_id: str | None = None, - ) -> SessionCheckout: - """Create a branch at an ancestor without rewriting its source.""" - - async def prepare_retry( - self, - session_id: str, - *, - run_id: str, - source_branch: str = DEFAULT_SESSION_BRANCH, - branch_id: str | None = None, - ) -> SessionRetry: - """Branch before one Run and return its original task.""" diff --git a/src/ejagent/session/tree.py b/src/ejagent/session/tree.py deleted file mode 100644 index 2a27c27..0000000 --- a/src/ejagent/session/tree.py +++ /dev/null @@ -1,62 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass -from enum import StrEnum - -from ejagent.session.journal import SessionRecord -from ejagent.session.types import AgentSession - - -class SessionBranchIntent(StrEnum): - """Audited reason for creating a Session branch.""" - - FORK = "fork" - ROLLBACK = "rollback" - RETRY = "retry" - - -@dataclass(frozen=True, slots=True) -class SessionBranch: - """One named branch and its current immutable journal head.""" - - branch_id: str - head_record_id: str - base_record_id: str | None - source_branch_id: str | None - intent: SessionBranchIntent | None - created_revision: int - - def __post_init__(self) -> None: - if not self.branch_id: - raise ValueError("branch_id must not be empty") - if not self.head_record_id: - raise ValueError("head_record_id must not be empty") - if self.created_revision <= 0: - raise ValueError("created_revision must be greater than zero") - - -@dataclass(frozen=True, slots=True) -class SessionCheckout: - """Detached Session projection at one branch or record head.""" - - session: AgentSession - branch: SessionBranch - head: SessionRecord - - def __post_init__(self) -> None: - object.__setattr__(self, "session", self.session.snapshot()) - - -@dataclass(frozen=True, slots=True) -class SessionRetry: - """A prepared retry branch and the original Run task to execute.""" - - checkout: SessionCheckout - task: str - retried_run_id: str - - def __post_init__(self) -> None: - if not self.task: - raise ValueError("task must not be empty") - if not self.retried_run_id: - raise ValueError("retried_run_id must not be empty") diff --git a/src/ejagent/session/types.py b/src/ejagent/session/types.py deleted file mode 100644 index bf846d4..0000000 --- a/src/ejagent/session/types.py +++ /dev/null @@ -1,290 +0,0 @@ -from __future__ import annotations - -from copy import deepcopy -from dataclasses import dataclass, field, replace -from enum import StrEnum - -from ejagent.agent.compaction import SummaryEntry -from ejagent.agent.result import AgentRunResult -from ejagent.agent.types import AgentMessage - - -class SessionRunIntent(StrEnum): - """Why one Session Run entered the agent runtime.""" - - TASK = "task" - CONTINUE = "continue" - - -@dataclass(frozen=True, slots=True) -class SessionMessage: - """One persistent conversation message associated with an agent run.""" - - run_id: str - sequence: int - message: AgentMessage - - def __post_init__(self) -> None: - if not self.run_id: - raise ValueError("run_id must not be empty") - if self.sequence <= 0: - raise ValueError("sequence must be greater than zero") - object.__setattr__(self, "message", deepcopy(self.message)) - - -@dataclass(frozen=True, slots=True) -class SessionRun: - """Persistent boundary and terminal result for one agent run.""" - - run_id: str - task: str | None - start_sequence: int - intent: SessionRunIntent = SessionRunIntent.TASK - finish_sequence: int | None = None - result: AgentRunResult | None = None - - def __post_init__(self) -> None: - if not self.run_id: - raise ValueError("run_id must not be empty") - if self.start_sequence <= 0: - raise ValueError("start_sequence must be greater than zero") - if not isinstance(self.intent, SessionRunIntent): - raise TypeError("intent must be a SessionRunIntent") - if self.intent is SessionRunIntent.TASK: - if self.task is None: - raise ValueError("task Run requires a task") - elif self.task is not None: - raise ValueError("Continue Run must not contain a task") - if (self.finish_sequence is None) != (self.result is None): - raise ValueError("finish_sequence and result must be set together") - if ( - self.finish_sequence is not None - and self.finish_sequence <= self.start_sequence - ): - raise ValueError("finish_sequence must be greater than start_sequence") - - @property - def finished(self) -> bool: - return self.result is not None - - -@dataclass(frozen=True, slots=True) -class SessionCompaction: - """One compacted conversation projection with retained audit history.""" - - operation_id: str - sequence: int - summary: SummaryEntry - messages: tuple[AgentMessage, ...] - covered_entry_count: int - - def __post_init__(self) -> None: - if not self.operation_id: - raise ValueError("operation_id must not be empty") - if self.sequence <= 0: - raise ValueError("sequence must be greater than zero") - if not self.messages: - raise ValueError("compaction messages must not be empty") - if self.covered_entry_count < 0: - raise ValueError("covered_entry_count must not be negative") - object.__setattr__(self, "messages", deepcopy(self.messages)) - - -@dataclass(slots=True) -class AgentSession: - """Linear conversation history spanning one or more agent runs.""" - - session_id: str - agent_id: str | None = None - entries: list[SessionMessage] = field(default_factory=list) - runs: list[SessionRun] = field(default_factory=list) - compactions: list[SessionCompaction] = field(default_factory=list) - - def __post_init__(self) -> None: - self.session_id = self.session_id.strip() - if not self.session_id: - raise ValueError("session_id must not be empty") - if self.agent_id is not None: - self.agent_id = self.agent_id.strip() - if not self.agent_id: - raise ValueError("agent_id must not be empty") - self.entries = [ - SessionMessage( - run_id=entry.run_id, - sequence=entry.sequence, - message=entry.message, - ) - for entry in self.entries - ] - self.runs = list(self.runs) - self.compactions = [ - SessionCompaction( - operation_id=compaction.operation_id, - sequence=compaction.sequence, - summary=compaction.summary, - messages=compaction.messages, - covered_entry_count=compaction.covered_entry_count, - ) - for compaction in self.compactions - ] - if any( - compaction.covered_entry_count > len(self.entries) - for compaction in self.compactions - ): - raise ValueError("compaction covers unavailable session entries") - - @property - def messages(self) -> list[AgentMessage]: - """Return an independent history suitable for ``BaseAgent.reset``.""" - - if not self.compactions: - return [deepcopy(entry.message) for entry in self.entries] - - latest = self.compactions[-1] - compacted = [deepcopy(message) for message in latest.messages] - compacted.extend( - deepcopy(entry.message) - for entry in self.entries[latest.covered_entry_count :] - ) - return compacted - - def bind_agent(self, agent_id: str) -> None: - """Bind a new Session to one logical agent identity.""" - - agent_id = agent_id.strip() - if not agent_id: - raise ValueError("agent_id must not be empty") - if self.agent_id is None: - self.agent_id = agent_id - elif self.agent_id != agent_id: - raise ValueError( - f"session {self.session_id!r} belongs to agent " - f"{self.agent_id!r}, not {agent_id!r}" - ) - - def begin_run(self, run_id: str, task: str, sequence: int) -> None: - """Open one run and persist its user task message.""" - - if any(run.run_id == run_id for run in self.runs): - raise ValueError(f"run {run_id!r} already exists") - if any(not run.finished for run in self.runs): - raise ValueError("session already has an unfinished run") - run = SessionRun( - run_id=run_id, - task=task, - start_sequence=sequence, - intent=SessionRunIntent.TASK, - ) - self.runs.append(run) - self.append_message( - run_id, - sequence, - {"role": "user", "content": task}, - ) - - def begin_continue(self, run_id: str, sequence: int) -> None: - """Open one Continue Run without appending a user message.""" - - if any(run.run_id == run_id for run in self.runs): - raise ValueError(f"run {run_id!r} already exists") - if any(not run.finished for run in self.runs): - raise ValueError("session already has an unfinished run") - self.runs.append( - SessionRun( - run_id=run_id, - task=None, - start_sequence=sequence, - intent=SessionRunIntent.CONTINUE, - ) - ) - - def append_message( - self, - run_id: str, - sequence: int, - message: AgentMessage, - ) -> None: - """Append a persistent message to an unfinished run.""" - - run = self._get_run(run_id) - if run.finished: - raise ValueError(f"run {run_id!r} is already finished") - if sequence < self._last_sequence(run_id): - raise ValueError(f"event sequence moved backwards for run {run_id!r}") - self.entries.append( - SessionMessage( - run_id=run_id, - sequence=sequence, - message=message, - ) - ) - - def finish_run( - self, - run_id: str, - sequence: int, - result: AgentRunResult, - ) -> None: - """Attach the existing structured result to an unfinished run.""" - - index = self._run_index(run_id) - run = self.runs[index] - if run.finished: - raise ValueError(f"run {run_id!r} is already finished") - if sequence <= self._last_sequence(run_id): - raise ValueError(f"finish sequence must follow messages for run {run_id!r}") - self.runs[index] = replace( - run, - finish_sequence=sequence, - result=result, - ) - - def apply_compaction( - self, - operation_id: str, - sequence: int, - summary: SummaryEntry, - messages: tuple[AgentMessage, ...], - ) -> None: - """Record a new compacted projection without deleting audit entries.""" - - if any( - compaction.operation_id == operation_id for compaction in self.compactions - ): - raise ValueError(f"compaction operation {operation_id!r} already exists") - self.compactions.append( - SessionCompaction( - operation_id=operation_id, - sequence=sequence, - summary=summary, - messages=messages, - covered_entry_count=len(self.entries), - ) - ) - - def snapshot(self) -> AgentSession: - """Return a detached copy safe for storage and callers.""" - - return AgentSession( - session_id=self.session_id, - agent_id=self.agent_id, - entries=list(self.entries), - runs=list(self.runs), - compactions=list(self.compactions), - ) - - def _get_run(self, run_id: str) -> SessionRun: - return self.runs[self._run_index(run_id)] - - def _run_index(self, run_id: str) -> int: - for index, run in enumerate(self.runs): - if run.run_id == run_id: - return index - raise ValueError(f"unknown run {run_id!r}") - - def _last_sequence(self, run_id: str) -> int: - run = self._get_run(run_id) - return max( - (entry.sequence for entry in self.entries if entry.run_id == run_id), - default=run.start_sequence, - ) diff --git a/src/ejagent/skills/__init__.py b/src/ejagent/skills/__init__.py new file mode 100644 index 0000000..a4691ac --- /dev/null +++ b/src/ejagent/skills/__init__.py @@ -0,0 +1,5 @@ +"""Local Skill discovery for disposable Context projections.""" + +from ejagent.skills.catalog import Skill, SkillCatalog + +__all__ = ["Skill", "SkillCatalog"] diff --git a/src/ejagent/plugins/skill/skill_manager.py b/src/ejagent/skills/catalog.py similarity index 67% rename from src/ejagent/plugins/skill/skill_manager.py rename to src/ejagent/skills/catalog.py index 8161433..340db7a 100644 --- a/src/ejagent/plugins/skill/skill_manager.py +++ b/src/ejagent/skills/catalog.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import re from dataclasses import dataclass from pathlib import Path @@ -12,8 +14,10 @@ _FRONTMATTER_PATTERN = re.compile(r"^---\s*\n(.*?)\n---", re.DOTALL) -@dataclass(frozen=True) +@dataclass(frozen=True, slots=True) class Skill: + """Discovered local Skill resources and compact metadata.""" + name: str skill_md: Path description: str = "" @@ -21,19 +25,15 @@ class Skill: sample_md: Path | None = None -class SkillManager: - """Discover local skills and build explicit context projections.""" +class SkillCatalog: + """Discover local Skills and return provider-neutral instruction text.""" - def __init__(self, skills_root: str | Path): - """ - Args: - skills_root: Root directory whose child folders may contain SKILL.md. - """ + def __init__(self, skills_root: str | Path) -> None: self.skills_root = Path(skills_root) self._skills: dict[str, Skill] = {} self._discovered = False - self._index_message: dict[str, str] | None = None - self._skill_context_messages: dict[str, dict[str, str]] = {} + self._index_content: str | None = None + self._skill_content: dict[str, str] = {} logger.info("Skill registry initialized root=%s", self.skills_root) @property @@ -45,32 +45,26 @@ def discovered(self) -> bool: return self._discovered async def discover(self) -> None: - """Scan child directories containing SKILL.md and build an index.""" + """Scan child directories containing SKILL.md exactly once.""" if self._discovered: return - if not self.skills_root.exists(): raise FileNotFoundError(f"skills root not found: {self.skills_root}") skills: dict[str, Skill] = {} - for child in sorted(self.skills_root.iterdir()): if not child.is_dir(): continue - skill_md = child / "SKILL.md" if not skill_md.exists(): continue - frontmatter = self._read_frontmatter(skill_md) name = self._skill_name(child, frontmatter) template_md = child / "template.md" sample_md = child / "examples" / "sample.md" - if name in skills: raise ValueError(f"duplicate skill name: {name!r}") - skills[name] = Skill( name=name, skill_md=skill_md, @@ -81,28 +75,23 @@ async def discover(self) -> None: self._skills = skills self._discovered = True - self._index_message = None - self._skill_context_messages.clear() - - if not self._skills: - logger.debug( - "Skill discovery completed with no skills root=%s", self.skills_root - ) + self._index_content = None + self._skill_content.clear() + if skills: + logger.info("Discovered %d skill(s): %s", len(skills), list(skills)) else: - logger.info( - "Discovered %d skill(s): %s", - len(self._skills), - list(self._skills.keys()), + logger.debug( + "Skill discovery completed with no skills root=%s", + self.skills_root, ) - def build_index_message(self) -> dict[str, str] | None: - """Return compact skill metadata for the model context.""" + def build_index_content(self) -> str | None: + """Return compact catalog instructions for a ContextView.""" if not self._skills: return None - if self._index_message is not None: - return dict(self._index_message) - + if self._index_content is not None: + return self._index_content lines = [ "Local skills are available. Use a skill when its description matches", "the user's task. An agent with a file-reading tool can read the full", @@ -115,34 +104,23 @@ def build_index_message(self) -> dict[str, str] | None: for skill in self._skills.values(): description = skill.description or "No description provided." lines.extend( - [ + ( f"- name: {skill.name}", f" description: {description}", f" location: {skill.skill_md.resolve()}", - ] + ) ) - self._index_message = { - "role": "system", - "content": "\n".join(lines), - } - return dict(self._index_message) - - def build_skill_context_message( - self, - skill_name: str, - ) -> dict[str, str]: - """Load full skill instructions for provider context.""" - - if skill_name in self._skill_context_messages: - return dict(self._skill_context_messages[skill_name]) - - skill = self.get(skill_name) - message = { - "role": "system", - "content": "\n".join(self._skill_content_parts(skill)), - } - self._skill_context_messages[skill_name] = message - return dict(message) + self._index_content = "\n".join(lines) + return self._index_content + + def build_skill_context_content(self, skill_name: str) -> str: + """Return full instructions and optional resources for one Skill.""" + + if skill_name in self._skill_content: + return self._skill_content[skill_name] + content = "\n".join(self._skill_content_parts(self.get(skill_name))) + self._skill_content[skill_name] = content + return content def get(self, skill_name: str) -> Skill: try: @@ -153,16 +131,13 @@ def get(self, skill_name: str) -> Skill: f"unknown skill {skill_name!r}; available skills: {available}" ) from exc - def select_explicit_skill( - self, - messages: list[dict[str, Any]], - ) -> str | None: - """Return a locally named skill from the latest user message.""" + def select_explicit_skill_from_text(self, task: str) -> str | None: + """Return a local Skill explicitly referenced by task text.""" - task = self._latest_user_task(messages) + if not isinstance(task, str): + raise TypeError("task must be text") if not task: return None - for name in self._skills: if re.search(rf"(? dict[str, Any]: match = _FRONTMATTER_PATTERN.match(text) if match is None: return {} - data = yaml.safe_load(match.group(1)) or {} if not isinstance(data, dict): raise ValueError(f"SKILL.md frontmatter must be a mapping: {skill_md}") return dict(data) @staticmethod - def _skill_name( - skill_dir: Path, - frontmatter: dict[str, Any], - ) -> str: + def _skill_name(skill_dir: Path, frontmatter: dict[str, Any]) -> str: raw_name = frontmatter.get("name") or skill_dir.name if not isinstance(raw_name, str) or not raw_name.strip(): raise ValueError(f"skill name must be a non-empty string: {skill_dir}") return raw_name.strip() - @staticmethod - def _latest_user_task(messages: list[dict[str, Any]]) -> str: - for message in reversed(messages): - if message.get("role") == "user": - content = message.get("content", "") - return content if isinstance(content, str) else str(content) - return "" - @staticmethod def _skill_content_parts(skill: Skill) -> list[str]: parts = [ @@ -212,22 +175,20 @@ def _skill_content_parts(skill: Skill) -> list[str]: "[SKILL.md]", skill.skill_md.read_text(encoding="utf-8").strip(), ] - if skill.template_md is not None: parts.extend( - [ + ( "", "[template.md]", skill.template_md.read_text(encoding="utf-8").strip(), - ] + ) ) - if skill.sample_md is not None: parts.extend( - [ + ( "", "[examples/sample.md]", skill.sample_md.read_text(encoding="utf-8").strip(), - ] + ) ) return parts diff --git a/src/ejagent/storage/__init__.py b/src/ejagent/storage/__init__.py new file mode 100644 index 0000000..34a63e8 --- /dev/null +++ b/src/ejagent/storage/__init__.py @@ -0,0 +1,8 @@ +"""Durable adapters for EJAgent Core contracts.""" + +from ejagent.storage.jsonl import STORE_SCHEMA_VERSION, JsonlSessionStore + +__all__ = [ + "JsonlSessionStore", + "STORE_SCHEMA_VERSION", +] diff --git a/src/ejagent/session/_file_lock.py b/src/ejagent/storage/_file_lock.py similarity index 79% rename from src/ejagent/session/_file_lock.py rename to src/ejagent/storage/_file_lock.py index 50b0ff3..98fbe7a 100644 --- a/src/ejagent/session/_file_lock.py +++ b/src/ejagent/storage/_file_lock.py @@ -11,20 +11,20 @@ from pathlib import Path from typing import Any -from ejagent.session.errors import ( - SessionLockTimeoutError, - SessionStorageError, +from ejagent.contracts.session import ( + SessionStoreError, + SessionStoreLockTimeoutError, ) _fcntl: Any | None try: _fcntl = importlib.import_module("fcntl") -except ImportError: # pragma: no cover - exercised only on non-POSIX platforms +except ImportError: # pragma: no cover - non-POSIX only _fcntl = None -class _SessionLockCancelled(Exception): - """Internal signal used to stop a worker still waiting for a file lock.""" +class _StoreLockCancelled(Exception): + pass @dataclass(slots=True) @@ -37,8 +37,8 @@ class _ProcessLockEntry: _PROCESS_LOCKS_GUARD = threading.Lock() -class SessionFileLock: - """Process-local plus POSIX advisory lock for one stable sidecar file.""" +class StoreFileLock: + """Process-local plus POSIX advisory lock for one Store sidecar file.""" def __init__( self, @@ -57,12 +57,10 @@ def acquire( self, cancelled: threading.Event | None = None, ) -> Iterator[None]: - """Acquire both lock layers within one shared timeout budget.""" - fcntl = _fcntl if fcntl is None: - raise SessionStorageError( - "JsonlSessionStorage cross-process locking requires a POSIX platform" + raise SessionStoreError( + "JsonlSessionStore cross-process locking requires POSIX" ) deadline = None if self.timeout is None else time.monotonic() + self.timeout entry = self._retain_process_lock() @@ -75,14 +73,10 @@ def acquire( self._raise_if_cancelled(cancelled) try: self.path.parent.mkdir(parents=True, exist_ok=True) - descriptor = os.open( - self.path, - os.O_RDWR | os.O_CREAT, - 0o600, - ) + descriptor = os.open(self.path, os.O_RDWR | os.O_CREAT, 0o600) except OSError as exc: - raise SessionStorageError( - f"failed to open Session journal lock {self.path}" + raise SessionStoreError( + f"failed to open SessionStore lock {self.path}" ) from exc self._acquire_file_lock(descriptor, deadline, cancelled, fcntl) file_acquired = True @@ -127,8 +121,8 @@ def _acquire_file_lock( return except OSError as exc: if exc.errno not in {errno.EACCES, errno.EAGAIN}: - raise SessionStorageError( - f"failed to acquire Session journal lock {self.path}" + raise SessionStoreError( + f"failed to acquire SessionStore lock {self.path}" ) from exc self._raise_if_timed_out(deadline) time.sleep(self._next_wait(deadline)) @@ -143,14 +137,14 @@ def _next_wait(self, deadline: float | None) -> float: def _raise_if_timed_out(self, deadline: float | None) -> None: if deadline is not None and time.monotonic() >= deadline: - raise SessionLockTimeoutError( - f"timed out acquiring Session journal lock {self.path}" + raise SessionStoreLockTimeoutError( + f"timed out acquiring SessionStore lock {self.path}" ) @staticmethod def _raise_if_cancelled(cancelled: threading.Event | None) -> None: if cancelled is not None and cancelled.is_set(): - raise _SessionLockCancelled + raise _StoreLockCancelled def _retain_process_lock(self) -> _ProcessLockEntry: with _PROCESS_LOCKS_GUARD: diff --git a/src/ejagent/storage/_legacy.py b/src/ejagent/storage/_legacy.py new file mode 100644 index 0000000..ef6d991 --- /dev/null +++ b/src/ejagent/storage/_legacy.py @@ -0,0 +1,316 @@ +from __future__ import annotations + +import json +from collections.abc import Mapping +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path +from typing import Any + +from ejagent.contracts.runs import RunStatus, StopReason +from ejagent.contracts.session import SessionMigrationError +from ejagent.contracts.usage import RunUsage + +_LEGACY_JOURNAL_VERSION = 1 +_LEGACY_SESSION_VERSION = 1 +_MAIN_BRANCH = "main" +_REMEDIATION = "repair or export the legacy JSONL Session before migration" + + +@dataclass(frozen=True, slots=True) +class LegacyRunResult: + status: RunStatus + stop_reason: StopReason + turns: int + output: str | None + error: str | None + usage: RunUsage + + +@dataclass(slots=True) +class LegacyRun: + run_id: str + result: LegacyRunResult | None = None + + +@dataclass(frozen=True, slots=True) +class LegacySessionData: + session_id: str + agent_id: str | None + entries: tuple[Mapping[str, Any], ...] + runs: tuple[LegacyRun, ...] + + +def load_legacy_session( + root: Path, + session_id: str, +) -> LegacySessionData | None: + """Project the main branch of one legacy JSONL Session journal.""" + + digest = sha256(session_id.encode("utf-8")).hexdigest() + path = root / f"{digest}.jsonl" + try: + content = path.read_bytes() + except FileNotFoundError: + return None + except OSError as exc: + raise _migration_error(f"failed to read legacy Session {path}") from exc + + records: list[Mapping[str, Any]] = [] + expected_revision = 1 + record_ids: set[str] = set() + for index, line in enumerate(content.splitlines(keepends=True)): + if not line.endswith(b"\n"): + if records: + break + raise _migration_error(f"legacy Session {path} has no complete record") + try: + value = json.loads(line) + except (UnicodeDecodeError, json.JSONDecodeError) as exc: + raise _migration_error( + f"legacy Session {path} has invalid JSON at line {index + 1}" + ) from exc + record = _mapping(value, f"journal line {index + 1}") + if ( + _integer(record.get("journal_schema_version"), "journal version") + != _LEGACY_JOURNAL_VERSION + ): + raise _migration_error("unsupported legacy journal schema version") + revision = _integer(record.get("revision"), "journal revision") + if revision != expected_revision: + raise _migration_error("legacy journal revisions are not contiguous") + expected_revision += 1 + record_id = _string(record.get("record_id"), "journal record_id") + if record_id in record_ids: + raise _migration_error("legacy journal repeats a record_id") + record_ids.add(record_id) + if _string(record.get("session_id"), "journal session_id") != session_id: + raise _migration_error("legacy journal contains another session_id") + records.append(record) + + if not records: + return None + return _project_main(records, session_id=session_id) + + +def _project_main( + records: list[Mapping[str, Any]], + *, + session_id: str, +) -> LegacySessionData: + agent_id: str | None = None + entries: list[Mapping[str, Any]] = [] + runs: list[LegacyRun] = [] + previous_id: str | None = None + + for record in records: + if _string(record.get("branch_id"), "journal branch_id") != _MAIN_BRANCH: + continue + parent = record.get("parent_id") + if parent is not None and not isinstance(parent, str): + raise _migration_error("legacy journal parent_id is invalid") + if parent != previous_id: + raise _migration_error("legacy main branch parent chain is broken") + previous_id = _string(record.get("record_id"), "journal record_id") + kind = _string(record.get("type"), "journal type") + data = _mapping(record.get("data"), f"{kind} data") + raw_agent = record.get("agent_id") + record_agent = ( + None if raw_agent is None else _string(raw_agent, "journal agent_id") + ) + + if kind == "checkpoint": + checkpoint = _decode_checkpoint(data, session_id=session_id) + agent_id = checkpoint.agent_id + entries = list(checkpoint.entries) + runs = list(checkpoint.runs) + if agent_id != record_agent: + raise _migration_error("legacy checkpoint agent_id does not match") + continue + if record_agent is None: + raise _migration_error(f"legacy {kind} record has no agent_id") + if agent_id is None: + agent_id = record_agent + elif agent_id != record_agent: + raise _migration_error("legacy journal contains multiple agent IDs") + + if kind == "run_started": + run_id = _string(data.get("run_id"), "run_started run_id") + _require_new_run(runs, run_id) + task = _string(data.get("task"), "run_started task") + runs.append(LegacyRun(run_id)) + entries.append({"role": "user", "content": task}) + elif kind == "run_continued": + run_id = _string(data.get("run_id"), "run_continued run_id") + _require_new_run(runs, run_id) + runs.append(LegacyRun(run_id)) + elif kind == "message_appended": + _require_open_run(runs, data) + entries.append(_mapping(data.get("message"), "message_appended message")) + elif kind == "messages_appended": + _require_open_run(runs, data) + messages = _list(data.get("messages"), "messages_appended messages") + entries.extend( + _mapping(message, f"messages_appended messages[{index}]") + for index, message in enumerate(messages) + ) + elif kind == "steering_applied": + _require_open_run(runs, data) + entries.append( + { + "role": "user", + "content": _string(data.get("content"), "steering content"), + } + ) + elif kind == "run_finished": + run = _require_open_run(runs, data) + run.result = _decode_result(data.get("result"), "run_finished result") + elif kind in {"compaction_applied", "branch_created"}: + continue + else: + raise _migration_error(f"unsupported legacy record type {kind!r}") + + return LegacySessionData( + session_id=session_id, + agent_id=agent_id, + entries=tuple(entries), + runs=tuple(runs), + ) + + +def _decode_checkpoint( + data: Mapping[str, Any], + *, + session_id: str, +) -> LegacySessionData: + document = _mapping(data.get("document"), "checkpoint document") + if ( + _integer(document.get("schema_version"), "session schema_version") + != _LEGACY_SESSION_VERSION + ): + raise _migration_error("unsupported legacy Session schema version") + payload = _mapping(document.get("session"), "checkpoint session") + if _string(payload.get("session_id"), "checkpoint session_id") != session_id: + raise _migration_error("legacy checkpoint session_id does not match") + raw_agent = payload.get("agent_id") + agent_id = None if raw_agent is None else _string(raw_agent, "checkpoint agent_id") + raw_entries = _list(payload.get("entries"), "checkpoint entries") + entries = tuple( + _mapping( + _mapping(item, f"checkpoint entries[{index}]").get("message"), + f"checkpoint entries[{index}].message", + ) + for index, item in enumerate(raw_entries) + ) + raw_runs = _list(payload.get("runs"), "checkpoint runs") + runs: list[LegacyRun] = [] + for index, value in enumerate(raw_runs): + item = _mapping(value, f"checkpoint runs[{index}]") + run_id = _string(item.get("run_id"), f"checkpoint runs[{index}].run_id") + _require_new_run(runs, run_id) + raw_result = item.get("result") + runs.append( + LegacyRun( + run_id, + None + if raw_result is None + else _decode_result(raw_result, f"checkpoint runs[{index}].result"), + ) + ) + return LegacySessionData(session_id, agent_id, entries, tuple(runs)) + + +def _decode_result(value: Any, label: str) -> LegacyRunResult: + item = _mapping(value, label) + return LegacyRunResult( + status=RunStatus(_string(item.get("status"), f"{label}.status")), + stop_reason=StopReason( + _string(item.get("stop_reason"), f"{label}.stop_reason") + ), + turns=_integer(item.get("turns"), f"{label}.turns"), + output=_optional_string(item.get("output"), f"{label}.output"), + error=_optional_string(item.get("error"), f"{label}.error"), + usage=_decode_usage(item.get("usage"), f"{label}.usage"), + ) + + +def _decode_usage(value: Any, label: str) -> RunUsage: + item = _mapping(value, label) + return RunUsage( + input_tokens=_integer(item.get("input_tokens"), f"{label}.input_tokens"), + output_tokens=_integer(item.get("output_tokens"), f"{label}.output_tokens"), + total_tokens=_integer(item.get("total_tokens"), f"{label}.total_tokens"), + request_count=_integer(item.get("request_count"), f"{label}.request_count"), + reported_request_count=_integer( + item.get("reported_request_count"), + f"{label}.reported_request_count", + ), + cache_read_tokens=_optional_integer( + item.get("cache_read_tokens"), + f"{label}.cache_read_tokens", + ), + cache_write_tokens=_optional_integer( + item.get("cache_write_tokens"), + f"{label}.cache_write_tokens", + ), + reasoning_tokens=_optional_integer( + item.get("reasoning_tokens"), + f"{label}.reasoning_tokens", + ), + ) + + +def _require_new_run(runs: list[LegacyRun], run_id: str) -> None: + if any(run.run_id == run_id for run in runs): + raise _migration_error(f"legacy journal repeats run_id {run_id!r}") + if any(run.result is None for run in runs): + raise _migration_error("legacy journal starts a Run before finishing another") + + +def _require_open_run( + runs: list[LegacyRun], + data: Mapping[str, Any], +) -> LegacyRun: + run_id = _string(data.get("run_id"), "record run_id") + for run in runs: + if run.run_id == run_id: + if run.result is not None: + raise _migration_error(f"legacy Run {run_id!r} is already finished") + return run + raise _migration_error(f"legacy journal references unknown Run {run_id!r}") + + +def _mapping(value: Any, label: str) -> Mapping[str, Any]: + if not isinstance(value, Mapping) or any(not isinstance(key, str) for key in value): + raise _migration_error(f"{label} must be an object with string keys") + return value + + +def _list(value: Any, label: str) -> list[Any]: + if not isinstance(value, list): + raise _migration_error(f"{label} must be an array") + return value + + +def _string(value: Any, label: str) -> str: + if not isinstance(value, str): + raise _migration_error(f"{label} must be text") + return value + + +def _optional_string(value: Any, label: str) -> str | None: + return None if value is None else _string(value, label) + + +def _integer(value: Any, label: str) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise _migration_error(f"{label} must be an integer") + return value + + +def _optional_integer(value: Any, label: str) -> int | None: + return None if value is None else _integer(value, label) + + +def _migration_error(message: str) -> SessionMigrationError: + return SessionMigrationError(message, remediation=_REMEDIATION) diff --git a/src/ejagent/storage/codec.py b/src/ejagent/storage/codec.py new file mode 100644 index 0000000..80dfbf1 --- /dev/null +++ b/src/ejagent/storage/codec.py @@ -0,0 +1,453 @@ +from __future__ import annotations + +from collections.abc import Mapping +from datetime import datetime +from typing import Any + +from ejagent.contracts.audit import RunAudit +from ejagent.contracts.conversation import ConversationSnapshot +from ejagent.contracts.json import thaw_json_value +from ejagent.contracts.messages import ( + AssistantMessage, + ConversationMessage, + SystemMessage, + ToolCall, + ToolResultMessage, + UserMessage, +) +from ejagent.contracts.runs import ( + AuditRecord, + FailureCode, + RunDelta, + RunFailure, + RunOutcome, + RunPhase, + RunResult, + RunStatus, + StopReason, +) +from ejagent.contracts.session import ( + SessionCommit, + SessionSnapshot, + SessionStoreSerializationError, +) +from ejagent.contracts.usage import RunUsage + + +def session_commit_to_dict(commit: SessionCommit) -> dict[str, Any]: + """Encode one idempotent commit into a stable JSON-compatible object.""" + + return { + "agent_id": commit.agent_id, + "base": conversation_to_dict(commit.base), + "outcome": _outcome_to_dict(commit.outcome), + } + + +def session_commit_from_dict(value: Any) -> SessionCommit: + """Decode and validate one stable SessionCommit object.""" + + try: + item = _mapping(value, "commit") + return SessionCommit( + agent_id=_string(item.get("agent_id"), "commit.agent_id"), + base=conversation_from_dict(item.get("base"), label="commit.base"), + outcome=_outcome_from_dict(item.get("outcome")), + ) + except SessionStoreSerializationError: + raise + except (TypeError, ValueError) as exc: + raise SessionStoreSerializationError( + f"invalid SessionCommit payload: {exc}" + ) from exc + + +def session_snapshot_to_dict(snapshot: SessionSnapshot) -> dict[str, Any]: + """Encode Harness recovery state without Audit data.""" + + return { + "agent_id": snapshot.agent_id, + "conversation": conversation_to_dict(snapshot.conversation), + "last_result": ( + _result_to_dict(snapshot.last_result) + if snapshot.last_result is not None + else None + ), + } + + +def session_snapshot_from_dict(value: Any) -> SessionSnapshot: + """Decode and validate Harness recovery state.""" + + try: + item = _mapping(value, "snapshot") + raw_result = item.get("last_result") + return SessionSnapshot( + agent_id=_string(item.get("agent_id"), "snapshot.agent_id"), + conversation=conversation_from_dict( + item.get("conversation"), + label="snapshot.conversation", + ), + last_result=( + _result_from_dict(raw_result, label="snapshot.last_result") + if raw_result is not None + else None + ), + ) + except SessionStoreSerializationError: + raise + except (TypeError, ValueError) as exc: + raise SessionStoreSerializationError( + f"invalid SessionSnapshot payload: {exc}" + ) from exc + + +def conversation_to_dict(conversation: ConversationSnapshot) -> dict[str, Any]: + return { + "revision": conversation.revision, + "messages": [_message_to_dict(message) for message in conversation.messages], + } + + +def conversation_from_dict( + value: Any, + *, + label: str = "conversation", +) -> ConversationSnapshot: + item = _mapping(value, label) + messages = _list(item.get("messages"), f"{label}.messages") + return ConversationSnapshot( + revision=_integer(item.get("revision"), f"{label}.revision"), + messages=tuple( + _message_from_dict(message, f"{label}.messages[{index}]") + for index, message in enumerate(messages) + ), + ) + + +def run_audit_to_dict(audit: RunAudit) -> dict[str, Any]: + return { + "result": _result_to_dict(audit.result), + "base_revision": audit.base_revision, + "resulting_revision": audit.resulting_revision, + "committed": audit.committed, + "records": [_audit_record_to_dict(record) for record in audit.records], + "failure": ( + _failure_to_dict(audit.failure) if audit.failure is not None else None + ), + } + + +def run_audit_from_dict(value: Any, *, label: str = "audit") -> RunAudit: + try: + item = _mapping(value, label) + records = _list(item.get("records"), f"{label}.records") + raw_failure = item.get("failure") + return RunAudit( + result=_result_from_dict(item.get("result"), label=f"{label}.result"), + base_revision=_integer( + item.get("base_revision"), + f"{label}.base_revision", + ), + resulting_revision=_integer( + item.get("resulting_revision"), + f"{label}.resulting_revision", + ), + committed=_boolean(item.get("committed"), f"{label}.committed"), + records=tuple( + _audit_record_from_dict(record, f"{label}.records[{index}]") + for index, record in enumerate(records) + ), + failure=( + _failure_from_dict(raw_failure, label=f"{label}.failure") + if raw_failure is not None + else None + ), + ) + except SessionStoreSerializationError: + raise + except (TypeError, ValueError) as exc: + raise SessionStoreSerializationError( + f"invalid RunAudit payload: {exc}" + ) from exc + + +def _message_to_dict(message: ConversationMessage) -> dict[str, Any]: + if isinstance(message, SystemMessage): + return {"type": "system", "content": message.content} + if isinstance(message, UserMessage): + return {"type": "user", "content": message.content} + if isinstance(message, AssistantMessage): + return { + "type": "assistant", + "content": message.content, + "tool_calls": [ + { + "id": call.id, + "name": call.name, + "arguments": thaw_json_value(call.arguments), + } + for call in message.tool_calls + ], + } + if isinstance(message, ToolResultMessage): + return { + "type": "tool_result", + "tool_call_id": message.tool_call_id, + "tool_name": message.tool_name, + "result": thaw_json_value(message.result), + "is_error": message.is_error, + } + raise TypeError(f"unsupported Conversation message {type(message).__name__}") + + +def _message_from_dict(value: Any, label: str) -> ConversationMessage: + item = _mapping(value, label) + kind = _string(item.get("type"), f"{label}.type") + if kind == "system": + return SystemMessage(_string(item.get("content"), f"{label}.content")) + if kind == "user": + return UserMessage(_string(item.get("content"), f"{label}.content")) + if kind == "assistant": + raw_calls = _list(item.get("tool_calls", []), f"{label}.tool_calls") + content = _optional_string(item.get("content"), f"{label}.content") + return AssistantMessage( + content=content, + tool_calls=tuple( + _tool_call_from_dict(call, f"{label}.tool_calls[{index}]") + for index, call in enumerate(raw_calls) + ), + ) + if kind == "tool_result": + return ToolResultMessage( + tool_call_id=_string( + item.get("tool_call_id"), + f"{label}.tool_call_id", + ), + tool_name=_string(item.get("tool_name"), f"{label}.tool_name"), + result=item.get("result"), + is_error=_boolean(item.get("is_error"), f"{label}.is_error"), + ) + raise SessionStoreSerializationError(f"{label}.type has unsupported value {kind!r}") + + +def _tool_call_from_dict(value: Any, label: str) -> ToolCall: + item = _mapping(value, label) + arguments = _mapping(item.get("arguments"), f"{label}.arguments") + return ToolCall( + id=_string(item.get("id"), f"{label}.id"), + name=_string(item.get("name"), f"{label}.name"), + arguments=arguments, + ) + + +def _outcome_to_dict(outcome: RunOutcome) -> dict[str, Any]: + return { + "result": _result_to_dict(outcome.result), + "delta": { + "base_revision": outcome.delta.base_revision, + "messages": [ + _message_to_dict(message) for message in outcome.delta.messages + ], + }, + "audit_records": [ + _audit_record_to_dict(record) for record in outcome.audit_records + ], + "failure": ( + _failure_to_dict(outcome.failure) if outcome.failure is not None else None + ), + } + + +def _outcome_from_dict(value: Any) -> RunOutcome: + item = _mapping(value, "commit.outcome") + delta = _mapping(item.get("delta"), "commit.outcome.delta") + raw_messages = _list( + delta.get("messages"), + "commit.outcome.delta.messages", + ) + raw_records = _list( + item.get("audit_records"), + "commit.outcome.audit_records", + ) + raw_failure = item.get("failure") + return RunOutcome( + result=_result_from_dict( + item.get("result"), + label="commit.outcome.result", + ), + delta=RunDelta( + base_revision=_integer( + delta.get("base_revision"), + "commit.outcome.delta.base_revision", + ), + messages=tuple( + _message_from_dict( + message, + f"commit.outcome.delta.messages[{index}]", + ) + for index, message in enumerate(raw_messages) + ), + ), + audit_records=tuple( + _audit_record_from_dict( + record, + f"commit.outcome.audit_records[{index}]", + ) + for index, record in enumerate(raw_records) + ), + failure=( + _failure_from_dict(raw_failure, label="commit.outcome.failure") + if raw_failure is not None + else None + ), + ) + + +def _result_to_dict(result: RunResult) -> dict[str, Any]: + return { + "run_id": result.run_id, + "status": result.status.value, + "stop_reason": result.stop_reason.value, + "turns": result.turns, + "output": result.output, + "usage": result.usage.to_dict(), + } + + +def _result_from_dict(value: Any, *, label: str) -> RunResult: + item = _mapping(value, label) + return RunResult( + run_id=_string(item.get("run_id"), f"{label}.run_id"), + status=RunStatus(_string(item.get("status"), f"{label}.status")), + stop_reason=StopReason( + _string(item.get("stop_reason"), f"{label}.stop_reason") + ), + turns=_integer(item.get("turns"), f"{label}.turns"), + output=_optional_string(item.get("output"), f"{label}.output"), + usage=_usage_from_dict(item.get("usage"), f"{label}.usage"), + ) + + +def _usage_from_dict(value: Any, label: str) -> RunUsage: + item = _mapping(value, label) + return RunUsage( + input_tokens=_integer(item.get("input_tokens"), f"{label}.input_tokens"), + output_tokens=_integer( + item.get("output_tokens"), + f"{label}.output_tokens", + ), + total_tokens=_integer(item.get("total_tokens"), f"{label}.total_tokens"), + request_count=_integer( + item.get("request_count"), + f"{label}.request_count", + ), + reported_request_count=_integer( + item.get("reported_request_count"), + f"{label}.reported_request_count", + ), + cache_read_tokens=_optional_integer( + item.get("cache_read_tokens"), + f"{label}.cache_read_tokens", + ), + cache_write_tokens=_optional_integer( + item.get("cache_write_tokens"), + f"{label}.cache_write_tokens", + ), + reasoning_tokens=_optional_integer( + item.get("reasoning_tokens"), + f"{label}.reasoning_tokens", + ), + ) + + +def _failure_to_dict(failure: RunFailure) -> dict[str, Any]: + return { + "phase": failure.phase.value, + "code": failure.code.value, + "message": failure.message, + "retryable": failure.retryable, + } + + +def _failure_from_dict(value: Any, *, label: str) -> RunFailure: + item = _mapping(value, label) + return RunFailure( + phase=RunPhase(_string(item.get("phase"), f"{label}.phase")), + code=FailureCode(_string(item.get("code"), f"{label}.code")), + message=_string(item.get("message"), f"{label}.message"), + retryable=_boolean(item.get("retryable"), f"{label}.retryable"), + ) + + +def _audit_record_to_dict(record: AuditRecord) -> dict[str, Any]: + return { + "run_id": record.run_id, + "sequence": record.sequence, + "kind": record.kind, + "occurred_at": record.occurred_at.isoformat(), + "payload": thaw_json_value(record.payload), + } + + +def _audit_record_from_dict(value: Any, label: str) -> AuditRecord: + item = _mapping(value, label) + raw_time = _string(item.get("occurred_at"), f"{label}.occurred_at") + try: + occurred_at = datetime.fromisoformat(raw_time) + except ValueError as exc: + raise SessionStoreSerializationError( + f"{label}.occurred_at must be an ISO datetime" + ) from exc + payload = _mapping(item.get("payload"), f"{label}.payload") + return AuditRecord( + run_id=_string(item.get("run_id"), f"{label}.run_id"), + sequence=_integer(item.get("sequence"), f"{label}.sequence"), + kind=_string(item.get("kind"), f"{label}.kind"), + occurred_at=occurred_at, + payload=payload, + ) + + +def _mapping(value: Any, label: str) -> Mapping[str, Any]: + if not isinstance(value, Mapping): + raise SessionStoreSerializationError(f"{label} must be an object") + if any(not isinstance(key, str) for key in value): + raise SessionStoreSerializationError(f"{label} keys must be strings") + return value + + +def _list(value: Any, label: str) -> list[Any]: + if not isinstance(value, list): + raise SessionStoreSerializationError(f"{label} must be an array") + return value + + +def _string(value: Any, label: str) -> str: + if not isinstance(value, str): + raise SessionStoreSerializationError(f"{label} must be a string") + return value + + +def _optional_string(value: Any, label: str) -> str | None: + if value is None: + return None + return _string(value, label) + + +def _integer(value: Any, label: str) -> int: + if isinstance(value, bool) or not isinstance(value, int): + raise SessionStoreSerializationError(f"{label} must be an integer") + return value + + +def _optional_integer(value: Any, label: str) -> int | None: + if value is None: + return None + return _integer(value, label) + + +def _boolean(value: Any, label: str) -> bool: + if not isinstance(value, bool): + raise SessionStoreSerializationError(f"{label} must be a boolean") + return value diff --git a/src/ejagent/storage/jsonl.py b/src/ejagent/storage/jsonl.py new file mode 100644 index 0000000..946a8d9 --- /dev/null +++ b/src/ejagent/storage/jsonl.py @@ -0,0 +1,616 @@ +from __future__ import annotations + +import asyncio +import json +import math +import os +import threading +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from hashlib import sha256 +from pathlib import Path +from typing import Any, TypeVar + +from ejagent.contracts.audit import AuditReader, RunAudit +from ejagent.contracts.session import ( + SessionCommit, + SessionConflictError, + SessionMigrationError, + SessionSnapshot, + SessionStore, + SessionStoreError, + SessionStoreSerializationError, +) +from ejagent.storage._file_lock import StoreFileLock +from ejagent.storage._legacy import load_legacy_session +from ejagent.storage.codec import ( + run_audit_from_dict, + run_audit_to_dict, + session_commit_from_dict, + session_commit_to_dict, + session_snapshot_from_dict, + session_snapshot_to_dict, +) +from ejagent.storage.migration import LegacySessionMigration, migrate_legacy_session + +STORE_SCHEMA_VERSION = 1 +_ResultT = TypeVar("_ResultT") + + +@dataclass(slots=True) +class _StoreIndex: + snapshot: SessionSnapshot | None + audit: list[RunAudit] + commits: dict[str, tuple[SessionCommit, SessionSnapshot]] + known_run_ids: set[str] + record_count: int = 0 + valid_length: int = 0 + + @classmethod + def empty(cls) -> _StoreIndex: + return cls(None, [], {}, set()) + + +class JsonlSessionStore(SessionStore, AuditReader): + """Durable append-only SessionStore with cross-process compare-and-commit.""" + + def __init__( + self, + root: str | Path, + *, + lock_timeout: float | None = 10.0, + legacy_session_id: str | None = None, + legacy_root: str | Path | None = None, + ) -> None: + if lock_timeout is not None: + if isinstance(lock_timeout, bool) or not isinstance( + lock_timeout, (int, float) + ): + raise TypeError("lock_timeout must be a number or None") + if lock_timeout < 0 or not math.isfinite(lock_timeout): + raise ValueError("lock_timeout must be finite and non-negative") + lock_timeout = float(lock_timeout) + if legacy_session_id is not None: + if not isinstance(legacy_session_id, str) or not legacy_session_id.strip(): + raise ValueError("legacy_session_id must not be empty") + legacy_session_id = legacy_session_id.strip() + self.root = Path(root).expanduser() + self.lock_timeout = lock_timeout + self.legacy_session_id = legacy_session_id + self.legacy_root = ( + Path(legacy_root).expanduser() if legacy_root is not None else self.root + ) + self._lock = asyncio.Lock() + + async def load(self, agent_id: str) -> SessionSnapshot | None: + normalized = self._normalize_agent_id(agent_id) + path = self._path_for(normalized) + async with self._lock: + index = await self._read_index(path, normalized) + if index.snapshot is not None or self.legacy_session_id is None: + return index.snapshot + migration = await self._decode_legacy( + normalized, + session_id=self.legacy_session_id, + root=self.legacy_root, + ) + if migration is None: + return None + return await self._run_locked( + path, + self._seed_sync, + path, + normalized, + migration, + ) + + async def commit(self, commit: SessionCommit) -> SessionSnapshot: + if not isinstance(commit, SessionCommit): + raise TypeError("commit must be a SessionCommit") + normalized = self._normalize_agent_id(commit.agent_id) + path = self._path_for(normalized) + self._validate_json(session_commit_to_dict(commit), label="SessionCommit") + async with self._lock: + return await self._run_locked( + path, + self._commit_sync, + path, + commit, + ) + + async def load_audit(self, agent_id: str) -> tuple[RunAudit, ...]: + normalized = self._normalize_agent_id(agent_id) + if self.legacy_session_id is not None: + await self.load(normalized) + path = self._path_for(normalized) + async with self._lock: + index = await self._read_index(path, normalized) + return tuple(index.audit) + + async def migrate_legacy( + self, + agent_id: str, + *, + session_id: str, + root: str | Path | None = None, + ) -> SessionSnapshot: + """Explicitly import one legacy JSONL Session projection once.""" + + normalized = self._normalize_agent_id(agent_id) + session_id = session_id.strip() + if not session_id: + raise ValueError("session_id must not be empty") + source_root = self._expand_root(root) if root is not None else self.legacy_root + migration = await self._decode_legacy( + normalized, + session_id=session_id, + root=source_root, + ) + if migration is None: + raise SessionMigrationError( + f"legacy Session {session_id!r} does not exist in {source_root}", + remediation="verify the legacy root and session_id", + ) + path = self._path_for(normalized) + async with self._lock: + return await self._run_locked( + path, + self._seed_sync, + path, + normalized, + migration, + ) + + async def _decode_legacy( + self, + agent_id: str, + *, + session_id: str, + root: Path, + ) -> LegacySessionMigration | None: + try: + session = await asyncio.to_thread( + load_legacy_session, + root, + session_id, + ) + except SessionMigrationError: + raise + except Exception as exc: + raise SessionMigrationError( + f"failed to decode legacy Session {session_id!r}: {exc}", + remediation="repair or export the legacy journal before migration", + ) from exc + if session is None: + return None + return migrate_legacy_session(session, agent_id=agent_id) + + def _commit_sync( + self, + path: Path, + commit: SessionCommit, + ) -> SessionSnapshot: + index = self._read_sync(path, commit.agent_id) + existing = index.commits.get(commit.run_id) + if existing is not None: + previous, snapshot = existing + if previous != commit: + raise SessionConflictError( + f"run_id {commit.run_id!r} already identifies a different commit" + ) + return snapshot + if commit.run_id in index.known_run_ids: + raise SessionConflictError( + f"run_id {commit.run_id!r} already exists in migrated Audit" + ) + + current = index.snapshot + if current is None: + if commit.base_revision != 0: + raise SessionConflictError( + f"agent {commit.agent_id!r} has no revision {commit.base_revision}" + ) + current = SessionSnapshot( + agent_id=commit.agent_id, + conversation=commit.base, + ) + self._validate_base(current, commit) + snapshot = self._resulting_snapshot(current, commit) + record = self._record( + sequence=index.record_count + 1, + agent_id=commit.agent_id, + kind="commit", + data={"commit": session_commit_to_dict(commit)}, + ) + self._write_record_sync(path, record, index.valid_length) + return snapshot + + def _seed_sync( + self, + path: Path, + agent_id: str, + migration: LegacySessionMigration, + ) -> SessionSnapshot: + index = self._read_sync(path, agent_id) + if index.snapshot is not None: + if ( + index.snapshot == migration.snapshot + and tuple(index.audit) == migration.audit + ): + return index.snapshot + raise SessionConflictError( + f"agent {agent_id!r} already has durable Core state" + ) + record = self._record( + sequence=1, + agent_id=agent_id, + kind="legacy_seed", + data={ + "source_session_id": migration.source_session_id, + "snapshot": session_snapshot_to_dict(migration.snapshot), + "audit": [run_audit_to_dict(item) for item in migration.audit], + }, + ) + self._validate_json(record, label="legacy migration") + self._write_record_sync(path, record, index.valid_length) + return migration.snapshot + + async def _read_index(self, path: Path, agent_id: str) -> _StoreIndex: + return await self._run_worker(self._read_locked_sync, path, agent_id) + + async def _run_locked( + self, + path: Path, + operation: Callable[..., _ResultT], + *args: Any, + ) -> _ResultT: + return await self._run_worker( + self._call_locked_sync, + path, + operation, + args, + ) + + async def _run_worker( + self, + operation: Callable[..., _ResultT], + *args: Any, + ) -> _ResultT: + cancelled = threading.Event() + worker = asyncio.create_task(asyncio.to_thread(operation, *args, cancelled)) + try: + return await asyncio.shield(worker) + except asyncio.CancelledError: + cancelled.set() + while True: + try: + await asyncio.shield(worker) + break + except asyncio.CancelledError: + continue + except Exception: + break + raise + + def _read_locked_sync( + self, + path: Path, + agent_id: str, + cancelled: threading.Event, + ) -> _StoreIndex: + lock_path = self._lock_path_for(path) + if not path.exists() and not lock_path.exists(): + return _StoreIndex.empty() + return self._call_locked_sync( + path, + self._read_sync, + (path, agent_id), + cancelled, + ) + + def _call_locked_sync( + self, + path: Path, + operation: Callable[..., _ResultT], + args: tuple[Any, ...], + cancelled: threading.Event, + ) -> _ResultT: + lock = StoreFileLock(self._lock_path_for(path), timeout=self.lock_timeout) + with lock.acquire(cancelled): + return operation(*args) + + def _read_sync(self, path: Path, agent_id: str) -> _StoreIndex: + try: + content = path.read_bytes() + except FileNotFoundError: + return _StoreIndex.empty() + except OSError as exc: + raise SessionStoreError(f"failed to read SessionStore {path}") from exc + + index = _StoreIndex.empty() + lines = content.splitlines(keepends=True) + for line_index, encoded_line in enumerate(lines): + line_number = line_index + 1 + if not encoded_line.endswith(b"\n"): + if line_index == len(lines) - 1: + break + try: + raw: Any = json.loads(encoded_line) + except (UnicodeDecodeError, json.JSONDecodeError) as exc: + raise SessionStoreSerializationError( + f"SessionStore {path} has invalid JSON at line {line_number}" + ) from exc + record = self._parse_record(raw, path=path, line_number=line_number) + if record["agent_id"] != agent_id: + raise SessionStoreSerializationError( + f"SessionStore {path} line {line_number} contains agent " + f"{record['agent_id']!r}, expected {agent_id!r}" + ) + if record["sequence"] != index.record_count + 1: + raise SessionStoreSerializationError( + f"SessionStore sequence jumped at {path}:{line_number}" + ) + self._apply_record(index, record, path=path, line_number=line_number) + index.record_count += 1 + index.valid_length += len(encoded_line) + return index + + def _apply_record( + self, + index: _StoreIndex, + record: dict[str, Any], + *, + path: Path, + line_number: int, + ) -> None: + kind = record["type"] + data = record["data"] + if kind == "legacy_seed": + if index.record_count != 0 or index.snapshot is not None: + raise SessionStoreSerializationError( + f"legacy seed must be the first record at {path}:{line_number}" + ) + snapshot = session_snapshot_from_dict(data.get("snapshot")) + if snapshot.agent_id != record["agent_id"]: + raise SessionStoreSerializationError( + f"legacy seed agent mismatch at {path}:{line_number}" + ) + raw_audit = data.get("audit") + if not isinstance(raw_audit, list): + raise SessionStoreSerializationError( + f"legacy seed audit must be an array at {path}:{line_number}" + ) + audit = tuple( + run_audit_from_dict( + item, + label=f"legacy_seed.audit[{audit_index}]", + ) + for audit_index, item in enumerate(raw_audit) + ) + run_ids = [item.run_id for item in audit] + if len(run_ids) != len(set(run_ids)): + raise SessionStoreSerializationError( + f"legacy seed repeats run_id at {path}:{line_number}" + ) + expected_revision = 0 + for item in audit: + if item.base_revision != expected_revision: + raise SessionStoreSerializationError( + f"legacy seed Audit revisions are discontinuous " + f"at {path}:{line_number}" + ) + expected_revision = item.resulting_revision + if audit: + if audit[-1].resulting_revision != snapshot.revision: + raise SessionStoreSerializationError( + f"legacy seed revision mismatch at {path}:{line_number}" + ) + elif snapshot.revision != 0: + raise SessionStoreSerializationError( + f"legacy seed has revision without Audit at {path}:{line_number}" + ) + index.snapshot = snapshot + index.audit.extend(audit) + index.known_run_ids.update(run_ids) + return + if kind != "commit": + raise SessionStoreSerializationError( + f"unsupported record type {kind!r} at {path}:{line_number}" + ) + commit = session_commit_from_dict(data.get("commit")) + if commit.agent_id != record["agent_id"]: + raise SessionStoreSerializationError( + f"commit agent mismatch at {path}:{line_number}" + ) + if commit.run_id in index.known_run_ids: + raise SessionStoreSerializationError( + f"SessionStore repeats run_id {commit.run_id!r} at {path}:{line_number}" + ) + current = index.snapshot + if current is None: + if commit.base_revision != 0: + raise SessionStoreSerializationError( + f"first commit has nonzero base revision at {path}:{line_number}" + ) + current = SessionSnapshot( + agent_id=commit.agent_id, + conversation=commit.base, + ) + try: + self._validate_base(current, commit) + except SessionConflictError as exc: + raise SessionStoreSerializationError( + f"stale stored commit at {path}:{line_number}: {exc}" + ) from exc + snapshot = self._resulting_snapshot(current, commit) + index.snapshot = snapshot + index.audit.append(commit.audit) + index.commits[commit.run_id] = (commit, snapshot) + index.known_run_ids.add(commit.run_id) + + @staticmethod + def _validate_base(current: SessionSnapshot, commit: SessionCommit) -> None: + if current.revision != commit.base_revision: + raise SessionConflictError( + f"agent {commit.agent_id!r} is at revision " + f"{current.revision}, not {commit.base_revision}" + ) + if current.messages != commit.base_messages: + raise SessionConflictError( + f"agent {commit.agent_id!r} Conversation does not match " + f"revision {commit.base_revision}" + ) + + @staticmethod + def _resulting_snapshot( + current: SessionSnapshot, + commit: SessionCommit, + ) -> SessionSnapshot: + return SessionSnapshot( + agent_id=commit.agent_id, + conversation=commit.resulting_conversation, + last_result=( + commit.outcome.result + if commit.advances_revision + else current.last_result + ), + ) + + def _write_record_sync( + self, + path: Path, + record: Mapping[str, Any], + valid_length: int, + ) -> None: + self._validate_json(record, label="SessionStore record") + try: + self.root.mkdir(parents=True, exist_ok=True) + line = ( + json.dumps( + record, + ensure_ascii=False, + separators=(",", ":"), + sort_keys=True, + allow_nan=False, + ) + + "\n" + ).encode("utf-8") + descriptor = os.open(path, os.O_RDWR | os.O_CREAT | os.O_APPEND, 0o600) + try: + current_size = os.fstat(descriptor).st_size + if valid_length < current_size: + os.ftruncate(descriptor, valid_length) + os.lseek(descriptor, 0, os.SEEK_END) + written = os.write(descriptor, line) + if written != len(line): + raise OSError( + f"partial SessionStore write: {written}/{len(line)} bytes" + ) + os.fsync(descriptor) + finally: + os.close(descriptor) + self._fsync_directory() + except OSError as exc: + raise SessionStoreError(f"failed to append SessionStore {path}") from exc + + @staticmethod + def _record( + *, + sequence: int, + agent_id: str, + kind: str, + data: Mapping[str, Any], + ) -> dict[str, Any]: + return { + "store_schema_version": STORE_SCHEMA_VERSION, + "sequence": sequence, + "agent_id": agent_id, + "type": kind, + "data": dict(data), + } + + @staticmethod + def _parse_record( + value: Any, + *, + path: Path, + line_number: int, + ) -> dict[str, Any]: + if not isinstance(value, dict): + raise SessionStoreSerializationError( + f"SessionStore {path} line {line_number} must be an object" + ) + version = value.get("store_schema_version") + if version != STORE_SCHEMA_VERSION: + raise SessionStoreSerializationError( + f"unsupported store_schema_version {version!r} " + f"at {path}:{line_number}; expected {STORE_SCHEMA_VERSION}" + ) + sequence = value.get("sequence") + if isinstance(sequence, bool) or not isinstance(sequence, int) or sequence <= 0: + raise SessionStoreSerializationError( + f"SessionStore sequence must be a positive integer " + f"at {path}:{line_number}" + ) + agent_id = value.get("agent_id") + kind = value.get("type") + data = value.get("data") + if not isinstance(agent_id, str) or not agent_id: + raise SessionStoreSerializationError( + f"SessionStore agent_id is invalid at {path}:{line_number}" + ) + if not isinstance(kind, str) or not kind: + raise SessionStoreSerializationError( + f"SessionStore type is invalid at {path}:{line_number}" + ) + if not isinstance(data, dict): + raise SessionStoreSerializationError( + f"SessionStore data must be an object at {path}:{line_number}" + ) + return { + "sequence": sequence, + "agent_id": agent_id, + "type": kind, + "data": data, + } + + @staticmethod + def _validate_json(value: Any, *, label: str) -> None: + try: + json.dumps(value, allow_nan=False, sort_keys=True) + except (TypeError, ValueError) as exc: + raise SessionStoreSerializationError( + f"{label} is not JSON-compatible: {exc}" + ) from exc + + @staticmethod + def _normalize_agent_id(agent_id: str) -> str: + if not isinstance(agent_id, str): + raise TypeError("agent_id must be a string") + if not agent_id.strip(): + raise ValueError("agent_id must not be empty") + return agent_id + + @staticmethod + def _expand_root(root: str | Path) -> Path: + return Path(root).expanduser() + + def _path_for(self, agent_id: str) -> Path: + digest = sha256(agent_id.encode("utf-8")).hexdigest() + return self.root / f"{digest}.core.jsonl" + + @staticmethod + def _lock_path_for(path: Path) -> Path: + return path.with_name(f"{path.name}.lock") + + def _fsync_directory(self) -> None: + try: + descriptor = os.open(self.root, os.O_RDONLY) + except OSError: + return + try: + os.fsync(descriptor) + except OSError: + pass + finally: + os.close(descriptor) diff --git a/src/ejagent/storage/migration.py b/src/ejagent/storage/migration.py new file mode 100644 index 0000000..f6bbe6e --- /dev/null +++ b/src/ejagent/storage/migration.py @@ -0,0 +1,225 @@ +from __future__ import annotations + +import json +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +from ejagent.contracts.audit import RunAudit +from ejagent.contracts.conversation import ConversationSnapshot +from ejagent.contracts.json import freeze_json_value +from ejagent.contracts.messages import ( + AssistantMessage, + ConversationMessage, + SystemMessage, + ToolCall, + ToolResultMessage, + UserMessage, +) +from ejagent.contracts.runs import ( + FailureCode, + RunFailure, + RunPhase, + RunResult, + RunStatus, +) +from ejagent.contracts.session import SessionMigrationError, SessionSnapshot +from ejagent.storage._legacy import LegacySessionData + +_REMEDIATION = ( + "export the legacy Session with only text, function tool calls, and " + "finished Runs, then retry migration" +) + + +@dataclass(frozen=True, slots=True) +class LegacySessionMigration: + """Typed recovery and Audit values decoded from one legacy Session.""" + + source_session_id: str + snapshot: SessionSnapshot + audit: tuple[RunAudit, ...] + + +def migrate_legacy_session( + session: LegacySessionData, + *, + agent_id: str | None = None, +) -> LegacySessionMigration: + """Convert projected legacy Session data without using compacted views.""" + + target_agent_id = agent_id or session.agent_id + if target_agent_id is None: + raise SessionMigrationError( + f"legacy Session {session.session_id!r} has no agent_id", + remediation="supply agent_id explicitly when importing the Session", + ) + target_agent_id = target_agent_id.strip() + if not target_agent_id: + raise SessionMigrationError( + "migration agent_id is empty", + remediation="supply a non-empty agent_id", + ) + if session.agent_id is not None and session.agent_id != target_agent_id: + raise SessionMigrationError( + f"legacy Session belongs to agent {session.agent_id!r}, not " + f"{target_agent_id!r}", + remediation="use the original agent_id or migrate into a new store key", + ) + + unfinished = [run.run_id for run in session.runs if run.result is None] + if unfinished: + raise SessionMigrationError( + "legacy Session contains unfinished Run(s): " + ", ".join(unfinished), + remediation="finish or explicitly discard those Runs before migration", + ) + + call_names: dict[str, str] = {} + messages: list[ConversationMessage] = [] + for index, message in enumerate(session.entries): + try: + messages.append(_legacy_message(message, call_names=call_names)) + except SessionMigrationError: + raise + except (KeyError, TypeError, ValueError) as exc: + raise SessionMigrationError( + f"legacy message {index} cannot be represented: {exc}", + remediation=_REMEDIATION, + ) from exc + + revision = 0 + audit: list[RunAudit] = [] + last_result: RunResult | None = None + for run in session.runs: + assert run.result is not None + result = RunResult( + run_id=run.run_id, + status=run.result.status, + stop_reason=run.result.stop_reason, + turns=run.result.turns, + output=run.result.output, + usage=run.result.usage, + ) + committed = result.status is RunStatus.COMPLETED + failure = ( + RunFailure( + phase=RunPhase.RUNTIME, + code=FailureCode.RUNTIME_ERROR, + message=run.result.error or result.stop_reason.value, + ) + if result.status is RunStatus.FAILED + else None + ) + next_revision = revision + int(committed) + audit.append( + RunAudit( + result=result, + base_revision=revision, + resulting_revision=next_revision, + committed=committed, + failure=failure, + ) + ) + revision = next_revision + if committed: + last_result = result + + return LegacySessionMigration( + source_session_id=session.session_id, + snapshot=SessionSnapshot( + agent_id=target_agent_id, + conversation=ConversationSnapshot( + revision=revision, + messages=tuple(messages), + ), + last_result=last_result, + ), + audit=tuple(audit), + ) + + +def _legacy_message( + message: Mapping[str, Any], + *, + call_names: dict[str, str], +) -> ConversationMessage: + role = _required_string(message.get("role"), "message.role") + if role == "system": + return SystemMessage(_required_string(message.get("content"), "content")) + if role == "user": + return UserMessage(_required_string(message.get("content"), "content")) + if role == "assistant": + content = message.get("content") + if content is not None and not isinstance(content, str): + raise TypeError("assistant content must be text or null") + raw_calls = message.get("tool_calls", []) + if not isinstance(raw_calls, list): + raise TypeError("assistant tool_calls must be an array") + calls = tuple( + _legacy_tool_call(value, index=index) + for index, value in enumerate(raw_calls) + ) + for call in calls: + call_names[call.id] = call.name + return AssistantMessage(content=content, tool_calls=calls) + if role == "tool": + call_id = _required_string(message.get("tool_call_id"), "tool_call_id") + raw_name = message.get("name") + if raw_name is None: + name = call_names.get(call_id) + if name is None: + raise ValueError( + f"tool result {call_id!r} has no name and no preceding call" + ) + else: + name = _required_string(raw_name, "tool name") + content = freeze_json_value(message.get("content"), label="tool content") + raw_error = message.get("is_error", False) + if not isinstance(raw_error, bool): + raise TypeError("tool is_error must be a boolean") + return ToolResultMessage( + tool_call_id=call_id, + tool_name=name, + result=content, + is_error=raw_error, + ) + raise SessionMigrationError( + f"legacy message role {role!r} is unsupported", + remediation=_REMEDIATION, + ) + + +def _legacy_tool_call(value: Any, *, index: int) -> ToolCall: + if not isinstance(value, Mapping): + raise TypeError(f"tool_calls[{index}] must be an object") + call_id = _required_string(value.get("id"), f"tool_calls[{index}].id") + function = value.get("function") + if function is not None: + if not isinstance(function, Mapping): + raise TypeError(f"tool_calls[{index}].function must be an object") + name = _required_string( + function.get("name"), + f"tool_calls[{index}].function.name", + ) + raw_arguments = function.get("arguments", "{}") + else: + name = _required_string(value.get("name"), f"tool_calls[{index}].name") + raw_arguments = value.get("arguments", {}) + if isinstance(raw_arguments, str): + try: + arguments = json.loads(raw_arguments) + except json.JSONDecodeError as exc: + raise ValueError( + f"tool_calls[{index}] arguments contain invalid JSON" + ) from exc + else: + arguments = raw_arguments + if not isinstance(arguments, Mapping): + raise TypeError(f"tool_calls[{index}] arguments must decode to an object") + return ToolCall(id=call_id, name=name, arguments=arguments) + + +def _required_string(value: Any, label: str) -> str: + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"{label} must be non-empty text") + return value diff --git a/src/ejagent/tools/__init__.py b/src/ejagent/tools/__init__.py new file mode 100644 index 0000000..e840edb --- /dev/null +++ b/src/ejagent/tools/__init__.py @@ -0,0 +1,18 @@ +"""Composable ToolExecutor adapters for EJAgent Core.""" + +from ejagent.tools.executor import ( + CompositeToolExecutor, + FunctionTool, + FunctionToolExecutor, + ToolFunction, +) +from ejagent.tools.mcp import McpManager, McpToolExecutor + +__all__ = [ + "CompositeToolExecutor", + "FunctionTool", + "FunctionToolExecutor", + "McpManager", + "McpToolExecutor", + "ToolFunction", +] diff --git a/src/ejagent/plugins/mcp/mcp_manager.py b/src/ejagent/tools/_mcp_manager.py similarity index 54% rename from src/ejagent/plugins/mcp/mcp_manager.py rename to src/ejagent/tools/_mcp_manager.py index b8ffc36..2f9335c 100644 --- a/src/ejagent/plugins/mcp/mcp_manager.py +++ b/src/ejagent/tools/_mcp_manager.py @@ -1,3 +1,5 @@ +from __future__ import annotations + import json from contextlib import AsyncExitStack from dataclasses import dataclass @@ -16,14 +18,9 @@ class _McpToolRoute: class McpServerManager: - """Load, connect, and route tools across configured MCP services.""" - - def __init__(self, path: str | Path): - """Initialize the manager. + """Connect configured MCP services and route namespaced Tool calls.""" - Args: - path: MCP configuration JSON path. - """ + def __init__(self, path: str | Path) -> None: self.path = Path(path) self._clients_by_service: dict[str, Any] = {} self._tool_routes: dict[str, _McpToolRoute] = {} @@ -31,10 +28,8 @@ def __init__(self, path: str | Path): self._exit_stack: AsyncExitStack | None = None async def startup(self) -> None: - """Connect configured MCP services and index their tools. + """Connect every configured service that can start successfully.""" - One service failure is logged and does not block other services. - """ logger.info("Starting MCP server manager") try: from fastmcp import Client @@ -47,80 +42,40 @@ async def startup(self) -> None: stack = AsyncExitStack() with self.path.open(encoding="utf-8") as config_file: - mcp_configs = MCPConfig.from_dict(json.load(config_file)) - logger.info( - "Loaded %d MCP service config(s)", - len(mcp_configs.mcpServers), - ) - for service_name, server_model in mcp_configs.mcpServers.items(): + configs = MCPConfig.from_dict(json.load(config_file)) + for service_name, server_model in configs.mcpServers.items(): service_stack = AsyncExitStack() try: - logger.info("Connecting MCP service %s", service_name) - mcp_client = await service_stack.enter_async_context( + client = await service_stack.enter_async_context( Client({service_name: server_model.model_dump()}) ) - tools = await mcp_client.list_tools() - self._register_service_tools(service_name, mcp_client, tools) + tools = await client.list_tools() + self._register_service_tools(service_name, client, tools) stack.push_async_callback(service_stack.aclose) - logger.info( - "MCP service %s connected with %d tool(s)", - service_name, - len(tools), - ) - except Exception as e: + except Exception as exc: await service_stack.aclose() logger.error( - "Failed to connect MCP service %s: %s", - service_name, - e, + "Failed to connect MCP service %s: %s", service_name, exc ) self._exit_stack = stack - logger.info( - "MCP server manager started with %d online service(s)", - len(self._clients_by_service), - ) async def shutdown(self) -> None: - """Close all MCP service connections.""" - logger.info("Shutting down MCP server manager") if self._exit_stack is not None: await self._exit_stack.aclose() self._exit_stack = None - for service_name in self._clients_by_service: - logger.info("MCP service %s disconnected", service_name) self._clients_by_service.clear() self._tool_routes.clear() self._openai_tools.clear() - logger.info("MCP server manager stopped") async def call_tool(self, tool_name: str, args: dict[str, object]) -> str: - """Call a routed MCP tool. - - Args: - tool_name: Prefixed tool name, such as "playwright__browser_navigate". - args: Tool arguments. - - Returns: - String representation of the tool result. - - Raises: - ValueError: If the tool is not registered. - """ try: route = self._tool_routes[tool_name] except KeyError as exc: raise ValueError(f"unknown MCP tool: {tool_name}") from exc - result = await route.client.call_tool(route.raw_name, args) return str(result) def get_openai_tools(self) -> list[dict[str, Any]]: - """Return connected service tools in OpenAI tool format. - - Returns: - OpenAI tools. - """ - logger.info("Returning %d OpenAI MCP tool(s)", len(self._openai_tools)) return list(self._openai_tools) def _register_service_tools( @@ -130,17 +85,13 @@ def _register_service_tools( tools: list[Any], ) -> None: routes: dict[str, _McpToolRoute] = {} - openai_tools: list[dict[str, Any]] = [] - + definitions: list[dict[str, Any]] = [] for tool in tools: prefixed_name = f"{service_name}__{tool.name}" if prefixed_name in self._tool_routes or prefixed_name in routes: raise ValueError(f"duplicate MCP tool {prefixed_name!r}") - routes[prefixed_name] = _McpToolRoute( - raw_name=tool.name, - client=client, - ) - openai_tools.append( + routes[prefixed_name] = _McpToolRoute(tool.name, client) + definitions.append( { "type": "function", "function": { @@ -150,7 +101,6 @@ def _register_service_tools( }, } ) - self._clients_by_service[service_name] = client self._tool_routes.update(routes) - self._openai_tools.extend(openai_tools) + self._openai_tools.extend(definitions) diff --git a/src/ejagent/tools/executor.py b/src/ejagent/tools/executor.py new file mode 100644 index 0000000..935eae5 --- /dev/null +++ b/src/ejagent/tools/executor.py @@ -0,0 +1,152 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Iterable, Sequence +from dataclasses import dataclass + +from ejagent.contracts.control import CancellationToken +from ejagent.contracts.lifecycle import ManagedResource +from ejagent.contracts.messages import ToolCall +from ejagent.contracts.tools import ( + ToolDefinition, + ToolExecutionError, + ToolExecutionResult, + ToolExecutor, + ToolProtocolError, +) + +ToolFunction = Callable[ + [ToolCall, CancellationToken], + Awaitable[ToolExecutionResult], +] + + +@dataclass(frozen=True, slots=True) +class FunctionTool: + """Pair one provider-neutral definition with its async implementation.""" + + definition: ToolDefinition + function: ToolFunction + + def __post_init__(self) -> None: + if not isinstance(self.definition, ToolDefinition): + raise TypeError("function tool definition must be a ToolDefinition") + if not callable(self.function): + raise TypeError("function tool implementation must be callable") + + +class FunctionToolExecutor(ToolExecutor): + """Validate and dispatch a fixed set of in-process async tools.""" + + def __init__(self, tools: Iterable[FunctionTool] = ()) -> None: + items = tuple(tools) + if not all(isinstance(item, FunctionTool) for item in items): + raise TypeError("tools must contain FunctionTool values") + names = [item.definition.name for item in items] + if len(names) != len(set(names)): + raise ValueError("function tools must have unique names") + self._tools = {item.definition.name: item for item in items} + self._definitions = tuple(item.definition for item in items) + + @property + def definitions(self) -> Sequence[ToolDefinition]: + return self._definitions + + async def execute( + self, + call: ToolCall, + *, + cancellation: CancellationToken, + ) -> ToolExecutionResult: + cancellation.raise_if_cancelled() + try: + tool = self._tools[call.name] + except KeyError as exc: + raise ToolExecutionError(f"unknown function tool {call.name!r}") from exc + result = await tool.function(call, cancellation) + if not isinstance(result, ToolExecutionResult): + raise ToolProtocolError( + f"function tool {call.name!r} returned " + f"{type(result).__name__}, expected ToolExecutionResult" + ) + return result + + +class CompositeToolExecutor(ToolExecutor): + """Present multiple ToolExecutors as one lifecycle-managed namespace.""" + + def __init__(self, executors: Iterable[ToolExecutor]) -> None: + self._executors = tuple(executors) + if not self._executors: + raise ValueError("executors must not be empty") + self._routes: dict[str, ToolExecutor] = {} + self._definitions: tuple[ToolDefinition, ...] = () + self._started: tuple[ManagedResource, ...] = () + self._is_started = False + self._refresh_routes() + + @property + def definitions(self) -> Sequence[ToolDefinition]: + return self._definitions + + async def start(self) -> None: + if self._is_started: + return + started: list[ManagedResource] = [] + try: + for executor in self._executors: + if isinstance(executor, ManagedResource): + await executor.start() + started.append(executor) + self._refresh_routes() + except BaseException: + for resource in reversed(started): + try: + await resource.shutdown() + except BaseException: + pass + raise + self._started = tuple(started) + self._is_started = True + + async def shutdown(self) -> None: + errors: list[Exception] = [] + for resource in reversed(self._started): + try: + await resource.shutdown() + except Exception as exc: + errors.append(exc) + self._started = () + self._is_started = False + self._refresh_routes() + if errors: + raise ExceptionGroup("ToolExecutor shutdown failed", errors) + + async def execute( + self, + call: ToolCall, + *, + cancellation: CancellationToken, + ) -> ToolExecutionResult: + try: + executor = self._routes[call.name] + except KeyError as exc: + raise ToolExecutionError(f"unknown composed tool {call.name!r}") from exc + return await executor.execute(call, cancellation=cancellation) + + def _refresh_routes(self) -> None: + routes: dict[str, ToolExecutor] = {} + definitions: list[ToolDefinition] = [] + for executor in self._executors: + for definition in executor.definitions: + if not isinstance(definition, ToolDefinition): + raise ToolProtocolError( + "composed executor exposed a non-ToolDefinition value" + ) + if definition.name in routes: + raise ToolProtocolError( + f"duplicate composed tool name {definition.name!r}" + ) + routes[definition.name] = executor + definitions.append(definition) + self._routes = routes + self._definitions = tuple(definitions) diff --git a/src/ejagent/tools/mcp.py b/src/ejagent/tools/mcp.py new file mode 100644 index 0000000..ce35ca6 --- /dev/null +++ b/src/ejagent/tools/mcp.py @@ -0,0 +1,139 @@ +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from pathlib import Path +from typing import Any, Protocol + +from ejagent.contracts.control import CancellationToken, RunCancelledError +from ejagent.contracts.json import JsonObject +from ejagent.contracts.messages import ToolCall +from ejagent.contracts.tools import ( + ToolDefinition, + ToolExecutionError, + ToolExecutionResult, + ToolExecutor, + ToolSemantics, +) +from ejagent.tools._mcp_manager import McpServerManager + + +class McpManager(Protocol): + async def startup(self) -> None: ... + + async def shutdown(self) -> None: ... + + def get_openai_tools(self) -> list[dict[str, Any]]: ... + + async def call_tool(self, tool_name: str, args: dict[str, object]) -> str: ... + + +class McpToolExecutor(ToolExecutor): + """Expose an MCP manager through the provider-neutral ToolExecutor seam.""" + + def __init__( + self, + config_path: str | Path | None = None, + *, + manager: McpManager | None = None, + semantics: Mapping[str, ToolSemantics] | None = None, + ) -> None: + if (config_path is None) == (manager is None): + raise ValueError("provide exactly one of config_path or manager") + if manager is None: + assert config_path is not None + self._manager: McpManager = McpServerManager(config_path) + else: + self._manager = manager + self._semantics = dict(semantics or {}) + if not all( + isinstance(item, ToolSemantics) for item in self._semantics.values() + ): + raise TypeError("MCP semantics values must be ToolSemantics") + self._definitions: tuple[ToolDefinition, ...] = () + self._started = False + + @property + def definitions(self) -> Sequence[ToolDefinition]: + return self._definitions + + async def start(self) -> None: + if self._started: + return + await self._manager.startup() + try: + definitions = tuple( + self._definition_from_openai(item) + for item in self._manager.get_openai_tools() + ) + names = [definition.name for definition in definitions] + if len(names) != len(set(names)): + raise ValueError("MCP manager returned duplicate tool names") + unknown = self._semantics.keys() - set(names) + if unknown: + raise ValueError( + "MCP semantics reference unknown tools: " + + ", ".join(sorted(unknown)) + ) + self._definitions = definitions + except BaseException: + await self._manager.shutdown() + raise + self._started = True + + async def shutdown(self) -> None: + if not self._started: + return + try: + await self._manager.shutdown() + finally: + self._definitions = () + self._started = False + + async def execute( + self, + call: ToolCall, + *, + cancellation: CancellationToken, + ) -> ToolExecutionResult: + if not self._started: + raise ToolExecutionError("MCP ToolExecutor is not started") + if call.name not in {definition.name for definition in self._definitions}: + raise ToolExecutionError(f"unknown MCP tool {call.name!r}") + try: + result = await cancellation.run( + self._manager.call_tool( + call.name, + {key: value for key, value in call.arguments.items()}, + ) + ) + except RunCancelledError: + raise + except Exception as exc: + raise ToolExecutionError(f"MCP tool {call.name!r} failed: {exc}") from exc + return ToolExecutionResult(result) + + def _definition_from_openai(self, value: Any) -> ToolDefinition: + if not isinstance(value, Mapping) or value.get("type") != "function": + raise ValueError("MCP tool must use the OpenAI function shape") + function = value.get("function") + if not isinstance(function, Mapping): + raise ValueError("MCP tool must contain a function object") + name = function.get("name") + description = function.get("description") + parameters = function.get("parameters", {}) + if not isinstance(name, str): + raise ValueError("MCP function name must be text") + if description is not None and not isinstance(description, str): + raise ValueError("MCP function description must be text or null") + if not isinstance(parameters, Mapping): + raise ValueError("MCP function parameters must be an object") + return ToolDefinition( + name=name, + description=description, + input_schema=self._json_object(parameters), + semantics=self._semantics.get(name, ToolSemantics()), + ) + + @staticmethod + def _json_object(value: Mapping[str, Any]) -> JsonObject: + return dict(value) diff --git a/tests/package_artifact_smoke.py b/tests/package_artifact_smoke.py index 10eaba1..6f1ad57 100644 --- a/tests/package_artifact_smoke.py +++ b/tests/package_artifact_smoke.py @@ -3,7 +3,6 @@ from __future__ import annotations import sys -import tarfile import zipfile from pathlib import Path @@ -22,18 +21,7 @@ def main() -> None: raise SystemExit("usage: package_artifact_smoke.py DIST_DIRECTORY") dist_dir = Path(sys.argv[1]) - wheel = only_match(dist_dir, "*.whl") - sdist = only_match(dist_dir, "*.tar.gz") - - with tarfile.open(sdist, "r:gz") as archive: - sdist_names = archive.getnames() - bundled_pi = [ - name for name in sdist_names if "/pi/" in name or name.endswith("/pi") - ] - if bundled_pi: - raise AssertionError("sdist unexpectedly contains the pi reference repo") - if any(name.endswith("/.gitmodules") for name in sdist_names): - raise AssertionError("sdist unexpectedly contains submodule metadata") + wheel = only_match(dist_dir, "ejagent_core-*.whl") with zipfile.ZipFile(wheel) as archive: names = archive.namelist() @@ -52,6 +40,8 @@ def main() -> None: raise AssertionError("wheel metadata has the wrong distribution version") if "Provides-Extra: mcp" not in metadata: raise AssertionError("wheel metadata does not declare the mcp extra") + if "Provides-Extra: anthropic" not in metadata: + raise AssertionError("wheel metadata does not declare the anthropic extra") fastmcp_requirements = [ line for line in metadata.splitlines() @@ -61,16 +51,30 @@ def main() -> None: "extra == 'mcp'" in line for line in fastmcp_requirements ): raise AssertionError("fastmcp must only be required by the mcp extra") - for dependency in ("jsonschema", "referencing"): - requirements = [ - line - for line in metadata.splitlines() - if line.startswith(f"Requires-Dist: {dependency}") - ] - if not requirements or any("extra ==" in line for line in requirements): - raise AssertionError( - f"{dependency} must be declared as a core wheel dependency" - ) + anthropic_requirements = [ + line + for line in metadata.splitlines() + if line.startswith("Requires-Dist: anthropic") + ] + if not anthropic_requirements or not all( + "extra == 'anthropic'" in line for line in anthropic_requirements + ): + raise AssertionError("anthropic must only be required by its extra") + for removed_package in ( + "ejagent/agent/", + "ejagent/handlers/", + "ejagent/middleware/", + "ejagent/plugins/", + "ejagent/session/", + ): + if any(name.startswith(removed_package) for name in names): + raise AssertionError(f"wheel unexpectedly contains {removed_package}") + for removed_module in ( + "ejagent/providers/base.py", + "ejagent/providers/openai.py", + ): + if removed_module in names: + raise AssertionError(f"wheel unexpectedly contains {removed_module}") if __name__ == "__main__": diff --git a/tests/public_api_smoke.py b/tests/public_api_smoke.py index 881692a..1857d04 100644 --- a/tests/public_api_smoke.py +++ b/tests/public_api_smoke.py @@ -16,8 +16,10 @@ def main() -> None: if args.expect_no_mcp and importlib.util.find_spec("fastmcp") is not None: raise AssertionError("fastmcp must not be installed by the core package") + if args.expect_no_mcp and importlib.util.find_spec("anthropic") is not None: + raise AssertionError("anthropic must not be installed by the core package") if importlib.util.find_spec("simagentplg") is not None: - raise AssertionError("the legacy simagentplg import must not be installed") + raise AssertionError("the former import package must not be installed") import ejagent @@ -28,71 +30,39 @@ def main() -> None: raise AssertionError( f"public exports do not resolve: {', '.join(missing_attributes)}" ) - required_exports = { + "AgentHarness", + "AnthropicConfig", + "AnthropicModelPort", + "RuntimeKernel", + "OpenAIModelPort", + "ModelConfig", + "FunctionToolExecutor", + "CompositeToolExecutor", + "McpToolExecutor", + "IdentityContextPipeline", + "DerivedCompactionPipeline", + "SkillsContextPipeline", + "MemorySessionStore", + "JsonlSessionStore", + "SkillCatalog", + } + if missing := required_exports.difference(ejagent.__all__): + raise AssertionError( + f"required public exports are missing: {', '.join(sorted(missing))}" + ) + forbidden_exports = { "BaseAgent", - "BehaviorAction", - "BehaviorDecision", - "BehaviorHook", - "BehaviorHookError", - "TurnSnapshot", "AgentOrchestrator", - "AgentRunResult", "ModelAdapter", "OpenAIModelAdapter", + "BaseHandler", "ToolMiddleware", - "ToolExecutionPolicy", - "RuleBasedToolPolicy", - "ToolPolicyMiddleware", - "ToolPolicyRule", - "ToolPolicyAction", - "ToolPolicyDecision", - "ToolApprover", - "ToolApprovalRequest", - "ToolApprovalDecision", - "ToolSchemaConfigurationError", - "ToolSchemaValidationMiddleware", - "ToolDefinition", - "ToolEffect", - "McpToolHandler", - "McpServerManager", - "SkillManager", "SessionStorage", - "Compactor", - "AutoCompactionPolicy", - "ContextOverflowError", - "ControlInput", - "ControlInputKind", - "ControlReceipt", - "ControlStatus", - "ContinueRejectedReason", - "ContinueRejectedError", - "FollowUpFailurePolicy", - "FollowUpDiscardReason", - "FollowUpHandle", - "FollowUpError", - "FollowUpRejectedError", - "FollowUpDiscardedError", - "CompactionTrigger", - "ModelCompactor", - "JsonlSessionStorage", - "SessionTreeStorage", - "SessionLockTimeoutError", - "SessionRecord", - "SessionBranch", - "SessionCheckout", - "SessionRetry", - "SteeringApplied", - "SteeringDiscarded", - "AgentContinued", - "SessionRunIntent", - "session_to_dict", - "session_from_dict", } - missing_exports = required_exports.difference(ejagent.__all__) - if missing_exports: + if present := forbidden_exports.intersection(ejagent.__all__): raise AssertionError( - f"required public exports are missing: {', '.join(sorted(missing_exports))}" + f"legacy exports remain public: {', '.join(sorted(present))}" ) if not files("ejagent").joinpath("py.typed").is_file(): @@ -103,9 +73,9 @@ def main() -> None: if args.expect_no_mcp: async def check_missing_mcp_message() -> None: - manager = ejagent.McpServerManager("unused.json") + executor = ejagent.McpToolExecutor("unused.json") try: - await manager.startup() + await executor.start() except RuntimeError as exc: if "ejagent-core[mcp]" not in str(exc): raise AssertionError( @@ -116,6 +86,22 @@ async def check_missing_mcp_message() -> None: asyncio.run(check_missing_mcp_message()) + async def check_missing_anthropic_message() -> None: + model = ejagent.AnthropicModelPort( + ejagent.AnthropicConfig("unused", "unused") + ) + try: + await model.start() + except RuntimeError as exc: + if "ejagent-core[anthropic]" not in str(exc): + raise AssertionError( + "missing Anthropic dependency produced no install guidance" + ) from exc + else: + raise AssertionError("Anthropic startup unexpectedly succeeded") + + asyncio.run(check_missing_anthropic_message()) + if __name__ == "__main__": main() diff --git a/tests/test_agent.py b/tests/test_agent.py deleted file mode 100644 index 0554c80..0000000 --- a/tests/test_agent.py +++ /dev/null @@ -1,1603 +0,0 @@ -import asyncio -import inspect -import json -import tempfile -import unittest -from dataclasses import dataclass -from pathlib import Path -from types import SimpleNamespace -from typing import Any -from unittest.mock import patch - -from ejagent import ( - AgentContextBuilder, - AgentOrchestrator, - AgentRunError, - AgentRunResult, - AgentState, - AgentStatus, - BaseAgent, - CancellationToken, - McpToolHandler, - MethodToolHandler, - Middleware, - ModelAdapter, - ModelConfig, - OpenAIModelAdapter, - RunStatus, - RuntimePolicy, - StepOutcome, - StopReason, - ToolCallContext, - ToolControl, - ToolMiddleware, - ToolNext, -) -from ejagent.agent.base import ( - DEFAULT_SYSTEM_PROMPT, - EXPLICIT_FINISH_PROTOCOL_PROMPT, - TOOL_PROTOCOL_PROMPT, -) -from ejagent.agent.orchestrator import TOOL_COMPLETION_RETRY_PROMPT - -TEST_CONFIG = ModelConfig( - model="test-model", - api_key="test-key", - base_url="https://example.invalid", -) - -ECHO_TOOL = { - "type": "function", - "function": { - "name": "echo", - "description": "Return the provided text.", - "parameters": { - "type": "object", - "properties": {"text": {"type": "string"}}, - "required": ["text"], - }, - }, -} - -DONE_TOOL = { - "type": "function", - "function": { - "name": "done", - "description": "Finish the current test task.", - "parameters": { - "type": "object", - "properties": {"summary": {"type": "string"}}, - "required": ["summary"], - }, - }, -} - - -@dataclass -class FakeFunction: - name: str - arguments: str - - -@dataclass -class FakeToolCall: - id: str - function: FakeFunction - type: str = "function" - - @property - def name(self) -> str: - return self.function.name - - @property - def arguments(self) -> str: - return self.function.arguments - - def to_agent_message(self) -> dict[str, Any]: - return { - "id": self.id, - "type": self.type, - "function": { - "name": self.name, - "arguments": self.arguments, - }, - } - - -class FakeMessage: - def __init__( - self, - content: str | None = None, - tool_calls: list[FakeToolCall] | None = None, - ) -> None: - self.content = content - self.tool_calls = tool_calls - - def to_agent_message(self) -> dict[str, Any]: - message: dict[str, Any] = { - "role": "assistant", - "content": self.content, - } - if self.tool_calls: - message["tool_calls"] = [ - tool_call.to_agent_message() for tool_call in self.tool_calls - ] - return message - - -class FakeCompletions: - def __init__(self, responses: list[FakeMessage]) -> None: - self.responses = list(responses) - self.calls: list[dict[str, Any]] = [] - - async def create(self, **kwargs: Any) -> FakeMessage: - self.calls.append(kwargs) - return self.responses.pop(0) - - -class FakeModelAdapter(ModelAdapter): - def __init__(self, responses: list[FakeMessage]) -> None: - self.completions = FakeCompletions(responses) - self.started = 0 - self.stopped = 0 - - async def startup(self) -> None: - self.started += 1 - - async def shutdown(self) -> None: - self.stopped += 1 - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> FakeMessage: - return await self.completions.create( - messages=context.llm_messages, - tools=context.tools or None, - ) - - -class EchoHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((ECHO_TOOL,)) - self.started = 0 - self.stopped = 0 - self.task_starts = 0 - self.calls = 0 - - async def startup(self) -> None: - self.started += 1 - - async def shutdown(self) -> None: - self.stopped += 1 - - async def on_task_start(self) -> None: - self.task_starts += 1 - - async def do_echo( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - self.calls += 1 - text = arguments.get("text") - if not isinstance(text, str): - return StepOutcome({"status": "error", "error": "text is required"}) - return StepOutcome({"status": "success", "text": text}) - - -class DoneHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((DONE_TOOL,)) - - async def do_done( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - return StepOutcome( - {"summary": arguments.get("summary", "")}, - control=ToolControl.COMPLETE, - ) - - -class FakeMcpManager: - def __init__(self) -> None: - self.started = False - self.stopped = False - self.calls: list[tuple[str, dict[str, Any]]] = [] - - async def startup(self) -> None: - self.started = True - - async def shutdown(self) -> None: - self.stopped = True - - def get_openai_tools(self) -> list[dict[str, Any]]: - return [ - { - "type": "function", - "function": { - "name": "demo__lookup", - "description": "Lookup a value.", - "parameters": {"type": "object", "properties": {}}, - }, - } - ] - - async def call_tool( - self, - tool_name: str, - arguments: dict[str, Any], - ) -> str: - self.calls.append((tool_name, arguments)) - return "mcp-result" - - -class RecordingToolMiddleware(ToolMiddleware): - def __init__( - self, - outcome: StepOutcome | None = None, - *, - enabled: bool = True, - ) -> None: - super().__init__(enabled=enabled) - self.outcome = outcome - self.started = 0 - self.stopped = 0 - self.task_starts = 0 - self.calls: list[tuple[str, dict[str, Any]]] = [] - self.contexts: list[ToolCallContext] = [] - - async def startup(self) -> None: - self.started += 1 - - async def shutdown(self) -> None: - self.stopped += 1 - - async def on_task_start(self) -> None: - self.task_starts += 1 - - async def __call__( - self, - context: ToolCallContext, - call_next: ToolNext, - ) -> StepOutcome: - self.calls.append((context.tool_name, dict(context.arguments))) - self.contexts.append(context) - if self.outcome is not None: - return self.outcome - return await call_next(context) - - -def make_agent( - *, - model: ModelAdapter | None = None, - **kwargs: Any, -) -> BaseAgent: - """Build a core agent while keeping repetitive test setup compact.""" - - return BaseAgent( - model or FakeModelAdapter([]), - **kwargs, - ) - - -class AgentTests(unittest.IsolatedAsyncioTestCase): - async def test_base_agent_composes_public_orchestrator(self) -> None: - agent = make_agent( - agent_id="orchestrated", - model=FakeModelAdapter([]), - ) - - self.assertIsInstance(agent.orchestrator, AgentOrchestrator) - self.assertIs(agent.orchestrator.state, agent.state) - - async def test_middleware_base_class_is_exported_with_standard_spelling( - self, - ) -> None: - self.assertTrue(issubclass(ToolMiddleware, Middleware)) - - async def test_base_agent_public_api_uses_model_and_runtime_policy(self) -> None: - parameters = inspect.signature(BaseAgent).parameters - - self.assertIn("model", parameters) - self.assertIn("runtime_policy", parameters) - self.assertNotIn("client", parameters) - self.assertNotIn("max_steps", parameters) - self.assertNotIn( - "has_handler_tools", - inspect.signature(AgentOrchestrator).parameters, - ) - - async def test_base_agent_owns_model_adapter_lifecycle(self) -> None: - model = FakeModelAdapter([]) - agent = BaseAgent(model, agent_id="model-lifecycle") - - await agent.startup() - await agent.startup() - await agent.shutdown() - await agent.shutdown() - - self.assertEqual(model.started, 1) - self.assertEqual(model.stopped, 1) - - async def test_model_config_reads_chat_model_from_env(self) -> None: - with ( - patch("ejagent.providers.openai.load_dotenv"), - patch.dict( - "os.environ", - { - "CHAT_MODEL": "chat-model", - "MODEL_API_KEY": "key", - "MODEL_URL": "https://model.example", - "LLM_TIMEOUT": "12", - "LLM_TEMPERATURE": "0.2", - }, - clear=True, - ), - ): - config = ModelConfig.from_env() - - self.assertEqual(config.model, "chat-model") - self.assertEqual(config.api_key, "key") - self.assertEqual(config.base_url, "https://model.example") - self.assertEqual(config.timeout, 12) - self.assertEqual(config.temperature, 0.2) - - async def test_openai_adapter_converts_provider_response(self) -> None: - class OpenAICompletions: - def __init__(self) -> None: - self.calls: list[dict[str, Any]] = [] - - async def create(self, **kwargs: Any) -> Any: - self.calls.append(kwargs) - return SimpleNamespace( - choices=[ - SimpleNamespace( - message=FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-1", - function=FakeFunction( - "lookup", - '{"query": "value"}', - ), - ) - ] - ) - ) - ] - ) - - completions = OpenAICompletions() - client = SimpleNamespace(chat=SimpleNamespace(completions=completions)) - adapter = OpenAIModelAdapter( - TEST_CONFIG, - client=client, # type: ignore[arg-type] - ) - context = AgentContextBuilder().build(AgentState()) - - message = await adapter.complete(context) - - self.assertIsNone(message.content) - self.assertEqual(message.tool_calls[0].id, "call-1") - self.assertEqual(message.tool_calls[0].name, "lookup") - self.assertEqual( - message.tool_calls[0].arguments, - '{"query": "value"}', - ) - self.assertEqual(completions.calls[0]["model"], "test-model") - - async def test_agents_use_independent_models_and_messages(self) -> None: - first_model = FakeModelAdapter([FakeMessage("first")]) - second_model = FakeModelAdapter([FakeMessage("second")]) - first = make_agent( - agent_id="first", - model=first_model, - ) - second = make_agent( - agent_id="second", - model=second_model, - ) - - await first.runtime(task="only first sees this") - await second.runtime(task="only second sees this") - - self.assertIsNot(first.model, second.model) - self.assertNotEqual(first.messages, second.messages) - self.assertFalse( - any( - message.get("content") == "only first sees this" - for message in second.messages - ) - ) - - async def test_runtime_keeps_memory_and_reset_clears_it(self) -> None: - client = FakeModelAdapter([FakeMessage("one"), FakeMessage("two")]) - agent = make_agent( - agent_id="memory", - model=client, - ) - - await agent.runtime(task="first task") - await agent.runtime(task="second task") - - second_call_messages = client.completions.calls[1]["messages"] - self.assertTrue( - any(message.get("content") == "one" for message in second_call_messages) - ) - - agent.reset([{"role": "user", "content": "seed"}]) - self.assertEqual( - agent.messages, - [ - {"role": "system", "content": agent.system_prompt}, - {"role": "user", "content": "seed"}, - ], - ) - - async def test_runtime_calls_are_serialized_per_agent(self) -> None: - agent = make_agent( - agent_id="serialized-runtime", - model=FakeModelAdapter([]), - ) - active = 0 - maximum = 0 - - async def run_loop( - cancellation: CancellationToken, - steering: object | None = None, - *, - run_id: str, - ) -> AgentRunResult: - nonlocal active, maximum - active += 1 - maximum = max(maximum, active) - try: - await asyncio.sleep(0.01) - return AgentRunResult( - status=RunStatus.COMPLETED, - stop_reason=StopReason.TEXT_RESPONSE, - turns=1, - output="done", - ) - finally: - active -= 1 - - agent.orchestrator._run_loop = run_loop # type: ignore[method-assign] - - results = await asyncio.gather( - agent.runtime(task="first"), - agent.runtime(task="second"), - ) - - self.assertEqual(results, ["done", "done"]) - self.assertEqual(maximum, 1) - - async def test_default_system_prompt_is_plain_chat_prompt(self) -> None: - agent = make_agent( - agent_id="default-prompt", - model=FakeModelAdapter([]), - ) - - self.assertEqual(agent.system_prompt, DEFAULT_SYSTEM_PROMPT) - self.assertEqual( - agent.messages, - [{"role": "system", "content": DEFAULT_SYSTEM_PROMPT}], - ) - - async def test_tool_mode_injects_tool_protocol_explicitly(self) -> None: - agent = make_agent( - agent_id="tool-protocol", - system_prompt="You are a custom coding agent.", - handlers=[DoneHandler()], - model=FakeModelAdapter([]), - ) - - self.assertEqual( - agent.messages, - [ - {"role": "system", "content": "You are a custom coding agent."}, - {"role": "system", "content": TOOL_PROTOCOL_PROMPT}, - ], - ) - - async def test_plain_chat_has_no_tool_protocol_message(self) -> None: - agent = make_agent( - agent_id="plain-protocol", - system_prompt="You are a plain chat agent.", - model=FakeModelAdapter([]), - ) - - self.assertEqual( - agent.messages, - [ - {"role": "system", "content": "You are a plain chat agent."}, - ], - ) - - async def test_explicit_finish_policy_injects_completion_protocol(self) -> None: - agent = make_agent( - agent_id="explicit-finish-protocol", - handlers=[DoneHandler()], - runtime_policy=RuntimePolicy(require_explicit_finish=True), - model=FakeModelAdapter([]), - ) - - self.assertEqual( - agent.messages[-1], - {"role": "system", "content": EXPLICIT_FINISH_PROTOCOL_PROMPT}, - ) - - async def test_tools_do_not_require_explicit_finish_by_default(self) -> None: - agent = make_agent( - agent_id="tool-text-completion", - handlers=[EchoHandler()], - model=FakeModelAdapter([FakeMessage("completed with plain text")]), - ) - - result = await agent.runtime(task="answer after using tools if needed") - - self.assertEqual(result, "completed with plain text") - - async def test_run_returns_structured_result(self) -> None: - agent = make_agent( - agent_id="structured-run", - model=FakeModelAdapter([FakeMessage("done")]), - ) - - result = await agent.run(task="complete") - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(result.stop_reason, StopReason.TEXT_RESPONSE) - self.assertEqual(result.output, "done") - self.assertEqual(result.turns, 1) - - async def test_plain_chat_returns_text_without_retry_prompt(self) -> None: - client = FakeModelAdapter([FakeMessage("plain text")]) - agent = make_agent( - agent_id="plain-retry", - runtime_policy=RuntimePolicy(max_steps=2), - model=client, - ) - - result = await agent.runtime(task="plain chat") - - self.assertEqual(result, "plain text") - self.assertFalse( - any( - message.get("content") == TOOL_COMPLETION_RETRY_PROMPT - for message in agent.messages - ) - ) - - async def test_plain_chat_rejects_empty_completion(self) -> None: - client = FakeModelAdapter([FakeMessage(None)]) - agent = make_agent( - agent_id="plain-empty", - model=client, - ) - - with self.assertRaisesRegex(RuntimeError, "empty content"): - await agent.runtime(task="plain chat") - - async def test_plain_chat_raises_when_step_limit_is_exhausted(self) -> None: - unknown_call = FakeToolCall( - id="call-1", - function=FakeFunction("unknown", "{}"), - ) - client = FakeModelAdapter( - [ - FakeMessage(tool_calls=[unknown_call]), - FakeMessage(tool_calls=[unknown_call]), - ] - ) - agent = make_agent( - agent_id="plain-step-limit", - runtime_policy=RuntimePolicy(max_steps=2), - model=client, - ) - - with self.assertRaisesRegex(RuntimeError, "did not finish within 2"): - await agent.runtime(task="plain chat") - - async def test_context_builder_can_filter_internal_context(self) -> None: - class FilteringContextBuilder(AgentContextBuilder): - def convert_to_llm_messages( - self, - messages: list[dict[str, Any]], - ) -> list[dict[str, Any]]: - return [ - dict(message) - for message in messages - if not message.get("exclude_from_llm") - ] - - client = FakeModelAdapter([FakeMessage("visible")]) - agent = make_agent( - agent_id="context", - model=client, - context_builder=FilteringContextBuilder(), - ) - agent.messages.append( - { - "role": "user", - "content": "internal note", - "exclude_from_llm": True, - } - ) - - result = await agent.runtime(task="real task") - - self.assertEqual(result, "visible") - sent_messages = client.completions.calls[0]["messages"] - self.assertFalse( - any(message.get("content") == "internal note" for message in sent_messages) - ) - self.assertTrue( - any(message.get("content") == "real task" for message in sent_messages) - ) - - async def test_agent_state_records_completed_task(self) -> None: - agent = make_agent( - agent_id="state-complete", - model=FakeModelAdapter([FakeMessage("done")]), - ) - - result = await agent.runtime(task="finish task") - - self.assertEqual(result, "done") - self.assertEqual(agent.state.task, "finish task") - self.assertEqual(agent.state.status, AgentStatus.COMPLETED) - self.assertEqual(agent.state.turn, 1) - self.assertEqual(agent.state.result, "done") - self.assertIsNone(agent.state.error) - - async def test_agent_state_records_failed_task(self) -> None: - agent = make_agent( - agent_id="state-failed", - model=FakeModelAdapter([FakeMessage(None)]), - ) - - with self.assertRaisesRegex(RuntimeError, "empty content"): - await agent.runtime(task="fail task") - - self.assertEqual(agent.state.task, "fail task") - self.assertEqual(agent.state.status, AgentStatus.FAILED) - self.assertIsNone(agent.state.result) - self.assertIn("empty content", agent.state.error or "") - - async def test_context_builder_does_not_mutate_agent_state(self) -> None: - state = AgentState(messages=[{"role": "user", "content": "saved"}]) - - result = AgentContextBuilder().build( - state, - tools=[ - { - "type": "function", - "function": {"name": "temporary_tool"}, - } - ], - transient_messages=[{"role": "system", "content": "temporary"}], - ) - - self.assertEqual( - state.messages, - [{"role": "user", "content": "saved"}], - ) - self.assertEqual( - result.llm_messages, - ( - {"role": "user", "content": "saved"}, - {"role": "system", "content": "temporary"}, - ), - ) - self.assertEqual( - result.tools, - ( - { - "type": "function", - "function": {"name": "temporary_tool"}, - }, - ), - ) - - async def test_skill_metadata_is_injected_without_llm_router(self) -> None: - with tempfile.TemporaryDirectory() as temp_dir: - skills_dir = Path(temp_dir) - skill_dir = skills_dir / "release_notes" - skill_dir.mkdir() - (skill_dir / "SKILL.md").write_text( - "\n".join( - [ - "---", - "name: release_notes", - "description: Write user-facing release notes.", - "---", - "", - "# Release Notes", - "", - "This full instruction should load only on demand.", - ] - ), - encoding="utf-8", - ) - client = FakeModelAdapter([FakeMessage("ok")]) - agent = make_agent( - agent_id="skills-index", - skills_dir=skills_dir, - model=client, - ) - - with patch.dict("os.environ", {}, clear=True): - result = await agent.runtime(task="Write release notes") - - self.assertEqual(result, "ok") - sent_messages = client.completions.calls[0]["messages"] - joined = "\n".join(str(message.get("content", "")) for message in sent_messages) - self.assertIn("Available skills:", joined) - self.assertIn("name: release_notes", joined) - self.assertIn("description: Write user-facing release notes.", joined) - self.assertIn( - f"location: {(skill_dir / 'SKILL.md').resolve()}", - joined, - ) - self.assertNotIn("This full instruction should load only on demand.", joined) - - async def test_skill_discovery_is_not_repeated_between_runtime_calls(self) -> None: - with tempfile.TemporaryDirectory() as temp_dir: - skills_dir = Path(temp_dir) - skill_dir = skills_dir / "release_notes" - skill_dir.mkdir() - (skill_dir / "SKILL.md").write_text( - "\n".join( - [ - "---", - "name: release_notes", - "description: Write release notes.", - "---", - ] - ), - encoding="utf-8", - ) - client = FakeModelAdapter([FakeMessage("first"), FakeMessage("second")]) - agent = make_agent( - agent_id="skills-once", - skills_dir=skills_dir, - model=client, - ) - - await agent.runtime(task="first") - new_skill_dir = skills_dir / "hot_loaded" - new_skill_dir.mkdir() - (new_skill_dir / "SKILL.md").write_text( - "\n".join( - [ - "---", - "name: hot_loaded", - "description: Should not appear without refresh.", - "---", - ] - ), - encoding="utf-8", - ) - await agent.runtime(task="second") - - self.assertIsNone(client.completions.calls[1]["tools"]) - second_messages = client.completions.calls[1]["messages"] - joined = "\n".join( - str(message.get("content", "")) for message in second_messages - ) - self.assertIn("name: release_notes", joined) - self.assertNotIn("name: hot_loaded", joined) - - async def test_explicit_skill_name_loads_full_skill_context(self) -> None: - with tempfile.TemporaryDirectory() as temp_dir: - skills_dir = Path(temp_dir) - skill_dir = skills_dir / "release_notes" - examples_dir = skill_dir / "examples" - examples_dir.mkdir(parents=True) - (skill_dir / "SKILL.md").write_text( - "\n".join( - [ - "---", - "name: release_notes", - "description: Write user-facing release notes.", - "---", - "", - "# Release Notes", - "", - "FULL SKILL RULES", - ] - ), - encoding="utf-8", - ) - (skill_dir / "template.md").write_text("TEMPLATE BODY", encoding="utf-8") - (examples_dir / "sample.md").write_text("SAMPLE BODY", encoding="utf-8") - client = FakeModelAdapter([FakeMessage("ok")]) - agent = make_agent( - agent_id="skills-load", - skills_dir=skills_dir, - model=client, - ) - - await agent.runtime(task="$release_notes Write release notes") - - sent_messages = client.completions.calls[0]["messages"] - joined = "\n".join(str(message.get("content", "")) for message in sent_messages) - self.assertIn('You are executing the local skill "release_notes".', joined) - self.assertIn("FULL SKILL RULES", joined) - self.assertIn("TEMPLATE BODY", joined) - self.assertIn("SAMPLE BODY", joined) - - async def test_skills_do_not_register_an_internal_tool(self) -> None: - with tempfile.TemporaryDirectory() as temp_dir: - skills_dir = Path(temp_dir) - skill_dir = skills_dir / "release_notes" - skill_dir.mkdir() - (skill_dir / "SKILL.md").write_text( - "\n".join( - [ - "---", - "name: release_notes", - "description: Write user-facing release notes.", - "---", - "", - "# Release Notes", - "", - "FULL SKILL RULES", - ] - ), - encoding="utf-8", - ) - client = FakeModelAdapter([FakeMessage("release notes")]) - agent = make_agent( - agent_id="skills-resource", - skills_dir=skills_dir, - model=client, - ) - - result = await agent.runtime(task="Write release notes") - - self.assertEqual(result, "release notes") - self.assertEqual(agent.tools, []) - self.assertIsNone(client.completions.calls[0]["tools"]) - self.assertFalse( - any(message.get("role") == "tool" for message in agent.messages) - ) - - async def test_full_skill_context_file_reads_are_cached(self) -> None: - with tempfile.TemporaryDirectory() as temp_dir: - skills_dir = Path(temp_dir) - skill_dir = skills_dir / "release_notes" - examples_dir = skill_dir / "examples" - examples_dir.mkdir(parents=True) - (skill_dir / "SKILL.md").write_text( - "\n".join( - [ - "---", - "name: release_notes", - "description: Write user-facing release notes.", - "---", - "", - "# Release Notes", - "", - "FULL SKILL RULES", - ] - ), - encoding="utf-8", - ) - (skill_dir / "template.md").write_text("TEMPLATE BODY", encoding="utf-8") - (examples_dir / "sample.md").write_text("SAMPLE BODY", encoding="utf-8") - client = FakeModelAdapter([FakeMessage("release notes")]) - agent = make_agent( - agent_id="skills-cache", - skills_dir=skills_dir, - model=client, - ) - await agent.startup() - - original_read_text = Path.read_text - read_counts: dict[Path, int] = {} - - def count_read_text(path: Path, *args: Any, **kwargs: Any) -> str: - read_counts[path] = read_counts.get(path, 0) + 1 - return original_read_text(path, *args, **kwargs) - - with patch.object(Path, "read_text", autospec=True) as read_text: - read_text.side_effect = count_read_text - result = await agent.runtime(task="$release_notes Write release notes") - - self.assertEqual(result, "release notes") - self.assertEqual(read_counts[skill_dir / "SKILL.md"], 1) - self.assertEqual(read_counts[skill_dir / "template.md"], 1) - self.assertEqual(read_counts[examples_dir / "sample.md"], 1) - - async def test_agent_id_is_normalized_and_read_only(self) -> None: - agent = make_agent( - agent_id=" assistant ", - model=FakeModelAdapter([]), - ) - - self.assertEqual(agent.agent_id, "assistant") - with self.assertRaises(AttributeError): - agent.agent_id = "renamed" # type: ignore[misc] - - async def test_empty_agent_id_is_rejected(self) -> None: - for agent_id in ("", " "): - with self.subTest(agent_id=agent_id): - with self.assertRaisesRegex(ValueError, "agent_id"): - make_agent( - agent_id=agent_id, - model=FakeModelAdapter([]), - ) - - async def test_agent_id_is_required(self) -> None: - with self.assertRaisesRegex(TypeError, "agent_id"): - BaseAgent(FakeModelAdapter([])) # type: ignore[call-arg] - - async def test_method_handler_dispatches_atomic_tool(self) -> None: - handler = EchoHandler() - result = await handler.dispatch("echo", {"text": "hello"}) - self.assertEqual( - result.data, - {"status": "success", "text": "hello"}, - ) - - async def test_tool_mode_uses_only_explicit_handlers(self) -> None: - echo = EchoHandler() - agent = make_agent( - agent_id="tools", - handlers=[echo], - model=FakeModelAdapter([]), - ) - - self.assertEqual(agent.handlers, [echo]) - self.assertEqual( - [tool["function"]["name"] for tool in agent.tools], - ["echo"], - ) - - async def test_tool_disabled_mode_does_not_add_bash_handler(self) -> None: - agent = make_agent( - agent_id="chat", - model=FakeModelAdapter([]), - ) - - self.assertEqual(agent.handlers, []) - self.assertEqual(agent.tools, []) - - async def test_duplicate_tool_names_fail_during_startup(self) -> None: - model = FakeModelAdapter([]) - agent = make_agent( - agent_id="duplicate-tools", - handlers=[EchoHandler(), EchoHandler()], - model=model, - ) - - with self.assertRaisesRegex(ValueError, "duplicate tool 'echo'"): - await agent.startup() - - self.assertEqual(model.started, 1) - self.assertEqual(model.stopped, 1) - - async def test_unknown_tool_has_explicit_error(self) -> None: - agent = make_agent( - agent_id="unknown-tool", - handlers=[EchoHandler()], - model=FakeModelAdapter([]), - ) - - with self.assertRaisesRegex(KeyError, "unknown tool 'missing'"): - await agent.dispatch("missing", {}) - - async def test_invalid_json_tool_arguments_are_returned_to_model(self) -> None: - client = FakeModelAdapter( - [ - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-1", - function=FakeFunction("echo", "[1, 2]"), - ) - ] - ), - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-2", - function=FakeFunction( - "done", - '{"summary": "recovered"}', - ), - ) - ] - ), - ] - ) - agent = make_agent( - agent_id="invalid-json", - handlers=[EchoHandler(), DoneHandler()], - model=client, - ) - - result = await agent.runtime(task="use echo") - - self.assertEqual(json.loads(result or "")["summary"], "recovered") - tool_message = next( - message for message in agent.messages if message["role"] == "tool" - ) - payload = json.loads(tool_message["content"]) - self.assertEqual(payload["status"], "error") - self.assertIn("JSON object", payload["error"]) - - async def test_tool_mode_corrects_plain_text_until_exit_tool(self) -> None: - client = FakeModelAdapter( - [ - FakeMessage("premature answer"), - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-1", - function=FakeFunction( - "done", - '{"summary": "done"}', - ), - ) - ] - ), - ] - ) - agent = make_agent( - agent_id="finish-required", - handlers=[DoneHandler()], - runtime_policy=RuntimePolicy(require_explicit_finish=True), - model=client, - ) - - result = await agent.runtime(task="complete the task") - - self.assertEqual(json.loads(result or "")["summary"], "done") - second_messages = client.completions.calls[1]["messages"] - self.assertTrue( - any( - message.get("role") == "system" - and message.get("content") == TOOL_COMPLETION_RETRY_PROMPT - for message in second_messages - ) - ) - - async def test_tool_mode_retries_plain_text_with_diagnostic_limit(self) -> None: - client = FakeModelAdapter( - [ - FakeMessage("first plain text"), - FakeMessage("second plain text"), - FakeMessage("third plain text"), - ] - ) - agent = make_agent( - agent_id="tool-retry-limit", - handlers=[DoneHandler()], - runtime_policy=RuntimePolicy( - max_steps=5, - require_explicit_finish=True, - ), - model=client, - ) - - with self.assertRaisesRegex(RuntimeError, "without a completing tool call"): - await agent.runtime(task="complete the task") - - retry_prompts = [ - message["content"] - for call in client.completions.calls[1:] - for message in call["messages"] - if message.get("role") == "system" - and message.get("content", "").startswith( - "Explicit-finish mode requires a completing tool call" - ) - ] - self.assertEqual(len(retry_prompts), 2) - self.assertEqual(len(set(retry_prompts)), 2) - self.assertIn("Retry 2/3", retry_prompts[1]) - self.assertFalse( - any( - message.get("content", "").startswith( - "Explicit-finish mode requires a completing tool call" - ) - for message in agent.state.messages - ) - ) - - async def test_tool_lifecycle_is_not_duplicated_in_logs(self) -> None: - client = FakeModelAdapter( - [ - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-1", - function=FakeFunction( - "echo", - '{"text": "hello"}', - ), - ) - ] - ), - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-2", - function=FakeFunction( - "done", - '{"summary": "logged"}', - ), - ) - ] - ), - ] - ) - agent = make_agent( - agent_id="log-tools", - handlers=[EchoHandler(), DoneHandler()], - model=client, - ) - - with self.assertLogs("log-tools", level="INFO") as captured: - result = await agent.runtime(task="use echo then finish") - - self.assertEqual(json.loads(result or "")["summary"], "logged") - logs = "\n".join(captured.output) - self.assertIn("registered tools: done, echo", logs) - self.assertNotIn("Calling tool", logs) - self.assertNotIn("Tool echo completed", logs) - self.assertNotIn("Tool done completed", logs) - - async def test_task_start_hook_runs_for_each_tool_task(self) -> None: - handler = EchoHandler() - client = FakeModelAdapter( - [ - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-1", - function=FakeFunction( - "done", - '{"summary": "first"}', - ), - ) - ] - ), - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-2", - function=FakeFunction( - "done", - '{"summary": "second"}', - ), - ) - ] - ), - ] - ) - agent = make_agent( - agent_id="task-hooks", - handlers=[handler, DoneHandler()], - model=client, - ) - - await agent.runtime(task="first") - await agent.runtime(task="second") - - self.assertEqual(handler.task_starts, 2) - - async def test_startup_before_runtime_does_not_repeat_startup(self) -> None: - with tempfile.TemporaryDirectory() as temp_dir: - skills_dir = Path(temp_dir) - skill_dir = skills_dir / "release_notes" - skill_dir.mkdir() - (skill_dir / "SKILL.md").write_text( - "\n".join( - [ - "---", - "name: release_notes", - "description: Write release notes.", - "---", - ] - ), - encoding="utf-8", - ) - handler = EchoHandler() - client = FakeModelAdapter( - [ - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-1", - function=FakeFunction( - "done", - '{"summary": "done"}', - ), - ) - ] - ) - ] - ) - agent = make_agent( - agent_id="startup-once", - handlers=[handler, DoneHandler()], - skills_dir=skills_dir, - model=client, - ) - - await agent.startup() - result = await agent.runtime(task="finish") - - self.assertEqual(json.loads(result or "")["summary"], "done") - self.assertEqual(handler.started, 1) - self.assertEqual(handler.task_starts, 1) - - async def test_tool_middleware_lifecycle_and_low_risk_execution(self) -> None: - handler = EchoHandler() - middleware = RecordingToolMiddleware() - agent = make_agent( - agent_id="middleware-low-risk", - handlers=[handler], - middlewares=[middleware], - model=FakeModelAdapter([]), - ) - - outcome = await agent.dispatch("echo", {"text": "allowed"}) - await agent.shutdown() - - self.assertEqual(outcome.data, {"status": "success", "text": "allowed"}) - self.assertEqual(handler.calls, 1) - self.assertEqual(middleware.started, 1) - self.assertEqual(middleware.stopped, 1) - self.assertEqual(middleware.calls, [("echo", {"text": "allowed"})]) - - async def test_tool_middleware_task_start_hook_runs_for_each_task(self) -> None: - middleware = RecordingToolMiddleware() - client = FakeModelAdapter( - [ - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-1", - function=FakeFunction("done", '{"summary": "first"}'), - ) - ] - ), - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-2", - function=FakeFunction("done", '{"summary": "second"}'), - ) - ] - ), - ] - ) - agent = make_agent( - agent_id="middleware-task-hooks", - handlers=[DoneHandler()], - middlewares=[middleware], - model=client, - ) - - await agent.runtime(task="first") - await agent.runtime(task="second") - - self.assertEqual(middleware.task_starts, 2) - - async def test_middlewares_do_not_run_without_handlers(self) -> None: - middleware = RecordingToolMiddleware() - client = FakeModelAdapter([FakeMessage("plain chat")]) - agent = make_agent( - agent_id="middleware-disabled", - middlewares=[middleware], - model=client, - ) - - result = await agent.runtime(task="hello") - - self.assertEqual(result, "plain chat") - self.assertEqual(middleware.started, 0) - self.assertEqual(middleware.task_starts, 0) - self.assertEqual(middleware.calls, []) - - async def test_tool_middleware_rejection_has_distinct_run_status(self) -> None: - handler = EchoHandler() - middleware = RecordingToolMiddleware( - StepOutcome( - { - "status": "rejected", - "tool": "echo", - "reason": "human rejected tool execution", - }, - control=ToolControl.REJECT, - ) - ) - client = FakeModelAdapter( - [ - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-1", - function=FakeFunction("echo", '{"text": "blocked"}'), - ) - ] - ) - ] - ) - agent = make_agent( - agent_id="middleware-reject", - handlers=[handler], - middlewares=[middleware], - model=client, - ) - - result = await agent.run(task="try echo") - - self.assertIsInstance(result, AgentRunResult) - self.assertEqual(result.status, RunStatus.REJECTED) - self.assertEqual(result.stop_reason, StopReason.TOOL_REJECTED) - payload = json.loads(result.output or "") - self.assertEqual(payload["status"], "rejected") - self.assertEqual(payload["tool"], "echo") - self.assertEqual(handler.calls, 0) - self.assertEqual(middleware.contexts[0].tool_call_id, "call-1") - self.assertIs(middleware.contexts[0].state, agent.state) - - async def test_runtime_raises_structured_error_for_rejection(self) -> None: - middleware = RecordingToolMiddleware( - StepOutcome({"status": "rejected"}, control=ToolControl.REJECT) - ) - agent = make_agent( - agent_id="middleware-reject-compatibility", - handlers=[EchoHandler()], - middlewares=[middleware], - model=FakeModelAdapter( - [ - FakeMessage( - tool_calls=[ - FakeToolCall( - id="call-1", - function=FakeFunction("echo", '{"text": "blocked"}'), - ) - ] - ) - ] - ), - ) - - with self.assertRaises(AgentRunError) as raised: - await agent.runtime(task="try echo") - - self.assertEqual(raised.exception.result.status, RunStatus.REJECTED) - - async def test_middlewares_short_circuit_in_order(self) -> None: - first = RecordingToolMiddleware( - StepOutcome({"status": "blocked"}, control=ToolControl.REJECT) - ) - second = RecordingToolMiddleware() - handler = EchoHandler() - agent = make_agent( - agent_id="middleware-chain", - handlers=[handler], - middlewares=[first, second], - model=FakeModelAdapter([]), - ) - - outcome = await agent.dispatch("echo", {"text": "blocked"}) - - self.assertEqual(outcome.data, {"status": "blocked"}) - self.assertEqual(first.calls, [("echo", {"text": "blocked"})]) - self.assertEqual(second.calls, []) - self.assertEqual(handler.calls, 0) - - async def test_tool_middlewares_wrap_handler_in_declaration_order(self) -> None: - events: list[str] = [] - - class TracingToolMiddleware(ToolMiddleware): - def __init__(self, name: str) -> None: - super().__init__(name=name) - - async def __call__( - self, - context: ToolCallContext, - call_next: ToolNext, - ) -> StepOutcome: - events.append(f"{self.name}:before") - try: - return await call_next(context) - finally: - events.append(f"{self.name}:after") - - agent = make_agent( - agent_id="middleware-wrap-order", - handlers=[EchoHandler()], - middlewares=[ - TracingToolMiddleware("first"), - TracingToolMiddleware("second"), - ], - model=FakeModelAdapter([]), - ) - - outcome = await agent.dispatch("echo", {"text": "wrapped"}) - - self.assertEqual(outcome.data, {"status": "success", "text": "wrapped"}) - self.assertEqual( - events, - ["first:before", "second:before", "second:after", "first:after"], - ) - - async def test_outer_tool_middleware_observes_inner_short_circuit(self) -> None: - events: list[str] = [] - - class ObservingToolMiddleware(ToolMiddleware): - async def __call__( - self, - context: ToolCallContext, - call_next: ToolNext, - ) -> StepOutcome: - events.append("outer:before") - try: - return await call_next(context) - finally: - events.append("outer:after") - - handler = EchoHandler() - blocking = RecordingToolMiddleware( - StepOutcome({"status": "blocked"}, control=ToolControl.REJECT) - ) - agent = make_agent( - agent_id="middleware-observe-short-circuit", - handlers=[handler], - middlewares=[ObservingToolMiddleware(), blocking], - model=FakeModelAdapter([]), - ) - - outcome = await agent.dispatch("echo", {"text": "blocked"}) - - self.assertEqual(outcome.data, {"status": "blocked"}) - self.assertEqual(events, ["outer:before", "outer:after"]) - self.assertEqual(handler.calls, 0) - - async def test_outer_tool_middleware_observes_handler_exception(self) -> None: - events: list[str] = [] - - class FailingEchoHandler(EchoHandler): - async def do_echo( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - self.calls += 1 - raise RuntimeError("handler failed") - - class ObservingToolMiddleware(ToolMiddleware): - async def __call__( - self, - context: ToolCallContext, - call_next: ToolNext, - ) -> StepOutcome: - events.append("outer:before") - try: - return await call_next(context) - finally: - events.append("outer:after") - - handler = FailingEchoHandler() - agent = make_agent( - agent_id="middleware-observe-error", - handlers=[handler], - middlewares=[ObservingToolMiddleware()], - model=FakeModelAdapter([]), - ) - - with self.assertRaisesRegex(RuntimeError, "handler failed"): - await agent.dispatch("echo", {"text": "fail"}) - - self.assertEqual(events, ["outer:before", "outer:after"]) - self.assertEqual(handler.calls, 1) - - async def test_started_tool_middleware_still_shuts_down_if_disabled(self) -> None: - middleware = RecordingToolMiddleware() - agent = make_agent( - agent_id="middleware-toggle", - handlers=[EchoHandler()], - middlewares=[middleware], - model=FakeModelAdapter([]), - ) - - await agent.startup() - middleware.enabled = False - await agent.shutdown() - - self.assertEqual(middleware.started, 1) - self.assertEqual(middleware.stopped, 1) - - async def test_third_identical_tool_call_fails_before_execution(self) -> None: - handler = EchoHandler() - repeated_call = FakeToolCall( - id="call", - function=FakeFunction("echo", '{"text": "same"}'), - ) - client = FakeModelAdapter( - [ - FakeMessage(tool_calls=[repeated_call]), - FakeMessage(tool_calls=[repeated_call]), - FakeMessage(tool_calls=[repeated_call]), - ] - ) - agent = make_agent( - agent_id="repeat-guard", - handlers=[handler], - runtime_policy=RuntimePolicy(max_steps=3), - model=client, - ) - - with self.assertRaisesRegex(RuntimeError, "consecutive times"): - await agent.runtime(task="repeat") - - self.assertEqual(handler.calls, 2) - - async def test_tool_mode_raises_when_finish_is_never_called(self) -> None: - client = FakeModelAdapter([FakeMessage("not finished"), FakeMessage(None)]) - agent = make_agent( - agent_id="step-limit", - handlers=[DoneHandler()], - runtime_policy=RuntimePolicy( - max_steps=2, - require_explicit_finish=True, - ), - model=client, - ) - - with self.assertRaisesRegex(RuntimeError, "did not finish within 2"): - await agent.runtime(task="never finish") - - async def test_plain_chat_mode_does_not_start_handlers(self) -> None: - client = FakeModelAdapter([FakeMessage("plain chat")]) - agent = make_agent( - agent_id="plain-chat", - model=client, - ) - - result = await agent.runtime(task="hello") - - self.assertEqual(result, "plain chat") - self.assertIsNone(client.completions.calls[0]["tools"]) - - async def test_mcp_handler_uses_the_same_handler_contract(self) -> None: - manager = FakeMcpManager() - handler = McpToolHandler(manager=manager) - - await handler.startup() - outcome = await handler.dispatch("demo__lookup", {"query": "value"}) - await handler.shutdown() - - self.assertEqual(outcome.data, "mcp-result") - self.assertEqual( - manager.calls, - [("demo__lookup", {"query": "value"})], - ) - self.assertTrue(manager.started) - self.assertTrue(manager.stopped) - - async def test_mcp_handler_requires_one_explicit_source(self) -> None: - with self.assertRaisesRegex(ValueError, "config_path is required"): - McpToolHandler() - with self.assertRaisesRegex(ValueError, "either config_path or manager"): - McpToolHandler("mcp.json", manager=FakeMcpManager()) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_agent_cancellation.py b/tests/test_agent_cancellation.py deleted file mode 100644 index e86ae11..0000000 --- a/tests/test_agent_cancellation.py +++ /dev/null @@ -1,345 +0,0 @@ -import asyncio -import json -import unittest -from typing import Any - -from ejagent import ( - AgentEvent, - AgentFinished, - AgentStarted, - AgentStatus, - AssistantMessage, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - McpToolHandler, - MemorySessionStorage, - MethodToolHandler, - ModelAdapter, - ModelToolCall, - RunStatus, - SessionRecorder, - StepOutcome, - StopReason, - ToolCompleted, - ToolStarted, - TurnCompleted, - TurnStarted, -) - -WAIT_TOOL = { - "type": "function", - "function": { - "name": "wait", - "description": "Wait until the operation is cancelled.", - "parameters": {"type": "object", "properties": {}}, - }, -} - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class BlockingThenCompleteModel(ModelAdapter): - def __init__(self) -> None: - self.started = asyncio.Event() - self.interrupted = asyncio.Event() - self.calls = 0 - self.tokens: list[CancellationToken] = [] - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - assert cancellation is not None - self.tokens.append(cancellation) - self.calls += 1 - if self.calls > 1: - return AssistantMessage(content="reused") - - self.started.set() - try: - await asyncio.Future() - finally: - self.interrupted.set() - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - return self.responses.pop(0) - - -class BlockingToolHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((WAIT_TOOL,)) - self.started = asyncio.Event() - self.interrupted = asyncio.Event() - self.calls = 0 - self.token: CancellationToken | None = None - - async def do_wait( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - assert cancellation is not None - self.calls += 1 - self.token = cancellation - self.started.set() - try: - await asyncio.Future() - finally: - self.interrupted.set() - - -class BlockingMcpManager: - def __init__(self) -> None: - self.started = asyncio.Event() - self.interrupted = asyncio.Event() - - async def startup(self) -> None: - pass - - async def shutdown(self) -> None: - pass - - def get_openai_tools(self) -> list[dict[str, Any]]: - return [ - { - "type": "function", - "function": { - "name": "mcp__wait", - "description": "Wait in an MCP server.", - "parameters": {"type": "object", "properties": {}}, - }, - } - ] - - async def call_tool( - self, - tool_name: str, - arguments: dict[str, Any], - ) -> str: - self.started.set() - try: - await asyncio.Future() - finally: - self.interrupted.set() - - -class SlowFinishSink(RecordingSink): - def __init__(self) -> None: - super().__init__() - self.finish_started = asyncio.Event() - self.release = asyncio.Event() - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - if isinstance(event.payload, AgentFinished): - self.finish_started.set() - await self.release.wait() - - -def payloads(sink: RecordingSink) -> list[object]: - return [event.payload for event in sink.events] - - -class AgentCancellationTests(unittest.IsolatedAsyncioTestCase): - async def test_abort_interrupts_model_and_agent_can_be_reused(self) -> None: - model = BlockingThenCompleteModel() - sink = RecordingSink() - agent = BaseAgent(model, agent_id="abort-model", event_sink=sink) - - self.assertFalse(agent.abort()) - first_run = asyncio.create_task(agent.run(task="block")) - await model.started.wait() - - self.assertTrue(agent.abort("stopped by user")) - self.assertFalse(agent.abort("duplicate")) - first_result = await first_run - await agent.wait_for_idle() - - self.assertTrue(model.interrupted.is_set()) - self.assertEqual(first_result.status, RunStatus.CANCELLED) - self.assertEqual(first_result.stop_reason, StopReason.EXTERNAL_ABORT) - self.assertEqual(first_result.error, "stopped by user") - self.assertEqual(agent.state.status, AgentStatus.CANCELLED) - self.assertEqual( - [type(payload) for payload in payloads(sink)], - [ - AgentStarted, - TurnStarted, - TurnCompleted, - AgentFinished, - ], - ) - - second_result = await agent.run(task="run again") - - self.assertEqual(second_result.status, RunStatus.COMPLETED) - self.assertEqual(second_result.output, "reused") - self.assertEqual(model.calls, 2) - self.assertIsNot(model.tokens[0], model.tokens[1]) - self.assertTrue(model.tokens[0].cancelled) - self.assertFalse(model.tokens[1].cancelled) - self.assertFalse(agent.abort()) - - async def test_tool_abort_closes_all_tool_calls_and_session_run(self) -> None: - first_call = ModelToolCall( - id="call-1", - name="wait", - arguments="{}", - ) - second_call = ModelToolCall( - id="call-2", - name="wait", - arguments="{}", - ) - handler = BlockingToolHandler() - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="cancel-tool", storage=storage) - observer = RecordingSink() - agent = BaseAgent( - SequenceModel([AssistantMessage(tool_calls=(first_call, second_call))]), - agent_id="abort-tool", - handlers=[handler], - event_sink=CompositeAgentEventSink([recorder, observer]), - ) - - run = asyncio.create_task(agent.run(task="use tools")) - await handler.started.wait() - self.assertTrue(agent.abort()) - result = await run - session = await recorder.load() - - self.assertEqual(result.status, RunStatus.CANCELLED) - self.assertEqual(result.stop_reason, StopReason.EXTERNAL_ABORT) - self.assertEqual(handler.calls, 1) - self.assertTrue(handler.interrupted.is_set()) - assert handler.token is not None - self.assertTrue(handler.token.cancelled) - - tool_started = [ - payload - for payload in payloads(observer) - if isinstance(payload, ToolStarted) - ] - tool_completed = [ - payload - for payload in payloads(observer) - if isinstance(payload, ToolCompleted) - ] - self.assertEqual(len(tool_started), 2) - self.assertEqual(len(tool_completed), 2) - self.assertTrue(all(payload.result.cancelled for payload in tool_completed)) - self.assertEqual( - [message["role"] for message in agent.messages[-3:]], - ["assistant", "tool", "tool"], - ) - for message in agent.messages[-2:]: - self.assertEqual( - json.loads(message["content"])["status"], - "cancelled", - ) - - assert session is not None - self.assertTrue(session.runs[0].finished) - assert session.runs[0].result is not None - self.assertEqual(session.runs[0].result.status, RunStatus.CANCELLED) - self.assertEqual( - [message["role"] for message in session.messages], - ["user", "assistant", "tool", "tool"], - ) - - async def test_wait_for_idle_includes_agent_finished_sink(self) -> None: - model = BlockingThenCompleteModel() - sink = SlowFinishSink() - agent = BaseAgent(model, agent_id="idle-settlement", event_sink=sink) - run = asyncio.create_task(agent.run(task="block")) - await model.started.wait() - - agent.abort() - await sink.finish_started.wait() - idle = asyncio.create_task(agent.wait_for_idle()) - await asyncio.sleep(0) - - self.assertFalse(idle.done()) - self.assertFalse(run.done()) - sink.release.set() - await idle - result = await run - self.assertEqual(result.status, RunStatus.CANCELLED) - - async def test_mcp_only_agent_can_abort_an_active_tool_call(self) -> None: - manager = BlockingMcpManager() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=( - ModelToolCall( - id="mcp-call", - name="mcp__wait", - arguments="{}", - ), - ) - ) - ] - ), - agent_id="abort-mcp", - handlers=[McpToolHandler(manager=manager)], - ) - - run = asyncio.create_task(agent.run(task="call MCP")) - await manager.started.wait() - agent.abort() - result = await run - - self.assertEqual(result.status, RunStatus.CANCELLED) - self.assertEqual(result.stop_reason, StopReason.EXTERNAL_ABORT) - self.assertTrue(manager.interrupted.is_set()) - self.assertEqual(agent.messages[-1]["role"], "tool") - - async def test_caller_task_cancellation_still_emits_terminal_events( - self, - ) -> None: - model = BlockingThenCompleteModel() - sink = RecordingSink() - agent = BaseAgent(model, agent_id="caller-cancel", event_sink=sink) - run = asyncio.create_task(agent.run(task="block")) - await model.started.wait() - - run.cancel() - with self.assertRaises(asyncio.CancelledError): - await run - await agent.wait_for_idle() - - event_payloads = payloads(sink) - self.assertIsInstance(event_payloads[-2], TurnCompleted) - self.assertIsInstance(event_payloads[-1], AgentFinished) - finish = event_payloads[-1] - assert isinstance(finish, AgentFinished) - self.assertEqual(finish.result.status, RunStatus.CANCELLED) - self.assertEqual(finish.result.stop_reason, StopReason.EXTERNAL_ABORT) - self.assertEqual(agent.state.status, AgentStatus.CANCELLED) - self.assertTrue(model.interrupted.is_set()) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_agent_continue.py b/tests/test_agent_continue.py deleted file mode 100644 index 38082ba..0000000 --- a/tests/test_agent_continue.py +++ /dev/null @@ -1,390 +0,0 @@ -import asyncio -import tempfile -import unittest -from typing import Any - -from ejagent import ( - AgentContinued, - AgentEvent, - AgentFinished, - AgentRunResult, - AgentSession, - AgentStarted, - AssistantMessage, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - ContinueRejectedError, - ContinueRejectedReason, - FollowUpDiscardedError, - FollowUpDiscardReason, - JsonlSessionStorage, - ModelAdapter, - RunStatus, - SessionRecorder, - SessionRecordKind, - SessionRunIntent, - StopReason, -) - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - self.contexts: list[tuple[dict[str, Any], ...]] = [] - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.contexts.append(context.agent_messages) - return self.responses.pop(0) - - -class BlockingModel(SequenceModel): - def __init__( - self, - responses: list[AssistantMessage], - *, - block_call: int, - ) -> None: - super().__init__(responses) - self.block_call = block_call - self.started = asyncio.Event() - self.release = asyncio.Event() - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - index = len(self.contexts) - self.contexts.append(context.agent_messages) - if index == self.block_call: - self.started.set() - await self.release.wait() - return self.responses.pop(0) - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class AgentContinueTests(unittest.IsolatedAsyncioTestCase): - async def test_continue_is_a_distinct_run_without_new_user_message(self) -> None: - model = SequenceModel( - [ - AssistantMessage(content="initial answer"), - AssistantMessage(content="continued answer"), - ] - ) - observer = RecordingSink() - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - recorder = SessionRecorder(session_id="continue-basic", storage=storage) - agent = BaseAgent( - model, - agent_id="continue-basic-agent", - event_sink=CompositeAgentEventSink([recorder, observer]), - ) - - initial = await agent.run(task="investigate") - history_before_continue = tuple(dict(message) for message in agent.messages) - self.assertTrue(agent.can_continue) - continued = await agent.continue_run() - - self.assertEqual(initial.output, "initial answer") - self.assertEqual(continued.output, "continued answer") - self.assertIsNone(agent.state.task) - self.assertEqual(model.contexts[1], history_before_continue) - self.assertEqual( - [message["role"] for message in agent.messages], - ["system", "user", "assistant", "assistant"], - ) - - starts = [ - event - for event in observer.events - if isinstance(event.payload, (AgentStarted, AgentContinued)) - ] - self.assertEqual(len(starts), 2) - self.assertIsInstance(starts[0].payload, AgentStarted) - self.assertIsInstance(starts[1].payload, AgentContinued) - self.assertNotEqual(starts[0].run_id, starts[1].run_id) - - session = await recorder.load() - assert session is not None - self.assertEqual( - [run.intent for run in session.runs], - [SessionRunIntent.TASK, SessionRunIntent.CONTINUE], - ) - self.assertEqual([run.task for run in session.runs], ["investigate", None]) - self.assertEqual( - [message["role"] for message in session.messages], - ["user", "assistant", "assistant"], - ) - records = await storage.records("continue-basic") - self.assertEqual( - [record.kind for record in records], - [ - SessionRecordKind.RUN_STARTED, - SessionRecordKind.MESSAGE_APPENDED, - SessionRecordKind.RUN_FINISHED, - SessionRecordKind.RUN_CONTINUED, - SessionRecordKind.MESSAGE_APPENDED, - SessionRecordKind.RUN_FINISHED, - ], - ) - - async def test_continue_rejection_reasons_are_explicit(self) -> None: - empty = BaseAgent(SequenceModel([]), agent_id="continue-empty") - self.assertEqual( - empty.continue_rejection_reason, - ContinueRejectedReason.NO_PREVIOUS_RUN, - ) - with self.assertRaises(ContinueRejectedError) as raised: - await empty.continue_run() - self.assertEqual( - raised.exception.reason, ContinueRejectedReason.NO_PREVIOUS_RUN - ) - with self.assertRaises(ContinueRejectedError) as raised: - await empty.orchestrator.continue_run() - self.assertEqual( - raised.exception.reason, ContinueRejectedReason.NO_PREVIOUS_RUN - ) - - active_model = BlockingModel( - [AssistantMessage(content="cancelled")], - block_call=0, - ) - active = BaseAgent(active_model, agent_id="continue-active") - run = asyncio.create_task(active.run(task="block")) - await active_model.started.wait() - with self.assertRaises(ContinueRejectedError) as raised: - await active.continue_run() - self.assertEqual(raised.exception.reason, ContinueRejectedReason.AGENT_ACTIVE) - - self.assertTrue(active.abort()) - cancelled = await run - self.assertEqual(cancelled.status, RunStatus.CANCELLED) - with self.assertRaises(ContinueRejectedError) as raised: - await active.continue_run() - self.assertEqual( - raised.exception.reason, - ContinueRejectedReason.UNSUPPORTED_STOP_REASON, - ) - - incomplete = AgentSession(session_id="incomplete") - incomplete.bind_agent("continue-incomplete") - incomplete.begin_run("run-1", "task", 1) - incomplete.append_message( - "run-1", - 2, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "missing-result", - "type": "function", - "function": {"name": "lookup", "arguments": "{}"}, - } - ], - }, - ) - incomplete.finish_run( - "run-1", - 3, - AgentRunResult( - status=RunStatus.COMPLETED, - stop_reason=StopReason.TEXT_RESPONSE, - turns=1, - output="incomplete", - ), - ) - restored = BaseAgent(SequenceModel([]), agent_id="continue-incomplete") - restored.restore_session(incomplete) - self.assertEqual( - restored.continue_rejection_reason, - ContinueRejectedReason.INCOMPLETE_TOOL_STATE, - ) - - async def test_steering_reaches_continue_safe_point(self) -> None: - model = BlockingModel( - [ - AssistantMessage(content="initial"), - AssistantMessage(content="superseded continuation"), - AssistantMessage(content="steered continuation"), - ], - block_call=1, - ) - agent = BaseAgent(model, agent_id="continue-steering") - await agent.run(task="initial task") - - continuing = asyncio.create_task(agent.continue_run()) - await model.started.wait() - receipt = await agent.steer("continue in a safer direction") - self.assertTrue(receipt.accepted) - model.release.set() - result = await continuing - - self.assertEqual(result.output, "steered continuation") - self.assertEqual(result.turns, 2) - self.assertEqual(model.contexts[2][-1]["role"], "user") - self.assertEqual( - model.contexts[2][-1]["content"], - "continue in a safer direction", - ) - - async def test_restored_session_can_continue(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - first_recorder = SessionRecorder( - session_id="continue-restore", storage=storage - ) - first = BaseAgent( - SequenceModel([AssistantMessage(content="saved")]), - agent_id="continue-restore-agent", - event_sink=first_recorder, - ) - await first.run(task="persist this") - saved = await first_recorder.load() - assert saved is not None - - second_recorder = SessionRecorder( - session_id="continue-restore", - storage=storage, - ) - model = SequenceModel([AssistantMessage(content="resumed")]) - restored = BaseAgent( - model, - agent_id="continue-restore-agent", - event_sink=second_recorder, - ) - restored.restore_session(saved) - self.assertTrue(restored.can_continue) - - result = await restored.continue_run() - - self.assertEqual(result.output, "resumed") - self.assertEqual( - [message["role"] for message in model.contexts[0]], - ["system", "user", "assistant"], - ) - reloaded = await second_recorder.load() - assert reloaded is not None - self.assertEqual(len(reloaded.runs), 2) - self.assertEqual(reloaded.runs[-1].intent, SessionRunIntent.CONTINUE) - - async def test_safe_limit_failure_can_continue(self) -> None: - session = AgentSession(session_id="continue-limit") - session.bind_agent("continue-limit-agent") - session.begin_run("limited-run", "long task", 1) - session.append_message( - "limited-run", - 2, - {"role": "assistant", "content": "partial progress"}, - ) - session.finish_run( - "limited-run", - 3, - AgentRunResult( - status=RunStatus.FAILED, - stop_reason=StopReason.MAX_STEPS, - turns=1, - error="step limit reached", - ), - ) - model = SequenceModel([AssistantMessage(content="finished")]) - agent = BaseAgent(model, agent_id="continue-limit-agent") - agent.restore_session(session) - - self.assertTrue(agent.can_continue) - self.assertEqual((await agent.continue_run()).output, "finished") - - async def test_continue_wait_for_idle_includes_terminal_sink(self) -> None: - model = BlockingModel( - [AssistantMessage(content="first"), AssistantMessage(content="second")], - block_call=1, - ) - observer = RecordingSink() - agent = BaseAgent(model, agent_id="continue-idle", event_sink=observer) - await agent.run(task="first") - continuing = asyncio.create_task(agent.continue_run()) - await model.started.wait() - - idle = asyncio.create_task(agent.wait_for_idle()) - await asyncio.sleep(0) - self.assertFalse(idle.done()) - model.release.set() - await continuing - await idle - self.assertIsInstance(observer.events[-1].payload, AgentFinished) - - async def test_continue_preserves_follow_up_and_shutdown_semantics(self) -> None: - model = BlockingModel( - [ - AssistantMessage(content="first"), - AssistantMessage(content="continued"), - AssistantMessage(content="follow-up"), - ], - block_call=1, - ) - agent = BaseAgent(model, agent_id="continue-chain") - await agent.run(task="first") - continuing = asyncio.create_task(agent.continue_run()) - await model.started.wait() - follow_up = await agent.follow_up("after continue") - self.assertTrue(follow_up.accepted) - - shutdown = asyncio.create_task(agent.shutdown()) - await asyncio.sleep(0) - with self.assertRaises(ContinueRejectedError) as raised: - await agent.continue_run() - self.assertEqual( - raised.exception.reason, - ContinueRejectedReason.AGENT_SHUTTING_DOWN, - ) - model.release.set() - - self.assertEqual((await continuing).output, "continued") - with self.assertRaises(FollowUpDiscardedError) as discarded: - await follow_up.wait() - self.assertEqual( - discarded.exception.reason, - FollowUpDiscardReason.AGENT_SHUTDOWN, - ) - await shutdown - self.assertEqual(len(model.contexts), 2) - - async def test_follow_up_runs_after_successful_continue(self) -> None: - model = BlockingModel( - [ - AssistantMessage(content="first"), - AssistantMessage(content="continued"), - AssistantMessage(content="followed up"), - ], - block_call=1, - ) - agent = BaseAgent(model, agent_id="continue-follow-up") - await agent.run(task="first") - continuing = asyncio.create_task(agent.continue_run()) - await model.started.wait() - follow_up = await agent.follow_up("next task") - model.release.set() - - self.assertEqual((await continuing).output, "continued") - self.assertEqual((await follow_up.wait()).output, "followed up") - self.assertEqual(model.contexts[2][-1]["role"], "user") - self.assertEqual(model.contexts[2][-1]["content"], "next task") - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_agent_events.py b/tests/test_agent_events.py deleted file mode 100644 index e7d07ca..0000000 --- a/tests/test_agent_events.py +++ /dev/null @@ -1,378 +0,0 @@ -import unittest -from dataclasses import FrozenInstanceError -from typing import Any, get_type_hints - -from ejagent import ( - AgentEvent, - AgentFinished, - AgentStarted, - AssistantMessage, - BaseAgent, - CancellationToken, - MessageCompleted, - MethodToolHandler, - ModelAdapter, - ModelToolCall, - RunStatus, - RuntimePolicy, - StepOutcome, - StopReason, - ToolCallResult, - ToolCompleted, - ToolControl, - ToolStarted, - TurnCompleted, - TurnStarted, -) - -ECHO_TOOL = { - "type": "function", - "function": { - "name": "echo", - "description": "Return the provided text.", - "parameters": { - "type": "object", - "properties": {"text": {"type": "string"}}, - "required": ["text"], - }, - }, -} - -FINISH_TOOL = { - "type": "function", - "function": { - "name": "finish", - "description": "Settle the current task.", - "parameters": { - "type": "object", - "properties": {"summary": {"type": "string"}}, - "required": ["summary"], - }, - }, -} - - -class RecordingEventSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class FailingEventSink: - def __init__(self) -> None: - self.calls = 0 - - async def emit(self, event: AgentEvent) -> None: - self.calls += 1 - raise RuntimeError(f"sink failed for {event.kind}") - - -class SequenceModel(ModelAdapter): - def __init__( - self, - responses: list[AssistantMessage | Exception], - ) -> None: - self.responses = list(responses) - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - response = self.responses.pop(0) - if isinstance(response, Exception): - raise response - return response - - -class EchoHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((ECHO_TOOL,)) - self.calls = 0 - - async def do_echo( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - self.calls += 1 - return StepOutcome({"text": arguments.get("text")}) - - -class FinishHandler(MethodToolHandler): - def __init__(self, control: ToolControl) -> None: - super().__init__((FINISH_TOOL,)) - self.control = control - - async def do_finish( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - return StepOutcome( - {"summary": arguments.get("summary")}, - control=self.control, - ) - - -def event_payloads(sink: RecordingEventSink) -> list[object]: - return [event.payload for event in sink.events] - - -class AgentEventTests(unittest.IsolatedAsyncioTestCase): - async def test_text_run_emits_ordered_correlated_events(self) -> None: - response = AssistantMessage(content="done") - sink = RecordingEventSink() - agent = BaseAgent( - SequenceModel([response]), - agent_id="text-events", - event_sink=sink, - ) - - result = await agent.run(task="complete") - - payloads = event_payloads(sink) - self.assertEqual( - [type(payload) for payload in payloads], - [ - AgentStarted, - TurnStarted, - MessageCompleted, - TurnCompleted, - AgentFinished, - ], - ) - self.assertIs(payloads[2].message, response) - self.assertIs(payloads[-1].result, result) - self.assertEqual( - [event.sequence for event in sink.events], - [1, 2, 3, 4, 5], - ) - self.assertEqual( - {event.agent_id for event in sink.events}, - {"text-events"}, - ) - self.assertEqual(len({event.run_id for event in sink.events}), 1) - - with self.assertRaises(FrozenInstanceError): - sink.events[0].sequence = 99 # type: ignore[misc] - self.assertIs( - get_type_hints(ToolCompleted)["result"], - ToolCallResult, - ) - - async def test_tool_run_emits_tool_events_between_message_and_turn(self) -> None: - tool_call = ModelToolCall( - id="call-1", - name="echo", - arguments='{"text": "hello"}', - ) - sink = RecordingEventSink() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(tool_calls=(tool_call,)), - AssistantMessage(content="finished"), - ] - ), - agent_id="tool-events", - handlers=[EchoHandler()], - event_sink=sink, - ) - - result = await agent.run(task="use echo") - - payloads = event_payloads(sink) - self.assertEqual( - [type(payload) for payload in payloads], - [ - AgentStarted, - TurnStarted, - MessageCompleted, - ToolStarted, - ToolCompleted, - TurnCompleted, - TurnStarted, - MessageCompleted, - TurnCompleted, - AgentFinished, - ], - ) - self.assertIs(payloads[3].tool_call, tool_call) - self.assertIs(payloads[4].tool_call, tool_call) - self.assertIsNone(payloads[4].result.error) - self.assertEqual(result.status, RunStatus.COMPLETED) - - async def test_terminal_tool_controls_reuse_agent_run_result(self) -> None: - statuses = { - ToolControl.COMPLETE: RunStatus.COMPLETED, - ToolControl.REJECT: RunStatus.REJECTED, - ToolControl.CANCEL: RunStatus.CANCELLED, - } - - for control, expected_status in statuses.items(): - with self.subTest(control=control): - sink = RecordingEventSink() - tool_call = ModelToolCall( - id=f"call-{control}", - name="finish", - arguments='{"summary": "settled"}', - ) - agent = BaseAgent( - SequenceModel([AssistantMessage(tool_calls=(tool_call,))]), - agent_id=f"terminal-{control}", - handlers=[FinishHandler(control)], - event_sink=sink, - ) - - result = await agent.run(task="settle") - - tool_event = next( - payload - for payload in event_payloads(sink) - if isinstance(payload, ToolCompleted) - ) - finish_event = event_payloads(sink)[-1] - self.assertEqual(tool_event.result.control, control) - self.assertIsInstance(finish_event, AgentFinished) - self.assertIs(finish_event.result, result) - self.assertEqual(result.status, expected_status) - - async def test_provider_failure_still_completes_turn_and_run_events(self) -> None: - sink = RecordingEventSink() - agent = BaseAgent( - SequenceModel([RuntimeError("provider unavailable")]), - agent_id="provider-failure", - event_sink=sink, - ) - - result = await agent.run(task="fail") - - self.assertEqual( - [type(payload) for payload in event_payloads(sink)], - [AgentStarted, TurnStarted, TurnCompleted, AgentFinished], - ) - self.assertEqual(result.status, RunStatus.FAILED) - self.assertEqual(result.stop_reason, StopReason.RUNTIME_ERROR) - - async def test_tool_argument_error_is_visible_in_tool_completed(self) -> None: - sink = RecordingEventSink() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=( - ModelToolCall( - id="invalid-json", - name="echo", - arguments="[]", - ), - ) - ), - AssistantMessage(content="recovered"), - ] - ), - agent_id="tool-error", - handlers=[EchoHandler()], - event_sink=sink, - ) - - result = await agent.run(task="recover") - - completed = next( - payload - for payload in event_payloads(sink) - if isinstance(payload, ToolCompleted) - ) - self.assertIn("JSON object", completed.result.error or "") - self.assertEqual(result.status, RunStatus.COMPLETED) - - async def test_repeated_tool_failure_keeps_tool_and_turn_events_paired( - self, - ) -> None: - call = ModelToolCall( - id="repeated", - name="echo", - arguments='{"text": "same"}', - ) - sink = RecordingEventSink() - handler = EchoHandler() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(tool_calls=(call,)), - AssistantMessage(tool_calls=(call,)), - ] - ), - agent_id="repeated-events", - handlers=[handler], - runtime_policy=RuntimePolicy( - max_steps=2, - max_repeated_tool_calls=2, - ), - event_sink=sink, - ) - - result = await agent.run(task="repeat") - - payloads = event_payloads(sink) - started = [p for p in payloads if isinstance(p, ToolStarted)] - completed = [p for p in payloads if isinstance(p, ToolCompleted)] - turns_started = [p for p in payloads if isinstance(p, TurnStarted)] - turns_completed = [p for p in payloads if isinstance(p, TurnCompleted)] - self.assertEqual(len(started), len(completed)) - self.assertEqual(len(turns_started), len(turns_completed)) - self.assertIn("consecutive times", completed[-1].result.error or "") - self.assertEqual(handler.calls, 1) - self.assertEqual(result.stop_reason, StopReason.REPEATED_TOOL_CALL) - - async def test_sink_failure_does_not_change_agent_result(self) -> None: - sink = FailingEventSink() - agent = BaseAgent( - SequenceModel([AssistantMessage(content="done")]), - agent_id="sink-failure", - event_sink=sink, - ) - - with self.assertLogs("sink-failure", level="WARNING"): - result = await agent.run(task="complete") - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(sink.calls, 5) - - async def test_each_run_has_a_new_id_and_sequence(self) -> None: - sink = RecordingEventSink() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(content="first"), - AssistantMessage(content="second"), - ] - ), - agent_id="multiple-runs", - event_sink=sink, - ) - - await agent.run(task="first") - await agent.run(task="second") - - first_run_id = sink.events[0].run_id - second_run_id = sink.events[5].run_id - self.assertNotEqual(first_run_id, second_run_id) - self.assertEqual( - [event.sequence for event in sink.events[:5]], - [1, 2, 3, 4, 5], - ) - self.assertEqual( - [event.sequence for event in sink.events[5:]], - [1, 2, 3, 4, 5], - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_agent_follow_up.py b/tests/test_agent_follow_up.py deleted file mode 100644 index d109645..0000000 --- a/tests/test_agent_follow_up.py +++ /dev/null @@ -1,278 +0,0 @@ -import asyncio -import tempfile -import unittest -from typing import Any - -from ejagent import ( - AgentEvent, - AgentFinished, - AgentStarted, - AssistantMessage, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - ControlStatus, - FollowUpDiscardedError, - FollowUpDiscardReason, - FollowUpFailurePolicy, - FollowUpRejectedError, - JsonlSessionStorage, - ModelAdapter, - RunStatus, - SessionRecorder, - SessionRecordKind, -) - - -class BlockingSequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - self.contexts: list[tuple[dict[str, Any], ...]] = [] - self.started = [asyncio.Event() for _ in responses] - self.release_first = asyncio.Event() - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - index = len(self.contexts) - self.contexts.append(context.agent_messages) - self.started[index].set() - if index == 0: - await self.release_first.wait() - return self.responses.pop(0) - - -class SlowFirstFinishSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - self.first_finish_started = asyncio.Event() - self.release_first_finish = asyncio.Event() - self._finish_count = 0 - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - if isinstance(event.payload, AgentFinished): - self._finish_count += 1 - if self._finish_count == 1: - self.first_finish_started.set() - await self.release_first_finish.wait() - - -class AgentFollowUpTests(unittest.IsolatedAsyncioTestCase): - async def test_follow_ups_are_independent_fifo_runs_after_terminal_sink( - self, - ) -> None: - model = BlockingSequenceModel( - [ - AssistantMessage(content="initial result"), - AssistantMessage(content="first follow-up result"), - AssistantMessage(content="second follow-up result"), - AssistantMessage(content="direct result"), - ] - ) - observer = SlowFirstFinishSink() - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - recorder = SessionRecorder(session_id="follow-up-order", storage=storage) - agent = BaseAgent( - model, - agent_id="follow-up-order-agent", - event_sink=CompositeAgentEventSink([recorder, observer]), - ) - initial = asyncio.create_task(agent.run(task="initial")) - await model.started[0].wait() - - first = await agent.follow_up("first follow-up") - second = await agent.follow_up("second follow-up") - direct = asyncio.create_task(agent.run(task="direct run")) - self.assertTrue(first.accepted) - self.assertTrue(second.accepted) - self.assertEqual(first.receipt.queue_size, 1) - self.assertEqual(second.receipt.queue_size, 2) - self.assertEqual(agent.pending_follow_up_count, 2) - - model.release_first.set() - await observer.first_finish_started.wait() - idle = asyncio.create_task(agent.wait_for_idle()) - await asyncio.sleep(0) - self.assertEqual(len(model.contexts), 1) - self.assertFalse(initial.done()) - self.assertFalse(first.done) - self.assertFalse(idle.done()) - - observer.release_first_finish.set() - initial_result = await initial - first_result = await first.wait() - second_result = await second.wait() - direct_result = await direct - await idle - - self.assertEqual(initial_result.output, "initial result") - self.assertEqual(first_result.output, "first follow-up result") - self.assertEqual(second_result.output, "second follow-up result") - self.assertEqual(direct_result.output, "direct result") - started_tasks = [ - event.payload.task - for event in observer.events - if isinstance(event.payload, AgentStarted) - ] - self.assertEqual( - started_tasks, - ["initial", "first follow-up", "second follow-up", "direct run"], - ) - records = await storage.records("follow-up-order") - terminal_kinds = [ - record.kind - for record in records - if record.kind - in {SessionRecordKind.RUN_STARTED, SessionRecordKind.RUN_FINISHED} - ] - self.assertEqual( - terminal_kinds, - [ - SessionRecordKind.RUN_STARTED, - SessionRecordKind.RUN_FINISHED, - ] - * 4, - ) - run_ids = [ - record.data["run_id"] - for record in records - if record.kind is SessionRecordKind.RUN_STARTED - ] - self.assertEqual(len(set(run_ids)), 4) - - async def test_idle_queue_full_and_empty_task_are_explicit(self) -> None: - model = BlockingSequenceModel( - [AssistantMessage(content="initial"), AssistantMessage(content="queued")] - ) - agent = BaseAgent( - model, - agent_id="follow-up-capacity", - follow_up_queue_capacity=1, - ) - - idle = await agent.follow_up("too early") - self.assertEqual(idle.receipt.status, ControlStatus.AGENT_IDLE) - with self.assertRaises(FollowUpRejectedError): - await idle.wait() - - initial = asyncio.create_task(agent.run(task="initial")) - await model.started[0].wait() - accepted = await agent.follow_up("accepted") - full = await agent.follow_up("full") - self.assertTrue(accepted.accepted) - self.assertEqual(full.receipt.status, ControlStatus.QUEUE_FULL) - self.assertEqual(full.receipt.queue_size, 1) - with self.assertRaises(FollowUpRejectedError): - await full.wait() - with self.assertRaisesRegex(ValueError, "must not be empty"): - await agent.follow_up(" ") - - model.release_first.set() - await initial - self.assertEqual((await accepted.wait()).output, "queued") - after = await agent.follow_up("too late") - self.assertEqual(after.receipt.status, ControlStatus.AGENT_IDLE) - - async def test_failure_policy_discards_or_continues_pending_runs(self) -> None: - discard_model = BlockingSequenceModel( - [AssistantMessage(), AssistantMessage(content="must not run")] - ) - discard_agent = BaseAgent(discard_model, agent_id="follow-up-discard") - failed = asyncio.create_task(discard_agent.run(task="fail")) - await discard_model.started[0].wait() - discarded = await discard_agent.follow_up("discard me") - discard_model.release_first.set() - - failed_result = await failed - self.assertEqual(failed_result.status, RunStatus.FAILED) - with self.assertRaises(FollowUpDiscardedError) as raised: - await discarded.wait() - self.assertEqual( - raised.exception.reason, - FollowUpDiscardReason.PREVIOUS_RUN_NOT_COMPLETED, - ) - self.assertEqual(len(discard_model.contexts), 1) - - continue_model = BlockingSequenceModel( - [AssistantMessage(), AssistantMessage(content="continued")] - ) - continue_agent = BaseAgent( - continue_model, - agent_id="follow-up-continue", - follow_up_failure_policy=FollowUpFailurePolicy.CONTINUE, - ) - failed = asyncio.create_task(continue_agent.run(task="fail")) - await continue_model.started[0].wait() - continued = await continue_agent.follow_up("continue anyway") - continue_model.release_first.set() - - self.assertEqual((await failed).status, RunStatus.FAILED) - self.assertEqual((await continued.wait()).output, "continued") - - async def test_aborted_run_discards_follow_up_by_default(self) -> None: - model = BlockingSequenceModel( - [AssistantMessage(content="unused"), AssistantMessage(content="unused")] - ) - agent = BaseAgent(model, agent_id="follow-up-abort") - initial = asyncio.create_task(agent.run(task="initial")) - await model.started[0].wait() - handle = await agent.follow_up("must not run") - - self.assertTrue(agent.abort("stop the chain")) - closing = await agent.follow_up("also must not run") - result = await initial - - self.assertEqual(result.status, RunStatus.CANCELLED) - self.assertEqual(closing.receipt.status, ControlStatus.RUN_CLOSING) - with self.assertRaises(FollowUpDiscardedError) as raised: - await handle.wait() - self.assertEqual( - raised.exception.reason, - FollowUpDiscardReason.PREVIOUS_RUN_NOT_COMPLETED, - ) - self.assertEqual(len(model.contexts), 1) - - async def test_shutdown_discards_queue_and_waiter_cancellation_is_local( - self, - ) -> None: - model = BlockingSequenceModel( - [AssistantMessage(content="initial"), AssistantMessage(content="unused")] - ) - agent = BaseAgent(model, agent_id="follow-up-shutdown") - initial = asyncio.create_task(agent.run(task="initial")) - await model.started[0].wait() - handle = await agent.follow_up("queued") - - waiter = asyncio.create_task(handle.wait()) - await asyncio.sleep(0) - waiter.cancel() - with self.assertRaises(asyncio.CancelledError): - await waiter - self.assertFalse(handle.done) - - shutdown = asyncio.create_task(agent.shutdown()) - await asyncio.sleep(0) - self.assertTrue(handle.done) - with self.assertRaises(FollowUpDiscardedError) as raised: - await handle.wait() - self.assertEqual( - raised.exception.reason, - FollowUpDiscardReason.AGENT_SHUTDOWN, - ) - closing = await agent.follow_up("too late") - self.assertEqual(closing.receipt.status, ControlStatus.RUN_CLOSING) - - model.release_first.set() - self.assertEqual((await initial).status, RunStatus.COMPLETED) - await shutdown - self.assertEqual(len(model.contexts), 1) - await agent.wait_for_idle() - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_agent_harness.py b/tests/test_agent_harness.py new file mode 100644 index 0000000..7a5ed0a --- /dev/null +++ b/tests/test_agent_harness.py @@ -0,0 +1,518 @@ +from __future__ import annotations + +import asyncio +import unittest +from collections.abc import AsyncIterator, Iterable, Sequence +from dataclasses import replace +from datetime import UTC, datetime + +from ejagent.contracts import ( + AssistantMessage, + CancellationToken, + FailureCode, + ModelCallError, + ModelPort, + ModelRequest, + ModelResponseCompleted, + ModelStreamEvent, + RunStatus, + SessionCommit, + SessionConflictError, + SessionSnapshot, + SessionStore, + SessionStoreError, + StopReason, + SystemMessage, + ToolCall, + ToolDefinition, + ToolExecutionResult, + ToolExecutor, + UserMessage, +) +from ejagent.harness import ( + AgentHarness, + HarnessClosedError, + HarnessStatus, + MemorySessionStore, +) + + +def fixed_clock() -> datetime: + return datetime(2026, 7, 31, 14, 0, tzinfo=UTC) + + +def ids(*values: str) -> Iterable[str]: + return iter(values) + + +class ScriptedModel(ModelPort): + def __init__( + self, + responses: Sequence[AssistantMessage | ModelCallError], + ) -> None: + self.responses = list(responses) + self.requests: list[ModelRequest] = [] + + async def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + self.requests.append(request) + response = self.responses.pop(0) + if isinstance(response, ModelCallError): + raise response + yield ModelResponseCompleted(response) + + +class BlockingModel(ModelPort): + def __init__(self) -> None: + self.started = asyncio.Event() + + async def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + self.started.set() + await asyncio.Event().wait() + yield ModelResponseCompleted(AssistantMessage(content="unreachable")) + + +class NoTools(ToolExecutor): + @property + def definitions(self) -> Sequence[ToolDefinition]: + return () + + async def execute( + self, + call: ToolCall, + *, + cancellation: CancellationToken, + ) -> ToolExecutionResult: + raise AssertionError("no tools are registered") + + +class FailingStore(SessionStore): + async def load(self, agent_id: str) -> SessionSnapshot | None: + return None + + async def commit(self, commit: SessionCommit) -> SessionSnapshot: + raise SessionStoreError("durable backend unavailable") + + +class RecordingResource: + def __init__( + self, + name: str, + events: list[str], + *, + fail_start: bool = False, + ) -> None: + self.name = name + self.events = events + self.fail_start = fail_start + + async def start(self) -> None: + self.events.append(f"start:{self.name}") + if self.fail_start: + raise RuntimeError(f"cannot start {self.name}") + + async def shutdown(self) -> None: + self.events.append(f"shutdown:{self.name}") + + +class ManagedScriptedModel(ScriptedModel): + def __init__(self, events: list[str]) -> None: + super().__init__([AssistantMessage(content="managed")]) + self.events = events + + async def start(self) -> None: + self.events.append("start:model") + + async def shutdown(self) -> None: + self.events.append("shutdown:model") + + +class ManagedNoTools(NoTools): + def __init__(self, events: list[str]) -> None: + self.events = events + + async def start(self) -> None: + self.events.append("start:tools") + + async def shutdown(self) -> None: + self.events.append("shutdown:tools") + + +class AgentHarnessTests(unittest.IsolatedAsyncioTestCase): + async def test_successful_run_commits_only_after_store_accepts(self) -> None: + store = MemorySessionStore() + model = ScriptedModel([AssistantMessage(content="done")]) + run_ids = ids("run-1") + harness = AgentHarness( + agent_id="agent-1", + model=model, + tools=NoTools(), + initial_messages=(SystemMessage("be precise"),), + store=store, + run_id_factory=lambda: next(run_ids), + clock=fixed_clock, + ) + + outcome = await harness.run("solve it") + + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertEqual(harness.revision, 1) + self.assertEqual( + harness.messages, + ( + SystemMessage("be precise"), + UserMessage("solve it"), + AssistantMessage(content="done"), + ), + ) + self.assertEqual(harness.last_result, outcome.result) + self.assertEqual((await store.load("agent-1")), harness.snapshot) + self.assertEqual(len(await store.commits("agent-1")), 1) + audit = await store.load_audit("agent-1") + self.assertEqual(len(audit), 1) + self.assertTrue(audit[0].committed) + self.assertEqual(audit[0].result, outcome.result) + self.assertFalse(hasattr(audit[0], "messages")) + await harness.shutdown() + + async def test_failed_run_is_audited_without_advancing_conversation(self) -> None: + store = MemorySessionStore() + model = ScriptedModel( + [ + ModelCallError( + FailureCode.RATE_LIMIT, + "slow down", + retryable=True, + ) + ] + ) + harness = AgentHarness( + agent_id="agent-failed", + model=model, + tools=NoTools(), + initial_messages=(SystemMessage("stable"),), + store=store, + run_id_factory=lambda: "failed-run", + clock=fixed_clock, + ) + + outcome = await harness.run("transient task") + + self.assertEqual(outcome.result.status, RunStatus.FAILED) + self.assertEqual(harness.revision, 0) + self.assertEqual(harness.messages, (SystemMessage("stable"),)) + commits = await store.commits("agent-failed") + self.assertEqual(len(commits), 1) + self.assertFalse(commits[0].advances_revision) + self.assertEqual(commits[0].outcome.failure, outcome.failure) + audit = await store.load_audit("agent-failed") + self.assertFalse(audit[0].committed) + self.assertEqual(audit[0].resulting_revision, 0) + self.assertEqual(audit[0].failure, outcome.failure) + + async def test_store_failure_cannot_report_success_or_advance_state(self) -> None: + harness = AgentHarness( + agent_id="agent-store-failure", + model=ScriptedModel([AssistantMessage(content="model succeeded")]), + tools=NoTools(), + initial_messages=(SystemMessage("unchanged"),), + store=FailingStore(), + run_id_factory=lambda: "uncommitted-run", + clock=fixed_clock, + ) + + outcome = await harness.run("do work") + + self.assertEqual(outcome.result.status, RunStatus.FAILED) + self.assertEqual(outcome.result.stop_reason, StopReason.PERSISTENCE_FAILED) + assert outcome.failure is not None + self.assertEqual(outcome.failure.phase.value, "commit") + self.assertEqual(outcome.failure.code, FailureCode.PERSISTENCE_FAILED) + self.assertEqual(harness.revision, 0) + self.assertEqual(harness.messages, (SystemMessage("unchanged"),)) + self.assertEqual(outcome.audit_records[-1].kind, "commit_failed") + + async def test_concurrent_calls_are_fifo_and_use_fresh_revisions(self) -> None: + model = ScriptedModel( + [ + AssistantMessage(content="first result"), + AssistantMessage(content="second result"), + ] + ) + run_ids = ids("run-1", "run-2") + harness = AgentHarness( + agent_id="serial-agent", + model=model, + tools=NoTools(), + run_id_factory=lambda: next(run_ids), + clock=fixed_clock, + ) + + first = asyncio.create_task(harness.run("first")) + await asyncio.sleep(0) + second = asyncio.create_task(harness.run("second")) + first_outcome, second_outcome = await asyncio.gather(first, second) + + self.assertEqual(first_outcome.delta.base_revision, 0) + self.assertEqual(second_outcome.delta.base_revision, 1) + self.assertEqual(harness.revision, 2) + self.assertEqual( + model.requests[1].messages, + ( + UserMessage("first"), + AssistantMessage(content="first result"), + UserMessage("second"), + ), + ) + + async def test_continue_run_uses_history_without_new_user_message(self) -> None: + store = MemorySessionStore() + model = ScriptedModel( + [ + AssistantMessage(content="initial"), + AssistantMessage(content="continued"), + ] + ) + run_ids = ids("task-run", "continue-run") + harness = AgentHarness( + agent_id="continue-agent", + model=model, + tools=NoTools(), + store=store, + run_id_factory=lambda: next(run_ids), + clock=fixed_clock, + ) + + await harness.run("begin") + outcome = await harness.continue_run() + + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertEqual( + outcome.delta.messages, + (AssistantMessage(content="continued"),), + ) + commits = await store.commits("continue-agent") + self.assertEqual(commits[-1].outcome.delta.base_revision, 1) + self.assertEqual(commits[-1].outcome.result.run_id, "continue-run") + + async def test_cancellation_is_structured_and_does_not_commit_delta(self) -> None: + model = BlockingModel() + harness = AgentHarness( + agent_id="cancel-agent", + model=model, + tools=NoTools(), + initial_messages=(SystemMessage("retain me"),), + run_id_factory=lambda: "cancelled-run", + clock=fixed_clock, + ) + running = asyncio.create_task(harness.run("never finish")) + await model.started.wait() + + self.assertTrue(harness.cancel("user stopped it")) + self.assertFalse(harness.cancel("second reason")) + outcome = await running + + self.assertEqual(outcome.result.status, RunStatus.CANCELLED) + self.assertEqual(harness.revision, 0) + self.assertEqual(harness.messages, (SystemMessage("retain me"),)) + self.assertEqual(harness.status, HarnessStatus.READY) + + async def test_startup_rolls_back_resources_in_reverse_order(self) -> None: + events: list[str] = [] + resources = ( + RecordingResource("one", events), + RecordingResource("two", events), + RecordingResource("broken", events, fail_start=True), + ) + harness = AgentHarness( + agent_id="resource-agent", + model=ScriptedModel([AssistantMessage(content="unused")]), + tools=NoTools(), + resources=resources, + ) + + with self.assertRaisesRegex(RuntimeError, "cannot start broken"): + await harness.start() + + self.assertEqual( + events, + [ + "start:one", + "start:two", + "start:broken", + "shutdown:two", + "shutdown:one", + ], + ) + self.assertEqual(harness.status, HarnessStatus.NEW) + + async def test_harness_owns_model_and_tool_lifecycle_without_duplicates( + self, + ) -> None: + events: list[str] = [] + model = ManagedScriptedModel(events) + tools = ManagedNoTools(events) + harness = AgentHarness( + agent_id="managed-dependencies", + model=model, + tools=tools, + resources=(model, tools), + run_id_factory=lambda: "managed-run", + clock=fixed_clock, + ) + + await harness.run("execute") + await harness.shutdown() + + self.assertEqual( + events, + [ + "start:model", + "start:tools", + "shutdown:tools", + "shutdown:model", + ], + ) + + async def test_shutdown_cancels_run_and_closes_resources_in_reverse(self) -> None: + events: list[str] = [] + model = BlockingModel() + harness = AgentHarness( + agent_id="shutdown-agent", + model=model, + tools=NoTools(), + resources=( + RecordingResource("one", events), + RecordingResource("two", events), + ), + run_id_factory=lambda: "shutdown-run", + clock=fixed_clock, + ) + running = asyncio.create_task(harness.run("wait")) + await model.started.wait() + + await harness.shutdown() + outcome = await running + + self.assertEqual(outcome.result.status, RunStatus.CANCELLED) + self.assertEqual(harness.status, HarnessStatus.CLOSED) + self.assertEqual( + events, + ["start:one", "start:two", "shutdown:two", "shutdown:one"], + ) + with self.assertRaises(HarnessClosedError): + await harness.run("too late") + + async def test_new_harness_restores_the_committed_snapshot(self) -> None: + store = MemorySessionStore() + first = AgentHarness( + agent_id="durable-agent", + model=ScriptedModel([AssistantMessage(content="saved")]), + tools=NoTools(), + store=store, + run_id_factory=lambda: "saved-run", + clock=fixed_clock, + ) + await first.run("persist") + await first.shutdown() + + second_model = ScriptedModel([AssistantMessage(content="restored")]) + second = AgentHarness( + agent_id="durable-agent", + model=second_model, + tools=NoTools(), + store=store, + run_id_factory=lambda: "restored-run", + clock=fixed_clock, + ) + + outcome = await second.continue_run() + + self.assertEqual(outcome.delta.base_revision, 1) + self.assertEqual(second.revision, 2) + self.assertEqual( + second_model.requests[0].messages, + (UserMessage("persist"), AssistantMessage(content="saved")), + ) + + async def test_memory_store_repeats_identical_commit_idempotently(self) -> None: + store = MemorySessionStore() + harness = AgentHarness( + agent_id="idempotent-agent", + model=ScriptedModel([AssistantMessage(content="once")]), + tools=NoTools(), + store=store, + run_id_factory=lambda: "stable-run-id", + clock=fixed_clock, + ) + await harness.run("commit") + commit = (await store.commits("idempotent-agent"))[0] + + repeated = await store.commit(commit) + + self.assertEqual(repeated, harness.snapshot) + self.assertEqual(len(await store.commits("idempotent-agent")), 1) + + async def test_memory_store_rejects_reused_run_id_with_new_content(self) -> None: + store = MemorySessionStore() + harness = AgentHarness( + agent_id="reuse-agent", + model=ScriptedModel([AssistantMessage(content="original")]), + tools=NoTools(), + store=store, + run_id_factory=lambda: "reused-id", + clock=fixed_clock, + ) + await harness.run("commit") + commit = (await store.commits("reuse-agent"))[0] + changed_result = replace(commit.outcome.result, output="changed") + changed_outcome = replace(commit.outcome, result=changed_result) + changed_commit = replace(commit, outcome=changed_outcome) + + with self.assertRaises(SessionConflictError): + await store.commit(changed_commit) + + async def test_stale_harness_commit_becomes_persistence_failure(self) -> None: + store = MemorySessionStore() + first = AgentHarness( + agent_id="shared-agent", + model=ScriptedModel([AssistantMessage(content="winner")]), + tools=NoTools(), + store=store, + run_id_factory=lambda: "winner-run", + clock=fixed_clock, + ) + stale = AgentHarness( + agent_id="shared-agent", + model=ScriptedModel([AssistantMessage(content="stale")]), + tools=NoTools(), + store=store, + run_id_factory=lambda: "stale-run", + clock=fixed_clock, + ) + await first.start() + await stale.start() + + await first.run("first writer") + outcome = await stale.run("stale writer") + + self.assertEqual(outcome.result.status, RunStatus.FAILED) + self.assertEqual(outcome.result.stop_reason, StopReason.PERSISTENCE_FAILED) + self.assertEqual(stale.revision, 0) + persisted = await store.load("shared-agent") + assert persisted is not None + self.assertEqual(persisted.revision, 1) + self.assertIn(AssistantMessage(content="winner"), persisted.messages) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_agent_session.py b/tests/test_agent_session.py deleted file mode 100644 index ed04fe5..0000000 --- a/tests/test_agent_session.py +++ /dev/null @@ -1,265 +0,0 @@ -import unittest -from typing import Any - -from ejagent import ( - AgentEvent, - AgentRunResult, - AgentSession, - AssistantMessage, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - MemorySessionStorage, - MethodToolHandler, - ModelAdapter, - ModelToolCall, - RunStatus, - SessionRecorder, - StepOutcome, - StopReason, -) - -ECHO_TOOL = { - "type": "function", - "function": { - "name": "echo", - "description": "Return the provided text.", - "parameters": { - "type": "object", - "properties": {"text": {"type": "string"}}, - "required": ["text"], - }, - }, -} - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - self.contexts: list[tuple[dict[str, Any], ...]] = [] - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.contexts.append(context.agent_messages) - return self.responses.pop(0) - - -class EchoHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((ECHO_TOOL,)) - - async def do_echo( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - return StepOutcome({"text": arguments.get("text")}) - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class FailingSink: - async def emit(self, event: AgentEvent) -> None: - raise RuntimeError(f"failed for {event.kind}") - - -class AgentSessionTests(unittest.IsolatedAsyncioTestCase): - async def test_recorder_persists_user_assistant_and_tool_messages( - self, - ) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="session-1", storage=storage) - tool_call = ModelToolCall( - id="call-1", - name="echo", - arguments='{"text": "hello"}', - ) - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(tool_calls=(tool_call,)), - AssistantMessage(content="finished"), - ] - ), - agent_id="session-agent", - handlers=[EchoHandler()], - event_sink=recorder, - ) - - result = await agent.run(task="use echo") - session = await recorder.load() - - self.assertIsNotNone(session) - assert session is not None - self.assertEqual(session.agent_id, "session-agent") - self.assertEqual( - [message["role"] for message in session.messages], - ["user", "assistant", "tool", "assistant"], - ) - self.assertFalse( - any(message["role"] == "system" for message in session.messages) - ) - self.assertEqual(len(session.runs), 1) - self.assertIs(session.runs[0].result, result) - self.assertTrue(session.runs[0].finished) - self.assertEqual( - [entry.sequence for entry in session.entries], - [1, 3, 5, 8], - ) - - async def test_one_session_records_multiple_runs(self) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="session-runs", storage=storage) - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(content="first answer"), - AssistantMessage(content="second answer"), - ] - ), - agent_id="session-agent", - event_sink=recorder, - ) - - first = await agent.run(task="first task") - second = await agent.run(task="second task") - session = await recorder.load() - - assert session is not None - self.assertEqual(len(session.runs), 2) - self.assertNotEqual(session.runs[0].run_id, session.runs[1].run_id) - self.assertIs(session.runs[0].result, first) - self.assertIs(session.runs[1].result, second) - self.assertEqual( - [message["content"] for message in session.messages], - ["first task", "first answer", "second task", "second answer"], - ) - - async def test_loaded_history_can_resume_in_a_new_agent(self) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="resume", storage=storage) - first_agent = BaseAgent( - SequenceModel([AssistantMessage(content="saved answer")]), - agent_id="resume-agent", - event_sink=recorder, - ) - await first_agent.run(task="saved task") - session = await recorder.load() - assert session is not None - - resumed_model = SequenceModel([AssistantMessage(content="continued answer")]) - resumed_agent = BaseAgent( - resumed_model, - agent_id="resume-agent", - ) - resumed_agent.reset(session.messages) - - result = await resumed_agent.run(task="continue") - - self.assertEqual(result.output, "continued answer") - self.assertEqual( - [message.get("content") for message in resumed_model.contexts[0]], - [ - resumed_agent.system_prompt, - "saved task", - "saved answer", - "continue", - ], - ) - - async def test_memory_storage_returns_detached_snapshots(self) -> None: - storage = MemorySessionStorage() - session = AgentSession(session_id="isolated") - session.bind_agent("agent") - session.begin_run("run-1", "task", 1) - session.append_message( - "run-1", - 2, - {"role": "assistant", "content": "original"}, - ) - session.finish_run( - "run-1", - 3, - AgentRunResult( - status=RunStatus.COMPLETED, - stop_reason=StopReason.TEXT_RESPONSE, - turns=1, - output="original", - ), - ) - await storage.save(session) - - loaded = await storage.load("isolated") - assert loaded is not None - loaded.entries[1].message["content"] = "mutated" - loaded.runs.clear() - reloaded = await storage.load("isolated") - - assert reloaded is not None - self.assertEqual(reloaded.messages[1]["content"], "original") - self.assertEqual(len(reloaded.runs), 1) - - async def test_different_session_ids_are_isolated(self) -> None: - storage = MemorySessionStorage() - first_recorder = SessionRecorder(session_id="first", storage=storage) - second_recorder = SessionRecorder(session_id="second", storage=storage) - first_agent = BaseAgent( - SequenceModel([AssistantMessage(content="one")]), - agent_id="first-agent", - event_sink=first_recorder, - ) - second_agent = BaseAgent( - SequenceModel([AssistantMessage(content="two")]), - agent_id="second-agent", - event_sink=second_recorder, - ) - - await first_agent.run(task="first task") - await second_agent.run(task="second task") - first_session = await first_recorder.load() - second_session = await second_recorder.load() - - assert first_session is not None - assert second_session is not None - self.assertEqual(first_session.messages[0]["content"], "first task") - self.assertEqual(second_session.messages[0]["content"], "second task") - - async def test_composite_sink_continues_after_one_observer_fails( - self, - ) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="composite", storage=storage) - observer = RecordingSink() - sink = CompositeAgentEventSink([FailingSink(), recorder, observer]) - agent = BaseAgent( - SequenceModel([AssistantMessage(content="done")]), - agent_id="composite-agent", - event_sink=sink, - ) - - with self.assertLogs("composite-agent", level="WARNING"): - result = await agent.run(task="record despite failure") - session = await recorder.load() - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(len(observer.events), 5) - assert session is not None - self.assertEqual( - [message["content"] for message in session.messages], - ["record despite failure", "done"], - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_agent_steering.py b/tests/test_agent_steering.py deleted file mode 100644 index 93da522..0000000 --- a/tests/test_agent_steering.py +++ /dev/null @@ -1,331 +0,0 @@ -import asyncio -import tempfile -import unittest -from typing import Any - -from ejagent import ( - AgentEvent, - AssistantMessage, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - ControlStatus, - JsonlSessionStorage, - MethodToolHandler, - ModelAdapter, - ModelToolCall, - RunStatus, - SessionRecorder, - SessionRecordKind, - SteeringApplied, - SteeringDiscarded, - StepOutcome, - StopReason, - ToolCompleted, - TurnCompleted, -) - -WAIT_TOOL = { - "type": "function", - "function": { - "name": "wait", - "description": "Wait for the test to release this tool.", - "parameters": {"type": "object", "properties": {}}, - }, -} - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class BlockingFirstModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - self.contexts: list[tuple[dict[str, Any], ...]] = [] - self.first_started = asyncio.Event() - self.release_first = asyncio.Event() - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.contexts.append(context.agent_messages) - if len(self.contexts) == 1: - self.first_started.set() - await self.release_first.wait() - return self.responses.pop(0) - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - self.contexts: list[tuple[dict[str, Any], ...]] = [] - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.contexts.append(context.agent_messages) - return self.responses.pop(0) - - -class NeverCompletingModel(ModelAdapter): - def __init__(self) -> None: - self.started = asyncio.Event() - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.started.set() - await asyncio.Future() - - -class BlockingToolHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((WAIT_TOOL,)) - self.started = asyncio.Event() - self.release = asyncio.Event() - - async def do_wait( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - self.started.set() - await self.release.wait() - return StepOutcome({"released": True}) - - -def payloads(sink: RecordingSink) -> list[object]: - return [event.payload for event in sink.events] - - -class AgentSteeringTests(unittest.IsolatedAsyncioTestCase): - async def test_steering_during_response_forces_next_turn_in_fifo_order( - self, - ) -> None: - model = BlockingFirstModel( - [ - AssistantMessage(content="superseded answer"), - AssistantMessage(content="corrected answer"), - ] - ) - sink = RecordingSink() - agent = BaseAgent( - model, - agent_id="steering-fifo", - steering_queue_capacity=2, - event_sink=sink, - ) - run = asyncio.create_task(agent.run(task="initial task")) - await model.first_started.wait() - - first = await agent.steer("first correction") - second = await agent.steer("second correction") - self.assertTrue(first.accepted) - self.assertTrue(second.accepted) - self.assertEqual(first.queue_size, 1) - self.assertEqual(second.queue_size, 2) - self.assertEqual(agent.pending_steering_count, 2) - model.release_first.set() - - result = await run - - self.assertEqual(result.output, "corrected answer") - self.assertEqual(result.turns, 2) - self.assertEqual(agent.pending_steering_count, 0) - self.assertEqual( - [message.get("content") for message in model.contexts[1]], - [ - agent.system_prompt, - "initial task", - "superseded answer", - "first correction", - "second correction", - ], - ) - applied = [ - payload - for payload in payloads(sink) - if isinstance(payload, SteeringApplied) - ] - self.assertEqual( - [event.control.input_id for event in applied], - [first.control.input_id, second.control.input_id], - ) - self.assertEqual([event.target_turn for event in applied], [2, 2]) - - async def test_queue_capacity_and_idle_submission_are_explicit(self) -> None: - model = BlockingFirstModel( - [AssistantMessage(content="old"), AssistantMessage(content="new")] - ) - agent = BaseAgent( - model, - agent_id="steering-capacity", - steering_queue_capacity=1, - ) - - idle = await agent.steer("too early") - self.assertEqual(idle.status, ControlStatus.AGENT_IDLE) - run = asyncio.create_task(agent.run(task="task")) - await model.first_started.wait() - accepted = await agent.steer("accepted") - full = await agent.steer("rejected because full") - self.assertEqual(accepted.status, ControlStatus.ACCEPTED) - self.assertEqual(full.status, ControlStatus.QUEUE_FULL) - self.assertEqual(full.queue_size, 1) - model.release_first.set() - await run - - after = await agent.steer("too late") - self.assertEqual(after.status, ControlStatus.AGENT_IDLE) - with self.assertRaisesRegex(ValueError, "must not be empty"): - await agent.steer(" ") - - async def test_tool_finishes_before_steering_reaches_next_model_call( - self, - ) -> None: - tool_call = ModelToolCall(id="wait-1", name="wait", arguments="{}") - model = SequenceModel( - [ - AssistantMessage(tool_calls=(tool_call,)), - AssistantMessage(content="finished after steering"), - ] - ) - handler = BlockingToolHandler() - observer = RecordingSink() - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - recorder = SessionRecorder(session_id="steered-tool", storage=storage) - agent = BaseAgent( - model, - agent_id="steered-tool-agent", - handlers=[handler], - event_sink=CompositeAgentEventSink([recorder, observer]), - ) - run = asyncio.create_task(agent.run(task="use the tool")) - await handler.started.wait() - - receipt = await agent.steer("change direction after the tool") - - self.assertTrue(receipt.accepted) - self.assertFalse(run.done()) - self.assertFalse(handler.release.is_set()) - handler.release.set() - result = await run - - self.assertEqual(result.status, RunStatus.COMPLETED) - second_context = model.contexts[1] - roles = [message["role"] for message in second_context] - self.assertEqual( - roles[-4:], - ["user", "assistant", "tool", "user"], - ) - self.assertEqual( - second_context[-1]["content"], - "change direction after the tool", - ) - observed = payloads(observer) - tool_completed_at = next( - index - for index, payload in enumerate(observed) - if isinstance(payload, ToolCompleted) - ) - first_turn_completed_at = next( - index - for index, payload in enumerate(observed) - if isinstance(payload, TurnCompleted) and payload.turn == 1 - ) - steering_at = next( - index - for index, payload in enumerate(observed) - if isinstance(payload, SteeringApplied) - ) - self.assertLess(tool_completed_at, first_turn_completed_at) - self.assertLess(first_turn_completed_at, steering_at) - - session = await recorder.load() - assert session is not None - self.assertEqual( - [message["role"] for message in session.messages], - ["user", "assistant", "tool", "user", "assistant"], - ) - records = await storage.records("steered-tool") - self.assertEqual( - [record.kind for record in records], - [ - SessionRecordKind.RUN_STARTED, - SessionRecordKind.MESSAGE_APPENDED, - SessionRecordKind.MESSAGES_APPENDED, - SessionRecordKind.STEERING_APPLIED, - SessionRecordKind.MESSAGE_APPENDED, - SessionRecordKind.RUN_FINISHED, - ], - ) - steering_record = records[3] - self.assertEqual(steering_record.data["input_id"], receipt.control.input_id) - self.assertEqual(steering_record.data["target_turn"], 2) - - async def test_cancelled_run_discards_unapplied_steering_without_persisting_it( - self, - ) -> None: - model = NeverCompletingModel() - observer = RecordingSink() - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - recorder = SessionRecorder(session_id="discarded", storage=storage) - agent = BaseAgent( - model, - agent_id="discarded-agent", - event_sink=CompositeAgentEventSink([recorder, observer]), - ) - run = asyncio.create_task(agent.run(task="wait")) - await model.started.wait() - accepted = await agent.steer("never applied") - self.assertTrue(accepted.accepted) - - self.assertTrue(agent.abort("cancel before next model call")) - closing = await agent.steer("also too late") - self.assertEqual(closing.status, ControlStatus.RUN_CLOSING) - result = await run - - self.assertEqual(result.status, RunStatus.CANCELLED) - discarded = [ - payload - for payload in payloads(observer) - if isinstance(payload, SteeringDiscarded) - ] - self.assertEqual(len(discarded), 1) - self.assertEqual( - discarded[0].control.input_id, - accepted.control.input_id, - ) - self.assertEqual(discarded[0].reason, StopReason.EXTERNAL_ABORT) - self.assertFalse( - any( - isinstance(payload, SteeringApplied) - for payload in payloads(observer) - ) - ) - records = await storage.records("discarded") - self.assertNotIn( - SessionRecordKind.STEERING_APPLIED, - [record.kind for record in records], - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_agent_streaming.py b/tests/test_agent_streaming.py deleted file mode 100644 index fdd22bf..0000000 --- a/tests/test_agent_streaming.py +++ /dev/null @@ -1,577 +0,0 @@ -import asyncio -import unittest -from collections.abc import AsyncIterator -from types import SimpleNamespace -from typing import Any - -from ejagent import ( - AgentContextBuilder, - AgentEvent, - AgentFinished, - AgentStarted, - AgentState, - AssistantMessage, - AssistantTextDelta, - AssistantThinkingDelta, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - MemorySessionStorage, - MessageCompleted, - MethodToolHandler, - ModelAdapter, - ModelConfig, - ModelResponseCompleted, - ModelStreamEvent, - ModelTextDelta, - ModelThinkingDelta, - ModelToolCall, - OpenAIModelAdapter, - RunStatus, - SessionRecorder, - StepOutcome, - StopReason, - ToolControl, - TurnCompleted, - TurnStarted, -) - -FINISH_TOOL = { - "type": "function", - "function": { - "name": "finish", - "description": "Finish the current task.", - "parameters": { - "type": "object", - "properties": {"summary": {"type": "string"}}, - "required": ["summary"], - }, - }, -} - -TEST_CONFIG = ModelConfig( - model="test-model", - api_key="test-key", - base_url="https://example.invalid", -) - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class StreamingTextModel(ModelAdapter): - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - raise AssertionError("stream() should be used by the orchestrator") - - async def stream( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - yield ModelTextDelta("Hel") - yield ModelTextDelta("lo") - yield ModelResponseCompleted(AssistantMessage(content="Hello")) - - -class CompleteOnlyModel(ModelAdapter): - def __init__(self) -> None: - self.calls = 0 - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.calls += 1 - return AssistantMessage(content="fallback") - - -class StreamingThinkingModel(ModelAdapter): - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - raise AssertionError("stream() should be used by the orchestrator") - - async def stream( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - yield ModelThinkingDelta("inspect ") - yield ModelThinkingDelta("context") - yield ModelTextDelta("final") - yield ModelResponseCompleted(AssistantMessage(content="final")) - - -class BlockingStreamModel(ModelAdapter): - def __init__(self) -> None: - self.blocked = asyncio.Event() - self.closed = asyncio.Event() - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - raise AssertionError("stream() should be used by the orchestrator") - - async def stream( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - yield ModelTextDelta("partial") - self.blocked.set() - try: - await asyncio.Future() - finally: - self.closed.set() - - -class BlockingThinkingStreamModel(ModelAdapter): - def __init__(self) -> None: - self.blocked = asyncio.Event() - self.closed = asyncio.Event() - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - raise AssertionError("stream() should be used by the orchestrator") - - async def stream( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - yield ModelThinkingDelta("unfinished reasoning") - self.blocked.set() - try: - await asyncio.Future() - finally: - self.closed.set() - - -class ToolStreamModel(ModelAdapter): - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - raise AssertionError("stream() should be used by the orchestrator") - - async def stream( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - yield ModelResponseCompleted( - AssistantMessage( - tool_calls=( - ModelToolCall( - id="finish-1", - name="finish", - arguments='{"summary":"streamed tool"}', - ), - ) - ) - ) - - -class MissingTerminalModel(ModelAdapter): - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - raise AssertionError("stream() should be used by the orchestrator") - - async def stream( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - yield ModelTextDelta("unfinished") - - -class FinishHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((FINISH_TOOL,)) - - async def do_finish( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - return StepOutcome( - {"summary": arguments["summary"]}, - control=ToolControl.COMPLETE, - ) - - -class FakeOpenAIStream: - def __init__(self, chunks: list[Any]) -> None: - self.chunks = list(chunks) - self.closed = False - - def __aiter__(self) -> "FakeOpenAIStream": - return self - - async def __anext__(self) -> Any: - if not self.chunks: - raise StopAsyncIteration - return self.chunks.pop(0) - - async def close(self) -> None: - self.closed = True - - -class FakeOpenAICompletions: - def __init__(self, response: FakeOpenAIStream) -> None: - self.response = response - self.calls: list[dict[str, Any]] = [] - - async def create(self, **kwargs: Any) -> FakeOpenAIStream: - self.calls.append(kwargs) - return self.response - - -def chunk( - *, - content: str | None = None, - tool_calls: list[Any] | None = None, - finish_reason: str | None = None, - reasoning_content: str | None = None, - reasoning: str | None = None, - reasoning_text: str | None = None, -) -> Any: - return SimpleNamespace( - choices=[ - SimpleNamespace( - finish_reason=finish_reason, - delta=SimpleNamespace( - content=content, - tool_calls=tool_calls, - reasoning_content=reasoning_content, - reasoning=reasoning, - reasoning_text=reasoning_text, - ), - ) - ] - ) - - -def payloads(sink: RecordingSink) -> list[object]: - return [event.payload for event in sink.events] - - -class AgentStreamingTests(unittest.IsolatedAsyncioTestCase): - async def test_text_deltas_are_observed_before_atomic_message_commit( - self, - ) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="stream", storage=storage) - observer = RecordingSink() - agent = BaseAgent( - StreamingTextModel(), - agent_id="stream-agent", - event_sink=CompositeAgentEventSink([observer, recorder]), - ) - - result = await agent.run(task="stream text") - session = await recorder.load() - - self.assertEqual(result.output, "Hello") - self.assertEqual( - [type(payload) for payload in payloads(observer)], - [ - AgentStarted, - TurnStarted, - AssistantTextDelta, - AssistantTextDelta, - MessageCompleted, - TurnCompleted, - AgentFinished, - ], - ) - deltas = [ - payload.delta - for payload in payloads(observer) - if isinstance(payload, AssistantTextDelta) - ] - self.assertEqual(deltas, ["Hel", "lo"]) - self.assertEqual(agent.messages[-1]["content"], "Hello") - assert session is not None - self.assertEqual( - [message["content"] for message in session.messages], - ["stream text", "Hello"], - ) - self.assertEqual( - [entry.sequence for entry in session.entries], - [1, 5], - ) - - async def test_complete_only_adapter_uses_default_stream_fallback( - self, - ) -> None: - model = CompleteOnlyModel() - sink = RecordingSink() - agent = BaseAgent(model, agent_id="fallback", event_sink=sink) - - result = await agent.run(task="fallback") - - self.assertEqual(result.output, "fallback") - self.assertEqual(model.calls, 1) - self.assertFalse( - any(isinstance(payload, AssistantTextDelta) for payload in payloads(sink)) - ) - - async def test_thinking_deltas_are_observable_but_not_persisted( - self, - ) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="thinking", storage=storage) - observer = RecordingSink() - agent = BaseAgent( - StreamingThinkingModel(), - agent_id="thinking-agent", - event_sink=CompositeAgentEventSink([observer, recorder]), - ) - - result = await agent.run(task="reason then answer") - session = await recorder.load() - - self.assertEqual(result.output, "final") - self.assertEqual( - [type(payload) for payload in payloads(observer)], - [ - AgentStarted, - TurnStarted, - AssistantThinkingDelta, - AssistantThinkingDelta, - AssistantTextDelta, - MessageCompleted, - TurnCompleted, - AgentFinished, - ], - ) - thinking = [ - payload.delta - for payload in payloads(observer) - if isinstance(payload, AssistantThinkingDelta) - ] - self.assertEqual(thinking, ["inspect ", "context"]) - self.assertNotIn("thinking", agent.messages[-1]) - assert session is not None - self.assertEqual( - [message["content"] for message in session.messages], - ["reason then answer", "final"], - ) - - async def test_abort_discards_partial_message_but_finishes_session( - self, - ) -> None: - model = BlockingStreamModel() - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="abort-stream", storage=storage) - observer = RecordingSink() - agent = BaseAgent( - model, - agent_id="abort-stream", - event_sink=CompositeAgentEventSink([observer, recorder]), - ) - run = asyncio.create_task(agent.run(task="abort partial")) - await model.blocked.wait() - - agent.abort("stop streaming") - result = await run - session = await recorder.load() - - self.assertEqual(result.status, RunStatus.CANCELLED) - self.assertEqual(result.stop_reason, StopReason.EXTERNAL_ABORT) - self.assertTrue(model.closed.is_set()) - self.assertFalse( - any(isinstance(payload, MessageCompleted) for payload in payloads(observer)) - ) - self.assertIsInstance(payloads(observer)[-2], TurnCompleted) - self.assertIsInstance(payloads(observer)[-1], AgentFinished) - self.assertEqual( - [message["role"] for message in agent.messages[-1:]], - ["user"], - ) - assert session is not None - self.assertEqual( - [message["role"] for message in session.messages], - ["user"], - ) - self.assertTrue(session.runs[0].finished) - - async def test_streamed_final_tool_call_executes_normally(self) -> None: - agent = BaseAgent( - ToolStreamModel(), - agent_id="stream-tool", - handlers=[FinishHandler()], - ) - - result = await agent.run(task="finish with a tool") - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(result.stop_reason, StopReason.TOOL_COMPLETION) - self.assertIn("streamed tool", result.output or "") - - async def test_abort_during_thinking_discards_provisional_reasoning( - self, - ) -> None: - model = BlockingThinkingStreamModel() - storage = MemorySessionStorage() - recorder = SessionRecorder( - session_id="abort-thinking", - storage=storage, - ) - observer = RecordingSink() - agent = BaseAgent( - model, - agent_id="abort-thinking", - event_sink=CompositeAgentEventSink([observer, recorder]), - ) - run = asyncio.create_task(agent.run(task="abort reasoning")) - await model.blocked.wait() - - agent.abort("stop thinking") - result = await run - session = await recorder.load() - - self.assertEqual(result.status, RunStatus.CANCELLED) - self.assertTrue(model.closed.is_set()) - self.assertTrue( - any( - isinstance(payload, AssistantThinkingDelta) - for payload in payloads(observer) - ) - ) - self.assertFalse( - any(isinstance(payload, MessageCompleted) for payload in payloads(observer)) - ) - assert session is not None - self.assertEqual( - [message["role"] for message in session.messages], - ["user"], - ) - - async def test_missing_stream_terminal_event_is_runtime_failure(self) -> None: - agent = BaseAgent(MissingTerminalModel(), agent_id="missing-terminal") - - result = await agent.run(task="malformed stream") - - self.assertEqual(result.status, RunStatus.FAILED) - self.assertEqual(result.stop_reason, StopReason.RUNTIME_ERROR) - self.assertIn("without a completed response", result.error or "") - self.assertEqual(agent.messages[-1]["role"], "user") - - async def test_openai_stream_normalizes_text_and_tool_call_chunks( - self, - ) -> None: - first_tool_delta = SimpleNamespace( - index=0, - id="call-1", - function=SimpleNamespace(name="fin", arguments='{"sum'), - ) - second_tool_delta = SimpleNamespace( - index=0, - id=None, - function=SimpleNamespace( - name="ish", - arguments='mary":"ok"}', - ), - ) - response = FakeOpenAIStream( - [ - chunk( - reasoning_content="reason ", - reasoning="reason ", - ), - chunk(reasoning_text="carefully"), - chunk(content="Hello "), - chunk(content="world"), - chunk(tool_calls=[first_tool_delta]), - chunk(tool_calls=[second_tool_delta]), - chunk(finish_reason="tool_calls"), - SimpleNamespace(choices=[]), - ] - ) - completions = FakeOpenAICompletions(response) - client = SimpleNamespace(chat=SimpleNamespace(completions=completions)) - adapter = OpenAIModelAdapter( - TEST_CONFIG, - client=client, # type: ignore[arg-type] - ) - context = AgentContextBuilder().build(AgentState()) - - events = [event async for event in adapter.stream(context)] - - self.assertEqual( - [event.delta for event in events if isinstance(event, ModelTextDelta)], - ["Hello ", "world"], - ) - self.assertEqual( - [event.delta for event in events if isinstance(event, ModelThinkingDelta)], - ["reason ", "carefully"], - ) - terminal = events[-1] - assert isinstance(terminal, ModelResponseCompleted) - self.assertEqual(terminal.message.content, "Hello world") - self.assertEqual(terminal.message.tool_calls[0].id, "call-1") - self.assertEqual(terminal.message.tool_calls[0].name, "finish") - self.assertEqual( - terminal.message.tool_calls[0].arguments, - '{"summary":"ok"}', - ) - self.assertTrue(completions.calls[0]["stream"]) - self.assertTrue(response.closed) - - async def test_openai_stream_rejects_silent_incomplete_response( - self, - ) -> None: - response = FakeOpenAIStream([chunk(content="partial")]) - completions = FakeOpenAICompletions(response) - client = SimpleNamespace(chat=SimpleNamespace(completions=completions)) - adapter = OpenAIModelAdapter( - TEST_CONFIG, - client=client, # type: ignore[arg-type] - ) - context = AgentContextBuilder().build(AgentState()) - - with self.assertRaisesRegex(RuntimeError, "without finish_reason"): - _ = [event async for event in adapter.stream(context)] - - self.assertTrue(response.closed) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_anthropic_provider.py b/tests/test_anthropic_provider.py new file mode 100644 index 0000000..71dcac4 --- /dev/null +++ b/tests/test_anthropic_provider.py @@ -0,0 +1,391 @@ +from __future__ import annotations + +import unittest +from collections.abc import Sequence +from types import SimpleNamespace +from typing import Any + +from ejagent.contracts import ( + AssistantMessage, + CancellationSource, + ContextSummary, + FailureCode, + ModelCallError, + ModelProtocolError, + ModelRequest, + ModelResponseCompleted, + ModelTextDelta, + ModelThinkingDelta, + SystemMessage, + ToolCall, + ToolDefinition, + ToolResultMessage, + TransientInstruction, + UserMessage, +) +from ejagent.providers import ( + AnthropicConfig, + AnthropicModelPort, + ModelConfig, + OpenAIModelPort, +) + + +class FakeAnthropicStream: + def __init__(self, events: Sequence[Any]) -> None: + self._events = iter(events) + + def __aiter__(self) -> FakeAnthropicStream: + return self + + async def __anext__(self) -> Any: + try: + return next(self._events) + except StopIteration as exc: + raise StopAsyncIteration from exc + + +class FakeAnthropicManager: + def __init__(self, result: Sequence[Any] | Exception) -> None: + self.result = result + self.exited = False + + async def __aenter__(self) -> FakeAnthropicStream: + if isinstance(self.result, Exception): + raise self.result + return FakeAnthropicStream(self.result) + + async def __aexit__(self, *_: Any) -> None: + self.exited = True + + +class FakeAnthropicMessages: + def __init__(self, results: Sequence[Sequence[Any] | Exception]) -> None: + self.results = list(results) + self.requests: list[dict[str, Any]] = [] + self.managers: list[FakeAnthropicManager] = [] + + def stream(self, **kwargs: Any) -> FakeAnthropicManager: + self.requests.append(kwargs) + manager = FakeAnthropicManager(self.results.pop(0)) + self.managers.append(manager) + return manager + + +class FakeAnthropicClient: + def __init__(self, messages: FakeAnthropicMessages) -> None: + self.messages = messages + + +def _anthropic_config() -> AnthropicConfig: + return AnthropicConfig( + model="test-claude", + api_key="test-key", + max_tokens=512, + base_url="https://anthropic.invalid", + ) + + +def _text_events(text: str = "ok") -> list[dict[str, Any]]: + return [ + { + "type": "message_start", + "message": {"usage": {"input_tokens": 2, "output_tokens": 1}}, + }, + { + "type": "content_block_start", + "index": 0, + "content_block": {"type": "text", "text": ""}, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "text_delta", "text": text}, + }, + {"type": "content_block_stop", "index": 0}, + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": {"output_tokens": 2}, + }, + {"type": "message_stop"}, + ] + + +class AnthropicModelPortTests(unittest.IsolatedAsyncioTestCase): + async def test_serializes_content_blocks_and_normalizes_stream(self) -> None: + events = [ + { + "type": "message_start", + "message": { + "usage": { + "input_tokens": 7, + "output_tokens": 1, + "cache_read_input_tokens": 2, + "cache_creation_input_tokens": 1, + } + }, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": {"type": "thinking_delta", "thinking": "想"}, + }, + { + "type": "content_block_delta", + "index": 1, + "delta": {"type": "text_delta", "text": "答"}, + }, + { + "type": "content_block_start", + "index": 2, + "content_block": { + "type": "tool_use", + "id": "toolu-2", + "name": "lookup", + "input": {}, + }, + }, + { + "type": "content_block_delta", + "index": 2, + "delta": { + "type": "input_json_delta", + "partial_json": '{"q":', + }, + }, + { + "type": "content_block_delta", + "index": 2, + "delta": { + "type": "input_json_delta", + "partial_json": '"值"}', + }, + }, + { + "type": "message_delta", + "delta": {"stop_reason": "tool_use"}, + "usage": {"output_tokens": 5}, + }, + {"type": "message_stop"}, + ] + messages = FakeAnthropicMessages([events]) + port = AnthropicModelPort( + _anthropic_config(), + client=FakeAnthropicClient(messages), + ) + await port.start() + request = ModelRequest( + messages=( + SystemMessage("stable"), + UserMessage("hello"), + AssistantMessage( + tool_calls=( + ToolCall("toolu-0", "lookup", {"q": "first"}), + ToolCall("toolu-1", "lookup", {"q": "second"}), + ) + ), + ToolResultMessage("toolu-0", "lookup", {"found": True}), + ToolResultMessage("toolu-1", "lookup", "missing", is_error=True), + ContextSummary(1, 3, "summary", "compact-v1"), + TransientInstruction("focus", "steering"), + ), + tools=( + ToolDefinition( + "lookup", + description="Lookup a value.", + input_schema={"type": "object"}, + ), + ), + ) + + normalized = [ + event + async for event in port.stream( + request, + cancellation=CancellationSource().token, + ) + ] + + self.assertEqual(normalized[0], ModelThinkingDelta("想")) + self.assertEqual(normalized[1], ModelTextDelta("答")) + completed = normalized[2] + self.assertIsInstance(completed, ModelResponseCompleted) + assert isinstance(completed, ModelResponseCompleted) + self.assertEqual( + completed.message, + AssistantMessage( + "答", + (ToolCall("toolu-2", "lookup", {"q": "值"}),), + ), + ) + assert completed.usage is not None + self.assertEqual(completed.usage.input_tokens, 10) + self.assertEqual(completed.usage.output_tokens, 5) + self.assertEqual(completed.usage.cache_read_tokens, 2) + self.assertEqual(completed.usage.cache_write_tokens, 1) + + sent = messages.requests[0] + self.assertEqual(sent["system"].split("\n\n")[0], "stable") + self.assertIn("Derived summary", sent["system"]) + self.assertIn("Transient instruction", sent["system"]) + self.assertEqual(sent["messages"][0]["role"], "user") + assistant_blocks = sent["messages"][1]["content"] + self.assertEqual(assistant_blocks[0]["type"], "tool_use") + self.assertEqual(assistant_blocks[0]["input"], {"q": "first"}) + result_blocks = sent["messages"][2]["content"] + self.assertEqual(len(result_blocks), 2) + self.assertEqual(result_blocks[0]["content"], '{"found":true}') + self.assertTrue(result_blocks[1]["is_error"]) + self.assertEqual(sent["tools"][0]["input_schema"], {"type": "object"}) + self.assertTrue(messages.managers[0].exited) + + async def test_rejects_incomplete_tool_input_as_protocol_error(self) -> None: + events = [ + { + "type": "content_block_start", + "index": 0, + "content_block": { + "type": "tool_use", + "id": "toolu-0", + "name": "lookup", + "input": {}, + }, + }, + { + "type": "content_block_delta", + "index": 0, + "delta": { + "type": "input_json_delta", + "partial_json": "not-json", + }, + }, + { + "type": "message_delta", + "delta": {"stop_reason": "tool_use"}, + }, + {"type": "message_stop"}, + ] + port = AnthropicModelPort( + _anthropic_config(), + client=FakeAnthropicClient(FakeAnthropicMessages([events])), + ) + await port.start() + + with self.assertRaises(ModelProtocolError): + async for _ in port.stream( + ModelRequest(messages=(UserMessage("lookup"),)), + cancellation=CancellationSource().token, + ): + pass + + async def test_normalizes_provider_timeout(self) -> None: + port = AnthropicModelPort( + _anthropic_config(), + client=FakeAnthropicClient(FakeAnthropicMessages([TimeoutError("late")])), + ) + await port.start() + + with self.assertRaises(ModelCallError) as raised: + async for _ in port.stream( + ModelRequest(messages=(UserMessage("hello"),)), + cancellation=CancellationSource().token, + ): + pass + + self.assertEqual(raised.exception.code, FailureCode.TIMEOUT) + self.assertTrue(raised.exception.retryable) + + +class _FakeOpenAIStream: + def __init__(self) -> None: + self._chunks = iter( + [ + SimpleNamespace( + choices=[ + SimpleNamespace( + finish_reason="stop", + delta=SimpleNamespace( + content="ok", + reasoning_content=None, + tool_calls=(), + ), + ) + ], + usage=None, + ) + ] + ) + + def __aiter__(self) -> _FakeOpenAIStream: + return self + + async def __anext__(self) -> Any: + try: + return next(self._chunks) + except StopIteration as exc: + raise StopAsyncIteration from exc + + async def aclose(self) -> None: + pass + + +class _FakeOpenAICompletions: + async def create(self, **_: Any) -> _FakeOpenAIStream: + return _FakeOpenAIStream() + + +async def _assert_model_port_contract( + case: unittest.TestCase, + port: Any, +) -> None: + request = ModelRequest(messages=(UserMessage("hello"),)) + with case.assertRaises(RuntimeError): + async for _ in port.stream( + request, + cancellation=CancellationSource().token, + ): + pass + + await port.start() + events = [ + event + async for event in port.stream( + request, + cancellation=CancellationSource().token, + ) + ] + case.assertEqual(events[0], ModelTextDelta("ok")) + case.assertIsInstance(events[-1], ModelResponseCompleted) + completed = events[-1] + assert isinstance(completed, ModelResponseCompleted) + case.assertEqual(completed.message, AssistantMessage("ok")) + await port.shutdown() + with case.assertRaises(RuntimeError): + async for _ in port.stream( + request, + cancellation=CancellationSource().token, + ): + pass + + +class ModelPortContractTests(unittest.IsolatedAsyncioTestCase): + async def test_openai_and_anthropic_share_kernel_stream_contract(self) -> None: + openai = OpenAIModelPort( + ModelConfig("test", "key", "https://openai.invalid"), + client=SimpleNamespace( + chat=SimpleNamespace(completions=_FakeOpenAICompletions()) + ), + ) + anthropic = AnthropicModelPort( + _anthropic_config(), + client=FakeAnthropicClient(FakeAnthropicMessages([_text_events()])), + ) + + for port in (openai, anthropic): + with self.subTest(port=type(port).__name__): + await _assert_model_port_contract(self, port) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_architecture_boundaries.py b/tests/test_architecture_boundaries.py new file mode 100644 index 0000000..8a16c0e --- /dev/null +++ b/tests/test_architecture_boundaries.py @@ -0,0 +1,49 @@ +from __future__ import annotations + +import ast +import unittest +from pathlib import Path + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +SOURCE_ROOT = PROJECT_ROOT / "src" / "ejagent" +LOW_LEVEL_PACKAGES = ( + SOURCE_ROOT / "contracts", + SOURCE_ROOT / "context", + SOURCE_ROOT / "kernel", + SOURCE_ROOT / "harness", +) +FORBIDDEN_IMPORTS = ( + "ejagent.agent", + "ejagent.handlers", + "ejagent.middleware", + "ejagent.plugins", + "ejagent.providers", + "ejagent.session", +) + + +class ArchitectureBoundaryTests(unittest.TestCase): + def test_new_core_does_not_import_legacy_layers(self) -> None: + violations: list[str] = [] + + for package in LOW_LEVEL_PACKAGES: + for path in sorted(package.rglob("*.py")): + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + for node in ast.walk(tree): + modules: tuple[str, ...] + if isinstance(node, ast.ImportFrom): + modules = (node.module or "",) + elif isinstance(node, ast.Import): + modules = tuple(alias.name for alias in node.names) + else: + continue + for module in modules: + if module.startswith(FORBIDDEN_IMPORTS): + relative = path.relative_to(PROJECT_ROOT) + violations.append(f"{relative}:{node.lineno}: {module}") + + self.assertEqual(violations, []) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_auto_compaction.py b/tests/test_auto_compaction.py deleted file mode 100644 index 27673ee..0000000 --- a/tests/test_auto_compaction.py +++ /dev/null @@ -1,457 +0,0 @@ -import asyncio -import unittest -from collections.abc import AsyncIterator, Mapping, Sequence -from typing import Any - -from ejagent import ( - AgentEvent, - AssistantMessage, - AssistantTextDelta, - AutoCompactionPolicy, - BaseAgent, - CancellationToken, - CompactionCompleted, - CompactionFailed, - CompactionPolicy, - CompactionRequest, - CompactionStarted, - CompactionTrigger, - CompactorOutput, - ContextBudget, - ContextOverflowError, - ContextPressureEvaluated, - MemorySessionStorage, - ModelAdapter, - ModelResponseCompleted, - ModelStreamEvent, - ModelTextDelta, - RunStatus, - SessionRecorder, - SteeringApplied, - StopReason, -) - - -class MarkerTokenEstimator: - def estimate_message(self, message: Mapping[str, Any]) -> int: - return int(message.get("_tokens", 0)) - - def estimate_tools(self, tools: Sequence[Mapping[str, Any]]) -> int: - return 0 - - -class ScriptedStreamModel(ModelAdapter): - def __init__( - self, - outcomes: list[Exception | list[ModelStreamEvent | Exception]], - ) -> None: - self.outcomes = list(outcomes) - self.contexts: list[Any] = [] - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - raise AssertionError("stream() should be used") - - async def stream( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - self.contexts.append(context) - outcome = self.outcomes.pop(0) - if isinstance(outcome, Exception): - raise outcome - for item in outcome: - if isinstance(item, Exception): - raise item - yield item - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class StaticCompactor: - def __init__(self) -> None: - self.requests: list[CompactionRequest] = [] - - async def compact( - self, - request: CompactionRequest, - *, - cancellation: CancellationToken | None = None, - ) -> CompactorOutput: - self.requests.append(request) - return CompactorOutput("summary of old turns", "test-compactor") - - -class FailingCompactor: - async def compact( - self, - request: CompactionRequest, - *, - cancellation: CancellationToken | None = None, - ) -> CompactorOutput: - raise RuntimeError("summary provider failed") - - -class BlockingCompactor: - def __init__(self) -> None: - self.started = asyncio.Event() - - async def compact( - self, - request: CompactionRequest, - *, - cancellation: CancellationToken | None = None, - ) -> CompactorOutput: - self.started.set() - await asyncio.Event().wait() - return CompactorOutput("unreachable", "test-compactor") - - -class ReleasableCompactor: - def __init__(self) -> None: - self.started = asyncio.Event() - self.release = asyncio.Event() - - async def compact( - self, - request: CompactionRequest, - *, - cancellation: CancellationToken | None = None, - ) -> CompactorOutput: - self.started.set() - await self.release.wait() - return CompactorOutput("summary before steering", "test-compactor") - - -def compaction_policy() -> CompactionPolicy: - return CompactionPolicy( - ContextBudget( - context_window=100, - reserve_tokens=10, - keep_recent_tokens=20, - ) - ) - - -def history() -> list[dict[str, Any]]: - return [ - {"role": "user", "content": "old task", "_tokens": 60}, - {"role": "assistant", "content": "old answer", "_tokens": 60}, - {"role": "user", "content": "recent task", "_tokens": 10}, - {"role": "assistant", "content": "recent answer", "_tokens": 10}, - ] - - -def completed(text: str = "done") -> list[ModelStreamEvent | Exception]: - return [ModelResponseCompleted(AssistantMessage(content=text))] - - -class AutoCompactionTests(unittest.IsolatedAsyncioTestCase): - async def test_automatic_behavior_is_off_by_default(self) -> None: - model = ScriptedStreamModel([completed()]) - compactor = StaticCompactor() - agent = BaseAgent( - model, - agent_id="auto-default-off", - compaction_policy=compaction_policy(), - compactor=compactor, - context_token_estimator=MarkerTokenEstimator(), - ) - agent.reset(history=history()) - - result = await agent.run(task="continue") - - self.assertTrue(result.succeeded) - self.assertEqual(compactor.requests, []) - self.assertIn("old task", repr(model.contexts[0].agent_messages)) - - async def test_pressure_compacts_and_rebuilds_context_in_same_run(self) -> None: - model = ScriptedStreamModel([completed()]) - compactor = StaticCompactor() - sink = RecordingSink() - agent = BaseAgent( - model, - agent_id="auto-pressure", - compaction_policy=compaction_policy(), - compactor=compactor, - auto_compaction_policy=AutoCompactionPolicy(), - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.reset(history=history()) - - result = await agent.run(task="continue") - - self.assertTrue(result.succeeded) - self.assertEqual(len(compactor.requests), 1) - self.assertEqual(compactor.requests[0].trigger, CompactionTrigger.PRESSURE) - self.assertNotIn("old task", repr(model.contexts[0].agent_messages)) - self.assertIn("Conversation summary", repr(model.contexts[0].llm_messages)) - relevant = [ - event - for event in sink.events - if isinstance( - event.payload, - (ContextPressureEvaluated, CompactionStarted, CompactionCompleted), - ) - ] - self.assertEqual( - [type(event.payload) for event in relevant], - [ - ContextPressureEvaluated, - CompactionStarted, - CompactionCompleted, - ContextPressureEvaluated, - ], - ) - started = relevant[1].payload - finished = relevant[2].payload - assert isinstance(started, CompactionStarted) - assert isinstance(finished, CompactionCompleted) - self.assertEqual( - started.request.operation_id, - finished.result.operation_id, - ) - self.assertEqual(len({event.run_id for event in sink.events}), 1) - self.assertEqual( - [event.sequence for event in sink.events], - list(range(1, len(sink.events) + 1)), - ) - - async def test_pressure_compactor_failure_stops_safely(self) -> None: - model = ScriptedStreamModel([completed()]) - agent = BaseAgent( - model, - agent_id="auto-failure", - compaction_policy=compaction_policy(), - compactor=FailingCompactor(), - auto_compaction_policy=AutoCompactionPolicy(), - context_token_estimator=MarkerTokenEstimator(), - ) - agent.reset(history=history()) - - result = await agent.run(task="continue") - - self.assertEqual(result.status, RunStatus.FAILED) - self.assertEqual(result.stop_reason, StopReason.COMPACTION_FAILED) - self.assertEqual(result.error, "summary provider failed") - self.assertEqual(model.contexts, []) - self.assertIn("old task", repr(agent.messages)) - self.assertNotIn("_ejagent_summary", repr(agent.messages)) - - async def test_abort_cancels_automatic_compaction(self) -> None: - model = ScriptedStreamModel([completed()]) - compactor = BlockingCompactor() - sink = RecordingSink() - agent = BaseAgent( - model, - agent_id="auto-abort", - compaction_policy=compaction_policy(), - compactor=compactor, - auto_compaction_policy=AutoCompactionPolicy(), - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.reset(history=history()) - - task = asyncio.create_task(agent.run(task="continue")) - await compactor.started.wait() - self.assertTrue(agent.abort("stop automatic compaction")) - result = await task - - self.assertEqual(result.status, RunStatus.CANCELLED) - self.assertEqual(result.stop_reason, StopReason.EXTERNAL_ABORT) - self.assertEqual(model.contexts, []) - failures = [ - event.payload - for event in sink.events - if isinstance(event.payload, CompactionFailed) - ] - self.assertEqual(len(failures), 1) - - async def test_context_overflow_compacts_and_retries_once(self) -> None: - model = ScriptedStreamModel( - [ContextOverflowError("too large"), completed("recovered")] - ) - compactor = StaticCompactor() - agent = BaseAgent( - model, - agent_id="overflow-recovery", - compaction_policy=compaction_policy(), - compactor=compactor, - auto_compaction_policy=AutoCompactionPolicy(compact_on_pressure=False), - context_token_estimator=MarkerTokenEstimator(), - ) - agent.reset(history=history()) - - result = await agent.run(task="continue") - - self.assertTrue(result.succeeded) - self.assertEqual(result.output, "recovered") - self.assertEqual(len(model.contexts), 2) - self.assertEqual(len(compactor.requests), 1) - self.assertEqual(compactor.requests[0].trigger, CompactionTrigger.OVERFLOW) - self.assertIn("old task", repr(model.contexts[0].agent_messages)) - self.assertNotIn("old task", repr(model.contexts[1].agent_messages)) - self.assertEqual(result.usage.request_count, 2) - - async def test_steering_during_overflow_compaction_reaches_retry(self) -> None: - model = ScriptedStreamModel( - [ContextOverflowError("too large"), completed("steered retry")] - ) - compactor = ReleasableCompactor() - sink = RecordingSink() - agent = BaseAgent( - model, - agent_id="overflow-steering", - compaction_policy=compaction_policy(), - compactor=compactor, - auto_compaction_policy=AutoCompactionPolicy(compact_on_pressure=False), - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.reset(history=history()) - run = asyncio.create_task(agent.run(task="continue")) - await compactor.started.wait() - - receipt = await agent.steer("use the corrected direction") - compactor.release.set() - result = await run - - self.assertTrue(receipt.accepted) - self.assertEqual(result.output, "steered retry") - self.assertNotIn( - "use the corrected direction", - repr(model.contexts[0].agent_messages), - ) - self.assertIn( - "use the corrected direction", - repr(model.contexts[1].agent_messages), - ) - applied = [ - event.payload - for event in sink.events - if isinstance(event.payload, SteeringApplied) - ] - self.assertEqual(len(applied), 1) - self.assertEqual(applied[0].target_turn, 1) - - async def test_second_context_overflow_is_structured_failure(self) -> None: - model = ScriptedStreamModel( - [ContextOverflowError("first"), ContextOverflowError("second")] - ) - compactor = StaticCompactor() - agent = BaseAgent( - model, - agent_id="overflow-limit", - compaction_policy=compaction_policy(), - compactor=compactor, - auto_compaction_policy=AutoCompactionPolicy(compact_on_pressure=False), - context_token_estimator=MarkerTokenEstimator(), - ) - agent.reset(history=history()) - - result = await agent.run(task="continue") - - self.assertEqual(result.status, RunStatus.FAILED) - self.assertEqual(result.stop_reason, StopReason.CONTEXT_OVERFLOW) - self.assertEqual(len(model.contexts), 2) - self.assertEqual(len(compactor.requests), 1) - - async def test_overflow_after_delta_is_not_retried(self) -> None: - model = ScriptedStreamModel( - [ - [ - ModelTextDelta("partial"), - ContextOverflowError("late overflow"), - ] - ] - ) - compactor = StaticCompactor() - sink = RecordingSink() - agent = BaseAgent( - model, - agent_id="overflow-after-delta", - compaction_policy=compaction_policy(), - compactor=compactor, - auto_compaction_policy=AutoCompactionPolicy(compact_on_pressure=False), - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.reset(history=history()) - - result = await agent.run(task="continue") - - self.assertEqual(result.stop_reason, StopReason.CONTEXT_OVERFLOW) - self.assertEqual(len(model.contexts), 1) - self.assertEqual(compactor.requests, []) - self.assertEqual( - len( - [ - event - for event in sink.events - if isinstance(event.payload, AssistantTextDelta) - ] - ), - 1, - ) - - async def test_session_records_automatic_compaction_operation(self) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="auto-session", storage=storage) - model = ScriptedStreamModel([completed()]) - agent = BaseAgent( - model, - agent_id="auto-session-agent", - compaction_policy=compaction_policy(), - compactor=StaticCompactor(), - auto_compaction_policy=AutoCompactionPolicy(), - context_token_estimator=MarkerTokenEstimator(), - event_sink=recorder, - ) - agent.reset(history=history()) - - result = await agent.run(task="continue") - session = await recorder.load() - - self.assertTrue(result.succeeded) - assert session is not None - self.assertEqual(len(session.compactions), 1) - self.assertTrue(session.compactions[0].operation_id) - self.assertNotEqual( - session.compactions[0].operation_id, - session.runs[0].run_id, - ) - - async def test_enabled_policy_requires_compaction_dependencies(self) -> None: - model = ScriptedStreamModel([completed()]) - with self.assertRaisesRegex(ValueError, "CompactionPolicy"): - BaseAgent( - model, - agent_id="auto-missing-policy", - compactor=StaticCompactor(), - auto_compaction_policy=AutoCompactionPolicy(), - ) - with self.assertRaisesRegex(ValueError, "Compactor"): - BaseAgent( - model, - agent_id="auto-missing-compactor", - compaction_policy=compaction_policy(), - auto_compaction_policy=AutoCompactionPolicy(), - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_behavior_hooks.py b/tests/test_behavior_hooks.py deleted file mode 100644 index 529f151..0000000 --- a/tests/test_behavior_hooks.py +++ /dev/null @@ -1,364 +0,0 @@ -import asyncio -import tempfile -import unittest -from typing import Any - -from ejagent import ( - AgentEvent, - AgentFinished, - AssistantMessage, - BaseAgent, - BehaviorDecision, - CancellationToken, - CompositeAgentEventSink, - JsonlSessionStorage, - MessageCompleted, - MethodToolHandler, - ModelAdapter, - ModelToolCall, - RunStatus, - SessionRecorder, - StepOutcome, - StopReason, - ToolCompleted, - TurnCompleted, - TurnSnapshot, -) - -ECHO_TOOL = { - "type": "function", - "function": { - "name": "echo", - "description": "Return a test value.", - "parameters": { - "type": "object", - "properties": {"value": {"type": "string"}}, - "required": ["value"], - }, - }, -} - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - self.contexts: list[tuple[dict[str, Any], ...]] = [] - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.contexts.append(context.agent_messages) - return self.responses.pop(0) - - -class EchoHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((ECHO_TOOL,)) - - async def do_echo( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - return StepOutcome({"value": arguments["value"]}) - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class StopHook: - name = "stop-guard" - - def __init__(self, sink: RecordingSink | None = None) -> None: - self.sink = sink - self.snapshots: list[TurnSnapshot] = [] - - async def after_turn( - self, - snapshot: TurnSnapshot, - *, - cancellation: CancellationToken, - ) -> BehaviorDecision: - if self.sink is not None: - assert isinstance(self.sink.events[-1].payload, TurnCompleted) - self.snapshots.append(snapshot) - detached = snapshot.messages - detached[-1]["content"] = "mutated outside AgentState" - return BehaviorDecision.stop("stopped by behavior policy") - - -class OrderedHook: - def __init__( - self, - name: str, - calls: list[str], - decision: BehaviorDecision | None, - ) -> None: - self.name = name - self.calls = calls - self.decision = decision - self.last_content: Any = None - - async def after_turn( - self, - snapshot: TurnSnapshot, - *, - cancellation: CancellationToken, - ) -> BehaviorDecision | None: - self.calls.append(self.name) - self.last_content = snapshot.messages[-1]["content"] - return self.decision - - -class RaisingHook: - name = "broken-hook" - - async def after_turn( - self, - snapshot: TurnSnapshot, - *, - cancellation: CancellationToken, - ) -> BehaviorDecision: - raise ValueError("hook exploded") - - -class InvalidHook: - async def after_turn( - self, - snapshot: TurnSnapshot, - *, - cancellation: CancellationToken, - ) -> Any: - return "invalid" - - -class BlockingHook: - def __init__(self) -> None: - self.started = asyncio.Event() - self.interrupted = asyncio.Event() - - async def after_turn( - self, - snapshot: TurnSnapshot, - *, - cancellation: CancellationToken, - ) -> BehaviorDecision: - self.started.set() - try: - await asyncio.Future() - finally: - self.interrupted.set() - - -def tool_turn(call_id: str = "echo-1") -> AssistantMessage: - return AssistantMessage( - tool_calls=( - ModelToolCall( - id=call_id, - name="echo", - arguments='{"value": "ready"}', - ), - ) - ) - - -class BehaviorHookTests(unittest.IsolatedAsyncioTestCase): - async def test_stop_occurs_after_tool_result_and_persists_terminal_result( - self, - ) -> None: - model = SequenceModel( - [tool_turn(), AssistantMessage(content="must not be requested")] - ) - observer = RecordingSink() - hook = StopHook(observer) - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - recorder = SessionRecorder(session_id="behavior-stop", storage=storage) - agent = BaseAgent( - model, - agent_id="behavior-stop-agent", - handlers=[EchoHandler()], - behavior_hooks=[hook], - event_sink=CompositeAgentEventSink([recorder, observer]), - ) - - result = await agent.run(task="use the tool") - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(result.stop_reason, StopReason.BEHAVIOR_STOP) - self.assertEqual(result.output, "stopped by behavior policy") - self.assertEqual(result.turns, 1) - self.assertTrue(agent.can_continue) - self.assertEqual(len(model.contexts), 1) - self.assertEqual(len(hook.snapshots), 1) - snapshot = hook.snapshots[0] - self.assertEqual(snapshot.agent_id, agent.agent_id) - self.assertEqual(snapshot.turn, 1) - self.assertEqual(snapshot.task, "use the tool") - self.assertEqual( - [message["role"] for message in snapshot.messages[-3:]], - ["user", "assistant", "tool"], - ) - self.assertNotEqual( - agent.messages[-1]["content"], - "mutated outside AgentState", - ) - - payloads = [event.payload for event in observer.events] - tool_completed = next( - index - for index, payload in enumerate(payloads) - if isinstance(payload, ToolCompleted) - ) - turn_completed = next( - index - for index, payload in enumerate(payloads) - if isinstance(payload, TurnCompleted) - ) - self.assertLess(tool_completed, turn_completed) - self.assertIsInstance(payloads[-1], AgentFinished) - session = await recorder.load() - assert session is not None - assert session.runs[-1].result is not None - self.assertEqual( - session.runs[-1].result.stop_reason, - StopReason.BEHAVIOR_STOP, - ) - - async def test_hooks_are_ordered_and_stop_short_circuits_remaining_hooks( - self, - ) -> None: - calls: list[str] = [] - first = OrderedHook("first", calls, None) - second = OrderedHook("second", calls, BehaviorDecision.stop()) - third = OrderedHook("third", calls, BehaviorDecision.continue_run()) - agent = BaseAgent( - SequenceModel([tool_turn(), AssistantMessage(content="unused")]), - agent_id="behavior-order", - handlers=[EchoHandler()], - behavior_hooks=[first, second, third], - ) - - result = await agent.run(task="ordered") - - self.assertEqual(result.stop_reason, StopReason.BEHAVIOR_STOP) - self.assertEqual(calls, ["first", "second"]) - self.assertEqual(first.last_content, second.last_content) - - async def test_continue_decision_allows_the_next_provider_turn(self) -> None: - calls: list[str] = [] - hook = OrderedHook( - "continue", - calls, - BehaviorDecision.continue_run(), - ) - model = SequenceModel( - [tool_turn(), AssistantMessage(content="finished normally")] - ) - agent = BaseAgent( - model, - agent_id="behavior-proceed", - handlers=[EchoHandler()], - behavior_hooks=[hook], - ) - - result = await agent.run(task="proceed") - - self.assertEqual(result.stop_reason, StopReason.TEXT_RESPONSE) - self.assertEqual(result.turns, 2) - self.assertEqual(len(model.contexts), 2) - self.assertEqual(calls, ["continue"]) - - async def test_terminal_text_response_does_not_invoke_after_turn(self) -> None: - hook = RaisingHook() - sink = RecordingSink() - agent = BaseAgent( - SequenceModel([AssistantMessage(content="already terminal")]), - agent_id="behavior-terminal", - behavior_hooks=[hook], - event_sink=sink, - ) - - result = await agent.run(task="finish normally") - - self.assertEqual(result.stop_reason, StopReason.TEXT_RESPONSE) - self.assertEqual( - sum(isinstance(event.payload, MessageCompleted) for event in sink.events), - 1, - ) - - async def test_hook_exception_and_invalid_decision_fail_the_run(self) -> None: - for hook, expected in [ - (RaisingHook(), "hook exploded"), - (InvalidHook(), "BehaviorDecision or None"), - ]: - with self.subTest(hook=type(hook).__name__): - agent = BaseAgent( - SequenceModel([tool_turn()]), - agent_id=f"behavior-{type(hook).__name__}", - handlers=[EchoHandler()], - behavior_hooks=[hook], - ) - - result = await agent.run(task="fail in hook") - - self.assertEqual(result.status, RunStatus.FAILED) - self.assertEqual(result.stop_reason, StopReason.RUNTIME_ERROR) - self.assertIn(expected, result.error or "") - - async def test_abort_interrupts_slow_hook_and_idle_waits_for_it(self) -> None: - hook = BlockingHook() - sink = RecordingSink() - agent = BaseAgent( - SequenceModel([tool_turn()]), - agent_id="behavior-cancel", - handlers=[EchoHandler()], - behavior_hooks=[hook], - event_sink=sink, - ) - run = asyncio.create_task(agent.run(task="block in hook")) - await hook.started.wait() - idle = asyncio.create_task(agent.wait_for_idle()) - await asyncio.sleep(0) - - self.assertFalse(run.done()) - self.assertFalse(idle.done()) - self.assertTrue(agent.abort("cancel slow behavior")) - result = await run - await idle - - self.assertEqual(result.status, RunStatus.CANCELLED) - self.assertEqual(result.stop_reason, StopReason.EXTERNAL_ABORT) - self.assertTrue(hook.interrupted.is_set()) - self.assertIsInstance(sink.events[-1].payload, AgentFinished) - - async def test_continue_run_uses_the_same_behavior_hooks(self) -> None: - hook = StopHook() - model = SequenceModel( - [AssistantMessage(content="initial"), tool_turn("continue-tool")] - ) - agent = BaseAgent( - model, - agent_id="behavior-continue", - handlers=[EchoHandler()], - behavior_hooks=[hook], - ) - await agent.run(task="initial") - - result = await agent.continue_run() - - self.assertEqual(result.stop_reason, StopReason.BEHAVIOR_STOP) - self.assertIsNone(hook.snapshots[-1].task) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_context_compaction.py b/tests/test_context_compaction.py deleted file mode 100644 index bad662a..0000000 --- a/tests/test_context_compaction.py +++ /dev/null @@ -1,477 +0,0 @@ -import asyncio -import unittest -from collections.abc import Mapping, Sequence -from copy import deepcopy -from typing import Any - -from ejagent import ( - AgentEvent, - AssistantMessage, - BaseAgent, - CancellationToken, - CompactionCompleted, - CompactionFailed, - CompactionPolicy, - CompactionRequest, - CompactionStarted, - CompactionStatus, - CompactorOutput, - ContextBudget, - MemorySessionStorage, - ModelAdapter, - SessionRecorder, - SummaryEntry, -) - - -class MarkerTokenEstimator: - def estimate_message(self, message: Mapping[str, Any]) -> int: - return int(message.get("_tokens", 0)) - - def estimate_tools(self, tools: Sequence[Mapping[str, Any]]) -> int: - return 0 - - -class ConstantTokenEstimator: - def estimate_message(self, message: Mapping[str, Any]) -> int: - return 100 - - def estimate_tools(self, tools: Sequence[Mapping[str, Any]]) -> int: - return 0 - - -class QueueModel(ModelAdapter): - def __init__(self, responses: list[str] | None = None) -> None: - self.responses = list(responses or ["done"]) - self.contexts: list[Any] = [] - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.contexts.append(context) - return AssistantMessage(content=self.responses.pop(0)) - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class StaticCompactor: - def __init__(self, outputs: list[str] | None = None) -> None: - self.outputs = list(outputs or ["old work was summarized"]) - self.requests: list[CompactionRequest] = [] - - async def compact( - self, - request: CompactionRequest, - *, - cancellation: CancellationToken | None = None, - ) -> CompactorOutput: - self.requests.append(request) - return CompactorOutput( - content=self.outputs.pop(0), - source="static-test-compactor", - ) - - -class FailingCompactor: - async def compact( - self, - request: CompactionRequest, - *, - cancellation: CancellationToken | None = None, - ) -> CompactorOutput: - raise RuntimeError("summary provider failed") - - -class BlockingCompactor: - def __init__(self) -> None: - self.started = asyncio.Event() - self.stopped = asyncio.Event() - - async def compact( - self, - request: CompactionRequest, - *, - cancellation: CancellationToken | None = None, - ) -> CompactorOutput: - self.started.set() - try: - await asyncio.Event().wait() - finally: - self.stopped.set() - return CompactorOutput("unreachable", "blocking") - - -def policy() -> CompactionPolicy: - return CompactionPolicy( - ContextBudget( - context_window=1000, - reserve_tokens=100, - keep_recent_tokens=50, - ) - ) - - -def history() -> list[dict[str, Any]]: - return [ - {"role": "user", "content": "old task", "_tokens": 100}, - { - "role": "assistant", - "content": None, - "tool_calls": [{"id": "old-call"}], - "_tokens": 10, - }, - { - "role": "tool", - "tool_call_id": "old-call", - "content": "large old output", - "_tokens": 100, - }, - {"role": "assistant", "content": "old answer", "_tokens": 10}, - {"role": "user", "content": "recent task", "_tokens": 30}, - {"role": "assistant", "content": "recent answer", "_tokens": 30}, - ] - - -class ContextCompactionTests(unittest.IsolatedAsyncioTestCase): - async def test_explicit_compaction_requires_policy_and_compactor(self) -> None: - agent = BaseAgent(QueueModel(), agent_id="compact-config") - - with self.assertRaisesRegex(RuntimeError, "requires a Compactor"): - await agent.compact() - with self.assertRaisesRegex(RuntimeError, "requires a CompactionPolicy"): - await agent.compact(compactor=StaticCompactor()) - await agent.wait_for_idle() - - async def test_explicit_compaction_atomically_replaces_old_turns(self) -> None: - model = QueueModel() - compactor = StaticCompactor() - sink = RecordingSink() - agent = BaseAgent( - model, - agent_id="compact-success", - compaction_policy=policy(), - compactor=compactor, - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.reset(history=history()) - - result = await agent.compact() - - self.assertEqual(result.status, CompactionStatus.COMPLETED) - self.assertTrue(result.completed) - self.assertIsInstance(result.summary, SummaryEntry) - assert result.summary is not None - self.assertEqual(result.summary.summarized_message_count, 4) - self.assertEqual(result.summary.tokens_before, 280) - self.assertEqual( - [message["role"] for message in agent.messages], - ["system", "system", "user", "assistant"], - ) - summary_message = agent.messages[1] - self.assertIn("_ejagent_summary", summary_message) - self.assertNotIn("old-call", repr(agent.messages)) - self.assertEqual( - [type(event.payload) for event in sink.events], - [CompactionStarted, CompactionCompleted], - ) - self.assertEqual([event.sequence for event in sink.events], [1, 2]) - self.assertEqual(len({event.run_id for event in sink.events}), 1) - self.assertEqual( - [message["role"] for message in result.messages], - ["system", "user", "assistant"], - ) - - await agent.run(task="continue from the summary") - sent_summary = model.contexts[0].llm_messages[1] - self.assertNotIn("_ejagent_summary", sent_summary) - self.assertIn("Conversation summary", sent_summary["content"]) - - async def test_compactor_failure_keeps_history_unchanged(self) -> None: - sink = RecordingSink() - agent = BaseAgent( - QueueModel(), - agent_id="compact-failure", - compaction_policy=policy(), - compactor=FailingCompactor(), - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.reset(history=history()) - before = deepcopy(agent.messages) - - result = await agent.compact() - - self.assertEqual(result.status, CompactionStatus.FAILED) - self.assertEqual(result.error, "summary provider failed") - self.assertEqual(agent.messages, before) - self.assertEqual( - [type(event.payload) for event in sink.events], - [CompactionStarted, CompactionFailed], - ) - - async def test_abort_cancels_compactor_and_preserves_history(self) -> None: - compactor = BlockingCompactor() - sink = RecordingSink() - agent = BaseAgent( - QueueModel(), - agent_id="compact-abort", - compaction_policy=policy(), - compactor=compactor, - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.reset(history=history()) - before = deepcopy(agent.messages) - - task = asyncio.create_task(agent.compact()) - await compactor.started.wait() - self.assertTrue(agent.abort("stop compaction")) - result = await task - await agent.wait_for_idle() - - self.assertEqual(result.status, CompactionStatus.CANCELLED) - self.assertEqual(result.error, "stop compaction") - self.assertEqual(agent.messages, before) - self.assertTrue(compactor.stopped.is_set()) - self.assertEqual( - [type(event.payload) for event in sink.events], - [CompactionStarted, CompactionFailed], - ) - - async def test_caller_cancellation_emits_failure_and_preserves_history( - self, - ) -> None: - compactor = BlockingCompactor() - sink = RecordingSink() - agent = BaseAgent( - QueueModel(), - agent_id="compact-caller-cancel", - compaction_policy=policy(), - compactor=compactor, - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.reset(history=history()) - before = deepcopy(agent.messages) - - task = asyncio.create_task(agent.compact()) - await compactor.started.wait() - task.cancel() - with self.assertRaises(asyncio.CancelledError): - await task - await agent.wait_for_idle() - - self.assertEqual(agent.messages, before) - self.assertTrue(compactor.stopped.is_set()) - failed = sink.events[-1].payload - self.assertIsInstance(failed, CompactionFailed) - assert isinstance(failed, CompactionFailed) - self.assertEqual(failed.result.status, CompactionStatus.CANCELLED) - - async def test_single_turn_is_safely_skipped(self) -> None: - compactor = StaticCompactor() - sink = RecordingSink() - agent = BaseAgent( - QueueModel(), - agent_id="compact-skip", - compaction_policy=policy(), - compactor=compactor, - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.reset( - history=[ - {"role": "user", "content": "only", "_tokens": 100}, - {"role": "assistant", "content": "turn", "_tokens": 100}, - ] - ) - before = deepcopy(agent.messages) - - result = await agent.compact() - - self.assertEqual(result.status, CompactionStatus.SKIPPED) - self.assertEqual(agent.messages, before) - self.assertEqual(compactor.requests, []) - self.assertEqual( - [type(event.payload) for event in sink.events], - [CompactionStarted, CompactionCompleted], - ) - - async def test_repeated_compaction_replaces_and_merges_summary(self) -> None: - compactor = StaticCompactor(["first summary", "merged summary"]) - agent = BaseAgent( - QueueModel(), - agent_id="compact-repeat", - compaction_policy=policy(), - compactor=compactor, - context_token_estimator=MarkerTokenEstimator(), - ) - agent.reset(history=history()) - - first = await agent.compact() - agent.state.add_messages( - [ - {"role": "user", "content": "next old", "_tokens": 100}, - { - "role": "assistant", - "content": "next old answer", - "_tokens": 100, - }, - {"role": "user", "content": "newest", "_tokens": 30}, - { - "role": "assistant", - "content": "newest answer", - "_tokens": 30, - }, - ] - ) - second = await agent.compact() - - self.assertEqual(first.status, CompactionStatus.COMPLETED) - self.assertEqual(second.status, CompactionStatus.COMPLETED) - self.assertIsNotNone(compactor.requests[1].previous_summary) - self.assertEqual( - compactor.requests[1].previous_summary.content, - "first summary", - ) - summary_messages = [ - message for message in agent.messages if "_ejagent_summary" in message - ] - self.assertEqual(len(summary_messages), 1) - assert second.summary is not None - self.assertGreater(second.summary.summarized_message_count, 4) - self.assertEqual(second.summary.content, "merged summary") - - async def test_session_projection_preserves_late_context_barrier(self) -> None: - agent = BaseAgent( - QueueModel(), - agent_id="compact-barrier", - compaction_policy=policy(), - compactor=StaticCompactor(), - context_token_estimator=MarkerTokenEstimator(), - ) - agent.reset( - history=[ - {"role": "user", "content": "protected old", "_tokens": 100}, - { - "role": "assistant", - "content": "protected answer", - "_tokens": 100, - }, - {"role": "system", "content": "late barrier", "_tokens": 1}, - {"role": "user", "content": "summarize", "_tokens": 100}, - { - "role": "assistant", - "content": "summarize answer", - "_tokens": 100, - }, - {"role": "user", "content": "recent", "_tokens": 30}, - { - "role": "assistant", - "content": "recent answer", - "_tokens": 30, - }, - ] - ) - - result = await agent.compact() - - self.assertEqual(result.status, CompactionStatus.COMPLETED) - self.assertEqual( - [message.get("content") for message in result.messages[:3]], - ["protected old", "protected answer", "late barrier"], - ) - - async def test_session_restores_compacted_projection_and_new_runs( - self, - ) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder( - session_id="compacted-session", - storage=storage, - ) - model = QueueModel(["continued"]) - agent = BaseAgent( - model, - agent_id="compact-session-agent", - compaction_policy=policy(), - compactor=StaticCompactor(), - context_token_estimator=MarkerTokenEstimator(), - event_sink=recorder, - ) - agent.reset(history=history()) - - compacted = await agent.compact() - await agent.run(task="new task") - session = await recorder.load() - - self.assertEqual(compacted.status, CompactionStatus.COMPLETED) - assert session is not None - self.assertEqual(len(session.compactions), 1) - self.assertEqual(len(session.entries), 2) - self.assertEqual( - [message["role"] for message in session.messages], - ["system", "user", "assistant", "user", "assistant"], - ) - self.assertIn("_ejagent_summary", session.messages[0]) - - resumed_model = QueueModel(["resumed"]) - resumed = BaseAgent( - resumed_model, - agent_id="compact-session-agent", - ) - resumed.reset(session.messages) - await resumed.run(task="resume") - - llm_summary = resumed_model.contexts[0].llm_messages[1] - self.assertIn("Conversation summary", llm_summary["content"]) - self.assertNotIn("_ejagent_summary", llm_summary) - - async def test_session_keeps_audit_entries_covered_by_compaction( - self, - ) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder( - session_id="compaction-audit", - storage=storage, - ) - agent = BaseAgent( - QueueModel(["one", "two", "three", "four"]), - agent_id="compaction-audit-agent", - compaction_policy=policy(), - compactor=StaticCompactor(), - context_token_estimator=ConstantTokenEstimator(), - event_sink=recorder, - ) - - await agent.run(task="first") - await agent.run(task="second") - await agent.run(task="third") - result = await agent.compact() - await agent.run(task="fourth") - session = await recorder.load() - - self.assertEqual(result.status, CompactionStatus.COMPLETED) - assert session is not None - self.assertEqual(len(session.entries), 8) - self.assertEqual(len(session.runs), 4) - self.assertEqual(len(session.compactions), 1) - self.assertEqual(session.compactions[0].covered_entry_count, 6) - self.assertEqual( - [message["content"] for message in session.messages[1:]], - ["third", "three", "fourth", "four"], - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_context_management.py b/tests/test_context_management.py deleted file mode 100644 index eb84db2..0000000 --- a/tests/test_context_management.py +++ /dev/null @@ -1,336 +0,0 @@ -import unittest -from collections.abc import Mapping, Sequence -from typing import Any - -from ejagent import ( - AgentContextBuilder, - AgentEvent, - AgentFinished, - AgentStarted, - AgentState, - AssistantMessage, - BaseAgent, - CancellationToken, - CompactionDecisionReason, - CompactionPolicy, - ContextBudget, - ContextPressureEvaluated, - ContextUsageSource, - HeuristicMessageTokenEstimator, - MessageCompleted, - ModelAdapter, - TurnCompleted, - TurnStarted, - estimate_context_usage, - prepare_compaction, -) - - -class MarkerTokenEstimator: - def estimate_message(self, message: Mapping[str, Any]) -> int: - return int(message.get("_tokens", 0)) - - def estimate_tools(self, tools: Sequence[Mapping[str, Any]]) -> int: - return sum(int(tool.get("_tokens", 0)) for tool in tools) - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class TextModel(ModelAdapter): - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - return AssistantMessage(content="done") - - -class ContextManagementTests(unittest.IsolatedAsyncioTestCase): - def test_context_builder_strips_current_and_legacy_internal_metadata(self) -> None: - state = AgentState( - messages=[ - { - "role": "system", - "content": "summary", - "_ejagent_summary": {"content": "current"}, - "_simagentplg_summary": {"content": "legacy"}, - } - ] - ) - - context = AgentContextBuilder().build(state) - - self.assertEqual( - context.agent_messages[0]["_ejagent_summary"], {"content": "current"} - ) - self.assertEqual( - context.agent_messages[0]["_simagentplg_summary"], - {"content": "legacy"}, - ) - self.assertNotIn("_ejagent_summary", context.llm_messages[0]) - self.assertNotIn("_simagentplg_summary", context.llm_messages[0]) - - def test_heuristic_is_utf8_aware_and_ignores_usage_metadata(self) -> None: - estimator = HeuristicMessageTokenEstimator() - english = {"role": "user", "content": "a" * 40} - chinese = {"role": "user", "content": "你" * 40} - with_usage = { - **english, - "usage": { - "input_tokens": 9999, - "output_tokens": 9999, - "total_tokens": 19998, - }, - } - with_internal_summary = { - **english, - "_ejagent_summary": { - "content": "internal metadata must not be counted" * 100, - }, - } - with_legacy_internal_summary = { - **english, - "_simagentplg_summary": { - "content": "legacy internal metadata must not be counted" * 100, - }, - } - - self.assertGreater( - estimator.estimate_message(chinese), - estimator.estimate_message(english), - ) - self.assertEqual( - estimator.estimate_message(english), - estimator.estimate_message(with_usage), - ) - self.assertEqual( - estimator.estimate_message(english), - estimator.estimate_message(with_internal_summary), - ) - self.assertEqual( - estimator.estimate_message(english), - estimator.estimate_message(with_legacy_internal_summary), - ) - - def test_estimate_without_usage_includes_messages_and_tools(self) -> None: - estimate = estimate_context_usage( - [ - {"role": "system", "_tokens": 10}, - {"role": "user", "_tokens": 20}, - ], - tools=[{"_tokens": 30}], - estimator=MarkerTokenEstimator(), - ) - - self.assertEqual(estimate.reported_tokens, 0) - self.assertEqual(estimate.trailing_tokens, 60) - self.assertEqual(estimate.heuristic_tokens, 60) - self.assertEqual(estimate.total_tokens, 60) - self.assertIsNone(estimate.last_usage_index) - self.assertEqual(estimate.source, ContextUsageSource.ESTIMATED) - - def test_latest_usage_is_combined_with_only_trailing_messages(self) -> None: - estimate = estimate_context_usage( - [ - {"role": "system", "_tokens": 10}, - { - "role": "assistant", - "_tokens": 10, - "usage": {"total_tokens": 500}, - }, - {"role": "tool", "_tokens": 30}, - {"role": "user", "_tokens": 20}, - ], - tools=[{"_tokens": 40}], - estimator=MarkerTokenEstimator(), - ) - - self.assertEqual(estimate.reported_tokens, 500) - self.assertEqual(estimate.trailing_tokens, 50) - self.assertEqual(estimate.heuristic_tokens, 110) - self.assertEqual(estimate.total_tokens, 550) - self.assertEqual(estimate.last_usage_index, 1) - self.assertEqual(estimate.source, ContextUsageSource.MIXED) - - def test_full_heuristic_is_a_lower_bound_for_changed_context(self) -> None: - estimate = estimate_context_usage( - [ - { - "role": "assistant", - "_tokens": 200, - "usage": {"total_tokens": 50}, - } - ], - estimator=MarkerTokenEstimator(), - ) - - self.assertEqual(estimate.usage_based_tokens, 50) - self.assertEqual(estimate.heuristic_tokens, 200) - self.assertEqual(estimate.total_tokens, 200) - self.assertEqual(estimate.source, ContextUsageSource.MIXED) - - def test_policy_is_independent_and_reports_its_reason(self) -> None: - budget = ContextBudget( - context_window=100, - reserve_tokens=20, - keep_recent_tokens=30, - ) - policy = CompactionPolicy(budget) - estimate = estimate_context_usage( - [{"role": "user", "_tokens": 80}], - estimator=MarkerTokenEstimator(), - ) - - decision = policy.evaluate(estimate) - - self.assertTrue(decision.should_compact) - self.assertEqual(decision.threshold_tokens, 80) - self.assertEqual(decision.pressure_ratio, 1.0) - self.assertEqual( - decision.reason, - CompactionDecisionReason.THRESHOLD_REACHED, - ) - - disabled = CompactionPolicy(budget, enabled=False).evaluate(estimate) - self.assertFalse(disabled.should_compact) - self.assertEqual(disabled.reason, CompactionDecisionReason.DISABLED) - - def test_budget_rejects_impossible_reservations(self) -> None: - with self.assertRaisesRegex(ValueError, "less than context_window"): - ContextBudget(100, 100, 10) - with self.assertRaisesRegex(ValueError, "context threshold"): - ContextBudget(100, 20, 81) - with self.assertRaisesRegex(ValueError, "greater than zero"): - ContextBudget(100, 20, 0) - - def test_preparation_summarizes_only_complete_old_turns(self) -> None: - messages: list[dict[str, Any]] = [ - {"role": "system", "content": "core", "_tokens": 1}, - {"role": "user", "content": "old", "_tokens": 10}, - { - "role": "assistant", - "content": None, - "tool_calls": [{"id": "call-1"}], - "_tokens": 5, - }, - { - "role": "tool", - "tool_call_id": "call-1", - "content": "large output", - "_tokens": 100, - }, - {"role": "assistant", "content": "observed", "_tokens": 5}, - {"role": "user", "content": "middle", "_tokens": 20}, - {"role": "assistant", "content": "middle answer", "_tokens": 10}, - {"role": "user", "content": "recent", "_tokens": 30}, - {"role": "assistant", "content": "recent answer", "_tokens": 20}, - ] - - preparation = prepare_compaction( - messages, - keep_recent_tokens=60, - estimator=MarkerTokenEstimator(), - ) - - self.assertTrue(preparation.can_compact) - self.assertEqual(preparation.history_start_index, 1) - self.assertEqual(preparation.first_kept_index, 5) - self.assertEqual( - [message["role"] for message in preparation.protected_messages], - ["system"], - ) - self.assertEqual( - [message["role"] for message in preparation.messages_to_summarize], - ["user", "assistant", "tool", "assistant"], - ) - self.assertEqual( - [message["role"] for message in preparation.messages_to_keep], - ["user", "assistant", "user", "assistant"], - ) - self.assertEqual(preparation.estimated_history_tokens, 200) - self.assertEqual(preparation.estimated_summarized_tokens, 120) - self.assertEqual(preparation.estimated_kept_tokens, 80) - - messages[2]["tool_calls"][0]["id"] = "mutated" - self.assertEqual( - preparation.messages_to_summarize[1]["tool_calls"][0]["id"], - "call-1", - ) - - def test_preparation_preserves_late_non_conversation_barrier(self) -> None: - messages = [ - {"role": "system", "_tokens": 1}, - {"role": "user", "_tokens": 100}, - {"role": "assistant", "_tokens": 100}, - {"role": "system", "content": "late policy", "_tokens": 1}, - {"role": "user", "_tokens": 100}, - {"role": "assistant", "_tokens": 100}, - ] - - preparation = prepare_compaction( - messages, - keep_recent_tokens=10, - estimator=MarkerTokenEstimator(), - ) - - self.assertFalse(preparation.can_compact) - self.assertEqual(preparation.history_start_index, 4) - self.assertEqual(len(preparation.protected_messages), 4) - self.assertEqual(len(preparation.messages_to_keep), 2) - - async def test_agent_emits_pressure_without_mutating_or_stopping(self) -> None: - sink = RecordingSink() - agent = BaseAgent( - TextModel(), - agent_id="context-pressure", - compaction_policy=CompactionPolicy( - ContextBudget( - context_window=100, - reserve_tokens=20, - keep_recent_tokens=20, - ) - ), - context_token_estimator=MarkerTokenEstimator(), - event_sink=sink, - ) - agent.messages[0]["_tokens"] = 90 - - result = await agent.run(task="pressure") - - payloads = [event.payload for event in sink.events] - pressure = next( - payload - for payload in payloads - if isinstance(payload, ContextPressureEvaluated) - ) - self.assertTrue(pressure.decision.should_compact) - self.assertIsNotNone(pressure.preparation) - assert pressure.preparation is not None - self.assertFalse(pressure.preparation.can_compact) - self.assertEqual( - [type(payload) for payload in payloads], - [ - AgentStarted, - TurnStarted, - ContextPressureEvaluated, - MessageCompleted, - TurnCompleted, - AgentFinished, - ], - ) - self.assertEqual(result.output, "done") - self.assertEqual( - [message["role"] for message in agent.messages], - ["system", "user", "assistant"], - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_context_pipeline.py b/tests/test_context_pipeline.py new file mode 100644 index 0000000..f0b893d --- /dev/null +++ b/tests/test_context_pipeline.py @@ -0,0 +1,363 @@ +from __future__ import annotations + +import unittest +from collections.abc import AsyncIterator, Sequence + +from ejagent.context import DerivedCompactionPipeline, IdentityContextPipeline +from ejagent.contracts import ( + AssistantMessage, + CancellationSource, + CancellationToken, + ContextBuildError, + ContextCompactionOutput, + ContextCompactionRequest, + ContextCompactor, + ContextCompactorError, + ContextPipeline, + ContextProtocolError, + ContextRequest, + ContextSummary, + ContextView, + FailureCode, + ModelPort, + ModelRequest, + ModelResponseCompleted, + ModelStreamEvent, + RunIntent, + RunPhase, + RunSpec, + RunStatus, + StopReason, + SystemMessage, + ToolCall, + ToolDefinition, + ToolExecutionResult, + ToolExecutor, + TransientInstruction, + UserMessage, +) +from ejagent.harness import AgentHarness, MemorySessionStore +from ejagent.kernel import RuntimeKernel + + +def context_request( + *, + committed: tuple[SystemMessage | UserMessage | AssistantMessage, ...] = (), + pending: tuple[SystemMessage | UserMessage | AssistantMessage, ...] = (), + revision: int = 3, + turn: int = 1, +) -> ContextRequest: + return ContextRequest( + run_id="context-run", + source_revision=revision, + turn=turn, + committed_messages=committed, + pending_messages=pending, + metadata={"tenant": "test"}, + ) + + +class StaticCompactor(ContextCompactor): + def __init__(self, content: str = "earlier work") -> None: + self.content = content + self.requests: list[ContextCompactionRequest] = [] + + async def compact( + self, + request: ContextCompactionRequest, + *, + cancellation: CancellationToken, + ) -> ContextCompactionOutput: + self.requests.append(request) + return ContextCompactionOutput(self.content, "static-compactor") + + +class FailingCompactor(ContextCompactor): + async def compact( + self, + request: ContextCompactionRequest, + *, + cancellation: CancellationToken, + ) -> ContextCompactionOutput: + raise ContextCompactorError("summary backend unavailable", retryable=True) + + +class BrokenCompactor(ContextCompactor): + async def compact( + self, + request: ContextCompactionRequest, + *, + cancellation: CancellationToken, + ) -> ContextCompactionOutput: + raise ValueError("adapter leaked implementation failure") + + +class RecordingModel(ModelPort): + def __init__(self, responses: Sequence[AssistantMessage]) -> None: + self.responses = list(responses) + self.requests: list[ModelRequest] = [] + + async def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + self.requests.append(request) + yield ModelResponseCompleted(self.responses.pop(0)) + + +class NoTools(ToolExecutor): + @property + def definitions(self) -> Sequence[ToolDefinition]: + return () + + async def execute( + self, + call: ToolCall, + *, + cancellation: CancellationToken, + ) -> ToolExecutionResult: + raise AssertionError("no tools are registered") + + +class SteeringProjection(ContextPipeline): + def __init__(self) -> None: + self.requests: list[ContextRequest] = [] + + async def build( + self, + request: ContextRequest, + *, + cancellation: CancellationToken, + ) -> ContextView: + self.requests.append(request) + return ContextView( + run_id=request.run_id, + source_revision=request.source_revision, + turn=request.turn, + messages=( + *request.messages, + TransientInstruction("change direction", "test-steering"), + ), + ) + + +class FailingPipeline(ContextPipeline): + async def build( + self, + request: ContextRequest, + *, + cancellation: CancellationToken, + ) -> ContextView: + raise ContextBuildError( + FailureCode.COMPACTION_FAILED, + "cannot summarize", + retryable=True, + ) + + +class ContextPipelineTests(unittest.IsolatedAsyncioTestCase): + async def test_identity_projection_preserves_messages_without_aliasing( + self, + ) -> None: + committed = (SystemMessage("stable"), UserMessage("old")) + pending = (UserMessage("new"),) + request = context_request(committed=committed, pending=pending) + + view = await IdentityContextPipeline().build( + request, + cancellation=CancellationSource().token, + ) + + self.assertEqual(view.messages, (*committed, *pending)) + self.assertEqual(view.source_revision, 3) + self.assertEqual(view.metadata["projection"], "identity") + self.assertEqual(request.committed_messages, committed) + self.assertEqual(request.pending_messages, pending) + + async def test_derived_compaction_only_changes_the_context_view(self) -> None: + committed = ( + SystemMessage("stable instruction"), + UserMessage("old task"), + AssistantMessage(content="old answer"), + ) + pending = (UserMessage("current task"),) + request = context_request(committed=committed, pending=pending) + compactor = StaticCompactor() + pipeline = DerivedCompactionPipeline(compactor, minimum_messages=2) + + view = await pipeline.build( + request, + cancellation=CancellationSource().token, + ) + + self.assertEqual(view.messages[0], committed[0]) + self.assertIsInstance(view.messages[1], ContextSummary) + summary = view.messages[1] + assert isinstance(summary, ContextSummary) + self.assertEqual(summary.source_revision_start, 1) + self.assertEqual(summary.source_revision_end, 3) + self.assertEqual(summary.content, "earlier work") + self.assertEqual(view.messages[2], pending[0]) + self.assertEqual(compactor.requests[0].messages, committed[1:]) + self.assertEqual(request.committed_messages, committed) + self.assertNotIn(summary, request.messages) + + async def test_compaction_projection_is_rebuildable_from_same_history(self) -> None: + request = context_request( + committed=( + UserMessage("one"), + AssistantMessage(content="two"), + ) + ) + pipeline = DerivedCompactionPipeline( + StaticCompactor("deterministic summary"), + minimum_messages=2, + ) + + first = await pipeline.build( + request, + cancellation=CancellationSource().token, + ) + second = await pipeline.build( + request, + cancellation=CancellationSource().token, + ) + + self.assertEqual(first, second) + self.assertEqual(request.messages[0], UserMessage("one")) + + async def test_below_threshold_uses_unmodified_projection(self) -> None: + compactor = StaticCompactor() + request = context_request(committed=(UserMessage("only one"),)) + + view = await DerivedCompactionPipeline( + compactor, + minimum_messages=2, + ).build(request, cancellation=CancellationSource().token) + + self.assertEqual(view.messages, request.messages) + self.assertEqual(view.metadata["projection"], "identity") + self.assertEqual(compactor.requests, []) + + async def test_declared_compactor_failure_becomes_context_failure(self) -> None: + pipeline = DerivedCompactionPipeline( + FailingCompactor(), + minimum_messages=1, + ) + + with self.assertRaises(ContextBuildError) as raised: + await pipeline.build( + context_request(committed=(UserMessage("old"),)), + cancellation=CancellationSource().token, + ) + + self.assertEqual(raised.exception.code, FailureCode.COMPACTION_FAILED) + self.assertTrue(raised.exception.retryable) + + async def test_undeclared_compactor_failure_is_a_protocol_error(self) -> None: + pipeline = DerivedCompactionPipeline( + BrokenCompactor(), + minimum_messages=1, + ) + + with self.assertRaises(ContextProtocolError): + await pipeline.build( + context_request(committed=(UserMessage("old"),)), + cancellation=CancellationSource().token, + ) + + async def test_kernel_projects_transient_context_without_committing_it( + self, + ) -> None: + model = RecordingModel([AssistantMessage(content="done")]) + context = SteeringProjection() + spec = RunSpec( + run_id="transient-run", + base_revision=2, + intent=RunIntent.TASK, + task="original task", + messages=(SystemMessage("stable"),), + ) + + outcome = await RuntimeKernel( + model=model, + tools=NoTools(), + context=context, + ).run(spec) + + self.assertIsInstance(model.requests[0].messages[-1], TransientInstruction) + self.assertNotIn( + TransientInstruction("change direction", "test-steering"), + outcome.delta.messages, + ) + self.assertEqual( + context.requests[0].pending_messages, (UserMessage("original task"),) + ) + self.assertIn( + "context_built", [record.kind for record in outcome.audit_records] + ) + + async def test_context_failure_is_a_structured_kernel_outcome(self) -> None: + spec = RunSpec( + run_id="context-failure", + base_revision=0, + intent=RunIntent.TASK, + task="task", + messages=(), + ) + kernel = RuntimeKernel( + model=RecordingModel([]), + tools=NoTools(), + context=FailingPipeline(), + ) + + outcome = await kernel.run(spec) + + self.assertEqual(outcome.result.status, RunStatus.FAILED) + self.assertEqual(outcome.result.stop_reason, StopReason.COMPACTION_FAILED) + assert outcome.failure is not None + self.assertEqual(outcome.failure.phase, RunPhase.CONTEXT) + self.assertEqual(outcome.failure.code, FailureCode.COMPACTION_FAILED) + self.assertEqual(outcome.delta.messages, (UserMessage("task"),)) + + async def test_harness_compaction_never_rewrites_conversation(self) -> None: + store = MemorySessionStore() + model = RecordingModel( + [ + AssistantMessage(content="first answer"), + AssistantMessage(content="second answer"), + ] + ) + run_ids = iter(("first-run", "second-run")) + harness = AgentHarness( + agent_id="derived-context-agent", + model=model, + tools=NoTools(), + context=DerivedCompactionPipeline( + StaticCompactor("first exchange summarized"), + minimum_messages=2, + ), + initial_messages=(SystemMessage("stable"),), + store=store, + run_id_factory=lambda: next(run_ids), + ) + + await harness.run("first task") + before_second = harness.messages + await harness.run("second task") + + self.assertIsInstance(model.requests[1].messages[1], ContextSummary) + self.assertEqual(model.requests[1].messages[-1], UserMessage("second task")) + self.assertEqual(harness.messages[: len(before_second)], before_second) + self.assertFalse( + any(isinstance(message, ContextSummary) for message in harness.messages) + ) + persisted = await store.load("derived-context-agent") + assert persisted is not None + self.assertEqual(persisted.messages, harness.messages) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_contracts.py b/tests/test_contracts.py new file mode 100644 index 0000000..fcf7e7a --- /dev/null +++ b/tests/test_contracts.py @@ -0,0 +1,231 @@ +from __future__ import annotations + +import unittest +from datetime import UTC, datetime + +from ejagent.contracts import ( + AssistantMessage, + AuditRecord, + ContextSummary, + ConversationSnapshot, + FailureCode, + RunDelta, + RunFailure, + RunIntent, + RunLimits, + RunOutcome, + RunPhase, + RunResult, + RunSpec, + RunStatus, + SessionCommit, + StopReason, + ToolCall, + ToolResultMessage, + TransientInstruction, + UserMessage, + is_context_message, + is_conversation_message, + thaw_json_value, +) + + +class MessageContractTests(unittest.TestCase): + def test_tool_call_recursively_freezes_detached_arguments(self) -> None: + source = {"path": "README.md", "options": {"lines": [1, 2]}} + + call = ToolCall(id="call-1", name="read", arguments=source) + source["path"] = "changed" + options = source["options"] + assert isinstance(options, dict) + lines = options["lines"] + assert isinstance(lines, list) + lines.append(3) + + self.assertEqual(call.arguments["path"], "README.md") + self.assertEqual( + thaw_json_value(call.arguments), + {"path": "README.md", "options": {"lines": [1, 2]}}, + ) + with self.assertRaises(TypeError): + call.arguments["path"] = "forbidden" # type: ignore[index] + + def test_assistant_requires_text_or_tool_calls(self) -> None: + with self.assertRaisesRegex(ValueError, "text or tool calls"): + AssistantMessage() + + message = AssistantMessage(tool_calls=(ToolCall(id="call-1", name="lookup"),)) + + self.assertEqual(message.tool_calls[0].name, "lookup") + + def test_tool_result_is_typed_and_immutable(self) -> None: + source = {"ok": True, "items": ["a"]} + + message = ToolResultMessage( + tool_call_id="call-1", + tool_name="lookup", + result=source, + ) + source["ok"] = False + + self.assertEqual( + thaw_json_value(message.result), + {"ok": True, "items": ["a"]}, + ) + + def test_context_summary_is_not_a_conversation_message(self) -> None: + summary = ContextSummary( + source_revision_start=1, + source_revision_end=4, + content="Earlier work was summarized.", + compactor_id="test", + ) + + self.assertEqual(summary.source_revision_end, 4) + self.assertTrue(is_context_message(summary)) + self.assertFalse(is_conversation_message(summary)) + + def test_transient_instruction_cannot_enter_conversation(self) -> None: + instruction = TransientInstruction("focus on safety", "steering") + + self.assertTrue(is_context_message(instruction)) + self.assertFalse(is_conversation_message(instruction)) + + +class DataDomainContractTests(unittest.TestCase): + def test_commit_produces_separate_conversation_and_audit_values(self) -> None: + base = ConversationSnapshot( + revision=2, + messages=(UserMessage("existing"),), + ) + result = RunResult( + run_id="run-3", + status=RunStatus.COMPLETED, + stop_reason=StopReason.TEXT_RESPONSE, + turns=1, + output="done", + ) + outcome = RunOutcome( + result=result, + delta=RunDelta( + base_revision=2, + messages=(AssistantMessage(content="done"),), + ), + ) + + commit = SessionCommit(agent_id="agent", base=base, outcome=outcome) + + self.assertEqual(commit.resulting_conversation.revision, 3) + self.assertEqual( + commit.resulting_conversation.messages, + (UserMessage("existing"), AssistantMessage(content="done")), + ) + self.assertEqual(base.messages, (UserMessage("existing"),)) + self.assertEqual(commit.audit.run_id, "run-3") + self.assertTrue(commit.audit.committed) + self.assertFalse(hasattr(commit.audit, "messages")) + + +class RunContractTests(unittest.TestCase): + def test_run_spec_freezes_inputs_at_one_revision(self) -> None: + messages = [UserMessage("existing input")] + metadata = {"trace": {"labels": ["contract"]}} + + spec = RunSpec( + run_id="run-1", + base_revision=7, + intent=RunIntent.TASK, + task="new task", + messages=messages, + configuration_revision="config-3", + metadata=metadata, + ) + messages.append(UserMessage("late mutation")) + trace = metadata["trace"] + assert isinstance(trace, dict) + labels = trace["labels"] + assert isinstance(labels, list) + labels.append("late") + + self.assertEqual(spec.messages, (UserMessage("existing input"),)) + self.assertEqual( + thaw_json_value(spec.metadata), + {"trace": {"labels": ["contract"]}}, + ) + + def test_continue_spec_rejects_a_task(self) -> None: + with self.assertRaisesRegex(ValueError, "must not contain a task"): + RunSpec( + run_id="run-1", + base_revision=0, + intent=RunIntent.CONTINUE, + task="unexpected", + messages=(), + ) + + def test_limits_reject_boolean_in_integer_fields(self) -> None: + with self.assertRaisesRegex(TypeError, "max_turns must be an integer"): + RunLimits(max_turns=True) # type: ignore[arg-type] + + def test_delta_proposes_the_next_revision(self) -> None: + delta = RunDelta( + base_revision=4, + messages=(UserMessage("task"),), + ) + + self.assertEqual(delta.next_revision, 5) + + def test_failed_outcome_requires_structured_failure(self) -> None: + result = RunResult( + run_id="run-1", + status=RunStatus.FAILED, + stop_reason=StopReason.RUNTIME_ERROR, + turns=1, + ) + + with self.assertRaisesRegex(ValueError, "requires a RunFailure"): + RunOutcome(result=result, delta=RunDelta(base_revision=0)) + + outcome = RunOutcome( + result=result, + delta=RunDelta(base_revision=0), + failure=RunFailure( + phase=RunPhase.RUNTIME, + code=FailureCode.RUNTIME_ERROR, + message="expected failure", + ), + ) + + self.assertEqual(outcome.failure.code, FailureCode.RUNTIME_ERROR) + + def test_audit_records_are_ordered_and_belong_to_the_run(self) -> None: + result = RunResult( + run_id="run-1", + status=RunStatus.COMPLETED, + stop_reason=StopReason.TEXT_RESPONSE, + turns=1, + output="done", + ) + later = AuditRecord( + run_id="run-1", + sequence=2, + kind="run_finished", + occurred_at=datetime.now(UTC), + ) + earlier = AuditRecord( + run_id="run-1", + sequence=1, + kind="run_started", + occurred_at=datetime.now(UTC), + ) + + with self.assertRaisesRegex(ValueError, "unique and ordered"): + RunOutcome( + result=result, + delta=RunDelta(base_revision=0), + audit_records=(later, earlier), + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_core_adapters.py b/tests/test_core_adapters.py new file mode 100644 index 0000000..2dad00f --- /dev/null +++ b/tests/test_core_adapters.py @@ -0,0 +1,515 @@ +from __future__ import annotations + +import tempfile +import unittest +from collections.abc import Sequence +from pathlib import Path +from types import SimpleNamespace +from typing import Any + +from ejagent.context import IdentityContextPipeline, SkillsContextPipeline +from ejagent.contracts import ( + AssistantMessage, + CancellationSource, + ContextRequest, + ContextSummary, + FailureCode, + ModelCallError, + ModelProtocolError, + ModelRequest, + ModelResponseCompleted, + ModelTextDelta, + ModelThinkingDelta, + RunStatus, + SystemMessage, + ToolCall, + ToolControl, + ToolDefinition, + ToolExecutionError, + ToolExecutionResult, + ToolProtocolError, + ToolResultMessage, + ToolSemantics, + TransientInstruction, + UserMessage, +) +from ejagent.harness import AgentHarness +from ejagent.providers import ModelConfig, OpenAIModelPort +from ejagent.tools import ( + CompositeToolExecutor, + FunctionTool, + FunctionToolExecutor, + McpToolExecutor, +) + + +def _delta( + *, + content: str | None = None, + finish_reason: str | None = None, + tool_calls: Sequence[Any] = (), + reasoning: str | None = None, +) -> Any: + return SimpleNamespace( + choices=[ + SimpleNamespace( + finish_reason=finish_reason, + delta=SimpleNamespace( + content=content, + reasoning_content=reasoning, + tool_calls=tool_calls, + ), + ) + ], + usage=None, + ) + + +def _tool_fragment( + *, + index: int = 0, + call_id: str | None = None, + name: str | None = None, + arguments: str | None = None, +) -> Any: + return SimpleNamespace( + index=index, + id=call_id, + function=SimpleNamespace(name=name, arguments=arguments), + ) + + +class FakeOpenAIStream: + def __init__(self, chunks: Sequence[Any]) -> None: + self._chunks = iter(chunks) + self.closed = False + + def __aiter__(self) -> FakeOpenAIStream: + return self + + async def __anext__(self) -> Any: + try: + return next(self._chunks) + except StopIteration as exc: + raise StopAsyncIteration from exc + + async def aclose(self) -> None: + self.closed = True + + +class FakeCompletions: + def __init__( + self, + responses: Sequence[FakeOpenAIStream | Exception], + ) -> None: + self.responses = list(responses) + self.requests: list[dict[str, Any]] = [] + + async def create(self, **kwargs: Any) -> FakeOpenAIStream: + self.requests.append(kwargs) + response = self.responses.pop(0) + if isinstance(response, Exception): + raise response + return response + + +class FakeOpenAIClient: + def __init__(self, completions: FakeCompletions) -> None: + self.chat = SimpleNamespace(completions=completions) + self.closed = False + + async def close(self) -> None: + self.closed = True + + +def _config() -> ModelConfig: + return ModelConfig( + model="test-model", + api_key="test-key", + base_url="https://model.invalid/v1", + ) + + +class OpenAIModelPortTests(unittest.IsolatedAsyncioTestCase): + async def test_serializes_typed_request_and_normalizes_stream(self) -> None: + usage = SimpleNamespace( + prompt_tokens=7, + completion_tokens=3, + prompt_tokens_details=SimpleNamespace( + cached_tokens=2, + cache_write_tokens=None, + ), + completion_tokens_details=SimpleNamespace(reasoning_tokens=1), + ) + stream = FakeOpenAIStream( + [ + _delta(content="答", reasoning="想"), + _delta(content="案", finish_reason="stop"), + SimpleNamespace(choices=[], usage=usage), + ] + ) + completions = FakeCompletions([stream]) + port = OpenAIModelPort( + _config(), + client=FakeOpenAIClient(completions), # type: ignore[arg-type] + ) + await port.start() + request = ModelRequest( + messages=( + SystemMessage("system"), + UserMessage("你好"), + AssistantMessage( + tool_calls=(ToolCall("call-0", "lookup", {"q": "值"}),) + ), + ToolResultMessage("call-0", "lookup", {"found": True}), + ContextSummary(1, 2, "summary", "compact-v1"), + TransientInstruction("focus", "steering"), + ), + tools=( + ToolDefinition( + name="lookup", + description="Lookup a value.", + input_schema={"type": "object"}, + ), + ), + ) + + events = [ + event + async for event in port.stream( + request, + cancellation=CancellationSource().token, + ) + ] + + self.assertEqual(events[0], ModelTextDelta("答")) + self.assertEqual(events[1], ModelThinkingDelta("想")) + self.assertEqual(events[2], ModelTextDelta("案")) + completed = events[3] + self.assertIsInstance(completed, ModelResponseCompleted) + assert isinstance(completed, ModelResponseCompleted) + self.assertEqual(completed.message, AssistantMessage("答案")) + self.assertEqual(completed.usage.input_tokens, 7) # type: ignore[union-attr] + sent = completions.requests[0] + self.assertEqual(sent["messages"][1], {"role": "user", "content": "你好"}) + self.assertEqual( + sent["messages"][2]["tool_calls"][0]["function"]["arguments"], + '{"q":"值"}', + ) + self.assertEqual(sent["messages"][3]["content"], '{"found":true}') + self.assertIn("Derived summary", sent["messages"][4]["content"]) + self.assertIn("Transient instruction", sent["messages"][5]["content"]) + self.assertEqual( + sent["tools"][0]["function"]["parameters"], + {"type": "object"}, + ) + self.assertTrue(stream.closed) + + async def test_reassembles_and_decodes_streamed_tool_call(self) -> None: + stream = FakeOpenAIStream( + [ + _delta( + tool_calls=( + _tool_fragment( + call_id="call-1", + name="weather", + arguments='{"city":', + ), + ) + ), + _delta( + finish_reason="tool_calls", + tool_calls=(_tool_fragment(arguments='"杭州"}'),), + ), + ] + ) + port = OpenAIModelPort( + _config(), + client=FakeOpenAIClient(FakeCompletions([stream])), # type: ignore[arg-type] + ) + await port.start() + + events = [ + event + async for event in port.stream( + ModelRequest(messages=(UserMessage("weather"),)), + cancellation=CancellationSource().token, + ) + ] + + self.assertEqual( + events, + [ + ModelResponseCompleted( + AssistantMessage( + tool_calls=(ToolCall("call-1", "weather", {"city": "杭州"}),) + ) + ) + ], + ) + + async def test_normalizes_expected_provider_failure(self) -> None: + port = OpenAIModelPort( + _config(), + client=FakeOpenAIClient( # type: ignore[arg-type] + FakeCompletions([TimeoutError("late")]) + ), + ) + await port.start() + + with self.assertRaises(ModelCallError) as raised: + async for _ in port.stream( + ModelRequest(messages=(UserMessage("hello"),)), + cancellation=CancellationSource().token, + ): + pass + + self.assertEqual(raised.exception.code, FailureCode.TIMEOUT) + self.assertTrue(raised.exception.retryable) + + async def test_rejects_invalid_tool_argument_json_as_protocol_error(self) -> None: + stream = FakeOpenAIStream( + [ + _delta( + finish_reason="tool_calls", + tool_calls=( + _tool_fragment( + call_id="call-1", + name="lookup", + arguments="not-json", + ), + ), + ) + ] + ) + port = OpenAIModelPort( + _config(), + client=FakeOpenAIClient(FakeCompletions([stream])), # type: ignore[arg-type] + ) + await port.start() + + with self.assertRaises(ModelProtocolError): + async for _ in port.stream( + ModelRequest(messages=(UserMessage("lookup"),)), + cancellation=CancellationSource().token, + ): + pass + + +class FakeMcpManager: + def __init__(self) -> None: + self.events: list[str] = [] + self.calls: list[tuple[str, dict[str, object]]] = [] + + async def startup(self) -> None: + self.events.append("start") + + async def shutdown(self) -> None: + self.events.append("shutdown") + + def get_openai_tools(self) -> list[dict[str, Any]]: + return [ + { + "type": "function", + "function": { + "name": "docs__search", + "description": "Search docs.", + "parameters": { + "type": "object", + "properties": {"q": {"type": "string"}}, + }, + }, + } + ] + + async def call_tool(self, tool_name: str, args: dict[str, object]) -> str: + self.calls.append((tool_name, args)) + return "found" + + +class ToolExecutorAdapterTests(unittest.IsolatedAsyncioTestCase): + async def test_function_executor_dispatches_normalized_result(self) -> None: + async def add( + call: ToolCall, + cancellation: Any, + ) -> ToolExecutionResult: + cancellation.raise_if_cancelled() + return ToolExecutionResult( + {"value": call.arguments["left"] + call.arguments["right"]}, + control=ToolControl.COMPLETE, + output="42", + ) + + executor = FunctionToolExecutor( + ( + FunctionTool( + ToolDefinition( + "add", + input_schema={"type": "object"}, + semantics=ToolSemantics.read_only(), + ), + add, + ), + ) + ) + + result = await executor.execute( + ToolCall("call-1", "add", {"left": 19, "right": 23}), + cancellation=CancellationSource().token, + ) + + self.assertEqual(result.result["value"], 42) # type: ignore[index] + self.assertEqual(result.control, ToolControl.COMPLETE) + with self.assertRaises(ToolExecutionError): + await executor.execute( + ToolCall("call-2", "missing"), + cancellation=CancellationSource().token, + ) + + async def test_function_executor_rejects_invalid_return_type(self) -> None: + async def invalid(call: ToolCall, cancellation: Any) -> Any: + return {"not": "normalized"} + + executor = FunctionToolExecutor( + (FunctionTool(ToolDefinition("invalid"), invalid),) + ) + with self.assertRaises(ToolProtocolError): + await executor.execute( + ToolCall("call-1", "invalid"), + cancellation=CancellationSource().token, + ) + + async def test_mcp_executor_owns_lifecycle_and_normalizes_metadata(self) -> None: + manager = FakeMcpManager() + executor = McpToolExecutor( + manager=manager, + semantics={"docs__search": ToolSemantics.read_only()}, + ) + + await executor.start() + result = await executor.execute( + ToolCall("call-1", "docs__search", {"q": "kernel"}), + cancellation=CancellationSource().token, + ) + await executor.shutdown() + + self.assertEqual(result, ToolExecutionResult("found")) + self.assertEqual(manager.calls, [("docs__search", {"q": "kernel"})]) + self.assertEqual(manager.events, ["start", "shutdown"]) + self.assertEqual(executor.definitions, ()) + + async def test_composite_refreshes_dynamic_routes(self) -> None: + manager = FakeMcpManager() + mcp = McpToolExecutor(manager=manager) + local = FunctionToolExecutor() + composite = CompositeToolExecutor((local, mcp)) + + self.assertEqual(composite.definitions, ()) + await composite.start() + self.assertEqual( + tuple(item.name for item in composite.definitions), + ("docs__search",), + ) + await composite.shutdown() + self.assertEqual(composite.definitions, ()) + + +class SkillsContextPipelineTests(unittest.IsolatedAsyncioTestCase): + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.root = Path(temporary.name) + skill = self.root / "release_notes" + skill.mkdir() + (skill / "SKILL.md").write_text( + "---\nname: release_notes\ndescription: Write releases.\n---\n" + "# Instructions\nUse concise bullets.\n", + encoding="utf-8", + ) + + async def test_injects_index_and_explicit_skill_only_into_context(self) -> None: + pipeline = SkillsContextPipeline( + self.root, + base=IdentityContextPipeline(), + ) + await pipeline.start() + request = ContextRequest( + run_id="run-1", + source_revision=0, + turn=1, + committed_messages=(SystemMessage("system"),), + pending_messages=(UserMessage("$release_notes draft this"),), + ) + + view = await pipeline.build( + request, + cancellation=CancellationSource().token, + ) + await pipeline.shutdown() + + self.assertEqual(request.transient_instructions, ()) + additions = [ + item for item in view.messages if isinstance(item, TransientInstruction) + ] + self.assertEqual( + tuple(item.source for item in additions), + ("skills:index", "skills:release_notes"), + ) + self.assertIn("Use concise bullets", additions[1].content) + self.assertEqual(view.metadata["selected_skill"], "release_notes") + + +class CoreAdapterIntegrationTests(unittest.IsolatedAsyncioTestCase): + async def test_harness_runs_openai_tool_round_trip(self) -> None: + first = FakeOpenAIStream( + [ + _delta( + finish_reason="tool_calls", + tool_calls=( + _tool_fragment( + call_id="call-add", + name="add", + arguments='{"left":20,"right":22}', + ), + ), + ) + ] + ) + second = FakeOpenAIStream([_delta(content="42", finish_reason="stop")]) + completions = FakeCompletions([first, second]) + model = OpenAIModelPort( + _config(), + client=FakeOpenAIClient(completions), # type: ignore[arg-type] + ) + + async def add(call: ToolCall, cancellation: Any) -> ToolExecutionResult: + return ToolExecutionResult( + { + "value": call.arguments["left"] + call.arguments["right"], + } + ) + + tools = FunctionToolExecutor((FunctionTool(ToolDefinition("add"), add),)) + harness = AgentHarness( + agent_id="calculator", + model=model, + tools=tools, + initial_messages=(SystemMessage("Use tools."),), + run_id_factory=lambda: "run-1", + ) + + async with harness: + outcome = await harness.run("add 20 and 22") + + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertEqual(outcome.result.output, "42") + self.assertEqual(harness.revision, 1) + self.assertEqual(len(completions.requests), 2) + self.assertEqual( + completions.requests[1]["messages"][-1], + { + "role": "tool", + "tool_call_id": "call-add", + "content": '{"value":42}', + }, + ) diff --git a/tests/test_examples.py b/tests/test_examples.py index 3fbbce6..ee8cab8 100644 --- a/tests/test_examples.py +++ b/tests/test_examples.py @@ -2,7 +2,10 @@ import unittest from pathlib import Path -from ejagent import MethodToolHandler, OpenAIModelAdapter, SkillManager +from ejagent.context import SkillsContextPipeline +from ejagent.contracts import CancellationSource, ToolCall +from ejagent.providers import AnthropicModelPort, OpenAIModelPort +from ejagent.tools import FunctionTool, FunctionToolExecutor EXAMPLES_DIR = Path(__file__).parents[1] / "examples" @@ -19,59 +22,56 @@ async def test_custom_tool_example_dispatches(self) -> None: EXAMPLES_DIR / "02_custom_tool.py", run_name="example_test", ) - handler = namespace["MathHandler"]() + executor = FunctionToolExecutor( + (FunctionTool(namespace["ADD_TOOL"], namespace["add"]),) + ) - self.assertIsInstance(handler, MethodToolHandler) - outcome = await handler.dispatch("add", {"left": 19.5, "right": 22.5}) + outcome = await executor.execute( + ToolCall("test-call", "add", {"left": 19.5, "right": 22.5}), + cancellation=CancellationSource().token, + ) self.assertEqual( - outcome.data, + outcome.result, {"status": "success", "value": 42.0}, ) async def test_skill_example_discovers_local_skill(self) -> None: namespace = runpy.run_path( - EXAMPLES_DIR / "06_skill.py", + EXAMPLES_DIR / "05_skill.py", run_name="example_test", ) - manager = SkillManager(namespace["SKILLS_DIR"]) + pipeline = SkillsContextPipeline(namespace["SKILLS_DIR"]) - await manager.discover() + await pipeline.start() + manager = pipeline.catalog - self.assertIn("release_notes", manager._skills) - skill = manager._skills["release_notes"] + self.assertIn("release_notes", tuple(item.name for item in manager.skills)) + skill = manager.get("release_notes") self.assertIsNotNone(skill.template_md) self.assertIsNotNone(skill.sample_md) + await pipeline.shutdown() def test_harness_examples_use_real_provider_adapter(self) -> None: for filename in ( - "07_event_observers.py", - "08_session_resume.py", - "09_runtime_control.py", - "10_composed_harness.py", - "11_streaming_events.py", - "12_tool_progress.py", - "13_usage_budget.py", - "14_context_pressure.py", - "15_explicit_compaction.py", - "16_durable_session.py", + "01_stateful_chat.py", + "02_custom_tool.py", + "04_mcp_tools.py", + "05_skill.py", + "06_session_resume.py", + "07_durable_session.py", ): with self.subTest(example=filename): namespace = runpy.run_path( EXAMPLES_DIR / filename, run_name="example_test", ) - if filename == "09_runtime_control.py": - adapter = namespace["ObservableOpenAIModelAdapter"] - self.assertTrue(issubclass(adapter, OpenAIModelAdapter)) - else: - self.assertIs( - namespace["OpenAIModelAdapter"], - OpenAIModelAdapter, - ) + self.assertIs(namespace["OpenAIModelPort"], OpenAIModelPort) - def test_skill_manager_requires_an_explicit_root(self) -> None: - with self.assertRaisesRegex(TypeError, "skills_root"): - SkillManager() # type: ignore[call-arg] + namespace = runpy.run_path( + EXAMPLES_DIR / "03_anthropic_chat.py", + run_name="example_test", + ) + self.assertIs(namespace["AnthropicModelPort"], AnthropicModelPort) if __name__ == "__main__": diff --git a/tests/test_harness_controls.py b/tests/test_harness_controls.py new file mode 100644 index 0000000..c12a7d8 --- /dev/null +++ b/tests/test_harness_controls.py @@ -0,0 +1,486 @@ +from __future__ import annotations + +import asyncio +import unittest +from collections.abc import AsyncIterator, Sequence + +from ejagent.contracts import ( + AssistantMessage, + CancellationToken, + ControlStatus, + FailureCode, + ModelCallError, + ModelPort, + ModelRequest, + ModelResponseCompleted, + ModelStreamEvent, + RunAudit, + RunStatus, + SessionCommit, + SessionSnapshot, + ToolCall, + ToolDefinition, + ToolExecutionResult, + ToolExecutor, + ToolSemantics, + TransientInstruction, + UserMessage, +) +from ejagent.harness import ( + AgentHarness, + FollowUpDiscardedError, + FollowUpRejectedError, + MemorySessionStore, +) + + +class ScriptedModel(ModelPort): + def __init__( + self, + responses: Sequence[AssistantMessage | ModelCallError], + *, + block_first: bool = False, + ) -> None: + self.responses = list(responses) + self.requests: list[ModelRequest] = [] + self.block_first = block_first + self.started = asyncio.Event() + self.release = asyncio.Event() + + async def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + self.requests.append(request) + if self.block_first and len(self.requests) == 1: + self.started.set() + await self.release.wait() + response = self.responses.pop(0) + if isinstance(response, ModelCallError): + raise response + yield ModelResponseCompleted(response) + + +class NoTools(ToolExecutor): + @property + def definitions(self) -> Sequence[ToolDefinition]: + return () + + async def execute( + self, + call: ToolCall, + *, + cancellation: CancellationToken, + ) -> ToolExecutionResult: + raise AssertionError("no tools are registered") + + +class BlockingTools(ToolExecutor): + def __init__(self) -> None: + self.started = asyncio.Event() + self.release = asyncio.Event() + + @property + def definitions(self) -> Sequence[ToolDefinition]: + return ( + ToolDefinition( + name="wait", + semantics=ToolSemantics.read_only(), + ), + ) + + async def execute( + self, + call: ToolCall, + *, + cancellation: CancellationToken, + ) -> ToolExecutionResult: + self.started.set() + await self.release.wait() + return ToolExecutionResult({"released": True}) + + +class BlockingObserver: + def __init__(self) -> None: + self.started = asyncio.Event() + self.release = asyncio.Event() + self.audits: list[RunAudit] = [] + + async def observe(self, audit: RunAudit) -> None: + self.audits.append(audit) + self.started.set() + await self.release.wait() + + +class FailingObserver: + def __init__(self) -> None: + self.called = asyncio.Event() + + async def observe(self, audit: RunAudit) -> None: + self.called.set() + raise RuntimeError("observer unavailable") + + +class StoreCheckingObserver: + def __init__(self, store: MemorySessionStore) -> None: + self.store = store + self.called = asyncio.Event() + self.persisted_revision: int | None = None + + async def observe(self, audit: RunAudit) -> None: + snapshot = await self.store.load("observed-agent") + self.persisted_revision = snapshot.revision if snapshot is not None else None + self.called.set() + + +class BlockingCommitStore(MemorySessionStore): + def __init__(self) -> None: + super().__init__() + self.started = asyncio.Event() + self.release = asyncio.Event() + + async def commit(self, commit: SessionCommit) -> SessionSnapshot: + self.started.set() + await self.release.wait() + return await super().commit(commit) + + +class LifecycleObserver: + def __init__(self, events: list[str]) -> None: + self.events = events + self.called = asyncio.Event() + + async def start(self) -> None: + self.events.append("observer:start") + + async def observe(self, audit: RunAudit) -> None: + self.events.append("observer:observe") + self.called.set() + + async def shutdown(self) -> None: + self.events.append("observer:shutdown") + + +class HarnessControlTests(unittest.IsolatedAsyncioTestCase): + async def test_steering_reaches_next_model_call_without_entering_history( + self, + ) -> None: + call = ToolCall(id="wait-1", name="wait") + model = ScriptedModel( + [ + AssistantMessage(tool_calls=(call,)), + AssistantMessage(content="steered result"), + ] + ) + tools = BlockingTools() + store = MemorySessionStore() + harness = AgentHarness( + agent_id="steered-agent", + model=model, + tools=tools, + store=store, + run_id_factory=lambda: "steered-run", + ) + running = asyncio.create_task(harness.run("original task")) + await tools.started.wait() + + receipt = harness.steer("use the safer approach") + tools.release.set() + outcome = await running + + self.assertTrue(receipt.accepted) + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertEqual( + model.requests[1].messages[-1], + TransientInstruction("use the safer approach", "steering"), + ) + self.assertFalse( + any( + isinstance(message, TransientInstruction) + for message in harness.messages + ) + ) + applied = [ + record + for record in outcome.audit_records + if record.kind == "steering_applied" + ] + self.assertEqual(applied[0].payload["input_id"], receipt.input_id) + self.assertEqual(applied[0].payload["turn"], 2) + durable = await store.load_audit("steered-agent") + self.assertIn( + "steering_applied", [record.kind for record in durable[0].records] + ) + + async def test_steering_without_another_safe_point_is_audited_as_discarded( + self, + ) -> None: + model = ScriptedModel( + [AssistantMessage(content="terminal")], + block_first=True, + ) + harness = AgentHarness( + agent_id="discarded-steering", + model=model, + tools=NoTools(), + run_id_factory=lambda: "discarded-run", + ) + running = asyncio.create_task(harness.run("task")) + await model.started.wait() + + receipt = harness.steer("arrived after the safe point") + model.release.set() + outcome = await running + + self.assertTrue(receipt.accepted) + discarded = [ + record + for record in outcome.audit_records + if record.kind == "steering_discarded" + ] + self.assertEqual(discarded[0].payload["input_id"], receipt.input_id) + self.assertEqual(discarded[0].payload["reason"], "run_finished") + + async def test_steering_admission_reports_idle_full_and_closed(self) -> None: + model = ScriptedModel( + [AssistantMessage(content="done")], + block_first=True, + ) + harness = AgentHarness( + agent_id="steering-capacity", + model=model, + tools=NoTools(), + steering_capacity=1, + ) + self.assertEqual(harness.steer("idle").status, ControlStatus.NOT_RUNNING) + running = asyncio.create_task(harness.run("task")) + await model.started.wait() + + self.assertEqual(harness.steer("accepted").status, ControlStatus.ACCEPTED) + self.assertEqual(harness.steer("full").status, ControlStatus.QUEUE_FULL) + model.release.set() + await running + await harness.shutdown() + + self.assertEqual(harness.steer("closed").status, ControlStatus.CLOSED) + + async def test_steering_after_last_safe_point_is_rejected_as_too_late( + self, + ) -> None: + store = BlockingCommitStore() + harness = AgentHarness( + agent_id="too-late-steering", + model=ScriptedModel([AssistantMessage(content="done")]), + tools=NoTools(), + store=store, + ) + running = asyncio.create_task(harness.run("task")) + await store.started.wait() + + receipt = harness.steer("cannot reach a model call") + + self.assertEqual(receipt.status, ControlStatus.TOO_LATE) + store.release.set() + await running + + async def test_follow_ups_run_as_independent_fifo_revisions(self) -> None: + model = ScriptedModel( + [ + AssistantMessage(content="initial result"), + AssistantMessage(content="first result"), + AssistantMessage(content="second result"), + ], + block_first=True, + ) + run_ids = iter(("initial-run", "first-run", "second-run")) + harness = AgentHarness( + agent_id="follow-up-agent", + model=model, + tools=NoTools(), + run_id_factory=lambda: next(run_ids), + ) + running = asyncio.create_task(harness.run("initial")) + await model.started.wait() + + first = harness.follow_up("first") + second = harness.follow_up("second") + self.assertTrue(first.accepted) + self.assertTrue(second.accepted) + self.assertEqual(harness.pending_follow_up_count, 2) + model.release.set() + + initial = await running + first_outcome = await first.wait() + second_outcome = await second.wait() + + self.assertEqual(initial.delta.base_revision, 0) + self.assertEqual(first_outcome.delta.base_revision, 1) + self.assertEqual(second_outcome.delta.base_revision, 2) + self.assertEqual(harness.revision, 3) + user_tasks = [ + message.content + for message in harness.messages + if isinstance(message, UserMessage) + ] + self.assertEqual(user_tasks, ["initial", "first", "second"]) + self.assertEqual(harness.pending_follow_up_count, 0) + + async def test_follow_up_capacity_and_idle_rejection_are_explicit(self) -> None: + model = ScriptedModel( + [ + AssistantMessage(content="initial"), + AssistantMessage(content="followed"), + ], + block_first=True, + ) + harness = AgentHarness( + agent_id="follow-up-capacity", + model=model, + tools=NoTools(), + follow_up_capacity=1, + ) + idle = harness.follow_up("idle") + with self.assertRaises(FollowUpRejectedError): + await idle.wait() + + running = asyncio.create_task(harness.run("initial")) + await model.started.wait() + accepted = harness.follow_up("accepted") + full = harness.follow_up("full") + + self.assertEqual(accepted.receipt.status, ControlStatus.ACCEPTED) + self.assertEqual(full.receipt.status, ControlStatus.QUEUE_FULL) + with self.assertRaises(FollowUpRejectedError): + await full.wait() + model.release.set() + await running + await accepted.wait() + + async def test_follow_up_runs_even_when_originating_run_fails(self) -> None: + model = ScriptedModel( + [ + ModelCallError(FailureCode.RATE_LIMIT, "busy", retryable=True), + AssistantMessage(content="recovered"), + ], + block_first=True, + ) + harness = AgentHarness( + agent_id="follow-up-after-failure", + model=model, + tools=NoTools(), + ) + running = asyncio.create_task(harness.run("initial")) + await model.started.wait() + follow_up = harness.follow_up("retry independently") + model.release.set() + + initial = await running + followed = await follow_up.wait() + + self.assertEqual(initial.result.status, RunStatus.FAILED) + self.assertEqual(followed.result.status, RunStatus.COMPLETED) + self.assertEqual(followed.delta.base_revision, 0) + self.assertEqual(harness.revision, 1) + + async def test_shutdown_discards_follow_up_that_has_not_started(self) -> None: + model = ScriptedModel( + [AssistantMessage(content="unreachable")], + block_first=True, + ) + harness = AgentHarness( + agent_id="follow-up-shutdown", + model=model, + tools=NoTools(), + ) + running = asyncio.create_task(harness.run("blocking")) + await model.started.wait() + follow_up = harness.follow_up("discard me") + + await harness.shutdown() + outcome = await running + + self.assertEqual(outcome.result.status, RunStatus.CANCELLED) + with self.assertRaises(FollowUpDiscardedError): + await follow_up.wait() + + async def test_observer_is_not_on_run_critical_path(self) -> None: + observer = BlockingObserver() + harness = AgentHarness( + agent_id="observer-agent", + model=ScriptedModel([AssistantMessage(content="done")]), + tools=NoTools(), + observers=(observer,), + ) + + outcome = await asyncio.wait_for(harness.run("task"), timeout=0.2) + await observer.started.wait() + + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertFalse(observer.release.is_set()) + observer.release.set() + await harness.shutdown() + self.assertEqual(observer.audits[0].result, outcome.result) + + async def test_observer_failure_cannot_change_run_or_commit(self) -> None: + observer = FailingObserver() + store = MemorySessionStore() + harness = AgentHarness( + agent_id="failing-observer", + model=ScriptedModel([AssistantMessage(content="done")]), + tools=NoTools(), + store=store, + observers=(observer,), + ) + + outcome = await harness.run("task") + await observer.called.wait() + await asyncio.sleep(0) + + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertEqual(harness.revision, 1) + persisted = await store.load("failing-observer") + assert persisted is not None + self.assertEqual(persisted.revision, 1) + await harness.shutdown() + + async def test_observer_runs_after_store_decision(self) -> None: + store = MemorySessionStore() + observer = StoreCheckingObserver(store) + harness = AgentHarness( + agent_id="observed-agent", + model=ScriptedModel([AssistantMessage(content="done")]), + tools=NoTools(), + store=store, + observers=(observer,), + ) + + await harness.run("task") + await observer.called.wait() + + self.assertEqual(observer.persisted_revision, 1) + await harness.shutdown() + + async def test_harness_owns_observer_lifecycle(self) -> None: + events: list[str] = [] + observer = LifecycleObserver(events) + harness = AgentHarness( + agent_id="observer-lifecycle", + model=ScriptedModel([AssistantMessage(content="done")]), + tools=NoTools(), + observers=(observer,), + ) + + await harness.run("task") + await observer.called.wait() + await harness.shutdown() + + self.assertEqual( + events, + ["observer:start", "observer:observe", "observer:shutdown"], + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_mcp_manager.py b/tests/test_mcp_manager.py index 6096da4..cacdf15 100644 --- a/tests/test_mcp_manager.py +++ b/tests/test_mcp_manager.py @@ -3,7 +3,7 @@ from types import SimpleNamespace from typing import Any -from ejagent.plugins.mcp.mcp_manager import McpServerManager +from ejagent.tools._mcp_manager import McpServerManager class FakeMcpClient: diff --git a/tests/test_model_compactor.py b/tests/test_model_compactor.py deleted file mode 100644 index a6d24b1..0000000 --- a/tests/test_model_compactor.py +++ /dev/null @@ -1,132 +0,0 @@ -import unittest - -from ejagent import ( - AgentContextBuilder, - AgentState, - AssistantMessage, - CancellationSource, - CancellationToken, - CompactionRequest, - ContextBuildResult, - ModelAdapter, - ModelCompactor, - prepare_compaction, -) - - -class SummaryModel(ModelAdapter): - def __init__(self, content: str | None) -> None: - self.content = content - self.contexts: list[ContextBuildResult] = [] - self.cancellations: list[CancellationToken | None] = [] - - async def complete( - self, - context: ContextBuildResult, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.contexts.append(context) - self.cancellations.append(cancellation) - return AssistantMessage(content=self.content) - - -def request() -> CompactionRequest: - preparation = prepare_compaction( - [ - {"role": "user", "content": "old"}, - {"role": "assistant", "content": "answer"}, - {"role": "user", "content": "recent"}, - ], - keep_recent_tokens=1, - ) - return CompactionRequest(preparation) - - -class ModelCompactorTests(unittest.IsolatedAsyncioTestCase): - async def test_injected_builder_and_model_produce_compactor_output(self) -> None: - model = SummaryModel(" concise summary ") - built_requests: list[CompactionRequest] = [] - - def build_context(compaction: CompactionRequest) -> ContextBuildResult: - built_requests.append(compaction) - state = AgentState( - messages=[ - {"role": "system", "content": "Summarize reliably."}, - { - "role": "user", - "content": repr(compaction.preparation.messages_to_summarize), - }, - ] - ) - return AgentContextBuilder().build(state) - - compactor = ModelCompactor( - model, - context_builder=build_context, - source="summary-model:test", - ) - source = CancellationSource() - active_request = request() - - output = await compactor.compact( - active_request, - cancellation=source.token, - ) - - self.assertEqual(output.content, "concise summary") - self.assertEqual(output.source, "summary-model:test") - self.assertEqual(built_requests, [active_request]) - self.assertEqual(len(model.contexts), 1) - self.assertIs(model.cancellations[0], source.token) - - async def test_empty_model_response_is_rejected(self) -> None: - model = SummaryModel(None) - context = AgentContextBuilder().build(AgentState()) - compactor = ModelCompactor( - model, - context_builder=lambda _: context, - source="summary-model:test", - ) - - with self.assertRaisesRegex(RuntimeError, "empty content"): - await compactor.compact(request()) - - async def test_invalid_context_builder_result_is_rejected(self) -> None: - model = SummaryModel("summary") - compactor = ModelCompactor( - model, - context_builder=lambda _: object(), # type: ignore[arg-type,return-value] - source="summary-model:test", - ) - - with self.assertRaisesRegex(TypeError, "ContextBuildResult"): - await compactor.compact(request()) - self.assertEqual(model.contexts, []) - - async def test_cancelled_request_does_not_call_model(self) -> None: - model = SummaryModel("summary") - context = AgentContextBuilder().build(AgentState()) - compactor = ModelCompactor( - model, - context_builder=lambda _: context, - source="summary-model:test", - ) - source = CancellationSource() - source.cancel("stop summary") - - with self.assertRaisesRegex(RuntimeError, "stop summary"): - await compactor.compact(request(), cancellation=source.token) - self.assertEqual(model.contexts, []) - - def test_source_is_validated(self) -> None: - with self.assertRaisesRegex(ValueError, "source"): - ModelCompactor( - SummaryModel("summary"), - context_builder=lambda _: AgentContextBuilder().build(AgentState()), - source=" ", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_parallel_tool_calls.py b/tests/test_parallel_tool_calls.py deleted file mode 100644 index bbff2b4..0000000 --- a/tests/test_parallel_tool_calls.py +++ /dev/null @@ -1,533 +0,0 @@ -import asyncio -import json -import unittest -from typing import Any - -from ejagent import ( - AgentEvent, - AssistantMessage, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - MemorySessionStorage, - MethodToolHandler, - ModelAdapter, - ModelToolCall, - RunStatus, - RuntimePolicy, - SessionRecorder, - StepOutcome, - StopReason, - ToolCompleted, - ToolControl, - ToolDefinitionError, - ToolEffect, - ToolProgressed, - ToolProgressReporter, - ToolProgressUpdate, - ToolStarted, -) - - -def tool_definition(name: str) -> dict[str, Any]: - return { - "type": "function", - "function": { - "name": name, - "description": f"Run {name}.", - "parameters": { - "type": "object", - "properties": {"label": {"type": "string"}}, - "required": ["label"], - }, - }, - } - - -READ_TOOL = tool_definition("read") -WRITE_TOOL = tool_definition("write") - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - return self.responses.pop(0) - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class ControlledHandler(MethodToolHandler): - def __init__( - self, - *, - read_only: bool = True, - failures: set[str] | None = None, - controls: dict[str, ToolControl] | None = None, - ) -> None: - effects = {"read": ToolEffect.READ_ONLY} if read_only else None - super().__init__((READ_TOOL, WRITE_TOOL), tool_effects=effects) - self.started: dict[str, asyncio.Event] = {} - self.release: dict[str, asyncio.Event] = {} - self.finished: dict[str, asyncio.Event] = {} - self.failures = failures or set() - self.controls = controls or {} - self.active = 0 - self.max_active = 0 - self.calls: list[str] = [] - - def started_for(self, label: str) -> asyncio.Event: - return self.started.setdefault(label, asyncio.Event()) - - def release_for(self, label: str) -> asyncio.Event: - return self.release.setdefault(label, asyncio.Event()) - - def finished_for(self, label: str) -> asyncio.Event: - return self.finished.setdefault(label, asyncio.Event()) - - async def do_read( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - progress: ToolProgressReporter | None = None, - ) -> StepOutcome: - return await self._execute(arguments, progress=progress) - - async def do_write( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - progress: ToolProgressReporter | None = None, - ) -> StepOutcome: - return await self._execute(arguments, progress=progress) - - async def _execute( - self, - arguments: dict[str, Any], - *, - progress: ToolProgressReporter | None, - ) -> StepOutcome: - label = str(arguments["label"]) - self.calls.append(label) - self.active += 1 - self.max_active = max(self.max_active, self.active) - self.started_for(label).set() - if progress is not None: - await progress.report(ToolProgressUpdate(f"started {label}")) - try: - await self.release_for(label).wait() - if label in self.failures: - raise RuntimeError(f"failed {label}") - return StepOutcome( - {"label": label}, - control=self.controls.get(label, ToolControl.CONTINUE), - ) - finally: - self.active -= 1 - self.finished_for(label).set() - - -def call(call_id: str, name: str, label: str) -> ModelToolCall: - return ModelToolCall( - id=call_id, - name=name, - arguments=json.dumps({"label": label}), - ) - - -def tool_payloads(sink: RecordingSink, payload_type: type[Any]) -> list[Any]: - return [ - event.payload - for event in sink.events - if isinstance(event.payload, payload_type) - ] - - -class ParallelToolCallTests(unittest.IsolatedAsyncioTestCase): - async def test_default_policy_keeps_declared_read_only_calls_sequential( - self, - ) -> None: - handler = ControlledHandler() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=( - call("one", "read", "one"), - call("two", "read", "two"), - ) - ), - AssistantMessage(content="done"), - ] - ), - agent_id="parallel-default-off", - handlers=[handler], - ) - - run = asyncio.create_task(agent.run(task="read twice")) - await handler.started_for("one").wait() - await asyncio.sleep(0) - self.assertFalse(handler.started_for("two").is_set()) - handler.release_for("one").set() - await handler.started_for("two").wait() - handler.release_for("two").set() - result = await run - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(handler.max_active, 1) - - async def test_read_only_calls_run_concurrently_but_commit_in_source_order( - self, - ) -> None: - handler = ControlledHandler() - sink = RecordingSink() - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="parallel-order", storage=storage) - calls = ( - call("slow", "read", "slow"), - call("fast", "read", "fast"), - ) - agent = BaseAgent( - SequenceModel( - [AssistantMessage(tool_calls=calls), AssistantMessage(content="done")] - ), - agent_id="parallel-order", - handlers=[handler], - runtime_policy=RuntimePolicy(parallel_tool_calls=True), - event_sink=CompositeAgentEventSink([recorder, sink]), - ) - - run = asyncio.create_task(agent.run(task="read concurrently")) - await asyncio.gather( - handler.started_for("slow").wait(), - handler.started_for("fast").wait(), - ) - handler.release_for("fast").set() - await handler.finished_for("fast").wait() - self.assertFalse(handler.finished_for("slow").is_set()) - handler.release_for("slow").set() - result = await run - session = await recorder.load() - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(handler.max_active, 2) - self.assertEqual( - [payload.tool_call.id for payload in tool_payloads(sink, ToolStarted)], - ["slow", "fast"], - ) - self.assertEqual( - [payload.tool_call.id for payload in tool_payloads(sink, ToolCompleted)], - ["slow", "fast"], - ) - progressed = tool_payloads(sink, ToolProgressed) - self.assertCountEqual( - [payload.tool_call.id for payload in progressed], - ["slow", "fast"], - ) - first_completion = next( - index - for index, event in enumerate(sink.events) - if isinstance(event.payload, ToolCompleted) - ) - self.assertTrue( - all( - index < first_completion - for index, event in enumerate(sink.events) - if isinstance(event.payload, ToolProgressed) - ) - ) - self.assertEqual( - [ - message["tool_call_id"] - for message in agent.messages - if message["role"] == "tool" - ], - ["slow", "fast"], - ) - assert session is not None - self.assertEqual( - [ - message["tool_call_id"] - for message in session.messages - if message["role"] == "tool" - ], - ["slow", "fast"], - ) - - async def test_side_effecting_call_is_a_barrier_between_read_batches(self) -> None: - handler = ControlledHandler() - calls = ( - call("r1", "read", "r1"), - call("r2", "read", "r2"), - call("w", "write", "w"), - call("r3", "read", "r3"), - call("r4", "read", "r4"), - ) - agent = BaseAgent( - SequenceModel( - [AssistantMessage(tool_calls=calls), AssistantMessage(content="done")] - ), - agent_id="parallel-barrier", - handlers=[handler], - runtime_policy=RuntimePolicy(parallel_tool_calls=True), - ) - - run = asyncio.create_task(agent.run(task="mixed tools")) - await asyncio.gather( - handler.started_for("r1").wait(), - handler.started_for("r2").wait(), - ) - self.assertFalse(handler.started_for("w").is_set()) - handler.release_for("r1").set() - handler.release_for("r2").set() - await handler.started_for("w").wait() - self.assertFalse(handler.started_for("r3").is_set()) - handler.release_for("w").set() - await asyncio.gather( - handler.started_for("r3").wait(), - handler.started_for("r4").wait(), - ) - handler.release_for("r3").set() - handler.release_for("r4").set() - result = await run - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(handler.calls, ["r1", "r2", "w", "r3", "r4"]) - self.assertEqual(handler.max_active, 2) - - async def test_parallel_limit_bounds_each_read_batch(self) -> None: - handler = ControlledHandler() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=tuple( - call(label, "read", label) - for label in ("one", "two", "three") - ) - ), - AssistantMessage(content="done"), - ] - ), - agent_id="parallel-limit", - handlers=[handler], - runtime_policy=RuntimePolicy( - parallel_tool_calls=True, - max_parallel_tool_calls=2, - ), - ) - - run = asyncio.create_task(agent.run(task="bounded reads")) - await asyncio.gather( - handler.started_for("one").wait(), - handler.started_for("two").wait(), - ) - self.assertFalse(handler.started_for("three").is_set()) - handler.release_for("one").set() - handler.release_for("two").set() - await handler.started_for("three").wait() - handler.release_for("three").set() - await run - - self.assertEqual(handler.max_active, 2) - - async def test_parallel_error_becomes_ordered_tool_result_and_loop_continues( - self, - ) -> None: - handler = ControlledHandler(failures={"bad"}) - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=( - call("bad", "read", "bad"), - call("good", "read", "good"), - ) - ), - AssistantMessage(content="recovered"), - ] - ), - agent_id="parallel-error", - handlers=[handler], - runtime_policy=RuntimePolicy(parallel_tool_calls=True), - ) - - run = asyncio.create_task(agent.run(task="recover")) - await asyncio.gather( - handler.started_for("bad").wait(), - handler.started_for("good").wait(), - ) - handler.release_for("good").set() - handler.release_for("bad").set() - result = await run - tool_messages = [ - message for message in agent.messages if message["role"] == "tool" - ] - - self.assertEqual(result.output, "recovered") - self.assertEqual( - [message["tool_call_id"] for message in tool_messages], - ["bad", "good"], - ) - self.assertIn("failed bad", tool_messages[0]["content"]) - - async def test_terminal_read_batch_settles_peers_before_stopping(self) -> None: - handler = ControlledHandler( - controls={ - "finish": ToolControl.COMPLETE, - "peer": ToolControl.REJECT, - } - ) - sink = RecordingSink() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=( - call("finish", "read", "finish"), - call("peer", "read", "peer"), - call("late", "read", "late"), - call("write", "write", "write"), - ) - ) - ] - ), - agent_id="parallel-terminal", - handlers=[handler], - runtime_policy=RuntimePolicy( - parallel_tool_calls=True, - max_parallel_tool_calls=2, - ), - event_sink=sink, - ) - - run = asyncio.create_task(agent.run(task="finish safely")) - await asyncio.gather( - handler.started_for("finish").wait(), - handler.started_for("peer").wait(), - ) - handler.release_for("peer").set() - handler.release_for("finish").set() - result = await run - - self.assertEqual(result.stop_reason, StopReason.TOOL_COMPLETION) - self.assertEqual(json.loads(result.output or "{}"), {"label": "finish"}) - self.assertEqual(handler.calls, ["finish", "peer"]) - self.assertEqual( - [payload.tool_call.id for payload in tool_payloads(sink, ToolCompleted)], - ["finish", "peer", "late"], - ) - self.assertTrue(tool_payloads(sink, ToolCompleted)[-1].result.cancelled) - - async def test_abort_settles_parallel_and_pending_calls_in_source_order( - self, - ) -> None: - handler = ControlledHandler() - sink = RecordingSink() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=( - call("one", "read", "one"), - call("two", "read", "two"), - call("write", "write", "write"), - ) - ) - ] - ), - agent_id="parallel-cancel", - handlers=[handler], - runtime_policy=RuntimePolicy(parallel_tool_calls=True), - event_sink=sink, - ) - - run = asyncio.create_task(agent.run(task="cancel reads")) - await asyncio.gather( - handler.started_for("one").wait(), - handler.started_for("two").wait(), - ) - self.assertTrue(agent.abort("stop parallel calls")) - result = await run - - self.assertEqual(result.status, RunStatus.CANCELLED) - self.assertEqual(result.stop_reason, StopReason.EXTERNAL_ABORT) - self.assertEqual(handler.calls, ["one", "two"]) - self.assertEqual( - [payload.tool_call.id for payload in tool_payloads(sink, ToolStarted)], - ["one", "two", "write"], - ) - completed = tool_payloads(sink, ToolCompleted) - self.assertEqual( - [payload.tool_call.id for payload in completed], - ["one", "two", "write"], - ) - self.assertTrue(all(payload.result.cancelled for payload in completed)) - self.assertEqual( - [ - message["tool_call_id"] - for message in agent.messages - if message["role"] == "tool" - ], - ["one", "two", "write"], - ) - - async def test_repetition_guard_falls_back_to_ordered_execution(self) -> None: - handler = ControlledHandler() - repeated_calls = ( - call("one", "read", "same"), - call("two", "read", "same"), - ) - agent = BaseAgent( - SequenceModel([AssistantMessage(tool_calls=repeated_calls)]), - agent_id="parallel-repeat", - handlers=[handler], - runtime_policy=RuntimePolicy( - max_repeated_tool_calls=2, - parallel_tool_calls=True, - ), - ) - - run = asyncio.create_task(agent.run(task="repeat")) - await handler.started_for("same").wait() - handler.release_for("same").set() - result = await run - - self.assertEqual(result.stop_reason, StopReason.REPEATED_TOOL_CALL) - self.assertEqual(handler.max_active, 1) - - def test_effect_and_parallel_limit_validation(self) -> None: - with self.assertRaises(ToolDefinitionError): - MethodToolHandler( - (READ_TOOL,), - tool_effects={"missing": ToolEffect.READ_ONLY}, - ) - with self.assertRaises(ToolDefinitionError): - MethodToolHandler( - (READ_TOOL,), - tool_effects={"read": "read_only"}, # type: ignore[dict-item] - ) - with self.assertRaises(ValueError): - RuntimePolicy(max_parallel_tool_calls=0) - with self.assertRaises(TypeError): - RuntimePolicy(parallel_tool_calls="yes") # type: ignore[arg-type] - with self.assertRaises(TypeError): - RuntimePolicy(max_parallel_tool_calls=True) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_provider_errors.py b/tests/test_provider_errors.py deleted file mode 100644 index cdd49cf..0000000 --- a/tests/test_provider_errors.py +++ /dev/null @@ -1,94 +0,0 @@ -import unittest - -import httpx -from openai import APITimeoutError, AuthenticationError, RateLimitError - -from ejagent import ( - ContextOverflowError, - ModelAuthenticationError, - ModelErrorKind, - ModelProviderError, - ModelRateLimitError, - ModelTimeoutError, -) -from ejagent.providers.openai import _normalize_provider_error - - -class StructuredProviderError(Exception): - def __init__(self, code: str, message: str) -> None: - self.code = code - super().__init__(message) - - -class ProviderErrorTests(unittest.TestCase): - def test_structured_context_code_is_normalized(self) -> None: - error = _normalize_provider_error( - StructuredProviderError( - "context_length_exceeded", - "request rejected", - ), - operation="chat completion", - ) - - self.assertIsInstance(error, ContextOverflowError) - self.assertEqual(error.kind, ModelErrorKind.CONTEXT_OVERFLOW) - - def test_compatible_context_message_fallback_is_centralized(self) -> None: - error = _normalize_provider_error( - RuntimeError("maximum context length is 128000 tokens"), - operation="chat completion stream", - ) - - self.assertIsInstance(error, ContextOverflowError) - - def test_openai_rate_auth_and_timeout_errors_are_normalized(self) -> None: - request = httpx.Request("POST", "https://example.invalid/chat") - auth_response = httpx.Response(401, request=request) - rate_response = httpx.Response(429, request=request) - cases = [ - ( - AuthenticationError( - "bad key", - response=auth_response, - body=None, - ), - ModelAuthenticationError, - ModelErrorKind.AUTHENTICATION, - ), - ( - RateLimitError( - "slow down", - response=rate_response, - body=None, - ), - ModelRateLimitError, - ModelErrorKind.RATE_LIMIT, - ), - ( - APITimeoutError(request), - ModelTimeoutError, - ModelErrorKind.TIMEOUT, - ), - ] - - for source, error_type, kind in cases: - with self.subTest(kind=kind): - error = _normalize_provider_error( - source, - operation="chat completion", - ) - self.assertIsInstance(error, error_type) - self.assertEqual(error.kind, kind) - - def test_unknown_error_remains_provider_error(self) -> None: - error = _normalize_provider_error( - ValueError("unexpected response"), - operation="chat completion", - ) - - self.assertIs(type(error), ModelProviderError) - self.assertEqual(error.kind, ModelErrorKind.PROVIDER_ERROR) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_public_api.py b/tests/test_public_api.py new file mode 100644 index 0000000..b672fed --- /dev/null +++ b/tests/test_public_api.py @@ -0,0 +1,48 @@ +from __future__ import annotations + +import unittest + +import ejagent + + +class PublicApiTests(unittest.TestCase): + def test_top_level_exports_only_new_composition_surface(self) -> None: + expected = { + "AgentHarness", + "AnthropicConfig", + "AnthropicModelPort", + "CompositeToolExecutor", + "DerivedCompactionPipeline", + "FunctionTool", + "FunctionToolExecutor", + "IdentityContextPipeline", + "JsonlSessionStore", + "McpToolExecutor", + "MemorySessionStore", + "ModelConfig", + "OpenAIModelPort", + "RuntimeKernel", + "Skill", + "SkillCatalog", + "SkillsContextPipeline", + } + + self.assertEqual(set(ejagent.__all__), expected) + self.assertTrue(all(hasattr(ejagent, name) for name in expected)) + + def test_legacy_entry_points_are_absent(self) -> None: + for name in ( + "BaseAgent", + "AgentOrchestrator", + "ModelAdapter", + "OpenAIModelAdapter", + "BaseHandler", + "ToolMiddleware", + "SessionStorage", + ): + with self.subTest(name=name): + self.assertFalse(hasattr(ejagent, name)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_runtime_kernel.py b/tests/test_runtime_kernel.py new file mode 100644 index 0000000..f76d246 --- /dev/null +++ b/tests/test_runtime_kernel.py @@ -0,0 +1,432 @@ +from __future__ import annotations + +import asyncio +import unittest +from collections.abc import AsyncIterator, Sequence +from datetime import UTC, datetime + +from ejagent.contracts import ( + AssistantMessage, + CancellationSource, + CancellationToken, + FailureCode, + ModelCallError, + ModelPort, + ModelProtocolError, + ModelRequest, + ModelResponseCompleted, + ModelStreamEvent, + ModelTextDelta, + ModelUsage, + RunIntent, + RunLimits, + RunPhase, + RunSpec, + RunStatus, + StopReason, + SystemMessage, + ToolCall, + ToolControl, + ToolDefinition, + ToolExecutionError, + ToolExecutionResult, + ToolExecutor, + ToolProtocolError, + ToolResultMessage, + ToolSemantics, + UserMessage, + thaw_json_value, +) +from ejagent.kernel import RuntimeKernel + + +def fixed_clock() -> datetime: + return datetime(2026, 7, 31, 12, 0, tzinfo=UTC) + + +def task_spec( + *, + run_id: str = "run-1", + base_revision: int = 4, + task: str = "new task", + limits: RunLimits | None = None, +) -> RunSpec: + return RunSpec( + run_id=run_id, + base_revision=base_revision, + intent=RunIntent.TASK, + task=task, + messages=(SystemMessage("You are a test agent."),), + limits=limits or RunLimits(), + configuration_revision="config-1", + ) + + +class ScriptedModel(ModelPort): + def __init__( + self, + responses: Sequence[ + tuple[AssistantMessage, ModelUsage | None] | ModelCallError + ], + ) -> None: + self.responses = list(responses) + self.requests: list[ModelRequest] = [] + + async def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + self.requests.append(request) + response = self.responses.pop(0) + if isinstance(response, ModelCallError): + raise response + message, usage = response + if message.content: + yield ModelTextDelta(message.content) + yield ModelResponseCompleted(message, usage) + + +class IncompleteModel(ModelPort): + async def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + yield ModelTextDelta("partial") + + +class BlockingModel(ModelPort): + def __init__(self) -> None: + self.started = asyncio.Event() + self.release = asyncio.Event() + + async def stream( + self, + request: ModelRequest, + *, + cancellation: CancellationToken, + ) -> AsyncIterator[ModelStreamEvent]: + self.started.set() + await self.release.wait() + yield ModelResponseCompleted(AssistantMessage(content="late")) + + +class RecordingTools(ToolExecutor): + def __init__( + self, + definitions: Sequence[ToolDefinition] = (), + *, + results: dict[str, ToolExecutionResult | ToolExecutionError] | None = None, + ) -> None: + self._definitions = tuple(definitions) + self.results = results or {} + self.calls: list[ToolCall] = [] + + @property + def definitions(self) -> Sequence[ToolDefinition]: + return self._definitions + + async def execute( + self, + call: ToolCall, + *, + cancellation: CancellationToken, + ) -> ToolExecutionResult: + self.calls.append(call) + result = self.results.get(call.id, ToolExecutionResult({"ok": True})) + if isinstance(result, ToolExecutionError): + raise result + return result + + +def lookup_definition() -> ToolDefinition: + return ToolDefinition( + name="lookup", + description="Look up one value.", + input_schema={ + "type": "object", + "properties": {"key": {"type": "string"}}, + "required": ["key"], + }, + semantics=ToolSemantics.read_only(), + ) + + +def lookup_call(call_id: str = "lookup-1") -> ToolCall: + return ToolCall( + id=call_id, + name="lookup", + arguments={"key": "answer"}, + ) + + +class RuntimeKernelTests(unittest.IsolatedAsyncioTestCase): + async def test_text_run_returns_delta_without_mutating_spec(self) -> None: + spec = task_spec() + original_messages = spec.messages + model = ScriptedModel( + [ + ( + AssistantMessage(content="done"), + ModelUsage( + input_tokens=10, + output_tokens=2, + total_tokens=12, + ), + ) + ] + ) + kernel = RuntimeKernel( + model=model, + tools=RecordingTools(), + clock=fixed_clock, + ) + + outcome = await kernel.run(spec) + + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertEqual(outcome.result.stop_reason, StopReason.TEXT_RESPONSE) + self.assertEqual(outcome.result.output, "done") + self.assertEqual(outcome.result.usage.total_tokens, 12) + self.assertEqual( + outcome.delta.messages, + (UserMessage("new task"), AssistantMessage(content="done")), + ) + self.assertEqual(outcome.delta.base_revision, 4) + self.assertEqual(spec.messages, original_messages) + self.assertEqual( + model.requests[0].messages, + (*original_messages, UserMessage("new task")), + ) + self.assertEqual( + [record.kind for record in outcome.audit_records], + [ + "run_started", + "turn_started", + "context_built", + "model_text_delta", + "assistant_message", + "turn_completed", + "run_finished", + ], + ) + + async def test_tool_result_is_committed_before_the_next_model_request( + self, + ) -> None: + call = lookup_call() + model = ScriptedModel( + [ + (AssistantMessage(tool_calls=(call,)), None), + (AssistantMessage(content="answer is 42"), None), + ] + ) + tools = RecordingTools( + (lookup_definition(),), + results={call.id: ToolExecutionResult({"value": 42})}, + ) + kernel = RuntimeKernel(model=model, tools=tools, clock=fixed_clock) + + outcome = await kernel.run(task_spec()) + + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertEqual(len(model.requests), 2) + tool_message = model.requests[1].messages[-1] + self.assertIsInstance(tool_message, ToolResultMessage) + assert isinstance(tool_message, ToolResultMessage) + self.assertEqual(thaw_json_value(tool_message.result), {"value": 42}) + self.assertEqual( + outcome.delta.messages, + ( + UserMessage("new task"), + AssistantMessage(tool_calls=(call,)), + tool_message, + AssistantMessage(content="answer is 42"), + ), + ) + + async def test_completing_tool_terminates_without_another_model_request( + self, + ) -> None: + call = lookup_call() + model = ScriptedModel([(AssistantMessage(tool_calls=(call,)), None)]) + tools = RecordingTools( + (lookup_definition(),), + results={ + call.id: ToolExecutionResult( + {"value": 42}, + control=ToolControl.COMPLETE, + output="finished", + ) + }, + ) + + outcome = await RuntimeKernel( + model=model, + tools=tools, + clock=fixed_clock, + ).run(task_spec()) + + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertEqual(outcome.result.stop_reason, StopReason.TOOL_COMPLETION) + self.assertEqual(outcome.result.output, "finished") + self.assertEqual(len(model.requests), 1) + + async def test_unknown_tool_becomes_an_error_result_for_the_model(self) -> None: + call = ToolCall(id="missing-1", name="missing") + model = ScriptedModel( + [ + (AssistantMessage(tool_calls=(call,)), None), + (AssistantMessage(content="recovered"), None), + ] + ) + tools = RecordingTools() + + outcome = await RuntimeKernel( + model=model, + tools=tools, + clock=fixed_clock, + ).run(task_spec()) + + self.assertEqual(outcome.result.status, RunStatus.COMPLETED) + self.assertEqual(tools.calls, []) + result = model.requests[1].messages[-1] + self.assertIsInstance(result, ToolResultMessage) + assert isinstance(result, ToolResultMessage) + self.assertTrue(result.is_error) + + async def test_model_operational_error_returns_structured_failure(self) -> None: + model = ScriptedModel( + [ + ModelCallError( + FailureCode.TIMEOUT, + "provider timed out", + retryable=True, + ) + ] + ) + + outcome = await RuntimeKernel( + model=model, + tools=RecordingTools(), + clock=fixed_clock, + ).run(task_spec()) + + self.assertEqual(outcome.result.status, RunStatus.FAILED) + self.assertEqual(outcome.failure.phase, RunPhase.MODEL) + self.assertEqual(outcome.failure.code, FailureCode.TIMEOUT) + self.assertTrue(outcome.failure.retryable) + self.assertEqual(outcome.delta.messages, (UserMessage("new task"),)) + + async def test_tool_infrastructure_error_returns_structured_failure(self) -> None: + call = lookup_call() + model = ScriptedModel([(AssistantMessage(tool_calls=(call,)), None)]) + tools = RecordingTools( + (lookup_definition(),), + results={ + call.id: ToolExecutionError("executor unavailable", retryable=True) + }, + ) + + outcome = await RuntimeKernel( + model=model, + tools=tools, + clock=fixed_clock, + ).run(task_spec()) + + self.assertEqual(outcome.result.status, RunStatus.FAILED) + self.assertEqual(outcome.failure.phase, RunPhase.TOOL) + self.assertTrue(outcome.failure.retryable) + self.assertEqual( + outcome.delta.messages, + (UserMessage("new task"), AssistantMessage(tool_calls=(call,))), + ) + + async def test_incomplete_model_stream_raises_protocol_error(self) -> None: + kernel = RuntimeKernel( + model=IncompleteModel(), + tools=RecordingTools(), + clock=fixed_clock, + ) + + with self.assertRaisesRegex(ModelProtocolError, "without completion"): + await kernel.run(task_spec()) + + async def test_cancellation_returns_partial_auditable_outcome(self) -> None: + model = BlockingModel() + source = CancellationSource() + kernel = RuntimeKernel( + model=model, + tools=RecordingTools(), + clock=fixed_clock, + ) + + task = asyncio.create_task(kernel.run(task_spec(), cancellation=source.token)) + await model.started.wait() + source.cancel("stop now") + outcome = await task + + self.assertEqual(outcome.result.status, RunStatus.CANCELLED) + self.assertEqual(outcome.result.stop_reason, StopReason.EXTERNAL_ABORT) + self.assertEqual(outcome.result.output, "stop now") + self.assertEqual(outcome.delta.messages, (UserMessage("new task"),)) + + async def test_max_turns_returns_failure_with_partial_delta(self) -> None: + call = lookup_call() + model = ScriptedModel([(AssistantMessage(tool_calls=(call,)), None)]) + kernel = RuntimeKernel( + model=model, + tools=RecordingTools((lookup_definition(),)), + clock=fixed_clock, + ) + + outcome = await kernel.run(task_spec(limits=RunLimits(max_turns=1))) + + self.assertEqual(outcome.result.status, RunStatus.FAILED) + self.assertEqual(outcome.result.stop_reason, StopReason.MAX_STEPS) + self.assertEqual(outcome.failure.code, FailureCode.BUDGET_EXCEEDED) + + async def test_repeated_tool_guard_rejects_before_third_execution(self) -> None: + calls = tuple(lookup_call(f"lookup-{index}") for index in range(3)) + model = ScriptedModel( + [(AssistantMessage(tool_calls=(call,)), None) for call in calls] + ) + tools = RecordingTools((lookup_definition(),)) + kernel = RuntimeKernel(model=model, tools=tools, clock=fixed_clock) + + outcome = await kernel.run(task_spec()) + + self.assertEqual(outcome.result.stop_reason, StopReason.REPEATED_TOOL_CALL) + self.assertEqual(len(tools.calls), 2) + + async def test_missing_usage_stops_before_second_budgeted_request(self) -> None: + call = lookup_call() + model = ScriptedModel([(AssistantMessage(tool_calls=(call,)), None)]) + kernel = RuntimeKernel( + model=model, + tools=RecordingTools((lookup_definition(),)), + clock=fixed_clock, + ) + + outcome = await kernel.run(task_spec(limits=RunLimits(max_tokens=10))) + + self.assertEqual(outcome.result.stop_reason, StopReason.USAGE_UNAVAILABLE) + self.assertEqual(len(model.requests), 1) + + async def test_duplicate_tool_definitions_are_protocol_error(self) -> None: + definition = lookup_definition() + kernel = RuntimeKernel( + model=ScriptedModel([(AssistantMessage(content="unused"), None)]), + tools=RecordingTools((definition, definition)), + clock=fixed_clock, + ) + + with self.assertRaisesRegex(ToolProtocolError, "duplicate"): + await kernel.run(task_spec()) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_session_persistence.py b/tests/test_session_persistence.py deleted file mode 100644 index a2c3184..0000000 --- a/tests/test_session_persistence.py +++ /dev/null @@ -1,947 +0,0 @@ -import asyncio -import json -import subprocess -import sys -import tempfile -import unittest -from pathlib import Path -from typing import Any -from unittest.mock import patch - -from ejagent import ( - SESSION_SCHEMA_VERSION, - AgentRunResult, - AgentSession, - AssistantMessage, - BaseAgent, - CancellationToken, - JsonlSessionStorage, - ModelAdapter, - RunStatus, - RunUsage, - SessionBranchIntent, - SessionConflictError, - SessionLockTimeoutError, - SessionRecordDraft, - SessionRecorder, - SessionRecordKind, - SessionSerializationError, - SessionStorageError, - SessionTreeStorage, - StopReason, - SummaryEntry, - session_from_dict, - session_to_dict, -) - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[str]) -> None: - self.responses = list(responses) - self.contexts: list[Any] = [] - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - self.contexts.append(context) - return AssistantMessage(content=self.responses.pop(0)) - - -def completed_result(output: str = "完成") -> AgentRunResult: - return AgentRunResult( - status=RunStatus.COMPLETED, - stop_reason=StopReason.TEXT_RESPONSE, - turns=2, - output=output, - usage=RunUsage( - input_tokens=12, - output_tokens=3, - total_tokens=15, - request_count=1, - reported_request_count=1, - cache_read_tokens=2, - reasoning_tokens=1, - ), - ) - - -def durable_session(session_id: str = "持久会话") -> AgentSession: - session = AgentSession(session_id=session_id) - session.bind_agent("durable-agent") - session.begin_run("run-1", "记住 café", 1) - session.append_message( - "run-1", - 2, - { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "call-1", - "type": "function", - "function": { - "name": "lookup", - "arguments": '{"query":"北京"}', - }, - } - ], - }, - ) - session.append_message( - "run-1", - 3, - { - "role": "tool", - "tool_call_id": "call-1", - "content": '{"value":"晴天"}', - }, - ) - session.append_message( - "run-1", - 4, - { - "role": "assistant", - "content": "完成", - "usage": { - "input_tokens": 12, - "output_tokens": 3, - "total_tokens": 15, - "cache_read_tokens": 2, - "cache_write_tokens": None, - "reasoning_tokens": 1, - }, - }, - ) - session.finish_run("run-1", 5, completed_result()) - summary = SummaryEntry( - content="用户要求记住 café;北京查询结果为晴天。", - source="summary-model:test", - history_start_index=1, - first_kept_index=3, - summarized_message_count=2, - tokens_before=15, - ) - session.apply_compaction( - "compact-1", - 6, - summary, - ( - summary.to_agent_message(), - {"role": "assistant", "content": "完成"}, - ), - ) - return session - - -def json_files(directory: str) -> list[Path]: - return list(Path(directory).glob("*.jsonl")) - - -def jsonl_lock_files(directory: str) -> list[Path]: - return list(Path(directory).glob("*.jsonl.lock")) - - -def write_text(path: Path, content: str) -> None: - path.write_text(content, encoding="utf-8") - - -def directory_entries(directory: str) -> list[Path]: - return list(Path(directory).iterdir()) - - -def append_bytes(path: Path, content: bytes) -> None: - with path.open("ab") as stream: - stream.write(content) - - -def read_json_lines(path: Path) -> list[dict[str, Any]]: - return [json.loads(line) for line in path.read_text(encoding="utf-8").splitlines()] - - -def write_json_lines(path: Path, records: list[dict[str, Any]]) -> None: - path.write_text( - "".join(json.dumps(record) + "\n" for record in records), - encoding="utf-8", - ) - - -def start_append_process( - directory: str, - session_id: str, - expected_head_id: str, - run_id: str, - branch_id: str, - start_path: Path, -) -> subprocess.Popen[str]: - script = """ -import asyncio -import json -import sys -import time -from pathlib import Path -from ejagent import ( - JsonlSessionStorage, - SessionConflictError, - SessionRecordDraft, -) - -async def main(): - root, session_id, expected_head, run_id, branch_id, start_path = sys.argv[1:] - while not Path(start_path).exists(): - time.sleep(0.005) - storage = JsonlSessionStorage(root) - draft = SessionRecordDraft.run_started( - session_id=session_id, - agent_id="durable-agent", - sequence=1, - run_id=run_id, - task=run_id, - branch_id=branch_id, - ) - try: - record = await storage.append( - draft, - expected_head_id=expected_head, - check_head=True, - ) - except SessionConflictError: - result = {"status": "conflict"} - else: - result = { - "status": "success", - "record_id": record.record_id, - "parent_id": record.parent_id, - "revision": record.revision, - } - print(json.dumps(result)) - -asyncio.run(main()) -""" - return subprocess.Popen( - [ - sys.executable, - "-c", - script, - directory, - session_id, - expected_head_id, - run_id, - branch_id, - str(start_path), - ], - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - text=True, - ) - - -async def collect_process(process: subprocess.Popen[str]) -> dict[str, Any]: - stdout, stderr = await asyncio.to_thread(process.communicate, timeout=10) - if process.returncode != 0: - raise AssertionError(f"subprocess failed: {stderr}") - return json.loads(stdout) - - -def start_python_process(script: str, *args: str) -> subprocess.Popen[str]: - return subprocess.Popen( - [sys.executable, "-c", script, *args], - stdout=subprocess.PIPE, - stderr=subprocess.PIPE, - text=True, - ) - - -class SessionPersistenceTests(unittest.IsolatedAsyncioTestCase): - def test_versioned_codec_round_trips_complete_session(self) -> None: - session = durable_session() - - payload = session_to_dict(session) - restored = session_from_dict(payload) - - self.assertEqual(payload["schema_version"], SESSION_SCHEMA_VERSION) - self.assertEqual(restored, session) - self.assertEqual(restored.runs[0].result, completed_result()) - self.assertEqual( - restored.compactions[0].summary.content, - session.compactions[0].summary.content, - ) - - async def test_jsonl_storage_round_trips_between_instances(self) -> None: - with tempfile.TemporaryDirectory() as directory: - writer = JsonlSessionStorage(directory) - reader = JsonlSessionStorage(directory) - session = durable_session() - - await writer.save(session) - restored = await reader.load(session.session_id) - - self.assertEqual(restored, session) - files = await asyncio.to_thread(json_files, directory) - self.assertEqual(len(files), 1) - self.assertNotIn("持久会话", files[0].name) - - async def test_session_id_cannot_escape_storage_root(self) -> None: - with tempfile.TemporaryDirectory() as directory: - root = Path(directory) / "sessions" - storage = JsonlSessionStorage(root) - session = durable_session("../../outside") - - await storage.save(session) - - self.assertEqual(len(list(root.glob("*.jsonl"))), 1) - self.assertFalse((Path(directory) / "outside.json").exists()) - - async def test_corrupt_and_unknown_schema_are_explicit_failures(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - session = durable_session("broken") - await storage.save(session) - path = (await asyncio.to_thread(json_files, directory))[0] - - await asyncio.to_thread(write_text, path, "{not-json\n") - with self.assertRaises(SessionSerializationError): - await storage.load("broken") - - await asyncio.to_thread( - write_text, - path, - json.dumps({"journal_schema_version": 999}) + "\n", - ) - with self.assertRaisesRegex( - SessionSerializationError, - "unsupported.*999", - ): - await storage.load("broken") - - async def test_incomplete_tail_is_ignored_and_repaired_on_append(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - original = durable_session("partial-tail") - await storage.save(original) - path = (await asyncio.to_thread(json_files, directory))[0] - await asyncio.to_thread(append_bytes, path, b'{"partial":') - - restored = await storage.load("partial-tail") - self.assertEqual(restored, original) - - changed = original.snapshot() - changed.begin_run("run-2", "new task", 1) - changed.append_message( - "run-2", - 2, - {"role": "assistant", "content": "new answer"}, - ) - changed.finish_run("run-2", 3, completed_result("new answer")) - await storage.save(changed) - - self.assertEqual(await storage.load("partial-tail"), changed) - self.assertEqual(len(await storage.records("partial-tail")), 2) - - async def test_broken_tree_parent_is_rejected(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - session = durable_session("broken-parent") - await storage.save(session) - await storage.save(session) - path = (await asyncio.to_thread(json_files, directory))[0] - records = await asyncio.to_thread(read_json_lines, path) - records[1]["parent_id"] = "not-the-first-record" - await asyncio.to_thread(write_json_lines, path, records) - - with self.assertRaisesRegex( - SessionSerializationError, - "parent changed", - ): - await storage.load("broken-parent") - - async def test_non_json_message_is_rejected_without_creating_file(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - session = AgentSession(session_id="invalid-json") - session.bind_agent("durable-agent") - session.begin_run("run-1", "task", 1) - session.append_message( - "run-1", - 2, - {"role": "assistant", "content": {"not-json"}}, - ) - - with self.assertRaises(SessionSerializationError): - await storage.save(session) - self.assertEqual( - await asyncio.to_thread(directory_entries, directory), - [], - ) - - async def test_failed_append_preserves_previous_journal(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - original = durable_session("atomic") - await storage.save(original) - changed = original.snapshot() - changed.begin_run("run-2", "new task", 1) - changed.append_message( - "run-2", - 2, - {"role": "assistant", "content": "new answer"}, - ) - changed.finish_run("run-2", 3, completed_result("new answer")) - - with ( - patch( - "ejagent.session.jsonl.os.write", - side_effect=OSError("disk failure"), - ), - self.assertRaises(SessionStorageError), - ): - await storage.save(changed) - - restored = await storage.load("atomic") - self.assertEqual(restored, original) - - async def test_separate_python_process_loads_saved_session(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - await storage.save(durable_session("cross-process")) - script = """ -import asyncio -import json -import sys -from ejagent import JsonlSessionStorage - -async def main(): - session = await JsonlSessionStorage(sys.argv[1]).load(sys.argv[2]) - if session is None: - raise RuntimeError("missing session") - print(json.dumps({ - "agent_id": session.agent_id, - "roles": [message["role"] for message in session.messages], - "runs": len(session.runs), - }, ensure_ascii=False)) - -asyncio.run(main()) -""" - - process = await asyncio.to_thread( - subprocess.run, - [ - sys.executable, - "-c", - script, - directory, - "cross-process", - ], - check=True, - capture_output=True, - text=True, - ) - loaded = json.loads(process.stdout) - - self.assertEqual(loaded["agent_id"], "durable-agent") - self.assertEqual(loaded["roles"], ["system", "assistant"]) - self.assertEqual(loaded["runs"], 1) - - async def test_storage_implements_backend_neutral_tree_protocol(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - - self.assertIsInstance(storage, SessionTreeStorage) - - async def test_missing_read_does_not_create_storage_directory(self) -> None: - with tempfile.TemporaryDirectory() as directory: - root = Path(directory) / "sessions" - storage = JsonlSessionStorage(root) - - self.assertIsNone(await storage.load("missing")) - self.assertFalse(root.exists()) - - async def test_competing_processes_conditionally_append_only_once(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - await storage.save(durable_session("process-race")) - head = await storage.head("process-race") - assert head is not None - start_path = Path(directory) / "start" - processes = [ - start_append_process( - directory, - "process-race", - head.record_id, - f"competing-run-{index}", - "main", - start_path, - ) - for index in range(2) - ] - await asyncio.to_thread(write_text, start_path, "go") - - results = await asyncio.gather( - *(collect_process(process) for process in processes) - ) - - self.assertEqual( - sorted(result["status"] for result in results), - ["conflict", "success"], - ) - records = await storage.records("process-race") - self.assertEqual([record.revision for record in records], [1, 2]) - self.assertEqual(records[1].parent_id, head.record_id) - - async def test_processes_append_different_branches_with_correct_parents( - self, - ) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - await storage.save(durable_session("branch-race")) - main_head = await storage.head("branch-race") - assert main_head is not None - branch = await storage.fork( - "branch-race", - branch_id="experiment", - ) - start_path = Path(directory) / "start" - processes = [ - start_append_process( - directory, - "branch-race", - main_head.record_id, - "main-run", - "main", - start_path, - ), - start_append_process( - directory, - "branch-race", - branch.head.record_id, - "branch-run", - "experiment", - start_path, - ), - ] - await asyncio.to_thread(write_text, start_path, "go") - - results = await asyncio.gather( - *(collect_process(process) for process in processes) - ) - - self.assertEqual( - [result["status"] for result in results], - ["success", "success"], - ) - records = await storage.records("branch-race") - appended = { - record.data.get("run_id"): record - for record in records - if record.kind is SessionRecordKind.RUN_STARTED - } - self.assertEqual(appended["main-run"].parent_id, main_head.record_id) - self.assertEqual( - appended["branch-run"].parent_id, - branch.head.record_id, - ) - self.assertEqual( - [record.revision for record in records], - list(range(1, len(records) + 1)), - ) - - async def test_lock_timeout_is_explicit_and_wait_is_cancellable(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - await storage.save(durable_session("held-lock")) - lock_path = (await asyncio.to_thread(jsonl_lock_files, directory))[0] - release_path = Path(directory) / "release" - script = """ -import fcntl -import os -import sys -import time -from pathlib import Path - -descriptor = os.open(sys.argv[1], os.O_RDWR | os.O_CREAT, 0o600) -fcntl.flock(descriptor, fcntl.LOCK_EX) -print("locked", flush=True) -while not Path(sys.argv[2]).exists(): - time.sleep(0.005) -fcntl.flock(descriptor, fcntl.LOCK_UN) -os.close(descriptor) -""" - holder = await asyncio.to_thread( - start_python_process, - script, - str(lock_path), - str(release_path), - ) - assert holder.stdout is not None - self.assertEqual( - await asyncio.to_thread(holder.stdout.readline), - "locked\n", - ) - try: - with self.assertRaises(SessionLockTimeoutError): - await JsonlSessionStorage( - directory, - lock_timeout=0.05, - ).load("held-lock") - - waiting = asyncio.create_task( - JsonlSessionStorage( - directory, - lock_timeout=None, - ).load("held-lock") - ) - await asyncio.sleep(0.03) - waiting.cancel() - with self.assertRaises(asyncio.CancelledError): - await asyncio.wait_for(waiting, timeout=1) - finally: - await asyncio.to_thread(write_text, release_path, "release") - _, stderr = await asyncio.to_thread(holder.communicate, timeout=10) - self.assertEqual(holder.returncode, 0, stderr) - - async def test_crashed_lock_holder_releases_and_tail_is_repaired(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - original = durable_session("crashed-writer") - await storage.save(original) - journal_path = (await asyncio.to_thread(json_files, directory))[0] - lock_path = (await asyncio.to_thread(jsonl_lock_files, directory))[0] - script = """ -import fcntl -import os -import sys - -lock_descriptor = os.open(sys.argv[1], os.O_RDWR | os.O_CREAT, 0o600) -fcntl.flock(lock_descriptor, fcntl.LOCK_EX) -journal_descriptor = os.open(sys.argv[2], os.O_WRONLY | os.O_APPEND) -os.write(journal_descriptor, b'{"partial":') -os.fsync(journal_descriptor) -os._exit(23) -""" - crashed = await asyncio.to_thread( - subprocess.run, - [ - sys.executable, - "-c", - script, - str(lock_path), - str(journal_path), - ], - check=False, - ) - self.assertEqual(crashed.returncode, 23) - self.assertEqual(await storage.load("crashed-writer"), original) - head = await storage.head("crashed-writer") - assert head is not None - appended = await storage.append( - SessionRecordDraft.run_started( - session_id="crashed-writer", - agent_id="durable-agent", - sequence=1, - run_id="after-crash", - task="continue", - ), - expected_head_id=head.record_id, - check_head=True, - ) - - self.assertEqual(appended.revision, 2) - self.assertEqual(len(await storage.records("crashed-writer")), 2) - await asyncio.to_thread(read_json_lines, journal_path) - - async def test_restored_agent_continues_and_updates_durable_session(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - recorder = SessionRecorder(session_id="resume", storage=storage) - first_agent = BaseAgent( - SequenceModel(["saved answer"]), - agent_id="durable-agent", - event_sink=recorder, - ) - await first_agent.run(task="saved task") - - saved = await JsonlSessionStorage(directory).load("resume") - assert saved is not None - resumed_model = SequenceModel(["continued answer"]) - resumed_agent = BaseAgent( - resumed_model, - agent_id="durable-agent", - event_sink=SessionRecorder( - session_id="resume", - storage=JsonlSessionStorage(directory), - ), - ) - resumed_agent.restore_session(saved) - - result = await resumed_agent.run(task="continue") - updated = await storage.load("resume") - - self.assertEqual(result.output, "continued answer") - assert updated is not None - self.assertEqual(len(updated.runs), 2) - records = await storage.records("resume") - self.assertEqual( - [record.kind for record in records], - [ - SessionRecordKind.RUN_STARTED, - SessionRecordKind.MESSAGE_APPENDED, - SessionRecordKind.RUN_FINISHED, - SessionRecordKind.RUN_STARTED, - SessionRecordKind.MESSAGE_APPENDED, - SessionRecordKind.RUN_FINISHED, - ], - ) - self.assertEqual( - [record.revision for record in records], - list(range(1, 7)), - ) - self.assertEqual( - [record.parent_id for record in records[1:]], - [record.record_id for record in records[:-1]], - ) - self.assertTrue(all(record.branch_id == "main" for record in records)) - self.assertEqual( - [ - message.get("content") - for message in resumed_model.contexts[0].agent_messages - ], - [ - resumed_agent.system_prompt, - "saved task", - "saved answer", - "continue", - ], - ) - - async def test_fork_continues_without_changing_main_branch(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - first_agent = BaseAgent( - SequenceModel(["main answer"]), - agent_id="durable-agent", - event_sink=SessionRecorder(session_id="tree", storage=storage), - ) - await first_agent.run(task="main task") - main_head = await storage.head("tree") - assert main_head is not None - - forked = await storage.fork("tree", branch_id="experiment") - self.assertEqual(forked.branch.intent, SessionBranchIntent.FORK) - self.assertEqual(forked.branch.base_record_id, main_head.record_id) - self.assertEqual(forked.head.parent_id, main_head.record_id) - self.assertEqual(len(forked.session.runs), 1) - - branch_agent = BaseAgent( - SequenceModel(["branch answer"]), - agent_id="durable-agent", - event_sink=SessionRecorder( - session_id="tree", - storage=storage, - branch_id="experiment", - ), - ) - branch_agent.restore_session(forked.session) - await branch_agent.run(task="branch task") - - main = await storage.load("tree") - experiment = await storage.checkout("tree", branch_id="experiment") - assert main is not None - assert experiment is not None - self.assertEqual([run.task for run in main.runs], ["main task"]) - self.assertEqual( - [run.task for run in experiment.session.runs], - ["main task", "branch task"], - ) - branches = await storage.list_branches("tree") - self.assertEqual( - [branch.branch_id for branch in branches], - ["main", "experiment"], - ) - records = await storage.records("tree") - self.assertEqual( - [record.revision for record in records], - list(range(1, len(records) + 1)), - ) - - async def test_rollback_requires_ancestor_and_preserves_source(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - recorder = SessionRecorder(session_id="rollback", storage=storage) - agent = BaseAgent( - SequenceModel(["one", "two"]), - agent_id="durable-agent", - event_sink=recorder, - ) - await agent.run(task="first") - first_finish = (await storage.records("rollback"))[-1] - await agent.run(task="second") - - rolled_back = await storage.rollback( - "rollback", - to_record_id=first_finish.record_id, - branch_id="before-second", - ) - self.assertEqual(rolled_back.branch.intent, SessionBranchIntent.ROLLBACK) - self.assertEqual([run.task for run in rolled_back.session.runs], ["first"]) - main = await storage.load("rollback") - assert main is not None - self.assertEqual([run.task for run in main.runs], ["first", "second"]) - - unrelated = await storage.fork( - "rollback", - branch_id="unrelated", - ) - with self.assertRaisesRegex(ValueError, "not an ancestor"): - await storage.rollback( - "rollback", - to_record_id=unrelated.head.record_id, - source_branch="main", - branch_id="invalid-rollback", - ) - - async def test_prepare_retry_reuses_task_from_before_run(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - agent = BaseAgent( - SequenceModel(["first result", "original result"]), - agent_id="durable-agent", - event_sink=SessionRecorder(session_id="retry", storage=storage), - ) - await agent.run(task="first task") - await agent.run(task="repeat me") - main = await storage.load("retry") - assert main is not None - retried_run_id = main.runs[-1].run_id - - retry = await storage.prepare_retry( - "retry", - run_id=retried_run_id, - branch_id="retry-second", - ) - self.assertEqual(retry.task, "repeat me") - self.assertEqual(retry.checkout.branch.intent, SessionBranchIntent.RETRY) - self.assertEqual( - [run.task for run in retry.checkout.session.runs], - ["first task"], - ) - - retry_agent = BaseAgent( - SequenceModel(["new result"]), - agent_id="durable-agent", - event_sink=SessionRecorder( - session_id="retry", - storage=storage, - branch_id="retry-second", - ), - ) - retry_agent.restore_session(retry.checkout.session) - await retry_agent.run(task=retry.task) - - retry_checkout = await storage.checkout( - "retry", - branch_id="retry-second", - ) - assert retry_checkout is not None - self.assertEqual( - [run.result.output for run in retry_checkout.session.runs], - ["first result", "new result"], - ) - unchanged_main = await storage.load("retry") - assert unchanged_main is not None - self.assertEqual(unchanged_main.runs[-1].result.output, "original result") - - async def test_first_run_retry_starts_from_bound_empty_session(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - agent = BaseAgent( - SequenceModel(["original"]), - agent_id="durable-agent", - event_sink=SessionRecorder(session_id="retry-first", storage=storage), - ) - await agent.run(task="first task") - main = await storage.load("retry-first") - assert main is not None - - retry = await storage.prepare_retry( - "retry-first", - run_id=main.runs[0].run_id, - branch_id="retry-root", - ) - - self.assertEqual(retry.task, "first task") - self.assertEqual(retry.checkout.session.agent_id, "durable-agent") - self.assertEqual(retry.checkout.session.runs, []) - self.assertIsNone(retry.checkout.head.parent_id) - - async def test_conditional_append_rejects_stale_branch_head(self) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - session = durable_session("conflict") - await storage.save(session) - head = await storage.head("conflict") - assert head is not None - draft = SessionRecordDraft.run_started( - session_id="conflict", - agent_id="durable-agent", - sequence=1, - run_id="run-2", - task="new task", - ) - - with self.assertRaises(SessionConflictError): - await storage.append( - draft, - expected_head_id="stale-head", - check_head=True, - ) - self.assertEqual((await storage.head("conflict")), head) - - async def test_checkout_allows_audit_of_unfinished_node_but_fork_rejects_it( - self, - ) -> None: - with tempfile.TemporaryDirectory() as directory: - storage = JsonlSessionStorage(directory) - agent = BaseAgent( - SequenceModel(["answer"]), - agent_id="durable-agent", - event_sink=SessionRecorder(session_id="audit-node", storage=storage), - ) - await agent.run(task="task") - started = (await storage.records("audit-node"))[0] - - audited = await storage.checkout( - "audit-node", - record_id=started.record_id, - ) - assert audited is not None - self.assertFalse(audited.session.runs[0].finished) - with self.assertRaisesRegex(ValueError, "unfinished run"): - await storage.fork( - "audit-node", - from_record_id=started.record_id, - branch_id="unsafe", - ) - - def test_restore_rejects_wrong_agent_and_unfinished_run(self) -> None: - wrong_agent = BaseAgent( - SequenceModel(["unused"]), - agent_id="other-agent", - ) - with self.assertRaisesRegex(ValueError, "belongs to agent"): - wrong_agent.restore_session(durable_session()) - - unfinished = AgentSession(session_id="unfinished") - unfinished.bind_agent("durable-agent") - unfinished.begin_run("run-open", "task", 1) - matching_agent = BaseAgent( - SequenceModel(["unused"]), - agent_id="durable-agent", - ) - with self.assertRaisesRegex(ValueError, "unfinished"): - matching_agent.restore_session(unfinished) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_session_store_contract.py b/tests/test_session_store_contract.py new file mode 100644 index 0000000..70b448e --- /dev/null +++ b/tests/test_session_store_contract.py @@ -0,0 +1,542 @@ +from __future__ import annotations + +import asyncio +import json +import tempfile +import unittest +from collections.abc import Callable +from datetime import UTC, datetime +from hashlib import sha256 +from pathlib import Path +from typing import Any + +from ejagent.contracts import ( + AssistantMessage, + AuditReader, + AuditRecord, + ConversationSnapshot, + FailureCode, + RunDelta, + RunFailure, + RunOutcome, + RunPhase, + RunResult, + RunStatus, + RunUsage, + SessionCommit, + SessionConflictError, + SessionMigrationError, + SessionSnapshot, + SessionStore, + SessionStoreSerializationError, + StopReason, + SystemMessage, + ToolCall, + ToolResultMessage, + UserMessage, +) +from ejagent.harness import MemorySessionStore +from ejagent.storage import JsonlSessionStore +from ejagent.storage._legacy import ( + LegacyRun, + LegacyRunResult, + LegacySessionData, +) +from ejagent.storage.codec import ( + session_commit_from_dict, + session_commit_to_dict, +) +from ejagent.storage.migration import migrate_legacy_session + +_NOW = datetime(2026, 8, 1, 8, 30, tzinfo=UTC) + + +def _legacy_result(*, output: str | None = None) -> LegacyRunResult: + return LegacyRunResult( + status=RunStatus.COMPLETED, + stop_reason=StopReason.TEXT_RESPONSE, + turns=2, + output=output, + error=None, + usage=RunUsage(), + ) + + +def _write_legacy_checkpoint( + root: Path, + *, + session_id: str, + agent_id: str, + messages: list[dict[str, Any]], + run_id: str, + output: str | None, +) -> None: + usage = RunUsage().to_dict() + document = { + "schema_version": 1, + "session": { + "session_id": session_id, + "agent_id": agent_id, + "entries": [ + {"run_id": run_id, "sequence": index + 1, "message": message} + for index, message in enumerate(messages) + ], + "runs": [ + { + "run_id": run_id, + "task": messages[0]["content"], + "intent": "task", + "start_sequence": 1, + "finish_sequence": len(messages) + 1, + "result": { + "status": "completed", + "stop_reason": "text_response", + "turns": 2, + "output": output, + "error": None, + "usage": usage, + }, + } + ], + "compactions": [], + }, + } + record = { + "journal_schema_version": 1, + "record_id": "legacy-checkpoint", + "parent_id": None, + "branch_id": "main", + "revision": 1, + "session_id": session_id, + "agent_id": agent_id, + "sequence": 0, + "type": "checkpoint", + "data": {"document": document}, + } + digest = sha256(session_id.encode("utf-8")).hexdigest() + (root / f"{digest}.jsonl").write_text( + json.dumps(record, ensure_ascii=False) + "\n", + encoding="utf-8", + ) + + +def _write_legacy_event_journal(root: Path) -> None: + usage = RunUsage().to_dict() + drafts = ( + ("run_started", 1, {"run_id": "event-run", "task": "remember blue"}), + ( + "message_appended", + 2, + { + "run_id": "event-run", + "message": {"role": "assistant", "content": "noted"}, + }, + ), + ( + "run_finished", + 3, + { + "run_id": "event-run", + "result": { + "status": "completed", + "stop_reason": "text_response", + "turns": 1, + "output": "noted", + "error": None, + "usage": usage, + }, + }, + ), + ) + records: list[dict[str, Any]] = [] + parent_id: str | None = None + for revision, (kind, sequence, data) in enumerate(drafts, start=1): + record_id = f"record-{revision}" + records.append( + { + "journal_schema_version": 1, + "record_id": record_id, + "parent_id": parent_id, + "branch_id": "main", + "revision": revision, + "session_id": "event-session", + "agent_id": "event-agent", + "sequence": sequence, + "type": kind, + "data": data, + } + ) + parent_id = record_id + digest = sha256(b"event-session").hexdigest() + (root / f"{digest}.jsonl").write_text( + "".join(json.dumps(record) + "\n" for record in records), + encoding="utf-8", + ) + + +class _ContractStore(SessionStore, AuditReader): + pass + + +def _success( + run_id: str, + *, + agent_id: str = "agent-1", + base: ConversationSnapshot | None = None, + output: str = "done", +) -> SessionCommit: + current = base or ConversationSnapshot() + result = RunResult( + run_id=run_id, + status=RunStatus.COMPLETED, + stop_reason=StopReason.TEXT_RESPONSE, + turns=1, + output=output, + ) + return SessionCommit( + agent_id=agent_id, + base=current, + outcome=RunOutcome( + result=result, + delta=RunDelta( + base_revision=current.revision, + messages=(UserMessage(f"task:{run_id}"), AssistantMessage(output)), + ), + audit_records=( + AuditRecord( + run_id=run_id, + sequence=1, + kind="model.completed", + occurred_at=_NOW, + payload={"unicode": "你好"}, + ), + ), + ), + ) + + +def _failure( + run_id: str, + *, + agent_id: str = "agent-1", + base: ConversationSnapshot | None = None, +) -> SessionCommit: + current = base or ConversationSnapshot() + failure = RunFailure( + phase=RunPhase.MODEL, + code=FailureCode.RATE_LIMIT, + message="slow down", + retryable=True, + ) + return SessionCommit( + agent_id=agent_id, + base=current, + outcome=RunOutcome( + result=RunResult( + run_id=run_id, + status=RunStatus.FAILED, + stop_reason=StopReason.RUNTIME_ERROR, + turns=0, + ), + delta=RunDelta(base_revision=current.revision), + failure=failure, + ), + ) + + +class _SessionStoreContract: + store_factory: Callable[[], _ContractStore] + store: _ContractStore + + async def asyncSetUp(self) -> None: + self.store = self.store_factory() + + async def test_commit_load_and_audit(self) -> None: + commit = _success("run-1") + + snapshot = await self.store.commit(commit) + + self.assertEqual(snapshot.revision, 1) + self.assertEqual(snapshot.messages, commit.resulting_conversation.messages) + self.assertEqual(snapshot.last_result, commit.outcome.result) + self.assertEqual(await self.store.load("agent-1"), snapshot) + self.assertEqual(await self.store.load_audit("agent-1"), (commit.audit,)) + + async def test_failed_run_is_audited_without_advancing(self) -> None: + commit = _failure("run-failed") + + snapshot = await self.store.commit(commit) + + self.assertEqual(snapshot.revision, 0) + self.assertEqual(snapshot.messages, ()) + self.assertIsNone(snapshot.last_result) + self.assertEqual(await self.store.load_audit("agent-1"), (commit.audit,)) + + async def test_repeating_same_commit_is_idempotent(self) -> None: + commit = _success("run-idempotent") + + first = await self.store.commit(commit) + second = await self.store.commit(commit) + + self.assertEqual(second, first) + self.assertEqual(len(await self.store.load_audit("agent-1")), 1) + + async def test_reusing_run_id_for_different_commit_conflicts(self) -> None: + await self.store.commit(_success("run-reused", output="first")) + + with self.assertRaises(SessionConflictError): + await self.store.commit(_success("run-reused", output="different")) + + async def test_only_one_stale_base_commit_wins(self) -> None: + commits = (_success("run-a"), _success("run-b")) + + results = await asyncio.gather( + *(self.store.commit(commit) for commit in commits), + return_exceptions=True, + ) + + self.assertEqual(sum(isinstance(item, SessionSnapshot) for item in results), 1) + self.assertEqual( + sum(isinstance(item, SessionConflictError) for item in results), + 1, + ) + + +class MemorySessionStoreContractTests( + _SessionStoreContract, + unittest.IsolatedAsyncioTestCase, +): + store_factory = MemorySessionStore + + +class JsonlSessionStoreContractTests( + _SessionStoreContract, + unittest.IsolatedAsyncioTestCase, +): + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + root = Path(temporary.name) + self.store_factory = lambda: JsonlSessionStore(root) + + +class JsonlSessionStoreTests(unittest.IsolatedAsyncioTestCase): + def setUp(self) -> None: + temporary = tempfile.TemporaryDirectory() + self.addCleanup(temporary.cleanup) + self.root = Path(temporary.name) + + async def test_round_trips_rich_typed_commit_across_instances(self) -> None: + base = ConversationSnapshot( + messages=(SystemMessage("answer briefly"),), + ) + result = RunResult( + run_id="rich-run", + status=RunStatus.COMPLETED, + stop_reason=StopReason.TOOL_COMPLETION, + turns=2, + output="晴天", + ) + commit = SessionCommit( + agent_id="rich-agent", + base=base, + outcome=RunOutcome( + result=result, + delta=RunDelta( + base_revision=0, + messages=( + AssistantMessage( + tool_calls=( + ToolCall( + id="call-1", + name="weather", + arguments={"city": "杭州"}, + ), + ) + ), + ToolResultMessage( + tool_call_id="call-1", + tool_name="weather", + result={"temperature": 31}, + ), + AssistantMessage("晴天"), + ), + ), + ), + ) + self.assertEqual( + session_commit_from_dict(session_commit_to_dict(commit)), + commit, + ) + + expected = await JsonlSessionStore(self.root).commit(commit) + + restored = await JsonlSessionStore(self.root).load("rich-agent") + self.assertEqual(restored, expected) + + async def test_partial_tail_is_ignored_then_truncated_on_append(self) -> None: + store = JsonlSessionStore(self.root) + first = await store.commit(_success("run-1")) + path = next(self.root.glob("*.core.jsonl")) + with path.open("ab") as stream: + stream.write(b'{"partial":') + + restarted = JsonlSessionStore(self.root) + self.assertEqual(await restarted.load("agent-1"), first) + await restarted.commit( + _success("run-2", base=first.conversation, output="again") + ) + + for line in path.read_text(encoding="utf-8").splitlines(): + self.assertIsInstance(json.loads(line), dict) + self.assertEqual(len(path.read_text(encoding="utf-8").splitlines()), 2) + + async def test_partial_first_record_is_replaced_by_first_commit(self) -> None: + digest = sha256(b"agent-1").hexdigest() + path = self.root / f"{digest}.core.jsonl" + path.write_bytes(b'{"partial":') + + snapshot = await JsonlSessionStore(self.root).commit(_success("run-1")) + + self.assertEqual(snapshot.revision, 1) + self.assertEqual(len(path.read_text(encoding="utf-8").splitlines()), 1) + self.assertIsInstance(json.loads(path.read_text(encoding="utf-8")), dict) + + async def test_complete_corrupt_record_is_rejected(self) -> None: + store = JsonlSessionStore(self.root) + await store.commit(_success("run-1")) + path = next(self.root.glob("*.core.jsonl")) + with path.open("ab") as stream: + stream.write(b"not-json\n") + + with self.assertRaises(SessionStoreSerializationError): + await JsonlSessionStore(self.root).load("agent-1") + + async def test_separate_instances_compare_and_commit_atomically(self) -> None: + first_store = JsonlSessionStore(self.root) + second_store = JsonlSessionStore(self.root) + self.assertIsNone(await first_store.load("agent-1")) + self.assertIsNone(await second_store.load("agent-1")) + + results = await asyncio.gather( + first_store.commit(_success("run-a")), + second_store.commit(_success("run-b")), + return_exceptions=True, + ) + + self.assertEqual(sum(isinstance(item, SessionSnapshot) for item in results), 1) + self.assertEqual( + sum(isinstance(item, SessionConflictError) for item in results), + 1, + ) + + async def test_auto_migrates_legacy_jsonl_once(self) -> None: + _write_legacy_checkpoint( + self.root, + session_id="legacy-session", + agent_id="legacy-agent", + run_id="legacy-run", + output="已记住", + messages=[ + {"role": "user", "content": "记住 café"}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call-1", + "type": "function", + "function": { + "name": "remember", + "arguments": '{"text":"你好"}', + }, + } + ], + }, + {"role": "tool", "tool_call_id": "call-1", "content": "ok"}, + {"role": "assistant", "content": "已记住"}, + ], + ) + + migrating = JsonlSessionStore( + self.root, + legacy_session_id="legacy-session", + ) + snapshot = await migrating.load("legacy-agent") + + self.assertIsNotNone(snapshot) + assert snapshot is not None + self.assertEqual(snapshot.revision, 1) + self.assertEqual(snapshot.messages[0], UserMessage("记住 café")) + self.assertEqual( + snapshot.messages[1], + AssistantMessage( + tool_calls=( + ToolCall( + id="call-1", + name="remember", + arguments={"text": "你好"}, + ), + ) + ), + ) + self.assertEqual( + snapshot.messages[2], + ToolResultMessage( + tool_call_id="call-1", + tool_name="remember", + result="ok", + ), + ) + self.assertEqual(len(await migrating.load_audit("legacy-agent")), 1) + self.assertEqual( + await JsonlSessionStore(self.root).load("legacy-agent"), + snapshot, + ) + self.assertEqual(len(tuple(self.root.glob("*.core.jsonl"))), 1) + + async def test_migrates_event_journal_without_legacy_runtime(self) -> None: + _write_legacy_event_journal(self.root) + store = JsonlSessionStore( + self.root, + legacy_session_id="event-session", + ) + + snapshot = await store.load("event-agent") + + self.assertIsNotNone(snapshot) + assert snapshot is not None + self.assertEqual( + snapshot.messages, + (UserMessage("remember blue"), AssistantMessage("noted")), + ) + self.assertEqual(snapshot.revision, 1) + self.assertEqual(snapshot.last_result.output, "noted") # type: ignore[union-attr] + + async def test_migration_errors_are_actionable(self) -> None: + unfinished = LegacySessionData( + session_id="unfinished", + agent_id="agent-1", + entries=({"role": "user", "content": "work"},), + runs=(LegacyRun("run-open"),), + ) + + with self.assertRaises(SessionMigrationError) as raised: + migrate_legacy_session(unfinished) + + self.assertIn("unfinished", str(raised.exception)) + self.assertIn("Remediation:", str(raised.exception)) + self.assertTrue(raised.exception.remediation) + + unsupported = LegacySessionData( + session_id="unsupported", + agent_id="agent-1", + entries=( + {"role": "user", "content": "work"}, + {"role": "developer", "content": "hidden"}, + ), + runs=(LegacyRun("run-1", _legacy_result()),), + ) + with self.assertRaises(SessionMigrationError) as unsupported_error: + migrate_legacy_session(unsupported) + self.assertIn("unsupported", str(unsupported_error.exception)) diff --git a/tests/test_tool_definitions.py b/tests/test_tool_definitions.py deleted file mode 100644 index fcb26e7..0000000 --- a/tests/test_tool_definitions.py +++ /dev/null @@ -1,347 +0,0 @@ -import unittest -from types import SimpleNamespace -from typing import Any, cast - -from ejagent import ( - AgentContextBuilder, - AgentState, - AssistantMessage, - BaseAgent, - CancellationToken, - ContextBuildResult, - MethodToolHandler, - ModelAdapter, - ModelConfig, - OpenAIModelAdapter, - StepOutcome, - ToolDefinition, - ToolDefinitionError, - ToolEffect, -) - -LEGACY_WRITE_TOOL = { - "type": "function", - "function": { - "name": "write_value", - "description": "Write one value.", - "parameters": { - "type": "object", - "properties": {"value": {"type": "string"}}, - "required": ["value"], - }, - }, -} - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - return self.responses.pop(0) - - -class MixedDefinitionHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__( - ( - ToolDefinition( - name="lookup", - description="Look up one value.", - parameters={ - "type": "object", - "properties": {"key": {"type": "string"}}, - "required": ["key"], - }, - effect=ToolEffect.READ_ONLY, - strict=True, - ), - LEGACY_WRITE_TOOL, - ) - ) - - async def do_lookup( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - return StepOutcome({"key": arguments["key"]}) - - async def do_write_value( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - return StepOutcome({"value": arguments["value"]}) - - -class FakeCompletions: - def __init__(self) -> None: - self.calls: list[dict[str, Any]] = [] - - async def create(self, **kwargs: Any) -> Any: - self.calls.append(kwargs) - return SimpleNamespace( - choices=[ - SimpleNamespace(message=SimpleNamespace(content="done", tool_calls=[])) - ] - ) - - -class FakeClient: - def __init__(self) -> None: - self.chat = SimpleNamespace(completions=FakeCompletions()) - - -class ToolDefinitionTests(unittest.IsolatedAsyncioTestCase): - def test_definition_serializes_openai_fields_but_not_core_effect(self) -> None: - parameters = { - "type": "object", - "properties": {"query": {"type": "string"}}, - } - definition = ToolDefinition( - name="search", - description="Search documents.", - parameters=parameters, - effect=ToolEffect.READ_ONLY, - strict=True, - ) - parameters["properties"]["query"]["type"] = "integer" - - serialized = definition.to_openai_tool() - - self.assertEqual( - serialized, - { - "type": "function", - "function": { - "name": "search", - "description": "Search documents.", - "parameters": { - "type": "object", - "properties": {"query": {"type": "string"}}, - }, - "strict": True, - }, - }, - ) - self.assertNotIn("effect", serialized["function"]) - serialized["function"]["parameters"]["properties"]["query"]["type"] = "boolean" - self.assertEqual( - definition.to_openai_tool()["function"]["parameters"]["properties"][ - "query" - ]["type"], - "string", - ) - - def test_legacy_dictionary_round_trips_supported_openai_fields(self) -> None: - legacy = { - "type": "function", - "function": { - "name": "lookup", - "description": "Lookup.", - "parameters": {"type": "object"}, - "strict": False, - }, - } - - definition = ToolDefinition.from_openai_tool( - legacy, - effect=ToolEffect.READ_ONLY, - ) - - self.assertEqual(definition.name, "lookup") - self.assertEqual(definition.effect, ToolEffect.READ_ONLY) - self.assertEqual(definition.to_openai_tool(), legacy) - - def test_optional_openai_fields_remain_omitted(self) -> None: - definition = ToolDefinition.from_openai_tool( - {"type": "function", "function": {"name": "minimal"}} - ) - - self.assertEqual( - definition.to_openai_tool(), - {"type": "function", "function": {"name": "minimal"}}, - ) - - def test_definition_rejects_invalid_core_and_openai_values(self) -> None: - cases = [ - lambda: ToolDefinition(name=""), - lambda: ToolDefinition(name="invalid name"), - lambda: ToolDefinition(name="x" * 65), - lambda: ToolDefinition(name="tool", description=1), - lambda: ToolDefinition(name="tool", parameters=[]), - lambda: ToolDefinition(name="tool", parameters={"bad": {1, 2}}), - lambda: ToolDefinition(name="tool", parameters={"bad": float("nan")}), - lambda: ToolDefinition(name="tool", effect="read_only"), - lambda: ToolDefinition(name="tool", strict="yes"), - lambda: ToolDefinition.from_openai_tool({"type": "custom"}), - lambda: ToolDefinition.from_openai_tool( - {"type": "function", "function": []} - ), - lambda: ToolDefinition.from_openai_tool( - {"type": "function", "function": {}} - ), - ] - for create in cases: - with self.subTest(create=create), self.assertRaises(ToolDefinitionError): - create() - - def test_parameters_are_deeply_immutable(self) -> None: - definition = ToolDefinition( - name="immutable", - parameters={ - "type": "object", - "properties": {"value": {"type": "string"}}, - }, - ) - assert definition.parameters is not None - - with self.assertRaises(TypeError): - definition.parameters["type"] = "array" # type: ignore[index] - properties = definition.parameters["properties"] - with self.assertRaises(TypeError): - properties["value"]["type"] = "number" - - def test_method_handler_accepts_canonical_and_legacy_definitions(self) -> None: - handler = MixedDefinitionHandler() - - self.assertEqual(handler.tool_names, ("lookup", "write_value")) - self.assertTrue( - all(isinstance(tool, ToolDefinition) for tool in handler.tool_definitions) - ) - self.assertEqual( - [tool.effect for tool in handler.tool_definitions], - [ToolEffect.READ_ONLY, ToolEffect.SIDE_EFFECTING], - ) - self.assertTrue(handler.tools[0]["function"]["strict"]) - - detached = handler.tools - detached[0]["function"]["name"] = "mutated" - self.assertEqual(handler.tool_definitions[0].name, "lookup") - - def test_legacy_effect_mapping_remains_compatible(self) -> None: - handler = MethodToolHandler( - (LEGACY_WRITE_TOOL,), - tool_effects={"write_value": ToolEffect.READ_ONLY}, - ) - - self.assertEqual( - handler.tool_definitions[0].effect, - ToolEffect.READ_ONLY, - ) - - def test_mixed_duplicate_names_are_rejected(self) -> None: - with self.assertRaisesRegex(ToolDefinitionError, "duplicate tool names"): - MethodToolHandler( - ( - ToolDefinition(name="duplicate"), - { - "type": "function", - "function": {"name": "duplicate"}, - }, - ) - ) - - def test_context_builder_keeps_canonical_source_and_legacy_view(self) -> None: - definition = ToolDefinition( - name="lookup", - parameters={"type": "object"}, - effect=ToolEffect.READ_ONLY, - ) - - context = AgentContextBuilder().build( - AgentState(messages=[{"role": "user", "content": "hello"}]), - tools=(definition,), - ) - - self.assertEqual(context.tool_definitions, (definition,)) - self.assertEqual( - context.tools, - ( - { - "type": "function", - "function": { - "name": "lookup", - "parameters": {"type": "object"}, - }, - }, - ), - ) - - async def test_agent_routes_canonical_definition_and_exposes_legacy_view( - self, - ) -> None: - handler = MixedDefinitionHandler() - agent = BaseAgent( - SequenceModel([AssistantMessage(content="done")]), - agent_id="canonical-agent", - handlers=[handler], - ) - - result = await agent.run(task="inspect tools") - - self.assertEqual(result.output, "done") - self.assertEqual( - [tool.name for tool in agent.tool_definitions], - ["lookup", "write_value"], - ) - self.assertEqual( - [tool["function"]["name"] for tool in agent.tools], - ["lookup", "write_value"], - ) - - async def test_openai_adapter_prefers_canonical_context_definitions(self) -> None: - canonical = ToolDefinition( - name="canonical", - parameters={"type": "object"}, - strict=True, - ) - context = ContextBuildResult( - agent_messages=({"role": "user", "content": "hello"},), - llm_messages=({"role": "user", "content": "hello"},), - tools=( - { - "type": "function", - "function": {"name": "legacy-view"}, - }, - ), - tool_definitions=(canonical,), - ) - client = FakeClient() - adapter = OpenAIModelAdapter( - ModelConfig( - model="test", - api_key="test", - base_url="https://example.invalid", - ), - client=cast(Any, client), - ) - - response = await adapter.complete(context) - - self.assertEqual(response.content, "done") - self.assertEqual( - client.chat.completions.calls[0]["tools"], - ( - { - "type": "function", - "function": { - "name": "canonical", - "parameters": {"type": "object"}, - "strict": True, - }, - }, - ), - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_tool_policy.py b/tests/test_tool_policy.py deleted file mode 100644 index 1259166..0000000 --- a/tests/test_tool_policy.py +++ /dev/null @@ -1,480 +0,0 @@ -from __future__ import annotations - -import asyncio -import json -import unittest -from typing import Any - -from ejagent import ( - AssistantMessage, - BaseAgent, - CancellationToken, - MethodToolHandler, - ModelAdapter, - ModelToolCall, - RuleBasedToolPolicy, - RunStatus, - RuntimePolicy, - StepOutcome, - StopReason, - ToolApprovalDecision, - ToolApprovalRequest, - ToolCallContext, - ToolDefinition, - ToolEffect, - ToolExecutionPolicy, - ToolPolicyAction, - ToolPolicyDecision, - ToolPolicyMiddleware, - ToolPolicyRule, -) -from ejagent.agent.state import AgentState - -READ_TOOL = ToolDefinition( - name="read", - parameters={ - "type": "object", - "properties": {"path": {"type": "string"}}, - "required": ["path"], - }, - effect=ToolEffect.READ_ONLY, -) -WRITE_TOOL = ToolDefinition( - name="write", - parameters={ - "type": "object", - "properties": {"path": {"type": "string"}}, - "required": ["path"], - }, -) - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - return self.responses.pop(0) - - -class RecordingHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((READ_TOOL, WRITE_TOOL)) - self.calls: list[tuple[str, dict[str, Any]]] = [] - - async def do_read( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken, - ) -> StepOutcome: - self.calls.append(("read", arguments)) - return StepOutcome({"status": "success", "tool": "read"}) - - async def do_write( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken, - ) -> StepOutcome: - self.calls.append(("write", arguments)) - return StepOutcome({"status": "success", "tool": "write"}) - - -class RecordingApprover: - def __init__(self, decision: ToolApprovalDecision) -> None: - self.decision = decision - self.requests: list[ToolApprovalRequest] = [] - - async def approve(self, request: ToolApprovalRequest) -> ToolApprovalDecision: - self.requests.append(request) - return self.decision - - -class RaisingApprover: - async def approve(self, request: ToolApprovalRequest) -> ToolApprovalDecision: - raise RuntimeError("approval service unavailable") - - -class LifecyclePolicy(ToolExecutionPolicy): - def __init__(self, decision: ToolPolicyDecision) -> None: - self.decision = decision - self.started = 0 - self.stopped = 0 - self.task_starts = 0 - self.contexts: list[ToolCallContext] = [] - - async def startup(self) -> None: - self.started += 1 - - async def shutdown(self) -> None: - self.stopped += 1 - - async def on_task_start(self) -> None: - self.task_starts += 1 - - async def evaluate(self, context: ToolCallContext) -> ToolPolicyDecision: - self.contexts.append(context) - return self.decision - - -class RaisingPolicy(ToolExecutionPolicy): - async def evaluate(self, context: ToolCallContext) -> ToolPolicyDecision: - raise RuntimeError("policy backend unavailable") - - -def call(call_id: str, name: str, path: str) -> ModelToolCall: - return ModelToolCall( - id=call_id, - name=name, - arguments=json.dumps({"path": path}), - ) - - -def context( - name: str = "read", - *, - effect: ToolEffect = ToolEffect.READ_ONLY, -) -> ToolCallContext: - state = AgentState() - state.reset([]) - return ToolCallContext( - state=state, - tool_name=name, - arguments={"path": "/workspace/file.txt"}, - tool_call_id="call-1", - tool_definition=ToolDefinition(name=name, effect=effect), - ) - - -class ToolPolicyTests(unittest.IsolatedAsyncioTestCase): - async def test_default_deny_short_circuits_handler_and_agent_run(self) -> None: - handler = RecordingHandler() - policy = RuleBasedToolPolicy(()) - agent = BaseAgent( - SequenceModel( - [AssistantMessage(tool_calls=(call("one", "write", "/tmp/a"),))] - ), - agent_id="policy-default-deny", - handlers=[handler], - middlewares=[ToolPolicyMiddleware(policy)], - ) - - result = await agent.run(task="write a file") - - self.assertEqual(result.status, RunStatus.REJECTED) - self.assertEqual(result.stop_reason, StopReason.TOOL_REJECTED) - self.assertEqual(handler.calls, []) - payload = json.loads(result.output or "") - self.assertEqual(payload["status"], "rejected") - self.assertEqual(payload["tool"], "write") - self.assertNotIn("arguments", payload) - - async def test_effect_rule_uses_effective_canonical_definition(self) -> None: - handler = RecordingHandler() - policy = LifecyclePolicy(ToolPolicyDecision(ToolPolicyAction.ALLOW)) - agent = BaseAgent( - SequenceModel([]), - agent_id="policy-definition", - handlers=[handler], - middlewares=[ToolPolicyMiddleware(policy)], - ) - - outcome = await agent.dispatch("read", {"path": "/workspace/a"}) - - self.assertEqual(outcome.data["status"], "success") - definition = policy.contexts[0].tool_definition - self.assertIsNotNone(definition) - assert definition is not None - self.assertEqual(definition.name, "read") - self.assertIs(definition.effect, ToolEffect.READ_ONLY) - - async def test_ordered_predicate_rule_denies_matching_arguments(self) -> None: - policy = RuleBasedToolPolicy( - ( - ToolPolicyRule( - rule_id="deny-private", - action=ToolPolicyAction.DENY, - tool_names=frozenset({"read"}), - when=lambda item: str(item.arguments["path"]).startswith( - "/private" - ), - reason="private paths are forbidden", - ), - ToolPolicyRule( - rule_id="allow-read", - action=ToolPolicyAction.ALLOW, - tool_names=frozenset({"read"}), - ), - ) - ) - - denied = await policy.evaluate( - ToolCallContext( - state=AgentState(), - tool_name="read", - arguments={"path": "/private/secret"}, - tool_definition=READ_TOOL, - ) - ) - allowed = await policy.evaluate(context()) - - self.assertIs(denied.action, ToolPolicyAction.DENY) - self.assertEqual(denied.rule_id, "deny-private") - self.assertIs(allowed.action, ToolPolicyAction.ALLOW) - self.assertEqual(allowed.rule_id, "allow-read") - - async def test_effect_selector_can_require_side_effect_approval(self) -> None: - policy = RuleBasedToolPolicy( - ( - ToolPolicyRule( - rule_id="approve-writes", - action=ToolPolicyAction.REQUIRE_APPROVAL, - effects=frozenset({ToolEffect.SIDE_EFFECTING}), - ), - ToolPolicyRule( - rule_id="allow-reads", - action=ToolPolicyAction.ALLOW, - effects=frozenset({ToolEffect.READ_ONLY}), - ), - ) - ) - - read = await policy.evaluate(context()) - write = await policy.evaluate( - context("write", effect=ToolEffect.SIDE_EFFECTING) - ) - - self.assertIs(read.action, ToolPolicyAction.ALLOW) - self.assertIs(write.action, ToolPolicyAction.REQUIRE_APPROVAL) - - async def test_required_approval_without_approver_fails_closed(self) -> None: - handler = RecordingHandler() - policy = LifecyclePolicy( - ToolPolicyDecision( - ToolPolicyAction.REQUIRE_APPROVAL, - rule_id="approve-write", - ) - ) - agent = BaseAgent( - SequenceModel([]), - agent_id="policy-no-approver", - handlers=[handler], - middlewares=[ToolPolicyMiddleware(policy)], - ) - - outcome = await agent.dispatch("write", {"path": "/tmp/a"}) - - self.assertEqual(outcome.control.value, "reject") - self.assertEqual(outcome.data["rule_id"], "approve-write") - self.assertEqual(handler.calls, []) - - async def test_approved_call_executes_and_request_is_immutable(self) -> None: - handler = RecordingHandler() - policy = LifecyclePolicy( - ToolPolicyDecision( - ToolPolicyAction.REQUIRE_APPROVAL, - reason="write changes external state", - rule_id="approve-write", - ) - ) - approver = RecordingApprover(ToolApprovalDecision(True, "approved")) - agent = BaseAgent( - SequenceModel([]), - agent_id="policy-approved", - handlers=[handler], - middlewares=[ToolPolicyMiddleware(policy, approver=approver)], - ) - - outcome = await agent.dispatch("write", {"path": "/tmp/a"}) - - self.assertEqual(outcome.data["status"], "success") - self.assertEqual(handler.calls, [("write", {"path": "/tmp/a"})]) - request = approver.requests[0] - self.assertEqual(request.tool_name, "write") - self.assertEqual(request.arguments["path"], "/tmp/a") - self.assertEqual(request.rule_id, "approve-write") - with self.assertRaises(TypeError): - request.arguments["path"] = "/tmp/changed" # type: ignore[index] - - async def test_denied_approval_uses_approver_reason(self) -> None: - handler = RecordingHandler() - policy = LifecyclePolicy( - ToolPolicyDecision( - ToolPolicyAction.REQUIRE_APPROVAL, - reason="approval required", - rule_id="approve-write", - ) - ) - approver = RecordingApprover( - ToolApprovalDecision(False, "operator denied the write") - ) - agent = BaseAgent( - SequenceModel([]), - agent_id="policy-approval-denied", - handlers=[handler], - middlewares=[ToolPolicyMiddleware(policy, approver=approver)], - ) - - outcome = await agent.dispatch("write", {"path": "/tmp/a"}) - - self.assertEqual(outcome.control.value, "reject") - self.assertEqual(outcome.data["reason"], "operator denied the write") - self.assertEqual(handler.calls, []) - - async def test_approval_exception_fails_closed(self) -> None: - handler = RecordingHandler() - policy = LifecyclePolicy( - ToolPolicyDecision( - ToolPolicyAction.REQUIRE_APPROVAL, - rule_id="approve-write", - ) - ) - agent = BaseAgent( - SequenceModel([]), - agent_id="policy-approval-error", - handlers=[handler], - middlewares=[ToolPolicyMiddleware(policy, approver=RaisingApprover())], - ) - - with self.assertLogs("ejagent.middleware.policy", level="ERROR"): - outcome = await agent.dispatch("write", {"path": "/tmp/a"}) - - self.assertEqual(outcome.control.value, "reject") - self.assertEqual(outcome.data["reason"], "tool approval failed") - self.assertEqual(handler.calls, []) - - async def test_policy_exception_fails_closed(self) -> None: - handler = RecordingHandler() - agent = BaseAgent( - SequenceModel([]), - agent_id="policy-error", - handlers=[handler], - middlewares=[ToolPolicyMiddleware(RaisingPolicy())], - ) - - with self.assertLogs("ejagent.middleware.policy", level="ERROR"): - outcome = await agent.dispatch("read", {"path": "/tmp/a"}) - - self.assertEqual(outcome.control.value, "reject") - self.assertEqual(outcome.data["reason"], "tool policy evaluation failed") - self.assertEqual(handler.calls, []) - - async def test_policy_lifecycle_delegates_through_middleware(self) -> None: - policy = LifecyclePolicy(ToolPolicyDecision(ToolPolicyAction.ALLOW)) - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(content="first"), - AssistantMessage(content="second"), - ] - ), - agent_id="policy-lifecycle", - handlers=[RecordingHandler()], - middlewares=[ToolPolicyMiddleware(policy)], - ) - - await agent.run(task="first") - await agent.run(task="second") - await agent.shutdown() - - self.assertEqual(policy.started, 1) - self.assertEqual(policy.task_starts, 2) - self.assertEqual(policy.stopped, 1) - - async def test_call_limit_reservation_is_atomic_and_resets_per_run(self) -> None: - policy = RuleBasedToolPolicy( - ( - ToolPolicyRule( - rule_id="read-budget", - action=ToolPolicyAction.ALLOW, - tool_names=frozenset({"read"}), - max_calls_per_run=3, - ), - ) - ) - await policy.on_task_start() - - decisions = await asyncio.gather( - *(policy.evaluate(context()) for _ in range(20)) - ) - - self.assertEqual( - sum(item.action is ToolPolicyAction.ALLOW for item in decisions), - 3, - ) - self.assertEqual( - sum(item.action is ToolPolicyAction.DENY for item in decisions), - 17, - ) - self.assertTrue( - all( - item.rule_id == "read-budget" - for item in decisions - if item.action is ToolPolicyAction.DENY - ) - ) - - await policy.on_task_start() - reset_decision = await policy.evaluate(context()) - - self.assertIs(reset_decision.action, ToolPolicyAction.ALLOW) - - async def test_parallel_tool_calls_cannot_overrun_policy_limit(self) -> None: - handler = RecordingHandler() - policy = RuleBasedToolPolicy( - ( - ToolPolicyRule( - rule_id="parallel-read-budget", - action=ToolPolicyAction.ALLOW, - tool_names=frozenset({"read"}), - max_calls_per_run=2, - ), - ) - ) - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=tuple( - call(str(index), "read", f"/workspace/{index}") - for index in range(4) - ) - ) - ] - ), - agent_id="policy-parallel-limit", - handlers=[handler], - middlewares=[ToolPolicyMiddleware(policy)], - runtime_policy=RuntimePolicy(parallel_tool_calls=True), - ) - - result = await agent.run(task="read four files") - - self.assertEqual(result.status, RunStatus.REJECTED) - self.assertEqual(result.stop_reason, StopReason.TOOL_REJECTED) - self.assertEqual(len(handler.calls), 2) - - async def test_invalid_rule_configuration_is_rejected(self) -> None: - with self.assertRaises(ValueError): - RuleBasedToolPolicy( - ( - ToolPolicyRule("duplicate", ToolPolicyAction.ALLOW), - ToolPolicyRule("duplicate", ToolPolicyAction.DENY), - ) - ) - with self.assertRaises(ValueError): - ToolPolicyRule( - "deny-limit", - ToolPolicyAction.DENY, - max_calls_per_run=1, - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_tool_progress.py b/tests/test_tool_progress.py deleted file mode 100644 index 1af7bc8..0000000 --- a/tests/test_tool_progress.py +++ /dev/null @@ -1,326 +0,0 @@ -import asyncio -import unittest -from dataclasses import FrozenInstanceError -from typing import Any - -from ejagent import ( - AgentEvent, - AssistantMessage, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - MemorySessionStorage, - MethodToolHandler, - ModelAdapter, - ModelToolCall, - RunStatus, - SessionRecorder, - StepOutcome, - StopReason, - ToolCallContext, - ToolCompleted, - ToolMiddleware, - ToolNext, - ToolProgressed, - ToolProgressReporter, - ToolProgressUpdate, - ToolStarted, -) - -WORK_TOOL = { - "type": "function", - "function": { - "name": "work", - "description": "Perform observable work.", - "parameters": {"type": "object", "properties": {}}, - }, -} - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - return self.responses.pop(0) - - -class ProgressHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((WORK_TOOL,)) - self.reporter: ToolProgressReporter | None = None - - async def do_work( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - progress: ToolProgressReporter | None = None, - ) -> StepOutcome: - assert progress is not None - self.reporter = progress - await progress.report( - ToolProgressUpdate("scanning", {"completed": 1, "total": 2}) - ) - await progress.report( - ToolProgressUpdate("finished", {"completed": 2, "total": 2}) - ) - return StepOutcome({"files": 2}) - - -class LateProgressHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((WORK_TOOL,)) - self.release = asyncio.Event() - self.late_report: asyncio.Task[None] | None = None - - async def do_work( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - progress: ToolProgressReporter | None = None, - ) -> StepOutcome: - assert progress is not None - await progress.report(ToolProgressUpdate("accepted")) - - async def report_late() -> None: - await self.release.wait() - await progress.report(ToolProgressUpdate("too late")) - - self.late_report = asyncio.create_task(report_late()) - return StepOutcome({"status": "done"}) - - -class BlockingProgressHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((WORK_TOOL,)) - self.started = asyncio.Event() - self.interrupted = asyncio.Event() - self.late_report: asyncio.Task[None] | None = None - - async def do_work( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - progress: ToolProgressReporter | None = None, - ) -> StepOutcome: - assert progress is not None - await progress.report(ToolProgressUpdate("waiting")) - self.started.set() - try: - await asyncio.Future() - finally: - self.interrupted.set() - self.late_report = asyncio.create_task( - progress.report(ToolProgressUpdate("interrupted")) - ) - - -class MiddlewareProgress(ToolMiddleware): - async def __call__( - self, - context: ToolCallContext, - call_next: ToolNext, - ) -> StepOutcome: - assert context.progress is not None - await context.progress.report(ToolProgressUpdate("middleware")) - return await call_next(context) - - -def tool_call() -> ModelToolCall: - return ModelToolCall(id="work-1", name="work", arguments="{}") - - -def payloads(sink: RecordingSink) -> list[object]: - return [event.payload for event in sink.events] - - -class ToolProgressTests(unittest.IsolatedAsyncioTestCase): - async def test_progress_is_ordered_and_correlated_to_tool_call(self) -> None: - call = tool_call() - sink = RecordingSink() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(tool_calls=(call,)), - AssistantMessage(content="done"), - ] - ), - agent_id="progress-events", - handlers=[ProgressHandler()], - event_sink=sink, - ) - - result = await agent.run(task="work") - - tool_payloads = [ - payload - for payload in payloads(sink) - if isinstance(payload, (ToolStarted, ToolProgressed, ToolCompleted)) - ] - self.assertEqual( - [type(payload) for payload in tool_payloads], - [ToolStarted, ToolProgressed, ToolProgressed, ToolCompleted], - ) - progress = tool_payloads[1:3] - self.assertEqual( - [payload.update.message for payload in progress], - ["scanning", "finished"], - ) - self.assertTrue(all(payload.tool_call is call for payload in tool_payloads)) - self.assertEqual( - [event.sequence for event in sink.events], - list(range(1, len(sink.events) + 1)), - ) - self.assertEqual(len({event.run_id for event in sink.events}), 1) - self.assertEqual(result.status, RunStatus.COMPLETED) - - async def test_progress_is_not_persisted_in_session(self) -> None: - call = tool_call() - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="progress", storage=storage) - observer = RecordingSink() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(tool_calls=(call,)), - AssistantMessage(content="done"), - ] - ), - agent_id="progress-session", - handlers=[ProgressHandler()], - event_sink=CompositeAgentEventSink([recorder, observer]), - ) - - await agent.run(task="work") - session = await recorder.load() - - self.assertEqual( - len( - [ - payload - for payload in payloads(observer) - if isinstance(payload, ToolProgressed) - ] - ), - 2, - ) - assert session is not None - self.assertEqual( - [message["role"] for message in session.messages], - ["user", "assistant", "tool", "assistant"], - ) - self.assertNotIn("scanning", repr(session)) - self.assertNotIn("finished", repr(session)) - - async def test_middleware_can_report_on_the_same_tool_call(self) -> None: - sink = RecordingSink() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(tool_calls=(tool_call(),)), - AssistantMessage(content="done"), - ] - ), - agent_id="middleware-progress", - handlers=[ProgressHandler()], - middlewares=[MiddlewareProgress()], - event_sink=sink, - ) - - await agent.run(task="work") - - updates = [ - payload.update.message - for payload in payloads(sink) - if isinstance(payload, ToolProgressed) - ] - self.assertEqual(updates, ["middleware", "scanning", "finished"]) - - async def test_late_progress_is_ignored_after_tool_completion(self) -> None: - handler = LateProgressHandler() - sink = RecordingSink() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage(tool_calls=(tool_call(),)), - AssistantMessage(content="done"), - ] - ), - agent_id="late-progress", - handlers=[handler], - event_sink=sink, - ) - - await agent.run(task="work") - event_count = len(sink.events) - handler.release.set() - assert handler.late_report is not None - await handler.late_report - - self.assertEqual(len(sink.events), event_count) - updates = [ - payload.update.message - for payload in payloads(sink) - if isinstance(payload, ToolProgressed) - ] - self.assertEqual(updates, ["accepted"]) - - async def test_abort_rejects_progress_after_cancellation(self) -> None: - handler = BlockingProgressHandler() - sink = RecordingSink() - agent = BaseAgent( - SequenceModel([AssistantMessage(tool_calls=(tool_call(),))]), - agent_id="abort-progress", - handlers=[handler], - event_sink=sink, - ) - - run = asyncio.create_task(agent.run(task="work")) - await handler.started.wait() - agent.abort("stop progress") - result = await run - assert handler.late_report is not None - await handler.late_report - - tool_payloads = [ - payload - for payload in payloads(sink) - if isinstance(payload, (ToolStarted, ToolProgressed, ToolCompleted)) - ] - self.assertEqual( - [type(payload) for payload in tool_payloads], - [ToolStarted, ToolProgressed, ToolCompleted], - ) - self.assertEqual(tool_payloads[1].update.message, "waiting") - self.assertTrue(tool_payloads[-1].result.cancelled) - self.assertTrue(handler.interrupted.is_set()) - self.assertEqual(result.status, RunStatus.CANCELLED) - self.assertEqual(result.stop_reason, StopReason.EXTERNAL_ABORT) - - async def test_progress_update_is_validated_and_immutable(self) -> None: - with self.assertRaises(ValueError): - ToolProgressUpdate("") - - update = ToolProgressUpdate("working") - with self.assertRaises(FrozenInstanceError): - update.message = "changed" # type: ignore[misc] - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_tool_validation.py b/tests/test_tool_validation.py deleted file mode 100644 index ce8714f..0000000 --- a/tests/test_tool_validation.py +++ /dev/null @@ -1,543 +0,0 @@ -from __future__ import annotations - -import json -import unittest -from typing import Any - -from ejagent import ( - AgentRunResult, - AssistantMessage, - BaseAgent, - CancellationToken, - MemorySessionStorage, - MethodToolHandler, - ModelAdapter, - ModelToolCall, - RuleBasedToolPolicy, - RunStatus, - RuntimePolicy, - SessionRecorder, - StepOutcome, - ToolDefinition, - ToolEffect, - ToolExecutionPolicy, - ToolPolicyAction, - ToolPolicyDecision, - ToolPolicyMiddleware, - ToolPolicyRule, - ToolSchemaConfigurationError, - ToolSchemaValidationMiddleware, -) - -TRANSFER_TOOL = ToolDefinition( - name="transfer", - parameters={ - "type": "object", - "properties": { - "to": {"type": "string"}, - "amount": {"type": "number", "minimum": 0}, - "currency": {"type": "string", "enum": ["USD", "CNY"]}, - }, - "required": ["to", "amount", "currency"], - "additionalProperties": False, - }, -) -INSPECT_TOOL = ToolDefinition( - name="inspect_items", - parameters={ - "type": "object", - "properties": { - "items": { - "type": "array", - "items": { - "type": "object", - "properties": {"count": {"type": "integer"}}, - "required": ["count"], - "additionalProperties": False, - }, - } - }, - "required": ["items"], - "additionalProperties": False, - }, - effect=ToolEffect.READ_ONLY, -) -PING_TOOL = ToolDefinition(name="ping", effect=ToolEffect.READ_ONLY) - - -class SequenceModel(ModelAdapter): - def __init__(self, responses: list[AssistantMessage]) -> None: - self.responses = list(responses) - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - return self.responses.pop(0) - - -class ValidationHandler(MethodToolHandler): - def __init__( - self, - tools: tuple[ToolDefinition | dict[str, Any], ...] = ( - TRANSFER_TOOL, - INSPECT_TOOL, - PING_TOOL, - ), - ) -> None: - super().__init__(tools) - self.calls: list[tuple[str, dict[str, Any]]] = [] - self.started = 0 - self.stopped = 0 - - async def startup(self) -> None: - self.started += 1 - - async def shutdown(self) -> None: - self.stopped += 1 - - async def do_transfer( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken, - ) -> StepOutcome: - self.calls.append(("transfer", arguments)) - return StepOutcome({"status": "success", "tool": "transfer"}) - - async def do_inspect_items( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken, - ) -> StepOutcome: - self.calls.append(("inspect_items", arguments)) - return StepOutcome({"status": "success", "tool": "inspect_items"}) - - async def do_ping( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken, - ) -> StepOutcome: - self.calls.append(("ping", arguments)) - return StepOutcome({"status": "success", "tool": "ping"}) - - -class CountingPolicy(ToolExecutionPolicy): - def __init__(self) -> None: - self.calls: list[str] = [] - - async def evaluate(self, context: Any) -> ToolPolicyDecision: - self.calls.append(context.tool_name) - return ToolPolicyDecision(ToolPolicyAction.ALLOW) - - -def tool_call(call_id: str, name: str, arguments: dict[str, Any]) -> ModelToolCall: - return ModelToolCall( - id=call_id, - name=name, - arguments=json.dumps(arguments), - ) - - -class ToolSchemaValidationTests(unittest.IsolatedAsyncioTestCase): - async def test_valid_arguments_reach_handler(self) -> None: - handler = ValidationHandler() - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-valid", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - arguments = {"to": "alice", "amount": 25.5, "currency": "USD"} - - outcome = await agent.dispatch("transfer", arguments) - - self.assertEqual(outcome.data["status"], "success") - self.assertEqual(handler.calls, [("transfer", arguments)]) - - async def test_type_error_is_safe_and_short_circuits_handler(self) -> None: - handler = ValidationHandler() - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-type-error", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - secret = "secret-account-token" - - outcome = await agent.dispatch( - "transfer", - {"to": "alice", "amount": secret, "currency": "USD"}, - ) - - self.assertEqual(outcome.control.value, "continue") - self.assertEqual(outcome.data["code"], "invalid_tool_arguments") - self.assertEqual(outcome.data["errors"][0]["path"], "/amount") - self.assertEqual(outcome.data["errors"][0]["keyword"], "type") - self.assertNotIn(secret, json.dumps(outcome.data)) - self.assertEqual(handler.calls, []) - - async def test_required_error_points_to_missing_property(self) -> None: - handler = ValidationHandler() - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-required", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - - outcome = await agent.dispatch( - "transfer", - {"amount": 10, "currency": "CNY"}, - ) - - self.assertEqual(outcome.data["errors"][0]["path"], "/to") - self.assertEqual(outcome.data["errors"][0]["keyword"], "required") - self.assertEqual(handler.calls, []) - - async def test_multiple_required_errors_point_to_distinct_properties( - self, - ) -> None: - handler = ValidationHandler() - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-required-many", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - - outcome = await agent.dispatch("transfer", {}) - - paths = {error["path"] for error in outcome.data["errors"]} - self.assertEqual(paths, {"/amount", "/currency", "/to"}) - self.assertEqual(handler.calls, []) - - async def test_nested_array_error_uses_json_pointer(self) -> None: - handler = ValidationHandler() - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-nested", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - - outcome = await agent.dispatch( - "inspect_items", - {"items": [{"count": "many"}]}, - ) - - self.assertEqual(outcome.data["errors"][0]["path"], "/items/0/count") - self.assertEqual(outcome.data["errors"][0]["keyword"], "type") - self.assertEqual(handler.calls, []) - - async def test_enum_and_additional_property_errors_do_not_echo_values( - self, - ) -> None: - handler = ValidationHandler() - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-multiple-errors", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - secret = "secret-field-value" - - outcome = await agent.dispatch( - "transfer", - { - "to": "alice", - "amount": 10, - "currency": "EUR", - "unexpected": secret, - }, - ) - - keywords = {error["keyword"] for error in outcome.data["errors"]} - self.assertEqual(keywords, {"additionalProperties", "enum"}) - self.assertNotIn("EUR", json.dumps(outcome.data)) - self.assertNotIn(secret, json.dumps(outcome.data)) - self.assertEqual(handler.calls, []) - - async def test_error_output_is_bounded(self) -> None: - handler = ValidationHandler() - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-error-limit", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware(max_errors=1)], - ) - - outcome = await agent.dispatch("transfer", {}) - - self.assertEqual(len(outcome.data["errors"]), 1) - self.assertTrue(outcome.data["truncated"]) - self.assertEqual(handler.calls, []) - - async def test_tool_without_parameters_schema_passes_through(self) -> None: - handler = ValidationHandler() - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-absent", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - - outcome = await agent.dispatch("ping", {"anything": "is accepted"}) - - self.assertEqual(outcome.data["status"], "success") - self.assertEqual( - handler.calls, - [("ping", {"anything": "is accepted"})], - ) - - async def test_legacy_openai_definition_is_normalized_and_validated( - self, - ) -> None: - legacy_tool = { - "type": "function", - "function": { - "name": "ping", - "parameters": { - "type": "object", - "properties": {"value": {"type": "integer"}}, - "required": ["value"], - "additionalProperties": False, - }, - }, - } - handler = ValidationHandler((legacy_tool,)) - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-legacy", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - - outcome = await agent.dispatch("ping", {"value": "wrong"}) - - self.assertEqual(outcome.data["code"], "invalid_tool_arguments") - self.assertEqual(handler.calls, []) - - async def test_invalid_schema_fails_agent_startup_and_rolls_back_handler( - self, - ) -> None: - invalid_tool = ToolDefinition( - name="ping", - parameters={"type": "not-a-json-schema-type"}, - ) - handler = ValidationHandler((invalid_tool,)) - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-invalid-config", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - - with self.assertRaises(ToolSchemaConfigurationError) as raised: - await agent.startup() - - self.assertEqual(raised.exception.tool_name, "ping") - self.assertEqual(handler.started, 1) - self.assertEqual(handler.stopped, 1) - - async def test_disabled_validation_preserves_existing_schema_behavior( - self, - ) -> None: - invalid_tool = ToolDefinition( - name="ping", - parameters={"type": "not-a-json-schema-type"}, - ) - handler = ValidationHandler((invalid_tool,)) - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-disabled", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware(enabled=False)], - ) - - outcome = await agent.dispatch("ping", {"legacy": "arguments"}) - - self.assertEqual(outcome.data["status"], "success") - self.assertEqual( - handler.calls, - [("ping", {"legacy": "arguments"})], - ) - - async def test_remote_schema_reference_is_not_retrieved(self) -> None: - remote_tool = ToolDefinition( - name="ping", - parameters={"$ref": "https://schemas.example.invalid/tool.json"}, - ) - handler = ValidationHandler((remote_tool,)) - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-no-remote", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - - with self.assertRaises(ToolSchemaConfigurationError) as raised: - await agent.dispatch("ping", {}) - - self.assertEqual(raised.exception.tool_name, "ping") - self.assertEqual(handler.calls, []) - - async def test_local_schema_reference_is_supported(self) -> None: - referenced_tool = ToolDefinition( - name="ping", - parameters={ - "$defs": {"count": {"type": "integer", "minimum": 1}}, - "type": "object", - "properties": {"count": {"$ref": "#/$defs/count"}}, - "required": ["count"], - "additionalProperties": False, - }, - ) - handler = ValidationHandler((referenced_tool,)) - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-local-ref", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - ) - - invalid = await agent.dispatch("ping", {"count": 0}) - valid = await agent.dispatch("ping", {"count": 2}) - - self.assertEqual(invalid.data["errors"][0]["keyword"], "minimum") - self.assertEqual(valid.data["status"], "success") - self.assertEqual(handler.calls, [("ping", {"count": 2})]) - - async def test_validation_precedes_policy_and_approval_path(self) -> None: - handler = ValidationHandler() - policy = CountingPolicy() - agent = BaseAgent( - SequenceModel([]), - agent_id="schema-before-policy", - handlers=[handler], - middlewares=[ - ToolSchemaValidationMiddleware(), - ToolPolicyMiddleware(policy), - ], - ) - - invalid = await agent.dispatch( - "transfer", - {"to": "alice", "amount": "wrong", "currency": "USD"}, - ) - valid = await agent.dispatch( - "transfer", - {"to": "alice", "amount": 10, "currency": "USD"}, - ) - - self.assertEqual(invalid.data["code"], "invalid_tool_arguments") - self.assertEqual(valid.data["status"], "success") - self.assertEqual(policy.calls, ["transfer"]) - self.assertEqual(len(handler.calls), 1) - - async def test_parallel_calls_validate_independently_before_policy(self) -> None: - handler = ValidationHandler() - policy = RuleBasedToolPolicy( - ( - ToolPolicyRule( - rule_id="allow-inspect", - action=ToolPolicyAction.ALLOW, - tool_names=frozenset({"inspect_items"}), - max_calls_per_run=1, - ), - ) - ) - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=( - tool_call( - "valid", - "inspect_items", - {"items": [{"count": 1}]}, - ), - tool_call( - "invalid", - "inspect_items", - {"items": [{"count": "wrong"}]}, - ), - ) - ), - AssistantMessage(content="done"), - ] - ), - agent_id="schema-parallel", - handlers=[handler], - middlewares=[ - ToolSchemaValidationMiddleware(), - ToolPolicyMiddleware(policy), - ], - runtime_policy=RuntimePolicy(parallel_tool_calls=True), - ) - - result = await agent.run(task="inspect two inputs") - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(len(handler.calls), 1) - tool_messages = [ - message for message in agent.messages if message.get("role") == "tool" - ] - self.assertEqual(len(tool_messages), 2) - self.assertIn("invalid_tool_arguments", tool_messages[1]["content"]) - - async def test_validation_error_is_persisted_as_normal_tool_result(self) -> None: - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="schema-session", storage=storage) - handler = ValidationHandler() - agent = BaseAgent( - SequenceModel( - [ - AssistantMessage( - tool_calls=( - tool_call( - "invalid", - "transfer", - { - "to": "alice", - "amount": "wrong", - "currency": "USD", - }, - ), - ) - ), - AssistantMessage(content="corrected later"), - ] - ), - agent_id="schema-session-agent", - handlers=[handler], - middlewares=[ToolSchemaValidationMiddleware()], - event_sink=recorder, - ) - - result = await agent.run(task="transfer") - restored = await recorder.load() - - self.assertIsInstance(result, AgentRunResult) - self.assertIsNotNone(restored) - assert restored is not None - tool_messages = [ - message for message in restored.messages if message.get("role") == "tool" - ] - self.assertEqual(len(tool_messages), 1) - self.assertIn("invalid_tool_arguments", tool_messages[0]["content"]) - self.assertEqual(handler.calls, []) - - async def test_constructor_rejects_invalid_error_limit(self) -> None: - with self.assertRaises(TypeError): - ToolSchemaValidationMiddleware(max_errors=True) - with self.assertRaises(ValueError): - ToolSchemaValidationMiddleware(max_errors=0) - - -if __name__ == "__main__": - unittest.main() diff --git a/tests/test_usage_budget.py b/tests/test_usage_budget.py deleted file mode 100644 index a0e6bcb..0000000 --- a/tests/test_usage_budget.py +++ /dev/null @@ -1,487 +0,0 @@ -import unittest -from collections.abc import AsyncIterator -from types import SimpleNamespace -from typing import Any - -from ejagent import ( - AgentContextBuilder, - AgentEvent, - AgentState, - AssistantMessage, - BaseAgent, - CancellationToken, - CompositeAgentEventSink, - MemorySessionStorage, - MessageCompleted, - MethodToolHandler, - ModelAdapter, - ModelConfig, - ModelResponseCompleted, - ModelStreamEvent, - ModelToolCall, - ModelUsage, - OpenAIModelAdapter, - RunStatus, - RuntimePolicy, - RunUsage, - SessionRecorder, - StepOutcome, - StopReason, -) - -CONTINUE_TOOL = { - "type": "function", - "function": { - "name": "continue_work", - "description": "Record one completed unit of work.", - "parameters": {"type": "object", "properties": {}}, - }, -} - - -def usage( - input_tokens: int, - output_tokens: int, - *, - cache_read_tokens: int | None = None, - reasoning_tokens: int | None = None, -) -> ModelUsage: - return ModelUsage( - input_tokens=input_tokens, - output_tokens=output_tokens, - total_tokens=input_tokens + output_tokens, - cache_read_tokens=cache_read_tokens, - reasoning_tokens=reasoning_tokens, - ) - - -class RecordingSink: - def __init__(self) -> None: - self.events: list[AgentEvent] = [] - - async def emit(self, event: AgentEvent) -> None: - self.events.append(event) - - -class UsageSequenceModel(ModelAdapter): - def __init__( - self, - responses: list[tuple[AssistantMessage, ModelUsage | None]], - ) -> None: - self.responses = list(responses) - self.contexts: list[Any] = [] - self.calls = 0 - - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - raise AssertionError("stream() should be used") - - async def stream( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AsyncIterator[ModelStreamEvent]: - self.calls += 1 - self.contexts.append(context) - message, response_usage = self.responses.pop(0) - yield ModelResponseCompleted(message, response_usage) - - -class CompleteOnlyModel(ModelAdapter): - async def complete( - self, - context: Any, - *, - cancellation: CancellationToken | None = None, - ) -> AssistantMessage: - return AssistantMessage(content="legacy complete") - - -class ContinueHandler(MethodToolHandler): - def __init__(self) -> None: - super().__init__((CONTINUE_TOOL,)) - self.calls = 0 - - async def do_continue_work( - self, - arguments: dict[str, Any], - *, - cancellation: CancellationToken | None = None, - ) -> StepOutcome: - self.calls += 1 - return StepOutcome({"completed": self.calls}) - - -class FakeOpenAIStream: - def __init__(self, chunks: list[Any]) -> None: - self.chunks = list(chunks) - self.closed = False - - def __aiter__(self) -> "FakeOpenAIStream": - return self - - async def __anext__(self) -> Any: - if not self.chunks: - raise StopAsyncIteration - return self.chunks.pop(0) - - async def close(self) -> None: - self.closed = True - - -class FakeOpenAICompletions: - def __init__(self, response: FakeOpenAIStream) -> None: - self.response = response - self.calls: list[dict[str, Any]] = [] - - async def create(self, **kwargs: Any) -> FakeOpenAIStream: - self.calls.append(kwargs) - return self.response - - -def tool_call() -> ModelToolCall: - return ModelToolCall( - id="continue-1", - name="continue_work", - arguments="{}", - ) - - -class UsageBudgetTests(unittest.IsolatedAsyncioTestCase): - async def test_usage_is_aggregated_persisted_and_removed_from_llm( - self, - ) -> None: - first_usage = usage( - 100, - 20, - cache_read_tokens=30, - reasoning_tokens=5, - ) - second_usage = usage( - 150, - 30, - cache_read_tokens=40, - ) - model = UsageSequenceModel( - [ - (AssistantMessage(tool_calls=(tool_call(),)), first_usage), - (AssistantMessage(content="done"), second_usage), - ] - ) - handler = ContinueHandler() - storage = MemorySessionStorage() - recorder = SessionRecorder(session_id="usage", storage=storage) - observer = RecordingSink() - agent = BaseAgent( - model, - agent_id="usage-aggregate", - handlers=[handler], - event_sink=CompositeAgentEventSink([recorder, observer]), - ) - - result = await agent.run(task="work") - session = await recorder.load() - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(result.usage.input_tokens, 250) - self.assertEqual(result.usage.output_tokens, 50) - self.assertEqual(result.usage.total_tokens, 300) - self.assertEqual(result.usage.request_count, 2) - self.assertEqual(result.usage.reported_request_count, 2) - self.assertEqual(result.usage.cache_read_tokens, 70) - self.assertIsNone(result.usage.reasoning_tokens) - self.assertTrue(result.usage.complete) - - completed = [ - event.payload - for event in observer.events - if isinstance(event.payload, MessageCompleted) - ] - self.assertEqual( - [payload.usage for payload in completed], - [first_usage, second_usage], - ) - - agent_assistant = next( - message for message in agent.messages if message["role"] == "assistant" - ) - self.assertEqual(agent_assistant["usage"], first_usage.to_dict()) - context_assistant_index = next( - index - for index, message in enumerate(model.contexts[1].agent_messages) - if message["role"] == "assistant" - ) - self.assertIn( - "usage", - model.contexts[1].agent_messages[context_assistant_index], - ) - self.assertNotIn( - "usage", - model.contexts[1].llm_messages[context_assistant_index], - ) - self.assertFalse( - any("usage" in message for message in model.contexts[1].llm_messages) - ) - - assert session is not None - session_assistant_index = next( - index - for index, message in enumerate(session.messages) - if message["role"] == "assistant" - ) - self.assertEqual( - session.messages[session_assistant_index]["usage"], - first_usage.to_dict(), - ) - resumed_state = AgentState(messages=session.messages) - resumed_context = AgentContextBuilder().build(resumed_state) - self.assertIn( - "usage", - resumed_context.agent_messages[session_assistant_index], - ) - self.assertNotIn( - "usage", - resumed_context.llm_messages[session_assistant_index], - ) - self.assertEqual(session.runs[0].result, result) - - async def test_budget_stops_before_next_request_after_tools_settle( - self, - ) -> None: - model = UsageSequenceModel( - [ - ( - AssistantMessage(tool_calls=(tool_call(),)), - usage(80, 20), - ), - (AssistantMessage(content="must not run"), usage(10, 5)), - ] - ) - handler = ContinueHandler() - agent = BaseAgent( - model, - agent_id="usage-budget", - handlers=[handler], - runtime_policy=RuntimePolicy(max_run_tokens=100), - ) - - result = await agent.run(task="work") - - self.assertEqual(result.status, RunStatus.FAILED) - self.assertEqual( - result.stop_reason, - StopReason.TOKEN_BUDGET_EXCEEDED, - ) - self.assertEqual(model.calls, 1) - self.assertEqual(handler.calls, 1) - self.assertEqual(result.turns, 1) - self.assertEqual(result.usage.total_tokens, 100) - self.assertEqual( - [message["role"] for message in agent.messages[-3:]], - ["user", "assistant", "tool"], - ) - - async def test_final_response_can_complete_after_crossing_budget(self) -> None: - model = UsageSequenceModel([(AssistantMessage(content="done"), usage(100, 50))]) - agent = BaseAgent( - model, - agent_id="final-over-budget", - runtime_policy=RuntimePolicy(max_run_tokens=100), - ) - - result = await agent.run(task="finish") - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(result.stop_reason, StopReason.TEXT_RESPONSE) - self.assertEqual(result.usage.total_tokens, 150) - self.assertEqual(model.calls, 1) - - async def test_missing_usage_stops_only_when_another_request_is_needed( - self, - ) -> None: - model = UsageSequenceModel( - [ - (AssistantMessage(tool_calls=(tool_call(),)), None), - (AssistantMessage(content="must not run"), usage(10, 5)), - ] - ) - handler = ContinueHandler() - agent = BaseAgent( - model, - agent_id="missing-usage-budget", - handlers=[handler], - runtime_policy=RuntimePolicy(max_run_tokens=100), - ) - - result = await agent.run(task="work") - - self.assertEqual(result.status, RunStatus.FAILED) - self.assertEqual(result.stop_reason, StopReason.USAGE_UNAVAILABLE) - self.assertEqual(model.calls, 1) - self.assertEqual(handler.calls, 1) - self.assertEqual(result.usage.request_count, 1) - self.assertEqual(result.usage.reported_request_count, 0) - self.assertFalse(result.usage.complete) - - async def test_complete_only_adapter_remains_compatible_without_budget( - self, - ) -> None: - agent = BaseAgent(CompleteOnlyModel(), agent_id="legacy-usage") - - result = await agent.run(task="finish") - - self.assertEqual(result.status, RunStatus.COMPLETED) - self.assertEqual(result.usage.request_count, 1) - self.assertEqual(result.usage.reported_request_count, 0) - self.assertFalse(result.usage.complete) - - async def test_usage_accumulator_is_reset_for_each_run(self) -> None: - model = UsageSequenceModel( - [ - (AssistantMessage(content="first"), usage(10, 5)), - (AssistantMessage(content="second"), usage(20, 5)), - ] - ) - agent = BaseAgent(model, agent_id="usage-reset") - - first = await agent.run(task="first") - second = await agent.run(task="second") - - self.assertEqual(first.usage.total_tokens, 15) - self.assertEqual(second.usage.total_tokens, 25) - self.assertEqual(first.usage.request_count, 1) - self.assertEqual(second.usage.request_count, 1) - - async def test_openai_stream_normalizes_usage_only_terminal_chunk( - self, - ) -> None: - response = FakeOpenAIStream( - [ - SimpleNamespace( - usage=None, - choices=[ - SimpleNamespace( - finish_reason="stop", - delta=SimpleNamespace( - content="done", - tool_calls=None, - reasoning_content=None, - reasoning=None, - reasoning_text=None, - ), - ) - ], - ), - SimpleNamespace( - choices=[], - usage=SimpleNamespace( - prompt_tokens=120, - completion_tokens=30, - prompt_tokens_details=SimpleNamespace( - cached_tokens=40, - cache_write_tokens=10, - ), - completion_tokens_details=SimpleNamespace( - reasoning_tokens=12, - ), - ), - ), - ] - ) - completions = FakeOpenAICompletions(response) - client = SimpleNamespace(chat=SimpleNamespace(completions=completions)) - adapter = OpenAIModelAdapter( - ModelConfig( - model="test-model", - api_key="test-key", - base_url="https://example.invalid", - ), - client=client, # type: ignore[arg-type] - ) - context = AgentContextBuilder().build(AgentState()) - - events = [event async for event in adapter.stream(context)] - - terminal = events[-1] - assert isinstance(terminal, ModelResponseCompleted) - self.assertEqual( - terminal.usage, - ModelUsage( - input_tokens=120, - output_tokens=30, - total_tokens=150, - cache_read_tokens=40, - cache_write_tokens=10, - reasoning_tokens=12, - ), - ) - self.assertEqual( - completions.calls[0]["stream_options"], - {"include_usage": True}, - ) - self.assertTrue(response.closed) - - async def test_openai_usage_collection_can_be_disabled(self) -> None: - response = FakeOpenAIStream( - [ - SimpleNamespace( - usage=None, - choices=[ - SimpleNamespace( - finish_reason="stop", - delta=SimpleNamespace( - content="done", - tool_calls=None, - reasoning_content=None, - reasoning=None, - reasoning_text=None, - ), - ) - ], - ) - ] - ) - completions = FakeOpenAICompletions(response) - client = SimpleNamespace(chat=SimpleNamespace(completions=completions)) - adapter = OpenAIModelAdapter( - ModelConfig( - model="test-model", - api_key="test-key", - base_url="https://example.invalid", - include_usage=False, - ), - client=client, # type: ignore[arg-type] - ) - - _ = [ - event - async for event in adapter.stream(AgentContextBuilder().build(AgentState())) - ] - - self.assertNotIn("stream_options", completions.calls[0]) - - async def test_usage_and_budget_values_are_validated(self) -> None: - with self.assertRaises(ValueError): - ModelUsage(10, 5, 14) - with self.assertRaises(ValueError): - ModelUsage(10, 5, 15, reasoning_tokens=6) - with self.assertRaises(ValueError): - RuntimePolicy(max_run_tokens=0) - with self.assertRaises(ValueError): - RunUsage( - input_tokens=1, - output_tokens=1, - total_tokens=2, - request_count=1, - reported_request_count=1, - reasoning_tokens=2, - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/uv.lock b/uv.lock index f646d7f..bc95fff 100644 --- a/uv.lock +++ b/uv.lock @@ -27,6 +27,25 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/78/b6/6307fbef88d9b5ee7421e68d78a9f162e0da4900bc5f5793f6d3d0e34fb8/annotated_types-0.7.0-py3-none-any.whl", hash = "sha256:1f02e8b43a8fbbc3f3e0d4f0f4bfc8131bcb4eebe8849b8e5c773f3a1c582a53", size = 13643, upload-time = "2024-05-20T21:33:24.1Z" }, ] +[[package]] +name = "anthropic" +version = "0.120.2" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "anyio" }, + { name = "distro" }, + { name = "docstring-parser" }, + { name = "httpx" }, + { name = "jiter" }, + { name = "pydantic" }, + { name = "sniffio" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/d7/10/4ca013cb166f226bd89e0aeb0fcaff94f45ddf716d4925ce89475d3c587b/anthropic-0.120.2.tar.gz", hash = "sha256:9722efc10c27a30a69f5338ddacdb35bc6a64297a4e4ba729bf83af873d5fb3a", size = 1008421, upload-time = "2026-07-28T17:38:26.986Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/63/af/0f5db57b9397a0f3b7fc204cbef143401a7cadaf982330f97f1ce3d39f34/anthropic-0.120.2-py3-none-any.whl", hash = "sha256:0f0bc2b381dc0eb41c8d886b815d79c2041cd2374f83aed36f574b6dc9c579c1", size = 1022851, upload-time = "2026-07-28T17:38:25.466Z" }, +] + [[package]] name = "anyio" version = "4.13.0" @@ -329,14 +348,15 @@ name = "ejagent-core" version = "0.6.0" source = { editable = "." } dependencies = [ - { name = "jsonschema" }, { name = "openai" }, { name = "python-dotenv" }, { name = "pyyaml" }, - { name = "referencing" }, ] [package.optional-dependencies] +anthropic = [ + { name = "anthropic" }, +] mcp = [ { name = "fastmcp" }, ] @@ -345,26 +365,23 @@ mcp = [ dev = [ { name = "mypy" }, { name = "ruff" }, - { name = "types-jsonschema" }, { name = "types-pyyaml" }, ] [package.metadata] requires-dist = [ + { name = "anthropic", marker = "extra == 'anthropic'", specifier = ">=0.104.1" }, { name = "fastmcp", marker = "extra == 'mcp'", specifier = ">=3.4.2" }, - { name = "jsonschema", specifier = ">=4.26.0" }, { name = "openai", specifier = ">=2.41.0" }, { name = "python-dotenv", specifier = ">=1.2.2" }, { name = "pyyaml", specifier = ">=6.0.3" }, - { name = "referencing", specifier = ">=0.37.0" }, ] -provides-extras = ["mcp"] +provides-extras = ["anthropic", "mcp"] [package.metadata.requires-dev] dev = [ { name = "mypy", specifier = ">=1.18.2" }, { name = "ruff", specifier = ">=0.14.0" }, - { name = "types-jsonschema", specifier = ">=4.26.0.20260518" }, { name = "types-pyyaml", specifier = ">=6.0.12.20250915" }, ] @@ -1466,18 +1483,6 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/47/aa/218a0eb34de1f753c83e4d0d1c8e7c4cef27f20dcb8342e024f63a80dc86/tqdm-4.68.1-py3-none-any.whl", hash = "sha256:fea4a90e4023f764914569f7802a297277c5ab1a66be5144143e142e1a4031d8", size = 78354, upload-time = "2026-06-05T17:23:13.654Z" }, ] -[[package]] -name = "types-jsonschema" -version = "4.26.0.20260518" -source = { registry = "https://pypi.org/simple" } -dependencies = [ - { name = "referencing" }, -] -sdist = { url = "https://files.pythonhosted.org/packages/dd/46/73b6a5d61a61015c4248030a8cb07e5bdddb4041430fae9e585a68692578/types_jsonschema-4.26.0.20260518.tar.gz", hash = "sha256:e1dd53dc97a64f5eccdd6fa9839666e09bb500a8ebba2db6fdaf1789faea81a6", size = 16638, upload-time = "2026-05-18T06:06:44.106Z" } -wheels = [ - { url = "https://files.pythonhosted.org/packages/07/d5/134f8a147dcecda10db7f60cfc6af0578a25a5c53c87b3907a64385e0184/types_jsonschema-4.26.0.20260518-py3-none-any.whl", hash = "sha256:30b30a518c7fe335df85c919fcbcc631b69c03d4a4b5b632fa916bea03065307", size = 16072, upload-time = "2026-05-18T06:06:43.264Z" }, -] - [[package]] name = "types-pyyaml" version = "6.0.12.20260518"