1use std::sync::Arc;
2
3use thiserror::Error;
4
5use super::types::{
6 ApiKeyCredential, AuthContext, AuthModel, AuthResult, Credential, CredentialStore, OAuthCredential, ProviderAuth,
7};
8use crate::types::ProviderEnv;
9
10#[derive(Debug, Clone, Copy, PartialEq, Eq)]
11pub enum ModelsErrorCode {
12 ModelSource,
13 ModelValidation,
14 Provider,
15 Stream,
16 Auth,
17 OAuth,
18}
19
20#[derive(Debug, Error)]
21#[error("{code:?}: {message}")]
22pub struct ModelsError {
23 pub code: ModelsErrorCode,
24 pub message: String,
25 #[source]
26 pub cause: Option<anyhow::Error>,
27}
28
29impl ModelsError {
30 pub fn new(code: ModelsErrorCode, message: impl Into<String>) -> Self {
31 Self {
32 code,
33 message: message.into(),
34 cause: None,
35 }
36 }
37
38 pub fn with_cause(code: ModelsErrorCode, message: impl Into<String>, cause: anyhow::Error) -> Self {
39 Self {
40 code,
41 message: message.into(),
42 cause: Some(cause),
43 }
44 }
45}
46
47pub struct AuthResolutionOverrides {
48 pub api_key: Option<String>,
49 pub env: Option<ProviderEnv>,
50}
51
52pub async fn resolve_provider_auth(
53 provider: &ProviderAuthHolder,
54 model: AuthModel,
55 credentials: &dyn CredentialStore,
56 auth_context: Arc<dyn AuthContext>,
57 overrides: Option<AuthResolutionOverrides>,
58) -> Result<Option<AuthResult>, ModelsError> {
59 let ctx = if let Some(env) = overrides.as_ref().and_then(|o| o.env.clone()) {
60 Arc::new(OverlayAuthContext {
61 base: auth_context.clone(),
62 env,
63 }) as Arc<dyn AuthContext>
64 } else {
65 auth_context
66 };
67
68 if let Some(key) = overrides.as_ref().and_then(|o| o.api_key.clone())
69 && let Some(api_key) = &provider.auth.api_key
70 {
71 return resolve_api_key(
72 ctx,
73 api_key,
74 model,
75 Some(ApiKeyCredential::new(key)),
76 overrides.as_ref().and_then(|o| o.env.clone()),
77 )
78 .await;
79 }
80
81 let stored = credentials.read(&provider.id).await;
82 if let Some(stored) = stored {
83 return match stored {
84 Credential::OAuth(cred) => {
85 if let Some(oauth) = &provider.auth.oauth {
86 resolve_stored_oauth(credentials, &provider.id, oauth, cred).await
87 } else {
88 Ok(None)
89 }
90 }
91 Credential::ApiKey(cred) => {
92 if let Some(api_key) = &provider.auth.api_key {
93 let merged = if let Some(env) = overrides.as_ref().and_then(|o| o.env.clone()) {
94 let mut c = cred.clone();
95 c.env = Some(c.env.unwrap_or_default().into_iter().chain(env).collect());
96 c
97 } else {
98 cred
99 };
100 Ok(resolve_api_key(ctx, api_key, model, Some(merged), None).await?)
101 } else {
102 Ok(None)
103 }
104 }
105 };
106 }
107
108 if let Some(api_key) = &provider.auth.api_key {
109 return resolve_api_key(ctx, api_key, model, None, overrides.and_then(|o| o.env)).await;
110 }
111
112 Ok(None)
113}
114
115pub struct ProviderAuthHolder {
116 pub id: String,
117 pub auth: ProviderAuth,
118}
119
120struct OverlayAuthContext {
121 base: Arc<dyn AuthContext>,
122 env: ProviderEnv,
123}
124
125#[async_trait::async_trait]
126impl AuthContext for OverlayAuthContext {
127 async fn env(&self, name: &str) -> Option<String> {
128 self.env.get(name).cloned().or(self.base.env(name).await)
129 }
130
131 async fn file_exists(&self, path: &str) -> bool {
132 self.base.file_exists(path).await
133 }
134}
135
136async fn resolve_stored_oauth(
137 credentials: &dyn CredentialStore,
138 provider_id: &str,
139 oauth: &super::types::OAuthAuth,
140 mut stored: OAuthCredential,
141) -> Result<Option<AuthResult>, ModelsError> {
142 if chrono::Utc::now().timestamp_millis() >= stored.expires {
143 let oauth = oauth.clone();
144 let refresh_result = (oauth.refresh)(stored.clone()).await;
145 let refreshed = match refresh_result {
146 Ok(next) => next,
147 Err(e) => {
148 return Err(ModelsError::with_cause(
149 ModelsErrorCode::OAuth,
150 format!("OAuth refresh failed for {provider_id}"),
151 e,
152 ));
153 }
154 };
155 let post = credentials
156 .modify(
157 provider_id,
158 Box::new(move |current| {
159 let refreshed = refreshed.clone();
160 Box::pin(async move {
161 let Some(Credential::OAuth(current)) = current else {
162 return None;
163 };
164 if chrono::Utc::now().timestamp_millis() < current.expires {
165 return None;
166 }
167 Some(Credential::OAuth(refreshed))
168 })
169 }),
170 )
171 .await;
172
173 if let Some(Credential::OAuth(cred)) = post {
174 stored = cred;
175 } else {
176 return Ok(None);
177 }
178 }
179
180 match (oauth.to_auth)(stored).await {
181 Ok(auth) => Ok(Some(AuthResult {
182 auth,
183 env: None,
184 source: Some("OAuth".to_string()),
185 })),
186 Err(e) => Err(ModelsError::with_cause(
187 ModelsErrorCode::OAuth,
188 format!("OAuth auth derivation failed for {provider_id}"),
189 e,
190 )),
191 }
192}
193
194async fn resolve_api_key(
195 ctx: Arc<dyn AuthContext>,
196 auth: &super::types::ApiKeyAuth,
197 model: AuthModel,
198 credential: Option<ApiKeyCredential>,
199 env_override: Option<ProviderEnv>,
200) -> Result<Option<AuthResult>, ModelsError> {
201 let input = super::types::AuthResolveInput { model, ctx, credential };
202 if let Some(mut result) = (auth.resolve)(input).await {
203 if let Some(env) = env_override {
204 result.env = Some(result.env.unwrap_or_default().into_iter().chain(env).collect());
205 }
206 Ok(Some(result))
207 } else {
208 Ok(None)
209 }
210}