refactor: isolate dispatch scheduling core

This commit is contained in:
fawney19
2026-05-12 13:15:41 +08:00
parent 81ee27cdea
commit fa73655134
50 changed files with 4467 additions and 3202 deletions

View File

@@ -469,7 +469,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if created.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(created)
}
@@ -485,7 +485,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if created.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(created)
}
@@ -500,7 +500,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -515,7 +515,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(deleted)
}
@@ -550,7 +550,7 @@ impl AppState {
}
}
if !endpoint_ids.is_empty() || !key_ids.is_empty() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(())
}
@@ -565,7 +565,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if created.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(created)
}
@@ -580,7 +580,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -595,7 +595,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(deleted)
}
@@ -610,7 +610,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -631,7 +631,7 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -672,7 +672,7 @@ impl AppState {
);
}
}
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(deleted)
}
@@ -806,8 +806,128 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use crate::cache::SchedulerAffinityTarget;
use crate::data::GatewayDataState;
use crate::AppState;
fn sample_provider() -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
"provider-1".to_string(),
"Provider 1".to_string(),
Some("https://example.com".to_string()),
"openai".to_string(),
)
.expect("provider should build")
}
fn sample_endpoint() -> StoredProviderCatalogEndpoint {
StoredProviderCatalogEndpoint::new(
"endpoint-1".to_string(),
"provider-1".to_string(),
"openai:chat".to_string(),
Some("openai".to_string()),
Some("chat".to_string()),
true,
)
.expect("endpoint should build")
.with_transport_fields(
"https://api.example.com/v1".to_string(),
None,
None,
None,
None,
None,
None,
None,
)
.expect("endpoint transport should build")
}
fn sample_key() -> StoredProviderCatalogKey {
StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"Key 1".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
}
#[tokio::test]
async fn provider_catalog_update_invalidates_scheduler_affinity_and_transport_snapshot_cache() {
let provider = sample_provider();
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![provider.clone()],
vec![sample_endpoint()],
vec![sample_key()],
));
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(repository)
.with_encryption_key_for_tests("test-encryption-key"),
);
let snapshot = state
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should read")
.expect("provider transport should exist");
assert!(!snapshot.provider.keep_priority_on_conversion);
let cache_key = "scheduler_affinity:api-key-1:openai:chat:gpt-5";
let ttl = Duration::from_secs(300);
state.remember_scheduler_affinity_target(
cache_key,
SchedulerAffinityTarget {
provider_id: "provider-1".to_string(),
endpoint_id: "endpoint-1".to_string(),
key_id: "key-1".to_string(),
},
ttl,
128,
);
assert!(state
.read_scheduler_affinity_target(cache_key, ttl)
.is_some());
let initial_epoch = state.scheduler_affinity_epoch();
let mut updated_provider = provider;
updated_provider.keep_priority_on_conversion = true;
updated_provider.provider_priority = -10;
state
.update_provider_catalog_provider(&updated_provider)
.await
.expect("provider update should succeed")
.expect("provider should update");
assert!(state.scheduler_affinity_epoch() > initial_epoch);
assert!(state
.read_scheduler_affinity_target(cache_key, ttl)
.is_none());
let snapshot = state
.read_provider_transport_snapshot("provider-1", "endpoint-1", "key-1")
.await
.expect("provider transport should read after update")
.expect("provider transport should exist after update");
assert!(snapshot.provider.keep_priority_on_conversion);
}
}

View File

@@ -59,6 +59,30 @@ use crate::maintenance::spawn_usage_cleanup_worker;
use crate::maintenance::spawn_wallet_daily_usage_aggregation_worker;
const SYSTEM_CONFIG_CACHE_TTL: Duration = Duration::from_secs(3);
const SCHEDULER_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] = &[
"enable_format_conversion",
"keep_priority_on_conversion",
"provider_priority_mode",
"scheduling_mode",
];
const AUTH_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] =
&[crate::constants::DEFAULT_USER_GROUP_CONFIG_KEY];
const FRONTDOOR_RPM_AFFECTING_SYSTEM_CONFIG_KEYS: &[&str] = &["rate_limit_per_minute"];
fn system_config_key_affects_scheduler(key: &str) -> bool {
let key = key.trim();
SCHEDULER_AFFECTING_SYSTEM_CONFIG_KEYS.contains(&key)
}
fn system_config_key_affects_auth(key: &str) -> bool {
let key = key.trim();
AUTH_AFFECTING_SYSTEM_CONFIG_KEYS.contains(&key)
}
fn system_config_key_affects_frontdoor_rpm(key: &str) -> bool {
let key = key.trim();
FRONTDOOR_RPM_AFFECTING_SYSTEM_CONFIG_KEYS.contains(&key)
}
impl AppState {
fn usage_worker_queue_for(
@@ -143,7 +167,10 @@ impl AppState {
pub(crate) fn replace_data_state(&mut self, data: Arc<GatewayDataState>) {
self.clear_provider_transport_snapshot_cache();
self.invalidate_scheduler_affinity_cache();
self.invalidate_auth_context_cache();
self.system_config_cache.clear();
self.frontdoor_user_rpm.clear_system_default_cache();
let data = Arc::new(
(*data)
.clone()
@@ -492,11 +519,7 @@ impl AppState {
.upsert_system_config_value(key, value, description)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.system_config_cache.insert(
key.to_string(),
Some(value.clone()),
SYSTEM_CONFIG_CACHE_TTL,
);
self.remember_system_config_write(key, Some(value.clone()));
Ok(value)
}
@@ -515,10 +538,13 @@ impl AppState {
value: &serde_json::Value,
description: Option<&str>,
) -> Result<crate::data::state::StoredSystemConfigEntry, GatewayError> {
self.data
let entry = self
.data
.upsert_system_config_entry(key, value, description)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.remember_system_config_write(entry.key.as_str(), Some(entry.value.clone()));
Ok(entry)
}
pub(crate) async fn delete_system_config_value(&self, key: &str) -> Result<bool, GatewayError> {
@@ -529,9 +555,41 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.system_config_cache
.insert(key.to_string(), None, SYSTEM_CONFIG_CACHE_TTL);
if deleted && system_config_key_affects_scheduler(key) {
self.invalidate_scheduler_affinity_cache();
}
if deleted && system_config_key_affects_auth(key) {
self.invalidate_auth_context_cache();
}
if deleted && system_config_key_affects_frontdoor_rpm(key) {
self.frontdoor_user_rpm.clear_system_default_cache();
}
Ok(deleted)
}
pub(crate) fn invalidate_provider_routing_caches(&self) {
self.clear_provider_transport_snapshot_cache();
self.invalidate_scheduler_affinity_cache();
}
pub(crate) fn invalidate_auth_context_cache(&self) {
self.auth_context_cache.clear();
}
fn remember_system_config_write(&self, key: &str, value: Option<serde_json::Value>) {
self.system_config_cache
.insert(key.to_string(), value, SYSTEM_CONFIG_CACHE_TTL);
if system_config_key_affects_scheduler(key) {
self.invalidate_scheduler_affinity_cache();
}
if system_config_key_affects_auth(key) {
self.invalidate_auth_context_cache();
}
if system_config_key_affects_frontdoor_rpm(key) {
self.frontdoor_user_rpm.clear_system_default_cache();
}
}
pub(crate) async fn read_admin_system_stats(
&self,
) -> Result<aether_data::repository::system::AdminSystemStats, GatewayError> {
@@ -558,7 +616,7 @@ impl AppState {
| aether_data::repository::system::AdminSystemPurgeTarget::Stats
) {
self.system_config_cache.clear();
self.clear_provider_transport_snapshot_cache();
self.invalidate_provider_routing_caches();
}
Ok(summary)
}
@@ -1190,6 +1248,7 @@ mod tests {
use serde_json::json;
use super::AppState;
use crate::cache::SchedulerAffinityTarget;
use crate::data::GatewayDataState;
#[tokio::test]
@@ -1237,6 +1296,91 @@ mod tests {
);
}
#[tokio::test]
async fn system_config_entry_write_refreshes_cache_and_scheduler_affinity_for_routing_keys() {
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::disabled().with_system_config_values_for_tests([(
"keep_priority_on_conversion".to_string(),
json!(false),
)]),
);
let cache_key = "scheduler_affinity:api-key-1:openai:chat:gpt-5";
let ttl = std::time::Duration::from_secs(300);
assert_eq!(
state
.read_system_config_json_value("keep_priority_on_conversion")
.await
.expect("system config read should succeed"),
Some(json!(false))
);
state.remember_scheduler_affinity_target(
cache_key,
SchedulerAffinityTarget {
provider_id: "provider-old".to_string(),
endpoint_id: "endpoint-old".to_string(),
key_id: "key-old".to_string(),
},
ttl,
128,
);
assert!(state
.read_scheduler_affinity_target(cache_key, ttl)
.is_some());
let initial_epoch = state.scheduler_affinity_epoch();
state
.upsert_system_config_entry("keep_priority_on_conversion", &json!(true), None)
.await
.expect("admin config write should succeed");
assert_eq!(
state
.read_system_config_json_value("keep_priority_on_conversion")
.await
.expect("system config read should use refreshed cache"),
Some(json!(true))
);
assert!(state.scheduler_affinity_epoch() > initial_epoch);
assert_eq!(state.read_scheduler_affinity_target(cache_key, ttl), None);
}
#[tokio::test]
async fn system_config_write_refreshes_frontdoor_rpm_default_cache() {
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(
GatewayDataState::disabled().with_system_config_values_for_tests([(
"rate_limit_per_minute".to_string(),
json!(1),
)]),
);
assert_eq!(
state
.frontdoor_user_rpm()
.current_system_default_limit(&state)
.await
.expect("default rpm limit should read"),
1
);
state
.upsert_system_config_entry("rate_limit_per_minute", &json!(0), None)
.await
.expect("rpm system config should update");
assert_eq!(
state
.frontdoor_user_rpm()
.current_system_default_limit(&state)
.await
.expect("default rpm limit should use refreshed value"),
0
);
}
#[tokio::test]
async fn replacing_data_state_clears_system_config_cache() {
let mut state = AppState::new()

View File

@@ -182,10 +182,15 @@ impl AppState {
record: aether_data::repository::auth::CreateUserApiKeyRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.create_user_api_key(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn create_standalone_api_key(
@@ -193,10 +198,15 @@ impl AppState {
record: aether_data::repository::auth::CreateStandaloneApiKeyRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.create_standalone_api_key(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn update_user_api_key_basic(
@@ -204,10 +214,15 @@ impl AppState {
record: aether_data::repository::auth::UpdateUserApiKeyBasicRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.update_user_api_key_basic(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn update_standalone_api_key_basic(
@@ -215,10 +230,15 @@ impl AppState {
record: aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.update_standalone_api_key_basic(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn set_user_api_key_active(
@@ -228,10 +248,15 @@ impl AppState {
is_active: bool,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.set_user_api_key_active(user_id, api_key_id, is_active)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn set_standalone_api_key_active(
@@ -240,10 +265,15 @@ impl AppState {
is_active: bool,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.set_standalone_api_key_active(api_key_id, is_active)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn set_user_api_key_locked(
@@ -252,10 +282,15 @@ impl AppState {
api_key_id: &str,
is_locked: bool,
) -> Result<bool, GatewayError> {
self.data
let updated = self
.data
.set_user_api_key_locked(user_id, api_key_id, is_locked)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.invalidate_auth_context_cache();
}
Ok(updated)
}
pub(crate) async fn set_user_api_key_allowed_providers(
@@ -265,10 +300,15 @@ impl AppState {
allowed_providers: Option<Vec<String>>,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.set_user_api_key_allowed_providers(user_id, api_key_id, allowed_providers)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn set_user_api_key_force_capabilities(
@@ -278,10 +318,15 @@ impl AppState {
force_capabilities: Option<serde_json::Value>,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
self.data
let api_key = self
.data
.set_user_api_key_force_capabilities(user_id, api_key_id, force_capabilities)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if api_key.is_some() {
self.invalidate_auth_context_cache();
}
Ok(api_key)
}
pub(crate) async fn delete_user_api_key(
@@ -289,19 +334,29 @@ impl AppState {
user_id: &str,
api_key_id: &str,
) -> Result<bool, GatewayError> {
self.data
let deleted = self
.data
.delete_user_api_key(user_id, api_key_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.invalidate_auth_context_cache();
}
Ok(deleted)
}
pub(crate) async fn delete_standalone_api_key(
&self,
api_key_id: &str,
) -> Result<bool, GatewayError> {
self.data
let deleted = self
.data
.delete_standalone_api_key(api_key_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.invalidate_auth_context_cache();
}
Ok(deleted)
}
}

View File

@@ -248,10 +248,15 @@ impl AppState {
&self,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
let group = self
.data
.create_user_group(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if group.is_some() {
self.invalidate_auth_context_cache();
}
Ok(group)
}
pub(crate) async fn update_user_group(
@@ -259,17 +264,27 @@ impl AppState {
group_id: &str,
record: aether_data::repository::users::UpsertUserGroupRecord,
) -> Result<Option<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
let group = self
.data
.update_user_group(group_id, record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if group.is_some() {
self.invalidate_auth_context_cache();
}
Ok(group)
}
pub(crate) async fn delete_user_group(&self, group_id: &str) -> Result<bool, GatewayError> {
self.data
let deleted = self
.data
.delete_user_group(group_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if deleted {
self.invalidate_auth_context_cache();
}
Ok(deleted)
}
pub(crate) async fn list_user_group_members(
@@ -287,10 +302,13 @@ impl AppState {
group_id: &str,
user_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroupMember>, GatewayError> {
self.data
let members = self
.data
.replace_user_group_members(group_id, user_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.invalidate_auth_context_cache();
Ok(members)
}
pub(crate) async fn list_user_groups_for_user(
@@ -318,10 +336,13 @@ impl AppState {
user_id: &str,
group_ids: &[String],
) -> Result<Vec<aether_data::repository::users::StoredUserGroup>, GatewayError> {
self.data
let groups = self
.data
.replace_user_groups_for_user(user_id, group_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.invalidate_auth_context_cache();
Ok(groups)
}
pub(crate) async fn add_user_to_group(
@@ -329,10 +350,15 @@ impl AppState {
group_id: &str,
user_id: &str,
) -> Result<bool, GatewayError> {
self.data
let added = self
.data
.add_user_to_group(group_id, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if added {
self.invalidate_auth_context_cache();
}
Ok(added)
}
pub(crate) async fn is_other_user_auth_email_taken(
@@ -417,13 +443,19 @@ impl AppState {
.lock()
.expect("auth user store should lock")
.insert(user.id.clone(), user.clone());
self.invalidate_auth_context_cache();
return Ok(Some(user));
}
self.data
let user = self
.data
.update_local_auth_user_profile(user_id, email, username)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if user.is_some() {
self.invalidate_auth_context_cache();
}
Ok(user)
}
pub(crate) async fn update_local_auth_user_password_hash(
@@ -606,10 +638,14 @@ impl AppState {
user.is_active = is_active;
}
let _ = (rate_limit_present, rate_limit);
return Ok(Some(user.clone()));
let user = user.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(user));
}
self.data
let user = self
.data
.update_local_auth_user_admin_fields(
user_id,
role,
@@ -624,7 +660,11 @@ impl AppState {
is_active,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if user.is_some() {
self.invalidate_auth_context_cache();
}
Ok(user)
}
pub(crate) async fn update_local_auth_user_policy_modes(
@@ -651,10 +691,14 @@ impl AppState {
user.allowed_models_mode = mode;
}
let _ = rate_limit_mode;
return Ok(Some(user.clone()));
let user = user.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(user));
}
self.data
let user = self
.data
.update_local_auth_user_policy_modes(
user_id,
allowed_providers_mode,
@@ -663,7 +707,11 @@ impl AppState {
rate_limit_mode,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if user.is_some() {
self.invalidate_auth_context_cache();
}
Ok(user)
}
pub(crate) async fn touch_auth_user_last_login(
@@ -821,3 +869,90 @@ fn normalized_user_group_ids(group_ids: &[String]) -> BTreeSet<String> {
.map(ToOwned::to_owned)
.collect()
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use aether_data::repository::users::{InMemoryUserReadRepository, UpsertUserGroupRecord};
use crate::control::GatewayControlAuthContext;
use crate::data::GatewayDataState;
use crate::AppState;
fn user_group_record(
allowed_models: Option<Vec<&str>>,
allowed_models_mode: &str,
) -> UpsertUserGroupRecord {
UpsertUserGroupRecord {
name: "Team".to_string(),
description: None,
priority: 0,
allowed_providers: None,
allowed_providers_mode: "unrestricted".to_string(),
allowed_api_formats: None,
allowed_api_formats_mode: "unrestricted".to_string(),
allowed_models: allowed_models.map(|values| {
values
.into_iter()
.map(ToOwned::to_owned)
.collect::<Vec<_>>()
}),
allowed_models_mode: allowed_models_mode.to_string(),
rate_limit: None,
rate_limit_mode: "inherit".to_string(),
}
}
fn cached_auth_context() -> GatewayControlAuthContext {
GatewayControlAuthContext {
user_id: "user-1".to_string(),
api_key_id: "key-1".to_string(),
username: Some("alice".to_string()),
api_key_name: Some("default".to_string()),
balance_remaining: None,
access_allowed: true,
user_rate_limit: None,
api_key_rate_limit: None,
api_key_is_standalone: false,
admin_bypass_limits: false,
local_rejection: None,
allowed_models: Some(vec!["gpt-4.1".to_string()]),
}
}
#[tokio::test]
async fn updating_user_group_invalidates_cached_auth_context() {
let repository = Arc::new(InMemoryUserReadRepository::default());
let state = AppState::new()
.expect("state should build")
.with_data_state_for_tests(GatewayDataState::with_user_reader_for_tests(repository));
let group = state
.create_user_group(user_group_record(Some(vec!["gpt-4.1"]), "specific"))
.await
.expect("group should create")
.expect("group should exist");
let cache_key = "auth-context-cache-key".to_string();
let ttl = Duration::from_secs(60);
state
.auth_context_cache
.insert(cache_key.clone(), cached_auth_context(), ttl, 10);
assert!(state
.auth_context_cache
.get_fresh(&cache_key, ttl)
.is_some());
state
.update_user_group(&group.id, user_group_record(None, "unrestricted"))
.await
.expect("group should update")
.expect("group should exist after update");
assert!(state
.auth_context_cache
.get_fresh(&cache_key, ttl)
.is_none());
}
}

View File

@@ -234,13 +234,19 @@ impl AppState {
.lock()
.expect("auth wallet store should lock")
.insert(wallet.id.clone(), wallet.clone());
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.initialize_auth_user_wallet(user_id, initial_gift_usd, unlimited)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
pub(crate) async fn initialize_auth_api_key_wallet(
@@ -284,13 +290,19 @@ impl AppState {
.lock()
.expect("auth wallet store should lock")
.insert(wallet.id.clone(), wallet.clone());
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.initialize_auth_api_key_wallet(api_key_id, initial_gift_usd, unlimited)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
pub(crate) async fn update_auth_user_wallet_limit_mode(
@@ -310,13 +322,21 @@ impl AppState {
let _ = wallet_id;
wallet.limit_mode = limit_mode.to_string();
wallet.updated_at_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
return Ok(Some(wallet.clone()));
let wallet = wallet.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.update_auth_user_wallet_limit_mode(user_id, limit_mode)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
pub(crate) async fn update_auth_api_key_wallet_limit_mode(
@@ -336,13 +356,21 @@ impl AppState {
let _ = wallet_id;
wallet.limit_mode = limit_mode.to_string();
wallet.updated_at_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
return Ok(Some(wallet.clone()));
let wallet = wallet.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.update_auth_api_key_wallet_limit_mode(api_key_id, limit_mode)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
#[allow(clippy::too_many_arguments)]
@@ -381,10 +409,14 @@ impl AppState {
if let Some(updated_at_unix_secs) = updated_at_unix_secs {
wallet.updated_at_unix_secs = updated_at_unix_secs;
}
return Ok(Some(wallet.clone()));
let wallet = wallet.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.update_auth_user_wallet_snapshot(
user_id,
balance,
@@ -399,7 +431,11 @@ impl AppState {
updated_at_unix_secs,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
#[allow(clippy::too_many_arguments)]
@@ -438,10 +474,14 @@ impl AppState {
if let Some(updated_at_unix_secs) = updated_at_unix_secs {
wallet.updated_at_unix_secs = updated_at_unix_secs;
}
return Ok(Some(wallet.clone()));
let wallet = wallet.clone();
drop(guard);
self.invalidate_auth_context_cache();
return Ok(Some(wallet));
}
self.data
let wallet = self
.data
.update_auth_api_key_wallet_snapshot(
api_key_id,
balance,
@@ -456,6 +496,10 @@ impl AppState {
updated_at_unix_secs,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if wallet.is_some() {
self.invalidate_auth_context_cache();
}
Ok(wallet)
}
}