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
Generated
+22 -4
View File
@@ -126,7 +126,6 @@ dependencies = [
"chrono-tz", "chrono-tz",
"flate2", "flate2",
"futures-util", "futures-util",
"redis",
"serde", "serde",
"serde_json", "serde_json",
"sha2", "sha2",
@@ -178,6 +177,7 @@ dependencies = [
"aether-oauth", "aether-oauth",
"aether-provider-transport", "aether-provider-transport",
"aether-runtime", "aether-runtime",
"aether-runtime-state",
"aether-scheduler-core", "aether-scheduler-core",
"aether-testkit", "aether-testkit",
"aether-usage-runtime", "aether-usage-runtime",
@@ -199,7 +199,6 @@ dependencies = [
"http", "http",
"ldap3", "ldap3",
"parking_lot", "parking_lot",
"redis",
"regex", "regex",
"reqwest", "reqwest",
"rustls 0.23.37", "rustls 0.23.37",
@@ -273,9 +272,9 @@ dependencies = [
"aether-ai-formats", "aether-ai-formats",
"aether-contracts", "aether-contracts",
"aether-crypto", "aether-crypto",
"aether-data",
"aether-data-contracts", "aether-data-contracts",
"aether-oauth", "aether-oauth",
"aether-runtime-state",
"aether-video-tasks-core", "aether-video-tasks-core",
"async-trait", "async-trait",
"axum", "axum",
@@ -300,6 +299,7 @@ dependencies = [
"aether-gateway", "aether-gateway",
"aether-http", "aether-http",
"aether-runtime", "aether-runtime",
"aether-runtime-state",
"anyhow", "anyhow",
"arc-swap", "arc-swap",
"axum", "axum",
@@ -342,13 +342,29 @@ dependencies = [
"axum", "axum",
"chrono", "chrono",
"futures-util", "futures-util",
"redis",
"serde_json", "serde_json",
"sha2", "sha2",
"thiserror 2.0.18", "thiserror 2.0.18",
"tokio", "tokio",
"tracing", "tracing",
"tracing-subscriber", "tracing-subscriber",
"uuid",
]
[[package]]
name = "aether-runtime-state"
version = "0.1.0"
dependencies = [
"aether-cache",
"aether-data-contracts",
"aether-runtime",
"async-trait",
"redis",
"serde",
"serde_json",
"thiserror 2.0.18",
"tokio",
"tracing",
"url", "url",
"uuid", "uuid",
] ]
@@ -376,6 +392,7 @@ dependencies = [
"aether-gateway", "aether-gateway",
"aether-http", "aether-http",
"aether-runtime", "aether-runtime",
"aether-runtime-state",
"async-stream", "async-stream",
"axum", "axum",
"bytes", "bytes",
@@ -397,6 +414,7 @@ dependencies = [
"aether-contracts", "aether-contracts",
"aether-data", "aether-data",
"aether-data-contracts", "aether-data-contracts",
"aether-runtime-state",
"async-trait", "async-trait",
"base64 0.22.1", "base64 0.22.1",
"serde", "serde",
+2
View File
@@ -16,6 +16,7 @@ members = [
"crates/aether-oauth", "crates/aether-oauth",
"crates/aether-provider-transport", "crates/aether-provider-transport",
"crates/aether-scheduler-core", "crates/aether-scheduler-core",
"crates/aether-runtime-state",
"crates/aether-usage-runtime", "crates/aether-usage-runtime",
"crates/aether-video-tasks-core", "crates/aether-video-tasks-core",
"apps/aether-gateway", "apps/aether-gateway",
@@ -46,6 +47,7 @@ aether-model-fetch = { path = "crates/aether-model-fetch" }
aether-oauth = { path = "crates/aether-oauth" } aether-oauth = { path = "crates/aether-oauth" }
aether-provider-transport = { path = "crates/aether-provider-transport" } aether-provider-transport = { path = "crates/aether-provider-transport" }
aether-scheduler-core = { path = "crates/aether-scheduler-core" } aether-scheduler-core = { path = "crates/aether-scheduler-core" }
aether-runtime-state = { path = "crates/aether-runtime-state" }
aether-usage-runtime = { path = "crates/aether-usage-runtime" } aether-usage-runtime = { path = "crates/aether-usage-runtime" }
aether-video-tasks-core = { path = "crates/aether-video-tasks-core" } aether-video-tasks-core = { path = "crates/aether-video-tasks-core" }
aether-gateway = { path = "apps/aether-gateway" } aether-gateway = { path = "apps/aether-gateway" }
+2 -2
View File
@@ -205,8 +205,8 @@ Aether Proxy 是配套的正向代理节点,部署在海外 VPS 上,为墙
- `APP_PORT`:`aether-gateway` 唯一监听端口,固定绑定 `0.0.0.0:${APP_PORT}` - `APP_PORT`:`aether-gateway` 唯一监听端口,固定绑定 `0.0.0.0:${APP_PORT}`
- `AETHER_DATABASE_DRIVER` / `AETHER_DATABASE_URL`:二进制单机部署可用 `sqlite`,例如 `sqlite:///opt/aether/data/aether.db` - `AETHER_DATABASE_DRIVER` / `AETHER_DATABASE_URL`:二进制单机部署可用 `sqlite`,例如 `sqlite:///opt/aether/data/aether.db`
- `DATABASE_URL` / `REDIS_URL`:`aether-gateway` 直接读取的共享后端连接串;多节点必须配置共享数据库和 Redis - `DATABASE_URL` / `REDIS_URL`:共享后端连接串;多节点必须配置共享数据库和 Redis
- `AETHER_RUNTIME_BACKEND=memory|redis`:单机 SQLite 默认用 `memory`;多节点必须用 Redis - `AETHER_RUNTIME_BACKEND=memory|redis`:运行时缓存/协调后端。单机 SQLite 默认用 `memory`,不会连接 Redis;显式设为 `redis` 或多节点部署才会把 `REDIS_URL` 注入运行时 Redis 后端
- `AETHER_GATEWAY_AUTO_PREPARE_DATABASE`:常规启动前自动执行挂起的 schema migration 和 backfill;仓库自带的 `docker-compose.yml` 和 `docker-compose.build.yml` 默认开启 - `AETHER_GATEWAY_AUTO_PREPARE_DATABASE`:常规启动前自动执行挂起的 schema migration 和 backfill;仓库自带的 `docker-compose.yml` 和 `docker-compose.build.yml` 默认开启
- `JWT_SECRET_KEY` / `ENCRYPTION_KEY`:认证和敏感数据加密所需密钥 - `JWT_SECRET_KEY` / `ENCRYPTION_KEY`:认证和敏感数据加密所需密钥
- `API_KEY_PREFIX`:用户和管理员新建 API Key 时使用的前缀,默认 `sk` - `API_KEY_PREFIX`:用户和管理员新建 API Key 时使用的前缀,默认 `sk`
+1 -1
View File
@@ -22,6 +22,7 @@ aether-oauth.workspace = true
aether-provider-transport.workspace = true aether-provider-transport.workspace = true
aether-scheduler-core.workspace = true aether-scheduler-core.workspace = true
aether-runtime.workspace = true aether-runtime.workspace = true
aether-runtime-state.workspace = true
aether-usage-runtime.workspace = true aether-usage-runtime.workspace = true
aether-video-tasks-core.workspace = true aether-video-tasks-core.workspace = true
aether-wallet.workspace = true aether-wallet.workspace = true
@@ -42,7 +43,6 @@ http.workspace = true
ldap3 = { version = "0.11", default-features = false, features = ["sync", "tls-rustls"] } ldap3 = { version = "0.11", default-features = false, features = ["sync", "tls-rustls"] }
parking_lot = "0.12" parking_lot = "0.12"
regex.workspace = true regex.workspace = true
redis.workspace = true
reqwest.workspace = true reqwest.workspace = true
rustls.workspace = true rustls.workspace = true
serde.workspace = true serde.workspace = true
@@ -4,10 +4,8 @@ use clap::Parser;
use tracing::info; use tracing::info;
use aether_gateway::{serve_execution_runtime_tcp, serve_execution_runtime_unix}; use aether_gateway::{serve_execution_runtime_tcp, serve_execution_runtime_unix};
use aether_runtime::{ use aether_runtime::{init_service_runtime, ServiceRuntimeConfig};
init_service_runtime, DistributedConcurrencyGate, RedisDistributedConcurrencyConfig, use aether_runtime_state::{RedisClientConfig, RuntimeSemaphoreConfig, RuntimeState};
ServiceRuntimeConfig,
};
#[derive(Parser, Debug)] #[derive(Parser, Debug)]
#[command( #[command(
@@ -96,10 +94,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
"AETHER_EXECUTION_RUNTIME_DISTRIBUTED_REQUEST_REDIS_URL is required when distributed request limit is enabled", "AETHER_EXECUTION_RUNTIME_DISTRIBUTED_REQUEST_REDIS_URL is required when distributed request limit is enabled",
) )
})?; })?;
Some(DistributedConcurrencyGate::new_redis( let runtime = RuntimeState::redis(
"execution_runtime_requests_distributed", RedisClientConfig {
limit,
RedisDistributedConcurrencyConfig {
url: redis_url.to_string(), url: redis_url.to_string(),
key_prefix: args key_prefix: args
.distributed_request_redis_key_prefix .distributed_request_redis_key_prefix
@@ -107,6 +103,14 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
.map(str::trim) .map(str::trim)
.filter(|value| !value.is_empty()) .filter(|value| !value.is_empty())
.map(ToOwned::to_owned), .map(ToOwned::to_owned),
},
Some(args.distributed_request_command_timeout_ms.max(1)),
)
.await?;
Some(runtime.semaphore(
"execution_runtime_requests_distributed",
limit,
RuntimeSemaphoreConfig {
lease_ttl_ms: args.distributed_request_lease_ttl_ms.max(1), lease_ttl_ms: args.distributed_request_lease_ttl_ms.max(1),
renew_interval_ms: args.distributed_request_renew_interval_ms.max(1), renew_interval_ms: args.distributed_request_renew_interval_ms.max(1),
command_timeout_ms: Some(args.distributed_request_command_timeout_ms.max(1)), command_timeout_ms: Some(args.distributed_request_command_timeout_ms.max(1)),
@@ -5,10 +5,8 @@ use aether_gateway::{
build_tunnel_runtime_router_with_state, TunnelConnConfig, TunnelControlPlaneClient, build_tunnel_runtime_router_with_state, TunnelConnConfig, TunnelControlPlaneClient,
TunnelRuntimeState, TunnelRuntimeState,
}; };
use aether_runtime::{ use aether_runtime::{init_service_runtime, ServiceRuntimeConfig};
init_service_runtime, DistributedConcurrencyGate, RedisDistributedConcurrencyConfig, use aether_runtime_state::{RedisClientConfig, RuntimeSemaphoreConfig, RuntimeState};
ServiceRuntimeConfig,
};
use clap::Parser; use clap::Parser;
use tracing::info; use tracing::info;
@@ -123,10 +121,8 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
"AETHER_TUNNEL_STANDALONE_DISTRIBUTED_REQUEST_REDIS_URL is required when distributed request limit is enabled", "AETHER_TUNNEL_STANDALONE_DISTRIBUTED_REQUEST_REDIS_URL is required when distributed request limit is enabled",
) )
})?; })?;
state = state.with_distributed_request_gate(DistributedConcurrencyGate::new_redis( let runtime = RuntimeState::redis(
"tunnel_requests_distributed", RedisClientConfig {
limit,
RedisDistributedConcurrencyConfig {
url: redis_url.to_string(), url: redis_url.to_string(),
key_prefix: args key_prefix: args
.distributed_request_redis_key_prefix .distributed_request_redis_key_prefix
@@ -134,6 +130,14 @@ async fn main() -> Result<(), Box<dyn std::error::Error>> {
.map(str::trim) .map(str::trim)
.filter(|value| !value.is_empty()) .filter(|value| !value.is_empty())
.map(ToOwned::to_owned), .map(ToOwned::to_owned),
},
Some(args.distributed_request_command_timeout_ms.max(1)),
)
.await?;
state = state.with_distributed_request_gate(runtime.semaphore(
"tunnel_requests_distributed",
limit,
RuntimeSemaphoreConfig {
lease_ttl_ms: args.distributed_request_lease_ttl_ms.max(1), lease_ttl_ms: args.distributed_request_lease_ttl_ms.max(1),
renew_interval_ms: args.distributed_request_renew_interval_ms.max(1), renew_interval_ms: args.distributed_request_renew_interval_ms.max(1),
command_timeout_ms: Some(args.distributed_request_command_timeout_ms.max(1)), command_timeout_ms: Some(args.distributed_request_command_timeout_ms.max(1)),
@@ -106,21 +106,19 @@ async fn schedule_pool_page_candidates(
let key_context_by_id = read_pool_catalog_key_contexts_by_id(state, &candidates).await; let key_context_by_id = read_pool_catalog_key_contexts_by_id(state, &candidates).await;
let mut runtime_by_provider = BTreeMap::new(); let mut runtime_by_provider = BTreeMap::new();
let redis_runner = state.app().redis_kv_runner();
for (provider_id, (pool_config, key_ids)) in provider_runtime_requirements { for (provider_id, (pool_config, key_ids)) in provider_runtime_requirements {
let key_ids = key_ids.into_iter().collect::<Vec<_>>(); let key_ids = key_ids.into_iter().collect::<Vec<_>>();
let runtime = match redis_runner.as_ref() { let runtime = if key_ids.is_empty() {
Some(runner) if !key_ids.is_empty() => { AdminProviderPoolRuntimeState::default()
} else {
read_admin_provider_pool_runtime_state( read_admin_provider_pool_runtime_state(
runner, state.app().runtime_state.as_ref(),
provider_id.as_str(), provider_id.as_str(),
&key_ids, &key_ids,
&pool_config, &pool_config,
sticky_session_token, sticky_session_token,
) )
.await .await
}
_ => AdminProviderPoolRuntimeState::default(),
}; };
runtime_by_provider.insert(provider_id, runtime); runtime_by_provider.insert(provider_id, runtime);
} }
+9 -27
View File
@@ -1,5 +1,4 @@
use aether_data::driver::postgres::PostgresPoolConfig; use aether_data::driver::postgres::PostgresPoolConfig;
use aether_data::driver::redis::RedisClientConfig;
use aether_data::{DataLayerConfig, SqlDatabaseConfig}; use aether_data::{DataLayerConfig, SqlDatabaseConfig};
use std::fmt; use std::fmt;
@@ -7,7 +6,6 @@ use std::fmt;
pub struct GatewayDataConfig { pub struct GatewayDataConfig {
database: Option<SqlDatabaseConfig>, database: Option<SqlDatabaseConfig>,
postgres: Option<PostgresPoolConfig>, postgres: Option<PostgresPoolConfig>,
redis: Option<RedisClientConfig>,
encryption_key: Option<String>, encryption_key: Option<String>,
} }
@@ -16,7 +14,6 @@ impl fmt::Debug for GatewayDataConfig {
f.debug_struct("GatewayDataConfig") f.debug_struct("GatewayDataConfig")
.field("database", &self.database) .field("database", &self.database)
.field("postgres", &self.postgres) .field("postgres", &self.postgres)
.field("redis", &self.redis)
.field("has_encryption_key", &self.encryption_key.is_some()) .field("has_encryption_key", &self.encryption_key.is_some())
.finish() .finish()
} }
@@ -31,7 +28,6 @@ impl GatewayDataConfig {
Self { Self {
database: Some(SqlDatabaseConfig::from_postgres_config(postgres.clone())), database: Some(SqlDatabaseConfig::from_postgres_config(postgres.clone())),
postgres: Some(postgres), postgres: Some(postgres),
redis: None,
encryption_key: None, encryption_key: None,
} }
} }
@@ -41,7 +37,6 @@ impl GatewayDataConfig {
Self { Self {
database: Some(database), database: Some(database),
postgres, postgres,
redis: None,
encryption_key: None, encryption_key: None,
} }
} }
@@ -61,26 +56,6 @@ impl GatewayDataConfig {
self.database.as_ref() self.database.as_ref()
} }
pub fn redis(&self) -> Option<&RedisClientConfig> {
self.redis.as_ref()
}
pub fn with_redis_config(mut self, redis: RedisClientConfig) -> Self {
self.redis = Some(redis);
self
}
pub fn with_redis_url(
self,
url: impl Into<String>,
key_prefix: Option<impl Into<String>>,
) -> Self {
self.with_redis_config(RedisClientConfig {
url: url.into(),
key_prefix: key_prefix.map(Into::into),
})
}
pub fn with_encryption_key(mut self, encryption_key: impl Into<String>) -> Self { pub fn with_encryption_key(mut self, encryption_key: impl Into<String>) -> Self {
let encryption_key = encryption_key.into(); let encryption_key = encryption_key.into();
let encryption_key = encryption_key.trim(); let encryption_key = encryption_key.trim();
@@ -96,15 +71,22 @@ impl GatewayDataConfig {
self.encryption_key.as_deref() self.encryption_key.as_deref()
} }
pub fn with_redis_url(
self,
_url: impl Into<String>,
_key_prefix: Option<impl Into<String>>,
) -> Self {
self
}
pub fn is_enabled(&self) -> bool { pub fn is_enabled(&self) -> bool {
self.database.is_some() || self.postgres.is_some() || self.redis.is_some() self.database.is_some() || self.postgres.is_some()
} }
pub fn to_data_layer_config(&self) -> DataLayerConfig { pub fn to_data_layer_config(&self) -> DataLayerConfig {
DataLayerConfig { DataLayerConfig {
database: self.database.clone(), database: self.database.clone(),
postgres: self.postgres.clone(), postgres: self.postgres.clone(),
redis: self.redis.clone(),
} }
} }
} }
@@ -195,27 +195,6 @@ impl GatewayDataState {
} }
} }
pub(crate) async fn cache_set_string_with_ttl(
&self,
key: &str,
value: &str,
ttl_seconds: u64,
) -> Result<(), DataLayerError> {
let Some(runner) = self.kv_runner() else {
return Ok(());
};
runner.setex(key, value, Some(ttl_seconds)).await?;
Ok(())
}
pub(crate) async fn cache_delete_key(&self, key: &str) -> Result<(), DataLayerError> {
let Some(runner) = self.kv_runner() else {
return Ok(());
};
let _deleted = runner.del(key).await?;
Ok(())
}
pub(crate) async fn list_provider_catalog_providers_by_ids( pub(crate) async fn list_provider_catalog_providers_by_ids(
&self, &self,
provider_ids: &[String], provider_ids: &[String],
+15 -26
View File
@@ -1,5 +1,6 @@
use aether_data::driver::redis::{RedisKvRunner, RedisKvRunnerConfig, RedisLockRunner};
use aether_data::{DataBackends, DataLayerError, DatabaseDriver}; use aether_data::{DataBackends, DataLayerError, DatabaseDriver};
use aether_runtime_state::RuntimeQueueStore;
use std::sync::Arc;
use super::{GatewayDataConfig, GatewayDataState, StoredSystemConfigEntry}; use super::{GatewayDataConfig, GatewayDataState, StoredSystemConfigEntry};
@@ -48,7 +49,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -86,7 +87,7 @@ impl GatewayDataState {
let usage_reader = backends.read().usage(); let usage_reader = backends.read().usage();
let usage_writer = backends.write().usage(); let usage_writer = backends.write().usage();
let user_reader = backends.read().users(); let user_reader = backends.read().users();
let usage_worker_runner = backends.workers().redis(); let usage_worker_queue = None;
let video_task_reader = backends.read().video_tasks(); let video_task_reader = backends.read().video_tasks();
let video_task_writer = backends.write().video_tasks(); let video_task_writer = backends.write().video_tasks();
let wallet_reader = backends.read().wallets(); let wallet_reader = backends.read().wallets();
@@ -124,7 +125,7 @@ impl GatewayDataState {
usage_writer, usage_writer,
user_reader, user_reader,
user_preferences: None, user_preferences: None,
usage_worker_runner, usage_worker_queue,
video_task_reader, video_task_reader,
video_task_writer, video_task_writer,
wallet_reader, wallet_reader,
@@ -138,6 +139,14 @@ impl GatewayDataState {
self.backends.is_some() self.backends.is_some()
} }
pub(crate) fn with_usage_worker_queue(
mut self,
queue: Option<Arc<dyn RuntimeQueueStore>>,
) -> Self {
self.usage_worker_queue = queue;
self
}
pub(crate) fn has_database_maintenance_backend(&self) -> bool { pub(crate) fn has_database_maintenance_backend(&self) -> bool {
self.backends self.backends
.as_ref() .as_ref()
@@ -219,13 +228,6 @@ impl GatewayDataState {
self.global_model_writer.is_some() self.global_model_writer.is_some()
} }
pub(crate) fn has_redis_backend(&self) -> bool {
self.backends
.as_ref()
.and_then(|backends| backends.redis())
.is_some()
}
#[allow(dead_code)] #[allow(dead_code)]
pub(crate) fn has_minimal_candidate_selection_reader(&self) -> bool { pub(crate) fn has_minimal_candidate_selection_reader(&self) -> bool {
self.minimal_candidate_selection_reader.is_some() self.minimal_candidate_selection_reader.is_some()
@@ -263,19 +265,6 @@ impl GatewayDataState {
.is_some_and(|backends| backends.has_system_config_backend()) .is_some_and(|backends| backends.has_system_config_backend())
} }
pub(crate) fn oauth_refresh_lock_runner(&self) -> Option<RedisLockRunner> {
self.backends
.as_ref()
.and_then(|backends| backends.locks().redis())
}
pub(crate) fn kv_runner(&self) -> Option<RedisKvRunner> {
self.backends
.as_ref()
.and_then(|backends| backends.redis())
.and_then(|backend| backend.kv_runner(RedisKvRunnerConfig::default()).ok())
}
pub(crate) fn database_driver(&self) -> Option<DatabaseDriver> { pub(crate) fn database_driver(&self) -> Option<DatabaseDriver> {
self.backends self.backends
.as_ref() .as_ref()
@@ -298,8 +287,8 @@ impl GatewayDataState {
self.usage_writer.is_some() self.usage_writer.is_some()
} }
pub(crate) fn has_usage_worker_runner(&self) -> bool { pub(crate) fn has_usage_worker_queue(&self) -> bool {
self.usage_worker_runner.is_some() self.usage_worker_queue.is_some()
} }
pub(crate) fn has_video_task_reader(&self) -> bool { pub(crate) fn has_video_task_reader(&self) -> bool {
@@ -1,6 +1,5 @@
use aether_billing::enrich_usage_event_with_billing; use aether_billing::enrich_usage_event_with_billing;
use aether_billing::BillingModelContextLookup; use aether_billing::BillingModelContextLookup;
use aether_data::driver::redis::RedisStreamRunner;
use aether_data::repository::audit::RequestAuditReader; use aether_data::repository::audit::RequestAuditReader;
use aether_data::repository::auth::{ use aether_data::repository::auth::{
AuthApiKeyLookupKey, ResolvedAuthApiKeySnapshotReader, StoredAuthApiKeySnapshot, AuthApiKeyLookupKey, ResolvedAuthApiKeySnapshotReader, StoredAuthApiKeySnapshot,
@@ -15,6 +14,7 @@ use aether_data_contracts::repository::provider_catalog::{
use aether_data_contracts::repository::settlement::{StoredUsageSettlement, UsageSettlementInput}; use aether_data_contracts::repository::settlement::{StoredUsageSettlement, UsageSettlementInput};
use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UpsertUsageRecord}; use aether_data_contracts::repository::usage::{StoredRequestUsageAudit, UpsertUsageRecord};
use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskLookupKey}; use aether_data_contracts::repository::video_tasks::{StoredVideoTask, VideoTaskLookupKey};
use aether_runtime_state::RuntimeQueueStore;
use aether_usage_runtime::{ use aether_usage_runtime::{
UsageBillingEventEnricher, UsageBodyCapturePolicy, UsageEvent, UsageRecordWriter, UsageBillingEventEnricher, UsageBodyCapturePolicy, UsageEvent, UsageRecordWriter,
UsageRequestRecordLevel, UsageRuntimeAccess, UsageSettlementWriter, UsageRequestRecordLevel, UsageRuntimeAccess, UsageSettlementWriter,
@@ -242,12 +242,12 @@ impl UsageRuntimeAccess for GatewayDataState {
GatewayDataState::has_usage_writer(self) GatewayDataState::has_usage_writer(self)
} }
fn has_usage_worker_runner(&self) -> bool { fn has_usage_worker_queue(&self) -> bool {
GatewayDataState::has_usage_worker_runner(self) GatewayDataState::has_usage_worker_queue(self)
} }
fn usage_worker_runner(&self) -> Option<RedisStreamRunner> { fn usage_worker_queue(&self) -> Option<std::sync::Arc<dyn RuntimeQueueStore>> {
GatewayDataState::usage_worker_runner(self) GatewayDataState::usage_worker_queue(self)
} }
async fn body_capture_policy(&self) -> Result<UsageBodyCapturePolicy, DataLayerError> { async fn body_capture_policy(&self) -> Result<UsageBodyCapturePolicy, DataLayerError> {
+3 -8
View File
@@ -11,9 +11,6 @@ use crate::provider_transport::{
read_provider_transport_snapshot, GatewayProviderTransportSnapshot, read_provider_transport_snapshot, GatewayProviderTransportSnapshot,
}; };
use crate::video_tasks::LocalVideoTaskReadResponse; use crate::video_tasks::LocalVideoTaskReadResponse;
use aether_data::driver::redis::{
RedisKvRunner, RedisKvRunnerConfig, RedisLockRunner, RedisStreamRunner,
};
use aether_data::repository::announcements::{ use aether_data::repository::announcements::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository, AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord, CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,
@@ -122,6 +119,7 @@ use aether_data_contracts::repository::video_tasks::{
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount, StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskStatusCount, VideoTaskWriteRepository, VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskStatusCount, VideoTaskWriteRepository,
}; };
use aether_runtime_state::RuntimeQueueStore;
#[derive(Clone, Default)] #[derive(Clone, Default)]
pub(crate) struct GatewayDataState { pub(crate) struct GatewayDataState {
@@ -155,7 +153,7 @@ pub(crate) struct GatewayDataState {
usage_writer: Option<Arc<dyn UsageWriteRepository>>, usage_writer: Option<Arc<dyn UsageWriteRepository>>,
user_reader: Option<Arc<dyn UserReadRepository>>, user_reader: Option<Arc<dyn UserReadRepository>>,
user_preferences: Option<Arc<RwLock<BTreeMap<String, StoredUserPreferenceRecord>>>>, user_preferences: Option<Arc<RwLock<BTreeMap<String, StoredUserPreferenceRecord>>>>,
usage_worker_runner: Option<RedisStreamRunner>, usage_worker_queue: Option<Arc<dyn RuntimeQueueStore>>,
video_task_reader: Option<Arc<dyn VideoTaskReadRepository>>, video_task_reader: Option<Arc<dyn VideoTaskReadRepository>>,
video_task_writer: Option<Arc<dyn VideoTaskWriteRepository>>, video_task_writer: Option<Arc<dyn VideoTaskWriteRepository>>,
wallet_reader: Option<Arc<dyn WalletReadRepository>>, wallet_reader: Option<Arc<dyn WalletReadRepository>>,
@@ -253,10 +251,7 @@ impl fmt::Debug for GatewayDataState {
.field("has_usage_reader", &self.usage_reader.is_some()) .field("has_usage_reader", &self.usage_reader.is_some())
.field("has_usage_writer", &self.usage_writer.is_some()) .field("has_usage_writer", &self.usage_writer.is_some())
.field("has_user_preferences", &self.user_preferences.is_some()) .field("has_user_preferences", &self.user_preferences.is_some())
.field( .field("has_usage_worker_queue", &self.usage_worker_queue.is_some())
"has_usage_worker_runner",
&self.usage_worker_runner.is_some(),
)
.field("has_video_task_reader", &self.video_task_reader.is_some()) .field("has_video_task_reader", &self.video_task_reader.is_some())
.field("has_video_task_writer", &self.video_task_writer.is_some()) .field("has_video_task_writer", &self.video_task_writer.is_some())
.field("has_wallet_reader", &self.wallet_reader.is_some()) .field("has_wallet_reader", &self.wallet_reader.is_some())
+18 -17
View File
@@ -13,27 +13,28 @@ use super::{
DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput,
GatewayDataState, GatewayProviderTransportSnapshot, LocalVideoTaskReadResponse, GatewayDataState, GatewayProviderTransportSnapshot, LocalVideoTaskReadResponse,
ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome,
RedeemWalletCodeInput, RedeemWalletCodeOutcome, RedisStreamRunner, RequestAuditBundle, RedeemWalletCodeInput, RedeemWalletCodeOutcome, RequestAuditBundle, RequestCandidateTrace,
RequestCandidateTrace, StoredAdminAuditLogPage, StoredAdminPaymentCallbackPage, StoredAdminAuditLogPage, StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder,
StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCodeBatch, StoredAdminPaymentOrderPage, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage,
StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage, StoredAdminWalletListPage,
StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestPage,
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction, StoredAdminWalletTransaction, StoredAdminWalletTransactionPage, StoredAnnouncement,
StoredAdminWalletTransactionPage, StoredAnnouncement, StoredAnnouncementPage, StoredAnnouncementPage, StoredBillingModelContext, StoredProviderQuotaSnapshot,
StoredBillingModelContext, StoredProviderQuotaSnapshot, StoredProviderUsageSummary, StoredProviderUsageSummary, StoredRequestUsageAudit, StoredSuspiciousActivity,
StoredRequestUsageAudit, StoredSuspiciousActivity, StoredUsageSettlement, StoredUsageSettlement, StoredUserAuditLogPage, StoredUserAuthRecord, StoredUserExportRow,
StoredUserAuditLogPage, StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary, StoredUserSummary, StoredVideoTask, StoredWalletDailyUsageLedger,
StoredVideoTask, StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage, StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAnnouncementRecord,
StoredWalletSnapshot, UpdateAnnouncementRecord, UpsertUsageRecord, UpsertVideoTask, UpsertUsageRecord, UpsertVideoTask, UsageSettlementInput, VideoTaskLookupKey,
UsageSettlementInput, VideoTaskLookupKey, VideoTaskModelCount, VideoTaskQueryFilter, VideoTaskModelCount, VideoTaskQueryFilter, VideoTaskStatusCount,
VideoTaskStatusCount, WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult, WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult, WalletLookupKey,
WalletLookupKey, WalletMutationOutcome, WalletMutationOutcome,
}; };
use aether_data_contracts::repository::usage::{ use aether_data_contracts::repository::usage::{
PendingUsageCleanupSummary, ProviderApiKeyWindowUsageRequest, PendingUsageCleanupSummary, ProviderApiKeyWindowUsageRequest,
StoredProviderApiKeyWindowUsageSummary, StoredUsageDailySummary, UsageAuditListQuery, StoredProviderApiKeyWindowUsageSummary, StoredUsageDailySummary, UsageAuditListQuery,
UsageCleanupSummary, UsageCleanupWindow, UsageDailyHeatmapQuery, UsageCleanupSummary, UsageCleanupWindow, UsageDailyHeatmapQuery,
}; };
use aether_runtime_state::RuntimeQueueStore;
use aether_video_tasks_core::read_data_backed_video_task_response; use aether_video_tasks_core::read_data_backed_video_task_response;
impl GatewayDataState { impl GatewayDataState {
@@ -1477,8 +1478,8 @@ impl GatewayDataState {
} }
} }
pub(crate) fn usage_worker_runner(&self) -> Option<RedisStreamRunner> { pub(crate) fn usage_worker_queue(&self) -> Option<std::sync::Arc<dyn RuntimeQueueStore>> {
self.usage_worker_runner.clone() self.usage_worker_queue.clone()
} }
pub(crate) async fn find_billing_model_context( pub(crate) async fn find_billing_model_context(
@@ -40,7 +40,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -90,7 +90,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -73,7 +73,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -122,7 +122,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -167,7 +167,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -300,7 +300,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -361,7 +361,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -415,7 +415,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -478,7 +478,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -523,7 +523,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -569,7 +569,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -626,7 +626,7 @@ impl GatewayDataState {
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -685,7 +685,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -728,7 +728,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -786,7 +786,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: Some(repository), user_reader: Some(repository),
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -837,7 +837,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: Some(user_repository), user_reader: Some(user_repository),
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: Some(wallet_reader), wallet_reader: Some(wallet_reader),
@@ -893,7 +893,7 @@ impl GatewayDataState {
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: Some(user_repository), user_reader: Some(user_repository),
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: Some(wallet_reader), wallet_reader: Some(wallet_reader),
@@ -950,7 +950,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: Some(user_repository), user_reader: Some(user_repository),
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: Some(wallet_reader), wallet_reader: Some(wallet_reader),
@@ -1006,7 +1006,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1051,7 +1051,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1096,7 +1096,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1153,7 +1153,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1215,7 +1215,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1260,7 +1260,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1310,7 +1310,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1377,7 +1377,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1439,7 +1439,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1485,7 +1485,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1531,7 +1531,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1579,7 +1579,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1625,7 +1625,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1671,7 +1671,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1725,7 +1725,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1780,7 +1780,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1838,7 +1838,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1902,7 +1902,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -1967,7 +1967,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -2036,7 +2036,7 @@ impl GatewayDataState {
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -2112,7 +2112,7 @@ impl GatewayDataState {
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: Some(wallet_reader), wallet_reader: Some(wallet_reader),
@@ -2170,7 +2170,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -2219,7 +2219,7 @@ impl GatewayDataState {
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -2264,7 +2264,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -2315,7 +2315,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: Some(wallet_reader), wallet_reader: Some(wallet_reader),
@@ -2370,7 +2370,7 @@ impl GatewayDataState {
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: Some(wallet_reader), wallet_reader: Some(wallet_reader),
@@ -2426,7 +2426,7 @@ impl GatewayDataState {
usage_writer: Some(usage_writer), usage_writer: Some(usage_writer),
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: Some(wallet_reader), wallet_reader: Some(wallet_reader),
@@ -2475,7 +2475,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: None, video_task_reader: None,
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -44,7 +44,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: Some(repository), video_task_reader: Some(repository),
video_task_writer: None, video_task_writer: None,
wallet_reader: None, wallet_reader: None,
@@ -96,7 +96,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: Some(video_task_reader), video_task_reader: Some(video_task_reader),
video_task_writer: Some(video_task_writer), video_task_writer: Some(video_task_writer),
wallet_reader: None, wallet_reader: None,
@@ -145,7 +145,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: Some(video_task_reader), video_task_reader: Some(video_task_reader),
video_task_writer: Some(video_task_writer), video_task_writer: Some(video_task_writer),
wallet_reader: None, wallet_reader: None,
@@ -198,7 +198,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: Some(video_task_reader), video_task_reader: Some(video_task_reader),
video_task_writer: Some(video_task_writer), video_task_writer: Some(video_task_writer),
wallet_reader: None, wallet_reader: None,
@@ -255,7 +255,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: Some(video_task_reader), video_task_reader: Some(video_task_reader),
video_task_writer: Some(video_task_writer), video_task_writer: Some(video_task_writer),
wallet_reader: None, wallet_reader: None,
@@ -321,7 +321,7 @@ impl GatewayDataState {
usage_writer: None, usage_writer: None,
user_reader: None, user_reader: None,
user_preferences: None, user_preferences: None,
usage_worker_runner: None, usage_worker_queue: None,
video_task_reader: Some(video_task_reader), video_task_reader: Some(video_task_reader),
video_task_writer: Some(video_task_writer), video_task_writer: Some(video_task_writer),
wallet_reader: None, wallet_reader: None,
@@ -4,10 +4,9 @@ use std::sync::Arc;
use aether_contracts::ExecutionPlan; use aether_contracts::ExecutionPlan;
use aether_runtime::{ use aether_runtime::{
maybe_hold_axum_response_permit, prometheus_response, service_up_sample, AdmissionPermit, maybe_hold_axum_response_permit, prometheus_response, service_up_sample, AdmissionPermit,
ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot, DistributedConcurrencyError, ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot, MetricKind, MetricLabel, MetricSample,
DistributedConcurrencyGate, DistributedConcurrencySnapshot, MetricKind, MetricLabel,
MetricSample,
}; };
use aether_runtime_state::{RuntimeSemaphore, RuntimeSemaphoreError, RuntimeSemaphoreSnapshot};
use axum::body::{to_bytes, Body}; use axum::body::{to_bytes, Body};
use axum::extract::{Request, State}; use axum::extract::{Request, State};
use axum::http::StatusCode; use axum::http::StatusCode;
@@ -30,7 +29,7 @@ const DISTRIBUTED_REQUEST_GATE_NAME: &str = "execution_runtime_requests_distribu
struct ExecutionRuntimeAppState { struct ExecutionRuntimeAppState {
execution_runtime: DirectSyncExecutionRuntime, execution_runtime: DirectSyncExecutionRuntime,
request_gate: Option<Arc<ConcurrencyGate>>, request_gate: Option<Arc<ConcurrencyGate>>,
distributed_request_gate: Option<Arc<DistributedConcurrencyGate>>, distributed_request_gate: Option<Arc<RuntimeSemaphore>>,
} }
impl ExecutionRuntimeAppState { impl ExecutionRuntimeAppState {
@@ -44,7 +43,7 @@ impl ExecutionRuntimeAppState {
} }
} }
fn with_distributed_request_gate(mut self, gate: DistributedConcurrencyGate) -> Self { fn with_distributed_request_gate(mut self, gate: RuntimeSemaphore) -> Self {
self.distributed_request_gate = Some(Arc::new(gate)); self.distributed_request_gate = Some(Arc::new(gate));
self self
} }
@@ -55,7 +54,7 @@ impl ExecutionRuntimeAppState {
async fn distributed_request_concurrency_snapshot( async fn distributed_request_concurrency_snapshot(
&self, &self,
) -> Result<Option<DistributedConcurrencySnapshot>, DistributedConcurrencyError> { ) -> Result<Option<RuntimeSemaphoreSnapshot>, RuntimeSemaphoreError> {
match self.distributed_request_gate.as_ref() { match self.distributed_request_gate.as_ref() {
Some(gate) => gate.snapshot().await.map(Some), Some(gate) => gate.snapshot().await.map(Some),
None => Ok(None), None => Ok(None),
@@ -122,7 +121,7 @@ pub fn build_execution_runtime_router_with_request_concurrency_limit(
pub fn build_execution_runtime_router_with_request_gates( pub fn build_execution_runtime_router_with_request_gates(
limit: Option<usize>, limit: Option<usize>,
distributed_gate: Option<DistributedConcurrencyGate>, distributed_gate: Option<RuntimeSemaphore>,
) -> Router { ) -> Router {
let state = match distributed_gate { let state = match distributed_gate {
Some(gate) => ExecutionRuntimeAppState::with_request_concurrency_limit(limit) Some(gate) => ExecutionRuntimeAppState::with_request_concurrency_limit(limit)
@@ -142,7 +141,7 @@ pub fn build_execution_runtime_router_with_request_gates(
pub async fn serve_execution_runtime_tcp( pub async fn serve_execution_runtime_tcp(
bind: &str, bind: &str,
max_in_flight_requests: Option<usize>, max_in_flight_requests: Option<usize>,
distributed_request_gate: Option<DistributedConcurrencyGate>, distributed_request_gate: Option<RuntimeSemaphore>,
) -> Result<(), Box<dyn std::error::Error>> { ) -> Result<(), Box<dyn std::error::Error>> {
let listener = tokio::net::TcpListener::bind(bind).await?; let listener = tokio::net::TcpListener::bind(bind).await?;
axum::serve( axum::serve(
@@ -159,7 +158,7 @@ pub async fn serve_execution_runtime_tcp(
pub async fn serve_execution_runtime_unix( pub async fn serve_execution_runtime_unix(
socket_path: &Path, socket_path: &Path,
max_in_flight_requests: Option<usize>, max_in_flight_requests: Option<usize>,
distributed_request_gate: Option<DistributedConcurrencyGate>, distributed_request_gate: Option<RuntimeSemaphore>,
) -> Result<(), Box<dyn std::error::Error>> { ) -> Result<(), Box<dyn std::error::Error>> {
if let Some(parent) = socket_path.parent() { if let Some(parent) = socket_path.parent() {
std::fs::create_dir_all(parent)?; std::fs::create_dir_all(parent)?;
@@ -262,11 +261,11 @@ async fn acquire_request_permit(
match state.try_acquire_request_permit().await { match state.try_acquire_request_permit().await {
Ok(permit) => Ok(permit), Ok(permit) => Ok(permit),
Err(RequestAdmissionError::Local(ConcurrencyError::Saturated { gate, limit })) Err(RequestAdmissionError::Local(ConcurrencyError::Saturated { gate, limit }))
| Err(RequestAdmissionError::Distributed(DistributedConcurrencyError::Saturated { | Err(RequestAdmissionError::Distributed(RuntimeSemaphoreError::Saturated {
gate, gate,
limit, limit,
})) }))
| Err(RequestAdmissionError::Distributed(DistributedConcurrencyError::Unavailable { | Err(RequestAdmissionError::Distributed(RuntimeSemaphoreError::Unavailable {
gate, gate,
limit, limit,
.. ..
@@ -278,9 +277,9 @@ async fn acquire_request_permit(
"execution runtime request concurrency gate {gate} is closed" "execution runtime request concurrency gate {gate} is closed"
))), ))),
), ),
Err(RequestAdmissionError::Distributed( Err(RequestAdmissionError::Distributed(RuntimeSemaphoreError::InvalidConfiguration(
DistributedConcurrencyError::InvalidConfiguration(message), message,
)) => Err(ExecutionRuntimeAppError( ))) => Err(ExecutionRuntimeAppError(
ExecutionRuntimeServerError::RequestRead(message), ExecutionRuntimeServerError::RequestRead(message),
)), )),
} }
@@ -289,7 +288,7 @@ async fn acquire_request_permit(
#[derive(Debug)] #[derive(Debug)]
enum RequestAdmissionError { enum RequestAdmissionError {
Local(ConcurrencyError), Local(ConcurrencyError),
Distributed(DistributedConcurrencyError), Distributed(RuntimeSemaphoreError),
} }
async fn parse_request_json<T>(request: Request) -> Result<T, ExecutionRuntimeAppError> async fn parse_request_json<T>(request: Request) -> Result<T, ExecutionRuntimeAppError>
@@ -380,6 +379,9 @@ mod tests {
build_execution_runtime_router_with_request_gates, DISTRIBUTED_REQUEST_GATE_NAME, build_execution_runtime_router_with_request_gates, DISTRIBUTED_REQUEST_GATE_NAME,
}; };
use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody}; use aether_contracts::{ExecutionPlan, ExecutionTimeouts, RequestBody};
use aether_runtime_state::{
MemoryRuntimeStateConfig, RuntimeSemaphore, RuntimeSemaphoreConfig, RuntimeState,
};
use axum::body::{Body, Bytes}; use axum::body::{Body, Bytes};
use axum::response::Response; use axum::response::Response;
use axum::routing::any; use axum::routing::any;
@@ -389,6 +391,12 @@ mod tests {
use std::sync::atomic::{AtomicUsize, Ordering}; use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc; use std::sync::Arc;
fn distributed_gate(gate: &'static str, limit: usize) -> RuntimeSemaphore {
RuntimeState::memory(MemoryRuntimeStateConfig::default())
.semaphore(gate, limit, RuntimeSemaphoreConfig::default())
.expect("distributed semaphore")
}
async fn start_server(app: Router) -> (String, tokio::task::JoinHandle<()>) { async fn start_server(app: Router) -> (String, tokio::task::JoinHandle<()>) {
let listener = crate::test_support::bind_loopback_listener() let listener = crate::test_support::bind_loopback_listener()
.await .await
@@ -517,10 +525,7 @@ mod tests {
}), }),
); );
let (upstream_url, upstream_handle) = start_server(upstream).await; let (upstream_url, upstream_handle) = start_server(upstream).await;
let distributed_gate = aether_runtime::DistributedConcurrencyGate::new_in_memory( let distributed_gate = distributed_gate(DISTRIBUTED_REQUEST_GATE_NAME, 1);
DISTRIBUTED_REQUEST_GATE_NAME,
1,
);
let runtime_a = let runtime_a =
build_execution_runtime_router_with_request_gates(None, Some(distributed_gate.clone())); build_execution_runtime_router_with_request_gates(None, Some(distributed_gate.clone()));
let runtime_b = let runtime_b =
@@ -571,10 +576,7 @@ mod tests {
async fn execution_runtime_exposes_request_concurrency_metrics() { async fn execution_runtime_exposes_request_concurrency_metrics() {
let runtime = build_execution_runtime_router_with_request_gates( let runtime = build_execution_runtime_router_with_request_gates(
Some(4), Some(4),
Some(aether_runtime::DistributedConcurrencyGate::new_in_memory( Some(distributed_gate(DISTRIBUTED_REQUEST_GATE_NAME, 6)),
DISTRIBUTED_REQUEST_GATE_NAME,
6,
)),
); );
let (runtime_url, runtime_handle) = start_server(runtime).await; let (runtime_url, runtime_handle) = start_server(runtime).await;
@@ -25,19 +25,16 @@ async fn store_admin_external_models_cache(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
payload: &serde_json::Value, payload: &serde_json::Value,
) -> Result<(), GatewayError> { ) -> Result<(), GatewayError> {
let Some(runner) = state.redis_kv_runner() else {
return Ok(());
};
let serialized = let serialized =
serde_json::to_string(payload).map_err(|err| GatewayError::Internal(err.to_string()))?; serde_json::to_string(payload).map_err(|err| GatewayError::Internal(err.to_string()))?;
runner state
.setex( .as_ref()
.runtime_kv_setex(
ADMIN_EXTERNAL_MODELS_CACHE_KEY, ADMIN_EXTERNAL_MODELS_CACHE_KEY,
&serialized, &serialized,
Some(ADMIN_EXTERNAL_MODELS_CACHE_TTL_SECS), ADMIN_EXTERNAL_MODELS_CACHE_TTL_SECS,
) )
.await .await?;
.map_err(|err| GatewayError::Internal(err.to_string()))?;
Ok(()) Ok(())
} }
@@ -64,21 +61,15 @@ async fn fetch_admin_external_models_from_source(
pub(crate) async fn read_admin_external_models_cache( pub(crate) async fn read_admin_external_models_cache(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
) -> Result<Option<serde_json::Value>, GatewayError> { ) -> Result<Option<serde_json::Value>, GatewayError> {
if let Some(runner) = state.redis_kv_runner() { if let Some(raw) = state
match runner.client().get_multiplexed_async_connection().await { .as_ref()
Ok(mut connection) => { .runtime_kv_get(ADMIN_EXTERNAL_MODELS_CACHE_KEY)
let namespaced_key = runner.keyspace().key(ADMIN_EXTERNAL_MODELS_CACHE_KEY); .await?
match redis::cmd("GET")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
{ {
Ok(Some(raw)) => match serde_json::from_str::<serde_json::Value>(&raw) { match serde_json::from_str::<serde_json::Value>(&raw) {
Ok(payload) => { Ok(payload) => {
let payload = normalize_admin_external_models_payload(payload); let payload = normalize_admin_external_models_payload(payload);
if let Err(err) = if let Err(err) = store_admin_external_models_cache(state, &payload).await {
store_admin_external_models_cache(state, &payload).await
{
warn!(error = ?err, "failed to refresh external models cache ttl"); warn!(error = ?err, "failed to refresh external models cache ttl");
} }
return Ok(Some(payload)); return Ok(Some(payload));
@@ -86,16 +77,6 @@ pub(crate) async fn read_admin_external_models_cache(
Err(err) => { Err(err) => {
warn!(error = %err, "failed to parse cached external models payload"); warn!(error = %err, "failed to parse cached external models payload");
} }
},
Ok(None) => {}
Err(err) => {
warn!(error = %err, "failed to read external models cache");
}
}
}
Err(err) => {
warn!(error = %err, "failed to connect to redis for external models cache");
}
} }
} }
@@ -116,19 +97,13 @@ pub(crate) async fn read_admin_external_models_cache(
pub(crate) async fn clear_admin_external_models_cache( pub(crate) async fn clear_admin_external_models_cache(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
) -> Result<serde_json::Value, GatewayError> { ) -> Result<serde_json::Value, GatewayError> {
let Some(runner) = state.redis_kv_runner() else { let deleted = state
return Ok(json!({ .as_ref()
"cleared": false, .runtime_kv_del(ADMIN_EXTERNAL_MODELS_CACHE_KEY)
"message": "Redis 未启用", .await?;
}));
};
let deleted = runner
.del(ADMIN_EXTERNAL_MODELS_CACHE_KEY)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
Ok(json!({ Ok(json!({
"cleared": deleted > 0, "cleared": deleted,
"message": if deleted > 0 { "缓存已清除" } else { "缓存不存在" }, "message": if deleted { "缓存已清除" } else { "缓存不存在" },
})) }))
} }
@@ -426,22 +426,13 @@ pub(super) async fn delete_admin_monitoring_cache_affinity_raw_keys(
return Ok(0); return Ok(0);
} }
if let Some(runner) = state.redis_kv_runner() { let deleted = state
let mut connection = runner .runtime_state()
.client() .kv_delete_many(raw_keys)
.get_multiplexed_async_connection()
.await .await
.map_err(|err| { .map_err(|err| GatewayError::Internal(format!("runtime cache delete failed: {err}")))?;
GatewayError::Internal(format!("admin monitoring redis connect failed: {err}")) if deleted > 0 {
})?; return Ok(deleted);
let deleted = redis::cmd("DEL")
.arg(raw_keys)
.query_async::<i64>(&mut connection)
.await
.map_err(|err| {
GatewayError::Internal(format!("admin monitoring redis delete failed: {err}"))
})?;
return Ok(usize::try_from(deleted).unwrap_or(0));
} }
Ok(delete_admin_monitoring_cache_affinity_entries_for_tests( Ok(delete_admin_monitoring_cache_affinity_entries_for_tests(
@@ -1,7 +1,5 @@
use super::cache_config::ADMIN_MONITORING_REDIS_CACHE_CATEGORIES; use super::cache_config::ADMIN_MONITORING_REDIS_CACHE_CATEGORIES;
use super::cache_store::{ use super::cache_store::list_admin_monitoring_namespaced_keys;
admin_monitoring_has_test_redis_keys, list_admin_monitoring_namespaced_keys,
};
use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError; use crate::GatewayError;
use axum::{ use axum::{
@@ -14,17 +12,6 @@ use serde_json::json;
pub(super) async fn build_admin_monitoring_model_mapping_stats_response( pub(super) async fn build_admin_monitoring_model_mapping_stats_response(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
) -> Result<Response<Body>, GatewayError> { ) -> Result<Response<Body>, GatewayError> {
if state.redis_kv_runner().is_none() && !admin_monitoring_has_test_redis_keys(state) {
return Ok(Json(json!({
"status": "ok",
"data": {
"available": false,
"message": "Redis 未启用,模型映射缓存不可用",
}
}))
.into_response());
};
let model_id_keys = list_admin_monitoring_namespaced_keys(state, "model:id:*").await?; let model_id_keys = list_admin_monitoring_namespaced_keys(state, "model:id:*").await?;
let global_model_id_keys = let global_model_id_keys =
list_admin_monitoring_namespaced_keys(state, "global_model:id:*").await?; list_admin_monitoring_namespaced_keys(state, "global_model:id:*").await?;
@@ -49,6 +36,7 @@ pub(super) async fn build_admin_monitoring_model_mapping_stats_response(
"status": "ok", "status": "ok",
"data": { "data": {
"available": true, "available": true,
"backend": state.runtime_state().backend_kind().as_str(),
"ttl_seconds": 300, "ttl_seconds": 300,
"total_keys": total_keys, "total_keys": total_keys,
"breakdown": { "breakdown": {
@@ -69,17 +57,6 @@ pub(super) async fn build_admin_monitoring_model_mapping_stats_response(
pub(super) async fn build_admin_monitoring_redis_cache_categories_response( pub(super) async fn build_admin_monitoring_redis_cache_categories_response(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
) -> Result<Response<Body>, GatewayError> { ) -> Result<Response<Body>, GatewayError> {
if state.redis_kv_runner().is_none() && !admin_monitoring_has_test_redis_keys(state) {
return Ok(Json(json!({
"status": "ok",
"data": {
"available": false,
"message": "Redis 未启用",
}
}))
.into_response());
};
let mut categories = Vec::with_capacity(ADMIN_MONITORING_REDIS_CACHE_CATEGORIES.len()); let mut categories = Vec::with_capacity(ADMIN_MONITORING_REDIS_CACHE_CATEGORIES.len());
let mut total_keys = 0usize; let mut total_keys = 0usize;
@@ -101,6 +78,7 @@ pub(super) async fn build_admin_monitoring_redis_cache_categories_response(
"status": "ok", "status": "ok",
"data": { "data": {
"available": true, "available": true,
"backend": state.runtime_state().backend_kind().as_str(),
"categories": categories, "categories": categories,
"total_keys": total_keys, "total_keys": total_keys,
} }
@@ -3,15 +3,8 @@ use super::super::cache_affinity::{
delete_admin_monitoring_cache_affinity_raw_keys, delete_admin_monitoring_cache_affinity_raw_keys,
}; };
use super::super::cache_identity::admin_monitoring_list_export_api_key_records_by_ids; use super::super::cache_identity::admin_monitoring_list_export_api_key_records_by_ids;
use super::super::cache_route_helpers::{ use super::super::cache_route_helpers::admin_monitoring_cache_affinity_delete_params_from_path;
admin_monitoring_cache_affinity_delete_params_from_path, use super::super::cache_store::list_admin_monitoring_cache_affinity_records_by_affinity_keys;
admin_monitoring_cache_affinity_unavailable_response,
};
use super::super::cache_store::{
admin_monitoring_has_runtime_scheduler_affinity_entries,
list_admin_monitoring_cache_affinity_records_by_affinity_keys,
load_admin_monitoring_cache_affinity_entries_for_tests,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError; use crate::GatewayError;
use aether_admin::observability::monitoring::{ use aether_admin::observability::monitoring::{
@@ -76,13 +69,6 @@ pub(in super::super) async fn build_admin_monitoring_cache_affinity_delete_respo
)); ));
}; };
if state.redis_kv_runner().is_none()
&& load_admin_monitoring_cache_affinity_entries_for_tests(state).is_empty()
&& !admin_monitoring_has_runtime_scheduler_affinity_entries(state)
{
return Ok(admin_monitoring_cache_affinity_unavailable_response());
}
let target_affinity_keys = let target_affinity_keys =
std::iter::once(affinity_key.clone()).collect::<std::collections::BTreeSet<_>>(); std::iter::once(affinity_key.clone()).collect::<std::collections::BTreeSet<_>>();
let delete_filter = admin_monitoring_cache_affinity_delete_filter_from_query( let delete_filter = admin_monitoring_cache_affinity_delete_filter_from_query(
@@ -2,7 +2,6 @@ use super::super::cache_affinity::{
clear_admin_monitoring_scheduler_affinity_entries, clear_admin_monitoring_scheduler_affinity_entries,
delete_admin_monitoring_cache_affinity_raw_keys, delete_admin_monitoring_cache_affinity_raw_keys,
}; };
use super::super::cache_route_helpers::admin_monitoring_cache_affinity_unavailable_response;
use super::super::cache_store::list_admin_monitoring_cache_affinity_records; use super::super::cache_store::list_admin_monitoring_cache_affinity_records;
use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::request::AdminAppState;
use crate::GatewayError; use crate::GatewayError;
@@ -13,10 +12,6 @@ pub(in super::super) async fn build_admin_monitoring_cache_flush_response(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
) -> Result<Response<Body>, GatewayError> { ) -> Result<Response<Body>, GatewayError> {
let raw_affinities = list_admin_monitoring_cache_affinity_records(state).await?; let raw_affinities = list_admin_monitoring_cache_affinity_records(state).await?;
if state.redis_kv_runner().is_none() && raw_affinities.is_empty() {
return Ok(admin_monitoring_cache_affinity_unavailable_response());
}
let raw_keys = raw_affinities let raw_keys = raw_affinities
.iter() .iter()
.map(|item| item.raw_key.clone()) .map(|item| item.raw_key.clone())
@@ -1,10 +1,9 @@
use super::super::cache_route_helpers::{ use super::super::cache_route_helpers::{
admin_monitoring_cache_model_mapping_provider_params_from_path, admin_monitoring_cache_model_mapping_provider_params_from_path,
admin_monitoring_cache_model_name_from_path, admin_monitoring_redis_unavailable_response, admin_monitoring_cache_model_name_from_path,
}; };
use super::super::cache_store::{ use super::super::cache_store::{
admin_monitoring_has_test_redis_keys, delete_admin_monitoring_namespaced_keys, delete_admin_monitoring_namespaced_keys, list_admin_monitoring_namespaced_keys,
list_admin_monitoring_namespaced_keys,
}; };
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError; use crate::GatewayError;
@@ -19,10 +18,6 @@ use axum::{body::Body, response::Response};
pub(in super::super) async fn build_admin_monitoring_model_mapping_delete_response( pub(in super::super) async fn build_admin_monitoring_model_mapping_delete_response(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
) -> Result<Response<Body>, GatewayError> { ) -> Result<Response<Body>, GatewayError> {
if state.redis_kv_runner().is_none() && !admin_monitoring_has_test_redis_keys(state) {
return Ok(admin_monitoring_redis_unavailable_response());
}
let mut raw_keys = list_admin_monitoring_namespaced_keys(state, "model:*").await?; let mut raw_keys = list_admin_monitoring_namespaced_keys(state, "model:*").await?;
raw_keys.extend(list_admin_monitoring_namespaced_keys(state, "global_model:*").await?); raw_keys.extend(list_admin_monitoring_namespaced_keys(state, "global_model:*").await?);
raw_keys.sort(); raw_keys.sort();
@@ -41,10 +36,6 @@ pub(in super::super) async fn build_admin_monitoring_model_mapping_delete_model_
else { else {
return Ok(admin_monitoring_bad_request_response("缺少 model_name")); return Ok(admin_monitoring_bad_request_response("缺少 model_name"));
}; };
if state.redis_kv_runner().is_none() && !admin_monitoring_has_test_redis_keys(state) {
return Ok(admin_monitoring_redis_unavailable_response());
}
let candidate_keys = [ let candidate_keys = [
format!("global_model:resolve:{model_name}"), format!("global_model:resolve:{model_name}"),
format!("global_model:name:{model_name}"), format!("global_model:name:{model_name}"),
@@ -85,10 +76,6 @@ pub(in super::super) async fn build_admin_monitoring_model_mapping_delete_provid
"缺少 provider_id 或 global_model_id", "缺少 provider_id 或 global_model_id",
)); ));
}; };
if state.redis_kv_runner().is_none() && !admin_monitoring_has_test_redis_keys(state) {
return Ok(admin_monitoring_redis_unavailable_response());
}
let candidate_keys = [ let candidate_keys = [
format!("model:provider_global:{provider_id}:{global_model_id}"), format!("model:provider_global:{provider_id}:{global_model_id}"),
format!("model:provider_global:hits:{provider_id}:{global_model_id}"), format!("model:provider_global:hits:{provider_id}:{global_model_id}"),
@@ -2,10 +2,7 @@ use super::super::cache_affinity::{
clear_admin_monitoring_scheduler_affinity_entries, clear_admin_monitoring_scheduler_affinity_entries,
delete_admin_monitoring_cache_affinity_raw_keys, delete_admin_monitoring_cache_affinity_raw_keys,
}; };
use super::super::cache_route_helpers::{ use super::super::cache_route_helpers::admin_monitoring_cache_provider_id_from_path;
admin_monitoring_cache_affinity_unavailable_response,
admin_monitoring_cache_provider_id_from_path,
};
use super::super::cache_store::list_admin_monitoring_cache_affinity_records; use super::super::cache_store::list_admin_monitoring_cache_affinity_records;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError; use crate::GatewayError;
@@ -26,10 +23,6 @@ pub(in super::super) async fn build_admin_monitoring_cache_provider_delete_respo
}; };
let raw_affinities = list_admin_monitoring_cache_affinity_records(state).await?; let raw_affinities = list_admin_monitoring_cache_affinity_records(state).await?;
if state.redis_kv_runner().is_none() && raw_affinities.is_empty() {
return Ok(admin_monitoring_cache_affinity_unavailable_response());
}
let target_affinities = raw_affinities let target_affinities = raw_affinities
.into_iter() .into_iter()
.filter(|item| item.provider_id.as_deref() == Some(provider_id.as_str())) .filter(|item| item.provider_id.as_deref() == Some(provider_id.as_str()))
@@ -1,10 +1,7 @@
use super::super::cache_config::ADMIN_MONITORING_REDIS_CACHE_CATEGORIES; use super::super::cache_config::ADMIN_MONITORING_REDIS_CACHE_CATEGORIES;
use super::super::cache_route_helpers::{ use super::super::cache_route_helpers::admin_monitoring_cache_redis_category_from_path;
admin_monitoring_cache_redis_category_from_path, admin_monitoring_redis_unavailable_response,
};
use super::super::cache_store::{ use super::super::cache_store::{
admin_monitoring_has_test_redis_keys, delete_admin_monitoring_namespaced_keys, delete_admin_monitoring_namespaced_keys, list_admin_monitoring_namespaced_keys,
list_admin_monitoring_namespaced_keys,
}; };
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError; use crate::GatewayError;
@@ -31,10 +28,6 @@ pub(in super::super) async fn build_admin_monitoring_redis_keys_delete_response(
return Ok(admin_monitoring_unknown_cache_category_response(&category)); return Ok(admin_monitoring_unknown_cache_category_response(&category));
}; };
if state.redis_kv_runner().is_none() && !admin_monitoring_has_test_redis_keys(state) {
return Ok(admin_monitoring_redis_unavailable_response());
}
let raw_keys = list_admin_monitoring_namespaced_keys(state, pattern).await?; let raw_keys = list_admin_monitoring_namespaced_keys(state, pattern).await?;
let deleted_count = delete_admin_monitoring_namespaced_keys(state, &raw_keys).await?; let deleted_count = delete_admin_monitoring_namespaced_keys(state, &raw_keys).await?;
@@ -6,15 +6,10 @@ use super::super::cache_identity::{
admin_monitoring_find_user_summary_by_id, admin_monitoring_list_export_api_key_records_by_ids, admin_monitoring_find_user_summary_by_id, admin_monitoring_list_export_api_key_records_by_ids,
}; };
use super::super::cache_route_helpers::{ use super::super::cache_route_helpers::{
admin_monitoring_cache_affinity_unavailable_response,
admin_monitoring_cache_users_not_found_response, admin_monitoring_cache_users_not_found_response,
admin_monitoring_cache_users_user_identifier_from_path, admin_monitoring_cache_users_user_identifier_from_path,
}; };
use super::super::cache_store::{ use super::super::cache_store::list_admin_monitoring_cache_affinity_records_by_affinity_keys;
admin_monitoring_has_runtime_scheduler_affinity_entries,
list_admin_monitoring_cache_affinity_records_by_affinity_keys,
load_admin_monitoring_cache_affinity_entries_for_tests,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext}; use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::GatewayError; use crate::GatewayError;
use aether_admin::observability::monitoring::{ use aether_admin::observability::monitoring::{
@@ -36,13 +31,6 @@ pub(in super::super) async fn build_admin_monitoring_cache_users_delete_response
)); ));
}; };
if state.redis_kv_runner().is_none()
&& load_admin_monitoring_cache_affinity_entries_for_tests(state).is_empty()
&& !admin_monitoring_has_runtime_scheduler_affinity_entries(state)
{
return Ok(admin_monitoring_cache_affinity_unavailable_response());
}
let direct_api_key_by_id = admin_monitoring_list_export_api_key_records_by_ids( let direct_api_key_by_id = admin_monitoring_list_export_api_key_records_by_ids(
state, state,
std::slice::from_ref(&user_identifier), std::slice::from_ref(&user_identifier),
@@ -22,41 +22,6 @@ async fn count_admin_monitoring_cache_affinity_entries(state: &AdminAppState<'_>
}) })
} }
async fn scan_admin_monitoring_namespaced_keys(
runner: &aether_data::driver::redis::RedisKvRunner,
pattern: &str,
) -> Result<Vec<String>, GatewayError> {
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| {
GatewayError::Internal(format!("admin monitoring redis connect failed: {err}"))
})?;
let namespaced_pattern = runner.keyspace().key(pattern);
let mut cursor = 0u64;
let mut keys = Vec::new();
loop {
let (next_cursor, batch) = redis::cmd("SCAN")
.arg(cursor)
.arg("MATCH")
.arg(&namespaced_pattern)
.arg("COUNT")
.arg(200)
.query_async::<(u64, Vec<String>)>(&mut connection)
.await
.map_err(|err| {
GatewayError::Internal(format!("admin monitoring redis scan failed: {err}"))
})?;
keys.extend(batch);
if next_cursor == 0 {
break;
}
cursor = next_cursor;
}
Ok(keys)
}
#[cfg(test)] #[cfg(test)]
pub(super) fn load_admin_monitoring_cache_affinity_entries_for_tests( pub(super) fn load_admin_monitoring_cache_affinity_entries_for_tests(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
@@ -116,8 +81,13 @@ pub(super) async fn list_admin_monitoring_namespaced_keys(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
pattern: &str, pattern: &str,
) -> Result<Vec<String>, GatewayError> { ) -> Result<Vec<String>, GatewayError> {
if let Some(runner) = state.redis_kv_runner() { let keys = state
return scan_admin_monitoring_namespaced_keys(&runner, pattern).await; .runtime_state()
.scan_keys(pattern, 200)
.await
.map_err(|err| GatewayError::Internal(format!("runtime cache scan failed: {err}")))?;
if !keys.is_empty() {
return Ok(keys);
} }
let mut keys = load_admin_monitoring_redis_keys_for_tests(state) let mut keys = load_admin_monitoring_redis_keys_for_tests(state)
@@ -136,22 +106,13 @@ pub(super) async fn delete_admin_monitoring_namespaced_keys(
return Ok(0); return Ok(0);
} }
if let Some(runner) = state.redis_kv_runner() { let deleted = state
let mut connection = runner .runtime_state()
.client() .kv_delete_many(raw_keys)
.get_multiplexed_async_connection()
.await .await
.map_err(|err| { .map_err(|err| GatewayError::Internal(format!("runtime cache delete failed: {err}")))?;
GatewayError::Internal(format!("admin monitoring redis connect failed: {err}")) if deleted > 0 {
})?; return Ok(deleted);
let deleted = redis::cmd("DEL")
.arg(raw_keys)
.query_async::<i64>(&mut connection)
.await
.map_err(|err| {
GatewayError::Internal(format!("admin monitoring redis delete failed: {err}"))
})?;
return Ok(usize::try_from(deleted).unwrap_or(0));
} }
Ok(delete_admin_monitoring_redis_keys_for_tests( Ok(delete_admin_monitoring_redis_keys_for_tests(
@@ -197,62 +158,45 @@ async fn list_admin_monitoring_cache_affinity_records_matching(
} }
}; };
if let Some(runner) = state.redis_kv_runner() { {
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| {
GatewayError::Internal(format!("admin monitoring redis connect failed: {err}"))
})?;
let patterns = affinity_keys let patterns = affinity_keys
.map(|keys| { .map(|keys| {
keys.iter() keys.iter()
.flat_map(|affinity_key| { .flat_map(|affinity_key| {
[ [
runner format!("cache_affinity:{affinity_key}:*"),
.keyspace() format!("scheduler_affinity:{affinity_key}:*"),
.key(&format!("cache_affinity:{affinity_key}:*")), format!("scheduler_affinity:v2:{affinity_key}:*"),
runner
.keyspace()
.key(&format!("scheduler_affinity:{affinity_key}:*")),
runner
.keyspace()
.key(&format!("scheduler_affinity:v2:{affinity_key}:*")),
] ]
}) })
.collect::<Vec<_>>() .collect::<Vec<_>>()
}) })
.unwrap_or_else(|| { .unwrap_or_else(|| {
vec![ vec![
runner.keyspace().key("cache_affinity:*"), "cache_affinity:*".to_string(),
runner.keyspace().key("scheduler_affinity:*"), "scheduler_affinity:*".to_string(),
] ]
}); });
for pattern in patterns { for pattern in patterns {
let mut cursor = 0u64; let keys = state
loop { .runtime_state()
let (next_cursor, keys) = redis::cmd("SCAN") .scan_keys(&pattern, 200)
.arg(cursor)
.arg("MATCH")
.arg(&pattern)
.arg("COUNT")
.arg(200)
.query_async::<(u64, Vec<String>)>(&mut connection)
.await .await
.map_err(|err| { .map_err(|err| {
GatewayError::Internal(format!("admin monitoring redis scan failed: {err}")) GatewayError::Internal(format!("runtime cache scan failed: {err}"))
})?; })?;
if !keys.is_empty() { if !keys.is_empty() {
let values = redis::cmd("MGET") let raw_keys = keys
.arg(&keys) .iter()
.query_async::<Vec<Option<String>>>(&mut connection) .map(|key| state.runtime_state().strip_namespace(key).to_string())
.collect::<Vec<_>>();
let values = state
.runtime_state()
.kv_get_many(&raw_keys)
.await .await
.map_err(|err| { .map_err(|err| {
GatewayError::Internal(format!( GatewayError::Internal(format!("runtime cache mget failed: {err}"))
"admin monitoring redis mget failed: {err}"
))
})?; })?;
for (key, raw_value) in keys.into_iter().zip(values) { for (key, raw_value) in keys.into_iter().zip(values) {
let Some(raw_value) = raw_value else { let Some(raw_value) = raw_value else {
@@ -272,11 +216,6 @@ async fn list_admin_monitoring_cache_affinity_records_matching(
push_record(record); push_record(record);
} }
} }
if next_cursor == 0 {
break;
}
cursor = next_cursor;
}
} }
} }
@@ -350,11 +289,7 @@ pub(super) async fn build_admin_monitoring_cache_snapshot(
round_to(cache_hits as f64 / usage_summary.total_requests as f64, 4) round_to(cache_hits as f64 / usage_summary.total_requests as f64, 4)
}; };
let total_affinities = count_admin_monitoring_cache_affinity_entries(state).await; let total_affinities = count_admin_monitoring_cache_affinity_entries(state).await;
let storage_type = if state.redis_kv_runner().is_some() { let storage_type = state.runtime_state().backend_kind().as_str();
"redis"
} else {
"memory"
};
let scheduler_name = if scheduling_mode == "cache_affinity" { let scheduler_name = if scheduling_mode == "cache_affinity" {
"cache_aware".to_string() "cache_aware".to_string()
} else { } else {
@@ -1,4 +1,3 @@
use super::super::cache_config::ADMIN_MONITORING_REDIS_REQUIRED_DETAIL;
use super::super::test_support::{request_context, sample_key, sample_provider, sample_usage}; use super::super::test_support::{request_context, sample_key, sample_provider, sample_usage};
use super::local_monitoring_response; use super::local_monitoring_response;
use crate::AppState; use crate::AppState;
@@ -69,7 +68,7 @@ fn admin_monitoring_matches_cache_delete_shapes_and_trailing_slashes() {
} }
#[tokio::test] #[tokio::test]
async fn admin_monitoring_model_mapping_delete_requires_redis_without_runtime_or_test_entries() { async fn admin_monitoring_model_mapping_delete_returns_empty_runtime_payload_without_entries() {
let state = AppState::new().expect("state should build"); let state = AppState::new().expect("state should build");
let context = request_context( let context = request_context(
http::Method::DELETE, http::Method::DELETE,
@@ -80,15 +79,13 @@ async fn admin_monitoring_model_mapping_delete_requires_redis_without_runtime_or
.expect("handler should not error") .expect("handler should not error")
.expect("monitoring route should be handled locally"); .expect("monitoring route should be handled locally");
assert_eq!(response.status(), http::StatusCode::SERVICE_UNAVAILABLE); assert_eq!(response.status(), http::StatusCode::OK);
let body = to_bytes(response.into_body(), usize::MAX) let body = to_bytes(response.into_body(), usize::MAX)
.await .await
.expect("body should read"); .expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!( assert_eq!(payload["status"], json!("ok"));
payload, assert_eq!(payload["deleted_count"], json!(0));
json!({ "detail": ADMIN_MONITORING_REDIS_REQUIRED_DETAIL })
);
} }
#[tokio::test] #[tokio::test]
@@ -1086,11 +1086,9 @@ async fn admin_monitoring_model_mapping_stats_returns_local_payload_without_redi
.expect("body should read"); .expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(payload["status"], json!("ok")); assert_eq!(payload["status"], json!("ok"));
assert_eq!(payload["data"]["available"], json!(false)); assert_eq!(payload["data"]["available"], json!(true));
assert_eq!( assert_eq!(payload["data"]["backend"], json!("memory"));
payload["data"]["message"], assert_eq!(payload["data"]["total_keys"], json!(0));
json!("Redis 未启用,模型映射缓存不可用")
);
} }
#[tokio::test] #[tokio::test]
@@ -1193,12 +1191,13 @@ async fn admin_monitoring_redis_keys_returns_local_payload_without_redis() {
.expect("body should read"); .expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(payload["status"], json!("ok")); assert_eq!(payload["status"], json!("ok"));
assert_eq!(payload["data"]["available"], json!(false)); assert_eq!(payload["data"]["available"], json!(true));
assert_eq!(payload["data"]["message"], json!("Redis 未启用")); assert_eq!(payload["data"]["backend"], json!("memory"));
assert_eq!(payload["data"]["total_keys"], json!(0));
} }
#[tokio::test] #[tokio::test]
async fn admin_monitoring_redis_keys_delete_returns_unavailable_without_redis() { async fn admin_monitoring_redis_keys_delete_returns_empty_runtime_payload_without_redis() {
let state = AppState::new().expect("state should build"); let state = AppState::new().expect("state should build");
let context = request_context( let context = request_context(
http::Method::DELETE, http::Method::DELETE,
@@ -1210,12 +1209,14 @@ async fn admin_monitoring_redis_keys_delete_returns_unavailable_without_redis()
.expect("handler should not error") .expect("handler should not error")
.expect("route should be handled locally"); .expect("route should be handled locally");
assert_eq!(response.status(), http::StatusCode::SERVICE_UNAVAILABLE); assert_eq!(response.status(), http::StatusCode::OK);
let body = to_bytes(response.into_body(), usize::MAX) let body = to_bytes(response.into_body(), usize::MAX)
.await .await
.expect("body should read"); .expect("body should read");
let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse"); let payload: serde_json::Value = serde_json::from_slice(&body).expect("json body should parse");
assert_eq!(payload["detail"], json!("Redis 未启用")); assert_eq!(payload["status"], json!("ok"));
assert_eq!(payload["category"], json!("upstream_models"));
assert_eq!(payload["deleted_count"], json!(0));
} }
#[tokio::test] #[tokio::test]
@@ -49,27 +49,11 @@ pub(super) async fn read_admin_provider_ops_balance_cache(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
provider_id: &str, provider_id: &str,
) -> AdminProviderOpsBalanceCacheLookup { ) -> AdminProviderOpsBalanceCacheLookup {
let Some(runner) = state.redis_kv_runner() else { let raw_key = format!("{ADMIN_PROVIDER_OPS_BALANCE_CACHE_PREFIX}{provider_id}");
return AdminProviderOpsBalanceCacheLookup::Unavailable; let raw = match state.runtime_state().kv_get(&raw_key).await {
};
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
{
Ok(raw) => raw, Ok(raw) => raw,
Err(err) => { 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; 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 { let Some(ttl_seconds) = balance_cache_ttl_seconds(payload) else {
return; return;
}; };
let Some(runner) = state.redis_kv_runner() else {
return;
};
let serialized = match serde_json::to_string(payload) { let serialized = match serde_json::to_string(payload) {
Ok(serialized) => serialized, Ok(serialized) => serialized,
Err(err) => { Err(err) => {
@@ -107,11 +88,12 @@ pub(super) async fn store_admin_provider_ops_balance_cache(
return; return;
} }
}; };
if let Err(err) = runner if let Err(err) = state
.setex( .runtime_state()
.kv_set(
&format!("{ADMIN_PROVIDER_OPS_BALANCE_CACHE_PREFIX}{provider_id}"), &format!("{ADMIN_PROVIDER_OPS_BALANCE_CACHE_PREFIX}{provider_id}"),
&serialized, serialized,
Some(ttl_seconds), Some(Duration::from_secs(ttl_seconds)),
) )
.await .await
{ {
@@ -123,11 +105,9 @@ pub(super) async fn clear_admin_provider_ops_balance_cache(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
provider_id: &str, provider_id: &str,
) { ) {
let Some(runner) = state.redis_kv_runner() else { if let Err(err) = state
return; .runtime_state()
}; .kv_delete(&format!(
if let Err(err) = runner
.del(&format!(
"{ADMIN_PROVIDER_OPS_BALANCE_CACHE_PREFIX}{provider_id}" "{ADMIN_PROVIDER_OPS_BALANCE_CACHE_PREFIX}{provider_id}"
)) ))
.await .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 { 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}"); let raw_key = format!("{ADMIN_PROVIDER_OPS_BALANCE_REFRESH_PREFIX}{provider_id}");
if let Some(runner) = state.redis_kv_runner() {
format!( format!(
"{:p}:{}", "{:p}:{}",
state.app(), state.app(),
runner.keyspace().key(raw_key.as_str()) state.runtime_state().namespace_key(raw_key.as_str())
) )
} else {
format!("{:p}:{raw_key}", state.app())
}
} }
async fn finish_refresh_provider(refresh_key: &str) { async fn finish_refresh_provider(refresh_key: &str) {
@@ -102,7 +102,9 @@ pub(super) async fn handle_admin_provider_ops_action(
cached cached
} }
AdminProviderOpsBalanceCacheLookup::Miss => { 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; spawn_admin_provider_ops_balance_refresh(state, provider_id).await;
admin_provider_ops_pending_balance_response("余额数据加载中,请稍后刷新") admin_provider_ops_pending_balance_response("余额数据加载中,请稍后刷新")
} else { } else {
@@ -86,8 +86,25 @@ pub(super) async fn handle_admin_provider_ops_batch_balance(
cached cached
} }
AdminProviderOpsBalanceCacheLookup::Miss => { AdminProviderOpsBalanceCacheLookup::Miss => {
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; spawn_admin_provider_ops_balance_refresh(state, &provider_id).await;
admin_provider_ops_pending_balance_response("余额数据加载中,请稍后刷新") admin_provider_ops_pending_balance_response(
"余额数据加载中,请稍后刷新",
)
}
} }
AdminProviderOpsBalanceCacheLookup::Unavailable => { AdminProviderOpsBalanceCacheLookup::Unavailable => {
let payload = admin_provider_ops_local_action_response( let payload = admin_provider_ops_local_action_response(
@@ -1,51 +1,33 @@
use aether_data::driver::redis::RedisKeyspace; pub(super) fn pool_sticky_pattern(provider_id: &str) -> String {
format!("ap:{provider_id}:sticky:*")
pub(super) fn pool_sticky_pattern(keyspace: &RedisKeyspace, provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:sticky:*"))
} }
pub(super) fn pool_sticky_key( pub(super) fn pool_sticky_key(provider_id: &str, session_token: &str) -> String {
keyspace: &RedisKeyspace, format!("ap:{provider_id}:sticky:{session_token}")
provider_id: &str,
session_token: &str,
) -> String {
keyspace.key(&format!("ap:{provider_id}:sticky:{session_token}"))
} }
pub(super) fn pool_lru_key(keyspace: &RedisKeyspace, provider_id: &str) -> String { pub(super) fn pool_lru_key(provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:lru")) format!("ap:{provider_id}:lru")
} }
pub(super) fn pool_cooldown_key( pub(super) fn pool_cooldown_key(provider_id: &str, key_id: &str) -> String {
keyspace: &RedisKeyspace, format!("ap:{provider_id}:cooldown:{key_id}")
provider_id: &str,
key_id: &str,
) -> String {
keyspace.key(&format!("ap:{provider_id}:cooldown:{key_id}"))
} }
pub(super) fn pool_cooldown_index_key(keyspace: &RedisKeyspace, provider_id: &str) -> String { pub(super) fn pool_cooldown_index_key(provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:cooldown_idx")) format!("ap:{provider_id}:cooldown_idx")
} }
pub(super) fn pool_cost_key(keyspace: &RedisKeyspace, provider_id: &str, key_id: &str) -> String { pub(super) fn pool_cost_key(provider_id: &str, key_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:cost:{key_id}")) format!("ap:{provider_id}:cost:{key_id}")
} }
pub(super) fn pool_latency_key( pub(super) fn pool_latency_key(provider_id: &str, key_id: &str) -> String {
keyspace: &RedisKeyspace, format!("ap:{provider_id}:latency:{key_id}")
provider_id: &str,
key_id: &str,
) -> String {
keyspace.key(&format!("ap:{provider_id}:latency:{key_id}"))
} }
pub(super) fn pool_stream_timeout_key( pub(super) fn pool_stream_timeout_key(provider_id: &str, key_id: &str) -> String {
keyspace: &RedisKeyspace, format!("ap:{provider_id}:stream_timeout:{key_id}")
provider_id: &str,
key_id: &str,
) -> String {
keyspace.key(&format!("ap:{provider_id}:stream_timeout:{key_id}"))
} }
pub(super) fn parse_pool_cost_member(member: &str) -> u64 { 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) .unwrap_or(0)
} }
pub(super) fn pool_cooldown_keys( pub(super) fn pool_cooldown_keys(provider_id: &str, key_ids: &[String]) -> Vec<String> {
keyspace: &RedisKeyspace,
provider_id: &str,
key_ids: &[String],
) -> Vec<String> {
key_ids key_ids
.iter() .iter()
.map(|key_id| pool_cooldown_key(keyspace, provider_id, key_id)) .map(|key_id| pool_cooldown_key(provider_id, key_id))
.collect() .collect()
} }
pub(super) fn pool_cost_keys( pub(super) fn pool_cost_keys(provider_id: &str, key_ids: &[String]) -> Vec<String> {
keyspace: &RedisKeyspace,
provider_id: &str,
key_ids: &[String],
) -> Vec<String> {
key_ids key_ids
.iter() .iter()
.map(|key_id| pool_cost_key(keyspace, provider_id, key_id)) .map(|key_id| pool_cost_key(provider_id, key_id))
.collect() .collect()
} }
pub(super) fn pool_latency_keys( pub(super) fn pool_latency_keys(provider_id: &str, key_ids: &[String]) -> Vec<String> {
keyspace: &RedisKeyspace,
provider_id: &str,
key_ids: &[String],
) -> Vec<String> {
key_ids key_ids
.iter() .iter()
.map(|key_id| pool_latency_key(keyspace, provider_id, key_id)) .map(|key_id| pool_latency_key(provider_id, key_id))
.collect() .collect()
} }
@@ -1,29 +1,18 @@
use super::keys::{pool_cooldown_index_key, pool_cooldown_key}; use super::keys::{pool_cooldown_index_key, pool_cooldown_key};
use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::request::AdminAppState;
use tracing::warn;
pub(crate) async fn clear_admin_provider_pool_cooldown( pub(crate) async fn clear_admin_provider_pool_cooldown(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
provider_id: &str, provider_id: &str,
key_id: &str, key_id: &str,
) { ) {
let Some(runner) = state.redis_kv_runner() else { let _ = state
return; .runtime_state()
}; .kv_delete(&pool_cooldown_key(provider_id, key_id))
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else { .await;
warn!("gateway admin provider pool: failed to connect redis to clear cooldown for key {key_id}"); let _ = state
return; .runtime_state()
}; .set_remove(&pool_cooldown_index_key(provider_id), key_id)
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)
.await; .await;
} }
@@ -32,8 +21,8 @@ pub(crate) async fn reset_admin_provider_pool_cost(
provider_id: &str, provider_id: &str,
key_id: &str, key_id: &str,
) { ) {
let Some(runner) = state.redis_kv_runner() else { let _ = state
return; .runtime_state()
}; .score_remove_by_score(&format!("ap:{provider_id}:cost:{key_id}"), f64::INFINITY)
let _ = runner.del(&format!("ap:{provider_id}:cost:{key_id}")).await; .await;
} }
@@ -4,10 +4,9 @@ use super::keys::{
pool_sticky_pattern, pool_sticky_pattern,
}; };
use crate::handlers::admin::provider::shared::support::{ use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolRuntimeState, ADMIN_PROVIDER_POOL_SCAN_BATCH, AdminProviderPoolConfig, AdminProviderPoolRuntimeState,
}; };
use crate::GatewayError; use aether_runtime_state::RuntimeState;
use aether_data::driver::redis::RedisKvRunner;
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH}; use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn; use tracing::warn;
@@ -19,165 +18,78 @@ fn current_unix_secs() -> u64 {
.as_secs() .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( pub(crate) async fn read_admin_provider_pool_cooldown_counts(
runner: &RedisKvRunner, runtime: &RuntimeState,
provider_ids: &[String], provider_ids: &[String],
) -> BTreeMap<String, usize> { ) -> BTreeMap<String, usize> {
if provider_ids.is_empty() { let mut counts = BTreeMap::new();
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();
for provider_id in provider_ids { for provider_id in provider_ids {
pipeline let count = runtime
.cmd("SCARD") .set_len(&pool_cooldown_index_key(provider_id))
.arg(pool_cooldown_index_key(&keyspace, provider_id)); .await
} .unwrap_or(0);
counts.insert(provider_id.clone(), count);
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()
}
} }
counts
} }
pub(crate) async fn read_admin_provider_pool_runtime_state( pub(crate) async fn read_admin_provider_pool_runtime_state(
runner: &RedisKvRunner, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
key_ids: &[String], key_ids: &[String],
pool_config: &AdminProviderPoolConfig, pool_config: &AdminProviderPoolConfig,
sticky_session_token: Option<&str>, sticky_session_token: Option<&str>,
) -> AdminProviderPoolRuntimeState { ) -> AdminProviderPoolRuntimeState {
let mut runtime = AdminProviderPoolRuntimeState::default(); let mut state = AdminProviderPoolRuntimeState::default();
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else { let cooldown_keys = pool_cooldown_keys(provider_id, key_ids);
warn!("gateway admin provider pool: failed to connect redis for provider {provider_id}"); let cost_keys = pool_cost_keys(provider_id, key_ids);
return runtime; let latency_keys = pool_latency_keys(provider_id, key_ids);
};
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);
if let Some(sticky_session_token) = sticky_session_token if let Some(sticky_session_token) = sticky_session_token
.map(str::trim) .map(str::trim)
.filter(|value| !value.is_empty()) .filter(|value| !value.is_empty())
.filter(|_| pool_config.sticky_session_ttl_seconds > 0) .filter(|_| pool_config.sticky_session_ttl_seconds > 0)
{ {
let sticky_key = pool_sticky_key(&keyspace, provider_id, sticky_session_token); let sticky_key = pool_sticky_key(provider_id, sticky_session_token);
let sticky_bound_key_id = redis::cmd("GET") if let Ok(Some(bound_key_id)) = runtime.kv_get(&sticky_key).await {
.arg(&sticky_key) let cooldown_key = pool_cooldown_key(provider_id, &bound_key_id);
.query_async::<Option<String>>(&mut connection) match runtime.kv_exists(&cooldown_key).await {
.await Ok(false) => {
.unwrap_or_else(|err| { let _ = runtime
warn!( .key_expire(
"gateway admin provider pool: failed to read sticky binding for provider {provider_id}: {:?}", &sticky_key,
err std::time::Duration::from_secs(pool_config.sticky_session_ttl_seconds),
); )
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)
.await; .await;
Some(bound_key_id) state.sticky_bound_key_id = Some(bound_key_id);
} }
Ok(_) => { Ok(true) => {
let _: Result<i64, _> = redis::cmd("DEL") let _ = runtime.kv_delete(&sticky_key).await;
.arg(&sticky_key)
.query_async(&mut connection)
.await;
None
} }
Err(err) => { Err(err) => {
warn!( warn!(
"gateway admin provider pool: failed to validate sticky cooldown for provider {provider_id}: {:?}", "gateway admin provider pool: failed to validate sticky cooldown for provider {provider_id}: {:?}",
err err
); );
Some(bound_key_id) state.sticky_bound_key_id = Some(bound_key_id);
}
} }
};
} }
} }
let sticky_keys = match scan_redis_keys( let sticky_keys = runtime
&mut connection, .scan_keys(&pool_sticky_pattern(provider_id), 200)
&pool_sticky_pattern(&keyspace, provider_id),
)
.await .await
{ .unwrap_or_default();
Ok(keys) => keys, state.total_sticky_sessions = sticky_keys.len();
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();
if !sticky_keys.is_empty() { if !sticky_keys.is_empty() {
for chunk in sticky_keys.chunks(ADMIN_PROVIDER_POOL_SCAN_BATCH as usize) { let raw_keys = sticky_keys
let values = redis::cmd("MGET") .iter()
.arg(chunk) .map(|key| runtime.strip_namespace(key).to_string())
.query_async::<Vec<Option<String>>>(&mut connection) .collect::<Vec<_>>();
.await; if let Ok(values) = runtime.kv_get_many(&raw_keys).await {
let Ok(values) = values else {
warn!(
"gateway admin provider pool: failed to read sticky bindings for provider {provider_id}"
);
break;
};
for bound_key_id in values.into_iter().flatten() { for bound_key_id in values.into_iter().flatten() {
*runtime *state
.sticky_sessions_by_key .sticky_sessions_by_key
.entry(bound_key_id) .entry(bound_key_id)
.or_insert(0) += 1; .or_insert(0) += 1;
@@ -186,45 +98,20 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
} }
if !cooldown_keys.is_empty() { if !cooldown_keys.is_empty() {
let cooldown_reasons = redis::cmd("MGET") let cooldown_reasons = runtime
.arg(&cooldown_keys) .kv_get_many(&cooldown_keys)
.query_async::<Vec<Option<String>>>(&mut connection)
.await .await
.unwrap_or_else(|err| { .unwrap_or_else(|_| vec![None; cooldown_keys.len()]);
warn!( for (key_id, (cooldown_key, reason)) in key_ids
"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
.iter() .iter()
.zip(cooldown_keys.iter()) .zip(cooldown_keys.iter().zip(cooldown_reasons))
.zip(cooldown_reasons)
.zip(cooldown_ttls)
{ {
if let Some(reason) = reason { if let Some(reason) = reason {
runtime state.cooldown_reason_by_key.insert(key_id.clone(), reason);
.cooldown_reason_by_key if let Ok(Some(ttl)) = runtime.kv_ttl_seconds(cooldown_key).await {
.insert(key_id.clone(), reason);
if let Ok(ttl_seconds) = u64::try_from(ttl) { if let Ok(ttl_seconds) = u64::try_from(ttl) {
if ttl_seconds > 0 { if ttl_seconds > 0 {
runtime state
.cooldown_ttl_by_key .cooldown_ttl_by_key
.insert(key_id.clone(), ttl_seconds); .insert(key_id.clone(), ttl_seconds);
} }
@@ -232,60 +119,29 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
} }
} }
} }
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 .await
.unwrap_or_else(|err| { .unwrap_or_default()
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() .iter()
.map(|member| parse_pool_cost_member(member)) .map(|member| parse_pool_cost_member(member))
.sum::<u64>(); .sum::<u64>();
runtime if total > 0 {
.cost_window_usage_by_key state.cost_window_usage_by_key.insert(key_id.clone(), total);
.insert(key_id.clone(), total);
} }
} }
if !latency_keys.is_empty() { for (key_id, latency_key) in key_ids.iter().zip(latency_keys) {
let window_start = current_unix_secs().saturating_sub(pool_config.latency_window_seconds); let window_start = now.saturating_sub(pool_config.latency_window_seconds) as f64;
let mut latency_pipeline = redis::pipe(); let samples = runtime
for latency_key in &latency_keys { .score_range_by_min(&latency_key, window_start)
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)
.await .await
.unwrap_or_else(|err| { .unwrap_or_default()
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() .iter()
.map(|member| parse_pool_latency_member(member)) .map(|member| parse_pool_latency_member(member))
.filter(|value| *value > 0) .filter(|value| *value > 0)
@@ -296,10 +152,7 @@ pub(crate) async fn read_admin_provider_pool_runtime_state(
let total = samples.iter().sum::<u64>() as f64; let total = samples.iter().sum::<u64>() as f64;
let average = total / samples.len() as f64; let average = total / samples.len() as f64;
if average.is_finite() && average >= 0.0 { if average.is_finite() && average >= 0.0 {
runtime state.latency_avg_ms_by_key.insert(key_id.clone(), average);
.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)) .any(|item| item.enabled))
&& !key_ids.is_empty() && !key_ids.is_empty()
{ {
let mut command = redis::cmd("ZMSCORE"); if let Ok(scores) = runtime
command.arg(pool_lru_key(&keyspace, provider_id)); .score_many(&pool_lru_key(provider_id), key_ids)
for key_id in key_ids {
command.arg(key_id);
}
if let Ok(scores) = command
.query_async::<Vec<Option<f64>>>(&mut connection)
.await .await
{ {
for (key_id, score) in key_ids.iter().zip(scores) { for (key_id, score) in key_ids.iter().zip(scores) {
if let Some(score) = score { 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( pub(crate) async fn read_admin_provider_pool_cooldown_count(
runner: &RedisKvRunner, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
) -> usize { ) -> usize {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else { runtime
warn!("gateway admin provider pool: failed to connect redis for provider {provider_id}"); .set_len(&pool_cooldown_index_key(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)
.await .await
.map(|value| value as usize)
.unwrap_or(0) .unwrap_or(0)
} }
pub(crate) async fn read_admin_provider_pool_cooldown_key_ids( pub(crate) async fn read_admin_provider_pool_cooldown_key_ids(
runner: &RedisKvRunner, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
) -> Vec<String> { ) -> Vec<String> {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else { runtime
warn!("gateway admin provider pool: failed to connect redis for provider {provider_id}"); .set_members(&pool_cooldown_index_key(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)
.await .await
.unwrap_or_default() .unwrap_or_default()
} }
@@ -1,6 +1,5 @@
use super::reads::read_admin_provider_pool_runtime_state; 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::pool::config::admin_provider_pool_config;
use crate::handlers::admin::provider::shared::support::AdminProviderPoolRuntimeState;
use crate::handlers::admin::request::AdminAppState; use crate::handlers::admin::request::AdminAppState;
use serde_json::json; use serde_json::json;
@@ -34,19 +33,14 @@ pub(crate) async fn build_admin_provider_pool_status_payload(
.ok() .ok()
.unwrap_or_default(); .unwrap_or_default();
let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>(); let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
let runtime = match state.redis_kv_runner() { let runtime = read_admin_provider_pool_runtime_state(
Some(runner) => { state.runtime_state(),
read_admin_provider_pool_runtime_state(
&runner,
&provider.id, &provider.id,
&key_ids, &key_ids,
&pool_config, &pool_config,
None, None,
) )
.await .await;
}
None => AdminProviderPoolRuntimeState::default(),
};
let key_payloads = keys let key_payloads = keys
.into_iter() .into_iter()
.map(|key| { .map(|key| {
@@ -5,7 +5,7 @@ use super::keys::{
use crate::handlers::admin::provider::shared::support::{ use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolUnschedulableRule, AdminProviderPoolConfig, AdminProviderPoolUnschedulableRule,
}; };
use aether_data::driver::redis::RedisKvRunner; use aether_runtime_state::RuntimeState;
use regex::Regex; use regex::Regex;
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH}; use std::time::{SystemTime, UNIX_EPOCH};
@@ -277,7 +277,7 @@ fn resolve_transient_cooldown_ttl(
} }
async fn set_pool_cooldown( async fn set_pool_cooldown(
runner: &RedisKvRunner, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
key_id: &str, key_id: &str,
reason: &str, reason: &str,
@@ -288,39 +288,32 @@ async fn set_pool_cooldown(
} }
let ttl_seconds = ttl_seconds.min(MAX_POOL_COOLDOWN_SECONDS); let ttl_seconds = ttl_seconds.min(MAX_POOL_COOLDOWN_SECONDS);
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else { if let Err(err) = runtime
warn!( .kv_set(
"gateway admin provider pool: failed to connect redis to set cooldown for key {key_id}" &pool_cooldown_key(provider_id, key_id),
); reason.to_string(),
return; Some(std::time::Duration::from_secs(ttl_seconds)),
}; )
let keyspace = runner.keyspace().clone(); .await
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 {
warn!( warn!(
"gateway admin provider pool: failed to set cooldown for provider {provider_id} key {key_id}: {:?}", "gateway admin provider pool: failed to set cooldown for provider {provider_id} key {key_id}: {:?}",
err 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) { async fn invalidate_pool_oauth_cache(runtime: &RuntimeState, key_id: &str) {
if let Err(err) = runner.del(&oauth_cache_key(key_id)).await { if let Err(err) = runtime.kv_delete(&oauth_cache_key(key_id)).await {
warn!( warn!(
"gateway admin provider pool: failed to invalidate oauth cache for key {key_id}: {:?}", "gateway admin provider pool: failed to invalidate oauth cache for key {key_id}: {:?}",
err err
@@ -339,7 +332,7 @@ fn matching_unschedulable_rule<'a>(
} }
pub(crate) async fn record_admin_provider_pool_success( pub(crate) async fn record_admin_provider_pool_success(
runner: &RedisKvRunner, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
key_id: &str, key_id: &str,
pool_config: &AdminProviderPoolConfig, pool_config: &AdminProviderPoolConfig,
@@ -347,111 +340,72 @@ pub(crate) async fn record_admin_provider_pool_success(
tokens_used: u64, tokens_used: u64,
ttfb_ms: Option<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 now = current_unix_secs_f64();
let mut pipeline = redis::pipe();
let mut has_commands = false;
if let Some(sticky_session_token) = sticky_session_token if let Some(sticky_session_token) = sticky_session_token
.map(str::trim) .map(str::trim)
.filter(|value| !value.is_empty()) .filter(|value| !value.is_empty())
.filter(|_| pool_config.sticky_session_ttl_seconds > 0) .filter(|_| pool_config.sticky_session_ttl_seconds > 0)
{ {
pipeline let _ = runtime
.cmd("SETEX") .kv_set(
.arg(pool_sticky_key( &pool_sticky_key(provider_id, sticky_session_token),
&keyspace, key_id.to_string(),
provider_id, Some(std::time::Duration::from_secs(
sticky_session_token, pool_config.sticky_session_ttl_seconds,
)) )),
.arg(pool_config.sticky_session_ttl_seconds) )
.arg(key_id) .await;
.ignore();
has_commands = true;
} }
if should_touch_lru(pool_config) { if should_touch_lru(pool_config) {
pipeline let _ = runtime
.cmd("ZADD") .score_set(&pool_lru_key(provider_id), key_id, now)
.arg(pool_lru_key(&keyspace, provider_id)) .await;
.arg(now)
.arg(key_id)
.ignore();
has_commands = true;
} }
if tokens_used > 0 && pool_config.cost_limit_per_key_tokens.is_some() { 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 window_seconds = pool_config.cost_window_seconds.max(1);
let member = format!("{}:{tokens_used}", Uuid::new_v4().simple()); let member = format!("{}:{tokens_used}", Uuid::new_v4().simple());
pipeline let _ = runtime.score_set(&cost_key, &member, now).await;
.cmd("ZADD") let _ = runtime
.arg(&cost_key) .score_remove_by_score(&cost_key, now - window_seconds as f64)
.arg(now) .await;
.arg(member) let _ = runtime
.ignore() .key_expire(
.cmd("ZREMRANGEBYSCORE") &cost_key,
.arg(&cost_key) std::time::Duration::from_secs(window_seconds.saturating_add(600)),
.arg("-inf") )
.arg(now - window_seconds as f64) .await;
.ignore()
.cmd("EXPIRE")
.arg(&cost_key)
.arg(window_seconds.saturating_add(600))
.ignore();
has_commands = true;
} }
if let Some(ttfb_ms) = ttfb_ms if let Some(ttfb_ms) = ttfb_ms
.filter(|value| should_record_latency(pool_config)) .filter(|value| should_record_latency(pool_config))
.filter(|_| pool_config.latency_window_seconds > 0) .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 window_seconds = pool_config.latency_window_seconds.max(1);
let sample_limit = pool_config.latency_sample_limit.max(1); let sample_limit = pool_config.latency_sample_limit.max(1);
let member = format!("{}:{ttfb_ms}", Uuid::new_v4().simple()); let member = format!("{}:{ttfb_ms}", Uuid::new_v4().simple());
pipeline let _ = runtime.score_set(&latency_key, &member, now).await;
.cmd("ZADD") let _ = runtime
.arg(&latency_key) .score_remove_by_score(&latency_key, now - window_seconds as f64)
.arg(now) .await;
.arg(member) let _ = runtime
.ignore() .score_remove_by_rank(&latency_key, 0, -((sample_limit as i64) + 1))
.cmd("ZREMRANGEBYSCORE") .await;
.arg(&latency_key) let _ = runtime
.arg("-inf") .key_expire(
.arg(now - window_seconds as f64) &latency_key,
.ignore() std::time::Duration::from_secs(window_seconds.saturating_add(600)),
.cmd("ZREMRANGEBYRANK") )
.arg(&latency_key) .await;
.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
);
} }
} }
pub(crate) async fn record_admin_provider_pool_error( pub(crate) async fn record_admin_provider_pool_error(
runner: &RedisKvRunner, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
key_id: &str, key_id: &str,
pool_config: &AdminProviderPoolConfig, 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(); let error_message = extract_error_message(error_body).to_ascii_lowercase();
if status_code == 401 { if status_code == 401 {
invalidate_pool_oauth_cache(runner, key_id).await; invalidate_pool_oauth_cache(runtime, key_id).await;
return; return;
} }
@@ -482,7 +436,7 @@ pub(crate) async fn record_admin_provider_pool_error(
return; return;
} }
set_pool_cooldown( set_pool_cooldown(
runner, runtime,
provider_id, provider_id,
key_id, key_id,
"forbidden_403", "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); let ttl_seconds = (rule.duration_minutes.max(1)).saturating_mul(60).max(60);
set_pool_cooldown( set_pool_cooldown(
runner, runtime,
provider_id, provider_id,
key_id, key_id,
&format!("rule:{}", rule.keyword), &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)), .or_else(|| parse_google_quota_cooldown_seconds(error_body)),
pool_config, 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; return;
} }
if status_code == 529 { if status_code == 529 {
set_pool_cooldown( set_pool_cooldown(
runner, runtime,
provider_id, provider_id,
key_id, key_id,
"overloaded_529", "overloaded_529",
@@ -555,12 +516,12 @@ pub(crate) async fn record_admin_provider_pool_error(
parse_retry_after_seconds(response_headers), parse_retry_after_seconds(response_headers),
pool_config, 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( pub(crate) async fn record_admin_provider_pool_stream_timeout(
runner: &RedisKvRunner, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
key_id: &str, key_id: &str,
pool_config: &AdminProviderPoolConfig, pool_config: &AdminProviderPoolConfig,
@@ -569,49 +530,25 @@ pub(crate) async fn record_admin_provider_pool_stream_timeout(
return; return;
} }
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else { let timeout_key = pool_stream_timeout_key(provider_id, key_id);
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 now = current_unix_secs_f64(); let now = current_unix_secs_f64();
let window_seconds = pool_config.stream_timeout_window_seconds.max(1); let window_seconds = pool_config.stream_timeout_window_seconds.max(1);
let member = Uuid::new_v4().simple().to_string(); let member = Uuid::new_v4().simple().to_string();
let results = redis::pipe() let _ = runtime
.cmd("ZREMRANGEBYSCORE") .score_remove_by_score(&timeout_key, now - window_seconds as f64)
.arg(&timeout_key) .await;
.arg("-inf") let _ = runtime.score_set(&timeout_key, &member, now).await;
.arg(now - window_seconds as f64) let count = runtime.score_len(&timeout_key).await.unwrap_or(0) as u64;
.cmd("ZADD") let _ = runtime
.arg(&timeout_key) .key_expire(
.arg(now) &timeout_key,
.arg(member) std::time::Duration::from_secs(window_seconds.saturating_add(60)),
.cmd("ZCARD") )
.arg(&timeout_key)
.cmd("EXPIRE")
.arg(&timeout_key)
.arg(window_seconds.saturating_add(60))
.query_async::<Vec<redis::Value>>(&mut connection)
.await; .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 { if count >= pool_config.stream_timeout_threshold {
set_pool_cooldown( set_pool_cooldown(
runner, runtime,
provider_id, provider_id,
key_id, key_id,
&format!("stream_timeout_x{count}"), &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_error, record_admin_provider_pool_stream_timeout,
record_admin_provider_pool_success, 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::pool::runtime::reads::read_admin_provider_pool_runtime_state;
use crate::handlers::admin::provider::shared::support::{ use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolSchedulingPreset, AdminProviderPoolConfig, AdminProviderPoolSchedulingPreset,
AdminProviderPoolUnschedulableRule, AdminProviderPoolUnschedulableRule,
}; };
use crate::AppState; use crate::AppState;
use aether_runtime_state::{RedisClientConfig, RuntimeState, RuntimeStateConfig};
use aether_testkit::ManagedRedisServer; use aether_testkit::ManagedRedisServer;
use std::collections::BTreeMap; use std::collections::BTreeMap;
@@ -682,14 +619,17 @@ mod tests {
} }
} }
fn build_runner_app(redis_url: &str, key_prefix: &str) -> AppState { async fn build_runner_app(redis_url: &str, key_prefix: &str) -> AppState {
let data_state = GatewayDataState::from_config( let runtime_state =
GatewayDataConfig::disabled().with_redis_url(redis_url, Some(key_prefix)), RuntimeState::from_config(RuntimeStateConfig::redis(RedisClientConfig {
) url: redis_url.to_string(),
.expect("data state should build"); key_prefix: Some(key_prefix.to_string()),
}))
.await
.expect("runtime state should build");
AppState::new() AppState::new()
.expect("app state should build") .expect("app state should build")
.with_data_state_for_tests(data_state) .with_runtime_state(std::sync::Arc::new(runtime_state))
} }
#[test] #[test]
@@ -798,13 +738,13 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else { let Some(redis) = start_managed_redis_or_skip().await else {
return; return;
}; };
let app = build_runner_app(redis.redis_url(), "pool_runtime_success_feedback"); let app = build_runner_app(redis.redis_url(), "pool_runtime_success_feedback").await;
let runner = app.redis_kv_runner().expect("redis runner should exist"); let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config(); let pool_config = sample_pool_config();
let key_ids = vec!["key-1".to_string()]; let key_ids = vec!["key-1".to_string()];
record_admin_provider_pool_success( record_admin_provider_pool_success(
&runner, runtime,
"provider-1", "provider-1",
"key-1", "key-1",
&pool_config, &pool_config,
@@ -815,7 +755,7 @@ mod tests {
.await; .await;
let runtime = read_admin_provider_pool_runtime_state( let runtime = read_admin_provider_pool_runtime_state(
&runner, runtime,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &pool_config,
@@ -836,14 +776,15 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else { let Some(redis) = start_managed_redis_or_skip().await else {
return; return;
}; };
let app = build_runner_app(redis.redis_url(), "pool_runtime_no_sticky_without_affinity"); let app =
let runner = app.redis_kv_runner().expect("redis runner should exist"); 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(); let mut pool_config = sample_pool_config();
pool_config.sticky_session_ttl_seconds = 0; pool_config.sticky_session_ttl_seconds = 0;
let key_ids = vec!["key-1".to_string()]; let key_ids = vec!["key-1".to_string()];
record_admin_provider_pool_success( record_admin_provider_pool_success(
&runner, runtime,
"provider-1", "provider-1",
"key-1", "key-1",
&pool_config, &pool_config,
@@ -854,7 +795,7 @@ mod tests {
.await; .await;
let runtime = read_admin_provider_pool_runtime_state( let runtime = read_admin_provider_pool_runtime_state(
&runner, runtime,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &pool_config,
@@ -874,13 +815,13 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else { let Some(redis) = start_managed_redis_or_skip().await else {
return; return;
}; };
let app = build_runner_app(redis.redis_url(), "pool_runtime_error_feedback"); let app = build_runner_app(redis.redis_url(), "pool_runtime_error_feedback").await;
let runner = app.redis_kv_runner().expect("redis runner should exist"); let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config(); let pool_config = sample_pool_config();
let key_ids = vec!["key-2".to_string()]; let key_ids = vec!["key-2".to_string()];
record_admin_provider_pool_error( record_admin_provider_pool_error(
&runner, runtime,
"provider-1", "provider-1",
"key-2", "key-2",
&pool_config, &pool_config,
@@ -894,7 +835,7 @@ mod tests {
.await; .await;
let runtime = read_admin_provider_pool_runtime_state( let runtime = read_admin_provider_pool_runtime_state(
&runner, runtime,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &pool_config,
@@ -920,13 +861,13 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else { let Some(redis) = start_managed_redis_or_skip().await else {
return; return;
}; };
let app = build_runner_app(redis.redis_url(), "pool_runtime_google_quota_cooldown"); let app = build_runner_app(redis.redis_url(), "pool_runtime_google_quota_cooldown").await;
let runner = app.redis_kv_runner().expect("redis runner should exist"); let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config(); let pool_config = sample_pool_config();
let key_ids = vec!["key-google-429".to_string()]; let key_ids = vec!["key-google-429".to_string()];
record_admin_provider_pool_error( record_admin_provider_pool_error(
&runner, runtime,
"provider-1", "provider-1",
"key-google-429", "key-google-429",
&pool_config, &pool_config,
@@ -949,7 +890,7 @@ mod tests {
.await; .await;
let runtime = read_admin_provider_pool_runtime_state( let runtime = read_admin_provider_pool_runtime_state(
&runner, runtime,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &pool_config,
@@ -975,13 +916,13 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else { let Some(redis) = start_managed_redis_or_skip().await else {
return; return;
}; };
let app = build_runner_app(redis.redis_url(), "pool_runtime_capped_cooldown"); let app = build_runner_app(redis.redis_url(), "pool_runtime_capped_cooldown").await;
let runner = app.redis_kv_runner().expect("redis runner should exist"); let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config(); let pool_config = sample_pool_config();
let key_ids = vec!["key-long-cooldown".to_string()]; let key_ids = vec!["key-long-cooldown".to_string()];
record_admin_provider_pool_error( record_admin_provider_pool_error(
&runner, runtime,
"provider-1", "provider-1",
"key-long-cooldown", "key-long-cooldown",
&pool_config, &pool_config,
@@ -995,7 +936,7 @@ mod tests {
.await; .await;
let runtime = read_admin_provider_pool_runtime_state( let runtime = read_admin_provider_pool_runtime_state(
&runner, runtime,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &pool_config,
@@ -1021,8 +962,8 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else { let Some(redis) = start_managed_redis_or_skip().await else {
return; return;
}; };
let app = build_runner_app(redis.redis_url(), "pool_runtime_circuit_no_cooldown"); let app = build_runner_app(redis.redis_url(), "pool_runtime_circuit_no_cooldown").await;
let runner = app.redis_kv_runner().expect("redis runner should exist"); let runtime = app.runtime_state.as_ref();
let pool_config = sample_pool_config(); let pool_config = sample_pool_config();
let key_ids = vec!["key-account-disabled".to_string()]; let key_ids = vec!["key-account-disabled".to_string()];
@@ -1035,7 +976,7 @@ mod tests {
Some("account_deactivated_401") Some("account_deactivated_401")
); );
record_admin_provider_pool_error( record_admin_provider_pool_error(
&runner, runtime,
"provider-1", "provider-1",
"key-account-disabled", "key-account-disabled",
&pool_config, &pool_config,
@@ -1046,7 +987,7 @@ mod tests {
.await; .await;
let runtime = read_admin_provider_pool_runtime_state( let runtime = read_admin_provider_pool_runtime_state(
&runner, runtime,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &pool_config,
@@ -1067,8 +1008,8 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else { let Some(redis) = start_managed_redis_or_skip().await else {
return; return;
}; };
let app = build_runner_app(redis.redis_url(), "pool_runtime_unschedulable_rule"); let app = build_runner_app(redis.redis_url(), "pool_runtime_unschedulable_rule").await;
let runner = app.redis_kv_runner().expect("redis runner should exist"); let runtime = app.runtime_state.as_ref();
let mut pool_config = sample_pool_config(); let mut pool_config = sample_pool_config();
pool_config.unschedulable_rules = vec![AdminProviderPoolUnschedulableRule { pool_config.unschedulable_rules = vec![AdminProviderPoolUnschedulableRule {
keyword: "review required".to_string(), keyword: "review required".to_string(),
@@ -1077,7 +1018,7 @@ mod tests {
let key_ids = vec!["key-3".to_string()]; let key_ids = vec!["key-3".to_string()];
record_admin_provider_pool_error( record_admin_provider_pool_error(
&runner, runtime,
"provider-1", "provider-1",
"key-3", "key-3",
&pool_config, &pool_config,
@@ -1088,7 +1029,7 @@ mod tests {
.await; .await;
let runtime = read_admin_provider_pool_runtime_state( let runtime = read_admin_provider_pool_runtime_state(
&runner, runtime,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &pool_config,
@@ -1114,8 +1055,8 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else { let Some(redis) = start_managed_redis_or_skip().await else {
return; return;
}; };
let app = build_runner_app(redis.redis_url(), "pool_runtime_ignore_400"); let app = build_runner_app(redis.redis_url(), "pool_runtime_ignore_400").await;
let runner = app.redis_kv_runner().expect("redis runner should exist"); let runtime = app.runtime_state.as_ref();
let mut pool_config = sample_pool_config(); let mut pool_config = sample_pool_config();
pool_config.unschedulable_rules = vec![AdminProviderPoolUnschedulableRule { pool_config.unschedulable_rules = vec![AdminProviderPoolUnschedulableRule {
keyword: "review required".to_string(), keyword: "review required".to_string(),
@@ -1124,7 +1065,7 @@ mod tests {
let key_ids = vec!["key-client-400".to_string()]; let key_ids = vec!["key-client-400".to_string()];
record_admin_provider_pool_error( record_admin_provider_pool_error(
&runner, runtime,
"provider-1", "provider-1",
"key-client-400", "key-client-400",
&pool_config, &pool_config,
@@ -1135,7 +1076,7 @@ mod tests {
.await; .await;
let runtime = read_admin_provider_pool_runtime_state( let runtime = read_admin_provider_pool_runtime_state(
&runner, runtime,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &pool_config,
@@ -1154,21 +1095,31 @@ mod tests {
let Some(redis) = start_managed_redis_or_skip().await else { let Some(redis) = start_managed_redis_or_skip().await else {
return; return;
}; };
let app = build_runner_app(redis.redis_url(), "pool_runtime_stream_timeout"); let app = build_runner_app(redis.redis_url(), "pool_runtime_stream_timeout").await;
let runner = app.redis_kv_runner().expect("redis runner should exist"); let runtime_state = app.runtime_state.as_ref();
let mut pool_config = sample_pool_config(); let mut pool_config = sample_pool_config();
pool_config.stream_timeout_threshold = 2; pool_config.stream_timeout_threshold = 2;
pool_config.stream_timeout_window_seconds = 300; pool_config.stream_timeout_window_seconds = 300;
pool_config.stream_timeout_cooldown_seconds = 90; pool_config.stream_timeout_cooldown_seconds = 90;
let key_ids = vec!["key-4".to_string()]; let key_ids = vec!["key-4".to_string()];
record_admin_provider_pool_stream_timeout(&runner, "provider-1", "key-4", &pool_config) record_admin_provider_pool_stream_timeout(
runtime_state,
"provider-1",
"key-4",
&pool_config,
)
.await; .await;
record_admin_provider_pool_stream_timeout(&runner, "provider-1", "key-4", &pool_config) record_admin_provider_pool_stream_timeout(
runtime_state,
"provider-1",
"key-4",
&pool_config,
)
.await; .await;
let mut runtime = read_admin_provider_pool_runtime_state( let mut runtime = read_admin_provider_pool_runtime_state(
&runner, runtime_state,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &pool_config,
@@ -1186,7 +1137,7 @@ mod tests {
} }
tokio::time::sleep(std::time::Duration::from_millis(10)).await; tokio::time::sleep(std::time::Duration::from_millis(10)).await;
runtime = read_admin_provider_pool_runtime_state( runtime = read_admin_provider_pool_runtime_state(
&runner, runtime_state,
"provider-1", "provider-1",
&key_ids, &key_ids,
&pool_config, &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 page_offset = page.saturating_sub(1).saturating_mul(page_size);
let (keys, total) = if status == "cooldown" { let (keys, total) = if status == "cooldown" {
let cooldown_key_ids = if let Some(runner) = state.redis_kv_runner() { let cooldown_key_ids =
read_admin_provider_pool_cooldown_key_ids(&runner, &provider.id).await read_admin_provider_pool_cooldown_key_ids(state.runtime_state(), &provider.id).await;
} else {
Vec::new()
};
let mut keys = if cooldown_key_ids.is_empty() { let mut keys = if cooldown_key_ids.is_empty() {
Vec::new() Vec::new()
} else { } else {
@@ -240,10 +237,10 @@ pub(super) async fn build_admin_pool_list_keys_response(
let endpoints = state let endpoints = state
.list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id)) .list_provider_catalog_endpoints_by_provider_ids(std::slice::from_ref(&provider.id))
.await?; .await?;
let runtime = match (state.redis_kv_runner(), pool_config.as_ref()) { let runtime = match pool_config.as_ref() {
(Some(runner), Some(pool_config)) if !key_ids.is_empty() => { Some(pool_config) if !key_ids.is_empty() => {
read_admin_provider_pool_runtime_state( read_admin_provider_pool_runtime_state(
&runner, state.runtime_state(),
&provider.id, &provider.id,
&key_ids, &key_ids,
pool_config, pool_config,
@@ -1,3 +1,5 @@
use std::collections::BTreeMap;
use super::{ use super::{
admin_provider_pool_config, build_admin_pool_error_response, admin_provider_pool_config, build_admin_pool_error_response,
read_admin_provider_pool_cooldown_counts, read_admin_provider_pool_cooldown_counts,
@@ -12,7 +14,6 @@ use axum::{
response::{IntoResponse, Response}, response::{IntoResponse, Response},
Json, Json,
}; };
use std::collections::BTreeMap;
pub(super) async fn build_admin_pool_overview_response( pub(super) async fn build_admin_pool_overview_response(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
@@ -35,7 +36,6 @@ pub(super) async fn build_admin_pool_overview_response(
.iter() .iter()
.map(|(provider, _)| provider.id.clone()) .map(|(provider, _)| provider.id.clone())
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let redis_runner = state.redis_kv_runner();
let (key_stats_result, cooldown_counts_by_provider) = tokio::join!( let (key_stats_result, cooldown_counts_by_provider) = tokio::join!(
async { async {
if provider_ids.is_empty() { if provider_ids.is_empty() {
@@ -47,11 +47,10 @@ pub(super) async fn build_admin_pool_overview_response(
} }
}, },
async { async {
match redis_runner.as_ref() { if provider_ids.is_empty() {
Some(runner) if !provider_ids.is_empty() => { std::collections::BTreeMap::new()
read_admin_provider_pool_cooldown_counts(runner, &provider_ids).await } else {
} read_admin_provider_pool_cooldown_counts(state.runtime_state(), &provider_ids).await
_ => BTreeMap::new(),
} }
}, },
); );
@@ -1553,20 +1553,8 @@ async fn provider_query_read_cached_models(
provider_id: &str, provider_id: &str,
key_id: &str, key_id: &str,
) -> Option<Vec<Value>> { ) -> Option<Vec<Value>> {
let runner = state.app().redis_kv_runner()?; let cache_key = format!("upstream_models:{provider_id}:{key_id}");
let cache_key = runner let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
.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 parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?; let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
Some(aggregate_models_for_cache(&parsed)) Some(aggregate_models_for_cache(&parsed))
} }
@@ -1575,20 +1563,8 @@ async fn provider_query_read_provider_cached_models(
state: &AdminAppState<'_>, state: &AdminAppState<'_>,
provider_id: &str, provider_id: &str,
) -> Option<Vec<Value>> { ) -> Option<Vec<Value>> {
let runner = state.app().redis_kv_runner()?; let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
let cache_key = runner.keyspace().key(&format!( let raw = state.runtime_state().kv_get(&cache_key).await.ok()??;
"{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 parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?; let parsed = serde_json::from_str::<Vec<Value>>(&raw).ok()?;
Some(aggregate_models_for_cache(&parsed)) Some(aggregate_models_for_cache(&parsed))
} }
@@ -1598,18 +1574,18 @@ async fn provider_query_write_provider_cached_models(
provider_id: &str, provider_id: &str,
models: &[Value], 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 { let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(models)) else {
return; return;
}; };
let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}"); let cache_key = format!("{ANTIGRAVITY_PROVIDER_CACHE_KEY_PREFIX}{provider_id}");
let _ = runner let _ = state
.setex( .runtime_state()
.kv_set(
&cache_key, &cache_key,
&serialized, serialized,
Some(aether_model_fetch::model_fetch_interval_minutes().saturating_mul(60)), Some(std::time::Duration::from_secs(
aether_model_fetch::model_fetch_interval_minutes().saturating_mul(60),
)),
) )
.await; .await;
} }
@@ -116,8 +116,8 @@ impl<'a> AdminAppState<'a> {
self.app.mark_provider_key_rpm_reset(key_id, now_unix_secs) self.app.mark_provider_key_rpm_reset(key_id, now_unix_secs)
} }
pub(crate) fn redis_kv_runner(&self) -> Option<aether_data::driver::redis::RedisKvRunner> { pub(crate) fn runtime_state(&self) -> &aether_runtime_state::RuntimeState {
self.app.redis_kv_runner() self.app.runtime_state.as_ref()
} }
pub(crate) fn provider_key_rpm_reset_at( pub(crate) fn provider_key_rpm_reset_at(
@@ -84,22 +84,12 @@ impl<'a> AdminAppState<'a> {
}); });
let key = provider_oauth_state_storage_key(&nonce); let key = provider_oauth_state_storage_key(&nonce);
let value = payload.to_string(); let value = payload.to_string();
if let Some(runner) = self.redis_kv_runner() { self.as_ref()
runner .runtime_kv_setex(&key, &value, PROVIDER_OAUTH_STATE_TTL_SECS)
.setex(&key, &value, Some(PROVIDER_OAUTH_STATE_TTL_SECS)) .await?;
.await self.as_ref()
.map_err(|err| GatewayError::Internal(err.to_string()))?; .save_provider_oauth_state_for_tests(&key, &value);
return Ok(nonce); Ok(nonce)
}
if self
.as_ref()
.save_provider_oauth_state_for_tests(&key, &value)
{
return Ok(nonce);
}
Err(GatewayError::Internal(
"provider oauth redis unavailable".to_string(),
))
} }
pub(crate) async fn consume_provider_oauth_state( pub(crate) async fn consume_provider_oauth_state(
@@ -107,21 +97,7 @@ impl<'a> AdminAppState<'a> {
nonce: &str, nonce: &str,
) -> Result<Option<StoredAdminProviderOAuthState>, GatewayError> { ) -> Result<Option<StoredAdminProviderOAuthState>, GatewayError> {
let key = provider_oauth_state_storage_key(nonce); let key = provider_oauth_state_storage_key(nonce);
let raw = if let Some(runner) = self.redis_kv_runner() { let raw = self.as_ref().runtime_kv_getdel(&key).await?;
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(&key);
redis::cmd("GETDEL")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
self.as_ref().take_provider_oauth_state_for_tests(&key)
};
raw.map(|value| { raw.map(|value| {
serde_json::from_str::<StoredAdminProviderOAuthState>(&value) serde_json::from_str::<StoredAdminProviderOAuthState>(&value)
.map_err(|err| GatewayError::Internal(err.to_string())) .map_err(|err| GatewayError::Internal(err.to_string()))
@@ -172,35 +148,12 @@ impl<'a> AdminAppState<'a> {
let serialized = serde_json::to_string(task_state) let serialized = serde_json::to_string(task_state)
.map_err(|err| GatewayError::Internal(err.to_string()))?; .map_err(|err| GatewayError::Internal(err.to_string()))?;
if let Some(runner) = self.redis_kv_runner() { self.as_ref()
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await .runtime_kv_setex(&key, &serialized, PROVIDER_OAUTH_BATCH_TASK_TTL_SECS)
else { .await?;
return Err(GatewayError::Internal( self.as_ref()
"provider oauth batch task redis unavailable".to_string(), .save_provider_oauth_batch_task_for_tests(&key, &serialized);
)); Ok(())
};
let redis_key = runner.keyspace().key(&key);
redis::cmd("SET")
.arg(redis_key)
.arg(&serialized)
.arg("EX")
.arg(PROVIDER_OAUTH_BATCH_TASK_TTL_SECS)
.query_async::<()>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(());
}
if self
.as_ref()
.save_provider_oauth_batch_task_for_tests(&key, &serialized)
{
return Ok(());
}
Err(GatewayError::Internal(
"provider oauth batch task redis unavailable".to_string(),
))
} }
pub(crate) async fn read_provider_oauth_batch_task_payload( pub(crate) async fn read_provider_oauth_batch_task_payload(
@@ -209,22 +162,7 @@ impl<'a> AdminAppState<'a> {
task_id: &str, task_id: &str,
) -> Result<Option<serde_json::Value>, GatewayError> { ) -> Result<Option<serde_json::Value>, GatewayError> {
let key = provider_oauth_batch_task_storage_key(task_id); let key = provider_oauth_batch_task_storage_key(task_id);
let raw = if let Some(runner) = self.redis_kv_runner() { let raw = self.as_ref().runtime_kv_get(&key).await?;
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await
else {
return Err(GatewayError::Internal(
"provider oauth batch task redis unavailable".to_string(),
));
};
let redis_key = runner.keyspace().key(&key);
redis::cmd("GET")
.arg(redis_key)
.query_async(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
self.as_ref().load_provider_oauth_batch_task_for_tests(&key)
};
let Some(raw) = raw else { let Some(raw) = raw else {
return Ok(None); return Ok(None);
}; };
@@ -262,9 +200,8 @@ impl<'a> AdminAppState<'a> {
"provider oauth redis unavailable", "provider oauth redis unavailable",
) )
})?; })?;
if let Some(runner) = self.redis_kv_runner() { self.as_ref()
runner .runtime_kv_setex(&key, &value, ttl_seconds)
.setex(&key, &value, Some(ttl_seconds))
.await .await
.map_err(|_| { .map_err(|_| {
build_internal_control_error_response( build_internal_control_error_response(
@@ -272,18 +209,9 @@ impl<'a> AdminAppState<'a> {
"provider oauth redis unavailable", "provider oauth redis unavailable",
) )
})?; })?;
return Ok(()); self.as_ref()
} .save_provider_oauth_device_session_for_tests(&key, &value);
if self Ok(())
.as_ref()
.save_provider_oauth_device_session_for_tests(&key, &value)
{
return Ok(());
}
Err(build_internal_control_error_response(
http::StatusCode::SERVICE_UNAVAILABLE,
"provider oauth redis unavailable",
))
} }
pub(crate) async fn read_provider_oauth_device_session( pub(crate) async fn read_provider_oauth_device_session(
@@ -291,22 +219,7 @@ impl<'a> AdminAppState<'a> {
session_id: &str, session_id: &str,
) -> Result<Option<StoredAdminProviderOAuthDeviceSession>, GatewayError> { ) -> Result<Option<StoredAdminProviderOAuthDeviceSession>, GatewayError> {
let key = provider_oauth_device_session_storage_key(session_id); let key = provider_oauth_device_session_storage_key(session_id);
let raw = if let Some(runner) = self.redis_kv_runner() { let raw = self.as_ref().runtime_kv_get(&key).await?;
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(&key);
redis::cmd("GET")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
self.as_ref()
.load_provider_oauth_device_session_for_tests(&key)
};
raw.map(|value| { raw.map(|value| {
serde_json::from_str::<StoredAdminProviderOAuthDeviceSession>(&value) serde_json::from_str::<StoredAdminProviderOAuthDeviceSession>(&value)
.map_err(|err| GatewayError::Internal(err.to_string())) .map_err(|err| GatewayError::Internal(err.to_string()))
@@ -689,10 +689,10 @@ pub(crate) async fn proxy_request(
))); )));
} }
Err(RequestAdmissionError::Distributed( Err(RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::Saturated { gate, limit }, aether_runtime_state::RuntimeSemaphoreError::Saturated { gate, limit },
)) ))
| Err(RequestAdmissionError::Distributed( | Err(RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::Unavailable { gate, limit, .. }, aether_runtime_state::RuntimeSemaphoreError::Unavailable { gate, limit, .. },
)) => { )) => {
let trace_id = extract_or_generate_trace_id(request.headers()); let trace_id = extract_or_generate_trace_id(request.headers());
let response = build_local_overloaded_response(&trace_id, None, gate, limit)?; let response = build_local_overloaded_response(&trace_id, None, gate, limit)?;
@@ -714,7 +714,7 @@ pub(crate) async fn proxy_request(
)); ));
} }
Err(RequestAdmissionError::Distributed( Err(RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::InvalidConfiguration(message), aether_runtime_state::RuntimeSemaphoreError::InvalidConfiguration(message),
)) => return Err(GatewayError::Internal(message)), )) => return Err(GatewayError::Internal(message)),
}; };
let request_admission_ms = started_at.elapsed().as_millis() as u64; let request_admission_ms = started_at.elapsed().as_millis() as u64;
@@ -41,67 +41,6 @@ pub(super) fn auth_email_verified_key(email: &str) -> String {
format!("{AUTH_EMAIL_VERIFIED_PREFIX}{email}") format!("{AUTH_EMAIL_VERIFIED_PREFIX}{email}")
} }
pub(super) fn load_auth_email_verification_entry_for_tests(
_state: &AppState,
_key: &str,
) -> Option<String> {
#[cfg(test)]
{
return _state
.auth_email_verification_store
.as_ref()
.and_then(|store| {
store
.lock()
.expect("auth email verification store should lock")
.get(_key)
.cloned()
});
}
#[allow(unreachable_code)]
None
}
pub(super) fn save_auth_email_verification_entry_for_tests(
_state: &AppState,
_key: &str,
_value: &str,
) -> bool {
#[cfg(test)]
{
if let Some(store) = _state.auth_email_verification_store.as_ref() {
store
.lock()
.expect("auth email verification store should lock")
.insert(_key.to_string(), _value.to_string());
return true;
}
}
false
}
pub(super) fn delete_auth_email_verification_entries_for_tests(
_state: &AppState,
_keys: &[String],
) -> bool {
#[cfg(test)]
{
if let Some(store) = _state.auth_email_verification_store.as_ref() {
let mut guard = store
.lock()
.expect("auth email verification store should lock");
for key in _keys {
guard.remove(key);
}
return true;
}
}
false
}
pub(super) fn record_auth_email_delivery_for_tests( pub(super) fn record_auth_email_delivery_for_tests(
_state: &AppState, _state: &AppState,
_payload: serde_json::Value, _payload: serde_json::Value,
@@ -418,21 +357,7 @@ pub(super) async fn read_auth_email_verification_code(
email: &str, email: &str,
) -> Result<Option<StoredAuthEmailVerificationCode>, GatewayError> { ) -> Result<Option<StoredAuthEmailVerificationCode>, GatewayError> {
let key = auth_email_verification_key(email); let key = auth_email_verification_key(email);
let raw = if let Some(runner) = state.redis_kv_runner() { let raw = state.runtime_kv_get(&key).await?;
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(&key);
redis::cmd("GET")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
load_auth_email_verification_entry_for_tests(state, &key)
};
raw.map(|value| { raw.map(|value| {
serde_json::from_str::<StoredAuthEmailVerificationCode>(&value) serde_json::from_str::<StoredAuthEmailVerificationCode>(&value)
.map_err(|err| GatewayError::Internal(err.to_string())) .map_err(|err| GatewayError::Internal(err.to_string()))
@@ -445,21 +370,7 @@ pub(super) async fn auth_email_is_verified(
email: &str, email: &str,
) -> Result<bool, GatewayError> { ) -> Result<bool, GatewayError> {
let key = auth_email_verified_key(email); let key = auth_email_verified_key(email);
if let Some(runner) = state.redis_kv_runner() { state.runtime_kv_exists(&key).await
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(&key);
let exists = redis::cmd("EXISTS")
.arg(&namespaced_key)
.query_async::<i64>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(exists > 0);
}
Ok(load_auth_email_verification_entry_for_tests(state, &key).is_some())
} }
pub(super) async fn mark_auth_email_verified( pub(super) async fn mark_auth_email_verified(
@@ -467,16 +378,10 @@ pub(super) async fn mark_auth_email_verified(
email: &str, email: &str,
) -> Result<bool, GatewayError> { ) -> Result<bool, GatewayError> {
let key = auth_email_verified_key(email); let key = auth_email_verified_key(email);
if let Some(runner) = state.redis_kv_runner() { state
runner .runtime_kv_setex(&key, "verified", AUTH_EMAIL_VERIFIED_TTL_SECS)
.setex(&key, "verified", Some(AUTH_EMAIL_VERIFIED_TTL_SECS)) .await?;
.await Ok(true)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(true);
}
Ok(save_auth_email_verification_entry_for_tests(
state, &key, "verified",
))
} }
pub(super) async fn clear_auth_email_pending_code( pub(super) async fn clear_auth_email_pending_code(
@@ -484,17 +389,7 @@ pub(super) async fn clear_auth_email_pending_code(
email: &str, email: &str,
) -> Result<bool, GatewayError> { ) -> Result<bool, GatewayError> {
let verification_key = auth_email_verification_key(email); let verification_key = auth_email_verification_key(email);
if let Some(runner) = state.redis_kv_runner() { state.runtime_kv_del(&verification_key).await
let _ = runner
.del(&verification_key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(true);
}
Ok(delete_auth_email_verification_entries_for_tests(
state,
&[verification_key],
))
} }
pub(super) async fn clear_auth_email_verification( pub(super) async fn clear_auth_email_verification(
@@ -503,21 +398,9 @@ pub(super) async fn clear_auth_email_verification(
) -> Result<bool, GatewayError> { ) -> Result<bool, GatewayError> {
let verification_key = auth_email_verification_key(email); let verification_key = auth_email_verification_key(email);
let verified_key = auth_email_verified_key(email); let verified_key = auth_email_verified_key(email);
if let Some(runner) = state.redis_kv_runner() { let deleted_pending = state.runtime_kv_del(&verification_key).await?;
let _ = runner let deleted_verified = state.runtime_kv_del(&verified_key).await?;
.del(&verification_key) Ok(deleted_pending || deleted_verified)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let _ = runner
.del(&verified_key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(true);
}
Ok(delete_auth_email_verification_entries_for_tests(
state,
&[verification_key, verified_key],
))
} }
pub(super) async fn store_auth_email_verification_code( pub(super) async fn store_auth_email_verification_code(
@@ -533,16 +416,8 @@ pub(super) async fn store_auth_email_verification_code(
"created_at": created_at.to_rfc3339(), "created_at": created_at.to_rfc3339(),
}) })
.to_string(); .to_string();
if let Some(runner) = state.redis_kv_runner() { state.runtime_kv_setex(&key, &value, ttl_seconds).await?;
runner Ok(true)
.setex(&key, &value, Some(ttl_seconds))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(true);
}
Ok(save_auth_email_verification_entry_for_tests(
state, &key, &value,
))
} }
pub(super) async fn read_auth_smtp_config( pub(super) async fn read_auth_smtp_config(
+133 -43
View File
@@ -3,12 +3,12 @@
static GLOBAL: tikv_jemallocator::Jemalloc = tikv_jemallocator::Jemalloc; static GLOBAL: tikv_jemallocator::Jemalloc = tikv_jemallocator::Jemalloc;
use std::path::PathBuf; use std::path::PathBuf;
use std::sync::Arc;
use clap::{Args as ClapArgs, Parser, Subcommand, ValueEnum}; use clap::{Args as ClapArgs, Parser, Subcommand, ValueEnum};
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
use aether_crypto::warm_python_fernet_secret; use aether_crypto::warm_python_fernet_secret;
use aether_data::driver::redis::RedisClientConfig;
use aether_data::lifecycle::export::{export_database_jsonl, import_database_jsonl, ExportDomain}; use aether_data::lifecycle::export::{export_database_jsonl, import_database_jsonl, ExportDomain};
use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, DEFAULT_SQLITE_DATABASE_URL}; use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, DEFAULT_SQLITE_DATABASE_URL};
use aether_gateway::{ use aether_gateway::{
@@ -17,8 +17,12 @@ use aether_gateway::{
VideoTaskTruthSourceMode, VideoTaskTruthSourceMode,
}; };
use aether_runtime::{ use aether_runtime::{
init_service_runtime, DistributedConcurrencyGate, FileLoggingConfig, LogDestination, LogFormat, init_service_runtime, FileLoggingConfig, LogDestination, LogFormat, LogRotation,
LogRotation, RedisDistributedConcurrencyConfig, ServiceRuntimeConfig, ServiceRuntimeConfig,
};
use aether_runtime_state::{
RedisClientConfig, RuntimeSemaphoreConfig, RuntimeState, RuntimeStateBackendMode,
RuntimeStateConfig,
}; };
#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)] #[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)]
@@ -126,6 +130,7 @@ impl NodeRoleArg {
#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)] #[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)]
enum RuntimeBackendArg { enum RuntimeBackendArg {
Auto,
Redis, Redis,
Memory, Memory,
} }
@@ -133,10 +138,19 @@ enum RuntimeBackendArg {
impl RuntimeBackendArg { impl RuntimeBackendArg {
const fn as_str(self) -> &'static str { const fn as_str(self) -> &'static str {
match self { match self {
Self::Auto => "auto",
Self::Redis => "redis", Self::Redis => "redis",
Self::Memory => "memory", Self::Memory => "memory",
} }
} }
const fn to_runtime_state_backend(self) -> RuntimeStateBackendMode {
match self {
Self::Auto => RuntimeStateBackendMode::Auto,
Self::Redis => RuntimeStateBackendMode::Redis,
Self::Memory => RuntimeStateBackendMode::Memory,
}
}
} }
#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)] #[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)]
@@ -394,25 +408,12 @@ impl GatewayDataArgs {
fn to_config(&self) -> GatewayDataConfig { fn to_config(&self) -> GatewayDataConfig {
let database = self.effective_sql_database_config(); let database = self.effective_sql_database_config();
let redis_url = self.effective_redis_url();
let mut config = match database { let config = match database {
Some(database) => GatewayDataConfig::from_database_config(database), Some(database) => GatewayDataConfig::from_database_config(database),
None => GatewayDataConfig::disabled(), None => GatewayDataConfig::disabled(),
}; };
if let Some(redis_url) = redis_url.as_deref() {
config = config.with_redis_config(RedisClientConfig {
url: redis_url.to_string(),
key_prefix: self
.redis_key_prefix
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
});
}
match self.effective_encryption_key() { match self.effective_encryption_key() {
Some(value) => { Some(value) => {
warm_python_fernet_secret(&value); warm_python_fernet_secret(&value);
@@ -754,6 +755,12 @@ struct Args {
#[arg(long, env = "AETHER_RUNTIME_BACKEND", value_enum)] #[arg(long, env = "AETHER_RUNTIME_BACKEND", value_enum)]
runtime_backend: Option<RuntimeBackendArg>, runtime_backend: Option<RuntimeBackendArg>,
#[arg(long, env = "AETHER_RUNTIME_REDIS_URL")]
runtime_redis_url: Option<String>,
#[arg(long, env = "AETHER_RUNTIME_REDIS_KEY_PREFIX")]
runtime_redis_key_prefix: Option<String>,
#[command(flatten)] #[command(flatten)]
data: GatewayDataArgs, data: GatewayDataArgs,
@@ -777,8 +784,10 @@ impl Args {
data_redis_url: Option<&str>, data_redis_url: Option<&str>,
) -> RuntimeBackendArg { ) -> RuntimeBackendArg {
if let Some(runtime_backend) = self.runtime_backend { if let Some(runtime_backend) = self.runtime_backend {
if !matches!(runtime_backend, RuntimeBackendArg::Auto) {
return runtime_backend; return runtime_backend;
} }
}
if matches!(self.deployment_topology, DeploymentTopologyArg::MultiNode) { if matches!(self.deployment_topology, DeploymentTopologyArg::MultiNode) {
return RuntimeBackendArg::Redis; return RuntimeBackendArg::Redis;
} }
@@ -792,6 +801,49 @@ impl Args {
} }
} }
fn effective_runtime_redis_url(&self, data_redis_url: Option<&str>) -> Option<String> {
self.runtime_redis_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| data_redis_url.map(ToOwned::to_owned))
}
fn effective_runtime_redis_key_prefix(&self) -> Option<String> {
self.runtime_redis_key_prefix
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.or_else(|| {
self.data
.redis_key_prefix
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
})
}
fn runtime_state_config(
&self,
runtime_backend: RuntimeBackendArg,
data_redis_url: Option<&str>,
) -> RuntimeStateConfig {
let redis = self
.effective_runtime_redis_url(data_redis_url)
.map(|url| RedisClientConfig {
url,
key_prefix: self.effective_runtime_redis_key_prefix(),
});
RuntimeStateConfig {
backend: runtime_backend.to_runtime_state_backend(),
redis,
..RuntimeStateConfig::default()
}
}
fn runtime_config(&self) -> Result<ServiceRuntimeConfig, std::io::Error> { fn runtime_config(&self) -> Result<ServiceRuntimeConfig, std::io::Error> {
let default_log_filter = if self.command.is_some() let default_log_filter = if self.command.is_some()
|| self.migrate || self.migrate
@@ -990,12 +1042,20 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
let data_redis_url = args.data.effective_redis_url(); let data_redis_url = args.data.effective_redis_url();
let runtime_backend = let runtime_backend =
args.effective_runtime_backend(sql_database_config.as_ref(), data_redis_url.as_deref()); args.effective_runtime_backend(sql_database_config.as_ref(), data_redis_url.as_deref());
let runtime_redis_url = args.effective_runtime_redis_url(data_redis_url.as_deref());
validate_deployment_topology( validate_deployment_topology(
&args, &args,
sql_database_config.as_ref(), sql_database_config.as_ref(),
data_redis_url.as_deref(), runtime_redis_url.as_deref(),
runtime_backend, runtime_backend,
)?; )?;
let runtime_state = Arc::new(
RuntimeState::from_config(
args.runtime_state_config(runtime_backend, data_redis_url.as_deref()),
)
.await
.map_err(|err| std::io::Error::new(std::io::ErrorKind::InvalidInput, err.to_string()))?,
);
let data_config = args.data.to_config(); let data_config = args.data.to_config();
let rate_limit_config = if matches!(args.deployment_topology, DeploymentTopologyArg::MultiNode) let rate_limit_config = if matches!(args.deployment_topology, DeploymentTopologyArg::MultiNode)
{ {
@@ -1045,7 +1105,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
distributed_request_redis_configured = args distributed_request_redis_configured = args
.distributed_request_redis_url .distributed_request_redis_url
.as_deref() .as_deref()
.or(data_redis_url.as_deref()) .or(runtime_redis_url.as_deref())
.is_some(), .is_some(),
data_database_configured = sql_database_config.is_some(), data_database_configured = sql_database_config.is_some(),
data_database_driver = sql_database_config data_database_driver = sql_database_config
@@ -1053,13 +1113,15 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.map(|database| database.driver.as_str()) .map(|database| database.driver.as_str())
.unwrap_or("-"), .unwrap_or("-"),
data_postgres_configured = data_postgres_url.is_some(), data_postgres_configured = data_postgres_url.is_some(),
data_redis_configured = data_redis_url.is_some(), runtime_redis_configured = matches!(runtime_backend, RuntimeBackendArg::Redis),
data_redis_url_supplied = data_redis_url.is_some(),
data_has_encryption_key = data_config.encryption_key().is_some(), data_has_encryption_key = data_config.encryption_key().is_some(),
data_postgres_require_ssl = args.data.postgres_require_ssl, data_postgres_require_ssl = args.data.postgres_require_ssl,
"aether-gateway startup configuration" "aether-gateway startup configuration"
); );
let mut state = AppState::new()? let mut state = AppState::new()?
.with_runtime_state(runtime_state)
.with_data_config(data_config)? .with_data_config(data_config)?
.with_usage_runtime_config(args.usage.to_config())? .with_usage_runtime_config(args.usage.to_config())?
.with_video_task_truth_source_mode(args.video_task_truth_source_mode.into()); .with_video_task_truth_source_mode(args.video_task_truth_source_mode.into());
@@ -1088,35 +1150,21 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
state = state.with_request_concurrency_limit(limit); state = state.with_request_concurrency_limit(limit);
} }
if let Some(limit) = args.distributed_request_limit.filter(|limit| *limit > 0) { if let Some(limit) = args.distributed_request_limit.filter(|limit| *limit > 0) {
let redis_url = args let distributed_gate = state
.distributed_request_redis_url .runtime_state()
.as_deref() .semaphore(
.map(str::trim)
.filter(|value| !value.is_empty())
.or(data_redis_url.as_deref())
.ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"AETHER_GATEWAY_DISTRIBUTED_REQUEST_REDIS_URL or REDIS_URL/AETHER_GATEWAY_DATA_REDIS_URL is required when distributed request limit is enabled",
)
})?;
state =
state.with_distributed_request_concurrency_gate(DistributedConcurrencyGate::new_redis(
"gateway_requests_distributed", "gateway_requests_distributed",
limit, limit,
RedisDistributedConcurrencyConfig { RuntimeSemaphoreConfig {
url: redis_url.to_string(),
key_prefix: args
.distributed_request_redis_key_prefix
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned),
lease_ttl_ms: args.distributed_request_lease_ttl_ms.max(1), lease_ttl_ms: args.distributed_request_lease_ttl_ms.max(1),
renew_interval_ms: args.distributed_request_renew_interval_ms.max(1), renew_interval_ms: args.distributed_request_renew_interval_ms.max(1),
command_timeout_ms: Some(args.distributed_request_command_timeout_ms.max(1)), command_timeout_ms: Some(args.distributed_request_command_timeout_ms.max(1)),
}, },
)?); )
.map_err(|err| {
std::io::Error::new(std::io::ErrorKind::InvalidInput, err.to_string())
})?;
state = state.with_distributed_request_concurrency_gate(distributed_gate);
} }
if matches!(args.deployment_topology, DeploymentTopologyArg::MultiNode) if matches!(args.deployment_topology, DeploymentTopologyArg::MultiNode)
&& !state.has_usage_data_writer() && !state.has_usage_data_writer()
@@ -1531,6 +1579,8 @@ mod tests {
distributed_request_renew_interval_ms: 10_000, distributed_request_renew_interval_ms: 10_000,
distributed_request_command_timeout_ms: 1_000, distributed_request_command_timeout_ms: 1_000,
runtime_backend: None, runtime_backend: None,
runtime_redis_url: None,
runtime_redis_key_prefix: None,
data: GatewayDataArgs { data: GatewayDataArgs {
database_driver: None, database_driver: None,
database_url: None, database_url: None,
@@ -1649,6 +1699,46 @@ mod tests {
); );
} }
#[test]
fn memory_runtime_data_config_keeps_redis_out_of_data_layer() {
let mut args = test_args();
args.data.database_driver = Some(DatabaseDriverArg::Sqlite);
args.data.database_url = Some("sqlite://./data/aether.db".to_string());
args.data.redis_url = Some("redis://127.0.0.1/0".to_string());
let config = args.data.to_config();
assert_eq!(
config
.database()
.expect("database should be configured")
.driver,
DatabaseDriver::Sqlite
);
}
#[test]
fn redis_runtime_config_owns_redis_connection() {
let mut args = test_args();
args.data.database_driver = Some(DatabaseDriverArg::Postgres);
args.data.database_url = Some("postgres://postgres:postgres@localhost/aether".to_string());
args.data.redis_url = Some("redis://127.0.0.1/0".to_string());
let config = args.runtime_state_config(
RuntimeBackendArg::Redis,
args.data.effective_redis_url().as_deref(),
);
assert_eq!(
config
.redis
.as_ref()
.expect("redis should be configured for runtime state")
.url,
"redis://127.0.0.1/0"
);
}
#[test] #[test]
fn redis_url_defaults_to_redis_runtime_backend_for_server_database() { fn redis_url_defaults_to_redis_runtime_backend_for_server_database() {
let args = test_args(); let args = test_args();
@@ -1,10 +1,10 @@
use std::collections::BTreeMap; use std::collections::BTreeMap;
use std::time::{Duration, SystemTime, UNIX_EPOCH}; use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_data::driver::redis::{RedisKvRunner, RedisLockLease, RedisLockRunner};
use aether_data_contracts::repository::provider_catalog::{ use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider, StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
}; };
use aether_runtime_state::{RuntimeLockLease, RuntimeState};
use serde_json::Value; use serde_json::Value;
use tracing::{debug, info, warn}; use tracing::{debug, info, warn};
@@ -177,31 +177,20 @@ fn probe_stamp_key(provider_id: &str, key_id: &str) -> String {
} }
async fn load_probe_timestamps( async fn load_probe_timestamps(
runner: Option<&RedisKvRunner>, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
key_ids: &[String], key_ids: &[String],
) -> BTreeMap<String, u64> { ) -> BTreeMap<String, u64> {
let Some(runner) = runner else {
return BTreeMap::new();
};
if key_ids.is_empty() { if key_ids.is_empty() {
return BTreeMap::new(); return BTreeMap::new();
} }
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else { let runtime_keys = key_ids
debug!("gateway pool quota probe: failed to connect redis for stamp read");
return BTreeMap::new();
};
let redis_keys = key_ids
.iter() .iter()
.map(|key_id| runner.keyspace().key(&probe_stamp_key(provider_id, key_id))) .map(|key_id| probe_stamp_key(provider_id, key_id))
.collect::<Vec<_>>(); .collect::<Vec<_>>();
let Ok(values) = redis::cmd("MGET") let Ok(values) = runtime.kv_get_many(&runtime_keys).await else {
.arg(redis_keys) debug!("gateway pool quota probe: failed to read runtime probe stamps");
.query_async::<Vec<Option<String>>>(&mut connection)
.await
else {
debug!("gateway pool quota probe: failed to read redis probe stamps");
return BTreeMap::new(); return BTreeMap::new();
}; };
@@ -215,55 +204,43 @@ async fn load_probe_timestamps(
} }
async fn mark_probe_timestamps( async fn mark_probe_timestamps(
runner: Option<&RedisKvRunner>, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
key_ids: &[String], key_ids: &[String],
now_ts: u64, now_ts: u64,
interval_seconds: u64, interval_seconds: u64,
) { ) {
let Some(runner) = runner else {
return;
};
if key_ids.is_empty() { if key_ids.is_empty() {
return; return;
} }
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await else {
debug!("gateway pool quota probe: failed to connect redis for stamp write");
return;
};
let ttl_seconds = interval_seconds.saturating_mul(2).max(120); let ttl_seconds = interval_seconds.saturating_mul(2).max(120);
let value = now_ts.to_string(); let value = now_ts.to_string();
let mut pipeline = redis::pipe();
for key_id in key_ids { for key_id in key_ids {
pipeline if runtime
.cmd("SETEX") .kv_set(
.arg(runner.keyspace().key(&probe_stamp_key(provider_id, key_id))) &probe_stamp_key(provider_id, key_id),
.arg(ttl_seconds) value.clone(),
.arg(&value) Some(Duration::from_secs(ttl_seconds)),
.ignore(); )
.await
.is_err()
{
debug!("gateway pool quota probe: failed to write runtime probe stamp");
} }
if pipeline.query_async::<()>(&mut connection).await.is_err() {
debug!("gateway pool quota probe: failed to write redis probe stamps");
} }
} }
async fn acquire_provider_probe_lock( async fn acquire_provider_probe_lock(
runner: Option<&RedisLockRunner>, runtime: &RuntimeState,
provider_id: &str, provider_id: &str,
) -> Option<RedisLockLease> { ) -> Option<RuntimeLockLease> {
let Some(runner) = runner else {
return None;
};
let lock_key = runner
.keyspace()
.lock_key(&format!("pool_quota_probe:{provider_id}"));
let owner = format!("aether-gateway-pool-probe-{}", std::process::id()); let owner = format!("aether-gateway-pool-probe-{}", std::process::id());
match runner match runtime
.try_acquire( .lock_try_acquire(
&lock_key, &format!("pool_quota_probe:{provider_id}"),
&owner, &owner,
Some(POOL_QUOTA_PROBE_PROVIDER_LOCK_TTL_MS), Duration::from_millis(POOL_QUOTA_PROBE_PROVIDER_LOCK_TTL_MS),
) )
.await .await
{ {
@@ -272,40 +249,36 @@ async fn acquire_provider_probe_lock(
debug!( debug!(
provider_id, provider_id,
error = %err, error = %err,
"gateway pool quota probe: failed to acquire redis provider lock" "gateway pool quota probe: failed to acquire runtime provider lock"
); );
None None
} }
} }
} }
async fn release_provider_probe_lock( async fn release_provider_probe_lock(runtime: &RuntimeState, lease: Option<RuntimeLockLease>) {
runner: Option<&RedisLockRunner>, let Some(lease) = lease else {
lease: Option<RedisLockLease>,
) {
let (Some(runner), Some(lease)) = (runner, lease) else {
return; return;
}; };
if let Err(err) = runner.release(&lease).await { if let Err(err) = runtime.lock_release(&lease).await {
debug!( debug!(
error = %err, error = %err,
"gateway pool quota probe: failed to release redis provider lock" "gateway pool quota probe: failed to release runtime provider lock"
); );
} }
} }
async fn select_keys_for_provider( async fn select_keys_for_provider(
state: &AppState, state: &AppState,
redis_kv: Option<&RedisKvRunner>, runtime: &RuntimeState,
redis_lock: Option<&RedisLockRunner>,
provider: &StoredProviderCatalogProvider, provider: &StoredProviderCatalogProvider,
provider_type: &str, provider_type: &str,
interval_seconds: u64, interval_seconds: u64,
max_keys_per_provider: usize, max_keys_per_provider: usize,
now_ts: u64, now_ts: u64,
) -> Result<Vec<StoredProviderCatalogKey>, GatewayError> { ) -> Result<Vec<StoredProviderCatalogKey>, GatewayError> {
let lease = acquire_provider_probe_lock(redis_lock, &provider.id).await; let lease = acquire_provider_probe_lock(runtime, &provider.id).await;
if redis_lock.is_some() && lease.is_none() { if lease.is_none() {
return Ok(Vec::new()); return Ok(Vec::new());
} }
@@ -321,7 +294,7 @@ async fn select_keys_for_provider(
} }
let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>(); let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
let probe_stamps = load_probe_timestamps(redis_kv, &provider.id, &key_ids).await; let probe_stamps = load_probe_timestamps(runtime, &provider.id, &key_ids).await;
let selected_ids = select_pool_quota_probe_key_ids( let selected_ids = select_pool_quota_probe_key_ids(
&keys, &keys,
provider_type, provider_type,
@@ -335,7 +308,7 @@ async fn select_keys_for_provider(
} }
mark_probe_timestamps( mark_probe_timestamps(
redis_kv, runtime,
&provider.id, &provider.id,
&selected_ids, &selected_ids,
now_ts, now_ts,
@@ -354,7 +327,7 @@ async fn select_keys_for_provider(
} }
.await; .await;
release_provider_probe_lock(redis_lock, lease).await; release_provider_probe_lock(runtime, lease).await;
result result
} }
@@ -461,8 +434,6 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
.push(endpoint); .push(endpoint);
} }
let redis_kv = state.redis_kv_runner();
let redis_lock = state.data.oauth_refresh_lock_runner();
let admin_state = AdminAppState::new(state); let admin_state = AdminAppState::new(state);
let now_ts = now_unix_secs(); let now_ts = now_unix_secs();
let mut summary = PoolQuotaProbeRunSummary { let mut summary = PoolQuotaProbeRunSummary {
@@ -487,8 +458,7 @@ pub(crate) async fn perform_pool_quota_probe_once_with_config(
let interval_seconds = interval_minutes.clamp(1, 1440).saturating_mul(60); let interval_seconds = interval_minutes.clamp(1, 1440).saturating_mul(60);
let keys = select_keys_for_provider( let keys = select_keys_for_provider(
state, state,
redis_kv.as_ref(), state.runtime_state.as_ref(),
redis_lock.as_ref(),
&provider, &provider,
&provider_type, &provider_type,
interval_seconds, interval_seconds,
+3 -27
View File
@@ -75,19 +75,9 @@ pub(crate) async fn save_identity_oauth_state(
let key = identity_oauth_state_storage_key(&record.nonce); let key = identity_oauth_state_storage_key(&record.nonce);
let value = let value =
serde_json::to_string(record).map_err(|err| GatewayError::Internal(err.to_string()))?; serde_json::to_string(record).map_err(|err| GatewayError::Internal(err.to_string()))?;
if let Some(runner) = state.redis_kv_runner() { state
runner .runtime_kv_setex(&key, &value, IDENTITY_OAUTH_STATE_TTL_SECS)
.setex(&key, &value, Some(IDENTITY_OAUTH_STATE_TTL_SECS))
.await .await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(());
}
if state.save_provider_oauth_state_for_tests(&key, &value) {
return Ok(());
}
Err(GatewayError::Internal(
"identity oauth state store unavailable".to_string(),
))
} }
pub(crate) async fn consume_identity_oauth_state( pub(crate) async fn consume_identity_oauth_state(
@@ -95,21 +85,7 @@ pub(crate) async fn consume_identity_oauth_state(
nonce: &str, nonce: &str,
) -> Result<Option<StoredIdentityOAuthState>, GatewayError> { ) -> Result<Option<StoredIdentityOAuthState>, GatewayError> {
let key = identity_oauth_state_storage_key(nonce); let key = identity_oauth_state_storage_key(nonce);
let raw = if let Some(runner) = state.redis_kv_runner() { let raw = state.runtime_kv_getdel(&key).await?;
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(&key);
redis::cmd("GETDEL")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
state.take_provider_oauth_state_for_tests(&key)
};
raw.map(|value| { raw.map(|value| {
serde_json::from_str::<StoredIdentityOAuthState>(&value) serde_json::from_str::<StoredIdentityOAuthState>(&value)
.map_err(|err| GatewayError::Internal(err.to_string())) .map_err(|err| GatewayError::Internal(err.to_string()))
@@ -96,7 +96,6 @@ pub(crate) enum LocalExecutionEffect<'a> {
} }
struct PoolFeedbackContext { struct PoolFeedbackContext {
runner: aether_data::driver::redis::RedisKvRunner,
pool_config: AdminProviderPoolConfig, pool_config: AdminProviderPoolConfig,
sticky_session_token: Option<String>, sticky_session_token: Option<String>,
} }
@@ -251,10 +250,6 @@ async fn resolve_pool_feedback_context(
state: &AppState, state: &AppState,
context: LocalExecutionEffectContext<'_>, context: LocalExecutionEffectContext<'_>,
) -> Option<PoolFeedbackContext> { ) -> Option<PoolFeedbackContext> {
let Some(runner) = state.redis_kv_runner() else {
return None;
};
let plan = context.plan; let plan = context.plan;
let transport = match state let transport = match state
.read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id) .read_provider_transport_snapshot(&plan.provider_id, &plan.endpoint_id, &plan.key_id)
@@ -281,7 +276,6 @@ async fn resolve_pool_feedback_context(
.and_then(extract_pool_sticky_session_token); .and_then(extract_pool_sticky_session_token);
Some(PoolFeedbackContext { Some(PoolFeedbackContext {
runner,
pool_config, pool_config,
sticky_session_token, sticky_session_token,
}) })
@@ -333,7 +327,7 @@ async fn record_sync_pool_success_effect(
let usage_outcome = let usage_outcome =
build_sync_terminal_usage_outcome(context.plan, context.report_context, payload); build_sync_terminal_usage_outcome(context.plan, context.report_context, payload);
record_admin_provider_pool_success( record_admin_provider_pool_success(
&pool_context.runner, state.runtime_state.as_ref(),
&context.plan.provider_id, &context.plan.provider_id,
&context.plan.key_id, &context.plan.key_id,
&pool_context.pool_config, &pool_context.pool_config,
@@ -552,7 +546,7 @@ async fn record_stream_pool_success_effect(
let usage_outcome = let usage_outcome =
build_stream_terminal_usage_outcome(context.plan, context.report_context, payload); build_stream_terminal_usage_outcome(context.plan, context.report_context, payload);
record_admin_provider_pool_success( record_admin_provider_pool_success(
&pool_context.runner, state.runtime_state.as_ref(),
&context.plan.provider_id, &context.plan.provider_id,
&context.plan.key_id, &context.plan.key_id,
&pool_context.pool_config, &pool_context.pool_config,
@@ -588,7 +582,7 @@ async fn record_pool_error_effect(
} }
record_admin_provider_pool_error( record_admin_provider_pool_error(
&pool_context.runner, state.runtime_state.as_ref(),
&context.plan.provider_id, &context.plan.provider_id,
&context.plan.key_id, &context.plan.key_id,
&pool_context.pool_config, &pool_context.pool_config,
@@ -747,7 +741,7 @@ async fn record_pool_stream_timeout_effect(
}; };
record_admin_provider_pool_stream_timeout( record_admin_provider_pool_stream_timeout(
&pool_context.runner, state.runtime_state.as_ref(),
&context.plan.provider_id, &context.plan.provider_id,
&context.plan.key_id, &context.plan.key_id,
&pool_context.pool_config, &pool_context.pool_config,
+34 -104
View File
@@ -4,56 +4,13 @@ use std::sync::Mutex as StdMutex;
use std::time::{Duration, SystemTime, UNIX_EPOCH}; use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_cache::ExpiringMap; use aether_cache::ExpiringMap;
use redis::Script; use aether_runtime_state::{RateLimitCheck, RateLimitInput, RateLimitScope};
use tokio::sync::Mutex; use tokio::sync::Mutex;
use tracing::warn; use tracing::warn;
use crate::control::GatewayControlDecision; use crate::control::GatewayControlDecision;
use crate::{AppState, GatewayError}; use crate::{AppState, GatewayError};
const RPM_CHECK_AND_CONSUME_SCRIPT: &str = r#"
local user_key = KEYS[1]
local key_key = KEYS[2]
local user_limit = tonumber(ARGV[1])
local key_limit = tonumber(ARGV[2])
local ttl = tonumber(ARGV[3])
local retry_after = tonumber(ARGV[4])
local user_count = 0
if user_limit > 0 then
user_count = tonumber(redis.call('GET', user_key) or '0')
if user_count >= user_limit then
return {0, 1, user_limit, 0, retry_after}
end
end
local key_count = 0
if key_limit > 0 then
key_count = tonumber(redis.call('GET', key_key) or '0')
if key_count >= key_limit then
return {0, 2, key_limit, 0, retry_after}
end
end
local remaining = -1
if user_limit > 0 then
user_count = redis.call('INCR', user_key)
redis.call('EXPIRE', user_key, ttl)
remaining = user_limit - user_count
end
if key_limit > 0 then
key_count = redis.call('INCR', key_key)
redis.call('EXPIRE', key_key, ttl)
local key_remaining = key_limit - key_count
if remaining == -1 or key_remaining < remaining then
remaining = key_remaining
end
end
return {1, 0, 0, remaining, 0}
"#;
const SYSTEM_RPM_CONFIG_KEY: &str = "rate_limit_per_minute"; const SYSTEM_RPM_CONFIG_KEY: &str = "rate_limit_per_minute";
const SYSTEM_RPM_CONFIG_CACHE_TTL: Duration = Duration::from_secs(15); const SYSTEM_RPM_CONFIG_CACHE_TTL: Duration = Duration::from_secs(15);
const SYSTEM_RPM_CONFIG_CACHE_MAX_ENTRIES: usize = 8; const SYSTEM_RPM_CONFIG_CACHE_MAX_ENTRIES: usize = 8;
@@ -189,19 +146,11 @@ impl FrontdoorUserRpmLimiter {
scope_key: &str, scope_key: &str,
bucket: u64, bucket: u64,
) -> Result<u32, GatewayError> { ) -> Result<u32, GatewayError> {
if let Some(runner) = state.redis_kv_runner() { if !state.runtime_state.is_memory() {
let mut connection = runner let raw = state.runtime_state.kv_get(scope_key).await.map_err(|err| {
.client() GatewayError::Internal(format!("frontdoor user rpm runtime read failed: {err}"))
.get_multiplexed_async_connection() })?;
.await return Ok(raw.and_then(|value| value.parse::<u32>().ok()).unwrap_or(0));
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let namespaced_key = runner.keyspace().key(scope_key);
let raw = redis::cmd("GET")
.arg(&namespaced_key)
.query_async::<Option<u32>>(&mut connection)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
return Ok(raw.unwrap_or(0));
} }
let counts = self.memory_counts.lock().await; let counts = self.memory_counts.lock().await;
@@ -234,8 +183,8 @@ impl FrontdoorUserRpmLimiter {
return Ok(FrontdoorUserRpmOutcome::Allowed); return Ok(FrontdoorUserRpmOutcome::Allowed);
} }
if let Some(runner) = state.data.kv_runner() { if !state.runtime_state.is_memory() {
match self.check_and_consume_redis(&runner, &plan).await { match self.check_and_consume_runtime(state, &plan).await {
Ok(outcome) => return Ok(outcome), Ok(outcome) => return Ok(outcome),
Err(err) => { Err(err) => {
warn!( warn!(
@@ -249,7 +198,7 @@ impl FrontdoorUserRpmLimiter {
} }
if !self.config.allow_local_fallback() { if !self.config.allow_local_fallback() {
return Err(GatewayError::Internal( return Err(GatewayError::Internal(
"frontdoor user rpm redis backend is unavailable and local fallback is disabled for the current deployment mode".to_string(), "frontdoor user rpm runtime backend is unavailable and local fallback is disabled for the current deployment mode".to_string(),
)); ));
} }
} }
@@ -261,7 +210,8 @@ impl FrontdoorUserRpmLimiter {
return Ok(FrontdoorUserRpmOutcome::NotApplicable); return Ok(FrontdoorUserRpmOutcome::NotApplicable);
} }
return Err(GatewayError::Internal( return Err(GatewayError::Internal(
"frontdoor user rpm requires redis in the current deployment mode".to_string(), "frontdoor user rpm requires shared runtime state in the current deployment mode"
.to_string(),
)); ));
} }
@@ -306,59 +256,39 @@ impl FrontdoorUserRpmLimiter {
self self
} }
async fn check_and_consume_redis( async fn check_and_consume_runtime(
&self, &self,
runner: &aether_data::driver::redis::RedisKvRunner, state: &AppState,
plan: &RpmPlan, plan: &RpmPlan,
) -> Result<FrontdoorUserRpmOutcome, GatewayError> { ) -> Result<FrontdoorUserRpmOutcome, GatewayError> {
let user_key = runner.keyspace().key(&plan.user_rpm_key); let result = state
let key_key = runner.keyspace().key(&plan.key_rpm_key); .runtime_state
let mut connection = runner .check_and_consume_rate_limit(RateLimitInput {
.client() user_key: &plan.user_rpm_key,
.get_multiplexed_async_connection() key_key: &plan.key_rpm_key,
.await bucket: plan.bucket,
.map_err(|err| GatewayError::Internal(err.to_string()))?; user_limit: plan.user_rpm_limit,
let raw_result: Vec<i64> = Script::new(RPM_CHECK_AND_CONSUME_SCRIPT) key_limit: plan.key_rpm_limit,
.key(user_key) ttl_seconds: self.config.key_ttl_seconds(),
.key(key_key) })
.arg(i64::from(plan.user_rpm_limit))
.arg(i64::from(plan.key_rpm_limit))
.arg(i64::try_from(self.config.key_ttl_seconds()).unwrap_or(i64::MAX))
.arg(i64::try_from(plan.retry_after).unwrap_or(i64::MAX))
.invoke_async::<Vec<i64>>(&mut connection)
.await .await
.map_err(|err| GatewayError::Internal(err.to_string()))?; .map_err(|err| GatewayError::Internal(err.to_string()))?;
let allowed = raw_result.first().copied().unwrap_or_default() == 1; if matches!(result, RateLimitCheck::Allowed { .. }) {
if allowed {
return Ok(FrontdoorUserRpmOutcome::Allowed); return Ok(FrontdoorUserRpmOutcome::Allowed);
} }
let scope = match raw_result.get(1).copied().unwrap_or_default() { let RateLimitCheck::Rejected { scope, limit } = result else {
2 => "key", unreachable!("allowed returned above");
_ => "user",
}; };
let limit = raw_result
.get(2)
.copied()
.and_then(|value| u32::try_from(value).ok())
.unwrap_or_else(|| {
if scope == "key" {
plan.key_rpm_limit
} else {
plan.user_rpm_limit
}
});
let retry_after = raw_result
.get(4)
.copied()
.and_then(|value| u64::try_from(value).ok())
.unwrap_or(plan.retry_after);
Ok(FrontdoorUserRpmOutcome::Rejected( Ok(FrontdoorUserRpmOutcome::Rejected(
FrontdoorUserRpmRejection { FrontdoorUserRpmRejection {
scope, scope: match scope {
RateLimitScope::User => "user",
RateLimitScope::Key => "key",
},
limit, limit,
retry_after, retry_after: plan.retry_after,
}, },
)) ))
} }
@@ -630,7 +560,7 @@ mod tests {
} }
#[tokio::test] #[tokio::test]
async fn limiter_rejects_missing_redis_when_local_fallback_disabled() { async fn limiter_rejects_missing_shared_runtime_when_local_fallback_disabled() {
let limiter = FrontdoorUserRpmLimiter::new( let limiter = FrontdoorUserRpmLimiter::new(
FrontdoorUserRpmConfig::new(60, 120, false).with_local_fallback(false), FrontdoorUserRpmConfig::new(60, 120, false).with_local_fallback(false),
); );
@@ -652,10 +582,10 @@ mod tests {
let err = limiter let err = limiter
.check_and_consume(&state, Some(&decision)) .check_and_consume(&state, Some(&decision))
.await .await
.expect_err("missing redis should fail in strict mode"); .expect_err("missing shared runtime should fail in strict mode");
match err { match err {
crate::GatewayError::Internal(message) => { crate::GatewayError::Internal(message) => {
assert!(message.contains("requires redis")); assert!(message.contains("requires shared runtime state"));
} }
other => panic!("expected internal error, got {other:?}"), other => panic!("expected internal error, got {other:?}"),
} }
+3 -2
View File
@@ -9,7 +9,8 @@ use tower::ServiceExt;
use tower_http::services::{ServeDir, ServeFile}; use tower_http::services::{ServeDir, ServeFile};
use tracing::warn; use tracing::warn;
use aether_runtime::{prometheus_response, ConcurrencyError, DistributedConcurrencyError}; use aether_runtime::{prometheus_response, ConcurrencyError};
use aether_runtime_state::RuntimeSemaphoreError;
use super::{api, handlers::proxy::proxy_request, middleware, state::AppState}; use super::{api, handlers::proxy::proxy_request, middleware, state::AppState};
@@ -124,7 +125,7 @@ pub(crate) async fn metrics(
#[derive(Debug)] #[derive(Debug)]
pub(crate) enum RequestAdmissionError { pub(crate) enum RequestAdmissionError {
Local(ConcurrencyError), Local(ConcurrencyError),
Distributed(DistributedConcurrencyError), Distributed(RuntimeSemaphoreError),
} }
pub async fn serve_tcp(bind: &str) -> Result<(), Box<dyn std::error::Error>> { pub async fn serve_tcp(bind: &str) -> Result<(), Box<dyn std::error::Error>> {
+4 -2
View File
@@ -2,7 +2,8 @@ use std::collections::HashMap;
use std::sync::Arc; use std::sync::Arc;
use std::sync::Mutex as StdMutex; use std::sync::Mutex as StdMutex;
use aether_runtime::{ConcurrencyGate, DistributedConcurrencyGate}; use aether_runtime::ConcurrencyGate;
use aether_runtime_state::{RuntimeSemaphore, RuntimeState};
use super::super::async_task::{VideoTaskPollerConfig, VideoTaskService}; use super::super::async_task::{VideoTaskPollerConfig, VideoTaskService};
use super::super::cache::{ use super::super::cache::{
@@ -47,11 +48,12 @@ pub struct AppState {
#[cfg(test)] #[cfg(test)]
pub(crate) execution_runtime_sync_override: Option<TestExecutionRuntimeSyncOverride>, pub(crate) execution_runtime_sync_override: Option<TestExecutionRuntimeSyncOverride>,
pub(crate) data: Arc<GatewayDataState>, pub(crate) data: Arc<GatewayDataState>,
pub(crate) runtime_state: Arc<RuntimeState>,
pub(crate) usage_runtime: Arc<usage::UsageRuntime>, pub(crate) usage_runtime: Arc<usage::UsageRuntime>,
pub(crate) video_tasks: Arc<VideoTaskService>, pub(crate) video_tasks: Arc<VideoTaskService>,
pub(crate) video_task_poller: Option<VideoTaskPollerConfig>, pub(crate) video_task_poller: Option<VideoTaskPollerConfig>,
pub(crate) request_gate: Option<Arc<ConcurrencyGate>>, pub(crate) request_gate: Option<Arc<ConcurrencyGate>>,
pub(crate) distributed_request_gate: Option<Arc<DistributedConcurrencyGate>>, pub(crate) distributed_request_gate: Option<Arc<RuntimeSemaphore>>,
pub(crate) client: reqwest::Client, pub(crate) client: reqwest::Client,
pub(crate) auth_context_cache: Arc<AuthContextCache>, pub(crate) auth_context_cache: Arc<AuthContextCache>,
pub(crate) auth_api_key_last_used_cache: Arc<AuthApiKeyLastUsedCache>, pub(crate) auth_api_key_last_used_cache: Arc<AuthApiKeyLastUsedCache>,
+144 -60
View File
@@ -9,9 +9,12 @@ use aether_data::repository::proxy_nodes::{
}; };
use aether_http::{build_http_client, HttpClientConfig}; use aether_http::{build_http_client, HttpClientConfig};
use aether_runtime::{ use aether_runtime::{
service_up_sample, AdmissionPermit, ConcurrencyGate, ConcurrencySnapshot, service_up_sample, AdmissionPermit, ConcurrencyGate, ConcurrencySnapshot, MetricKind,
DistributedConcurrencyError, DistributedConcurrencyGate, DistributedConcurrencySnapshot, MetricLabel, MetricSample,
MetricKind, MetricLabel, MetricSample, };
use aether_runtime_state::{
MemoryRuntimeStateConfig, RuntimeQueueStore, RuntimeSemaphore, RuntimeSemaphoreError,
RuntimeSemaphoreSnapshot, RuntimeState,
}; };
use aether_scheduler_core::PROVIDER_KEY_RPM_WINDOW_SECS; use aether_scheduler_core::PROVIDER_KEY_RPM_WINDOW_SECS;
use tokio::task::JoinHandle; use tokio::task::JoinHandle;
@@ -53,21 +56,32 @@ use crate::maintenance::spawn_wallet_daily_usage_aggregation_worker;
const SYSTEM_CONFIG_CACHE_TTL: Duration = Duration::from_secs(3); const SYSTEM_CONFIG_CACHE_TTL: Duration = Duration::from_secs(3);
impl AppState { impl AppState {
fn usage_worker_queue_for(
runtime_state: &Arc<RuntimeState>,
) -> Option<Arc<dyn RuntimeQueueStore>> {
if runtime_state.is_redis() {
let queue: Arc<dyn RuntimeQueueStore> = runtime_state.clone();
Some(queue)
} else {
None
}
}
fn spawn_scheduler_affinity_redis_write( fn spawn_scheduler_affinity_redis_write(
&self, &self,
cache_key: &str, cache_key: &str,
target: &SchedulerAffinityTarget, target: &SchedulerAffinityTarget,
ttl: Duration, ttl: Duration,
) { ) {
let Some(runner) = self.redis_kv_runner() else { if self.runtime_state.is_memory() {
return; return;
}; }
let Ok(handle) = tokio::runtime::Handle::try_current() else { let Ok(handle) = tokio::runtime::Handle::try_current() else {
return; return;
}; };
let cache_key = cache_key.to_string(); let cache_key = cache_key.to_string();
let namespaced_cache_key = runner.keyspace().key(&cache_key); let runtime_state = self.runtime_state.clone();
let provider_id = target.provider_id.clone(); let provider_id = target.provider_id.clone();
let endpoint_id = target.endpoint_id.clone(); let endpoint_id = target.endpoint_id.clone();
let key_id = target.key_id.clone(); let key_id = target.key_id.clone();
@@ -76,56 +90,55 @@ impl AppState {
let expire_at = now_unix_secs.saturating_add(ttl_seconds); let expire_at = now_unix_secs.saturating_add(ttl_seconds);
handle.spawn(async move { handle.spawn(async move {
let Ok(mut connection) = runner.client().get_multiplexed_async_connection().await let existing = runtime_state
else { .kv_get(&cache_key)
return; .await
}; .ok()
let script = r#" .flatten()
local existing = redis.call('GET', KEYS[1]) .and_then(|raw| serde_json::from_str::<serde_json::Value>(&raw).ok());
local request_count = 0 let request_count = existing
local created_at = tonumber(ARGV[4]) .as_ref()
if existing then .and_then(|value| value.get("request_count"))
local ok, payload = pcall(cjson.decode, existing) .and_then(serde_json::Value::as_u64)
if ok and type(payload) == 'table' then .unwrap_or_default()
if type(payload['request_count']) == 'number' then .saturating_add(1);
request_count = payload['request_count'] let created_at = existing
end .as_ref()
if type(payload['created_at']) == 'number' then .and_then(|value| value.get("created_at"))
created_at = payload['created_at'] .and_then(serde_json::Value::as_u64)
end .unwrap_or(now_unix_secs);
end let payload = serde_json::json!({
end "provider_id": provider_id,
request_count = request_count + 1 "endpoint_id": endpoint_id,
local payload = { "key_id": key_id,
provider_id = ARGV[1], "created_at": created_at,
endpoint_id = ARGV[2], "expire_at": expire_at,
key_id = ARGV[3], "request_count": request_count,
created_at = created_at, });
expire_at = tonumber(ARGV[5]), if let Ok(serialized) = serde_json::to_string(&payload) {
request_count = request_count let _ = runtime_state
} .kv_set(
redis.call('SETEX', KEYS[1], tonumber(ARGV[6]), cjson.encode(payload)) &cache_key,
return request_count serialized,
"#; Some(Duration::from_secs(ttl_seconds)),
let _ = redis::cmd("EVAL") )
.arg(script)
.arg(1)
.arg(&namespaced_cache_key)
.arg(&provider_id)
.arg(&endpoint_id)
.arg(&key_id)
.arg(now_unix_secs)
.arg(expire_at)
.arg(ttl_seconds)
.query_async::<i64>(&mut connection)
.await; .await;
}
}); });
} }
pub(crate) fn replace_data_state(&mut self, data: Arc<GatewayDataState>) { pub(crate) fn replace_data_state(&mut self, data: Arc<GatewayDataState>) {
self.clear_provider_transport_snapshot_cache(); self.clear_provider_transport_snapshot_cache();
self.system_config_cache.clear(); self.system_config_cache.clear();
self.tunnel = crate::tunnel::EmbeddedTunnelState::with_data(Arc::clone(&data)); let data = Arc::new(
(*data)
.clone()
.with_usage_worker_queue(Self::usage_worker_queue_for(&self.runtime_state)),
);
self.tunnel = crate::tunnel::EmbeddedTunnelState::with_data_and_runtime_state(
Arc::clone(&data),
self.runtime_state.clone(),
);
self.data = data; self.data = data;
} }
@@ -153,7 +166,11 @@ return request_count
} }
fn build(execution_runtime_override_base_url: Option<String>) -> Result<Self, reqwest::Error> { fn build(execution_runtime_override_base_url: Option<String>) -> Result<Self, reqwest::Error> {
let data = Arc::new(GatewayDataState::disabled()); let runtime_state = Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default()));
let data = Arc::new(
GatewayDataState::disabled()
.with_usage_worker_queue(Self::usage_worker_queue_for(&runtime_state)),
);
let client = build_http_client(&HttpClientConfig { let client = build_http_client(&HttpClientConfig {
connect_timeout_ms: Some(10_000), connect_timeout_ms: Some(10_000),
request_timeout_ms: Some(300_000), request_timeout_ms: Some(300_000),
@@ -168,6 +185,7 @@ return request_count
#[cfg(test)] #[cfg(test)]
execution_runtime_sync_override: None, execution_runtime_sync_override: None,
data: Arc::clone(&data), data: Arc::clone(&data),
runtime_state: runtime_state.clone(),
usage_runtime: Arc::new(usage::UsageRuntime::disabled()), usage_runtime: Arc::new(usage::UsageRuntime::disabled()),
video_tasks: Arc::new(VideoTaskService::new( video_tasks: Arc::new(VideoTaskService::new(
VideoTaskTruthSourceMode::PythonSyncReport, VideoTaskTruthSourceMode::PythonSyncReport,
@@ -188,7 +206,10 @@ return request_count
frontdoor_user_rpm: Arc::new(FrontdoorUserRpmLimiter::new( frontdoor_user_rpm: Arc::new(FrontdoorUserRpmLimiter::new(
FrontdoorUserRpmConfig::default(), FrontdoorUserRpmConfig::default(),
)), )),
tunnel: crate::tunnel::EmbeddedTunnelState::with_data(data), tunnel: crate::tunnel::EmbeddedTunnelState::with_data_and_runtime_state(
data,
runtime_state.clone(),
),
provider_transport_snapshot_cache: Arc::new(StdMutex::new(HashMap::new())), provider_transport_snapshot_cache: Arc::new(StdMutex::new(HashMap::new())),
provider_key_rpm_resets: Arc::new(StdMutex::new(HashMap::new())), provider_key_rpm_resets: Arc::new(StdMutex::new(HashMap::new())),
local_execution_runtime_miss_diagnostics: Arc::new(StdMutex::new(HashMap::new())), local_execution_runtime_miss_diagnostics: Arc::new(StdMutex::new(HashMap::new())),
@@ -261,11 +282,12 @@ return request_count
instance_id: impl Into<String>, instance_id: impl Into<String>,
relay_base_url: Option<impl Into<String>>, relay_base_url: Option<impl Into<String>>,
) -> Self { ) -> Self {
self.tunnel = crate::tunnel::EmbeddedTunnelState::with_data_and_identity( self.tunnel = crate::tunnel::EmbeddedTunnelState::with_data_identity_and_runtime_state(
Arc::clone(&self.data), Arc::clone(&self.data),
instance_id, instance_id,
relay_base_url, relay_base_url,
90, 90,
self.runtime_state.clone(),
); );
self self
} }
@@ -334,10 +356,21 @@ return request_count
self self
} }
pub fn with_distributed_request_concurrency_gate( pub fn with_runtime_state(mut self, runtime_state: Arc<RuntimeState>) -> Self {
mut self, self.runtime_state = runtime_state;
gate: DistributedConcurrencyGate, self.data = Arc::new(
) -> Self { (*self.data)
.clone()
.with_usage_worker_queue(Self::usage_worker_queue_for(&self.runtime_state)),
);
self.tunnel = crate::tunnel::EmbeddedTunnelState::with_data_and_runtime_state(
Arc::clone(&self.data),
self.runtime_state.clone(),
);
self
}
pub fn with_distributed_request_concurrency_gate(mut self, gate: RuntimeSemaphore) -> Self {
self.distributed_request_gate = Some(Arc::new(gate)); self.distributed_request_gate = Some(Arc::new(gate));
self self
} }
@@ -624,7 +657,7 @@ return request_count
pub(crate) async fn distributed_request_concurrency_snapshot( pub(crate) async fn distributed_request_concurrency_snapshot(
&self, &self,
) -> Result<Option<DistributedConcurrencySnapshot>, DistributedConcurrencyError> { ) -> Result<Option<RuntimeSemaphoreSnapshot>, RuntimeSemaphoreError> {
match self.distributed_request_gate.as_ref() { match self.distributed_request_gate.as_ref() {
Some(gate) => gate.snapshot().await.map(Some), Some(gate) => gate.snapshot().await.map(Some),
None => Ok(None), None => Ok(None),
@@ -767,11 +800,62 @@ return request_count
} }
pub fn has_redis_data_backend(&self) -> bool { pub fn has_redis_data_backend(&self) -> bool {
self.data.has_redis_backend() self.runtime_state.is_redis()
} }
pub(crate) fn redis_kv_runner(&self) -> Option<aether_data::driver::redis::RedisKvRunner> { pub(crate) fn runtime_state_backend(&self) -> &'static str {
self.data.kv_runner() self.runtime_state.backend_kind().as_str()
}
pub fn runtime_state(&self) -> &RuntimeState {
self.runtime_state.as_ref()
}
pub(crate) async fn runtime_kv_setex(
&self,
key: &str,
value: &str,
ttl_seconds: u64,
) -> Result<(), GatewayError> {
self.runtime_state
.kv_set(
key,
value.to_string(),
Some(Duration::from_secs(ttl_seconds)),
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn runtime_kv_get(&self, key: &str) -> Result<Option<String>, GatewayError> {
self.runtime_state
.kv_get(key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn runtime_kv_getdel(
&self,
key: &str,
) -> Result<Option<String>, GatewayError> {
self.runtime_state
.kv_take(key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn runtime_kv_del(&self, key: &str) -> Result<bool, GatewayError> {
self.runtime_state
.kv_delete(key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn runtime_kv_exists(&self, key: &str) -> Result<bool, GatewayError> {
self.runtime_state
.kv_exists(key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
} }
pub(crate) fn remove_scheduler_affinity_cache_entry(&self, cache_key: &str) -> bool { pub(crate) fn remove_scheduler_affinity_cache_entry(&self, cache_key: &str) -> bool {
@@ -159,19 +159,19 @@ impl ModelFetchRuntimeState for AppState {
key_id: &str, key_id: &str,
cached_models: &[Value], cached_models: &[Value],
) { ) {
let Some(runner) = AppState::redis_kv_runner(self) else {
return;
};
let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(cached_models)) let Ok(serialized) = serde_json::to_string(&aggregate_models_for_cache(cached_models))
else { else {
return; return;
}; };
let cache_key = format!("upstream_models:{provider_id}:{key_id}"); let cache_key = format!("upstream_models:{provider_id}:{key_id}");
if let Err(err) = runner if let Err(err) = self
.setex( .runtime_state
.kv_set(
&cache_key, &cache_key,
&serialized, serialized,
Some(model_fetch_interval_minutes().saturating_mul(60)), Some(std::time::Duration::from_secs(
model_fetch_interval_minutes().saturating_mul(60),
)),
) )
.await .await
{ {
+4 -4
View File
@@ -833,7 +833,7 @@ impl AppState {
&self, &self,
transport: &provider_transport::GatewayProviderTransportSnapshot, transport: &provider_transport::GatewayProviderTransportSnapshot,
) -> Result<Option<provider_transport::LocalResolvedOAuthRequestAuth>, GatewayError> { ) -> Result<Option<provider_transport::LocalResolvedOAuthRequestAuth>, GatewayError> {
let distributed_lock = self.data.oauth_refresh_lock_runner(); let distributed_lock = self.runtime_state.as_ref();
let lock_owner = format!("aether-gateway-{}", std::process::id()); let lock_owner = format!("aether-gateway-{}", std::process::id());
let mut current_transport = transport.clone(); let mut current_transport = transport.clone();
let executor = GatewayLocalOAuthHttpExecutor { state: self }; let executor = GatewayLocalOAuthHttpExecutor { state: self };
@@ -844,7 +844,7 @@ impl AppState {
.resolve_with_result( .resolve_with_result(
&executor, &executor,
&current_transport, &current_transport,
distributed_lock.as_ref(), Some(distributed_lock),
Some(lock_owner.as_str()), Some(lock_owner.as_str()),
) )
.await .await
@@ -929,7 +929,7 @@ impl AppState {
Option<provider_transport::CachedOAuthEntry>, Option<provider_transport::CachedOAuthEntry>,
provider_transport::LocalOAuthRefreshError, provider_transport::LocalOAuthRefreshError,
> { > {
let distributed_lock = self.data.oauth_refresh_lock_runner(); let distributed_lock = self.runtime_state.as_ref();
let lock_owner = format!("aether-gateway-admin-{}", std::process::id()); let lock_owner = format!("aether-gateway-admin-{}", std::process::id());
let mut current_transport = transport.clone(); let mut current_transport = transport.clone();
current_transport.key.decrypted_api_key = "__placeholder__".to_string(); current_transport.key.decrypted_api_key = "__placeholder__".to_string();
@@ -958,7 +958,7 @@ impl AppState {
.force_refresh_with_result( .force_refresh_with_result(
&executor, &executor,
&current_transport, &current_transport,
distributed_lock.as_ref(), Some(distributed_lock),
Some(lock_owner.as_str()), Some(lock_owner.as_str()),
) )
.await?; .await?;
@@ -77,16 +77,21 @@ impl AppState {
value: &str, value: &str,
ttl_seconds: u64, ttl_seconds: u64,
) -> Result<(), GatewayError> { ) -> Result<(), GatewayError> {
self.data self.runtime_state
.cache_set_string_with_ttl(key, value, ttl_seconds) .kv_set(
key,
value.to_string(),
Some(std::time::Duration::from_secs(ttl_seconds)),
)
.await .await
.map_err(|err| GatewayError::Internal(err.to_string())) .map_err(|err| GatewayError::Internal(err.to_string()))
} }
pub(crate) async fn cache_delete_key(&self, key: &str) -> Result<(), GatewayError> { pub(crate) async fn cache_delete_key(&self, key: &str) -> Result<(), GatewayError> {
self.data self.runtime_state
.cache_delete_key(key) .kv_delete(key)
.await .await
.map(|_| ())
.map_err(|err| GatewayError::Internal(err.to_string())) .map_err(|err| GatewayError::Internal(err.to_string()))
} }
} }
+1 -1
View File
@@ -67,7 +67,7 @@ impl AppState {
} }
pub fn has_usage_worker_backend(&self) -> bool { pub fn has_usage_worker_backend(&self) -> bool {
self.data.has_usage_worker_runner() self.data.has_usage_worker_queue()
} }
pub fn has_wallet_data_reader(&self) -> bool { pub fn has_wallet_data_reader(&self) -> bool {
+53 -253
View File
@@ -10,41 +10,16 @@ impl AppState {
) -> Result<bool, GatewayError> { ) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:"; const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
if let Some(runner) = self.redis_kv_runner() { let key = format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}{ip_address}");
let mut connection = match runner.client().get_multiplexed_async_connection().await { self.runtime_state
Ok(value) => value, .kv_set(
Err(_) => return Ok(false), &key,
}; reason.to_string(),
let key = runner ttl_seconds.map(std::time::Duration::from_secs),
.keyspace() )
.key(&format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}{ip_address}"));
let result = if let Some(ttl_seconds) = ttl_seconds {
redis::cmd("SETEX")
.arg(&key)
.arg(ttl_seconds)
.arg(reason)
.query_async::<String>(&mut connection)
.await .await
} else { .map(|_| true)
redis::cmd("SET") .map_err(|err| GatewayError::Internal(err.to_string()))
.arg(&key)
.arg(reason)
.query_async::<String>(&mut connection)
.await
};
return Ok(result.is_ok());
}
#[cfg(test)]
if let Some(store) = self.admin_security_blacklist_store.as_ref() {
store
.lock()
.expect("admin security blacklist store should lock")
.insert(ip_address.to_string(), reason.to_string());
return Ok(true);
}
Ok(false)
} }
pub(crate) async fn remove_admin_security_blacklist( pub(crate) async fn remove_admin_security_blacklist(
@@ -53,36 +28,11 @@ impl AppState {
) -> Result<bool, GatewayError> { ) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:"; const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
if let Some(runner) = self.redis_kv_runner() { let key = format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}{ip_address}");
let mut connection = match runner.client().get_multiplexed_async_connection().await { self.runtime_state
Ok(value) => value, .kv_delete(&key)
Err(_) => return Ok(false),
};
let key = runner
.keyspace()
.key(&format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}{ip_address}"));
let deleted = match redis::cmd("DEL")
.arg(&key)
.query_async::<i64>(&mut connection)
.await .await
{ .map_err(|err| GatewayError::Internal(err.to_string()))
Ok(value) => value,
Err(_) => return Ok(false),
};
return Ok(deleted > 0);
}
#[cfg(test)]
if let Some(store) = self.admin_security_blacklist_store.as_ref() {
let removed = store
.lock()
.expect("admin security blacklist store should lock")
.remove(ip_address)
.is_some();
return Ok(removed);
}
Ok(false)
} }
pub(crate) async fn admin_security_blacklist_stats( pub(crate) async fn admin_security_blacklist_stats(
@@ -90,48 +40,13 @@ impl AppState {
) -> Result<(bool, usize, Option<String>), GatewayError> { ) -> Result<(bool, usize, Option<String>), GatewayError> {
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:"; const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
if let Some(runner) = self.redis_kv_runner() { let total = self
let mut connection = match runner.client().get_multiplexed_async_connection().await { .runtime_state
Ok(value) => value, .scan_keys(&format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}*"), 100)
Err(_) => return Ok((false, 0, Some("Redis 不可用".to_string()))),
};
let pattern = runner
.keyspace()
.key(&format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}*"));
let mut cursor = 0u64;
let mut total = 0usize;
loop {
let (next_cursor, keys) = match redis::cmd("SCAN")
.arg(cursor)
.arg("MATCH")
.arg(&pattern)
.arg("COUNT")
.arg(100)
.query_async::<(u64, Vec<String>)>(&mut connection)
.await .await
{ .map(|keys| keys.len())
Ok(value) => value, .map_err(|err| GatewayError::Internal(err.to_string()))?;
Err(err) => return Ok((false, 0, Some(err.to_string()))), Ok((true, total, None))
};
total += keys.len();
if next_cursor == 0 {
break;
}
cursor = next_cursor;
}
return Ok((true, total, None));
}
#[cfg(test)]
if let Some(store) = self.admin_security_blacklist_store.as_ref() {
let total = store
.lock()
.expect("admin security blacklist store should lock")
.len();
return Ok((true, total, None));
}
Ok((false, 0, Some("Redis 不可用".to_string())))
} }
pub(crate) async fn list_admin_security_blacklist( pub(crate) async fn list_admin_security_blacklist(
@@ -139,83 +54,40 @@ impl AppState {
) -> Result<Vec<AdminSecurityBlacklistEntry>, GatewayError> { ) -> Result<Vec<AdminSecurityBlacklistEntry>, GatewayError> {
const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:"; const ADMIN_SECURITY_BLACKLIST_PREFIX: &str = "ip:blacklist:";
if let Some(runner) = self.redis_kv_runner() { let keys = self
let mut connection = match runner.client().get_multiplexed_async_connection().await { .runtime_state
Ok(value) => value, .scan_keys(&format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}*"), 100)
Err(_) => return Ok(Vec::new()), .await
}; .map_err(|err| GatewayError::Internal(err.to_string()))?;
let pattern = runner
.keyspace()
.key(&format!("{ADMIN_SECURITY_BLACKLIST_PREFIX}*"));
let prefix = runner.keyspace().key(ADMIN_SECURITY_BLACKLIST_PREFIX);
let mut cursor = 0u64;
let mut entries = Vec::new(); let mut entries = Vec::new();
loop {
let (next_cursor, keys) = match redis::cmd("SCAN")
.arg(cursor)
.arg("MATCH")
.arg(&pattern)
.arg("COUNT")
.arg(100)
.query_async::<(u64, Vec<String>)>(&mut connection)
.await
{
Ok(value) => value,
Err(_) => break,
};
for full_key in keys { for full_key in keys {
let ip_address = full_key let raw_key = self.runtime_state.strip_namespace(&full_key);
.strip_prefix(prefix.as_str()) let ip_address = raw_key
.map(|value| value.to_string()) .strip_prefix(ADMIN_SECURITY_BLACKLIST_PREFIX)
.unwrap_or_else(|| full_key.clone()); .unwrap_or(raw_key)
let reason: Result<String, _> = redis::cmd("GET") .to_string();
.arg(&full_key) let Some(reason) = self
.query_async(&mut connection) .runtime_state
.await; .kv_get(raw_key)
let reason = match reason {
Ok(value) => value,
Err(_) => continue,
};
let ttl = match redis::cmd("TTL")
.arg(&full_key)
.query_async::<i64>(&mut connection)
.await .await
{ .map_err(|err| GatewayError::Internal(err.to_string()))?
Ok(value) if value >= 0 => Some(value), else {
_ => None, continue;
}; };
let ttl_seconds = self
.runtime_state
.kv_ttl_seconds(raw_key)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
.filter(|ttl| *ttl >= 0);
entries.push(AdminSecurityBlacklistEntry { entries.push(AdminSecurityBlacklistEntry {
ip_address, ip_address,
reason, reason,
ttl_seconds: ttl, ttl_seconds,
}); });
} }
if next_cursor == 0 {
break;
}
cursor = next_cursor;
}
entries.sort_by(|a, b| a.ip_address.cmp(&b.ip_address)); entries.sort_by(|a, b| a.ip_address.cmp(&b.ip_address));
return Ok(entries); Ok(entries)
}
#[cfg(test)]
if let Some(store) = self.admin_security_blacklist_store.as_ref() {
let mut entries = store
.lock()
.expect("admin security blacklist store should lock")
.iter()
.map(|(ip, reason)| AdminSecurityBlacklistEntry {
ip_address: ip.clone(),
reason: reason.clone(),
ttl_seconds: None,
})
.collect::<Vec<_>>();
entries.sort_by(|a, b| a.ip_address.cmp(&b.ip_address));
return Ok(entries);
}
Ok(Vec::new())
} }
pub(crate) async fn add_admin_security_whitelist( pub(crate) async fn add_admin_security_whitelist(
@@ -224,34 +96,11 @@ impl AppState {
) -> Result<bool, GatewayError> { ) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist"; const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist";
if let Some(runner) = self.redis_kv_runner() { self.runtime_state
let mut connection = match runner.client().get_multiplexed_async_connection().await { .set_add(ADMIN_SECURITY_WHITELIST_KEY, ip_address)
Ok(value) => value,
Err(_) => return Ok(false),
};
let key = runner.keyspace().key(ADMIN_SECURITY_WHITELIST_KEY);
let added = match redis::cmd("SADD")
.arg(&key)
.arg(ip_address)
.query_async::<i64>(&mut connection)
.await .await
{ .map(|_| true)
Ok(value) => value, .map_err(|err| GatewayError::Internal(err.to_string()))
Err(_) => return Ok(false),
};
return Ok(added >= 0);
}
#[cfg(test)]
if let Some(store) = self.admin_security_whitelist_store.as_ref() {
store
.lock()
.expect("admin security whitelist store should lock")
.insert(ip_address.to_string());
return Ok(true);
}
Ok(false)
} }
pub(crate) async fn remove_admin_security_whitelist( pub(crate) async fn remove_admin_security_whitelist(
@@ -260,67 +109,18 @@ impl AppState {
) -> Result<bool, GatewayError> { ) -> Result<bool, GatewayError> {
const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist"; const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist";
if let Some(runner) = self.redis_kv_runner() { self.runtime_state
let mut connection = match runner.client().get_multiplexed_async_connection().await { .set_remove(ADMIN_SECURITY_WHITELIST_KEY, ip_address)
Ok(value) => value,
Err(_) => return Ok(false),
};
let key = runner.keyspace().key(ADMIN_SECURITY_WHITELIST_KEY);
let removed = match redis::cmd("SREM")
.arg(&key)
.arg(ip_address)
.query_async::<i64>(&mut connection)
.await .await
{ .map_err(|err| GatewayError::Internal(err.to_string()))
Ok(value) => value,
Err(_) => return Ok(false),
};
return Ok(removed > 0);
}
#[cfg(test)]
if let Some(store) = self.admin_security_whitelist_store.as_ref() {
let removed = store
.lock()
.expect("admin security whitelist store should lock")
.remove(ip_address);
return Ok(removed);
}
Ok(false)
} }
pub(crate) async fn list_admin_security_whitelist(&self) -> Result<Vec<String>, GatewayError> { pub(crate) async fn list_admin_security_whitelist(&self) -> Result<Vec<String>, GatewayError> {
const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist"; const ADMIN_SECURITY_WHITELIST_KEY: &str = "ip:whitelist";
if let Some(runner) = self.redis_kv_runner() { self.runtime_state
let mut connection = match runner.client().get_multiplexed_async_connection().await { .set_members(ADMIN_SECURITY_WHITELIST_KEY)
Ok(value) => value,
Err(_) => return Ok(Vec::new()),
};
let key = runner.keyspace().key(ADMIN_SECURITY_WHITELIST_KEY);
let mut whitelist = match redis::cmd("SMEMBERS")
.arg(&key)
.query_async::<Vec<String>>(&mut connection)
.await .await
{ .map_err(|err| GatewayError::Internal(err.to_string()))
Ok(value) => value,
Err(_) => return Ok(Vec::new()),
};
whitelist.sort();
return Ok(whitelist);
}
#[cfg(test)]
if let Some(store) = self.admin_security_whitelist_store.as_ref() {
return Ok(store
.lock()
.expect("admin security whitelist store should lock")
.iter()
.cloned()
.collect());
}
Ok(Vec::new())
} }
} }
+41
View File
@@ -1,5 +1,6 @@
use std::collections::HashMap; use std::collections::HashMap;
use std::sync::{Arc, Mutex as StdMutex}; use std::sync::{Arc, Mutex as StdMutex};
use std::time::Duration;
use aether_contracts::{ExecutionPlan, ExecutionResult}; use aether_contracts::{ExecutionPlan, ExecutionResult};
use aether_data_contracts::repository::candidates::RequestCandidateReadRepository; use aether_data_contracts::repository::candidates::RequestCandidateReadRepository;
@@ -207,6 +208,13 @@ impl AppState {
.lock() .lock()
.expect("provider oauth state store should lock") .expect("provider oauth state store should lock")
.insert(format!("provider_oauth_state:{nonce}"), payload.to_string()); .insert(format!("provider_oauth_state:{nonce}"), payload.to_string());
self.runtime_state.kv_set_local_nowait(
&format!("provider_oauth_state:{nonce}"),
payload.to_string(),
Some(Duration::from_secs(
aether_data::repository::provider_oauth::PROVIDER_OAUTH_STATE_TTL_SECS,
)),
);
self self
} }
@@ -225,6 +233,11 @@ impl AppState {
format!("device_auth_session:{session_id}"), format!("device_auth_session:{session_id}"),
payload.to_string(), payload.to_string(),
); );
self.runtime_state.kv_set_local_nowait(
&format!("device_auth_session:{session_id}"),
payload.to_string(),
Some(Duration::from_secs(3600)),
);
self self
} }
@@ -243,6 +256,13 @@ impl AppState {
format!("provider_oauth_batch_task:{task_id}"), format!("provider_oauth_batch_task:{task_id}"),
payload.to_string(), payload.to_string(),
); );
self.runtime_state.kv_set_local_nowait(
&format!("provider_oauth_batch_task:{task_id}"),
payload.to_string(),
Some(Duration::from_secs(
aether_data::repository::provider_oauth::PROVIDER_OAUTH_BATCH_TASK_TTL_SECS,
)),
);
self self
} }
@@ -417,6 +437,11 @@ impl AppState {
.lock() .lock()
.expect("admin security blacklist store should lock"); .expect("admin security blacklist store should lock");
for (ip_address, reason) in entries { for (ip_address, reason) in entries {
self.runtime_state.kv_set_local_nowait(
&format!("ip:blacklist:{ip_address}"),
reason.clone(),
None,
);
guard.insert(ip_address, reason); guard.insert(ip_address, reason);
} }
drop(guard); drop(guard);
@@ -434,6 +459,8 @@ impl AppState {
.lock() .lock()
.expect("admin security whitelist store should lock"); .expect("admin security whitelist store should lock");
for ip_address in entries { for ip_address in entries {
self.runtime_state
.set_add_local_nowait("ip:whitelist", &ip_address);
guard.insert(ip_address); guard.insert(ip_address);
} }
drop(guard); drop(guard);
@@ -552,6 +579,15 @@ impl AppState {
}) })
.to_string(), .to_string(),
); );
self.runtime_state.kv_set_local_nowait(
&format!("email:verification:{}", email.trim().to_ascii_lowercase()),
json!({
"code": code,
"created_at": created_at.to_rfc3339(),
})
.to_string(),
Some(Duration::from_secs(600)),
);
self self
} }
@@ -566,6 +602,11 @@ impl AppState {
format!("email:verified:{}", email.trim().to_ascii_lowercase()), format!("email:verified:{}", email.trim().to_ascii_lowercase()),
"verified".to_string(), "verified".to_string(),
); );
self.runtime_state.kv_set_local_nowait(
&format!("email:verified:{}", email.trim().to_ascii_lowercase()),
"verified".to_string(),
Some(Duration::from_secs(3600)),
);
self self
} }
@@ -2,6 +2,15 @@ use std::path::{Path, PathBuf};
use super::*; use super::*;
fn production_workspace_source(path: &Path) -> String {
let source = std::fs::read_to_string(path).expect("source file should be readable");
source
.split("#[cfg(test)]")
.next()
.unwrap_or(&source)
.to_string()
}
#[test] #[test]
fn gateway_small_runtime_shims_stay_deleted() { fn gateway_small_runtime_shims_stay_deleted() {
for path in [ for path in [
@@ -79,6 +88,120 @@ fn gateway_small_runtime_shims_stay_deleted() {
} }
} }
#[test]
fn runtime_state_owns_redis_runtime_boundaries() {
let forbidden_business_patterns = [
"use redis::",
"redis::cmd",
"::redis::cmd",
"::redis::Script",
"aether_data::driver::redis",
"RedisKvRunner",
"RedisLockRunner",
"RedisStreamRunner",
"redis_kv_runner(",
];
let mut violations = Vec::new();
for root in [
"apps/aether-gateway/src",
"crates/aether-runtime/src",
"crates/aether-usage-runtime/src",
"crates/aether-provider-transport/src",
] {
for path in collect_workspace_rust_files(root) {
if path
.components()
.any(|component| component.as_os_str() == "tests")
{
continue;
}
let source = production_workspace_source(&path);
let hits = forbidden_business_patterns
.iter()
.filter(|pattern| source.contains(**pattern))
.copied()
.collect::<Vec<_>>();
if !hits.is_empty() {
violations.push(format!("{} -> {}", path.display(), hits.join(", ")));
}
}
}
assert!(
violations.is_empty(),
"business/runtime crates must use aether-runtime-state instead of Redis directly:\n{}",
violations.join("\n")
);
let mut runtime_state_violations = Vec::new();
for path in collect_workspace_rust_files("crates/aether-runtime-state/src") {
if path
.components()
.any(|component| component.as_os_str() == "redis")
{
continue;
}
let source = production_workspace_source(&path);
let hits = [
"use redis::",
"redis::cmd",
"::redis::cmd",
"::redis::Script",
]
.iter()
.filter(|pattern| source.contains(**pattern))
.copied()
.collect::<Vec<_>>();
if !hits.is_empty() {
runtime_state_violations.push(format!("{} -> {}", path.display(), hits.join(", ")));
}
}
assert!(
runtime_state_violations.is_empty(),
"only crates/aether-runtime-state/src/redis may depend on the redis crate directly:\n{}",
runtime_state_violations.join("\n")
);
}
#[test]
fn aether_data_stays_free_of_redis_runtime_backends() {
let cargo = read_workspace_file("crates/aether-data/Cargo.toml");
assert!(
!cargo.contains("redis.workspace"),
"aether-data should not depend on redis; runtime Redis belongs to aether-runtime-state"
);
for removed_path in [
"crates/aether-data/src/backend/redis.rs",
"crates/aether-data/src/backend/locks.rs",
"crates/aether-data/src/backend/workers.rs",
"crates/aether-data/src/driver/redis/mod.rs",
] {
assert!(
!workspace_file_exists(removed_path),
"{removed_path} should stay removed from aether-data"
);
}
for path in collect_workspace_rust_files("crates/aether-data/src") {
let source = production_workspace_source(&path);
for forbidden in [
"pub mod redis",
"driver::redis",
"RedisBackend",
"DataLockBackends",
"DataWorkerBackends",
"redis::cmd",
"use redis::",
] {
assert!(
!source.contains(forbidden),
"{} should not keep Redis runtime backend surface {forbidden}",
path.display()
);
}
}
}
#[test] #[test]
fn gateway_request_candidate_trace_type_is_owned_by_aether_data_contracts() { fn gateway_request_candidate_trace_type_is_owned_by_aether_data_contracts() {
let gateway_candidates = read_workspace_file("apps/aether-gateway/src/data/candidates.rs"); let gateway_candidates = read_workspace_file("apps/aether-gateway/src/data/candidates.rs");
+12 -8
View File
@@ -17,9 +17,18 @@ use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelect
use aether_data::repository::candidates::InMemoryRequestCandidateRepository; use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::usage::InMemoryUsageReadRepository; use aether_data::repository::usage::InMemoryUsageReadRepository;
use aether_runtime_state::{
MemoryRuntimeStateConfig, RuntimeSemaphore, RuntimeSemaphoreConfig, RuntimeState,
};
use crate::data::GatewayDataState; use crate::data::GatewayDataState;
fn memory_runtime_semaphore(gate: &'static str, limit: usize) -> RuntimeSemaphore {
RuntimeState::memory(MemoryRuntimeStateConfig::default())
.semaphore(gate, limit, RuntimeSemaphoreConfig::default())
.expect("memory runtime semaphore should build")
}
fn sample_decision() -> crate::control::GatewayControlDecision { fn sample_decision() -> crate::control::GatewayControlDecision {
crate::control::GatewayControlDecision { crate::control::GatewayControlDecision {
public_path: "/v1/chat/completions".to_string(), public_path: "/v1/chat/completions".to_string(),
@@ -104,10 +113,7 @@ async fn gateway_rejects_second_in_flight_stream_request_with_distributed_overlo
); );
let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await; let (execution_runtime_url, execution_runtime_handle) = start_server(execution_runtime).await;
let distributed_gate = aether_runtime::DistributedConcurrencyGate::new_in_memory( let distributed_gate = memory_runtime_semaphore("gateway_requests_distributed", 1);
"gateway_requests_distributed",
1,
);
let gateway_a = build_router_with_state( let gateway_a = build_router_with_state(
build_local_openai_gateway_state(execution_runtime_url.clone()) build_local_openai_gateway_state(execution_runtime_url.clone())
.with_distributed_request_concurrency_gate(distributed_gate.clone()), .with_distributed_request_concurrency_gate(distributed_gate.clone()),
@@ -262,12 +268,10 @@ async fn gateway_exposes_request_concurrency_metrics() {
AppState::new() AppState::new()
.expect("gateway state should build") .expect("gateway state should build")
.with_request_concurrency_limit(3) .with_request_concurrency_limit(3)
.with_distributed_request_concurrency_gate( .with_distributed_request_concurrency_gate(memory_runtime_semaphore(
aether_runtime::DistributedConcurrencyGate::new_in_memory(
"gateway_requests_distributed", "gateway_requests_distributed",
5, 5,
), )),
),
); );
let (gateway_url, gateway_handle) = start_server(gateway).await; let (gateway_url, gateway_handle) = start_server(gateway).await;
@@ -130,7 +130,7 @@ async fn gateway_clears_admin_external_models_cache_locally_with_trusted_admin_p
assert_eq!(response.status(), StatusCode::OK); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["cleared"], false); assert_eq!(payload["cleared"], false);
assert_eq!(payload["message"], "Redis 未启用"); assert_eq!(payload["message"], "缓存不存在");
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort(); gateway_handle.abort();
@@ -1307,9 +1307,11 @@ async fn gateway_handles_admin_monitoring_cache_redis_keys_delete_locally_with_t
.await .await
.expect("request should succeed"); .expect("request should succeed");
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["detail"], json!("Redis 未启用")); assert_eq!(payload["status"], json!("ok"));
assert_eq!(payload["category"], json!("upstream_models"));
assert_eq!(payload["deleted_count"], json!(0));
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort(); gateway_handle.abort();
@@ -1485,11 +1487,9 @@ async fn gateway_handles_admin_monitoring_model_mapping_stats_locally_with_trust
assert_eq!(response.status(), StatusCode::OK); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["status"], json!("ok")); assert_eq!(payload["status"], json!("ok"));
assert_eq!(payload["data"]["available"], json!(false)); assert_eq!(payload["data"]["available"], json!(true));
assert_eq!( assert_eq!(payload["data"]["backend"], json!("memory"));
payload["data"]["message"], assert_eq!(payload["data"]["total_keys"], json!(0));
json!("Redis 未启用,模型映射缓存不可用")
);
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort(); gateway_handle.abort();
@@ -1707,8 +1707,9 @@ async fn gateway_handles_admin_monitoring_redis_keys_locally_with_trusted_admin_
assert_eq!(response.status(), StatusCode::OK); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["status"], json!("ok")); assert_eq!(payload["status"], json!("ok"));
assert_eq!(payload["data"]["available"], json!(false)); assert_eq!(payload["data"]["available"], json!(true));
assert_eq!(payload["data"]["message"], json!("Redis 未启用")); assert_eq!(payload["data"]["backend"], json!("memory"));
assert_eq!(payload["data"]["total_keys"], json!(0));
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort(); gateway_handle.abort();
@@ -1118,9 +1118,16 @@ async fn gateway_handles_admin_provider_oauth_start_key_locally_with_trusted_adm
.await .await
.expect("request should succeed"); .expect("request should succeed");
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["detail"], "provider oauth redis unavailable"); assert_eq!(payload["provider_type"], "codex");
assert_eq!(
payload["redirect_uri"],
"http://localhost:1455/auth/callback"
);
assert!(payload["authorization_url"]
.as_str()
.is_some_and(|url| url.contains("state=")));
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort(); gateway_handle.abort();
@@ -1174,9 +1181,16 @@ async fn gateway_handles_admin_provider_oauth_start_provider_locally_with_truste
.await .await
.expect("request should succeed"); .expect("request should succeed");
assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); assert_eq!(response.status(), StatusCode::OK);
let payload: serde_json::Value = response.json().await.expect("json body should parse"); let payload: serde_json::Value = response.json().await.expect("json body should parse");
assert_eq!(payload["detail"], "provider oauth redis unavailable"); assert_eq!(payload["provider_type"], "codex");
assert_eq!(
payload["redirect_uri"],
"http://localhost:1455/auth/callback"
);
assert!(payload["authorization_url"]
.as_str()
.is_some_and(|url| url.contains("state=")));
assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0); assert_eq!(*upstream_hits.lock().expect("mutex should lock"), 0);
gateway_handle.abort(); gateway_handle.abort();
@@ -9,6 +9,7 @@ use aether_crypto::{
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository; use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository; use aether_data::repository::proxy_nodes::InMemoryProxyNodeRepository;
use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository; use aether_data_contracts::repository::provider_catalog::ProviderCatalogReadRepository;
use aether_runtime_state::{RedisClientConfig, RuntimeState};
use aether_testkit::ManagedRedisServer; use aether_testkit::ManagedRedisServer;
use axum::body::to_bytes; use axum::body::to_bytes;
use axum::body::Body; use axum::body::Body;
@@ -42,6 +43,23 @@ async fn start_managed_redis_or_skip() -> Option<ManagedRedisServer> {
} }
} }
async fn redis_runtime_state_for_test(
redis: &ManagedRedisServer,
key_prefix: &str,
) -> Arc<RuntimeState> {
Arc::new(
RuntimeState::redis(
RedisClientConfig {
url: redis.redis_url().to_string(),
key_prefix: Some(key_prefix.to_string()),
},
None,
)
.await
.expect("redis runtime state should build"),
)
}
#[tokio::test] #[tokio::test]
async fn gateway_handles_admin_provider_ops_architectures_locally_with_trusted_admin_principal() { async fn gateway_handles_admin_provider_ops_architectures_locally_with_trusted_admin_principal() {
let upstream_hits = Arc::new(Mutex::new(0usize)); let upstream_hits = Arc::new(Mutex::new(0usize));
@@ -3529,16 +3547,16 @@ async fn gateway_handles_admin_provider_ops_balance_cache_refresh_modes_with_red
vec![], vec![],
)); ));
let data_state = GatewayDataState::from_config( let data_state = GatewayDataState::from_config(
GatewayDataConfig::disabled() GatewayDataConfig::disabled().with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
.with_redis_url(redis.redis_url(), Some("provider_ops_balance_cache"))
.with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
) )
.expect("data state should build") .expect("data state should build")
.attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository)); .attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository));
let runtime_state = redis_runtime_state_for_test(&redis, "provider_ops_balance_cache").await;
let gateway = build_router_with_state( let gateway = build_router_with_state(
AppState::new() AppState::new()
.expect("gateway should build") .expect("gateway should build")
.with_data_state_for_tests(data_state), .with_data_state_for_tests(data_state)
.with_runtime_state(runtime_state),
); );
let (gateway_url, gateway_handle) = start_server(gateway).await; let (gateway_url, gateway_handle) = start_server(gateway).await;
let client = reqwest::Client::new(); let client = reqwest::Client::new();
@@ -3683,19 +3701,17 @@ async fn gateway_handles_admin_provider_ops_balance_cache_miss_without_refresh_r
vec![], vec![],
)); ));
let data_state = GatewayDataState::from_config( let data_state = GatewayDataState::from_config(
GatewayDataConfig::disabled() GatewayDataConfig::disabled().with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
.with_redis_url(
redis.redis_url(),
Some("provider_ops_balance_cache_sync_miss"),
)
.with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
) )
.expect("data state should build") .expect("data state should build")
.attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository)); .attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository));
let runtime_state =
redis_runtime_state_for_test(&redis, "provider_ops_balance_cache_sync_miss").await;
let gateway = build_router_with_state( let gateway = build_router_with_state(
AppState::new() AppState::new()
.expect("gateway should build") .expect("gateway should build")
.with_data_state_for_tests(data_state), .with_data_state_for_tests(data_state)
.with_runtime_state(runtime_state),
); );
let (gateway_url, gateway_handle) = start_server(gateway).await; let (gateway_url, gateway_handle) = start_server(gateway).await;
let client = reqwest::Client::new(); let client = reqwest::Client::new();
@@ -3841,19 +3857,17 @@ async fn gateway_clears_admin_provider_ops_balance_cache_after_config_save_with_
vec![], vec![],
)); ));
let data_state = GatewayDataState::from_config( let data_state = GatewayDataState::from_config(
GatewayDataConfig::disabled() GatewayDataConfig::disabled().with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
.with_redis_url(
redis.redis_url(),
Some("provider_ops_balance_cache_config_save"),
)
.with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
) )
.expect("data state should build") .expect("data state should build")
.attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository)); .attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository));
let runtime_state =
redis_runtime_state_for_test(&redis, "provider_ops_balance_cache_config_save").await;
let gateway = build_router_with_state( let gateway = build_router_with_state(
AppState::new() AppState::new()
.expect("gateway should build") .expect("gateway should build")
.with_data_state_for_tests(data_state), .with_data_state_for_tests(data_state)
.with_runtime_state(runtime_state),
); );
let (gateway_url, gateway_handle) = start_server(gateway).await; let (gateway_url, gateway_handle) = start_server(gateway).await;
let client = reqwest::Client::new(); let client = reqwest::Client::new();
@@ -4058,19 +4072,17 @@ async fn gateway_verify_does_not_pollute_balance_cache_and_balance_uses_saved_ac
vec![], vec![],
)); ));
let data_state = GatewayDataState::from_config( let data_state = GatewayDataState::from_config(
GatewayDataConfig::disabled() GatewayDataConfig::disabled().with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
.with_redis_url(
redis.redis_url(),
Some("provider_ops_verify_cache_isolation"),
)
.with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
) )
.expect("data state should build") .expect("data state should build")
.attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository)); .attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository));
let runtime_state =
redis_runtime_state_for_test(&redis, "provider_ops_verify_cache_isolation").await;
let gateway = build_router_with_state( let gateway = build_router_with_state(
AppState::new() AppState::new()
.expect("gateway should build") .expect("gateway should build")
.with_data_state_for_tests(data_state), .with_data_state_for_tests(data_state)
.with_runtime_state(runtime_state),
); );
let (gateway_url, gateway_handle) = start_server(gateway).await; let (gateway_url, gateway_handle) = start_server(gateway).await;
let client = reqwest::Client::new(); let client = reqwest::Client::new();
@@ -4232,16 +4244,16 @@ async fn gateway_handles_admin_provider_ops_batch_balance_with_pending_cache_hit
vec![], vec![],
)); ));
let data_state = GatewayDataState::from_config( let data_state = GatewayDataState::from_config(
GatewayDataConfig::disabled() GatewayDataConfig::disabled().with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
.with_redis_url(redis.redis_url(), Some("provider_ops_batch_balance"))
.with_encryption_key(DEVELOPMENT_ENCRYPTION_KEY),
) )
.expect("data state should build") .expect("data state should build")
.attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository)); .attach_provider_catalog_repository_for_tests(Arc::clone(&provider_catalog_repository));
let runtime_state = redis_runtime_state_for_test(&redis, "provider_ops_batch_balance").await;
let gateway = build_router_with_state( let gateway = build_router_with_state(
AppState::new() AppState::new()
.expect("gateway should build") .expect("gateway should build")
.with_data_state_for_tests(data_state), .with_data_state_for_tests(data_state)
.with_runtime_state(runtime_state),
); );
let (gateway_url, gateway_handle) = start_server(gateway).await; let (gateway_url, gateway_handle) = start_server(gateway).await;
let client = reqwest::Client::new(); let client = reqwest::Client::new();
@@ -134,16 +134,16 @@ fn map_request_admission_error(error: super::RequestAdmissionError) -> String {
.. ..
}) })
| super::RequestAdmissionError::Distributed( | super::RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::Saturated { .. }, aether_runtime_state::RuntimeSemaphoreError::Saturated { .. },
) )
| super::RequestAdmissionError::Distributed( | super::RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::Unavailable { .. }, aether_runtime_state::RuntimeSemaphoreError::Unavailable { .. },
) => "overloaded: hub relay overloaded".to_string(), ) => "overloaded: hub relay overloaded".to_string(),
super::RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Closed { super::RequestAdmissionError::Local(aether_runtime::ConcurrencyError::Closed {
.. ..
}) => "overloaded: hub relay gate closed".to_string(), }) => "overloaded: hub relay gate closed".to_string(),
super::RequestAdmissionError::Distributed( super::RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::InvalidConfiguration(_), aether_runtime_state::RuntimeSemaphoreError::InvalidConfiguration(_),
) => "overloaded: hub relay distributed gate invalid".to_string(), ) => "overloaded: hub relay distributed gate invalid".to_string(),
} }
} }
@@ -174,10 +174,10 @@ pub async fn relay_request(
.. ..
})) }))
| Err(super::RequestAdmissionError::Distributed( | Err(super::RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::Saturated { .. }, aether_runtime_state::RuntimeSemaphoreError::Saturated { .. },
)) ))
| Err(super::RequestAdmissionError::Distributed( | Err(super::RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::Unavailable { .. }, aether_runtime_state::RuntimeSemaphoreError::Unavailable { .. },
)) => { )) => {
return tunnel_error_response( return tunnel_error_response(
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
@@ -195,7 +195,7 @@ pub async fn relay_request(
); );
} }
Err(super::RequestAdmissionError::Distributed( Err(super::RequestAdmissionError::Distributed(
aether_runtime::DistributedConcurrencyError::InvalidConfiguration(_), aether_runtime_state::RuntimeSemaphoreError::InvalidConfiguration(_),
)) => { )) => {
return tunnel_error_response( return tunnel_error_response(
StatusCode::SERVICE_UNAVAILABLE, StatusCode::SERVICE_UNAVAILABLE,
+13 -16
View File
@@ -8,10 +8,9 @@ use std::sync::Arc;
use aether_runtime::{ use aether_runtime::{
hold_admission_permit_until, prometheus_response, service_up_sample, AdmissionPermit, hold_admission_permit_until, prometheus_response, service_up_sample, AdmissionPermit,
ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot, DistributedConcurrencyError, ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot, MetricKind, MetricLabel, MetricSample,
DistributedConcurrencyGate, DistributedConcurrencySnapshot, MetricKind, MetricLabel,
MetricSample,
}; };
use aether_runtime_state::{RuntimeSemaphore, RuntimeSemaphoreError, RuntimeSemaphoreSnapshot};
use axum::extract::ws::WebSocketUpgrade; use axum::extract::ws::WebSocketUpgrade;
use axum::extract::State; use axum::extract::State;
use axum::response::{IntoResponse, Json}; use axum::response::{IntoResponse, Json};
@@ -33,13 +32,13 @@ pub struct AppState {
pub max_streams: usize, pub max_streams: usize,
data: Arc<GatewayDataState>, data: Arc<GatewayDataState>,
request_gate: Option<Arc<ConcurrencyGate>>, request_gate: Option<Arc<ConcurrencyGate>>,
distributed_request_gate: Option<Arc<DistributedConcurrencyGate>>, distributed_request_gate: Option<Arc<RuntimeSemaphore>>,
} }
#[derive(Debug)] #[derive(Debug)]
enum RequestAdmissionError { enum RequestAdmissionError {
Local(ConcurrencyError), Local(ConcurrencyError),
Distributed(DistributedConcurrencyError), Distributed(RuntimeSemaphoreError),
} }
impl AppState { impl AppState {
@@ -70,7 +69,7 @@ impl AppState {
self self
} }
pub fn with_distributed_request_gate(mut self, gate: DistributedConcurrencyGate) -> Self { pub fn with_distributed_request_gate(mut self, gate: RuntimeSemaphore) -> Self {
self.distributed_request_gate = Some(Arc::new(gate)); self.distributed_request_gate = Some(Arc::new(gate));
self self
} }
@@ -81,7 +80,7 @@ impl AppState {
async fn distributed_request_concurrency_snapshot( async fn distributed_request_concurrency_snapshot(
&self, &self,
) -> Result<Option<DistributedConcurrencySnapshot>, DistributedConcurrencyError> { ) -> Result<Option<RuntimeSemaphoreSnapshot>, RuntimeSemaphoreError> {
match self.distributed_request_gate.as_ref() { match self.distributed_request_gate.as_ref() {
Some(gate) => gate.snapshot().await.map(Some), Some(gate) => gate.snapshot().await.map(Some),
None => Ok(None), None => Ok(None),
@@ -225,12 +224,10 @@ pub async fn ws_proxy(
let request_permit = match state.try_acquire_request_permit().await { let request_permit = match state.try_acquire_request_permit().await {
Ok(permit) => permit, Ok(permit) => permit,
Err(RequestAdmissionError::Local(ConcurrencyError::Saturated { .. })) Err(RequestAdmissionError::Local(ConcurrencyError::Saturated { .. }))
| Err(RequestAdmissionError::Distributed(DistributedConcurrencyError::Saturated { | Err(RequestAdmissionError::Distributed(RuntimeSemaphoreError::Saturated { .. }))
.. | Err(RequestAdmissionError::Distributed(RuntimeSemaphoreError::Unavailable { .. })) => {
})) return axum::http::StatusCode::SERVICE_UNAVAILABLE.into_response()
| Err(RequestAdmissionError::Distributed(DistributedConcurrencyError::Unavailable { }
..
})) => return axum::http::StatusCode::SERVICE_UNAVAILABLE.into_response(),
Err(RequestAdmissionError::Local(ConcurrencyError::Closed { gate })) => { Err(RequestAdmissionError::Local(ConcurrencyError::Closed { gate })) => {
warn!( warn!(
gate = gate, gate = gate,
@@ -238,9 +235,9 @@ pub async fn ws_proxy(
); );
return axum::http::StatusCode::SERVICE_UNAVAILABLE.into_response(); return axum::http::StatusCode::SERVICE_UNAVAILABLE.into_response();
} }
Err(RequestAdmissionError::Distributed( Err(RequestAdmissionError::Distributed(RuntimeSemaphoreError::InvalidConfiguration(
DistributedConcurrencyError::InvalidConfiguration(message), message,
)) => { ))) => {
warn!( warn!(
error = %message, error = %message,
"standalone tunnel relay distributed request gate is invalid" "standalone tunnel relay distributed request gate is invalid"
+52 -28
View File
@@ -14,6 +14,7 @@ use aether_data::repository::proxy_nodes::{
ProxyNodeHeartbeatMutation, ProxyNodeTunnelStatusMutation, StoredProxyNode, ProxyNodeHeartbeatMutation, ProxyNodeTunnelStatusMutation, StoredProxyNode,
}; };
use aether_runtime::MetricSample; use aether_runtime::MetricSample;
use aether_runtime_state::{MemoryRuntimeStateConfig, RuntimeState};
use async_stream::stream; use async_stream::stream;
use axum::body::{Body, Bytes}; use axum::body::{Body, Bytes};
use axum::extract::ws::WebSocketUpgrade; use axum::extract::ws::WebSocketUpgrade;
@@ -102,6 +103,7 @@ pub(crate) struct TunnelAttachmentRecord {
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub(crate) struct TunnelAttachmentDirectory { pub(crate) struct TunnelAttachmentDirectory {
identity: Arc<TunnelInstanceIdentity>, identity: Arc<TunnelInstanceIdentity>,
runtime_state: Arc<RuntimeState>,
} }
impl TunnelAttachmentDirectory { impl TunnelAttachmentDirectory {
@@ -118,6 +120,7 @@ impl TunnelAttachmentDirectory {
.map(|value| value.clamp(15, 3600)) .map(|value| value.clamp(15, 3600))
.unwrap_or(DEFAULT_ATTACHMENT_TTL_SECS), .unwrap_or(DEFAULT_ATTACHMENT_TTL_SECS),
}), }),
runtime_state: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())),
} }
} }
@@ -132,9 +135,15 @@ impl TunnelAttachmentDirectory {
relay_base_url: relay_base_url.map(Into::into), relay_base_url: relay_base_url.map(Into::into),
attachment_ttl_secs, attachment_ttl_secs,
}), }),
runtime_state: Arc::new(RuntimeState::memory(MemoryRuntimeStateConfig::default())),
} }
} }
fn with_runtime_state(mut self, runtime_state: Arc<RuntimeState>) -> Self {
self.runtime_state = runtime_state;
self
}
#[cfg(test)] #[cfg(test)]
pub(crate) fn for_tests( pub(crate) fn for_tests(
instance_id: &str, instance_id: &str,
@@ -257,7 +266,7 @@ impl TunnelAttachmentDirectory {
data: &GatewayDataState, data: &GatewayDataState,
node_id: &str, node_id: &str,
) -> Result<Option<TunnelAttachmentRecord>, String> { ) -> Result<Option<TunnelAttachmentRecord>, String> {
match self.read_attachment_record_from_redis(data, node_id).await { match self.read_attachment_record_from_runtime(node_id).await {
Ok(Some(record)) => return Ok(Some(record)), Ok(Some(record)) => return Ok(Some(record)),
Ok(None) => {} Ok(None) => {}
Err(error) => { Err(error) => {
@@ -272,28 +281,18 @@ impl TunnelAttachmentDirectory {
.await .await
} }
async fn read_attachment_record_from_redis( async fn read_attachment_record_from_runtime(
&self, &self,
data: &GatewayDataState,
node_id: &str, node_id: &str,
) -> Result<Option<TunnelAttachmentRecord>, String> { ) -> Result<Option<TunnelAttachmentRecord>, String> {
let Some(runner) = data.kv_runner() else { let raw = self
return Ok(None); .runtime_state
}; .kv_get(&tunnel_attachment_redis_key(node_id))
let mut connection = runner
.client()
.get_multiplexed_async_connection()
.await .await
.map_err(|err| format!("attachment redis connect failed: {err}"))?; .map_err(|err| format!("attachment runtime read failed: {err}"))?;
let namespaced_key = runner.keyspace().key(&tunnel_attachment_redis_key(node_id));
let raw = redis::cmd("GET")
.arg(&namespaced_key)
.query_async::<Option<String>>(&mut connection)
.await
.map_err(|err| format!("attachment redis read failed: {err}"))?;
raw.map(|value| { raw.map(|value| {
serde_json::from_str::<TunnelAttachmentRecord>(&value) serde_json::from_str::<TunnelAttachmentRecord>(&value)
.map_err(|err| format!("invalid redis tunnel attachment record: {err}")) .map_err(|err| format!("invalid runtime tunnel attachment record: {err}"))
}) })
.transpose() .transpose()
} }
@@ -323,22 +322,21 @@ impl TunnelAttachmentDirectory {
) -> Result<(), String> { ) -> Result<(), String> {
let serialized = serde_json::to_string(record) let serialized = serde_json::to_string(record)
.map_err(|err| format!("attachment serialization failed: {err}"))?; .map_err(|err| format!("attachment serialization failed: {err}"))?;
if let Some(runner) = data.kv_runner() { if let Err(error) = self
if let Err(error) = runner .runtime_state
.setex( .kv_set(
&tunnel_attachment_redis_key(node_id), &tunnel_attachment_redis_key(node_id),
&serialized, serialized.clone(),
Some(self.identity.attachment_ttl_secs), Some(Duration::from_secs(self.identity.attachment_ttl_secs)),
) )
.await .await
{ {
warn!( warn!(
error = %error, error = %error,
node_id = %node_id, node_id = %node_id,
"failed to write tunnel attachment to redis; keeping system_config shadow only" "failed to write tunnel attachment to runtime state; keeping system_config shadow only"
); );
} }
}
let value = serde_json::to_value(record) let value = serde_json::to_value(record)
.map_err(|err| format!("attachment serialization failed: {err}"))?; .map_err(|err| format!("attachment serialization failed: {err}"))?;
data.upsert_system_config_value(&tunnel_attachment_key(node_id), &value, None) data.upsert_system_config_value(&tunnel_attachment_key(node_id), &value, None)
@@ -352,15 +350,17 @@ impl TunnelAttachmentDirectory {
data: &GatewayDataState, data: &GatewayDataState,
node_id: &str, node_id: &str,
) -> Result<(), String> { ) -> Result<(), String> {
if let Some(runner) = data.kv_runner() { if let Err(error) = self
if let Err(error) = runner.del(&tunnel_attachment_redis_key(node_id)).await { .runtime_state
.kv_delete(&tunnel_attachment_redis_key(node_id))
.await
{
warn!( warn!(
error = %error, error = %error,
node_id = %node_id, node_id = %node_id,
"failed to delete tunnel attachment from redis; clearing system_config shadow anyway" "failed to delete tunnel attachment from runtime state; clearing system_config shadow anyway"
); );
} }
}
data.delete_system_config_value(&tunnel_attachment_key(node_id)) data.delete_system_config_value(&tunnel_attachment_key(node_id))
.await .await
.map(|_| ()) .map(|_| ())
@@ -396,6 +396,16 @@ impl EmbeddedTunnelState {
Self::with_data_and_directory(data, TunnelAttachmentDirectory::from_environment()) Self::with_data_and_directory(data, TunnelAttachmentDirectory::from_environment())
} }
pub(crate) fn with_data_and_runtime_state(
data: Arc<GatewayDataState>,
runtime_state: Arc<RuntimeState>,
) -> Self {
Self::with_data_and_directory(
data,
TunnelAttachmentDirectory::from_environment().with_runtime_state(runtime_state),
)
}
pub(crate) fn with_data_and_identity( pub(crate) fn with_data_and_identity(
data: Arc<GatewayDataState>, data: Arc<GatewayDataState>,
instance_id: impl Into<String>, instance_id: impl Into<String>,
@@ -408,6 +418,20 @@ impl EmbeddedTunnelState {
) )
} }
pub(crate) fn with_data_identity_and_runtime_state(
data: Arc<GatewayDataState>,
instance_id: impl Into<String>,
relay_base_url: Option<impl Into<String>>,
attachment_ttl_secs: u64,
runtime_state: Arc<RuntimeState>,
) -> Self {
Self::with_data_and_directory(
data,
TunnelAttachmentDirectory::from_parts(instance_id, relay_base_url, attachment_ttl_secs)
.with_runtime_state(runtime_state),
)
}
pub(crate) fn with_data_and_directory( pub(crate) fn with_data_and_directory(
data: Arc<GatewayDataState>, data: Arc<GatewayDataState>,
attachment_directory: TunnelAttachmentDirectory, attachment_directory: TunnelAttachmentDirectory,
+1
View File
@@ -8,6 +8,7 @@ description = "Tunnel proxy for Aether"
aether-contracts.workspace = true aether-contracts.workspace = true
aether-http.workspace = true aether-http.workspace = true
aether-runtime.workspace = true aether-runtime.workspace = true
aether-runtime-state.workspace = true
tokio = { version = "1", features = ["full"] } tokio = { version = "1", features = ["full"] }
reqwest.workspace = true reqwest.workspace = true
hyper = { version = "1", features = ["client", "http1", "http2"] } hyper = { version = "1", features = ["client", "http1", "http2"] }
+12 -8
View File
@@ -6,10 +6,8 @@ use std::sync::{Arc, RwLock};
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use aether_http::{jittered_delay_for_retry, HttpRetryConfig}; use aether_http::{jittered_delay_for_retry, HttpRetryConfig};
use aether_runtime::{ use aether_runtime::{init_reloadable_service_tracing, wait_for_shutdown_signal, ConcurrencyGate};
init_reloadable_service_tracing, wait_for_shutdown_signal, ConcurrencyGate, use aether_runtime_state::{RedisClientConfig, RuntimeSemaphoreConfig, RuntimeState};
DistributedConcurrencyGate, RedisDistributedConcurrencyConfig,
};
use arc_swap::ArcSwap; use arc_swap::ArcSwap;
use tokio::sync::{watch, Mutex}; use tokio::sync::{watch, Mutex};
use tokio::task::JoinHandle; use tokio::task::JoinHandle;
@@ -225,12 +223,18 @@ pub async fn run(mut config: Config, servers: Vec<ServerEntry>) -> anyhow::Resul
.distributed_stream_redis_url .distributed_stream_redis_url
.clone() .clone()
.expect("distributed stream redis url should be validated"); .expect("distributed stream redis url should be validated");
let distributed_gate = DistributedConcurrencyGate::new_redis( let runtime = RuntimeState::redis(
"proxy_streams_distributed", RedisClientConfig {
limit,
RedisDistributedConcurrencyConfig {
url: redis_url, url: redis_url,
key_prefix: state.config.distributed_stream_redis_key_prefix.clone(), key_prefix: state.config.distributed_stream_redis_key_prefix.clone(),
},
Some(state.config.distributed_stream_command_timeout_ms),
)
.await?;
let distributed_gate = runtime.semaphore(
"proxy_streams_distributed",
limit,
RuntimeSemaphoreConfig {
lease_ttl_ms: state.config.distributed_stream_lease_ttl_ms, lease_ttl_ms: state.config.distributed_stream_lease_ttl_ms,
renew_interval_ms: state.config.distributed_stream_renew_interval_ms, renew_interval_ms: state.config.distributed_stream_renew_interval_ms,
command_timeout_ms: Some(state.config.distributed_stream_command_timeout_ms), command_timeout_ms: Some(state.config.distributed_stream_command_timeout_ms),
+8 -13
View File
@@ -4,10 +4,8 @@ use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, RwLock}; use std::sync::{Arc, RwLock};
use std::time::Duration; use std::time::Duration;
use aether_runtime::{ use aether_runtime::{AdmissionPermit, ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot};
AdmissionPermit, ConcurrencyError, ConcurrencyGate, ConcurrencySnapshot, use aether_runtime_state::{RuntimeSemaphore, RuntimeSemaphoreError, RuntimeSemaphoreSnapshot};
DistributedConcurrencyError, DistributedConcurrencyGate, DistributedConcurrencySnapshot,
};
use crate::config::Config; use crate::config::Config;
use crate::registration::client::AetherClient; use crate::registration::client::AetherClient;
@@ -27,7 +25,7 @@ pub struct AppState {
/// Optional per-process stream admission gate. /// Optional per-process stream admission gate.
pub stream_gate: Option<Arc<ConcurrencyGate>>, pub stream_gate: Option<Arc<ConcurrencyGate>>,
/// Optional cross-instance stream admission gate. /// Optional cross-instance stream admission gate.
pub distributed_stream_gate: Option<Arc<DistributedConcurrencyGate>>, pub distributed_stream_gate: Option<Arc<RuntimeSemaphore>>,
} }
/// Per-server state: one instance per Aether server connection. /// Per-server state: one instance per Aether server connection.
@@ -103,10 +101,7 @@ impl AppState {
self self
} }
pub fn with_distributed_stream_concurrency_gate( pub fn with_distributed_stream_concurrency_gate(mut self, gate: Arc<RuntimeSemaphore>) -> Self {
mut self,
gate: Arc<DistributedConcurrencyGate>,
) -> Self {
self.distributed_stream_gate = Some(gate); self.distributed_stream_gate = Some(gate);
self self
} }
@@ -117,7 +112,7 @@ impl AppState {
pub async fn distributed_stream_concurrency_snapshot( pub async fn distributed_stream_concurrency_snapshot(
&self, &self,
) -> Result<Option<DistributedConcurrencySnapshot>, DistributedConcurrencyError> { ) -> Result<Option<RuntimeSemaphoreSnapshot>, RuntimeSemaphoreError> {
match &self.distributed_stream_gate { match &self.distributed_stream_gate {
Some(gate) => gate.snapshot().await.map(Some), Some(gate) => gate.snapshot().await.map(Some),
None => Ok(None), None => Ok(None),
@@ -150,10 +145,10 @@ impl AppState {
let distributed = match &self.distributed_stream_gate { let distributed = match &self.distributed_stream_gate {
Some(gate) => Some(gate.try_acquire().await.map_err(|err| { Some(gate) => Some(gate.try_acquire().await.map_err(|err| {
match err { match err {
DistributedConcurrencyError::Saturated { gate, limit } => { RuntimeSemaphoreError::Saturated { gate, limit } => {
ProxyAdmissionError::Saturated { gate, limit } ProxyAdmissionError::Saturated { gate, limit }
} }
DistributedConcurrencyError::Unavailable { RuntimeSemaphoreError::Unavailable {
gate, gate,
limit, limit,
message, message,
@@ -162,7 +157,7 @@ impl AppState {
limit, limit,
message, message,
}, },
DistributedConcurrencyError::InvalidConfiguration(message) => { RuntimeSemaphoreError::InvalidConfiguration(message) => {
ProxyAdmissionError::Unavailable { ProxyAdmissionError::Unavailable {
gate: "proxy_streams_distributed", gate: "proxy_streams_distributed",
limit: self limit: self
+12 -4
View File
@@ -1515,7 +1515,10 @@ mod tests {
use std::sync::{Mutex, Once}; use std::sync::{Mutex, Once};
use std::task::{Context, Poll}; use std::task::{Context, Poll};
use aether_runtime::{ConcurrencyGate, DistributedConcurrencyGate}; use aether_runtime::ConcurrencyGate;
use aether_runtime_state::{
MemoryRuntimeStateConfig, RuntimeSemaphore, RuntimeSemaphoreConfig, RuntimeState,
};
use arc_swap::ArcSwap; use arc_swap::ArcSwap;
use axum::body::Body; use axum::body::Body;
use axum::http::{header, HeaderMap, Response, StatusCode}; use axum::http::{header, HeaderMap, Response, StatusCode};
@@ -2206,10 +2209,15 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn rejects_stream_when_distributed_admission_gate_is_saturated() { async fn rejects_stream_when_distributed_admission_gate_is_saturated() {
let gate = Arc::new(DistributedConcurrencyGate::new_in_memory( let gate = Arc::new(
RuntimeState::memory(MemoryRuntimeStateConfig::default())
.semaphore(
"proxy_streams_distributed", "proxy_streams_distributed",
1, 1,
)); RuntimeSemaphoreConfig::default(),
)
.expect("distributed semaphore"),
);
let _permit = gate.try_acquire().await.expect("first permit"); let _permit = gate.try_acquire().await.expect("first permit");
let state = sample_state(None, Some(gate)); let state = sample_state(None, Some(gate));
let server = sample_server(&state); let server = sample_server(&state);
@@ -2264,7 +2272,7 @@ mod tests {
fn sample_state( fn sample_state(
stream_gate: Option<Arc<ConcurrencyGate>>, stream_gate: Option<Arc<ConcurrencyGate>>,
distributed_stream_gate: Option<Arc<DistributedConcurrencyGate>>, distributed_stream_gate: Option<Arc<RuntimeSemaphore>>,
) -> Arc<AppState> { ) -> Arc<AppState> {
ensure_rustls_provider(); ensure_rustls_provider();
let config = Arc::new(sample_config()); let config = Arc::new(sample_config());
-1
View File
@@ -17,7 +17,6 @@ chrono.workspace = true
chrono-tz.workspace = true chrono-tz.workspace = true
futures-util.workspace = true futures-util.workspace = true
flate2.workspace = true flate2.workspace = true
redis.workspace = true
serde.workspace = true serde.workspace = true
serde_json.workspace = true serde_json.workspace = true
sha2.workspace = true sha2.workspace = true
+5 -7
View File
@@ -1,7 +1,7 @@
# aether-data # aether-data
`aether-data` is the runtime data-access crate. It owns concrete database and `aether-data` is the runtime data-access crate. It owns concrete SQL/database
Redis clients, concrete repository implementations, migration/backfill/export drivers, concrete repository implementations, migration/backfill/export
workflows, and the composition layer that wires those pieces into the rest of workflows, and the composition layer that wires those pieces into the rest of
the application. the application.
@@ -14,10 +14,9 @@ task crates live in `../aether-data-contracts`.
| Path | Responsibility | | Path | Responsibility |
|---|---| |---|---|
| `src/database.rs` | Logical SQL driver selection and shared pool configuration. | | `src/database.rs` | Logical SQL driver selection and shared pool configuration. |
| `src/config.rs` | Data-layer config that combines SQL and Redis settings. | | `src/config.rs` | Data-layer config for SQL drivers and repository wiring. |
| `src/maintenance.rs` | Maintenance DTOs and aggregation summaries used by backend dispatch and runtime maintenance entrypoints. | | `src/maintenance.rs` | Maintenance DTOs and aggregation summaries used by backend dispatch and runtime maintenance entrypoints. |
| `src/driver/{postgres,mysql,sqlite}` | Low-level SQL driver primitives such as pools, transactions, and leases. These modules should not contain domain repository logic. | | `src/driver/{postgres,mysql,sqlite}` | Low-level SQL driver primitives such as pools, transactions, and leases. These modules should not contain domain repository logic. |
| `src/driver/redis` | Low-level Redis clients, locks, streams, and namespaces. |
| `src/repository` | Domain repository traits/types re-exported from contracts plus concrete in-memory/Postgres/MySQL/SQLite implementations. | | `src/repository` | Domain repository traits/types re-exported from contracts plus concrete in-memory/Postgres/MySQL/SQLite implementations. |
| `src/backend` | Composition root. Builds concrete driver backends and exposes app-facing read/write/worker/lock/lease handles. | | `src/backend` | Composition root. Builds concrete driver backends and exposes app-facing read/write/worker/lock/lease handles. |
| `src/backend/{maintenance,stats,wallet,system}.rs` | Backend-owned maintenance, aggregation, wallet ledger, and system config workflows that are not normal request-path repositories. | | `src/backend/{maintenance,stats,wallet,system}.rs` | Backend-owned maintenance, aggregation, wallet ledger, and system config workflows that are not normal request-path repositories. |
@@ -40,9 +39,8 @@ The crate is easiest to read as five layers:
1. Contracts: DTOs, input structs, repository traits, and `DataLayerError`. 1. Contracts: DTOs, input structs, repository traits, and `DataLayerError`.
Prefer `aether-data-contracts` for anything that another crate needs to Prefer `aether-data-contracts` for anything that another crate needs to
compile against. compile against.
2. Driver primitives: `driver/postgres`, `driver/mysql`, `driver/sqlite`, and 2. Driver primitives: `driver/postgres`, `driver/mysql`, and `driver/sqlite`
`driver/redis` connect to connect to infrastructure and expose pools/runners.
infrastructure and expose pools/runners.
3. Repository implementations: `repository/<domain>/{sql,mysql,sqlite,memory}` 3. Repository implementations: `repository/<domain>/{sql,mysql,sqlite,memory}`
translate contract types to driver-specific SQL. translate contract types to driver-specific SQL.
4. Backend composition: `backend` chooses one SQL driver from config and wires 4. Backend composition: `backend` chooses one SQL driver from config and wires
-58
View File
@@ -1,58 +0,0 @@
use std::fmt;
use super::RedisBackend;
use crate::driver::redis::{RedisLockRunner, RedisLockRunnerConfig};
use crate::DataLayerError;
#[derive(Clone, Default)]
pub struct DataLockBackends {
redis: Option<RedisLockRunner>,
}
impl fmt::Debug for DataLockBackends {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DataLockBackends")
.field("has_redis", &self.redis.is_some())
.finish()
}
}
impl DataLockBackends {
pub(crate) fn from_redis(redis: Option<&RedisBackend>) -> Result<Self, DataLayerError> {
Ok(Self {
redis: redis
.map(|backend| backend.lock_runner(RedisLockRunnerConfig::default()))
.transpose()?,
})
}
pub fn redis(&self) -> Option<RedisLockRunner> {
self.redis.clone()
}
pub fn has_any(&self) -> bool {
self.redis.is_some()
}
}
#[cfg(test)]
mod tests {
use super::DataLockBackends;
use crate::backend::RedisBackend;
use crate::driver::redis::RedisClientConfig;
#[test]
fn builds_redis_lock_runner_from_backend() {
let backend = RedisBackend::from_config(RedisClientConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
})
.expect("redis backend should build");
let locks =
DataLockBackends::from_redis(Some(&backend)).expect("lock backends should build");
assert!(locks.has_any());
assert!(locks.redis().is_some());
}
}
+1 -71
View File
@@ -2,37 +2,31 @@
//! //!
//! `DataBackends` chooses the configured SQL driver, builds low-level pools, //! `DataBackends` chooses the configured SQL driver, builds low-level pools,
//! instantiates concrete repositories, and exposes app-facing read/write, //! instantiates concrete repositories, and exposes app-facing read/write,
//! lease, lock, worker, and maintenance handles. Request-path repository SQL //! lease, transaction, and maintenance handles. Request-path repository SQL
//! belongs in `repository/*`; backend-owned maintenance SQL lives in focused //! belongs in `repository/*`; backend-owned maintenance SQL lives in focused
//! modules such as `stats`, `wallet`, and `system`. Pool/client primitives //! modules such as `stats`, `wallet`, and `system`. Pool/client primitives
//! belong in `driver/*`. //! belong in `driver/*`.
mod leases; mod leases;
mod locks;
mod maintenance; mod maintenance;
mod mysql; mod mysql;
mod postgres; mod postgres;
mod read; mod read;
mod redis;
mod sqlite; mod sqlite;
mod stats; mod stats;
mod stats_common; mod stats_common;
mod system; mod system;
mod transactions; mod transactions;
mod wallet; mod wallet;
mod workers;
mod write; mod write;
use crate::maintenance::DatabasePoolSummary; use crate::maintenance::DatabasePoolSummary;
pub use leases::DataLeaseBackends; pub use leases::DataLeaseBackends;
pub use locks::DataLockBackends;
pub use mysql::MysqlBackend; pub use mysql::MysqlBackend;
pub use postgres::PostgresBackend; pub use postgres::PostgresBackend;
pub use read::DataReadRepositories; pub use read::DataReadRepositories;
pub use redis::RedisBackend;
pub use sqlite::SqliteBackend; pub use sqlite::SqliteBackend;
pub use transactions::DataTransactionBackends; pub use transactions::DataTransactionBackends;
pub use workers::DataWorkerBackends;
pub use write::DataWriteRepositories; pub use write::DataWriteRepositories;
use crate::database::DatabaseDriver; use crate::database::DatabaseDriver;
@@ -51,12 +45,9 @@ pub struct DataBackends {
postgres: Option<PostgresBackend>, postgres: Option<PostgresBackend>,
mysql: Option<MysqlBackend>, mysql: Option<MysqlBackend>,
sqlite: Option<SqliteBackend>, sqlite: Option<SqliteBackend>,
redis: Option<RedisBackend>,
leases: DataLeaseBackends, leases: DataLeaseBackends,
locks: DataLockBackends,
read: DataReadRepositories, read: DataReadRepositories,
transactions: DataTransactionBackends, transactions: DataTransactionBackends,
workers: DataWorkerBackends,
write: DataWriteRepositories, write: DataWriteRepositories,
} }
@@ -111,17 +102,10 @@ impl DataBackends {
} }
_ => None, _ => None,
}; };
let redis = config
.redis
.clone()
.map(RedisBackend::from_config)
.transpose()?;
let leases = DataLeaseBackends::from_postgres(postgres.as_ref())?; let leases = DataLeaseBackends::from_postgres(postgres.as_ref())?;
let locks = DataLockBackends::from_redis(redis.as_ref())?;
let read = let read =
DataReadRepositories::from_backends(postgres.as_ref(), mysql.as_ref(), sqlite.as_ref()); DataReadRepositories::from_backends(postgres.as_ref(), mysql.as_ref(), sqlite.as_ref());
let transactions = DataTransactionBackends::from_postgres(postgres.as_ref()); let transactions = DataTransactionBackends::from_postgres(postgres.as_ref());
let workers = DataWorkerBackends::from_redis(redis.as_ref())?;
let write = DataWriteRepositories::from_backends( let write = DataWriteRepositories::from_backends(
postgres.as_ref(), postgres.as_ref(),
mysql.as_ref(), mysql.as_ref(),
@@ -133,12 +117,9 @@ impl DataBackends {
postgres, postgres,
mysql, mysql,
sqlite, sqlite,
redis,
leases, leases,
locks,
read, read,
transactions, transactions,
workers,
write, write,
}) })
} }
@@ -165,10 +146,6 @@ impl DataBackends {
self.sqlite.as_ref() self.sqlite.as_ref()
} }
pub fn redis(&self) -> Option<&RedisBackend> {
self.redis.as_ref()
}
pub fn read(&self) -> &DataReadRepositories { pub fn read(&self) -> &DataReadRepositories {
&self.read &self.read
} }
@@ -177,18 +154,10 @@ impl DataBackends {
&self.leases &self.leases
} }
pub fn locks(&self) -> &DataLockBackends {
&self.locks
}
pub fn transactions(&self) -> &DataTransactionBackends { pub fn transactions(&self) -> &DataTransactionBackends {
&self.transactions &self.transactions
} }
pub fn workers(&self) -> &DataWorkerBackends {
&self.workers
}
pub fn write(&self) -> &DataWriteRepositories { pub fn write(&self) -> &DataWriteRepositories {
&self.write &self.write
} }
@@ -197,12 +166,9 @@ impl DataBackends {
self.postgres.is_some() self.postgres.is_some()
|| self.mysql.is_some() || self.mysql.is_some()
|| self.sqlite.is_some() || self.sqlite.is_some()
|| self.redis.is_some()
|| self.leases.has_any() || self.leases.has_any()
|| self.locks.has_any()
|| self.read.has_any() || self.read.has_any()
|| self.transactions.has_any() || self.transactions.has_any()
|| self.workers.has_any()
|| self.write.has_any() || self.write.has_any()
} }
} }
@@ -224,9 +190,7 @@ mod tests {
assert!(backends.postgres().is_none()); assert!(backends.postgres().is_none());
assert!(backends.mysql().is_none()); assert!(backends.mysql().is_none());
assert!(backends.sqlite().is_none()); assert!(backends.sqlite().is_none());
assert!(backends.redis().is_none());
assert!(backends.leases().postgres().is_none()); assert!(backends.leases().postgres().is_none());
assert!(backends.locks().redis().is_none());
assert!(backends.read().auth_api_keys().is_none()); assert!(backends.read().auth_api_keys().is_none());
assert!(backends.read().auth_modules().is_none()); assert!(backends.read().auth_modules().is_none());
assert!(backends.read().billing().is_none()); assert!(backends.read().billing().is_none());
@@ -241,7 +205,6 @@ mod tests {
assert!(backends.read().usage().is_none()); assert!(backends.read().usage().is_none());
assert!(backends.read().video_tasks().is_none()); assert!(backends.read().video_tasks().is_none());
assert!(backends.transactions().postgres().is_none()); assert!(backends.transactions().postgres().is_none());
assert!(backends.workers().redis().is_none());
assert!(backends.write().settlement().is_none()); assert!(backends.write().settlement().is_none());
assert!(backends.write().usage().is_none()); assert!(backends.write().usage().is_none());
} }
@@ -260,7 +223,6 @@ mod tests {
statement_cache_capacity: 64, statement_cache_capacity: 64,
require_ssl: false, require_ssl: false,
}), }),
redis: None,
}) })
.expect("postgres backend should build"); .expect("postgres backend should build");
@@ -308,7 +270,6 @@ mod tests {
pool: SqlPoolConfig::default(), pool: SqlPoolConfig::default(),
}), }),
postgres: None, postgres: None,
redis: None,
}) })
.expect("mysql backend should build"); .expect("mysql backend should build");
@@ -360,7 +321,6 @@ mod tests {
pool: SqlPoolConfig::default(), pool: SqlPoolConfig::default(),
}), }),
postgres: None, postgres: None,
redis: None,
}) })
.expect("sqlite backend should build"); .expect("sqlite backend should build");
@@ -401,34 +361,4 @@ mod tests {
assert!(backends.write().wallets().is_some()); assert!(backends.write().wallets().is_some());
assert!(backends.config().effective_database().is_some()); assert!(backends.config().effective_database().is_some());
} }
#[test]
fn builds_redis_backend_from_config() {
let backends = DataBackends::from_config(DataLayerConfig {
database: None,
postgres: None,
redis: Some(crate::driver::redis::RedisClientConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
}),
})
.expect("redis backend should build");
assert!(backends.has_runtime_backends());
assert!(backends.postgres().is_none());
assert!(backends.mysql().is_none());
assert!(backends.sqlite().is_none());
assert!(backends.redis().is_some());
assert!(backends.leases().postgres().is_none());
assert!(backends.locks().redis().is_some());
assert!(backends.workers().redis().is_some());
assert!(backends.read().auth_api_keys().is_none());
assert!(backends.read().auth_modules().is_none());
assert!(backends.read().global_models().is_none());
assert!(backends.read().oauth_providers().is_none());
assert!(backends.transactions().postgres().is_none());
assert!(backends.write().settlement().is_none());
assert!(backends.write().usage().is_none());
assert!(backends.config().redis.is_some());
}
} }
-86
View File
@@ -1,86 +0,0 @@
use crate::driver::redis::{
RedisClient, RedisClientConfig, RedisClientFactory, RedisKeyspace, RedisKvRunner,
RedisKvRunnerConfig, RedisLockRunner, RedisLockRunnerConfig, RedisStreamRunner,
RedisStreamRunnerConfig,
};
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct RedisBackend {
config: RedisClientConfig,
client: RedisClient,
}
impl RedisBackend {
pub fn from_config(config: RedisClientConfig) -> Result<Self, DataLayerError> {
let factory = RedisClientFactory::new(config.clone())?;
let client = factory.connect_lazy()?;
Ok(Self { config, client })
}
pub fn config(&self) -> &RedisClientConfig {
&self.config
}
pub fn client(&self) -> &RedisClient {
&self.client
}
pub fn client_clone(&self) -> RedisClient {
self.client.clone()
}
pub fn keyspace(&self) -> RedisKeyspace {
self.config.keyspace()
}
pub fn lock_runner(
&self,
config: RedisLockRunnerConfig,
) -> Result<RedisLockRunner, DataLayerError> {
RedisLockRunner::new(self.client_clone(), self.keyspace(), config)
}
pub fn stream_runner(
&self,
config: RedisStreamRunnerConfig,
) -> Result<RedisStreamRunner, DataLayerError> {
RedisStreamRunner::new(self.client_clone(), self.keyspace(), config)
}
pub fn kv_runner(&self, config: RedisKvRunnerConfig) -> Result<RedisKvRunner, DataLayerError> {
RedisKvRunner::new(self.client_clone(), self.keyspace(), config)
}
}
#[cfg(test)]
mod tests {
use super::RedisBackend;
use crate::driver::redis::{
RedisClientConfig, RedisKvRunnerConfig, RedisLockRunnerConfig, RedisStreamRunnerConfig,
};
#[test]
fn backend_retains_config_client_and_shared_runners() {
let config = RedisClientConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
};
let backend = RedisBackend::from_config(config.clone()).expect("backend should build");
assert_eq!(backend.config(), &config);
assert_eq!(backend.keyspace().key("audit"), "aether:audit");
let _client_ref = backend.client();
let _client_clone = backend.client_clone();
let _lock_runner = backend
.lock_runner(RedisLockRunnerConfig::default())
.expect("lock runner should build");
let _stream_runner = backend
.stream_runner(RedisStreamRunnerConfig::default())
.expect("stream runner should build");
let _kv_runner = backend
.kv_runner(RedisKvRunnerConfig::default())
.expect("kv runner should build");
}
}
-58
View File
@@ -1,58 +0,0 @@
use std::fmt;
use super::RedisBackend;
use crate::driver::redis::{RedisStreamRunner, RedisStreamRunnerConfig};
use crate::DataLayerError;
#[derive(Clone, Default)]
pub struct DataWorkerBackends {
redis: Option<RedisStreamRunner>,
}
impl fmt::Debug for DataWorkerBackends {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DataWorkerBackends")
.field("has_redis", &self.redis.is_some())
.finish()
}
}
impl DataWorkerBackends {
pub(crate) fn from_redis(redis: Option<&RedisBackend>) -> Result<Self, DataLayerError> {
Ok(Self {
redis: redis
.map(|backend| backend.stream_runner(RedisStreamRunnerConfig::default()))
.transpose()?,
})
}
pub fn redis(&self) -> Option<RedisStreamRunner> {
self.redis.clone()
}
pub fn has_any(&self) -> bool {
self.redis.is_some()
}
}
#[cfg(test)]
mod tests {
use super::DataWorkerBackends;
use crate::backend::RedisBackend;
use crate::driver::redis::RedisClientConfig;
#[test]
fn builds_redis_stream_runner_from_backend() {
let backend = RedisBackend::from_config(RedisClientConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
})
.expect("redis backend should build");
let workers =
DataWorkerBackends::from_redis(Some(&backend)).expect("worker backends should build");
assert!(workers.has_any());
assert!(workers.redis().is_some());
}
}
+1 -15
View File
@@ -1,13 +1,11 @@
use crate::database::SqlDatabaseConfig; use crate::database::SqlDatabaseConfig;
use crate::driver::postgres::PostgresPoolConfig; use crate::driver::postgres::PostgresPoolConfig;
use crate::driver::redis::RedisClientConfig;
use crate::DataLayerError; use crate::DataLayerError;
#[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize, PartialEq, Eq)] #[derive(Debug, Clone, Default, serde::Serialize, serde::Deserialize, PartialEq, Eq)]
pub struct DataLayerConfig { pub struct DataLayerConfig {
pub database: Option<SqlDatabaseConfig>, pub database: Option<SqlDatabaseConfig>,
pub postgres: Option<PostgresPoolConfig>, pub postgres: Option<PostgresPoolConfig>,
pub redis: Option<RedisClientConfig>,
} }
impl DataLayerConfig { impl DataLayerConfig {
@@ -15,7 +13,6 @@ impl DataLayerConfig {
Self { Self {
database: Some(database), database: Some(database),
postgres: None, postgres: None,
redis: None,
} }
} }
@@ -23,7 +20,6 @@ impl DataLayerConfig {
Self { Self {
database: Some(SqlDatabaseConfig::from_postgres_config(postgres)), database: Some(SqlDatabaseConfig::from_postgres_config(postgres)),
postgres: None, postgres: None,
redis: None,
} }
} }
@@ -42,14 +38,11 @@ impl DataLayerConfig {
if let Some(postgres) = &self.postgres { if let Some(postgres) = &self.postgres {
postgres.validate()?; postgres.validate()?;
} }
if let Some(redis) = &self.redis {
redis.validate()?;
}
Ok(()) Ok(())
} }
pub fn has_persistent_backends(&self) -> bool { pub fn has_persistent_backends(&self) -> bool {
self.effective_database().is_some() || self.redis.is_some() self.effective_database().is_some()
} }
} }
@@ -58,7 +51,6 @@ mod tests {
use super::DataLayerConfig; use super::DataLayerConfig;
use crate::database::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig}; use crate::database::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
use crate::driver::postgres::PostgresPoolConfig; use crate::driver::postgres::PostgresPoolConfig;
use crate::driver::redis::RedisClientConfig;
#[test] #[test]
fn validates_nested_backend_configs() { fn validates_nested_backend_configs() {
@@ -74,10 +66,6 @@ mod tests {
statement_cache_capacity: 64, statement_cache_capacity: 64,
require_ssl: false, require_ssl: false,
}), }),
redis: Some(RedisClientConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
}),
}; };
assert!(config.validate().is_ok()); assert!(config.validate().is_ok());
@@ -98,7 +86,6 @@ mod tests {
statement_cache_capacity: 64, statement_cache_capacity: 64,
require_ssl: false, require_ssl: false,
}), }),
redis: None,
}; };
assert!(config.validate().is_err()); assert!(config.validate().is_err());
@@ -122,7 +109,6 @@ mod tests {
statement_cache_capacity: 64, statement_cache_capacity: 64,
require_ssl: false, require_ssl: false,
}), }),
redis: None,
}; };
let effective = config let effective = config
+1 -2
View File
@@ -1,10 +1,9 @@
//! Low-level data driver primitives. //! Low-level data driver primitives.
//! //!
//! These modules own pools, transactions, leases, and Redis client helpers. //! These modules own pools, transactions, and lease helpers.
//! Domain repository logic belongs in `repository/*`, and application-facing //! Domain repository logic belongs in `repository/*`, and application-facing
//! composition belongs in `backend`. //! composition belongs in `backend`.
pub mod mysql; pub mod mysql;
pub mod postgres; pub mod postgres;
pub mod redis;
pub mod sqlite; pub mod sqlite;
-14
View File
@@ -4,10 +4,6 @@ pub(crate) fn postgres_error(error: impl std::fmt::Display) -> DataLayerError {
DataLayerError::postgres(error) DataLayerError::postgres(error)
} }
pub(crate) fn redis_error(error: impl std::fmt::Display) -> DataLayerError {
DataLayerError::redis(error)
}
pub(crate) fn sql_error(error: impl std::fmt::Display) -> DataLayerError { pub(crate) fn sql_error(error: impl std::fmt::Display) -> DataLayerError {
DataLayerError::sql(error) DataLayerError::sql(error)
} }
@@ -22,16 +18,6 @@ impl<T> SqlxResultExt<T> for Result<T, sqlx::Error> {
} }
} }
pub(crate) trait RedisResultExt<T> {
fn map_redis_err(self) -> Result<T, DataLayerError>;
}
impl<T> RedisResultExt<T> for Result<T, redis::RedisError> {
fn map_redis_err(self) -> Result<T, DataLayerError> {
self.map_err(redis_error)
}
}
pub(crate) trait SqlResultExt<T> { pub(crate) trait SqlResultExt<T> {
fn map_sql_err(self) -> Result<T, DataLayerError>; fn map_sql_err(self) -> Result<T, DataLayerError>;
} }
+3 -4
View File
@@ -1,6 +1,6 @@
//! Runtime data access for Aether. //! Runtime data access for Aether.
//! //!
//! This crate contains concrete database/Redis clients, repository //! This crate contains concrete database clients, repository
//! implementations, migration/backfill/export workflows, and the backend //! implementations, migration/backfill/export workflows, and the backend
//! composition layer. Shared repository contracts that other crates compile //! composition layer. Shared repository contracts that other crates compile
//! against live in `aether-data-contracts`. //! against live in `aether-data-contracts`.
@@ -17,9 +17,8 @@ pub mod maintenance;
pub mod repository; pub mod repository;
pub use backend::{ pub use backend::{
DataBackends, DataLeaseBackends, DataLockBackends, DataReadRepositories, DataBackends, DataLeaseBackends, DataReadRepositories, DataTransactionBackends,
DataTransactionBackends, DataWorkerBackends, DataWriteRepositories, PostgresBackend, DataWriteRepositories, PostgresBackend,
RedisBackend,
}; };
pub use config::DataLayerConfig; pub use config::DataLayerConfig;
pub use database::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, DEFAULT_SQLITE_DATABASE_URL}; pub use database::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, DEFAULT_SQLITE_DATABASE_URL};
@@ -211,7 +211,7 @@ fn relative_path(path: &Path, workspace_root: &Path) -> PathBuf {
fn grouped_import_scanner_allows_nested_new_paths() { fn grouped_import_scanner_allows_nested_new_paths() {
let source = r#" let source = r#"
use aether_data::{ use aether_data::{
driver::{postgres::PostgresPool, redis::RedisKvRunner}, driver::{postgres::PostgresPool, mysql::MySqlPool},
lifecycle::{backfill::PendingBackfillInfo, migrate::PendingMigrationInfo}, lifecycle::{backfill::PendingBackfillInfo, migrate::PendingMigrationInfo},
}; };
"#; "#;
+1 -1
View File
@@ -10,9 +10,9 @@ description = "Provider transport core extracted from aether-gateway"
aether-ai-formats.workspace = true aether-ai-formats.workspace = true
aether-contracts.workspace = true aether-contracts.workspace = true
aether-crypto.workspace = true aether-crypto.workspace = true
aether-data.workspace = true
aether-data-contracts.workspace = true aether-data-contracts.workspace = true
aether-oauth.workspace = true aether-oauth.workspace = true
aether-runtime-state.workspace = true
aether-video-tasks-core.workspace = true aether-video-tasks-core.workspace = true
async-trait.workspace = true async-trait.workspace = true
http.workspace = true http.workspace = true
@@ -2,12 +2,12 @@ use std::collections::BTreeMap;
use std::fmt; use std::fmt;
use std::sync::Arc; use std::sync::Arc;
use aether_data::driver::redis::{RedisLockKey, RedisLockRunner};
use aether_oauth::core::OAuthError; use aether_oauth::core::OAuthError;
use aether_oauth::network::{ use aether_oauth::network::{
OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse, OAuthNetworkContext, OAuthHttpExecutor, OAuthHttpRequest, OAuthHttpResponse, OAuthNetworkContext,
}; };
use aether_oauth::provider::ProviderOAuthTransportContext; use aether_oauth::provider::ProviderOAuthTransportContext;
use aether_runtime_state::RuntimeState;
use async_trait::async_trait; use async_trait::async_trait;
use serde_json::Value; use serde_json::Value;
use thiserror::Error; use thiserror::Error;
@@ -374,7 +374,7 @@ impl LocalOAuthRefreshCoordinator {
&self, &self,
executor: &dyn LocalOAuthHttpExecutor, executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot, transport: &GatewayProviderTransportSnapshot,
distributed_lock: Option<&RedisLockRunner>, distributed_lock: Option<&RuntimeState>,
distributed_owner: Option<&str>, distributed_owner: Option<&str>,
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> { ) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
self.resolve_with_result_mode( self.resolve_with_result_mode(
@@ -391,7 +391,7 @@ impl LocalOAuthRefreshCoordinator {
&self, &self,
executor: &dyn LocalOAuthHttpExecutor, executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot, transport: &GatewayProviderTransportSnapshot,
distributed_lock: Option<&RedisLockRunner>, distributed_lock: Option<&RuntimeState>,
distributed_owner: Option<&str>, distributed_owner: Option<&str>,
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> { ) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
self.resolve_with_result_mode( self.resolve_with_result_mode(
@@ -408,7 +408,7 @@ impl LocalOAuthRefreshCoordinator {
&self, &self,
executor: &dyn LocalOAuthHttpExecutor, executor: &dyn LocalOAuthHttpExecutor,
transport: &GatewayProviderTransportSnapshot, transport: &GatewayProviderTransportSnapshot,
distributed_lock: Option<&RedisLockRunner>, distributed_lock: Option<&RuntimeState>,
distributed_owner: Option<&str>, distributed_owner: Option<&str>,
force_refresh: bool, force_refresh: bool,
) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> { ) -> Result<Option<LocalOAuthResolution>, LocalOAuthRefreshError> {
@@ -465,12 +465,11 @@ impl LocalOAuthRefreshCoordinator {
let distributed_lease = match (distributed_lock, distributed_owner) { let distributed_lease = match (distributed_lock, distributed_owner) {
(Some(lock), Some(owner)) if !owner.trim().is_empty() => { (Some(lock), Some(owner)) if !owner.trim().is_empty() => {
let lock_key = RedisLockKey(format!("provider_oauth_refresh_lock:{key_id}"));
match lock match lock
.try_acquire( .lock_try_acquire(
&lock_key, &format!("provider_oauth_refresh_lock:{key_id}"),
owner, owner,
Some(Self::DISTRIBUTED_REFRESH_LOCK_TTL_MS), std::time::Duration::from_millis(Self::DISTRIBUTED_REFRESH_LOCK_TTL_MS),
) )
.await .await
{ {
@@ -497,7 +496,7 @@ impl LocalOAuthRefreshCoordinator {
let refresh_entry = cached_entry.as_ref(); let refresh_entry = cached_entry.as_ref();
let refresh_result = adapter.refresh(executor, transport, refresh_entry).await; let refresh_result = adapter.refresh(executor, transport, refresh_entry).await;
if let (Some(lock), Some(lease)) = (distributed_lock, distributed_lease.as_ref()) { if let (Some(lock), Some(lease)) = (distributed_lock, distributed_lease.as_ref()) {
if let Err(err) = lock.release(lease).await { if let Err(err) = lock.lock_release(lease).await {
tracing::warn!( tracing::warn!(
key_id = %key_id, key_id = %key_id,
provider_type = adapter.provider_type(), provider_type = adapter.provider_type(),
+21
View File
@@ -0,0 +1,21 @@
[package]
name = "aether-runtime-state"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
description = "Runtime state backends for Aether services"
[dependencies]
aether-cache.workspace = true
aether-data-contracts.workspace = true
aether-runtime.workspace = true
async-trait.workspace = true
redis.workspace = true
serde.workspace = true
serde_json.workspace = true
thiserror.workspace = true
tokio.workspace = true
tracing.workspace = true
url.workspace = true
uuid.workspace = true
+15
View File
@@ -0,0 +1,15 @@
pub use aether_data_contracts::DataLayerError;
pub(crate) fn redis_error(error: impl std::fmt::Display) -> DataLayerError {
DataLayerError::redis(error)
}
pub(crate) trait RedisResultExt<T> {
fn map_redis_err(self) -> Result<T, DataLayerError>;
}
impl<T> RedisResultExt<T> for Result<T, redis::RedisError> {
fn map_redis_err(self) -> Result<T, DataLayerError> {
self.map_err(redis_error)
}
}
File diff suppressed because it is too large Load Diff
+537
View File
@@ -0,0 +1,537 @@
use std::collections::{BTreeMap, BTreeSet, HashMap, VecDeque};
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{Duration, Instant};
use tokio::sync::Mutex;
use crate::{RuntimeQueueEntry, RuntimeQueueReclaimConfig};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct MemoryRuntimeStateConfig {
pub max_kv_entries: usize,
}
impl Default for MemoryRuntimeStateConfig {
fn default() -> Self {
Self {
max_kv_entries: 10_000,
}
}
}
#[derive(Debug, Clone)]
pub(crate) struct MemoryKvEntry {
pub(crate) value: String,
pub(crate) inserted_at: Instant,
pub(crate) expires_at: Option<Instant>,
}
impl MemoryKvEntry {
fn is_expired(&self, now: Instant) -> bool {
self.expires_at.is_some_and(|expires_at| now >= expires_at)
}
}
#[derive(Debug, Default)]
pub(crate) struct MemoryRuntimeBackend {
config: MemoryRuntimeStateConfig,
kv: Mutex<HashMap<String, MemoryKvEntry>>,
counters: Mutex<HashMap<String, MemoryCounterEntry>>,
sets: Mutex<HashMap<String, BTreeSet<String>>>,
scores: Mutex<HashMap<String, BTreeMap<String, f64>>>,
queues: Mutex<HashMap<String, VecDeque<RuntimeQueueEntry>>>,
queue_seq: AtomicU64,
locks: Mutex<HashMap<String, MemoryLockEntry>>,
semaphores: Mutex<HashMap<String, BTreeMap<String, u64>>>,
}
#[derive(Debug, Clone)]
struct MemoryCounterEntry {
value: u32,
bucket: u64,
expires_at: Instant,
}
#[derive(Debug, Clone)]
pub(crate) struct MemoryLockEntry {
pub(crate) token: String,
#[allow(dead_code)]
pub(crate) owner: String,
pub(crate) expires_at: Instant,
}
impl MemoryRuntimeBackend {
pub(crate) fn new(config: MemoryRuntimeStateConfig) -> Self {
Self {
config,
..Self::default()
}
}
pub(crate) async fn kv_set(&self, key: &str, value: String, ttl: Option<Duration>) {
let mut kv = self.kv.lock().await;
let now = Instant::now();
if ttl.is_some_and(|ttl| ttl.is_zero()) {
kv.remove(key);
return;
}
prune_kv(&mut kv, now);
while kv.len() >= self.config.max_kv_entries.max(1) {
let Some(oldest_key) = kv
.iter()
.min_by_key(|(_, entry)| entry.inserted_at)
.map(|(key, _)| key.clone())
else {
break;
};
kv.remove(&oldest_key);
}
kv.insert(
key.to_string(),
MemoryKvEntry {
value,
inserted_at: now,
expires_at: ttl.map(|ttl| now + ttl),
},
);
}
pub(crate) fn kv_set_nowait(&self, key: &str, value: String, ttl: Option<Duration>) -> bool {
let Ok(mut kv) = self.kv.try_lock() else {
return false;
};
let now = Instant::now();
if ttl.is_some_and(|ttl| ttl.is_zero()) {
kv.remove(key);
return true;
}
prune_kv(&mut kv, now);
while kv.len() >= self.config.max_kv_entries.max(1) {
let Some(oldest_key) = kv
.iter()
.min_by_key(|(_, entry)| entry.inserted_at)
.map(|(key, _)| key.clone())
else {
break;
};
kv.remove(&oldest_key);
}
kv.insert(
key.to_string(),
MemoryKvEntry {
value,
inserted_at: now,
expires_at: ttl.map(|ttl| now + ttl),
},
);
true
}
pub(crate) async fn kv_get(&self, key: &str) -> Option<String> {
let mut kv = self.kv.lock().await;
get_fresh_locked(&mut kv, key, Instant::now())
}
pub(crate) async fn kv_take(&self, key: &str) -> Option<String> {
let mut kv = self.kv.lock().await;
let now = Instant::now();
let entry = kv.remove(key)?;
if entry.is_expired(now) {
return None;
}
Some(entry.value)
}
pub(crate) async fn kv_delete(&self, key: &str) -> bool {
self.kv.lock().await.remove(key).is_some()
}
pub(crate) async fn kv_delete_many(&self, keys: &[String]) -> usize {
let mut kv = self.kv.lock().await;
keys.iter().filter(|key| kv.remove(*key).is_some()).count()
}
pub(crate) async fn kv_exists(&self, key: &str) -> bool {
self.kv_get(key).await.is_some()
}
pub(crate) async fn kv_ttl_seconds(&self, key: &str) -> Option<i64> {
let mut kv = self.kv.lock().await;
let now = Instant::now();
let entry = kv.get(key).cloned()?;
if entry.is_expired(now) {
kv.remove(key);
return None;
}
Some(
entry
.expires_at
.map(|expires_at| {
expires_at
.saturating_duration_since(now)
.as_secs()
.try_into()
.unwrap_or(i64::MAX)
})
.unwrap_or(-1),
)
}
pub(crate) async fn kv_scan(&self, pattern: &str) -> Vec<String> {
let mut kv = self.kv.lock().await;
prune_kv(&mut kv, Instant::now());
let mut keys = kv
.keys()
.filter(|key| key_matches_pattern(key, pattern))
.cloned()
.collect::<Vec<_>>();
keys.sort();
keys
}
pub(crate) async fn check_and_consume_rate_limit(
&self,
user_key: &str,
key_key: &str,
bucket: u64,
user_limit: u32,
key_limit: u32,
ttl: Duration,
) -> Result<crate::RateLimitCheck, crate::DataLayerError> {
let mut counters = self.counters.lock().await;
let now = Instant::now();
counters.retain(|_, entry| entry.expires_at > now && entry.bucket >= bucket);
if user_limit > 0 {
let user_count = counters
.get(user_key)
.filter(|entry| entry.bucket == bucket)
.map(|entry| entry.value)
.unwrap_or_default();
if user_count >= user_limit {
return Ok(crate::RateLimitCheck::Rejected {
scope: crate::RateLimitScope::User,
limit: user_limit,
});
}
}
if key_limit > 0 {
let key_count = counters
.get(key_key)
.filter(|entry| entry.bucket == bucket)
.map(|entry| entry.value)
.unwrap_or_default();
if key_count >= key_limit {
return Ok(crate::RateLimitCheck::Rejected {
scope: crate::RateLimitScope::Key,
limit: key_limit,
});
}
}
let mut remaining = None::<u32>;
let expires_at = now + ttl;
if user_limit > 0 {
let next = counters
.entry(user_key.to_string())
.and_modify(|entry| {
entry.bucket = bucket;
entry.value = entry.value.saturating_add(1);
entry.expires_at = expires_at;
})
.or_insert(MemoryCounterEntry {
value: 1,
bucket,
expires_at,
})
.value;
remaining = Some(user_limit.saturating_sub(next));
}
if key_limit > 0 {
let next = counters
.entry(key_key.to_string())
.and_modify(|entry| {
entry.bucket = bucket;
entry.value = entry.value.saturating_add(1);
entry.expires_at = expires_at;
})
.or_insert(MemoryCounterEntry {
value: 1,
bucket,
expires_at,
})
.value;
let key_remaining = key_limit.saturating_sub(next);
remaining = Some(remaining.map_or(key_remaining, |value| value.min(key_remaining)));
}
Ok(crate::RateLimitCheck::Allowed {
remaining: remaining.unwrap_or(0),
})
}
pub(crate) async fn set_add(&self, key: &str, member: &str) -> bool {
self.sets
.lock()
.await
.entry(key.to_string())
.or_default()
.insert(member.to_string())
}
pub(crate) fn set_add_nowait(&self, key: &str, member: &str) -> bool {
let Ok(mut sets) = self.sets.try_lock() else {
return false;
};
sets.entry(key.to_string())
.or_default()
.insert(member.to_string())
}
pub(crate) async fn set_remove(&self, key: &str, member: &str) -> bool {
self.sets
.lock()
.await
.get_mut(key)
.is_some_and(|set| set.remove(member))
}
pub(crate) async fn set_members(&self, key: &str) -> Vec<String> {
self.sets
.lock()
.await
.get(key)
.map(|set| set.iter().cloned().collect())
.unwrap_or_default()
}
pub(crate) async fn set_len(&self, key: &str) -> usize {
self.sets.lock().await.get(key).map_or(0, BTreeSet::len)
}
pub(crate) async fn score_set(&self, key: &str, member: &str, score: f64) {
self.scores
.lock()
.await
.entry(key.to_string())
.or_default()
.insert(member.to_string(), score);
}
pub(crate) async fn score_many(&self, key: &str, members: &[String]) -> Vec<Option<f64>> {
let scores = self.scores.lock().await;
members
.iter()
.map(|member| scores.get(key).and_then(|set| set.get(member)).copied())
.collect()
}
pub(crate) async fn score_range_by_min(&self, key: &str, min_score: f64) -> Vec<String> {
let scores = self.scores.lock().await;
scores
.get(key)
.map(|set| {
set.iter()
.filter(|(_, score)| **score >= min_score)
.map(|(member, _)| member.clone())
.collect()
})
.unwrap_or_default()
}
pub(crate) async fn score_remove_by_score(&self, key: &str, max_score: f64) -> usize {
let mut scores = self.scores.lock().await;
let Some(set) = scores.get_mut(key) else {
return 0;
};
let before = set.len();
set.retain(|_, score| *score > max_score);
before.saturating_sub(set.len())
}
pub(crate) async fn score_len(&self, key: &str) -> usize {
self.scores.lock().await.get(key).map_or(0, BTreeMap::len)
}
pub(crate) async fn queue_append(
&self,
stream: &str,
fields: BTreeMap<String, String>,
maxlen: Option<usize>,
) -> String {
let id = format!(
"{}-0",
self.queue_seq
.fetch_add(1, Ordering::Relaxed)
.saturating_add(1)
);
let mut queues = self.queues.lock().await;
let queue = queues.entry(stream.to_string()).or_default();
queue.push_back(RuntimeQueueEntry {
id: id.clone(),
fields,
});
if let Some(maxlen) = maxlen.filter(|value| *value > 0) {
while queue.len() > maxlen {
queue.pop_front();
}
}
id
}
pub(crate) async fn queue_read(&self, stream: &str, count: usize) -> Vec<RuntimeQueueEntry> {
let mut queues = self.queues.lock().await;
let Some(queue) = queues.get_mut(stream) else {
return Vec::new();
};
let mut entries = Vec::new();
for _ in 0..count.max(1) {
let Some(entry) = queue.pop_front() else {
break;
};
entries.push(entry);
}
entries
}
pub(crate) async fn queue_claim_stale(
&self,
_stream: &str,
_config: RuntimeQueueReclaimConfig,
) -> Vec<RuntimeQueueEntry> {
Vec::new()
}
pub(crate) async fn queue_delete(&self, _stream: &str, _ids: &[String]) -> usize {
0
}
pub(crate) async fn lock_try_acquire(
&self,
key: &str,
owner: &str,
token: String,
ttl: Duration,
) -> bool {
let mut locks = self.locks.lock().await;
let now = Instant::now();
locks.retain(|_, entry| entry.expires_at > now);
if locks.contains_key(key) {
return false;
}
locks.insert(
key.to_string(),
MemoryLockEntry {
token,
owner: owner.to_string(),
expires_at: now + ttl,
},
);
true
}
pub(crate) async fn lock_release(&self, key: &str, token: &str) -> bool {
let mut locks = self.locks.lock().await;
if locks.get(key).is_some_and(|entry| entry.token == token) {
locks.remove(key);
return true;
}
false
}
pub(crate) async fn lock_renew(&self, key: &str, token: &str, ttl: Duration) -> bool {
let mut locks = self.locks.lock().await;
if let Some(entry) = locks.get_mut(key) {
if entry.token == token {
entry.expires_at = Instant::now() + ttl;
return true;
}
}
false
}
pub(crate) async fn semaphore_try_acquire(
&self,
key: &str,
token: String,
limit: usize,
ttl_ms: u64,
) -> Result<usize, usize> {
let now_ms = unix_time_ms();
let expires_at = now_ms.saturating_add(ttl_ms);
let mut semaphores = self.semaphores.lock().await;
let holders = semaphores.entry(key.to_string()).or_default();
holders.retain(|_, expires| *expires > now_ms);
let count = holders.len();
if count >= limit {
return Err(count);
}
holders.insert(token, expires_at);
Ok(holders.len())
}
pub(crate) async fn semaphore_renew(&self, key: &str, token: &str, ttl_ms: u64) -> bool {
let now_ms = unix_time_ms();
let mut semaphores = self.semaphores.lock().await;
let Some(holders) = semaphores.get_mut(key) else {
return false;
};
holders.retain(|_, expires| *expires > now_ms);
if let Some(expires) = holders.get_mut(token) {
*expires = now_ms.saturating_add(ttl_ms);
return true;
}
false
}
pub(crate) async fn semaphore_release(&self, key: &str, token: &str) {
let mut semaphores = self.semaphores.lock().await;
if let Some(holders) = semaphores.get_mut(key) {
holders.remove(token);
if holders.is_empty() {
semaphores.remove(key);
}
}
}
pub(crate) async fn semaphore_live_count(&self, key: &str) -> usize {
let now_ms = unix_time_ms();
let mut semaphores = self.semaphores.lock().await;
let Some(holders) = semaphores.get_mut(key) else {
return 0;
};
holders.retain(|_, expires| *expires > now_ms);
holders.len()
}
}
fn get_fresh_locked(
kv: &mut HashMap<String, MemoryKvEntry>,
key: &str,
now: Instant,
) -> Option<String> {
let entry = kv.get(key).cloned()?;
if entry.is_expired(now) {
kv.remove(key);
return None;
}
Some(entry.value)
}
fn prune_kv(kv: &mut HashMap<String, MemoryKvEntry>, now: Instant) {
kv.retain(|_, entry| !entry.is_expired(now));
}
pub(crate) fn key_matches_pattern(key: &str, pattern: &str) -> bool {
match pattern.strip_suffix('*') {
Some(prefix) => key.starts_with(prefix),
None => key == pattern,
}
}
fn unix_time_ms() -> u64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64
}
@@ -1,5 +1,5 @@
use crate::driver::redis::RedisKeyspace;
use crate::error::RedisResultExt; use crate::error::RedisResultExt;
use crate::redis::RedisKeyspace;
use crate::DataLayerError; use crate::DataLayerError;
pub type RedisClient = redis::Client; pub type RedisClient = redis::Client;
@@ -1,8 +1,8 @@
use std::future::Future; use std::future::Future;
use std::time::Duration; use std::time::Duration;
use crate::driver::redis::{RedisClient, RedisKeyspace};
use crate::error::RedisResultExt; use crate::error::RedisResultExt;
use crate::redis::{RedisClient, RedisKeyspace};
use crate::DataLayerError; use crate::DataLayerError;
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
@@ -96,6 +96,59 @@ impl RedisKvRunner {
.await .await
} }
pub async fn get(&self, key: &str) -> Result<Option<String>, DataLayerError> {
let namespaced_key = self.keyspace.key(key);
self.run_with_timeout("redis kv get", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
redis::cmd("GET")
.arg(&namespaced_key)
.query_async(&mut connection)
.await
.map_redis_err()
})
.await
}
pub async fn getdel(&self, key: &str) -> Result<Option<String>, DataLayerError> {
let namespaced_key = self.keyspace.key(key);
self.run_with_timeout("redis kv getdel", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
redis::cmd("GETDEL")
.arg(&namespaced_key)
.query_async(&mut connection)
.await
.map_redis_err()
})
.await
}
pub async fn exists(&self, key: &str) -> Result<bool, DataLayerError> {
let namespaced_key = self.keyspace.key(key);
let exists = self
.run_with_timeout("redis kv exists", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_redis_err()?;
redis::cmd("EXISTS")
.arg(&namespaced_key)
.query_async::<i64>(&mut connection)
.await
.map_redis_err()
})
.await?;
Ok(exists > 0)
}
pub async fn del(&self, key: &str) -> Result<i64, DataLayerError> { pub async fn del(&self, key: &str) -> Result<i64, DataLayerError> {
let namespaced_key = self.keyspace.key(key); let namespaced_key = self.keyspace.key(key);
self.run_with_timeout("redis kv del", async { self.run_with_timeout("redis kv del", async {
@@ -136,7 +189,7 @@ impl RedisKvRunner {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{RedisKvRunner, RedisKvRunnerConfig}; use super::{RedisKvRunner, RedisKvRunnerConfig};
use crate::driver::redis::{RedisClientConfig, RedisClientFactory, RedisKeyspace}; use crate::redis::{RedisClientConfig, RedisClientFactory, RedisKeyspace};
fn build_runner() -> RedisKvRunner { fn build_runner() -> RedisKvRunner {
let config = RedisClientConfig { let config = RedisClientConfig {
@@ -1,8 +1,8 @@
use std::future::Future; use std::future::Future;
use std::time::Duration; use std::time::Duration;
use crate::driver::redis::{RedisClient, RedisKeyspace};
use crate::error::RedisResultExt; use crate::error::RedisResultExt;
use crate::redis::{RedisClient, RedisKeyspace};
use crate::DataLayerError; use crate::DataLayerError;
use uuid::Uuid; use uuid::Uuid;
@@ -242,7 +242,7 @@ fn validate_lease(lease: &RedisLockLease) -> Result<(), DataLayerError> {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{RedisLockKey, RedisLockLease, RedisLockRunner, RedisLockRunnerConfig}; use super::{RedisLockKey, RedisLockLease, RedisLockRunner, RedisLockRunnerConfig};
use crate::driver::redis::{RedisClientConfig, RedisClientFactory}; use crate::redis::{RedisClientConfig, RedisClientFactory};
fn sample_runner() -> RedisLockRunner { fn sample_runner() -> RedisLockRunner {
let client = RedisClientFactory::new(RedisClientConfig { let client = RedisClientFactory::new(RedisClientConfig {
@@ -12,3 +12,14 @@ pub use stream::{
RedisConsumerGroup, RedisConsumerName, RedisStreamEntry, RedisStreamName, RedisConsumerGroup, RedisConsumerName, RedisStreamEntry, RedisStreamName,
RedisStreamReclaimConfig, RedisStreamReclaimResult, RedisStreamRunner, RedisStreamRunnerConfig, RedisStreamReclaimConfig, RedisStreamReclaimResult, RedisStreamRunner, RedisStreamRunnerConfig,
}; };
pub(crate) type RedisCmd = redis::Cmd;
pub(crate) type RedisScript = redis::Script;
pub(crate) fn cmd(name: &str) -> RedisCmd {
redis::cmd(name)
}
pub(crate) fn script(source: &str) -> RedisScript {
redis::Script::new(source)
}
@@ -1,6 +1,6 @@
use aether_cache::CacheKeyNamespace; use aether_cache::CacheKeyNamespace;
use crate::driver::redis::{RedisLockKey, RedisStreamName}; use crate::redis::{RedisLockKey, RedisStreamName};
#[derive(Debug, Clone, PartialEq, Eq)] #[derive(Debug, Clone, PartialEq, Eq)]
pub struct RedisKeyspace { pub struct RedisKeyspace {
@@ -6,8 +6,8 @@ use redis::from_redis_value;
use redis::streams::StreamReadReply; use redis::streams::StreamReadReply;
use redis::Value as RedisValue; use redis::Value as RedisValue;
use crate::driver::redis::{RedisClient, RedisKeyspace};
use crate::error::{redis_error, RedisResultExt}; use crate::error::{redis_error, RedisResultExt};
use crate::redis::{RedisClient, RedisKeyspace};
use crate::DataLayerError; use crate::DataLayerError;
#[derive(Debug, Clone, PartialEq, Eq, Hash)] #[derive(Debug, Clone, PartialEq, Eq, Hash)]
@@ -575,7 +575,7 @@ mod tests {
RedisStreamReclaimConfig, RedisStreamReclaimResult, RedisStreamRunner, RedisStreamReclaimConfig, RedisStreamReclaimResult, RedisStreamRunner,
RedisStreamRunnerConfig, RedisStreamRunnerConfig,
}; };
use crate::driver::redis::{RedisClientConfig, RedisClientFactory}; use crate::redis::{RedisClientConfig, RedisClientFactory};
use redis::Value as RedisValue; use redis::Value as RedisValue;
fn sample_runner() -> RedisStreamRunner { fn sample_runner() -> RedisStreamRunner {
-2
View File
@@ -11,12 +11,10 @@ async-stream.workspace = true
axum = { version = "0.8" } axum = { version = "0.8" }
chrono.workspace = true chrono.workspace = true
futures-util.workspace = true futures-util.workspace = true
redis.workspace = true
serde_json.workspace = true serde_json.workspace = true
sha2.workspace = true sha2.workspace = true
thiserror.workspace = true thiserror.workspace = true
tokio.workspace = true tokio.workspace = true
tracing.workspace = true tracing.workspace = true
tracing-subscriber.workspace = true tracing-subscriber.workspace = true
url.workspace = true
uuid.workspace = true uuid.workspace = true
+18 -21
View File
@@ -4,25 +4,32 @@ use axum::http::Response;
use futures_util::StreamExt; use futures_util::StreamExt;
use crate::concurrency::ConcurrencyPermit; use crate::concurrency::ConcurrencyPermit;
use crate::distributed::DistributedConcurrencyPermit;
#[derive(Debug)]
pub struct AdmissionPermit { pub struct AdmissionPermit {
_local: Option<ConcurrencyPermit>, _local: Option<ConcurrencyPermit>,
_distributed: Option<DistributedConcurrencyPermit>, _distributed: Option<Box<dyn Send + Sync>>,
}
impl std::fmt::Debug for AdmissionPermit {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("AdmissionPermit")
.field("has_local", &self._local.is_some())
.field("has_distributed", &self._distributed.is_some())
.finish()
}
} }
impl AdmissionPermit { impl AdmissionPermit {
pub fn from_parts( pub fn from_parts<D: Send + Sync + 'static>(
local: Option<ConcurrencyPermit>, local: Option<ConcurrencyPermit>,
distributed: Option<DistributedConcurrencyPermit>, distributed: Option<D>,
) -> Option<Self> { ) -> Option<Self> {
if local.is_none() && distributed.is_none() { if local.is_none() && distributed.is_none() {
None None
} else { } else {
Some(Self { Some(Self {
_local: local, _local: local,
_distributed: distributed, _distributed: distributed.map(|permit| Box::new(permit) as Box<dyn Send + Sync>),
}) })
} }
} }
@@ -70,7 +77,7 @@ fn hold_axum_response_permit(response: Response<Body>, permit: AdmissionPermit)
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{hold_admission_permit_until, maybe_hold_axum_response_permit, AdmissionPermit}; use super::{hold_admission_permit_until, maybe_hold_axum_response_permit, AdmissionPermit};
use crate::{ConcurrencyGate, DistributedConcurrencyGate}; use crate::ConcurrencyGate;
use axum::body::{to_bytes, Body}; use axum::body::{to_bytes, Body};
use axum::http::Response; use axum::http::Response;
@@ -96,12 +103,9 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn holds_combined_local_and_distributed_permit_until_future_finishes() { async fn holds_combined_local_and_distributed_permit_until_future_finishes() {
let local_gate = ConcurrencyGate::new("local", 1); let local_gate = ConcurrencyGate::new("local", 1);
let distributed_gate = DistributedConcurrencyGate::new_in_memory("distributed", 1);
let local = local_gate.try_acquire().expect("local permit"); let local = local_gate.try_acquire().expect("local permit");
let distributed = distributed_gate let distributed_gate = ConcurrencyGate::new("distributed", 1);
.try_acquire() let distributed = distributed_gate.try_acquire().expect("distributed permit");
.await
.expect("distributed permit");
let task = tokio::spawn(hold_admission_permit_until( let task = tokio::spawn(hold_admission_permit_until(
AdmissionPermit::from_parts(Some(local), Some(distributed)), AdmissionPermit::from_parts(Some(local), Some(distributed)),
@@ -116,19 +120,12 @@ mod tests {
"local permit should still be held" "local permit should still be held"
); );
assert!( assert!(
distributed_gate.try_acquire().await.is_err(), distributed_gate.try_acquire().is_err(),
"distributed permit should still be held" "distributed permit should still be held"
); );
task.await.expect("task should complete"); task.await.expect("task should complete");
assert_eq!(local_gate.snapshot().in_flight, 0); assert_eq!(local_gate.snapshot().in_flight, 0);
assert_eq!( assert_eq!(distributed_gate.snapshot().in_flight, 0);
distributed_gate
.snapshot()
.await
.expect("snapshot should build")
.in_flight,
0
);
} }
} }
+9 -482
View File
@@ -1,10 +1,4 @@
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc; use std::sync::Arc;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use tokio::task::JoinHandle;
use tracing::warn;
use uuid::Uuid;
use crate::concurrency::{ConcurrencyGate, ConcurrencyPermit}; use crate::concurrency::{ConcurrencyGate, ConcurrencyPermit};
use crate::metrics::{MetricKind, MetricLabel, MetricSample}; use crate::metrics::{MetricKind, MetricLabel, MetricSample};
@@ -68,80 +62,11 @@ impl DistributedConcurrencySnapshot {
} }
} }
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RedisDistributedConcurrencyConfig {
pub url: String,
pub key_prefix: Option<String>,
pub lease_ttl_ms: u64,
pub renew_interval_ms: u64,
pub command_timeout_ms: Option<u64>,
}
impl Default for RedisDistributedConcurrencyConfig {
fn default() -> Self {
Self {
url: String::new(),
key_prefix: None,
lease_ttl_ms: 30_000,
renew_interval_ms: 10_000,
command_timeout_ms: Some(1_000),
}
}
}
impl RedisDistributedConcurrencyConfig {
fn validate(&self) -> Result<(), DistributedConcurrencyError> {
let raw = self.url.trim();
if raw.is_empty() {
return Err(DistributedConcurrencyError::InvalidConfiguration(
"distributed concurrency redis url cannot be empty".to_string(),
));
}
url::Url::parse(raw).map_err(|err| {
DistributedConcurrencyError::InvalidConfiguration(format!(
"invalid distributed concurrency redis url: {err}"
))
})?;
if self.lease_ttl_ms == 0 {
return Err(DistributedConcurrencyError::InvalidConfiguration(
"distributed concurrency lease_ttl_ms must be positive".to_string(),
));
}
if self.renew_interval_ms == 0 {
return Err(DistributedConcurrencyError::InvalidConfiguration(
"distributed concurrency renew_interval_ms must be positive".to_string(),
));
}
if self.renew_interval_ms >= self.lease_ttl_ms {
return Err(DistributedConcurrencyError::InvalidConfiguration(
"distributed concurrency renew_interval_ms must be smaller than lease_ttl_ms"
.to_string(),
));
}
if matches!(self.command_timeout_ms, Some(0)) {
return Err(DistributedConcurrencyError::InvalidConfiguration(
"distributed concurrency command_timeout_ms must be positive".to_string(),
));
}
Ok(())
}
fn semaphore_key(&self, gate: &'static str) -> String {
prefixed_key(self.key_prefix.as_deref(), &format!("admission:{gate}"))
}
}
#[derive(Debug)]
enum DistributedConcurrencyBackend {
InMemory(Arc<ConcurrencyGate>),
Redis(Arc<RedisDistributedState>),
}
#[derive(Debug)] #[derive(Debug)]
struct DistributedConcurrencyState { struct DistributedConcurrencyState {
gate: &'static str, gate: &'static str,
limit: usize, limit: usize,
backend: DistributedConcurrencyBackend, gate_impl: Arc<ConcurrencyGate>,
} }
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -159,49 +84,11 @@ impl DistributedConcurrencyGate {
state: Arc::new(DistributedConcurrencyState { state: Arc::new(DistributedConcurrencyState {
gate, gate,
limit, limit,
backend: DistributedConcurrencyBackend::InMemory(Arc::new(ConcurrencyGate::new( gate_impl: Arc::new(ConcurrencyGate::new(gate, limit)),
gate, limit,
))),
}), }),
} }
} }
pub fn new_redis(
gate: &'static str,
limit: usize,
config: RedisDistributedConcurrencyConfig,
) -> Result<Self, DistributedConcurrencyError> {
if limit == 0 {
return Err(DistributedConcurrencyError::InvalidConfiguration(
"distributed concurrency gate limit must be positive".to_string(),
));
}
config.validate()?;
let client = redis::Client::open(config.url.clone()).map_err(|err| {
DistributedConcurrencyError::InvalidConfiguration(format!(
"failed to build distributed concurrency redis client: {err}"
))
})?;
Ok(Self {
state: Arc::new(DistributedConcurrencyState {
gate,
limit,
backend: DistributedConcurrencyBackend::Redis(Arc::new(RedisDistributedState {
gate,
limit,
client,
key: config.semaphore_key(gate),
lease_ttl_ms: config.lease_ttl_ms,
renew_interval_ms: config.renew_interval_ms,
command_timeout_ms: config.command_timeout_ms,
high_watermark: AtomicUsize::new(0),
rejected: AtomicU64::new(0),
})),
}),
})
}
pub fn gate(&self) -> &'static str { pub fn gate(&self) -> &'static str {
self.state.gate self.state.gate
} }
@@ -213,10 +100,10 @@ impl DistributedConcurrencyGate {
pub async fn try_acquire( pub async fn try_acquire(
&self, &self,
) -> Result<DistributedConcurrencyPermit, DistributedConcurrencyError> { ) -> Result<DistributedConcurrencyPermit, DistributedConcurrencyError> {
match &self.state.backend { self.state
DistributedConcurrencyBackend::InMemory(gate) => gate .gate_impl
.try_acquire() .try_acquire()
.map(DistributedConcurrencyPermit::from_in_memory) .map(|permit| DistributedConcurrencyPermit { _permit: permit })
.map_err(|err| match err { .map_err(|err| match err {
crate::ConcurrencyError::Saturated { gate, limit } => { crate::ConcurrencyError::Saturated { gate, limit } => {
DistributedConcurrencyError::Saturated { gate, limit } DistributedConcurrencyError::Saturated { gate, limit }
@@ -228,17 +115,13 @@ impl DistributedConcurrencyGate {
message: "in-memory distributed concurrency gate is closed".to_string(), message: "in-memory distributed concurrency gate is closed".to_string(),
} }
} }
}), })
DistributedConcurrencyBackend::Redis(state) => state.try_acquire().await,
}
} }
pub async fn snapshot( pub async fn snapshot(
&self, &self,
) -> Result<DistributedConcurrencySnapshot, DistributedConcurrencyError> { ) -> Result<DistributedConcurrencySnapshot, DistributedConcurrencyError> {
match &self.state.backend { let snapshot = self.state.gate_impl.snapshot();
DistributedConcurrencyBackend::InMemory(gate) => {
let snapshot = gate.snapshot();
Ok(DistributedConcurrencySnapshot { Ok(DistributedConcurrencySnapshot {
limit: snapshot.limit, limit: snapshot.limit,
in_flight: snapshot.in_flight, in_flight: snapshot.in_flight,
@@ -247,330 +130,16 @@ impl DistributedConcurrencyGate {
rejected: snapshot.rejected, rejected: snapshot.rejected,
}) })
} }
DistributedConcurrencyBackend::Redis(state) => state.snapshot().await,
}
}
} }
#[derive(Debug)] #[derive(Debug)]
pub struct DistributedConcurrencyPermit { pub struct DistributedConcurrencyPermit {
inner: DistributedConcurrencyPermitInner, _permit: ConcurrencyPermit,
}
#[derive(Debug)]
enum DistributedConcurrencyPermitInner {
InMemory(ConcurrencyPermit),
Redis {
state: Arc<RedisDistributedState>,
token: String,
renew_task: JoinHandle<()>,
},
}
impl DistributedConcurrencyPermit {
fn from_in_memory(permit: ConcurrencyPermit) -> Self {
Self {
inner: DistributedConcurrencyPermitInner::InMemory(permit),
}
}
fn from_redis(
state: Arc<RedisDistributedState>,
token: String,
renew_task: JoinHandle<()>,
) -> Self {
Self {
inner: DistributedConcurrencyPermitInner::Redis {
state,
token,
renew_task,
},
}
}
}
impl Drop for DistributedConcurrencyPermit {
fn drop(&mut self) {
match &mut self.inner {
DistributedConcurrencyPermitInner::InMemory(_permit) => {}
DistributedConcurrencyPermitInner::Redis {
state,
token,
renew_task,
} => {
renew_task.abort();
let state = Arc::clone(state);
let token = token.clone();
tokio::spawn(async move {
if let Err(err) = state.release(&token).await {
warn!(
gate = state.gate,
error = %err,
"failed to release distributed concurrency permit"
);
}
});
}
}
}
}
#[derive(Debug)]
struct RedisDistributedState {
gate: &'static str,
limit: usize,
client: redis::Client,
key: String,
lease_ttl_ms: u64,
renew_interval_ms: u64,
command_timeout_ms: Option<u64>,
high_watermark: AtomicUsize,
rejected: AtomicU64,
}
impl RedisDistributedState {
async fn try_acquire(
self: &Arc<Self>,
) -> Result<DistributedConcurrencyPermit, DistributedConcurrencyError> {
let token = format!("{}:{}", self.gate, Uuid::new_v4());
let now_ms = unix_time_ms();
let expires_at_ms = now_ms.saturating_add(self.lease_ttl_ms);
let key = self.key.clone();
let result: (i64, i64) = self
.run_with_timeout("acquire", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|err| self.unavailable(format!("connect failed: {err}")))?;
redis::Script::new(
"redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]); \
local count = redis.call('ZCARD', KEYS[1]); \
if count >= tonumber(ARGV[3]) then \
redis.call('PEXPIRE', KEYS[1], ARGV[5]); \
return {0, count}; \
end; \
redis.call('ZADD', KEYS[1], ARGV[2], ARGV[4]); \
count = redis.call('ZCARD', KEYS[1]); \
redis.call('PEXPIRE', KEYS[1], ARGV[5]); \
return {1, count};",
)
.key(&key)
.arg(now_ms as i64)
.arg(expires_at_ms as i64)
.arg(self.limit as i64)
.arg(&token)
.arg(self.lease_ttl_ms as i64)
.invoke_async::<(i64, i64)>(&mut connection)
.await
.map_err(|err| self.unavailable(format!("acquire failed: {err}")))
})
.await?;
let acquired = result.0 > 0;
let in_flight = result.1.max(0) as usize;
self.observe_in_flight(in_flight);
if !acquired {
self.rejected.fetch_add(1, Ordering::Relaxed);
return Err(DistributedConcurrencyError::Saturated {
gate: self.gate,
limit: self.limit,
});
}
let renew_state = Arc::clone(self);
let renew_token = token.clone();
let renew_task = tokio::spawn(async move {
let interval = Duration::from_millis(renew_state.renew_interval_ms);
loop {
tokio::time::sleep(interval).await;
if let Err(err) = renew_state.renew(&renew_token).await {
warn!(
gate = renew_state.gate,
error = %err,
"failed to renew distributed concurrency permit"
);
break;
}
}
});
Ok(DistributedConcurrencyPermit::from_redis(
Arc::clone(self),
token,
renew_task,
))
}
async fn snapshot(
&self,
) -> Result<DistributedConcurrencySnapshot, DistributedConcurrencyError> {
let in_flight = self.live_count().await?;
Ok(DistributedConcurrencySnapshot {
limit: self.limit,
in_flight,
available_permits: self.limit.saturating_sub(in_flight),
high_watermark: self.high_watermark.load(Ordering::Relaxed),
rejected: self.rejected.load(Ordering::Relaxed),
})
}
async fn renew(&self, token: &str) -> Result<(), DistributedConcurrencyError> {
let now_ms = unix_time_ms();
let expires_at_ms = now_ms.saturating_add(self.lease_ttl_ms);
let key = self.key.clone();
let renewed = self
.run_with_timeout("renew", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|err| self.unavailable(format!("connect failed: {err}")))?;
redis::Script::new(
"redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]); \
local score = redis.call('ZSCORE', KEYS[1], ARGV[2]); \
if not score then \
return 0; \
end; \
redis.call('ZADD', KEYS[1], 'XX', ARGV[3], ARGV[2]); \
redis.call('PEXPIRE', KEYS[1], ARGV[4]); \
return 1;",
)
.key(&key)
.arg(now_ms as i64)
.arg(token)
.arg(expires_at_ms as i64)
.arg(self.lease_ttl_ms as i64)
.invoke_async::<i64>(&mut connection)
.await
.map_err(|err| self.unavailable(format!("renew failed: {err}")))
})
.await?;
if renewed == 0 {
return Err(self.unavailable("lease token expired".to_string()));
}
Ok(())
}
async fn release(&self, token: &str) -> Result<(), DistributedConcurrencyError> {
let key = self.key.clone();
self.run_with_timeout("release", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|err| self.unavailable(format!("connect failed: {err}")))?;
redis::Script::new(
"local removed = redis.call('ZREM', KEYS[1], ARGV[1]); \
if removed > 0 and redis.call('ZCARD', KEYS[1]) == 0 then \
redis.call('DEL', KEYS[1]); \
end; \
return removed;",
)
.key(&key)
.arg(token)
.invoke_async::<i64>(&mut connection)
.await
.map_err(|err| self.unavailable(format!("release failed: {err}")))?;
Ok(())
})
.await
}
async fn live_count(&self) -> Result<usize, DistributedConcurrencyError> {
let now_ms = unix_time_ms();
let key = self.key.clone();
let count = self
.run_with_timeout("snapshot", async {
let mut connection = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|err| self.unavailable(format!("connect failed: {err}")))?;
redis::Script::new(
"redis.call('ZREMRANGEBYSCORE', KEYS[1], '-inf', ARGV[1]); \
return redis.call('ZCARD', KEYS[1]);",
)
.key(&key)
.arg(now_ms as i64)
.invoke_async::<i64>(&mut connection)
.await
.map_err(|err| self.unavailable(format!("snapshot failed: {err}")))
})
.await?
.max(0) as usize;
self.observe_in_flight(count);
Ok(count)
}
async fn run_with_timeout<T, F>(
&self,
operation: &'static str,
future: F,
) -> Result<T, DistributedConcurrencyError>
where
F: std::future::Future<Output = Result<T, DistributedConcurrencyError>>,
{
if let Some(timeout_ms) = self.command_timeout_ms {
tokio::time::timeout(Duration::from_millis(timeout_ms), future)
.await
.map_err(|_| {
self.unavailable(format!(
"{operation} exceeded {timeout_ms}ms command timeout"
))
})?
} else {
future.await
}
}
fn unavailable(&self, message: String) -> DistributedConcurrencyError {
DistributedConcurrencyError::Unavailable {
gate: self.gate,
limit: self.limit,
message,
}
}
fn observe_in_flight(&self, in_flight: usize) {
let mut observed = self.high_watermark.load(Ordering::Acquire);
while in_flight > observed {
match self.high_watermark.compare_exchange_weak(
observed,
in_flight,
Ordering::AcqRel,
Ordering::Acquire,
) {
Ok(_) => break,
Err(next) => observed = next,
}
}
}
}
fn prefixed_key(prefix: Option<&str>, raw_key: &str) -> String {
let prefix = prefix.unwrap_or_default().trim().trim_matches(':');
if prefix.is_empty() {
raw_key.trim_matches(':').to_string()
} else {
format!("{prefix}:{}", raw_key.trim_matches(':'))
}
}
fn unix_time_ms() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as u64
} }
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::{ use super::{DistributedConcurrencyError, DistributedConcurrencyGate};
DistributedConcurrencyError, DistributedConcurrencyGate, RedisDistributedConcurrencyConfig,
};
#[tokio::test] #[tokio::test]
async fn shared_in_memory_gate_rejects_second_acquire() { async fn shared_in_memory_gate_rejects_second_acquire() {
@@ -599,46 +168,4 @@ mod tests {
let snapshot = gate.snapshot().await.expect("snapshot should build"); let snapshot = gate.snapshot().await.expect("snapshot should build");
assert_eq!(snapshot.in_flight, 0); assert_eq!(snapshot.in_flight, 0);
} }
#[test]
fn rejects_invalid_redis_config() {
let error = DistributedConcurrencyGate::new_redis(
"shared",
1,
RedisDistributedConcurrencyConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
lease_ttl_ms: 10_000,
renew_interval_ms: 10_000,
command_timeout_ms: Some(1_000),
},
)
.expect_err("equal renew interval should fail");
assert_eq!(
error,
DistributedConcurrencyError::InvalidConfiguration(
"distributed concurrency renew_interval_ms must be smaller than lease_ttl_ms"
.to_string()
)
);
}
#[test]
fn builds_redis_gate_without_touching_network() {
let gate = DistributedConcurrencyGate::new_redis(
"gateway_requests_distributed",
2,
RedisDistributedConcurrencyConfig {
url: "redis://127.0.0.1/0".to_string(),
key_prefix: Some("aether".to_string()),
lease_ttl_ms: 15_000,
renew_interval_ms: 5_000,
command_timeout_ms: Some(1_000),
},
)
.expect("redis gate should build");
assert_eq!(gate.gate(), "gateway_requests_distributed");
assert_eq!(gate.limit(), 2);
}
} }
+1 -1
View File
@@ -20,7 +20,7 @@ pub use concurrency::{ConcurrencyError, ConcurrencyGate, ConcurrencyPermit, Conc
pub use config::ServiceRuntimeConfig; pub use config::ServiceRuntimeConfig;
pub use distributed::{ pub use distributed::{
DistributedConcurrencyError, DistributedConcurrencyGate, DistributedConcurrencyPermit, DistributedConcurrencyError, DistributedConcurrencyGate, DistributedConcurrencyPermit,
DistributedConcurrencySnapshot, RedisDistributedConcurrencyConfig, DistributedConcurrencySnapshot,
}; };
pub use error::RuntimeBootstrapError; pub use error::RuntimeBootstrapError;
pub use metrics::{prometheus_response, service_up_sample, MetricKind, MetricLabel, MetricSample}; pub use metrics::{prometheus_response, service_up_sample, MetricKind, MetricLabel, MetricSample};
+1
View File
@@ -13,6 +13,7 @@ aether-contracts.workspace = true
aether-gateway.workspace = true aether-gateway.workspace = true
aether-http.workspace = true aether-http.workspace = true
aether-runtime.workspace = true aether-runtime.workspace = true
aether-runtime-state.workspace = true
axum.workspace = true axum.workspace = true
bytes.workspace = true bytes.workspace = true
http.workspace = true http.workspace = true
@@ -6,12 +6,12 @@ use aether_data::driver::postgres::{
DatabaseRecordId, PostgresLeaseClaimOptions, PostgresLeaseClaimSpec, PostgresLeaseRunnerConfig, DatabaseRecordId, PostgresLeaseClaimOptions, PostgresLeaseClaimSpec, PostgresLeaseRunnerConfig,
PostgresPoolConfig, PostgresPoolConfig,
}; };
use aether_data::driver::redis::{ use aether_data::PostgresBackend;
RedisClientConfig, RedisConsumerGroup, RedisConsumerName, RedisLockLease, RedisLockRunner, use aether_runtime_state::{
RedisLockRunnerConfig, RedisStreamName, RedisStreamReclaimConfig, RedisStreamRunner, RedisClientConfig, RedisClientFactory, RedisConsumerGroup, RedisConsumerName, RedisKeyspace,
RedisStreamRunnerConfig, RedisLockLease, RedisLockRunner, RedisLockRunnerConfig, RedisStreamName,
RedisStreamReclaimConfig, RedisStreamRunner, RedisStreamRunnerConfig,
}; };
use aether_data::{PostgresBackend, RedisBackend};
use aether_testkit::{init_test_runtime_for, ManagedPostgresServer, ManagedRedisServer}; use aether_testkit::{init_test_runtime_for, ManagedPostgresServer, ManagedRedisServer};
use futures_util::stream::{self, StreamExt}; use futures_util::stream::{self, StreamExt};
use serde::Serialize; use serde::Serialize;
@@ -189,10 +189,12 @@ async fn run_suite(
}) })
.expect("postgres url should resolve"); .expect("postgres url should resolve");
let redis_backend = RedisBackend::from_config(RedisClientConfig { let redis_factory = RedisClientFactory::new(RedisClientConfig {
url: redis_url.clone(), url: redis_url.clone(),
key_prefix: Some(format!("aether-dependency-pressure-{}", std::process::id())), key_prefix: Some(format!("aether-dependency-pressure-{}", std::process::id())),
})?; })?;
let redis_client = redis_factory.connect_lazy()?;
let redis_keyspace = redis_factory.config().keyspace();
let postgres_backend = PostgresBackend::from_config(PostgresPoolConfig { let postgres_backend = PostgresBackend::from_config(PostgresPoolConfig {
database_url: postgres_url.clone(), database_url: postgres_url.clone(),
min_connections: 1, min_connections: 1,
@@ -206,22 +208,30 @@ async fn run_suite(
bootstrap_postgres_lease_table(postgres_backend.pool_clone(), config).await?; bootstrap_postgres_lease_table(postgres_backend.pool_clone(), config).await?;
let lock_runner = redis_backend.lock_runner(RedisLockRunnerConfig { let lock_runner = RedisLockRunner::new(
redis_client.clone(),
redis_keyspace.clone(),
RedisLockRunnerConfig {
command_timeout_ms: Some(config.timeout.as_millis() as u64), command_timeout_ms: Some(config.timeout.as_millis() as u64),
default_ttl_ms: 5_000, default_ttl_ms: 5_000,
})?; },
let stream_runner = redis_backend.stream_runner(RedisStreamRunnerConfig { )?;
let stream_runner = RedisStreamRunner::new(
redis_client.clone(),
redis_keyspace.clone(),
RedisStreamRunnerConfig {
command_timeout_ms: Some(config.timeout.as_millis() as u64), command_timeout_ms: Some(config.timeout.as_millis() as u64),
read_block_ms: Some(10), read_block_ms: Some(10),
read_count: 64, read_count: 64,
})?; },
)?;
let lease_runner = postgres_backend.lease_runner(PostgresLeaseRunnerConfig { let lease_runner = postgres_backend.lease_runner(PostgresLeaseRunnerConfig {
statement_timeout_ms: Some(config.timeout.as_millis() as u64), statement_timeout_ms: Some(config.timeout.as_millis() as u64),
lock_timeout_ms: Some(1_000), lock_timeout_ms: Some(1_000),
})?; })?;
let redis_lock = benchmark_redis_lock(&redis_backend, &lock_runner, config).await?; let redis_lock = benchmark_redis_lock(&redis_keyspace, &lock_runner, config).await?;
let redis_stream = benchmark_redis_stream(&redis_backend, &stream_runner, config).await?; let redis_stream = benchmark_redis_stream(&redis_keyspace, &stream_runner, config).await?;
let postgres_lease = benchmark_postgres_lease(&lease_runner, config).await?; let postgres_lease = benchmark_postgres_lease(&lease_runner, config).await?;
Ok(DependencyPressureBaselineReport { Ok(DependencyPressureBaselineReport {
@@ -263,7 +273,7 @@ async fn bootstrap_postgres_lease_table(
} }
async fn benchmark_redis_lock( async fn benchmark_redis_lock(
backend: &RedisBackend, keyspace: &RedisKeyspace,
runner: &RedisLockRunner, runner: &RedisLockRunner,
config: &DependencyPressureBaselineConfig, config: &DependencyPressureBaselineConfig,
) -> Result<RedisLockPressureReport, Box<dyn std::error::Error>> { ) -> Result<RedisLockPressureReport, Box<dyn std::error::Error>> {
@@ -271,7 +281,6 @@ async fn benchmark_redis_lock(
let renew = Arc::new(SummaryCollector::default()); let renew = Arc::new(SummaryCollector::default());
let release = Arc::new(SummaryCollector::default()); let release = Arc::new(SummaryCollector::default());
let next = Arc::new(std::sync::atomic::AtomicUsize::new(0)); let next = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let keyspace = backend.keyspace();
stream::iter(0..config.redis_lock_concurrency) stream::iter(0..config.redis_lock_concurrency)
.for_each_concurrent(config.redis_lock_concurrency, |_| { .for_each_concurrent(config.redis_lock_concurrency, |_| {
@@ -338,11 +347,11 @@ async fn record_redis_lock_follow_up(
} }
async fn benchmark_redis_stream( async fn benchmark_redis_stream(
backend: &RedisBackend, keyspace: &RedisKeyspace,
runner: &RedisStreamRunner, runner: &RedisStreamRunner,
config: &DependencyPressureBaselineConfig, config: &DependencyPressureBaselineConfig,
) -> Result<RedisStreamPressureReport, Box<dyn std::error::Error>> { ) -> Result<RedisStreamPressureReport, Box<dyn std::error::Error>> {
let stream = backend.keyspace().stream_name("dependency-pressure"); let stream = keyspace.stream_name("dependency-pressure");
let group = RedisConsumerGroup("dependency-group".to_string()); let group = RedisConsumerGroup("dependency-group".to_string());
let consumer_a = RedisConsumerName("consumer-a".to_string()); let consumer_a = RedisConsumerName("consumer-a".to_string());
let consumer_b = RedisConsumerName("consumer-b".to_string()); let consumer_b = RedisConsumerName("consumer-b".to_string());

Some files were not shown because too many files have changed in this diff Show More