| use std::sync::Arc; |
| use std::time::Duration; |
|
|
| use forge_app::{AuthStrategy, ProviderAuthService, StrategyFactory}; |
| use forge_domain::{ |
| AuthContextRequest, AuthContextResponse, AuthMethod, Provider, ProviderId, ProviderRepository, |
| }; |
|
|
| |
| #[derive(Clone)] |
| pub struct ForgeProviderAuthService<I> { |
| infra: Arc<I>, |
| } |
|
|
| impl<I> ForgeProviderAuthService<I> { |
| |
| pub fn new(infra: Arc<I>) -> Self { |
| Self { infra } |
| } |
| } |
|
|
| #[async_trait::async_trait] |
| impl<I> ProviderAuthService for ForgeProviderAuthService<I> |
| where |
| I: StrategyFactory + ProviderRepository + Send + Sync + 'static, |
| { |
| |
| async fn init_provider_auth( |
| &self, |
| provider_id: ProviderId, |
| auth_method: AuthMethod, |
| ) -> anyhow::Result<AuthContextRequest> { |
| |
| let required_params = if matches!( |
| auth_method, |
| AuthMethod::ApiKey | AuthMethod::GoogleAdc | AuthMethod::AwsProfile |
| ) { |
| |
| |
| let providers = self.infra.get_all_providers().await?; |
| let provider = providers |
| .iter() |
| .find(|p| p.id() == provider_id) |
| .ok_or_else(|| forge_domain::Error::provider_not_available(provider_id.clone()))?; |
| provider.url_params().to_vec() |
| } else { |
| vec![] |
| }; |
|
|
| |
| let strategy = self.infra.create_auth_strategy( |
| provider_id.clone(), |
| auth_method.clone(), |
| required_params, |
| )?; |
| let mut request = strategy.init().await?; |
|
|
| |
| if let AuthContextRequest::ApiKey(ref mut api_key_request) = request |
| && let Ok(Some(existing_credential)) = self.infra.get_credential(&provider_id).await |
| { |
| api_key_request.existing_params = Some(existing_credential.url_params.into()); |
|
|
| |
| |
| |
| if !matches!(auth_method, AuthMethod::GoogleAdc | AuthMethod::AwsProfile) |
| && let Some(key) = existing_credential.auth_details.api_key() |
| { |
| let is_adc_marker = key.as_ref() == "google_adc_marker"; |
| if !is_adc_marker { |
| api_key_request.api_key = Some(key.clone()); |
| } |
| } |
| } |
|
|
| Ok(request) |
| } |
|
|
| |
| async fn complete_provider_auth( |
| &self, |
| provider_id: ProviderId, |
| auth_context_response: AuthContextResponse, |
| _timeout: Duration, |
| ) -> anyhow::Result<()> { |
| |
| |
| let auth_method = match &auth_context_response { |
| AuthContextResponse::ApiKey(response) => { |
| |
| let is_vertex_provider = provider_id == forge_domain::ProviderId::VERTEX_AI |
| || provider_id == forge_domain::ProviderId::VERTEX_AI_ANTHROPIC; |
| if is_vertex_provider && response.response.api_key.as_ref() == "google_adc_marker" { |
| |
| forge_domain::AuthMethod::google_adc() |
| } else if response.response.api_key.as_ref() == "aws_profile_marker" { |
| |
| forge_domain::AuthMethod::AwsProfile |
| } else { |
| |
| forge_domain::AuthMethod::ApiKey |
| } |
| } |
| AuthContextResponse::Code(ctx) => { |
| AuthMethod::OAuthCode(ctx.request.oauth_config.clone()) |
| } |
| AuthContextResponse::DeviceCode(ctx) => { |
| if provider_id == forge_domain::ProviderId::CODEX { |
| AuthMethod::CodexDevice(ctx.request.oauth_config.clone()) |
| } else { |
| AuthMethod::OAuthDevice(ctx.request.oauth_config.clone()) |
| } |
| } |
| }; |
|
|
| |
| let required_params = if matches!(auth_method, AuthMethod::ApiKey) { |
| |
| |
| let providers = self.infra.get_all_providers().await?; |
| let provider = providers |
| .iter() |
| .find(|p| p.id() == provider_id) |
| .ok_or_else(|| forge_domain::Error::provider_not_available(provider_id.clone()))?; |
| provider.url_params().to_vec() |
| } else { |
| vec![] |
| }; |
|
|
| |
| let strategy = |
| self.infra |
| .create_auth_strategy(provider_id.clone(), auth_method, required_params)?; |
| let credential = strategy.complete(auth_context_response).await?; |
|
|
| |
| self.infra.upsert_credential(credential).await |
| } |
|
|
| |
| |
| |
| |
| |
| async fn refresh_provider_credential( |
| &self, |
| mut provider: Provider<url::Url>, |
| ) -> anyhow::Result<Provider<url::Url>> { |
| |
| if let Some(credential) = &provider.credential { |
| let buffer = chrono::Duration::minutes(5); |
|
|
| if credential.needs_refresh(buffer) { |
| |
| for auth_method in &provider.auth_methods { |
| match auth_method { |
| AuthMethod::OAuthDevice(_) |
| | AuthMethod::OAuthCode(_) |
| | AuthMethod::CodexDevice(_) |
| | AuthMethod::GoogleAdc => { |
| |
| let existing_credential = |
| self.infra.get_credential(&provider.id).await?.ok_or_else( |
| || forge_domain::Error::ProviderNotAvailable { |
| provider: provider.id.clone(), |
| }, |
| )?; |
|
|
| |
| let required_params = if matches!(auth_method, AuthMethod::ApiKey) { |
| provider.url_params.clone() |
| } else { |
| vec![] |
| }; |
|
|
| |
| if let Ok(strategy) = self.infra.create_auth_strategy( |
| provider.id.clone(), |
| auth_method.clone(), |
| required_params, |
| ) { |
| match strategy.refresh(&existing_credential).await { |
| Ok(refreshed) => { |
| |
| if self |
| .infra |
| .upsert_credential(refreshed.clone()) |
| .await |
| .is_err() |
| { |
| continue; |
| } |
|
|
| |
| provider.credential = Some(refreshed); |
| break; |
| } |
| Err(_) => { |
| |
| |
| } |
| } |
| } |
| } |
| _ => {} |
| } |
| } |
| } |
| } |
|
|
| Ok(provider) |
| } |
| } |
|
|