Add usage queue worker autoscaling

This commit is contained in:
elky
2026-06-26 14:02:57 +08:00
parent 6c5e70ccb1
commit 7e9424008f
21 changed files with 1801 additions and 85 deletions
@@ -15,7 +15,8 @@ use self::plan::{build_codex_quota_request_spec, execute_codex_quota_plan};
use super::shared::{
build_quota_snapshot_payload, extract_execution_error_message,
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
provider_auto_remove_banned_keys, quota_key_auto_removed, quota_refresh_success_invalid_state,
provider_auto_remove_banned_keys, provider_auto_remove_quota_exhausted_keys,
quota_key_auto_removed, quota_refresh_success_invalid_state,
should_auto_remove_structured_reason, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::AdminAppState;
@@ -61,6 +62,8 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
proxy_override: Option<ProxySnapshot>,
) -> Result<Option<serde_json::Value>, GatewayError> {
let auto_remove_abnormal_keys = provider_auto_remove_banned_keys(provider.config.as_ref());
let auto_remove_quota_exhausted_keys =
provider_auto_remove_quota_exhausted_keys(provider.config.as_ref());
let mut results = Vec::new();
let mut success_count = 0usize;
let mut failed_count = 0usize;
@@ -312,7 +315,7 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}));
continue;
}
let auto_removed = if auto_remove_candidate {
let auto_removed_hard_banned = if auto_remove_candidate {
state
.cleanup_provider_catalog_key_if_current(provider, &key.id, |latest_key| {
should_auto_remove_structured_reason(latest_key.oauth_invalid_reason.as_deref())
@@ -321,10 +324,28 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
} else {
false
};
if auto_removed {
if auto_removed_hard_banned {
auto_removed_count += 1;
auto_removed_hard_banned_count += 1;
}
let auto_removed_quota_exhausted =
if !auto_removed_hard_banned && auto_remove_quota_exhausted_keys {
state
.cleanup_provider_catalog_key_if_current(provider, &key.id, |latest_key| {
aether_admin::provider::pool::admin_pool_key_account_quota_exhausted(
latest_key,
provider.provider_type.as_str(),
)
})
.await?
} else {
false
};
if auto_removed_quota_exhausted {
auto_removed_count += 1;
status = "quota_exhausted".to_string();
}
let auto_removed = auto_removed_hard_banned || auto_removed_quota_exhausted;
let refresh_fixed =
status == "success" && had_oauth_refresh_issue && oauth_invalid_reason.is_none();
if refresh_fixed {
@@ -370,8 +391,13 @@ pub(crate) async fn refresh_codex_provider_quota_locally(
}
if auto_removed {
payload.insert("auto_removed".to_string(), json!(true));
}
if auto_removed_hard_banned {
payload.insert("auto_removed_hard_banned".to_string(), json!(true));
}
if auto_removed_quota_exhausted {
payload.insert("auto_removed_quota_exhausted".to_string(), json!(true));
}
if refresh_fixed {
payload.insert("refresh_fixed".to_string(), json!(true));
}
@@ -6,8 +6,8 @@ use self::plan::execute_kiro_quota_plan;
use super::shared::{
build_quota_snapshot_payload, extract_execution_error_message,
oauth_refresh_auto_removed_result, persist_provider_quota_refresh_state,
persist_quota_oauth_refresh_failure_state, quota_refresh_success_invalid_state,
ProviderQuotaExecutionOutcome,
persist_quota_oauth_refresh_failure_state, provider_auto_remove_quota_exhausted_keys,
quota_refresh_success_invalid_state, ProviderQuotaExecutionOutcome,
};
use crate::handlers::admin::request::{AdminAppState, AdminLocalOAuthRefreshError};
use crate::GatewayError;
@@ -97,6 +97,8 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
let mut success_count = 0usize;
let mut failed_count = 0usize;
let mut auto_removed_count = 0usize;
let auto_remove_quota_exhausted_keys =
provider_auto_remove_quota_exhausted_keys(provider.config.as_ref());
for key in keys {
let transport = match state
@@ -304,6 +306,23 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
continue;
}
let auto_removed_quota_exhausted = if auto_remove_quota_exhausted_keys {
state
.cleanup_provider_catalog_key_if_current(provider, &key.id, |latest_key| {
aether_admin::provider::pool::admin_pool_key_account_quota_exhausted(
latest_key,
provider.provider_type.as_str(),
)
})
.await?
} else {
false
};
if auto_removed_quota_exhausted {
auto_removed_count += 1;
status = "quota_exhausted".to_string();
}
if status == "success" {
success_count += 1;
} else {
@@ -331,6 +350,10 @@ pub(crate) async fn refresh_kiro_provider_quota_locally(
) {
payload.insert("quota_snapshot".to_string(), quota_snapshot);
}
if auto_removed_quota_exhausted {
payload.insert("auto_removed".to_string(), json!(true));
payload.insert("auto_removed_quota_exhausted".to_string(), json!(true));
}
results.push(serde_json::Value::Object(payload));
}
@@ -67,6 +67,12 @@ pub(crate) fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value
admin_provider_quota_pure::provider_auto_remove_banned_keys(config)
}
pub(crate) fn provider_auto_remove_quota_exhausted_keys(
config: Option<&serde_json::Value>,
) -> bool {
admin_provider_quota_pure::provider_auto_remove_quota_exhausted_keys(config)
}
pub(super) fn should_auto_remove_structured_reason(reason: Option<&str>) -> bool {
admin_provider_quota_pure::should_auto_remove_structured_reason(reason)
}
@@ -16,6 +16,8 @@ use std::collections::{BTreeMap, BTreeSet};
const FIXED_PROVIDER_TEMPLATE_METADATA_KEY: &str = "_aether_fixed_provider_template";
const OVERRIDE_BODY_RULES: &str = "body_rules";
const OVERRIDE_FORMAT_ACCEPTANCE_CONFIG: &str = "format_acceptance_config";
const OVERRIDE_BASE_URL: &str = "base_url";
const OVERRIDE_CUSTOM_PATH: &str = "custom_path";
const OVERRIDE_HEADER_RULES: &str = "header_rules";
const OVERRIDE_IS_ACTIVE: &str = "is_active";
const OVERRIDE_MAX_RETRIES: &str = "max_retries";
@@ -158,6 +160,20 @@ pub(crate) fn apply_admin_fixed_provider_endpoint_template_overrides(
.unwrap_or_else(|| managed_fixed_provider_endpoint_metadata(template, endpoint_template));
let mut overrides = metadata.overrides.clone();
sync_override_if_changed(
&mut overrides,
OVERRIDE_BASE_URL,
&existing_endpoint.base_url,
&updated_endpoint.base_url,
&defaults.base_url,
);
sync_override_if_changed(
&mut overrides,
OVERRIDE_CUSTOM_PATH,
&existing_endpoint.custom_path,
&updated_endpoint.custom_path,
&defaults.custom_path,
);
sync_override_if_changed(
&mut overrides,
OVERRIDE_HEADER_RULES,
@@ -249,9 +265,13 @@ fn reconcile_fixed_provider_endpoint(
updated.api_format = defaults.api_format.clone();
updated.api_family = Some(defaults.api_family.clone());
updated.endpoint_kind = Some(defaults.endpoint_kind.clone());
updated.base_url = defaults.base_url;
updated.custom_path = defaults.custom_path;
if !metadata.overrides.contains(OVERRIDE_BASE_URL) {
updated.base_url = defaults.base_url;
}
if !metadata.overrides.contains(OVERRIDE_CUSTOM_PATH) {
updated.custom_path = defaults.custom_path;
}
if !metadata.overrides.contains(OVERRIDE_HEADER_RULES) {
updated.header_rules = defaults.header_rules;
}
@@ -534,3 +554,74 @@ fn sync_override_if_changed<T>(
}
sync_override(overrides, key, actual, desired);
}
#[cfg(test)]
mod tests {
use super::{
apply_admin_fixed_provider_endpoint_template_overrides, fixed_provider_endpoint_metadata,
reconcile_fixed_provider_endpoint,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogProvider,
};
use aether_provider_transport::provider_types::fixed_provider_template;
fn sample_codex_provider() -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
"provider-codex".to_string(),
"Codex".to_string(),
None,
"codex".to_string(),
)
.expect("provider should build")
}
fn sample_codex_endpoint(base_url: &str) -> StoredProviderCatalogEndpoint {
StoredProviderCatalogEndpoint::new(
"endpoint-codex-responses".to_string(),
"provider-codex".to_string(),
"openai:responses".to_string(),
Some("openai".to_string()),
Some("responses".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
base_url.to_string(),
None,
None,
Some(2),
None,
None,
None,
None,
)
.expect("endpoint transport should build")
}
#[test]
fn fixed_provider_endpoint_reconcile_preserves_base_url_override() {
let provider = sample_codex_provider();
let template = fixed_provider_template("codex").expect("codex template should exist");
let endpoint_template = template
.endpoints
.iter()
.find(|endpoint| endpoint.api_format == "openai:responses")
.expect("responses endpoint template should exist");
let existing = sample_codex_endpoint("https://chatgpt.com/backend-api/codex");
let mut updated = existing.clone();
updated.base_url = "http://127.0.0.1:18181/v1".to_string();
apply_admin_fixed_provider_endpoint_template_overrides(&provider, &existing, &mut updated)
.expect("override metadata should apply");
let metadata = fixed_provider_endpoint_metadata(&updated)
.expect("fixed provider metadata should exist");
assert!(metadata.overrides.contains("base_url"));
let reconciled =
reconcile_fixed_provider_endpoint(&provider, &updated, template, endpoint_template)
.expect("endpoint should reconcile");
assert_eq!(reconciled.base_url, "http://127.0.0.1:18181/v1");
}
}
@@ -315,14 +315,6 @@ impl<'a> AdminAppState<'a> {
return Err("Gemini CLI Endpoint 由系统固定管理,不允许修改".to_string());
}
if self.provider_type_is_fixed(&provider.provider_type)
&& (fields.contains("base_url") || fields.contains("custom_path"))
{
return Err(
"固定类型 Provider 的 Endpoint 不允许修改 base_url/custom_path".to_string(),
);
}
let mut update_fields = admin_provider_endpoints_pure::AdminProviderEndpointUpdateFields {
base_url: payload.base_url,
custom_path: payload.custom_path,
@@ -1,4 +1,8 @@
use super::*;
use crate::ai_serving::provider_key_pool_score_scope;
use aether_data_contracts::repository::pool_scores::{
ListPoolMemberScoresQuery, PoolMemberHardState, POOL_KIND_PROVIDER_KEY_POOL,
};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
@@ -266,6 +270,100 @@ impl<'a> AdminAppState<'a> {
Ok(affected)
}
pub(crate) async fn cleanup_quota_exhausted_provider_catalog_keys(
&self,
provider: &StoredProviderCatalogProvider,
provider_type: &str,
) -> Result<usize, GatewayError> {
use aether_admin::provider::pool as admin_provider_pool_pure;
let keys = self
.list_provider_catalog_keys_by_provider_ids(std::slice::from_ref(&provider.id))
.await?;
if keys.is_empty() {
return Ok(0);
}
let known_key_ids = keys
.iter()
.map(|key| key.id.as_str())
.collect::<std::collections::BTreeSet<_>>();
let mut exhausted_key_ids = keys
.iter()
.filter(|key| {
admin_provider_pool_pure::admin_pool_key_account_quota_exhausted(key, provider_type)
})
.map(|key| key.id.clone())
.collect::<std::collections::BTreeSet<_>>();
if self.app().data.has_pool_score_reader() {
let scope = provider_key_pool_score_scope();
let page_size = 10_000usize;
let mut offset = 0usize;
loop {
let scores = self
.app()
.data
.list_pool_member_scores(&ListPoolMemberScoresQuery {
pool_kind: POOL_KIND_PROVIDER_KEY_POOL.to_string(),
pool_id: provider.id.clone(),
capability: Some(scope.capability.clone()),
scope_kind: Some(scope.scope_kind.clone()),
scope_id: scope.scope_id.clone(),
hard_states: vec![PoolMemberHardState::QuotaExhausted],
probe_statuses: None,
offset,
limit: page_size,
})
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if scores.is_empty() {
break;
}
let page_len = scores.len();
for score in scores {
if known_key_ids.contains(score.member_id.as_str()) {
exhausted_key_ids.insert(score.member_id);
}
}
if page_len < page_size {
break;
}
offset = offset.saturating_add(page_size);
}
}
let exhausted_keys = keys
.iter()
.filter(|key| exhausted_key_ids.contains(&key.id))
.collect::<Vec<_>>();
if exhausted_keys.is_empty() {
return Ok(0);
}
let deleted_key_ids = exhausted_keys
.iter()
.map(|key| key.id.clone())
.collect::<Vec<_>>();
for key in exhausted_keys {
self.clear_admin_provider_pool_cooldown(&provider.id, &key.id)
.await;
self.reset_admin_provider_pool_cost(&provider.id, &key.id)
.await;
}
let mut affected = 0usize;
for key_id in &deleted_key_ids {
if self.delete_provider_catalog_key(key_id).await? {
affected += 1;
}
}
self.cleanup_deleted_provider_catalog_refs(&provider.id, false, &[], &deleted_key_ids)
.await?;
Ok(affected)
}
pub(crate) async fn cleanup_provider_catalog_key_if_current<F>(
&self,
provider: &StoredProviderCatalogProvider,
+450 -19
View File
@@ -245,6 +245,12 @@ const AUTO_SERVER_SQL_POOL_MIN_CONNECTIONS_FLOOR: u32 = 4;
const AUTO_SERVER_SQL_POOL_MIN_CONNECTIONS_CAP: u32 = 16;
const AUTO_SERVER_SQL_POOL_MAX_CONNECTIONS_FLOOR: u32 = 20;
const AUTO_SERVER_SQL_POOL_MAX_CONNECTIONS_CAP: u32 = 100;
const DEFAULT_USAGE_QUEUE_WORKERS_CAP: usize = 8;
const AUTO_USAGE_QUEUE_WORKERS_MIN: usize = 2;
const AUTO_USAGE_QUEUE_WORKERS_REQUESTS_PER_WORKER: usize = 128;
const AUTO_USAGE_QUEUE_WORKERS_DB_SHARE_ALL: usize = 4;
const AUTO_USAGE_QUEUE_WORKERS_DB_SHARE_BACKGROUND: usize = 2;
const MAX_USAGE_QUEUE_WORKERS: usize = 64;
const DEFAULT_GATEWAY_LISTEN_BACKLOG: i32 = 65_535;
const MIN_GATEWAY_LISTEN_BACKLOG: i32 = 128;
const MAX_GATEWAY_LISTEN_BACKLOG: i32 = 65_535;
@@ -261,12 +267,99 @@ fn env_var_trimmed(name: &str) -> Option<String> {
}
fn available_parallelism_u32() -> u32 {
std::thread::available_parallelism()
.map(|value| u32::try_from(value.get()).unwrap_or(u32::MAX))
.unwrap_or(AUTO_SERVER_SQL_POOL_MIN_CONNECTIONS_FLOOR)
u32::try_from(available_parallelism_usize())
.unwrap_or(u32::MAX)
.max(1)
}
fn available_parallelism_usize() -> usize {
std::thread::available_parallelism()
.map(|value| value.get())
.unwrap_or(AUTO_SERVER_SQL_POOL_MIN_CONNECTIONS_FLOOR as usize)
.max(1)
}
fn usage_queue_request_concurrency_hint(
max_in_flight_requests: Option<usize>,
distributed_request_limit: Option<usize>,
) -> Option<usize> {
match (
max_in_flight_requests.filter(|limit| *limit > 0),
distributed_request_limit.filter(|limit| *limit > 0),
) {
(Some(local), Some(distributed)) => Some(local.min(distributed)),
(Some(local), None) => Some(local),
(None, Some(distributed)) => Some(distributed),
(None, None) => None,
}
}
fn usage_queue_workers_for_request_concurrency(request_concurrency: usize) -> usize {
let workers = request_concurrency
.saturating_add(AUTO_USAGE_QUEUE_WORKERS_REQUESTS_PER_WORKER - 1)
/ AUTO_USAGE_QUEUE_WORKERS_REQUESTS_PER_WORKER;
workers.clamp(AUTO_USAGE_QUEUE_WORKERS_MIN, MAX_USAGE_QUEUE_WORKERS)
}
fn usage_queue_worker_database_cap(
node_role: NodeRoleArg,
database: Option<&SqlDatabaseConfig>,
) -> usize {
let Some(database) = database else {
return MAX_USAGE_QUEUE_WORKERS;
};
if database.driver == DatabaseDriver::Sqlite {
return 1;
}
let divisor = if matches!(node_role, NodeRoleArg::Background) {
AUTO_USAGE_QUEUE_WORKERS_DB_SHARE_BACKGROUND
} else {
AUTO_USAGE_QUEUE_WORKERS_DB_SHARE_ALL
};
let max_connections = database.pool.max_connections.max(1) as usize;
max_connections
.saturating_add(divisor - 1)
.checked_div(divisor)
.unwrap_or(1)
.clamp(1, MAX_USAGE_QUEUE_WORKERS)
}
fn automatic_usage_queue_workers_for_parallelism(
parallelism: usize,
node_role: NodeRoleArg,
max_in_flight_requests: Option<usize>,
distributed_request_limit: Option<usize>,
database: Option<&SqlDatabaseConfig>,
) -> usize {
let cpu_default = parallelism.max(1).clamp(
AUTO_USAGE_QUEUE_WORKERS_MIN,
DEFAULT_USAGE_QUEUE_WORKERS_CAP,
);
let requested =
usage_queue_request_concurrency_hint(max_in_flight_requests, distributed_request_limit)
.map(usage_queue_workers_for_request_concurrency)
.unwrap_or(cpu_default);
requested
.min(usage_queue_worker_database_cap(node_role, database))
.clamp(1, MAX_USAGE_QUEUE_WORKERS)
}
fn automatic_usage_queue_workers(
node_role: NodeRoleArg,
max_in_flight_requests: Option<usize>,
distributed_request_limit: Option<usize>,
database: Option<&SqlDatabaseConfig>,
) -> usize {
automatic_usage_queue_workers_for_parallelism(
available_parallelism_usize(),
node_role,
max_in_flight_requests,
distributed_request_limit,
database,
)
}
fn automatic_sql_pool_config(driver: DatabaseDriver) -> SqlPoolConfig {
automatic_sql_pool_config_for_parallelism(driver, available_parallelism_u32())
}
@@ -525,6 +618,37 @@ struct GatewayUsageArgs {
)]
queue_lifecycle_events: bool,
#[arg(long, env = "AETHER_GATEWAY_USAGE_QUEUE_WORKERS", value_name = "COUNT")]
queue_workers: Option<usize>,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_QUEUE_WORKER_AUTOSCALE_ENABLED",
default_value_t = true
)]
queue_worker_autoscale_enabled: bool,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_QUEUE_WORKER_MAX_COUNT",
value_name = "COUNT"
)]
queue_worker_max_count: Option<usize>,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_QUEUE_WORKER_SCALE_INTERVAL_MS",
default_value_t = 1_000
)]
queue_worker_scale_interval_ms: u64,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_QUEUE_WORKER_IDLE_SCALE_DOWN_TICKS",
default_value_t = 30
)]
queue_worker_idle_scale_down_ticks: u64,
#[arg(
long,
env = "AETHER_GATEWAY_USAGE_QUEUE_STREAM_KEY",
@@ -639,11 +763,66 @@ struct GatewayUsageArgs {
}
impl GatewayUsageArgs {
fn to_config(&self) -> UsageRuntimeConfig {
fn effective_queue_workers(
&self,
node_role: NodeRoleArg,
max_in_flight_requests: Option<usize>,
distributed_request_limit: Option<usize>,
database: Option<&SqlDatabaseConfig>,
) -> usize {
if let Some(queue_workers) = self.queue_workers {
return queue_workers.clamp(1, MAX_USAGE_QUEUE_WORKERS);
}
if !self.queue_terminal_events && !self.queue_lifecycle_events {
return 1;
}
automatic_usage_queue_workers(
node_role,
max_in_flight_requests,
distributed_request_limit,
database,
)
}
fn effective_queue_worker_max_count(
&self,
node_role: NodeRoleArg,
database: Option<&SqlDatabaseConfig>,
worker_count: usize,
) -> usize {
if !self.queue_worker_autoscale_enabled {
return worker_count.max(1).min(MAX_USAGE_QUEUE_WORKERS);
}
self.queue_worker_max_count
.unwrap_or_else(|| usage_queue_worker_database_cap(node_role, database))
.clamp(worker_count.max(1), MAX_USAGE_QUEUE_WORKERS)
}
fn runtime_state_blocking_stream_lanes(
&self,
node_role: NodeRoleArg,
database: Option<&SqlDatabaseConfig>,
worker_max_count: usize,
) -> Option<usize> {
if !node_role.spawns_background_tasks()
|| (!self.queue_terminal_events && !self.queue_lifecycle_events)
|| database.is_none()
{
return None;
}
Some(worker_max_count.clamp(1, MAX_USAGE_QUEUE_WORKERS))
}
fn to_config(&self, worker_count: usize, worker_max_count: usize) -> UsageRuntimeConfig {
UsageRuntimeConfig {
enabled: true,
queue_terminal_events: self.queue_terminal_events,
queue_lifecycle_events: self.queue_lifecycle_events,
worker_count: worker_count.clamp(1, MAX_USAGE_QUEUE_WORKERS),
worker_autoscale_enabled: self.queue_worker_autoscale_enabled,
worker_max_count: worker_max_count.clamp(worker_count.max(1), MAX_USAGE_QUEUE_WORKERS),
worker_scale_interval_ms: self.queue_worker_scale_interval_ms.max(1),
worker_idle_scale_down_ticks: self.queue_worker_idle_scale_down_ticks.max(1),
stream_key: self.queue_stream_key.trim().to_string(),
consumer_group: self.queue_group.trim().to_string(),
dlq_stream_key: self.queue_dlq_stream_key.trim().to_string(),
@@ -1041,6 +1220,7 @@ impl Args {
&self,
runtime_backend: RuntimeBackendArg,
data_redis_url: Option<&str>,
blocking_stream_lanes: Option<usize>,
) -> RuntimeStateConfig {
let redis = self
.effective_runtime_redis_url(data_redis_url)
@@ -1052,6 +1232,7 @@ impl Args {
backend: runtime_backend.to_runtime_state_backend(),
redis,
command_timeout_ms: Some(self.runtime_command_timeout_ms.max(1)),
blocking_stream_lanes,
..RuntimeStateConfig::default()
}
}
@@ -1397,10 +1578,35 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
runtime_redis_url.as_deref(),
runtime_backend,
)?;
let usage_queue_request_concurrency_hint = usage_queue_request_concurrency_hint(
args.max_in_flight_requests,
args.distributed_request_limit,
);
let usage_queue_workers = args.usage.effective_queue_workers(
args.node_role,
args.max_in_flight_requests,
args.distributed_request_limit,
sql_database_config.as_ref(),
);
let usage_queue_worker_max_count = args.usage.effective_queue_worker_max_count(
args.node_role,
sql_database_config.as_ref(),
usage_queue_workers,
);
let usage_config = args
.usage
.to_config(usage_queue_workers, usage_queue_worker_max_count);
let usage_blocking_stream_lanes = args.usage.runtime_state_blocking_stream_lanes(
args.node_role,
sql_database_config.as_ref(),
usage_config.worker_max_count,
);
let runtime_state = Arc::new(
RuntimeState::from_config(
args.runtime_state_config(runtime_backend, data_redis_url.as_deref()),
)
RuntimeState::from_config(args.runtime_state_config(
runtime_backend,
data_redis_url.as_deref(),
usage_blocking_stream_lanes,
))
.await
.map_err(|err| std::io::Error::new(std::io::ErrorKind::InvalidInput, err.to_string()))?,
);
@@ -1425,6 +1631,16 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
deployment_topology = args.deployment_topology.as_str(),
node_role = args.node_role.as_str(),
runtime_backend = runtime_backend.as_str(),
usage_queue_workers = usage_config.worker_count,
usage_queue_worker_autoscale_enabled = usage_config.worker_autoscale_enabled,
usage_queue_worker_max_count = usage_config.worker_max_count,
usage_queue_request_concurrency_hint =
usage_queue_request_concurrency_hint.unwrap_or_default(),
usage_queue_request_concurrency_hint_source = if usage_queue_request_concurrency_hint.is_some() {
"explicit"
} else {
"none"
},
frontdoor_mode = "compatibility_frontdoor",
log_format = ?args.logging.log_format,
log_destination = args.logging.log_destination.as_str(),
@@ -1448,6 +1664,22 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
video_task_poller_interval_ms = args.video_task_poller_interval_ms,
video_task_poller_batch_size = args.video_task_poller_batch_size,
video_task_store_path = args.video_task_store_path.as_deref().unwrap_or("-"),
usage_queue_workers = usage_config.worker_count,
usage_queue_workers_source = if args.usage.queue_workers.is_some() {
"explicit"
} else {
"auto"
},
usage_queue_worker_autoscale_enabled = usage_config.worker_autoscale_enabled,
usage_queue_worker_max_count = usage_config.worker_max_count,
usage_queue_request_concurrency_hint =
usage_queue_request_concurrency_hint.unwrap_or_default(),
usage_queue_request_concurrency_hint_source =
if usage_queue_request_concurrency_hint.is_some() {
"explicit"
} else {
"none"
},
max_in_flight_requests = args.max_in_flight_requests.unwrap_or_default(),
distributed_request_limit = args.distributed_request_limit.unwrap_or_default(),
distributed_request_redis_configured = args
@@ -1479,7 +1711,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
let mut state = AppState::new()?
.with_runtime_state(runtime_state)
.with_data_config(data_config)?
.with_usage_runtime_config(args.usage.to_config())?
.with_usage_runtime_config(usage_config)?
.with_video_task_truth_source_mode(args.video_task_truth_source_mode.into());
if let Some(cors_config) = args.frontdoor.cors_config() {
state = state.with_frontdoor_cors_config(cors_config);
@@ -1977,9 +2209,7 @@ fn pending_schema_error(
) -> std::io::Error {
std::io::Error::other(format!(
"database schema is behind by {} migration(s); next pending migration is {} ({})\nrun `aether-gateway --migrate` before starting the service",
pending_count,
next_version,
next_description
pending_count, next_version, next_description
))
}
@@ -1990,9 +2220,7 @@ fn pending_backfills_error(
) -> std::io::Error {
std::io::Error::other(format!(
"database backfills are behind by {} backfill(s); next pending backfill is {} ({})\nrun `aether-gateway --apply-backfills` before starting the service",
pending_count,
next_version,
next_description
pending_count, next_version, next_description
))
}
@@ -2000,11 +2228,11 @@ fn pending_backfills_error(
mod tests {
use super::{
automatic_sql_pool_config, automatic_sql_pool_config_for_parallelism,
ensure_database_backfills_are_current, ensure_database_schema_is_current,
pending_backfills_error, pending_schema_error, resolve_healthcheck_url, Args,
DatabaseDriverArg, DeploymentTopologyArg, GatewayDataArgs, GatewayFrontdoorArgs,
GatewayLogDestinationArg, GatewayLogFormatArg, GatewayLogRotationArg, GatewayLoggingArgs,
GatewayRateLimitArgs, GatewayUsageArgs, NodeRoleArg, RuntimeBackendArg,
automatic_usage_queue_workers_for_parallelism, ensure_database_backfills_are_current,
ensure_database_schema_is_current, pending_backfills_error, pending_schema_error,
resolve_healthcheck_url, Args, DatabaseDriverArg, DeploymentTopologyArg, GatewayDataArgs,
GatewayFrontdoorArgs, GatewayLogDestinationArg, GatewayLogFormatArg, GatewayLogRotationArg,
GatewayLoggingArgs, GatewayRateLimitArgs, GatewayUsageArgs, NodeRoleArg, RuntimeBackendArg,
VideoTaskTruthSourceArg, DEFAULT_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS,
DEFAULT_GATEWAY_LISTENER_SHARDS, DEFAULT_GATEWAY_LISTEN_BACKLOG,
MAX_GATEWAY_HTTP2_MAX_CONCURRENT_STREAMS, MAX_GATEWAY_LISTENER_SHARDS,
@@ -2062,6 +2290,11 @@ mod tests {
usage: GatewayUsageArgs {
queue_terminal_events: true,
queue_lifecycle_events: true,
queue_workers: Some(4),
queue_worker_autoscale_enabled: true,
queue_worker_max_count: None,
queue_worker_scale_interval_ms: 1_000,
queue_worker_idle_scale_down_ticks: 30,
queue_stream_key: "usage:events".to_string(),
queue_group: "usage_consumers".to_string(),
queue_dlq_stream_key: "usage:events:dlq".to_string(),
@@ -2100,6 +2333,25 @@ mod tests {
}
}
fn test_database(driver: DatabaseDriver, max_connections: u32) -> SqlDatabaseConfig {
let url = match driver {
DatabaseDriver::Sqlite => "sqlite://./data/aether.db",
DatabaseDriver::Mysql => "mysql://root:root@localhost/aether",
DatabaseDriver::Postgres => "postgres://postgres:postgres@localhost/aether",
};
let max_connections = max_connections.max(1);
SqlDatabaseConfig::new(
driver,
url,
SqlPoolConfig {
min_connections: 1,
max_connections,
..SqlPoolConfig::default()
},
)
.expect("test database config should build")
}
#[test]
fn resolves_healthcheck_url_from_app_port() {
assert_eq!(
@@ -2260,6 +2512,183 @@ mod tests {
assert_eq!(many_cpu.max_connections, 100);
}
#[test]
fn gateway_usage_queue_workers_manual_override_wins_and_is_capped() {
let mut args = test_args();
args.usage.queue_workers = Some(72);
let database = test_database(DatabaseDriver::Postgres, 100);
let workers = args.usage.effective_queue_workers(
NodeRoleArg::All,
Some(10_000),
None,
Some(&database),
);
assert_eq!(workers, 64);
assert_eq!(args.usage.to_config(workers, 64).worker_count, 64);
}
#[test]
fn gateway_usage_queue_workers_auto_uses_cpu_default_without_concurrency_hint() {
let database = test_database(DatabaseDriver::Postgres, 100);
let workers = automatic_usage_queue_workers_for_parallelism(
4,
NodeRoleArg::All,
None,
None,
Some(&database),
);
assert_eq!(workers, 4);
}
#[test]
fn gateway_usage_queue_worker_autoscale_max_uses_database_cap() {
let mut args = test_args();
args.usage.queue_workers = None;
let database = test_database(DatabaseDriver::Postgres, 40);
let workers =
args.usage
.effective_queue_workers(args.node_role, None, None, Some(&database));
let max_workers =
args.usage
.effective_queue_worker_max_count(args.node_role, Some(&database), workers);
assert_eq!(workers, 8);
assert_eq!(max_workers, 10);
}
#[test]
fn gateway_usage_queue_worker_autoscale_max_respects_explicit_override() {
let mut args = test_args();
args.usage.queue_workers = None;
args.usage.queue_worker_max_count = Some(32);
let database = test_database(DatabaseDriver::Postgres, 100);
let workers =
args.usage
.effective_queue_workers(args.node_role, None, None, Some(&database));
let max_workers =
args.usage
.effective_queue_worker_max_count(args.node_role, Some(&database), workers);
assert_eq!(workers, 8);
assert_eq!(max_workers, 32);
}
#[test]
fn gateway_usage_queue_blocking_stream_lanes_only_expand_when_worker_can_spawn() {
let database = test_database(DatabaseDriver::Postgres, 100);
let args = test_args();
assert_eq!(
args.usage
.runtime_state_blocking_stream_lanes(NodeRoleArg::All, Some(&database), 10,),
Some(10)
);
assert_eq!(
args.usage.runtime_state_blocking_stream_lanes(
NodeRoleArg::Frontdoor,
Some(&database),
10,
),
None
);
assert_eq!(
args.usage
.runtime_state_blocking_stream_lanes(NodeRoleArg::All, None, 10),
None
);
let mut disabled_queue_args = args;
disabled_queue_args.usage.queue_terminal_events = false;
disabled_queue_args.usage.queue_lifecycle_events = false;
assert_eq!(
disabled_queue_args
.usage
.runtime_state_blocking_stream_lanes(NodeRoleArg::All, Some(&database), 10,),
None
);
}
#[test]
fn gateway_usage_queue_workers_auto_scales_from_request_concurrency() {
let database = test_database(DatabaseDriver::Postgres, 100);
let workers = automatic_usage_queue_workers_for_parallelism(
8,
NodeRoleArg::All,
Some(1_536),
None,
Some(&database),
);
assert_eq!(workers, 12);
}
#[test]
fn gateway_usage_queue_workers_auto_respects_effective_request_limit() {
let database = test_database(DatabaseDriver::Postgres, 100);
let workers = automatic_usage_queue_workers_for_parallelism(
8,
NodeRoleArg::All,
Some(2_048),
Some(256),
Some(&database),
);
assert_eq!(workers, 2);
}
#[test]
fn gateway_usage_queue_workers_auto_is_capped_by_database_pool() {
let database = test_database(DatabaseDriver::Postgres, 20);
let workers = automatic_usage_queue_workers_for_parallelism(
16,
NodeRoleArg::All,
Some(5_000),
None,
Some(&database),
);
assert_eq!(workers, 5);
}
#[test]
fn gateway_usage_queue_workers_auto_gives_background_nodes_more_pool_budget() {
let database = test_database(DatabaseDriver::Postgres, 20);
let workers = automatic_usage_queue_workers_for_parallelism(
16,
NodeRoleArg::Background,
Some(5_000),
None,
Some(&database),
);
assert_eq!(workers, 10);
}
#[test]
fn gateway_usage_queue_workers_auto_uses_single_worker_for_sqlite() {
let database = test_database(DatabaseDriver::Sqlite, 1);
let workers = automatic_usage_queue_workers_for_parallelism(
16,
NodeRoleArg::All,
Some(5_000),
None,
Some(&database),
);
assert_eq!(workers, 1);
}
#[test]
fn gateway_data_pool_explicit_values_override_auto_sizing() {
let mut args = test_args();
@@ -2355,8 +2784,10 @@ mod tests {
let config = args.runtime_state_config(
RuntimeBackendArg::Redis,
args.data.effective_redis_url().as_deref(),
Some(7),
);
assert_eq!(config.blocking_stream_lanes, Some(7));
assert_eq!(
config
.redis
@@ -1255,6 +1255,23 @@ async fn perform_pool_quota_probe_for_provider(
);
}
}
if aether_admin::provider::quota::provider_auto_remove_quota_exhausted_keys(
provider.config.as_ref(),
) {
let auto_removed = admin_state
.cleanup_quota_exhausted_provider_catalog_keys(provider, provider_type)
.await?;
if auto_removed > 0 {
summary.auto_removed += auto_removed;
info!(
event_name = "auto_removed_quota_exhausted",
provider_id = %provider_short_id,
provider_type,
auto_removed,
"gateway pool quota probe auto-cleaned quota-exhausted provider keys"
);
}
}
let Some(endpoint) = endpoint_for_probe_with_reconcile(
state,
+39 -4
View File
@@ -1357,6 +1357,15 @@ impl AppState {
definition.trigger,
));
};
if let Some(handle) = self
.usage_runtime
.spawn_worker_supervisor(self.data.clone())
{
supervisor.supervise_handle(crate::task_runtime::TASK_KEY_USAGE_QUEUE_WORKER, handle);
record_boot(crate::task_runtime::TASK_KEY_USAGE_QUEUE_WORKER);
}
let mut supervise_worker =
|task_key: &'static str, handle: Option<tokio::task::JoinHandle<()>>| {
if let Some(handle) = handle {
@@ -1365,10 +1374,6 @@ impl AppState {
}
};
supervise_worker(
crate::task_runtime::TASK_KEY_USAGE_QUEUE_WORKER,
self.usage_runtime.spawn_worker(self.data.clone()),
);
supervise_worker(
crate::task_runtime::TASK_KEY_USAGE_COUNTER_FLUSH,
spawn_usage_counter_flush_worker(self.data.clone()),
@@ -1551,6 +1556,36 @@ fn usage_runtime_metric_samples(
MetricKind::Gauge,
u64::from(snapshot.queue_lifecycle_events),
),
MetricSample::new(
"usage_runtime_queue_worker_count",
"Minimum configured number of usage queue worker consumers.",
MetricKind::Gauge,
snapshot.worker_count as u64,
),
MetricSample::new(
"usage_runtime_queue_worker_autoscale_enabled",
"Whether usage queue worker autoscaling is enabled.",
MetricKind::Gauge,
u64::from(snapshot.worker_autoscale_enabled),
),
MetricSample::new(
"usage_runtime_queue_worker_max_count",
"Maximum configured number of elastic usage queue worker consumers.",
MetricKind::Gauge,
snapshot.worker_max_count as u64,
),
MetricSample::new(
"usage_runtime_queue_worker_active_count",
"Current active usage queue worker consumers managed by the supervisor.",
MetricKind::Gauge,
snapshot.worker_active_count as u64,
),
MetricSample::new(
"usage_runtime_queue_worker_desired_count",
"Current desired usage queue worker consumers selected by autoscaling.",
MetricKind::Gauge,
snapshot.worker_desired_count as u64,
),
MetricSample::new(
"usage_runtime_retry_deferred_lifecycle_events_enabled",
"Whether deferred lifecycle usage events are scheduled for local enqueue retry.",
@@ -667,6 +667,71 @@ async fn gateway_updates_admin_provider_endpoint_locally_with_trusted_admin_prin
upstream_handle.abort();
}
#[tokio::test]
async fn gateway_updates_fixed_provider_endpoint_base_url_as_template_override() {
let mut provider = sample_provider("provider-codex", "codex", 10);
provider.provider_type = "codex".to_string();
let provider_catalog_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider],
vec![sample_endpoint(
"endpoint-codex-responses",
"provider-codex",
"openai:responses",
"https://chatgpt.com/backend-api/codex",
)],
vec![],
));
let gateway = build_router_with_state(
AppState::new()
.expect("gateway should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(
provider_catalog_repository.clone(),
),
),
);
let (gateway_url, gateway_handle) = start_server(gateway).await;
let response = reqwest::Client::new()
.put(format!(
"{gateway_url}/api/admin/endpoints/endpoint-codex-responses"
))
.header(crate::constants::GATEWAY_HEADER, "rust-phase3b")
.header(TRUSTED_ADMIN_USER_ID_HEADER, "admin-user-123")
.header(TRUSTED_ADMIN_USER_ROLE_HEADER, "admin")
.header(TRUSTED_ADMIN_SESSION_ID_HEADER, "session-123")
.json(&json!({
"base_url": "http://127.0.0.1:18181/v1",
"max_retries": 0
}))
.send()
.await
.expect("request should succeed");
let status = response.status();
let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(status, StatusCode::OK, "payload={payload}");
assert_eq!(payload["base_url"], "http://127.0.0.1:18181/v1");
let endpoints = provider_catalog_repository
.list_endpoints_by_ids(&["endpoint-codex-responses".to_string()])
.await
.expect("endpoints should read");
assert_eq!(endpoints.len(), 1);
assert_eq!(endpoints[0].base_url, "http://127.0.0.1:18181/v1");
assert_eq!(endpoints[0].max_retries, Some(0));
assert!(endpoints[0]
.config
.as_ref()
.and_then(|value| value.get("_aether_fixed_provider_template"))
.and_then(|value| value.get("overrides"))
.and_then(serde_json::Value::as_array)
.is_some_and(|items| items.iter().any(|item| item.as_str() == Some("base_url"))));
gateway_handle.abort();
}
#[tokio::test]
async fn gateway_deletes_admin_provider_endpoint_locally_with_trusted_admin_principal() {
let upstream_hits = Arc::new(Mutex::new(0usize));
+29 -3
View File
@@ -20,6 +20,15 @@ pub fn provider_auto_remove_banned_keys(config: Option<&serde_json::Value>) -> b
.unwrap_or(false)
}
pub fn provider_auto_remove_quota_exhausted_keys(config: Option<&serde_json::Value>) -> bool {
config
.and_then(|value| value.get("pool_advanced"))
.and_then(serde_json::Value::as_object)
.and_then(|object| object.get("auto_remove_quota_exhausted_keys"))
.and_then(serde_json::Value::as_bool)
.unwrap_or(false)
}
pub fn should_auto_remove_structured_reason(reason: Option<&str>) -> bool {
provider_status::should_auto_remove_account_state(&provider_status::resolve_pool_account_state(
None, None, reason,
@@ -1767,13 +1776,30 @@ mod tests {
parse_codex_wham_usage_response, parse_gemini_cli_retrieve_user_quota_response,
parse_gemini_cli_v1internal_credits_response, parse_windsurf_model_configs_response,
parse_windsurf_rate_limit_response, parse_windsurf_user_status_response,
quota_refresh_success_invalid_state, should_auto_remove_structured_reason,
OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX, OAUTH_REFRESH_FAILED_PREFIX,
OAUTH_REQUEST_FAILED_PREFIX,
provider_auto_remove_quota_exhausted_keys, quota_refresh_success_invalid_state,
should_auto_remove_structured_reason, OAUTH_ACCOUNT_BLOCK_PREFIX, OAUTH_EXPIRED_PREFIX,
OAUTH_REFRESH_FAILED_PREFIX, OAUTH_REQUEST_FAILED_PREFIX,
};
use aether_data_contracts::repository::provider_catalog::StoredProviderCatalogKey;
use serde_json::json;
#[test]
fn provider_auto_remove_quota_exhausted_keys_defaults_to_false() {
assert!(!provider_auto_remove_quota_exhausted_keys(None));
assert!(!provider_auto_remove_quota_exhausted_keys(Some(&json!({
"pool_advanced": {}
}))));
}
#[test]
fn provider_auto_remove_quota_exhausted_keys_reads_pool_advanced_flag() {
assert!(provider_auto_remove_quota_exhausted_keys(Some(&json!({
"pool_advanced": {
"auto_remove_quota_exhausted_keys": true
}
}))));
}
#[test]
fn codex_runtime_invalid_reason_marks_401_as_expired() {
assert_eq!(
+69 -2
View File
@@ -62,6 +62,7 @@ pub struct RuntimeStateConfig {
pub redis: Option<RedisClientConfig>,
pub memory: MemoryRuntimeStateConfig,
pub command_timeout_ms: Option<u64>,
pub blocking_stream_lanes: Option<usize>,
}
impl Default for RuntimeStateConfig {
@@ -71,6 +72,7 @@ impl Default for RuntimeStateConfig {
redis: None,
memory: MemoryRuntimeStateConfig::default(),
command_timeout_ms: Some(DEFAULT_COMMAND_TIMEOUT_MS),
blocking_stream_lanes: None,
}
}
}
@@ -138,6 +140,11 @@ impl RuntimeStateConfig {
"runtime state command_timeout_ms must be positive".to_string(),
));
}
if matches!(self.blocking_stream_lanes, Some(0)) {
return Err(DataLayerError::InvalidConfiguration(
"runtime state blocking_stream_lanes must be positive".to_string(),
));
}
Ok(())
}
}
@@ -194,7 +201,12 @@ impl RuntimeState {
let redis = config.redis.clone().ok_or_else(|| {
DataLayerError::InvalidConfiguration("runtime redis config missing".to_string())
})?;
Self::redis(redis, config.command_timeout_ms).await
Self::redis_with_blocking_stream_lanes(
redis,
config.command_timeout_ms,
config.blocking_stream_lanes,
)
.await
}
RuntimeStateBackendMode::Auto => unreachable!("auto resolved above"),
}
@@ -211,10 +223,20 @@ impl RuntimeState {
pub async fn redis(
config: RedisClientConfig,
command_timeout_ms: Option<u64>,
) -> Result<Self, DataLayerError> {
Self::redis_with_blocking_stream_lanes(config, command_timeout_ms, None).await
}
pub async fn redis_with_blocking_stream_lanes(
config: RedisClientConfig,
command_timeout_ms: Option<u64>,
blocking_stream_lanes: Option<usize>,
) -> Result<Self, DataLayerError> {
let factory = redis::RedisClientFactory::new(config)?;
let keyspace = factory.config().keyspace();
let connections = factory.connect_router(command_timeout_ms).await?;
let connections = factory
.connect_router_with_blocking_stream_lanes(command_timeout_ms, blocking_stream_lanes)
.await?;
let runtime = redis::RedisRuntimeRunner::new(
connections.clone(),
keyspace.clone(),
@@ -1597,6 +1619,51 @@ mod tests {
let _ = blocking.await.expect("blocking task join");
}
#[tokio::test]
async fn redis_concurrent_blocking_stream_reads_do_not_share_single_connection() {
let Some(redis) = TestRedisServer::start().await else {
return;
};
let runtime = RuntimeState::redis(
RedisClientConfig {
url: redis.redis_url.clone(),
key_prefix: Some(format!("aether-block-pool-test-{}", std::process::id())),
},
Some(1_000),
)
.await
.expect("runtime should connect");
RuntimeQueueStore::ensure_consumer_group(&runtime, "blocking-stream", "workers", "0-0")
.await
.expect("consumer group");
let mut handles = Vec::new();
for index in 0..4 {
let blocking_runtime = runtime.clone();
handles.push(tokio::spawn(async move {
let consumer = format!("consumer-{index}");
RuntimeQueueStore::read_group(
&blocking_runtime,
"blocking-stream",
"workers",
&consumer,
1,
Some(600),
)
.await
}));
}
for handle in handles {
let result = handle.await.expect("blocking task join");
assert!(
!matches!(result, Err(DataLayerError::TimedOut(_))),
"concurrent blocking stream reads should not queue behind one connection"
);
assert!(result.expect("blocking read should succeed").is_empty());
}
}
#[tokio::test]
async fn redis_connection_manager_recovers_after_restart() {
let Some(mut redis) = TestRedisServer::start().await else {
+109 -13
View File
@@ -1,7 +1,7 @@
use crate::error::RedisResultExt;
use crate::redis::RedisKeyspace;
use crate::DataLayerError;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tracing::info;
@@ -9,6 +9,10 @@ use tracing::info;
pub(crate) type RedisClient = redis::Client;
pub(crate) type RedisManagedConnection = redis::aio::ConnectionManager;
const DEFAULT_BLOCKING_STREAM_LANES_FALLBACK: usize = 4;
const DEFAULT_BLOCKING_STREAM_LANES_CAP: usize = 16;
const MAX_BLOCKING_STREAM_LANES_CAP: usize = 64;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
pub struct RedisClientConfig {
pub url: String,
@@ -57,7 +61,21 @@ impl RedisClientFactory {
&self,
command_timeout_ms: Option<u64>,
) -> Result<RedisConnectionRouter, DataLayerError> {
RedisConnectionRouter::connect(self.connect_lazy()?, command_timeout_ms).await
self.connect_router_with_blocking_stream_lanes(command_timeout_ms, None)
.await
}
pub(crate) async fn connect_router_with_blocking_stream_lanes(
&self,
command_timeout_ms: Option<u64>,
blocking_stream_lanes: Option<usize>,
) -> Result<RedisConnectionRouter, DataLayerError> {
RedisConnectionRouter::connect(
self.connect_lazy()?,
command_timeout_ms,
blocking_stream_lanes,
)
.await
}
}
@@ -84,7 +102,8 @@ impl RedisConnectionLane {
pub(crate) struct RedisConnectionRouter {
fast: RedisManagedConnection,
stream: RedisManagedConnection,
blocking_stream: RedisManagedConnection,
blocking_stream: Arc<Vec<RedisManagedConnection>>,
blocking_stream_next: Arc<AtomicUsize>,
admin: RedisManagedConnection,
metrics: Arc<RedisConnectionMetrics>,
}
@@ -93,6 +112,7 @@ impl std::fmt::Debug for RedisConnectionRouter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RedisConnectionRouter")
.field("lanes", &["fast", "stream", "blocking_stream", "admin"])
.field("blocking_stream_lanes", &self.blocking_stream.len())
.finish()
}
}
@@ -101,6 +121,7 @@ impl RedisConnectionRouter {
pub(crate) async fn connect(
client: RedisClient,
command_timeout_ms: Option<u64>,
blocking_stream_lanes: Option<usize>,
) -> Result<Self, DataLayerError> {
let fast = connect_lane(
&client,
@@ -116,13 +137,9 @@ impl RedisConnectionRouter {
command_timeout_ms,
)
.await?;
let blocking_stream = connect_lane(
&client,
connection_manager_config(command_timeout_ms),
RedisConnectionLane::BlockingStream,
command_timeout_ms,
)
.await?;
let blocking_stream =
connect_blocking_stream_lanes(&client, command_timeout_ms, blocking_stream_lanes)
.await?;
let admin = connect_lane(
&client,
connection_manager_config(command_timeout_ms),
@@ -130,14 +147,17 @@ impl RedisConnectionRouter {
command_timeout_ms,
)
.await?;
let blocking_stream_lanes = blocking_stream.len();
info!(
redis_lanes = "fast,stream,blocking_stream,admin",
redis_blocking_stream_lanes = blocking_stream_lanes,
"runtime redis connection lanes initialized"
);
Ok(Self {
fast,
stream,
blocking_stream,
blocking_stream: Arc::new(blocking_stream),
blocking_stream_next: Arc::new(AtomicUsize::new(0)),
admin,
metrics: Arc::new(RedisConnectionMetrics::default()),
})
@@ -147,7 +167,11 @@ impl RedisConnectionRouter {
match lane {
RedisConnectionLane::Fast => self.fast.clone(),
RedisConnectionLane::Stream => self.stream.clone(),
RedisConnectionLane::BlockingStream => self.blocking_stream.clone(),
RedisConnectionLane::BlockingStream => {
let index = self.blocking_stream_next.fetch_add(1, Ordering::Relaxed)
% self.blocking_stream.len();
self.blocking_stream[index].clone()
}
RedisConnectionLane::Admin => self.admin.clone(),
}
}
@@ -228,6 +252,51 @@ fn connection_manager_config(
config
}
async fn connect_blocking_stream_lanes(
client: &RedisClient,
command_timeout_ms: Option<u64>,
requested_lanes: Option<usize>,
) -> Result<Vec<RedisManagedConnection>, DataLayerError> {
let lane_count = blocking_stream_lane_count(requested_lanes)?;
let mut lanes = Vec::with_capacity(lane_count);
for _ in 0..lane_count {
lanes.push(
connect_lane(
client,
connection_manager_config(command_timeout_ms),
RedisConnectionLane::BlockingStream,
command_timeout_ms,
)
.await?,
);
}
Ok(lanes)
}
fn blocking_stream_lane_count(requested_lanes: Option<usize>) -> Result<usize, DataLayerError> {
if matches!(requested_lanes, Some(0)) {
return Err(DataLayerError::InvalidConfiguration(
"runtime redis blocking_stream_lanes must be positive".to_string(),
));
}
let default_lanes = default_blocking_stream_lane_count();
Ok(requested_lanes
.map(|lanes| lanes.max(default_lanes))
.unwrap_or(default_lanes)
.clamp(1, MAX_BLOCKING_STREAM_LANES_CAP))
}
fn default_blocking_stream_lane_count() -> usize {
std::thread::available_parallelism()
.map(|value| value.get())
.unwrap_or(DEFAULT_BLOCKING_STREAM_LANES_FALLBACK)
.clamp(
DEFAULT_BLOCKING_STREAM_LANES_FALLBACK,
DEFAULT_BLOCKING_STREAM_LANES_CAP,
)
}
async fn connect_lane(
client: &RedisClient,
config: redis::aio::ConnectionManagerConfig,
@@ -259,7 +328,10 @@ async fn connect_lane(
#[cfg(test)]
mod tests {
use super::{RedisClientConfig, RedisClientFactory};
use super::{
blocking_stream_lane_count, default_blocking_stream_lane_count, RedisClientConfig,
RedisClientFactory, MAX_BLOCKING_STREAM_LANES_CAP,
};
#[test]
fn factory_builds_lazy_client_from_valid_config() {
@@ -274,4 +346,28 @@ mod tests {
.connect_lazy()
.expect("lazy redis client should build");
}
#[test]
fn blocking_stream_lane_count_uses_requested_as_floor() {
let default_lanes = default_blocking_stream_lane_count();
assert_eq!(
blocking_stream_lane_count(None).expect("default lanes"),
default_lanes
);
assert_eq!(
blocking_stream_lane_count(Some(1)).expect("requested below default"),
default_lanes
);
assert_eq!(
blocking_stream_lane_count(Some(default_lanes + 1)).expect("requested above default"),
default_lanes + 1
);
assert_eq!(
blocking_stream_lane_count(Some(MAX_BLOCKING_STREAM_LANES_CAP + 1))
.expect("requested above cap"),
MAX_BLOCKING_STREAM_LANES_CAP
);
assert!(blocking_stream_lane_count(Some(0)).is_err());
}
}
+38 -21
View File
@@ -289,27 +289,11 @@ impl RedisStreamRunner {
command.arg("STREAMS").arg(&stream.0).arg(">");
let reply = command
.query_async::<StreamReadReply>(&mut connection)
.query_async::<RedisValue>(&mut connection)
.await
.map_redis_err()?;
Ok(reply
.keys
.into_iter()
.flat_map(|key| key.ids.into_iter())
.map(|id| RedisStreamEntry {
id: id.id,
fields: id
.map
.into_iter()
.filter_map(|(field, value)| {
redis::from_redis_value::<String>(&value)
.ok()
.map(|text| (field, text))
})
.collect(),
})
.collect())
parse_stream_read_entries(reply)
})
.await
}
@@ -455,6 +439,31 @@ fn validate_stream_position(position: &str) -> Result<(), DataLayerError> {
Ok(())
}
fn parse_stream_read_entries(value: RedisValue) -> Result<Vec<RedisStreamEntry>, DataLayerError> {
if matches!(value, RedisValue::Nil) {
return Ok(Vec::new());
}
let reply = from_redis_value::<StreamReadReply>(&value).map_err(redis_error)?;
Ok(reply
.keys
.into_iter()
.flat_map(|key| key.ids.into_iter())
.map(|id| RedisStreamEntry {
id: id.id,
fields: id
.map
.into_iter()
.filter_map(|(field, value)| {
redis::from_redis_value::<String>(&value)
.ok()
.map(|text| (field, text))
})
.collect(),
})
.collect())
}
fn parse_reclaim_result(value: RedisValue) -> Result<RedisStreamReclaimResult, DataLayerError> {
let RedisValue::Array(parts) = value else {
return Err(DataLayerError::UnexpectedValue(
@@ -573,9 +582,9 @@ mod tests {
use std::collections::BTreeMap;
use super::{
parse_reclaim_result, validate_consumer, validate_group, validate_stream_name,
validate_stream_position, RedisConsumerName, RedisStreamName, RedisStreamReclaimConfig,
RedisStreamReclaimResult, RedisStreamRunnerConfig,
parse_reclaim_result, parse_stream_read_entries, validate_consumer, validate_group,
validate_stream_name, validate_stream_position, RedisConsumerName, RedisStreamName,
RedisStreamReclaimConfig, RedisStreamReclaimResult, RedisStreamRunnerConfig,
};
use redis::Value as RedisValue;
@@ -639,6 +648,14 @@ mod tests {
assert!(validate_stream_position("").is_err());
}
#[test]
fn parses_empty_blocking_read_as_no_entries() {
let parsed = parse_stream_read_entries(RedisValue::Nil)
.expect("nil stream read reply should be empty");
assert!(parsed.is_empty());
}
#[test]
fn parses_reclaim_result_with_deleted_ids() {
let parsed = parse_reclaim_result(RedisValue::Array(vec![
+30
View File
@@ -5,6 +5,11 @@ pub struct UsageRuntimeConfig {
pub enabled: bool,
pub queue_terminal_events: bool,
pub queue_lifecycle_events: bool,
pub worker_count: usize,
pub worker_autoscale_enabled: bool,
pub worker_max_count: usize,
pub worker_scale_interval_ms: u64,
pub worker_idle_scale_down_ticks: u64,
pub stream_key: String,
pub consumer_group: String,
pub dlq_stream_key: String,
@@ -29,6 +34,11 @@ impl Default for UsageRuntimeConfig {
enabled: false,
queue_terminal_events: false,
queue_lifecycle_events: false,
worker_count: 4,
worker_autoscale_enabled: true,
worker_max_count: 64,
worker_scale_interval_ms: 1_000,
worker_idle_scale_down_ticks: 30,
stream_key: "usage:events".to_string(),
consumer_group: "usage_consumers".to_string(),
dlq_stream_key: "usage:events:dlq".to_string(),
@@ -74,6 +84,26 @@ impl UsageRuntimeConfig {
"usage runtime dlq_stream_key cannot be empty".to_string(),
));
}
if self.worker_count == 0 {
return Err(DataLayerError::InvalidConfiguration(
"usage runtime worker_count must be positive".to_string(),
));
}
if self.worker_max_count == 0 {
return Err(DataLayerError::InvalidConfiguration(
"usage runtime worker_max_count must be positive".to_string(),
));
}
if self.worker_scale_interval_ms == 0 {
return Err(DataLayerError::InvalidConfiguration(
"usage runtime worker_scale_interval_ms must be positive".to_string(),
));
}
if self.worker_idle_scale_down_ticks == 0 {
return Err(DataLayerError::InvalidConfiguration(
"usage runtime worker_idle_scale_down_ticks must be positive".to_string(),
));
}
if self.stream_maxlen == 0 {
return Err(DataLayerError::InvalidConfiguration(
"usage runtime stream_maxlen must be positive".to_string(),
+597 -3
View File
@@ -1,6 +1,7 @@
use std::collections::BTreeMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
@@ -9,9 +10,10 @@ use aether_data_contracts::DataLayerError;
use aether_runtime_state::RuntimeQueueStore;
use async_trait::async_trait;
use tokio::sync::mpsc;
use tracing::warn;
use tracing::{info, warn};
use crate::executor::spawn_on_usage_background_runtime;
use crate::worker::{UsageWorkerControl, UsageWorkerObservation};
use crate::{
apply_usage_body_capture_policy_to_event, build_stream_terminal_usage_seed,
build_sync_terminal_usage_seed, build_terminal_usage_event_from_seed,
@@ -80,15 +82,27 @@ pub struct UsageRuntime {
config: UsageRuntimeConfig,
body_policy_cache: Arc<tokio::sync::Mutex<Option<UsageBodyCapturePolicyCacheEntry>>>,
enqueue_retry: Arc<UsageEnqueueRetryDispatcher>,
worker_supervisor_state: Arc<UsageWorkerSupervisorState>,
terminal_enqueue_state: Arc<LifecycleEnqueueState>,
lifecycle_enqueue_state: Arc<LifecycleEnqueueState>,
}
#[derive(Debug, Default)]
struct UsageWorkerSupervisorState {
active_count: AtomicUsize,
desired_count: AtomicUsize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct UsageRuntimeMetricsSnapshot {
pub enabled: bool,
pub queue_terminal_events: bool,
pub queue_lifecycle_events: bool,
pub worker_count: usize,
pub worker_autoscale_enabled: bool,
pub worker_max_count: usize,
pub worker_active_count: usize,
pub worker_desired_count: usize,
pub retry_deferred_lifecycle_events: bool,
pub terminal_enqueue_in_flight: u64,
pub terminal_enqueue_deferred_total: u64,
@@ -118,6 +132,7 @@ impl UsageRuntime {
config: UsageRuntimeConfig::disabled(),
body_policy_cache: Arc::new(tokio::sync::Mutex::new(None)),
enqueue_retry: UsageEnqueueRetryDispatcher::disabled(),
worker_supervisor_state: Arc::new(UsageWorkerSupervisorState::default()),
terminal_enqueue_state: Arc::new(LifecycleEnqueueState::default()),
lifecycle_enqueue_state: Arc::new(LifecycleEnqueueState::default()),
}
@@ -130,6 +145,7 @@ impl UsageRuntime {
config,
body_policy_cache: Arc::new(tokio::sync::Mutex::new(None)),
enqueue_retry,
worker_supervisor_state: Arc::new(UsageWorkerSupervisorState::default()),
terminal_enqueue_state: Arc::new(LifecycleEnqueueState::default()),
lifecycle_enqueue_state: Arc::new(LifecycleEnqueueState::default()),
})
@@ -144,6 +160,17 @@ impl UsageRuntime {
enabled: self.config.enabled,
queue_terminal_events: self.config.queue_terminal_events,
queue_lifecycle_events: self.config.queue_lifecycle_events,
worker_count: self.config.worker_count,
worker_autoscale_enabled: self.config.worker_autoscale_enabled,
worker_max_count: self.config.worker_max_count,
worker_active_count: self
.worker_supervisor_state
.active_count
.load(Ordering::Acquire),
worker_desired_count: self
.worker_supervisor_state
.desired_count
.load(Ordering::Acquire),
retry_deferred_lifecycle_events: self.config.retry_deferred_lifecycle_events,
terminal_enqueue_in_flight: self.terminal_enqueue_state.in_flight(),
terminal_enqueue_deferred_total: self.terminal_enqueue_state.deferred_total(),
@@ -182,10 +209,61 @@ impl UsageRuntime {
return None;
}
let runner = data.usage_worker_queue()?;
let worker = build_usage_queue_worker(runner, data, self.config.clone()).ok()?;
let worker = build_usage_queue_worker(runner, data, self.config.clone(), None).ok()?;
Some(worker.spawn())
}
pub fn spawn_workers<T>(&self, data: Arc<T>) -> Vec<tokio::task::JoinHandle<()>>
where
T: UsageRuntimeAccess + 'static,
{
if !self.can_spawn_worker(data.as_ref()) {
return Vec::new();
}
let Some(runner) = data.usage_worker_queue() else {
return Vec::new();
};
let worker_count = self.config.worker_count.max(1);
let mut handles = Vec::with_capacity(worker_count);
for worker_index in 0..worker_count {
let Ok(worker) = build_usage_queue_worker(
Arc::clone(&runner),
Arc::clone(&data),
self.config.clone(),
Some(worker_index),
) else {
warn!(
event_name = "usage_worker_build_failed",
log_type = "ops",
worker_index,
"usage runtime failed to build usage queue worker"
);
continue;
};
handles.push(worker.spawn());
}
handles
}
pub fn spawn_worker_supervisor<T>(&self, data: Arc<T>) -> Option<tokio::task::JoinHandle<()>>
where
T: UsageRuntimeAccess + 'static,
{
if !self.can_spawn_worker(data.as_ref()) {
return None;
}
let runner = data.usage_worker_queue()?;
Some(spawn_on_usage_background_runtime(
run_usage_worker_supervisor(
runner,
data,
self.config.clone(),
Arc::clone(&self.worker_supervisor_state),
),
))
}
pub fn record_pending<T>(&self, data: &T, seed: LifecycleUsageSeed)
where
T: UsageRuntimeAccess + Clone + 'static,
@@ -785,6 +863,258 @@ where
}
}
struct ManagedUsageWorker {
control: UsageWorkerControl,
stopping: bool,
}
async fn run_usage_worker_supervisor<T>(
runner: Arc<dyn RuntimeQueueStore>,
data: Arc<T>,
config: UsageRuntimeConfig,
state: Arc<UsageWorkerSupervisorState>,
) where
T: UsageRuntimeAccess + 'static,
{
let min_workers = config.worker_count.max(1);
let max_workers = if config.worker_autoscale_enabled {
config.worker_max_count.max(min_workers)
} else {
min_workers
};
let mut desired_workers = min_workers;
let mut next_worker_index = 0usize;
let mut workers = BTreeMap::<usize, ManagedUsageWorker>::new();
let mut worker_task_indexes = BTreeMap::<tokio::task::Id, usize>::new();
let mut join_set = tokio::task::JoinSet::<usize>::new();
let (telemetry_tx, mut telemetry_rx) =
mpsc::channel::<UsageWorkerObservation>(max_workers.saturating_mul(4).clamp(16, 1024));
let mut scale_interval = tokio::time::interval(Duration::from_millis(
config.worker_scale_interval_ms.max(1),
));
scale_interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
let mut full_reads = 0usize;
let mut busy_reads = 0usize;
let mut idle_reads = 0usize;
let mut idle_ticks = 0u64;
state
.desired_count
.store(desired_workers, Ordering::Release);
reconcile_usage_workers(
&runner,
&data,
&config,
&telemetry_tx,
&mut join_set,
&mut worker_task_indexes,
&mut workers,
&mut next_worker_index,
desired_workers,
);
state.active_count.store(workers.len(), Ordering::Release);
loop {
tokio::select! {
Some(observation) = telemetry_rx.recv() => {
if observation.entries_read == 0 {
idle_reads = idle_reads.saturating_add(1);
} else {
busy_reads = busy_reads.saturating_add(1);
if observation.entries_read >= observation.batch_size {
full_reads = full_reads.saturating_add(1);
}
}
}
_ = scale_interval.tick() => {
drain_finished_usage_workers(
&mut join_set,
&mut worker_task_indexes,
&mut workers,
);
if config.worker_autoscale_enabled {
let active_workers = workers.len().max(1);
let high_pressure = full_reads > 0
|| (busy_reads >= active_workers.saturating_mul(2) && idle_reads == 0);
if high_pressure && desired_workers < max_workers {
let grow_by = (active_workers + 1) / 2;
let next = desired_workers
.saturating_add(grow_by.max(1))
.clamp(min_workers, max_workers);
if next > desired_workers {
info!(
event_name = "usage_worker_autoscale_up",
log_type = "ops",
desired_workers = next,
previous_desired_workers = desired_workers,
active_workers = workers.len(),
max_workers,
full_reads,
busy_reads,
idle_reads,
"usage worker supervisor scaled up"
);
desired_workers = next;
idle_ticks = 0;
}
} else if desired_workers > min_workers
&& busy_reads == 0
&& full_reads == 0
&& idle_reads >= active_workers
{
idle_ticks = idle_ticks.saturating_add(1);
if idle_ticks >= config.worker_idle_scale_down_ticks {
let next = desired_workers
.saturating_sub((desired_workers + 1) / 2)
.max(min_workers);
if next < desired_workers {
info!(
event_name = "usage_worker_autoscale_down",
log_type = "ops",
desired_workers = next,
previous_desired_workers = desired_workers,
active_workers = workers.len(),
min_workers,
idle_ticks,
"usage worker supervisor scaled down"
);
desired_workers = next;
}
idle_ticks = 0;
}
} else if busy_reads > 0 || full_reads > 0 {
idle_ticks = 0;
}
}
state.desired_count.store(desired_workers, Ordering::Release);
reconcile_usage_workers(
&runner,
&data,
&config,
&telemetry_tx,
&mut join_set,
&mut worker_task_indexes,
&mut workers,
&mut next_worker_index,
desired_workers,
);
state
.active_count
.store(workers.len(), Ordering::Release);
full_reads = 0;
busy_reads = 0;
idle_reads = 0;
}
}
}
}
fn reconcile_usage_workers<T>(
runner: &Arc<dyn RuntimeQueueStore>,
data: &Arc<T>,
config: &UsageRuntimeConfig,
telemetry_tx: &mpsc::Sender<UsageWorkerObservation>,
join_set: &mut tokio::task::JoinSet<usize>,
worker_task_indexes: &mut BTreeMap<tokio::task::Id, usize>,
workers: &mut BTreeMap<usize, ManagedUsageWorker>,
next_worker_index: &mut usize,
desired_workers: usize,
) where
T: UsageRuntimeAccess + 'static,
{
while workers.len() < desired_workers {
let worker_index = *next_worker_index;
*next_worker_index = (*next_worker_index).saturating_add(1);
let control = UsageWorkerControl::default();
let Ok(worker) = build_usage_queue_worker(
Arc::clone(runner),
Arc::clone(data),
config.clone(),
Some(worker_index),
) else {
warn!(
event_name = "usage_worker_build_failed",
log_type = "ops",
worker_index,
"usage runtime failed to build elastic usage queue worker"
);
break;
};
let worker = worker.with_supervisor(control.clone(), telemetry_tx.clone());
let handle = join_set.spawn(async move {
worker.run().await;
worker_index
});
worker_task_indexes.insert(handle.id(), worker_index);
workers.insert(
worker_index,
ManagedUsageWorker {
control,
stopping: false,
},
);
}
let mut excess = workers.len().saturating_sub(desired_workers);
for worker in workers.values_mut().rev() {
if excess == 0 {
break;
}
if worker.stopping {
continue;
}
worker.control.request_shutdown();
worker.stopping = true;
excess -= 1;
}
}
fn drain_finished_usage_workers(
join_set: &mut tokio::task::JoinSet<usize>,
worker_task_indexes: &mut BTreeMap<tokio::task::Id, usize>,
workers: &mut BTreeMap<usize, ManagedUsageWorker>,
) {
while let Some(result) = join_set.try_join_next_with_id() {
match result {
Ok((task_id, worker_index)) => {
worker_task_indexes.remove(&task_id);
let stopping = workers
.remove(&worker_index)
.is_some_and(|worker| worker.stopping);
if !stopping {
warn!(
event_name = "usage_worker_unexpected_exit",
log_type = "ops",
worker_index,
"usage worker exited before supervisor requested shutdown"
);
}
}
Err(err) => {
let worker_index = worker_task_indexes.remove(&err.id());
if let Some(worker_index) = worker_index {
workers.remove(&worker_index);
warn!(
event_name = "usage_worker_join_failed",
log_type = "ops",
worker_index,
error = %err,
"usage worker task failed"
);
continue;
}
warn!(
event_name = "usage_worker_join_failed",
log_type = "ops",
error = %err,
"usage worker task failed"
);
}
}
}
}
#[derive(Debug, Clone, Copy)]
struct UsageBodyCapturePolicyCacheEntry {
cached_at: Instant,
@@ -1301,6 +1631,11 @@ mod tests {
queue: Arc<dyn RuntimeQueueStore>,
}
struct PanicOnceQueueConfiguredUsageStore {
inner: CloneQueueConfiguredUsageStore,
remaining_panics: AtomicUsize,
}
struct EnrichmentCountingQueueStore {
records: Mutex<Vec<UpsertUsageRecord>>,
queue: Arc<dyn RuntimeQueueStore>,
@@ -1495,6 +1830,73 @@ mod tests {
}
}
#[async_trait]
impl UsageRecordWriter for PanicOnceQueueConfiguredUsageStore {
async fn upsert_usage_record(
&self,
record: UpsertUsageRecord,
) -> Result<Option<StoredRequestUsageAudit>, DataLayerError> {
let should_panic = self
.remaining_panics
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |current| {
(current > 0).then(|| current - 1)
})
.is_ok();
if should_panic {
panic!("forced usage writer panic");
}
self.inner.upsert_usage_record(record).await
}
}
#[async_trait]
impl UsageSettlementWriter for PanicOnceQueueConfiguredUsageStore {
fn has_usage_settlement_writer(&self) -> bool {
false
}
async fn settle_usage(
&self,
_input: UsageSettlementInput,
) -> Result<Option<StoredUsageSettlement>, DataLayerError> {
Ok(None)
}
}
#[async_trait]
impl UsageBillingEventEnricher for PanicOnceQueueConfiguredUsageStore {
async fn enrich_usage_event(&self, _event: &mut UsageEvent) -> Result<(), DataLayerError> {
Ok(())
}
}
#[async_trait]
impl ManualProxyNodeCounter for PanicOnceQueueConfiguredUsageStore {
async fn increment_manual_proxy_node_requests(
&self,
_node_id: &str,
_total_delta: i64,
_failed_delta: i64,
_latency_ms: Option<i64>,
) -> Result<(), DataLayerError> {
Ok(())
}
}
impl UsageRuntimeAccess for PanicOnceQueueConfiguredUsageStore {
fn has_usage_writer(&self) -> bool {
true
}
fn has_usage_worker_queue(&self) -> bool {
true
}
fn usage_worker_queue(&self) -> Option<Arc<dyn RuntimeQueueStore>> {
Some(Arc::clone(&self.inner.queue))
}
}
#[async_trait]
impl UsageRecordWriter for EnrichmentCountingQueueStore {
async fn upsert_usage_record(
@@ -1831,6 +2233,198 @@ mod tests {
panic!("pending lifecycle usage event was not enqueued");
}
#[test]
fn spawn_workers_uses_configured_worker_count() {
let config = UsageRuntimeConfig {
enabled: true,
queue_terminal_events: true,
worker_count: 3,
stream_key: "usage:events:test:worker-count".to_string(),
consumer_group: "usage_consumers_test_worker_count".to_string(),
..UsageRuntimeConfig::default()
};
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
let store = CloneQueueConfiguredUsageStore {
records: Arc::new(Mutex::new(Vec::new())),
queue: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())),
};
let handles = runtime.spawn_workers(Arc::new(store));
assert_eq!(handles.len(), 3);
for handle in handles {
handle.abort();
}
}
#[tokio::test]
async fn worker_supervisor_scales_up_when_reads_stay_full() {
let config = UsageRuntimeConfig {
enabled: true,
queue_terminal_events: true,
worker_count: 1,
worker_autoscale_enabled: true,
worker_max_count: 4,
worker_scale_interval_ms: 10,
worker_idle_scale_down_ticks: 100,
stream_key: "usage:events:test:worker-autoscale-up".to_string(),
consumer_group: "usage_consumers_test_worker_autoscale_up".to_string(),
consumer_batch_size: 1,
consumer_block_ms: 1,
..UsageRuntimeConfig::default()
};
let queue_runner: Arc<dyn RuntimeQueueStore> =
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let queue = UsageQueue::new(Arc::clone(&queue_runner), config.clone())
.expect("usage queue should build");
queue
.ensure_consumer_group()
.await
.expect("consumer group should initialize");
for index in 0..32 {
queue
.enqueue(&UsageEvent::new(
UsageEventType::Completed,
format!("req-worker-autoscale-up-{index}"),
UsageEventData {
provider_name: "openai".to_string(),
model: "gpt-5".to_string(),
total_tokens: Some(12),
status_code: Some(200),
..UsageEventData::default()
},
))
.await
.expect("usage event should enqueue");
}
let store = CloneQueueConfiguredUsageStore {
records: Arc::new(Mutex::new(Vec::new())),
queue: queue_runner,
};
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
let supervisor = runtime
.spawn_worker_supervisor(Arc::new(store))
.expect("supervisor should spawn");
for _ in 0..100 {
let snapshot = runtime.metrics_snapshot();
if snapshot.worker_desired_count > 1 {
supervisor.abort();
return;
}
sleep(Duration::from_millis(10)).await;
}
supervisor.abort();
let snapshot = runtime.metrics_snapshot();
assert!(
snapshot.worker_desired_count > 1,
"usage worker supervisor should scale up after repeated full reads: {snapshot:?}"
);
}
#[tokio::test]
async fn worker_supervisor_replaces_worker_after_panic() {
let config = UsageRuntimeConfig {
enabled: true,
queue_terminal_events: true,
worker_count: 1,
worker_autoscale_enabled: false,
worker_max_count: 1,
worker_scale_interval_ms: 10,
stream_key: "usage:events:test:worker-panic-recovery".to_string(),
consumer_group: "usage_consumers_test_worker_panic_recovery".to_string(),
consumer_batch_size: 1,
consumer_block_ms: 1,
reclaim_idle_ms: 60_000,
reclaim_interval_ms: 60_000,
..UsageRuntimeConfig::default()
};
let queue_runner: Arc<dyn RuntimeQueueStore> =
Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let queue = UsageQueue::new(Arc::clone(&queue_runner), config.clone())
.expect("usage queue should build");
queue
.ensure_consumer_group()
.await
.expect("consumer group should initialize");
queue
.enqueue(&UsageEvent::new(
UsageEventType::Completed,
"req-worker-panic-first",
UsageEventData {
provider_name: "openai".to_string(),
model: "gpt-5".to_string(),
total_tokens: Some(12),
status_code: Some(200),
..UsageEventData::default()
},
))
.await
.expect("first usage event should enqueue");
let records = Arc::new(Mutex::new(Vec::new()));
let store = Arc::new(PanicOnceQueueConfiguredUsageStore {
inner: CloneQueueConfiguredUsageStore {
records: Arc::clone(&records),
queue: queue_runner,
},
remaining_panics: AtomicUsize::new(1),
});
let runtime = UsageRuntime::new(config).expect("usage runtime should build");
let supervisor = runtime
.spawn_worker_supervisor(Arc::clone(&store))
.expect("supervisor should spawn");
for _ in 0..100 {
if store.remaining_panics.load(Ordering::Acquire) == 0 {
break;
}
sleep(Duration::from_millis(10)).await;
}
assert_eq!(
store.remaining_panics.load(Ordering::Acquire),
0,
"first worker should panic while processing the first event"
);
queue
.enqueue(&UsageEvent::new(
UsageEventType::Completed,
"req-worker-panic-second",
UsageEventData {
provider_name: "openai".to_string(),
model: "gpt-5".to_string(),
total_tokens: Some(24),
status_code: Some(200),
..UsageEventData::default()
},
))
.await
.expect("second usage event should enqueue");
for _ in 0..100 {
let recorded_second = records
.lock()
.expect("records lock")
.iter()
.any(|record| record.request_id == "req-worker-panic-second");
if recorded_second {
supervisor.abort();
return;
}
sleep(Duration::from_millis(10)).await;
}
supervisor.abort();
let records = records.lock().expect("records lock");
assert!(
records
.iter()
.any(|record| record.request_id == "req-worker-panic-second"),
"replacement worker should consume events after the first worker panics: {records:?}"
);
}
#[tokio::test]
async fn lifecycle_queue_append_failure_does_not_write_directly() {
let config = UsageRuntimeConfig {
+86 -5
View File
@@ -1,3 +1,4 @@
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::time::Duration;
@@ -5,6 +6,7 @@ use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UpsertUs
use aether_data_contracts::DataLayerError;
use aether_runtime_state::{RuntimeQueueEntry, RuntimeQueueStore};
use async_trait::async_trait;
use tokio::sync::mpsc;
use tracing::warn;
use crate::executor::spawn_on_usage_background_runtime;
@@ -69,29 +71,72 @@ pub struct UsageQueueWorker {
queue: UsageQueue,
recorder: Arc<dyn UsageEventRecorder>,
consumer: String,
worker_index: Option<usize>,
control: Option<UsageWorkerControl>,
telemetry: Option<mpsc::Sender<UsageWorkerObservation>>,
config: UsageRuntimeConfig,
}
#[derive(Clone, Default)]
pub(crate) struct UsageWorkerControl {
shutdown: Arc<AtomicBool>,
}
impl UsageWorkerControl {
pub(crate) fn request_shutdown(&self) {
self.shutdown.store(true, Ordering::Release);
}
fn should_shutdown(&self) -> bool {
self.shutdown.load(Ordering::Acquire)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct UsageWorkerObservation {
pub worker_index: Option<usize>,
pub entries_read: usize,
pub batch_size: usize,
}
impl UsageQueueWorker {
pub fn new(
runner: Arc<dyn RuntimeQueueStore>,
recorder: Arc<dyn UsageEventRecorder>,
config: UsageRuntimeConfig,
worker_index: Option<usize>,
) -> Result<Self, DataLayerError> {
let queue = UsageQueue::new(runner, config.clone())?;
let consumer = consumer_name();
let consumer = consumer_name(worker_index);
Ok(Self {
queue,
recorder,
consumer,
worker_index,
control: None,
telemetry: None,
config,
})
}
pub(crate) fn with_supervisor(
mut self,
control: UsageWorkerControl,
telemetry: mpsc::Sender<UsageWorkerObservation>,
) -> Self {
self.control = Some(control);
self.telemetry = Some(telemetry);
self
}
pub fn spawn(self) -> tokio::task::JoinHandle<()> {
spawn_on_usage_background_runtime(async move { self.run_forever().await })
}
pub(crate) async fn run(self) {
self.run_forever().await;
}
async fn run_forever(self) {
if let Err(err) = self.queue.ensure_consumer_group().await {
warn!(
@@ -111,6 +156,9 @@ impl UsageQueueWorker {
reclaim_interval.tick().await;
loop {
if self.should_shutdown() {
break;
}
tokio::select! {
_ = reclaim_interval.tick() => {
match self.queue.claim_stale(&self.consumer, "0-0").await {
@@ -139,6 +187,10 @@ impl UsageQueueWorker {
result = self.queue.read_group(&self.consumer) => {
match result {
Ok(entries) => {
self.report_read(entries.len());
if entries.is_empty() && self.should_shutdown() {
break;
}
if let Err(err) = self.process_entries(entries).await {
warn!(
event_name = "usage_worker_process_failed",
@@ -150,6 +202,9 @@ impl UsageQueueWorker {
);
tokio::time::sleep(Duration::from_millis(250)).await;
}
if self.should_shutdown() {
break;
}
}
Err(err) => {
warn!(
@@ -168,6 +223,23 @@ impl UsageQueueWorker {
}
}
fn should_shutdown(&self) -> bool {
self.control
.as_ref()
.is_some_and(UsageWorkerControl::should_shutdown)
}
fn report_read(&self, entries_read: usize) {
let Some(telemetry) = &self.telemetry else {
return;
};
let _ = telemetry.try_send(UsageWorkerObservation {
worker_index: self.worker_index,
entries_read,
batch_size: self.config.consumer_batch_size.max(1),
});
}
async fn process_entries(&self, entries: Vec<RuntimeQueueEntry>) -> Result<(), DataLayerError> {
if entries.is_empty() {
return Ok(());
@@ -282,6 +354,7 @@ pub fn build_usage_queue_worker<T>(
runner: Arc<dyn RuntimeQueueStore>,
data: Arc<T>,
config: UsageRuntimeConfig,
worker_index: Option<usize>,
) -> Result<UsageQueueWorker, DataLayerError>
where
T: UsageRecordWriter
@@ -292,7 +365,12 @@ where
+ Sync
+ 'static,
{
UsageQueueWorker::new(runner, Arc::new(UsageDataEventRecorder::new(data)), config)
UsageQueueWorker::new(
runner,
Arc::new(UsageDataEventRecorder::new(data)),
config,
worker_index,
)
}
pub async fn write_event_record<T>(data: &T, event: &UsageEvent) -> Result<(), DataLayerError>
@@ -380,13 +458,16 @@ fn extract_manual_proxy_node_id(event: &UsageEvent) -> Option<String> {
.map(String::from)
}
fn consumer_name() -> String {
fn consumer_name(worker_index: Option<usize>) -> String {
let host = std::env::var("HOSTNAME")
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.unwrap_or_else(|| "aether-gateway".to_string());
format!("{host}:{}", std::process::id())
match worker_index {
Some(worker_index) => format!("{host}:{}:{worker_index}", std::process::id()),
None => format!("{host}:{}", std::process::id()),
}
}
#[cfg(test)]
@@ -662,7 +743,7 @@ mod tests {
consumer_block_ms: 1,
..UsageRuntimeConfig::default()
};
let worker = UsageQueueWorker::new(queue_runner, recorder.clone(), config)
let worker = UsageQueueWorker::new(queue_runner, recorder.clone(), config, None)
.expect("worker should build");
worker
.queue
@@ -808,6 +808,7 @@ export interface PoolAdvancedConfig {
account_self_check_interval_minutes?: number | null
account_self_check_concurrency?: number | null
auto_remove_banned_keys?: boolean
auto_remove_quota_exhausted_keys?: boolean
}
function isPlainObject(value: unknown): value is Record<string, unknown> {
@@ -588,6 +588,7 @@ const form = ref({
account_self_check_interval_minutes: null as number | null | undefined,
account_self_check_concurrency: null as number | null | undefined,
auto_remove_banned_keys: false,
auto_remove_quota_exhausted_keys: false,
skip_exhausted_accounts: false,
})
@@ -625,6 +626,8 @@ function getHealthToggleValue(key: PoolHealthToggleKey): boolean {
return form.value.account_self_check_enabled
case 'auto_remove_banned_keys':
return form.value.auto_remove_banned_keys
case 'auto_remove_quota_exhausted_keys':
return form.value.auto_remove_quota_exhausted_keys
case 'skip_exhausted_accounts':
return form.value.skip_exhausted_accounts
}
@@ -641,6 +644,9 @@ function updateHealthToggleValue(key: PoolHealthToggleKey, value: boolean): void
case 'auto_remove_banned_keys':
form.value.auto_remove_banned_keys = value
return
case 'auto_remove_quota_exhausted_keys':
form.value.auto_remove_quota_exhausted_keys = value
return
case 'skip_exhausted_accounts':
form.value.skip_exhausted_accounts = value
}
@@ -675,6 +681,7 @@ watch(() => props.modelValue, (open) => {
account_self_check_interval_minutes: cfg?.account_self_check_interval_minutes ?? null,
account_self_check_concurrency: cfg?.account_self_check_concurrency ?? null,
auto_remove_banned_keys: cfg?.auto_remove_banned_keys ?? false,
auto_remove_quota_exhausted_keys: cfg?.auto_remove_quota_exhausted_keys ?? false,
skip_exhausted_accounts: cfg?.skip_exhausted_accounts ?? false,
}
@@ -751,6 +758,7 @@ async function handleSave() {
? (form.value.account_self_check_concurrency ?? undefined)
: undefined,
auto_remove_banned_keys: form.value.auto_remove_banned_keys,
auto_remove_quota_exhausted_keys: form.value.auto_remove_quota_exhausted_keys,
skip_exhausted_accounts: form.value.skip_exhausted_accounts,
}
@@ -11,6 +11,7 @@ describe('poolAdvancedDialog', () => {
'probing_enabled',
'account_self_check_enabled',
'auto_remove_banned_keys',
'auto_remove_quota_exhausted_keys',
'skip_exhausted_accounts',
])
})
@@ -32,6 +33,11 @@ describe('poolAdvancedDialog', () => {
label: '异常自动清除',
description: '检测到不可恢复账号异常,或 RT 与 AT 均失效时自动从号池移除。',
},
{
key: 'auto_remove_quota_exhausted_keys',
label: '自动清理额度耗尽',
description: '探测到黑色“额度耗尽”账号后自动从号池移除。',
},
{
key: 'skip_exhausted_accounts',
label: '跳过额度耗尽账号',
@@ -2,6 +2,7 @@ export type PoolHealthToggleKey =
| 'probing_enabled'
| 'account_self_check_enabled'
| 'auto_remove_banned_keys'
| 'auto_remove_quota_exhausted_keys'
| 'skip_exhausted_accounts'
export interface PoolHealthToggleCard {
@@ -32,6 +33,11 @@ export function buildPoolHealthToggleCards(): PoolHealthToggleCard[] {
label: '异常自动清除',
description: '检测到不可恢复账号异常,或 RT 与 AT 均失效时自动从号池移除。',
},
{
key: 'auto_remove_quota_exhausted_keys',
label: '自动清理额度耗尽',
description: '探测到黑色“额度耗尽”账号后自动从号池移除。',
},
{
key: 'skip_exhausted_accounts',
label: '跳过额度耗尽账号',