use salvor_llm::{ContentBlock, Message, MessageRequest, Tool, ToolChoice};
use serde_json::Value;
use crate::agent::Agent;
use crate::compact::FailureTracker;
use crate::ctx::{Resumption, RunCtx, ToolCallResult, Waking};
use crate::error::RuntimeError;
use crate::runtime::ParkReason;
use crate::validate::validate_against_schema;
use crate::wire::{content_string, slept_output};
use salvor_core::Effect;
use salvor_core::{BudgetExtensions, BudgetObservations};
use time::OffsetDateTime;
pub const ANSWER_TOOL: &str = "salvor_answer";
const ANSWER_TOOL_DESCRIPTION: &str = "Deliver your final reply by calling this tool; its input is \
the reply itself, and nothing else you write is read as the answer.";
const ANSWER_NOT_ALONE: &str = "`salvor_answer` was called alongside other tools; it ends the turn, \
so call it alone once the tool work it depends on has come back.";
const NO_TOOL_CALL_REASK: &str = "That turn called no tool. Call a tool, or deliver your final \
reply by calling `salvor_answer`.";
fn violation_content(violation: &str) -> String {
format!(
"`salvor_answer` was called with input that does not match its schema: {violation}. Call \
it again with input in the declared shape."
)
}
#[derive(Debug, Clone)]
pub enum LoopOutcome {
Completed(Value),
Parked(ParkReason),
}
pub(crate) async fn drive(
ctx: &mut RunCtx,
agent: &Agent,
initial_input: &Value,
) -> Result<LoopOutcome, RuntimeError> {
let input = begin(ctx, agent, initial_input).await?;
let outcome = drive_loop_inner(ctx, agent, &input, agent.output_schema()).await?;
if let LoopOutcome::Completed(output) = &outcome {
ctx.complete_run(output).await?;
}
Ok(outcome)
}
pub(crate) async fn begin(
ctx: &mut RunCtx,
agent: &Agent,
initial_input: &Value,
) -> Result<Value, RuntimeError> {
ctx.begin(agent.def_hash(), initial_input).await
}
pub async fn drive_loop(
ctx: &mut RunCtx,
agent: &Agent,
input: &Value,
) -> Result<LoopOutcome, RuntimeError> {
drive_loop_inner(ctx, agent, input, None).await
}
pub async fn drive_loop_structured(
ctx: &mut RunCtx,
agent: &Agent,
input: &Value,
schema: &Value,
) -> Result<LoopOutcome, RuntimeError> {
drive_loop_inner(ctx, agent, input, Some(schema)).await
}
async fn drive_loop_inner(
ctx: &mut RunCtx,
agent: &Agent,
input: &Value,
output_schema: Option<&Value>,
) -> Result<LoopOutcome, RuntimeError> {
let mut conversation: Vec<Message> = vec![Message::user(content_string(input))];
let mut llm_tools: Vec<Tool> = agent
.tools()
.descriptors()
.into_iter()
.map(|descriptor| Tool {
name: descriptor.name,
description: Some(descriptor.description),
input_schema: descriptor.input_schema,
})
.collect();
if let Some(schema) = output_schema {
if llm_tools.iter().any(|tool| tool.name == ANSWER_TOOL) {
return Err(RuntimeError::AnswerToolNameTaken);
}
llm_tools.push(Tool {
name: ANSWER_TOOL.to_owned(),
description: Some(ANSWER_TOOL_DESCRIPTION.to_owned()),
input_schema: schema.clone(),
});
}
let mut steps: u64 = 0;
let mut input_tokens: u64 = 0;
let mut output_tokens: u64 = 0;
let mut started_at: Option<OffsetDateTime> = None;
let mut extensions = BudgetExtensions::default();
let mut failures = FailureTracker::new();
loop {
let now = ctx.now().await?;
let start = *started_at.get_or_insert(now);
loop {
let observations = BudgetObservations {
steps,
input_tokens,
output_tokens,
elapsed_seconds: (now - start).as_seconds_f64(),
};
let Some((budget, observed)) =
agent
.budgets()
.first_crossing(&extensions, agent.pricing(), &observations)
else {
break;
};
ctx.budget_exceeded(budget, observed).await?;
match ctx.await_resume().await? {
Resumption::Parked => {
return Ok(LoopOutcome::Parked(ParkReason::BudgetExceeded {
budget,
observed,
}));
}
Resumption::Resumed(resume_input) => extensions.absorb(&resume_input),
}
}
let mut request = MessageRequest::new(agent.model(), agent.max_response_tokens())
.with_messages(conversation.clone());
if let Some(system) = agent.system_prompt() {
request = request.with_system(system);
}
if !llm_tools.is_empty() {
request = request.with_tools(llm_tools.clone());
}
if output_schema.is_some() {
request = request.with_tool_choice(ToolChoice::any());
}
let turn = ctx.model_call(agent.client(), &request).await?;
steps += 1;
input_tokens = input_tokens.saturating_add(u64::from(turn.usage.input_tokens));
output_tokens = output_tokens.saturating_add(u64::from(turn.usage.output_tokens));
let tool_uses: Vec<(String, String, Value)> = turn
.response
.tool_uses()
.into_iter()
.map(|(id, name, tool_input)| (id.to_owned(), name.to_owned(), tool_input.clone()))
.collect();
conversation.push(Message::assistant_blocks(turn.response.content.clone()));
if tool_uses.is_empty() {
if output_schema.is_none() {
let output = Value::String(turn.response.text());
return Ok(LoopOutcome::Completed(output));
}
conversation.push(Message::user(NO_TOOL_CALL_REASK));
continue;
}
if let Some(schema) = output_schema
&& let [(tool_use_id, name, answer)] = tool_uses.as_slice()
&& name == ANSWER_TOOL
{
match validate_against_schema(answer, schema) {
Ok(()) => return Ok(LoopOutcome::Completed(answer.clone())),
Err(violation) => {
let content =
failures.content_for_failure(ANSWER_TOOL, &violation_content(&violation));
conversation.push(Message::user_blocks(vec![ContentBlock::tool_error(
tool_use_id.clone(),
content,
)]));
continue;
}
}
}
let mut result_blocks: Vec<ContentBlock> = Vec::with_capacity(tool_uses.len());
for (tool_use_id, name, tool_input) in tool_uses {
if output_schema.is_some() && name == ANSWER_TOOL {
result_blocks.push(ContentBlock::tool_error(tool_use_id, ANSWER_NOT_ALONE));
continue;
}
let Some(tool) = agent.tools().get(&name) else {
result_blocks.push(ContentBlock::tool_error(
tool_use_id,
format!("unknown tool `{name}`"),
));
continue;
};
let idempotency_key = match tool.effect() {
Effect::Idempotent => Some(format!("{:016x}", ctx.random().await?)),
Effect::Read | Effect::Write => None,
};
match ctx
.tool_call(tool, &tool_input, idempotency_key.as_deref())
.await?
{
ToolCallResult::Output(output) => {
failures.record_success();
result_blocks.push(ContentBlock::tool_result(
tool_use_id,
content_string(&output),
));
}
ToolCallResult::Failed(failure) => {
let content = failures.content_for_failure(&name, &failure.message);
result_blocks.push(ContentBlock::tool_error(tool_use_id, content));
}
ToolCallResult::Suspended(suspension) => {
ctx.suspend_with_kind(
&suspension.reason,
&suspension.input_schema,
suspension.kind,
)
.await?;
match ctx.await_resume().await? {
Resumption::Parked => {
return Ok(LoopOutcome::Parked(ParkReason::Suspended {
reason: suspension.reason,
input_schema: suspension.input_schema,
kind: suspension.kind,
}));
}
Resumption::Resumed(resume_input) => {
failures.record_success();
result_blocks.push(ContentBlock::tool_result(
tool_use_id,
content_string(&resume_input),
));
}
}
}
ToolCallResult::Sleeping(sleep) => {
ctx.sleep_until(sleep.wake_at).await?;
match ctx.await_wake().await? {
Waking::Asleep { wake_at } => {
return Ok(LoopOutcome::Parked(ParkReason::Sleeping { wake_at }));
}
Waking::Woken => {
failures.record_success();
result_blocks.push(ContentBlock::tool_result(
tool_use_id,
content_string(&slept_output(sleep.wake_at)),
));
}
}
}
}
}
conversation.push(Message::user_blocks(result_blocks));
}
}