mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
Refactor provider transport modules
This commit is contained in:
352
crates/aether-provider-transport/src/generic_oauth/mod.rs
Normal file
352
crates/aether-provider-transport/src/generic_oauth/mod.rs
Normal file
@@ -0,0 +1,352 @@
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
use aether_oauth::provider::providers::{
|
||||
GenericProviderOAuthAdapter, GENERIC_PROVIDER_OAUTH_TEMPLATES,
|
||||
};
|
||||
use aether_oauth::provider::{ProviderOAuthAccount, ProviderOAuthAdapter, ProviderOAuthTokenSet};
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::oauth_refresh::{
|
||||
oauth_error_to_local_refresh_error, provider_oauth_transport_context_from_snapshot,
|
||||
CachedOAuthEntry, LocalOAuthHttpExecutor, LocalOAuthRefreshAdapter, LocalOAuthRefreshError,
|
||||
LocalResolvedOAuthRequestAuth, ProviderOAuthLocalHttpExecutor,
|
||||
};
|
||||
use super::snapshot::GatewayProviderTransportSnapshot;
|
||||
|
||||
const AUTH_HEADER_NAME: &str = "authorization";
|
||||
const OAUTH_REFRESH_SKEW_SECS: u64 = 120;
|
||||
const PLACEHOLDER_API_KEY: &str = "__placeholder__";
|
||||
|
||||
pub fn supports_local_generic_oauth_request_auth_resolution(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> bool {
|
||||
transport.key.auth_type.trim().eq_ignore_ascii_case("oauth")
|
||||
&& generic_provider_type(transport.provider.provider_type.as_str()).is_some()
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct GenericOAuthRefreshAdapter {
|
||||
token_url_overrides: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
impl GenericOAuthRefreshAdapter {
|
||||
pub fn with_token_url_for_tests(
|
||||
mut self,
|
||||
provider_type: &str,
|
||||
token_url: impl Into<String>,
|
||||
) -> Self {
|
||||
self.token_url_overrides
|
||||
.insert(provider_type.trim().to_ascii_lowercase(), token_url.into());
|
||||
self
|
||||
}
|
||||
|
||||
fn adapter_for_provider_type(
|
||||
&self,
|
||||
provider_type: &'static str,
|
||||
) -> Option<GenericProviderOAuthAdapter> {
|
||||
let adapter = GenericProviderOAuthAdapter::for_provider_type(provider_type)?;
|
||||
if let Some(token_url) = self.token_url_overrides.get(provider_type) {
|
||||
return Some(adapter.with_token_url_override(token_url.clone()));
|
||||
}
|
||||
Some(adapter)
|
||||
}
|
||||
|
||||
fn auth_config_from_transport(transport: &GatewayProviderTransportSnapshot) -> Option<Value> {
|
||||
transport
|
||||
.key
|
||||
.decrypted_auth_config
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(|value| serde_json::from_str::<Value>(value).ok())
|
||||
}
|
||||
|
||||
fn auth_config_from_entry(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<Value> {
|
||||
entry
|
||||
.metadata
|
||||
.as_ref()
|
||||
.filter(|_| {
|
||||
entry
|
||||
.provider_type
|
||||
.eq_ignore_ascii_case(transport.provider.provider_type.as_str())
|
||||
})
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn auth_config_updated_at(auth_config: &Value) -> Option<u64> {
|
||||
auth_config
|
||||
.as_object()
|
||||
.and_then(|object| object.get("updated_at"))
|
||||
.and_then(|value| parse_u64_value(Some(value)))
|
||||
}
|
||||
|
||||
fn base_auth_config(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<Value> {
|
||||
let cached = entry.and_then(|cached| Self::auth_config_from_entry(transport, cached));
|
||||
let transport_auth = Self::auth_config_from_transport(transport);
|
||||
|
||||
match (cached, transport_auth) {
|
||||
(Some(cached), Some(transport_auth)) => {
|
||||
let cached_updated_at = Self::auth_config_updated_at(&cached);
|
||||
let transport_updated_at = Self::auth_config_updated_at(&transport_auth);
|
||||
if transport_updated_at > cached_updated_at {
|
||||
Some(transport_auth)
|
||||
} else {
|
||||
Some(cached)
|
||||
}
|
||||
}
|
||||
(Some(cached), None) => Some(cached),
|
||||
(None, Some(transport_auth)) => Some(transport_auth),
|
||||
(None, None) => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_direct_header(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
if !supports_local_generic_oauth_request_auth_resolution(transport) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
if secret.is_empty() || secret == PLACEHOLDER_API_KEY {
|
||||
return None;
|
||||
}
|
||||
|
||||
let auth_config = Self::auth_config_from_transport(transport);
|
||||
let refreshable = auth_config
|
||||
.as_ref()
|
||||
.and_then(refresh_token_from_auth_config)
|
||||
.is_some();
|
||||
if refreshable && auth_config_expires_soon(auth_config.as_ref()) {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: AUTH_HEADER_NAME.to_string(),
|
||||
value: format!("Bearer {secret}"),
|
||||
})
|
||||
}
|
||||
|
||||
fn build_cached_entry(
|
||||
provider_type: &'static str,
|
||||
refreshed: ProviderOAuthTokenSet,
|
||||
) -> CachedOAuthEntry {
|
||||
CachedOAuthEntry {
|
||||
provider_type: provider_type.to_string(),
|
||||
auth_header_name: AUTH_HEADER_NAME.to_string(),
|
||||
auth_header_value: refreshed.token_set.bearer_header_value(),
|
||||
expires_at_unix_secs: refreshed.token_set.expires_at_unix_secs,
|
||||
metadata: Some(refreshed.auth_config),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl LocalOAuthRefreshAdapter for GenericOAuthRefreshAdapter {
|
||||
fn provider_type(&self) -> &'static str {
|
||||
"generic_oauth"
|
||||
}
|
||||
|
||||
fn supports(&self, transport: &GatewayProviderTransportSnapshot) -> bool {
|
||||
supports_local_generic_oauth_request_auth_resolution(transport)
|
||||
}
|
||||
|
||||
fn resolve_cached(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: &CachedOAuthEntry,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
if !entry
|
||||
.provider_type
|
||||
.eq_ignore_ascii_case(transport.provider.provider_type.as_str())
|
||||
{
|
||||
return None;
|
||||
}
|
||||
if expires_at_requires_refresh(entry.expires_at_unix_secs) {
|
||||
return None;
|
||||
}
|
||||
|
||||
let name = entry.auth_header_name.trim();
|
||||
let value = entry.auth_header_value.trim();
|
||||
if name.is_empty() || value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
|
||||
Some(LocalResolvedOAuthRequestAuth::Header {
|
||||
name: name.to_ascii_lowercase(),
|
||||
value: value.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
fn resolve_without_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
) -> Option<LocalResolvedOAuthRequestAuth> {
|
||||
self.resolve_direct_header(transport)
|
||||
}
|
||||
|
||||
fn should_refresh(
|
||||
&self,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> bool {
|
||||
if !supports_local_generic_oauth_request_auth_resolution(transport) {
|
||||
return false;
|
||||
}
|
||||
if entry
|
||||
.and_then(|cached| self.resolve_cached(transport, cached))
|
||||
.is_some()
|
||||
|| self.resolve_direct_header(transport).is_some()
|
||||
{
|
||||
return false;
|
||||
}
|
||||
|
||||
self.base_auth_config(transport, entry)
|
||||
.as_ref()
|
||||
.and_then(refresh_token_from_auth_config)
|
||||
.is_some()
|
||||
}
|
||||
|
||||
async fn refresh(
|
||||
&self,
|
||||
executor: &dyn LocalOAuthHttpExecutor,
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Result<Option<CachedOAuthEntry>, LocalOAuthRefreshError> {
|
||||
let Some(provider_type) = generic_provider_type(transport.provider.provider_type.as_str())
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(auth_config) = self.base_auth_config(transport, entry) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(refresh_token) = refresh_token_from_auth_config(&auth_config) else {
|
||||
tracing::warn!(
|
||||
key_id = %transport.key.id,
|
||||
provider_id = %transport.provider.id,
|
||||
provider_type,
|
||||
"gateway generic oauth refresh skipped because auth_config has no refresh_token"
|
||||
);
|
||||
return Ok(None);
|
||||
};
|
||||
let Some(adapter) = self.adapter_for_provider_type(provider_type) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
tracing::info!(
|
||||
key_id = %transport.key.id,
|
||||
provider_id = %transport.provider.id,
|
||||
endpoint_id = %transport.endpoint.id,
|
||||
provider_type,
|
||||
request_refresh_token_len = refresh_token.len(),
|
||||
"gateway generic oauth refresh delegated to provider oauth adapter"
|
||||
);
|
||||
|
||||
let oauth_executor =
|
||||
ProviderOAuthLocalHttpExecutor::new(provider_type, transport, executor);
|
||||
let ctx = provider_oauth_transport_context_from_snapshot(transport);
|
||||
let account = ProviderOAuthAccount {
|
||||
provider_type: provider_type.to_string(),
|
||||
access_token: current_access_token(transport, entry).unwrap_or_default(),
|
||||
expires_at_unix_secs: auth_config_expires_at(&auth_config),
|
||||
auth_config,
|
||||
identity: BTreeMap::new(),
|
||||
};
|
||||
let refreshed = adapter
|
||||
.refresh(&oauth_executor, &ctx, &account)
|
||||
.await
|
||||
.map_err(|error| oauth_error_to_local_refresh_error(provider_type, error))?;
|
||||
|
||||
tracing::info!(
|
||||
key_id = %transport.key.id,
|
||||
provider_id = %transport.provider.id,
|
||||
endpoint_id = %transport.endpoint.id,
|
||||
provider_type,
|
||||
expires_at_unix_secs = ?refreshed.token_set.expires_at_unix_secs,
|
||||
response_has_refresh_token = refreshed.token_set.refresh_token.is_some(),
|
||||
"gateway generic oauth refresh succeeded"
|
||||
);
|
||||
|
||||
Ok(Some(Self::build_cached_entry(provider_type, refreshed)))
|
||||
}
|
||||
}
|
||||
|
||||
fn generic_provider_type(provider_type: &str) -> Option<&'static str> {
|
||||
let normalized = provider_type.trim();
|
||||
GENERIC_PROVIDER_OAUTH_TEMPLATES
|
||||
.iter()
|
||||
.find(|template| normalized.eq_ignore_ascii_case(template.provider_type))
|
||||
.map(|template| template.provider_type)
|
||||
}
|
||||
|
||||
fn refresh_token_from_auth_config(auth_config: &Value) -> Option<String> {
|
||||
auth_config
|
||||
.as_object()
|
||||
.and_then(|object| object.get("refresh_token"))
|
||||
.and_then(non_empty_string)
|
||||
}
|
||||
|
||||
fn auth_config_expires_at(auth_config: &Value) -> Option<u64> {
|
||||
auth_config
|
||||
.as_object()
|
||||
.and_then(|object| object.get("expires_at"))
|
||||
.and_then(|value| parse_u64_value(Some(value)))
|
||||
}
|
||||
|
||||
fn auth_config_expires_soon(auth_config: Option<&Value>) -> bool {
|
||||
expires_at_requires_refresh(auth_config.and_then(auth_config_expires_at))
|
||||
}
|
||||
|
||||
fn expires_at_requires_refresh(expires_at_unix_secs: Option<u64>) -> bool {
|
||||
expires_at_unix_secs
|
||||
.map(|expires_at_unix_secs| {
|
||||
aether_oauth::core::current_unix_secs()
|
||||
>= expires_at_unix_secs.saturating_sub(OAUTH_REFRESH_SKEW_SECS)
|
||||
})
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn parse_u64_value(value: Option<&Value>) -> Option<u64> {
|
||||
match value? {
|
||||
Value::Number(number) => number.as_u64(),
|
||||
Value::String(string) => string.trim().parse::<u64>().ok(),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn non_empty_string(value: &Value) -> Option<String> {
|
||||
value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
}
|
||||
|
||||
fn current_access_token(
|
||||
transport: &GatewayProviderTransportSnapshot,
|
||||
entry: Option<&CachedOAuthEntry>,
|
||||
) -> Option<String> {
|
||||
entry
|
||||
.and_then(|entry| {
|
||||
entry
|
||||
.auth_header_value
|
||||
.trim()
|
||||
.strip_prefix("Bearer ")
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
.or_else(|| {
|
||||
let secret = transport.key.decrypted_api_key.trim();
|
||||
(!secret.is_empty() && secret != PLACEHOLDER_API_KEY).then(|| secret.to_string())
|
||||
})
|
||||
}
|
||||
Reference in New Issue
Block a user