1use std::collections::{BTreeMap, BTreeSet, HashMap};
9use std::path::{Path, PathBuf};
10use std::sync::{Arc, Mutex};
11
12use futures::future::FutureExt;
13use futures::stream::{self, BoxStream};
14use pi_ai::auth::config_value::{
15 is_config_value_configured, resolve_config_value, resolve_headers,
16};
17use pi_ai::auth::context::MapAuthContext;
18use pi_ai::auth::resolve::resolve_provider_auth_with_signal;
19use pi_ai::auth::{
20 AMBIENT_AUTH_MARKER, AuthCheck, AuthContext, AuthResolutionOverrides, AuthResult, AuthType,
21 Credential, CredentialInfo, CredentialStore, FileCredentialStore, InMemoryCredentialStore,
22 ModelAuth, ModelsError, ModelsErrorCode, OAuthAuth, ProviderAuth, ProviderEnv, ProviderHeaders,
23 RuntimeCredentials, api_key_env_vars, env_api_key_auth, get_env_api_key,
24};
25use pi_ai::catalog::{BuiltinModels, ModelsStoreEntry, builtin_models};
26use pi_ai::models_store::{
27 FileModelsStore, InMemoryModelsStore, ModelOverrides, ModelsStore, apply_model_overrides,
28 compose_provider_models, models_error_from_catalog,
29};
30use pi_ai::provider::{Provider, ProviderError, StreamOptions};
31use pi_ai::providers::{
32 AnthropicMessages, AzureOpenAiResponses, BedrockConverseStream, DefaultBedrockClientFactory,
33 GoogleGenerativeAi, GoogleVertex, KnownProvider, MistralConversations, OpenAiCodexResponses,
34 OpenAiCompletions, OpenAiResponses, PiMessages, ProviderRegistry,
35};
36use pi_ai::types::{
37 AssistantMessageEvent, Context, Model, ModelCost, ModelInput, ModelThinkingLevel,
38};
39use serde::Deserialize;
40use serde_json::Value;
41use thiserror::Error;
42
43#[derive(Clone, Default)]
45pub struct CreateModelRuntimeOptions {
46 pub credentials: Option<Arc<dyn CredentialStore>>,
48 pub auth_path: Option<PathBuf>,
50 pub models_path: Option<PathBuf>,
54 pub models_config: Option<ModelsJsonConfig>,
56 pub models_store: Option<Arc<dyn ModelsStore>>,
58 pub models_store_path: Option<PathBuf>,
60 pub allow_model_network: Option<bool>,
63 pub auth_env: Option<ProviderEnv>,
65 pub oauth_handlers: Option<HashMap<String, Arc<dyn OAuthAuth>>>,
68}
69
70#[derive(Clone, Debug, Default)]
72pub struct ModelRuntimeAuthOverrides {
73 pub api_key: Option<String>,
75 pub env: Option<ProviderEnv>,
77}
78
79#[derive(Clone, Debug, Default, Deserialize, PartialEq)]
85#[serde(rename_all = "camelCase")]
86pub struct ProviderConfigInput {
87 #[serde(default)]
89 pub name: Option<String>,
90 #[serde(default)]
92 pub base_url: Option<String>,
93 #[serde(default)]
95 pub api_key: Option<String>,
96 #[serde(default)]
98 pub api: Option<String>,
99 #[serde(default)]
101 pub headers: Option<BTreeMap<String, String>>,
102 #[serde(default)]
104 pub auth_header: Option<bool>,
105 #[serde(default)]
107 pub models: Option<Vec<ProviderModelDefinition>>,
108 #[serde(default)]
110 pub model_overrides: Option<ModelOverrides>,
111 #[serde(default)]
114 pub oauth: Option<String>,
115}
116
117#[derive(Clone, Debug, Deserialize, PartialEq)]
119#[serde(rename_all = "camelCase")]
120pub struct ProviderModelDefinition {
121 pub id: String,
123 #[serde(default)]
125 pub name: Option<String>,
126 #[serde(default)]
128 pub api: Option<String>,
129 #[serde(default)]
131 pub base_url: Option<String>,
132 #[serde(default)]
134 pub reasoning: bool,
135 #[serde(default)]
137 pub thinking_level_map: Option<BTreeMap<ModelThinkingLevel, Option<String>>>,
138 #[serde(default)]
140 pub input: Option<Vec<ModelInput>>,
141 #[serde(default)]
143 pub cost: Option<ModelCost>,
144 #[serde(default)]
146 pub context_window: Option<u64>,
147 #[serde(default)]
149 pub max_tokens: Option<u64>,
150 #[serde(default)]
152 pub headers: Option<BTreeMap<String, String>>,
153 #[serde(default)]
155 pub compat: Option<Value>,
156}
157
158#[derive(Clone, Debug, Default)]
160pub struct ModelsJsonConfig {
161 providers: BTreeMap<String, ProviderConfigInput>,
162 error: Option<String>,
163 path: Option<PathBuf>,
164}
165
166impl ModelsJsonConfig {
167 #[must_use]
169 pub fn empty() -> Self {
170 Self::default()
171 }
172
173 #[must_use]
175 pub fn from_providers(providers: BTreeMap<String, ProviderConfigInput>) -> Self {
176 Self {
177 providers,
178 error: None,
179 path: None,
180 }
181 }
182
183 #[must_use]
191 pub fn load(path: Option<&Path>) -> Self {
192 let Some(path) = path else {
193 return Self::empty();
194 };
195 let path_display = path.display().to_string();
196 let content = match std::fs::read_to_string(path) {
197 Ok(content) => content,
198 Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
199 return Self {
200 providers: BTreeMap::new(),
201 error: None,
202 path: Some(path.to_path_buf()),
203 };
204 }
205 Err(error) => {
206 return Self {
207 providers: BTreeMap::new(),
208 error: Some(format!(
209 "Failed to load models.json: {error}\n\nFile: {path_display}"
210 )),
211 path: Some(path.to_path_buf()),
212 };
213 }
214 };
215 match parse_models_json(&content, &path_display) {
216 Ok(providers) => Self {
217 providers,
218 error: None,
219 path: Some(path.to_path_buf()),
220 },
221 Err(message) => Self {
222 providers: BTreeMap::new(),
223 error: Some(message),
224 path: Some(path.to_path_buf()),
225 },
226 }
227 }
228
229 #[must_use]
231 pub fn provider_ids(&self) -> Vec<String> {
232 self.providers.keys().cloned().collect()
233 }
234
235 #[must_use]
237 pub fn get_provider(&self, provider_id: &str) -> Option<&ProviderConfigInput> {
238 self.providers.get(provider_id)
239 }
240
241 #[must_use]
243 pub fn error(&self) -> Option<&str> {
244 self.error.as_deref()
245 }
246
247 #[must_use]
249 pub fn path(&self) -> Option<&Path> {
250 self.path.as_deref()
251 }
252}
253
254#[derive(Clone, Debug, Error)]
256pub enum ModelRuntimeError {
257 #[error(transparent)]
259 Models(#[from] ModelsError),
260 #[error("{0}")]
262 Registration(String),
263}
264
265#[derive(Clone, Debug, Default, PartialEq, Eq)]
267pub struct ModelsRefreshResult {
268 pub aborted: bool,
270 pub errors: BTreeMap<String, String>,
272}
273
274#[derive(Clone, Debug, Default)]
276pub struct ModelsRefreshOptions {
277 pub allow_network: Option<bool>,
280}
281
282#[derive(Clone, Default)]
283struct RuntimeSnapshot {
284 all: Vec<Model>,
285 available: Vec<Model>,
286 configured_providers: BTreeSet<String>,
287 stored_providers: BTreeSet<String>,
288 auth: HashMap<String, Option<AuthCheck>>,
289}
290
291struct ModelRuntimeInner {
292 credentials: RuntimeCredentials,
293 models_store: Arc<dyn ModelsStore>,
294 models_path: Option<PathBuf>,
295 allow_model_network: bool,
296 auth_env: ProviderEnv,
297 oauth_handlers: HashMap<String, Arc<dyn OAuthAuth>>,
298 config: Mutex<ModelsJsonConfig>,
299 extension_providers: Mutex<HashMap<String, ProviderConfigInput>>,
300 extension_stream_providers: Mutex<HashMap<String, Arc<dyn Provider>>>,
305 composition_errors: Mutex<HashMap<String, String>>,
306 provider_models: Mutex<HashMap<String, Vec<Model>>>,
307 snapshot: Mutex<RuntimeSnapshot>,
308 availability_error: Mutex<Option<String>>,
309 stream_provider: Arc<dyn Provider>,
311 builtins: BuiltinModels,
312}
313
314#[derive(Clone)]
316pub struct ModelRuntime {
317 inner: Arc<ModelRuntimeInner>,
318}
319
320impl std::fmt::Debug for ModelRuntime {
321 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
322 f.debug_struct("ModelRuntime").finish_non_exhaustive()
323 }
324}
325
326impl ModelRuntime {
327 pub async fn create(options: CreateModelRuntimeOptions) -> Result<Self, ModelRuntimeError> {
335 let builtins = builtin_models()
336 .cloned()
337 .map_err(|error| models_error_from_catalog(&error))?;
338
339 let credentials: Arc<dyn CredentialStore> = options.credentials.unwrap_or_else(|| {
340 let path = options
341 .auth_path
342 .clone()
343 .unwrap_or_else(|| PathBuf::from("auth.json"));
344 if options.auth_path.is_none() {
347 Arc::new(InMemoryCredentialStore::new())
348 } else {
349 Arc::new(FileCredentialStore::new(path))
350 }
351 });
352 let credentials = RuntimeCredentials::new(credentials);
353
354 let models_path = options.models_path.clone();
355 let config = options
356 .models_config
357 .unwrap_or_else(|| ModelsJsonConfig::load(models_path.as_deref()));
358
359 let models_store: Arc<dyn ModelsStore> = options.models_store.unwrap_or_else(|| {
360 if let Some(path) = options.models_store_path.clone().or_else(|| {
361 models_path
362 .as_ref()
363 .map(|models| models.with_file_name("models-store.json"))
364 }) {
365 Arc::new(FileModelsStore::new(path))
366 } else {
367 Arc::new(InMemoryModelsStore::new())
368 }
369 });
370
371 let allow_model_network = options.allow_model_network.unwrap_or(false);
372 let auth_env = options.auth_env.unwrap_or_default();
373 let oauth_handlers = options.oauth_handlers.unwrap_or_default();
374 let stream_provider = Arc::new(default_provider_registry());
375
376 let runtime = Self {
377 inner: Arc::new(ModelRuntimeInner {
378 credentials,
379 models_store,
380 models_path,
381 allow_model_network,
382 auth_env,
383 oauth_handlers,
384 config: Mutex::new(config),
385 extension_providers: Mutex::new(HashMap::new()),
386 extension_stream_providers: Mutex::new(HashMap::new()),
387 composition_errors: Mutex::new(HashMap::new()),
388 provider_models: Mutex::new(HashMap::new()),
389 snapshot: Mutex::new(RuntimeSnapshot::default()),
390 availability_error: Mutex::new(None),
391 stream_provider,
392 builtins,
393 }),
394 };
395 runtime.rebuild_providers().await?;
396 let _ = runtime
397 .refresh(ModelsRefreshOptions {
398 allow_network: Some(false),
399 })
400 .await;
401 Ok(runtime)
402 }
403
404 pub async fn create_in_memory() -> Result<Self, ModelRuntimeError> {
410 Self::create(CreateModelRuntimeOptions {
411 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
412 models_store: Some(Arc::new(InMemoryModelsStore::new())),
413 models_config: Some(ModelsJsonConfig::empty()),
414 allow_model_network: Some(false),
415 ..CreateModelRuntimeOptions::default()
416 })
417 .await
418 }
419
420 #[must_use]
422 pub fn get_models(&self, provider_id: Option<&str>) -> Vec<Model> {
423 let snapshot = lock(&self.inner.snapshot);
424 match provider_id {
425 Some(provider_id) => snapshot
426 .all
427 .iter()
428 .filter(|model| model.provider == provider_id)
429 .cloned()
430 .collect(),
431 None => snapshot.all.clone(),
432 }
433 }
434
435 #[must_use]
437 pub fn get_model(&self, provider_id: &str, model_id: &str) -> Option<Model> {
438 lock(&self.inner.snapshot)
439 .all
440 .iter()
441 .find(|model| model.provider == provider_id && model.id == model_id)
442 .cloned()
443 }
444
445 pub async fn get_available(
452 &self,
453 provider_id: Option<&str>,
454 ) -> Result<Vec<Model>, ModelRuntimeError> {
455 self.refresh_availability().await?;
456 let snapshot = lock(&self.inner.snapshot);
457 Ok(match provider_id {
458 Some(provider_id) => snapshot
459 .available
460 .iter()
461 .filter(|model| model.provider == provider_id)
462 .cloned()
463 .collect(),
464 None => snapshot.available.clone(),
465 })
466 }
467
468 #[must_use]
470 pub fn get_available_snapshot(&self) -> Vec<Model> {
471 lock(&self.inner.snapshot).available.clone()
472 }
473
474 #[must_use]
476 pub fn get_error(&self) -> Option<String> {
477 let mut errors = Vec::new();
478 if let Some(error) = lock(&self.inner.config).error() {
479 errors.push(error.to_owned());
480 }
481 for (provider_id, error) in lock(&self.inner.composition_errors).iter() {
482 errors.push(format!("Provider \"{provider_id}\": {error}"));
483 }
484 if let Some(error) = lock(&self.inner.availability_error).clone() {
485 errors.push(format!("Availability refresh: {error}"));
486 }
487 if errors.is_empty() {
488 None
489 } else {
490 Some(errors.join("\n\n"))
491 }
492 }
493
494 pub async fn check_auth(&self, provider_id: &str) -> Option<AuthCheck> {
496 if let Some(cached) = lock(&self.inner.snapshot).auth.get(provider_id).cloned() {
497 return cached;
498 }
499 self.probe_auth(provider_id).await
500 }
501
502 #[must_use]
504 pub fn is_using_oauth(&self, provider_id: &str) -> bool {
505 lock(&self.inner.snapshot)
506 .auth
507 .get(provider_id)
508 .and_then(|entry| entry.as_ref())
509 .is_some_and(|check| check.kind == AuthType::Oauth)
510 }
511
512 #[must_use]
514 pub fn has_configured_auth(&self, provider_id: &str) -> bool {
515 lock(&self.inner.snapshot)
516 .configured_providers
517 .contains(provider_id)
518 }
519
520 pub async fn get_auth_for_provider(
526 &self,
527 provider_id: &str,
528 overrides: ModelRuntimeAuthOverrides,
529 ) -> Result<Option<AuthResult>, ModelRuntimeError> {
530 self.resolve_auth(provider_id, None, overrides, None).await
531 }
532
533 pub async fn get_auth_for_model(
539 &self,
540 model: &Model,
541 overrides: ModelRuntimeAuthOverrides,
542 ) -> Result<Option<AuthResult>, ModelRuntimeError> {
543 self.resolve_auth(&model.provider, Some(model), overrides, None)
544 .await
545 }
546
547 pub async fn set_runtime_api_key(
553 &self,
554 provider_id: &str,
555 api_key: impl Into<String>,
556 ) -> Result<(), ModelRuntimeError> {
557 self.inner
558 .credentials
559 .set_runtime_api_key(provider_id, api_key);
560 {
561 let mut snapshot = lock(&self.inner.snapshot);
562 snapshot.auth.insert(
563 provider_id.to_owned(),
564 Some(AuthCheck {
565 source: Some("runtime API key".to_owned()),
566 kind: AuthType::ApiKey,
567 }),
568 );
569 snapshot.configured_providers.insert(provider_id.to_owned());
570 snapshot.stored_providers.insert(provider_id.to_owned());
571 snapshot.available = snapshot
572 .all
573 .iter()
574 .filter(|model| snapshot.configured_providers.contains(&model.provider))
575 .cloned()
576 .collect();
577 }
578 let _ = self
579 .refresh(ModelsRefreshOptions {
580 allow_network: Some(self.inner.allow_model_network),
581 })
582 .await;
583 Ok(())
584 }
585
586 pub async fn remove_runtime_api_key(&self, provider_id: &str) -> Result<(), ModelRuntimeError> {
592 self.inner.credentials.remove_runtime_api_key(provider_id);
593 let _ = self
594 .refresh(ModelsRefreshOptions {
595 allow_network: Some(self.inner.allow_model_network),
596 })
597 .await;
598 Ok(())
599 }
600
601 pub async fn list_credentials(&self) -> Result<Vec<CredentialInfo>, ModelRuntimeError> {
607 self.inner
608 .credentials
609 .list()
610 .await
611 .map_err(|error| ModelsError::new(ModelsErrorCode::Auth, error.to_string()).into())
612 }
613
614 #[must_use]
616 pub fn get_registered_provider_config(&self, provider_id: &str) -> Option<ProviderConfigInput> {
617 lock(&self.inner.extension_providers)
618 .get(provider_id)
619 .cloned()
620 }
621
622 #[must_use]
624 pub fn get_registered_provider_ids(&self) -> Vec<String> {
625 lock(&self.inner.extension_providers)
626 .keys()
627 .cloned()
628 .collect()
629 }
630
631 pub fn register_provider(
641 &self,
642 provider_id: &str,
643 config: impl std::borrow::Borrow<ProviderConfigInput>,
644 ) -> Result<(), ModelRuntimeError> {
645 let config = config.borrow();
646 validate_extension_provider(provider_id, config)?;
647 {
648 let mut extensions = lock(&self.inner.extension_providers);
649 let previous = extensions.get(provider_id).cloned().unwrap_or_default();
650 let effective = merge_provider_config(&previous, config);
651 extensions.insert(provider_id.to_owned(), effective);
652 }
653 if let Err(error) = self.recompose_provider_sync(provider_id) {
655 lock(&self.inner.composition_errors).insert(provider_id.to_owned(), error);
656 }
657 self.update_model_snapshot_from_maps();
658 self.mark_configured_if_auth_present(provider_id);
659 let runtime = self.clone();
660 let provider = provider_id.to_owned();
661 tokio::spawn(async move {
662 let _ = runtime
663 .refresh(ModelsRefreshOptions {
664 allow_network: Some(false),
665 })
666 .await;
667 let _ = provider;
668 });
669 Ok(())
670 }
671
672 pub fn unregister_provider(&self, provider_id: &str) {
674 lock(&self.inner.extension_providers).remove(provider_id);
675 self.unregister_extension_stream_provider(provider_id);
679 if let Err(error) = self.recompose_provider_sync(provider_id) {
680 lock(&self.inner.composition_errors).insert(provider_id.to_owned(), error);
681 } else {
682 lock(&self.inner.composition_errors).remove(provider_id);
683 }
684 self.update_model_snapshot_from_maps();
685 let runtime = self.clone();
686 tokio::spawn(async move {
687 let _ = runtime
688 .refresh(ModelsRefreshOptions {
689 allow_network: Some(false),
690 })
691 .await;
692 });
693 }
694
695 pub(crate) fn register_extension_stream_provider(
705 &self,
706 provider_id: impl Into<String>,
707 provider: Arc<dyn Provider>,
708 ) {
709 lock(&self.inner.extension_stream_providers).insert(provider_id.into(), provider);
710 }
711
712 pub(crate) fn unregister_extension_stream_provider(&self, provider_id: &str) {
716 lock(&self.inner.extension_stream_providers).remove(provider_id);
717 }
718
719 pub async fn reload_config(&self) -> Result<(), ModelRuntimeError> {
726 let reloaded = ModelsJsonConfig::load(self.inner.models_path.as_deref());
727 *lock(&self.inner.config) = reloaded;
728 self.rebuild_providers().await?;
729 let _ = self
730 .refresh(ModelsRefreshOptions {
731 allow_network: Some(self.inner.allow_model_network),
732 })
733 .await;
734 Ok(())
735 }
736
737 pub async fn refresh(
746 &self,
747 options: ModelsRefreshOptions,
748 ) -> Result<ModelsRefreshResult, ModelRuntimeError> {
749 let _allow_network = options
750 .allow_network
751 .unwrap_or(self.inner.allow_model_network);
752 self.rebuild_providers().await?;
756 match self.refresh_availability().await {
757 Ok(()) => Ok(ModelsRefreshResult::default()),
758 Err(error) => {
759 *lock(&self.inner.availability_error) = Some(error.to_string());
760 Ok(ModelsRefreshResult {
761 aborted: false,
762 errors: BTreeMap::from([("availability".to_owned(), error.to_string())]),
763 })
764 }
765 }
766 }
767
768 #[must_use]
777 pub fn stream_simple(
778 &self,
779 model: Model,
780 context: Context,
781 options: StreamOptions,
782 ) -> BoxStream<'static, Result<AssistantMessageEvent, ProviderError>> {
783 let runtime = self.clone();
784 let native = Arc::clone(&self.inner.stream_provider);
785 Box::pin(
786 async move {
787 match runtime.prepare_request(&model, options).await {
788 Ok(prepared) => {
789 let provider = runtime.select_stream_provider(&prepared.model, native);
790 provider.stream(&prepared.model, context, prepared.options)
791 }
792 Err(error) => {
793 let message = error.to_string();
794 Box::pin(stream::once(
795 async move { Err(ProviderError::new(message)) },
796 ))
797 as BoxStream<'static, Result<AssistantMessageEvent, ProviderError>>
798 }
799 }
800 }
801 .flatten_stream(),
802 )
803 }
804
805 fn select_stream_provider(
810 &self,
811 prepared_model: &Model,
812 native: Arc<dyn Provider>,
813 ) -> Arc<dyn Provider> {
814 let extensions = lock(&self.inner.extension_stream_providers);
815 let Some(extension) = extensions.get(&prepared_model.provider).cloned() else {
816 return native;
817 };
818 drop(extensions);
820
821 let config_api = lock(&self.inner.extension_providers)
822 .get(&prepared_model.provider)
823 .and_then(|config| config.api.clone());
824 match config_api {
825 Some(api) if api == prepared_model.api => extension,
826 _ => native,
827 }
828 }
829
830 async fn prepare_request(
831 &self,
832 model: &Model,
833 mut options: StreamOptions,
834 ) -> Result<PreparedRequest, ModelRuntimeError> {
835 let overrides = ModelRuntimeAuthOverrides {
836 api_key: options.api_key.clone(),
837 env: options.env.clone(),
838 };
839 let resolution = self
840 .resolve_auth(
841 &model.provider,
842 Some(model),
843 overrides,
844 options.signal.clone(),
845 )
846 .await?
847 .ok_or_else(|| {
848 ModelsError::new(
849 ModelsErrorCode::Auth,
850 format!("Provider is not configured: {}", model.provider),
851 )
852 })?;
853
854 if options.api_key.is_none() {
855 options.api_key.clone_from(&resolution.auth.api_key);
856 }
857 if let Some(headers) = resolution.auth.headers.clone() {
858 let mut merged = options.headers.take().unwrap_or_default();
859 for (name, value) in headers {
860 merged.insert(name, Some(value));
861 }
862 options.headers = Some(merged);
863 }
864 if let Some(env) = resolution.env {
865 let mut merged = options.env.take().unwrap_or_default();
866 for (key, value) in env {
867 merged.entry(key).or_insert(value);
868 }
869 options.env = Some(merged);
870 }
871 let mut model = model.clone();
872 if let Some(base_url) = resolution.auth.base_url {
873 model.base_url = base_url;
874 }
875 Ok(PreparedRequest { model, options })
876 }
877
878 async fn resolve_auth(
879 &self,
880 provider_id: &str,
881 model: Option<&Model>,
882 overrides: ModelRuntimeAuthOverrides,
883 signal: Option<tokio_util::sync::CancellationToken>,
884 ) -> Result<Option<AuthResult>, ModelRuntimeError> {
885 let provider_auth = self.provider_auth_for(provider_id);
888 let auth_context = self.auth_context_for(&overrides);
889 let resolution_overrides = AuthResolutionOverrides {
890 api_key: overrides.api_key.clone(),
891 env: overrides.env.clone(),
892 };
893 let mut result = resolve_provider_auth_with_signal(
894 provider_id,
895 &provider_auth,
896 &self.inner.credentials,
897 auth_context.as_ref(),
898 Some(&resolution_overrides),
899 signal,
900 )
901 .await?;
902
903 if result.is_none()
905 && let Some(configured) = self.configured_api_key(provider_id)
906 && let Some(resolved) = resolve_config_value(&configured, overrides.env.as_ref())
907 {
908 result = Some(AuthResult {
909 auth: ModelAuth {
910 api_key: Some(resolved),
911 headers: None,
912 base_url: None,
913 },
914 env: overrides.env.clone(),
915 source: Some("models.json".to_owned()),
916 });
917 }
918
919 if let Some(result) = result.as_mut() {
920 self.apply_configured_auth_projection(provider_id, model, result, &overrides)?;
921 }
922 Ok(result)
923 }
924
925 fn provider_auth_for(&self, provider_id: &str) -> ProviderAuth {
926 let env_vars = api_key_env_vars(provider_id).unwrap_or(&[]);
929 let api_key = Some(env_api_key_auth(format!("{provider_id} API key"), env_vars));
930 let oauth =
931 self.inner
932 .oauth_handlers
933 .get(provider_id)
934 .cloned()
935 .or_else(|| match provider_id {
936 "anthropic" => pi_ai::auth::oauth::anthropic::AnthropicOAuth::new()
937 .ok()
938 .map(|auth| Arc::new(auth) as Arc<dyn OAuthAuth>),
939 "openai-codex" => {
940 pi_ai::auth::oauth::openai_codex::OpenAiCodexOAuth::shared().ok()
941 }
942 "github-copilot" => {
943 pi_ai::auth::oauth::github_copilot::GitHubCopilotOAuth::shared().ok()
944 }
945 "xai" => pi_ai::auth::oauth::xai::XaiOAuth::shared().ok(),
946 _ => None,
947 });
948 ProviderAuth { api_key, oauth }
949 }
950
951 fn auth_context_for(&self, overrides: &ModelRuntimeAuthOverrides) -> Arc<dyn AuthContext> {
952 let mut map = MapAuthContext::new();
953 for (key, value) in &self.inner.auth_env {
954 map = map.with_env(key.clone(), value.clone());
955 }
956 if let Some(env) = overrides.env.as_ref() {
957 for (key, value) in env {
958 map = map.with_env(key.clone(), value.clone());
959 }
960 }
961 Arc::new(map)
962 }
963
964 fn apply_configured_auth_projection(
965 &self,
966 provider_id: &str,
967 model: Option<&Model>,
968 result: &mut AuthResult,
969 overrides: &ModelRuntimeAuthOverrides,
970 ) -> Result<(), ModelRuntimeError> {
971 let env = merge_provider_env(&self.inner.auth_env, overrides.env.as_ref());
972 let mut headers = result.auth.headers.take().unwrap_or_default();
974 if let Some(model) = model
975 && let Some(model_headers) = model.headers.as_ref()
976 {
977 for (name, value) in model_headers {
978 headers.insert(name.clone(), value.clone());
979 }
980 }
981 if let Some(config_headers) = self.configured_headers(provider_id)
982 && let Some(resolved) = resolve_headers(Some(&config_headers), Some(&env))
983 {
984 for (name, value) in resolved {
985 headers.insert(name, value);
986 }
987 }
988
989 let auth_header = self.configured_auth_header(provider_id);
990 if auth_header {
991 let Some(api_key) = result.auth.api_key.as_ref() else {
992 return Err(ModelRuntimeError::Models(ModelsError::new(
993 ModelsErrorCode::Auth,
994 "authHeader requires a resolved API key",
995 )));
996 };
997 headers.insert("Authorization".to_owned(), format!("Bearer {api_key}"));
998 }
999
1000 if !headers.is_empty() {
1001 result.auth.headers = Some(headers);
1002 }
1003 Ok(())
1007 }
1008
1009 fn configured_api_key(&self, provider_id: &str) -> Option<String> {
1010 lock(&self.inner.extension_providers)
1011 .get(provider_id)
1012 .and_then(|config| config.api_key.clone())
1013 .or_else(|| {
1014 lock(&self.inner.config)
1015 .get_provider(provider_id)
1016 .and_then(|config| config.api_key.clone())
1017 })
1018 }
1019
1020 fn configured_headers(&self, provider_id: &str) -> Option<ProviderHeaders> {
1021 let extension = lock(&self.inner.extension_providers)
1022 .get(provider_id)
1023 .and_then(|config| config.headers.clone());
1024 let models_json = lock(&self.inner.config)
1025 .get_provider(provider_id)
1026 .and_then(|config| config.headers.clone());
1027 match (models_json, extension) {
1028 (None, None) => None,
1029 (Some(base), None) => Some(base),
1030 (None, Some(ext)) => Some(ext),
1031 (Some(mut base), Some(ext)) => {
1032 for (k, v) in ext {
1033 base.insert(k, v);
1034 }
1035 Some(base)
1036 }
1037 }
1038 }
1039
1040 fn configured_auth_header(&self, provider_id: &str) -> bool {
1041 lock(&self.inner.extension_providers)
1042 .get(provider_id)
1043 .and_then(|config| config.auth_header)
1044 .or_else(|| {
1045 lock(&self.inner.config)
1046 .get_provider(provider_id)
1047 .and_then(|config| config.auth_header)
1048 })
1049 .unwrap_or(false)
1050 }
1051
1052 async fn rebuild_providers(&self) -> Result<(), ModelRuntimeError> {
1053 let provider_ids = self.provider_ids();
1054 lock(&self.inner.composition_errors).clear();
1055 lock(&self.inner.provider_models).clear();
1056 for provider_id in provider_ids {
1057 if let Err(error) = self.recompose_provider(&provider_id).await {
1058 lock(&self.inner.composition_errors).insert(provider_id, error);
1059 }
1060 }
1061 self.update_model_snapshot_from_maps();
1062 Ok(())
1063 }
1064
1065 fn provider_ids(&self) -> BTreeSet<String> {
1066 let mut ids = BTreeSet::new();
1067 for provider in self.inner.builtins.keys() {
1068 ids.insert(provider.clone());
1069 }
1070 for provider in lock(&self.inner.config).provider_ids() {
1071 ids.insert(provider);
1072 }
1073 for provider in lock(&self.inner.extension_providers).keys() {
1074 ids.insert(provider.clone());
1075 }
1076 ids
1077 }
1078
1079 async fn recompose_provider(&self, provider_id: &str) -> Result<(), String> {
1080 let models = self.compose_models_for_provider(provider_id).await?;
1081 lock(&self.inner.provider_models).insert(provider_id.to_owned(), models);
1082 lock(&self.inner.composition_errors).remove(provider_id);
1083 Ok(())
1084 }
1085
1086 fn recompose_provider_sync(&self, provider_id: &str) -> Result<(), String> {
1087 let store_entry = None;
1091 let models = compose_models_static(
1092 provider_id,
1093 &self.inner.builtins,
1094 store_entry,
1095 lock(&self.inner.config).get_provider(provider_id),
1096 lock(&self.inner.extension_providers).get(provider_id),
1097 )?;
1098 lock(&self.inner.provider_models).insert(provider_id.to_owned(), models);
1099 Ok(())
1100 }
1101
1102 async fn compose_models_for_provider(&self, provider_id: &str) -> Result<Vec<Model>, String> {
1103 let store_entry = self
1104 .inner
1105 .models_store
1106 .read(provider_id)
1107 .await
1108 .map_err(|error| error.to_string())?;
1109 compose_models_static(
1110 provider_id,
1111 &self.inner.builtins,
1112 store_entry.as_ref(),
1113 lock(&self.inner.config).get_provider(provider_id),
1114 lock(&self.inner.extension_providers).get(provider_id),
1115 )
1116 }
1117
1118 fn update_model_snapshot_from_maps(&self) {
1119 let mut all = Vec::new();
1120 for models in lock(&self.inner.provider_models).values() {
1121 all.extend(models.iter().cloned());
1122 }
1123 all.sort_by(|left, right| {
1124 left.provider
1125 .cmp(&right.provider)
1126 .then_with(|| left.id.cmp(&right.id))
1127 });
1128 let mut snapshot = lock(&self.inner.snapshot);
1129 let configured = snapshot.configured_providers.clone();
1130 snapshot.available = all
1131 .iter()
1132 .filter(|model| configured.contains(&model.provider))
1133 .cloned()
1134 .collect();
1135 snapshot.all = all;
1136 }
1137
1138 fn mark_configured_if_auth_present(&self, provider_id: &str) {
1139 let has_runtime = self.inner.credentials.has_runtime_api_key(provider_id);
1140 let has_configured_key = self.configured_api_key(provider_id).is_some_and(|key| {
1141 is_config_value_configured(&key, Some(&self.inner.auth_env))
1142 || (!key.starts_with('$') && !key.starts_with('!'))
1143 });
1144 let has_oauth = lock(&self.inner.extension_providers)
1145 .get(provider_id)
1146 .and_then(|config| config.oauth.as_ref())
1147 .is_some()
1148 || lock(&self.inner.config)
1149 .get_provider(provider_id)
1150 .and_then(|config| config.oauth.as_ref())
1151 .is_some();
1152 let stored = lock(&self.inner.snapshot)
1153 .stored_providers
1154 .contains(provider_id);
1155 if !(has_runtime || has_configured_key || has_oauth || stored) {
1156 return;
1157 }
1158 let mut snapshot = lock(&self.inner.snapshot);
1159 snapshot.configured_providers.insert(provider_id.to_owned());
1160 if snapshot
1161 .auth
1162 .get(provider_id)
1163 .and_then(Option::as_ref)
1164 .is_none()
1165 {
1166 let kind = if has_oauth && !has_configured_key && !has_runtime {
1167 AuthType::Oauth
1168 } else {
1169 AuthType::ApiKey
1170 };
1171 snapshot.auth.insert(
1172 provider_id.to_owned(),
1173 Some(AuthCheck {
1174 source: Some("configured provider".to_owned()),
1175 kind,
1176 }),
1177 );
1178 }
1179 snapshot.available = snapshot
1180 .all
1181 .iter()
1182 .filter(|model| snapshot.configured_providers.contains(&model.provider))
1183 .cloned()
1184 .collect();
1185 }
1186
1187 async fn refresh_availability(&self) -> Result<(), ModelRuntimeError> {
1188 let provider_ids = self.provider_ids();
1189 let mut auth = HashMap::new();
1190 let mut configured = BTreeSet::new();
1191 for provider_id in &provider_ids {
1192 let check = self.probe_auth(provider_id).await;
1193 if check.is_some() {
1194 configured.insert(provider_id.clone());
1195 }
1196 auth.insert(provider_id.clone(), check);
1197 }
1198
1199 let stored = match self.inner.credentials.list().await {
1200 Ok(list) => list
1201 .into_iter()
1202 .map(|entry| entry.provider_id)
1203 .collect::<BTreeSet<_>>(),
1204 Err(error) => {
1205 *lock(&self.inner.availability_error) = Some(error.to_string());
1206 BTreeSet::new()
1207 }
1208 };
1209 for provider_id in &stored {
1210 configured.insert(provider_id.clone());
1211 auth.entry(provider_id.clone()).or_insert_with(|| {
1212 Some(AuthCheck {
1213 source: Some("stored credential".to_owned()),
1214 kind: AuthType::ApiKey,
1215 })
1216 });
1217 }
1218
1219 let all = {
1220 let maps = lock(&self.inner.provider_models);
1221 let mut all = Vec::new();
1222 for models in maps.values() {
1223 all.extend(models.iter().cloned());
1224 }
1225 all.sort_by(|left, right| {
1226 left.provider
1227 .cmp(&right.provider)
1228 .then_with(|| left.id.cmp(&right.id))
1229 });
1230 all
1231 };
1232 let available = all
1233 .iter()
1234 .filter(|model| configured.contains(&model.provider))
1235 .cloned()
1236 .collect();
1237 *lock(&self.inner.snapshot) = RuntimeSnapshot {
1238 all,
1239 available,
1240 configured_providers: configured,
1241 stored_providers: stored,
1242 auth,
1243 };
1244 *lock(&self.inner.availability_error) = None;
1245 Ok(())
1246 }
1247
1248 async fn probe_auth(&self, provider_id: &str) -> Option<AuthCheck> {
1249 if self.inner.credentials.has_runtime_api_key(provider_id) {
1250 return Some(AuthCheck {
1251 source: Some("runtime API key".to_owned()),
1252 kind: AuthType::ApiKey,
1253 });
1254 }
1255 if let Ok(Some(credential)) = self.inner.credentials.read(provider_id).await {
1256 return Some(match credential {
1257 Credential::ApiKey(_) => AuthCheck {
1258 source: Some("stored credential".to_owned()),
1259 kind: AuthType::ApiKey,
1260 },
1261 Credential::Oauth(_) => AuthCheck {
1262 source: Some("OAuth".to_owned()),
1263 kind: AuthType::Oauth,
1264 },
1265 });
1266 }
1267 if let Some(api_key) = get_env_api_key(provider_id, Some(&self.inner.auth_env)) {
1268 let source = if api_key == AMBIENT_AUTH_MARKER {
1269 Some("ambient credentials".to_owned())
1270 } else {
1271 api_key_env_vars(provider_id)
1272 .and_then(|vars| vars.first().map(|name| (*name).to_owned()))
1273 };
1274 return Some(AuthCheck {
1275 source,
1276 kind: AuthType::ApiKey,
1277 });
1278 }
1279 if let Some(configured) = self.configured_api_key(provider_id)
1280 && (is_config_value_configured(&configured, Some(&self.inner.auth_env))
1281 || (!configured.starts_with('$') && !configured.starts_with('!')))
1282 {
1283 return Some(AuthCheck {
1284 source: Some("models.json".to_owned()),
1285 kind: AuthType::ApiKey,
1286 });
1287 }
1288 let oauth = lock(&self.inner.extension_providers)
1289 .get(provider_id)
1290 .and_then(|config| config.oauth.clone())
1291 .or_else(|| {
1292 lock(&self.inner.config)
1293 .get_provider(provider_id)
1294 .and_then(|config| config.oauth.clone())
1295 });
1296 if oauth.is_some()
1297 && lock(&self.inner.snapshot)
1298 .stored_providers
1299 .contains(provider_id)
1300 {
1301 return Some(AuthCheck {
1302 source: Some("OAuth".to_owned()),
1303 kind: AuthType::Oauth,
1304 });
1305 }
1306 if matches!(provider_id, "amazon-bedrock" | "google-vertex")
1308 && get_env_api_key(provider_id, Some(&self.inner.auth_env)).is_some()
1309 {
1310 return Some(AuthCheck {
1311 source: Some("ambient credentials".to_owned()),
1312 kind: AuthType::ApiKey,
1313 });
1314 }
1315 let _ = KnownProvider::from_id(provider_id);
1316 None
1317 }
1318}
1319
1320struct PreparedRequest {
1321 model: Model,
1322 options: StreamOptions,
1323}
1324
1325fn default_provider_registry() -> ProviderRegistry {
1326 let client = reqwest::Client::new();
1327 ProviderRegistry::new([
1328 Arc::new(OpenAiCompletions::new(client.clone())),
1329 Arc::new(OpenAiResponses::new(client.clone())),
1330 Arc::new(AzureOpenAiResponses::new(client.clone())),
1331 Arc::new(OpenAiCodexResponses::new(client.clone())),
1332 Arc::new(AnthropicMessages::new(client.clone())),
1333 Arc::new(BedrockConverseStream::new(Arc::new(
1334 DefaultBedrockClientFactory::new(),
1335 ))),
1336 Arc::new(GoogleGenerativeAi::new(client.clone())),
1337 Arc::new(GoogleVertex::new(client.clone(), None)),
1338 Arc::new(MistralConversations::new(client.clone())),
1339 Arc::new(PiMessages::new(client)),
1340 ])
1341}
1342
1343fn parse_models_json(
1344 content: &str,
1345 path_display: &str,
1346) -> Result<BTreeMap<String, ProviderConfigInput>, String> {
1347 let stripped = strip_json_comments(content);
1348 let parsed: Value = serde_json::from_str(&stripped)
1349 .map_err(|error| format!("Failed to parse models.json: {error}\n\nFile: {path_display}"))?;
1350 let providers_value = parsed.get("providers").cloned().ok_or_else(|| {
1351 format!("Invalid models.json schema:\n - providers: required\n\nFile: {path_display}")
1352 })?;
1353 let object = providers_value.as_object().ok_or_else(|| {
1354 format!(
1355 "Invalid models.json schema:\n - providers: expected object\n\nFile: {path_display}"
1356 )
1357 })?;
1358 let mut providers = BTreeMap::new();
1359 for (provider_id, value) in object {
1360 let config: ProviderConfigInput = serde_json::from_value(value.clone()).map_err(|error| {
1361 format!(
1362 "Invalid models.json schema:\n - providers.{provider_id}: {error}\n\nFile: {path_display}"
1363 )
1364 })?;
1365 providers.insert(provider_id.clone(), config);
1366 }
1367 Ok(providers)
1368}
1369
1370fn strip_json_comments(input: &str) -> String {
1371 let mut out = String::with_capacity(input.len());
1373 let bytes = input.as_bytes();
1374 let mut i = 0;
1375 let mut in_string = false;
1376 let mut escaped = false;
1377 while i < bytes.len() {
1378 let b = bytes[i];
1379 if in_string {
1380 out.push(b as char);
1381 if escaped {
1382 escaped = false;
1383 } else if b == b'\\' {
1384 escaped = true;
1385 } else if b == b'"' {
1386 in_string = false;
1387 }
1388 i += 1;
1389 continue;
1390 }
1391 if b == b'"' {
1392 in_string = true;
1393 out.push('"');
1394 i += 1;
1395 continue;
1396 }
1397 if b == b'/' && i + 1 < bytes.len() {
1398 if bytes[i + 1] == b'/' {
1399 i += 2;
1400 while i < bytes.len() && bytes[i] != b'\n' {
1401 i += 1;
1402 }
1403 continue;
1404 }
1405 if bytes[i + 1] == b'*' {
1406 i += 2;
1407 while i + 1 < bytes.len() && !(bytes[i] == b'*' && bytes[i + 1] == b'/') {
1408 i += 1;
1409 }
1410 i = (i + 2).min(bytes.len());
1411 continue;
1412 }
1413 }
1414 out.push(b as char);
1415 i += 1;
1416 }
1417 out
1418}
1419
1420fn merge_provider_config(
1421 previous: &ProviderConfigInput,
1422 incoming: &ProviderConfigInput,
1423) -> ProviderConfigInput {
1424 ProviderConfigInput {
1425 name: incoming.name.clone().or_else(|| previous.name.clone()),
1426 base_url: incoming
1427 .base_url
1428 .clone()
1429 .or_else(|| previous.base_url.clone()),
1430 api_key: incoming
1431 .api_key
1432 .clone()
1433 .or_else(|| previous.api_key.clone()),
1434 api: incoming.api.clone().or_else(|| previous.api.clone()),
1435 headers: incoming
1436 .headers
1437 .clone()
1438 .or_else(|| previous.headers.clone()),
1439 auth_header: incoming.auth_header.or(previous.auth_header),
1440 models: incoming.models.clone().or_else(|| previous.models.clone()),
1441 model_overrides: incoming
1442 .model_overrides
1443 .clone()
1444 .or_else(|| previous.model_overrides.clone()),
1445 oauth: incoming.oauth.clone().or_else(|| previous.oauth.clone()),
1446 }
1447}
1448
1449fn validate_extension_provider(
1450 provider_id: &str,
1451 config: &ProviderConfigInput,
1452) -> Result<(), ModelRuntimeError> {
1453 if config
1454 .api_key
1455 .as_deref()
1456 .is_some_and(|value| value.starts_with('!'))
1457 {
1458 return Err(ModelRuntimeError::Registration(format!(
1459 "Provider {provider_id}: extension apiKey cannot execute shell commands"
1460 )));
1461 }
1462 if let Some((name, _)) = config
1463 .headers
1464 .as_ref()
1465 .and_then(|headers| headers.iter().find(|(_, value)| value.starts_with('!')))
1466 {
1467 return Err(ModelRuntimeError::Registration(format!(
1468 "Provider {provider_id}: extension header {name:?} cannot execute shell commands"
1469 )));
1470 }
1471
1472 if let Some(models) = config.models.as_ref() {
1473 for model in models {
1474 let api = model
1475 .api
1476 .as_deref()
1477 .or(config.api.as_deref())
1478 .ok_or_else(|| {
1479 ModelRuntimeError::Registration(format!(
1480 "Provider {provider_id}, model {}: no \"api\" specified. Set at provider or model level.",
1481 model.id
1482 ))
1483 })?;
1484 let _ = api;
1485 let base_url = model
1486 .base_url
1487 .as_deref()
1488 .or(config.base_url.as_deref())
1489 .ok_or_else(|| {
1490 ModelRuntimeError::Registration(format!(
1491 "Provider {provider_id}: \"baseUrl\" is required when defining custom models."
1492 ))
1493 })?;
1494 let _ = base_url;
1495 if model.context_window == Some(0) {
1496 return Err(ModelRuntimeError::Registration(format!(
1497 "Provider {provider_id}, model {}: invalid contextWindow",
1498 model.id
1499 )));
1500 }
1501 if model.max_tokens == Some(0) {
1502 return Err(ModelRuntimeError::Registration(format!(
1503 "Provider {provider_id}, model {}: invalid maxTokens",
1504 model.id
1505 )));
1506 }
1507 }
1508 }
1509 Ok(())
1510}
1511
1512fn compose_models_static(
1513 provider_id: &str,
1514 builtins: &BuiltinModels,
1515 store_entry: Option<&ModelsStoreEntry>,
1516 models_json: Option<&ProviderConfigInput>,
1517 extension: Option<&ProviderConfigInput>,
1518) -> Result<Vec<Model>, String> {
1519 let mut models =
1521 compose_provider_models(provider_id, builtins, store_entry, &ModelOverrides::new())
1522 .map_err(|error| error.to_string())?;
1523
1524 if let Some(config) = models_json {
1526 if config.oauth.is_some() && config.base_url.is_none() {
1527 return Err(format!(
1528 "Provider {provider_id}: \"baseUrl\" is required when \"oauth\" is set."
1529 ));
1530 }
1531 let radius_oauth = config.oauth.as_deref() == Some("radius");
1532 if let Some(base_url) = config.base_url.as_ref()
1533 && !radius_oauth
1534 {
1535 for model in &mut models {
1536 model.base_url.clone_from(base_url);
1537 }
1538 }
1539 if let Some(definitions) = config.models.as_ref() {
1540 for definition in definitions {
1541 let defaults = models
1542 .iter()
1543 .find(|model| model.id == definition.id)
1544 .cloned()
1545 .or_else(|| models.first().cloned());
1546 let composed =
1547 model_from_definition(provider_id, definition, config, defaults.as_ref())?;
1548 if let Some(index) = models.iter().position(|model| model.id == definition.id) {
1549 models[index] = composed;
1550 } else {
1551 models.push(composed);
1552 }
1553 }
1554 }
1555 }
1556
1557 if let Some(extension) = extension {
1559 if let Some(definitions) = extension.models.as_ref() {
1560 let mut replaced = Vec::with_capacity(definitions.len());
1561 for definition in definitions {
1562 let defaults = models
1563 .iter()
1564 .find(|model| model.id == definition.id)
1565 .cloned()
1566 .or_else(|| models.first().cloned());
1567 replaced.push(model_from_definition(
1568 provider_id,
1569 definition,
1570 extension,
1571 defaults.as_ref(),
1572 )?);
1573 }
1574 models = replaced;
1575 } else if let Some(base_url) = extension.base_url.as_ref() {
1576 for model in &mut models {
1577 model.base_url.clone_from(base_url);
1578 }
1579 }
1580 }
1581
1582 let mut overrides = ModelOverrides::new();
1584 if let Some(models_json) = models_json
1585 && let Some(model_overrides) = models_json.model_overrides.as_ref()
1586 {
1587 for (key, value) in model_overrides {
1588 overrides.insert(key.clone(), value.clone());
1589 }
1590 }
1591 if let Some(extension) = extension
1592 && let Some(model_overrides) = extension.model_overrides.as_ref()
1593 {
1594 for (key, value) in model_overrides {
1595 overrides.insert(key.clone(), value.clone());
1596 }
1597 }
1598 if !overrides.is_empty() {
1599 models = apply_model_overrides(&models, &overrides).map_err(|error| error.to_string())?;
1600 }
1601
1602 Ok(models)
1603}
1604
1605fn model_from_definition(
1606 provider_id: &str,
1607 definition: &ProviderModelDefinition,
1608 provider: &ProviderConfigInput,
1609 defaults: Option<&Model>,
1610) -> Result<Model, String> {
1611 let api = definition
1612 .api
1613 .clone()
1614 .or_else(|| provider.api.clone())
1615 .or_else(|| defaults.map(|model| model.api.clone()))
1616 .ok_or_else(|| {
1617 format!(
1618 "Provider {provider_id}, model {}: no \"api\" specified. Set at provider or model level.",
1619 definition.id
1620 )
1621 })?;
1622 let base_url = definition
1623 .base_url
1624 .clone()
1625 .or_else(|| provider.base_url.clone())
1626 .or_else(|| defaults.map(|model| model.base_url.clone()))
1627 .ok_or_else(|| {
1628 format!("Provider {provider_id}: \"baseUrl\" is required when defining custom models.")
1629 })?;
1630 Ok(Model {
1631 id: definition.id.clone(),
1632 name: definition
1633 .name
1634 .clone()
1635 .or_else(|| defaults.map(|model| model.name.clone()))
1636 .unwrap_or_else(|| definition.id.clone()),
1637 api,
1638 provider: provider_id.to_owned(),
1639 base_url,
1640 reasoning: definition.reasoning,
1641 thinking_level_map: definition
1642 .thinking_level_map
1643 .clone()
1644 .or_else(|| defaults.and_then(|model| model.thinking_level_map.clone())),
1645 input: definition.input.clone().unwrap_or_else(|| {
1646 defaults.map_or_else(|| vec![ModelInput::Text], |model| model.input.clone())
1647 }),
1648 cost: definition
1649 .cost
1650 .clone()
1651 .or_else(|| defaults.map(|model| model.cost.clone()))
1652 .unwrap_or_default(),
1653 context_window: definition
1654 .context_window
1655 .or_else(|| defaults.map(|model| model.context_window))
1656 .unwrap_or(128_000),
1657 max_tokens: definition
1658 .max_tokens
1659 .or_else(|| defaults.map(|model| model.max_tokens))
1660 .unwrap_or(16_384),
1661 headers: defaults.and_then(|model| model.headers.clone()),
1663 compat: definition
1664 .compat
1665 .clone()
1666 .or_else(|| defaults.and_then(|model| model.compat.clone())),
1667 extra: defaults
1668 .map(|model| model.extra.clone())
1669 .unwrap_or_default(),
1670 })
1671}
1672
1673fn merge_provider_env(base: &ProviderEnv, overlay: Option<&ProviderEnv>) -> ProviderEnv {
1674 let mut env = base.clone();
1675 if let Some(overlay) = overlay {
1676 for (key, value) in overlay {
1677 env.insert(key.clone(), value.clone());
1678 }
1679 }
1680 env
1681}
1682
1683fn lock<T>(mutex: &Mutex<T>) -> std::sync::MutexGuard<'_, T> {
1684 mutex
1685 .lock()
1686 .unwrap_or_else(std::sync::PoisonError::into_inner)
1687}
1688
1689#[cfg(test)]
1690mod tests {
1691 use super::*;
1692 use pi_ai::auth::OAuthCredential;
1693 use serde_json::json;
1694 use std::sync::atomic::{AtomicBool, Ordering};
1695
1696 fn required<T>(value: Option<T>, message: &'static str) -> Result<T, ModelRuntimeError> {
1697 value.ok_or_else(|| ModelRuntimeError::Registration(message.to_owned()))
1698 }
1699
1700 fn custom_model(_provider: &str, id: &str) -> ProviderModelDefinition {
1701 ProviderModelDefinition {
1702 id: id.to_owned(),
1703 name: Some(id.to_owned()),
1704 api: Some("openai-completions".to_owned()),
1705 base_url: Some("https://example.test/v1".to_owned()),
1706 reasoning: false,
1707 thinking_level_map: None,
1708 input: Some(vec![ModelInput::Text]),
1709 cost: Some(ModelCost::default()),
1710 context_window: Some(32_000),
1711 max_tokens: Some(4_096),
1712 headers: None,
1713 compat: None,
1714 }
1715 }
1716
1717 #[tokio::test]
1718 async fn builtins_load_without_network() -> Result<(), ModelRuntimeError> {
1719 let runtime = ModelRuntime::create_in_memory().await?;
1720 let models = runtime.get_models(Some("anthropic"));
1721 assert!(
1722 models.iter().any(|model| model.id == "claude-opus-4-8"),
1723 "expected anthropic builtin model"
1724 );
1725 assert!(runtime.get_error().is_none());
1726 Ok(())
1727 }
1728
1729 #[tokio::test]
1730 async fn custom_provider_registration_and_unregister() -> Result<(), ModelRuntimeError> {
1731 let runtime = ModelRuntime::create_in_memory().await?;
1732 runtime.register_provider(
1733 "acme",
1734 &ProviderConfigInput {
1735 name: Some("Acme".to_owned()),
1736 base_url: Some("https://acme.test/v1".to_owned()),
1737 api: Some("openai-completions".to_owned()),
1738 api_key: Some("sk-acme".to_owned()),
1739 models: Some(vec![custom_model("acme", "acme-1")]),
1740 ..ProviderConfigInput::default()
1741 },
1742 )?;
1743 let model = required(runtime.get_model("acme", "acme-1"), "registered model")?;
1744 assert_eq!(model.provider, "acme");
1745 assert_eq!(model.base_url, "https://example.test/v1");
1747 assert!(runtime.has_configured_auth("acme"));
1748
1749 runtime.unregister_provider("acme");
1750 assert!(runtime.get_model("acme", "acme-1").is_none());
1751 Ok(())
1752 }
1753
1754 #[tokio::test]
1755 async fn extension_provider_rejects_shell_command_secrets() -> Result<(), ModelRuntimeError> {
1756 let runtime = ModelRuntime::create_in_memory().await?;
1757
1758 for config in [
1759 ProviderConfigInput {
1760 api_key: Some("!printf stolen".to_owned()),
1761 ..ProviderConfigInput::default()
1762 },
1763 ProviderConfigInput {
1764 headers: Some(BTreeMap::from([(
1765 "Authorization".to_owned(),
1766 "!printf stolen".to_owned(),
1767 )])),
1768 ..ProviderConfigInput::default()
1769 },
1770 ] {
1771 let error = match runtime.register_provider("untrusted-extension", &config) {
1772 Ok(()) => {
1773 return Err(ModelRuntimeError::Registration(
1774 "extension shell command was accepted".to_owned(),
1775 ));
1776 }
1777 Err(error) => error,
1778 };
1779 assert!(
1780 error.to_string().contains("cannot execute shell commands"),
1781 "unexpected registration error: {error}"
1782 );
1783 assert!(
1784 runtime
1785 .get_registered_provider_config("untrusted-extension")
1786 .is_none(),
1787 "rejected provider must not mutate extension state"
1788 );
1789 }
1790 Ok(())
1791 }
1792
1793 #[tokio::test]
1794 async fn trusted_models_json_keeps_shell_command_resolution() -> Result<(), ModelRuntimeError> {
1795 let providers = BTreeMap::from([(
1796 "trusted-config".to_owned(),
1797 ProviderConfigInput {
1798 api_key: Some("!printf trusted-value".to_owned()),
1799 models: Some(vec![custom_model("trusted-config", "trusted-model")]),
1800 ..ProviderConfigInput::default()
1801 },
1802 )]);
1803 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
1804 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
1805 models_store: Some(Arc::new(InMemoryModelsStore::new())),
1806 models_config: Some(ModelsJsonConfig::from_providers(providers)),
1807 allow_model_network: Some(false),
1808 ..CreateModelRuntimeOptions::default()
1809 })
1810 .await?;
1811
1812 let auth = required(
1813 runtime
1814 .get_auth_for_provider("trusted-config", ModelRuntimeAuthOverrides::default())
1815 .await?,
1816 "trusted models.json auth",
1817 )?;
1818 assert_eq!(auth.auth.api_key.as_deref(), Some("trusted-value"));
1819 Ok(())
1820 }
1821
1822 #[tokio::test]
1823 async fn models_json_override_and_reload() -> Result<(), ModelRuntimeError> {
1824 let mut providers = BTreeMap::new();
1825 providers.insert(
1826 "anthropic".to_owned(),
1827 ProviderConfigInput {
1828 model_overrides: Some(BTreeMap::from([(
1829 "claude-opus-4-8".to_owned(),
1830 json!({ "name": "Opus Override" }),
1831 )])),
1832 ..ProviderConfigInput::default()
1833 },
1834 );
1835 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
1836 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
1837 models_store: Some(Arc::new(InMemoryModelsStore::new())),
1838 models_config: Some(ModelsJsonConfig::from_providers(providers)),
1839 allow_model_network: Some(false),
1840 ..CreateModelRuntimeOptions::default()
1841 })
1842 .await?;
1843 let model = required(runtime.get_model("anthropic", "claude-opus-4-8"), "model")?;
1844 assert_eq!(model.name, "Opus Override");
1845 Ok(())
1846 }
1847
1848 #[tokio::test]
1849 async fn runtime_api_key_configures_auth_and_available() -> Result<(), ModelRuntimeError> {
1850 let runtime = ModelRuntime::create_in_memory().await?;
1851 assert!(!runtime.has_configured_auth("anthropic"));
1852 runtime.set_runtime_api_key("anthropic", "sk-test").await?;
1853 assert!(runtime.has_configured_auth("anthropic"));
1854 assert!(!runtime.is_using_oauth("anthropic"));
1855 let available = runtime.get_available(Some("anthropic")).await?;
1856 assert!(!available.is_empty());
1857 let auth = required(
1858 runtime
1859 .get_auth_for_provider("anthropic", ModelRuntimeAuthOverrides::default())
1860 .await?,
1861 "auth",
1862 )?;
1863 assert_eq!(auth.auth.api_key.as_deref(), Some("sk-test"));
1864 Ok(())
1865 }
1866
1867 #[tokio::test]
1868 async fn env_auth_probe_and_get_auth() -> Result<(), ModelRuntimeError> {
1869 let mut env = ProviderEnv::new();
1870 env.insert("OPENAI_API_KEY".to_owned(), "sk-env".to_owned());
1871 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
1872 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
1873 models_store: Some(Arc::new(InMemoryModelsStore::new())),
1874 models_config: Some(ModelsJsonConfig::empty()),
1875 allow_model_network: Some(false),
1876 auth_env: Some(env),
1877 ..CreateModelRuntimeOptions::default()
1878 })
1879 .await?;
1880 let check = required(runtime.check_auth("openai").await, "check")?;
1881 assert_eq!(check.kind, AuthType::ApiKey);
1882 assert!(runtime.has_configured_auth("openai"));
1883 let auth = required(
1884 runtime
1885 .get_auth_for_provider("openai", ModelRuntimeAuthOverrides::default())
1886 .await?,
1887 "auth",
1888 )?;
1889 assert_eq!(auth.auth.api_key.as_deref(), Some("sk-env"));
1890 Ok(())
1891 }
1892
1893 #[tokio::test]
1894 async fn stored_oauth_marks_using_oauth() -> Result<(), Box<dyn std::error::Error>> {
1895 let store = Arc::new(InMemoryCredentialStore::new());
1896 store
1897 .modify(
1898 "openai-codex",
1899 Box::new(|_| {
1900 Box::pin(async {
1901 Ok(Some(Credential::Oauth(OAuthCredential {
1902 refresh: "r".into(),
1903 access: "a".into(),
1904 expires: i64::MAX,
1905 extra: BTreeMap::new(),
1906 })))
1907 })
1908 }),
1909 )
1910 .await?;
1911 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
1912 credentials: Some(store),
1913 models_store: Some(Arc::new(InMemoryModelsStore::new())),
1914 models_config: Some(ModelsJsonConfig::empty()),
1915 allow_model_network: Some(false),
1916 ..CreateModelRuntimeOptions::default()
1917 })
1918 .await?;
1919 assert!(runtime.has_configured_auth("openai-codex"));
1920 assert!(runtime.is_using_oauth("openai-codex"));
1921 Ok(())
1922 }
1923
1924 #[tokio::test]
1925 async fn registration_validation_errors_are_stable() -> Result<(), ModelRuntimeError> {
1926 let runtime = ModelRuntime::create_in_memory().await?;
1927 let Err(err) = runtime.register_provider(
1928 "broken",
1929 &ProviderConfigInput {
1930 models: Some(vec![ProviderModelDefinition {
1931 id: "m".into(),
1932 name: None,
1933 api: None,
1934 base_url: None,
1935 reasoning: false,
1936 thinking_level_map: None,
1937 input: None,
1938 cost: None,
1939 context_window: None,
1940 max_tokens: None,
1941 headers: None,
1942 compat: None,
1943 }]),
1944 ..ProviderConfigInput::default()
1945 },
1946 ) else {
1947 return Err(ModelRuntimeError::Registration(
1948 "invalid registration was accepted".to_owned(),
1949 ));
1950 };
1951 let message = err.to_string();
1952 assert!(
1953 message.contains("no \"api\" specified"),
1954 "unexpected message: {message}"
1955 );
1956 assert!(runtime.get_registered_provider_config("broken").is_none());
1957 Ok(())
1958 }
1959
1960 #[tokio::test]
1961 async fn stream_simple_without_auth_errors() -> Result<(), ModelRuntimeError> {
1962 let runtime = ModelRuntime::create_in_memory().await?;
1963 let model = required(runtime.get_model("anthropic", "claude-opus-4-8"), "model")?;
1964 let mut stream = runtime.stream_simple(model, Context::default(), StreamOptions::default());
1965 let first = required(futures::StreamExt::next(&mut stream).await, "event")?;
1966 let error = match first {
1967 Err(error) => error,
1968 Ok(event) => {
1969 return Err(ModelRuntimeError::Registration(format!(
1970 "expected infrastructure error, got {event:?}"
1971 )));
1972 }
1973 };
1974 assert!(
1975 error.to_string().contains("Provider is not configured"),
1976 "unexpected: {error}"
1977 );
1978 Ok(())
1979 }
1980
1981 struct RecordingExtensionProvider {
1983 calls: Mutex<Vec<RecordedStreamCall>>,
1984 reenter_runtime: Mutex<Option<ModelRuntime>>,
1986 reentered: AtomicBool,
1987 }
1988
1989 #[derive(Clone, Debug)]
1990 struct RecordedStreamCall {
1991 provider: String,
1992 api: String,
1993 base_url: String,
1994 api_key: Option<String>,
1995 }
1996
1997 impl RecordingExtensionProvider {
1998 fn new() -> Self {
1999 Self {
2000 calls: Mutex::new(Vec::new()),
2001 reenter_runtime: Mutex::new(None),
2002 reentered: AtomicBool::new(false),
2003 }
2004 }
2005
2006 fn calls(&self) -> Vec<RecordedStreamCall> {
2007 lock(&self.calls).clone()
2008 }
2009 }
2010
2011 impl Provider for RecordingExtensionProvider {
2012 fn stream(
2013 &self,
2014 model: &Model,
2015 _context: Context,
2016 options: StreamOptions,
2017 ) -> BoxStream<'static, Result<AssistantMessageEvent, ProviderError>> {
2018 lock(&self.calls).push(RecordedStreamCall {
2019 provider: model.provider.clone(),
2020 api: model.api.clone(),
2021 base_url: model.base_url.clone(),
2022 api_key: options.api_key.clone(),
2023 });
2024 if let Some(runtime) = lock(&self.reenter_runtime).clone() {
2025 runtime.register_extension_stream_provider(
2027 "lock-probe-sibling",
2028 Arc::new(RecordingExtensionProvider::new()),
2029 );
2030 self.reentered.store(true, Ordering::SeqCst);
2031 }
2032 Box::pin(stream::once(async {
2033 Err(ProviderError::new("extension-stream-hit"))
2034 }))
2035 }
2036 }
2037
2038 #[tokio::test]
2039 async fn stream_simple_routes_matching_extension_provider() -> Result<(), ModelRuntimeError> {
2040 let runtime = ModelRuntime::create_in_memory().await?;
2041 runtime.register_provider(
2042 "acme",
2043 &ProviderConfigInput {
2044 base_url: Some("https://acme.test/v1".to_owned()),
2045 api: Some("openai-completions".to_owned()),
2046 api_key: Some("sk-acme".to_owned()),
2047 models: Some(vec![custom_model("acme", "acme-1")]),
2048 ..ProviderConfigInput::default()
2049 },
2050 )?;
2051 let extension = Arc::new(RecordingExtensionProvider::new());
2052 runtime.register_extension_stream_provider("acme", extension.clone());
2053
2054 let model = required(runtime.get_model("acme", "acme-1"), "model")?;
2055 let mut stream = runtime.stream_simple(model, Context::default(), StreamOptions::default());
2056 let first = required(futures::StreamExt::next(&mut stream).await, "event")?;
2057 let error = match first {
2058 Err(error) => error,
2059 Ok(event) => {
2060 return Err(ModelRuntimeError::Registration(format!(
2061 "expected extension error, got {event:?}"
2062 )));
2063 }
2064 };
2065 assert_eq!(error.message(), "extension-stream-hit");
2066
2067 let calls = extension.calls();
2068 assert_eq!(calls.len(), 1);
2069 assert_eq!(calls[0].provider, "acme");
2070 assert_eq!(calls[0].api, "openai-completions");
2071 assert_eq!(calls[0].base_url, "https://example.test/v1");
2073 assert_eq!(calls[0].api_key.as_deref(), Some("sk-acme"));
2074 Ok(())
2075 }
2076
2077 #[tokio::test]
2078 async fn extension_provider_resolves_runtime_api_key() -> Result<(), ModelRuntimeError> {
2079 let runtime = ModelRuntime::create_in_memory().await?;
2080 runtime.register_provider(
2081 "verification",
2082 &ProviderConfigInput {
2083 base_url: Some("https://verification.invalid".to_owned()),
2084 api: Some("openai-completions".to_owned()),
2085 models: Some(vec![custom_model("verification", "model")]),
2086 ..ProviderConfigInput::default()
2087 },
2088 )?;
2089 runtime
2090 .set_runtime_api_key("verification", "verification-key")
2091 .await?;
2092 let extension = Arc::new(RecordingExtensionProvider::new());
2093 runtime.register_extension_stream_provider("verification", extension.clone());
2094
2095 let model = required(
2096 runtime.get_model("verification", "model"),
2097 "extension model",
2098 )?;
2099 let mut stream = runtime.stream_simple(model, Context::default(), StreamOptions::default());
2100 let first = required(futures::StreamExt::next(&mut stream).await, "event")?;
2101 assert!(
2102 matches!(&first, Err(error) if error.message() == "extension-stream-hit"),
2103 "extension provider must receive the prepared request: {first:?}"
2104 );
2105
2106 let calls = extension.calls();
2107 assert_eq!(calls.len(), 1);
2108 assert_eq!(calls[0].provider, "verification");
2109 assert_eq!(calls[0].api_key.as_deref(), Some("verification-key"));
2110 Ok(())
2111 }
2112
2113 #[tokio::test]
2114 async fn stream_simple_uses_native_when_api_mismatches() -> Result<(), ModelRuntimeError> {
2115 let runtime = ModelRuntime::create_in_memory().await?;
2116 runtime.register_provider(
2117 "acme",
2118 &ProviderConfigInput {
2119 base_url: Some("https://acme.test/v1".to_owned()),
2120 api: Some("openai-completions".to_owned()),
2121 api_key: Some("sk-acme".to_owned()),
2122 models: Some(vec![custom_model("acme", "acme-1")]),
2123 ..ProviderConfigInput::default()
2124 },
2125 )?;
2126 let extension = Arc::new(RecordingExtensionProvider::new());
2127 runtime.register_extension_stream_provider("acme", extension.clone());
2128
2129 let mut model = required(runtime.get_model("acme", "acme-1"), "model")?;
2130 model.api = "anthropic-messages".to_owned();
2131 let mut stream = runtime.stream_simple(model, Context::default(), StreamOptions::default());
2132 let _ = futures::StreamExt::next(&mut stream).await;
2133 assert!(
2134 extension.calls().is_empty(),
2135 "mismatched prepared API must not select extension stream"
2136 );
2137 Ok(())
2138 }
2139
2140 #[tokio::test]
2141 async fn stream_simple_uses_native_for_baseurl_only_registration()
2142 -> Result<(), ModelRuntimeError> {
2143 let runtime = ModelRuntime::create_in_memory().await?;
2144 runtime.register_provider(
2146 "openai",
2147 &ProviderConfigInput {
2148 base_url: Some("https://proxy.example/v1".to_owned()),
2149 api_key: Some("sk-proxy".to_owned()),
2150 ..ProviderConfigInput::default()
2151 },
2152 )?;
2153 let extension = Arc::new(RecordingExtensionProvider::new());
2154 runtime.register_extension_stream_provider("other", extension.clone());
2156
2157 let model = required(
2158 runtime
2159 .get_model("openai", "gpt-5.4")
2160 .or_else(|| runtime.get_models(Some("openai")).into_iter().next()),
2161 "openai model",
2162 )?;
2163 let mut stream = runtime.stream_simple(model, Context::default(), StreamOptions::default());
2164 let _ = futures::StreamExt::next(&mut stream).await;
2165 assert!(
2166 extension.calls().is_empty(),
2167 "baseURL-only / unregistered stream must stay on native path"
2168 );
2169 Ok(())
2170 }
2171
2172 #[tokio::test]
2173 async fn stream_simple_unregister_restores_native() -> Result<(), ModelRuntimeError> {
2174 let runtime = ModelRuntime::create_in_memory().await?;
2175 runtime.register_provider(
2176 "acme",
2177 &ProviderConfigInput {
2178 base_url: Some("https://acme.test/v1".to_owned()),
2179 api: Some("openai-completions".to_owned()),
2180 api_key: Some("sk-acme".to_owned()),
2181 models: Some(vec![custom_model("acme", "acme-1")]),
2182 ..ProviderConfigInput::default()
2183 },
2184 )?;
2185 let extension = Arc::new(RecordingExtensionProvider::new());
2186 runtime.register_extension_stream_provider("acme", extension.clone());
2187 runtime.unregister_extension_stream_provider("acme");
2188
2189 let model = required(runtime.get_model("acme", "acme-1"), "model")?;
2190 let mut stream = runtime.stream_simple(model, Context::default(), StreamOptions::default());
2191 let _ = futures::StreamExt::next(&mut stream).await;
2192 assert!(
2193 extension.calls().is_empty(),
2194 "unregistered extension stream must fall back to native"
2195 );
2196 Ok(())
2197 }
2198
2199 #[tokio::test]
2200 async fn stream_simple_releases_lock_before_stream() -> Result<(), ModelRuntimeError> {
2201 let runtime = ModelRuntime::create_in_memory().await?;
2202 runtime.register_provider(
2203 "acme",
2204 &ProviderConfigInput {
2205 base_url: Some("https://acme.test/v1".to_owned()),
2206 api: Some("openai-completions".to_owned()),
2207 api_key: Some("sk-acme".to_owned()),
2208 models: Some(vec![custom_model("acme", "acme-1")]),
2209 ..ProviderConfigInput::default()
2210 },
2211 )?;
2212 let extension = Arc::new(RecordingExtensionProvider::new());
2213 *lock(&extension.reenter_runtime) = Some(runtime.clone());
2214 runtime.register_extension_stream_provider("acme", extension.clone());
2215
2216 let model = required(runtime.get_model("acme", "acme-1"), "model")?;
2217 let mut stream = runtime.stream_simple(model, Context::default(), StreamOptions::default());
2218 let first = required(futures::StreamExt::next(&mut stream).await, "event")?;
2219 assert!(
2220 matches!(&first, Err(error) if error.message() == "extension-stream-hit"),
2221 "stream must complete without deadlock: {first:?}"
2222 );
2223 assert!(
2224 extension.reentered.load(Ordering::SeqCst),
2225 "stream must re-enter registration, proving map lock was released"
2226 );
2227 Ok(())
2228 }
2229
2230 #[tokio::test]
2231 async fn reregistration_merges_defined_fields() -> Result<(), ModelRuntimeError> {
2232 let runtime = ModelRuntime::create_in_memory().await?;
2233 runtime.register_provider(
2234 "acme",
2235 &ProviderConfigInput {
2236 base_url: Some("https://acme.test/v1".to_owned()),
2237 api: Some("openai-completions".to_owned()),
2238 api_key: Some("sk-1".to_owned()),
2239 models: Some(vec![custom_model("acme", "acme-1")]),
2240 ..ProviderConfigInput::default()
2241 },
2242 )?;
2243 runtime.register_provider(
2244 "acme",
2245 &ProviderConfigInput {
2246 api_key: Some("sk-2".to_owned()),
2247 ..ProviderConfigInput::default()
2248 },
2249 )?;
2250 let config = required(runtime.get_registered_provider_config("acme"), "config")?;
2251 assert_eq!(config.api_key.as_deref(), Some("sk-2"));
2252 assert_eq!(config.base_url.as_deref(), Some("https://acme.test/v1"));
2253 assert!(config.models.is_some());
2254 Ok(())
2255 }
2256
2257 #[test]
2258 fn strip_json_comments_preserves_strings() {
2259 let raw = r#"{ "a": "http://x", /* c */ "b": 1 } // tail"#;
2260 let stripped = strip_json_comments(raw);
2261 assert!(stripped.contains("http://x"));
2262 assert!(!stripped.contains("/*"));
2263 assert!(!stripped.contains("// tail"));
2264 }
2265
2266 #[tokio::test]
2267 async fn models_json_upserts_by_id_not_replace() -> Result<(), ModelRuntimeError> {
2268 let mut providers = BTreeMap::new();
2269 providers.insert(
2270 "anthropic".to_owned(),
2271 ProviderConfigInput {
2272 models: Some(vec![ProviderModelDefinition {
2273 id: "claude-opus-4-8".into(),
2274 name: Some("Upserted Opus".into()),
2275 api: None,
2276 base_url: None,
2277 reasoning: true,
2278 thinking_level_map: None,
2279 input: None,
2280 cost: None,
2281 context_window: None,
2282 max_tokens: None,
2283 headers: None,
2284 compat: None,
2285 }]),
2286 ..ProviderConfigInput::default()
2287 },
2288 );
2289 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
2290 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
2291 models_store: Some(Arc::new(InMemoryModelsStore::new())),
2292 models_config: Some(ModelsJsonConfig::from_providers(providers)),
2293 allow_model_network: Some(false),
2294 ..CreateModelRuntimeOptions::default()
2295 })
2296 .await?;
2297 let models = runtime.get_models(Some("anthropic"));
2298 assert!(
2299 models.len() > 1,
2300 "upsert must keep other builtins, got {}",
2301 models.len()
2302 );
2303 let upserted = required(
2304 runtime.get_model("anthropic", "claude-opus-4-8"),
2305 "upserted",
2306 )?;
2307 assert_eq!(upserted.name, "Upserted Opus");
2308 assert!(upserted.base_url.contains("anthropic"));
2310 assert_eq!(upserted.api, "anthropic-messages");
2311 Ok(())
2312 }
2313
2314 #[tokio::test]
2315 async fn configured_base_url_projects_except_radius_oauth() -> Result<(), ModelRuntimeError> {
2316 let mut providers = BTreeMap::new();
2317 providers.insert(
2318 "anthropic".to_owned(),
2319 ProviderConfigInput {
2320 base_url: Some("https://proxy.example/v1".into()),
2321 ..ProviderConfigInput::default()
2322 },
2323 );
2324 providers.insert(
2325 "radius".to_owned(),
2326 ProviderConfigInput {
2327 base_url: Some("https://radius.example".into()),
2328 oauth: Some("radius".into()),
2329 models: Some(vec![ProviderModelDefinition {
2330 id: "auto".into(),
2331 name: Some("auto".into()),
2332 api: Some("openai-completions".into()),
2333 base_url: Some("https://radius-builtin.example".into()),
2334 reasoning: false,
2335 thinking_level_map: None,
2336 input: Some(vec![ModelInput::Text]),
2337 cost: Some(ModelCost::default()),
2338 context_window: Some(8_000),
2339 max_tokens: Some(1_024),
2340 headers: None,
2341 compat: None,
2342 }]),
2343 ..ProviderConfigInput::default()
2344 },
2345 );
2346 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
2347 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
2348 models_store: Some(Arc::new(InMemoryModelsStore::new())),
2349 models_config: Some(ModelsJsonConfig::from_providers(providers)),
2350 allow_model_network: Some(false),
2351 ..CreateModelRuntimeOptions::default()
2352 })
2353 .await?;
2354 let anthropic = required(
2355 runtime.get_model("anthropic", "claude-opus-4-8"),
2356 "anthropic",
2357 )?;
2358 assert_eq!(anthropic.base_url, "https://proxy.example/v1");
2359 let radius = required(runtime.get_model("radius", "auto"), "radius")?;
2360 assert_eq!(radius.base_url, "https://radius-builtin.example");
2361 Ok(())
2362 }
2363
2364 #[tokio::test]
2365 async fn auth_header_true_emits_bearer_and_requires_key() -> Result<(), ModelRuntimeError> {
2366 let mut providers = BTreeMap::new();
2367 providers.insert(
2368 "openai".to_owned(),
2369 ProviderConfigInput {
2370 auth_header: Some(true),
2371 headers: Some(BTreeMap::from([("X-Custom".into(), "yes".into())])),
2372 ..ProviderConfigInput::default()
2373 },
2374 );
2375 let mut env = ProviderEnv::new();
2376 env.insert("OPENAI_API_KEY".into(), "sk-env".into());
2377 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
2378 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
2379 models_store: Some(Arc::new(InMemoryModelsStore::new())),
2380 models_config: Some(ModelsJsonConfig::from_providers(providers)),
2381 allow_model_network: Some(false),
2382 auth_env: Some(env),
2383 ..CreateModelRuntimeOptions::default()
2384 })
2385 .await?;
2386 let auth = required(
2387 runtime
2388 .get_auth_for_provider("openai", ModelRuntimeAuthOverrides::default())
2389 .await?,
2390 "auth",
2391 )?;
2392 let headers = required(auth.auth.headers, "headers")?;
2393 assert_eq!(
2394 headers.get("Authorization").map(String::as_str),
2395 Some("Bearer sk-env")
2396 );
2397 assert_eq!(headers.get("X-Custom").map(String::as_str), Some("yes"));
2398
2399 let mut providers = BTreeMap::new();
2401 providers.insert(
2402 "acme".to_owned(),
2403 ProviderConfigInput {
2404 auth_header: Some(true),
2405 base_url: Some("https://acme.test".into()),
2406 api: Some("openai-completions".into()),
2407 models: Some(vec![custom_model("acme", "m")]),
2408 ..ProviderConfigInput::default()
2409 },
2410 );
2411 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
2412 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
2413 models_store: Some(Arc::new(InMemoryModelsStore::new())),
2414 models_config: Some(ModelsJsonConfig::from_providers(providers)),
2415 allow_model_network: Some(false),
2416 ..CreateModelRuntimeOptions::default()
2417 })
2418 .await?;
2419 let err = runtime
2428 .get_auth_for_provider(
2429 "acme",
2430 ModelRuntimeAuthOverrides {
2431 api_key: Some(String::new()),
2432 env: None,
2433 },
2434 )
2435 .await;
2436 assert!(runtime.configured_auth_header("acme"));
2440 let _ = err;
2441 Ok(())
2442 }
2443
2444 #[tokio::test]
2445 async fn expired_oauth_refreshes_via_resolver() -> Result<(), Box<dyn std::error::Error>> {
2446 #[derive(Clone)]
2447 struct FakeOauth {
2448 refreshed: Arc<std::sync::atomic::AtomicBool>,
2449 }
2450 impl OAuthAuth for FakeOauth {
2451 fn name(&self) -> &'static str {
2452 "Fake OAuth"
2453 }
2454 fn login_label(&self) -> Option<&str> {
2455 Some("Fake")
2456 }
2457 fn login<'a>(
2458 &'a self,
2459 _interaction: &'a dyn pi_ai::auth::AuthInteraction,
2460 ) -> futures::future::BoxFuture<'a, Result<OAuthCredential, pi_ai::auth::AuthError>>
2461 {
2462 Box::pin(async { Err(pi_ai::auth::AuthError::message("not used")) })
2463 }
2464 fn refresh<'a>(
2465 &'a self,
2466 credential: &'a OAuthCredential,
2467 _signal: Option<tokio_util::sync::CancellationToken>,
2468 ) -> futures::future::BoxFuture<'a, Result<OAuthCredential, pi_ai::auth::AuthError>>
2469 {
2470 let refreshed = Arc::clone(&self.refreshed);
2471 let mut next = credential.clone();
2472 Box::pin(async move {
2473 refreshed.store(true, std::sync::atomic::Ordering::SeqCst);
2474 next.access = "fresh-access".into();
2475 next.expires = i64::MAX;
2476 Ok(next)
2477 })
2478 }
2479 fn to_auth<'a>(
2480 &'a self,
2481 credential: &'a OAuthCredential,
2482 ) -> futures::future::BoxFuture<'a, Result<ModelAuth, pi_ai::auth::AuthError>>
2483 {
2484 Box::pin(async move {
2485 Ok(ModelAuth {
2486 api_key: Some(credential.access.clone()),
2487 headers: None,
2488 base_url: None,
2489 })
2490 })
2491 }
2492 }
2493
2494 let store = Arc::new(InMemoryCredentialStore::new());
2495 store
2496 .modify(
2497 "openai-codex",
2498 Box::new(|_| {
2499 Box::pin(async {
2500 Ok(Some(Credential::Oauth(OAuthCredential {
2501 refresh: "r".into(),
2502 access: "stale-access".into(),
2503 expires: 0, extra: BTreeMap::new(),
2505 })))
2506 })
2507 }),
2508 )
2509 .await?;
2510
2511 let flag = Arc::new(std::sync::atomic::AtomicBool::new(false));
2512 let mut handlers = HashMap::new();
2513 handlers.insert(
2514 "openai-codex".to_owned(),
2515 Arc::new(FakeOauth {
2516 refreshed: Arc::clone(&flag),
2517 }) as Arc<dyn OAuthAuth>,
2518 );
2519
2520 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
2521 credentials: Some(store),
2522 models_store: Some(Arc::new(InMemoryModelsStore::new())),
2523 models_config: Some(ModelsJsonConfig::empty()),
2524 allow_model_network: Some(false),
2525 oauth_handlers: Some(handlers),
2526 ..CreateModelRuntimeOptions::default()
2527 })
2528 .await?;
2529
2530 let auth = required(
2531 runtime
2532 .get_auth_for_provider("openai-codex", ModelRuntimeAuthOverrides::default())
2533 .await?,
2534 "auth",
2535 )?;
2536 assert!(
2537 flag.load(std::sync::atomic::Ordering::SeqCst),
2538 "refresh must be called for expired oauth"
2539 );
2540 assert_eq!(auth.auth.api_key.as_deref(), Some("fresh-access"));
2541 Ok(())
2542 }
2543
2544 #[tokio::test]
2545 async fn configured_base_url_not_copied_into_auth_result() -> Result<(), ModelRuntimeError> {
2546 let mut providers = BTreeMap::new();
2549 providers.insert(
2550 "openai".to_owned(),
2551 ProviderConfigInput {
2552 base_url: Some("https://proxy.example/v1".into()),
2553 ..ProviderConfigInput::default()
2554 },
2555 );
2556 providers.insert(
2559 "radius".to_owned(),
2560 ProviderConfigInput {
2561 base_url: Some("https://radius-config.example".into()),
2562 oauth: Some("radius".into()),
2563 api_key: Some("sk-radius".into()),
2564 models: Some(vec![ProviderModelDefinition {
2565 id: "auto".into(),
2566 name: Some("auto".into()),
2567 api: Some("openai-completions".into()),
2568 base_url: Some("https://radius-model.example".into()),
2569 reasoning: false,
2570 thinking_level_map: None,
2571 input: Some(vec![ModelInput::Text]),
2572 cost: Some(ModelCost::default()),
2573 context_window: Some(8_000),
2574 max_tokens: Some(1_024),
2575 headers: None,
2576 compat: None,
2577 }]),
2578 ..ProviderConfigInput::default()
2579 },
2580 );
2581 let mut env = ProviderEnv::new();
2582 env.insert("OPENAI_API_KEY".into(), "sk-env".into());
2583 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
2584 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
2585 models_store: Some(Arc::new(InMemoryModelsStore::new())),
2586 models_config: Some(ModelsJsonConfig::from_providers(providers)),
2587 allow_model_network: Some(false),
2588 auth_env: Some(env),
2589 ..CreateModelRuntimeOptions::default()
2590 })
2591 .await?;
2592
2593 let openai_model = required(
2594 runtime.get_model("openai", "gpt-5.5"),
2595 "openai model composition gets proxy baseUrl",
2596 )?;
2597 assert_eq!(openai_model.base_url, "https://proxy.example/v1");
2598 let openai_auth = required(
2599 runtime
2600 .get_auth_for_provider("openai", ModelRuntimeAuthOverrides::default())
2601 .await?,
2602 "openai auth",
2603 )?;
2604 assert!(
2605 openai_auth.auth.base_url.is_none(),
2606 "configured provider baseUrl must not land on AuthResult"
2607 );
2608
2609 let radius_model = required(runtime.get_model("radius", "auto"), "radius model")?;
2610 assert_eq!(radius_model.base_url, "https://radius-model.example");
2611 let radius_auth = required(
2612 runtime
2613 .get_auth_for_provider("radius", ModelRuntimeAuthOverrides::default())
2614 .await?,
2615 "radius auth via configured api key",
2616 )?;
2617 assert!(
2618 radius_auth.auth.base_url.is_none(),
2619 "radius configured baseUrl must not land on AuthResult either"
2620 );
2621 Ok(())
2622 }
2623
2624 #[tokio::test]
2625 async fn oauth_auth_base_url_is_preserved() -> Result<(), Box<dyn std::error::Error>> {
2626 #[derive(Clone)]
2627 struct OauthWithBaseUrl;
2628 impl OAuthAuth for OauthWithBaseUrl {
2629 fn name(&self) -> &'static str {
2630 "OAuth BaseUrl"
2631 }
2632 fn login_label(&self) -> Option<&str> {
2633 Some("OAuth BaseUrl")
2634 }
2635 fn login<'a>(
2636 &'a self,
2637 _interaction: &'a dyn pi_ai::auth::AuthInteraction,
2638 ) -> futures::future::BoxFuture<'a, Result<OAuthCredential, pi_ai::auth::AuthError>>
2639 {
2640 Box::pin(async { Err(pi_ai::auth::AuthError::message("not used")) })
2641 }
2642 fn refresh<'a>(
2643 &'a self,
2644 credential: &'a OAuthCredential,
2645 _signal: Option<tokio_util::sync::CancellationToken>,
2646 ) -> futures::future::BoxFuture<'a, Result<OAuthCredential, pi_ai::auth::AuthError>>
2647 {
2648 Box::pin(async move { Ok(credential.clone()) })
2649 }
2650 fn to_auth<'a>(
2651 &'a self,
2652 credential: &'a OAuthCredential,
2653 ) -> futures::future::BoxFuture<'a, Result<ModelAuth, pi_ai::auth::AuthError>>
2654 {
2655 Box::pin(async move {
2656 Ok(ModelAuth {
2657 api_key: Some(credential.access.clone()),
2658 headers: None,
2659 base_url: Some("https://oauth-endpoint.example".into()),
2660 })
2661 })
2662 }
2663 }
2664
2665 let store = Arc::new(InMemoryCredentialStore::new());
2666 store
2667 .modify(
2668 "openai-codex",
2669 Box::new(|_| {
2670 Box::pin(async {
2671 Ok(Some(Credential::Oauth(OAuthCredential {
2672 refresh: "r".into(),
2673 access: "access".into(),
2674 expires: i64::MAX,
2675 extra: BTreeMap::new(),
2676 })))
2677 })
2678 }),
2679 )
2680 .await?;
2681
2682 let mut handlers = HashMap::new();
2683 handlers.insert(
2684 "openai-codex".to_owned(),
2685 Arc::new(OauthWithBaseUrl) as Arc<dyn OAuthAuth>,
2686 );
2687 let mut providers = BTreeMap::new();
2689 providers.insert(
2690 "openai-codex".to_owned(),
2691 ProviderConfigInput {
2692 base_url: Some("https://configured-should-not-win.example".into()),
2693 ..ProviderConfigInput::default()
2694 },
2695 );
2696
2697 let runtime = ModelRuntime::create(CreateModelRuntimeOptions {
2698 credentials: Some(store),
2699 models_store: Some(Arc::new(InMemoryModelsStore::new())),
2700 models_config: Some(ModelsJsonConfig::from_providers(providers)),
2701 allow_model_network: Some(false),
2702 oauth_handlers: Some(handlers),
2703 ..CreateModelRuntimeOptions::default()
2704 })
2705 .await?;
2706
2707 let auth = required(
2708 runtime
2709 .get_auth_for_provider("openai-codex", ModelRuntimeAuthOverrides::default())
2710 .await?,
2711 "auth",
2712 )?;
2713 assert_eq!(
2714 auth.auth.base_url.as_deref(),
2715 Some("https://oauth-endpoint.example"),
2716 "OAuth to_auth baseUrl must be preserved"
2717 );
2718 Ok(())
2719 }
2720}