827 lines
29 KiB
Rust
827 lines
29 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>,
|
|
/// Pricing for `model`, from models.dev metadata. `None` leaves cost at 0.
|
|
pub cost: Option<crate::types::ModelCost>,
|
|
/// Whether to append the background job board to requests (primary/delegating agents).
|
|
pub inject_job_board: bool,
|
|
/// Optional user-provided reminder injected at the start of every turn (off by default).
|
|
pub reminder_turn_start: Option<String>,
|
|
/// Optional user-provided reminder injected on the turn after a file tool ran.
|
|
pub reminder_after_file_tool: Option<String>,
|
|
}
|
|
|
|
/// Adds a step's usage/cost onto the persisted session and republishes it. Cost accounting is
|
|
/// best-effort: a store error here is logged, not surfaced as a run failure.
|
|
async fn accumulate_session_usage(
|
|
ctx: &StepContext,
|
|
usage: &crate::types::TokenUsage,
|
|
cost: f64,
|
|
now: i64,
|
|
) {
|
|
match ctx.store.session(ctx.session_id.clone()).await {
|
|
Ok(Some(mut session)) => {
|
|
session.usage.add(usage);
|
|
session.cost += cost;
|
|
session.updated_at = now;
|
|
if let Err(e) = ctx.store.upsert_session(session.clone()).await {
|
|
tracing::warn!(error = %e, "failed to persist session usage");
|
|
return;
|
|
}
|
|
ctx.bus
|
|
.publish(crate::event::AppEvent::SessionUpdated { session });
|
|
}
|
|
Ok(None) => {}
|
|
Err(e) => tracing::warn!(error = %e, "failed to load session for usage accounting"),
|
|
}
|
|
}
|
|
|
|
/// 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;
|
|
// Whether the previous step ran a file tool, gating the `after_file_tool` reminder.
|
|
let mut prev_used_file_tool = false;
|
|
|
|
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));
|
|
}
|
|
|
|
// Collect synthetic (non-persisted) blocks to append to the last user message this
|
|
// turn: the optional turn-start reminder, the job board, and — if the previous step
|
|
// ran a file tool — the optional after-file-tool reminder. docs/04-multiagent.md.
|
|
let mut synthetic: Vec<String> = Vec::new();
|
|
if let Some(reminder) = &run_config.reminder_turn_start {
|
|
synthetic.push(reminder.clone());
|
|
}
|
|
if run_config.inject_job_board {
|
|
if let Some(board) = &ctx.job_board {
|
|
if let Some(block) = board.format_for_prompt() {
|
|
synthetic.push(block);
|
|
}
|
|
}
|
|
}
|
|
if prev_used_file_tool {
|
|
if let Some(reminder) = &run_config.reminder_after_file_tool {
|
|
synthetic.push(reminder.clone());
|
|
}
|
|
}
|
|
if !synthetic.is_empty() {
|
|
if let Some(last_user) = wire_messages
|
|
.iter_mut()
|
|
.rev()
|
|
.find(|m| m.role == WireRole::User)
|
|
{
|
|
for text in synthetic {
|
|
last_user.content.push(WireContent::Text { text });
|
|
}
|
|
}
|
|
}
|
|
|
|
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,
|
|
run_config.cost,
|
|
&mut doomloop,
|
|
)
|
|
.await;
|
|
|
|
match step {
|
|
Ok(outcome) if outcome.aborted => {
|
|
accumulate_session_usage(&ctx, &outcome.usage, outcome.cost, now_fn()).await;
|
|
return RunOutcome::Aborted;
|
|
}
|
|
Ok(outcome) => {
|
|
accumulate_session_usage(&ctx, &outcome.usage, outcome.cost, now_fn()).await;
|
|
prev_used_file_tool = outcome.used_file_tool;
|
|
// A completed step means the orchestrator has now seen any terminal jobs
|
|
// that were on the board this turn; mark them reconciled.
|
|
if run_config.inject_job_board {
|
|
if let Some(board) = &ctx.job_board {
|
|
if let Err(e) = board.reconcile_terminal(now_fn()).await {
|
|
tracing::warn!(error = %e, "failed to reconcile job board");
|
|
}
|
|
}
|
|
}
|
|
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())),
|
|
parent_rules: Vec::new(),
|
|
session_id,
|
|
cwd: cwd.clone(),
|
|
data_dir: cwd.join("tool-output"),
|
|
cancel: CancellationToken::new(),
|
|
now: 1,
|
|
spawner: None,
|
|
job_board: None,
|
|
context_reporter: None,
|
|
diagnostics: None,
|
|
}
|
|
}
|
|
|
|
#[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(),
|
|
// $3/1M input, $15/1M output.
|
|
cost: Some(crate::types::ModelCost {
|
|
input: 3.0,
|
|
output: 15.0,
|
|
..Default::default()
|
|
}),
|
|
inject_job_board: false,
|
|
reminder_turn_start: None,
|
|
reminder_after_file_tool: None,
|
|
};
|
|
|
|
let outcome = run_session(Arc::new(provider), ctx, &run_config, || 2).await;
|
|
assert!(matches!(outcome, RunOutcome::Stopped));
|
|
|
|
// Both steps' usage and cost accumulate onto the session.
|
|
let session = store.session(session_id.clone()).await.unwrap().unwrap();
|
|
assert_eq!(session.usage.input, 30);
|
|
assert_eq!(session.usage.output, 13);
|
|
// (10*3 + 5*15)/1e6 + (20*3 + 8*15)/1e6 = 0.000105 + 0.00018
|
|
assert!(
|
|
(session.cost - 0.000_285).abs() < 1e-9,
|
|
"cost = {}",
|
|
session.cost
|
|
);
|
|
|
|
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(),
|
|
cost: None,
|
|
inject_job_board: false,
|
|
reminder_turn_start: None,
|
|
reminder_after_file_tool: None,
|
|
};
|
|
|
|
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);
|
|
}
|
|
|
|
/// Records the last request it was asked to stream so tests can assert on prompt content.
|
|
struct CapturingProvider {
|
|
last: StdMutex<Option<LlmRequest>>,
|
|
}
|
|
|
|
#[async_trait]
|
|
impl Provider for CapturingProvider {
|
|
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> {
|
|
*self.last.lock().unwrap() = Some(req);
|
|
let events = vec![
|
|
Ok(LlmEvent::TextStart { id: "t".into() }),
|
|
Ok(LlmEvent::TextDelta {
|
|
id: "t".into(),
|
|
text: "ok".into(),
|
|
}),
|
|
Ok(LlmEvent::TextEnd { id: "t".into() }),
|
|
Ok(LlmEvent::Finish {
|
|
reason: FinishReason::Stop,
|
|
usage: usage(1, 1),
|
|
}),
|
|
];
|
|
Ok(Box::pin(futures::stream::iter(events)))
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn job_board_is_injected_into_the_last_user_message() {
|
|
use crate::engine::jobs::{JobBoard, LaunchSpec};
|
|
|
|
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: "carry on".into(),
|
|
synthetic: false,
|
|
},
|
|
})
|
|
.await
|
|
.unwrap();
|
|
|
|
// A board with one running job for this session.
|
|
let board = std::sync::Arc::new(
|
|
JobBoard::load(store.clone(), bus.clone(), &session_id, 2)
|
|
.await
|
|
.unwrap(),
|
|
);
|
|
board
|
|
.register_launch(
|
|
LaunchSpec {
|
|
task_id: "t1".into(),
|
|
parent_session: session_id.clone(),
|
|
child_session: SessionId::new(),
|
|
agent: "explorer".into(),
|
|
description: "map auth".into(),
|
|
objective: Some("map the auth flow".into()),
|
|
},
|
|
1,
|
|
)
|
|
.await
|
|
.unwrap();
|
|
|
|
let cwd = tempfile::tempdir().unwrap();
|
|
let mut ctx = make_ctx(store, bus, session_id, cwd.path().to_path_buf()).await;
|
|
ctx.job_board = Some(board);
|
|
let run_config = RunConfig {
|
|
agent_name: "orchestrator".into(),
|
|
agent_prompt: "You orchestrate.".into(),
|
|
model,
|
|
temperature: None,
|
|
max_steps: 1,
|
|
instructions: Vec::new(),
|
|
cost: None,
|
|
inject_job_board: true,
|
|
reminder_turn_start: None,
|
|
reminder_after_file_tool: None,
|
|
};
|
|
|
|
let provider = std::sync::Arc::new(CapturingProvider {
|
|
last: StdMutex::new(None),
|
|
});
|
|
let outcome = run_session(provider.clone(), ctx, &run_config, || 2).await;
|
|
assert!(matches!(outcome, RunOutcome::Stopped));
|
|
|
|
let req = provider.last.lock().unwrap().clone().expect("a request");
|
|
let last_user = req
|
|
.messages
|
|
.iter()
|
|
.rev()
|
|
.find(|m| m.role == WireRole::User)
|
|
.expect("a user message");
|
|
let text: String = last_user
|
|
.content
|
|
.iter()
|
|
.filter_map(|c| match c {
|
|
WireContent::Text { text } => Some(text.as_str()),
|
|
_ => None,
|
|
})
|
|
.collect::<Vec<_>>()
|
|
.join("\n");
|
|
assert!(text.contains("Background Job Board"), "got: {text}");
|
|
assert!(text.contains("exp-1"), "got: {text}");
|
|
assert!(text.contains("map the auth flow"), "got: {text}");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn turn_start_reminder_is_injected_into_the_request() {
|
|
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: "do the thing".into(),
|
|
synthetic: false,
|
|
},
|
|
})
|
|
.await
|
|
.unwrap();
|
|
|
|
let cwd = tempfile::tempdir().unwrap();
|
|
let ctx = make_ctx(store, bus, session_id, cwd.path().to_path_buf()).await;
|
|
let run_config = RunConfig {
|
|
agent_name: "orchestrator".into(),
|
|
agent_prompt: "You orchestrate.".into(),
|
|
model,
|
|
temperature: None,
|
|
max_steps: 1,
|
|
instructions: Vec::new(),
|
|
cost: None,
|
|
inject_job_board: false,
|
|
reminder_turn_start: Some("REMEMBER: stay on task.".into()),
|
|
reminder_after_file_tool: None,
|
|
};
|
|
|
|
let provider = std::sync::Arc::new(CapturingProvider {
|
|
last: StdMutex::new(None),
|
|
});
|
|
let outcome = run_session(provider.clone(), ctx, &run_config, || 2).await;
|
|
assert!(matches!(outcome, RunOutcome::Stopped));
|
|
|
|
let req = provider.last.lock().unwrap().clone().expect("a request");
|
|
let last_user = req
|
|
.messages
|
|
.iter()
|
|
.rev()
|
|
.find(|m| m.role == WireRole::User)
|
|
.expect("a user message");
|
|
let has_reminder = last_user.content.iter().any(
|
|
|c| matches!(c, WireContent::Text { text } if text.contains("REMEMBER: stay on task.")),
|
|
);
|
|
assert!(has_reminder, "turn-start reminder should be injected");
|
|
}
|
|
}
|