Files
ai-harness/crates/harness-core/src/engine/loop.rs
T
Erik Simon a68ca02894 M1 core: tool trait, permission service, config, Provider trait, engine loop
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.
2026-07-08 17:05:37 +02:00

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);
}
}