use std::collections::HashMap;
use std::sync::Arc;
use reqwest::Client;
use serde_json::Value;
use tracing::info;
use super::api_format::ApiFormat;
use super::client_registry::ClientRegistry;
use super::error_classify::AttemptDecision;
use super::errors::{ConduitError, ErrorKind};
use super::message_norm::normalize_messages_for_api;
use super::provider_runtime::ProviderRuntime;
use super::response_parser::TransportResponse;
use crate::clients::parsing::TransportKind;
#[derive(Debug, Clone)]
pub enum ApiKeyConfig {
None,
Single(String),
PerProvider(HashMap<String, String>),
}
#[derive(Debug, Clone)]
pub enum ApiBaseConfig {
None,
Single(String),
PerProvider(HashMap<String, String>),
}
pub struct LLMCore {
provider: String,
model: String,
fallback_models: Vec<String>,
max_retries: u32,
api_key: ApiKeyConfig,
api_base: ApiBaseConfig,
client_registry: ClientRegistry,
api_format: ApiFormat,
verbose: u32,
#[allow(clippy::type_complexity)]
error_classifier: Option<Box<dyn Fn(&ConduitError) -> Option<ErrorKind> + Send + Sync>>,
}
fn split_model_id<'a>(
model: &'a str,
message: &'static str,
) -> Result<(&'a str, &'a str), ConduitError> {
let (provider, model_id) = model
.split_once(':')
.ok_or_else(|| ConduitError::new(ErrorKind::InvalidInput, message))?;
if provider.is_empty() || model_id.is_empty() {
return Err(ConduitError::new(ErrorKind::InvalidInput, message));
}
Ok((provider, model_id))
}
impl LLMCore {
#[allow(clippy::too_many_arguments)]
pub fn new(
provider: String,
model: String,
fallback_models: Vec<String>,
max_retries: u32,
api_key: ApiKeyConfig,
api_base: ApiBaseConfig,
api_format: impl Into<ApiFormat>,
verbose: u32,
) -> Self {
Self {
provider,
model,
fallback_models,
max_retries,
api_key,
api_base,
client_registry: ClientRegistry::new(),
api_format: api_format.into(),
verbose,
error_classifier: None,
}
}
pub fn with_error_classifier(
mut self,
classifier: impl Fn(&ConduitError) -> Option<ErrorKind> + Send + Sync + 'static,
) -> Self {
self.error_classifier = Some(Box::new(classifier));
self
}
pub fn provider(&self) -> &str {
&self.provider
}
pub fn model(&self) -> &str {
&self.model
}
pub fn fallback_models(&self) -> &[String] {
&self.fallback_models
}
pub fn max_retries(&self) -> u32 {
self.max_retries
}
pub fn api_key_config(&self) -> &ApiKeyConfig {
&self.api_key
}
pub fn api_base_config(&self) -> &ApiBaseConfig {
&self.api_base
}
pub fn api_format(&self) -> ApiFormat {
self.api_format
}
pub fn verbose(&self) -> u32 {
self.verbose
}
pub fn max_attempts(&self) -> u32 {
1u32.max(1 + self.max_retries)
}
fn retry_attempts(&self) -> std::ops::Range<u32> {
0..self.max_attempts()
}
pub(crate) fn custom_classify(&self, error: &ConduitError) -> Option<ErrorKind> {
if let Some(ref classifier) = self.error_classifier
&& let Some(kind) = classifier(error)
{
return Some(kind);
}
None
}
pub fn resolve_model_provider(
model: &str,
provider: Option<&str>,
) -> Result<(String, String), ConduitError> {
if let Some(p) = provider {
if model.contains(':') {
return Err(ConduitError::new(
ErrorKind::InvalidInput,
"When provider is specified, model must not include a provider prefix.",
));
}
return Ok((p.to_owned(), model.to_owned()));
}
let (prov, mdl) = split_model_id(model, "Model must be in 'provider:model' format.")?;
Ok((prov.to_owned(), mdl.to_owned()))
}
pub fn resolve_fallback(&self, model: &str) -> Result<(String, String), ConduitError> {
if model.contains(':') {
let (prov, mdl) =
split_model_id(model, "Fallback models must be in 'provider:model' format.")?;
return Ok((prov.to_owned(), mdl.to_owned()));
}
if !self.provider.is_empty() {
return Ok((self.provider.clone(), model.to_owned()));
}
Err(ConduitError::new(
ErrorKind::InvalidInput,
"Fallback models must include provider or LLM must be initialized with a provider.",
))
}
pub fn model_candidates(
&self,
override_model: Option<&str>,
override_provider: Option<&str>,
) -> Result<Vec<(String, String)>, ConduitError> {
if let Some(om) = override_model {
let (p, m) = Self::resolve_model_provider(om, override_provider)?;
return Ok(vec![(p, m)]);
}
let mut candidates = vec![(self.provider.clone(), self.model.clone())];
for fallback in &self.fallback_models {
candidates.push(self.resolve_fallback(fallback)?);
}
Ok(candidates)
}
pub fn resolve_api_key(&self, provider: &str) -> Option<String> {
match &self.api_key {
ApiKeyConfig::None => None,
ApiKeyConfig::Single(key) => Some(key.clone()),
ApiKeyConfig::PerProvider(map) => map.get(provider).cloned(),
}
}
pub fn resolve_api_base(&self, provider: &str) -> Option<String> {
match &self.api_base {
ApiBaseConfig::None => None,
ApiBaseConfig::Single(base) => Some(base.clone()),
ApiBaseConfig::PerProvider(map) => map.get(provider).cloned(),
}
}
pub fn get_client(&mut self, provider: &str) -> Arc<Client> {
let api_key = self.resolve_api_key(provider);
let api_base = self.resolve_api_base(provider);
self.client_registry
.get_or_create(provider, api_key.as_deref(), api_base.as_deref())
}
#[allow(clippy::too_many_arguments)]
pub async fn run_chat<T, F>(
&mut self,
messages_payload: Vec<Value>,
tools_payload: Option<Vec<Value>>,
model: Option<&str>,
provider: Option<&str>,
max_tokens: Option<u32>,
stream: bool,
reasoning_effort: Option<Value>,
kwargs: serde_json::Map<String, Value>,
on_response: F,
) -> Result<T, ConduitError>
where
F: Fn(TransportResponse, &str, &str, u32) -> Result<T, Option<ConduitError>>,
{
let candidates = self.model_candidates(model, provider)?;
let mut last_error: Option<ConduitError> = None;
for (provider_name, model_id) in &candidates {
let client = self.get_client(provider_name);
let api_base = self.resolve_api_base(provider_name);
let api_key = self.resolve_api_key(provider_name);
let runtime = ProviderRuntime::new(
provider_name,
model_id,
api_key.as_deref(),
api_base.as_deref(),
self.api_format,
);
for attempt in self.retry_attempts() {
let transport = match runtime.selected_transport(
tools_payload.as_deref(),
false, None,
) {
Ok(t) => t,
Err(e) => {
last_error = Some(e);
break;
}
};
let resolved_api_base = runtime.resolved_api_base();
let normalized_messages =
normalize_messages_for_api(messages_payload.clone(), transport);
let request = super::request_builder::TransportCallRequest {
client: Arc::clone(&client),
provider_name: provider_name.clone(),
model_id: model_id.clone(),
api_base: Some(resolved_api_base.clone()),
messages_payload: normalized_messages,
tools_payload: tools_payload.clone(),
max_tokens,
stream,
reasoning_effort: reasoning_effort.clone(),
kwargs: kwargs.clone(),
is_anthropic_oauth: runtime.is_anthropic_oauth(),
};
let url = Self::build_request_url(&resolved_api_base, transport);
let body = match transport {
TransportKind::Responses => Self::build_responses_body(&request)?,
TransportKind::Messages => Self::build_messages_body(&request)?,
TransportKind::Completion => {
Self::build_completion_body(&request, provider_name)?
}
};
info!(
target: "eli_trace",
provider = %provider_name,
model = %model_id,
transport = ?transport,
stream = body.get("stream").and_then(|v| v.as_bool()).unwrap_or(false),
request_body = %body,
"llm.request"
);
let http_result = client.post(&url).json(&body).send().await;
match http_result {
Ok(resp) => {
let status = resp.status();
if !status.is_success() {
let error_body = resp.text().await.unwrap_or_default();
let kind = Self::classify_http_status(status.as_u16())
.unwrap_or(ErrorKind::Provider);
let error = ConduitError::new(
kind,
format!(
"{}:{}: HTTP {} - {}",
provider_name, model_id, status, error_body
),
);
let outcome =
self.handle_attempt_error(error, provider_name, model_id, attempt);
last_error = Some(outcome.error);
if outcome.decision == AttemptDecision::RetrySameModel {
continue;
}
break;
}
let body_forced_stream = !stream
&& body
.get("stream")
.and_then(|v| v.as_bool())
.unwrap_or(false);
let payload: Value = if body_forced_stream {
match Self::collect_sse_response(resp, transport).await {
Ok(v) => v,
Err(e) => {
let outcome = self.handle_attempt_error(
e,
provider_name,
model_id,
attempt,
);
last_error = Some(outcome.error);
if outcome.decision == AttemptDecision::RetrySameModel {
continue;
}
break;
}
}
} else {
match resp.json().await {
Ok(v) => v,
Err(e) => {
let error = ConduitError::new(
ErrorKind::Provider,
format!(
"{}:{}: failed to parse response: {}",
provider_name, model_id, e
),
);
let outcome = self.handle_attempt_error(
error,
provider_name,
model_id,
attempt,
);
last_error = Some(outcome.error);
if outcome.decision == AttemptDecision::RetrySameModel {
continue;
}
break;
}
}
};
info!(
target: "eli_trace",
provider = %provider_name,
model = %model_id,
transport = ?transport,
response_payload = %payload,
"llm.response"
);
let transport_response = TransportResponse { transport, payload };
match on_response(transport_response, provider_name, model_id, attempt) {
Ok(result) => return Ok(result),
Err(Some(e)) => {
last_error = Some(e);
break;
}
Err(None) => {
continue;
}
}
}
Err(e) => {
let kind = if e.is_timeout() {
ErrorKind::Temporary
} else if e.is_connect() {
ErrorKind::Provider
} else {
ErrorKind::Unknown
};
let error = ConduitError::new(
kind,
format!("{}:{}: {}", provider_name, model_id, e),
);
let outcome =
self.handle_attempt_error(error, provider_name, model_id, attempt);
last_error = Some(outcome.error);
if outcome.decision == AttemptDecision::RetrySameModel {
continue;
}
break;
}
}
}
}
Err(last_error.unwrap_or_else(|| {
ConduitError::new(ErrorKind::Temporary, "LLM call failed after retries")
}))
}
#[allow(clippy::too_many_arguments)]
pub async fn run_chat_stream(
&mut self,
messages_payload: Vec<Value>,
tools_payload: Option<Vec<Value>>,
model: Option<&str>,
provider: Option<&str>,
max_tokens: Option<u32>,
reasoning_effort: Option<Value>,
kwargs: serde_json::Map<String, Value>,
) -> Result<(reqwest::Response, TransportKind, String, String), ConduitError> {
let candidates = self.model_candidates(model, provider)?;
let mut last_error: Option<ConduitError> = None;
for (provider_name, model_id) in &candidates {
let client = self.get_client(provider_name);
let api_base = self.resolve_api_base(provider_name);
let api_key = self.resolve_api_key(provider_name);
let runtime = ProviderRuntime::new(
provider_name,
model_id,
api_key.as_deref(),
api_base.as_deref(),
self.api_format,
);
for attempt in self.retry_attempts() {
let transport =
match runtime.selected_transport(tools_payload.as_deref(), false, None) {
Ok(t) => t,
Err(e) => {
last_error = Some(e);
break;
}
};
let resolved_api_base = runtime.resolved_api_base();
let normalized_messages =
normalize_messages_for_api(messages_payload.clone(), transport);
let request = super::request_builder::TransportCallRequest {
client: Arc::clone(&client),
provider_name: provider_name.clone(),
model_id: model_id.clone(),
api_base: Some(resolved_api_base.clone()),
messages_payload: normalized_messages,
tools_payload: tools_payload.clone(),
max_tokens,
stream: true,
reasoning_effort: reasoning_effort.clone(),
kwargs: kwargs.clone(),
is_anthropic_oauth: runtime.is_anthropic_oauth(),
};
let url = Self::build_request_url(&resolved_api_base, transport);
let body = match transport {
TransportKind::Responses => Self::build_responses_body(&request)?,
TransportKind::Messages => Self::build_messages_body(&request)?,
TransportKind::Completion => {
Self::build_completion_body(&request, provider_name)?
}
};
let http_result = client.post(&url).json(&body).send().await;
match http_result {
Ok(resp) => {
let status = resp.status();
if !status.is_success() {
let error_body = resp.text().await.unwrap_or_default();
let kind = Self::classify_http_status(status.as_u16())
.unwrap_or(ErrorKind::Provider);
let error = ConduitError::new(
kind,
format!(
"{}:{}: HTTP {} - {}",
provider_name, model_id, status, error_body
),
);
let outcome =
self.handle_attempt_error(error, provider_name, model_id, attempt);
last_error = Some(outcome.error);
if outcome.decision == AttemptDecision::RetrySameModel {
continue;
}
break;
}
return Ok((resp, transport, provider_name.clone(), model_id.clone()));
}
Err(e) => {
let kind = if e.is_timeout() {
ErrorKind::Temporary
} else if e.is_connect() {
ErrorKind::Provider
} else {
ErrorKind::Unknown
};
let error = ConduitError::new(
kind,
format!("{}:{}: {}", provider_name, model_id, e),
);
let outcome =
self.handle_attempt_error(error, provider_name, model_id, attempt);
last_error = Some(outcome.error);
if outcome.decision == AttemptDecision::RetrySameModel {
continue;
}
break;
}
}
}
}
Err(last_error.unwrap_or_else(|| {
ConduitError::new(
ErrorKind::Temporary,
"LLM streaming call failed after retries",
)
}))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::anthropic_messages;
use crate::core::error_classify::classify_by_text_signature;
use crate::core::message_norm::{enforce_anthropic_message_rules, prune_orphan_tool_messages};
use serde_json::json;
use super::super::request_builder::TransportCallRequest;
#[test]
fn test_resolve_model_provider() {
let (p, m) = LLMCore::resolve_model_provider("openai:gpt-4", None).unwrap();
assert_eq!(p, "openai");
assert_eq!(m, "gpt-4");
}
#[test]
fn test_resolve_model_provider_with_override() {
let (p, m) = LLMCore::resolve_model_provider("gpt-4", Some("openai")).unwrap();
assert_eq!(p, "openai");
assert_eq!(m, "gpt-4");
}
#[test]
fn test_resolve_model_provider_error() {
let result = LLMCore::resolve_model_provider("gpt-4", None);
assert!(result.is_err());
}
#[test]
fn test_classify_http_status() {
assert_eq!(LLMCore::classify_http_status(401), Some(ErrorKind::Config));
assert_eq!(
LLMCore::classify_http_status(429),
Some(ErrorKind::Temporary)
);
assert_eq!(
LLMCore::classify_http_status(500),
Some(ErrorKind::Provider)
);
assert_eq!(LLMCore::classify_http_status(200), None);
}
#[test]
fn test_split_messages_for_responses() {
let messages = vec![
json!({"role": "system", "content": "You are helpful."}),
json!({"role": "user", "content": "Hello"}),
];
let (instructions, items) = LLMCore::split_messages_for_responses(&messages);
assert_eq!(instructions.unwrap(), "You are helpful.");
assert_eq!(items.len(), 1);
assert_eq!(items[0]["role"], "user");
}
#[test]
fn test_convert_tools_for_responses() {
let tools = vec![json!({
"type": "function",
"function": {
"name": "greet",
"description": "Say hello",
"parameters": {"type": "object"}
}
})];
let converted = LLMCore::convert_tools_for_responses(Some(&tools)).unwrap();
assert_eq!(converted.len(), 1);
assert_eq!(converted[0]["name"], "greet");
assert!(converted[0].get("function").is_none());
}
#[test]
fn test_normalize_anthropic_messages_keeps_multiple_tool_results() {
let messages = vec![
json!({
"role": "user",
"content": "find latest papers"
}),
json!({
"role": "assistant",
"content": [
{"type": "tool_use", "id": "toolu_1", "name": "bash", "input": {"cmd": "pwd"}},
{"type": "tool_use", "id": "toolu_2", "name": "skill", "input": {"name": "active-research"}}
]
}),
json!({
"role": "user",
"content": [
{"type": "tool_result", "tool_use_id": "toolu_1", "content": "ok-1"}
]
}),
json!({
"role": "user",
"content": [
{"type": "tool_result", "tool_use_id": "toolu_2", "content": "ok-2"}
]
}),
];
let normalized = anthropic_messages::normalize_messages(messages);
assert_eq!(normalized.len(), 3);
assert_eq!(normalized[0]["role"], "user");
assert_eq!(normalized[1]["role"], "assistant");
assert_eq!(normalized[2]["role"], "user");
let content = normalized[2]["content"].as_array().unwrap();
assert_eq!(content.len(), 2);
assert_eq!(content[0]["tool_use_id"], "toolu_1");
assert_eq!(content[1]["tool_use_id"], "toolu_2");
}
#[test]
fn test_normalize_anthropic_messages_merges_text_with_tool_results() {
let messages = vec![
json!({
"role": "user",
"content": [
{"type": "tool_result", "tool_use_id": "toolu_1", "content": "ok"}
]
}),
json!({
"role": "user",
"content": "continue with a summary"
}),
];
let normalized = anthropic_messages::normalize_messages(messages);
assert_eq!(normalized.len(), 1);
let content = normalized[0]["content"].as_array().unwrap();
assert_eq!(content.len(), 2);
assert_eq!(content[0]["type"], "tool_result");
assert_eq!(content[1]["type"], "text");
assert_eq!(content[1]["text"], "continue with a summary");
}
#[test]
fn test_build_messages_body_keeps_tool_results_immediately_after_tool_use() {
let request = TransportCallRequest {
client: Arc::new(reqwest::Client::new()),
provider_name: "anthropic".to_owned(),
model_id: "claude-test".to_owned(),
api_base: None,
messages_payload: vec![
json!({"role": "system", "content": "system rules"}),
json!({"role": "user", "content": "find latest papers"}),
json!({
"role": "assistant",
"tool_calls": [
{"id": "toolu_1", "type": "function", "function": {"name": "bash", "arguments": "{\"cmd\":\"pwd\"}"}},
{"id": "toolu_2", "type": "function", "function": {"name": "skill", "arguments": "{\"name\":\"active-research\"}"}}
]
}),
json!({"role": "tool", "tool_call_id": "toolu_1", "content": "ok-1"}),
json!({"role": "tool", "tool_call_id": "toolu_2", "content": "ok-2"}),
json!({"role": "user", "content": "\u{7ee7}\u{7eed}"}),
],
tools_payload: None,
max_tokens: Some(512),
stream: false,
reasoning_effort: None,
kwargs: serde_json::Map::new(),
is_anthropic_oauth: false,
};
let body = LLMCore::build_messages_body(&request).unwrap();
let messages = body["messages"].as_array().unwrap();
assert_eq!(messages.len(), 3);
assert_eq!(messages[1]["role"], "assistant");
assert_eq!(messages[2]["role"], "user");
let user_blocks = messages[2]["content"].as_array().unwrap();
assert_eq!(user_blocks[0]["type"], "tool_result");
assert_eq!(user_blocks[0]["tool_use_id"], "toolu_1");
assert_eq!(user_blocks[1]["type"], "tool_result");
assert_eq!(user_blocks[1]["tool_use_id"], "toolu_2");
assert_eq!(user_blocks[2]["type"], "text");
assert_eq!(user_blocks[2]["text"], "\u{7ee7}\u{7eed}");
}
#[test]
fn test_max_attempts() {
let core = LLMCore::new(
"openai".into(),
"gpt-4".into(),
vec![],
2,
ApiKeyConfig::None,
ApiBaseConfig::None,
TransportKind::Completion,
0,
);
assert_eq!(core.max_attempts(), 3);
}
#[test]
fn test_build_request_url() {
assert_eq!(
LLMCore::build_request_url("https://api.openai.com/v1", TransportKind::Completion),
"https://api.openai.com/v1/chat/completions"
);
assert_eq!(
LLMCore::build_request_url("https://api.openai.com/v1", TransportKind::Responses),
"https://api.openai.com/v1/responses"
);
assert_eq!(
LLMCore::build_request_url("https://api.anthropic.com/v1", TransportKind::Messages),
"https://api.anthropic.com/v1/messages"
);
}
#[test]
fn test_build_error_payload_basic() {
let core = LLMCore::new(
"openai".into(),
"gpt-4".into(),
vec![],
2,
ApiKeyConfig::None,
ApiBaseConfig::None,
TransportKind::Completion,
0,
);
let error = ConduitError::new(ErrorKind::Provider, "server error");
let payload = core.build_error_payload(&error, "openai", "gpt-4", 1, None);
assert_eq!(payload.kind, ErrorKind::Provider);
assert_eq!(payload.message, "server error");
let details = payload.details.unwrap();
assert_eq!(details["provider"], "openai");
assert_eq!(details["model"], "gpt-4");
assert_eq!(details["attempt"], 2); assert_eq!(details["max_attempts"], 3); assert!(details.get("http_status").is_none());
}
#[test]
fn test_build_error_payload_with_http_status() {
let core = LLMCore::new(
"openai".into(),
"gpt-4".into(),
vec![],
3,
ApiKeyConfig::None,
ApiBaseConfig::None,
TransportKind::Completion,
0,
);
let error = ConduitError::new(ErrorKind::Temporary, "rate limited");
let payload = core.build_error_payload(&error, "openai", "gpt-4", 0, Some(429));
assert_eq!(payload.kind, ErrorKind::Temporary);
let details = payload.details.unwrap();
assert_eq!(details["http_status"], 429);
assert_eq!(details["attempt"], 1);
assert_eq!(details["max_attempts"], 4);
}
#[test]
fn test_build_error_payload_different_provider() {
let core = LLMCore::new(
"anthropic".into(),
"claude-3".into(),
vec![],
1,
ApiKeyConfig::None,
ApiBaseConfig::None,
TransportKind::Completion,
0,
);
let error = ConduitError::new(ErrorKind::Config, "auth failed");
let payload = core.build_error_payload(&error, "anthropic", "claude-3", 0, Some(401));
assert_eq!(payload.kind, ErrorKind::Config);
let details = payload.details.unwrap();
assert_eq!(details["provider"], "anthropic");
assert_eq!(details["model"], "claude-3");
assert_eq!(details["http_status"], 401);
}
#[test]
fn test_api_key_config_accessor() {
let core = LLMCore::new(
"openai".into(),
"gpt-4".into(),
vec![],
3,
ApiKeyConfig::Single("my-key".into()),
ApiBaseConfig::None,
TransportKind::Completion,
0,
);
match core.api_key_config() {
ApiKeyConfig::Single(key) => assert_eq!(key, "my-key"),
_ => panic!("Expected Single key config"),
}
}
#[test]
fn test_api_base_config_accessor() {
let core = LLMCore::new(
"openai".into(),
"gpt-4".into(),
vec![],
3,
ApiKeyConfig::None,
ApiBaseConfig::Single("https://custom.api.com".into()),
TransportKind::Completion,
0,
);
match core.api_base_config() {
ApiBaseConfig::Single(base) => assert_eq!(base, "https://custom.api.com"),
_ => panic!("Expected Single base config"),
}
}
#[test]
fn test_api_format_accessor() {
let core = LLMCore::new(
"openai".into(),
"gpt-4".into(),
vec![],
3,
ApiKeyConfig::None,
ApiBaseConfig::None,
TransportKind::Responses,
0,
);
assert_eq!(core.api_format(), ApiFormat::Responses);
}
#[test]
fn test_verbose_accessor() {
let core = LLMCore::new(
"openai".into(),
"gpt-4".into(),
vec![],
3,
ApiKeyConfig::None,
ApiBaseConfig::None,
TransportKind::Completion,
2,
);
assert_eq!(core.verbose(), 2);
}
#[test]
fn test_classify_auth_errors() {
assert_eq!(
classify_by_text_signature("Unauthorized request"),
Some(ErrorKind::Config)
);
assert_eq!(
classify_by_text_signature("invalid api key provided"),
Some(ErrorKind::Config)
);
assert_eq!(
classify_by_text_signature("Invalid key format"),
Some(ErrorKind::Config)
);
}
#[test]
fn test_classify_rate_limit_errors() {
assert_eq!(
classify_by_text_signature("Rate limit exceeded"),
Some(ErrorKind::Temporary)
);
assert_eq!(
classify_by_text_signature("HTTP 429 Too Many Requests"),
Some(ErrorKind::Temporary)
);
assert_eq!(
classify_by_text_signature("Quota exceeded for this month"),
Some(ErrorKind::Temporary)
);
}
#[test]
fn test_classify_not_found_errors() {
assert_eq!(
classify_by_text_signature("Model not found"),
Some(ErrorKind::NotFound)
);
assert_eq!(
classify_by_text_signature("HTTP 404"),
Some(ErrorKind::NotFound)
);
}
#[test]
fn test_classify_timeout_errors() {
assert_eq!(
classify_by_text_signature("Request timeout"),
Some(ErrorKind::Temporary)
);
assert_eq!(
classify_by_text_signature("Connection timed out"),
Some(ErrorKind::Temporary)
);
}
#[test]
fn test_classify_server_errors() {
assert_eq!(
classify_by_text_signature("Internal server error"),
Some(ErrorKind::Temporary)
);
assert_eq!(
classify_by_text_signature("HTTP 502 Bad Gateway"),
Some(ErrorKind::Temporary)
);
assert_eq!(
classify_by_text_signature("HTTP 503 Service Unavailable"),
Some(ErrorKind::Temporary)
);
}
#[test]
fn test_classify_unknown_message() {
assert_eq!(classify_by_text_signature("Something went wrong"), None);
assert_eq!(classify_by_text_signature(""), None);
}
#[test]
fn test_prune_orphan_tool_result() {
let messages = vec![
json!({"role": "user", "content": "hello"}),
json!({"role": "tool", "tool_call_id": "call_orphan", "content": "result"}),
json!({"role": "assistant", "content": "hi"}),
];
let result = prune_orphan_tool_messages(messages);
assert_eq!(result.len(), 2);
assert_eq!(result[0]["role"], "user");
assert_eq!(result[1]["role"], "assistant");
}
#[test]
fn test_prune_orphan_tool_call() {
let messages = vec![
json!({"role": "user", "content": "hello"}),
json!({"role": "assistant", "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "foo", "arguments": "{}"}}]}),
];
let result = prune_orphan_tool_messages(messages);
assert_eq!(result.len(), 1);
assert_eq!(result[0]["role"], "user");
}
#[test]
fn test_keep_matched_tool_pair() {
let messages = vec![
json!({"role": "user", "content": "hello"}),
json!({"role": "assistant", "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "foo", "arguments": "{}"}}]}),
json!({"role": "tool", "tool_call_id": "call_1", "content": "result"}),
json!({"role": "assistant", "content": "done"}),
];
let result = prune_orphan_tool_messages(messages);
assert_eq!(result.len(), 4);
}
#[test]
fn test_anthropic_merge_consecutive_user() {
let messages = vec![
json!({"role": "user", "content": "first"}),
json!({"role": "user", "content": "second"}),
json!({"role": "assistant", "content": "reply"}),
];
let result = enforce_anthropic_message_rules(messages);
assert_eq!(result.len(), 3);
assert_eq!(result[0]["content"], "first\n\nsecond");
assert_eq!(result[1]["role"], "assistant");
assert_eq!(result[2]["role"], "user");
assert_eq!(result[2]["content"], "Continue.");
}
#[test]
fn test_anthropic_insert_user_at_start() {
let messages = vec![json!({"role": "assistant", "content": "hi"})];
let result = enforce_anthropic_message_rules(messages);
assert_eq!(result.len(), 3); assert_eq!(result[0]["role"], "user");
assert_eq!(result[0]["content"], "Continue.");
assert_eq!(result[1]["role"], "assistant");
assert_eq!(result[2]["role"], "user");
}
#[test]
fn test_anthropic_append_user_at_end() {
let messages = vec![
json!({"role": "user", "content": "hello"}),
json!({"role": "assistant", "content": "reply"}),
];
let result = enforce_anthropic_message_rules(messages);
assert_eq!(result.len(), 3);
assert_eq!(result[2]["role"], "user");
assert_eq!(result[2]["content"], "Continue.");
}
#[test]
fn test_anthropic_no_append_when_user_last() {
let messages = vec![
json!({"role": "user", "content": "hello"}),
json!({"role": "assistant", "content": "reply"}),
json!({"role": "user", "content": "more"}),
];
let result = enforce_anthropic_message_rules(messages);
assert_eq!(result.len(), 3);
assert_eq!(result[2]["content"], "more");
}
#[test]
fn test_anthropic_system_preserved() {
let messages = vec![
json!({"role": "system", "content": "system prompt"}),
json!({"role": "user", "content": "hello"}),
];
let result = enforce_anthropic_message_rules(messages);
assert_eq!(result.len(), 2);
assert_eq!(result[0]["role"], "system");
assert_eq!(result[1]["role"], "user");
}
#[test]
fn test_normalize_empty() {
let result = normalize_messages_for_api(vec![], TransportKind::Messages);
assert!(result.is_empty());
}
#[test]
fn test_normalize_completion_skips_anthropic_rules() {
let messages = vec![
json!({"role": "user", "content": "first"}),
json!({"role": "user", "content": "second"}),
];
let result = normalize_messages_for_api(messages, TransportKind::Completion);
assert_eq!(result.len(), 2);
}
#[test]
fn test_normalize_messages_keeps_multiple_tool_results_for_messages_transport() {
let messages = vec![
json!({"role": "user", "content": "hello"}),
json!({
"role": "assistant",
"tool_calls": [
{"id": "call_1", "type": "function", "function": {"name": "bash", "arguments": "{}"}},
{"id": "call_2", "type": "function", "function": {"name": "skill", "arguments": "{}"}}
]
}),
json!({"role": "tool", "tool_call_id": "call_1", "content": "result-1"}),
json!({"role": "tool", "tool_call_id": "call_2", "content": "result-2"}),
];
let result = normalize_messages_for_api(messages, TransportKind::Messages);
assert_eq!(result.len(), 4);
assert_eq!(result[2]["tool_call_id"], "call_1");
assert_eq!(result[3]["tool_call_id"], "call_2");
}
#[test]
fn test_normalize_messages_canonicalizes_responses_style_tool_calls() {
let messages = vec![
json!({"role": "user", "content": "hello"}),
json!({
"role": "assistant",
"tool_calls": [{
"type": "function_call",
"call_id": "call_123",
"name": "tape_info",
"arguments": "{}"
}]
}),
json!({"role": "tool", "tool_call_id": "call_123", "content": "ok"}),
];
let result = normalize_messages_for_api(messages, TransportKind::Messages);
assert_eq!(result.len(), 3);
assert_eq!(result[1]["tool_calls"][0]["id"], "call_123");
assert_eq!(result[1]["tool_calls"][0]["function"]["name"], "tape_info");
}
}