diff --git a/Cargo.lock b/Cargo.lock index 8356839..b758415 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -4242,6 +4242,7 @@ dependencies = [ "rand 0.10.1", "regex", "reqwest 0.13.4", + "rustls", "serde", "serde_json", "serde_yaml", diff --git a/crates/agent-loop/src/events.rs b/crates/agent-loop/src/events.rs index aafe754..3e7236a 100644 --- a/crates/agent-loop/src/events.rs +++ b/crates/agent-loop/src/events.rs @@ -123,7 +123,10 @@ pub struct EventSink { } impl EventSink { - pub(crate) fn new(conversation: ConversationId, tx: broadcast::Sender>) -> Self { + /// Wrap a bus sender for one conversation. Public so hosts can build + /// sinks in their own tests and adapters; the kernel builds them via the + /// manager. + pub fn new(conversation: ConversationId, tx: broadcast::Sender>) -> Self { Self { conversation, tx } } diff --git a/crates/agent-loop/src/gate.rs b/crates/agent-loop/src/gate.rs index 650efcb..4b8a378 100644 --- a/crates/agent-loop/src/gate.rs +++ b/crates/agent-loop/src/gate.rs @@ -27,6 +27,10 @@ pub struct PendingCall { pub enum GateDecision { Allow, Reject { reason: String }, + /// The gate was waiting for a human and the channel closed: the turn ends + /// and the call STAYS `AwaitingHuman` (the gate marked it before + /// suspending) — the same semantics as `ToolFailure::Suspend`. + Suspend, } #[async_trait] diff --git a/crates/agent-loop/src/kernel.rs b/crates/agent-loop/src/kernel.rs index 87c6a3c..4d7e74f 100644 --- a/crates/agent-loop/src/kernel.rs +++ b/crates/agent-loop/src/kernel.rs @@ -348,6 +348,7 @@ async fn run_sequential( continue; } PreExecution::TurnCancelled => return Ok(Some(TurnOutcome::Cancelled)), + PreExecution::Suspended => return Ok(Some(TurnOutcome::Cancelled)), }; let ctx = ToolCtx { @@ -468,6 +469,7 @@ async fn phase2_one<'a>( } Ok(PreExecution::Resolved(outcome)) => Phase2::Done(outcome), Ok(PreExecution::TurnCancelled) => Phase2::Done(CallOutcome::Cancelled), + Ok(PreExecution::Suspended) => Phase2::Suspended, Err(e) => Phase2::Done(CallOutcome::Failed(format!("pre-execution error: {e}"))), }; (idx, phase) @@ -507,6 +509,9 @@ enum PreExecution { Run(Arc), Resolved(CallOutcome), TurnCancelled, + /// The gate suspended awaiting a human: the call STAYS `AwaitingHuman` + /// (never resolved) and the turn ends. + Suspended, } /// Gate + hooks.pre + tool lookup — shared by sequential and fan-out paths. @@ -530,8 +535,12 @@ async fn pre_execution( _ = 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 })); + match decision { + GateDecision::Reject { reason } => { + return Ok(PreExecution::Resolved(CallOutcome::Rejected { reason })); + } + GateDecision::Suspend => return Ok(PreExecution::Suspended), + GateDecision::Allow => {} } let mut ptc_mut = ptc.clone(); diff --git a/crates/agent-loop/src/lib.rs b/crates/agent-loop/src/lib.rs index 62fdea5..1c9ca70 100644 --- a/crates/agent-loop/src/lib.rs +++ b/crates/agent-loop/src/lib.rs @@ -27,6 +27,10 @@ pub mod store_memory; pub mod testing; pub mod tool; +/// Re-exported so implementors of the crate's async traits can write +/// `#[agent_loop::async_trait]` without a direct dependency. +pub use async_trait::async_trait; + /// Application name sent as the `X-Title` header by the shipped clients /// (OpenRouter rankings). Clients accept an override. pub const APP_NAME: &str = "Skald"; diff --git a/crates/skald-core/Cargo.toml b/crates/skald-core/Cargo.toml index bfe7d88..c930948 100644 --- a/crates/skald-core/Cargo.toml +++ b/crates/skald-core/Cargo.toml @@ -81,3 +81,8 @@ honcho-client = { path = "../honcho-client" } agent-loop = { path = "../agent-loop" } core-api = { path = "../core-api" } mcp-client = { path = "../mcp-client" } + +[dev-dependencies] +# Tests that build reqwest clients (rustls-no-provider) need a process-wide +# crypto provider, installed in main() in production. +rustls = { version = "0.23", default-features = false, features = ["ring", "std", "tls12", "logging"] } diff --git a/crates/skald-core/src/lib.rs b/crates/skald-core/src/lib.rs index 726bb0d..88275b9 100644 --- a/crates/skald-core/src/lib.rs +++ b/crates/skald-core/src/lib.rs @@ -29,6 +29,7 @@ pub mod inbox; pub mod latex; pub mod llm; pub mod location; +pub mod loop_adapters; pub mod memory; pub mod mcp; pub mod notification; diff --git a/crates/skald-core/src/loop_adapters/activation.rs b/crates/skald-core/src/loop_adapters/activation.rs new file mode 100644 index 0000000..d2dbb7b --- /dev/null +++ b/crates/skald-core/src/loop_adapters/activation.rs @@ -0,0 +1,315 @@ +//! DTL activation adapters (blueprint D15): the crate owns the wire protocol, +//! Skald owns the catalog (MCP servers + the reserved `config` group) and the +//! persistence (`activated_tools`, anchored at the triggering message). + +use std::collections::HashSet; +use std::sync::{Arc, RwLock}; + +use agent_loop::activation::{Activation, ActivationSource, ToolActivator}; +use agent_loop::ids::{FrameId, MessageId}; +use agent_loop::tool::{ToolCtx, ToolFailure}; +use serde_json::Value; +use sqlx::SqlitePool; + +use crate::db::{activated_tools, chat_llm_tools}; +use crate::mcp::McpProvider; +use crate::tools::tool_names::CONFIG_GROUP; + +// ── ActivationSource ───────────────────────────────────────────────────────── + +/// Reads the durable activations of one scope (root session or sub-agent +/// frame) and resolves them to OpenAI tool defs for the assembler's DTL +/// injection. Port of `MessageBuilder::resolve_activation_defs`. +pub struct SkaldActivationSource { + pool: Arc, + mcp: Arc, + config_defs: Arc>, + session_id: i64, + /// `None` = root (session scope); `Some(stack_id)` = sub-agent frame. + stack: Option, +} + +impl SkaldActivationSource { + pub fn new( + pool: Arc, + mcp: Arc, + config_defs: Arc>, + session_id: i64, + stack: Option, + ) -> Self { + Self { pool, mcp, config_defs, session_id, stack } + } +} + +#[agent_loop::async_trait] +impl ActivationSource for SkaldActivationSource { + async fn activations(&self, _frame: FrameId) -> agent_loop::Result> { + let rows = activated_tools::list_active_at(&self.pool, self.session_id, self.stack, i64::MAX).await?; + + // Group by anchor, dedup tool names per anchor (a server may reappear). + let mut out: Vec = Vec::new(); + for row in rows { + let defs: Vec = if row.kind == "builtin" && row.ref_ == CONFIG_GROUP { + self.config_defs.as_ref().clone() + } else { + self.mcp + .tools_for(std::slice::from_ref(&row.ref_)) + .iter() + .map(|t| t.to_openai_definition()) + .collect() + }; + let anchor = MessageId(row.message_id); + match out.iter_mut().find(|a| a.anchor == anchor) { + Some(existing) => { + for d in defs { + let name = d["function"]["name"].as_str().unwrap_or(""); + if !existing.defs.iter().any(|e| e["function"]["name"].as_str() == Some(name)) { + existing.defs.push(d); + } + } + } + None => out.push(Activation { anchor, defs }), + } + } + Ok(out) + } +} + +// ── ToolActivator ──────────────────────────────────────────────────────────── + +/// Backend of the crate's shipped `activate_tools` tool: validates the groups +/// against the catalog, updates the in-memory grant set **immediately** (next +/// round sees the tools), and persists the activation anchored at the +/// triggering assistant message (derived from the call's `chat_llm_tools` +/// row). Unifies what today lives split between `tools/activate_tools.rs` +/// (grants) and `llm_loop.rs` (persistence). +pub struct SkaldToolActivator { + pool: Arc, + mcp: Arc, + grants: Arc>>, + session_id: i64, + stack: Option, +} + +impl SkaldToolActivator { + pub fn new( + pool: Arc, + mcp: Arc, + grants: Arc>>, + session_id: i64, + stack: Option, + ) -> Self { + Self { pool, mcp, grants, session_id, stack } + } +} + +#[agent_loop::async_trait] +impl ToolActivator for SkaldToolActivator { + async fn activate(&self, groups: Vec, ctx: &ToolCtx) -> Result { + if groups.is_empty() { + return Err(ToolFailure::Failed("activate_tools: `groups` is empty".into())); + } + + let available: HashSet = self.mcp.tools().iter().map(|t| t.server_name.clone()).collect(); + + // Immediate in-memory effect (the defs re-read at the next round picks + // the new grants up for free). + { + let mut set = self.grants.write().map_err(|_| ToolFailure::Failed("activate_tools: lock poisoned".into()))?; + for g in &groups { + set.insert(g.clone()); + } + } + + // Durable effect, anchored at the triggering assistant message. The + // anchor is derived from the call row — the crate's ToolCtx carries + // the call id, the message id is one lookup away. + let call = chat_llm_tools::get(&self.pool, ctx.call_id.get()) + .await + .map_err(|e| ToolFailure::Failed(format!("activate_tools: anchor lookup failed: {e}")))? + .ok_or_else(|| ToolFailure::Failed("activate_tools: call row not found".into()))?; + for g in &groups { + let kind = if g == CONFIG_GROUP { "builtin" } else { "mcp" }; + activated_tools::grant(&self.pool, self.session_id, self.stack, call.message_id, kind, g) + .await + .map_err(|e| ToolFailure::Failed(format!("activate_tools: grant failed: {e}")))?; + } + + let activated: Vec = groups + .iter() + .map(|n| { + if n == CONFIG_GROUP || available.contains(n) { + format!("{n} ✓") + } else { + format!("{n} (registered but not yet running — tools will appear after reconnect)") + } + }) + .collect(); + let scope = match self.stack { + None => "session".to_string(), + Some(s) => format!("stack {s}"), + }; + Ok(format!( + "Tool groups activated for this {scope}: {}. \ + Their tools are available from the next tool-call round.", + activated.join(", ") + )) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + + use agent_loop::store::HistoryStore; + use agent_loop::tool::ToolOutput; + use mcp_client::McpTool; + + use crate::db::{chat_history, chat_sessions_stack}; + use crate::loop_adapters::history::SqliteHistory; + use crate::tools::ToolResult; + + struct FakeMcp { + tools: Vec, + } + + impl FakeMcp { + fn with_server(name: &str, tool_names: &[&str]) -> Self { + Self { + tools: tool_names + .iter() + .map(|t| McpTool { + server_name: name.to_string(), + name: t.to_string(), + description: String::new(), + input_schema: serde_json::json!({"type":"object"}), + title: None, + output_schema: None, + annotations: None, + task_support: None, + }) + .collect(), + } + } + } + + #[async_trait::async_trait] + impl McpProvider for FakeMcp { + fn tools(&self) -> Vec { self.tools.clone() } + fn tools_for(&self, names: &[String]) -> Vec { + self.tools.iter().filter(|t| names.contains(&t.server_name)).cloned().collect() + } + fn server_descriptions(&self) -> HashMap> { HashMap::new() } + fn server_infos(&self) -> Vec { Vec::new() } + fn tool_display_name(&self, _server: &str, _tool: &str) -> Option { None } + async fn call(&self, _server: &str, _tool: &str, _args: Value) -> anyhow::Result { + unimplemented!() + } + } + + + fn temp_db_path(tag: &str) -> String { + let mut p = std::env::temp_dir(); + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH).unwrap().as_nanos(); + p.push(format!("skald-test-{tag}-{}-{nanos}.db", std::process::id())); + p.to_string_lossy().into_owned() + } + + fn cleanup(path: &str) { + for suffix in ["", "-wal", "-shm"] { + let _ = std::fs::remove_file(format!("{path}{suffix}")); + } + } + + struct Fixture { + pool: Arc, + frame: FrameId, + msg: MessageId, + call: agent_loop::ids::ToolCallId, + path: String, + } + + async fn fixture(tag: &str) -> Fixture { + let path = temp_db_path(tag); + let pool = Arc::new(crate::db::init_system_pool(&path).await.unwrap()); + sqlx::query("INSERT INTO chat_sessions (id) VALUES (1)").execute(&*pool).await.unwrap(); + let frame_row = chat_sessions_stack::create(&pool, 1, "assistant", None, 0, None).await.unwrap(); + let msg = chat_history::append(&pool, frame_row.id, &chat_history::Role::Assistant, "activating", false, None) + .await + .unwrap(); + let call = chat_llm_tools::append(&pool, msg, "activate_tools", "{}").await.unwrap(); + Fixture { + pool, + frame: FrameId(frame_row.id), + msg: MessageId(msg), + call: agent_loop::ids::ToolCallId(call), + path, + } + } + + #[tokio::test] + async fn activate_grants_in_memory_and_persists_anchored() { + let f = fixture("act-grant").await; + let mcp: Arc = Arc::new(FakeMcp::with_server("gmail", &["send", "read"])); + let grants = Arc::new(RwLock::new(HashSet::new())); + let activator = SkaldToolActivator::new(f.pool.clone(), mcp, grants.clone(), 1, None); + + let ctx = ToolCtx { + conversation: agent_loop::ids::ConversationId::new("session:1"), + frame: f.frame, + agent: "assistant".into(), + call_id: f.call, + cancel: tokio_util::sync::CancellationToken::new(), + extensions: Default::default(), + }; + let text = activator.activate(vec!["gmail".into(), CONFIG_GROUP.into()], &ctx).await.unwrap(); + assert!(text.contains("gmail ✓")); + + // In-memory effect. + assert!(grants.read().unwrap().contains("gmail")); + assert!(grants.read().unwrap().contains(CONFIG_GROUP)); + + // Durable effect, anchored at the assistant message. + let refs = activated_tools::list_refs_session(&f.pool, 1).await.unwrap(); + assert_eq!(refs.len(), 2); + let acts = activated_tools::list_active_at(&f.pool, 1, None, i64::MAX).await.unwrap(); + assert!(acts.iter().all(|a| a.message_id == f.msg.get())); + + f.pool.close().await; + cleanup(&f.path); + } + + #[tokio::test] + async fn activation_source_resolves_defs_per_anchor() { + let f = fixture("act-src").await; + let mcp: Arc = Arc::new(FakeMcp::with_server("gmail", &["send", "read"])); + activated_tools::grant(&f.pool, 1, None, f.msg.get(), "mcp", "gmail").await.unwrap(); + activated_tools::grant(&f.pool, 1, None, f.msg.get(), "builtin", CONFIG_GROUP).await.unwrap(); + + let config_defs = Arc::new(vec![serde_json::json!({ + "type":"function","function":{"name":"cron_list","parameters":{"type":"object"}} + })]); + let src = SkaldActivationSource::new(f.pool.clone(), mcp, config_defs, 1, None); + let acts = src.activations(f.frame).await.unwrap(); + + assert_eq!(acts.len(), 1, "same anchor → one merged entry"); + let names: Vec<&str> = acts[0] + .defs + .iter() + .filter_map(|d| d["function"]["name"].as_str()) + .collect(); + assert!(names.contains(&"send") || names.iter().any(|n| n.contains("send")), "{names:?}"); + assert!(names.contains(&"cron_list"), "{names:?}"); + + // The SqliteHistory + LinearAssembler path agrees on the anchor type. + let store = SqliteHistory::new(f.pool.clone()); + let history = store.load(f.frame).await.unwrap(); + assert_eq!(history[0].id, f.msg); + let _ = ToolOutput::Text("unused".into()); + + f.pool.close().await; + cleanup(&f.path); + } +} diff --git a/crates/skald-core/src/loop_adapters/gate.rs b/crates/skald-core/src/loop_adapters/gate.rs new file mode 100644 index 0000000..11f2cf1 --- /dev/null +++ b/crates/skald-core/src/loop_adapters/gate.rs @@ -0,0 +1,383 @@ +//! `ApprovalGate` — Skald's approval flow behind the crate's `Gate` trait +//! (port of `handler/gate.rs::run_approval_gate`, blueprint §10): +//! +//! 1. `pre_approved` short-circuit (post-restart manual resolve); +//! 2. the approval engine decides (explicit Allow/Deny rules win); +//! 3. the RunContext fast-path relaxes `Require` to `Allow` for pre-authorized +//! fs paths (never overrides a Deny); +//! 4. `Require` → auto-deny, or mark `AwaitingHuman` + register + emit +//! `ApprovalRequired` + block on the human decision; a closed channel maps +//! to `GateDecision::Suspend` (the call stays `AwaitingHuman`, the turn +//! ends) — the old `GateOutcome::ChannelClosed`. + +use std::collections::HashSet; +use std::sync::atomic::{AtomicBool, Ordering}; +use std::sync::{Arc, Mutex, RwLock}; + +use agent_loop::events::{EventSink, LoopEvent}; +use agent_loop::gate::{Gate, GateDecision, PendingCall}; +use agent_loop::store::{CallState, HistoryStore}; + +use crate::approval::{ApprovalManager, GateResult}; +use crate::run_context::RunContext; +use crate::session::handler::ApprovalDecision; +use crate::tools::{ToolRegistry, is_file_read_tool, is_file_write_tool}; + +/// Everything the gate needs that the current loop keeps on the handler. +/// Shared by reference so phase-2 wiring shares the same cells. +pub struct ApprovalGate { + approval: Arc, + store: Arc, + tools: Arc, + session_id: i64, + source: String, + group_id: Option, + run_context: Arc>>, + pre_approved: Arc>>, + auto_deny: Arc, + context_label: Arc>>, +} + +impl ApprovalGate { + #[allow(clippy::too_many_arguments)] + pub fn new( + approval: Arc, + store: Arc, + tools: Arc, + session_id: i64, + source: impl Into, + group_id: Option, + run_context: Arc>>, + pre_approved: Arc>>, + auto_deny: Arc, + context_label: Arc>>, + ) -> Self { + Self { + approval, + store, + tools, + session_id, + source: source.into(), + group_id, + run_context, + pre_approved, + auto_deny, + context_label, + } + } +} + +#[agent_loop::async_trait] +impl Gate for ApprovalGate { + async fn check(&self, call: &PendingCall, events: &EventSink) -> GateDecision { + // Post-restart manual resolve: already approved via a resolve endpoint. + if self.pre_approved.lock().unwrap().remove(&call.id.get()) { + return GateDecision::Allow; + } + + let category = self.tools.category_of(&call.name); + + // The approval engine decides first: an explicit Deny/Allow rule wins. + let mut gate = self + .approval + .check( + self.session_id, + category, + &call.agent, + &self.source, + &call.name, + &call.args, + self.group_id.as_deref(), + ) + .await; + + // RunContext fast-path: relax `Require` for pre-authorized fs paths + // (never overrides a Deny). + if matches!(gate, GateResult::Require) { + let path = call.args["path"].as_str().unwrap_or(""); + let guard = self.run_context.read().map(|g| g.clone()).unwrap_or_default(); + let dflt = RunContext::default(); + let rc = guard.as_ref().unwrap_or(&dflt); + let pre_allowed = if is_file_read_tool(&call.name) { + rc.is_read_allowed(path) + } else if is_file_write_tool(&call.name) { + rc.is_write_allowed(path) + } else { + false + }; + if pre_allowed { + gate = GateResult::Allow; + } + } + + match gate { + GateResult::Allow => GateDecision::Allow, + GateResult::Deny => GateDecision::Reject { + reason: "Tool call denied by approval policy.".to_string(), + }, + GateResult::Require => { + if self.auto_deny.load(Ordering::Relaxed) { + return GateDecision::Reject { + reason: "Tool call auto-denied: this session does not support approval requests." + .to_string(), + }; + } + + // Durability FIRST: the call must survive a crash as pending. + if let Err(e) = self.store.set_call_state(call.id, CallState::AwaitingHuman).await { + return GateDecision::Reject { + reason: format!("approval: failed to mark call pending: {e}"), + }; + } + + let label = self.context_label.read().ok().and_then(|g| g.clone()); + let (request_id, approve_rx) = self + .approval + .register( + self.session_id, + call.id.get(), + &call.name, + call.args.clone(), + &call.agent, + &self.source, + label.as_deref(), + category, + ) + .await; + events.emit(call.frame, None, LoopEvent::ApprovalRequired { + id: call.id, + name: call.name.clone(), + args: call.args.clone(), + }); + let _ = request_id; + + match approve_rx.await { + Ok(ApprovalDecision::Approved) => GateDecision::Allow, + Ok(ApprovalDecision::Rejected { note }) => GateDecision::Reject { + reason: ApprovalDecision::rejection_message(¬e), + }, + Err(_) => GateDecision::Suspend, + } + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use agent_loop::events::EventSink; + use agent_loop::ids::{ConversationId, FrameId, ToolCallId}; + use agent_loop::tool::Extensions; + use serde_json::json; + use sqlx::SqlitePool; + + use crate::approval::{NewApprovalRule, RuleAction}; + use crate::db::{chat_history, chat_llm_tools, chat_sessions_stack}; + use crate::loop_adapters::history::SqliteHistory; + + fn temp_db_path(tag: &str) -> String { + let mut p = std::env::temp_dir(); + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH).unwrap().as_nanos(); + p.push(format!("skald-test-{tag}-{}-{nanos}.db", std::process::id())); + p.to_string_lossy().into_owned() + } + + fn cleanup(path: &str) { + for suffix in ["", "-wal", "-shm"] { + let _ = std::fs::remove_file(format!("{path}{suffix}")); + } + } + + struct Fixture { + gate: ApprovalGate, + events: EventSink, + pool: Arc, + call: PendingCall, + path: String, + approval: Arc, + } + + async fn fixture(tag: &str) -> Fixture { + let path = temp_db_path(tag); + let pool = Arc::new(crate::db::init_system_pool(&path).await.unwrap()); + // The `default` permission group is a FK target for approval_rules.group_id. + sqlx::query("INSERT INTO tool_permission_groups (id, name) VALUES ('default', 'Default')") + .execute(&*pool) + .await + .unwrap(); + sqlx::query("INSERT INTO chat_sessions (id) VALUES (1)").execute(&*pool).await.unwrap(); + let frame = chat_sessions_stack::create(&pool, 1, "assistant", None, 0, None).await.unwrap(); + let msg = chat_history::append(&pool, frame.id, &chat_history::Role::Assistant, "a", false, None) + .await + .unwrap(); + let call_id = chat_llm_tools::append(&pool, msg, "some_tool", "{}").await.unwrap(); + + let (tx, _) = tokio::sync::broadcast::channel(16); + let approval = Arc::new(ApprovalManager::new(pool.clone(), tx)); + let store: Arc = Arc::new(SqliteHistory::new(pool.clone())); + let tools = Arc::new(ToolRegistry::new()); + let gate = ApprovalGate::new( + approval.clone(), + store, + tools, + 1, + "web", + None, + Arc::new(RwLock::new(None)), + Arc::new(Mutex::new(HashSet::new())), + Arc::new(AtomicBool::new(false)), + Arc::new(RwLock::new(None)), + ); + let (bus, _) = tokio::sync::broadcast::channel(16); + let events = EventSink::new(ConversationId::new("session:1"), bus); + let call = PendingCall { + id: ToolCallId(call_id), + name: "some_tool".into(), + args: json!({}), + frame: FrameId(frame.id), + agent: "assistant".into(), + extensions: Extensions::new(), + }; + Fixture { gate, events, pool, call, path, approval } + } + + #[tokio::test] + async fn explicit_deny_rule_rejects() { + let f = fixture("gate-deny").await; + f.approval + .add_rule(NewApprovalRule { + agent_id: None, + source: None, + tool_pattern: "some_tool".into(), + path_pattern: None, + action: RuleAction::Deny, + note: None, + priority: Some(1), + group_id: None, + }) + .await + .unwrap(); + + let d = f.gate.check(&f.call, &f.events).await; + assert!(matches!(d, GateDecision::Reject { .. })); + + f.pool.close().await; + cleanup(&f.path); + } + + #[tokio::test] + async fn auto_deny_rejects_require() { + let path = temp_db_path("gate-autodeny"); + let pool = Arc::new(crate::db::init_system_pool(&path).await.unwrap()); + sqlx::query("INSERT INTO chat_sessions (id) VALUES (1)").execute(&*pool).await.unwrap(); + let frame = chat_sessions_stack::create(&pool, 1, "assistant", None, 0, None).await.unwrap(); + let msg = chat_history::append(&pool, frame.id, &chat_history::Role::Assistant, "a", false, None) + .await + .unwrap(); + let call_id = chat_llm_tools::append(&pool, msg, "some_tool", "{}").await.unwrap(); + + let (tx, _) = tokio::sync::broadcast::channel(16); + let approval = Arc::new(ApprovalManager::new(pool.clone(), tx)); + let gate = ApprovalGate::new( + approval, + Arc::new(SqliteHistory::new(pool.clone())), + Arc::new(ToolRegistry::new()), + 1, + "cron", // background source: auto-deny + None, + Arc::new(RwLock::new(None)), + Arc::new(Mutex::new(HashSet::new())), + Arc::new(AtomicBool::new(true)), + Arc::new(RwLock::new(None)), + ); + let (bus, _) = tokio::sync::broadcast::channel(16); + let events = EventSink::new(ConversationId::new("session:1"), bus); + let call = PendingCall { + id: ToolCallId(call_id), + name: "some_tool".into(), + args: json!({}), + frame: FrameId(frame.id), + agent: "assistant".into(), + extensions: Extensions::new(), + }; + + // No rules at all → the seeded-less default is Require; auto-deny rejects. + let d = gate.check(&call, &events).await; + assert!(matches!(d, GateDecision::Reject { .. })); + + pool.close().await; + cleanup(&path); + } + + #[tokio::test] + async fn human_approval_allows_and_marks_pending_first() { + let f = fixture("gate-human").await; + let approval = f.approval.clone(); + let gate = Arc::new(f.gate); + let events = f.events.clone(); + let call = f.call.clone(); + + let check = tokio::spawn(async move { gate.check(&call, &events).await }); + + // Wait for the request to register, then approve it. + let request_id = tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + let pending = approval.list_pending().await; + if let Some(p) = pending.first() { + break p.request_id; + } + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + } + }) + .await + .unwrap(); + + // The call is durably pending while the human decides. + let row = chat_llm_tools::get(&f.pool, f.call.id.get()).await.unwrap().unwrap(); + assert_eq!(row.status, "pending"); + + approval.resolve(request_id, ApprovalDecision::Approved).await; + let d = check.await.unwrap(); + assert!(matches!(d, GateDecision::Allow)); + + f.pool.close().await; + cleanup(&f.path); + } + + #[tokio::test] + async fn human_rejection_rejects_with_note() { + let f = fixture("gate-reject").await; + let approval = f.approval.clone(); + let gate = Arc::new(f.gate); + let events = f.events.clone(); + let call = f.call.clone(); + + let check = tokio::spawn(async move { gate.check(&call, &events).await }); + + let request_id = tokio::time::timeout(std::time::Duration::from_secs(5), async { + loop { + let pending = approval.list_pending().await; + if let Some(p) = pending.first() { + break p.request_id; + } + tokio::time::sleep(std::time::Duration::from_millis(20)).await; + } + }) + .await + .unwrap(); + + approval + .resolve(request_id, ApprovalDecision::Rejected { note: "too risky".into() }) + .await; + let d = check.await.unwrap(); + match d { + GateDecision::Reject { reason } => assert!(reason.contains("too risky")), + other => panic!("expected Reject, got {other:?}"), + } + + f.pool.close().await; + cleanup(&f.path); + } +} diff --git a/crates/skald-core/src/loop_adapters/history.rs b/crates/skald-core/src/loop_adapters/history.rs new file mode 100644 index 0000000..9ab9d22 --- /dev/null +++ b/crates/skald-core/src/loop_adapters/history.rs @@ -0,0 +1,517 @@ +//! `SqliteHistory` — `HistoryStore` over the EXISTING Skald tables (no +//! migration, blueprint §0/§10): +//! +//! | crate concept | Skald table | +//! |---|---| +//! | conversation `"session:{id}"` | `chat_sessions.id` (the id rides in the `ConversationId` string) | +//! | frame | `chat_sessions_stack` (`terminated_at IS NULL` = active) | +//! | message | `chat_history` (`status='failed'` = failed orphan) | +//! | tool call | `chat_llm_tools` (status strings map 1:1 on `CallState`) | +//! | summary | `chat_summaries` (`covers_up_to_message_id`) | +//! +//! The store is built on an **owner pool** (one per user, §11): all ids are +//! pool-local, so the adapter needs no user scoping. The wire tool-call id is +//! synthesized as `tc_{row_id}`, exactly like the current message builder. + +use std::sync::Arc; + +use agent_loop::model::Usage; +use agent_loop::store::{ + CallOutcome, CallState, FrameRecord, FrameSpec, HistoryStore, NewCall, NewMessage, NewSummary, + Role, StoredCall, StoredMessage, StoredSummary, +}; +use agent_loop::tool::ToolOutput; +use agent_loop::ids::{ConversationId, FrameId, MessageId, SummaryId, ToolCallId}; +use serde_json::Value; +use sqlx::SqlitePool; + +use crate::db::{chat_history, chat_llm_tools, chat_sessions_stack, chat_summaries}; + +/// `HistoryStore` on a Skald owner pool. +pub struct SqliteHistory { + pool: Arc, +} + +impl SqliteHistory { + pub fn new(pool: Arc) -> Self { Self { pool } } + + /// Parse `"session:{id}"` (the adapter's conversation encoding). + fn session_id(conv: &ConversationId) -> anyhow::Result { + conv.as_str() + .strip_prefix("session:") + .and_then(|s| s.parse::().ok()) + .ok_or_else(|| anyhow::anyhow!("SqliteHistory: conversation id must be \"session:\", got '{conv}'")) + } + + fn map_role(role: Role) -> anyhow::Result { + match role { + Role::User => Ok(chat_history::Role::User), + Role::Assistant => Ok(chat_history::Role::Assistant), + Role::Agent => Ok(chat_history::Role::Agent), + // chat_history has no system role: system context is BUILT, never + // stored. Failing loudly beats silently mis-filing a message. + Role::System => anyhow::bail!( + "SqliteHistory: Role::System is not persistable — system context is not stored" + ), + } + } + + fn unmap_role(role: &chat_history::Role) -> Role { + match role { + chat_history::Role::User => Role::User, + chat_history::Role::Assistant => Role::Assistant, + chat_history::Role::Agent => Role::Agent, + } + } + + fn map_state(state: CallState) -> &'static str { + match state { + CallState::Running => "running", + CallState::AwaitingHuman => "pending", + CallState::Done => "done", + CallState::Failed => "failed", + CallState::Cancelled => "cancelled", + CallState::Rejected => "rejected", + } + } + + fn unmap_state(status: &str) -> CallState { + match status { + "pending" => CallState::AwaitingHuman, + "done" => CallState::Done, + "failed" => CallState::Failed, + "cancelled" => CallState::Cancelled, + "rejected" => CallState::Rejected, + _ => CallState::Running, + } + } + + fn stored_call(c: chat_llm_tools::LlmToolCall) -> StoredCall { + let arguments: Value = c + .arguments + .as_deref() + .and_then(|s| serde_json::from_str(s).ok()) + .unwrap_or(Value::Object(Default::default())); + // preview/media ride in `extras` (host free-form), mirroring how the + // current loop reads them back for the history projection. + let extras = serde_json::json!({ + "preview_old": c.preview_old, + "preview_new": c.preview_new, + "media": c.media, + }); + StoredCall { + id: ToolCallId(c.id), + message_id: MessageId(c.message_id), + provider_id: format!("tc_{}", c.id), + name: c.name, + arguments, + state: Self::unmap_state(&c.status), + result: c.result, + result_kind: c.result_type, + extras, + } + } + + fn stored_message(m: chat_history::ChatMessage, calls: Vec) -> StoredMessage { + StoredMessage { + id: MessageId(m.id), + role: Self::unmap_role(&m.role), + content: m.content, + reasoning: m.reasoning_content, + synthetic: m.is_synthetic, + failed: m.status == "failed", + metadata: m.metadata.map(|meta| { + serde_json::to_value(meta).unwrap_or(Value::Null) + }), + usage: Usage { + input_tokens: m.input_tokens.map(|n| n as u32), + output_tokens: m.output_tokens.map(|n| n as u32), + cache_read: None, + cache_write: None, + cost_usd: m.cost, + truncated: false, + }, + calls, + } + } + + async fn with_calls(&self, msgs: Vec) -> anyhow::Result> { + let mut out = Vec::with_capacity(msgs.len()); + for m in msgs { + let calls = chat_llm_tools::for_message(&self.pool, m.id) + .await? + .into_iter() + .map(Self::stored_call) + .collect(); + out.push(Self::stored_message(m, calls)); + } + Ok(out) + } +} + +#[agent_loop::async_trait] +impl HistoryStore for SqliteHistory { + // ── frames ── + + async fn open_frame( + &self, + conv: &ConversationId, + parent: Option, + spec: FrameSpec, + ) -> agent_loop::Result { + let session_id = Self::session_id(conv)?; + // Root frame: reuse the session's existing root stack row when present + // (sessions are provisioned with one), create it otherwise. + if parent.is_none() + && let Some(root) = chat_sessions_stack::main_for_session(&self.pool, session_id).await? + { + return Ok(FrameId(root.id)); + } + let frame = chat_sessions_stack::create( + &self.pool, + session_id, + &spec.agent, + spec.prompt.as_deref(), + spec.depth as i64, + spec.parent_call.map(|c| c.get()), + ) + .await?; + Ok(FrameId(frame.id)) + } + + async fn close_frame(&self, frame: FrameId) -> agent_loop::Result<()> { + chat_sessions_stack::terminate(&self.pool, frame.get()).await?; + Ok(()) + } + + async fn active_frames(&self, conv: &ConversationId) -> agent_loop::Result> { + let session_id = Self::session_id(conv)?; + let rows = sqlx::query_as::<_, (i64, i64, String, Option, i64, Option)>( + "SELECT id, session_id, agent_id, agent_prompt, depth, parent_tool_call_id + FROM chat_sessions_stack + WHERE session_id = ? AND terminated_at IS NULL + ORDER BY depth ASC", + ) + .bind(session_id) + .fetch_all(&*self.pool) + .await?; + Ok(rows + .into_iter() + .map(|(id, sid, agent, prompt, depth, parent_call)| FrameRecord { + id: FrameId(id), + conversation: ConversationId::new(format!("session:{sid}")), + // The parent frame id is not stored directly (only the parent + // tool call); recovery walks the call when it needs the link. + parent: None, + spec: FrameSpec { + agent, + prompt, + depth: depth as u32, + parent_call: parent_call.map(ToolCallId), + meta: Value::Null, + }, + active: true, + }) + .collect()) + } + + async fn deepest_active(&self, conv: &ConversationId) -> agent_loop::Result> { + Ok(self + .active_frames(conv) + .await? + .into_iter() + .max_by_key(|f| f.spec.depth)) + } + + // ── messages ── + + async fn append(&self, frame: FrameId, msg: NewMessage) -> agent_loop::Result { + let role = Self::map_role(msg.role)?; + // chat_history.metadata is a typed MessageMetadata column; the crate's + // free-form Value only round-trips when it parses back as one. + let metadata = msg + .metadata + .as_ref() + .and_then(|v| serde_json::from_value::(v.clone()).ok()); + let id = chat_history::append_with_metadata( + &self.pool, + frame.get(), + &role, + &msg.content, + msg.synthetic, + msg.reasoning.as_deref(), + metadata.as_ref(), + ) + .await?; + Ok(MessageId(id)) + } + + async fn set_usage(&self, msg: MessageId, usage: &Usage) -> agent_loop::Result<()> { + if let (Some(i), Some(o)) = (usage.input_tokens, usage.output_tokens) { + chat_history::set_usage(&self.pool, msg.get(), i, o, 0, usage.cost_usd).await?; + } + Ok(()) + } + + async fn load(&self, frame: FrameId) -> agent_loop::Result> { + let msgs = chat_history::for_stack(&self.pool, frame.get()).await?; + self.with_calls(msgs).await + } + + async fn load_since(&self, frame: FrameId, after: MessageId) -> agent_loop::Result> { + let msgs = chat_history::for_stack_since(&self.pool, frame.get(), after.get()).await?; + self.with_calls(msgs).await + } + + async fn last(&self, frame: FrameId) -> agent_loop::Result> { + let Some(m) = chat_history::last_message_for_stack(&self.pool, frame.get()).await? else { + return Ok(None); + }; + Ok(self.with_calls(vec![m]).await?.into_iter().next()) + } + + async fn mark_failed(&self, msg: MessageId) -> agent_loop::Result<()> { + chat_history::mark_failed(&self.pool, msg.get()).await?; + Ok(()) + } + + // ── tool calls ── + + async fn append_call(&self, msg: MessageId, call: NewCall) -> agent_loop::Result { + let args = serde_json::to_string(&call.arguments)?; + let id = chat_llm_tools::append(&self.pool, msg.get(), &call.name, &args).await?; + Ok(ToolCallId(id)) + } + + async fn resolve_call(&self, id: ToolCallId, outcome: &CallOutcome) -> agent_loop::Result<()> { + let pool = &self.pool; + match outcome { + CallOutcome::Completed(out) => { + chat_llm_tools::complete(pool, id.get(), &out.to_wire(), out.kind()).await?; + if let ToolOutput::Media { refs, .. } = out { + let media_json = serde_json::to_string(refs)?; + chat_llm_tools::set_media(pool, id.get(), &media_json).await?; + } + } + CallOutcome::Failed(e) => { + chat_llm_tools::fail(pool, id.get(), e).await?; + } + CallOutcome::Cancelled => { + chat_llm_tools::cancel(pool, id.get(), &outcome.result_text()).await?; + } + CallOutcome::Rejected { reason } => { + chat_llm_tools::reject(pool, id.get(), reason).await?; + } + } + Ok(()) + } + + async fn set_call_state(&self, id: ToolCallId, state: CallState) -> agent_loop::Result<()> { + anyhow::ensure!( + !state.is_terminal(), + "set_call_state is only for Running → AwaitingHuman, not terminal {state:?}" + ); + sqlx::query("UPDATE chat_llm_tools SET status = ? WHERE id = ?") + .bind(Self::map_state(state)) + .bind(id.get()) + .execute(&*self.pool) + .await?; + Ok(()) + } + + async fn calls_in_state(&self, frame: FrameId, states: &[CallState]) -> agent_loop::Result> { + // All calls of the frame, filtered in Rust: a frame's call set is + // bounded, and a static query keeps sqlx's dynamic-SQL audit happy. + let rows = sqlx::query_as::<_, (i64, i64, String, Option, Option, String, String)>( + "SELECT t.id, t.message_id, t.name, t.arguments, t.result, t.result_type, t.status + FROM chat_llm_tools t + JOIN chat_history h ON t.message_id = h.id + WHERE h.session_stack_id = ? + ORDER BY t.id ASC", + ) + .bind(frame.get()) + .fetch_all(&*self.pool) + .await?; + Ok(rows + .into_iter() + .map(|(id, message_id, name, arguments, result, result_type, status)| { + Self::stored_call(chat_llm_tools::LlmToolCall { + id, + message_id, + name, + arguments, + result, + result_type, + status, + preview_old: None, + preview_new: None, + media: None, + }) + }) + .filter(|c| states.contains(&c.state)) + .collect()) + } + + // ── summaries ── + + async fn save_summary(&self, frame: FrameId, s: NewSummary) -> agent_loop::Result { + let id = chat_summaries::save(&self.pool, frame.get(), &s.text, s.covered_up_to.get()).await?; + Ok(SummaryId(id)) + } + + async fn latest_summary(&self, frame: FrameId) -> agent_loop::Result> { + let Some(s) = chat_summaries::latest_for_stack(&self.pool, frame.get()).await? else { + return Ok(None); + }; + Ok(Some(StoredSummary { + id: SummaryId(s.id), + text: s.content, + covered_up_to: MessageId(s.covers_up_to_message_id), + })) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn temp_db_path(tag: &str) -> String { + let mut p = std::env::temp_dir(); + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH).unwrap().as_nanos(); + p.push(format!("skald-test-{tag}-{}-{nanos}.db", std::process::id())); + p.to_string_lossy().into_owned() + } + + fn cleanup(path: &str) { + for suffix in ["", "-wal", "-shm"] { + let _ = std::fs::remove_file(format!("{path}{suffix}")); + } + } + + async fn setup(tag: &str) -> (Arc, SqliteHistory, ConversationId, String) { + let path = temp_db_path(tag); + let pool = Arc::new(crate::db::init_system_pool(&path).await.unwrap()); + sqlx::query("INSERT INTO chat_sessions (id) VALUES (1)") + .execute(&*pool) + .await + .unwrap(); + // The session's root frame (created at provisioning time in production). + chat_sessions_stack::create(&pool, 1, "assistant", None, 0, None).await.unwrap(); + let store = SqliteHistory::new(pool.clone()); + (pool, store, ConversationId::new("session:1"), path) + } + + #[tokio::test] + async fn frames_open_reuse_root_and_close() { + let (pool, store, conv, path) = setup("hist-frames").await; + + // Root: reuses the provisioned root frame. + let root = store.open_frame(&conv, None, FrameSpec::root("assistant")).await.unwrap(); + // Child: creates a new frame at depth 1. + let child = store + .open_frame(&conv, Some(root), FrameSpec { + agent: "task".into(), + prompt: Some("do a thing".into()), + depth: 1, + parent_call: None, + meta: Value::Null, + }) + .await + .unwrap(); + assert_ne!(root, child); + + let active = store.active_frames(&conv).await.unwrap(); + assert_eq!(active.len(), 2); + assert_eq!(store.deepest_active(&conv).await.unwrap().unwrap().id, child); + + store.close_frame(child).await.unwrap(); + assert!(store.deepest_active(&conv).await.unwrap().unwrap().spec.depth == 0); + + pool.close().await; + cleanup(&path); + } + + #[tokio::test] + async fn messages_calls_and_states_round_trip() { + let (pool, store, conv, path) = setup("hist-msgs").await; + let frame = store.open_frame(&conv, None, FrameSpec::root("assistant")).await.unwrap(); + + store.append(frame, NewMessage::user("hi")).await.unwrap(); + let asst = store.append(frame, NewMessage::assistant("calling", Some("thinking…".into()))).await.unwrap(); + let call = store + .append_call(asst, NewCall::new("read_file", serde_json::json!({"path": "a.txt"}))) + .await + .unwrap(); + + // Running → AwaitingHuman (the only legal set_call_state). + store.set_call_state(call, CallState::AwaitingHuman).await.unwrap(); + assert!(store.set_call_state(call, CallState::Done).await.is_err()); + + store + .resolve_call(call, &CallOutcome::Completed(ToolOutput::Text("file contents".into()))) + .await + .unwrap(); + + let history = store.load(frame).await.unwrap(); + assert_eq!(history.len(), 2); + assert_eq!(history[1].reasoning.as_deref(), Some("thinking…")); + assert_eq!(history[1].calls.len(), 1); + let c = &history[1].calls[0]; + assert_eq!(c.state, CallState::Done); + assert_eq!(c.result.as_deref(), Some("file contents")); + assert_eq!(c.provider_id, format!("tc_{}", c.id.get())); + assert_eq!(c.arguments["path"], serde_json::json!("a.txt")); + + let done = store.calls_in_state(frame, &[CallState::Done]).await.unwrap(); + assert_eq!(done.len(), 1); + + // Orphan marking drops the message from the projection. + store.mark_failed(history[0].id).await.unwrap(); + assert_eq!(store.load(frame).await.unwrap().len(), 1); + + pool.close().await; + cleanup(&path); + } + + #[tokio::test] + async fn summaries_round_trip() { + let (pool, store, conv, path) = setup("hist-sum").await; + let frame = store.open_frame(&conv, None, FrameSpec::root("assistant")).await.unwrap(); + + let m1 = store.append(frame, NewMessage::user("old")).await.unwrap(); + store.append(frame, NewMessage::assistant("answer", None)).await.unwrap(); + let m3 = store.append(frame, NewMessage::user("new")).await.unwrap(); + + store + .save_summary(frame, NewSummary { text: "covered".into(), covered_up_to: m1 }) + .await + .unwrap(); + let latest = store.latest_summary(frame).await.unwrap().unwrap(); + assert_eq!(latest.text, "covered"); + assert_eq!(latest.covered_up_to, m1); + + let since = store.load_since(frame, latest.covered_up_to).await.unwrap(); + assert_eq!(since.len(), 2); + assert_eq!(since[1].id, m3); + + pool.close().await; + cleanup(&path); + } + + #[tokio::test] + async fn system_role_is_rejected() { + let (pool, store, conv, path) = setup("hist-sys").await; + let frame = store.open_frame(&conv, None, FrameSpec::root("assistant")).await.unwrap(); + let msg = NewMessage { + role: Role::System, + content: "nope".into(), + synthetic: true, + reasoning: None, + metadata: None, + }; + assert!(store.append(frame, msg).await.is_err()); + pool.close().await; + cleanup(&path); + } +} diff --git a/crates/skald-core/src/loop_adapters/mod.rs b/crates/skald-core/src/loop_adapters/mod.rs new file mode 100644 index 0000000..2d39df6 --- /dev/null +++ b/crates/skald-core/src/loop_adapters/mod.rs @@ -0,0 +1,22 @@ +//! Skald-side adapters implementing the `agent-loop` trait surface over the +//! existing infrastructure (blueprint §14 phase 1). **Unused by the current +//! loop** — they compile and are unit-tested here, and get wired in phase 2. +//! +//! - [`history::SqliteHistory`] — `HistoryStore` over the existing +//! `chat_sessions_stack` / `chat_history` / `chat_llm_tools` / `chat_summaries` +//! tables (no migration, §0). +//! - [`selector::SkaldSelector`] — `ModelSelector` over `LlmManager`, with the +//! agent's strength captured per-turn (D14). +//! - [`gate::ApprovalGate`] — `Gate` over `ApprovalManager` + the RunContext +//! fast-path + auto-deny + pre-approved (port of `handler/gate.rs`). +//! - [`toolset::SkaldToolSet`] — `ToolSet` over base/config defs + MCP grants + +//! memory/image/interface tools, with DTL rendering (port of +//! `AgentRunConfig::all_tool_defs`), plus the core-api→agent-loop tool bridge. +//! - [`activation`] — `ActivationSource` + `ToolActivator` over the +//! `activated_tools` table and the MCP provider (D15). + +pub mod activation; +pub mod gate; +pub mod history; +pub mod selector; +pub mod toolset; diff --git a/crates/skald-core/src/loop_adapters/selector.rs b/crates/skald-core/src/loop_adapters/selector.rs new file mode 100644 index 0000000..328d92f --- /dev/null +++ b/crates/skald-core/src/loop_adapters/selector.rs @@ -0,0 +1,171 @@ +//! `SkaldSelector` — `ModelSelector` over `LlmManager` (blueprint §10, D14). +//! +//! The agent's required **strength is captured at construction, per-turn** — +//! the crate never sees it: `hint` carries only an explicit pin, and the AUTO +//! path delegates to `LlmManager`'s strength tiering + priority ordering. + +use std::sync::Arc; + +use agent_loop::activation::ToolRendering; +use agent_loop::async_trait; +use agent_loop::model::{ModelHandle, ModelHint, ModelInfo, ModelSelector}; +use agent_loop::ids::ModelId; +use serde_json::Value; + +use crate::llm::{DtlMode, LlmEntry, LlmManager, LlmStrength}; + +/// Maps Skald's per-model DTL mode to the crate's wire protocol (D15). +pub fn tool_rendering_of(dtl: DtlMode) -> ToolRendering { + match dtl { + DtlMode::None => ToolRendering::Inline, + DtlMode::AnthropicToolReference => ToolRendering::DeferredToolReference, + DtlMode::KimiSystemTools => ToolRendering::SystemToolBlock, + } +} + +/// Builds the crate-side metadata for a resolved entry. `extras` stays empty: +/// the model's `extra_params` are already baked into the client at build time +/// (they would otherwise be merged into every request body a second time). +pub fn model_info_of(entry: &LlmEntry) -> ModelInfo { + ModelInfo { + prompt_cache: entry.prompt_cache, + capabilities: entry.capabilities.clone(), + tool_rendering: tool_rendering_of(entry.dtl), + extras: Value::Null, + } +} + +/// The selector handed to the loop manager for one turn: the manager's +/// strength tiering + health + priority, behind the crate's seam. +pub struct SkaldSelector { + manager: Arc, + strength: Option, +} + +impl SkaldSelector { + pub fn new(manager: Arc, strength: Option) -> Self { + Self { manager, strength } + } +} + +#[async_trait] +impl ModelSelector for SkaldSelector { + async fn select(&self, hint: &ModelHint, exclude: &[ModelId]) -> agent_loop::Result { + let (name, entry) = if exclude.is_empty() { + // First selection of the round: pin (hint.name) or AUTO by strength. + self.manager.resolve(hint.name.as_deref(), self.strength).await? + } else { + // Fallback: next healthy model in tier/priority order, skipping the + // ones already tried. The pin is intentionally dropped (it failed). + let excluded: Vec<&str> = exclude.iter().map(String::as_str).collect(); + self.manager.select_excluding(&excluded, self.strength).await? + }; + Ok(ModelHandle { + id: name, + model: entry.client.clone(), + info: model_info_of(&entry), + }) + } + + async fn report_success(&self, id: &ModelId) { + self.manager.mark_success(id).await; + } + + async fn report_failure(&self, id: &ModelId, err: &str) { + self.manager.mark_failure(id, err).await; + } +} + +#[cfg(test)] +mod tests { + use super::*; + use sqlx::SqlitePool; + + fn temp_db_path(tag: &str) -> String { + let mut p = std::env::temp_dir(); + let nanos = std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH).unwrap().as_nanos(); + p.push(format!("skald-test-{tag}-{}-{nanos}.db", std::process::id())); + p.to_string_lossy().into_owned() + } + + fn cleanup(path: &str) { + for suffix in ["", "-wal", "-shm"] { + let _ = std::fs::remove_file(format!("{path}{suffix}")); + } + } + + async fn manager_with_two_models(tag: &str) -> (Arc, Arc, String) { + // Building reqwest clients (rustls-no-provider) needs the process-wide + // crypto provider main() installs in production. Idempotent. + let _ = rustls::crypto::ring::default_provider().install_default(); + + let path = temp_db_path(tag); + let pool = Arc::new(crate::db::init_system_pool(&path).await.unwrap()); + sqlx::query("INSERT INTO llm_providers (id, name, type, api_key) VALUES (1, 'test', 'open_ai', 'sk-test')") + .execute(&*pool) + .await + .unwrap(); + // weak: low strength, better priority; strong: high strength. + sqlx::query("INSERT INTO llm_models (provider_id, model_id, name, strength, priority) VALUES + (1, 'weak-id', 'weak-model', 'low', 10), + (1, 'strong-id', 'strong-model', 'high', 20)") + .execute(&*pool) + .await + .unwrap(); + + let bus = Arc::new(core_api::system_bus::SystemEventBus::new()); + let mut registry = crate::provider::ProviderRegistry::new(bus); + registry.register_builtin(crate::llm::providers::openai::OpenAiProvider); + let manager = LlmManager::new(pool.clone(), Arc::new(registry), false).await.unwrap(); + (manager, pool, path) + } + + #[tokio::test] + async fn pin_resolves_exact_model() { + let (manager, pool, path) = manager_with_two_models("sel-pin").await; + let sel = SkaldSelector::new(manager, None); + + let h = sel.select(&ModelHint::name("weak-model"), &[]).await.unwrap(); + assert_eq!(h.id, "weak-model"); + + assert!(sel.select(&ModelHint::name("nope"), &[]).await.is_err()); + + pool.close().await; + cleanup(&path); + } + + #[tokio::test] + async fn auto_prefers_exact_strength_then_fallback_excludes() { + let (manager, pool, path) = manager_with_two_models("sel-auto").await; + let sel = SkaldSelector::new(manager, Some(LlmStrength::High)); + + // AUTO with strength High: the exact-tier model wins despite worse priority. + let h = sel.select(&ModelHint::default(), &[]).await.unwrap(); + assert_eq!(h.id, "strong-model"); + + // Fallback excludes it: the remaining one is served. + let h2 = sel.select(&ModelHint::default(), &["strong-model".to_string()]).await.unwrap(); + assert_eq!(h2.id, "weak-model"); + + pool.close().await; + cleanup(&path); + } + + #[tokio::test] + async fn health_reporting_degrades_and_recovers() { + let (manager, pool, path) = manager_with_two_models("sel-health").await; + let sel = SkaldSelector::new(manager, None); + + for _ in 0..5 { + sel.report_failure(&"weak-model".to_string(), "boom").await; + } + sel.report_success(&"weak-model".to_string()).await; + // Still resolvable after recovery. + let h = sel.select(&ModelHint::name("weak-model"), &[]).await.unwrap(); + assert_eq!(h.id, "weak-model"); + + pool.close().await; + cleanup(&path); + } +} diff --git a/crates/skald-core/src/loop_adapters/toolset.rs b/crates/skald-core/src/loop_adapters/toolset.rs new file mode 100644 index 0000000..4b72e06 --- /dev/null +++ b/crates/skald-core/src/loop_adapters/toolset.rs @@ -0,0 +1,420 @@ +//! `SkaldToolSet` — the crate's `ToolSet` over Skald's tool surface (port of +//! `AgentRunConfig::all_tool_defs`, blueprint §10), plus the bridges that let +//! core-api tools and MCP tools run inside the crate's kernel (the "double +//! Tool trait" seam of phase 1: bridged, not re-exported). + +use std::collections::HashSet; +use std::sync::{Arc, RwLock}; + +use agent_loop::activation::ToolRendering; +use agent_loop::async_trait; +use agent_loop::model::ModelInfo; +use agent_loop::tool::{ + MediaRef, RestartHint, Tool as LoopTool, ToolCtx, ToolExecution, ToolFailure, + ToolOutput, ToolSet, Visibility, +}; +use core_api::interface_tool::InterfaceTool; +use core_api::tool::{ExecutionOutcome as CoreOutcome, ToolExecutionState as CoreState}; +use core_api::user_fs::UserFs; +use serde_json::Value; +use sqlx::SqlitePool; + +use crate::mcp::McpProvider; +use crate::tools::tool_names::CONFIG_GROUP; + +// ── Extension keys ─────────────────────────────────────────────────────────── + +/// The calling user's id — tools that address per-user external stores key on +/// it. Inserted by the host at TurnParams construction. +#[derive(Debug, Clone)] +pub struct CallerUserId(pub String); + +/// Reads the `core_api::tool::ToolContext` pieces out of a `ToolCtx`: +/// owner pool + fs from the type-map, session id from the conversation. +fn core_tool_context(ctx: &ToolCtx) -> Result { + let pool = ctx.extensions.get::().ok_or_else(|| { + ToolFailure::Failed("tool bridge: no SqlitePool in extensions".into()) + })?; + let fs = ctx.extensions.get::().ok_or_else(|| { + ToolFailure::Failed("tool bridge: no UserFs in extensions".into()) + })?; + let user_id = ctx + .extensions + .get::() + .map(|u| u.0.clone()) + .unwrap_or_default(); + let session_id = ctx + .conversation + .as_str() + .strip_prefix("session:") + .and_then(|s| s.parse::().ok()) + .unwrap_or_default(); + Ok(core_api::tool::ToolContext { session_id, user_id, pool, fs }) +} + +/// Maps a core-api `ToolResult` to the crate's `ToolOutput`. +fn map_output(r: core_api::tool::ToolResult) -> ToolOutput { + match r { + core_api::tool::ToolResult::Text(s) => ToolOutput::Text(s), + core_api::tool::ToolResult::Json(v) => ToolOutput::Json(v), + core_api::tool::ToolResult::Media { text, media } => ToolOutput::Media { + text, + refs: media + .iter() + .map(|m| MediaRef { host_path: m.host_path.clone(), mime: m.mime.clone() }) + .collect(), + }, + } +} + +// ── BridgeExecution ────────────────────────────────────────────────────────── + +/// Wraps a core-api `ToolExecution` as the crate's `ToolExecution` (the two +/// state machines are structurally identical). +struct BridgeExecution<'a> { + inner: Box, +} + +impl ToolExecution for BridgeExecution<'_> { + fn state(&self) -> agent_loop::tool::ToolExecutionState { + match self.inner.state() { + CoreState::Pending | CoreState::AwaitingApproval | CoreState::Running => { + agent_loop::tool::ToolExecutionState::Running + } + CoreState::Completed => agent_loop::tool::ToolExecutionState::Completed, + CoreState::Failed => agent_loop::tool::ToolExecutionState::Failed, + CoreState::Cancelled | CoreState::Rejected => agent_loop::tool::ToolExecutionState::Cancelled, + } + } + + fn wait<'b>(&'b self) -> std::pin::Pin + Send + 'b>> { + Box::pin(async move { + match self.inner.wait().await { + CoreOutcome::Completed(r) => agent_loop::tool::ExecutionOutcome::Completed(map_output(r)), + CoreOutcome::Failed(e) => agent_loop::tool::ExecutionOutcome::Failed(e), + CoreOutcome::Cancelled => agent_loop::tool::ExecutionOutcome::Cancelled, + } + }) + } + + fn stop<'b>(&'b self) -> std::pin::Pin + Send + 'b>> { + self.inner.stop() + } +} + +// ── CoreToolBridge ─────────────────────────────────────────────────────────── + +/// Runs a core-api tool (`crate::tools::Tool`) inside the crate's kernel: +/// context from the type-map, execution bridged (kill/teardown preserved — +/// `execute_cmd`'s reaper keeps working through `stop`). +pub struct CoreToolBridge { + inner: Arc, +} + +impl CoreToolBridge { + pub fn new(inner: Arc) -> Self { Self { inner } } +} + +#[async_trait] +impl LoopTool for CoreToolBridge { + fn name(&self) -> &str { self.inner.name() } + + fn definition(&self) -> Value { self.inner.openai_definition() } + + async fn call(&self, args: Value, ctx: &ToolCtx) -> Result { + // Same path as `start`, driven to completion without a cancel token. + let exec = self.start(args, ctx); + match exec.wait().await { + agent_loop::tool::ExecutionOutcome::Completed(out) => Ok(out), + agent_loop::tool::ExecutionOutcome::Failed(e) => Err(ToolFailure::Failed(e)), + agent_loop::tool::ExecutionOutcome::Cancelled | + agent_loop::tool::ExecutionOutcome::Suspended => { + Err(ToolFailure::Failed("tool execution interrupted".into())) + } + } + } + + fn start<'a>(&'a self, args: Value, ctx: &'a ToolCtx) -> Box { + match core_tool_context(ctx) { + Ok(tool_ctx) => Box::new(BridgeExecution { inner: self.inner.run_with(&tool_ctx, args) }), + Err(e) => Box::new(agent_loop::tool::SimpleExecution::new(Box::pin(async move { Err(e) }))), + } + } + + fn restart_hint(&self) -> RestartHint { + // D7: shell commands are not idempotent — never re-run them on restart. + if self.inner.name() == "execute_cmd" { + RestartHint::MarkInterrupted + } else { + RestartHint::ReExecute + } + } + + fn visibility(&self) -> Visibility { + if self.inner.root_agent_only() { + Visibility::RootOnly + } else if self.inner.sub_agents_only() { + Visibility::SubAgentsOnly + } else if self.inner.interactive_only() { + Visibility::InteractiveOnly + } else { + Visibility::Always + } + } +} + +// ── McpToolBridge ──────────────────────────────────────────────────────────── + +/// Runs one MCP tool (`mcp__server__tool`) inside the crate's kernel. +pub struct McpToolBridge { + mcp: Arc, + server: String, + tool: String, + definition: Value, +} + +impl McpToolBridge { + pub fn new(mcp: Arc, server: impl Into, tool: impl Into, definition: Value) -> Self { + Self { mcp, server: server.into(), tool: tool.into(), definition } + } +} + +#[async_trait] +impl LoopTool for McpToolBridge { + fn name(&self) -> &str { self.definition["function"]["name"].as_str().unwrap_or("") } + + fn definition(&self) -> Value { self.definition.clone() } + + async fn call(&self, args: Value, _ctx: &ToolCtx) -> Result { + match self.mcp.call(&self.server, &self.tool, args).await { + Ok(r) => Ok(map_output(r)), + Err(e) => Err(ToolFailure::Failed(e.to_string())), + } + } +} + +// ── SkaldToolSet ───────────────────────────────────────────────────────────── + +/// The per-turn tool set: base built-ins + MCP grants + the lazy `config` +/// group + memory/image/interface tools, rendered per the model's +/// `ToolRendering` (D15). `defs` is re-read at every round/attempt — grants +/// activated at round N are visible at round N+1 for free. +pub struct SkaldToolSet { + base_defs: Vec, + config_defs: Vec, + mcp: Arc, + grants: Arc>>, + memory_tools: Vec>, + image_tools: Vec>, + /// Crate-native tools (ActivateToolsTool, aliases) — returned as-is. + interface_tools: Vec, + /// Core tools available for execution by name (the find() side). + core_tools: Vec>, + /// Extra crate-native tools for find() (bridge-free). + native_tools: Vec>, +} + +impl SkaldToolSet { + #[allow(clippy::too_many_arguments)] + pub fn new( + base_defs: Vec, + config_defs: Vec, + mcp: Arc, + grants: Arc>>, + memory_tools: Vec>, + image_tools: Vec>, + interface_tools: Vec, + core_tools: Vec>, + ) -> Self { + Self { + base_defs, + config_defs, + mcp, + grants, + memory_tools, + image_tools, + interface_tools, + core_tools, + native_tools: Vec::new(), + } + } + + pub fn with_native(mut self, tool: Arc) -> Self { + self.native_tools.push(tool); + self + } +} + +/// Tags an OpenAI tool definition as deferred (Anthropic tool search). +fn deferred(mut def: Value) -> Value { + def["defer_loading"] = Value::Bool(true); + def +} + +impl ToolSet for SkaldToolSet { + fn defs(&self, model: &ModelInfo) -> Vec { + let mut defs = self.base_defs.clone(); + + match model.tool_rendering { + // Declare EVERY accessible MCP tool + the config group as + // `defer_loading:true` — a stable, cache-safe set. + ToolRendering::DeferredToolReference => { + defs.extend(self.mcp.tools().iter().map(|t| deferred(t.to_openai_definition()))); + defs.extend(self.config_defs.iter().cloned().map(deferred)); + } + // Activated tools are injected as `system`+`tools` messages by the + // assembler — NOT in the top-level array. + ToolRendering::SystemToolBlock => {} + ToolRendering::Inline => { + let granted: HashSet = self.grants.read().map(|g| g.clone()).unwrap_or_default(); + let servers: Vec = granted + .iter() + .filter(|n| n.as_str() != CONFIG_GROUP) + .cloned() + .collect(); + if !servers.is_empty() { + defs.extend(self.mcp.tools_for(&servers).iter().map(|t| t.to_openai_definition())); + } + if granted.contains(CONFIG_GROUP) { + defs.extend(self.config_defs.iter().cloned()); + } + } + } + + defs.extend(self.memory_tools.iter().map(|t| t.openai_definition())); + defs.extend(self.image_tools.iter().map(|t| t.openai_definition())); + defs.extend(self.interface_tools.iter().map(|t| t.definition.clone())); + defs.extend(self.native_tools.iter().map(|t| t.definition())); + defs + } + + fn find(&self, name: &str) -> Option> { + if let Some(t) = self.native_tools.iter().find(|t| t.name() == name) { + return Some(t.clone()); + } + if let Some(t) = self.core_tools.iter().find(|t| t.name() == name) { + return Some(Arc::new(CoreToolBridge::new(t.clone()))); + } + if let Some(t) = self.memory_tools.iter().find(|t| t.name() == name) { + return Some(Arc::new(CoreToolBridge::new(t.clone()))); + } + if let Some(t) = self.image_tools.iter().find(|t| t.name() == name) { + return Some(Arc::new(CoreToolBridge::new(t.clone()))); + } + // MCP names are `mcp____`. + if let Some((server, tool)) = crate::mcp::parse_mcp_tool_name(name) { + let def = self + .mcp + .tools_for(&[server.to_string()]) + .into_iter() + .find(|t| t.name == tool) + .map(|t| t.to_openai_definition()); + if let Some(def) = def { + return Some(Arc::new(McpToolBridge::new(self.mcp.clone(), server, tool, def))); + } + } + None + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + + use mcp_client::McpTool; + + use crate::tools::ToolResult; + + fn fake_mcp(server: &str, tool_names: &[&str]) -> Arc { + struct Fake(Vec); + #[async_trait::async_trait] + impl McpProvider for Fake { + fn tools(&self) -> Vec { self.0.clone() } + fn tools_for(&self, names: &[String]) -> Vec { + self.0.iter().filter(|t| names.contains(&t.server_name)).cloned().collect() + } + fn server_descriptions(&self) -> HashMap> { HashMap::new() } + fn server_infos(&self) -> Vec { Vec::new() } + fn tool_display_name(&self, _s: &str, _t: &str) -> Option { None } + async fn call(&self, _s: &str, _t: &str, _a: Value) -> anyhow::Result { + unimplemented!() + } + } + Arc::new(Fake( + tool_names + .iter() + .map(|t| McpTool { + server_name: server.to_string(), + name: t.to_string(), + description: String::new(), + input_schema: serde_json::json!({"type":"object"}), + title: None, + output_schema: None, + annotations: None, + task_support: None, + }) + .collect(), + )) + } + + fn set(grants: &[&str]) -> Arc>> { + Arc::new(RwLock::new(grants.iter().map(|s| s.to_string()).collect())) + } + + fn toolset(grants: Arc>>) -> SkaldToolSet { + SkaldToolSet::new( + vec![serde_json::json!({"type":"function","function":{"name":"read_file","parameters":{}}})], + vec![serde_json::json!({"type":"function","function":{"name":"cron_list","parameters":{}}})], + fake_mcp("gmail", &["send"]), + grants, + vec![], + vec![], + vec![], + vec![], + ) + } + + #[test] + fn inline_renders_only_granted_groups() { + let ts = toolset(set(&[])); + let defs = ts.defs(&ModelInfo::default()); + let names: Vec<&str> = defs.iter().filter_map(|d| d["function"]["name"].as_str()).collect(); + assert_eq!(names, ["read_file"]); + + let ts = toolset(set(&["gmail", CONFIG_GROUP])); + let defs = ts.defs(&ModelInfo::default()); + let names: Vec<&str> = defs.iter().filter_map(|d| d["function"]["name"].as_str()).collect(); + assert!(names.contains(&"mcp__gmail__send"), "{names:?}"); + assert!(names.contains(&"cron_list")); + } + + #[test] + fn deferred_declares_everything_tagged() { + let ts = toolset(set(&[])); + let info = ModelInfo { tool_rendering: ToolRendering::DeferredToolReference, ..Default::default() }; + let defs = ts.defs(&info); + let gmail = defs.iter().find(|d| d["function"]["name"].as_str() == Some("mcp__gmail__send")).unwrap(); + assert_eq!(gmail["defer_loading"], serde_json::json!(true)); + let base = defs.iter().find(|d| d["function"]["name"].as_str() == Some("read_file")).unwrap(); + assert!(base.get("defer_loading").is_none()); + } + + #[test] + fn system_tool_block_keeps_array_stable() { + let ts = toolset(set(&["gmail"])); + let info = ModelInfo { tool_rendering: ToolRendering::SystemToolBlock, ..Default::default() }; + let defs = ts.defs(&info); + let names: Vec<&str> = defs.iter().filter_map(|d| d["function"]["name"].as_str()).collect(); + assert_eq!(names, ["read_file"], "activated tools must NOT be in the array in Kimi mode"); + } + + #[test] + fn find_bridges_mcp_names() { + let ts = toolset(set(&["gmail"])); + let t = ts.find("mcp__gmail__send").expect("mcp tool not bridged"); + assert_eq!(t.definition()["function"]["name"], serde_json::json!("mcp__gmail__send")); + assert!(ts.find("mcp__gmail__nope").is_none()); + assert!(ts.find("unknown_tool").is_none()); + } +}