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