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;
pub const DEFAULT_LLM_API_TIMEOUT: Duration = Duration::from_secs(600);
#[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: Some(DEFAULT_LLM_API_TIMEOUT),
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)
}
pub fn request_timeout(&self) -> Duration {
self.timeout.unwrap_or(DEFAULT_LLM_API_TIMEOUT)
}
}
#[cfg(test)]
mod tests {
use super::{SimpleStreamOptions, DEFAULT_LLM_API_TIMEOUT};
use std::time::Duration;
#[test]
fn llm_api_timeout_defaults_to_ten_minutes() {
let options = SimpleStreamOptions::default();
assert_eq!(options.timeout, Some(Duration::from_secs(600)));
assert_eq!(options.request_timeout(), DEFAULT_LLM_API_TIMEOUT);
}
#[test]
fn llm_api_timeout_allows_override_and_none_uses_provider_fallback() {
let mut options = SimpleStreamOptions::default();
options.timeout = Some(Duration::from_secs(15));
assert_eq!(options.request_timeout(), Duration::from_secs(15));
options.timeout = None;
assert_eq!(options.request_timeout(), Duration::from_secs(600));
}
}
#[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();
}
}
}