Skip to main content

elph_ai/auth/
resolve.rs

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}