mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-11 19:59:50 +08:00
refactor: extract provider pool abstractions
This commit is contained in:
@@ -1,7 +1,6 @@
|
||||
use super::quota::antigravity::refresh_antigravity_provider_quota_locally;
|
||||
use super::quota::chatgpt_web::refresh_chatgpt_web_provider_quota_locally;
|
||||
use super::quota::codex::refresh_codex_provider_quota_locally;
|
||||
use super::quota::kiro::refresh_kiro_provider_quota_locally;
|
||||
use super::quota::dispatch::refresh_provider_pool_quota_locally;
|
||||
use super::quota::shared::provider_quota_refresh_endpoint_for_provider;
|
||||
use super::quota::shared::provider_type_supports_quota_refresh;
|
||||
use crate::handlers::admin::provider::write::provider::reconcile_admin_fixed_provider_template_endpoints;
|
||||
use crate::handlers::admin::request::AdminAppState;
|
||||
use crate::provider_key_auth::provider_key_is_oauth_managed;
|
||||
@@ -113,16 +112,20 @@ pub(crate) struct ProviderOAuthRuntimeEndpoints {
|
||||
pub(crate) runtime_endpoint: Option<StoredProviderCatalogEndpoint>,
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_provider_oauth_runtime_endpoints(
|
||||
async fn resolve_provider_runtime_endpoints_with_selector(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
provider_type: &str,
|
||||
endpoint_selector: fn(
|
||||
&str,
|
||||
&[StoredProviderCatalogEndpoint],
|
||||
bool,
|
||||
) -> Option<StoredProviderCatalogEndpoint>,
|
||||
) -> Result<ProviderOAuthRuntimeEndpoints, GatewayError> {
|
||||
let mut endpoints = state
|
||||
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await?;
|
||||
let mut runtime_endpoint =
|
||||
provider_oauth_maintenance_endpoint_for_provider(provider_type, &endpoints);
|
||||
let mut runtime_endpoint = endpoint_selector(provider_type, &endpoints, true);
|
||||
if runtime_endpoint.is_none()
|
||||
&& state
|
||||
.fixed_provider_template(&provider.provider_type)
|
||||
@@ -133,8 +136,7 @@ pub(crate) async fn resolve_provider_oauth_runtime_endpoints(
|
||||
endpoints = state
|
||||
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
|
||||
.await?;
|
||||
runtime_endpoint =
|
||||
provider_oauth_maintenance_endpoint_for_provider(provider_type, &endpoints);
|
||||
runtime_endpoint = endpoint_selector(provider_type, &endpoints, true);
|
||||
}
|
||||
|
||||
Ok(ProviderOAuthRuntimeEndpoints {
|
||||
@@ -143,6 +145,34 @@ pub(crate) async fn resolve_provider_oauth_runtime_endpoints(
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) async fn resolve_provider_oauth_runtime_endpoints(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
provider_type: &str,
|
||||
) -> Result<ProviderOAuthRuntimeEndpoints, GatewayError> {
|
||||
resolve_provider_runtime_endpoints_with_selector(
|
||||
state,
|
||||
provider,
|
||||
provider_type,
|
||||
select_provider_oauth_runtime_endpoint,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn resolve_provider_quota_runtime_endpoints(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
provider_type: &str,
|
||||
) -> Result<ProviderOAuthRuntimeEndpoints, GatewayError> {
|
||||
resolve_provider_runtime_endpoints_with_selector(
|
||||
state,
|
||||
provider,
|
||||
provider_type,
|
||||
provider_quota_refresh_endpoint_for_provider,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) async fn refresh_provider_oauth_account_state_after_update(
|
||||
state: &AdminAppState<'_>,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -150,16 +180,13 @@ pub(crate) async fn refresh_provider_oauth_account_state_after_update(
|
||||
proxy_override: Option<&ProxySnapshot>,
|
||||
) -> Result<(bool, Option<String>), GatewayError> {
|
||||
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
||||
if !matches!(
|
||||
provider_type.as_str(),
|
||||
"codex" | "kiro" | "antigravity" | "chatgpt_web"
|
||||
) {
|
||||
if !provider_type_supports_quota_refresh(&provider_type) {
|
||||
return Ok((false, None));
|
||||
}
|
||||
|
||||
let ProviderOAuthRuntimeEndpoints {
|
||||
runtime_endpoint, ..
|
||||
} = resolve_provider_oauth_runtime_endpoints(state, provider, &provider_type).await?;
|
||||
} = resolve_provider_quota_runtime_endpoints(state, provider, &provider_type).await?;
|
||||
let Some(endpoint) = runtime_endpoint else {
|
||||
return Ok((false, None));
|
||||
};
|
||||
@@ -175,50 +202,15 @@ pub(crate) async fn refresh_provider_oauth_account_state_after_update(
|
||||
return Ok((false, None));
|
||||
}
|
||||
|
||||
let proxy_override = proxy_override.cloned();
|
||||
let payload = match provider_type.as_str() {
|
||||
"codex" => {
|
||||
refresh_codex_provider_quota_locally(
|
||||
state,
|
||||
provider,
|
||||
&endpoint,
|
||||
vec![key],
|
||||
proxy_override.clone(),
|
||||
)
|
||||
.await?
|
||||
}
|
||||
"kiro" => {
|
||||
refresh_kiro_provider_quota_locally(
|
||||
state,
|
||||
provider,
|
||||
&endpoint,
|
||||
vec![key],
|
||||
proxy_override.clone(),
|
||||
)
|
||||
.await?
|
||||
}
|
||||
"antigravity" => {
|
||||
refresh_antigravity_provider_quota_locally(
|
||||
state,
|
||||
provider,
|
||||
&endpoint,
|
||||
vec![key],
|
||||
proxy_override,
|
||||
)
|
||||
.await?
|
||||
}
|
||||
"chatgpt_web" => {
|
||||
refresh_chatgpt_web_provider_quota_locally(
|
||||
state,
|
||||
provider,
|
||||
&endpoint,
|
||||
vec![key],
|
||||
proxy_override,
|
||||
)
|
||||
.await?
|
||||
}
|
||||
_ => None,
|
||||
};
|
||||
let payload = refresh_provider_pool_quota_locally(
|
||||
state,
|
||||
provider,
|
||||
&endpoint,
|
||||
&provider_type,
|
||||
vec![key],
|
||||
proxy_override.cloned(),
|
||||
)
|
||||
.await?;
|
||||
let Some(payload) = payload else {
|
||||
return Ok((false, None));
|
||||
};
|
||||
|
||||
Reference in New Issue
Block a user