Skip to main content

elph_ai/images/
collection.rs

1use std::collections::HashMap;
2use std::sync::Arc;
3
4use crate::api::OpenRouterImagesApi;
5use crate::auth::{
6    AuthContext, AuthModel, AuthResolutionOverrides, AuthResult, CredentialStore, InMemoryCredentialStore, ModelsError,
7    ModelsErrorCode, ProviderAuth, ProviderAuthHolder, env_api_key_auth, resolve_provider_auth,
8};
9use crate::images::models::OPENROUTER_IMAGE_MODELS;
10use crate::types::{AssistantImages, ImagesContext, ImagesModel, ImagesOptions, ProviderImages};
11
12pub struct ImagesProvider {
13    pub id: String,
14    pub name: String,
15    pub auth: ProviderAuth,
16    models: Vec<ImagesModel>,
17    api: Arc<dyn ProviderImages>,
18}
19
20impl ImagesProvider {
21    pub fn get_models(&self) -> &[ImagesModel] {
22        &self.models
23    }
24}
25
26pub struct CreateImagesProviderOptions {
27    pub id: String,
28    pub name: Option<String>,
29    pub auth: ProviderAuth,
30    pub models: Vec<ImagesModel>,
31    pub api: Arc<dyn ProviderImages>,
32}
33
34pub fn create_images_provider(input: CreateImagesProviderOptions) -> ImagesProvider {
35    let id = input.id.clone();
36    ImagesProvider {
37        id: input.id,
38        name: input.name.unwrap_or(id),
39        auth: input.auth,
40        models: input.models,
41        api: input.api,
42    }
43}
44
45pub struct CreateImagesModelsOptions {
46    pub credentials: Option<Arc<dyn CredentialStore>>,
47    pub auth_context: Option<Arc<dyn AuthContext>>,
48}
49
50pub struct ImagesModels {
51    providers: HashMap<String, ImagesProvider>,
52    credentials: Arc<dyn CredentialStore>,
53    auth_context: Arc<dyn AuthContext>,
54}
55
56pub struct MutableImagesModels {
57    inner: ImagesModels,
58}
59
60impl ImagesModels {
61    pub fn get_providers(&self) -> Vec<&ImagesProvider> {
62        self.providers.values().collect()
63    }
64
65    pub fn get_provider(&self, id: &str) -> Option<&ImagesProvider> {
66        self.providers.get(id)
67    }
68
69    pub fn get_models(&self, provider: Option<&str>) -> Vec<ImagesModel> {
70        match provider {
71            Some(id) => self
72                .providers
73                .get(id)
74                .map(|p| p.get_models().to_vec())
75                .unwrap_or_default(),
76            None => self
77                .providers
78                .values()
79                .flat_map(|p| p.get_models().iter().cloned())
80                .collect(),
81        }
82    }
83
84    pub fn get_model(&self, provider: &str, id: &str) -> Option<ImagesModel> {
85        self.get_models(Some(provider)).into_iter().find(|m| m.id == id)
86    }
87
88    pub async fn get_auth(&self, model: &ImagesModel) -> Result<Option<AuthResult>, ModelsError> {
89        let provider = self.providers.get(&model.provider).ok_or_else(|| {
90            ModelsError::new(
91                ModelsErrorCode::Provider,
92                format!("Unknown provider: {}", model.provider),
93            )
94        })?;
95        resolve_provider_auth(
96            &ProviderAuthHolder {
97                id: provider.id.clone(),
98                auth: provider.auth.clone(),
99            },
100            AuthModel::Images(model.clone()),
101            self.credentials.as_ref(),
102            self.auth_context.clone(),
103            None,
104        )
105        .await
106    }
107
108    pub async fn generate_images(
109        &self,
110        model: &ImagesModel,
111        context: &ImagesContext,
112        options: Option<ImagesOptions>,
113    ) -> AssistantImages {
114        let Some(provider) = self.providers.get(&model.provider) else {
115            return AssistantImages {
116                api: model.api.clone(),
117                provider: model.provider.clone(),
118                model: model.id.clone(),
119                output: vec![],
120                response_id: None,
121                usage: None,
122                stop_reason: crate::types::StopReason::Error,
123                error_message: Some(format!("Unknown provider: {}", model.provider)),
124                timestamp: chrono::Utc::now().timestamp_millis(),
125            };
126        };
127
128        let overrides = options.as_ref().map(|o| AuthResolutionOverrides {
129            api_key: o.api_key.clone(),
130            env: o.env.clone(),
131        });
132        let resolution = match resolve_provider_auth(
133            &ProviderAuthHolder {
134                id: provider.id.clone(),
135                auth: provider.auth.clone(),
136            },
137            AuthModel::Images(model.clone()),
138            self.credentials.as_ref(),
139            self.auth_context.clone(),
140            overrides,
141        )
142        .await
143        {
144            Ok(r) => r,
145            Err(e) => {
146                return AssistantImages {
147                    api: model.api.clone(),
148                    provider: model.provider.clone(),
149                    model: model.id.clone(),
150                    output: vec![],
151                    response_id: None,
152                    usage: None,
153                    stop_reason: crate::types::StopReason::Error,
154                    error_message: Some(e.message),
155                    timestamp: chrono::Utc::now().timestamp_millis(),
156                };
157            }
158        };
159
160        let mut opts = options.unwrap_or(ImagesOptions {
161            api_key: None,
162            signal: None,
163            env: None,
164            headers: None,
165            timeout_ms: None,
166            max_retries: None,
167            on_payload: None,
168            on_response: None,
169        });
170        if let Some(res) = resolution {
171            if opts.api_key.is_none() {
172                opts.api_key = res.auth.api_key;
173            }
174            if let Some(headers) = res.auth.headers {
175                let mut merged = headers;
176                if let Some(request_headers) = opts.headers.take() {
177                    merged.extend(request_headers);
178                }
179                opts.headers = Some(merged);
180            }
181            if res.env.is_some() || opts.env.is_some() {
182                let mut merged = res.env.unwrap_or_default();
183                if let Some(request_env) = opts.env.take() {
184                    merged.extend(request_env);
185                }
186                opts.env = Some(merged);
187            }
188        }
189
190        provider.api.generate_images(model, context, Some(opts)).await
191    }
192}
193
194pub fn create_images_models(options: Option<CreateImagesModelsOptions>) -> MutableImagesModels {
195    MutableImagesModels {
196        inner: ImagesModels {
197            providers: HashMap::new(),
198            credentials: options
199                .as_ref()
200                .and_then(|o| o.credentials.clone())
201                .unwrap_or_else(|| Arc::new(InMemoryCredentialStore::new())),
202            auth_context: options
203                .as_ref()
204                .and_then(|o| o.auth_context.clone())
205                .unwrap_or_else(|| Arc::new(crate::auth::DefaultAuthContext::new())),
206        },
207    }
208}
209
210impl MutableImagesModels {
211    pub fn set_provider(&mut self, provider: ImagesProvider) {
212        self.inner.providers.insert(provider.id.clone(), provider);
213    }
214}
215
216impl std::ops::Deref for MutableImagesModels {
217    type Target = ImagesModels;
218    fn deref(&self) -> &Self::Target {
219        &self.inner
220    }
221}
222
223pub fn openrouter_images_provider() -> ImagesProvider {
224    create_images_provider(CreateImagesProviderOptions {
225        id: "openrouter".to_string(),
226        name: Some("OpenRouter".to_string()),
227        auth: ProviderAuth {
228            api_key: Some(env_api_key_auth("OpenRouter API key", vec!["OPENROUTER_API_KEY"])),
229            oauth: None,
230        },
231        models: OPENROUTER_IMAGE_MODELS.to_vec(),
232        api: Arc::new(OpenRouterImagesApi),
233    })
234}
235
236pub fn builtin_images_models(options: Option<CreateImagesModelsOptions>) -> MutableImagesModels {
237    let mut models = create_images_models(options);
238    models.set_provider(openrouter_images_provider());
239    models
240}
241
242pub async fn generate_images(
243    model: &ImagesModel,
244    context: &ImagesContext,
245    options: Option<ImagesOptions>,
246) -> AssistantImages {
247    let collection = builtin_images_models(None);
248    collection.generate_images(model, context, options).await
249}