use std::sync::OnceLock;
use std::sync::atomic::{AtomicUsize, Ordering};
use crate::config::LlmConfig;
use crate::llm::backend::{BackendFactory, ProviderBackend};
use crate::llm::cache::{Cache, CacheKey};
use crate::llm::concurrency::Limiter;
use crate::llm::error::{BackendErrorKind, LlmError};
use crate::llm::json_parsing::Extracted;
#[derive(Debug)]
pub struct Provider {
pub(crate) backend: ProviderBackend,
identity: String,
limiter: Limiter,
down: OnceLock<LlmError>,
served: AtomicUsize,
}
impl Provider {
pub fn model(&self) -> &str {
self.backend.model()
}
pub fn location(&self) -> &str {
self.backend.location()
}
pub fn limiter(&self) -> &Limiter {
&self.limiter
}
pub fn served(&self) -> usize {
self.served.load(Ordering::Relaxed)
}
pub fn is_down(&self) -> bool {
self.down.get().is_some()
}
fn down_reason(&self) -> Option<&LlmError> {
self.down.get()
}
fn mark_down(&self, err: LlmError) {
let _ = self.down.set(err);
}
fn record_served(&self) {
self.served.fetch_add(1, Ordering::Relaxed);
}
pub fn cache_key(&self, cache: &Cache, system_prompt: &str, user_content: &str) -> CacheKey {
cache.key(
system_prompt,
user_content,
&self.identity,
self.model(),
self.backend.request_identity(),
self.backend.temperature(),
)
}
#[cfg(test)]
pub(crate) fn for_test(backend: ProviderBackend, max_concurrent: usize) -> Self {
let identity = backend.identity();
Self {
backend,
identity,
limiter: Limiter::new(max_concurrent),
down: OnceLock::new(),
served: AtomicUsize::new(0),
}
}
}
#[derive(Debug)]
pub struct Attempt {
pub provider: usize,
pub model: String,
pub error: LlmError,
pub skipped: bool,
}
impl Attempt {
fn new(index: usize, provider: &Provider, error: LlmError, skipped: bool) -> Self {
Self {
provider: index,
model: provider.model().to_owned(),
error,
skipped,
}
}
}
#[derive(Debug)]
pub struct ChainError {
pub attempts: Vec<Attempt>,
pub chain_len: usize,
}
#[derive(Debug)]
pub struct Served {
pub provider: usize,
pub key: CacheKey,
pub extracted: Extracted,
pub from_cache: bool,
}
#[derive(Debug)]
pub struct ProviderChain {
pub(crate) providers: Vec<Provider>,
}
impl ProviderChain {
#[cfg(test)]
pub(crate) fn for_test(providers: impl IntoIterator<Item = (ProviderBackend, usize)>) -> Self {
Self {
providers: providers
.into_iter()
.map(|(backend, max_concurrent)| Provider::for_test(backend, max_concurrent))
.collect(),
}
}
pub fn new(cfgs: &[&LlmConfig]) -> Result<Self, LlmError> {
if cfgs.is_empty() {
return Err(LlmError::NotConfigured(
"no enabled `[[llm]]` provider".to_string(),
));
}
let mut providers = Vec::with_capacity(cfgs.len());
let mut factory = BackendFactory::new();
for (index, cfg) in cfgs.iter().enumerate() {
let backend = factory
.build(cfg)
.map_err(|err| LlmError::NotConfigured(format!("[[llm]] #{}: {err}", index + 1)))?;
let identity = backend.identity();
providers.push(Provider {
backend,
identity,
limiter: Limiter::new(cfg.max_concurrent),
down: OnceLock::new(),
served: AtomicUsize::new(0),
});
}
Ok(Self { providers })
}
pub fn providers(&self) -> &[Provider] {
&self.providers
}
pub async fn complete_json(
&self,
system_prompt: &str,
user_content: &str,
cache: &Cache,
) -> Result<Served, ChainError> {
let mut attempts: Vec<Attempt> = Vec::new();
for (index, provider) in self.providers.iter().enumerate() {
match try_provider(index, provider, system_prompt, user_content, cache).await {
ProviderOutcome::Served(served) => return Ok(served),
ProviderOutcome::Failed { attempt, advance } => {
attempts.push(attempt);
if !advance {
return Err(self.error(attempts));
}
}
}
}
Err(self.error(attempts))
}
pub fn cached_json(
&self,
system_prompt: &str,
user_content: &str,
cache: &Cache,
) -> Option<Served> {
for (index, provider) in self.providers.iter().enumerate() {
let key = provider.cache_key(cache, system_prompt, user_content);
if let Some(served) = cached_for(index, provider, &key, cache) {
return Some(served);
}
}
None
}
fn error(&self, attempts: Vec<Attempt>) -> ChainError {
ChainError {
attempts,
chain_len: self.providers.len(),
}
}
}
enum ProviderOutcome {
Served(Served),
Failed { attempt: Attempt, advance: bool },
}
impl ProviderOutcome {
fn demoted(index: usize, provider: &Provider, err: &LlmError) -> Self {
ProviderOutcome::Failed {
advance: should_failover(err),
attempt: Attempt::new(index, provider, err.clone(), true),
}
}
}
async fn try_provider(
index: usize,
provider: &Provider,
system_prompt: &str,
user_content: &str,
cache: &Cache,
) -> ProviderOutcome {
let key = provider.cache_key(cache, system_prompt, user_content);
if let Some(served) = cached_for(index, provider, &key, cache) {
return ProviderOutcome::Served(served);
}
if let Some(err) = provider.down_reason() {
return ProviderOutcome::demoted(index, provider, err);
}
let guard = provider.limiter.acquire().await;
if let Some(err) = provider.down_reason() {
drop(guard);
return ProviderOutcome::demoted(index, provider, err);
}
let outcome = provider
.backend
.complete_json(system_prompt, user_content)
.await;
drop(guard);
match outcome {
Ok(extracted) => {
provider.record_served();
ProviderOutcome::Served(Served {
provider: index,
key,
extracted,
from_cache: false,
})
}
Err(err) => {
if is_sticky(&err) {
provider.mark_down(err.clone());
}
ProviderOutcome::Failed {
advance: should_failover(&err),
attempt: Attempt::new(index, provider, err, false),
}
}
}
}
fn cached_for(index: usize, provider: &Provider, key: &CacheKey, cache: &Cache) -> Option<Served> {
let value = cache.get(key)?;
provider.record_served();
Some(Served {
provider: index,
key: key.clone(),
extracted: Extracted::Complete(value),
from_cache: true,
})
}
fn should_failover(err: &LlmError) -> bool {
match err {
LlmError::Transport { status: None, .. } => true,
LlmError::Transport {
status: Some(code), ..
} => is_retryable_status(*code),
LlmError::Unparseable(_) => true,
LlmError::ModelStopped { .. } => false,
LlmError::NotConfigured(_) => false,
LlmError::Backend { kind, .. } => matches!(kind, BackendErrorKind::UsageLimit),
}
}
fn is_sticky(err: &LlmError) -> bool {
(should_failover(err) && !matches!(err, LlmError::Unparseable(_)))
|| is_auth_failure(err)
|| is_sticky_backend_failure(err)
}
fn is_sticky_backend_failure(err: &LlmError) -> bool {
matches!(
err,
LlmError::Backend {
kind: BackendErrorKind::Contract | BackendErrorKind::Authentication,
..
}
)
}
fn is_auth_failure(err: &LlmError) -> bool {
matches!(
err,
LlmError::Transport {
status: Some(401 | 403),
..
}
)
}
fn is_retryable_status(code: u16) -> bool {
matches!(code, 408 | 429) || (500..=599).contains(&code)
}
#[cfg(test)]
mod tests;