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}