use crate::http_input::llm_client::{LlmClient, LlmProvider};
use crate::reasoning::conversation::Conversation;
use crate::reasoning::inference::*;
use async_trait::async_trait;
fn map_anthropic_stop_reason(stop_reason: &str, has_tool_calls: bool) -> FinishReason {
match stop_reason {
"tool_use" => FinishReason::ToolCalls,
"max_tokens" => FinishReason::MaxTokens,
"refusal" => FinishReason::Refusal,
_ if has_tool_calls => FinishReason::ToolCalls,
_ => FinishReason::Stop,
}
}
fn parse_token_usage(
value: Option<&serde_json::Value>,
anthropic: bool,
) -> Result<Usage, InferenceError> {
let Some(value) = value.filter(|value| !value.is_null()) else {
return Ok(Usage::default());
};
if !value.is_object() {
return Err(InferenceError::ParseError("usage must be an object".into()));
}
let count = |key: &str| -> Result<Option<u32>, InferenceError> {
value
.get(key)
.map(|value| {
value
.as_u64()
.and_then(|n| u32::try_from(n).ok())
.ok_or_else(|| {
InferenceError::ParseError(format!("invalid token count: {key}"))
})
})
.transpose()
};
let (input, output, total) = if anthropic {
let input = count("input_tokens")?;
let output = count("output_tokens")?;
let creation = count("cache_creation_input_tokens")?.unwrap_or(0);
let read = count("cache_read_input_tokens")?.unwrap_or(0);
let (Some(input), Some(output)) = (input, output) else {
return Ok(Usage::default());
};
let input = input
.checked_add(creation)
.and_then(|n| n.checked_add(read))
.ok_or_else(|| InferenceError::ParseError("input token count overflow".into()))?;
let total = input
.checked_add(output)
.ok_or_else(|| InferenceError::ParseError("total token count overflow".into()))?;
(input, output, total)
} else {
let (input, output, total) = (
count("prompt_tokens")?,
count("completion_tokens")?,
count("total_tokens")?,
);
let (Some(input), Some(output), Some(total)) = (input, output, total) else {
return Ok(Usage::default());
};
if input.checked_add(output).is_none_or(|sum| sum > total) {
return Err(InferenceError::ParseError(
"inconsistent total token count".into(),
));
}
(input, output, total)
};
Ok(Usage {
prompt_tokens: input,
completion_tokens: output,
total_tokens: total,
})
}
pub struct CloudInferenceProvider {
client: LlmClient,
}
impl CloudInferenceProvider {
pub fn new(client: LlmClient) -> Self {
Self { client }
}
pub fn from_env() -> Option<Self> {
LlmClient::from_env().map(|c| Self { client: c })
}
pub async fn from_env_or_secrets(
store: Option<std::sync::Arc<dyn crate::secrets::SecretStore + Send + Sync>>,
) -> Option<Self> {
LlmClient::from_env_or_secrets(store)
.await
.map(|c| Self { client: c })
}
fn build_openai_body(
&self,
conversation: &Conversation,
options: &InferenceOptions,
) -> serde_json::Value {
let model = options
.model
.as_deref()
.unwrap_or_else(|| self.client.model());
let mut body = serde_json::json!({
"model": model,
"messages": conversation.to_openai_messages(),
"max_tokens": options.max_tokens,
"temperature": options.temperature,
});
if !options.tool_definitions.is_empty() {
let tools: Vec<serde_json::Value> = options
.tool_definitions
.iter()
.map(|td| {
serde_json::json!({
"type": "function",
"function": {
"name": td.name,
"description": td.description,
"parameters": td.parameters,
}
})
})
.collect();
body["tools"] = serde_json::Value::Array(tools);
if let Some(choice) = &options.tool_choice {
body["tool_choice"] = match choice {
crate::reasoning::inference::ToolChoice::Auto => {
serde_json::Value::String("auto".into())
}
crate::reasoning::inference::ToolChoice::Any => {
serde_json::Value::String("required".into())
}
crate::reasoning::inference::ToolChoice::Tool { name } => {
serde_json::json!({
"type": "function",
"function": {"name": name}
})
}
};
}
}
match &options.response_format {
ResponseFormat::Text => {}
ResponseFormat::JsonObject => {
body["response_format"] = serde_json::json!({"type": "json_object"});
}
ResponseFormat::JsonSchema { schema, name } => {
body["response_format"] = serde_json::json!({
"type": "json_schema",
"json_schema": {
"name": name.as_deref().unwrap_or("response"),
"schema": schema,
}
});
}
}
body
}
fn build_anthropic_body(
&self,
conversation: &Conversation,
options: &InferenceOptions,
) -> serde_json::Value {
let model = options
.model
.as_deref()
.unwrap_or_else(|| self.client.model());
let (system, messages) = conversation.to_anthropic_messages();
let mut body = serde_json::json!({
"model": model,
"messages": messages,
"max_tokens": options.max_tokens,
});
if options.temperature > 0.0 {
body["temperature"] = serde_json::json!(options.temperature);
}
let has_system = system.is_some();
if let Some(sys) = system {
body["system"] = serde_json::json!([
{
"type": "text",
"text": sys,
"cache_control": { "type": "ephemeral" }
}
]);
}
if !options.tool_definitions.is_empty() {
let mut tools: Vec<serde_json::Value> = options
.tool_definitions
.iter()
.map(|td| {
serde_json::json!({
"name": td.name,
"description": td.description,
"input_schema": td.parameters,
})
})
.collect();
if !has_system {
if let Some(last_tool) = tools.last_mut().and_then(|t| t.as_object_mut()) {
last_tool.insert(
"cache_control".to_string(),
serde_json::json!({"type": "ephemeral"}),
);
}
}
body["tools"] = serde_json::Value::Array(tools);
if let Some(choice) = &options.tool_choice {
body["tool_choice"] = match choice {
crate::reasoning::inference::ToolChoice::Auto => {
serde_json::json!({"type": "auto"})
}
crate::reasoning::inference::ToolChoice::Any => {
serde_json::json!({"type": "any"})
}
crate::reasoning::inference::ToolChoice::Tool { name } => {
serde_json::json!({"type": "tool", "name": name})
}
};
}
}
for (k, v) in &options.extra {
body[k] = v.clone();
}
body
}
fn parse_openai_response(
&self,
resp: &serde_json::Value,
model: &str,
) -> Result<InferenceResponse, InferenceError> {
let choice = resp
.get("choices")
.and_then(|c| c.get(0))
.ok_or_else(|| InferenceError::ParseError("No choices in response".into()))?;
let message = choice
.get("message")
.ok_or_else(|| InferenceError::ParseError("No message in choice".into()))?;
let content = message
.get("content")
.and_then(|c| c.as_str())
.unwrap_or("")
.to_string();
let tool_calls = message
.get("tool_calls")
.and_then(|tc| tc.as_array())
.map(|arr| {
arr.iter()
.filter_map(|tc| {
let id = tc.get("id")?.as_str()?.to_string();
let func = tc.get("function")?;
let name = func.get("name")?.as_str()?.to_string();
let arguments = func.get("arguments")?.as_str()?.to_string();
Some(ToolCallRequest {
id,
name,
arguments,
})
})
.collect::<Vec<_>>()
})
.unwrap_or_default();
let finish_reason = match choice
.get("finish_reason")
.and_then(|f| f.as_str())
.unwrap_or("stop")
{
"tool_calls" => FinishReason::ToolCalls,
"length" => FinishReason::MaxTokens,
"content_filter" => FinishReason::ContentFilter,
_ => {
if tool_calls.is_empty() {
FinishReason::Stop
} else {
FinishReason::ToolCalls
}
}
};
let usage = parse_token_usage(resp.get("usage"), false)?;
let actual_model = resp
.get("model")
.and_then(|m| m.as_str())
.unwrap_or(model)
.to_string();
Ok(InferenceResponse {
content,
tool_calls,
finish_reason,
usage,
model: actual_model,
})
}
fn parse_anthropic_response(
&self,
resp: &serde_json::Value,
model: &str,
) -> Result<InferenceResponse, InferenceError> {
let content_blocks = resp
.get("content")
.and_then(|c| c.as_array())
.ok_or_else(|| InferenceError::ParseError("No content in response".into()))?;
let mut text_content = String::new();
let mut tool_calls = Vec::new();
for block in content_blocks {
match block.get("type").and_then(|t| t.as_str()) {
Some("text") => {
if let Some(text) = block.get("text").and_then(|t| t.as_str()) {
if !text_content.is_empty() {
text_content.push('\n');
}
text_content.push_str(text);
}
}
Some("tool_use") => {
if let (Some(id), Some(name), Some(input)) = (
block.get("id").and_then(|v| v.as_str()),
block.get("name").and_then(|v| v.as_str()),
block.get("input"),
) {
tool_calls.push(ToolCallRequest {
id: id.to_string(),
name: name.to_string(),
arguments: serde_json::to_string(input).unwrap_or_default(),
});
}
}
_ => {}
}
}
let stop_reason = resp
.get("stop_reason")
.and_then(|s| s.as_str())
.unwrap_or("end_turn");
let finish_reason = map_anthropic_stop_reason(stop_reason, !tool_calls.is_empty());
if finish_reason == FinishReason::Refusal {
tracing::warn!(
"Anthropic response was a refusal (stop_reason=refusal); returning FinishReason::Refusal"
);
} else if text_content.is_empty() && tool_calls.is_empty() {
tracing::warn!(
"Anthropic response produced no text and no tool calls (stop_reason={}); \
the turn made no progress — likely a thinking-only turn or an unparsed block type",
stop_reason
);
}
let usage = parse_token_usage(resp.get("usage"), true)?;
let actual_model = resp
.get("model")
.and_then(|m| m.as_str())
.unwrap_or(model)
.to_string();
Ok(InferenceResponse {
content: text_content,
tool_calls,
finish_reason,
usage,
model: actual_model,
})
}
}
#[async_trait]
impl InferenceProvider for CloudInferenceProvider {
fn input_token_reservation(
&self,
conversation: &Conversation,
options: &InferenceOptions,
) -> Result<u32, InferenceError> {
let body = if matches!(self.client.provider(), LlmProvider::Anthropic) {
self.build_anthropic_body(conversation, options)
} else {
self.build_openai_body(conversation, options)
};
let bytes = serde_json::to_vec(&body)
.map_err(|error| InferenceError::InvalidRequest(error.to_string()))?;
input_reservation_from_bytes(bytes.len())
}
async fn complete(
&self,
conversation: &Conversation,
options: &InferenceOptions,
) -> Result<InferenceResponse, InferenceError> {
let is_anthropic = matches!(self.client.provider(), LlmProvider::Anthropic);
let model = options
.model
.as_deref()
.unwrap_or_else(|| self.client.model());
#[cfg(feature = "bedrock")]
if matches!(self.client.provider(), LlmProvider::Bedrock) {
let (system_opt, messages) = conversation.to_anthropic_messages();
let system = system_opt.as_deref().unwrap_or("");
let tools: Vec<serde_json::Value> = options
.tool_definitions
.iter()
.map(|td| {
serde_json::json!({
"name": td.name,
"description": td.description,
"input_schema": td.parameters,
})
})
.collect();
let resp_json = self
.client
.bedrock_converse(
system,
&messages,
&tools,
options.temperature,
options.max_tokens,
)
.await
.map_err(|e| InferenceError::Provider(format!("Bedrock Converse error: {e}")))?;
return self.parse_anthropic_response(&resp_json, model);
}
let body = if is_anthropic {
self.build_anthropic_body(conversation, options)
} else {
self.build_openai_body(conversation, options)
};
let http_client = reqwest::Client::builder()
.timeout(std::time::Duration::from_secs(120))
.build()
.map_err(|e| InferenceError::Provider(format!("HTTP client error: {}", e)))?;
let base = self.client.base_url();
let api_key = self.client.api_key();
let (url, request_builder) = if is_anthropic {
let url = format!("{}/messages", base);
let rb = http_client
.post(&url)
.header("x-api-key", api_key)
.header("anthropic-version", "2023-06-01")
.header("content-type", "application/json")
.json(&body);
(url, rb)
} else {
let url = format!("{}/chat/completions", base);
let mut rb = http_client
.post(&url)
.header("authorization", format!("Bearer {}", api_key))
.header("content-type", "application/json");
if matches!(self.client.provider(), LlmProvider::OpenRouter) {
for (k, v) in crate::http_input::llm_client::openrouter_attribution_headers() {
rb = rb.header(k, v);
}
}
let rb = rb.json(&body);
(url, rb)
};
tracing::debug!(
"Cloud inference: provider={} model={} url={}",
self.provider_name(),
model,
url
);
tracing::debug!(
"Cloud request fingerprint: tool_choice={} tools={} system_chars={} msg_count={}",
body.get("tool_choice")
.map(|v| v.to_string())
.unwrap_or_else(|| "<absent>".into()),
body.get("tools")
.and_then(|v| v.as_array())
.map(|a| a.len())
.unwrap_or(0),
body.get("system")
.and_then(|v| v.as_str())
.map(|s| s.len())
.unwrap_or(0),
body.get("messages")
.and_then(|v| v.as_array())
.map(|a| a.len())
.unwrap_or(0),
);
let start = std::time::Instant::now();
let response = request_builder.send().await.map_err(|e| {
if e.is_timeout() {
InferenceError::Timeout(std::time::Duration::from_secs(120))
} else {
InferenceError::Provider(format!("Request failed: {}", e))
}
})?;
let status = response.status();
if status.as_u16() == 429 {
let retry_after = response
.headers()
.get("retry-after")
.and_then(|v| v.to_str().ok())
.and_then(|v| v.parse::<u64>().ok())
.unwrap_or(1000);
return Err(InferenceError::RateLimited {
retry_after_ms: retry_after * 1000,
});
}
if !status.is_success() {
let error_text = response
.text()
.await
.unwrap_or_else(|_| "Unknown error".into());
tracing::warn!(
"Cloud API non-success: status={} body={}",
status,
error_text.chars().take(400).collect::<String>()
);
return Err(InferenceError::Provider(format!(
"API error ({}): {}",
status, error_text
)));
}
let resp_json: serde_json::Value = response
.json()
.await
.map_err(|e| InferenceError::ParseError(format!("JSON parse error: {}", e)))?;
let latency = start.elapsed();
tracing::debug!("Cloud inference completed in {:?}", latency);
if is_anthropic {
let stop = resp_json
.get("stop_reason")
.and_then(|v| v.as_str())
.unwrap_or("<absent>");
let content_types: Vec<&str> = resp_json
.get("content")
.and_then(|v| v.as_array())
.map(|arr| {
arr.iter()
.filter_map(|c| c.get("type").and_then(|t| t.as_str()))
.collect()
})
.unwrap_or_default();
tracing::debug!(
"Cloud response fingerprint: stop_reason={} content_types={:?}",
stop,
content_types
);
}
if is_anthropic {
self.parse_anthropic_response(&resp_json, model)
} else {
self.parse_openai_response(&resp_json, model)
}
}
fn provider_name(&self) -> &str {
match self.client.provider() {
LlmProvider::OpenRouter => "openrouter",
LlmProvider::OpenAI => "openai",
LlmProvider::Anthropic => "anthropic",
#[cfg(feature = "bedrock")]
LlmProvider::Bedrock => "bedrock",
}
}
fn default_model(&self) -> &str {
self.client.model()
}
fn configuration_identity(&self) -> Option<String> {
#[cfg(feature = "bedrock")]
let region = self.client.region();
#[cfg(not(feature = "bedrock"))]
let region = "";
crate::reasoning::prepared::digest_json(&serde_json::json!([
self.provider_name(),
self.client.model(),
self.client.base_url(),
region
]))
.ok()
}
fn supports_native_tools(&self) -> bool {
true
}
fn supports_structured_output(&self) -> bool {
true
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn token_accounting_includes_cached_input_and_rejects_counter_truncation() {
let cached = serde_json::json!({"input_tokens":10,"output_tokens":5,"cache_creation_input_tokens":100,"cache_read_input_tokens":1000});
let usage = parse_token_usage(Some(&cached), true).unwrap();
assert_eq!(usage.prompt_tokens, 1110);
assert_eq!(usage.total_tokens, 1115);
for value in [
serde_json::json!({"prompt_tokens":4294967296u64,"completion_tokens":1,"total_tokens":4294967297u64}),
serde_json::json!({"prompt_tokens":10,"completion_tokens":5,"total_tokens":3}),
serde_json::json!({"prompt_tokens":-1,"completion_tokens":5,"total_tokens":3}),
] {
assert!(parse_token_usage(Some(&value), false).is_err());
}
let overflow = serde_json::json!({"input_tokens":4294967295u64,"output_tokens":1});
assert!(parse_token_usage(Some(&overflow), true).is_err());
let missing = serde_json::json!({"prompt_tokens":10});
assert_eq!(
parse_token_usage(Some(&missing), false)
.unwrap()
.total_tokens,
0
);
}
use crate::reasoning::conversation::{ConversationMessage, ToolCall};
use serial_test::serial;
#[test]
fn test_build_openai_body_basic() {
let openai_response = serde_json::json!({
"choices": [{
"message": {
"role": "assistant",
"content": "Hello!",
},
"finish_reason": "stop"
}],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 5,
"total_tokens": 15,
},
"model": "gpt-4o"
});
let choice = openai_response["choices"][0].clone();
let content = choice["message"]["content"].as_str().unwrap();
assert_eq!(content, "Hello!");
}
#[test]
fn test_parse_openai_response_with_tools() {
let resp = serde_json::json!({
"choices": [{
"message": {
"role": "assistant",
"content": null,
"tool_calls": [{
"id": "call_abc123",
"type": "function",
"function": {
"name": "web_search",
"arguments": "{\"query\": \"rust crates\"}"
}
}]
},
"finish_reason": "tool_calls"
}],
"usage": {
"prompt_tokens": 20,
"completion_tokens": 10,
"total_tokens": 30,
},
"model": "gpt-4o"
});
let tool_calls = resp["choices"][0]["message"]["tool_calls"]
.as_array()
.unwrap();
assert_eq!(tool_calls.len(), 1);
assert_eq!(tool_calls[0]["function"]["name"], "web_search");
}
#[test]
fn test_parse_anthropic_response() {
let resp = serde_json::json!({
"content": [
{"type": "text", "text": "I'll search for that."},
{
"type": "tool_use",
"id": "toolu_123",
"name": "web_search",
"input": {"query": "rust crates"}
}
],
"stop_reason": "tool_use",
"usage": {
"input_tokens": 15,
"output_tokens": 20,
},
"model": "claude-sonnet-4-5-20250514"
});
let content = resp["content"].as_array().unwrap();
assert_eq!(content.len(), 2);
assert_eq!(content[0]["type"], "text");
assert_eq!(content[1]["type"], "tool_use");
assert_eq!(content[1]["name"], "web_search");
}
#[test]
fn test_map_anthropic_stop_reason() {
use super::map_anthropic_stop_reason;
assert_eq!(
map_anthropic_stop_reason("refusal", false),
FinishReason::Refusal
);
assert_eq!(
map_anthropic_stop_reason("refusal", true),
FinishReason::Refusal
);
assert_eq!(
map_anthropic_stop_reason("tool_use", false),
FinishReason::ToolCalls
);
assert_eq!(
map_anthropic_stop_reason("max_tokens", false),
FinishReason::MaxTokens
);
assert_eq!(
map_anthropic_stop_reason("end_turn", true),
FinishReason::ToolCalls
);
assert_eq!(
map_anthropic_stop_reason("end_turn", false),
FinishReason::Stop
);
assert_eq!(
map_anthropic_stop_reason("pause_turn", false),
FinishReason::Stop
);
}
#[test]
fn test_conversation_to_openai_format() {
let mut conv = Conversation::with_system("sys");
conv.push(ConversationMessage::user("hello"));
conv.push(ConversationMessage::assistant_tool_calls(vec![ToolCall {
id: "tc1".into(),
name: "search".into(),
arguments: r#"{"q":"test"}"#.into(),
}]));
conv.push(ConversationMessage::tool_result("tc1", "search", "result"));
let msgs = conv.to_openai_messages();
assert_eq!(msgs.len(), 4);
assert_eq!(msgs[0]["role"], "system");
assert_eq!(msgs[2]["tool_calls"][0]["function"]["name"], "search");
assert_eq!(msgs[3]["tool_call_id"], "tc1");
}
#[serial]
#[test]
fn test_build_anthropic_body_cache_control_and_extra() {
use crate::reasoning::inference::{InferenceOptions, ToolDefinition};
std::env::set_var("ANTHROPIC_API_KEY", "test-key-not-real");
for k in ["OPENROUTER_API_KEY", "OPENAI_API_KEY", "BEDROCK_MODEL_ID"] {
std::env::remove_var(k);
}
let provider = CloudInferenceProvider::from_env()
.expect("CloudInferenceProvider should resolve via ANTHROPIC_API_KEY");
let conv = Conversation::with_system("You are a helpful assistant.");
let mut options = InferenceOptions {
tool_definitions: vec![ToolDefinition {
name: "search".into(),
description: "Search the web".into(),
parameters: serde_json::json!({"type": "object", "properties": {}}),
}],
..Default::default()
};
options.extra.insert(
"output_config".into(),
serde_json::json!({"effort": "high"}),
);
let body = provider.build_anthropic_body(&conv, &options);
let system = body["system"]
.as_array()
.expect("system should be a content-block array carrying cache_control");
assert_eq!(system.len(), 1);
assert_eq!(system[0]["type"], "text");
assert_eq!(system[0]["text"], "You are a helpful assistant.");
assert_eq!(system[0]["cache_control"]["type"], "ephemeral");
let tools = body["tools"].as_array().expect("tools array");
assert!(tools[0].get("cache_control").is_none());
assert_eq!(body["output_config"]["effort"], "high");
let conv_no_system = Conversation::new();
let options_no_system = InferenceOptions {
tool_definitions: vec![
ToolDefinition {
name: "first_tool".into(),
description: "d1".into(),
parameters: serde_json::json!({"type": "object", "properties": {}}),
},
ToolDefinition {
name: "last_tool".into(),
description: "d2".into(),
parameters: serde_json::json!({"type": "object", "properties": {}}),
},
],
..Default::default()
};
let body_no_system = provider.build_anthropic_body(&conv_no_system, &options_no_system);
assert!(body_no_system.get("system").is_none());
let tools_no_system = body_no_system["tools"].as_array().expect("tools array");
assert_eq!(tools_no_system.len(), 2);
assert!(tools_no_system[0].get("cache_control").is_none());
assert_eq!(tools_no_system[1]["cache_control"]["type"], "ephemeral");
std::env::remove_var("ANTHROPIC_API_KEY");
}
#[cfg(feature = "bedrock")]
#[serial]
#[test]
fn test_cloud_provider_name_bedrock() {
std::env::set_var(
"BEDROCK_MODEL_ID",
"anthropic.claude-3-5-sonnet-20241022-v2:0",
);
std::env::set_var("AWS_REGION", "us-east-1");
for k in ["OPENROUTER_API_KEY", "OPENAI_API_KEY", "ANTHROPIC_API_KEY"] {
std::env::remove_var(k);
}
let provider = CloudInferenceProvider::from_env()
.expect("CloudInferenceProvider should resolve via BEDROCK_MODEL_ID");
assert_eq!(provider.provider_name(), "bedrock");
std::env::remove_var("BEDROCK_MODEL_ID");
std::env::remove_var("AWS_REGION");
}
}