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
13pub type LlmResponseStream = Pin<Box<dyn Stream<Item = LlmResult<LlmResponse>> + Send>>;
20
21#[doc = include_str!("docs/provider_factory.md")]
22pub trait ProviderFactory: Sized {
23 fn from_env() -> impl Future<Output = LlmResult<Self>> + Send;
25
26 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 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 fn context_window(&self) -> Option<u32>;
46
47 fn model(&self) -> Option<LlmModel> {
51 None
52 }
53}
54
55pub 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}