Skip to main content

elph_ai/auth/
resolve.rs

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}