From b8cc6d263b61e4b82860de503760daeb611b468a Mon Sep 17 00:00:00 2001 From: xavix-yo Date: Sat, 25 Jul 2026 23:40:41 +0100 Subject: [PATCH] =?UTF-8?q?agent-loop:=20new=20crate=20=E2=80=94=20LLM=20l?= =?UTF-8?q?oop=20kernel=20+=20Model=20clients=20(phase=200)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Extract the LLM agent loop into a standalone workspace crate with zero deps on skald-core/core-api (blueprint project-loop.md, D13-D15): - kernel: round loop, model fallback with rebuild, parallel tool fan-out (ordered id alloc / bounded concurrent exec / ordered record), streaming deltas drained before outcomes, sticky cancellation - models: OpenAiModel/AnthropicModel/OllamaModel/LmStudioModel ported from llm-client onto the Model trait; ModelError carries the HTTP status; is_retriable default = the 401/403/404/422 rule - DTL as crate protocol (ToolRendering Inline/DeferredToolReference/ SystemToolBlock; Anthropic conversions + Kimi system+tools passthrough), host catalog behind ActivationSource/ToolActivator - HistoryStore durability contract + InMemoryStore; LinearAssembler with well-formed projection (incl. DTL injection, summary, crash survivors) - LoopManager singleton (broadcast bus + live registry), one live loop per conversation, orphan-marking on start_turn - 32 tests green (kernel §13 suite, assembler DTL, SSE/Anthropic ports), clippy clean --- Cargo.lock | 16 + Cargo.toml | 1 + crates/agent-loop/Cargo.toml | 21 + crates/agent-loop/src/activation.rs | 112 ++++ crates/agent-loop/src/context.rs | 343 ++++++++++ crates/agent-loop/src/events.rs | 170 +++++ crates/agent-loop/src/gate.rs | 77 +++ crates/agent-loop/src/hooks.rs | 50 ++ crates/agent-loop/src/human.rs | 105 +++ crates/agent-loop/src/ids.rs | 54 ++ crates/agent-loop/src/kernel.rs | 584 ++++++++++++++++ crates/agent-loop/src/lib.rs | 73 ++ crates/agent-loop/src/manager.rs | 411 ++++++++++++ crates/agent-loop/src/model.rs | 440 ++++++++++++ crates/agent-loop/src/models/anthropic.rs | 775 ++++++++++++++++++++++ crates/agent-loop/src/models/lm_studio.rs | 43 ++ crates/agent-loop/src/models/mod.rs | 42 ++ crates/agent-loop/src/models/ollama.rs | 102 +++ crates/agent-loop/src/models/openai.rs | 459 +++++++++++++ crates/agent-loop/src/models/sse.rs | 80 +++ crates/agent-loop/src/store.rs | 281 ++++++++ crates/agent-loop/src/store_memory.rs | 256 +++++++ crates/agent-loop/src/testing.rs | 130 ++++ crates/agent-loop/src/tool.rs | 365 ++++++++++ crates/agent-loop/tests/assembler.rs | 196 ++++++ crates/agent-loop/tests/kernel.rs | 498 ++++++++++++++ 26 files changed, 5684 insertions(+) create mode 100644 crates/agent-loop/Cargo.toml create mode 100644 crates/agent-loop/src/activation.rs create mode 100644 crates/agent-loop/src/context.rs create mode 100644 crates/agent-loop/src/events.rs create mode 100644 crates/agent-loop/src/gate.rs create mode 100644 crates/agent-loop/src/hooks.rs create mode 100644 crates/agent-loop/src/human.rs create mode 100644 crates/agent-loop/src/ids.rs create mode 100644 crates/agent-loop/src/kernel.rs create mode 100644 crates/agent-loop/src/lib.rs create mode 100644 crates/agent-loop/src/manager.rs create mode 100644 crates/agent-loop/src/model.rs create mode 100644 crates/agent-loop/src/models/anthropic.rs create mode 100644 crates/agent-loop/src/models/lm_studio.rs create mode 100644 crates/agent-loop/src/models/mod.rs create mode 100644 crates/agent-loop/src/models/ollama.rs create mode 100644 crates/agent-loop/src/models/openai.rs create mode 100644 crates/agent-loop/src/models/sse.rs create mode 100644 crates/agent-loop/src/store.rs create mode 100644 crates/agent-loop/src/store_memory.rs create mode 100644 crates/agent-loop/src/testing.rs create mode 100644 crates/agent-loop/src/tool.rs create mode 100644 crates/agent-loop/tests/assembler.rs create mode 100644 crates/agent-loop/tests/kernel.rs diff --git a/Cargo.lock b/Cargo.lock index 94c2c4d..1ca477a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -43,6 +43,22 @@ dependencies = [ "subtle", ] +[[package]] +name = "agent-loop" +version = "0.1.0" +dependencies = [ + "anyhow", + "async-trait", + "futures", + "futures-util", + "reqwest 0.13.4", + "serde", + "serde_json", + "tokio", + "tokio-util", + "tracing", +] + [[package]] name = "aho-corasick" version = "1.1.4" diff --git a/Cargo.toml b/Cargo.toml index 7ed15b8..a0325c4 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,7 @@ [workspace] members = [ ".", + "crates/agent-loop", "crates/skald-core", "crates/skald-setup", "crates/honcho-client", diff --git a/crates/agent-loop/Cargo.toml b/crates/agent-loop/Cargo.toml new file mode 100644 index 0000000..ab3a306 --- /dev/null +++ b/crates/agent-loop/Cargo.toml @@ -0,0 +1,21 @@ +[package] +name = "agent-loop" +version = "0.1.0" +edition = "2024" +description = "Reusable LLM agent loop kernel: round loop, tool calling, fallback, streaming, durability traits — no database, no host types." +license = "MIT" + +[dependencies] +tokio = { version = "1", features = ["sync", "rt", "time", "macros"] } +tokio-util = { version = "0.7" } +async-trait = "0.1" +serde = { version = "1", features = ["derive"] } +serde_json = "1" +tracing = "0.1" +anyhow = "1" +futures = "0.3" +futures-util = "0.3" +reqwest = { version = "0.13.4", default-features = false, features = ["rustls-no-provider", "charset", "http2", "system-proxy", "json", "stream"] } + +[dev-dependencies] +tokio = { version = "1", features = ["macros", "rt-multi-thread"] } diff --git a/crates/agent-loop/src/activation.rs b/crates/agent-loop/src/activation.rs new file mode 100644 index 0000000..2bfbfc1 --- /dev/null +++ b/crates/agent-loop/src/activation.rs @@ -0,0 +1,112 @@ +//! Dynamic tool loading (DTL) — the wire PROTOCOL lives in the crate +//! (blueprint D15), the catalog and persistence stay with the host. +//! +//! Three rendering modes ([`ToolRendering`]) decide how dynamically-activated +//! tools reach the model without invalidating the prompt-cache prefix: +//! +//! - `Inline`: active tools go in the `tools` array (every activation changes +//! the array — no cache). +//! - `DeferredToolReference`: all activatable tools are declared upfront with +//! `defer_loading: true`; an activation's tool result carries a +//! `_tool_references` marker the Anthropic client converts to +//! `tool_reference` blocks. +//! - `SystemToolBlock`: activated tools never touch the `tools` array; a +//! `{role:"system", tools:[…]}` message is appended after the activation's +//! tool-result group (Kimi/Moonshot speaks this natively). + +use std::sync::Arc; + +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::ids::MessageId; +use crate::tool::{Tool, ToolCtx, ToolFailure, ToolOutput}; + +/// How dynamically-activated tools are rendered on the wire. On +/// [`crate::model::ModelInfo`]; read by `ToolSet::defs` and assemblers, +/// consumed by the shipped clients. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub enum ToolRendering { + /// Only the currently-active tools in the `tools` array. + #[default] + Inline, + /// Anthropic: all activatable tools `defer_loading: true` + tool_reference + /// blocks in activation results. + DeferredToolReference, + /// Kimi K3: `{role:"system", tools:[defs]}` appended after the activation + /// (append-only, cache-safe). + SystemToolBlock, +} + +/// One activation: the defs of the groups activated at a given anchor message. +#[derive(Debug, Clone)] +pub struct Activation { + pub anchor: MessageId, + /// OpenAI-shaped tool defs of the groups activated at `anchor`. + pub defs: Vec, +} + +/// Catalog + persistence of activations — implemented by the host. Consulted +/// by assemblers (injection) and by host `ToolSet`s (array rendering). +#[async_trait] +pub trait ActivationSource: Send + Sync { + /// The activations in force for a frame, ordered by anchor. + async fn activations(&self, frame: crate::ids::FrameId) -> crate::Result>; +} + +/// Backend of the shipped [`ActivateToolsTool`]: validates the groups, mutates +/// the grants, persists the activation (anchored at the current message via +/// `ctx`). Returns the confirmation text shown to the model. +#[async_trait] +pub trait ToolActivator: Send + Sync { + async fn activate(&self, groups: Vec, ctx: &ToolCtx) -> Result; +} + +/// The shipped `activate_tools` tool. To the kernel it's a tool like any +/// other — the defs re-read at the next round makes the new grants visible. +pub struct ActivateToolsTool { + activator: Arc, +} + +impl ActivateToolsTool { + pub fn new(activator: Arc) -> Self { Self { activator } } +} + +#[async_trait] +impl Tool for ActivateToolsTool { + fn name(&self) -> &str { "activate_tools" } + + fn definition(&self) -> Value { + json!({ + "type": "function", + "function": { + "name": "activate_tools", + "description": "Load additional tool groups on demand. Activated tools \ + become available from the next step of this conversation.", + "parameters": { + "type": "object", + "properties": { + "groups": { + "type": "array", + "items": { "type": "string" }, + "description": "Names of the tool groups to activate" + } + }, + "required": ["groups"] + } + } + }) + } + + async fn call(&self, args: Value, ctx: &ToolCtx) -> Result { + let groups: Vec = args["groups"] + .as_array() + .map(|a| a.iter().filter_map(|v| v.as_str().map(str::to_string)).collect()) + .unwrap_or_default(); + if groups.is_empty() { + return Err(ToolFailure::Failed("activate_tools: no groups given".into())); + } + let text = self.activator.activate(groups, ctx).await?; + Ok(ToolOutput::Text(text)) + } +} diff --git a/crates/agent-loop/src/context.rs b/crates/agent-loop/src/context.rs new file mode 100644 index 0000000..475ee66 --- /dev/null +++ b/crates/agent-loop/src/context.rs @@ -0,0 +1,343 @@ +//! The system context (layered) and the `ContextAssembler` — from system + +//! history to wire messages. +//! +//! **Well-formedness contract** (every assembler MUST honor it): +//! +//! 1. Order: static system → compaction summary (if any) → messages after +//! `covered_up_to` → dynamic tail → tail reminder. +//! 2. Every assistant `tool_call` has a tool-result: `Done`→result, +//! `Failed`→error, `Cancelled`/`Rejected`→note, **`Running`/`AwaitingHuman` +//! surviving a crash → synthetic "interrupted" result**. +//! 3. No `failed` messages (orphans) — already filtered by the store. +//! 4. DTL injection (§4.10 of the blueprint): when `model.tool_rendering` is +//! not `Inline` and an `ActivationSource` is present, each activation is +//! projected at its anchor (marker vs system+tools block, append-only). + +use std::sync::Arc; + +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::activation::{ActivationSource, ToolRendering}; +use crate::ids::{ConversationId, FrameId}; +use crate::model::ModelInfo; +use crate::store::{HistoryStore, Role, StoredMessage}; + +// ── SystemContext ──────────────────────────────────────────────────────────── + +/// The system prompt as LAYERS (the static prefix is cacheable, the dynamic +/// tail is per-turn fresh). +#[derive(Debug, Clone, Default)] +pub struct SystemContext { + /// The agent's prompt (static, cacheable). + pub base: String, + /// Per-interface extras (e.g. output format rules). + pub extra_static: Vec, + /// Per-turn: date/time, memory, run context. + pub dynamic_tail: Vec, + pub tail_reminder: Option, +} + +impl SystemContext { + pub fn base(s: impl Into) -> Self { + Self { base: s.into(), ..Default::default() } + } + + pub fn with_dynamic(mut self, s: impl Into) -> Self { + self.dynamic_tail.push(s.into()); + self + } + + pub fn with_static(mut self, s: impl Into) -> Self { + self.extra_static.push(s.into()); + self + } + + pub fn with_reminder(mut self, s: impl Into) -> Self { + self.tail_reminder = Some(s.into()); + self + } +} + +// ── SystemContextSource ────────────────────────────────────────────────────── + +/// What the kernel knows about the current turn when asking for the system +/// context. +#[derive(Debug, Clone)] +pub struct TurnInfo { + pub conversation: ConversationId, + pub frame: FrameId, + pub agent: String, + /// The user message that opened the turn (None on resume). + pub user_message: Option, +} + +#[async_trait] +pub trait SystemContextSource: Send + Sync { + async fn system_context(&self, turn: &TurnInfo) -> crate::Result; +} + +/// A fixed system context (simple hosts, tests). +pub struct StaticSystemContext { + ctx: SystemContext, +} + +impl StaticSystemContext { + pub fn new(base: impl Into) -> Self { + Self { ctx: SystemContext::base(base) } + } +} + +#[async_trait] +impl SystemContextSource for StaticSystemContext { + async fn system_context(&self, _turn: &TurnInfo) -> crate::Result { + Ok(self.ctx.clone()) + } +} + +// ── ContextAssembler ───────────────────────────────────────────────────────── + +pub struct AssembleInput { + pub frame: FrameId, + pub system: SystemContext, + pub model: ModelInfo, + pub round: usize, +} + +#[async_trait] +pub trait ContextAssembler: Send + Sync { + async fn build( + &self, + store: &Arc, + input: &AssembleInput, + ) -> crate::Result>; +} + +// ── LinearAssembler ────────────────────────────────────────────────────────── + +/// The shipped assembler: system + summary + history, with an optional +/// message window and tool-result truncation. Honors the DTL injection +/// contract when given an `ActivationSource`. +pub struct LinearAssembler { + /// Keep at most this many history messages (cut at a User/Agent boundary, + /// never mid assistant+tool group). + pub max_messages: Option, + /// Truncate each tool result to this many chars. + pub max_tool_result_chars: Option, + /// DTL activations (only consulted when `tool_rendering != Inline`). + pub activation: Option>, +} + +impl LinearAssembler { + pub fn new() -> Self { + Self { max_messages: None, max_tool_result_chars: None, activation: None } + } + + pub fn with_max_messages(mut self, n: usize) -> Self { + self.max_messages = Some(n); + self + } + + pub fn with_tool_result_limit(mut self, n: usize) -> Self { + self.max_tool_result_chars = Some(n); + self + } + + pub fn with_activation(mut self, src: Arc) -> Self { + self.activation = Some(src); + self + } +} + +impl Default for LinearAssembler { + fn default() -> Self { Self::new() } +} + +/// The summary block is prefixed so the model understands what it is (Skald +/// keeps its own SUMMARY_PREFIX in its assembler). +pub const SUMMARY_PREFIX: &str = "[CONTEXT SUMMARY — earlier messages were compacted into this summary]"; + +#[async_trait] +impl ContextAssembler for LinearAssembler { + async fn build( + &self, + store: &Arc, + input: &AssembleInput, + ) -> crate::Result> { + let mut out: Vec = Vec::new(); + + // 1. static system + if !input.system.base.is_empty() { + out.push(json!({ "role": "system", "content": input.system.base })); + } + for s in &input.system.extra_static { + out.push(json!({ "role": "system", "content": s })); + } + + // 2. summary + surviving history + let summary = store.latest_summary(input.frame).await?; + if let Some(s) = &summary { + out.push(json!({ + "role": "system", + "content": format!("{SUMMARY_PREFIX}\n\n{}", s.text), + })); + } + let mut history = match &summary { + Some(s) => store.load_since(input.frame, s.covered_up_to).await?, + None => store.load(input.frame).await?, + }; + if let Some(max) = self.max_messages { + history = window(history, max); + } + + // 3. DTL activations (consulted only in non-Inline modes) + let activations = match (&self.activation, input.model.tool_rendering) { + (Some(src), ToolRendering::Inline) => { + let _ = src; + Vec::new() + } + (Some(src), _) => src.activations(input.frame).await.unwrap_or_default(), + (None, _) => Vec::new(), + }; + + for msg in &history { + project_message(&mut out, msg, self.max_tool_result_chars); + inject_activations(&mut out, msg, &activations, &input.model.tool_rendering); + } + + // 4. dynamic tail + reminder + for s in &input.system.dynamic_tail { + out.push(json!({ "role": "system", "content": s })); + } + if let Some(r) = &input.system.tail_reminder { + out.push(json!({ "role": "system", "content": r })); + } + + Ok(out) + } +} + +/// Cut the history to at most `max` messages, at a User/Agent boundary so an +/// assistant+tool group is never split. +fn window(history: Vec, max: usize) -> Vec { + if history.len() <= max { + return history; + } + let start = history.len() - max; + let cut = history[start..] + .iter() + .position(|m| matches!(m.role, Role::User | Role::Agent)) + .map(|p| start + p) + .unwrap_or(start); + history[cut..].to_vec() +} + +/// Project one stored message (and its tool results) to wire messages. +fn project_message(out: &mut Vec, msg: &StoredMessage, result_limit: Option) { + match msg.role { + Role::System => { + out.push(json!({ "role": "system", "content": msg.content })); + } + Role::User | Role::Agent => { + out.push(json!({ "role": "user", "content": msg.content })); + } + Role::Assistant => { + let mut wire = json!({ "role": "assistant", "content": msg.content }); + if let Some(r) = &msg.reasoning { + // Echoed under both names: DeepSeek expects reasoning_content, + // others reasoning (the clients normalize on read). + wire["reasoning_content"] = json!(r); + } + if !msg.calls.is_empty() { + let calls: Vec = msg + .calls + .iter() + .map(|c| { + json!({ + "id": c.provider_id, + "type": "function", + "function": { + "name": c.name, + "arguments": serde_json::to_string(&c.arguments) + .unwrap_or_else(|_| "{}".into()), + }, + }) + }) + .collect(); + wire["tool_calls"] = Value::Array(calls); + } + out.push(wire); + + for call in &msg.calls { + let mut content = match call.state { + crate::store::CallState::Running | crate::store::CallState::AwaitingHuman => { + "[interrupted: this tool call did not complete — the session restarted \ + before a result was recorded]" + .to_string() + } + _ => call.result.clone().unwrap_or_default(), + }; + if let Some(limit) = result_limit + && content.chars().count() > limit + { + content = format!( + "{}… [truncated]", + content.chars().take(limit).collect::() + ); + } + out.push(json!({ + "role": "tool", + "tool_call_id": call.provider_id, + "content": content, + })); + } + } + } +} + +/// DTL injection at an activation anchor (blueprint §4.10): +/// - `DeferredToolReference`: `_tool_references` marker on the FIRST tool +/// result of the anchored assistant message (the client converts it). +/// - `SystemToolBlock`: a `{role:"system", tools:[defs]}` message appended +/// right after the anchored message's tool-result group. +fn inject_activations( + out: &mut Vec, + msg: &StoredMessage, + activations: &[crate::activation::Activation], + mode: &ToolRendering, +) { + let acts: Vec<&crate::activation::Activation> = + activations.iter().filter(|a| a.anchor == msg.id).collect(); + if acts.is_empty() { + return; + } + match mode { + ToolRendering::Inline => {} + ToolRendering::DeferredToolReference => { + let names: Vec = acts + .iter() + .flat_map(|a| &a.defs) + .filter_map(|d| d["function"]["name"].as_str()) + .map(|n| json!(n)) + .collect(); + if names.is_empty() { + return; + } + // Attach to the first tool result just emitted for this message. + if let Some(tool_msg) = out + .iter_mut() + .rev() + .take(msg.calls.len()) + .find(|m| m["role"].as_str() == Some("tool")) + { + tool_msg["_tool_references"] = Value::Array(names); + } + } + ToolRendering::SystemToolBlock => { + let defs: Vec = acts.iter().flat_map(|a| a.defs.clone()).collect(); + if !defs.is_empty() { + out.push(json!({ "role": "system", "tools": defs })); + } + } + } +} diff --git a/crates/agent-loop/src/events.rs b/crates/agent-loop/src/events.rs new file mode 100644 index 0000000..aafe754 --- /dev/null +++ b/crates/agent-loop/src/events.rs @@ -0,0 +1,170 @@ +//! The loop event taxonomy and the broadcast bus. +//! +//! Every event is wrapped in [`Event`], tagged with the emitting conversation, +//! frame and parent frame — subscribers (a UI translator, a logger) reconstruct +//! nesting from the tags. Transport: `tokio::sync::broadcast` (multi-subscriber, +//! lag-tolerant). + +use serde_json::Value; +use tokio::sync::broadcast; + +use crate::ids::{ConversationId, FrameId, MessageId, ModelId, TaskId, ToolCallId}; +use crate::model::{ToolCall, Usage}; +use crate::store::CallOutcome; + +/// Whether a [`LoopEvent::TokenDelta`] carries visible answer text or reasoning. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum DeltaKind { + Content, + Reasoning, +} + +/// Events emitted by a running loop. Every variant is wrapped in [`Event`] +/// before hitting the bus, so conversation/frame tags are never optional. +#[derive(Debug, Clone)] +pub enum LoopEvent { + // ── turn ── + TurnStarted, + RoundStarted { + round: usize, + }, + UserMessage { + message_id: MessageId, + content: String, + synthetic: bool, + metadata: Option, + }, + TokenDelta { + kind: DeltaKind, + text: String, + }, + Thinking { + message_id: MessageId, + content: String, + usage: Usage, + reasoning: Option, + }, + Done { + message_id: MessageId, + content: String, + usage: Usage, + reasoning: Option, + }, + // ── tools ── + ToolCallStarted { + id: ToolCallId, + message_id: MessageId, + name: String, + args: Value, + }, + ToolCallFinished { + id: ToolCallId, + outcome: CallOutcome, + }, + ApprovalRequired { + id: ToolCallId, + name: String, + args: Value, + }, + // ── sub-agents (emitted by child loops; parent_frame in the tag) ── + AgentSpawned { + frame: FrameId, + agent: String, + depth: u32, + prompt_preview: String, + }, + AgentFinished { + frame: FrameId, + agent: String, + result_preview: String, + }, + AsyncResultReady { + task: TaskId, + }, + // ── infrastructure ── + ModelFallback { + from: ModelId, + to: ModelId, + reason: String, + }, + LlmFailed { + tried: Vec, + last_error: String, + }, + Compacted { + frame: FrameId, + covered_up_to: MessageId, + }, + Truncated { + output_tokens: Option, + }, + Error(String), + Cancelled, + /// Escape hatch for host-specific events (Skald: PendingWrite with diff, + /// SecurityGroupSelected, …). Other subscribers ignore it. + Host(Value), +} + +/// An event tagged with its emitting scope. +#[derive(Debug, Clone)] +pub struct Event { + pub conversation: ConversationId, + pub frame: FrameId, + pub parent_frame: Option, + pub inner: E, +} + +/// Thin wrapper over the manager's broadcast sender, handed to the kernel, +/// gates, tools and hooks for out-of-band emission. Cheap to clone. +#[derive(Clone)] +pub struct EventSink { + pub(crate) conversation: ConversationId, + pub(crate) tx: broadcast::Sender>, +} + +impl EventSink { + pub(crate) fn new(conversation: ConversationId, tx: broadcast::Sender>) -> Self { + Self { conversation, tx } + } + + /// Emit an event for a frame. Best-effort: with no subscribers the send + /// fails silently — events are never load-bearing for the loop's outcome. + pub fn emit(&self, frame: FrameId, parent_frame: Option, inner: LoopEvent) { + let _ = self.tx.send(Event { + conversation: self.conversation.clone(), + frame, + parent_frame, + inner, + }); + } + + pub fn conversation(&self) -> &ConversationId { &self.conversation } + + /// Recover the sink from a tool's extensions (the kernel inserts one into + /// every `ToolCtx` it builds, so shipped tools can emit out-of-band). + pub fn from_extensions(ext: &crate::tool::Extensions) -> Option { + ext.get::().map(|s| (*s).clone()) + } +} + +/// A running tool call, as passed to `LoopHooks::pre_tool_call` (mutable) and +/// `post_tool_call`. Distinct from the model's [`crate::model::ToolCall`]: +/// this one carries the store id allocated before execution. +#[derive(Debug, Clone)] +pub struct PendingToolCall { + pub id: ToolCallId, + pub message_id: MessageId, + pub provider_id: Option, + pub name: String, + pub arguments: Value, +} + +impl PendingToolCall { + pub fn wire_call(&self) -> ToolCall { + ToolCall { + id: self.provider_id.clone().unwrap_or_default(), + name: self.name.clone(), + arguments: self.arguments.clone(), + } + } +} diff --git a/crates/agent-loop/src/gate.rs b/crates/agent-loop/src/gate.rs new file mode 100644 index 0000000..650efcb --- /dev/null +++ b/crates/agent-loop/src/gate.rs @@ -0,0 +1,77 @@ +//! `Gate` — the pre-execution decision point (policy and/or human). It MAY +//! block waiting for a human: the implementation decides (oneshot, UI, …). +//! Before suspending, an implementation marks the call `AwaitingHuman` via the +//! store (durability) and emits `LoopEvent::ApprovalRequired`. + +use async_trait::async_trait; +use serde_json::Value; + +use crate::events::EventSink; +use crate::ids::{FrameId, ToolCallId}; +use crate::tool::Extensions; + +/// A tool call awaiting a gate decision. +#[derive(Debug, Clone)] +pub struct PendingCall { + pub id: ToolCallId, + pub name: String, + pub args: Value, + pub frame: FrameId, + pub agent: String, + /// Host free-form (source, permission group, …). + pub extensions: Extensions, +} + +/// The gate's verdict. +#[derive(Debug, Clone)] +pub enum GateDecision { + Allow, + Reject { reason: String }, +} + +#[async_trait] +pub trait Gate: Send + Sync { + /// Decide on a call. MAY block awaiting a human — in that case the + /// implementation marks the call `AwaitingHuman` first (via the store the + /// host gave it) and emits `ApprovalRequired` on `events`. + async fn check(&self, call: &PendingCall, events: &EventSink) -> GateDecision; +} + +/// Everything runs. The default for simple hosts and tests. +pub struct AllowAll; + +#[async_trait] +impl Gate for AllowAll { + async fn check(&self, _call: &PendingCall, _events: &EventSink) -> GateDecision { + GateDecision::Allow + } +} + +/// Reject calls whose name matches a pattern: exact, or `prefix*`. +pub struct DenyList { + patterns: Vec, +} + +impl DenyList { + pub fn new(patterns: impl IntoIterator>) -> Self { + Self { patterns: patterns.into_iter().map(Into::into).collect() } + } + + fn matches(&self, name: &str) -> bool { + self.patterns.iter().any(|p| match p.strip_suffix('*') { + Some(prefix) => name.starts_with(prefix), + None => name == p, + }) + } +} + +#[async_trait] +impl Gate for DenyList { + async fn check(&self, call: &PendingCall, _events: &EventSink) -> GateDecision { + if self.matches(&call.name) { + GateDecision::Reject { reason: format!("tool '{}' denied by policy", call.name) } + } else { + GateDecision::Allow + } + } +} diff --git a/crates/agent-loop/src/hooks.rs b/crates/agent-loop/src/hooks.rs new file mode 100644 index 0000000..9795ef6 --- /dev/null +++ b/crates/agent-loop/src/hooks.rs @@ -0,0 +1,50 @@ +//! `LoopHooks` — the passive/active interception seam. Every host special-case +//! (diff-preview bracketing, per-tool arg normalization, telemetry, discovery) +//! lives here, not in the kernel. All methods default to no-op. + +use std::sync::Arc; + +use async_trait::async_trait; + +use crate::events::{EventSink, PendingToolCall}; +use crate::ids::{ConversationId, FrameId, MessageId}; +use crate::kernel::TurnOutcome; +use crate::store::{CallOutcome, HistoryStore}; + +/// Verdict of `pre_tool_call`. +#[derive(Debug, Clone)] +pub enum HookVerdict { + Allow, + Reject { reason: String }, +} + +/// Context handed to every hook. +pub struct HookCtx { + pub conversation: ConversationId, + pub frame: FrameId, + pub agent: String, + pub store: Arc, + pub events: EventSink, +} + +#[async_trait] +pub trait LoopHooks: Send + Sync { + async fn before_round(&self, _round: usize, _ctx: &HookCtx) {} + async fn after_round(&self, _round: usize, _ctx: &HookCtx) {} + + /// May MUTATE the call's arguments or veto it (Reject). Covers diff-preview + /// bracketing and per-tool normalizations. + async fn pre_tool_call(&self, _call: &mut PendingToolCall, _ctx: &HookCtx) -> HookVerdict { + HookVerdict::Allow + } + + /// Covers persistence of activated tools, discovery, file-change + /// notifications, telemetry. + async fn post_tool_call(&self, _call: &PendingToolCall, _outcome: &CallOutcome, _ctx: &HookCtx) {} + + async fn on_turn_end(&self, _outcome: &TurnOutcome, _ctx: &HookCtx) {} + + /// Fired after a compaction (blueprint §9): hosts re-anchor DTL + /// activations to the first surviving message here. + async fn on_compacted(&self, _frame: FrameId, _covered: MessageId, _first_surviving: MessageId) {} +} diff --git a/crates/agent-loop/src/human.rs b/crates/agent-loop/src/human.rs new file mode 100644 index 0000000..7d2ba67 --- /dev/null +++ b/crates/agent-loop/src/human.rs @@ -0,0 +1,105 @@ +//! `HumanChannel` + the shipped `ask_user` tool: synchronous +//! question-to-a-human from inside a tool call. + +use std::sync::Arc; + +use async_trait::async_trait; +use serde_json::{Value, json}; + +use crate::events::EventSink; +use crate::ids::ToolCallId; +use crate::store::{CallState, HistoryStore}; +use crate::tool::{Tool, ToolCtx, ToolFailure, ToolOutput}; + +/// A question posed to a human. +#[derive(Debug, Clone)] +pub struct Question { + pub title: String, + pub question: String, + pub suggested: Vec, + /// The tool call asking (for UI correlation). + pub call: ToolCallId, +} + +/// The human channel closed while waiting (WS down, user gone). +#[derive(Debug, Clone, Copy)] +pub struct HumanGone; + +impl std::fmt::Display for HumanGone { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("human channel closed") + } +} +impl std::error::Error for HumanGone {} + +#[async_trait] +pub trait HumanChannel: Send + Sync { + /// Block until an answer arrives. `Err(HumanGone)` = the channel closed: + /// the tool returns [`ToolFailure::Suspend`] and the call stays + /// `AwaitingHuman` for a later resume. + async fn ask(&self, q: Question, events: &EventSink) -> Result; +} + +/// The shipped `ask_user` tool. Marks the call `AwaitingHuman` BEFORE +/// suspending (durability rule: a crash mid-question must be recoverable), +/// then blocks on the channel. +pub struct AskUserTool { + channel: Arc, + store: Arc, +} + +impl AskUserTool { + pub fn new(channel: Arc, store: Arc) -> Self { + Self { channel, store } + } +} + +#[async_trait] +impl Tool for AskUserTool { + fn name(&self) -> &str { "ask_user" } + + fn definition(&self) -> Value { + json!({ + "type": "function", + "function": { + "name": "ask_user", + "description": "Ask the user a clarifying question and wait for the answer.", + "parameters": { + "type": "object", + "properties": { + "title": { "type": "string", "description": "Short title of the question" }, + "question": { "type": "string", "description": "The question to ask" }, + "suggested": { "type": "array", "items": { "type": "string" }, + "description": "Optional suggested answers" } + }, + "required": ["question"] + } + } + }) + } + + async fn call(&self, args: Value, ctx: &ToolCtx) -> Result { + let q = Question { + title: args["title"].as_str().unwrap_or("Question").to_string(), + question: args["question"].as_str().unwrap_or("").to_string(), + suggested: args["suggested"] + .as_array() + .map(|a| a.iter().filter_map(|v| v.as_str().map(str::to_string)).collect()) + .unwrap_or_default(), + call: ctx.call_id, + }; + // Durability FIRST: the call must survive a crash as AwaitingHuman. + self.store + .set_call_state(ctx.call_id, CallState::AwaitingHuman) + .await + .map_err(|e| ToolFailure::Failed(format!("ask_user: store error: {e}")))?; + + let events = EventSink::from_extensions(&ctx.extensions) + .ok_or_else(|| ToolFailure::Failed("ask_user: no EventSink in extensions".into()))?; + + match self.channel.ask(q, &events).await { + Ok(answer) => Ok(ToolOutput::Text(answer)), + Err(HumanGone) => Err(ToolFailure::Suspend), + } + } +} diff --git a/crates/agent-loop/src/ids.rs b/crates/agent-loop/src/ids.rs new file mode 100644 index 0000000..3a1a16e --- /dev/null +++ b/crates/agent-loop/src/ids.rs @@ -0,0 +1,54 @@ +//! Opaque id newtypes. The store contract requires `MessageId` and `ToolCallId` +//! to be **monotonically increasing per frame**: a concurrent fan-out allocates +//! ids in call order BEFORE execution, and the model reconstructs results by id. + +use std::fmt; + +/// Identifies a conversation (Skald: `"session:42"`; InMemory: any string). +#[derive(Debug, Clone, PartialEq, Eq, Hash, PartialOrd, Ord)] +pub struct ConversationId(pub String); + +impl ConversationId { + pub fn new(s: impl Into) -> Self { Self(s.into()) } + pub fn as_str(&self) -> &str { &self.0 } +} + +impl fmt::Display for ConversationId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { f.write_str(&self.0) } +} + +impl From<&str> for ConversationId { + fn from(s: &str) -> Self { Self(s.to_string()) } +} +impl From for ConversationId { + fn from(s: String) -> Self { Self(s) } +} + +macro_rules! int_id { + ($name:ident, $doc:literal) => { + #[doc = $doc] + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)] + pub struct $name(pub i64); + + impl $name { + pub fn get(self) -> i64 { self.0 } + } + + impl fmt::Display for $name { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{}", self.0) } + } + + impl From for $name { + fn from(v: i64) -> Self { Self(v) } + } + }; +} + +int_id!(FrameId, "A conversation frame (root frame = the conversation; children = sub-agents)."); +int_id!(MessageId, "A stored message. Monotonically increasing per frame."); +int_id!(ToolCallId, "A stored tool call. Monotonically increasing per frame."); +int_id!(TaskId, "An async delegated task."); +int_id!(SummaryId, "A compaction summary."); + +/// Key of a model inside a `ModelSelector` ("kimi-k3", "claude-sonnet-4", …). +pub type ModelId = String; diff --git a/crates/agent-loop/src/kernel.rs b/crates/agent-loop/src/kernel.rs new file mode 100644 index 0000000..e820674 --- /dev/null +++ b/crates/agent-loop/src/kernel.rs @@ -0,0 +1,584 @@ +//! The kernel — `LlmLoop`. It owns ONLY control flow: round loop, model +//! fallback, tool fan-out, recording. It knows nothing about agents, approval +//! rules, MCP, compaction or recovery (blueprint §5). + +use std::collections::HashMap; +use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; + +use anyhow::anyhow; +use futures::StreamExt as _; +use tokio::sync::mpsc; +use tokio_util::sync::CancellationToken; +use tracing::warn; + +use crate::context::{AssembleInput, ContextAssembler}; +use crate::events::{EventSink, LoopEvent, PendingToolCall}; +use crate::gate::{Gate, GateDecision, PendingCall}; +use crate::hooks::{HookCtx, HookVerdict, LoopHooks}; +use crate::ids::{FrameId, MessageId, ModelId}; +use crate::manager::LoopParams; +use crate::model::{ + ModelHandle, ModelRequest, ModelResponse, ModelSelector, RetryPolicy, StreamDelta, Usage, +}; +use crate::store::{CallOutcome, HistoryStore, NewCall, NewMessage}; +use crate::tool::{ExecutionOutcome, ToolCtx, drive_execution}; + +/// The terminal outcome of a turn. +#[derive(Debug, Clone)] +pub enum TurnOutcome { + Final { + content: String, + message_id: MessageId, + usage: Usage, + reasoning: Option, + }, + Cancelled, + /// Round budget exhausted. + Exhausted, +} + +/// Shared dependencies the manager hands to every loop. +pub(crate) struct KernelDeps { + pub(crate) models: Arc, + pub(crate) store: Arc, + pub(crate) gate: Arc, + pub(crate) hooks: Vec>, + pub(crate) assembler: Arc, + pub(crate) max_rounds: usize, + pub(crate) max_parallel_calls: usize, + pub(crate) retry: RetryPolicy, +} + +static REQUEST_COUNTER: AtomicU64 = AtomicU64::new(0); + +/// Correlation id for host-side payload logging (one per attempt). +fn mint_request_id() -> String { + let n = REQUEST_COUNTER.fetch_add(1, Ordering::Relaxed); + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .map(|d| d.as_nanos()) + .unwrap_or(0); + format!("{nanos:032x}-{n:08x}") +} + +/// Run one loop to completion. Spawned by the manager; the `token` is cloned +/// by value through the whole call tree — never re-read from a field mid-turn. +pub(crate) async fn run( + deps: Arc, + params: LoopParams, + token: CancellationToken, + events: EventSink, +) -> crate::Result { + let frame = params.frame; + let parent = params.parent_frame; + let store = deps.store.clone(); + let assembler = params.assembler.clone().unwrap_or_else(|| deps.assembler.clone()); + + let hook_ctx = || HookCtx { + conversation: params.conversation.clone(), + frame, + agent: params.agent.clone(), + store: store.clone(), + events: events.clone(), + }; + + // ToolCtx extensions: host extensions + the event sink, so shipped tools + // (ask_user, activate_tools) can emit out-of-band. + let tool_extensions = || { + let mut ext = params.extensions.clone(); + ext.insert(Arc::new(events.clone())); + ext + }; + + events.emit(frame, parent, LoopEvent::TurnStarted); + + // First selection of the turn. + let mut handle: ModelHandle = match deps.models.select(¶ms.model_hint, &[]).await { + Ok(h) => h, + Err(e) => { + events.emit(frame, parent, LoopEvent::Error(format!("model selection failed: {e}"))); + return Err(e); + } + }; + + for round in 0..deps.max_rounds { + if token.is_cancelled() { + return finish(TurnOutcome::Cancelled, &deps, &hook_ctx(), &events, frame, parent).await; + } + for h in &deps.hooks { + h.before_round(round, &hook_ctx()).await; + } + events.emit(frame, parent, LoopEvent::RoundStarted { round }); + + // Live input (pull-based, blueprint D10): user messages queued mid-turn. + if let Some(input) = ¶ms.live_input { + for msg in input.drain().await { + let id = store.append(frame, msg.clone()).await?; + events.emit(frame, parent, LoopEvent::UserMessage { + message_id: id, + content: msg.content, + synthetic: msg.synthetic, + metadata: msg.metadata, + }); + } + } + + let turn_info = crate::context::TurnInfo { + conversation: params.conversation.clone(), + frame, + agent: params.agent.clone(), + user_message: params.meta.user_message.clone(), + }; + + let system = params.system.system_context(&turn_info).await?; + let mut messages = assembler + .build(&store, &AssembleInput { + frame, + system: system.clone(), + model: handle.info.clone(), + round, + }) + .await?; + let mut defs = params.tools.defs(&handle.info); + + // ── one LLM call with fallback ── + let mut tried: Vec = vec![handle.id.clone()]; + let response: ModelResponse = loop { + let (delta_tx, forwarder) = spawn_delta_forwarder(&events, frame, parent); + let req = ModelRequest { + messages: messages.clone(), + tools: defs.clone(), + model: handle.id.clone(), + max_tokens: None, + temperature: None, + request_id: mint_request_id(), + conversation: params.conversation.clone(), + frame, + extras: handle.info.extras.clone(), + }; + let result = tokio::select! { + biased; + _ = token.cancelled() => { + drop(forwarder); + return finish(TurnOutcome::Cancelled, &deps, &hook_ctx(), &events, frame, parent).await; + } + r = handle.model.complete(&req, Some(delta_tx)) => r, + }; + // Drain deltas BEFORE the round's outcome events (ordering). + let _ = forwarder.await; + + match result { + Ok(resp) => { + deps.models.report_success(&handle.id).await; + break resp; + } + Err(e) => { + deps.models.report_failure(&handle.id, &e.to_string()).await; + let retriable = handle.model.is_retriable(&e); + warn!(model = %handle.id, error = %e, retriable, "llm call failed"); + if !retriable || tried.len() >= deps.retry.max_attempts { + events.emit(frame, parent, LoopEvent::LlmFailed { + tried: tried.clone(), + last_error: e.to_string(), + }); + return Err(anyhow!("llm call failed on {}: {e}", handle.id)); + } + match deps.models.select(¶ms.model_hint, &tried).await { + Ok(next) => { + events.emit(frame, parent, LoopEvent::ModelFallback { + from: handle.id.clone(), + to: next.id.clone(), + reason: e.to_string(), + }); + handle = next; + tried.push(handle.id.clone()); + // Rebuild for the new model: prompt_cache / + // capabilities / DTL mode may differ. + messages = assembler + .build(&store, &AssembleInput { + frame, + system: system.clone(), + model: handle.info.clone(), + round, + }) + .await?; + defs = params.tools.defs(&handle.info); + } + Err(sel_err) => { + events.emit(frame, parent, LoopEvent::LlmFailed { + tried: tried.clone(), + last_error: format!("{e}; no fallback: {sel_err}"), + }); + return Err(anyhow!("llm call failed on {} and no fallback: {e}", handle.id)); + } + } + } + } + }; + + match response { + ModelResponse::Message { content, reasoning, usage, .. } => { + let id = store + .append(frame, NewMessage::assistant(content.clone(), reasoning.clone())) + .await?; + store.set_usage(id, &usage).await?; + if usage.truncated { + events.emit(frame, parent, LoopEvent::Truncated { output_tokens: usage.output_tokens }); + } + events.emit(frame, parent, LoopEvent::Done { + message_id: id, + content: content.clone(), + usage: usage.clone(), + reasoning: reasoning.clone(), + }); + let outcome = TurnOutcome::Final { content, message_id: id, usage, reasoning }; + return finish(outcome, &deps, &hook_ctx(), &events, frame, parent).await; + } + ModelResponse::ToolCalls { content, calls, reasoning, usage, .. } => { + let msg_id = store + .append(frame, NewMessage::assistant(content.clone(), reasoning.clone())) + .await?; + store.set_usage(msg_id, &usage).await?; + if !content.is_empty() || usage.is_present() { + events.emit(frame, parent, LoopEvent::Thinking { + message_id: msg_id, + content, + usage, + reasoning, + }); + } + + let fan_out = + calls.len() >= 2 && calls.iter().all(|c| { + params + .tools + .find(&c.name) + .is_some_and(|t| t.concurrency_safe(&c.arguments)) + }); + + if fan_out { + if let Some(outcome) = run_fan_out( + &deps, ¶ms, &events, &token, msg_id, &calls, tool_extensions(), + ) + .await? + { + return finish(outcome, &deps, &hook_ctx(), &events, frame, parent).await; + } + } else if let Some(outcome) = run_sequential( + &deps, ¶ms, &events, &token, msg_id, &calls, tool_extensions(), + ) + .await? + { + return finish(outcome, &deps, &hook_ctx(), &events, frame, parent).await; + } + } + } + + for h in &deps.hooks { + h.after_round(round, &hook_ctx()).await; + } + } + + finish(TurnOutcome::Exhausted, &deps, &hook_ctx(), &events, frame, parent).await +} + +/// Terminal helper: hooks.on_turn_end (+ Cancelled event) then return. +async fn finish( + outcome: TurnOutcome, + deps: &Arc, + ctx: &HookCtx, + events: &EventSink, + frame: FrameId, + parent: Option, +) -> crate::Result { + if matches!(outcome, TurnOutcome::Cancelled) { + events.emit(frame, parent, LoopEvent::Cancelled); + } + for h in &deps.hooks { + h.on_turn_end(&outcome, ctx).await; + } + Ok(outcome) +} + +/// Map streamed deltas to bus events; drained before the round's outcomes. +fn spawn_delta_forwarder( + events: &EventSink, + frame: FrameId, + parent: Option, +) -> (mpsc::Sender, tokio::task::JoinHandle<()>) { + let (tx, mut rx) = mpsc::channel::(256); + let events = events.clone(); + let handle = tokio::spawn(async move { + while let Some(delta) = rx.recv().await { + let (kind, text) = match delta { + StreamDelta::Text(t) => (crate::events::DeltaKind::Content, t), + StreamDelta::Reasoning(t) => (crate::events::DeltaKind::Reasoning, t), + }; + events.emit(frame, parent, LoopEvent::TokenDelta { kind, text }); + } + }); + (tx, handle) +} + +/// Sequential tool-call path (a lone call, or any mixed batch). Returns +/// `Ok(Some(outcome))` when the turn must end (cancel/suspend). +async fn run_sequential( + deps: &Arc, + params: &LoopParams, + events: &EventSink, + token: &CancellationToken, + msg_id: MessageId, + calls: &[crate::model::ToolCall], + ext: crate::tool::Extensions, +) -> crate::Result> { + let store = deps.store.clone(); + for call in calls { + if token.is_cancelled() { + return Ok(Some(TurnOutcome::Cancelled)); + } + let ptc = record_call(&store, events, params, msg_id, call).await?; + + let pre = pre_execution(deps, params, events, token, &ptc).await?; + let tool = match pre { + PreExecution::Run(tool) => tool, + PreExecution::Resolved(outcome) => { + record_outcome(deps, params, events, &store, &ptc, outcome).await?; + continue; + } + PreExecution::TurnCancelled => return Ok(Some(TurnOutcome::Cancelled)), + }; + + let ctx = ToolCtx { + conversation: params.conversation.clone(), + frame: params.frame, + agent: params.agent.clone(), + call_id: ptc.id, + cancel: token.clone(), + extensions: ext.clone(), + }; + let exec = tool.start(ptc.arguments.clone(), &ctx); + match drive_execution(&*exec, token).await { + ExecutionOutcome::Suspended => { + // The call STAYS AwaitingHuman (the tool marked it) — no resolve. + return Ok(Some(TurnOutcome::Cancelled)); + } + outcome => { + record_outcome(deps, params, events, &store, &ptc, outcome.into_call_outcome()) + .await?; + } + } + } + Ok(None) +} + +/// The concurrent fan-out (generalized sub-agent batch, blueprint §5): ids +/// allocated in order (phase 1), execution concurrent and bounded (phase 2), +/// recording in order (phase 3). +async fn run_fan_out( + deps: &Arc, + params: &LoopParams, + events: &EventSink, + token: &CancellationToken, + msg_id: MessageId, + calls: &[crate::model::ToolCall], + ext: crate::tool::Extensions, +) -> crate::Result> { + let store = deps.store.clone(); + + // ── Phase 1: sequential, in call order ── + let mut ptcs = Vec::with_capacity(calls.len()); + for call in calls { + ptcs.push(record_call(&store, events, params, msg_id, call).await?); + } + + // ── Phase 2: concurrent, bounded ── + let futs: Vec<_> = ptcs + .iter() + .enumerate() + .map(|(idx, ptc)| phase2_one(deps, params, events, token.clone(), ext.clone(), idx, ptc)) + .collect(); + let results: HashMap = futures::stream::iter(futs) + .buffer_unordered(deps.max_parallel_calls.max(1)) + .collect() + .await; + + // ── Phase 3: sequential, in call order ── + let mut suspended = false; + for (idx, ptc) in ptcs.iter().enumerate() { + match results.get(&idx) { + Some(Phase2::Suspended) => { + // Stays AwaitingHuman; the turn ends after recording the rest. + suspended = true; + } + Some(Phase2::Done(outcome)) => { + record_outcome(deps, params, events, &store, ptc, outcome.clone()).await?; + } + None => { + record_outcome( + deps, params, events, &store, ptc, + CallOutcome::Failed("internal: fan-out result missing".into()), + ) + .await?; + } + } + } + + if suspended { + return Ok(Some(TurnOutcome::Cancelled)); + } + if token.is_cancelled() { + return Ok(Some(TurnOutcome::Cancelled)); + } + Ok(None) +} + +enum Phase2 { + Done(CallOutcome), + Suspended, +} + +/// One fanned-out call: gate → hooks.pre → execute. An explicit async fn (not +/// a closure) so the futures are uniform and the borrows are higher-ranked. +async fn phase2_one<'a>( + deps: &'a Arc, + params: &'a LoopParams, + events: &'a EventSink, + token: CancellationToken, + ext: crate::tool::Extensions, + idx: usize, + ptc: &'a PendingToolCall, +) -> (usize, Phase2) { + let phase = match pre_execution(deps, params, events, &token, ptc).await { + Ok(PreExecution::Run(tool)) => { + let ctx = ToolCtx { + conversation: params.conversation.clone(), + frame: params.frame, + agent: params.agent.clone(), + call_id: ptc.id, + cancel: token.clone(), + extensions: ext, + }; + let exec = tool.start(ptc.arguments.clone(), &ctx); + match drive_execution(&*exec, &token).await { + ExecutionOutcome::Suspended => Phase2::Suspended, + outcome => Phase2::Done(outcome.into_call_outcome()), + } + } + Ok(PreExecution::Resolved(outcome)) => Phase2::Done(outcome), + Ok(PreExecution::TurnCancelled) => Phase2::Done(CallOutcome::Cancelled), + Err(e) => Phase2::Done(CallOutcome::Failed(format!("pre-execution error: {e}"))), + }; + (idx, phase) +} + +/// Phase-1 shared by both paths: allocate the id and emit `ToolCallStarted`. +async fn record_call( + store: &Arc, + events: &EventSink, + params: &LoopParams, + msg_id: MessageId, + call: &crate::model::ToolCall, +) -> crate::Result { + let id = store + .append_call(msg_id, NewCall { + provider_id: if call.id.is_empty() { None } else { Some(call.id.clone()) }, + name: call.name.clone(), + arguments: call.arguments.clone(), + }) + .await?; + events.emit(params.frame, params.parent_frame, LoopEvent::ToolCallStarted { + id, + message_id: msg_id, + name: call.name.clone(), + args: call.arguments.clone(), + }); + Ok(PendingToolCall { + id, + message_id: msg_id, + provider_id: Some(call.id.clone()).filter(|s| !s.is_empty()), + name: call.name.clone(), + arguments: call.arguments.clone(), + }) +} + +enum PreExecution { + Run(Arc), + Resolved(CallOutcome), + TurnCancelled, +} + +/// Gate + hooks.pre + tool lookup — shared by sequential and fan-out paths. +async fn pre_execution( + deps: &Arc, + params: &LoopParams, + events: &EventSink, + token: &CancellationToken, + ptc: &PendingToolCall, +) -> crate::Result { + let pending = PendingCall { + id: ptc.id, + name: ptc.name.clone(), + args: ptc.arguments.clone(), + frame: params.frame, + agent: params.agent.clone(), + extensions: params.extensions.clone(), + }; + let decision = tokio::select! { + biased; + _ = token.cancelled() => return Ok(PreExecution::TurnCancelled), + d = deps.gate.check(&pending, events) => d, + }; + if let GateDecision::Reject { reason } = decision { + return Ok(PreExecution::Resolved(CallOutcome::Rejected { reason })); + } + + let mut ptc_mut = ptc.clone(); + let hook_ctx = HookCtx { + conversation: params.conversation.clone(), + frame: params.frame, + agent: params.agent.clone(), + store: deps.store.clone(), + events: events.clone(), + }; + for h in &deps.hooks { + if let HookVerdict::Reject { reason } = h.pre_tool_call(&mut ptc_mut, &hook_ctx).await { + return Ok(PreExecution::Resolved(CallOutcome::Rejected { reason })); + } + } + + match params.tools.find(&ptc.name) { + Some(tool) => Ok(PreExecution::Run(tool)), + None => Ok(PreExecution::Resolved(CallOutcome::Failed(format!( + "unknown tool '{}' (not in this turn's tool set)", + ptc.name + )))), + } +} + +/// Phase-3 shared by both paths: hooks.post → resolve → emit. +async fn record_outcome( + deps: &Arc, + params: &LoopParams, + events: &EventSink, + store: &Arc, + ptc: &PendingToolCall, + outcome: CallOutcome, +) -> crate::Result<()> { + let hook_ctx = HookCtx { + conversation: params.conversation.clone(), + frame: params.frame, + agent: params.agent.clone(), + store: store.clone(), + events: events.clone(), + }; + for h in &deps.hooks { + h.post_tool_call(ptc, &outcome, &hook_ctx).await; + } + store.resolve_call(ptc.id, &outcome).await?; + events.emit(params.frame, params.parent_frame, LoopEvent::ToolCallFinished { + id: ptc.id, + outcome, + }); + Ok(()) +} diff --git a/crates/agent-loop/src/lib.rs b/crates/agent-loop/src/lib.rs new file mode 100644 index 0000000..62fdea5 --- /dev/null +++ b/crates/agent-loop/src/lib.rs @@ -0,0 +1,73 @@ +//! `agent-loop` — a reusable LLM agent-loop kernel. +//! +//! The crate owns the **control flow** of a tool-calling agent loop (round loop, +//! model fallback, parallel tool fan-out, streaming deltas, cancellation) and the +//! **LLM clients + protocols** (OpenAI-compatible, Anthropic, Ollama, LM Studio; +//! SSE; dynamic tool loading wire semantics). It knows nothing about databases, +//! agents, MCP, approval rules or Docker: the host implements the trait surface +//! (`Model`, `ModelSelector`, `HistoryStore`, `ContextAssembler`, +//! `SystemContextSource`, `Tool`, `ToolSet`, `Gate`, `LoopHooks`, `HumanChannel`, +//! `ActivationSource`, `ToolActivator`) or uses the shipped defaults. +//! +//! Design document: `blueprint/project-loop.md` (Skald workspace). + +pub mod activation; +pub mod context; +pub mod events; +pub mod gate; +pub mod hooks; +pub mod human; +pub mod ids; +pub mod kernel; +pub mod manager; +pub mod model; +pub mod models; +pub mod store; +pub mod store_memory; +pub mod testing; +pub mod tool; + +/// Application name sent as the `X-Title` header by the shipped clients +/// (OpenRouter rankings). Clients accept an override. +pub const APP_NAME: &str = "Skald"; + +/// Crate-wide result type for host-implemented traits. +pub type Result = anyhow::Result; + +pub mod prelude { + pub use crate::activation::{ + ActivateToolsTool, Activation, ActivationSource, ToolActivator, ToolRendering, + }; + pub use crate::context::{ + AssembleInput, ContextAssembler, LinearAssembler, StaticSystemContext, SystemContext, + SystemContextSource, TurnInfo, + }; + pub use crate::events::{DeltaKind, Event, EventSink, LoopEvent}; + pub use crate::gate::{AllowAll, DenyList, Gate, GateDecision, PendingCall}; + pub use crate::hooks::{HookCtx, HookVerdict, LoopHooks}; + pub use crate::human::{AskUserTool, HumanChannel, HumanGone, Question}; + pub use crate::ids::{ + ConversationId, FrameId, MessageId, ModelId, SummaryId, TaskId, ToolCallId, + }; + pub use crate::manager::{ + LiveInput, LoopManager, LoopManagerBuilder, LoopParams, StartError, TurnHandle, TurnMeta, + TurnParams, + }; + pub use crate::model::{ + Model, ModelError, ModelHandle, ModelHint, ModelInfo, ModelRequest, ModelResponse, + ModelSelector, RawMeta, RetryPolicy, SingleModel, StaticModels, StreamDelta, ToolCall, + Usage, + }; + pub use crate::store::{ + CallOutcome, CallState, FrameRecord, FrameSpec, HistoryStore, NewCall, NewMessage, + NewSummary, Role, StoredCall, StoredMessage, StoredSummary, + }; + pub use crate::tool::{ + Extensions, MediaRef, RestartHint, SimpleExecution, Tool, ToolCtx, ToolExecution, + ToolFailure, ToolOutput, ToolSet, Visibility, drive_execution, + }; + pub use crate::{APP_NAME, Result}; + pub use async_trait::async_trait; + pub use serde_json::{Value, json}; + pub use tokio_util::sync::CancellationToken; +} diff --git a/crates/agent-loop/src/manager.rs b/crates/agent-loop/src/manager.rs new file mode 100644 index 0000000..5ba3384 --- /dev/null +++ b/crates/agent-loop/src/manager.rs @@ -0,0 +1,411 @@ +//! `LoopManager` — the singleton (per tenant/user) that owns the event bus and +//! the registry of live loops, and spawns disposable `LlmLoop`s (blueprint D1). +//! +//! Policy: **one live loop per conversation** — `start_turn` rejects a second +//! one (anti double-driving). Serialization/queueing of user messages stays +//! with the host. + +use std::collections::HashMap; +use std::sync::{Arc, Mutex}; + +use async_trait::async_trait; +use tokio::sync::broadcast; +use tokio::task::JoinHandle; +use tokio_util::sync::CancellationToken; + +use crate::context::{ContextAssembler, LinearAssembler, SystemContextSource}; +use crate::events::{Event, EventSink, LoopEvent}; +use crate::gate::{AllowAll, Gate}; +use crate::hooks::LoopHooks; +use crate::human::HumanChannel; +use crate::ids::{ConversationId, FrameId}; +use crate::kernel::{KernelDeps, TurnOutcome}; +use crate::model::{ModelHint, ModelSelector, RetryPolicy}; +use crate::store::{FrameSpec, HistoryStore, NewMessage, Role}; +use crate::tool::{Extensions, ToolSet}; + +// ── LiveInput ──────────────────────────────────────────────────────────────── + +/// Pull-based live user input (blueprint D10): drained at round boundaries. +#[async_trait] +pub trait LiveInput: Send + Sync { + async fn drain(&self) -> Vec; +} + +// ── TurnMeta ───────────────────────────────────────────────────────────────── + +/// Per-turn metadata. +#[derive(Debug, Clone, Default)] +pub struct TurnMeta { + /// Synthetic turn (TIC/notify) — no user echo semantics. + pub synthetic: bool, + /// Interactive surface (web chat, telegram, …). + pub interactive: bool, + /// Label for UI/logging ("session 42", "cron job X"). + pub context_label: Option, + /// The user message that opened the turn (for `TurnInfo`). + pub user_message: Option, +} + +// ── TurnParams / LoopParams ────────────────────────────────────────────────── + +/// Parameters of a user turn (root frame). +pub struct TurnParams { + /// Root frame (opened by the host or via `LoopManager::open_root`). + pub frame: FrameId, + pub agent: String, + pub system: Arc, + /// Already filtered (visibility/approval). + pub tools: Arc, + pub model_hint: ModelHint, + /// None for sub-agents / cron / resume. + pub live_input: Option>, + /// Flows into `ToolCtx.extensions`. + pub extensions: Extensions, + pub meta: TurnMeta, + /// Per-turn assembler override (default: the manager's). + pub assembler: Option>, +} + +/// Parameters of a raw loop (DelegateTool, recovery, background runners). +pub struct LoopParams { + pub conversation: ConversationId, + pub frame: FrameId, + pub parent_frame: Option, + pub agent: String, + pub system: Arc, + pub tools: Arc, + pub model_hint: ModelHint, + pub live_input: Option>, + pub extensions: Extensions, + pub meta: TurnMeta, + pub assembler: Option>, +} + +// ── TurnHandle ─────────────────────────────────────────────────────────────── + +/// Handle of a spawned turn. +pub struct TurnHandle { + pub conversation: ConversationId, + pub frame: FrameId, + /// Clone; cancels THIS turn (sticky down the whole call tree). + pub cancel: CancellationToken, + join: JoinHandle>, +} + +impl TurnHandle { + pub async fn join(self) -> crate::Result { + self.join.await.map_err(|e| anyhow::anyhow!("loop task panicked: {e}"))? + } +} + +// ── StartError ─────────────────────────────────────────────────────────────── + +#[derive(Debug)] +pub enum StartError { + /// A loop is already live on this conversation (anti double-driving). + AlreadyRunning, + Store(anyhow::Error), +} + +impl std::fmt::Display for StartError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::AlreadyRunning => write!(f, "a loop is already running on this conversation"), + Self::Store(e) => write!(f, "store error: {e}"), + } + } +} +impl std::error::Error for StartError {} + +// ── RunningInfo ────────────────────────────────────────────────────────────── + +#[derive(Debug, Clone)] +pub struct RunningInfo { + pub conversation: ConversationId, + pub frame: FrameId, + pub agent: String, +} + +struct RunningEntry { + frame: FrameId, + agent: String, + cancel: CancellationToken, +} + +// ── LoopManager ────────────────────────────────────────────────────────────── + +pub struct LoopManager { + deps: Arc, + bus: broadcast::Sender>, + registry: Arc>>, + human: Option>, +} + +impl LoopManager { + pub fn builder() -> LoopManagerBuilder { LoopManagerBuilder::default() } + + /// Subscribe to the global event bus (every event tagged with + /// conversation/frame/parent_frame). + pub fn events(&self) -> broadcast::Receiver> { self.bus.subscribe() } + + /// The host-provided human channel, if any. + pub fn human(&self) -> Option> { self.human.clone() } + + /// Convenience: open a root frame on the store. + pub async fn open_root(&self, conv: &ConversationId, spec: FrameSpec) -> crate::Result { + self.deps.store.open_frame(conv, None, spec).await + } + + pub fn store(&self) -> Arc { self.deps.store.clone() } + + // ── user turns ── + + /// High-level entry point: + /// 1. rejects when a loop is already live on the conversation; + /// 2. marks a trailing orphan User/Agent message failed (alternation rule + /// for strict APIs); + /// 3. appends the user message + echo event; + /// 4. spawns the loop; returns the handle immediately. + pub async fn start_turn( + &self, + conv: ConversationId, + msg: NewMessage, + mut params: TurnParams, + ) -> Result { + { + let registry = self.registry.lock().unwrap(); + if registry.contains_key(&conv) { + return Err(StartError::AlreadyRunning); + } + } + + // Orphan rule: a trailing User/Agent message with no assistant reply + // breaks strict alternation — mark it failed before appending. + if let Some(last) = self.deps.store.last(params.frame).await.map_err(StartError::Store)? + && matches!(last.role, Role::User | Role::Agent) + { + self.deps.store.mark_failed(last.id).await.map_err(StartError::Store)?; + } + + let events = self.sink(conv.clone()); + let id = self.deps.store.append(params.frame, msg.clone()).await.map_err(StartError::Store)?; + events.emit(params.frame, None, LoopEvent::UserMessage { + message_id: id, + content: msg.content.clone(), + synthetic: msg.synthetic, + metadata: msg.metadata.clone(), + }); + + params.meta.user_message = Some(msg.content); + self.spawn(LoopParams { + conversation: conv, + frame: params.frame, + parent_frame: None, + agent: params.agent, + system: params.system, + tools: params.tools, + model_hint: params.model_hint, + live_input: params.live_input, + extensions: params.extensions, + meta: params.meta, + assembler: params.assembler, + }) + } + + // ── raw loops (DelegateTool, recovery, background runners) ── + + pub async fn start_loop(&self, params: LoopParams) -> Result { + { + let registry = self.registry.lock().unwrap(); + if registry.contains_key(¶ms.conversation) { + return Err(StartError::AlreadyRunning); + } + } + self.spawn(params) + } + + fn spawn(&self, params: LoopParams) -> Result { + let conv = params.conversation.clone(); + let frame = params.frame; + let agent = params.agent.clone(); + let token = CancellationToken::new(); + let events = self.sink(conv.clone()); + + { + let mut registry = self.registry.lock().unwrap(); + registry.insert(conv.clone(), RunningEntry { + frame, + agent, + cancel: token.clone(), + }); + } + + let deps = self.deps.clone(); + let registry = self.registry.clone(); + let turn_token = token.clone(); + let join_conv = conv.clone(); + let join = tokio::spawn(async move { + let outcome = crate::kernel::run(deps, params, turn_token, events).await; + registry.lock().unwrap().remove(&join_conv); + outcome + }); + + Ok(TurnHandle { conversation: conv, frame, cancel: token, join }) + } + + // ── control ── + + /// `/stop`: cancel the live loop on a conversation, if any. + pub fn cancel(&self, conv: &ConversationId) { + if let Some(entry) = self.registry.lock().unwrap().get(conv) { + entry.cancel.cancel(); + } + } + + pub fn is_running(&self, conv: &ConversationId) -> bool { + self.registry.lock().unwrap().contains_key(conv) + } + + /// Global view (UI "running agents"). + pub fn list_running(&self) -> Vec { + self.registry + .lock() + .unwrap() + .iter() + .map(|(conversation, e)| RunningInfo { + conversation: conversation.clone(), + frame: e.frame, + agent: e.agent.clone(), + }) + .collect() + } + + /// Cancel all live loops. Joins are detached — callers wanting a drain + /// should hold the handles. + pub async fn shutdown(&self) { + let tokens: Vec = self + .registry + .lock() + .unwrap() + .values() + .map(|e| e.cancel.clone()) + .collect(); + for t in tokens { + t.cancel(); + } + } + + fn sink(&self, conv: ConversationId) -> EventSink { + EventSink::new(conv, self.bus.clone()) + } +} + +// ── Builder ────────────────────────────────────────────────────────────────── + +pub struct LoopManagerBuilder { + models: Option>, + store: Option>, + gate: Option>, + hooks: Vec>, + human: Option>, + assembler: Option>, + max_rounds: usize, + max_parallel_calls: usize, + retry: RetryPolicy, + bus_capacity: usize, +} + +impl Default for LoopManagerBuilder { + fn default() -> Self { + Self { + models: None, + store: None, + gate: None, + hooks: Vec::new(), + human: None, + assembler: None, + max_rounds: 20, + max_parallel_calls: 4, + retry: RetryPolicy::default(), + bus_capacity: 512, + } + } +} + +impl LoopManagerBuilder { + pub fn models(mut self, models: Arc) -> Self { + self.models = Some(models); + self + } + + pub fn store(mut self, store: Arc) -> Self { + self.store = Some(store); + self + } + + pub fn gate(mut self, gate: impl Gate + 'static) -> Self { + self.gate = Some(Arc::new(gate)); + self + } + + pub fn gate_arc(mut self, gate: Arc) -> Self { + self.gate = Some(gate); + self + } + + pub fn hook(mut self, hook: Arc) -> Self { + self.hooks.push(hook); + self + } + + pub fn human(mut self, human: Arc) -> Self { + self.human = Some(human); + self + } + + pub fn assembler(mut self, assembler: Arc) -> Self { + self.assembler = Some(assembler); + self + } + + pub fn max_rounds(mut self, n: usize) -> Self { + self.max_rounds = n; + self + } + + pub fn max_parallel_calls(mut self, n: usize) -> Self { + self.max_parallel_calls = n; + self + } + + pub fn retry(mut self, retry: RetryPolicy) -> Self { + self.retry = retry; + self + } + + pub fn bus_capacity(mut self, n: usize) -> Self { + self.bus_capacity = n; + self + } + + pub fn build(self) -> crate::Result { + let deps = Arc::new(KernelDeps { + models: self.models.ok_or_else(|| anyhow::anyhow!("LoopManager: models required"))?, + store: self.store.ok_or_else(|| anyhow::anyhow!("LoopManager: store required"))?, + gate: self.gate.unwrap_or_else(|| Arc::new(AllowAll)), + hooks: self.hooks, + assembler: self.assembler.unwrap_or_else(|| Arc::new(LinearAssembler::new())), + max_rounds: self.max_rounds, + max_parallel_calls: self.max_parallel_calls, + retry: self.retry, + }); + let (bus, _) = broadcast::channel(self.bus_capacity); + Ok(LoopManager { + deps, + bus, + registry: Arc::new(Mutex::new(HashMap::new())), + human: self.human, + }) + } +} diff --git a/crates/agent-loop/src/model.rs b/crates/agent-loop/src/model.rs new file mode 100644 index 0000000..004b3d7 --- /dev/null +++ b/crates/agent-loop/src/model.rs @@ -0,0 +1,440 @@ +//! The `Model` trait (a stateless LLM client), the `ModelSelector` seam +//! (selection + health), and the shipped selectors. +//! +//! `Model` is the boundary the kernel talks to; the shipped clients live in +//! [`crate::models`]. The wire format at this boundary is OpenAI-shaped +//! `serde_json::Value` (blueprint D4) — the Anthropic client translates +//! internally. + +use std::sync::Arc; +use std::sync::atomic::{AtomicUsize, Ordering}; + +use async_trait::async_trait; +use serde_json::Value; +use tokio::sync::mpsc; + +use crate::activation::ToolRendering; +use crate::ids::{ConversationId, FrameId, ModelId}; + +// ── Usage ──────────────────────────────────────────────────────────────────── + +/// Token/cost accounting of one model call. All fields optional: providers +/// report different subsets (or nothing, e.g. Ollama cost). +#[derive(Debug, Default, Clone)] +pub struct Usage { + pub input_tokens: Option, + pub output_tokens: Option, + pub cache_read: Option, + pub cache_write: Option, + pub cost_usd: Option, + /// The model stopped at the token limit (`finish_reason == "length"` / + /// `stop_reason == "max_tokens"`). + pub truncated: bool, +} + +impl Usage { + pub fn is_present(&self) -> bool { + self.input_tokens.is_some() || self.output_tokens.is_some() + } +} + +// ── ToolCall ───────────────────────────────────────────────────────────────── + +/// A tool call requested by the model (wire level). +#[derive(Debug, Clone)] +pub struct ToolCall { + /// The provider's call id ("call_abc", "toolu_01…"). May be empty for + /// providers that don't assign one — the assembler then synthesizes one. + pub id: String, + pub name: String, + pub arguments: Value, +} + +// ── StreamDelta ────────────────────────────────────────────────────────────── + +/// An incremental piece of a streaming completion. Best-effort UI feedback: +/// senders use `try_send` and drop deltas when the channel is full — streaming +/// must never backpressure the HTTP read. The returned [`ModelResponse`] +/// remains the only authoritative result. +#[derive(Debug, Clone)] +pub enum StreamDelta { + Text(String), + Reasoning(String), +} + +// ── RawMeta ────────────────────────────────────────────────────────────────── + +/// Raw HTTP metadata captured during a provider call, for host-side payload +/// logging (a `LoggingModel` decorator persists it). Sensitive header values +/// are redacted by the clients before capture. +#[derive(Debug, Default, Clone)] +pub struct RawMeta { + pub request_headers: Option, + pub request_body: Option, + pub response_headers: Option, + pub response_body: Option, +} + +// ── ModelResponse ──────────────────────────────────────────────────────────── + +/// The authoritative outcome of one model call. +#[derive(Debug, Clone)] +pub enum ModelResponse { + Message { + content: String, + reasoning: Option, + usage: Usage, + raw: Option, + }, + ToolCalls { + content: String, + calls: Vec, + reasoning: Option, + usage: Usage, + raw: Option, + }, +} + +impl ModelResponse { + pub fn message(content: impl Into) -> Self { + Self::Message { content: content.into(), reasoning: None, usage: Usage::default(), raw: None } + } + + pub fn tool_calls(content: impl Into, calls: Vec) -> Self { + Self::ToolCalls { content: content.into(), calls, reasoning: None, usage: Usage::default(), raw: None } + } + + pub fn usage(&self) -> &Usage { + match self { + Self::Message { usage, .. } | Self::ToolCalls { usage, .. } => usage, + } + } + + pub fn usage_mut(&mut self) -> &mut Usage { + match self { + Self::Message { usage, .. } | Self::ToolCalls { usage, .. } => usage, + } + } + + pub fn content(&self) -> &str { + match self { + Self::Message { content, .. } | Self::ToolCalls { content, .. } => content, + } + } + + pub fn reasoning(&self) -> Option<&str> { + match self { + Self::Message { reasoning, .. } | Self::ToolCalls { reasoning, .. } => { + reasoning.as_deref() + } + } + } + + pub fn raw(&self) -> Option<&RawMeta> { + match self { + Self::Message { raw, .. } | Self::ToolCalls { raw, .. } => raw.as_ref(), + } + } +} + +// ── ModelError ─────────────────────────────────────────────────────────────── + +/// A structured model-call failure. The HTTP status lives in the type, never +/// in a substring of the message — a model id or token count containing +/// "404" must not mis-classify retriability. +#[derive(Debug, Clone)] +pub struct ModelError { + /// HTTP status, when the failure came from an HTTP response. `None` for + /// network/parse/cancellation failures — callers treat those as retriable. + pub status: Option, + pub message: String, + /// Request/response payload captured at the failing call, so the host's + /// debug log can show what was actually sent even when the provider + /// rejected it. `None` when there was no HTTP round-trip. + pub raw: Option, +} + +impl ModelError { + pub fn new(status: Option, message: impl Into) -> Self { + Self { status, message: message.into(), raw: None } + } + + pub fn with_raw(mut self, raw: RawMeta) -> Self { + self.raw = Some(raw); + self + } + + pub fn from_reqwest(err: reqwest::Error) -> Self { + let status = err.status().map(|s| s.as_u16()); + Self { status, message: err.to_string(), raw: None } + } +} + +impl std::fmt::Display for ModelError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self.status { + Some(s) => write!(f, "[HTTP {s}] {}", self.message), + None => f.write_str(&self.message), + } + } +} + +impl std::error::Error for ModelError {} + +// ── ModelRequest ───────────────────────────────────────────────────────────── + +/// One model call. `messages`/`tools` are OpenAI-shaped wire values (D4). +#[derive(Debug, Clone)] +pub struct ModelRequest { + pub messages: Vec, + pub tools: Vec, + /// Concrete model name ("kimi-k3", "claude-sonnet-4-5", …). + pub model: String, + pub max_tokens: Option, + pub temperature: Option, + /// Correlation id minted by the kernel at every attempt — for host-side + /// logging/telemetry only, ignored by the kernel itself. + pub request_id: String, + pub conversation: ConversationId, + pub frame: FrameId, + /// Host free-form per-request extras (e.g. reasoning knobs resolved for + /// this model). Merged last by the shipped clients. + pub extras: Value, +} + +// ── Model ──────────────────────────────────────────────────────────────────── + +/// A stateless LLM client. Implementations hold only connection config (base +/// URL, API key). No memory, no database, no session state. +#[async_trait] +pub trait Model: Send + Sync { + /// One completion. `deltas` is a best-effort side-channel for streaming: + /// implementations push [`StreamDelta`]s via `try_send` and never block on + /// it. The returned [`ModelResponse`] is the only authoritative result. + /// + /// Shipped clients retry the call buffered when the stream fails before + /// any delta was emitted (providers rejecting `stream` keep working); a + /// mid-stream failure propagates to the caller's fallback logic. + async fn complete( + &self, + req: &ModelRequest, + deltas: Option>, + ) -> Result; + + /// Retriability classification **for this model**. Default — the crate + /// owns the protocols (blueprint D13): 401/403/404/422 are NOT retriable; + /// 400/429/5xx and status-less failures (network, parse, cancel) are. + /// Hosts may override via a wrapping `Model`. + fn is_retriable(&self, err: &ModelError) -> bool { + !matches!(err.status, Some(401 | 403 | 404 | 422)) + } +} + +// ── ModelInfo / ModelHandle ────────────────────────────────────────────────── + +/// Metadata influencing build/serialization. Read by assemblers and `ToolSet`, +/// NEVER interpreted by the kernel (it passes them through). +#[derive(Debug, Clone, Default)] +pub struct ModelInfo { + /// Anthropic-style prompt-cache hints. + pub prompt_cache: bool, + /// "vision", "video", "tool_search", … + pub capabilities: Vec, + /// Dynamic-tool-loading wire protocol (blueprint §4.10). Default `Inline`. + pub tool_rendering: ToolRendering, + /// Host free-form (Skald: context_length, extra_params). + pub extras: Value, +} + +impl ModelInfo { + pub fn has_capability(&self, cap: &str) -> bool { + self.capabilities.iter().any(|c| c == cap) + } +} + +/// A selected model plus its metadata, as returned by a `ModelSelector`. +#[derive(Clone)] +pub struct ModelHandle { + pub id: ModelId, + pub model: Arc, + pub info: ModelInfo, +} + +// ── ModelHint ──────────────────────────────────────────────────────────────── + +/// Selection hint: only the explicit pin (blueprint D14). Strength/tiering/ +/// priority are host logic, resolved inside the host's `ModelSelector`. +#[derive(Debug, Clone, Default)] +pub struct ModelHint { + /// Explicit model pin — bypasses the host's AUTO selection. + pub name: Option, +} + +impl ModelHint { + pub fn name(name: impl Into) -> Self { + Self { name: Some(name.into()) } + } +} + +// ── ModelSelector ──────────────────────────────────────────────────────────── + +/// The selection seam. The kernel calls `select` once per round and again on +/// every fallback (`exclude` = models already tried in this round). +#[async_trait] +pub trait ModelSelector: Send + Sync { + async fn select(&self, hint: &ModelHint, exclude: &[ModelId]) -> crate::Result; + + /// Health reporting — default no-op. Hosts back these with circuit + /// breakers / status dashboards (Skald: LlmManager mark_success/failure). + async fn report_success(&self, _id: &ModelId) {} + async fn report_failure(&self, _id: &ModelId, _err: &str) {} +} + +// ── RetryPolicy ────────────────────────────────────────────────────────────── + +/// Fallback budget per round: how many DISTINCT models to try before +/// `LlmFailed`. Retriability classification lives on `Model::is_retriable`. +#[derive(Debug, Clone, Copy)] +pub struct RetryPolicy { + pub max_attempts: usize, +} + +impl Default for RetryPolicy { + fn default() -> Self { Self { max_attempts: 3 } } +} + +// ── Shipped selectors ──────────────────────────────────────────────────────── + +/// One model, no fallback. Pair it with a shipped client +/// (`models::OpenAiModel::new(...)`) for a complete agent in ~50 lines. +pub struct SingleModel { + handle: ModelHandle, +} + +impl SingleModel { + pub fn new(model: impl NamedModel) -> Self { + Self { handle: model.into_handle() } + } + + pub fn with_info(model: impl NamedModel, info: ModelInfo) -> Self { + let mut handle = model.into_handle(); + handle.info = info; + Self { handle } + } + + pub fn from_handle(handle: ModelHandle) -> Self { Self { handle } } +} + +#[async_trait] +impl ModelSelector for SingleModel { + async fn select(&self, _hint: &ModelHint, _exclude: &[ModelId]) -> crate::Result { + Ok(self.handle.clone()) + } +} + +/// A model with a self-assigned selector id — implemented by every shipped +/// client (the id defaults to the client's `default_model()`). +pub trait NamedModel: Model + 'static { + /// Selector id and default wire model name for this client. + fn default_model(&self) -> &str; + + fn into_handle(self) -> ModelHandle + where + Self: Sized, + { + ModelHandle { + id: self.default_model().to_string(), + model: Arc::new(self), + info: ModelInfo::default(), + } + } +} + +/// An ordered list of models: the first non-excluded entry wins, so the list +/// order IS the fallback order (blueprint D14 — "an ordered list given at +/// construction"). `hint.name` pins a list entry by id. +pub struct StaticModels { + handles: Vec, + cursor: AtomicUsize, +} + +impl StaticModels { + pub fn new(handles: Vec) -> Self { + assert!(!handles.is_empty(), "StaticModels requires at least one model"); + Self { handles, cursor: AtomicUsize::new(0) } + } + + pub fn from_clients(models: Vec) -> Self { + Self::new(models.into_iter().map(|m| m.into_handle()).collect()) + } +} + +#[async_trait] +impl ModelSelector for StaticModels { + async fn select(&self, hint: &ModelHint, exclude: &[ModelId]) -> crate::Result { + // Explicit pin on the first selection of a round: resolve by id. + // (A non-empty `exclude` means the pinned model already failed: + // fall through to the ordered list.) + if let Some(name) = &hint.name + && exclude.is_empty() + { + return self + .handles + .iter() + .find(|h| &h.id == name) + .cloned() + .ok_or_else(|| anyhow::anyhow!("unknown pinned model '{name}'")); + } + // Rotation start so concurrent conversations don't pile onto handle[0]. + let start = self.cursor.fetch_add(1, Ordering::Relaxed) % self.handles.len(); + self.handles + .iter() + .cycle() + .skip(start) + .take(self.handles.len()) + .find(|h| !exclude.iter().any(|e| e == &h.id)) + .cloned() + .ok_or_else(|| anyhow::anyhow!("no alternative models available (all excluded)")) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn default_retriability_classifies_on_status() { + struct M; + #[async_trait] + impl Model for M { + async fn complete( + &self, + _req: &ModelRequest, + _d: Option>, + ) -> Result { + unreachable!() + } + } + let m = M; + for non_retriable in [401, 403, 404, 422] { + assert!( + !m.is_retriable(&ModelError::new(Some(non_retriable), "x")), + "{non_retriable} must not retry" + ); + } + for retriable in [400, 429, 500, 502, 503] { + assert!( + m.is_retriable(&ModelError::new(Some(retriable), "x")), + "{retriable} must retry" + ); + } + assert!(m.is_retriable(&ModelError::new(None, "network down"))); + } + + #[test] + fn model_hint_is_only_a_pin() { + let h = ModelHint::name("kimi-k3"); + assert_eq!(h.name.as_deref(), Some("kimi-k3")); + assert!(ModelHint::default().name.is_none()); + } +} diff --git a/crates/agent-loop/src/models/anthropic.rs b/crates/agent-loop/src/models/anthropic.rs new file mode 100644 index 0000000..54f94df --- /dev/null +++ b/crates/agent-loop/src/models/anthropic.rs @@ -0,0 +1,775 @@ +//! Anthropic client (`/v1/messages`). Ported from `llm-client/src/anthropic.rs` +//! onto the `Model` trait — including the DTL conversions (blueprint §4.10): +//! `defer_loading`, `_tool_references` → `tool_reference` blocks, and the +//! `cache_control` breakpoint moved onto the last non-deferred tool. + +use std::collections::BTreeMap; + +use async_trait::async_trait; +use futures_util::StreamExt; +use serde_json::{Value, json}; +use tokio::sync::mpsc; +use tracing::{debug, info, trace, warn}; + +use super::{SseDecoder, error_response_body, headers_to_json, redact_key}; +use crate::APP_NAME; +use crate::model::{ + Model, ModelError, ModelRequest, ModelResponse, NamedModel, RawMeta, StreamDelta, ToolCall, + Usage, +}; + +const DEFAULT_BASE_URL: &str = "https://api.anthropic.com"; +const ANTHROPIC_VERSION: &str = "2023-06-01"; + +pub struct AnthropicModel { + base_url: String, + api_key: String, + default_model: String, + /// Extra top-level request-body keys merged into every request (e.g. the + /// `thinking` config for extended reasoning). + extra_body: Option, + app_name: String, + http: reqwest::Client, +} + +impl AnthropicModel { + pub fn new(api_key: impl Into, default_model: impl Into) -> Self { + Self::with_extra_body(api_key, default_model, None) + } + + pub fn with_base_url( + base_url: impl Into, + api_key: impl Into, + default_model: impl Into, + ) -> Self { + Self { + base_url: base_url.into(), + api_key: api_key.into(), + default_model: default_model.into(), + extra_body: None, + app_name: APP_NAME.to_string(), + http: reqwest::Client::new(), + } + } + + /// Like `new` but with extra request-body keys (e.g. `{"thinking": {...}}`). + pub fn with_extra_body( + api_key: impl Into, + default_model: impl Into, + extra_body: Option, + ) -> Self { + Self { + base_url: DEFAULT_BASE_URL.to_string(), + api_key: api_key.into(), + default_model: default_model.into(), + extra_body, + app_name: APP_NAME.to_string(), + http: reqwest::Client::new(), + } + } + + pub fn with_app_name(mut self, app_name: impl Into) -> Self { + self.app_name = app_name.into(); + self + } + + /// Merges `extra_body` (then the request's own `extras`) into `body` and + /// enforces Anthropic's extended-thinking constraints: when `thinking` is + /// enabled, `temperature` is not allowed and `max_tokens` must be strictly + /// greater than `budget_tokens`. + fn apply_extra(&self, body: &mut Value, req_extras: &Value) { + for extra in [self.extra_body.as_ref(), Some(req_extras).filter(|v| v.is_object())] + .into_iter() + .flatten() + { + let Some(extra) = extra.as_object() else { continue }; + let Some(obj) = body.as_object_mut() else { return }; + for (k, v) in extra { + obj.insert(k.clone(), v.clone()); + } + } + let Some(obj) = body.as_object_mut() else { return }; + if obj.get("thinking").map(|t| t["type"] == json!("enabled")).unwrap_or(false) { + obj.remove("temperature"); + let budget = obj["thinking"]["budget_tokens"].as_i64().unwrap_or(0); + let cur_max = obj.get("max_tokens").and_then(|v| v.as_i64()).unwrap_or(4096); + if budget > 0 && cur_max <= budget { + obj.insert("max_tokens".to_string(), json!(budget + 4096)); + } + } + } + + /// Converts OpenAI-format tool definitions to Anthropic format. + /// OpenAI: { "type": "function", "function": { "name", "description", "parameters" } } + /// Anthropic: { "name", "description", "input_schema" } + /// + /// DTL (`DeferredToolReference`): a top-level `defer_loading: true` on the + /// OpenAI tool object is carried through. When any tool is deferred, the + /// cache breakpoint is placed on the last **non-deferred** tool — a + /// deferred tool cannot carry `cache_control` (the API 400s). + fn convert_tools(tools: &[Value]) -> Vec { + let has_deferred = tools.iter().any(|t| t["defer_loading"].as_bool() == Some(true)); + let mut out: Vec = tools + .iter() + .filter_map(|t| { + let func = &t["function"]; + let name = func["name"].as_str()?; + let mut tool = json!({ + "name": name, + "description": func["description"].as_str().unwrap_or(""), + "input_schema": func["parameters"], + }); + if t["defer_loading"].as_bool() == Some(true) { + tool["defer_loading"] = json!(true); + } + Some(tool) + }) + .collect(); + if has_deferred + && let Some(t) = out.iter_mut().rev().find(|t| t["defer_loading"].as_bool() != Some(true)) + { + t["cache_control"] = json!({ "type": "ephemeral" }); + } + out + } + + /// Converts OpenAI-format messages to Anthropic format: system extracted + /// separately; assistant tool_calls → tool_use blocks; consecutive `tool` + /// messages grouped into one user message of tool_result blocks. + fn convert_messages(messages: &[Value]) -> Vec { + let mut out: Vec = Vec::new(); + let mut i = 0; + + while i < messages.len() { + let msg = &messages[i]; + let role = msg["role"].as_str().unwrap_or(""); + + match role { + "system" => { i += 1; } + + "user" => { + out.push(json!({ + "role": "user", + "content": convert_user_content(&msg["content"]), + })); + i += 1; + } + + "assistant" => { + if let Some(tool_calls) = msg["tool_calls"].as_array() { + let mut content: Vec = Vec::new(); + + let text = msg["content"].as_str().unwrap_or(""); + if !text.is_empty() { + content.push(json!({ "type": "text", "text": text })); + } + + for tc in tool_calls { + let id = tc["id"].as_str().unwrap_or(""); + let name = tc["function"]["name"].as_str().unwrap_or(""); + let args_str = tc["function"]["arguments"].as_str().unwrap_or("{}"); + let input: Value = serde_json::from_str(args_str) + .unwrap_or(Value::Object(Default::default())); + + content.push(json!({ + "type": "tool_use", + "id": id, + "name": name, + "input": input, + })); + } + + out.push(json!({ "role": "assistant", "content": content })); + } else { + out.push(json!({ + "role": "assistant", + "content": msg["content"].as_str().unwrap_or(""), + })); + } + i += 1; + } + + "tool" => { + // Group consecutive tool results into a single user message. + let mut results: Vec = Vec::new(); + while i < messages.len() && messages[i]["role"].as_str() == Some("tool") { + let tm = &messages[i]; + // DTL (`DeferredToolReference`): a tool result carrying + // `_tool_references` becomes a content array of + // `tool_reference` blocks, which the API expands into + // the deferred tools' full definitions. + let content: Value = match tm["_tool_references"].as_array() { + Some(refs) if !refs.is_empty() => Value::Array( + refs.iter() + .filter_map(|r| r.as_str()) + .map(|name| json!({ "type": "tool_reference", "tool_name": name })) + .collect(), + ), + _ => Value::String(tm["content"].as_str().unwrap_or("").to_string()), + }; + results.push(json!({ + "type": "tool_result", + "tool_use_id": tm["tool_call_id"].as_str().unwrap_or(""), + "content": content, + })); + i += 1; + } + out.push(json!({ "role": "user", "content": results })); + } + + _ => { i += 1; } + } + } + + out + } + + /// Shared `/v1/messages` body (the caller adds `stream` on top). + fn tools_body(&self, system: Option, messages: Vec, tools: Vec, req: &ModelRequest) -> Value { + let max_tokens = req.max_tokens.unwrap_or(4096); + let mut body = json!({ + "model": req.model, + "max_tokens": max_tokens, + "messages": messages, + "tools": tools, + }); + + if let Some(sys) = system { body["system"] = sys; } + if let Some(t) = req.temperature { body["temperature"] = t.into(); } + self.apply_extra(&mut body, &req.extras); + body + } + + /// Collects ALL system-role messages into the single `system` parameter. + /// Structured content (a text-block array with `cache_control`) is kept + /// in array form so the cache breakpoint survives. + fn merged_system(messages: &[Value]) -> Option { + let sys: Vec<&Value> = messages + .iter() + .filter(|m| m["role"].as_str() == Some("system")) + .collect(); + if sys.is_empty() { return None; } + + if !sys.iter().any(|m| m["content"].is_array()) { + let parts: Vec<&str> = sys.iter().filter_map(|m| m["content"].as_str()).collect(); + return if parts.is_empty() { None } else { Some(Value::String(parts.join("\n\n---\n\n"))) }; + } + + let mut blocks: Vec = Vec::new(); + for m in &sys { + match &m["content"] { + Value::String(s) if !s.is_empty() => blocks.push(json!({ "type": "text", "text": s })), + Value::Array(arr) => { + for b in arr { + if b["type"].as_str() == Some("text") { + blocks.push(b.clone()); + } + } + } + _ => {} + } + } + if blocks.is_empty() { None } else { Some(Value::Array(blocks)) } + } + + fn url(&self) -> String { + format!("{}/v1/messages", self.base_url.trim_end_matches('/')) + } + + fn logged_headers(&self) -> Value { + json!({ + "x-api-key": redact_key(&self.api_key), + "anthropic-version": ANTHROPIC_VERSION, + "content-type": "application/json", + }) + } + + /// Sends the request WITHOUT `error_for_status`, so the caller can read + /// the error body and attach the payload to the `ModelError`. + async fn send_request(&self, body: &Value) -> Result { + self.http + .post(self.url()) + .header("x-api-key", &self.api_key) + .header("anthropic-version", ANTHROPIC_VERSION) + .header("X-Title", &self.app_name) + .json(body) + .send() + .await + .map_err(ModelError::from_reqwest) + } + + /// Joined `thinking` blocks of a content array (extended thinking). + fn reasoning_of(content_blocks: &[Value]) -> Option { + let parts: Vec<&str> = content_blocks + .iter() + .filter(|b| b["type"].as_str() == Some("thinking")) + .filter_map(|b| b["thinking"].as_str()) + .collect(); + if parts.is_empty() { None } else { Some(parts.join("\n")) } + } + + /// The buffered path. + async fn buffered(&self, req: &ModelRequest) -> Result { + let system = Self::merged_system(&req.messages); + let anthropic_messages = Self::convert_messages(&req.messages); + let anthropic_tools = Self::convert_tools(&req.tools); + let body = self.tools_body(system, anthropic_messages, anthropic_tools, req); + + debug!(model = %req.model, tools = req.tools.len(), "anthropic: sending request"); + trace!(body = %body, "anthropic: request body"); + + let request_body = body.clone(); + let request_headers = self.logged_headers(); + + let http_resp = self.send_request(&body).await?; + + let response_headers = headers_to_json(http_resp.headers()); + let status = http_resp.status(); + let resp_text = http_resp.text().await.map_err(ModelError::from_reqwest)?; + if !status.is_success() { + return Err(ModelError { + status: Some(status.as_u16()), + message: format!("anthropic: HTTP {status} from {url}\nbody: {resp_text}", url = self.url()), + raw: Some(RawMeta { + request_headers: Some(request_headers), + request_body: Some(request_body), + response_headers: Some(response_headers), + response_body: Some(error_response_body(resp_text)), + }), + }); + } + let resp: Value = serde_json::from_str(&resp_text).map_err(|e| { + ModelError::new(None, format!("anthropic: failed to parse response JSON: {e}\nbody: {resp_text}")) + })?; + + let raw = RawMeta { + request_headers: Some(request_headers), + request_body: Some(request_body), + response_headers: Some(response_headers), + response_body: Some(resp.clone()), + }; + + let stop_reason = resp["stop_reason"].as_str().unwrap_or(""); + let mut usage = Usage { + input_tokens: resp["usage"]["input_tokens"].as_u64().map(|n| n as u32), + output_tokens: resp["usage"]["output_tokens"].as_u64().map(|n| n as u32), + cache_read: resp["usage"]["cache_read_input_tokens"].as_u64().map(|n| n as u32), + cache_write: resp["usage"]["cache_creation_input_tokens"].as_u64().map(|n| n as u32), + cost_usd: None, + truncated: stop_reason == "max_tokens", + }; + let content_blocks = resp["content"].as_array().cloned().unwrap_or_default(); + info!(model = %req.model, ?usage.input_tokens, ?usage.output_tokens, stop_reason, "anthropic: response received"); + if usage.truncated { + warn!(model = %req.model, ?usage.output_tokens, "anthropic: response truncated (max_tokens reached)"); + } + + let has_tool_use = content_blocks.iter().any(|b| b["type"].as_str() == Some("tool_use")); + let reasoning = Self::reasoning_of(&content_blocks); + + // Anthropic sometimes returns stop_reason "end_turn" even when + // tool_use blocks are present — check the blocks directly. + let mut resp_out = if stop_reason == "tool_use" || has_tool_use { + let text: String = content_blocks + .iter() + .filter(|b| b["type"].as_str() == Some("text")) + .filter_map(|b| b["text"].as_str()) + .collect::>() + .join("\n"); + usage.truncated = false; + let calls: Vec = content_blocks + .iter() + .filter(|b| b["type"].as_str() == Some("tool_use")) + .map(|b| ToolCall { + id: b["id"].as_str().unwrap_or("").to_string(), + name: b["name"].as_str().unwrap_or("").to_string(), + arguments: b["input"].clone(), + }) + .collect(); + ModelResponse::ToolCalls { content: text, calls, reasoning, usage, raw: None } + } else { + let content = content_blocks + .iter() + .find(|b| b["type"].as_str() == Some("text")) + .and_then(|b| b["text"].as_str()) + .unwrap_or("") + .to_string(); + ModelResponse::Message { content, reasoning, usage, raw: None } + }; + match &mut resp_out { + ModelResponse::Message { raw: r, .. } | ModelResponse::ToolCalls { raw: r, .. } => { + *r = Some(raw) + } + } + Ok(resp_out) + } + + /// SSE streaming path: Anthropic streams typed events (`message_start` / + /// `content_block_*` / `message_delta`); text and thinking deltas are + /// forwarded best-effort while blocks accumulate into the same + /// `ModelResponse` the buffered path returns. + #[allow(clippy::result_large_err)] + async fn stream_chat( + &self, + req: &ModelRequest, + delta_tx: &mpsc::Sender, + emitted: &mut bool, + ) -> Result { + let system = Self::merged_system(&req.messages); + let anthropic_messages = Self::convert_messages(&req.messages); + let anthropic_tools = Self::convert_tools(&req.tools); + let mut body = self.tools_body(system, anthropic_messages, anthropic_tools, req); + body["stream"] = json!(true); + + debug!(model = %req.model, tools = req.tools.len(), "anthropic: sending streaming request"); + trace!(body = %body, "anthropic: streaming request body"); + + let request_body = body.clone(); + let request_headers = self.logged_headers(); + + let http_resp = self.send_request(&body).await?; + let response_headers = headers_to_json(http_resp.headers()); + let status = http_resp.status(); + if !status.is_success() { + let resp_text = http_resp.text().await.map_err(ModelError::from_reqwest)?; + return Err(ModelError { + status: Some(status.as_u16()), + message: format!("anthropic: HTTP {status} from {url}\nbody: {resp_text}", url = self.url()), + raw: Some(RawMeta { + request_headers: Some(request_headers), + request_body: Some(request_body), + response_headers: Some(response_headers), + response_body: Some(error_response_body(resp_text)), + }), + }); + } + + /// One content block being accumulated by index. + #[derive(Default)] + struct Block { + kind: String, // "text" | "thinking" | "tool_use" + buf: String, // text/thinking content or input_json fragments + id: String, + name: String, + } + + let mut blocks: BTreeMap = BTreeMap::new(); + let mut stop_reason: Option = None; + let mut usage = json!({}); + let mut sse = SseDecoder::new(); + let mut byte_stream = http_resp.bytes_stream(); + + let mut handle_payload = |payload: &str, emitted: &mut bool| -> Result<(), ModelError> { + let Ok(v) = serde_json::from_str::(payload) else { return Ok(()) }; + match v["type"].as_str().unwrap_or("") { + "message_start" => { + if let Some(u) = v["message"]["usage"].as_object() { + for (k, val) in u { usage[k.clone()] = val.clone(); } + } + } + "content_block_start" => { + let idx = v["index"].as_u64().unwrap_or(0); + let cb = &v["content_block"]; + let block = blocks.entry(idx).or_default(); + block.kind = cb["type"].as_str().unwrap_or("").to_string(); + block.id = cb["id"].as_str().unwrap_or("").to_string(); + block.name = cb["name"].as_str().unwrap_or("").to_string(); + } + "content_block_delta" => { + let idx = v["index"].as_u64().unwrap_or(0); + let delta = &v["delta"]; + match delta["type"].as_str().unwrap_or("") { + "text_delta" => { + if let Some(t) = delta["text"].as_str().filter(|t| !t.is_empty()) { + blocks.entry(idx).or_default().buf.push_str(t); + *emitted = true; + let _ = delta_tx.try_send(StreamDelta::Text(t.to_string())); + } + } + "thinking_delta" => { + if let Some(t) = delta["thinking"].as_str().filter(|t| !t.is_empty()) { + blocks.entry(idx).or_default().buf.push_str(t); + *emitted = true; + let _ = delta_tx.try_send(StreamDelta::Reasoning(t.to_string())); + } + } + "input_json_delta" => { + if let Some(j) = delta["partial_json"].as_str() { + blocks.entry(idx).or_default().buf.push_str(j); + } + } + _ => {} + } + } + "message_delta" => { + if let Some(sr) = v["delta"]["stop_reason"].as_str() { + stop_reason = Some(sr.to_string()); + } + if let Some(u) = v["usage"].as_object() { + for (k, val) in u { usage[k.clone()] = val.clone(); } + } + } + "error" => { + return Err(ModelError::new(None, format!("anthropic: stream error event: {payload}"))); + } + _ => {} + } + Ok(()) + }; + + while let Some(chunk) = byte_stream.next().await { + let chunk = chunk.map_err(ModelError::from_reqwest)?; + for payload in sse.feed(&chunk) { + handle_payload(&payload, emitted)?; + } + } + for payload in sse.finish() { + handle_payload(&payload, emitted)?; + } + + let stop = stop_reason.as_deref().unwrap_or(""); + let usage_struct = Usage { + input_tokens: usage["input_tokens"].as_u64().map(|n| n as u32), + output_tokens: usage["output_tokens"].as_u64().map(|n| n as u32), + cache_read: usage["cache_read_input_tokens"].as_u64().map(|n| n as u32), + cache_write: usage["cache_creation_input_tokens"].as_u64().map(|n| n as u32), + cost_usd: None, + truncated: stop == "max_tokens", + }; + info!(model = %req.model, ?usage_struct.input_tokens, ?usage_struct.output_tokens, stop_reason = stop, "anthropic: streaming response completed"); + if usage_struct.truncated { + warn!(model = %req.model, "anthropic: response truncated (max_tokens reached)"); + } + + let text_of = |kind: &str| -> String { + blocks.values() + .filter(|b| b.kind == kind) + .map(|b| b.buf.as_str()) + .collect::>() + .join("\n") + }; + let reasoning_text = text_of("thinking"); + let reasoning = if reasoning_text.is_empty() { None } else { Some(reasoning_text) }; + let tool_blocks: Vec<&Block> = blocks.values().filter(|b| b.kind == "tool_use").collect(); + + // Buffered-shaped response body for the payload log. + let content_log: Vec = blocks.values().map(|b| match b.kind.as_str() { + "tool_use" => json!({"type": "tool_use", "id": b.id, "name": b.name, "input": serde_json::from_str::(&b.buf).unwrap_or(json!({}))}), + "thinking" => json!({"type": "thinking", "thinking": b.buf}), + _ => json!({"type": "text", "text": b.buf}), + }).collect(); + let raw = RawMeta { + request_headers: Some(request_headers), + request_body: Some(request_body), + response_headers: Some(response_headers), + response_body: Some(json!({ + "streamed": true, + "content": content_log, + "stop_reason": stop, + "usage": usage, + })), + }; + + let mut resp_out = if !tool_blocks.is_empty() { + let calls = tool_blocks + .iter() + .map(|b| ToolCall { + id: b.id.clone(), + name: b.name.clone(), + arguments: serde_json::from_str(&b.buf).unwrap_or(Value::Object(Default::default())), + }) + .collect(); + ModelResponse::ToolCalls { content: text_of("text"), calls, reasoning, usage: usage_struct, raw: None } + } else { + ModelResponse::Message { content: text_of("text"), reasoning, usage: usage_struct, raw: None } + }; + match &mut resp_out { + ModelResponse::Message { raw: r, .. } | ModelResponse::ToolCalls { raw: r, .. } => { + *r = Some(raw) + } + } + Ok(resp_out) + } +} + +impl NamedModel for AnthropicModel { + fn default_model(&self) -> &str { &self.default_model } +} + +#[async_trait] +impl Model for AnthropicModel { + async fn complete( + &self, + req: &ModelRequest, + deltas: Option>, + ) -> Result { + match deltas { + None => self.buffered(req).await, + Some(delta_tx) => { + let mut emitted = false; + match self.stream_chat(req, &delta_tx, &mut emitted).await { + Ok(ok) => Ok(ok), + // Pre-stream failure (nothing shown yet): retry buffered. + // A mid-stream failure propagates to the fallback logic. + Err(e) if !emitted => { + debug!(model = %req.model, error = %e, "anthropic: streaming failed before any delta; retrying buffered"); + self.buffered(req).await + } + Err(e) => Err(e), + } + } + } + } +} + +/// User content arrives either as a plain string or as an OpenAI-style parts +/// array (text + `image_url` data URLs + `file` PDF parts). Strings pass +/// through; parts become Anthropic blocks. Unknown parts are dropped with a +/// warning. +fn convert_user_content(content: &Value) -> Value { + let Some(parts) = content.as_array() else { + return Value::String(content.as_str().unwrap_or("").to_string()); + }; + let mut blocks = Vec::new(); + for p in parts { + match p["type"].as_str().unwrap_or("") { + "text" => blocks.push(json!({ + "type": "text", + "text": p["text"].as_str().unwrap_or(""), + })), + "image_url" => { + if let Some(block) = parse_data_image(&p["image_url"]) { + blocks.push(block); + } + } + "file" => { + if let Some(block) = parse_data_document(&p["file"]) { + blocks.push(block); + } + } + other => tracing::warn!(part_type = other, "dropping content part unsupported by Anthropic"), + } + } + Value::Array(blocks) +} + +/// `{"url": "data:;base64,"}` → an Anthropic base64 image block. +fn parse_data_image(image_url: &Value) -> Option { + let url = image_url["url"].as_str().or_else(|| image_url.as_str())?; + let (mime, data) = url.strip_prefix("data:")?.split_once(";base64,")?; + Some(json!({ + "type": "image", + "source": { "type": "base64", "media_type": mime, "data": data }, + })) +} + +/// `{"file_data": "data:application/pdf;base64,"}` → an Anthropic +/// base64 `document` block (the native PDF input). +fn parse_data_document(file: &Value) -> Option { + let url = file["file_data"].as_str()?; + let (mime, data) = url.strip_prefix("data:")?.split_once(";base64,")?; + Some(json!({ + "type": "document", + "source": { "type": "base64", "media_type": mime, "data": data }, + })) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn reasoning_of_joins_thinking_blocks() { + let blocks = vec![ + json!({"type": "thinking", "thinking": "first"}), + json!({"type": "text", "text": "answer"}), + json!({"type": "thinking", "thinking": "second"}), + ]; + assert_eq!( + AnthropicModel::reasoning_of(&blocks), + Some("first\nsecond".to_string()) + ); + assert_eq!(AnthropicModel::reasoning_of(&[]), None); + assert_eq!( + AnthropicModel::reasoning_of(&[json!({"type": "text", "text": "a"})]), + None + ); + } + + #[test] + fn convert_tools_carries_defer_loading_and_moves_cache_control() { + let tools = vec![ + json!({"type":"function","function":{"name":"a","description":"","parameters":{}}}), + json!({"type":"function","function":{"name":"b","description":"","parameters":{}},"defer_loading":true}), + json!({"type":"function","function":{"name":"c","description":"","parameters":{}},"defer_loading":true}), + ]; + let out = AnthropicModel::convert_tools(&tools); + assert_eq!(out[0]["cache_control"], json!({"type": "ephemeral"})); + assert!(out[0].get("defer_loading").is_none()); + assert_eq!(out[1]["defer_loading"], json!(true)); + assert!(out[1].get("cache_control").is_none()); + assert_eq!(out[2]["defer_loading"], json!(true)); + } + + #[test] + fn convert_messages_tool_references_become_blocks() { + let messages = vec![ + json!({"role":"assistant","content":"","tool_calls":[ + {"id":"t1","type":"function","function":{"name":"activate_tools","arguments":"{\"groups\":[\"gmail\"]}"}} + ]}), + json!({"role":"tool","tool_call_id":"t1","content":"ok","_tool_references":["mcp__gmail__send"]}), + ]; + let out = AnthropicModel::convert_messages(&messages); + assert_eq!(out.len(), 2); + let results = out[1]["content"].as_array().unwrap(); + assert_eq!( + results[0]["content"], + json!([{ "type": "tool_reference", "tool_name": "mcp__gmail__send" }]) + ); + } + + #[test] + fn user_content_string_passthrough() { + let v = convert_user_content(&json!("hello")); + assert_eq!(v, json!("hello")); + } + + #[test] + fn user_content_parts_become_anthropic_blocks() { + let v = convert_user_content(&json!([ + { "type": "text", "text": "what is this?" }, + { "type": "image_url", "image_url": { "url": "data:image/png;base64,QUJD" } }, + ])); + assert_eq!(v, json!([ + { "type": "text", "text": "what is this?" }, + { "type": "image", "source": { "type": "base64", "media_type": "image/png", "data": "QUJD" } }, + ])); + } + + #[test] + fn user_content_drops_video_and_non_data_urls() { + let v = convert_user_content(&json!([ + { "type": "text", "text": "t" }, + { "type": "video_url", "video_url": { "url": "data:video/mp4;base64,QUJD" } }, + { "type": "image_url", "image_url": { "url": "https://example.com/x.png" } }, + ])); + assert_eq!(v, json!([{ "type": "text", "text": "t" }])); + } + + #[test] + fn user_content_file_part_becomes_document_block() { + let v = convert_user_content(&json!([ + { "type": "text", "text": "read this" }, + { "type": "file", "file": { "filename": "a.pdf", "file_data": "data:application/pdf;base64,QUJD" } }, + ])); + assert_eq!(v, json!([ + { "type": "text", "text": "read this" }, + { "type": "document", "source": { "type": "base64", "media_type": "application/pdf", "data": "QUJD" } }, + ])); + + let v = convert_user_content(&json!([ + { "type": "file", "file": { "filename": "a.pdf", "file_data": "https://example.com/a.pdf" } }, + ])); + assert_eq!(v, json!([])); + } +} diff --git a/crates/agent-loop/src/models/lm_studio.rs b/crates/agent-loop/src/models/lm_studio.rs new file mode 100644 index 0000000..ca758c9 --- /dev/null +++ b/crates/agent-loop/src/models/lm_studio.rs @@ -0,0 +1,43 @@ +//! LM Studio client — a thin wrapper over [`OpenAiModel`] defaulting to +//! `http://localhost:1234/v1` with no API key. (LM Studio can also be served +//! by a YAML-declared provider; this client is kept for explicit use.) + +use async_trait::async_trait; +use tokio::sync::mpsc; + +use super::openai::OpenAiModel; +use crate::model::{Model, ModelError, ModelRequest, ModelResponse, NamedModel, StreamDelta}; + +pub struct LmStudioModel { + inner: OpenAiModel, +} + +impl LmStudioModel { + /// `base_url` defaults to `http://localhost:1234/v1` if `None`. + pub fn new(base_url: Option>, default_model: impl Into) -> Self { + let url = base_url + .map(|u| u.into()) + .unwrap_or_else(|| "http://localhost:1234/v1".to_string()); + Self { inner: OpenAiModel::new(url, "", default_model) } + } +} + +impl NamedModel for LmStudioModel { + fn default_model(&self) -> &str { self.inner.default_model() } +} + +#[async_trait] +impl Model for LmStudioModel { + /// LM Studio is OpenAI-compatible: everything forwards to the inner + /// client (its pre-delta buffered retry covers local builds rejecting + /// `stream_options`). + async fn complete( + &self, + req: &ModelRequest, + deltas: Option>, + ) -> Result { + self.inner.complete(req, deltas).await + } + + fn is_retriable(&self, err: &ModelError) -> bool { self.inner.is_retriable(err) } +} diff --git a/crates/agent-loop/src/models/mod.rs b/crates/agent-loop/src/models/mod.rs new file mode 100644 index 0000000..97ffbb5 --- /dev/null +++ b/crates/agent-loop/src/models/mod.rs @@ -0,0 +1,42 @@ +//! Shipped `Model` clients (blueprint D13): OpenAI-compatible, Anthropic, +//! Ollama, LM Studio — plus the shared SSE decoder and HTTP helpers. +//! +//! All clients are stateless (connection config only) and share the same +//! failure policy: if a stream dies BEFORE any delta, the client retries +//! buffered on the same model (providers rejecting `stream` keep working); a +//! mid-stream failure propagates to the caller's fallback logic. + +pub mod anthropic; +pub mod lm_studio; +pub mod ollama; +pub mod openai; +mod sse; + +pub use anthropic::AnthropicModel; +pub use lm_studio::LmStudioModel; +pub use ollama::OllamaModel; +pub use openai::OpenAiModel; +pub(crate) use sse::SseDecoder; + +use serde_json::Value; + +/// Converts a reqwest `HeaderMap` into a JSON object (for payload logging). +pub(crate) fn headers_to_json(headers: &reqwest::header::HeaderMap) -> Value { + let map: serde_json::Map = headers + .iter() + .map(|(k, v)| (k.as_str().to_string(), v.to_str().unwrap_or("").into())) + .collect(); + Value::Object(map) +} + +/// Raw error body → JSON for the payload log: parsed JSON when the provider +/// returned JSON, else the raw text wrapped as a JSON string so a non-JSON +/// body (HTML gateway page) is still preserved verbatim. +pub(crate) fn error_response_body(text: String) -> Value { + serde_json::from_str::(&text).unwrap_or(Value::String(text)) +} + +/// Redacted preview of an API key: first 7 chars + "***". +pub(crate) fn redact_key(key: &str) -> String { + if key.len() > 7 { format!("{}***", &key[..7]) } else { "***".to_string() } +} diff --git a/crates/agent-loop/src/models/ollama.rs b/crates/agent-loop/src/models/ollama.rs new file mode 100644 index 0000000..2f197ac --- /dev/null +++ b/crates/agent-loop/src/models/ollama.rs @@ -0,0 +1,102 @@ +//! Ollama client (native `/api/chat` endpoint). Ported from +//! `llm-client/src/ollama.rs`. No streaming, no tool support — tool-call +//! messages are flattened to text, mirroring the previous default behavior. + +use async_trait::async_trait; +use serde_json::{Value, json}; +use tokio::sync::mpsc; + +use crate::model::{Model, ModelError, ModelRequest, ModelResponse, NamedModel, StreamDelta, Usage}; + +/// Ollama client. Defaults to `http://localhost:11434`. No API key required. +pub struct OllamaModel { + base_url: String, + default_model: String, + http: reqwest::Client, +} + +impl OllamaModel { + /// `base_url` defaults to `http://localhost:11434` if `None`. + pub fn new(base_url: Option>, default_model: impl Into) -> Self { + let url = base_url + .map(|u| u.into()) + .unwrap_or_else(|| "http://localhost:11434".to_string()); + Self { base_url: url, default_model: default_model.into(), http: reqwest::Client::new() } + } +} + +impl NamedModel for OllamaModel { + fn default_model(&self) -> &str { &self.default_model } +} + +#[async_trait] +impl Model for OllamaModel { + async fn complete( + &self, + req: &ModelRequest, + _deltas: Option>, + ) -> Result { + // Flatten to plain text messages: tool results and assistant + // tool_calls are dropped (no native tool support on this path). + let msgs: Vec = req + .messages + .iter() + .filter_map(|m| { + let role = m["role"].as_str()?; + if !matches!(role, "system" | "user" | "assistant") { + return None; + } + let content = m["content"].as_str().unwrap_or("").to_string(); + Some(json!({ "role": role, "content": content })) + }) + .collect(); + + let mut options_obj = json!({}); + if let Some(t) = req.temperature { options_obj["temperature"] = t.into(); } + if let Some(n) = req.max_tokens { options_obj["num_predict"] = n.into(); } + + let body = json!({ + "model": req.model, + "messages": msgs, + "stream": false, + "options": options_obj, + }); + + let url = format!("{}/api/chat", self.base_url.trim_end_matches('/')); + + let http_resp = self + .http + .post(&url) + .json(&body) + .send() + .await + .map_err(ModelError::from_reqwest)?; + + let status = http_resp.status(); + if !status.is_success() { + let resp_text = http_resp.text().await.map_err(ModelError::from_reqwest)?; + return Err(ModelError::new( + Some(status.as_u16()), + format!("ollama: HTTP {status} from {url}\nbody: {resp_text}"), + )); + } + + let resp: Value = http_resp.json().await.map_err(ModelError::from_reqwest)?; + + let content = resp["message"]["content"] + .as_str() + .ok_or_else(|| ModelError::new(None, "ollama: missing content in response"))? + .to_string(); + + Ok(ModelResponse::Message { + content, + reasoning: None, + usage: Usage { + input_tokens: resp["prompt_eval_count"].as_u64().map(|n| n as u32), + output_tokens: resp["eval_count"].as_u64().map(|n| n as u32), + ..Usage::default() + }, + raw: None, + }) + } +} diff --git a/crates/agent-loop/src/models/openai.rs b/crates/agent-loop/src/models/openai.rs new file mode 100644 index 0000000..b161bbf --- /dev/null +++ b/crates/agent-loop/src/models/openai.rs @@ -0,0 +1,459 @@ +//! OpenAI-compatible client (OpenAI, OpenRouter, Moonshot/Kimi, and every +//! provider declared via YAML). Ported from `llm-client/src/openai.rs` onto +//! the `Model` trait. +//! +//! Kimi's `SystemToolBlock` DTL needs NO client code: messages are passed +//! through verbatim and the endpoint speaks the `{role:"system", tools:[…]}` +//! convention natively. + +use std::collections::BTreeMap; + +use async_trait::async_trait; +use futures_util::StreamExt; +use serde_json::{Value, json}; +use tokio::sync::mpsc; +use tracing::{debug, info, trace, warn}; + +use super::{SseDecoder, error_response_body, headers_to_json, redact_key}; +use crate::APP_NAME; +use crate::model::{ + Model, ModelError, ModelRequest, ModelResponse, NamedModel, RawMeta, StreamDelta, ToolCall, + Usage, +}; + +/// OpenAI ChatGPT client (also compatible with any OpenAI-spec endpoint). +pub struct OpenAiModel { + base_url: String, + api_key: String, + default_model: String, + extra_params: Option, + /// When true, Anthropic-compatible prompt-caching hints are injected + /// (OpenRouter routing to Anthropic models). + enable_prompt_cache: bool, + app_name: String, + http: reqwest::Client, +} + +impl OpenAiModel { + /// Minimal constructor: base URL + key + default model name (used as the + /// selector id by `SingleModel`). + pub fn new( + base_url: impl Into, + api_key: impl Into, + default_model: impl Into, + ) -> Self { + Self::with_options(base_url, api_key, default_model, None, false) + } + + pub fn with_options( + base_url: impl Into, + api_key: impl Into, + default_model: impl Into, + extra_params: Option, + enable_prompt_cache: bool, + ) -> Self { + Self { + base_url: base_url.into(), + api_key: api_key.into(), + default_model: default_model.into(), + extra_params, + enable_prompt_cache, + app_name: APP_NAME.to_string(), + http: reqwest::Client::new(), + } + } + + /// Override the `X-Title` header (OpenRouter rankings). + pub fn with_app_name(mut self, app_name: impl Into) -> Self { + self.app_name = app_name.into(); + self + } + + /// Merges extra top-level object keys into `body` (later maps win). + fn merge_extra(body: &mut Value, extra: Option<&Value>) { + if let Some(Value::Object(extra)) = extra + && let Some(b) = body.as_object_mut() + { + for (k, v) in extra { + b.insert(k.clone(), v.clone()); + } + } + } + + fn url(&self) -> String { + format!("{}/chat/completions", self.base_url.trim_end_matches('/')) + } + + /// Shared request body for the buffered and the streaming path. + fn base_body(&self, model: &str, messages: &[Value], tools: &[Value]) -> Value { + let mut body = json!({ + "model": model, + "messages": messages, + }); + + if !tools.is_empty() { + // When prompt caching is enabled, tag the last tool with cache_control + // so the entire tools array is included in the KV cache prefix. + let tools_value: Value = if self.enable_prompt_cache { + let mut tagged = tools.to_vec(); + if let Some(last) = tagged.last_mut() { + last["cache_control"] = json!({"type": "ephemeral"}); + } + tagged.into() + } else { + tools.into() + }; + body["tools"] = tools_value; + body["tool_choice"] = "auto".into(); + } + body + } + + fn finalize_body(&self, mut body: Value, req: &ModelRequest) -> Value { + if let Some(t) = req.max_tokens { body["max_tokens"] = t.into(); } + if let Some(t) = req.temperature { body["temperature"] = t.into(); } + Self::merge_extra(&mut body, self.extra_params.as_ref()); + Self::merge_extra(&mut body, Some(&req.extras)); + body + } + + /// Request metadata for logging (shared by buffered and streaming paths). + fn logged_headers(&self) -> Value { + let mut logged_headers = json!({ + "authorization": format!("Bearer {}", redact_key(&self.api_key)), + "content-type": "application/json", + }); + if self.enable_prompt_cache { + logged_headers["anthropic-beta"] = "prompt-caching-2024-07-31".into(); + } + logged_headers + } + + async fn send_request(&self, body: &Value) -> Result { + let mut req = self + .http + .post(self.url()) + .bearer_auth(&self.api_key) + .header("X-Title", &self.app_name); + if self.enable_prompt_cache { + req = req.header("anthropic-beta", "prompt-caching-2024-07-31"); + } + req.json(body).send().await.map_err(ModelError::from_reqwest) + } + + /// The buffered path. + async fn buffered(&self, req: &ModelRequest) -> Result { + let body = self.finalize_body(self.base_body(&req.model, &req.messages, &req.tools), req); + + debug!(model = %req.model, tools = req.tools.len(), prompt_cache = self.enable_prompt_cache, "openai: sending request"); + trace!(body = %body, "openai: request body"); + + let request_body = body.clone(); + let request_headers = self.logged_headers(); + + let http_resp = self.send_request(&body).await?; + + let response_headers = headers_to_json(http_resp.headers()); + let status = http_resp.status(); + let resp_text = http_resp.text().await.map_err(ModelError::from_reqwest)?; + + if !status.is_success() { + return Err(ModelError { + status: Some(status.as_u16()), + message: format!("openai: HTTP {status} from {url}\nbody: {resp_text}", url = self.url()), + raw: Some(RawMeta { + request_headers: Some(request_headers), + request_body: Some(request_body), + response_headers: Some(response_headers), + response_body: Some(error_response_body(resp_text)), + }), + }); + } + + let resp: Value = serde_json::from_str(&resp_text).map_err(|e| { + ModelError::new(None, format!("openai: failed to parse response JSON: {e}\nbody: {resp_text}")) + })?; + let response_body: Value = serde_json::from_str(&resp_text).unwrap_or(Value::Null); + + let raw = RawMeta { + request_headers: Some(request_headers), + request_body: Some(request_body), + response_headers: Some(response_headers), + response_body: Some(response_body), + }; + + Ok(parse_turn(&resp, &req.model).with_raw(raw)) + } + + /// SSE streaming path. Accumulates fragments into the same `ModelResponse` + /// the buffered path returns, forwarding deltas best-effort. `emitted` + /// tracks whether any delta was pushed, distinguishing a pre-stream + /// failure (safe to retry buffered) from a mid-stream one. + async fn stream_chat( + &self, + req: &ModelRequest, + delta_tx: &mpsc::Sender, + emitted: &mut bool, + ) -> Result { + let mut body = self.base_body(&req.model, &req.messages, &req.tools); + body["stream"] = json!(true); + body["stream_options"] = json!({ "include_usage": true }); + let body = self.finalize_body(body, req); + + debug!(model = %req.model, tools = req.tools.len(), prompt_cache = self.enable_prompt_cache, "openai: sending streaming request"); + trace!(body = %body, "openai: streaming request body"); + + let request_body = body.clone(); + let request_headers = self.logged_headers(); + + let http_resp = self.send_request(&body).await?; + + let response_headers = headers_to_json(http_resp.headers()); + let status = http_resp.status(); + if !status.is_success() { + let resp_text = http_resp.text().await.map_err(ModelError::from_reqwest)?; + return Err(ModelError { + status: Some(status.as_u16()), + message: format!("openai: HTTP {status} from {url}\nbody: {resp_text}", url = self.url()), + raw: Some(RawMeta { + request_headers: Some(request_headers), + request_body: Some(request_body), + response_headers: Some(response_headers), + response_body: Some(error_response_body(resp_text)), + }), + }); + } + + let mut content = String::new(); + let mut reasoning = String::new(); + // index → (id, name, arguments fragment buffer) + let mut tool_calls: BTreeMap = BTreeMap::new(); + let mut finish_reason: Option = None; + let mut usage: Option = None; + let mut sse = SseDecoder::new(); + let mut byte_stream = http_resp.bytes_stream(); + + let mut handle_payload = |payload: &str, emitted: &mut bool| { + if payload == "[DONE]" { + return; + } + let Ok(v) = serde_json::from_str::(payload) else { return }; + if let Some(u) = v.get("usage").filter(|u| !u.is_null()) { + usage = Some(u.clone()); + } + let Some(choice) = v["choices"].as_array().and_then(|a| a.first()) else { return }; + if let Some(fr) = choice["finish_reason"].as_str() { + finish_reason = Some(fr.to_string()); + } + let delta = &choice["delta"]; + if let Some(t) = delta["content"].as_str().filter(|t| !t.is_empty()) { + content.push_str(t); + *emitted = true; + let _ = delta_tx.try_send(StreamDelta::Text(t.to_string())); + } + // DeepSeek uses `reasoning_content`, MiniMax M3 and others `reasoning`. + if let Some(t) = delta["reasoning_content"].as_str() + .or_else(|| delta["reasoning"].as_str()) + .filter(|t| !t.is_empty()) + { + reasoning.push_str(t); + *emitted = true; + let _ = delta_tx.try_send(StreamDelta::Reasoning(t.to_string())); + } + if let Some(tc_arr) = delta["tool_calls"].as_array() { + for tc in tc_arr { + let idx = tc["index"].as_u64().unwrap_or(0); + let entry = tool_calls.entry(idx).or_default(); + if let Some(id) = tc["id"].as_str() { entry.0 = id.to_string(); } + if let Some(n) = tc["function"]["name"].as_str() { entry.1 = n.to_string(); } + if let Some(a) = tc["function"]["arguments"].as_str() { entry.2.push_str(a); } + } + } + }; + + while let Some(chunk) = byte_stream.next().await { + let chunk = chunk.map_err(ModelError::from_reqwest)?; + for payload in sse.feed(&chunk) { + handle_payload(&payload, emitted); + } + } + for payload in sse.finish() { + handle_payload(&payload, emitted); + } + + let finish = finish_reason.as_deref().unwrap_or("stop"); + let input_tokens = usage.as_ref().and_then(|u| u["prompt_tokens"].as_u64()).map(|n| n as u32); + let output_tokens = usage.as_ref().and_then(|u| u["completion_tokens"].as_u64()).map(|n| n as u32); + let cache_read = usage.as_ref() + .and_then(|u| u["prompt_tokens_details"]["cached_tokens"].as_u64()) + .map(|n| n as u32); + let cost_usd = usage.as_ref().and_then(|u| u["cost"].as_f64()); + let reasoning_content = if reasoning.is_empty() { None } else { Some(reasoning) }; + info!(model = %req.model, ?input_tokens, ?output_tokens, finish_reason = finish, "openai: streaming response completed"); + if finish == "length" { + warn!(model = %req.model, ?output_tokens, "openai: response truncated (max_tokens reached)"); + } + + let usage_struct = Usage { + input_tokens, + output_tokens, + cache_read, + cache_write: None, + cost_usd, + truncated: finish == "length", + }; + + // Reassemble the streamed message for the payload log (buffered shape). + let logged_tool_calls: Vec = tool_calls.iter() + .map(|(_idx, (id, name, args))| json!({ + "id": id, + "type": "function", + "function": { "name": name, "arguments": args }, + })) + .collect(); + let mut logged_message = json!({ "role": "assistant", "content": content.clone() }); + if let Some(rc) = &reasoning_content { + logged_message["reasoning_content"] = rc.clone().into(); + } + if !logged_tool_calls.is_empty() { + logged_message["tool_calls"] = Value::Array(logged_tool_calls); + } + let raw = RawMeta { + request_headers: Some(request_headers), + request_body: Some(request_body), + response_headers: Some(response_headers), + response_body: Some(json!({ + "streamed": true, + "choices": [{ "finish_reason": finish, "message": logged_message }], + "usage": usage, + })), + }; + + let mut resp = if !tool_calls.is_empty() { + let calls = tool_calls + .into_values() + .map(|(id, name, args)| ToolCall { + id, + name, + arguments: serde_json::from_str(&args).unwrap_or(Value::Object(Default::default())), + }) + .collect(); + ModelResponse::ToolCalls { content, calls, reasoning: reasoning_content, usage: usage_struct, raw: None } + } else { + ModelResponse::Message { content, reasoning: reasoning_content, usage: usage_struct, raw: None } + }; + set_raw(&mut resp, raw); + Ok(resp) + } +} + +impl NamedModel for OpenAiModel { + fn default_model(&self) -> &str { &self.default_model } +} + +#[async_trait] +impl Model for OpenAiModel { + async fn complete( + &self, + req: &ModelRequest, + deltas: Option>, + ) -> Result { + match deltas { + None => self.buffered(req).await, + Some(delta_tx) => { + let mut emitted = false; + match self.stream_chat(req, &delta_tx, &mut emitted).await { + Ok(ok) => Ok(ok), + // Nothing was ever streamed: some OpenAI-compatible + // providers reject `stream`/`stream_options` outright — + // retry buffered so they keep working. A mid-stream + // failure instead propagates to the fallback logic. + Err(e) if !emitted => { + debug!(model = %req.model, error = %e, "openai: streaming failed before any delta; retrying buffered"); + self.buffered(req).await + } + Err(e) => Err(e), + } + } + } + } +} + +// ── response parsing (shared by buffered and tests) ── + +trait WithRaw { + fn with_raw(self, raw: RawMeta) -> ModelResponse; +} + +impl WithRaw for ModelResponse { + fn with_raw(mut self, raw: RawMeta) -> ModelResponse { + set_raw(&mut self, raw); + self + } +} + +fn set_raw(resp: &mut ModelResponse, raw: RawMeta) { + match resp { + ModelResponse::Message { raw: r, .. } | ModelResponse::ToolCalls { raw: r, .. } => { + *r = Some(raw) + } + } +} + +/// Parse a buffered OpenAI response body into a `ModelResponse`. +fn parse_turn(resp: &Value, model: &str) -> ModelResponse { + let usage = Usage { + input_tokens: resp["usage"]["prompt_tokens"].as_u64().map(|n| n as u32), + output_tokens: resp["usage"]["completion_tokens"].as_u64().map(|n| n as u32), + cache_read: resp["usage"]["prompt_tokens_details"]["cached_tokens"].as_u64().map(|n| n as u32), + cache_write: None, + cost_usd: resp["usage"]["cost"].as_f64(), + truncated: false, + }; + + let choice = &resp["choices"][0]; + let message = &choice["message"]; + let finish = choice["finish_reason"].as_str().unwrap_or("stop"); + if finish == "length" { + warn!(model = %model, "openai: response truncated (max_tokens reached)"); + } + + let reasoning_content = message["reasoning_content"].as_str() + .or_else(|| message["reasoning"].as_str()) + .map(str::to_string); + + let tool_calls_array = message["tool_calls"].as_array().filter(|a| !a.is_empty()); + + // Some models (e.g. Qwen via OpenRouter) return finish_reason "stop" even + // when tool_calls are present, so check the array directly. + if finish == "tool_calls" || tool_calls_array.is_some() { + let content = message["content"].as_str().unwrap_or("").to_string(); + let calls = tool_calls_array + .map(|arr| { + arr.iter() + .map(|tc| ToolCall { + id: tc["id"].as_str().unwrap_or("").to_string(), + name: tc["function"]["name"].as_str().unwrap_or("").to_string(), + arguments: tc["function"]["arguments"] + .as_str() + .and_then(|s| serde_json::from_str(s).ok()) + .unwrap_or(Value::Object(Default::default())), + }) + .collect() + }) + .unwrap_or_default(); + ModelResponse::ToolCalls { content, calls, reasoning: reasoning_content, usage, raw: None } + } else { + // content can be null for thinking models or finish_reason="length". + let content = match message["content"].as_str() { + Some(s) => s.to_string(), + None => { + warn!(finish_reason = finish, raw_message = %message, "openai: response has null content"); + String::new() + } + }; + let mut usage = usage; + usage.truncated = finish == "length"; + ModelResponse::Message { content, reasoning: reasoning_content, usage, raw: None } + } +} diff --git a/crates/agent-loop/src/models/sse.rs b/crates/agent-loop/src/models/sse.rs new file mode 100644 index 0000000..9a0eaee --- /dev/null +++ b/crates/agent-loop/src/models/sse.rs @@ -0,0 +1,80 @@ +//! Incremental SSE decoder: feed raw response bytes, get back the payload of +//! every complete `data:` line seen (`[DONE]` included — callers decide). +//! Buffers partial lines across chunks; `event:` lines and comments are +//! skipped (both OpenAI and Anthropic put the event type inside the JSON). +//! +//! Ported verbatim from `llm-client`. + +#[derive(Default)] +pub(crate) struct SseDecoder { + buf: Vec, +} + +impl SseDecoder { + pub(crate) fn new() -> Self { Self::default() } + + pub(crate) fn feed(&mut self, bytes: &[u8]) -> Vec { + self.buf.extend_from_slice(bytes); + let mut out = Vec::new(); + while let Some(pos) = self.buf.iter().position(|&b| b == b'\n') { + let line: Vec = self.buf.drain(..=pos).collect(); + if let Some(payload) = parse_sse_line(&line) { + out.push(payload); + } + } + out + } + + /// Flush a trailing line not terminated by `\n` at end-of-stream. + pub(crate) fn finish(&mut self) -> Vec { + let rest = std::mem::take(&mut self.buf); + parse_sse_line(&rest).into_iter().collect() + } +} + +/// A complete SSE line is valid UTF-8 (a multibyte sequence never contains a +/// `\n` byte), but decode lossily anyway — a corrupt line is skipped, not fatal. +fn parse_sse_line(line: &[u8]) -> Option { + let line = String::from_utf8_lossy(line); + let line = line.trim_end_matches('\r').trim(); + let data = line.strip_prefix("data:")?.trim_start(); + if data.is_empty() { None } else { Some(data.to_string()) } +} + +#[cfg(test)] +mod tests { + use super::SseDecoder; + + #[test] + fn sse_decoder_buffers_partial_lines_across_chunks() { + let mut dec = SseDecoder::new(); + assert!(dec.feed(br#"data: {"a": 1"#).is_empty()); + assert_eq!(dec.feed(b"}\r\n").len(), 1); + } + + #[test] + fn sse_decoder_skips_events_comments_and_keeps_done() { + let mut dec = SseDecoder::new(); + let out = dec.feed(b"event: message_start\n: ping\n\ndata: {\"type\":\"ping\"}\ndata: [DONE]\n"); + assert_eq!(out, vec!["{\"type\":\"ping\"}".to_string(), "[DONE]".to_string()]); + assert!(dec.finish().is_empty()); + } + + #[test] + fn sse_decoder_finish_flushes_unterminated_tail() { + let mut dec = SseDecoder::new(); + assert!(dec.feed(b"data: tail-without-newline").is_empty()); + assert_eq!(dec.finish(), vec!["tail-without-newline".to_string()]); + } + + #[test] + fn sse_decoder_handles_multibyte_split() { + // "€" is 3 bytes in UTF-8; split across the chunk boundary. + let payload = "data: {\"t\":\"€\"}\n".as_bytes(); + let (a, b) = payload.split_at(12); + let mut dec = SseDecoder::new(); + let (first, second) = (dec.feed(a), dec.feed(b)); + assert!(first.is_empty()); + assert_eq!(second.len(), 1); + } +} diff --git a/crates/agent-loop/src/store.rs b/crates/agent-loop/src/store.rs new file mode 100644 index 0000000..06508a9 --- /dev/null +++ b/crates/agent-loop/src/store.rs @@ -0,0 +1,281 @@ +//! `HistoryStore` — the durability heart of the loop. +//! +//! Contract (enforced by doc, relied upon by recovery): +//! +//! 1. **Every state transition is an immediate write** — the kernel never +//! accumulates state in RAM. A crash loses only RAM, never truth. +//! 2. `MessageId`/`ToolCallId` are **monotonically increasing per frame**. +//! 3. `resolve_call` is the ONLY path to terminal states; `set_call_state` +//! is only for `Running → AwaitingHuman`. +//! 4. `load` returns calls nested inside their messages — the input of the +//! assembler's well-formed projection. + +use async_trait::async_trait; +use serde_json::Value; + +use crate::ids::{ConversationId, FrameId, MessageId, SummaryId, ToolCallId}; +use crate::model::Usage; + +// ── Role ───────────────────────────────────────────────────────────────────── + +/// Who produced a message. `Agent` is an injected agent-to-agent message +/// (sub-agent prompt, async result delivery); it projects to `user` on the wire. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Role { + System, + User, + Assistant, + Agent, +} + +// ── CallState ──────────────────────────────────────────────────────────────── + +/// Lifecycle of a tool call — semantics identical to Skald's +/// `chat_llm_tools.status`. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CallState { + /// Was executing at crash time → interrupted (NOT terminal). + Running, + /// 'pending': approval or clarification in flight (NOT terminal). + AwaitingHuman, + /// Terminal. + Done, + /// Terminal. + Failed, + /// Deliberate /stop — NEVER re-execute. + Cancelled, + /// Policy/human denial — NEVER re-execute. + Rejected, +} + +impl CallState { + pub fn is_terminal(self) -> bool { + matches!(self, Self::Done | Self::Failed | Self::Cancelled | Self::Rejected) + } +} + +// ── CallOutcome ────────────────────────────────────────────────────────────── + +/// The result of an execution, before recording. +#[derive(Debug, Clone)] +pub enum CallOutcome { + Completed(crate::tool::ToolOutput), + Failed(String), + Cancelled, + Rejected { reason: String }, +} + +impl CallOutcome { + pub fn state(&self) -> CallState { + match self { + Self::Completed(_) => CallState::Done, + Self::Failed(_) => CallState::Failed, + Self::Cancelled => CallState::Cancelled, + Self::Rejected { .. } => CallState::Rejected, + } + } + + /// Text persisted as the call's result (what the model will read back). + pub fn result_text(&self) -> String { + match self { + Self::Completed(out) => out.to_wire(), + Self::Failed(e) => format!("Error: {e}"), + Self::Cancelled => "Tool call cancelled by user.".to_string(), + Self::Rejected { reason } => format!("Tool call rejected: {reason}"), + } + } + + pub fn result_kind(&self) -> &'static str { + match self { + Self::Completed(out) => out.kind(), + Self::Failed(_) => "error", + Self::Cancelled => "cancelled", + Self::Rejected { .. } => "rejected", + } + } +} + +// ── Frames ─────────────────────────────────────────────────────────────────── + +/// What a frame is opened with (a sub-agent dispatch; the root carries the +/// conversation's entry agent). +#[derive(Debug, Clone)] +pub struct FrameSpec { + /// Agent id in the HOST's catalog (opaque to the crate). + pub agent: String, + /// The sub-agent's prompt (root: None). + pub prompt: Option, + pub depth: u32, + /// The parent frame's tool call that spawned this frame. + pub parent_call: Option, + /// Host free-form (run_context_json, …). + pub meta: Value, +} + +impl FrameSpec { + pub fn root(agent: impl Into) -> Self { + Self { + agent: agent.into(), + prompt: None, + depth: 0, + parent_call: None, + meta: Value::Null, + } + } +} + +/// A stored frame. +#[derive(Debug, Clone)] +pub struct FrameRecord { + pub id: FrameId, + pub conversation: ConversationId, + pub parent: Option, + pub spec: FrameSpec, + pub active: bool, +} + +// ── Messages ───────────────────────────────────────────────────────────────── + +#[derive(Debug, Clone)] +pub struct NewMessage { + pub role: Role, + pub content: String, + /// TIC/notify/injection: not echoed to the UI as a user message. + pub synthetic: bool, + pub reasoning: Option, + /// Attachments, command display, … (host free-form). + pub metadata: Option, +} + +impl NewMessage { + pub fn user(content: impl Into) -> Self { + Self { role: Role::User, content: content.into(), synthetic: false, reasoning: None, metadata: None } + } + + pub fn assistant(content: impl Into, reasoning: Option) -> Self { + Self { role: Role::Assistant, content: content.into(), synthetic: false, reasoning, metadata: None } + } + + pub fn agent(content: impl Into) -> Self { + Self { role: Role::Agent, content: content.into(), synthetic: false, reasoning: None, metadata: None } + } + + pub fn synthetic(mut self, synthetic: bool) -> Self { + self.synthetic = synthetic; + self + } + + pub fn with_metadata(mut self, metadata: Value) -> Self { + self.metadata = Some(metadata); + self + } +} + +/// A stored message with its tool calls nested. +#[derive(Debug, Clone)] +pub struct StoredMessage { + pub id: MessageId, + pub role: Role, + pub content: String, + pub reasoning: Option, + pub synthetic: bool, + /// Orphan of a cancelled turn — excluded from `load`. + pub failed: bool, + pub metadata: Option, + pub usage: Usage, + pub calls: Vec, +} + +// ── Tool calls ─────────────────────────────────────────────────────────────── + +/// What a call is recorded with, BEFORE execution (phase 1 of the fan-out). +#[derive(Debug, Clone)] +pub struct NewCall { + /// The model's wire call id ("call_abc", "toolu_…"), needed to rebuild + /// `tool_calls`/`tool` wire messages. Synthesized by the store when absent. + pub provider_id: Option, + pub name: String, + pub arguments: Value, +} + +impl NewCall { + pub fn new(name: impl Into, arguments: Value) -> Self { + Self { provider_id: None, name: name.into(), arguments } + } + + pub fn with_provider_id(mut self, id: impl Into) -> Self { + self.provider_id = Some(id.into()); + self + } +} + +/// A stored tool call. +#[derive(Debug, Clone)] +pub struct StoredCall { + pub id: ToolCallId, + pub message_id: MessageId, + /// The model's wire call id (see [`NewCall::provider_id`]). + pub provider_id: String, + pub name: String, + pub arguments: Value, + pub state: CallState, + pub result: Option, + pub result_kind: String, + /// Host free-form (Skald: preview_old/new, media refs). + pub extras: Value, +} + +// ── Summaries ──────────────────────────────────────────────────────────────── + +#[derive(Debug, Clone)] +pub struct NewSummary { + pub text: String, + /// Last message covered by the summary — the projection resumes after it. + pub covered_up_to: MessageId, +} + +#[derive(Debug, Clone)] +pub struct StoredSummary { + pub id: SummaryId, + pub text: String, + pub covered_up_to: MessageId, +} + +// ── HistoryStore ───────────────────────────────────────────────────────────── + +#[async_trait] +pub trait HistoryStore: Send + Sync { + // ── frames ── + async fn open_frame( + &self, + conv: &ConversationId, + parent: Option, + spec: FrameSpec, + ) -> crate::Result; + async fn close_frame(&self, frame: FrameId) -> crate::Result<()>; + /// All active frames of a conversation (recovery: batch detection, cascade). + async fn active_frames(&self, conv: &ConversationId) -> crate::Result>; + async fn deepest_active(&self, conv: &ConversationId) -> crate::Result>; + + // ── messages ── + async fn append(&self, frame: FrameId, msg: NewMessage) -> crate::Result; + async fn set_usage(&self, msg: MessageId, usage: &Usage) -> crate::Result<()>; + /// Frame history with calls nested per message. EXCLUDES failed messages + /// (orphans of cancelled turns). + async fn load(&self, frame: FrameId) -> crate::Result>; + async fn load_since(&self, frame: FrameId, after: MessageId) -> crate::Result>; + async fn last(&self, frame: FrameId) -> crate::Result>; + async fn mark_failed(&self, msg: MessageId) -> crate::Result<()>; + + // ── tool calls ── + async fn append_call(&self, msg: MessageId, call: NewCall) -> crate::Result; + /// The ONLY path to terminal states. + async fn resolve_call(&self, id: ToolCallId, outcome: &CallOutcome) -> crate::Result<()>; + /// Only `Running → AwaitingHuman`. + async fn set_call_state(&self, id: ToolCallId, state: CallState) -> crate::Result<()>; + async fn calls_in_state(&self, frame: FrameId, states: &[CallState]) -> crate::Result>; + + // ── summaries ── + async fn save_summary(&self, frame: FrameId, s: NewSummary) -> crate::Result; + async fn latest_summary(&self, frame: FrameId) -> crate::Result>; +} diff --git a/crates/agent-loop/src/store_memory.rs b/crates/agent-loop/src/store_memory.rs new file mode 100644 index 0000000..add64a8 --- /dev/null +++ b/crates/agent-loop/src/store_memory.rs @@ -0,0 +1,256 @@ +//! `InMemoryStore` — the shipped non-persistent store (chat not persisted; +//! testing; simple hosts). Monotonic ids per the store contract. + +use std::collections::HashMap; +use std::sync::Mutex; + +use async_trait::async_trait; + +use crate::ids::{ConversationId, FrameId, MessageId, SummaryId, ToolCallId}; +use crate::model::Usage; +use crate::store::{ + CallOutcome, CallState, FrameRecord, FrameSpec, HistoryStore, NewCall, NewMessage, NewSummary, + StoredCall, StoredMessage, StoredSummary, +}; + +#[derive(Default)] +struct Inner { + frames: HashMap, + messages: HashMap>, + calls: HashMap>, + summaries: HashMap>, + next_frame: i64, + next_msg: i64, + next_call: i64, + next_summary: i64, +} + +/// Non-persistent store. A "crash" loses everything — which is exactly why +/// it's also the natural target for recovery scenario tests (build the +/// post-crash state by hand). +pub struct InMemoryStore { + inner: Mutex, +} + +impl InMemoryStore { + pub fn new() -> Self { Self { inner: Mutex::new(Inner::default()) } } +} + +impl Default for InMemoryStore { + fn default() -> Self { Self::new() } +} + +#[async_trait] +impl HistoryStore for InMemoryStore { + async fn open_frame( + &self, + conv: &ConversationId, + parent: Option, + spec: FrameSpec, + ) -> crate::Result { + let mut i = self.inner.lock().unwrap(); + i.next_frame += 1; + let id = FrameId(i.next_frame); + i.frames.insert(id, FrameRecord { + id, + conversation: conv.clone(), + parent, + spec, + active: true, + }); + Ok(id) + } + + async fn close_frame(&self, frame: FrameId) -> crate::Result<()> { + let mut i = self.inner.lock().unwrap(); + if let Some(f) = i.frames.get_mut(&frame) { + f.active = false; + } + Ok(()) + } + + async fn active_frames(&self, conv: &ConversationId) -> crate::Result> { + let i = self.inner.lock().unwrap(); + Ok(i.frames.values().filter(|f| f.active && &f.conversation == conv).cloned().collect()) + } + + async fn deepest_active(&self, conv: &ConversationId) -> crate::Result> { + let i = self.inner.lock().unwrap(); + Ok(i.frames + .values() + .filter(|f| f.active && &f.conversation == conv) + .max_by_key(|f| f.spec.depth) + .cloned()) + } + + async fn append(&self, frame: FrameId, msg: NewMessage) -> crate::Result { + let mut i = self.inner.lock().unwrap(); + i.next_msg += 1; + let id = MessageId(i.next_msg); + i.messages.entry(frame).or_default().push(StoredMessage { + id, + role: msg.role, + content: msg.content, + reasoning: msg.reasoning, + synthetic: msg.synthetic, + failed: false, + metadata: msg.metadata, + usage: Usage::default(), + calls: Vec::new(), + }); + Ok(id) + } + + async fn set_usage(&self, msg: MessageId, usage: &Usage) -> crate::Result<()> { + let mut i = self.inner.lock().unwrap(); + for msgs in i.messages.values_mut() { + if let Some(m) = msgs.iter_mut().find(|m| m.id == msg) { + m.usage = usage.clone(); + return Ok(()); + } + } + Ok(()) + } + + async fn load(&self, frame: FrameId) -> crate::Result> { + let i = self.inner.lock().unwrap(); + Ok(load_frame(&i, frame, None)) + } + + async fn load_since(&self, frame: FrameId, after: MessageId) -> crate::Result> { + let i = self.inner.lock().unwrap(); + Ok(load_frame(&i, frame, Some(after))) + } + + async fn last(&self, frame: FrameId) -> crate::Result> { + let i = self.inner.lock().unwrap(); + Ok(load_frame(&i, frame, None).into_iter().last()) + } + + async fn mark_failed(&self, msg: MessageId) -> crate::Result<()> { + let mut i = self.inner.lock().unwrap(); + for msgs in i.messages.values_mut() { + if let Some(m) = msgs.iter_mut().find(|m| m.id == msg) { + m.failed = true; + return Ok(()); + } + } + Ok(()) + } + + async fn append_call(&self, msg: MessageId, call: NewCall) -> crate::Result { + let mut i = self.inner.lock().unwrap(); + i.next_call += 1; + let id = ToolCallId(i.next_call); + let provider_id = call.provider_id.unwrap_or_else(|| format!("call_{}", id.get())); + let stored = StoredCall { + id, + message_id: msg, + provider_id, + name: call.name, + arguments: call.arguments, + state: CallState::Running, + result: None, + result_kind: String::new(), + extras: serde_json::Value::Null, + }; + i.calls.entry(msg).or_default().push(stored.clone()); + // Keep the nested copy inside the message in sync. + for msgs in i.messages.values_mut() { + if let Some(m) = msgs.iter_mut().find(|m| m.id == msg) { + m.calls.push(stored); + break; + } + } + Ok(id) + } + + async fn resolve_call(&self, id: ToolCallId, outcome: &CallOutcome) -> crate::Result<()> { + let mut i = self.inner.lock().unwrap(); + update_call(&mut i, id, |c| { + c.state = outcome.state(); + c.result = Some(outcome.result_text()); + c.result_kind = outcome.result_kind().to_string(); + }); + Ok(()) + } + + async fn set_call_state(&self, id: ToolCallId, state: CallState) -> crate::Result<()> { + anyhow::ensure!( + !state.is_terminal(), + "set_call_state is only for Running → AwaitingHuman, not terminal {state:?}" + ); + let mut i = self.inner.lock().unwrap(); + update_call(&mut i, id, |c| c.state = state); + Ok(()) + } + + async fn calls_in_state(&self, frame: FrameId, states: &[CallState]) -> crate::Result> { + let i = self.inner.lock().unwrap(); + Ok(i.messages + .get(&frame) + .map(|msgs| { + msgs.iter() + .flat_map(|m| &m.calls) + .filter(|c| states.contains(&c.state)) + .cloned() + .collect() + }) + .unwrap_or_default()) + } + + async fn save_summary(&self, frame: FrameId, s: NewSummary) -> crate::Result { + let mut i = self.inner.lock().unwrap(); + i.next_summary += 1; + let id = SummaryId(i.next_summary); + i.summaries.entry(frame).or_default().push(StoredSummary { + id, + text: s.text, + covered_up_to: s.covered_up_to, + }); + Ok(id) + } + + async fn latest_summary(&self, frame: FrameId) -> crate::Result> { + let i = self.inner.lock().unwrap(); + Ok(i.summaries.get(&frame).and_then(|v| v.last()).cloned()) + } +} + +/// Load a frame's history with calls nested, excluding failed messages, +/// optionally only messages after `after`. +fn load_frame(i: &Inner, frame: FrameId, after: Option) -> Vec { + i.messages + .get(&frame) + .map(|msgs| { + msgs.iter() + .filter(|m| !m.failed) + .filter(|m| after.is_none_or(|a| m.id > a)) + .cloned() + .collect() + }) + .unwrap_or_default() +} + +/// Apply a mutation to a call both in the by-message index and in the nested +/// copy inside its message. +fn update_call(i: &mut Inner, id: ToolCallId, f: impl Fn(&mut StoredCall)) { + let mut msg_id = None; + for calls in i.calls.values_mut() { + if let Some(c) = calls.iter_mut().find(|c| c.id == id) { + f(c); + msg_id = Some(c.message_id); + break; + } + } + if let Some(msg_id) = msg_id { + for msgs in i.messages.values_mut() { + if let Some(m) = msgs.iter_mut().find(|m| m.id == msg_id) { + if let Some(c) = m.calls.iter_mut().find(|c| c.id == id) { + f(c); + } + break; + } + } + } +} diff --git a/crates/agent-loop/src/testing.rs b/crates/agent-loop/src/testing.rs new file mode 100644 index 0000000..8a1f62a --- /dev/null +++ b/crates/agent-loop/src/testing.rs @@ -0,0 +1,130 @@ +//! Test utilities: a scripted `FakeModel` + builders for kernel and recovery +//! scenarios. (Blueprint: will move behind a `test-util` feature if the crate +//! is ever published.) + +use std::collections::VecDeque; +use std::sync::Mutex; + +use async_trait::async_trait; +use tokio::sync::mpsc; + +use crate::model::{ + Model, ModelError, ModelRequest, ModelResponse, NamedModel, StreamDelta, ToolCall, Usage, +}; + +/// One scripted step: the response (or error) plus optional deltas to emit +/// before returning. +pub struct Step { + pub result: Result, + pub deltas: Vec, + /// Never return (cancellation tests). + pub pending: bool, +} + +impl Step { + pub fn message(content: impl Into) -> Self { + Self { result: Ok(ModelResponse::message(content)), deltas: Vec::new(), pending: false } + } + + pub fn message_with_usage(content: impl Into, input: u32, output: u32) -> Self { + let mut resp = ModelResponse::message(content); + *resp.usage_mut() = Usage { + input_tokens: Some(input), + output_tokens: Some(output), + ..Usage::default() + }; + Self { result: Ok(resp), deltas: Vec::new(), pending: false } + } + + pub fn tool_calls(content: impl Into, calls: Vec) -> Self { + Self { result: Ok(ModelResponse::tool_calls(content, calls)), deltas: Vec::new(), pending: false } + } + + pub fn error(status: Option, message: impl Into) -> Self { + Self { result: Err(ModelError::new(status, message)), deltas: Vec::new(), pending: false } + } + + /// Never completes — the only way out is cancelling the turn. + pub fn pending() -> Self { + Self { result: Ok(ModelResponse::message("")), deltas: Vec::new(), pending: true } + } + + /// Stream these deltas (in order) before returning the response. + pub fn with_deltas(mut self, deltas: Vec) -> Self { + self.deltas = deltas; + self + } +} + +/// A scripted model: pops one [`Step`] per `complete` call, records every +/// request for assertions. Clone the `Arc` around it to inspect afterwards. +pub struct FakeModel { + script: Mutex>, + requests: Mutex>, + default_model: String, +} + +impl FakeModel { + pub fn new(default_model: impl Into, script: Vec) -> Self { + Self { + script: Mutex::new(script.into()), + requests: Mutex::new(Vec::new()), + default_model: default_model.into(), + } + } + + /// All requests seen so far (one per attempt, fallback included). + pub fn requests(&self) -> Vec { + self.requests.lock().unwrap().clone() + } + + /// Steps not yet consumed (assert a script was fully driven). + pub fn remaining(&self) -> usize { + self.script.lock().unwrap().len() + } +} + +impl NamedModel for FakeModel { + fn default_model(&self) -> &str { &self.default_model } +} + +#[async_trait] +impl Model for FakeModel { + async fn complete( + &self, + req: &ModelRequest, + deltas: Option>, + ) -> Result { + self.requests.lock().unwrap().push(req.clone()); + let step = self + .script + .lock() + .unwrap() + .pop_front() + .unwrap_or_else(|| panic!("FakeModel: script exhausted (request for model {})", req.model)); + if let Some(tx) = deltas { + for d in step.deltas { + let _ = tx.try_send(d); + } + } + if step.pending { + std::future::pending::<()>().await; + } + step.result + } +} + +/// A `ModelHandle` over a shared `FakeModel` (tests keep the Arc to inspect +/// `requests()` afterwards). +pub fn handle(fake: &std::sync::Arc, id: &str) -> crate::model::ModelHandle { + crate::model::ModelHandle { + id: id.to_string(), + model: fake.clone(), + info: crate::model::ModelInfo::default(), + } +} + +/// Build a wire `ToolCall` compactly in tests. +pub fn call(id: &str, name: &str, args: serde_json::Value) -> ToolCall { + ToolCall { id: id.to_string(), name: name.to_string(), arguments: args } +} diff --git a/crates/agent-loop/src/tool.rs b/crates/agent-loop/src/tool.rs new file mode 100644 index 0000000..1ef2218 --- /dev/null +++ b/crates/agent-loop/src/tool.rs @@ -0,0 +1,365 @@ +//! The `Tool` trait, the type-erased [`ToolCtx`] (blueprint D3 — a type-map, +//! axum/tower style, not generics), and the cancellable execution machinery +//! (ported verbatim from Skald's core-api: it was already pure). + +use std::any::{Any, TypeId}; +use std::collections::HashMap; +use std::future::Future; +use std::pin::Pin; +use std::sync::{Arc, Mutex}; + +use async_trait::async_trait; +use serde_json::Value; +use tokio_util::sync::CancellationToken; + +use crate::ids::{ConversationId, FrameId, ToolCallId}; + +// ── Extensions ─────────────────────────────────────────────────────────────── + +/// A type-map of host values threaded into every tool call (axum/tower +/// style). Hosts insert in ONE place (turn construction) and read with typed +/// helpers — never scattered string keys. +#[derive(Clone, Default)] +pub struct Extensions { + map: HashMap>, +} + +impl Extensions { + pub fn new() -> Self { Self::default() } + + pub fn insert(&mut self, value: Arc) -> &mut Self { + self.map.insert(TypeId::of::(), value); + self + } + + pub fn get(&self) -> Option> { + self.map.get(&TypeId::of::())?.clone().downcast::().ok() + } + + pub fn contains(&self) -> bool { + self.map.contains_key(&TypeId::of::()) + } +} + +impl std::fmt::Debug for Extensions { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "Extensions({} entries)", self.map.len()) + } +} + +// ── ToolCtx ────────────────────────────────────────────────────────────────── + +/// Per-invocation execution context threaded into a tool call. +#[derive(Clone)] +pub struct ToolCtx { + pub conversation: ConversationId, + pub frame: FrameId, + /// Agent of the current frame (self-call check for delegation). + pub agent: String, + /// The call being executed (parent_call of any child frame). + pub call_id: ToolCallId, + pub cancel: CancellationToken, + pub extensions: Extensions, +} + +// ── ToolOutput / ToolFailure ───────────────────────────────────────────────── + +/// A reference to one media file a tool produced. The assembler decides +/// whether to inline it — the kernel only transports it. +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] +pub struct MediaRef { + /// Absolute host path, already containment-checked by the producing tool. + pub host_path: String, + /// Sniffed MIME (informational — pipelines re-sniff from bytes). + pub mime: String, +} + +/// The successful output of a tool. +#[derive(Debug, Clone)] +pub enum ToolOutput { + Text(String), + Json(Value), + /// A text note plus media refs; the wire message carries only `text`. + Media { text: String, refs: Vec }, +} + +impl ToolOutput { + /// Canonical string form persisted as the call result and replayed to the + /// model (both OpenAI and Anthropic encode tool results as text/JSON). + pub fn to_wire(&self) -> String { + match self { + Self::Text(s) => s.clone(), + Self::Json(v) => serde_json::to_string(v).unwrap_or_else(|_| "null".into()), + Self::Media { text, .. } => text.clone(), + } + } + + pub fn kind(&self) -> &'static str { + match self { + Self::Text(_) | Self::Media { .. } => "string", + Self::Json(_) => "json", + } + } + + pub fn media(&self) -> &[MediaRef] { + match self { + Self::Media { refs, .. } => refs, + _ => &[], + } + } +} + +impl From for ToolOutput { + fn from(s: String) -> Self { Self::Text(s) } +} +impl From<&str> for ToolOutput { + fn from(s: &str) -> Self { Self::Text(s.to_string()) } +} + +/// How a tool call can fail. +#[derive(Debug, Clone)] +pub enum ToolFailure { + Failed(String), + /// The tool suspended waiting for a human and the channel closed: the turn + /// ends, the call STAYS `AwaitingHuman` for the resume. (The tool marks + /// the call `AwaitingHuman` via the store BEFORE returning this.) + Suspend, +} + +impl std::fmt::Display for ToolFailure { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Failed(e) => write!(f, "{e}"), + Self::Suspend => write!(f, "tool suspended awaiting human input"), + } + } +} + +impl std::error::Error for ToolFailure {} + +// ── RestartHint / Visibility ───────────────────────────────────────────────── + +/// What recovery does with a call that was `Running` at crash (blueprint D7). +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub enum RestartHint { + /// Re-gate and re-execute (default — today's behavior; idempotent tools). + #[default] + ReExecute, + /// Resolve as Failed "interrupted" (tools with non-idempotent external + /// side effects, e.g. shell commands). + MarkInterrupted, +} + +/// Declared visibility — the HOST filters at `ToolSet` construction, the +/// kernel never filters. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub enum Visibility { + #[default] + Always, + InteractiveOnly, + RootOnly, + SubAgentsOnly, +} + +// ── Tool ───────────────────────────────────────────────────────────────────── + +/// A single LLM-callable tool. +#[async_trait] +pub trait Tool: Send + Sync { + fn name(&self) -> &str; + + /// OpenAI-shaped tool definition (`{"type":"function","function":{…}}`). + fn definition(&self) -> Value; + + /// The simple execution path. The kernel wraps it in a [`SimpleExecution`] + /// by default (drop of the future = stop) — override [`start`](Self::start) + /// for remote/child teardown instead. + async fn call(&self, args: Value, ctx: &ToolCtx) -> Result; + + /// May this call run in parallel with other concurrency-safe calls of the + /// same round? (Generalized sub-agent batch, blueprint §7.) Default false + /// → the sequential path. + fn concurrency_safe(&self, _args: &Value) -> bool { false } + + /// Recovery behavior when the call was `Running` at crash (D7). + fn restart_hint(&self) -> RestartHint { RestartHint::ReExecute } + + /// Declared visibility (host-side filtering only). + fn visibility(&self) -> Visibility { Visibility::Always } + + /// Start one execution, returning a live handle. The default wraps + /// [`call`](Self::call) in a [`SimpleExecution`]. Tools needing + /// remote/child teardown (kill a process group, POST an /interrupt) + /// override this with a bespoke [`ToolExecution::stop`]. + fn start<'a>(&'a self, args: Value, ctx: &'a ToolCtx) -> Box { + Box::new(SimpleExecution::new(Box::pin(self.call(args, ctx)))) + } +} + +// ── ToolSet ────────────────────────────────────────────────────────────────── + +/// The per-turn tool registry, ALREADY filtered by the host (visibility, +/// approval, interactive). `defs` is re-read at EVERY round and every +/// fallback attempt: grants activated at round N are visible at round N+1, +/// and a cross-mode DTL fallback re-shapes for free. +pub trait ToolSet: Send + Sync { + fn defs(&self, model: &crate::model::ModelInfo) -> Vec; + fn find(&self, name: &str) -> Option>; +} + +/// A trivial `ToolSet` from a list of tools (testing, simple hosts). +pub struct ToolRegistry { + tools: Vec>, +} + +impl ToolRegistry { + pub fn new() -> Self { Self { tools: Vec::new() } } + + pub fn with(mut self, tool: impl Tool + 'static) -> Self { + self.tools.push(Arc::new(tool)); + self + } + + pub fn with_arc(mut self, tool: Arc) -> Self { + self.tools.push(tool); + self + } + + pub fn into_toolset(self) -> Arc { Arc::new(self) } +} + +impl Default for ToolRegistry { + fn default() -> Self { Self::new() } +} + +impl ToolSet for ToolRegistry { + fn defs(&self, _model: &crate::model::ModelInfo) -> Vec { + self.tools.iter().map(|t| t.definition()).collect() + } + + fn find(&self, name: &str) -> Option> { + self.tools.iter().find(|t| t.name() == name).cloned() + } +} + +// ── ToolExecution ──────────────────────────────────────────────────────────── + +/// Lifecycle state of a single tool execution (in-memory, richer than the +/// persisted `CallState`). +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ToolExecutionState { + Pending, + Running, + Completed, + Failed, + Cancelled, +} + +/// Terminal outcome of [`ToolExecution::wait`]. +#[derive(Debug, Clone)] +pub enum ExecutionOutcome { + Completed(ToolOutput), + Failed(String), + Cancelled, + /// The tool suspended awaiting a human (`ToolFailure::Suspend`): the turn + /// ends and the call STAYS `AwaitingHuman` — never resolve it here. + Suspended, +} + +impl ExecutionOutcome { + pub fn into_call_outcome(self) -> crate::store::CallOutcome { + match self { + Self::Completed(out) => crate::store::CallOutcome::Completed(out), + Self::Failed(e) => crate::store::CallOutcome::Failed(e), + Self::Cancelled => crate::store::CallOutcome::Cancelled, + // Handled by the kernel before this conversion is reached. + Self::Suspended => crate::store::CallOutcome::Cancelled, + } + } +} + +/// A single live execution of a [`Tool`]. Pure: it never touches a store or a +/// transport — the kernel mirrors transitions to persistence and events. +pub trait ToolExecution: Send + Sync { + fn state(&self) -> ToolExecutionState; + /// Drive the work to its terminal outcome. Called exactly once. + fn wait<'a>(&'a self) -> Pin + Send + 'a>>; + /// Tool-specific cancellation. The default relies on the driver dropping + /// the `wait` future. + fn stop<'a>(&'a self) -> Pin + Send + 'a>> { + Box::pin(async {}) + } +} + +/// The boxed work unit inside a [`SimpleExecution`]. +pub type ToolWork<'a> = + Pin> + Send + 'a>>; + +/// Default [`ToolExecution`] for any tool that is a single async unit of work: +/// `wait` races the work against a stop-token, so `stop()` (or dropping +/// `wait`) aborts the in-flight I/O. +pub struct SimpleExecution<'a> { + state: Mutex, + stop: CancellationToken, + work: tokio::sync::Mutex>>, +} + +impl<'a> SimpleExecution<'a> { + pub fn new(work: ToolWork<'a>) -> Self { + Self { + state: Mutex::new(ToolExecutionState::Running), + stop: CancellationToken::new(), + work: tokio::sync::Mutex::new(Some(work)), + } + } +} + +impl ToolExecution for SimpleExecution<'_> { + fn state(&self) -> ToolExecutionState { *self.state.lock().unwrap() } + + fn wait<'b>(&'b self) -> Pin + Send + 'b>> { + Box::pin(async move { + let work = self.work.lock().await.take(); + let Some(work) = work else { return ExecutionOutcome::Cancelled }; + let outcome = tokio::select! { + biased; + _ = self.stop.cancelled() => ExecutionOutcome::Cancelled, + r = work => match r { + Ok(out) => ExecutionOutcome::Completed(out), + Err(ToolFailure::Failed(e)) => ExecutionOutcome::Failed(e), + Err(ToolFailure::Suspend) => ExecutionOutcome::Suspended, + }, + }; + *self.state.lock().unwrap() = match outcome { + ExecutionOutcome::Completed(_) => ToolExecutionState::Completed, + ExecutionOutcome::Failed(_) => ToolExecutionState::Failed, + ExecutionOutcome::Cancelled | ExecutionOutcome::Suspended => ToolExecutionState::Cancelled, + }; + outcome + }) + } + + fn stop<'b>(&'b self) -> Pin + Send + 'b>> { + Box::pin(async move { self.stop.cancel() }) + } +} + +/// Run a [`ToolExecution`] to completion honouring a cancellation token: on +/// cancel, `exec.stop()` is called once (tool-specific teardown), then `wait` +/// resolves. +pub async fn drive_execution(exec: &dyn ToolExecution, cancel: &CancellationToken) -> ExecutionOutcome { + let work = exec.wait(); + tokio::pin!(work); + + let mut stopped = false; + loop { + tokio::select! { + biased; + outcome = &mut work => return outcome, + _ = cancel.cancelled(), if !stopped => { + exec.stop().await; + stopped = true; + } + } + } +} diff --git a/crates/agent-loop/tests/assembler.rs b/crates/agent-loop/tests/assembler.rs new file mode 100644 index 0000000..e7fd12d --- /dev/null +++ b/crates/agent-loop/tests/assembler.rs @@ -0,0 +1,196 @@ +//! Assembler tests (blueprint §13): well-formed projection, DTL rendering +//! modes, summary, window, crash survivors. + +use std::sync::Arc; + +use agent_loop::activation::{Activation, ActivationSource, ToolRendering}; +use agent_loop::context::{AssembleInput, ContextAssembler, LinearAssembler, SystemContext}; +use agent_loop::ids::{ConversationId, FrameId}; +use agent_loop::model::ModelInfo; +use agent_loop::prelude::async_trait; +use agent_loop::store::{ + CallOutcome, FrameSpec, HistoryStore, NewCall, NewMessage, +}; +use agent_loop::store_memory::InMemoryStore; +use agent_loop::tool::ToolOutput; +use serde_json::{Value, json}; + +fn tool_def(name: &str) -> Value { + json!({"type":"function","function":{"name":name,"parameters":{"type":"object"}}}) +} + +struct StubActivations { + acts: Vec, +} + +#[async_trait] +impl ActivationSource for StubActivations { + async fn activations(&self, _frame: FrameId) -> agent_loop::Result> { + Ok(self.acts.clone()) + } +} + +fn model_info(mode: ToolRendering) -> ModelInfo { + ModelInfo { tool_rendering: mode, ..ModelInfo::default() } +} + +async fn input(store: &Arc, conv: &ConversationId, mode: ToolRendering) -> (FrameId, AssembleInput) { + let frame = store.open_frame(conv, None, FrameSpec::root("assistant")).await.unwrap(); + let input = AssembleInput { + frame, + system: SystemContext::base("BASE"), + model: model_info(mode), + round: 0, + }; + (frame, input) +} + +/// History: user → assistant with an activate_tools call (resolved) → final. +/// Returns the anchor (the assistant message id). +async fn seed_activation_history(store: &Arc, frame: FrameId) -> agent_loop::ids::MessageId { + store.append(frame, NewMessage::user("use gmail")).await.unwrap(); + let anchor = store.append(frame, NewMessage::assistant("activating", None)).await.unwrap(); + let call = store + .append_call(anchor, NewCall::new("activate_tools", json!({"groups":["gmail"]})).with_provider_id("c1")) + .await + .unwrap(); + store + .resolve_call(call, &CallOutcome::Completed(ToolOutput::Text("gmail activated".into()))) + .await + .unwrap(); + anchor +} + +#[tokio::test] +async fn inline_mode_injects_nothing() { + let store = Arc::new(InMemoryStore::new()); + let conv = ConversationId::new("a1"); + let (frame, input) = input(&store, &conv, ToolRendering::Inline).await; + let anchor = seed_activation_history(&store, frame).await; + + let assembler = LinearAssembler::new().with_activation(Arc::new(StubActivations { + acts: vec![Activation { anchor, defs: vec![tool_def("mcp__gmail__send")] }], + })); + let store_dyn: Arc = store; + let msgs = assembler.build(&store_dyn, &input).await.unwrap(); + + assert!(!msgs.iter().any(|m| m.get("tools").is_some()), "Inline must not inject system+tools"); + assert!(!msgs.iter().any(|m| m.get("_tool_references").is_some())); +} + +#[tokio::test] +async fn system_tool_block_appends_after_tool_results() { + let store = Arc::new(InMemoryStore::new()); + let conv = ConversationId::new("a2"); + let (frame, input) = input(&store, &conv, ToolRendering::SystemToolBlock).await; + let anchor = seed_activation_history(&store, frame).await; + + let assembler = LinearAssembler::new().with_activation(Arc::new(StubActivations { + acts: vec![Activation { anchor, defs: vec![tool_def("mcp__gmail__send")] }], + })); + let store_dyn: Arc = store; + let msgs = assembler.build(&store_dyn, &input).await.unwrap(); + + // [system BASE, user, assistant(tool_calls), tool(result), system+tools] + let block_idx = msgs + .iter() + .position(|m| m["role"].as_str() == Some("system") && m.get("tools").is_some()) + .expect("no system+tools block injected"); + assert_eq!(msgs[block_idx]["tools"][0]["function"]["name"], json!("mcp__gmail__send")); + assert!(msgs[block_idx].get("content").is_none(), "Kimi block has no content field"); + // It comes right after the tool result of the anchor group. + assert_eq!(msgs[block_idx - 1]["role"], json!("tool")); +} + +#[tokio::test] +async fn deferred_tool_reference_marks_first_tool_result() { + let store = Arc::new(InMemoryStore::new()); + let conv = ConversationId::new("a3"); + let (frame, input) = input(&store, &conv, ToolRendering::DeferredToolReference).await; + let anchor = seed_activation_history(&store, frame).await; + + let assembler = LinearAssembler::new().with_activation(Arc::new(StubActivations { + acts: vec![Activation { anchor, defs: vec![tool_def("mcp__gmail__send")] }], + })); + let store_dyn: Arc = store; + let msgs = assembler.build(&store_dyn, &input).await.unwrap(); + + let tool_msg = msgs + .iter() + .find(|m| m["role"].as_str() == Some("tool")) + .expect("no tool result projected"); + assert_eq!(tool_msg["_tool_references"], json!(["mcp__gmail__send"])); +} + +#[tokio::test] +async fn crash_survivors_get_synthetic_interrupted_results() { + let store = Arc::new(InMemoryStore::new()); + let conv = ConversationId::new("a4"); + let (frame, input) = input(&store, &conv, ToolRendering::Inline).await; + + store.append(frame, NewMessage::user("do it")).await.unwrap(); + let msg = store.append(frame, NewMessage::assistant("running", None)).await.unwrap(); + // Never resolved: still Running, as after a crash. + store.append_call(msg, NewCall::new("execute_cmd", json!({})).with_provider_id("c1")).await.unwrap(); + + let assembler = LinearAssembler::new(); + let store_dyn: Arc = store; + let msgs = assembler.build(&store_dyn, &input).await.unwrap(); + + let tool_msg = msgs.iter().find(|m| m["role"].as_str() == Some("tool")).unwrap(); + assert!( + tool_msg["content"].as_str().unwrap().contains("interrupted"), + "a Running survivor must project a synthetic interrupted result: {tool_msg}" + ); +} + +#[tokio::test] +async fn summary_replaces_covered_history() { + let store = Arc::new(InMemoryStore::new()); + let conv = ConversationId::new("a5"); + let (frame, input) = input(&store, &conv, ToolRendering::Inline).await; + + let m1 = store.append(frame, NewMessage::user("old question")).await.unwrap(); + store.append(frame, NewMessage::assistant("old answer", None)).await.unwrap(); + let m3 = store.append(frame, NewMessage::user("new question")).await.unwrap(); + + store + .save_summary(frame, agent_loop::store::NewSummary { + text: "User asked about old stuff.".into(), + covered_up_to: m1, + }) + .await + .unwrap(); + + let assembler = LinearAssembler::new(); + let store_dyn: Arc = store; + let msgs = assembler.build(&store_dyn, &input).await.unwrap(); + + let joined = msgs.iter().filter_map(|m| m["content"].as_str()).collect::>().join("\n"); + assert!(joined.contains("CONTEXT SUMMARY"), "summary block missing: {joined}"); + assert!(joined.contains("old answer"), "post-summary messages must survive"); + assert!(!joined.contains("old question"), "covered messages must be gone"); + let _ = m3; +} + +#[tokio::test] +async fn window_cuts_at_user_boundary_never_mid_tool_group() { + let store = Arc::new(InMemoryStore::new()); + let conv = ConversationId::new("a6"); + let (frame, input) = input(&store, &conv, ToolRendering::Inline).await; + + store.append(frame, NewMessage::user("first")).await.unwrap(); + let asst = store.append(frame, NewMessage::assistant("calling", None)).await.unwrap(); + let call = store.append_call(asst, NewCall::new("t", json!({})).with_provider_id("c1")).await.unwrap(); + store.resolve_call(call, &CallOutcome::Completed(ToolOutput::Text("r".into()))).await.unwrap(); + store.append(frame, NewMessage::user("second")).await.unwrap(); + + // Window of 2 would cut right before the assistant+tool group; the + // boundary rule must move the cut to "second". + let assembler = LinearAssembler::new().with_max_messages(2); + let store_dyn: Arc = store; + let msgs = assembler.build(&store_dyn, &input).await.unwrap(); + + let roles: Vec<&str> = msgs.iter().filter_map(|m| m["role"].as_str()).collect(); + assert_eq!(roles, ["system", "user"], "cut must land on the user boundary: {roles:?}"); +} diff --git a/crates/agent-loop/tests/kernel.rs b/crates/agent-loop/tests/kernel.rs new file mode 100644 index 0000000..c638b9a --- /dev/null +++ b/crates/agent-loop/tests/kernel.rs @@ -0,0 +1,498 @@ +//! Kernel test suite (blueprint §13) — against `FakeModel` + `InMemoryStore`, +//! no DB, no Docker, no network. + +use std::sync::{Arc, Mutex}; +use std::time::Duration; + +use agent_loop::gate::DenyList; +use agent_loop::ids::ConversationId; +use agent_loop::kernel::TurnOutcome; +use agent_loop::manager::{LoopManager, TurnMeta, TurnParams}; +use agent_loop::model::{ModelHint, StaticModels, StreamDelta}; +use agent_loop::prelude::async_trait; +use agent_loop::store::{CallState, FrameSpec, HistoryStore, NewMessage}; +use agent_loop::store_memory::InMemoryStore; +use agent_loop::testing::{self, FakeModel, Step}; +use agent_loop::tool::{Tool, ToolCtx, ToolFailure, ToolOutput, ToolRegistry}; +use agent_loop::context::StaticSystemContext; +use agent_loop::events::LoopEvent; +use serde_json::{Value, json}; +use tokio_util::sync::CancellationToken; + +// ── test tools ── + +struct WeatherTool; + +#[async_trait] +impl Tool for WeatherTool { + fn name(&self) -> &str { "get_weather" } + fn definition(&self) -> Value { + json!({"type":"function","function":{"name":"get_weather","parameters":{"type":"object","properties":{"city":{"type":"string"}}}}}) + } + async fn call(&self, args: Value, _ctx: &ToolCtx) -> Result { + Ok(ToolOutput::Text(format!("Sunny in {}", args["city"].as_str().unwrap_or("?")))) + } +} + +struct SlowTool; + +#[async_trait] +impl Tool for SlowTool { + fn name(&self) -> &str { "slow" } + fn definition(&self) -> Value { + json!({"type":"function","function":{"name":"slow","parameters":{"type":"object"}}}) + } + async fn call(&self, _args: Value, _ctx: &ToolCtx) -> Result { + tokio::time::sleep(Duration::from_secs(60)).await; + Ok(ToolOutput::Text("done".into())) + } +} + +/// Concurrency-safe tool rendezvousing on a barrier: proves the fan-out runs +/// concurrently (a sequential path would deadlock → timeout). +struct BarrierTool { + name: &'static str, + barrier: Arc, + log: Arc>>, +} + +#[async_trait] +impl Tool for BarrierTool { + fn name(&self) -> &str { self.name } + fn definition(&self) -> Value { + json!({"type":"function","function":{"name":self.name,"parameters":{"type":"object"}}}) + } + fn concurrency_safe(&self, _args: &Value) -> bool { true } + async fn call(&self, _args: Value, _ctx: &ToolCtx) -> Result { + self.log.lock().unwrap().push(format!("start:{}", self.name)); + self.barrier.wait().await; + self.log.lock().unwrap().push(format!("end:{}", self.name)); + Ok(ToolOutput::Text(format!("{} done", self.name))) + } +} + +/// Records start/end order in a shared log (sequentiality proofs). +struct OrderedTool { + name: &'static str, + safe: bool, + log: Arc>>, +} + +#[async_trait] +impl Tool for OrderedTool { + fn name(&self) -> &str { self.name } + fn definition(&self) -> Value { + json!({"type":"function","function":{"name":self.name,"parameters":{"type":"object"}}}) + } + fn concurrency_safe(&self, _args: &Value) -> bool { self.safe } + async fn call(&self, _args: Value, _ctx: &ToolCtx) -> Result { + self.log.lock().unwrap().push(format!("start:{}", self.name)); + tokio::task::yield_now().await; + self.log.lock().unwrap().push(format!("end:{}", self.name)); + Ok(ToolOutput::Text("ok".into())) + } +} + +/// Marks itself AwaitingHuman then suspends (ask_user semantics). +struct SuspendTool { + store: Arc, +} + +#[async_trait] +impl Tool for SuspendTool { + fn name(&self) -> &str { "suspend_me" } + fn definition(&self) -> Value { + json!({"type":"function","function":{"name":"suspend_me","parameters":{"type":"object"}}}) + } + async fn call(&self, _args: Value, ctx: &ToolCtx) -> Result { + self.store + .set_call_state(ctx.call_id, CallState::AwaitingHuman) + .await + .map_err(|e| ToolFailure::Failed(e.to_string()))?; + Err(ToolFailure::Suspend) + } +} + +// ── harness ── + +struct Harness { + manager: LoopManager, + store: Arc, +} + +fn harness_with(model: testing::FakeModel) -> Harness { + let store = Arc::new(InMemoryStore::new()); + let manager = LoopManager::builder() + .models(Arc::new(agent_loop::model::SingleModel::new(model))) + .store(store.clone()) + .build() + .unwrap(); + Harness { manager, store } +} + +async fn params( + manager: &LoopManager, + conv: &ConversationId, + tools: Arc, +) -> TurnParams { + let frame = manager.open_root(conv, FrameSpec::root("assistant")).await.unwrap(); + TurnParams { + frame, + agent: "assistant".into(), + system: Arc::new(StaticSystemContext::new("You are a test agent.")), + tools, + model_hint: ModelHint::default(), + live_input: None, + extensions: Default::default(), + meta: TurnMeta::default(), + assembler: None, + } +} + +// ── tests ── + +#[tokio::test] +async fn multi_round_text_tool_text_final() { + let model = FakeModel::new("m", vec![ + Step::tool_calls("let me check", vec![testing::call("c1", "get_weather", json!({"city":"Rome"}))]), + Step::message("It is sunny in Rome."), + ]); + let h = harness_with(model); + let conv = ConversationId::new("t1"); + let tools = ToolRegistry::new().with(WeatherTool).into_toolset(); + let p = params(&h.manager, &conv, tools).await; + + let handle = h.manager.start_turn(conv, NewMessage::user("weather?"), p).await.unwrap(); + let outcome = handle.join().await.unwrap(); + + let TurnOutcome::Final { content, .. } = outcome else { panic!("expected Final, got {outcome:?}") }; + assert_eq!(content, "It is sunny in Rome."); + + // The store recorded everything: user, assistant+tool_call, tool result, + // final assistant. + let frame = h.manager.store().active_frames(&ConversationId::new("t1")).await.unwrap()[0].id; + let history = h.store.load(frame).await.unwrap(); + assert_eq!(history.len(), 3); + assert_eq!(history[1].calls.len(), 1); + assert_eq!(history[1].calls[0].state, CallState::Done); + assert_eq!(history[1].calls[0].result.as_deref(), Some("Sunny in Rome")); +} + +#[tokio::test] +async fn exhausted_after_max_rounds() { + let model = FakeModel::new("m", vec![ + Step::tool_calls("", vec![testing::call("c1", "get_weather", json!({}))]), + Step::tool_calls("", vec![testing::call("c2", "get_weather", json!({}))]), + ]); + let store = Arc::new(InMemoryStore::new()); + let manager = LoopManager::builder() + .models(Arc::new(agent_loop::model::SingleModel::new(model))) + .store(store.clone()) + .max_rounds(2) + .build() + .unwrap(); + let conv = ConversationId::new("t2"); + let p = params(&manager, &conv, ToolRegistry::new().with(WeatherTool).into_toolset()).await; + + let handle = manager.start_turn(conv, NewMessage::user("loop forever"), p).await.unwrap(); + let outcome = handle.join().await.unwrap(); + assert!(matches!(outcome, TurnOutcome::Exhausted), "got {outcome:?}"); +} + +#[tokio::test] +async fn fallback_retriable_moves_to_second_model() { + let m1 = Arc::new(FakeModel::new("m1", vec![Step::error(Some(500), "boom")])); + let m2 = Arc::new(FakeModel::new("m2", vec![Step::message("recovered")])); + let store = Arc::new(InMemoryStore::new()); + let mut rx; + let manager = LoopManager::builder() + .models(Arc::new(StaticModels::new(vec![ + testing::handle(&m1, "m1"), + testing::handle(&m2, "m2"), + ]))) + .store(store.clone()) + .build() + .unwrap(); + rx = manager.events(); + let conv = ConversationId::new("t3"); + let p = params(&manager, &conv, ToolRegistry::new().into_toolset()).await; + + let handle = manager.start_turn(conv, NewMessage::user("hi"), p).await.unwrap(); + let outcome = handle.join().await.unwrap(); + assert!(matches!(outcome, TurnOutcome::Final { .. }), "got {outcome:?}"); + + assert_eq!(m1.requests().len(), 1); + assert_eq!(m2.requests().len(), 1); + + let mut saw_fallback = false; + while let Ok(ev) = rx.try_recv() { + if let LoopEvent::ModelFallback { from, to, .. } = ev.inner { + assert_eq!(from, "m1"); + assert_eq!(to, "m2"); + saw_fallback = true; + } + } + assert!(saw_fallback, "no ModelFallback event"); +} + +#[tokio::test] +async fn non_retriable_error_stops_without_fallback() { + let m1 = Arc::new(FakeModel::new("m1", vec![Step::error(Some(404), "no such model")])); + let m2 = Arc::new(FakeModel::new("m2", vec![Step::message("never reached")])); + let store = Arc::new(InMemoryStore::new()); + let manager = LoopManager::builder() + .models(Arc::new(StaticModels::new(vec![ + testing::handle(&m1, "m1"), + testing::handle(&m2, "m2"), + ]))) + .store(store.clone()) + .build() + .unwrap(); + let conv = ConversationId::new("t4"); + let p = params(&manager, &conv, ToolRegistry::new().into_toolset()).await; + + let handle = manager.start_turn(conv, NewMessage::user("hi"), p).await.unwrap(); + assert!(handle.join().await.is_err(), "404 must fail the turn"); + assert_eq!(m2.requests().len(), 0, "404 must not fall back"); +} + +#[tokio::test] +async fn cancel_during_llm_call() { + let model = FakeModel::new("m", vec![Step::pending()]); + let h = harness_with(model); + let conv = ConversationId::new("t5"); + let p = params(&h.manager, &conv, ToolRegistry::new().into_toolset()).await; + + let handle = h.manager.start_turn(conv.clone(), NewMessage::user("hi"), p).await.unwrap(); + let cancel: CancellationToken = handle.cancel.clone(); + tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(100)).await; + cancel.cancel(); + }); + let outcome = tokio::time::timeout(Duration::from_secs(5), handle.join()) + .await + .expect("join hung") + .unwrap(); + assert!(matches!(outcome, TurnOutcome::Cancelled), "got {outcome:?}"); + assert!(!h.manager.is_running(&conv)); +} + +#[tokio::test] +async fn cancel_during_slow_tool_marks_call_cancelled() { + let model = FakeModel::new("m", vec![ + Step::tool_calls("", vec![testing::call("c1", "slow", json!({}))]), + ]); + let h = harness_with(model); + let conv = ConversationId::new("t6"); + let p = params(&h.manager, &conv, ToolRegistry::new().with(SlowTool).into_toolset()).await; + let frame = p.frame; + + let handle = h.manager.start_turn(conv, NewMessage::user("run slow"), p).await.unwrap(); + let cancel = handle.cancel.clone(); + tokio::spawn(async move { + tokio::time::sleep(Duration::from_millis(150)).await; + cancel.cancel(); + }); + let outcome = tokio::time::timeout(Duration::from_secs(5), handle.join()) + .await + .expect("join hung") + .unwrap(); + assert!(matches!(outcome, TurnOutcome::Cancelled), "got {outcome:?}"); + + let calls = h.store.calls_in_state(frame, &[CallState::Cancelled]).await.unwrap(); + assert_eq!(calls.len(), 1, "the slow call must be recorded Cancelled, got {calls:?}"); +} + +#[tokio::test] +async fn fan_out_runs_concurrently_and_records_in_order() { + let barrier = Arc::new(tokio::sync::Barrier::new(3)); + let log = Arc::new(Mutex::new(Vec::new())); + let model = FakeModel::new("m", vec![ + Step::tool_calls("", vec![ + testing::call("c1", "p1", json!({})), + testing::call("c2", "p2", json!({})), + testing::call("c3", "p3", json!({})), + ]), + Step::message("all done"), + ]); + let h = harness_with(model); + let conv = ConversationId::new("t7"); + let p = params(&h.manager, &conv, ToolRegistry::new() + .with_arc(Arc::new(BarrierTool { name: "p1", barrier: barrier.clone(), log: log.clone() })) + .with_arc(Arc::new(BarrierTool { name: "p2", barrier: barrier.clone(), log: log.clone() })) + .with_arc(Arc::new(BarrierTool { name: "p3", barrier: barrier.clone(), log: log.clone() })) + .into_toolset()).await; + let frame = p.frame; + + let handle = h.manager.start_turn(conv, NewMessage::user("go"), p).await.unwrap(); + let outcome = tokio::time::timeout(Duration::from_secs(5), handle.join()) + .await + .expect("fan-out deadlocked (ran sequentially?)") + .unwrap(); + assert!(matches!(outcome, TurnOutcome::Final { .. }), "got {outcome:?}"); + + // All three started before any ended (true concurrency). + { + let log = log.lock().unwrap(); + let first_end = log.iter().position(|e| e.starts_with("end:")).unwrap(); + assert_eq!(log[..first_end].iter().filter(|e| e.starts_with("start:")).count(), 3, + "not all tools started before the first end: {log:?}"); + } + + // Ids are increasing in call order and all resolved Done. + let calls = h.store.calls_in_state(frame, &[CallState::Done]).await.unwrap(); + assert_eq!(calls.len(), 3); + let mut ids: Vec = calls.iter().map(|c| c.id.get()).collect(); + let sorted = ids.clone(); + ids.sort_unstable(); + // calls_in_state returns in message order; ids must already be ascending. + assert_eq!(ids, sorted); +} + +#[tokio::test] +async fn mixed_batch_stays_sequential() { + let log = Arc::new(Mutex::new(Vec::new())); + let model = FakeModel::new("m", vec![ + Step::tool_calls("", vec![ + testing::call("c1", "safe", json!({})), + testing::call("c2", "unsafe", json!({})), + ]), + Step::message("done"), + ]); + let h = harness_with(model); + let conv = ConversationId::new("t8"); + let p = params(&h.manager, &conv, ToolRegistry::new() + .with_arc(Arc::new(OrderedTool { name: "safe", safe: true, log: log.clone() })) + .with_arc(Arc::new(OrderedTool { name: "unsafe", safe: false, log: log.clone() })) + .into_toolset()).await; + + let handle = h.manager.start_turn(conv, NewMessage::user("go"), p).await.unwrap(); + let outcome = handle.join().await.unwrap(); + assert!(matches!(outcome, TurnOutcome::Final { .. }), "got {outcome:?}"); + + assert_eq!( + *log.lock().unwrap(), + vec!["start:safe", "end:safe", "start:unsafe", "end:unsafe"], + "mixed batch must run sequentially in order" + ); +} + +#[tokio::test] +async fn suspend_leaves_call_awaiting_human_and_ends_turn() { + let store = Arc::new(InMemoryStore::new()); + let model = FakeModel::new("m", vec![ + Step::tool_calls("", vec![testing::call("c1", "suspend_me", json!({}))]), + ]); + let manager = LoopManager::builder() + .models(Arc::new(agent_loop::model::SingleModel::new(model))) + .store(store.clone()) + .build() + .unwrap(); + let conv = ConversationId::new("t9"); + let suspend = SuspendTool { store: store.clone() }; + let p = params(&manager, &conv, ToolRegistry::new().with(suspend).into_toolset()).await; + let frame = p.frame; + + let handle = manager.start_turn(conv, NewMessage::user("ask something"), p).await.unwrap(); + let outcome = handle.join().await.unwrap(); + assert!(matches!(outcome, TurnOutcome::Cancelled), "got {outcome:?}"); + + let pending = store.calls_in_state(frame, &[CallState::AwaitingHuman]).await.unwrap(); + assert_eq!(pending.len(), 1, "the call must STAY AwaitingHuman"); + assert!(pending[0].result.is_none(), "no result recorded for a suspended call"); +} + +#[tokio::test] +async fn gate_reject_marks_rejected_and_loop_continues() { + let model = FakeModel::new("m", vec![ + Step::tool_calls("", vec![testing::call("c1", "blocked_tool", json!({}))]), + Step::message("after rejection"), + ]); + let store = Arc::new(InMemoryStore::new()); + let manager = LoopManager::builder() + .models(Arc::new(agent_loop::model::SingleModel::new(model))) + .store(store.clone()) + .gate(DenyList::new(["blocked_*"])) + .build() + .unwrap(); + let conv = ConversationId::new("t10"); + let p = params(&manager, &conv, ToolRegistry::new().with(WeatherTool).into_toolset()).await; + let frame = p.frame; + + let handle = manager.start_turn(conv, NewMessage::user("try it"), p).await.unwrap(); + let outcome = handle.join().await.unwrap(); + let TurnOutcome::Final { content, .. } = outcome else { panic!("expected Final, got {outcome:?}") }; + assert_eq!(content, "after rejection"); + + let rejected = store.calls_in_state(frame, &[CallState::Rejected]).await.unwrap(); + assert_eq!(rejected.len(), 1); +} + +#[tokio::test] +async fn streaming_deltas_precede_outcome_events() { + let model = FakeModel::new("m", vec![ + Step::message("hello").with_deltas(vec![ + StreamDelta::Text("he".into()), + StreamDelta::Text("llo".into()), + ]), + ]); + let h = harness_with(model); + let mut rx = h.manager.events(); + let conv = ConversationId::new("t11"); + let p = params(&h.manager, &conv, ToolRegistry::new().into_toolset()).await; + + let handle = h.manager.start_turn(conv, NewMessage::user("hi"), p).await.unwrap(); + let _ = handle.join().await.unwrap(); + + let mut events = Vec::new(); + while let Ok(ev) = rx.try_recv() { + events.push(ev.inner); + } + let done_idx = events.iter().position(|e| matches!(e, LoopEvent::Done { .. })).unwrap(); + let delta_count = events[..done_idx] + .iter() + .filter(|e| matches!(e, LoopEvent::TokenDelta { .. })) + .count(); + assert_eq!(delta_count, 2, "both deltas must precede Done: {events:?}"); +} + +#[tokio::test] +async fn orphan_user_message_marked_failed_on_new_turn() { + let model = FakeModel::new("m", vec![Step::message("reply")]); + let h = harness_with(model); + let conv = ConversationId::new("t12"); + let p = params(&h.manager, &conv, ToolRegistry::new().into_toolset()).await; + let frame = p.frame; + + // A previous user message with no assistant reply (crash mid-turn). + h.store.append(frame, NewMessage::user("orphan")).await.unwrap(); + + let handle = h.manager.start_turn(conv, NewMessage::user("fresh"), p).await.unwrap(); + let _ = handle.join().await.unwrap(); + + let history = h.store.load(frame).await.unwrap(); + assert!( + !history.iter().any(|m| m.content == "orphan"), + "the orphan must be excluded from the projection: {history:?}" + ); +} + +#[tokio::test] +async fn second_loop_on_same_conversation_rejected() { + let model = FakeModel::new("m", vec![Step::pending()]); + let h = harness_with(model); + let conv = ConversationId::new("t13"); + let p1 = params(&h.manager, &conv, ToolRegistry::new().into_toolset()).await; + + let handle = h.manager.start_turn(conv.clone(), NewMessage::user("first"), p1).await.unwrap(); + + let p2 = params(&h.manager, &conv, ToolRegistry::new().into_toolset()).await; + let second = h.manager.start_turn(conv.clone(), NewMessage::user("second"), p2).await; + assert!( + matches!(second, Err(agent_loop::manager::StartError::AlreadyRunning)), + "double-driving must be rejected" + ); + + handle.cancel.cancel(); + let _ = handle.join().await; +}