mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
989 lines
31 KiB
Rust
989 lines
31 KiB
Rust
use std::collections::BTreeMap;
|
|
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
|
|
|
use aether_data_contracts::repository::pool_scores::{
|
|
ListPoolMemberProbeCandidatesQuery, PoolMemberHardState, PoolMemberIdentity,
|
|
PoolMemberProbeAttempt, PoolMemberProbeResult, PoolMemberProbeStatus,
|
|
POOL_KIND_PROVIDER_KEY_POOL,
|
|
};
|
|
use aether_data_contracts::repository::provider_catalog::{
|
|
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
|
};
|
|
use aether_runtime_state::{RuntimeLockLease, RuntimeState};
|
|
use futures_util::{stream, StreamExt};
|
|
use serde_json::Value;
|
|
use tracing::{debug, info, warn};
|
|
|
|
use crate::admin_api::{
|
|
admin_provider_pool_config, provider_oauth_runtime_endpoint_for_provider,
|
|
provider_type_supports_quota_refresh, reconcile_admin_fixed_provider_template_endpoints,
|
|
refresh_antigravity_provider_quota_locally, refresh_chatgpt_web_provider_quota_locally,
|
|
refresh_codex_provider_quota_locally, refresh_kiro_provider_quota_locally, AdminAppState,
|
|
};
|
|
use crate::{AppState, GatewayError};
|
|
|
|
use super::pool_score_rebuild::ensure_provider_key_pool_scores_for_keys;
|
|
|
|
const POOL_QUOTA_PROBE_REDIS_PREFIX: &str = "ap:quota_probe:last";
|
|
const POOL_QUOTA_PROBE_DEFAULT_SCAN_INTERVAL_SECONDS: u64 = 60;
|
|
const POOL_QUOTA_PROBE_MIN_SCAN_INTERVAL_SECONDS: u64 = 15;
|
|
const POOL_QUOTA_PROBE_DEFAULT_MAX_KEYS_PER_PROVIDER: usize = 50;
|
|
const POOL_QUOTA_PROBE_DEFAULT_GLOBAL_CONCURRENCY: usize = 16;
|
|
const POOL_QUOTA_PROBE_PROVIDER_LOCK_TTL_MS: u64 = 30_000;
|
|
|
|
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
|
pub(crate) struct PoolQuotaProbeRunSummary {
|
|
pub(crate) providers_checked: usize,
|
|
pub(crate) providers_probed: usize,
|
|
pub(crate) providers_skipped: usize,
|
|
pub(crate) selected_keys: usize,
|
|
pub(crate) succeeded: usize,
|
|
pub(crate) failed: usize,
|
|
pub(crate) auto_removed: usize,
|
|
}
|
|
|
|
impl PoolQuotaProbeRunSummary {
|
|
const fn empty() -> Self {
|
|
Self {
|
|
providers_checked: 0,
|
|
providers_probed: 0,
|
|
providers_skipped: 0,
|
|
selected_keys: 0,
|
|
succeeded: 0,
|
|
failed: 0,
|
|
auto_removed: 0,
|
|
}
|
|
}
|
|
}
|
|
|
|
#[derive(Debug, Clone, Copy)]
|
|
pub(crate) struct PoolQuotaProbeWorkerConfig {
|
|
pub(crate) scan_interval: Duration,
|
|
pub(crate) max_keys_per_provider: usize,
|
|
pub(crate) global_concurrency: usize,
|
|
}
|
|
|
|
impl PoolQuotaProbeWorkerConfig {
|
|
fn from_env() -> Self {
|
|
let scan_interval_seconds = env_u64(
|
|
"POOL_QUOTA_PROBE_SCAN_INTERVAL_SECONDS",
|
|
POOL_QUOTA_PROBE_DEFAULT_SCAN_INTERVAL_SECONDS,
|
|
)
|
|
.max(POOL_QUOTA_PROBE_MIN_SCAN_INTERVAL_SECONDS);
|
|
let max_keys_per_provider = env_usize(
|
|
"POOL_QUOTA_PROBE_MAX_KEYS_PER_PROVIDER",
|
|
POOL_QUOTA_PROBE_DEFAULT_MAX_KEYS_PER_PROVIDER,
|
|
);
|
|
let global_concurrency = env_usize(
|
|
"POOL_QUOTA_PROBE_GLOBAL_CONCURRENCY",
|
|
POOL_QUOTA_PROBE_DEFAULT_GLOBAL_CONCURRENCY,
|
|
)
|
|
.clamp(1, 256);
|
|
Self {
|
|
scan_interval: Duration::from_secs(scan_interval_seconds),
|
|
max_keys_per_provider,
|
|
global_concurrency,
|
|
}
|
|
}
|
|
}
|
|
|
|
fn env_u64(name: &str, default_value: u64) -> u64 {
|
|
std::env::var(name)
|
|
.ok()
|
|
.and_then(|value| value.trim().parse::<u64>().ok())
|
|
.unwrap_or(default_value)
|
|
}
|
|
|
|
fn env_usize(name: &str, default_value: usize) -> usize {
|
|
std::env::var(name)
|
|
.ok()
|
|
.and_then(|value| value.trim().parse::<usize>().ok())
|
|
.unwrap_or(default_value)
|
|
}
|
|
|
|
fn now_unix_secs() -> u64 {
|
|
SystemTime::now()
|
|
.duration_since(UNIX_EPOCH)
|
|
.unwrap_or_default()
|
|
.as_secs()
|
|
}
|
|
|
|
fn provider_supports_quota_probe(provider_type: &str) -> bool {
|
|
provider_type_supports_quota_refresh(provider_type)
|
|
}
|
|
|
|
fn json_number(value: Option<&Value>) -> Option<f64> {
|
|
let value = value?;
|
|
if let Some(number) = value.as_f64() {
|
|
return Some(number);
|
|
}
|
|
value
|
|
.as_str()
|
|
.map(str::trim)
|
|
.filter(|value| !value.is_empty())
|
|
.and_then(|value| value.parse::<f64>().ok())
|
|
}
|
|
|
|
fn extract_quota_updated_at(provider_type: &str, upstream_metadata: Option<&Value>) -> Option<u64> {
|
|
let metadata = upstream_metadata?.as_object()?;
|
|
let bucket_name = match provider_type.trim().to_ascii_lowercase().as_str() {
|
|
"codex" => "codex",
|
|
"kiro" => "kiro",
|
|
"antigravity" => "antigravity",
|
|
"chatgpt_web" => "chatgpt_web",
|
|
_ => return None,
|
|
};
|
|
let bucket = metadata.get(bucket_name)?.as_object()?;
|
|
let mut updated_at = json_number(bucket.get("updated_at"))?;
|
|
if updated_at <= 0.0 {
|
|
return None;
|
|
}
|
|
if updated_at > 1_000_000_000_000.0 {
|
|
updated_at /= 1000.0;
|
|
}
|
|
Some(updated_at as u64)
|
|
}
|
|
|
|
fn parse_probe_stamp(raw_value: Option<&str>) -> Option<u64> {
|
|
let parsed = raw_value
|
|
.map(str::trim)
|
|
.filter(|value| !value.is_empty())
|
|
.and_then(|value| value.parse::<f64>().ok())?;
|
|
if parsed <= 0.0 {
|
|
return None;
|
|
}
|
|
Some(parsed as u64)
|
|
}
|
|
|
|
pub(crate) fn select_pool_quota_probe_key_ids(
|
|
keys: &[StoredProviderCatalogKey],
|
|
provider_type: &str,
|
|
now_ts: u64,
|
|
interval_seconds: u64,
|
|
last_probe_timestamps: &BTreeMap<String, u64>,
|
|
limit: usize,
|
|
) -> Vec<String> {
|
|
let mut stale = Vec::<(u64, String)>::new();
|
|
for key in keys {
|
|
if key.id.trim().is_empty() {
|
|
continue;
|
|
}
|
|
let quota_updated_ts =
|
|
extract_quota_updated_at(provider_type, key.upstream_metadata.as_ref());
|
|
let last_probe_ts = last_probe_timestamps.get(&key.id).copied();
|
|
let anchor_ts = quota_updated_ts
|
|
.unwrap_or(0)
|
|
.max(last_probe_ts.unwrap_or(0));
|
|
if anchor_ts == 0 || now_ts.saturating_sub(anchor_ts) >= interval_seconds {
|
|
stale.push((anchor_ts, key.id.clone()));
|
|
}
|
|
}
|
|
|
|
stale.sort_by(|left, right| left.0.cmp(&right.0).then_with(|| left.1.cmp(&right.1)));
|
|
if limit > 0 && stale.len() > limit {
|
|
stale.truncate(limit);
|
|
}
|
|
stale.into_iter().map(|(_, key_id)| key_id).collect()
|
|
}
|
|
|
|
async fn select_score_probe_key_ids(
|
|
state: &AppState,
|
|
provider_id: &str,
|
|
now_ts: u64,
|
|
interval_seconds: u64,
|
|
limit: usize,
|
|
) -> Vec<String> {
|
|
if limit == 0 {
|
|
return Vec::new();
|
|
}
|
|
let stale_before_unix_secs = now_ts.saturating_sub(interval_seconds);
|
|
let query = ListPoolMemberProbeCandidatesQuery {
|
|
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
|
|
pool_id: provider_id.to_string(),
|
|
capability: None,
|
|
stale_before_unix_secs,
|
|
limit: limit.saturating_mul(4).max(limit),
|
|
};
|
|
let scores = match state.data.list_pool_member_probe_candidates(&query).await {
|
|
Ok(scores) => scores,
|
|
Err(err) => {
|
|
debug!(
|
|
provider_id,
|
|
error = ?err,
|
|
"gateway pool quota probe: failed to read score probe candidates"
|
|
);
|
|
return Vec::new();
|
|
}
|
|
};
|
|
let mut selected = Vec::new();
|
|
let mut seen = std::collections::BTreeSet::new();
|
|
for score in scores {
|
|
if !seen.insert(score.member_id.clone()) {
|
|
continue;
|
|
}
|
|
selected.push(score.member_id);
|
|
if selected.len() >= limit {
|
|
break;
|
|
}
|
|
}
|
|
selected
|
|
}
|
|
|
|
fn probe_stamp_key(provider_id: &str, key_id: &str) -> String {
|
|
format!("{POOL_QUOTA_PROBE_REDIS_PREFIX}:{provider_id}:{key_id}")
|
|
}
|
|
|
|
async fn load_probe_timestamps(
|
|
runtime: &RuntimeState,
|
|
provider_id: &str,
|
|
key_ids: &[String],
|
|
) -> BTreeMap<String, u64> {
|
|
if key_ids.is_empty() {
|
|
return BTreeMap::new();
|
|
}
|
|
|
|
let runtime_keys = key_ids
|
|
.iter()
|
|
.map(|key_id| probe_stamp_key(provider_id, key_id))
|
|
.collect::<Vec<_>>();
|
|
let Ok(values) = runtime.kv_get_many(&runtime_keys).await else {
|
|
debug!("gateway pool quota probe: failed to read runtime probe stamps");
|
|
return BTreeMap::new();
|
|
};
|
|
|
|
key_ids
|
|
.iter()
|
|
.zip(values)
|
|
.filter_map(|(key_id, raw)| {
|
|
parse_probe_stamp(raw.as_deref()).map(|ts| (key_id.clone(), ts))
|
|
})
|
|
.collect()
|
|
}
|
|
|
|
async fn mark_probe_timestamps(
|
|
runtime: &RuntimeState,
|
|
provider_id: &str,
|
|
key_ids: &[String],
|
|
now_ts: u64,
|
|
interval_seconds: u64,
|
|
) {
|
|
if key_ids.is_empty() {
|
|
return;
|
|
}
|
|
|
|
let ttl_seconds = interval_seconds.saturating_mul(2).max(120);
|
|
let value = now_ts.to_string();
|
|
for key_id in key_ids {
|
|
if runtime
|
|
.kv_set(
|
|
&probe_stamp_key(provider_id, key_id),
|
|
value.clone(),
|
|
Some(Duration::from_secs(ttl_seconds)),
|
|
)
|
|
.await
|
|
.is_err()
|
|
{
|
|
debug!("gateway pool quota probe: failed to write runtime probe stamp");
|
|
}
|
|
}
|
|
}
|
|
|
|
async fn acquire_provider_probe_lock(
|
|
runtime: &RuntimeState,
|
|
provider_id: &str,
|
|
) -> Option<RuntimeLockLease> {
|
|
let owner = format!("aether-gateway-pool-probe-{}", std::process::id());
|
|
match runtime
|
|
.lock_try_acquire(
|
|
&format!("pool_quota_probe:{provider_id}"),
|
|
&owner,
|
|
Duration::from_millis(POOL_QUOTA_PROBE_PROVIDER_LOCK_TTL_MS),
|
|
)
|
|
.await
|
|
{
|
|
Ok(lease) => lease,
|
|
Err(err) => {
|
|
debug!(
|
|
provider_id,
|
|
error = %err,
|
|
"gateway pool quota probe: failed to acquire runtime provider lock"
|
|
);
|
|
None
|
|
}
|
|
}
|
|
}
|
|
|
|
async fn release_provider_probe_lock(runtime: &RuntimeState, lease: Option<RuntimeLockLease>) {
|
|
let Some(lease) = lease else {
|
|
return;
|
|
};
|
|
if let Err(err) = runtime.lock_release(&lease).await {
|
|
debug!(
|
|
error = %err,
|
|
"gateway pool quota probe: failed to release runtime provider lock"
|
|
);
|
|
}
|
|
}
|
|
|
|
async fn select_keys_for_provider(
|
|
state: &AppState,
|
|
runtime: &RuntimeState,
|
|
provider: &StoredProviderCatalogProvider,
|
|
provider_type: &str,
|
|
interval_seconds: u64,
|
|
max_keys_per_provider: usize,
|
|
now_ts: u64,
|
|
) -> Result<Vec<StoredProviderCatalogKey>, GatewayError> {
|
|
let lease = acquire_provider_probe_lock(runtime, &provider.id).await;
|
|
if lease.is_none() {
|
|
return Ok(Vec::new());
|
|
}
|
|
|
|
let result = async {
|
|
let keys = state
|
|
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
|
|
.await?
|
|
.into_iter()
|
|
.filter(|key| key.is_active)
|
|
.collect::<Vec<_>>();
|
|
if keys.is_empty() {
|
|
return Ok(Vec::new());
|
|
}
|
|
|
|
let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
|
|
let probe_stamps = load_probe_timestamps(runtime, &provider.id, &key_ids).await;
|
|
let mut selected_ids = select_score_probe_key_ids(
|
|
state,
|
|
&provider.id,
|
|
now_ts,
|
|
interval_seconds,
|
|
max_keys_per_provider,
|
|
)
|
|
.await
|
|
.into_iter()
|
|
.filter(|key_id| key_ids.iter().any(|known_id| known_id == key_id))
|
|
.filter(|key_id| {
|
|
probe_stamps.get(key_id).is_none_or(|last_probe_ts| {
|
|
now_ts.saturating_sub(*last_probe_ts) >= interval_seconds
|
|
})
|
|
})
|
|
.collect::<Vec<_>>();
|
|
let mut selected_seen = selected_ids
|
|
.iter()
|
|
.cloned()
|
|
.collect::<std::collections::BTreeSet<_>>();
|
|
let remaining = max_keys_per_provider.saturating_sub(selected_ids.len());
|
|
let fallback_selected_ids = select_pool_quota_probe_key_ids(
|
|
&keys,
|
|
provider_type,
|
|
now_ts,
|
|
interval_seconds,
|
|
&probe_stamps,
|
|
max_keys_per_provider,
|
|
);
|
|
for key_id in fallback_selected_ids {
|
|
if selected_ids.len() >= max_keys_per_provider || remaining == 0 {
|
|
break;
|
|
}
|
|
if selected_seen.insert(key_id.clone()) {
|
|
selected_ids.push(key_id);
|
|
}
|
|
}
|
|
if selected_ids.is_empty() {
|
|
return Ok(Vec::new());
|
|
}
|
|
|
|
mark_probe_timestamps(
|
|
runtime,
|
|
&provider.id,
|
|
&selected_ids,
|
|
now_ts,
|
|
interval_seconds,
|
|
)
|
|
.await;
|
|
|
|
let mut keys_by_id = keys
|
|
.into_iter()
|
|
.map(|key| (key.id.clone(), key))
|
|
.collect::<BTreeMap<_, _>>();
|
|
Ok(selected_ids
|
|
.into_iter()
|
|
.filter_map(|key_id| keys_by_id.remove(&key_id))
|
|
.collect::<Vec<_>>())
|
|
}
|
|
.await;
|
|
|
|
release_provider_probe_lock(runtime, lease).await;
|
|
result
|
|
}
|
|
|
|
fn endpoint_for_probe(
|
|
provider_type: &str,
|
|
endpoints: &[StoredProviderCatalogEndpoint],
|
|
) -> Option<StoredProviderCatalogEndpoint> {
|
|
provider_oauth_runtime_endpoint_for_provider(provider_type, endpoints)
|
|
}
|
|
|
|
async fn endpoint_for_probe_with_reconcile(
|
|
state: &AppState,
|
|
admin_state: &AdminAppState<'_>,
|
|
provider: &StoredProviderCatalogProvider,
|
|
provider_type: &str,
|
|
endpoints_by_provider: &mut BTreeMap<String, Vec<StoredProviderCatalogEndpoint>>,
|
|
) -> Result<Option<StoredProviderCatalogEndpoint>, GatewayError> {
|
|
let endpoints = endpoints_by_provider
|
|
.get(&provider.id)
|
|
.map(Vec::as_slice)
|
|
.unwrap_or(&[]);
|
|
if let Some(endpoint) = endpoint_for_probe(provider_type, endpoints) {
|
|
return Ok(Some(endpoint));
|
|
}
|
|
|
|
if admin_state
|
|
.fixed_provider_template(&provider.provider_type)
|
|
.is_none()
|
|
{
|
|
return Ok(None);
|
|
}
|
|
|
|
reconcile_admin_fixed_provider_template_endpoints(admin_state, provider).await?;
|
|
let refreshed = state
|
|
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
|
|
.await?;
|
|
let endpoint = endpoint_for_probe(provider_type, &refreshed);
|
|
endpoints_by_provider.insert(provider.id.clone(), refreshed);
|
|
Ok(endpoint)
|
|
}
|
|
|
|
async fn refresh_provider_probe_keys(
|
|
admin_state: &AdminAppState<'_>,
|
|
provider: &StoredProviderCatalogProvider,
|
|
endpoint: &StoredProviderCatalogEndpoint,
|
|
provider_type: &str,
|
|
keys: Vec<StoredProviderCatalogKey>,
|
|
) -> Result<Option<Value>, GatewayError> {
|
|
match provider_type {
|
|
"codex" => {
|
|
refresh_codex_provider_quota_locally(admin_state, provider, endpoint, keys, None).await
|
|
}
|
|
"kiro" => {
|
|
refresh_kiro_provider_quota_locally(admin_state, provider, endpoint, keys, None).await
|
|
}
|
|
"antigravity" => {
|
|
refresh_antigravity_provider_quota_locally(admin_state, provider, endpoint, keys, None)
|
|
.await
|
|
}
|
|
"chatgpt_web" => {
|
|
refresh_chatgpt_web_provider_quota_locally(admin_state, provider, endpoint, keys, None)
|
|
.await
|
|
}
|
|
_ => Ok(None),
|
|
}
|
|
}
|
|
|
|
fn update_summary_from_payload(
|
|
summary: &mut PoolQuotaProbeRunSummary,
|
|
selected_count: usize,
|
|
payload: Option<&Value>,
|
|
) {
|
|
let Some(payload) = payload else {
|
|
summary.failed += selected_count;
|
|
return;
|
|
};
|
|
summary.succeeded += payload.get("success").and_then(Value::as_u64).unwrap_or(0) as usize;
|
|
summary.failed += payload.get("failed").and_then(Value::as_u64).unwrap_or(0) as usize;
|
|
summary.auto_removed += payload
|
|
.get("auto_removed")
|
|
.and_then(Value::as_u64)
|
|
.unwrap_or(0) as usize;
|
|
}
|
|
|
|
async fn record_score_probe_results_from_payload(
|
|
state: &AppState,
|
|
provider_id: &str,
|
|
selected_key_ids: &[String],
|
|
payload: Option<&Value>,
|
|
attempted_at: u64,
|
|
) {
|
|
let mut recorded = std::collections::BTreeSet::new();
|
|
if let Some(results) = payload
|
|
.and_then(|value| value.get("results"))
|
|
.and_then(Value::as_array)
|
|
{
|
|
for item in results {
|
|
let Some(key_id) = item
|
|
.get("key_id")
|
|
.and_then(Value::as_str)
|
|
.map(str::trim)
|
|
.filter(|value| !value.is_empty())
|
|
else {
|
|
continue;
|
|
};
|
|
recorded.insert(key_id.to_string());
|
|
record_score_probe_result_for_key(
|
|
state,
|
|
provider_id,
|
|
key_id,
|
|
attempted_at,
|
|
probe_result_succeeded(item),
|
|
probe_result_hard_state(item).or_else(|| {
|
|
(!probe_result_succeeded(item)).then_some(PoolMemberHardState::Cooldown)
|
|
}),
|
|
serde_json::json!({
|
|
"last_probe": {
|
|
"source": "pool_quota_probe",
|
|
"status": item.get("status").cloned().unwrap_or(Value::Null),
|
|
"status_code": item.get("status_code").cloned().unwrap_or(Value::Null),
|
|
"message": item.get("message").cloned().unwrap_or(Value::Null),
|
|
"auto_removed": item.get("auto_removed").cloned().unwrap_or(Value::Null)
|
|
}
|
|
}),
|
|
)
|
|
.await;
|
|
}
|
|
}
|
|
|
|
for key_id in selected_key_ids {
|
|
if recorded.contains(key_id) {
|
|
continue;
|
|
}
|
|
record_score_probe_result_for_key(
|
|
state,
|
|
provider_id,
|
|
key_id,
|
|
attempted_at,
|
|
false,
|
|
Some(PoolMemberHardState::Cooldown),
|
|
serde_json::json!({
|
|
"last_probe": {
|
|
"source": "pool_quota_probe",
|
|
"status": "missing_result"
|
|
}
|
|
}),
|
|
)
|
|
.await;
|
|
}
|
|
}
|
|
|
|
async fn record_score_probe_result_for_key(
|
|
state: &AppState,
|
|
provider_id: &str,
|
|
key_id: &str,
|
|
attempted_at: u64,
|
|
succeeded: bool,
|
|
hard_state: Option<PoolMemberHardState>,
|
|
score_reason_patch: Value,
|
|
) {
|
|
let result = PoolMemberProbeResult {
|
|
identity: PoolMemberIdentity::provider_api_key(provider_id.to_string(), key_id.to_string()),
|
|
scope: None,
|
|
attempted_at,
|
|
succeeded,
|
|
hard_state,
|
|
probe_status: if succeeded {
|
|
PoolMemberProbeStatus::Ok
|
|
} else {
|
|
PoolMemberProbeStatus::Failed
|
|
},
|
|
score_reason_patch: Some(score_reason_patch),
|
|
};
|
|
if let Err(err) = state.data.record_pool_member_probe_result(result).await {
|
|
debug!(
|
|
provider_id,
|
|
key_id,
|
|
error = ?err,
|
|
"gateway pool quota probe: failed to record score probe result"
|
|
);
|
|
}
|
|
}
|
|
|
|
async fn record_score_probe_in_progress_for_key(
|
|
state: &AppState,
|
|
provider_id: &str,
|
|
key_id: &str,
|
|
attempted_at: u64,
|
|
) {
|
|
let attempt = PoolMemberProbeAttempt {
|
|
identity: PoolMemberIdentity::provider_api_key(provider_id.to_string(), key_id.to_string()),
|
|
scope: None,
|
|
attempted_at,
|
|
score_reason_patch: Some(serde_json::json!({
|
|
"last_probe": {
|
|
"source": "pool_quota_probe",
|
|
"status": "in_progress"
|
|
}
|
|
})),
|
|
};
|
|
if let Err(err) = state.data.mark_pool_member_probe_in_progress(attempt).await {
|
|
debug!(
|
|
provider_id,
|
|
key_id,
|
|
error = ?err,
|
|
"gateway pool quota probe: failed to mark score probe in progress"
|
|
);
|
|
}
|
|
}
|
|
|
|
fn probe_result_succeeded(item: &Value) -> bool {
|
|
item.get("status")
|
|
.and_then(Value::as_str)
|
|
.is_some_and(|status| status == "success")
|
|
}
|
|
|
|
fn probe_result_hard_state(item: &Value) -> Option<PoolMemberHardState> {
|
|
if probe_result_succeeded(item) {
|
|
return Some(PoolMemberHardState::Available);
|
|
}
|
|
if item
|
|
.get("auto_removed")
|
|
.and_then(Value::as_bool)
|
|
.unwrap_or(false)
|
|
{
|
|
return Some(PoolMemberHardState::Banned);
|
|
}
|
|
let status = item
|
|
.get("status")
|
|
.and_then(Value::as_str)
|
|
.unwrap_or_default()
|
|
.trim()
|
|
.to_ascii_lowercase();
|
|
match status.as_str() {
|
|
"auth_invalid" | "forbidden" => Some(PoolMemberHardState::AuthInvalid),
|
|
"workspace_deactivated" => Some(PoolMemberHardState::Banned),
|
|
"quota_exhausted" => Some(PoolMemberHardState::QuotaExhausted),
|
|
_ => match item.get("status_code").and_then(Value::as_u64) {
|
|
Some(401 | 403) => Some(PoolMemberHardState::AuthInvalid),
|
|
Some(402) => Some(PoolMemberHardState::QuotaExhausted),
|
|
Some(429 | 500..=599) => Some(PoolMemberHardState::Cooldown),
|
|
_ => None,
|
|
},
|
|
}
|
|
}
|
|
|
|
pub(crate) async fn perform_pool_quota_probe_once_with_config(
|
|
state: &AppState,
|
|
config: PoolQuotaProbeWorkerConfig,
|
|
) -> Result<PoolQuotaProbeRunSummary, GatewayError> {
|
|
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
|
|
return Ok(PoolQuotaProbeRunSummary::empty());
|
|
}
|
|
|
|
let providers = state
|
|
.list_provider_catalog_providers(true)
|
|
.await?
|
|
.into_iter()
|
|
.filter_map(|provider| {
|
|
let provider_type = provider.provider_type.trim().to_ascii_lowercase();
|
|
if provider_supports_quota_probe(&provider_type) {
|
|
Some((provider, provider_type))
|
|
} else {
|
|
None
|
|
}
|
|
})
|
|
.filter_map(|(provider, provider_type)| {
|
|
let pool_config = admin_provider_pool_config(&provider)?;
|
|
if pool_config.probing_enabled {
|
|
Some((provider, provider_type, pool_config))
|
|
} else {
|
|
None
|
|
}
|
|
})
|
|
.collect::<Vec<_>>();
|
|
|
|
if providers.is_empty() {
|
|
return Ok(PoolQuotaProbeRunSummary::empty());
|
|
}
|
|
|
|
let provider_ids = providers
|
|
.iter()
|
|
.map(|(provider, _, _)| provider.id.clone())
|
|
.collect::<Vec<_>>();
|
|
let mut endpoints_by_provider = BTreeMap::<String, Vec<StoredProviderCatalogEndpoint>>::new();
|
|
for endpoint in state
|
|
.list_provider_catalog_endpoints_by_provider_ids(&provider_ids)
|
|
.await?
|
|
{
|
|
endpoints_by_provider
|
|
.entry(endpoint.provider_id.clone())
|
|
.or_default()
|
|
.push(endpoint);
|
|
}
|
|
|
|
let admin_state = AdminAppState::new(state);
|
|
let now_ts = now_unix_secs();
|
|
let mut summary = PoolQuotaProbeRunSummary {
|
|
providers_checked: providers.len(),
|
|
..PoolQuotaProbeRunSummary::empty()
|
|
};
|
|
|
|
for (provider, provider_type, pool_config) in providers {
|
|
let Some(endpoint) = endpoint_for_probe_with_reconcile(
|
|
state,
|
|
&admin_state,
|
|
&provider,
|
|
&provider_type,
|
|
&mut endpoints_by_provider,
|
|
)
|
|
.await?
|
|
else {
|
|
summary.providers_skipped += 1;
|
|
debug!(
|
|
provider_id = %provider.id,
|
|
provider_type,
|
|
"gateway pool quota probe skipped provider without active quota endpoint"
|
|
);
|
|
continue;
|
|
};
|
|
let endpoints = endpoints_by_provider
|
|
.remove(&provider.id)
|
|
.unwrap_or_else(|| vec![endpoint.clone()]);
|
|
|
|
let interval_minutes = pool_config.probing_interval_minutes;
|
|
let interval_seconds = interval_minutes.clamp(1, 1440).saturating_mul(60);
|
|
let keys = select_keys_for_provider(
|
|
state,
|
|
state.runtime_state.as_ref(),
|
|
&provider,
|
|
&provider_type,
|
|
interval_seconds,
|
|
config.max_keys_per_provider,
|
|
now_ts,
|
|
)
|
|
.await?;
|
|
if keys.is_empty() {
|
|
continue;
|
|
}
|
|
|
|
let selected_count = keys.len();
|
|
summary.providers_probed += 1;
|
|
summary.selected_keys += selected_count;
|
|
|
|
let selected_key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
|
|
let score_ensure_budget = (pool_config.score_fallback_scan_limit as usize)
|
|
.min(50_000)
|
|
.max(selected_count.min(50_000));
|
|
match ensure_provider_key_pool_scores_for_keys(
|
|
state,
|
|
&provider,
|
|
&pool_config,
|
|
&endpoints,
|
|
&keys,
|
|
now_ts,
|
|
score_ensure_budget,
|
|
)
|
|
.await
|
|
{
|
|
Ok(upserted) if upserted > 0 => {
|
|
debug!(
|
|
provider_id = %provider.id,
|
|
key_count = selected_count,
|
|
scores_upserted = upserted,
|
|
"gateway pool quota probe: ensured score rows for selected probe keys"
|
|
);
|
|
}
|
|
Ok(_) => {}
|
|
Err(err) => {
|
|
warn!(
|
|
provider_id = %provider.id,
|
|
key_count = selected_count,
|
|
error = ?err,
|
|
"gateway pool quota probe: failed to ensure score rows for selected probe keys"
|
|
);
|
|
}
|
|
}
|
|
for key_id in &selected_key_ids {
|
|
record_score_probe_in_progress_for_key(state, &provider.id, key_id, now_ts).await;
|
|
}
|
|
|
|
let provider_short_id = provider.id.chars().take(8).collect::<String>();
|
|
let probe_concurrency = pool_config.probe_concurrency.clamp(1, 64) as usize;
|
|
let probe_concurrency = probe_concurrency.min(config.global_concurrency).max(1);
|
|
let probe_results = stream::iter(keys.into_iter().map(|key| {
|
|
let key_id = key.id.clone();
|
|
let admin_state = &admin_state;
|
|
let provider = &provider;
|
|
let endpoint = &endpoint;
|
|
let provider_type = provider_type.as_str();
|
|
async move {
|
|
let result = refresh_provider_probe_keys(
|
|
admin_state,
|
|
provider,
|
|
endpoint,
|
|
provider_type,
|
|
vec![key],
|
|
)
|
|
.await;
|
|
(key_id, result)
|
|
}
|
|
}))
|
|
.buffer_unordered(probe_concurrency)
|
|
.collect::<Vec<_>>()
|
|
.await;
|
|
|
|
let mut probe_success = 0usize;
|
|
let mut probe_failed = 0usize;
|
|
for (key_id, result) in probe_results {
|
|
match result {
|
|
Ok(payload) => {
|
|
update_summary_from_payload(&mut summary, 1, payload.as_ref());
|
|
probe_success += payload
|
|
.as_ref()
|
|
.and_then(|value| value.get("success"))
|
|
.and_then(Value::as_u64)
|
|
.unwrap_or(0) as usize;
|
|
probe_failed += payload
|
|
.as_ref()
|
|
.and_then(|value| value.get("failed"))
|
|
.and_then(Value::as_u64)
|
|
.unwrap_or(0) as usize;
|
|
record_score_probe_results_from_payload(
|
|
state,
|
|
&provider.id,
|
|
std::slice::from_ref(&key_id),
|
|
payload.as_ref(),
|
|
now_ts,
|
|
)
|
|
.await;
|
|
}
|
|
Err(err) => {
|
|
summary.failed += 1;
|
|
probe_failed += 1;
|
|
record_score_probe_result_for_key(
|
|
state,
|
|
&provider.id,
|
|
&key_id,
|
|
now_ts,
|
|
false,
|
|
Some(PoolMemberHardState::Cooldown),
|
|
serde_json::json!({
|
|
"last_probe": {
|
|
"source": "pool_quota_probe",
|
|
"status": "worker_error",
|
|
"message": format!("{err:?}")
|
|
}
|
|
}),
|
|
)
|
|
.await;
|
|
warn!(
|
|
provider_id = %provider_short_id,
|
|
provider_type,
|
|
key_id,
|
|
error = ?err,
|
|
"gateway pool quota probe failed"
|
|
);
|
|
}
|
|
}
|
|
}
|
|
info!(
|
|
provider_id = %provider_short_id,
|
|
provider_type,
|
|
selected = selected_count,
|
|
success = probe_success,
|
|
failed = probe_failed,
|
|
concurrency = probe_concurrency,
|
|
"gateway pool quota probe completed"
|
|
);
|
|
}
|
|
|
|
Ok(summary)
|
|
}
|
|
|
|
pub(crate) async fn perform_pool_quota_probe_once(
|
|
state: &AppState,
|
|
) -> Result<PoolQuotaProbeRunSummary, GatewayError> {
|
|
perform_pool_quota_probe_once_with_config(state, PoolQuotaProbeWorkerConfig::from_env()).await
|
|
}
|
|
|
|
pub(crate) fn spawn_pool_quota_probe_worker(
|
|
state: AppState,
|
|
) -> Option<tokio::task::JoinHandle<()>> {
|
|
if !state.has_provider_catalog_data_reader() || !state.has_provider_catalog_data_writer() {
|
|
return None;
|
|
}
|
|
|
|
let config = PoolQuotaProbeWorkerConfig::from_env();
|
|
Some(tokio::spawn(async move {
|
|
let mut interval = tokio::time::interval(config.scan_interval);
|
|
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
|
|
interval.tick().await;
|
|
loop {
|
|
interval.tick().await;
|
|
if let Err(err) = perform_pool_quota_probe_once_with_config(&state, config).await {
|
|
warn!(
|
|
error = ?err,
|
|
"gateway pool quota probe worker tick failed"
|
|
);
|
|
}
|
|
}
|
|
}))
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use serde_json::json;
|
|
|
|
fn key(
|
|
id: &str,
|
|
provider_id: &str,
|
|
upstream_metadata: Option<Value>,
|
|
) -> StoredProviderCatalogKey {
|
|
let mut key = StoredProviderCatalogKey::new(
|
|
id.to_string(),
|
|
provider_id.to_string(),
|
|
id.to_string(),
|
|
"oauth".to_string(),
|
|
None,
|
|
true,
|
|
)
|
|
.expect("key should build");
|
|
key.upstream_metadata = upstream_metadata;
|
|
key
|
|
}
|
|
|
|
#[test]
|
|
fn selects_stale_probe_keys_by_oldest_anchor() {
|
|
let keys = vec![
|
|
key(
|
|
"fresh",
|
|
"provider-1",
|
|
Some(json!({ "codex": { "updated_at": 1_990 } })),
|
|
),
|
|
key(
|
|
"old",
|
|
"provider-1",
|
|
Some(json!({ "codex": { "updated_at": 1_000 } })),
|
|
),
|
|
key("never", "provider-1", None),
|
|
key(
|
|
"stamped",
|
|
"provider-1",
|
|
Some(json!({ "codex": { "updated_at": 900 } })),
|
|
),
|
|
];
|
|
let stamps = BTreeMap::from([("stamped".to_string(), 1_950)]);
|
|
|
|
let selected = select_pool_quota_probe_key_ids(&keys, "codex", 2_000, 600, &stamps, 2);
|
|
|
|
assert_eq!(selected, vec!["never".to_string(), "old".to_string()]);
|
|
}
|
|
|
|
#[test]
|
|
fn parses_quota_updated_at_seconds_and_milliseconds() {
|
|
assert_eq!(
|
|
extract_quota_updated_at(
|
|
"codex",
|
|
Some(&json!({ "codex": { "updated_at": 1_700_000_000 } }))
|
|
),
|
|
Some(1_700_000_000)
|
|
);
|
|
assert_eq!(
|
|
extract_quota_updated_at(
|
|
"kiro",
|
|
Some(&json!({ "kiro": { "updated_at": 1_700_000_000_000_u64 } }))
|
|
),
|
|
Some(1_700_000_000)
|
|
);
|
|
}
|
|
}
|