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, pub max_steps: u32, pub instructions: Vec, } /// 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 { 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 { match message.role { Role::User => { let content: Vec = 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, 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 = 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>>>, } impl MockProvider { fn scripted(steps: Vec>>) -> 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, ProviderError> { Ok(vec![]) } async fn stream( &self, _req: LlmRequest, _cancel: CancellationToken, ) -> Result { 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 { 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); } }