1use std::collections::HashMap;
2use std::future::Future;
3use std::pin::Pin;
4use std::sync::Arc;
5
6use crate::auth::ProviderAuthHolder;
7use crate::auth::resolve::resolve_provider_auth;
8use crate::auth::resolve::{AuthResolutionOverrides, ModelsError, ModelsErrorCode};
9use crate::auth::{AuthContext, AuthModel, AuthResult, CredentialStore, InMemoryCredentialStore, ProviderAuth};
10use crate::types::{AssistantMessage, Context, Model, ProviderHeaders, SimpleStreamOptions, StreamOptions};
11use crate::utils::event_stream::AssistantMessageEventStream;
12
13pub trait ProviderStreamsDyn: Send + Sync {
14 fn stream(&self, model: &Model, context: &Context, options: Option<StreamOptions>) -> AssistantMessageEventStream;
15
16 fn stream_simple(
17 &self,
18 model: &Model,
19 context: &Context,
20 options: Option<SimpleStreamOptions>,
21 ) -> AssistantMessageEventStream;
22}
23
24pub enum ProviderApi {
25 Single(Arc<dyn ProviderStreamsDyn>),
26 Map(HashMap<String, Arc<dyn ProviderStreamsDyn>>),
27}
28
29pub struct Provider {
30 pub id: String,
31 pub name: String,
32 pub base_url: Option<String>,
33 pub headers: Option<ProviderHeaders>,
34 pub auth: ProviderAuth,
35 models: Vec<Model>,
36 refresh: Option<RefreshFn>,
37 api: ProviderApi,
38}
39
40type RefreshFn = Arc<dyn Fn() -> Pin<Box<dyn Future<Output = anyhow::Result<Vec<Model>>> + Send>> + Send + Sync>;
41
42impl Provider {
43 pub fn get_models(&self) -> &[Model] {
44 &self.models
45 }
46
47 pub fn stream(
48 &self,
49 model: &Model,
50 context: &Context,
51 options: Option<StreamOptions>,
52 ) -> AssistantMessageEventStream {
53 self.dispatch(model, |streams| streams.stream(model, context, options))
54 }
55
56 pub fn stream_simple(
57 &self,
58 model: &Model,
59 context: &Context,
60 options: Option<SimpleStreamOptions>,
61 ) -> AssistantMessageEventStream {
62 self.dispatch(model, |streams| streams.stream_simple(model, context, options))
63 }
64
65 pub async fn refresh_models(&self) -> Result<(), ModelsError> {
66 let Some(refresh) = &self.refresh else {
67 return Ok(());
68 };
69 refresh().await.map_err(|e| {
70 ModelsError::with_cause(ModelsErrorCode::ModelSource, format!("Model refresh failed for {}", self.id), e)
71 })?;
72 Ok(())
73 }
74
75 fn api_for(&self, model: &Model) -> Option<Arc<dyn ProviderStreamsDyn>> {
76 match &self.api {
77 ProviderApi::Single(api) => Some(api.clone()),
78 ProviderApi::Map(map) => map.get(&model.api).cloned(),
79 }
80 }
81
82 fn dispatch(
83 &self,
84 model: &Model,
85 run: impl FnOnce(Arc<dyn ProviderStreamsDyn>) -> AssistantMessageEventStream,
86 ) -> AssistantMessageEventStream {
87 match self.api_for(model) {
88 Some(api) => run(api),
89 None => AssistantMessageEventStream::failed(format!(
90 "Provider {} has no API implementation for \"{}\"",
91 self.id, model.api
92 )),
93 }
94 }
95}
96
97pub struct CreateProviderOptions {
98 pub id: String,
99 pub name: Option<String>,
100 pub base_url: Option<String>,
101 pub headers: Option<ProviderHeaders>,
102 pub auth: ProviderAuth,
103 pub models: Vec<Model>,
104 pub refresh_models: Option<RefreshFn>,
105 pub api: ProviderApi,
106}
107
108pub fn create_provider(input: CreateProviderOptions) -> Provider {
109 let id = input.id.clone();
110 Provider {
111 id: input.id,
112 name: input.name.unwrap_or(id),
113 base_url: input.base_url,
114 headers: input.headers,
115 auth: input.auth,
116 models: input.models,
117 refresh: input.refresh_models,
118 api: input.api,
119 }
120}
121
122pub struct CreateModelsOptions {
123 pub credentials: Option<Arc<dyn CredentialStore>>,
124 pub auth_context: Option<Arc<dyn AuthContext>>,
125}
126
127pub struct Models {
128 providers: HashMap<String, Provider>,
129 credentials: Arc<dyn CredentialStore>,
130 auth_context: Arc<dyn AuthContext>,
131}
132
133pub struct MutableModels {
134 inner: Models,
135}
136
137impl Models {
138 pub fn get_providers(&self) -> Vec<&Provider> {
139 self.providers.values().collect()
140 }
141
142 pub fn get_provider(&self, id: &str) -> Option<&Provider> {
143 self.providers.get(id)
144 }
145
146 pub fn get_models(&self, provider: Option<&str>) -> Vec<Model> {
147 match provider {
148 Some(id) => self
149 .providers
150 .get(id)
151 .map(|p| p.get_models().to_vec())
152 .unwrap_or_default(),
153 None => self
154 .providers
155 .values()
156 .flat_map(|p| p.get_models().iter().cloned())
157 .collect(),
158 }
159 }
160
161 pub fn get_model(&self, provider: &str, id: &str) -> Option<Model> {
162 self.get_models(Some(provider)).into_iter().find(|m| m.id == id)
163 }
164
165 pub async fn refresh(&self, provider: Option<&str>) -> Result<(), ModelsError> {
166 match provider {
167 Some(id) => {
168 let p = self
169 .providers
170 .get(id)
171 .ok_or_else(|| ModelsError::new(ModelsErrorCode::Provider, format!("Unknown provider: {id}")))?;
172 p.refresh_models().await
173 }
174 None => {
175 let mut errors = vec![];
176 for p in self.providers.values() {
177 if let Err(e) = p.refresh_models().await {
178 errors.push(e);
179 }
180 }
181 if let Some(e) = errors.into_iter().next() {
182 return Err(e);
183 }
184 Ok(())
185 }
186 }
187 }
188
189 pub async fn get_auth(&self, model: &Model) -> Result<Option<AuthResult>, ModelsError> {
190 let provider = self.providers.get(&model.provider).ok_or_else(|| {
191 ModelsError::new(ModelsErrorCode::Provider, format!("Unknown provider: {}", model.provider))
192 })?;
193 resolve_provider_auth(
194 &ProviderAuthHolder {
195 id: provider.id.clone(),
196 auth: provider.auth.clone(),
197 },
198 AuthModel::Chat(model.clone()),
199 self.credentials.as_ref(),
200 self.auth_context.clone(),
201 None,
202 )
203 .await
204 }
205
206 pub fn stream(
207 &self,
208 model: &Model,
209 context: &Context,
210 options: Option<StreamOptions>,
211 ) -> AssistantMessageEventStream {
212 let inner = self.clone_for_stream();
213 let model = model.clone();
214 let context = context.clone();
215 lazy_stream(model.clone(), move || async move {
216 let provider = inner.require_provider(&model)?;
217 let (request_model, request_options) = inner.apply_auth(&model, options).await?;
218 Ok(provider.stream(&request_model, &context, request_options))
219 })
220 }
221
222 pub async fn complete(&self, model: &Model, context: &Context, options: Option<StreamOptions>) -> AssistantMessage {
223 self.stream(model, context, options).result().await
224 }
225
226 pub fn stream_simple(
227 &self,
228 model: &Model,
229 context: &Context,
230 options: Option<SimpleStreamOptions>,
231 ) -> AssistantMessageEventStream {
232 let inner = self.clone_for_stream();
233 let model = model.clone();
234 let context = context.clone();
235 lazy_stream(model.clone(), move || async move {
236 let provider = inner.require_provider(&model)?;
237 let (request_model, request_options) = inner.apply_auth_simple(&model, options).await?;
238 Ok(provider.stream_simple(&request_model, &context, request_options))
239 })
240 }
241
242 pub async fn complete_simple(
243 &self,
244 model: &Model,
245 context: &Context,
246 options: Option<SimpleStreamOptions>,
247 ) -> AssistantMessage {
248 self.stream_simple(model, context, options).result().await
249 }
250
251 fn clone_for_stream(&self) -> Models {
252 Models {
253 providers: self.providers.clone(),
254 credentials: self.credentials.clone(),
255 auth_context: self.auth_context.clone(),
256 }
257 }
258
259 fn require_provider(&self, model: &Model) -> Result<&Provider, ModelsError> {
260 self.providers
261 .get(&model.provider)
262 .ok_or_else(|| ModelsError::new(ModelsErrorCode::Provider, format!("Unknown provider: {}", model.provider)))
263 }
264
265 async fn apply_auth(
266 &self,
267 model: &Model,
268 options: Option<StreamOptions>,
269 ) -> Result<(Model, Option<StreamOptions>), ModelsError> {
270 let provider = self.require_provider(model)?;
271 let overrides = options.as_ref().map(|o| AuthResolutionOverrides {
272 api_key: o.api_key.clone(),
273 env: o.env.clone(),
274 });
275 let resolution = resolve_provider_auth(
276 &ProviderAuthHolder {
277 id: provider.id.clone(),
278 auth: provider.auth.clone(),
279 },
280 AuthModel::Chat(model.clone()),
281 self.credentials.as_ref(),
282 self.auth_context.clone(),
283 overrides,
284 )
285 .await?;
286 Ok(merge_auth(model, options, resolution, provider))
287 }
288
289 async fn apply_auth_simple(
290 &self,
291 model: &Model,
292 options: Option<SimpleStreamOptions>,
293 ) -> Result<(Model, Option<SimpleStreamOptions>), ModelsError> {
294 let stream_opts = options.as_ref().map(|o| o.base.clone());
295 let (request_model, stream_opts) = self.apply_auth(model, stream_opts).await?;
296 let request_options = stream_opts.map(|base| SimpleStreamOptions {
297 base,
298 reasoning: options.as_ref().and_then(|o| o.reasoning),
299 thinking_budgets: options.as_ref().and_then(|o| o.thinking_budgets.clone()),
300 });
301 Ok((request_model, request_options))
302 }
303}
304
305impl Clone for Provider {
306 fn clone(&self) -> Self {
307 Self {
308 id: self.id.clone(),
309 name: self.name.clone(),
310 base_url: self.base_url.clone(),
311 headers: self.headers.clone(),
312 auth: self.auth.clone(),
313 models: self.models.clone(),
314 refresh: self.refresh.clone(),
315 api: match &self.api {
316 ProviderApi::Single(s) => ProviderApi::Single(s.clone()),
317 ProviderApi::Map(m) => ProviderApi::Map(m.clone()),
318 },
319 }
320 }
321}
322
323fn merge_auth(
324 model: &Model,
325 options: Option<StreamOptions>,
326 resolution: Option<AuthResult>,
327 provider: &Provider,
328) -> (Model, Option<StreamOptions>) {
329 let mut request_model = model.clone();
330 let mut request_options = options.unwrap_or_default();
331
332 if let Some(res) = resolution {
333 if let Some(url) = res.auth.base_url {
334 request_model.base_url = url;
335 }
336 if request_options.api_key.is_none() {
337 request_options.api_key = res.auth.api_key;
338 }
339 if let Some(headers) = res.auth.headers {
340 let mut merged = provider.headers.clone().unwrap_or_default();
341 merged.extend(headers);
342 if let Some(opts) = &request_options.headers {
343 merged.extend(opts.clone());
344 }
345 request_options.headers = Some(merged);
346 }
347 if let Some(env) = res.env {
348 let mut merged = request_options.env.unwrap_or_default();
349 merged.extend(env);
350 request_options.env = Some(merged);
351 }
352 }
353
354 (request_model, Some(request_options))
355}
356
357fn lazy_stream<F, Fut>(model: Model, setup: F) -> AssistantMessageEventStream
358where
359 F: FnOnce() -> Fut + Send + 'static,
360 Fut: Future<Output = Result<AssistantMessageEventStream, ModelsError>> + Send + 'static,
361{
362 let stream = AssistantMessageEventStream::new();
363 let output = stream.clone_handle();
364 let trace_model = model.clone();
365 crate::trace::spawn_stream(&trace_model, async move {
366 match setup().await {
367 Ok(mut inner) => {
368 while let Some(event) = inner.next_event().await {
369 let terminal = matches!(
370 &event,
371 crate::types::AssistantMessageEvent::Done { .. }
372 | crate::types::AssistantMessageEvent::Error { .. }
373 );
374 output.push(event);
375 if terminal {
376 break;
377 }
378 }
379 }
380 Err(e) => {
381 let mut partial = crate::types::AssistantMessage::empty(&model);
382 partial.stop_reason = crate::types::StopReason::Error;
383 partial.error_message = Some(e.message);
384 output.push(crate::types::AssistantMessageEvent::Error {
385 reason: crate::types::StopReason::Error,
386 error: partial,
387 });
388 }
389 }
390 output.end();
391 });
392 stream
393}
394
395pub fn create_models(options: Option<CreateModelsOptions>) -> MutableModels {
396 MutableModels {
397 inner: Models {
398 providers: HashMap::new(),
399 credentials: options
400 .as_ref()
401 .and_then(|o| o.credentials.clone())
402 .unwrap_or_else(|| Arc::new(InMemoryCredentialStore::new())),
403 auth_context: options
404 .as_ref()
405 .and_then(|o| o.auth_context.clone())
406 .unwrap_or_else(|| Arc::new(crate::auth::DefaultAuthContext::new())),
407 },
408 }
409}
410
411impl MutableModels {
412 pub fn into_arc(self) -> std::sync::Arc<Models> {
414 std::sync::Arc::new(self.inner)
415 }
416
417 pub fn set_provider(&mut self, provider: Provider) {
418 self.inner.providers.insert(provider.id.clone(), provider);
419 }
420
421 pub fn delete_provider(&mut self, id: &str) {
422 self.inner.providers.remove(id);
423 }
424
425 pub fn clear_providers(&mut self) {
426 self.inner.providers.clear();
427 }
428
429 pub fn inner(&self) -> &Models {
430 &self.inner
431 }
432
433 pub fn inner_mut(&mut self) -> &mut Models {
434 &mut self.inner
435 }
436}
437
438impl std::ops::Deref for MutableModels {
439 type Target = Models;
440 fn deref(&self) -> &Self::Target {
441 &self.inner
442 }
443}
444
445pub fn has_api(model: &Model, api: &str) -> bool {
446 model.api == api
447}
448
449pub fn models_are_equal(a: Option<&Model>, b: Option<&Model>) -> bool {
450 match (a, b) {
451 (Some(a), Some(b)) => a.id == b.id && a.provider == b.provider,
452 _ => false,
453 }
454}
455
456pub fn get_supported_thinking_levels(model: &Model) -> Vec<crate::types::ThinkingLevel> {
457 if !model.reasoning {
458 return vec![];
459 }
460 let levels = [
461 crate::types::ThinkingLevel::Minimal,
462 crate::types::ThinkingLevel::Low,
463 crate::types::ThinkingLevel::Medium,
464 crate::types::ThinkingLevel::High,
465 crate::types::ThinkingLevel::Xhigh,
466 crate::types::ThinkingLevel::Max,
467 ];
468 levels
469 .into_iter()
470 .filter(|level| {
471 if let Some(map) = &model.thinking_level_map {
472 let key = crate::models::thinking_level_to_str(*level);
473 if map.get(key) == Some(&None) {
474 return false;
475 }
476 if matches!(level, crate::types::ThinkingLevel::Xhigh | crate::types::ThinkingLevel::Max) {
478 return map.contains_key(key);
479 }
480 } else if matches!(level, crate::types::ThinkingLevel::Xhigh | crate::types::ThinkingLevel::Max) {
481 return false;
482 }
483 true
484 })
485 .collect()
486}
487
488pub fn clamp_thinking_level(model: &Model, level: crate::types::ThinkingLevel) -> crate::types::ThinkingLevel {
489 let available = get_supported_thinking_levels(model);
490 if available.contains(&level) {
491 return level;
492 }
493 let all = [
494 crate::types::ThinkingLevel::Minimal,
495 crate::types::ThinkingLevel::Low,
496 crate::types::ThinkingLevel::Medium,
497 crate::types::ThinkingLevel::High,
498 crate::types::ThinkingLevel::Xhigh,
499 crate::types::ThinkingLevel::Max,
500 ];
501 let idx = all.iter().position(|l| *l == level).unwrap_or(0);
502 for &candidate in &all[idx..] {
503 if available.contains(&candidate) {
504 return candidate;
505 }
506 }
507 for &candidate in all[..idx].iter().rev() {
508 if available.contains(&candidate) {
509 return candidate;
510 }
511 }
512 crate::types::ThinkingLevel::High
513}