use std::collections::BTreeMap;
use std::sync::Arc;
use base64::Engine as _;
use serde_json::{json, Value};
use crate::backends::loop_util::resolve_tool_args;
use crate::backends::anthropic::api::SharedClient;
use crate::backends::anthropic::wire::{
Block, BlockDelta, ImageSource, Message, MessagesRequest, Role, StopReason, StreamEvent,
ThinkingConfig, ToolDef, WireUsage, DEFAULT_MAX_TOKENS,
};
use crate::backends::turn_engine::{
self, DispatchedResult, EmitCtx, EngineDeps, ResolvedCall, StreamEnd, TurnProvider,
};
use crate::content::{Content, Part as ApiPart};
use crate::error::{Error, Result};
use crate::hooks::{HookRunner, SessionContext};
use crate::tools::ToolRunner;
use crate::types::{StepStatus, SystemInstructions, ThinkingLevel, UsageMetadata};
const MAX_PAUSE_RESUMES: u32 = 8;
#[derive(Clone)]
pub(crate) struct LoopConfig {
pub model: String,
pub system: Option<String>,
pub thinking: Option<ThinkingLevel>,
pub temperature: Option<f32>,
pub max_tokens: u32,
pub tool_declarations: Vec<ToolDef>,
pub compaction_threshold: Option<u32>,
pub compaction_epilogue: Option<String>,
}
impl LoopConfig {
#[allow(clippy::too_many_arguments)] pub fn from_system(
model: String,
system: Option<&SystemInstructions>,
thinking: Option<ThinkingLevel>,
temperature: Option<f32>,
max_tokens: Option<u32>,
tool_declarations: Vec<ToolDef>,
compaction_threshold: Option<u32>,
compaction_epilogue: Option<String>,
) -> Result<Self> {
let system = system.map(render_system);
Ok(Self {
model,
system,
thinking,
temperature,
max_tokens: max_tokens.unwrap_or(DEFAULT_MAX_TOKENS),
tool_declarations,
compaction_threshold,
compaction_epilogue,
})
}
}
pub(crate) use crate::backends::render_system;
pub(crate) type LoopState = crate::backends::state::LoopState<Message>;
pub(crate) fn to_wire_user_content(content: Content) -> Result<Message> {
let mut blocks: Vec<Block> = Vec::with_capacity(content.parts.len().max(1));
for p in content.parts {
match p {
ApiPart::Text(t) => blocks.push(Block::Text { text: t }),
ApiPart::Media(m) => blocks.push(Block::Image {
source: ImageSource {
source_type: "base64".to_string(),
media_type: m.mime_type,
data: base64::engine::general_purpose::STANDARD.encode(m.data.as_ref()),
},
}),
}
}
if blocks.is_empty() {
return Err(Error::config("empty content"));
}
Ok(Message {
role: Role::User,
content: blocks,
})
}
#[derive(Clone)]
pub(crate) struct TurnDeps {
pub client: SharedClient,
pub config: LoopConfig,
pub state: Arc<LoopState>,
pub tool_runner: Option<Arc<ToolRunner>>,
pub hook_runner: Option<Arc<HookRunner>>,
pub session_ctx: Option<SessionContext>,
}
#[derive(Default)]
struct ToolUseAccum {
id: String,
name: String,
args_json: String,
}
#[derive(Default)]
struct ThinkingAccum {
thinking: String,
signature: String,
}
#[derive(Default)]
pub(crate) struct RoundAccum {
thinking_blocks: BTreeMap<u32, ThinkingAccum>,
tool_blocks: BTreeMap<u32, ToolUseAccum>,
stop_reason: Option<StopReason>,
usage: WireUsage,
}
pub(crate) struct AnthropicProvider;
impl TurnProvider for AnthropicProvider {
type Message = Message;
type Config = LoopConfig;
type Request = MessagesRequest;
type Event = StreamEvent;
type Accum = RoundAccum;
fn build_request(config: &LoopConfig, history: &[Message]) -> MessagesRequest {
build_request(config, history)
}
fn compaction_threshold(config: &LoopConfig) -> Option<u32> {
config.compaction_threshold
}
fn fold_event(
acc: &mut RoundAccum,
ctx: &mut EmitCtx<'_, Message>,
ev: StreamEvent,
) -> Result<()> {
match ev {
StreamEvent::MessageStart { message } => {
if let Some(u) = message.usage {
accumulate_wire_usage(&mut acc.usage, &u);
}
}
StreamEvent::ContentBlockStart {
index,
content_block,
} => match content_block {
Block::ToolUse { id, name, .. } => {
acc.tool_blocks.insert(
index,
ToolUseAccum {
id,
name,
args_json: String::new(),
},
);
}
Block::Text { text } => ctx.push_text(&text),
_ => {}
},
StreamEvent::ContentBlockDelta { index, delta } => match delta {
BlockDelta::TextDelta { text } => ctx.push_text(&text),
BlockDelta::ThinkingDelta { thinking } => {
if !thinking.is_empty() {
acc.thinking_blocks
.entry(index)
.or_default()
.thinking
.push_str(&thinking);
ctx.push_thought(&thinking);
}
}
BlockDelta::SignatureDelta { signature } => {
acc.thinking_blocks.entry(index).or_default().signature = signature;
}
BlockDelta::InputJsonDelta { partial_json } => {
if let Some(a) = acc.tool_blocks.get_mut(&index) {
a.args_json.push_str(&partial_json);
}
}
_ => {}
},
StreamEvent::MessageDelta { delta, usage } => {
if let Some(r) = delta.stop_reason {
acc.stop_reason = Some(r);
}
if let Some(u) = usage {
accumulate_wire_usage(&mut acc.usage, &u);
}
}
StreamEvent::Error { error } => {
return Err(Error::other(format!(
"anthropic stream error [{}]: {}",
error.kind, error.message
)));
}
StreamEvent::ContentBlockStop { .. }
| StreamEvent::MessageStop
| StreamEvent::Ping
| StreamEvent::Unknown => {}
}
Ok(())
}
fn resolve_pending_calls(acc: &mut RoundAccum) -> Vec<ResolvedCall> {
std::mem::take(&mut acc.tool_blocks)
.into_values()
.map(|a| {
let (args, parse_error) = resolve_tool_args(&a.name, &a.args_json);
ResolvedCall {
id: Some(a.id),
name: a.name,
args,
parse_error,
}
})
.collect()
}
fn round_usage(acc: &RoundAccum) -> UsageMetadata {
acc.usage.clone().into()
}
fn map_finish_reason(acc: &RoundAccum) -> (StepStatus, &'static str) {
match acc.stop_reason {
Some(StopReason::Refusal) => (StepStatus::Error, "stopped by refusal"),
Some(StopReason::MaxTokens) => (StepStatus::Done, "stopped at max tokens"),
Some(StopReason::PauseTurn) => (StepStatus::Done, "paused (resume cap reached)"),
_ => (StepStatus::Done, ""),
}
}
fn assemble_assistant_message(
acc: RoundAccum,
text: &str,
calls: &[ResolvedCall],
) -> Option<Message> {
let mut blocks: Vec<Block> = Vec::new();
for (_idx, t) in acc.thinking_blocks {
if !t.thinking.is_empty() && !t.signature.is_empty() {
blocks.push(Block::Thinking {
thinking: t.thinking,
signature: Some(t.signature),
});
}
}
if !text.is_empty() {
blocks.push(Block::Text {
text: text.to_string(),
});
}
for c in calls {
blocks.push(Block::ToolUse {
id: c.id.clone().unwrap_or_default(),
name: c.name.clone(),
input: c.args.clone(),
});
}
(!blocks.is_empty()).then_some(Message {
role: Role::Assistant,
content: blocks,
})
}
fn tool_result_messages(results: Vec<DispatchedResult>) -> Vec<Message> {
if results.is_empty() {
return Vec::new();
}
let blocks: Vec<Block> = results
.into_iter()
.map(|r| Block::ToolResult {
tool_use_id: r.call.id.unwrap_or_default(),
content: tool_result_content(&r.value),
is_error: r.is_error.then_some(true),
})
.collect();
vec![Message {
role: Role::User,
content: blocks,
}]
}
fn on_stream_end(acc: &mut RoundAccum, pause_resumes: u32) -> StreamEnd {
if !matches!(acc.stop_reason, Some(StopReason::PauseTurn)) {
return StreamEnd::Proceed;
}
if pause_resumes < MAX_PAUSE_RESUMES {
StreamEnd::Resume
} else {
StreamEnd::ProceedAndEndTurn
}
}
fn on_cancel_with_pending_calls(calls: &[ResolvedCall]) -> Vec<Message> {
let blocks: Vec<Block> = calls
.iter()
.map(|c| Block::ToolResult {
tool_use_id: c.id.clone().unwrap_or_default(),
content: tool_result_content(&json!({ "error": "cancelled" })),
is_error: Some(true),
})
.collect();
vec![Message {
role: Role::User,
content: blocks,
}]
}
}
pub(crate) async fn run_turn(deps: TurnDeps, user: Message, prompt: Content) -> Result<()> {
let TurnDeps {
client,
config,
state,
tool_runner,
hook_runner,
session_ctx,
} = deps;
let model = config.model.clone();
let compact_epilogue = config.compaction_epilogue.clone();
let engine_deps = EngineDeps::<AnthropicProvider> {
config,
state: state.clone(),
tool_runner,
hook_runner,
session_ctx,
};
let open_client = client.clone();
turn_engine::run_turn::<AnthropicProvider, _, _, _, _, _>(
engine_deps,
user,
prompt,
move |req: MessagesRequest| {
let client = open_client.clone();
async move { client.stream_messages(&req).await }
},
move || async move {
crate::backends::anthropic::compaction::try_compact(
&state.history,
&client,
&model,
compact_epilogue.as_deref(),
)
.await;
},
)
.await
}
fn tool_result_content(v: &Value) -> Value {
match v {
Value::String(_) => v.clone(),
other => Value::String(other.to_string()),
}
}
pub(crate) fn build_request(config: &LoopConfig, history: &[Message]) -> MessagesRequest {
let (thinking, max_tokens) = match config.thinking.map(thinking_level_to_budget) {
Some(budget) => {
let max = config.max_tokens.max(budget + 1024);
(Some(ThinkingConfig::enabled(budget)), max)
}
None => (None, config.max_tokens),
};
MessagesRequest {
model: config.model.clone(),
max_tokens,
system: MessagesRequest::system_from(config.system.clone()),
messages: history.to_vec(),
tools: config.tool_declarations.clone(),
tool_choice: None,
stream: true,
temperature: if thinking.is_some() {
None
} else {
config.temperature
},
thinking,
}
}
fn thinking_level_to_budget(level: ThinkingLevel) -> u32 {
match level {
ThinkingLevel::Minimal => 1024,
ThinkingLevel::Low => 2048,
ThinkingLevel::Medium => 8192,
ThinkingLevel::High => 16384,
}
}
fn accumulate_wire_usage(acc: &mut WireUsage, other: &WireUsage) {
fn take_latest(a: &mut Option<i32>, b: Option<i32>) {
if b.is_some() {
*a = b;
}
}
take_latest(&mut acc.input_tokens, other.input_tokens);
take_latest(&mut acc.output_tokens, other.output_tokens);
take_latest(
&mut acc.cache_read_input_tokens,
other.cache_read_input_tokens,
);
take_latest(
&mut acc.cache_creation_input_tokens,
other.cache_creation_input_tokens,
);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::types::{CustomSystemInstructions, Step, StepStatus, StreamChunk, SystemInstructions};
use tokio::sync::broadcast;
#[test]
fn render_system_custom() {
let s = SystemInstructions::Custom(CustomSystemInstructions {
text: "be terse".into(),
});
assert_eq!(render_system(&s), "be terse");
}
#[test]
fn tool_result_content_objects_become_json_strings() {
let obj = json!({"contents": "fn main() {}", "lines": 1});
let wire = tool_result_content(&obj);
assert!(wire.is_string(), "object must serialize to a string, got {wire}");
let back: Value = serde_json::from_str(wire.as_str().unwrap()).unwrap();
assert_eq!(back, obj);
assert!(tool_result_content(&json!({"ok": true})).is_string());
assert!(tool_result_content(&json!({"error": "boom"})).is_string());
assert!(tool_result_content(&json!(["a", "b"])).is_string());
let s = json!("plain text result");
assert_eq!(tool_result_content(&s), json!("plain text result"));
}
#[test]
fn build_request_clamps_thinking_below_max_tokens() {
let config = LoopConfig {
model: "claude-haiku-4-5-20251001".into(),
system: None,
thinking: Some(ThinkingLevel::High),
temperature: Some(0.7),
max_tokens: 8192, tool_declarations: Vec::new(),
compaction_threshold: None,
compaction_epilogue: None,
};
let req = build_request(&config, &[Message::user_text("hi")]);
let thinking = req.thinking.expect("thinking enabled");
assert!(
req.max_tokens > thinking.budget_tokens,
"max_tokens ({}) must exceed budget ({})",
req.max_tokens,
thinking.budget_tokens
);
assert!(req.temperature.is_none());
}
#[test]
fn build_request_no_thinking_keeps_temperature() {
let config = LoopConfig {
model: "claude-haiku-4-5-20251001".into(),
system: Some("sys".into()),
thinking: None,
temperature: Some(0.3),
max_tokens: 4096,
tool_declarations: Vec::new(),
compaction_threshold: None,
compaction_epilogue: None,
};
let req = build_request(&config, &[Message::user_text("hi")]);
assert!(req.thinking.is_none());
assert_eq!(req.temperature, Some(0.3));
assert_eq!(req.max_tokens, 4096);
assert_eq!(req.system.len(), 1);
assert_eq!(req.system[0].text, "sys");
assert!(req.system[0].cache_control.is_some());
}
#[test]
fn usage_does_not_double_count_message_start_output_placeholder() {
let mut round = WireUsage::default();
accumulate_wire_usage(
&mut round,
&WireUsage {
input_tokens: Some(12),
output_tokens: Some(1), cache_read_input_tokens: Some(8),
cache_creation_input_tokens: None,
},
);
accumulate_wire_usage(
&mut round,
&WireUsage {
input_tokens: None,
output_tokens: Some(33),
cache_read_input_tokens: None,
cache_creation_input_tokens: None,
},
);
assert_eq!(round.input_tokens, Some(12), "input from message_start");
assert_eq!(
round.cache_read_input_tokens,
Some(8),
"cache_read from message_start"
);
assert_eq!(
round.output_tokens,
Some(33),
"output_tokens must be the cumulative message_delta value (33), \
not message_start placeholder (1) + 33 = 34"
);
let neutral: UsageMetadata = round.into();
assert_eq!(neutral.candidates_token_count, Some(33));
assert_eq!(neutral.prompt_token_count, Some(12));
assert_eq!(neutral.total_token_count, Some(45)); }
#[test]
fn assistant_turn_preserves_signed_thinking_block_before_tool_use() {
let mut acc = RoundAccum::default();
acc.thinking_blocks.insert(
0,
ThinkingAccum {
thinking: "Let me reason about this.".into(),
signature: "sig_abc".into(),
},
);
let pending_calls = vec![ResolvedCall {
id: Some("toolu_1".into()),
name: "view_file".into(),
args: json!({"path": "a.rs"}),
parse_error: None,
}];
let msg = AnthropicProvider::assemble_assistant_message(acc, "I'll read it.", &pending_calls)
.expect("assistant message assembled");
assert_eq!(msg.role, Role::Assistant);
let assistant_blocks = msg.content;
match &assistant_blocks[0] {
Block::Thinking { thinking, signature } => {
assert_eq!(thinking, "Let me reason about this.");
assert_eq!(signature.as_deref(), Some("sig_abc"));
}
other => panic!("expected leading Thinking block, got {other:?}"),
}
assert!(matches!(assistant_blocks[1], Block::Text { .. }));
assert!(matches!(assistant_blocks[2], Block::ToolUse { .. }));
let wire = serde_json::to_value(&assistant_blocks[0]).unwrap();
assert_eq!(wire["type"], "thinking");
assert_eq!(wire["thinking"], "Let me reason about this.");
assert_eq!(wire["signature"], "sig_abc");
}
#[test]
fn unsigned_thinking_block_is_dropped() {
let mut acc = RoundAccum::default();
acc.thinking_blocks.insert(
0,
ThinkingAccum {
thinking: "partial reasoning".into(),
signature: String::new(), },
);
let msg = AnthropicProvider::assemble_assistant_message(acc, "", &[]);
assert!(
msg.is_none(),
"unsigned thinking must not be persisted"
);
}
#[test]
fn cancelled_turn_balances_pending_tool_use_with_tool_results() {
let pending_calls = vec![
ResolvedCall {
id: Some("toolu_1".into()),
name: "view_file".into(),
args: json!({"path": "a.rs"}),
parse_error: None,
},
ResolvedCall {
id: Some("toolu_2".into()),
name: "list_directory".into(),
args: json!({}),
parse_error: None,
},
];
let balance = AnthropicProvider::on_cancel_with_pending_calls(&pending_calls);
assert_eq!(balance.len(), 1, "anthropic balances with one user turn");
assert_eq!(balance[0].role, Role::User);
let cancelled_blocks = &balance[0].content;
assert_eq!(cancelled_blocks.len(), 2);
let ids: Vec<&str> = cancelled_blocks
.iter()
.map(|b| match b {
Block::ToolResult {
tool_use_id,
is_error,
content,
} => {
assert_eq!(*is_error, Some(true), "cancelled result must be is_error");
assert!(content.is_string(), "tool_result.content must be a string");
assert!(
content.as_str().unwrap().contains("cancelled"),
"content should mark the call cancelled"
);
tool_use_id.as_str()
}
other => panic!("expected ToolResult block, got {other:?}"),
})
.collect();
assert_eq!(ids, vec!["toolu_1", "toolu_2"], "every tool_use id answered");
}
#[test]
fn usage_takes_latest_cumulative_output_across_multiple_deltas() {
let mut round = WireUsage::default();
accumulate_wire_usage(
&mut round,
&WireUsage {
input_tokens: Some(20),
output_tokens: Some(2),
..Default::default()
},
);
accumulate_wire_usage(
&mut round,
&WireUsage {
output_tokens: Some(10),
..Default::default()
},
);
accumulate_wire_usage(
&mut round,
&WireUsage {
output_tokens: Some(25),
..Default::default()
},
);
assert_eq!(
round.output_tokens,
Some(25),
"cumulative deltas: final reported output is the LAST value (25), not 2+10+25"
);
assert_eq!(round.input_tokens, Some(20));
}
#[tokio::test]
async fn inline_tool_call_step_is_done_so_dispatcher_skips_it() {
let (tx, mut rx) = broadcast::channel::<Step>(8);
let state = LoopState::new(tx);
state.emit_chunk_step(StreamChunk::ToolCall(crate::types::ToolCall {
name: "create_file".into(),
args: serde_json::Value::Null,
id: None,
canonical_path: None,
}));
let step = rx.recv().await.expect("a tool-call step was emitted");
assert_eq!(
step.status,
StepStatus::Done,
"inline-dispatched tool-call step must be Done, not Active",
);
}
#[test]
fn provider_fold_accumulates_index_keyed_deltas() {
let (tx, _rx) = broadcast::channel::<Step>(8);
let state = LoopState::new(tx);
let mut acc = RoundAccum::default();
turn_engine::test_fold_events::<AnthropicProvider>(
&state,
&mut acc,
vec![
StreamEvent::ContentBlockDelta {
index: 0,
delta: BlockDelta::ThinkingDelta { thinking: "reason ".into() },
},
StreamEvent::ContentBlockDelta {
index: 0,
delta: BlockDelta::ThinkingDelta { thinking: "more".into() },
},
StreamEvent::ContentBlockDelta {
index: 0,
delta: BlockDelta::SignatureDelta { signature: "sig_x".into() },
},
StreamEvent::ContentBlockStart {
index: 1,
content_block: Block::ToolUse {
id: "toolu_9".into(),
name: "view_file".into(),
input: Value::Null,
},
},
StreamEvent::ContentBlockDelta {
index: 1,
delta: BlockDelta::InputJsonDelta { partial_json: "{\"path\":".into() },
},
StreamEvent::ContentBlockDelta {
index: 1,
delta: BlockDelta::InputJsonDelta { partial_json: "\"a.rs\"}".into() },
},
StreamEvent::MessageDelta {
delta: crate::backends::anthropic::wire::MessageDeltaBody {
stop_reason: Some(StopReason::ToolUse),
stop_sequence: None,
},
usage: None,
},
],
);
let t = &acc.thinking_blocks[&0];
assert_eq!(t.thinking, "reason more", "thinking deltas concatenate per index");
assert_eq!(t.signature, "sig_x", "trailing signature lands on the same index");
let calls = AnthropicProvider::resolve_pending_calls(&mut acc);
assert_eq!(calls.len(), 1);
assert_eq!(calls[0].id.as_deref(), Some("toolu_9"));
assert_eq!(calls[0].name, "view_file");
assert_eq!(calls[0].args, json!({"path": "a.rs"}), "args fragments reassemble");
assert!(calls[0].parse_error.is_none());
assert_eq!(acc.stop_reason, Some(StopReason::ToolUse));
}
#[test]
fn pause_turn_resumes_until_cap_then_ends_turn() {
let mut acc = RoundAccum {
stop_reason: Some(StopReason::EndTurn),
..Default::default()
};
assert!(matches!(
AnthropicProvider::on_stream_end(&mut acc, 0),
StreamEnd::Proceed
));
acc.stop_reason = Some(StopReason::PauseTurn);
assert!(matches!(
AnthropicProvider::on_stream_end(&mut acc, 0),
StreamEnd::Resume
));
assert_eq!(
acc.stop_reason,
Some(StopReason::PauseTurn),
"retained so a cancelled pause still reports the paused finish reason"
);
assert!(matches!(
AnthropicProvider::on_stream_end(&mut acc, MAX_PAUSE_RESUMES - 1),
StreamEnd::Resume
));
assert!(matches!(
AnthropicProvider::on_stream_end(&mut acc, MAX_PAUSE_RESUMES),
StreamEnd::ProceedAndEndTurn
));
let (status, msg) = AnthropicProvider::map_finish_reason(&acc);
assert_eq!(status, StepStatus::Done);
assert_eq!(msg, "paused (resume cap reached)");
}
#[test]
fn tool_result_messages_batch_into_one_user_turn() {
let mk = |id: &str, value: Value, is_error: bool| DispatchedResult {
call: ResolvedCall {
id: Some(id.into()),
name: "t".into(),
args: json!({}),
parse_error: None,
},
value,
is_error,
};
let msgs = AnthropicProvider::tool_result_messages(vec![
mk("toolu_1", json!({"ok": true}), false),
mk("toolu_2", json!({"error": "boom"}), true),
]);
assert_eq!(msgs.len(), 1, "one batched user turn");
assert_eq!(msgs[0].role, Role::User);
match &msgs[0].content[0] {
Block::ToolResult { tool_use_id, is_error, content } => {
assert_eq!(tool_use_id, "toolu_1");
assert_eq!(*is_error, None, "success omits is_error");
assert!(content.is_string());
}
other => panic!("expected ToolResult, got {other:?}"),
}
match &msgs[0].content[1] {
Block::ToolResult { tool_use_id, is_error, .. } => {
assert_eq!(tool_use_id, "toolu_2");
assert_eq!(*is_error, Some(true), "failure marks is_error");
}
other => panic!("expected ToolResult, got {other:?}"),
}
}
}