mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-03 01:40:21 +08:00
Implement generic pool member scoring and probing
This commit is contained in:
@@ -1,10 +1,16 @@
|
||||
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};
|
||||
|
||||
@@ -19,6 +25,7 @@ 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)]
|
||||
@@ -50,6 +57,7 @@ impl PoolQuotaProbeRunSummary {
|
||||
pub(crate) struct PoolQuotaProbeWorkerConfig {
|
||||
pub(crate) scan_interval: Duration,
|
||||
pub(crate) max_keys_per_provider: usize,
|
||||
pub(crate) global_concurrency: usize,
|
||||
}
|
||||
|
||||
impl PoolQuotaProbeWorkerConfig {
|
||||
@@ -63,9 +71,15 @@ impl PoolQuotaProbeWorkerConfig {
|
||||
"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,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -172,6 +186,49 @@ pub(crate) fn select_pool_quota_probe_key_ids(
|
||||
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}")
|
||||
}
|
||||
@@ -295,7 +352,28 @@ async fn select_keys_for_provider(
|
||||
|
||||
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 selected_ids = select_pool_quota_probe_key_ids(
|
||||
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,
|
||||
@@ -303,6 +381,14 @@ async fn select_keys_for_provider(
|
||||
&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());
|
||||
}
|
||||
@@ -381,6 +467,166 @@ fn update_summary_from_payload(
|
||||
.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),
|
||||
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,
|
||||
None,
|
||||
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,
|
||||
@@ -408,6 +654,7 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
|
||||
provider,
|
||||
provider_type,
|
||||
pool_config.probing_interval_minutes,
|
||||
pool_config.probe_concurrency.clamp(1, 64) as usize,
|
||||
))
|
||||
} else {
|
||||
None
|
||||
@@ -421,7 +668,7 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
|
||||
|
||||
let provider_ids = providers
|
||||
.iter()
|
||||
.map(|(provider, _, _)| provider.id.clone())
|
||||
.map(|(provider, _, _, _)| provider.id.clone())
|
||||
.collect::<Vec<_>>();
|
||||
let mut endpoints_by_provider = BTreeMap::<String, Vec<StoredProviderCatalogEndpoint>>::new();
|
||||
for endpoint in state
|
||||
@@ -441,7 +688,7 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
|
||||
..PoolQuotaProbeRunSummary::empty()
|
||||
};
|
||||
|
||||
for (provider, provider_type, interval_minutes) in providers {
|
||||
for (provider, provider_type, interval_minutes, probe_concurrency) in providers {
|
||||
let endpoints = endpoints_by_provider
|
||||
.remove(&provider.id)
|
||||
.unwrap_or_default();
|
||||
@@ -474,42 +721,98 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
|
||||
summary.providers_probed += 1;
|
||||
summary.selected_keys += selected_count;
|
||||
|
||||
let selected_key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
|
||||
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>();
|
||||
match refresh_provider_probe_keys(&admin_state, &provider, &endpoint, &provider_type, keys)
|
||||
.await
|
||||
{
|
||||
Ok(payload) => {
|
||||
update_summary_from_payload(&mut summary, selected_count, payload.as_ref());
|
||||
let probe_success = payload
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("success"))
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or(0);
|
||||
let probe_failed = payload
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("failed"))
|
||||
.and_then(Value::as_u64)
|
||||
.unwrap_or(0);
|
||||
info!(
|
||||
provider_id = %provider_short_id,
|
||||
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,
|
||||
selected = selected_count,
|
||||
success = probe_success,
|
||||
failed = probe_failed,
|
||||
"gateway pool quota probe completed"
|
||||
);
|
||||
vec![key],
|
||||
)
|
||||
.await;
|
||||
(key_id, result)
|
||||
}
|
||||
Err(err) => {
|
||||
summary.failed += selected_count;
|
||||
warn!(
|
||||
provider_id = %provider_short_id,
|
||||
provider_type,
|
||||
selected = selected_count,
|
||||
error = ?err,
|
||||
"gateway pool quota probe failed"
|
||||
);
|
||||
}))
|
||||
.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)
|
||||
|
||||
@@ -0,0 +1,405 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
|
||||
use aether_data_contracts::repository::global_models::AdminProviderModelListQuery;
|
||||
use aether_data_contracts::repository::pool_scores::GetPoolMemberScoresByIdsQuery;
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
use crate::admin_api::admin_provider_pool_config;
|
||||
use crate::ai_serving::build_provider_key_pool_score_upsert;
|
||||
use crate::{AppState, GatewayError};
|
||||
|
||||
const POOL_SCORE_REBUILD_DEFAULT_INTERVAL_SECONDS: u64 = 300;
|
||||
const POOL_SCORE_REBUILD_MIN_INTERVAL_SECONDS: u64 = 30;
|
||||
const POOL_SCORE_REBUILD_DEFAULT_MAX_UPSERTS_PER_TICK: usize = 20_000;
|
||||
const POOL_SCORE_REBUILD_PROVIDER_CURSOR_KEY: &str = "ap:pool_score_rebuild:provider_cursor";
|
||||
const POOL_SCORE_REBUILD_PROVIDER_OFFSET_PREFIX: &str = "ap:pool_score_rebuild:provider_offset";
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub(crate) struct PoolScoreRebuildRunSummary {
|
||||
pub(crate) providers_checked: usize,
|
||||
pub(crate) providers_scored: usize,
|
||||
pub(crate) keys_seen: usize,
|
||||
pub(crate) scores_upserted: usize,
|
||||
}
|
||||
|
||||
impl PoolScoreRebuildRunSummary {
|
||||
const fn empty() -> Self {
|
||||
Self {
|
||||
providers_checked: 0,
|
||||
providers_scored: 0,
|
||||
keys_seen: 0,
|
||||
scores_upserted: 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
pub(crate) struct PoolScoreRebuildWorkerConfig {
|
||||
pub(crate) interval: Duration,
|
||||
pub(crate) max_upserts_per_tick: usize,
|
||||
}
|
||||
|
||||
impl PoolScoreRebuildWorkerConfig {
|
||||
fn from_env() -> Self {
|
||||
let interval_seconds = env_u64(
|
||||
"POOL_SCORE_REBUILD_INTERVAL_SECONDS",
|
||||
POOL_SCORE_REBUILD_DEFAULT_INTERVAL_SECONDS,
|
||||
)
|
||||
.max(POOL_SCORE_REBUILD_MIN_INTERVAL_SECONDS);
|
||||
let max_upserts_per_tick = env_usize(
|
||||
"POOL_SCORE_REBUILD_MAX_UPSERTS_PER_TICK",
|
||||
POOL_SCORE_REBUILD_DEFAULT_MAX_UPSERTS_PER_TICK,
|
||||
)
|
||||
.max(1);
|
||||
Self {
|
||||
interval: Duration::from_secs(interval_seconds),
|
||||
max_upserts_per_tick,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
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()
|
||||
}
|
||||
|
||||
async fn load_runtime_usize(state: &AppState, key: &str) -> usize {
|
||||
state
|
||||
.runtime_state
|
||||
.kv_get(key)
|
||||
.await
|
||||
.ok()
|
||||
.flatten()
|
||||
.and_then(|value| value.trim().parse::<usize>().ok())
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
async fn store_runtime_usize(state: &AppState, key: &str, value: usize) {
|
||||
if let Err(err) = state
|
||||
.runtime_state
|
||||
.kv_set(key, value.to_string(), None)
|
||||
.await
|
||||
{
|
||||
debug!(
|
||||
key,
|
||||
error = ?err,
|
||||
"gateway pool score rebuild: failed to store cursor"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_offset_cursor_key(provider_id: &str) -> String {
|
||||
format!("{POOL_SCORE_REBUILD_PROVIDER_OFFSET_PREFIX}:{provider_id}")
|
||||
}
|
||||
|
||||
fn score_combo_indices(
|
||||
flat_index: usize,
|
||||
model_count: usize,
|
||||
key_count: usize,
|
||||
) -> (usize, usize, usize) {
|
||||
let endpoint_stride = model_count.saturating_mul(key_count).max(1);
|
||||
let endpoint_index = flat_index / endpoint_stride;
|
||||
let remainder = flat_index % endpoint_stride;
|
||||
let model_index = remainder / key_count.max(1);
|
||||
let key_index = remainder % key_count.max(1);
|
||||
(endpoint_index, model_index, key_index)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct ProviderScoreBuildItem {
|
||||
endpoint_index: usize,
|
||||
model_index: usize,
|
||||
key_index: usize,
|
||||
score_id: String,
|
||||
}
|
||||
|
||||
pub(crate) async fn perform_pool_score_rebuild_once_with_config(
|
||||
state: &AppState,
|
||||
config: PoolScoreRebuildWorkerConfig,
|
||||
) -> Result<PoolScoreRebuildRunSummary, GatewayError> {
|
||||
if !state.has_provider_catalog_data_reader()
|
||||
|| !state.data.has_pool_score_reader()
|
||||
|| !state.data.has_pool_score_writer()
|
||||
{
|
||||
return Ok(PoolScoreRebuildRunSummary::empty());
|
||||
}
|
||||
|
||||
let mut providers = state
|
||||
.list_provider_catalog_providers(true)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|provider| admin_provider_pool_config(provider).is_some())
|
||||
.collect::<Vec<_>>();
|
||||
providers.sort_by(|left, right| left.id.cmp(&right.id));
|
||||
if providers.is_empty() {
|
||||
return Ok(PoolScoreRebuildRunSummary::empty());
|
||||
}
|
||||
|
||||
let provider_ids = providers
|
||||
.iter()
|
||||
.map(|provider| provider.id.clone())
|
||||
.collect::<Vec<_>>();
|
||||
let mut endpoints_by_provider = BTreeMap::new();
|
||||
for endpoint in state
|
||||
.list_provider_catalog_endpoints_by_provider_ids(&provider_ids)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|endpoint| endpoint.is_active)
|
||||
{
|
||||
endpoints_by_provider
|
||||
.entry(endpoint.provider_id.clone())
|
||||
.or_insert_with(Vec::new)
|
||||
.push(endpoint);
|
||||
}
|
||||
|
||||
let mut keys_by_provider = BTreeMap::new();
|
||||
for key in state
|
||||
.list_provider_catalog_keys_by_provider_ids(&provider_ids)
|
||||
.await?
|
||||
{
|
||||
keys_by_provider
|
||||
.entry(key.provider_id.clone())
|
||||
.or_insert_with(Vec::new)
|
||||
.push(key);
|
||||
}
|
||||
|
||||
let now = now_unix_secs();
|
||||
let mut summary = PoolScoreRebuildRunSummary {
|
||||
providers_checked: providers.len(),
|
||||
..PoolScoreRebuildRunSummary::empty()
|
||||
};
|
||||
|
||||
let start_provider_index =
|
||||
load_runtime_usize(state, POOL_SCORE_REBUILD_PROVIDER_CURSOR_KEY).await % providers.len();
|
||||
let mut last_provider_index = None;
|
||||
for provider_index in
|
||||
(0..providers.len()).map(|offset| (start_provider_index + offset) % providers.len())
|
||||
{
|
||||
if summary.scores_upserted >= config.max_upserts_per_tick {
|
||||
break;
|
||||
}
|
||||
last_provider_index = Some(provider_index);
|
||||
let provider = providers[provider_index].clone();
|
||||
let endpoints = endpoints_by_provider
|
||||
.remove(&provider.id)
|
||||
.unwrap_or_default();
|
||||
let keys = keys_by_provider.remove(&provider.id).unwrap_or_default();
|
||||
if endpoints.is_empty() || keys.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let models = state
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
provider_id: provider.id.clone(),
|
||||
is_active: Some(true),
|
||||
offset: 0,
|
||||
limit: 10_000,
|
||||
})
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|model| model.is_available)
|
||||
.collect::<Vec<_>>();
|
||||
if models.is_empty() {
|
||||
continue;
|
||||
}
|
||||
|
||||
let total_combinations = endpoints
|
||||
.len()
|
||||
.saturating_mul(models.len())
|
||||
.saturating_mul(keys.len());
|
||||
if total_combinations == 0 {
|
||||
continue;
|
||||
}
|
||||
let provider_cursor_key = provider_offset_cursor_key(&provider.id);
|
||||
let provider_cursor =
|
||||
load_runtime_usize(state, &provider_cursor_key).await % total_combinations.max(1);
|
||||
let remaining_budget = config
|
||||
.max_upserts_per_tick
|
||||
.saturating_sub(summary.scores_upserted);
|
||||
let provider_budget = remaining_budget.min(total_combinations);
|
||||
let mut build_items = Vec::with_capacity(provider_budget);
|
||||
for offset in 0..provider_budget {
|
||||
let flat_index = (provider_cursor + offset) % total_combinations;
|
||||
let (endpoint_index, model_index, key_index) =
|
||||
score_combo_indices(flat_index, models.len(), keys.len());
|
||||
let endpoint = &endpoints[endpoint_index];
|
||||
let api_format = endpoint.api_format.trim();
|
||||
if api_format.is_empty() {
|
||||
continue;
|
||||
}
|
||||
let model = &models[model_index];
|
||||
let key = &keys[key_index];
|
||||
let draft = build_provider_key_pool_score_upsert(
|
||||
key,
|
||||
provider.provider_type.as_str(),
|
||||
api_format,
|
||||
Some(model.id.as_str()),
|
||||
None,
|
||||
now,
|
||||
);
|
||||
build_items.push(ProviderScoreBuildItem {
|
||||
endpoint_index,
|
||||
model_index,
|
||||
key_index,
|
||||
score_id: draft.id,
|
||||
});
|
||||
}
|
||||
if build_items.is_empty() {
|
||||
store_runtime_usize(
|
||||
state,
|
||||
&provider_cursor_key,
|
||||
(provider_cursor + provider_budget) % total_combinations,
|
||||
)
|
||||
.await;
|
||||
continue;
|
||||
}
|
||||
let existing_scores = state
|
||||
.data
|
||||
.get_pool_member_scores_by_ids(&GetPoolMemberScoresByIdsQuery {
|
||||
ids: build_items
|
||||
.iter()
|
||||
.map(|item| item.score_id.clone())
|
||||
.collect(),
|
||||
})
|
||||
.await
|
||||
.unwrap_or_else(|err| {
|
||||
debug!(
|
||||
provider_id = %provider.id,
|
||||
error = ?err,
|
||||
"gateway pool score rebuild: failed to read existing scores by id"
|
||||
);
|
||||
Vec::new()
|
||||
})
|
||||
.into_iter()
|
||||
.map(|score| (score.id.clone(), score))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let mut provider_upserts = 0usize;
|
||||
summary.keys_seen = summary.keys_seen.saturating_add(keys.len());
|
||||
for item in &build_items {
|
||||
if summary.scores_upserted >= config.max_upserts_per_tick {
|
||||
break;
|
||||
}
|
||||
let endpoint = &endpoints[item.endpoint_index];
|
||||
let model = &models[item.model_index];
|
||||
let key = &keys[item.key_index];
|
||||
let existing = existing_scores.get(&item.score_id);
|
||||
let upsert = build_provider_key_pool_score_upsert(
|
||||
key,
|
||||
provider.provider_type.as_str(),
|
||||
endpoint.api_format.trim(),
|
||||
Some(model.id.as_str()),
|
||||
existing,
|
||||
now,
|
||||
);
|
||||
if state
|
||||
.data
|
||||
.upsert_pool_member_score(upsert)
|
||||
.await
|
||||
.map_err(|err| GatewayError::Internal(format!("{err:?}")))?
|
||||
.is_some()
|
||||
{
|
||||
summary.scores_upserted = summary.scores_upserted.saturating_add(1);
|
||||
provider_upserts = provider_upserts.saturating_add(1);
|
||||
}
|
||||
}
|
||||
store_runtime_usize(
|
||||
state,
|
||||
&provider_cursor_key,
|
||||
(provider_cursor + provider_budget) % total_combinations,
|
||||
)
|
||||
.await;
|
||||
if provider_upserts > 0 {
|
||||
summary.providers_scored = summary.providers_scored.saturating_add(1);
|
||||
}
|
||||
}
|
||||
if let Some(last_provider_index) = last_provider_index {
|
||||
store_runtime_usize(
|
||||
state,
|
||||
POOL_SCORE_REBUILD_PROVIDER_CURSOR_KEY,
|
||||
(last_provider_index + 1) % providers.len(),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
Ok(summary)
|
||||
}
|
||||
|
||||
pub(crate) async fn perform_pool_score_rebuild_once(
|
||||
state: &AppState,
|
||||
) -> Result<PoolScoreRebuildRunSummary, GatewayError> {
|
||||
perform_pool_score_rebuild_once_with_config(state, PoolScoreRebuildWorkerConfig::from_env())
|
||||
.await
|
||||
}
|
||||
|
||||
pub(crate) fn spawn_pool_score_rebuild_worker(
|
||||
state: AppState,
|
||||
) -> Option<tokio::task::JoinHandle<()>> {
|
||||
if !state.has_provider_catalog_data_reader()
|
||||
|| !state.data.has_pool_score_reader()
|
||||
|| !state.data.has_pool_score_writer()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let config = PoolScoreRebuildWorkerConfig::from_env();
|
||||
Some(tokio::spawn(async move {
|
||||
if let Err(err) = perform_pool_score_rebuild_once_with_config(&state, config).await {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway pool score rebuild initial tick failed"
|
||||
);
|
||||
}
|
||||
let mut interval = tokio::time::interval(config.interval);
|
||||
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
|
||||
loop {
|
||||
interval.tick().await;
|
||||
match perform_pool_score_rebuild_once_with_config(&state, config).await {
|
||||
Ok(summary) if summary.scores_upserted > 0 => {
|
||||
info!(
|
||||
providers_checked = summary.providers_checked,
|
||||
providers_scored = summary.providers_scored,
|
||||
keys_seen = summary.keys_seen,
|
||||
scores_upserted = summary.scores_upserted,
|
||||
"gateway pool score rebuild completed"
|
||||
);
|
||||
}
|
||||
Ok(_) => {}
|
||||
Err(err) => {
|
||||
warn!(
|
||||
error = ?err,
|
||||
"gateway pool score rebuild worker tick failed"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::score_combo_indices;
|
||||
|
||||
#[test]
|
||||
fn score_combo_indices_walks_endpoint_model_key_order() {
|
||||
assert_eq!(score_combo_indices(0, 2, 3), (0, 0, 0));
|
||||
assert_eq!(score_combo_indices(2, 2, 3), (0, 0, 2));
|
||||
assert_eq!(score_combo_indices(3, 2, 3), (0, 1, 0));
|
||||
assert_eq!(score_combo_indices(6, 2, 3), (1, 0, 0));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user