use std::time::Duration;
use async_trait::async_trait;
use thiserror::Error;
use crate::completion::{Completion, CompletionDelta};
use crate::request::ConversationRequest;
#[async_trait]
pub trait Provider: Send + Sync {
fn api_schema(&self) -> &str;
async fn complete(&self, request: &ConversationRequest) -> Result<Completion, ProviderError>;
async fn stream(
&self,
request: &ConversationRequest,
on_delta: &mut (dyn FnMut(CompletionDelta) + Send),
) -> Result<Completion, ProviderError> {
let completion = self.complete(request).await?;
if let Some(text) = completion.text() {
on_delta(CompletionDelta::Text(text));
}
Ok(completion)
}
}
#[derive(Debug, Error)]
pub enum ProviderError {
#[error("transport error: {0}")]
Transport(String),
#[error("rate limited")]
RateLimited {
retry_after: Option<Duration>,
},
#[error("api error (status {status}): {message}")]
Api {
status: u16,
message: String,
},
#[error("context window exceeded")]
ContextOverflow,
#[error("quota exceeded")]
Quota,
#[error("authentication error: {0}")]
Auth(String),
#[error("failed to decode provider response: {0}")]
Decode(String),
#[error("invalid configuration: {0}")]
Config(String),
}
impl ProviderError {
#[must_use]
pub fn retryable(&self) -> bool {
match self {
ProviderError::Transport(_) | ProviderError::RateLimited { .. } => true,
ProviderError::Api { status, .. } => {
matches!(status, 408 | 409 | 500..=599)
}
ProviderError::ContextOverflow
| ProviderError::Quota
| ProviderError::Auth(_)
| ProviderError::Decode(_)
| ProviderError::Config(_) => false,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::completion::StopReason;
use crate::request::{CacheHint, SamplingArgs};
use locode_protocol::{ContentBlock, Usage};
struct TextOnly(String);
#[async_trait]
impl Provider for TextOnly {
#[allow(clippy::unnecessary_literal_bound)]
fn api_schema(&self) -> &str {
"text-only"
}
async fn complete(
&self,
_request: &ConversationRequest,
) -> Result<Completion, ProviderError> {
Ok(Completion {
content: vec![ContentBlock::Text {
text: self.0.clone(),
}],
usage: Usage::default(),
stop: StopReason::EndTurn,
})
}
}
fn req() -> ConversationRequest {
ConversationRequest {
messages: vec![],
tools: vec![],
sampling_args: SamplingArgs::default(),
cache_hint: CacheHint::default(),
}
}
#[tokio::test]
async fn default_stream_emits_one_text_delta_and_same_completion() {
let provider = TextOnly("hello there".into());
let mut deltas = Vec::new();
let completion = provider
.stream(&req(), &mut |d| deltas.push(d))
.await
.expect("default stream ok");
assert_eq!(deltas, vec![CompletionDelta::Text("hello there".into())]);
assert_eq!(completion.text().as_deref(), Some("hello there"));
}
#[tokio::test]
async fn default_stream_with_empty_text_emits_nothing() {
let provider = TextOnly(String::new());
let mut deltas = Vec::new();
provider
.stream(&req(), &mut |d| deltas.push(d))
.await
.expect("ok");
assert!(
deltas.is_empty(),
"no text → no synthetic delta: {deltas:?}"
);
}
}