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