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::{
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}