refactor: extract runtime state backends

This commit is contained in:
fawney19
2026-05-08 00:18:12 +08:00
parent 6f620d92be
commit 6247ac3edc
111 changed files with 4358 additions and 3203 deletions
@@ -49,27 +49,11 @@ pub(super) async fn read_admin_provider_ops_balance_cache(
state: &AdminAppState<'_>,
provider_id: &str,
) -> AdminProviderOpsBalanceCacheLookup {
let Some(runner) = state.redis_kv_runner() else {
return AdminProviderOpsBalanceCacheLookup::Unavailable;
};
let mut connection = match runner.client().get_multiplexed_async_connection().await {
Ok(connection) => connection,
Err(err) => {
warn!(error = %err, provider_id, "failed to connect to redis for provider ops balance cache");
return AdminProviderOpsBalanceCacheLookup::Unavailable;
}
};
let namespaced_key = runner.keyspace().key(&format!(
"{ADMIN_PROVIDER_OPS_BALANCE_CACHE_PREFIX}{provider_id}"
));
let raw = match redis::cmd("GET")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
{
let raw_key = format!("{ADMIN_PROVIDER_OPS_BALANCE_CACHE_PREFIX}{provider_id}");
let raw = match state.runtime_state().kv_get(&raw_key).await {
Ok(raw) => raw,
Err(err) => {
warn!(error = %err, provider_id, "failed to read provider ops balance cache");
warn!(error = %err, provider_id, "failed to read provider ops balance runtime cache");
return AdminProviderOpsBalanceCacheLookup::Unavailable;
}
};
@@ -93,9 +77,6 @@ pub(super) async fn store_admin_provider_ops_balance_cache(
let Some(ttl_seconds) = balance_cache_ttl_seconds(payload) else {
return;
};
let Some(runner) = state.redis_kv_runner() else {
return;
};
let serialized = match serde_json::to_string(payload) {
Ok(serialized) => serialized,
Err(err) => {
@@ -107,11 +88,12 @@ pub(super) async fn store_admin_provider_ops_balance_cache(
return;
}
};
if let Err(err) = runner
.setex(
if let Err(err) = state
.runtime_state()
.kv_set(
&format!("{ADMIN_PROVIDER_OPS_BALANCE_CACHE_PREFIX}{provider_id}"),
&serialized,
Some(ttl_seconds),
serialized,
Some(Duration::from_secs(ttl_seconds)),
)
.await
{
@@ -123,11 +105,9 @@ pub(super) async fn clear_admin_provider_ops_balance_cache(
state: &AdminAppState<'_>,
provider_id: &str,
) {
let Some(runner) = state.redis_kv_runner() else {
return;
};
if let Err(err) = runner
.del(&format!(
if let Err(err) = state
.runtime_state()
.kv_delete(&format!(
"{ADMIN_PROVIDER_OPS_BALANCE_CACHE_PREFIX}{provider_id}"
))
.await
@@ -243,15 +223,11 @@ fn balance_cache_ttl_seconds(payload: &Value) -> Option<u64> {
fn admin_provider_ops_balance_refresh_key(state: &AdminAppState<'_>, provider_id: &str) -> String {
let raw_key = format!("{ADMIN_PROVIDER_OPS_BALANCE_REFRESH_PREFIX}{provider_id}");
if let Some(runner) = state.redis_kv_runner() {
format!(
"{:p}:{}",
state.app(),
runner.keyspace().key(raw_key.as_str())
)
} else {
format!("{:p}:{raw_key}", state.app())
}
format!(
"{:p}:{}",
state.app(),
state.runtime_state().namespace_key(raw_key.as_str())
)
}
async fn finish_refresh_provider(refresh_key: &str) {
@@ -102,7 +102,9 @@ pub(super) async fn handle_admin_provider_ops_action(
cached
}
AdminProviderOpsBalanceCacheLookup::Miss => {
if query_param_bool(query_string, "refresh", true) {
if query_param_bool(query_string, "refresh", true)
&& !state.runtime_state().is_memory()
{
spawn_admin_provider_ops_balance_refresh(state, provider_id).await;
admin_provider_ops_pending_balance_response("余额数据加载中,请稍后刷新")
} else {
@@ -86,8 +86,25 @@ pub(super) async fn handle_admin_provider_ops_batch_balance(
cached
}
AdminProviderOpsBalanceCacheLookup::Miss => {
spawn_admin_provider_ops_balance_refresh(state, &provider_id).await;
admin_provider_ops_pending_balance_response("余额数据加载中,请稍后刷新")
if state.runtime_state().is_memory() {
let payload = admin_provider_ops_local_action_response(
state,
&provider_id,
provider.as_ref(),
&provider_endpoints,
"query_balance",
None,
)
.await;
store_admin_provider_ops_balance_cache(state, &provider_id, &payload)
.await;
payload
} else {
spawn_admin_provider_ops_balance_refresh(state, &provider_id).await;
admin_provider_ops_pending_balance_response(
"余额数据加载中,请稍后刷新",
)
}
}
AdminProviderOpsBalanceCacheLookup::Unavailable => {
let payload = admin_provider_ops_local_action_response(
@@ -1,51 +1,33 @@
use aether_data::driver::redis::RedisKeyspace;
pub(super) fn pool_sticky_pattern(keyspace: &RedisKeyspace, provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:sticky:*"))
pub(super) fn pool_sticky_pattern(provider_id: &str) -> String {
format!("ap:{provider_id}:sticky:*")
}
pub(super) fn pool_sticky_key(
keyspace: &RedisKeyspace,
provider_id: &str,
session_token: &str,
) -> String {
keyspace.key(&format!("ap:{provider_id}:sticky:{session_token}"))
pub(super) fn pool_sticky_key(provider_id: &str, session_token: &str) -> String {
format!("ap:{provider_id}:sticky:{session_token}")
}
pub(super) fn pool_lru_key(keyspace: &RedisKeyspace, provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:lru"))
pub(super) fn pool_lru_key(provider_id: &str) -> String {
format!("ap:{provider_id}:lru")
}
pub(super) fn pool_cooldown_key(
keyspace: &RedisKeyspace,
provider_id: &str,
key_id: &str,
) -> String {
keyspace.key(&format!("ap:{provider_id}:cooldown:{key_id}"))
pub(super) fn pool_cooldown_key(provider_id: &str, key_id: &str) -> String {
format!("ap:{provider_id}:cooldown:{key_id}")
}
pub(super) fn pool_cooldown_index_key(keyspace: &RedisKeyspace, provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:cooldown_idx"))
pub(super) fn pool_cooldown_index_key(provider_id: &str) -> String {
format!("ap:{provider_id}:cooldown_idx")
}
pub(super) fn pool_cost_key(keyspace: &RedisKeyspace, provider_id: &str, key_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:cost:{key_id}"))
pub(super) fn pool_cost_key(provider_id: &str, key_id: &str) -> String {
format!("ap:{provider_id}:cost:{key_id}")
}
pub(super) fn pool_latency_key(
keyspace: &RedisKeyspace,
provider_id: &str,
key_id: &str,
) -> String {
keyspace.key(&format!("ap:{provider_id}:latency:{key_id}"))
pub(super) fn pool_latency_key(provider_id: &str, key_id: &str) -> String {
format!("ap:{provider_id}:latency:{key_id}")
}
pub(super) fn pool_stream_timeout_key(
keyspace: &RedisKeyspace,
provider_id: &str,
key_id: &str,
) -> String {
keyspace.key(&format!("ap:{provider_id}:stream_timeout:{key_id}"))
pub(super) fn pool_stream_timeout_key(provider_id: &str, key_id: &str) -> String {
format!("ap:{provider_id}:stream_timeout:{key_id}")
}
pub(super) fn parse_pool_cost_member(member: &str) -> u64 {
@@ -62,35 +44,23 @@ pub(super) fn parse_pool_latency_member(member: &str) -> u64 {
.unwrap_or(0)
}
pub(super) fn pool_cooldown_keys(
keyspace: &RedisKeyspace,
provider_id: &str,
key_ids: &[String],
) -> Vec<String> {
pub(super) fn pool_cooldown_keys(provider_id: &str, key_ids: &[String]) -> Vec<String> {
key_ids
.iter()
.map(|key_id| pool_cooldown_key(keyspace, provider_id, key_id))
.map(|key_id| pool_cooldown_key(provider_id, key_id))
.collect()
}
pub(super) fn pool_cost_keys(
keyspace: &RedisKeyspace,
provider_id: &str,
key_ids: &[String],
) -> Vec<String> {
pub(super) fn pool_cost_keys(provider_id: &str, key_ids: &[String]) -> Vec<String> {
key_ids
.iter()
.map(|key_id| pool_cost_key(keyspace, provider_id, key_id))
.map(|key_id| pool_cost_key(provider_id, key_id))
.collect()
}
pub(super) fn pool_latency_keys(
keyspace: &RedisKeyspace,
provider_id: &str,
key_ids: &[String],
) -> Vec<String> {
pub(super) fn pool_latency_keys(provider_id: &str, key_ids: &[String]) -> Vec<String> {
key_ids
.iter()
.map(|key_id| pool_latency_key(keyspace, provider_id, key_id))
.map(|key_id| pool_latency_key(provider_id, key_id))
.collect()
}
@@ -1,29 +1,18 @@
use super::keys::{pool_cooldown_index_key, pool_cooldown_key};
use crate::handlers::admin::request::AdminAppState;
use tracing::warn;
pub(crate) async fn clear_admin_provider_pool_cooldown(
state: &AdminAppState<'_>,
provider_id: &str,
key_id: &str,
) {
let Some(runner) = state.redis_kv_runner() else {
return;
};
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis to clear cooldown for key {key_id}");
return;
};
let keyspace = runner.keyspace().clone();
let _: Result<(), _> = redis::pipe()
.cmd("DEL")
.arg(pool_cooldown_key(&keyspace, provider_id, key_id))
.ignore()
.cmd("SREM")
.arg(pool_cooldown_index_key(&keyspace, provider_id))
.arg(key_id)
.ignore()
.query_async(&mut connection)
let _ = state
.runtime_state()
.kv_delete(&pool_cooldown_key(provider_id, key_id))
.await;
let _ = state
.runtime_state()
.set_remove(&pool_cooldown_index_key(provider_id), key_id)
.await;
}
@@ -32,8 +21,8 @@ pub(crate) async fn reset_admin_provider_pool_cost(
provider_id: &str,
key_id: &str,
) {
let Some(runner) = state.redis_kv_runner() else {
return;
};
let _ = runner.del(&format!("ap:{provider_id}:cost:{key_id}")).await;
let _ = state
.runtime_state()
.score_remove_by_score(&format!("ap:{provider_id}:cost:{key_id}"), f64::INFINITY)
.await;
}
@@ -4,10 +4,9 @@ use super::keys::{
pool_sticky_pattern,
};
use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolRuntimeState, ADMIN_PROVIDER_POOL_SCAN_BATCH,
AdminProviderPoolConfig, AdminProviderPoolRuntimeState,
};
use crate::GatewayError;
use aether_data::driver::redis::RedisKvRunner;
use aether_runtime_state::RuntimeState;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;
@@ -19,165 +18,78 @@ fn current_unix_secs() -> u64 {
.as_secs()
}
async fn scan_redis_keys(
connection: &mut redis::aio::MultiplexedConnection,
pattern: &str,
) -> Result<Vec<String>, GatewayError> {
let mut cursor = 0u64;
let mut keys = Vec::new();
loop {
let (next_cursor, batch): (u64, Vec<String>) = redis::cmd("SCAN")
.arg(cursor)
.arg("MATCH")
.arg(pattern)
.arg("COUNT")
.arg(ADMIN_PROVIDER_POOL_SCAN_BATCH)
.query_async(connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
keys.extend(batch);
if next_cursor == 0 {
break;
}
cursor = next_cursor;
}
Ok(keys)
}
pub(crate) async fn read_admin_provider_pool_cooldown_counts(
runner: &RedisKvRunner,
runtime: &RuntimeState,
provider_ids: &[String],
) -> BTreeMap<String, usize> {
if provider_ids.is_empty() {
return BTreeMap::new();
}
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis for cooldown counts");
return BTreeMap::new();
};
let keyspace = runner.keyspace().clone();
let mut pipeline = redis::pipe();
let mut counts = BTreeMap::new();
for provider_id in provider_ids {
pipeline
.cmd("SCARD")
.arg(pool_cooldown_index_key(&keyspace, provider_id));
}
match pipeline.query_async::<Vec<u64>>(&mut connection).await {
Ok(counts) => provider_ids
.iter()
.cloned()
.zip(counts)
.map(|(provider_id, count)| (provider_id, count as usize))
.collect(),
Err(err) => {
warn!(
"gateway admin provider pool: failed to batch read cooldown counts: {:?}",
err
);
BTreeMap::new()
}
let count = runtime
.set_len(&pool_cooldown_index_key(provider_id))
.await
.unwrap_or(0);
counts.insert(provider_id.clone(), count);
}
counts
}
pub(crate) async fn read_admin_provider_pool_runtime_state(
runner: &RedisKvRunner,
runtime: &RuntimeState,
provider_id: &str,
key_ids: &[String],
pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>,
) -> AdminProviderPoolRuntimeState {
let mut runtime = AdminProviderPoolRuntimeState::default();
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis for provider {provider_id}");
return runtime;
};
let keyspace = runner.keyspace().clone();
let cooldown_keys = pool_cooldown_keys(&keyspace, provider_id, key_ids);
let cost_keys = pool_cost_keys(&keyspace, provider_id, key_ids);
let latency_keys = pool_latency_keys(&keyspace, provider_id, key_ids);
let mut state = AdminProviderPoolRuntimeState::default();
let cooldown_keys = pool_cooldown_keys(provider_id, key_ids);
let cost_keys = pool_cost_keys(provider_id, key_ids);
let latency_keys = pool_latency_keys(provider_id, key_ids);
if let Some(sticky_session_token) = sticky_session_token
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|_| pool_config.sticky_session_ttl_seconds > 0)
{
let sticky_key = pool_sticky_key(&keyspace, provider_id, sticky_session_token);
let sticky_bound_key_id = redis::cmd("GET")
.arg(&sticky_key)
.query_async::<Option<String>>(&mut connection)
.await
.unwrap_or_else(|err| {
warn!(
"gateway admin provider pool: failed to read sticky binding for provider {provider_id}: {:?}",
err
);
None
});
if let Some(bound_key_id) = sticky_bound_key_id {
let cooldown_key = pool_cooldown_key(&keyspace, provider_id, &bound_key_id);
runtime.sticky_bound_key_id = match redis::cmd("EXISTS")
.arg(&cooldown_key)
.query_async::<u64>(&mut connection)
.await
{
Ok(0) => {
let _: Result<bool, _> = redis::cmd("EXPIRE")
.arg(&sticky_key)
.arg(pool_config.sticky_session_ttl_seconds)
.query_async(&mut connection)
let sticky_key = pool_sticky_key(provider_id, sticky_session_token);
if let Ok(Some(bound_key_id)) = runtime.kv_get(&sticky_key).await {
let cooldown_key = pool_cooldown_key(provider_id, &bound_key_id);
match runtime.kv_exists(&cooldown_key).await {
Ok(false) => {
let _ = runtime
.key_expire(
&sticky_key,
std::time::Duration::from_secs(pool_config.sticky_session_ttl_seconds),
)
.await;
Some(bound_key_id)
state.sticky_bound_key_id = Some(bound_key_id);
}
Ok(_) => {
let _: Result<i64, _> = redis::cmd("DEL")
.arg(&sticky_key)
.query_async(&mut connection)
.await;
None
Ok(true) => {
let _ = runtime.kv_delete(&sticky_key).await;
}
Err(err) => {
warn!(
"gateway admin provider pool: failed to validate sticky cooldown for provider {provider_id}: {:?}",
err
);
Some(bound_key_id)
state.sticky_bound_key_id = Some(bound_key_id);
}
};
}
}
}
let sticky_keys = match scan_redis_keys(
&mut connection,
&pool_sticky_pattern(&keyspace, provider_id),
)
.await
{
Ok(keys) => keys,
Err(err) => {
warn!(
"gateway admin provider pool: failed to scan sticky keys for provider {provider_id}: {:?}",
err
);
Vec::new()
}
};
runtime.total_sticky_sessions = sticky_keys.len();
let sticky_keys = runtime
.scan_keys(&pool_sticky_pattern(provider_id), 200)
.await
.unwrap_or_default();
state.total_sticky_sessions = sticky_keys.len();
if !sticky_keys.is_empty() {
for chunk in sticky_keys.chunks(ADMIN_PROVIDER_POOL_SCAN_BATCH as usize) {
let values = redis::cmd("MGET")
.arg(chunk)
.query_async::<Vec<Option<String>>>(&mut connection)
.await;
let Ok(values) = values else {
warn!(
"gateway admin provider pool: failed to read sticky bindings for provider {provider_id}"
);
break;
};
let raw_keys = sticky_keys
.iter()
.map(|key| runtime.strip_namespace(key).to_string())
.collect::<Vec<_>>();
if let Ok(values) = runtime.kv_get_many(&raw_keys).await {
for bound_key_id in values.into_iter().flatten() {
*runtime
*state
.sticky_sessions_by_key
.entry(bound_key_id)
.or_insert(0) += 1;
@@ -186,120 +98,61 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
}
if !cooldown_keys.is_empty() {
let cooldown_reasons = redis::cmd("MGET")
.arg(&cooldown_keys)
.query_async::<Vec<Option<String>>>(&mut connection)
let cooldown_reasons = runtime
.kv_get_many(&cooldown_keys)
.await
.unwrap_or_else(|err| {
warn!(
"gateway admin provider pool: failed to batch read cooldown reasons for provider {provider_id}: {:?}",
err
);
vec![None; cooldown_keys.len()]
});
let mut ttl_pipeline = redis::pipe();
for cooldown_key in &cooldown_keys {
ttl_pipeline.cmd("TTL").arg(cooldown_key);
}
let cooldown_ttls = ttl_pipeline
.query_async::<Vec<i64>>(&mut connection)
.await
.unwrap_or_else(|err| {
warn!(
"gateway admin provider pool: failed to batch read cooldown ttl for provider {provider_id}: {:?}",
err
);
vec![-1; cooldown_keys.len()]
});
for (((key_id, _cooldown_key), reason), ttl) in key_ids
.unwrap_or_else(|_| vec![None; cooldown_keys.len()]);
for (key_id, (cooldown_key, reason)) in key_ids
.iter()
.zip(cooldown_keys.iter())
.zip(cooldown_reasons)
.zip(cooldown_ttls)
.zip(cooldown_keys.iter().zip(cooldown_reasons))
{
if let Some(reason) = reason {
runtime
.cooldown_reason_by_key
.insert(key_id.clone(), reason);
if let Ok(ttl_seconds) = u64::try_from(ttl) {
if ttl_seconds > 0 {
runtime
.cooldown_ttl_by_key
.insert(key_id.clone(), ttl_seconds);
state.cooldown_reason_by_key.insert(key_id.clone(), reason);
if let Ok(Some(ttl)) = runtime.kv_ttl_seconds(cooldown_key).await {
if let Ok(ttl_seconds) = u64::try_from(ttl) {
if ttl_seconds > 0 {
state
.cooldown_ttl_by_key
.insert(key_id.clone(), ttl_seconds);
}
}
}
}
}
}
if !cost_keys.is_empty() {
let window_start = current_unix_secs().saturating_sub(pool_config.cost_window_seconds);
let mut cost_pipeline = redis::pipe();
for cost_key in &cost_keys {
cost_pipeline
.cmd("ZRANGEBYSCORE")
.arg(cost_key)
.arg(window_start)
.arg("+inf");
}
let members_by_key = cost_pipeline
.query_async::<Vec<Vec<String>>>(&mut connection)
let now = current_unix_secs();
for (key_id, cost_key) in key_ids.iter().zip(cost_keys) {
let window_start = now.saturating_sub(pool_config.cost_window_seconds) as f64;
let total = runtime
.score_range_by_min(&cost_key, window_start)
.await
.unwrap_or_else(|err| {
warn!(
"gateway admin provider pool: failed to batch read cost windows for provider {provider_id}: {:?}",
err
);
vec![Vec::new(); cost_keys.len()]
});
for (key_id, members) in key_ids.iter().zip(members_by_key) {
let total = members
.iter()
.map(|member| parse_pool_cost_member(member))
.sum::<u64>();
runtime
.cost_window_usage_by_key
.insert(key_id.clone(), total);
.unwrap_or_default()
.iter()
.map(|member| parse_pool_cost_member(member))
.sum::<u64>();
if total > 0 {
state.cost_window_usage_by_key.insert(key_id.clone(), total);
}
}
if !latency_keys.is_empty() {
let window_start = current_unix_secs().saturating_sub(pool_config.latency_window_seconds);
let mut latency_pipeline = redis::pipe();
for latency_key in &latency_keys {
latency_pipeline
.cmd("ZRANGEBYSCORE")
.arg(latency_key)
.arg(window_start)
.arg("+inf");
}
let members_by_key = latency_pipeline
.query_async::<Vec<Vec<String>>>(&mut connection)
for (key_id, latency_key) in key_ids.iter().zip(latency_keys) {
let window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64;
let samples = runtime
.score_range_by_min(&latency_key, window_start)
.await
.unwrap_or_else(|err| {
warn!(
"gateway admin provider pool: failed to batch read latency windows for provider {provider_id}: {:?}",
err
);
vec![Vec::new(); latency_keys.len()]
});
for (key_id, members) in key_ids.iter().zip(members_by_key) {
let samples = members
.iter()
.map(|member| parse_pool_latency_member(member))
.filter(|value| *value > 0)
.collect::<Vec<_>>();
if samples.is_empty() {
continue;
}
let total = samples.iter().sum::<u64>() as f64;
let average = total / samples.len() as f64;
if average.is_finite() && average >= 0.0 {
runtime
.latency_avg_ms_by_key
.insert(key_id.clone(), average);
}
.unwrap_or_default()
.iter()
.map(|member| parse_pool_latency_member(member))
.filter(|value| *value > 0)
.collect::<Vec<_>>();
if samples.is_empty() {
continue;
}
let total = samples.iter().sum::<u64>() as f64;
let average = total / samples.len() as f64;
if average.is_finite() && average >= 0.0 {
state.latency_avg_ms_by_key.insert(key_id.clone(), average);
}
}
@@ -310,55 +163,37 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
.any(|item| item.enabled))
&& !key_ids.is_empty()
{
let mut command = redis::cmd("ZMSCORE");
command.arg(pool_lru_key(&keyspace, provider_id));
for key_id in key_ids {
command.arg(key_id);
}
if let Ok(scores) = command
.query_async::<Vec<Option<f64>>>(&mut connection)
if let Ok(scores) = runtime
.score_many(&pool_lru_key(provider_id), key_ids)
.await
{
for (key_id, score) in key_ids.iter().zip(scores) {
if let Some(score) = score {
runtime.lru_score_by_key.insert(key_id.clone(), score);
state.lru_score_by_key.insert(key_id.clone(), score);
}
}
}
}
runtime
state
}
pub(crate) async fn read_admin_provider_pool_cooldown_count(
runner: &RedisKvRunner,
runtime: &RuntimeState,
provider_id: &str,
) -> usize {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis for provider {provider_id}");
return 0;
};
let keyspace = runner.keyspace().clone();
redis::cmd("SCARD")
.arg(pool_cooldown_index_key(&keyspace, provider_id))
.query_async::<u64>(&mut connection)
runtime
.set_len(&pool_cooldown_index_key(provider_id))
.await
.map(|value| value as usize)
.unwrap_or(0)
}
pub(crate) async fn read_admin_provider_pool_cooldown_key_ids(
runner: &RedisKvRunner,
runtime: &RuntimeState,
provider_id: &str,
) -> Vec<String> {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis for provider {provider_id}");
return Vec::new();
};
let keyspace = runner.keyspace().clone();
redis::cmd("SMEMBERS")
.arg(pool_cooldown_index_key(&keyspace, provider_id))
.query_async::<Vec<String>>(&mut connection)
runtime
.set_members(&pool_cooldown_index_key(provider_id))
.await
.unwrap_or_default()
}
@@ -1,6 +1,5 @@
use super::reads::read_admin_provider_pool_runtime_state;
use crate::handlers::admin::provider::pool::config::admin_provider_pool_config;
use crate::handlers::admin::provider::shared::support::AdminProviderPoolRuntimeState;
use crate::handlers::admin::request::AdminAppState;
use serde_json::json;
@@ -34,19 +33,14 @@ pub(crate) async fn build_admin_provider_pool_status_payload(
.ok()
.unwrap_or_default();
let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
let runtime = match state.redis_kv_runner() {
Some(runner) => {
read_admin_provider_pool_runtime_state(
&runner,
&provider.id,
&key_ids,
&pool_config,
None,
)
.await
}
None => AdminProviderPoolRuntimeState::default(),
};
let runtime = read_admin_provider_pool_runtime_state(
state.runtime_state(),
&provider.id,
&key_ids,
&pool_config,
None,
)
.await;
let key_payloads = keys
.into_iter()
.map(|key| {
@@ -5,7 +5,7 @@ use super::keys::{
use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolUnschedulableRule,
};
use aether_data::driver::redis::RedisKvRunner;
use aether_runtime_state::RuntimeState;
use regex::Regex;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
@@ -277,7 +277,7 @@ fn resolve_transient_cooldown_ttl(
}
async fn set_pool_cooldown(
runner: &RedisKvRunner,
runtime: &RuntimeState,
provider_id: &str,
key_id: &str,
reason: &str,
@@ -288,39 +288,32 @@ async fn set_pool_cooldown(
}
let ttl_seconds = ttl_seconds.min(MAX_POOL_COOLDOWN_SECONDS);
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!(
"gateway admin provider pool: failed to connect redis to set cooldown for key {key_id}"
);
return;
};
let keyspace = runner.keyspace().clone();
let result: Result<(), _> = redis::pipe()
.cmd("SETEX")
.arg(pool_cooldown_key(&keyspace, provider_id, key_id))
.arg(ttl_seconds)
.arg(reason)
.ignore()
.cmd("SADD")
.arg(pool_cooldown_index_key(&keyspace, provider_id))
.arg(key_id)
.ignore()
.cmd("EXPIRE")
.arg(pool_cooldown_index_key(&keyspace, provider_id))
.arg(ttl_seconds.saturating_add(60))
.ignore()
.query_async(&mut connection)
.await;
if let Err(err) = result {
if let Err(err) = runtime
.kv_set(
&pool_cooldown_key(provider_id, key_id),
reason.to_string(),
Some(std::time::Duration::from_secs(ttl_seconds)),
)
.await
{
warn!(
"gateway admin provider pool: failed to set cooldown for provider {provider_id} key {key_id}: {:?}",
err
);
}
let _ = runtime
.set_add(&pool_cooldown_index_key(provider_id), key_id)
.await;
let _ = runtime
.key_expire(
&pool_cooldown_index_key(provider_id),
std::time::Duration::from_secs(ttl_seconds.saturating_add(60)),
)
.await;
}
async fn invalidate_pool_oauth_cache(runner: &RedisKvRunner, key_id: &str) {
if let Err(err) = runner.del(&oauth_cache_key(key_id)).await {
async fn invalidate_pool_oauth_cache(runtime: &RuntimeState, key_id: &str) {
if let Err(err) = runtime.kv_delete(&oauth_cache_key(key_id)).await {
warn!(
"gateway admin provider pool: failed to invalidate oauth cache for key {key_id}: {:?}",
err
@@ -339,7 +332,7 @@ fn matching_unschedulable_rule<'a>(
}
pub(crate) async fn record_admin_provider_pool_success(
runner: &RedisKvRunner,
runtime: &RuntimeState,
provider_id: &str,
key_id: &str,
pool_config: &AdminProviderPoolConfig,
@@ -347,111 +340,72 @@ pub(crate) async fn record_admin_provider_pool_success(
tokens_used: u64,
ttfb_ms: Option<u64>,
) {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis to record success for key {key_id}");
return;
};
let keyspace = runner.keyspace().clone();
let now = current_unix_secs_f64();
let mut pipeline = redis::pipe();
let mut has_commands = false;
if let Some(sticky_session_token) = sticky_session_token
.map(str::trim)
.filter(|value| !value.is_empty())
.filter(|_| pool_config.sticky_session_ttl_seconds > 0)
{
pipeline
.cmd("SETEX")
.arg(pool_sticky_key(
&keyspace,
provider_id,
sticky_session_token,
))
.arg(pool_config.sticky_session_ttl_seconds)
.arg(key_id)
.ignore();
has_commands = true;
let _ = runtime
.kv_set(
&pool_sticky_key(provider_id, sticky_session_token),
key_id.to_string(),
Some(std::time::Duration::from_secs(
pool_config.sticky_session_ttl_seconds,
)),
)
.await;
}
if should_touch_lru(pool_config) {
pipeline
.cmd("ZADD")
.arg(pool_lru_key(&keyspace, provider_id))
.arg(now)
.arg(key_id)
.ignore();
has_commands = true;
let _ = runtime
.score_set(&pool_lru_key(provider_id), key_id, now)
.await;
}
if tokens_used > 0 && pool_config.cost_limit_per_key_tokens.is_some() {
let cost_key = pool_cost_key(&keyspace, provider_id, key_id);
let cost_key = pool_cost_key(provider_id, key_id);
let window_seconds = pool_config.cost_window_seconds.max(1);
let member = format!("{}:{tokens_used}", Uuid::new_v4().simple());
pipeline
.cmd("ZADD")
.arg(&cost_key)
.arg(now)
.arg(member)
.ignore()
.cmd("ZREMRANGEBYSCORE")
.arg(&cost_key)
.arg("-inf")
.arg(now - window_seconds as f64)
.ignore()
.cmd("EXPIRE")
.arg(&cost_key)
.arg(window_seconds.saturating_add(600))
.ignore();
has_commands = true;
let _ = runtime.score_set(&cost_key, &member, now).await;
let _ = runtime
.score_remove_by_score(&cost_key, now - window_seconds as f64)
.await;
let _ = runtime
.key_expire(
&cost_key,
std::time::Duration::from_secs(window_seconds.saturating_add(600)),
)
.await;
}
if let Some(ttfb_ms) = ttfb_ms
.filter(|value| should_record_latency(pool_config))
.filter(|_| pool_config.latency_window_seconds > 0)
{
let latency_key = pool_latency_key(&keyspace, provider_id, key_id);
let latency_key = pool_latency_key(provider_id, key_id);
let window_seconds = pool_config.latency_window_seconds.max(1);
let sample_limit = pool_config.latency_sample_limit.max(1);
let member = format!("{}:{ttfb_ms}", Uuid::new_v4().simple());
pipeline
.cmd("ZADD")
.arg(&latency_key)
.arg(now)
.arg(member)
.ignore()
.cmd("ZREMRANGEBYSCORE")
.arg(&latency_key)
.arg("-inf")
.arg(now - window_seconds as f64)
.ignore()
.cmd("ZREMRANGEBYRANK")
.arg(&latency_key)
.arg(0)
.arg(-((sample_limit as i64) + 1))
.ignore()
.cmd("EXPIRE")
.arg(&latency_key)
.arg(window_seconds.saturating_add(600))
.ignore();
has_commands = true;
}
if !has_commands {
return;
}
let result: Result<(), _> = pipeline.query_async(&mut connection).await;
if let Err(err) = result {
warn!(
"gateway admin provider pool: failed to record success feedback for provider {provider_id} key {key_id}: {:?}",
err
);
let _ = runtime.score_set(&latency_key, &member, now).await;
let _ = runtime
.score_remove_by_score(&latency_key, now - window_seconds as f64)
.await;
let _ = runtime
.score_remove_by_rank(&latency_key, 0, -((sample_limit as i64) + 1))
.await;
let _ = runtime
.key_expire(
&latency_key,
std::time::Duration::from_secs(window_seconds.saturating_add(600)),
)
.await;
}
}
pub(crate) async fn record_admin_provider_pool_error(
runner: &RedisKvRunner,
runtime: &RuntimeState,
provider_id: &str,
key_id: &str,
pool_config: &AdminProviderPoolConfig,
@@ -466,7 +420,7 @@ pub(crate) async fn record_admin_provider_pool_error(
let error_message = extract_error_message(error_body).to_ascii_lowercase();
if status_code == 401 {
invalidate_pool_oauth_cache(runner, key_id).await;
invalidate_pool_oauth_cache(runtime, key_id).await;
return;
}
@@ -482,7 +436,7 @@ pub(crate) async fn record_admin_provider_pool_error(
return;
}
set_pool_cooldown(
runner,
runtime,
provider_id,
key_id,
"forbidden_403",
@@ -503,7 +457,7 @@ pub(crate) async fn record_admin_provider_pool_error(
{
let ttl_seconds = (rule.duration_minutes.max(1)).saturating_mul(60).max(60);
set_pool_cooldown(
runner,
runtime,
provider_id,
key_id,
&format!("rule:{}", rule.keyword),
@@ -520,13 +474,20 @@ pub(crate) async fn record_admin_provider_pool_error(
.or_else(|| parse_google_quota_cooldown_seconds(error_body)),
pool_config,
);
set_pool_cooldown(runner, provider_id, key_id, "rate_limited_429", ttl_seconds).await;
set_pool_cooldown(
runtime,
provider_id,
key_id,
"rate_limited_429",
ttl_seconds,
)
.await;
return;
}
if status_code == 529 {
set_pool_cooldown(
runner,
runtime,
provider_id,
key_id,
"overloaded_529",
@@ -555,12 +516,12 @@ pub(crate) async fn record_admin_provider_pool_error(
parse_retry_after_seconds(response_headers),
pool_config,
);
set_pool_cooldown(runner, provider_id, key_id, &reason, ttl_seconds).await;
set_pool_cooldown(runtime, provider_id, key_id, &reason, ttl_seconds).await;
}
}
pub(crate) async fn record_admin_provider_pool_stream_timeout(
runner: &RedisKvRunner,
runtime: &RuntimeState,
provider_id: &str,
key_id: &str,
pool_config: &AdminProviderPoolConfig,
@@ -569,49 +530,25 @@ pub(crate) async fn record_admin_provider_pool_stream_timeout(
return;
}
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
warn!("gateway admin provider pool: failed to connect redis to record stream timeout for key {key_id}");
return;
};
let keyspace = runner.keyspace().clone();
let timeout_key = pool_stream_timeout_key(&keyspace, provider_id, key_id);
let timeout_key = pool_stream_timeout_key(provider_id, key_id);
let now = current_unix_secs_f64();
let window_seconds = pool_config.stream_timeout_window_seconds.max(1);
let member = Uuid::new_v4().simple().to_string();
let results = redis::pipe()
.cmd("ZREMRANGEBYSCORE")
.arg(&timeout_key)
.arg("-inf")
.arg(now - window_seconds as f64)
.cmd("ZADD")
.arg(&timeout_key)
.arg(now)
.arg(member)
.cmd("ZCARD")
.arg(&timeout_key)
.cmd("EXPIRE")
.arg(&timeout_key)
.arg(window_seconds.saturating_add(60))
.query_async::<Vec<redis::Value>>(&mut connection)
let _ = runtime
.score_remove_by_score(&timeout_key, now - window_seconds as f64)
.await;
let _ = runtime.score_set(&timeout_key, &member, now).await;
let count = runtime.score_len(&timeout_key).await.unwrap_or(0) as u64;
let _ = runtime
.key_expire(
&timeout_key,
std::time::Duration::from_secs(window_seconds.saturating_add(60)),
)
.await;
let count = match results
.ok()
.and_then(|values| values.get(2).cloned())
.and_then(|value| redis::from_redis_value::<u64>(&value).ok())
{
Some(count) => count,
None => {
warn!(
"gateway admin provider pool: failed to compute stream timeout count for provider {provider_id} key {key_id}"
);
return;
}
};
if count >= pool_config.stream_timeout_threshold {
set_pool_cooldown(
runner,
runtime,
provider_id,
key_id,
&format!("stream_timeout_x{count}"),
@@ -628,13 +565,13 @@ mod tests {
record_admin_provider_pool_error, record_admin_provider_pool_stream_timeout,
record_admin_provider_pool_success,
};
use crate::data::{GatewayDataConfig, GatewayDataState};
use crate::handlers::admin::provider::pool::runtime::reads::read_admin_provider_pool_runtime_state;
use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolSchedulingPreset,
AdminProviderPoolUnschedulableRule,
};
use crate::AppState;
use aether_runtime_state::{RedisClientConfig, RuntimeState, RuntimeStateConfig};
use aether_testkit::ManagedRedisServer;
use std::collections::BTreeMap;
@@ -682,14 +619,17 @@ mod tests {
}
}
fn build_runner_app(redis_url: &str, key_prefix: &str) -> AppState {
let data_state = GatewayDataState::from_config(
GatewayDataConfig::disabled().with_redis_url(redis_url, Some(key_prefix)),
)
.expect("data state should build");
async fn build_runner_app(redis_url: &str, key_prefix: &str) -> AppState {
let runtime_state =
RuntimeState::from_config(RuntimeStateConfig::redis(RedisClientConfig {
url: redis_url.to_string(),
key_prefix: Some(key_prefix.to_string()),
}))
.await
.expect("runtime state should build");
AppState::new()
.expect("app state should build")
.with_data_state_for_tests(data_state)
.with_runtime_state(std::sync::Arc::new(runtime_state))
}
#[test]
@@ -798,13 +738,13 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else {
return;
};
let app = build_runner_app(redis.redis_url(), "pool_runtime_success_feedback");
let runner = app.redis_kv_runner().expect("redis runner should exist");
let app = build_runner_app(redis.redis_url(), "pool_runtime_success_feedback").await;
let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config();
let key_ids = vec!["key-1".to_string()];
record_admin_provider_pool_success(
&runner,
runtime,
"provider-1",
"key-1",
&pool_config,
@@ -815,7 +755,7 @@ mod tests {
.await;
let runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime,
"provider-1",
&key_ids,
&pool_config,
@@ -836,14 +776,15 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else {
return;
};
let app = build_runner_app(redis.redis_url(), "pool_runtime_no_sticky_without_affinity");
let runner = app.redis_kv_runner().expect("redis runner should exist");
let app =
build_runner_app(redis.redis_url(), "pool_runtime_no_sticky_without_affinity").await;
let runtime = app.runtime_state.as_ref();
let mut pool_config = sample_pool_config();
pool_config.sticky_session_ttl_seconds = 0;
let key_ids = vec!["key-1".to_string()];
record_admin_provider_pool_success(
&runner,
runtime,
"provider-1",
"key-1",
&pool_config,
@@ -854,7 +795,7 @@ mod tests {
.await;
let runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime,
"provider-1",
&key_ids,
&pool_config,
@@ -874,13 +815,13 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else {
return;
};
let app = build_runner_app(redis.redis_url(), "pool_runtime_error_feedback");
let runner = app.redis_kv_runner().expect("redis runner should exist");
let app = build_runner_app(redis.redis_url(), "pool_runtime_error_feedback").await;
let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config();
let key_ids = vec!["key-2".to_string()];
record_admin_provider_pool_error(
&runner,
runtime,
"provider-1",
"key-2",
&pool_config,
@@ -894,7 +835,7 @@ mod tests {
.await;
let runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime,
"provider-1",
&key_ids,
&pool_config,
@@ -920,13 +861,13 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else {
return;
};
let app = build_runner_app(redis.redis_url(), "pool_runtime_google_quota_cooldown");
let runner = app.redis_kv_runner().expect("redis runner should exist");
let app = build_runner_app(redis.redis_url(), "pool_runtime_google_quota_cooldown").await;
let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config();
let key_ids = vec!["key-google-429".to_string()];
record_admin_provider_pool_error(
&runner,
runtime,
"provider-1",
"key-google-429",
&pool_config,
@@ -949,7 +890,7 @@ mod tests {
.await;
let runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime,
"provider-1",
&key_ids,
&pool_config,
@@ -975,13 +916,13 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else {
return;
};
let app = build_runner_app(redis.redis_url(), "pool_runtime_capped_cooldown");
let runner = app.redis_kv_runner().expect("redis runner should exist");
let app = build_runner_app(redis.redis_url(), "pool_runtime_capped_cooldown").await;
let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config();
let key_ids = vec!["key-long-cooldown".to_string()];
record_admin_provider_pool_error(
&runner,
runtime,
"provider-1",
"key-long-cooldown",
&pool_config,
@@ -995,7 +936,7 @@ mod tests {
.await;
let runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime,
"provider-1",
&key_ids,
&pool_config,
@@ -1021,8 +962,8 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else {
return;
};
let app = build_runner_app(redis.redis_url(), "pool_runtime_circuit_no_cooldown");
let runner = app.redis_kv_runner().expect("redis runner should exist");
let app = build_runner_app(redis.redis_url(), "pool_runtime_circuit_no_cooldown").await;
let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config();
let key_ids = vec!["key-account-disabled".to_string()];
@@ -1035,7 +976,7 @@ mod tests {
Some("account_deactivated_401")
);
record_admin_provider_pool_error(
&runner,
runtime,
"provider-1",
"key-account-disabled",
&pool_config,
@@ -1046,7 +987,7 @@ mod tests {
.await;
let runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime,
"provider-1",
&key_ids,
&pool_config,
@@ -1067,8 +1008,8 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else {
return;
};
let app = build_runner_app(redis.redis_url(), "pool_runtime_unschedulable_rule");
let runner = app.redis_kv_runner().expect("redis runner should exist");
let app = build_runner_app(redis.redis_url(), "pool_runtime_unschedulable_rule").await;
let runtime = app.runtime_state.as_ref();
let mut pool_config = sample_pool_config();
pool_config.unschedulable_rules = vec![AdminProviderPoolUnschedulableRule {
keyword: "review required".to_string(),
@@ -1077,7 +1018,7 @@ mod tests {
let key_ids = vec!["key-3".to_string()];
record_admin_provider_pool_error(
&runner,
runtime,
"provider-1",
"key-3",
&pool_config,
@@ -1088,7 +1029,7 @@ mod tests {
.await;
let runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime,
"provider-1",
&key_ids,
&pool_config,
@@ -1114,8 +1055,8 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else {
return;
};
let app = build_runner_app(redis.redis_url(), "pool_runtime_ignore_400");
let runner = app.redis_kv_runner().expect("redis runner should exist");
let app = build_runner_app(redis.redis_url(), "pool_runtime_ignore_400").await;
let runtime = app.runtime_state.as_ref();
let mut pool_config = sample_pool_config();
pool_config.unschedulable_rules = vec![AdminProviderPoolUnschedulableRule {
keyword: "review required".to_string(),
@@ -1124,7 +1065,7 @@ mod tests {
let key_ids = vec!["key-client-400".to_string()];
record_admin_provider_pool_error(
&runner,
runtime,
"provider-1",
"key-client-400",
&pool_config,
@@ -1135,7 +1076,7 @@ mod tests {
.await;
let runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime,
"provider-1",
&key_ids,
&pool_config,
@@ -1154,21 +1095,31 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else {
return;
};
let app = build_runner_app(redis.redis_url(), "pool_runtime_stream_timeout");
let runner = app.redis_kv_runner().expect("redis runner should exist");
let app = build_runner_app(redis.redis_url(), "pool_runtime_stream_timeout").await;
let runtime_state = app.runtime_state.as_ref();
let mut pool_config = sample_pool_config();
pool_config.stream_timeout_threshold = 2;
pool_config.stream_timeout_window_seconds = 300;
pool_config.stream_timeout_cooldown_seconds = 90;
let key_ids = vec!["key-4".to_string()];
record_admin_provider_pool_stream_timeout(&runner, "provider-1", "key-4", &pool_config)
.await;
record_admin_provider_pool_stream_timeout(&runner, "provider-1", "key-4", &pool_config)
.await;
record_admin_provider_pool_stream_timeout(
runtime_state,
"provider-1",
"key-4",
&pool_config,
)
.await;
record_admin_provider_pool_stream_timeout(
runtime_state,
"provider-1",
"key-4",
&pool_config,
)
.await;
let mut runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime_state,
"provider-1",
&key_ids,
&pool_config,
@@ -1186,7 +1137,7 @@ mod tests {
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
runtime = read_admin_provider_pool_runtime_state(
&runner,
runtime_state,
"provider-1",
&key_ids,
&pool_config,
@@ -139,11 +139,8 @@ pub(super) async fn build_admin_pool_list_keys_response(
let page_offset = page.saturating_sub(1).saturating_mul(page_size);
let (keys, total) = if status == "cooldown" {
let cooldown_key_ids = if let Some(runner) = state.redis_kv_runner() {
read_admin_provider_pool_cooldown_key_ids(&runner, &provider.id).await
} else {
Vec::new()
};
let cooldown_key_ids =
read_admin_provider_pool_cooldown_key_ids(state.runtime_state(), &provider.id).await;
let mut keys = if cooldown_key_ids.is_empty() {
Vec::new()
} else {
@@ -240,10 +237,10 @@ pub(super) async fn build_admin_pool_list_keys_response(
let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
let runtime = match (state.redis_kv_runner(), pool_config.as_ref()) {
(Some(runner), Some(pool_config)) if !key_ids.is_empty() => {
let runtime = match pool_config.as_ref() {
Some(pool_config) if !key_ids.is_empty() => {
read_admin_provider_pool_runtime_state(
&runner,
state.runtime_state(),
&provider.id,
&key_ids,
pool_config,
@@ -1,3 +1,5 @@
use std::collections::BTreeMap;
use super::{
admin_provider_pool_config, build_admin_pool_error_response,
read_admin_provider_pool_cooldown_counts,
@@ -12,7 +14,6 @@ use axum::{
response::{IntoResponse, Response},
Json,
};
use std::collections::BTreeMap;
pub(super) async fn build_admin_pool_overview_response(
state: &AdminAppState<'_>,
@@ -35,7 +36,6 @@ pub(super) async fn build_admin_pool_overview_response(
.iter()
.map(|(provider, _)| provider.id.clone())
.collect::<Vec<_>>();
let redis_runner = state.redis_kv_runner();
let (key_stats_result, cooldown_counts_by_provider) = tokio::join!(
async {
if provider_ids.is_empty() {
@@ -47,11 +47,10 @@ pub(super) async fn build_admin_pool_overview_response(
}
},
async {
match redis_runner.as_ref() {
Some(runner) if !provider_ids.is_empty() => {
read_admin_provider_pool_cooldown_counts(runner, &provider_ids).await
}
_ => BTreeMap::new(),
if provider_ids.is_empty() {
std::collections::BTreeMap::new()
} else {
read_admin_provider_pool_cooldown_counts(state.runtime_state(), &provider_ids).await
}
},
);
@@ -1553,20 +1553,8 @@ async fn provider_query_read_cached_models(
provider_id: &str,
key_id: &str,
) -> Option<Vec<Value>> {
let runner = state.app().redis_kv_runner()?;
let cache_key = runner
.keyspace()
.key(&format!("upstream_models:{provider_id}:{key_id}"));
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.ok()?;
let raw = redis::cmd("GET")
.arg(&cache_key)
.query_async::<Option<String>>(&mut connection)
.await
.ok()??;
let cache_key = format!("upstream_models:{provider_id}:{key_id}");
let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
Some(aggregate_models_for_cache(&parsed))
}
@@ -1575,20 +1563,8 @@ async fn provider_query_read_provider_cached_models(
state: &AdminAppState<'_>,
provider_id: &str,
) -> Option<Vec<Value>> {
let runner = state.app().redis_kv_runner()?;
let cache_key = runner.keyspace().key(&format!(
"{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}"
));
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.ok()?;
let raw = redis::cmd("GET")
.arg(&cache_key)
.query_async::<Option<String>>(&mut connection)
.await
.ok()??;
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
Some(aggregate_models_for_cache(&parsed))
}
@@ -1598,18 +1574,18 @@ async fn provider_query_write_provider_cached_models(
provider_id: &str,
models: &[Value],
) {
let Some(runner) = state.app().redis_kv_runner() else {
return;
};
let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(models)) else {
return;
};
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
let _ = runner
.setex(
let _ = state
.runtime_state()
.kv_set(
&cache_key,
&serialized,
Some(aether_model_fetch::model_fetch_interval_minutes().saturating_mul(60)),
serialized,
Some(std::time::Duration::from_secs(
aether_model_fetch::model_fetch_interval_minutes().saturating_mul(60),
)),
)
.await;
}