pub mod build;
mod client;
pub mod config;
pub mod error;
pub mod parse;
pub mod retry;
pub mod wire;
use std::sync::{Arc, Mutex, PoisonError};
use async_trait::async_trait;
pub use build::{build_request, count_cache_controls, normalize_input_schema};
pub use config::{ApiBackend, AuthScheme, DeveloperRendering, ModelConfig};
pub use error::{HttpFailure, classify, parse_retry_after};
pub use parse::response_to_completion;
pub use retry::{RetryPolicy, backoff, run_with_retry};
use crate::completion::{Completion, CompletionDelta};
use crate::provider::{Provider, ProviderError};
use crate::repair::repair_pairing;
use crate::request::ConversationRequest;
pub trait AuthRefresh: Send + Sync {
fn refresh(&self) -> Option<AuthScheme>;
}
pub struct AnthropicProvider {
http: reqwest::Client,
config: ModelConfig,
retry: RetryPolicy,
auth: Mutex<AuthScheme>,
auth_refresh: Option<Arc<dyn AuthRefresh>>,
}
impl AnthropicProvider {
pub fn new(config: ModelConfig) -> Result<Self, ProviderError> {
Ok(Self {
http: crate::http::build_http_client()?,
auth: Mutex::new(config.auth.clone()),
config,
retry: RetryPolicy::default(),
auth_refresh: None,
})
}
pub fn from_env() -> Result<Self, ProviderError> {
Self::new(ModelConfig::from_env()?)
}
#[must_use]
pub fn with_retry_policy(mut self, retry: RetryPolicy) -> Self {
self.retry = retry;
self
}
#[must_use]
pub fn with_auth_refresh(mut self, refresh: Arc<dyn AuthRefresh>) -> Self {
self.auth_refresh = Some(refresh);
self
}
#[must_use]
pub fn config(&self) -> &ModelConfig {
&self.config
}
pub fn config_mut(&mut self) -> &mut ModelConfig {
&mut self.config
}
fn current_auth(&self) -> AuthScheme {
self.auth
.lock()
.unwrap_or_else(PoisonError::into_inner)
.clone()
}
fn try_refresh(&self) -> bool {
let Some(refresher) = &self.auth_refresh else {
return false;
};
let Some(new_auth) = refresher.refresh() else {
return false;
};
let mut current = self.auth.lock().unwrap_or_else(PoisonError::into_inner);
if *current == new_auth {
return false;
}
*current = new_auth;
true
}
async fn exchange(&self, request: &wire::MessagesRequest) -> Result<Completion, ProviderError> {
let auth = self.current_auth();
run_with_retry(&self.retry, |_attempt| {
client::send_once(&self.http, &self.config, &auth, request)
})
.await
}
}
#[async_trait]
impl Provider for AnthropicProvider {
#[allow(clippy::unnecessary_literal_bound)] fn api_schema(&self) -> &str {
"anthropic"
}
async fn complete(&self, request: &ConversationRequest) -> Result<Completion, ProviderError> {
let mut repaired = request.clone();
let _ = repair_pairing(&mut repaired.messages);
let wire_request = build_request(&repaired, &self.config);
match self.exchange(&wire_request).await {
Err(ProviderError::Auth(message)) => {
if self.try_refresh() {
self.exchange(&wire_request).await
} else {
Err(ProviderError::Auth(message))
}
}
other => other,
}
}
async fn stream(
&self,
request: &ConversationRequest,
on_delta: &mut (dyn FnMut(CompletionDelta) + Send),
) -> Result<Completion, ProviderError> {
let mut repaired = request.clone();
let _ = repair_pairing(&mut repaired.messages);
let mut wire_request = build_request(&repaired, &self.config);
wire_request.stream = Some(true);
let first = client::send_once_streaming(
&self.http,
&self.config,
&self.current_auth(),
&wire_request,
on_delta,
)
.await;
match first {
Ok(completion) => Ok(completion),
Err(failure)
if matches!(failure.error, ProviderError::Auth(_)) && self.try_refresh() =>
{
client::send_once_streaming(
&self.http,
&self.config,
&self.current_auth(),
&wire_request,
on_delta,
)
.await
.map_err(|f| f.error)
}
Err(failure) => Err(failure.error),
}
}
}