Skip to main content

elph_ai/models/
collection.rs

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    /// Consume the mutable wrapper and share the underlying [`Models`] via [`Arc`].
413    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                // xhigh/max are opt-in via thinkingLevelMap; other levels default on.
477                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}