harness-core now has everything the headless agent loop needs:
- Tool trait/ToolCtx/ToolRegistry + 30k-char head+tail output truncation
- PermissionService: async ask over a oneshot + AppEvent::PermissionAsked,
Once/Always/Reject replies, an auto-approve stub for tests/headless runs
- Config: JSONC loading, bundled/global/project-chain/env precedence,
{env:VAR} and {file:path} interpolation
- llm.rs: LlmEvent/LlmRequest/Provider trait, wire message/content types
- engine/: outer loop (run_session), inner stream processor (persists
parts/messages as events arrive, executes tool calls inline), retry
policy (retries only the pre-first-event window), doom-loop guard,
system prompt assembly
Verified end-to-end against a scripted MockProvider: text -> tool call
(read) -> final text, with messages/parts persisted in the right shape,
plus a provider-error-before-any-event case surfacing as Errored (no
partial message left behind). 49 tests passing, clippy clean.
530 lines
18 KiB
Rust
530 lines
18 KiB
Rust
use std::sync::Arc;
|
|
|
|
use crate::event::RunOutcome;
|
|
use crate::llm::{
|
|
FinishReason, Initiator, LlmRequest, Provider, ProviderError, Role as WireRole, ToolSchema,
|
|
WireContent, WireMessage,
|
|
};
|
|
use crate::store::Store;
|
|
use crate::types::{Message, Part, PartBody, Role, SessionId, ToolState};
|
|
|
|
use super::doomloop::DoomLoopGuard;
|
|
use super::processor::{process_step, StepContext, StepResult};
|
|
use super::retry;
|
|
use super::system;
|
|
|
|
pub struct RunConfig {
|
|
pub agent_name: String,
|
|
pub agent_prompt: String,
|
|
pub model: crate::types::ModelRef,
|
|
pub temperature: Option<f32>,
|
|
pub max_steps: u32,
|
|
pub instructions: Vec<String>,
|
|
}
|
|
|
|
/// opencode's exit condition: keep stepping while the last assistant turn asked for more
|
|
/// tool calls; stop once it produced a final answer (or a fresh user message is waiting).
|
|
async fn should_continue(store: &Store, session_id: &SessionId) -> Result<bool, ProviderError> {
|
|
let messages = store
|
|
.messages(session_id.clone())
|
|
.await
|
|
.map_err(|e| ProviderError::Decode(e.to_string()))?;
|
|
Ok(match messages.last() {
|
|
None => false,
|
|
Some(msg) if msg.role == Role::User => true,
|
|
Some(msg) => matches!(msg.finished, Some(FinishReason::ToolCalls)),
|
|
})
|
|
}
|
|
|
|
fn convert_message(message: &Message, parts: &[Part]) -> Vec<WireMessage> {
|
|
match message.role {
|
|
Role::User => {
|
|
let content: Vec<WireContent> = parts
|
|
.iter()
|
|
.filter_map(|p| match &p.body {
|
|
PartBody::Text { text, .. } => Some(WireContent::Text { text: text.clone() }),
|
|
_ => None,
|
|
})
|
|
.collect();
|
|
if content.is_empty() {
|
|
vec![]
|
|
} else {
|
|
vec![WireMessage {
|
|
role: WireRole::User,
|
|
content,
|
|
}]
|
|
}
|
|
}
|
|
Role::Assistant => {
|
|
let mut assistant_content = Vec::new();
|
|
let mut tool_results = Vec::new();
|
|
for part in parts {
|
|
match &part.body {
|
|
PartBody::Text { text, .. } => {
|
|
assistant_content.push(WireContent::Text { text: text.clone() })
|
|
}
|
|
PartBody::Tool {
|
|
call_id,
|
|
name,
|
|
state,
|
|
} => match state {
|
|
ToolState::Completed { input, output, .. } => {
|
|
assistant_content.push(WireContent::ToolCall {
|
|
call_id: call_id.clone(),
|
|
name: name.clone(),
|
|
input: input.clone(),
|
|
});
|
|
tool_results.push(WireContent::ToolResult {
|
|
call_id: call_id.clone(),
|
|
output: output.clone(),
|
|
is_error: false,
|
|
});
|
|
}
|
|
ToolState::Error { input, error } => {
|
|
assistant_content.push(WireContent::ToolCall {
|
|
call_id: call_id.clone(),
|
|
name: name.clone(),
|
|
input: input.clone(),
|
|
});
|
|
tool_results.push(WireContent::ToolResult {
|
|
call_id: call_id.clone(),
|
|
output: error.clone(),
|
|
is_error: true,
|
|
});
|
|
}
|
|
ToolState::Running { input, .. } => {
|
|
assistant_content.push(WireContent::ToolCall {
|
|
call_id: call_id.clone(),
|
|
name: name.clone(),
|
|
input: input.clone(),
|
|
});
|
|
}
|
|
ToolState::Pending { .. } => {}
|
|
},
|
|
_ => {}
|
|
}
|
|
}
|
|
let mut out = Vec::new();
|
|
if !assistant_content.is_empty() {
|
|
out.push(WireMessage {
|
|
role: WireRole::Assistant,
|
|
content: assistant_content,
|
|
});
|
|
}
|
|
if !tool_results.is_empty() {
|
|
out.push(WireMessage {
|
|
role: WireRole::Tool,
|
|
content: tool_results,
|
|
});
|
|
}
|
|
out
|
|
}
|
|
}
|
|
}
|
|
|
|
/// The outer loop: one call per session turn until the model stops asking for tool calls.
|
|
/// `now_fn` supplies `created_at`/`StepContext::now` stamps (kept out of the loop body so
|
|
/// tests can drive deterministic timestamps).
|
|
pub async fn run_session(
|
|
provider: Arc<dyn Provider>,
|
|
mut ctx: StepContext,
|
|
run_config: &RunConfig,
|
|
now_fn: impl Fn() -> i64,
|
|
) -> RunOutcome {
|
|
let mut doomloop = DoomLoopGuard::new();
|
|
let mut steps = 0u32;
|
|
|
|
loop {
|
|
match should_continue(&ctx.store, &ctx.session_id).await {
|
|
Ok(true) => {}
|
|
Ok(false) => return RunOutcome::Stopped,
|
|
Err(e) => {
|
|
return RunOutcome::Errored {
|
|
message: e.to_string(),
|
|
}
|
|
}
|
|
}
|
|
if steps >= run_config.max_steps {
|
|
return RunOutcome::Stopped;
|
|
}
|
|
steps += 1;
|
|
ctx.now = now_fn();
|
|
|
|
let messages = match ctx.store.messages(ctx.session_id.clone()).await {
|
|
Ok(m) => m,
|
|
Err(e) => {
|
|
return RunOutcome::Errored {
|
|
message: e.to_string(),
|
|
}
|
|
}
|
|
};
|
|
let mut wire_messages = Vec::new();
|
|
for message in &messages {
|
|
let parts = match ctx.store.parts(message.id.clone()).await {
|
|
Ok(p) => p,
|
|
Err(e) => {
|
|
return RunOutcome::Errored {
|
|
message: e.to_string(),
|
|
}
|
|
}
|
|
};
|
|
wire_messages.extend(convert_message(message, &parts));
|
|
}
|
|
|
|
let system_blocks = system::assemble(
|
|
system::env_header(&ctx.cwd),
|
|
&run_config.agent_prompt,
|
|
&run_config.instructions,
|
|
);
|
|
let tools: Vec<ToolSchema> = ctx
|
|
.tools
|
|
.all()
|
|
.iter()
|
|
.map(|t| ToolSchema {
|
|
name: t.name().to_string(),
|
|
description: t.description().to_string(),
|
|
parameters: t.parameters(),
|
|
})
|
|
.collect();
|
|
|
|
let req = LlmRequest {
|
|
model: run_config.model.model_id.clone(),
|
|
system: system_blocks,
|
|
messages: wire_messages,
|
|
tools,
|
|
temperature: run_config.temperature,
|
|
max_tokens: None,
|
|
reasoning: None,
|
|
initiator: Initiator::User,
|
|
};
|
|
|
|
let stream_result = retry::with_retry(&ctx.cancel, || {
|
|
let provider = provider.clone();
|
|
let req = req.clone();
|
|
let cancel = ctx.cancel.child_token();
|
|
async move { provider.stream(req, cancel).await }
|
|
})
|
|
.await;
|
|
|
|
let stream = match stream_result {
|
|
Ok(s) => s,
|
|
Err(ProviderError::Cancelled) => return RunOutcome::Aborted,
|
|
Err(e) => {
|
|
return RunOutcome::Errored {
|
|
message: e.to_string(),
|
|
}
|
|
}
|
|
};
|
|
|
|
let step = process_step(
|
|
stream,
|
|
&ctx,
|
|
run_config.model.clone(),
|
|
&run_config.agent_name,
|
|
&mut doomloop,
|
|
)
|
|
.await;
|
|
|
|
match step {
|
|
Ok(outcome) if outcome.aborted => return RunOutcome::Aborted,
|
|
Ok(outcome) => match outcome.result {
|
|
StepResult::Continue => continue,
|
|
StepResult::Stop => return RunOutcome::Stopped,
|
|
StepResult::Compact => return RunOutcome::Stopped, // stub until M6
|
|
},
|
|
Err(step_err) if matches!(step_err.source, ProviderError::Cancelled) => {
|
|
return RunOutcome::Aborted;
|
|
}
|
|
Err(step_err) => {
|
|
return RunOutcome::Errored {
|
|
message: step_err.source.to_string(),
|
|
};
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use std::collections::VecDeque;
|
|
use std::sync::Mutex as StdMutex;
|
|
|
|
use async_trait::async_trait;
|
|
use tokio_util::sync::CancellationToken;
|
|
|
|
use super::*;
|
|
use crate::event::EventBus;
|
|
use crate::llm::{LlmEvent, LlmEventStream};
|
|
use crate::permission::{spawn_auto_approve, PermissionService};
|
|
use crate::store::Store;
|
|
use crate::tool::{Tool, ToolCtx, ToolError, ToolOutput, ToolRegistry};
|
|
use crate::types::{ModelInfo, ModelRef, Session, TokenUsage};
|
|
|
|
struct MockProvider {
|
|
steps: StdMutex<VecDeque<Vec<Result<LlmEvent, ProviderError>>>>,
|
|
}
|
|
|
|
impl MockProvider {
|
|
fn scripted(steps: Vec<Vec<Result<LlmEvent, ProviderError>>>) -> Self {
|
|
Self {
|
|
steps: StdMutex::new(steps.into_iter().collect()),
|
|
}
|
|
}
|
|
}
|
|
|
|
#[async_trait]
|
|
impl Provider for MockProvider {
|
|
fn id(&self) -> &str {
|
|
"mock"
|
|
}
|
|
async fn list_models(&self) -> Result<Vec<ModelInfo>, ProviderError> {
|
|
Ok(vec![])
|
|
}
|
|
async fn stream(
|
|
&self,
|
|
_req: LlmRequest,
|
|
_cancel: CancellationToken,
|
|
) -> Result<LlmEventStream, ProviderError> {
|
|
let events = self
|
|
.steps
|
|
.lock()
|
|
.unwrap()
|
|
.pop_front()
|
|
.ok_or_else(|| ProviderError::Decode("no more scripted steps".into()))?;
|
|
Ok(Box::pin(futures::stream::iter(events)))
|
|
}
|
|
}
|
|
|
|
struct MockReadTool;
|
|
|
|
#[async_trait]
|
|
impl Tool for MockReadTool {
|
|
fn name(&self) -> &str {
|
|
"read"
|
|
}
|
|
fn description(&self) -> &str {
|
|
"reads a file"
|
|
}
|
|
fn parameters(&self) -> serde_json::Value {
|
|
serde_json::json!({"type": "object", "properties": {"file_path": {"type": "string"}}})
|
|
}
|
|
async fn execute(
|
|
&self,
|
|
_input: serde_json::Value,
|
|
_ctx: ToolCtx,
|
|
) -> Result<ToolOutput, ToolError> {
|
|
Ok(ToolOutput::new("read foo.txt", "mock file content"))
|
|
}
|
|
}
|
|
|
|
fn usage(input: u64, output: u64) -> TokenUsage {
|
|
TokenUsage {
|
|
input,
|
|
output,
|
|
..Default::default()
|
|
}
|
|
}
|
|
|
|
async fn make_ctx(
|
|
store: Store,
|
|
bus: EventBus,
|
|
session_id: SessionId,
|
|
cwd: std::path::PathBuf,
|
|
) -> StepContext {
|
|
let mut tools = ToolRegistry::new();
|
|
tools.register(Arc::new(MockReadTool));
|
|
let permissions = Arc::new(PermissionService::new(bus.clone()));
|
|
spawn_auto_approve(bus.clone(), permissions.clone());
|
|
StepContext {
|
|
store,
|
|
bus,
|
|
tools,
|
|
permissions,
|
|
static_rules: Vec::new(),
|
|
extra_rules: Arc::new(std::sync::Mutex::new(Vec::new())),
|
|
session_id,
|
|
cwd: cwd.clone(),
|
|
data_dir: cwd.join("tool-output"),
|
|
cancel: CancellationToken::new(),
|
|
now: 1,
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn text_then_tool_call_then_final_text_persists_correctly() {
|
|
let store = Store::open_in_memory().unwrap();
|
|
let bus = EventBus::new();
|
|
let model = ModelRef::new("mock", "mock-model");
|
|
let session = Session::new_root("orchestrator", model.clone(), 1);
|
|
let session_id = session.id.clone();
|
|
store.upsert_session(session).await.unwrap();
|
|
|
|
let user_message = Message::new_user(session_id.clone(), 1);
|
|
store.upsert_message(user_message.clone()).await.unwrap();
|
|
store
|
|
.upsert_part(Part {
|
|
id: crate::types::PartId::new(),
|
|
message_id: user_message.id.clone(),
|
|
session_id: session_id.clone(),
|
|
idx: 0,
|
|
body: PartBody::Text {
|
|
text: "please read foo.txt".into(),
|
|
synthetic: false,
|
|
},
|
|
})
|
|
.await
|
|
.unwrap();
|
|
|
|
let provider = MockProvider::scripted(vec![
|
|
vec![
|
|
Ok(LlmEvent::TextStart { id: "t1".into() }),
|
|
Ok(LlmEvent::TextDelta {
|
|
id: "t1".into(),
|
|
text: "Let me check the file.".into(),
|
|
}),
|
|
Ok(LlmEvent::TextEnd { id: "t1".into() }),
|
|
Ok(LlmEvent::ToolCall {
|
|
call_id: "call_1".into(),
|
|
name: "read".into(),
|
|
input: serde_json::json!({"file_path": "foo.txt"}),
|
|
}),
|
|
Ok(LlmEvent::Finish {
|
|
reason: FinishReason::ToolCalls,
|
|
usage: usage(10, 5),
|
|
}),
|
|
],
|
|
vec![
|
|
Ok(LlmEvent::TextStart { id: "t2".into() }),
|
|
Ok(LlmEvent::TextDelta {
|
|
id: "t2".into(),
|
|
text: "The file contains: mock file content".into(),
|
|
}),
|
|
Ok(LlmEvent::TextEnd { id: "t2".into() }),
|
|
Ok(LlmEvent::Finish {
|
|
reason: FinishReason::Stop,
|
|
usage: usage(20, 8),
|
|
}),
|
|
],
|
|
]);
|
|
|
|
let cwd = tempfile::tempdir().unwrap();
|
|
let ctx = make_ctx(
|
|
store.clone(),
|
|
bus,
|
|
session_id.clone(),
|
|
cwd.path().to_path_buf(),
|
|
)
|
|
.await;
|
|
let run_config = RunConfig {
|
|
agent_name: "orchestrator".into(),
|
|
agent_prompt: "You are a helpful assistant.".into(),
|
|
model,
|
|
temperature: None,
|
|
max_steps: 10,
|
|
instructions: Vec::new(),
|
|
};
|
|
|
|
let outcome = run_session(Arc::new(provider), ctx, &run_config, || 2).await;
|
|
assert!(matches!(outcome, RunOutcome::Stopped));
|
|
|
|
let messages = store.messages(session_id.clone()).await.unwrap();
|
|
assert_eq!(
|
|
messages.len(),
|
|
3,
|
|
"user + tool-call assistant turn + final assistant turn"
|
|
);
|
|
assert_eq!(messages[0].role, Role::User);
|
|
assert_eq!(messages[1].role, Role::Assistant);
|
|
assert_eq!(messages[1].finished, Some(FinishReason::ToolCalls));
|
|
assert_eq!(messages[2].role, Role::Assistant);
|
|
assert_eq!(messages[2].finished, Some(FinishReason::Stop));
|
|
|
|
let first_parts = store.parts(messages[1].id.clone()).await.unwrap();
|
|
let tool_part = first_parts
|
|
.iter()
|
|
.find(|p| matches!(p.body, PartBody::Tool { .. }))
|
|
.expect("tool part persisted");
|
|
match &tool_part.body {
|
|
PartBody::Tool { name, state, .. } => {
|
|
assert_eq!(name, "read");
|
|
match state {
|
|
ToolState::Completed { output, .. } => assert_eq!(output, "mock file content"),
|
|
other => panic!("expected Completed tool state, got {other:?}"),
|
|
}
|
|
}
|
|
_ => unreachable!(),
|
|
}
|
|
let first_text = first_parts
|
|
.iter()
|
|
.find_map(|p| match &p.body {
|
|
PartBody::Text { text, .. } => Some(text.clone()),
|
|
_ => None,
|
|
})
|
|
.expect("text part persisted");
|
|
assert_eq!(first_text, "Let me check the file.");
|
|
|
|
let final_parts = store.parts(messages[2].id.clone()).await.unwrap();
|
|
let final_text = final_parts
|
|
.iter()
|
|
.find_map(|p| match &p.body {
|
|
PartBody::Text { text, .. } => Some(text.clone()),
|
|
_ => None,
|
|
})
|
|
.expect("final text part persisted");
|
|
assert_eq!(final_text, "The file contains: mock file content");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn provider_error_on_first_event_is_reported_and_run_errors() {
|
|
let store = Store::open_in_memory().unwrap();
|
|
let bus = EventBus::new();
|
|
let model = ModelRef::new("mock", "mock-model");
|
|
let session = Session::new_root("orchestrator", model.clone(), 1);
|
|
let session_id = session.id.clone();
|
|
store.upsert_session(session).await.unwrap();
|
|
|
|
let user_message = Message::new_user(session_id.clone(), 1);
|
|
store.upsert_message(user_message.clone()).await.unwrap();
|
|
store
|
|
.upsert_part(Part {
|
|
id: crate::types::PartId::new(),
|
|
message_id: user_message.id.clone(),
|
|
session_id: session_id.clone(),
|
|
idx: 0,
|
|
body: PartBody::Text {
|
|
text: "hi".into(),
|
|
synthetic: false,
|
|
},
|
|
})
|
|
.await
|
|
.unwrap();
|
|
|
|
// ContextOverflow is never retried, so this should surface immediately as Errored.
|
|
let provider = MockProvider::scripted(vec![vec![Err(ProviderError::ContextOverflow)]]);
|
|
|
|
let cwd = tempfile::tempdir().unwrap();
|
|
let ctx = make_ctx(
|
|
store.clone(),
|
|
bus,
|
|
session_id.clone(),
|
|
cwd.path().to_path_buf(),
|
|
)
|
|
.await;
|
|
let run_config = RunConfig {
|
|
agent_name: "orchestrator".into(),
|
|
agent_prompt: "You are a helpful assistant.".into(),
|
|
model,
|
|
temperature: None,
|
|
max_steps: 10,
|
|
instructions: Vec::new(),
|
|
};
|
|
|
|
let outcome = run_session(Arc::new(provider), ctx, &run_config, || 2).await;
|
|
assert!(matches!(outcome, RunOutcome::Errored { .. }));
|
|
|
|
// No assistant message should have been created — the error hit before any event.
|
|
let messages = store.messages(session_id).await.unwrap();
|
|
assert_eq!(messages.len(), 1);
|
|
}
|
|
}
|