Skip to main content

llm/
provider.rs

1use crate::LlmError;
2use crate::LlmModel;
3use crate::ProviderConnectionConfig;
4use crate::Result as LlmResult;
5use crate::catalog::ReasoningEffortError;
6use std::future::Future;
7use std::pin::Pin;
8use tokio_stream::Stream;
9use utils::ReasoningEffort;
10
11use super::{Context, LlmResponse};
12
13/// A stream of [`LlmResponse`] events from an LLM provider.
14///
15/// This is a pinned, boxed, `Send` stream used as the return type of
16/// [`StreamingModelProvider::stream_response`]. Boxing is required to support
17/// trait objects (`Vec<Box<dyn StreamingModelProvider>>`) in types like
18/// [`AlloyedModelProvider`](crate::alloyed::AlloyedModelProvider).
19pub type LlmResponseStream = Pin<Box<dyn Stream<Item = LlmResult<LlmResponse>> + Send>>;
20
21#[doc = include_str!("docs/provider_factory.md")]
22pub trait ProviderFactory: Sized {
23    /// Create provider from environment variables and default configuration
24    fn from_env() -> impl Future<Output = LlmResult<Self>> + Send;
25
26    /// Create provider from environment variables with provider connection overrides.
27    fn from_env_with_connection(connection: ProviderConnectionConfig) -> impl Future<Output = LlmResult<Self>> + Send {
28        async move {
29            let _ = connection;
30            Self::from_env().await
31        }
32    }
33
34    /// Set or update the model for this provider (builder pattern)
35    fn with_model(self, model: &str) -> Self;
36}
37
38#[doc = include_str!("docs/streaming_model_provider.md")]
39pub trait StreamingModelProvider: Send + Sync {
40    fn stream_response(&self, context: &Context) -> LlmResponseStream;
41    fn display_name(&self) -> String;
42
43    /// Context window size in tokens for the current model.
44    /// Returns `None` for unknown models (e.g. Ollama, `LlamaCpp`).
45    fn context_window(&self) -> Option<u32>;
46
47    /// The `LlmModel` this provider is currently configured to use.
48    /// Returns `None` for providers where the model is unknown at compile time
49    /// (e.g. test fakes).
50    fn model(&self) -> Option<LlmModel> {
51        None
52    }
53}
54
55/// Look up context window for a known provider + model ID combo via the catalog.
56///
57/// Returns `None` if the model is not in the catalog.
58pub fn get_context_window(provider: &str, model_id: &str) -> Option<u32> {
59    let key = format!("{provider}:{model_id}");
60    key.parse::<LlmModel>().ok().and_then(|m| m.context_window())
61}
62
63pub(crate) fn validate_reasoning(context: &Context, model: Option<&LlmModel>) -> LlmResult<()> {
64    if context.reasoning_effort() != ReasoningEffort::Disabled {
65        return Ok(());
66    }
67
68    let model = model.ok_or_else(|| ReasoningEffortError::Unsupported {
69        model: "unknown".to_string(),
70        effort: ReasoningEffort::Disabled,
71        supported: Vec::new(),
72    })?;
73
74    if !model.supports_reasoning_off() {
75        model.validate_reasoning_effort(ReasoningEffort::Disabled)?;
76    }
77
78    if !model.supports_reasoning_off_transport() {
79        return Err(LlmError::UnsupportedDisableTransport { model: model.to_string() });
80    }
81
82    Ok(())
83}
84
85impl StreamingModelProvider for Box<dyn StreamingModelProvider> {
86    fn stream_response(&self, context: &Context) -> LlmResponseStream {
87        (**self).stream_response(context)
88    }
89
90    fn display_name(&self) -> String {
91        (**self).display_name()
92    }
93
94    fn context_window(&self) -> Option<u32> {
95        (**self).context_window()
96    }
97
98    fn model(&self) -> Option<LlmModel> {
99        (**self).model()
100    }
101}
102
103impl<T: StreamingModelProvider + ?Sized> StreamingModelProvider for std::sync::Arc<T> {
104    fn stream_response(&self, context: &Context) -> LlmResponseStream {
105        (**self).stream_response(context)
106    }
107
108    fn display_name(&self) -> String {
109        (**self).display_name()
110    }
111
112    fn context_window(&self) -> Option<u32> {
113        (**self).context_window()
114    }
115
116    fn model(&self) -> Option<LlmModel> {
117        (**self).model()
118    }
119}
120
121#[cfg(test)]
122mod tests {
123    use super::*;
124
125    #[test]
126    fn lookup_context_window_known_model() {
127        assert_eq!(get_context_window("anthropic", "claude-opus-4-6"), Some(1_000_000));
128    }
129
130    #[test]
131    fn lookup_context_window_openrouter_model() {
132        let model = LlmModel::all()
133            .iter()
134            .find(|model| model.provider() == "openrouter" && model.context_window().is_some())
135            .expect("OpenRouter catalog should contain a model with a context window");
136
137        assert_eq!(get_context_window(model.provider(), &model.model_id()), model.context_window());
138    }
139
140    #[test]
141    fn lookup_context_window_unknown_model() {
142        assert_eq!(get_context_window("anthropic", "unknown-model-xyz"), None);
143    }
144
145    #[test]
146    fn lookup_context_window_unknown_provider() {
147        assert_eq!(get_context_window("unknown-provider", "some-model"), None);
148    }
149}