refactor: extract provider pool abstractions

This commit is contained in:
fawney19
2026-05-13 18:19:15 +08:00
parent 3c2497f019
commit 5d1460e051
55 changed files with 3469 additions and 2184 deletions
@@ -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));
};