use crate::event_stream::AssistantMessageEventStream;
use crate::model::Model;
use crate::types::{AssistantMessage, Context, ThinkingBudgets, ThinkingLevel};
use std::collections::BTreeMap;
use std::time::Duration;
use tokio_util::sync::CancellationToken;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum CacheRetention {
None,
#[default]
Short,
Long,
}
#[derive(Debug, Clone)]
pub struct SimpleStreamOptions {
pub api_key: Option<String>,
pub timeout: Option<Duration>,
pub max_retries: Option<u32>,
pub max_retry_delay: Option<Duration>,
pub headers: Option<BTreeMap<String, String>>,
pub metadata: Option<BTreeMap<String, String>>,
pub cache_retention: CacheRetention,
pub session_id: Option<String>,
pub signal: CancellationToken,
pub reasoning: Option<ThinkingLevel>,
pub max_tokens: Option<u64>,
pub temperature: Option<f64>,
pub thinking_budgets: Option<ThinkingBudgets>,
}
impl Default for SimpleStreamOptions {
fn default() -> Self {
Self {
api_key: None,
timeout: None,
max_retries: None,
max_retry_delay: None,
headers: None,
metadata: None,
cache_retention: CacheRetention::default(),
session_id: None,
signal: CancellationToken::new(),
reasoning: None,
max_tokens: None,
temperature: None,
thinking_budgets: None,
}
}
}
impl SimpleStreamOptions {
pub fn new() -> Self {
Self::default()
}
pub fn with_api_key(mut self, key: impl Into<String>) -> Self {
self.api_key = Some(key.into());
self
}
pub fn with_signal(mut self, signal: CancellationToken) -> Self {
self.signal = signal;
self
}
pub fn reasoning_level(&self) -> ThinkingLevel {
self.reasoning.unwrap_or(ThinkingLevel::Off)
}
}
#[async_trait::async_trait]
pub trait Provider: Send + Sync {
fn id(&self) -> &str;
fn models(&self) -> &[Model];
async fn stream_simple(
&self,
model: &Model,
ctx: &Context,
opts: &SimpleStreamOptions,
) -> AssistantMessageEventStream;
}
pub trait ProviderHooks: Send + Sync {
fn before_request(
&self,
_model: &Model,
_ctx: &Context,
_opts: &SimpleStreamOptions,
) -> Option<SimpleStreamOptionsPatch> {
None
}
fn after_response(&self, _model: &Model, _message: &AssistantMessage) {}
}
#[derive(Default)]
pub struct NoopProviderHooks;
impl ProviderHooks for NoopProviderHooks {}
#[derive(Debug, Clone, Default)]
pub struct SimpleStreamOptionsPatch {
pub timeout: Option<Option<Duration>>,
pub max_retries: Option<Option<u32>>,
pub max_retry_delay: Option<Option<Duration>>,
pub headers: Option<BTreeMap<String, Option<String>>>,
pub metadata: Option<BTreeMap<String, Option<String>>>,
pub cache_retention: Option<Option<CacheRetention>>,
pub max_tokens: Option<Option<u64>>,
pub temperature: Option<Option<f64>>,
pub session_id: Option<Option<String>>,
}
impl SimpleStreamOptionsPatch {
pub fn apply(&self, opts: &mut SimpleStreamOptions) {
if let Some(v) = self.timeout {
opts.timeout = v;
}
if let Some(v) = self.max_retries {
opts.max_retries = v;
}
if let Some(v) = self.max_retry_delay {
opts.max_retry_delay = v;
}
if let Some(patch) = &self.headers {
let mut map = opts.headers.take().unwrap_or_default();
for (k, v) in patch {
match v {
Some(val) => map.insert(k.clone(), val.clone()),
None => map.remove(k),
};
}
opts.headers = if map.is_empty() { None } else { Some(map) };
}
if let Some(patch) = &self.metadata {
let mut map = opts.metadata.take().unwrap_or_default();
for (k, v) in patch {
match v {
Some(val) => map.insert(k.clone(), val.clone()),
None => map.remove(k),
};
}
opts.metadata = if map.is_empty() { None } else { Some(map) };
}
if let Some(v) = self.cache_retention {
opts.cache_retention = v.unwrap_or_default();
}
if let Some(v) = self.max_tokens {
opts.max_tokens = v;
}
if let Some(v) = self.temperature {
opts.temperature = v;
}
if let Some(v) = &self.session_id {
opts.session_id = v.clone();
}
}
}