use std::collections::HashMap;
use std::num::NonZeroUsize;
use std::sync::Arc;
use std::time::Duration;
use misanthropic::model::{ModelInfo, Models};
use misanthropic::{batch, response};
use serde::{Deserialize, Serialize};
use crate::retry::client_error_recoverable;
use super::backend::Inference;
use super::inference::Quirks;
const DEFAULT_MAX_BATCH: usize = 1000;
const DEFAULT_POLL_PERIOD: Duration = Duration::from_secs(5);
const DUMMY_KEY: &str = "sk-ant-api03-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx";
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub enum EndpointVariant {
#[default]
Anthropic,
Ollama,
Blallama,
}
impl From<EndpointVariant> for Quirks {
fn from(variant: EndpointVariant) -> Self {
let mut quirks = Quirks::default();
match variant {
EndpointVariant::Anthropic => {}
EndpointVariant::Ollama => {
quirks.cache_markers_ignored = true;
quirks.tool_choice_not_respected = true;
quirks.cache_stats_unreported = true;
quirks.web_search_unsupported = true;
quirks.web_fetch_unsupported = true;
}
EndpointVariant::Blallama => {
quirks.breakpoint_after_assistant = true;
quirks.output_config_cache_safe = true;
quirks.web_search_unsupported = true;
quirks.web_fetch_unsupported = true;
}
}
quirks
}
}
#[derive(Deserialize)]
struct Tags {
#[serde(default)]
models: Vec<Tag>,
}
#[derive(Deserialize)]
struct Tag {
name: String,
#[serde(default)]
modified_at: Option<chrono::DateTime<chrono::Utc>>,
}
fn models_from_tags(body: &str) -> Result<Models, misanthropic::client::Error> {
let tags: Tags = serde_json::from_str(body)?;
Ok(tags
.models
.into_iter()
.map(|tag| ModelInfo {
id: tag.name.clone().into(),
display_name: tag.name.into(),
capabilities: Default::default(),
max_input_tokens: 0,
max_tokens: 0,
kind: Default::default(),
created_at: tag.modified_at.unwrap_or_default(),
})
.collect())
}
pub struct Client {
client: misanthropic::Client,
variant: EndpointVariant,
concurrency: NonZeroUsize,
max_batch: usize,
poll_period: Duration,
}
impl Client {
pub fn new(client: misanthropic::Client) -> Self {
Self {
client,
variant: EndpointVariant::default(),
concurrency: 1.try_into().unwrap(),
max_batch: DEFAULT_MAX_BATCH,
poll_period: DEFAULT_POLL_PERIOD,
}
}
pub fn with_variant(mut self, variant: EndpointVariant) -> Self {
self.variant = variant;
if !matches!(variant, EndpointVariant::Anthropic) {
self.client.key = Arc::new(
DUMMY_KEY
.to_string()
.try_into()
.expect("DUMMY_KEY has a valid key length"),
);
}
self
}
pub fn with_concurrency(mut self, n: NonZeroUsize) -> Self {
self.set_concurrency(n);
self
}
pub fn set_concurrency(&mut self, n: NonZeroUsize) {
self.concurrency = n;
}
pub fn with_batch(
mut self,
max_batch: usize,
poll_period: Duration,
) -> Self {
self.max_batch = max_batch;
self.poll_period = poll_period;
self
}
}
const MAX_BATCH_RETRIES: usize = 5;
async fn backoff(attempt: usize) {
let secs = 1u64 << attempt.min(5);
tokio::time::sleep(Duration::from_secs(secs.min(30))).await;
}
#[async_trait::async_trait]
impl Inference for Client {
type Error = misanthropic::client::Error;
async fn infer<P>(
&self,
prompt: P,
) -> Result<response::Message, Self::Error>
where
P: Serialize + Send,
{
self.client.message(prompt).await
}
async fn infer_batch<P>(
&self,
prompts: &[&P],
) -> Result<Vec<Result<response::Message, Self::Error>>, Self::Error>
where
P: Serialize + Send + Sync,
{
let chunk = self.max_batch.max(1);
let mut out: Vec<Option<Result<response::Message, Self::Error>>> =
(0..prompts.len()).map(|_| None).collect();
for start in (0..prompts.len()).step_by(chunk) {
let end = (start + chunk).min(prompts.len());
let mut id_to_idx: HashMap<batch::Id, usize> = HashMap::new();
let items: Vec<(batch::Id, &P)> = (start..end)
.map(|i| {
let id = batch::Id::default();
id_to_idx.insert(id, i);
(id, prompts[i])
})
.collect();
let mut pending = {
let mut attempt = 0usize;
loop {
match self.client.tagged_batch(items.clone()).await {
Ok(pending) => break pending,
Err(e) => {
if attempt >= MAX_BATCH_RETRIES
|| !client_error_recoverable(&e)
{
return Err(e);
}
tracing::warn!(
attempt = attempt + 1,
max = MAX_BATCH_RETRIES,
error = %e,
"batch submit failed, retrying"
);
backoff(attempt).await;
attempt += 1;
}
}
}
};
let ready = {
let mut attempt = 0usize;
loop {
match self.client.batch_poll(pending).await {
Ok(batch::Batch::Ready(ready)) => break ready,
Ok(batch::Batch::Pending(p)) => {
pending = p;
attempt = 0;
tokio::time::sleep(self.poll_period).await;
}
Err(batch::Error {
client_error,
pending: p,
}) => {
if attempt >= MAX_BATCH_RETRIES
|| !client_error_recoverable(&client_error)
{
return Err(client_error);
}
tracing::warn!(
attempt = attempt + 1,
max = MAX_BATCH_RETRIES,
batch_id = %p.meta().id,
error = %client_error,
"batch poll failed, retrying"
);
pending = p;
backoff(attempt).await;
attempt += 1;
}
}
}
};
let (_, results) = ready.decompose();
for (id, result) in results {
if let Some(&idx) = id_to_idx.get(&id) {
let r: Result<
response::Message,
misanthropic::client::AnthropicError,
> = result.into();
out[idx] = Some(r.map_err(Into::into));
}
}
}
Ok(out
.into_iter()
.map(|slot| {
slot.unwrap_or(Err(
misanthropic::client::Error::UnexpectedResponse {
message: "batch returned no result for prompt",
},
))
})
.collect())
}
async fn models(&self) -> Result<misanthropic::model::Models, Self::Error> {
match self.variant {
EndpointVariant::Anthropic | EndpointVariant::Blallama => {
self.client.models().await
}
EndpointVariant::Ollama => {
let url = self.client.messages_url.join("/api/tags").map_err(
|_| misanthropic::client::Error::UnexpectedResponse {
message: "cannot derive /api/tags from messages_url",
},
)?;
let body = self
.client
.inner
.get(url)
.send()
.await?
.error_for_status()?
.text()
.await?;
models_from_tags(&body)
}
}
}
fn quirks(&self) -> Quirks {
self.variant.into()
}
fn max_concurrency(&self) -> NonZeroUsize {
self.concurrency
}
}
#[cfg(test)]
mod tests {
use super::*;
use super::super::{COURTESY_BACKOFF, RetryAfter};
#[test]
fn variant_quirks_mapping() {
assert_eq!(Quirks::from(EndpointVariant::Anthropic), Quirks::default());
let ollama = Quirks::from(EndpointVariant::Ollama);
assert!(ollama.cache_markers_ignored);
assert!(ollama.tool_choice_not_respected);
assert!(ollama.cache_stats_unreported);
assert!(!ollama.breakpoint_after_assistant);
assert!(!ollama.output_config_cache_safe);
let blallama = Quirks::from(EndpointVariant::Blallama);
assert!(blallama.breakpoint_after_assistant);
assert!(blallama.output_config_cache_safe);
assert!(!blallama.cache_markers_ignored);
assert!(!blallama.tool_choice_not_respected);
assert!(!blallama.cache_stats_unreported);
assert!(!Quirks::default().web_search_unsupported);
assert!(!Quirks::default().web_fetch_unsupported);
assert!(ollama.web_search_unsupported);
assert!(ollama.web_fetch_unsupported);
assert!(blallama.web_search_unsupported);
assert!(blallama.web_fetch_unsupported);
}
#[test]
fn server_hints_win_over_the_courtesy_backoff() {
use misanthropic::client::{AnthropicError, Error};
let e = Error::Anthropic(AnthropicError::Overloaded {
message: "overloaded".into(),
retry_after: Some(3),
});
assert_eq!(
e.retry_after(),
Some(Duration::from_secs(3)),
"an explicit Retry-After must not be replaced by the default"
);
let e = Error::Anthropic(AnthropicError::Overloaded {
message: "Session is busy.".into(),
retry_after: None,
});
assert_eq!(e.retry_after(), Some(COURTESY_BACKOFF));
}
#[test]
fn transient_classes_are_retryable_without_a_header() {
use misanthropic::client::{AnthropicError, Error};
use std::num::NonZeroU16;
let cases: Vec<(&str, Error)> = vec![
(
"the 2026-08-21 gateway 503",
Error::NonJsonResponse {
status: 503,
body: "upstream connect error".into(),
},
),
(
"a non-JSON 429 challenge page",
Error::NonJsonResponse {
status: 429,
body: "<html>slow down</html>".into(),
},
),
(
"header-less 429",
Error::Anthropic(AnthropicError::RateLimit {
message: "slow down".into(),
retry_after: None,
}),
),
(
"5xx from the API itself",
Error::Anthropic(AnthropicError::API {
message: "boom".into(),
}),
),
(
"an unknown 5xx",
Error::Anthropic(AnthropicError::Unknown {
code: Some(NonZeroU16::new(502).unwrap()),
message: "bad gateway".into(),
}),
),
];
for (what, e) in cases {
assert_eq!(
e.retry_after(),
Some(COURTESY_BACKOFF),
"{what} must be retryable"
);
assert!(!e.is_fatal(), "{what} must not be fatal");
}
}
#[test]
fn caller_side_failures_stay_fatal() {
use misanthropic::client::{AnthropicError, Error};
use std::num::NonZeroU16;
let cases: Vec<(&str, Error)> = vec![
(
"a non-JSON 404 — wrong URL, not a blip",
Error::NonJsonResponse {
status: 404,
body: "<html>not found</html>".into(),
},
),
(
"bad request",
Error::Anthropic(AnthropicError::InvalidRequest {
message: "malformed".into(),
}),
),
(
"bad key",
Error::Anthropic(AnthropicError::Authentication {
message: "nope".into(),
}),
),
(
"an unknown 4xx",
Error::Anthropic(AnthropicError::Unknown {
code: Some(NonZeroU16::new(418).unwrap()),
message: "teapot".into(),
}),
),
(
"an unparseable success body",
Error::UnexpectedResponse {
message: "stream where a message was expected",
},
),
];
for (what, e) in cases {
assert_eq!(e.retry_after(), None, "{what} must stay fatal");
assert!(e.is_fatal(), "{what} must be fatal");
}
}
#[test]
fn the_retry_helper_agrees_with_retry_after() {
use crate::retry::client_error_recoverable;
use misanthropic::client::{AnthropicError, Error};
let transient = Error::NonJsonResponse {
status: 503,
body: "upstream connect error".into(),
};
assert!(client_error_recoverable(&transient));
assert!(!transient.is_fatal());
let fatal = Error::Anthropic(AnthropicError::Authentication {
message: "nope".into(),
});
assert!(!client_error_recoverable(&fatal));
assert!(fatal.is_fatal());
}
#[tokio::test(start_paused = true)]
async fn batch_submit_retries_a_transient_gateway_503() {
use httpmock::prelude::*;
use misanthropic::{Prompt, prompt::message::Role};
let server = MockServer::start();
let mock = server.mock(|when, then| {
when.method(POST).path("/v1/messages/batches/");
then.status(503).body(
"upstream connect error or disconnect/reset before headers. \
reset reason: connection termination",
);
});
let transport = Client::new(
misanthropic::Client::new("x".repeat(108))
.unwrap()
.base_url(server.base_url())
.unwrap(),
);
let prompt = Prompt::default()
.model(misanthropic::Id::Haiku45)
.max_tokens(std::num::NonZeroU32::new(16).unwrap())
.add_message((Role::User, "hi"))
.unwrap();
let error = transport
.infer_batch(&[&prompt])
.await
.expect_err("a permanent 503 must eventually surface");
assert!(
matches!(
error,
misanthropic::client::Error::NonJsonResponse {
status: 503,
..
}
),
"expected the edge 503 to surface unchanged, got: {error:?}"
);
mock.assert_hits(MAX_BATCH_RETRIES + 1);
}
#[tokio::test(start_paused = true)]
async fn batch_submit_does_not_retry_a_400() {
use httpmock::prelude::*;
use misanthropic::{Prompt, prompt::message::Role};
let server = MockServer::start();
let mock = server.mock(|when, then| {
when.method(POST).path("/v1/messages/batches/");
then.status(400).json_body(serde_json::json!({
"type": "error",
"error": { "type": "invalid_request_error", "message": "bad" }
}));
});
let transport = Client::new(
misanthropic::Client::new("x".repeat(108))
.unwrap()
.base_url(server.base_url())
.unwrap(),
);
let prompt = Prompt::default()
.model(misanthropic::Id::Haiku45)
.max_tokens(std::num::NonZeroU32::new(16).unwrap())
.add_message((Role::User, "hi"))
.unwrap();
transport
.infer_batch(&[&prompt])
.await
.expect_err("a 400 must surface");
mock.assert_hits(1);
}
#[tokio::test]
#[ignore = "hits the live Anthropic API (a search, so cents)"]
async fn live_batch_runs_server_tools() {
use misanthropic::prompt::message::Block;
use misanthropic::tool::{ServerMethodDef, WebSearch};
use misanthropic::{
Prompt, prompt::message::Role, response::StopReason,
};
let key = std::env::var("ANTHROPIC_API_KEY").unwrap_or_else(|_| {
let path = format!(
"{}/Projects/agora/secrets/anthropic_api_key",
std::env::var("HOME").expect("HOME")
);
std::fs::read_to_string(path)
.expect("no ANTHROPIC_API_KEY and no key file")
.trim()
.to_string()
});
let transport = Client::new(misanthropic::Client::new(key).unwrap());
let prompt = Prompt::default()
.model(misanthropic::Id::Haiku45)
.max_tokens(std::num::NonZeroU32::new(512).unwrap())
.add_message((
Role::User,
"Search and name one product Anthropic makes.",
))
.unwrap()
.add_tool(ServerMethodDef::web_search(WebSearch {
max_uses: Some(1),
allowed_domains: Some(vec!["anthropic.com".into()]),
..Default::default()
}));
let results = transport
.infer_batch(&[&prompt])
.await
.expect("batch submission accepted");
let response = results
.into_iter()
.next()
.expect("one prompt, one result")
.expect("the batch item itself succeeded");
let searched = response
.inner
.content
.iter()
.any(|b| matches!(b, Block::WebSearchToolResult { .. }));
let paused =
matches!(response.stop_reason, Some(StopReason::PauseTurn));
println!(
"batch server-tool run: stop={:?} searched={searched} \
usage={:?}\n{}",
response.stop_reason,
response.usage.server_tool_use,
response.inner.content
);
assert!(
searched || paused,
"the batch item ran no server tool: {:?}",
response.inner.content
);
}
#[tokio::test]
async fn blallama_discovers_via_v1_models() {
use httpmock::prelude::*;
let server = MockServer::start();
let mock = server.mock(|when, then| {
when.method(GET)
.path("/v1/models")
.header("x-api-key", DUMMY_KEY);
then.status(200)
.header("content-type", "application/json")
.body(include_str!(
"../../tests/fixtures/blallama_v1_models.json"
));
});
let transport = Client::new(
misanthropic::Client::new("x".repeat(108))
.unwrap()
.base_url(server.base_url())
.unwrap(),
)
.with_variant(EndpointVariant::Blallama);
let models = transport.models().await.unwrap();
mock.assert();
let qwen = models
.iter()
.find(|m| m.id.name() == "Qwen3.8-27B-UD-Q8_K_XL.gguf")
.expect("offered");
assert_eq!(qwen.display_name, "Qwen3.8-27B");
assert_eq!(qwen.max_input_tokens, 131072);
assert!(qwen.capabilities.structured_outputs.supported);
let stored: ModelInfo = serde_json::from_value(serde_json::json!({
"capabilities": {},
"created_at": "1970-01-01T00:00:00Z",
"display_name": "Qwen3.8-27B-UD-Q8_K_XL.gguf",
"id": "Qwen3.8-27B-UD-Q8_K_XL.gguf",
"max_input_tokens": 0,
"max_tokens": 0,
"type": "model"
}))
.unwrap();
assert!(models.iter().any(|m| m.satisfies(&stored)));
for offered in models.iter() {
let mut requested = offered.clone();
requested.display_name = requested.id.name().to_owned().into();
requested.capabilities = Default::default();
requested.max_input_tokens = 0;
requested.max_tokens = 0;
assert!(offered.satisfies(&requested), "{}", offered.id.name());
}
}
#[test]
fn models_from_tags_synthesizes() {
let body = r#"{
"models": [
{"name": "llama3.3:70b", "modified_at": "2026-01-01T00:00:00Z"},
{"name": "qwen3:32b"}
]
}"#;
let models = models_from_tags(body).unwrap();
let infos: Vec<&ModelInfo> = models.iter().collect();
assert_eq!(infos.len(), 2);
assert_eq!(infos[0].id.name(), "llama3.3:70b");
assert!(!infos[0].capabilities.batch.supported, "batch never");
assert_eq!(infos[0].max_tokens, 0, "ceiling unreported");
assert_eq!(infos[1].id.name(), "qwen3:32b");
}
}