mod fake;
mod retry;
pub use fake::{FakeProvider, FakeReply};
pub use retry::{Backoff, RetryPolicy, RetryProvider, Retryable};
use crate::message::Message;
use crate::run::{RunContext, RunMetadata};
use crate::tool::ToolSchema;
use futures::stream::BoxStream;
use serde::{Deserialize, Serialize};
use std::collections::BTreeMap;
use std::fmt;
use std::time::{Duration, Instant};
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ProviderCapabilities {
pub streaming: bool,
pub reasoning: bool,
pub tool_calls: bool,
pub parallel_tool_calls: bool,
pub structured_output: bool,
pub usage: bool,
pub context_cancellation: bool,
pub context_deadline: bool,
}
impl ProviderCapabilities {
pub fn baseline() -> Self {
Self::default()
}
}
#[derive(Clone)]
pub struct ProviderRequestContext {
pub run_id: String,
pub model_request_id: String,
pub cancellation: tokio_util::sync::CancellationToken,
pub deadline: Option<Instant>,
pub timeout: Option<Duration>,
pub metadata: RunMetadata,
}
impl fmt::Debug for ProviderRequestContext {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ProviderRequestContext")
.field("run_id", &self.run_id)
.field("model_request_id", &self.model_request_id)
.field("cancellation", &"CancellationToken")
.field("deadline", &self.deadline)
.field("timeout", &self.timeout)
.field("metadata", &self.metadata)
.finish()
}
}
impl ProviderRequestContext {
pub fn from_run_context(model_request_id: impl Into<String>, context: &RunContext) -> Self {
Self {
run_id: context.run_id.clone(),
model_request_id: model_request_id.into(),
cancellation: context.cancellation.clone(),
deadline: context.deadline,
timeout: context.remaining(),
metadata: context.metadata.clone(),
}
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout);
self
}
pub fn is_cancelled(&self) -> bool {
self.cancellation.is_cancelled()
}
pub fn is_expired(&self) -> bool {
self.deadline
.is_some_and(|deadline| Instant::now() >= deadline)
}
pub fn remaining(&self) -> Option<Duration> {
self.deadline
.map(|deadline| deadline.saturating_duration_since(Instant::now()))
}
}
#[async_trait::async_trait]
pub trait Provider: Send + Sync {
fn model(&self) -> Option<&str> {
None
}
fn capabilities(&self) -> ProviderCapabilities {
ProviderCapabilities::baseline()
}
async fn chat_with_context(
&self,
request: ChatRequest,
context: &ProviderRequestContext,
) -> Result<ChatResponse, ProviderError>;
async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, ProviderError> {
let run = RunContext::generated();
let context = ProviderRequestContext::from_run_context("direct-chat", &run);
self.chat_with_context(request, &context).await
}
async fn stream_chat(
&self,
request: ChatRequest,
) -> Result<BoxStream<'static, Result<StreamEvent, ProviderError>>, ProviderError> {
let run = RunContext::generated();
let context = ProviderRequestContext::from_run_context("direct-stream", &run);
self.stream_chat_with_context(request, &context).await
}
async fn stream_chat_with_context(
&self,
request: ChatRequest,
context: &ProviderRequestContext,
) -> Result<BoxStream<'static, Result<StreamEvent, ProviderError>>, ProviderError>;
}
#[async_trait::async_trait]
impl Provider for Box<dyn Provider> {
fn model(&self) -> Option<&str> {
self.as_ref().model()
}
fn capabilities(&self) -> ProviderCapabilities {
self.as_ref().capabilities()
}
async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, ProviderError> {
self.as_ref().chat(request).await
}
async fn chat_with_context(
&self,
request: ChatRequest,
context: &ProviderRequestContext,
) -> Result<ChatResponse, ProviderError> {
self.as_ref().chat_with_context(request, context).await
}
async fn stream_chat(
&self,
request: ChatRequest,
) -> Result<BoxStream<'static, Result<StreamEvent, ProviderError>>, ProviderError> {
self.as_ref().stream_chat(request).await
}
async fn stream_chat_with_context(
&self,
request: ChatRequest,
context: &ProviderRequestContext,
) -> Result<BoxStream<'static, Result<StreamEvent, ProviderError>>, ProviderError> {
self.as_ref()
.stream_chat_with_context(request, context)
.await
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct ChatRequest {
pub messages: Vec<Message>,
pub tools: Vec<ToolSchema>,
pub options: ModelOptions,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct ModelOptions {
pub temperature: Option<f32>,
pub max_tokens: Option<u32>,
pub extra: BTreeMap<String, serde_json::Value>,
pub structured: Option<serde_json::Value>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct Usage {
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub total_tokens: u32,
}
impl Usage {
pub fn new(prompt_tokens: u32, completion_tokens: u32) -> Self {
Self {
prompt_tokens,
completion_tokens,
total_tokens: prompt_tokens + completion_tokens,
}
}
}
impl std::ops::AddAssign for Usage {
fn add_assign(&mut self, rhs: Self) {
self.prompt_tokens += rhs.prompt_tokens;
self.completion_tokens += rhs.completion_tokens;
self.total_tokens += rhs.total_tokens;
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ChatResponse {
pub message: Message,
pub finish_reason: FinishReason,
pub usage: Option<Usage>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum FinishReason {
Stop,
Length,
Other(String),
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum StreamEvent {
Delta(String),
ToolCall {
id: String,
name: String,
arguments: String,
},
Reasoning(String),
Done {
reason: FinishReason,
usage: Option<Usage>,
},
}
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
#[non_exhaustive]
pub enum ProviderError {
#[error("api error (status {status}): {message}")]
Api {
status: u16,
code: Option<String>,
message: String,
},
#[error("provider protocol error: {message}")]
Protocol {
message: String,
},
#[error("response decode error: {message}")]
Decode {
message: String,
},
#[error("rate limited")]
RateLimited {
retry_after: Option<Duration>,
},
#[error("network error: {0}")]
Network(String),
#[error("request timed out during {0:?}")]
Timeout(TimeoutStage),
#[error("provider request cancelled")]
Cancelled,
#[error("response exceeded configured limit ({limit_bytes} bytes)")]
ResponseTooLarge {
limit_bytes: usize,
},
#[error("unsupported provider capability: {capability}")]
Unsupported {
capability: &'static str,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub enum TimeoutStage {
Connect,
Request,
Idle,
StreamTotal,
ResponseBody,
Transport,
}