feat(security): harden gateway boundaries and usage policies

Consolidate subscription usage policy enforcement, privacy-safe persistence, and gateway security hardening into one reviewable change.

Includes bounded HTTP and execution envelopes, header and protocol guards, DNS and relay validation, authentication and secret projection hardening, secure backup/install paths, and regression coverage.
This commit is contained in:
elky
2026-09-04 03:45:52 +08:00
parent ddcbeb3ae9
commit 579f2c7cc1
1019 changed files with 190437 additions and 26080 deletions
+3 -2
View File
@@ -6,6 +6,7 @@ pub(crate) use aether_data::repository::wallet::{
pub(crate) use aether_data_contracts::repository::billing::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
BillingPlanRecord, BillingPlanWriteInput, PaymentGatewayConfigRecord,
PaymentGatewayConfigWriteInput, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord,
BillingPlanRecord, BillingPlanWriteInput, PaymentGatewayConfigCasWriteInput,
PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate,
UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord,
};
+74 -5
View File
@@ -34,7 +34,6 @@ use super::{
ProviderTransportSnapshotFlight,
};
const DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 120_000;
const MIN_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 1_000;
const MAX_REQUEST_BODY_READ_TIMEOUT_MS: u64 = 600_000;
const REQUEST_BODY_READ_TIMEOUT_MS_ENV: &str = "AETHER_GATEWAY_REQUEST_BODY_READ_TIMEOUT_MS";
@@ -96,7 +95,7 @@ impl std::fmt::Debug for TestExecutionRuntimeSyncOverride {
#[derive(Debug, Clone)]
pub(crate) struct FrontdoorRuntimeGuardConfig {
pub(crate) request_body_read_timeout: Duration,
pub(crate) request_body_read_timeout: Option<Duration>,
pub(crate) request_body_buffer_budget_bytes: usize,
pub(crate) request_body_buffer_budget_permits: usize,
pub(crate) local_execution_planning_timeout: Duration,
@@ -113,9 +112,8 @@ pub(crate) const METRIC_SNAPSHOT_TTL: Duration = Duration::from_secs(2);
impl FrontdoorRuntimeGuardConfig {
pub(crate) fn from_env() -> Self {
Self {
request_body_read_timeout: env_duration_ms(
request_body_read_timeout: optional_env_duration_ms(
REQUEST_BODY_READ_TIMEOUT_MS_ENV,
DEFAULT_REQUEST_BODY_READ_TIMEOUT_MS,
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
),
@@ -148,7 +146,7 @@ impl FrontdoorRuntimeGuardConfig {
#[cfg(test)]
pub(crate) fn for_tests(
request_body_read_timeout: Duration,
request_body_read_timeout: Option<Duration>,
local_execution_planning_timeout: Duration,
) -> Self {
Self {
@@ -189,6 +187,19 @@ fn request_body_buffer_budget_permits_from_env() -> usize {
/ REQUEST_BODY_BUFFER_PERMIT_BYTES
}
fn optional_env_duration_ms(key: &str, min_ms: u64, max_ms: u64) -> Option<Duration> {
let raw = std::env::var(key).ok();
parse_optional_duration_ms(raw.as_deref(), min_ms, max_ms)
}
fn parse_optional_duration_ms(raw: Option<&str>, min_ms: u64, max_ms: u64) -> Option<Duration> {
let parsed = raw?.trim().parse::<u64>().ok()?;
if parsed == 0 {
return None;
}
Some(Duration::from_millis(parsed.clamp(min_ms, max_ms)))
}
fn env_duration_ms(key: &str, default_ms: u64, min_ms: u64, max_ms: u64) -> Duration {
let ms = std::env::var(key)
.ok()
@@ -371,6 +382,7 @@ pub struct AppState {
pub(crate) background_data: Arc<GatewayDataState>,
pub(crate) background_data_isolated: bool,
pub(crate) runtime_state: Arc<RuntimeState>,
pub(crate) internal_gateway_auth: Arc<crate::internal_gateway_auth::InternalGatewayAuthConfig>,
pub(crate) usage_runtime: Arc<usage::UsageRuntime>,
pub(crate) video_tasks: Arc<VideoTaskService>,
pub(crate) video_task_poller: Option<VideoTaskPollerConfig>,
@@ -397,6 +409,8 @@ pub struct AppState {
pub(crate) auth_api_key_feature_settings_cache: Arc<JsonValueCache<AuthApiKeyFeatureCacheKey>>,
pub(crate) auth_daily_quota_availability_cache:
Arc<ValueCache<String, UserDailyQuotaAvailabilityRecord>>,
pub(crate) auth_plan_usage_policy_cache:
Arc<ValueCache<String, crate::plan_usage_policy::EffectivePlanUsagePolicy>>,
pub(crate) auth_wallet_snapshot_cache:
Arc<ValueCache<String, aether_data::repository::wallet::StoredWalletSnapshot>>,
pub(crate) auth_request_cost_upper_bound_cache: Arc<ValueCache<String, f64>>,
@@ -518,6 +532,61 @@ mod tests {
fd_soft_limit: 1_048_576,
};
#[test]
fn request_body_read_timeout_parser_defaults_to_disabled() {
assert_eq!(
parse_optional_duration_ms(
None,
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
),
None
);
}
#[test]
fn request_body_read_timeout_parser_disables_zero_and_invalid_values() {
for value in ["", "invalid", "-1", "0", " 0 "] {
assert_eq!(
parse_optional_duration_ms(
Some(value),
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
),
None,
"{value:?} should disable the optional timeout"
);
}
}
#[test]
fn request_body_read_timeout_parser_clamps_nonzero_values() {
assert_eq!(
parse_optional_duration_ms(
Some("1"),
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
),
Some(Duration::from_millis(MIN_REQUEST_BODY_READ_TIMEOUT_MS))
);
assert_eq!(
parse_optional_duration_ms(
Some("120000"),
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
),
Some(Duration::from_millis(120_000))
);
assert_eq!(
parse_optional_duration_ms(
Some("900000"),
MIN_REQUEST_BODY_READ_TIMEOUT_MS,
MAX_REQUEST_BODY_READ_TIMEOUT_MS,
),
Some(Duration::from_millis(MAX_REQUEST_BODY_READ_TIMEOUT_MS))
);
}
#[test]
fn gate_limit_parser_defaults_to_auto() {
assert_eq!(
@@ -7,13 +7,24 @@ const BOOTSTRAP_ADMIN_EMAIL_ENVS: &[&str] = &["ADMIN_EMAIL"];
const BOOTSTRAP_ADMIN_USERNAME_ENVS: &[&str] = &["ADMIN_USERNAME"];
const BOOTSTRAP_ADMIN_PASSWORD_ENVS: &[&str] = &["ADMIN_PASSWORD"];
#[derive(Debug, Clone, PartialEq, Eq)]
#[derive(Clone, PartialEq, Eq)]
struct BootstrapAdminConfig {
email: Option<String>,
username: String,
password: String,
}
impl std::fmt::Debug for BootstrapAdminConfig {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("BootstrapAdminConfig")
.field("email", &self.email)
.field("username", &self.username)
.field("password", &"[REDACTED]")
.finish()
}
}
impl BootstrapAdminConfig {
fn from_env() -> Result<Option<Self>, GatewayError> {
Self::from_lookup(|key| {
@@ -495,6 +506,14 @@ mod tests {
);
}
#[test]
fn bootstrap_admin_config_debug_output_redacts_password() {
let config = bootstrap_config();
let debug = format!("{config:?}");
assert!(debug.contains("[REDACTED]"));
assert!(!debug.contains("Secret123!"));
}
#[test]
fn bootstrap_admin_config_rejects_partial_env() {
let vars =
+189 -34
View File
@@ -37,20 +37,24 @@ impl AppState {
&self,
active_only: bool,
) -> Result<Vec<provider_catalog::StoredProviderCatalogProvider>, GatewayError> {
self.data
let providers = self
.data
.list_provider_catalog_providers(active_only)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.open_provider_catalog_providers(providers).await
}
pub(crate) async fn list_provider_catalog_endpoints_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<provider_catalog::StoredProviderCatalogEndpoint>, GatewayError> {
self.data
let endpoints = self
.data
.list_provider_catalog_endpoints_by_provider_ids(provider_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.open_provider_catalog_endpoints(endpoints).await
}
pub(crate) async fn list_public_global_models(
@@ -128,6 +132,20 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn update_management_token_for_user(
&self,
record: &aether_data::repository::management_tokens::UpdateManagementTokenRecord,
user_id: &str,
) -> Result<
LocalMutationOutcome<aether_data::repository::management_tokens::StoredManagementToken>,
GatewayError,
> {
self.data
.update_management_token_for_user(record, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn delete_management_token(
&self,
token_id: &str,
@@ -138,6 +156,17 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn delete_management_token_for_user(
&self,
token_id: &str,
user_id: &str,
) -> Result<bool, GatewayError> {
self.data
.delete_management_token_for_user(token_id, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn record_management_token_usage(
&self,
token_id: &str,
@@ -166,6 +195,41 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn set_management_token_active_for_user(
&self,
token_id: &str,
user_id: &str,
is_active: bool,
) -> Result<
Option<aether_data::repository::management_tokens::StoredManagementToken>,
GatewayError,
> {
self.data
.set_management_token_active_for_user(token_id, user_id, is_active)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn activate_management_token_if_matches(
&self,
mutation: &aether_data::repository::management_tokens::ActivateManagementTokenIfMatches,
) -> Result<bool, GatewayError> {
self.data
.activate_management_token_if_matches(mutation)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn delete_inactive_management_token_if_matches(
&self,
mutation: &aether_data::repository::management_tokens::ActivateManagementTokenIfMatches,
) -> Result<bool, GatewayError> {
self.data
.delete_inactive_management_token_if_matches(mutation)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn regenerate_management_token_secret(
&self,
mutation: &aether_data::repository::management_tokens::RegenerateManagementTokenSecret,
@@ -179,6 +243,20 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn regenerate_management_token_secret_for_user(
&self,
mutation: &aether_data::repository::management_tokens::RegenerateManagementTokenSecret,
user_id: &str,
) -> Result<
LocalMutationOutcome<aether_data::repository::management_tokens::StoredManagementToken>,
GatewayError,
> {
self.data
.regenerate_management_token_secret_for_user(mutation, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn get_public_global_model_by_name(
&self,
model_name: &str,
@@ -443,20 +521,24 @@ impl AppState {
&self,
provider_ids: &[String],
) -> Result<Vec<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
self.data
let keys = self
.data
.list_provider_catalog_keys_by_provider_ids(provider_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.open_provider_catalog_keys(keys).await
}
pub(crate) async fn list_provider_catalog_key_summaries_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
self.data
let keys = self
.data
.list_provider_catalog_key_summaries_by_provider_ids(provider_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.open_provider_catalog_keys(keys).await
}
pub(crate) async fn list_provider_catalog_key_maintenance_summaries_by_provider_ids(
@@ -474,30 +556,37 @@ impl AppState {
&self,
key_ids: &[String],
) -> Result<Vec<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
self.data
let keys = self
.data
.list_provider_catalog_keys_by_ids(key_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.open_provider_catalog_keys(keys).await
}
pub(crate) async fn list_provider_catalog_keys_by_ids_strong(
&self,
key_ids: &[String],
) -> Result<Vec<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
self.data
let keys = self
.data
.list_provider_catalog_keys_by_ids_strong(key_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.open_provider_catalog_keys(keys).await
}
pub(crate) async fn list_provider_catalog_key_page(
&self,
query: &provider_catalog::ProviderCatalogKeyListQuery,
) -> Result<provider_catalog::StoredProviderCatalogKeyPage, GatewayError> {
self.data
let mut page = self
.data
.list_provider_catalog_key_page(query)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
page.items = self.open_provider_catalog_keys(page.items).await?;
Ok(page)
}
pub(crate) async fn list_provider_catalog_key_stats_by_provider_ids(
@@ -514,15 +603,19 @@ impl AppState {
&self,
key: &provider_catalog::StoredProviderCatalogKey,
) -> Result<Option<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
let protected = self.protect_provider_catalog_key(key)?;
let created = self
.data
.create_provider_catalog_key(key)
.create_provider_catalog_key(&protected)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if created.is_some() {
self.invalidate_provider_routing_caches();
}
Ok(created)
match created {
Some(key) => self.open_provider_catalog_key(key).await.map(Some),
None => Ok(None),
}
}
pub(crate) async fn create_provider_catalog_provider(
@@ -530,29 +623,58 @@ impl AppState {
provider: &provider_catalog::StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
) -> Result<Option<provider_catalog::StoredProviderCatalogProvider>, GatewayError> {
let protected = self.protect_provider_catalog_provider(provider)?;
let created = self
.data
.create_provider_catalog_provider(provider, shift_existing_priorities_from)
.create_provider_catalog_provider(&protected, shift_existing_priorities_from)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if created.is_some() {
self.invalidate_provider_routing_caches();
}
Ok(created)
match created {
Some(provider) => self
.open_provider_catalog_provider(provider)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn update_provider_catalog_provider(
&self,
provider: &provider_catalog::StoredProviderCatalogProvider,
) -> Result<Option<provider_catalog::StoredProviderCatalogProvider>, GatewayError> {
let protected = self.protect_provider_catalog_provider(provider)?;
let updated = self
.data
.update_provider_catalog_provider(provider)
.update_provider_catalog_provider(&protected)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
self.invalidate_provider_routing_caches();
}
match updated {
Some(provider) => self
.open_provider_catalog_provider(provider)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn compare_and_swap_provider_catalog_provider_config(
&self,
update: &provider_catalog::ProviderCatalogProviderConfigCasUpdate,
) -> Result<bool, GatewayError> {
let updated = self
.data
.compare_and_swap_provider_catalog_provider_config(update)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
self.invalidate_provider_routing_caches();
}
Ok(updated)
}
@@ -616,30 +738,44 @@ impl AppState {
&self,
endpoint: &provider_catalog::StoredProviderCatalogEndpoint,
) -> Result<Option<provider_catalog::StoredProviderCatalogEndpoint>, GatewayError> {
let protected = self.protect_provider_catalog_endpoint(endpoint)?;
let created = self
.data
.create_provider_catalog_endpoint(endpoint)
.create_provider_catalog_endpoint(&protected)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if created.is_some() {
self.invalidate_provider_routing_caches();
}
Ok(created)
match created {
Some(endpoint) => self
.open_provider_catalog_endpoint(endpoint)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn update_provider_catalog_endpoint(
&self,
endpoint: &provider_catalog::StoredProviderCatalogEndpoint,
) -> Result<Option<provider_catalog::StoredProviderCatalogEndpoint>, GatewayError> {
let protected = self.protect_provider_catalog_endpoint(endpoint)?;
let updated = self
.data
.update_provider_catalog_endpoint(endpoint)
.update_provider_catalog_endpoint(&protected)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
self.invalidate_provider_routing_caches();
}
Ok(updated)
match updated {
Some(endpoint) => self
.open_provider_catalog_endpoint(endpoint)
.await
.map(Some),
None => Ok(None),
}
}
pub(crate) async fn delete_provider_catalog_endpoint(
@@ -661,24 +797,30 @@ impl AppState {
&self,
key: &provider_catalog::StoredProviderCatalogKey,
) -> Result<Option<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
let protected = self.protect_provider_catalog_key(key)?;
let updated = self
.data
.update_provider_catalog_key(key)
.update_provider_catalog_key(&protected)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.is_some() {
self.invalidate_provider_routing_caches();
}
Ok(updated)
match updated {
Some(key) => self.open_provider_catalog_key(key).await.map(Some),
None => Ok(None),
}
}
pub(crate) async fn compare_and_update_provider_catalog_key_admin_state(
&self,
update: &provider_catalog::ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, GatewayError> {
let mut protected = update.clone();
protected.key = self.protect_provider_catalog_key(&update.key)?;
let updated = self
.data
.compare_and_update_provider_catalog_key_admin_state(update)
.compare_and_update_provider_catalog_key_admin_state(&protected)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
// A conflict means another instance changed credentials. Invalidate on
@@ -691,15 +833,22 @@ impl AppState {
&self,
keys: &[provider_catalog::StoredProviderCatalogKey],
) -> Result<Option<Vec<provider_catalog::StoredProviderCatalogKey>>, GatewayError> {
let protected = keys
.iter()
.map(|key| self.protect_provider_catalog_key(key))
.collect::<Result<Vec<_>, _>>()?;
let updated = self
.data
.update_provider_catalog_keys(keys)
.update_provider_catalog_keys(&protected)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated.as_ref().is_some_and(|keys| !keys.is_empty()) {
self.invalidate_provider_routing_caches();
}
Ok(updated)
match updated {
Some(keys) => self.open_provider_catalog_keys(keys).await.map(Some),
None => Ok(None),
}
}
pub(crate) async fn compare_and_update_provider_catalog_key_adaptive_state(
@@ -1041,30 +1190,36 @@ impl AppState {
&self,
provider_ids: &[String],
) -> Result<Vec<provider_catalog::StoredProviderCatalogProvider>, GatewayError> {
self.data
let providers = self
.data
.list_provider_catalog_providers_by_ids(provider_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.open_provider_catalog_providers(providers).await
}
pub(crate) async fn read_provider_catalog_endpoints_by_ids(
&self,
endpoint_ids: &[String],
) -> Result<Vec<provider_catalog::StoredProviderCatalogEndpoint>, GatewayError> {
self.data
let endpoints = self
.data
.list_provider_catalog_endpoints_by_ids(endpoint_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.open_provider_catalog_endpoints(endpoints).await
}
pub(crate) async fn read_provider_catalog_keys_by_ids(
&self,
key_ids: &[String],
) -> Result<Vec<provider_catalog::StoredProviderCatalogKey>, GatewayError> {
self.data
let keys = self
.data
.list_provider_catalog_keys_by_ids(key_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.open_provider_catalog_keys(keys).await
}
pub(crate) async fn update_provider_catalog_key_format_health(
@@ -0,0 +1,395 @@
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyCredentialsCasUpdate, StoredProviderCatalogKey,
};
use super::AppState;
use crate::handlers::shared::{
open_provider_catalog_credential, seal_provider_catalog_credential,
ProviderCatalogCredentialField, ProviderCatalogCredentialProjection,
};
use crate::GatewayError;
impl AppState {
pub(super) fn protect_provider_catalog_key_credentials(
&self,
key: &StoredProviderCatalogKey,
) -> Result<StoredProviderCatalogKey, GatewayError> {
let mut protected = key.clone();
protected.encrypted_api_key = self
.project_provider_catalog_key_credential(
key,
ProviderCatalogCredentialField::ApiKey,
key.encrypted_api_key.as_deref(),
)?
.map(|projection| projection.protected);
protected.encrypted_auth_config = self
.project_provider_catalog_key_credential(
key,
ProviderCatalogCredentialField::AuthConfig,
key.encrypted_auth_config.as_deref(),
)?
.map(|projection| projection.protected);
Ok(protected)
}
pub(super) async fn open_provider_catalog_key_credentials_once(
&self,
key: &mut StoredProviderCatalogKey,
) -> Result<bool, GatewayError> {
let observed_api_key = key.encrypted_api_key.clone();
let observed_auth_config = key.encrypted_auth_config.clone();
let api_key = self.project_provider_catalog_key_credential(
key,
ProviderCatalogCredentialField::ApiKey,
observed_api_key.as_deref(),
)?;
let auth_config = self.project_provider_catalog_key_credential(
key,
ProviderCatalogCredentialField::AuthConfig,
observed_auth_config.as_deref(),
)?;
let migration_required = api_key
.as_ref()
.is_some_and(|projection| projection.migration_required)
|| auth_config
.as_ref()
.is_some_and(|projection| projection.migration_required);
if !migration_required {
return Ok(true);
}
if !self.has_provider_catalog_data_writer() {
return Err(provider_catalog_credential_error(
"stored provider catalog credentials require migration but the catalog writer is unavailable",
));
}
let protected_api_key = api_key.map(|projection| projection.protected);
let protected_auth_config = auth_config.map(|projection| projection.protected);
let updated = self
.data
.compare_and_swap_provider_catalog_key_credentials(
&ProviderCatalogKeyCredentialsCasUpdate {
key_id: key.id.clone(),
expected_provider_id: key.provider_id.clone(),
expected_encrypted_api_key: observed_api_key,
expected_encrypted_auth_config: observed_auth_config,
encrypted_api_key: protected_api_key.clone(),
encrypted_auth_config: protected_auth_config.clone(),
},
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if updated {
key.encrypted_api_key = protected_api_key;
key.encrypted_auth_config = protected_auth_config;
}
Ok(updated)
}
pub(crate) fn decrypt_provider_catalog_key_api_key(
&self,
key: &StoredProviderCatalogKey,
) -> Result<Option<String>, GatewayError> {
self.project_provider_catalog_key_credential(
key,
ProviderCatalogCredentialField::ApiKey,
key.encrypted_api_key.as_deref(),
)
.map(|projection| projection.map(|projection| projection.plaintext))
}
pub(crate) fn decrypt_provider_catalog_key_auth_config(
&self,
key: &StoredProviderCatalogKey,
) -> Result<Option<String>, GatewayError> {
self.project_provider_catalog_key_credential(
key,
ProviderCatalogCredentialField::AuthConfig,
key.encrypted_auth_config.as_deref(),
)
.map(|projection| projection.map(|projection| projection.plaintext))
}
pub(crate) fn seal_provider_catalog_key_api_key(
&self,
provider_id: &str,
key_id: &str,
plaintext: &str,
) -> Result<String, GatewayError> {
seal_provider_catalog_credential(
self,
provider_id,
key_id,
ProviderCatalogCredentialField::ApiKey,
plaintext,
)
.map_err(provider_catalog_credential_error)
}
pub(crate) fn seal_provider_catalog_key_auth_config(
&self,
provider_id: &str,
key_id: &str,
plaintext: &str,
) -> Result<String, GatewayError> {
seal_provider_catalog_credential(
self,
provider_id,
key_id,
ProviderCatalogCredentialField::AuthConfig,
plaintext,
)
.map_err(provider_catalog_credential_error)
}
pub(super) fn validate_protected_provider_catalog_key_api_key(
&self,
provider_id: &str,
key_id: &str,
stored: &str,
) -> Result<(), GatewayError> {
self.validate_protected_provider_catalog_key_credential(
provider_id,
key_id,
ProviderCatalogCredentialField::ApiKey,
stored,
)
}
pub(super) fn validate_protected_provider_catalog_key_auth_config(
&self,
provider_id: &str,
key_id: &str,
stored: &str,
) -> Result<(), GatewayError> {
self.validate_protected_provider_catalog_key_credential(
provider_id,
key_id,
ProviderCatalogCredentialField::AuthConfig,
stored,
)
}
fn validate_protected_provider_catalog_key_credential(
&self,
provider_id: &str,
key_id: &str,
field: ProviderCatalogCredentialField,
stored: &str,
) -> Result<(), GatewayError> {
let projection = open_provider_catalog_credential(self, provider_id, key_id, field, stored)
.map_err(provider_catalog_credential_error)?;
if projection.migration_required || projection.protected != stored {
return Err(provider_catalog_credential_error(
"provider catalog credential write requires a bound v2 ciphertext",
));
}
Ok(())
}
fn project_provider_catalog_key_credential(
&self,
key: &StoredProviderCatalogKey,
field: ProviderCatalogCredentialField,
stored: Option<&str>,
) -> Result<Option<ProviderCatalogCredentialProjection>, GatewayError> {
let Some(stored) = stored else {
return Ok(None);
};
if stored.is_empty() {
return Err(provider_catalog_credential_error(
"stored provider catalog credential is empty",
));
}
open_provider_catalog_credential(self, &key.provider_id, &key.id, field, stored)
.map(Some)
.map_err(provider_catalog_credential_error)
}
}
fn provider_catalog_credential_error(message: &'static str) -> GatewayError {
GatewayError::Internal(message.to_string())
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogReadRepository,
ProviderCatalogWriteRepository, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use crate::{data::GatewayDataState, AppState};
fn sample_provider(id: &str) -> StoredProviderCatalogProvider {
StoredProviderCatalogProvider::new(
id.to_string(),
format!("Provider {id}"),
Some("https://example.test".to_string()),
"openai".to_string(),
)
.expect("provider should build")
}
fn sample_key(
id: &str,
provider_id: &str,
encrypted_api_key: Option<String>,
encrypted_auth_config: Option<String>,
) -> StoredProviderCatalogKey {
StoredProviderCatalogKey::new(
id.to_string(),
provider_id.to_string(),
format!("Key {id}"),
"oauth".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
None,
encrypted_api_key,
encrypted_auth_config,
None,
None,
None,
None,
None,
None,
)
.expect("key transport should build")
}
fn state_with_repository(repository: Arc<InMemoryProviderCatalogReadRepository>) -> AppState {
AppState::new()
.expect("test state should build")
.with_data_state_for_tests(
GatewayDataState::with_provider_catalog_repository_for_tests(repository)
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY),
)
}
#[tokio::test]
async fn app_state_migrates_both_legacy_fields_with_one_exact_cas() {
let legacy_api =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "legacy-api-key")
.expect("legacy API key should encrypt");
let legacy_auth = encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"refresh_token":"legacy-refresh"}"#,
)
.expect("legacy auth config should encrypt");
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-1")],
Vec::new(),
vec![sample_key(
"key-1",
"provider-1",
Some(legacy_api),
Some(legacy_auth),
)],
));
let state = state_with_repository(Arc::clone(&repository));
let opened = state
.list_provider_catalog_keys_by_ids(&["key-1".to_string()])
.await
.expect("legacy key should migrate")
.into_iter()
.next()
.expect("key should exist");
assert_eq!(
state
.decrypt_provider_catalog_key_api_key(&opened)
.expect("API key should open")
.as_deref(),
Some("legacy-api-key")
);
assert_eq!(
state
.decrypt_provider_catalog_key_auth_config(&opened)
.expect("auth config should open")
.as_deref(),
Some(r#"{"refresh_token":"legacy-refresh"}"#)
);
let stored = repository
.list_keys_by_ids(&["key-1".to_string()])
.await
.expect("stored key should read")
.into_iter()
.next()
.expect("stored key should exist");
assert!(stored
.encrypted_api_key
.as_deref()
.is_some_and(|value| value.starts_with("aether-provider-catalog-credential-v2:")));
assert!(stored
.encrypted_auth_config
.as_deref()
.is_some_and(|value| value.starts_with("aether-provider-catalog-credential-v2:")));
}
#[tokio::test]
async fn app_state_rejects_ciphertext_copied_to_another_key() {
let empty_repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-1")],
Vec::new(),
Vec::new(),
));
let bootstrap = state_with_repository(Arc::clone(&empty_repository));
let copied = bootstrap
.seal_provider_catalog_key_api_key("provider-1", "key-1", "secret")
.expect("credential should seal");
let repository = Arc::new(InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-1")],
Vec::new(),
vec![sample_key("key-2", "provider-1", Some(copied), None)],
));
let state = state_with_repository(repository);
assert!(state
.list_provider_catalog_keys_by_ids(&["key-2".to_string()])
.await
.is_err());
}
#[tokio::test]
async fn credential_cas_fences_provider_and_both_ciphertexts() {
let repository = InMemoryProviderCatalogReadRepository::seed(
vec![sample_provider("provider-1"), sample_provider("provider-2")],
Vec::new(),
vec![sample_key(
"key-1",
"provider-2",
Some("api-before".to_string()),
Some("auth-before".to_string()),
)],
);
let update = ProviderCatalogKeyCredentialsCasUpdate {
key_id: "key-1".to_string(),
expected_provider_id: "provider-1".to_string(),
expected_encrypted_api_key: Some("api-before".to_string()),
expected_encrypted_auth_config: Some("auth-before".to_string()),
encrypted_api_key: Some("api-after".to_string()),
encrypted_auth_config: Some("auth-after".to_string()),
};
assert!(!repository
.compare_and_swap_key_credentials(&update)
.await
.expect("provider-fenced CAS should execute"));
let stored = repository
.list_keys_by_ids(&["key-1".to_string()])
.await
.expect("key should read")
.into_iter()
.next()
.expect("key should exist");
assert_eq!(stored.encrypted_api_key.as_deref(), Some("api-before"));
assert_eq!(stored.encrypted_auth_config.as_deref(), Some("auth-before"));
}
}
File diff suppressed because it is too large Load Diff
+127 -23
View File
@@ -13,7 +13,7 @@ use aether_data::repository::proxy_nodes::{
use aether_data_contracts::repository::usage::{
UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot,
};
use aether_http::{build_http_client, HttpClientConfig};
use aether_http::{apply_http_client_config, HttpClientConfig};
use aether_runtime::{
service_up_sample, AdmissionPermit, ConcurrencyGate, ConcurrencySnapshot, MetricKind,
MetricLabel, MetricSample,
@@ -104,6 +104,24 @@ const USAGE_COUNTER_EXACT_HEALTH_METRICS_TTL: Duration = Duration::from_secs(5 *
const USAGE_COUNTER_EXACT_HEALTH_METRICS_MAX_STALENESS: Duration = Duration::from_secs(10 * 60);
const USAGE_COUNTER_EXACT_HEALTH_METRICS_RETRY_BACKOFF: Duration = Duration::from_secs(5);
const ADMIN_USAGE_AGGREGATE_INVALID_INPUT_DETAIL: &str =
"usage aggregate import payload is invalid";
fn admin_usage_aggregate_import_error(detail: String) -> GatewayError {
// Import validation errors can include source row IDs, table names, and
// adapter-specific details. Keep those details in process memory only;
// the public/admin response receives a stable client-safe message.
warn!(
event_name = "admin_usage_aggregate_import_rejected",
error_length = detail.len(),
"usage aggregate import input was rejected"
);
GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
message: ADMIN_USAGE_AGGREGATE_INVALID_INPUT_DETAIL.to_string(),
}
}
fn system_config_key_affects_scheduler(key: &str) -> bool {
let key = key.trim();
SCHEDULER_AFFECTING_SYSTEM_CONFIG_KEYS.contains(&key)
@@ -287,18 +305,38 @@ impl AppState {
GatewayDataState::disabled()
.with_usage_worker_queue(Self::usage_worker_queue_for(&runtime_state)),
);
let client = build_http_client(&HttpClientConfig {
connect_timeout_ms: Some(10_000),
request_timeout_ms: Some(300_000),
http2_adaptive_window: true,
..HttpClientConfig::default()
})?;
let owner_forward_client = build_http_client(&HttpClientConfig {
connect_timeout_ms: Some(10_000),
http2_adaptive_window: true,
..HttpClientConfig::default()
})?;
let client = apply_http_client_config(
reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none()),
&HttpClientConfig {
connect_timeout_ms: Some(10_000),
request_timeout_ms: Some(300_000),
http2_adaptive_window: true,
..HttpClientConfig::default()
},
)
.build()?;
let owner_forward_client = apply_http_client_config(
reqwest::Client::builder()
.no_proxy()
.redirect(reqwest::redirect::Policy::none()),
&HttpClientConfig {
connect_timeout_ms: Some(10_000),
http2_adaptive_window: true,
..HttpClientConfig::default()
},
)
.build()?;
let frontdoor_runtime_guards = Arc::new(FrontdoorRuntimeGuardConfig::from_env());
let internal_gateway_auth =
Arc::new(crate::internal_gateway_auth::InternalGatewayAuthConfig::for_process());
if internal_gateway_auth.status() == "misconfigured" {
warn!(
environment_variable = crate::internal_gateway_auth::INTERNAL_GATEWAY_AUTH_SECRET_ENV,
"internal gateway control plane is fail-closed because its authentication secret is invalid"
);
}
Ok(Self {
#[cfg(test)]
execution_runtime_override_base_url: execution_runtime_override_base_url
@@ -310,6 +348,7 @@ impl AppState {
background_data: Arc::clone(&data),
background_data_isolated: false,
runtime_state: runtime_state.clone(),
internal_gateway_auth,
usage_runtime: Arc::new(usage::UsageRuntime::disabled()),
video_tasks: Arc::new(VideoTaskService::new(
VideoTaskTruthSourceMode::PythonSyncReport,
@@ -349,6 +388,7 @@ impl AppState {
auth_api_key_force_capabilities_cache: Arc::new(JsonValueCache::default()),
auth_api_key_feature_settings_cache: Arc::new(JsonValueCache::default()),
auth_daily_quota_availability_cache: Arc::new(ValueCache::default()),
auth_plan_usage_policy_cache: Arc::new(ValueCache::default()),
auth_wallet_snapshot_cache: Arc::new(ValueCache::default()),
auth_request_cost_upper_bound_cache: Arc::new(ValueCache::default()),
provider_quota_snapshot_cache: Arc::new(ValueCache::default()),
@@ -451,6 +491,10 @@ impl AppState {
true
}
pub(crate) fn internal_gateway_auth_status(&self) -> &'static str {
self.internal_gateway_auth.status()
}
#[cfg(test)]
pub(crate) fn execution_runtime_override_base_url(&self) -> Option<&str> {
self.execution_runtime_override_base_url.as_deref()
@@ -755,6 +799,20 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn compare_and_set_system_config_string_value(
&self,
key: &str,
expected: &str,
replacement: &str,
) -> Result<bool, GatewayError> {
let result = self
.data
.compare_and_set_system_config_string_value(key, expected, replacement)
.await;
self.system_config_cache.invalidate(key);
result.map_err(|err| GatewayError::Internal(err.to_string()))
}
async fn read_system_config_json_value_with_cache_windows(
&self,
key: &str,
@@ -945,6 +1003,7 @@ impl AppState {
self.auth_api_key_force_capabilities_cache.clear();
self.auth_api_key_feature_settings_cache.clear();
self.auth_daily_quota_availability_cache.clear();
self.auth_plan_usage_policy_cache.clear();
self.auth_wallet_snapshot_cache.clear();
self.auth_request_cost_upper_bound_cache.clear();
self.provider_quota_snapshot_cache.clear();
@@ -1030,10 +1089,9 @@ impl AppState {
.import_admin_system_usage_aggregates(snapshot, user_id_map, api_key_id_map, mode)
.await
.map_err(|err| match err {
aether_data::DataLayerError::InvalidInput(detail) => GatewayError::Client {
status: http::StatusCode::BAD_REQUEST,
message: detail,
},
aether_data::DataLayerError::InvalidInput(detail) => {
admin_usage_aggregate_import_error(detail)
}
other => GatewayError::Internal(other.to_string()),
})
}
@@ -1124,18 +1182,27 @@ impl AppState {
&self,
mutation: &aether_data::repository::proxy_nodes::ProxyNodeRegistrationMutation,
) -> Result<Option<StoredProxyNode>, GatewayError> {
self.data
.register_proxy_node(mutation)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
self.register_proxy_node_with_bound_secrets(mutation).await
}
pub(crate) async fn create_manual_proxy_node(
&self,
mutation: &ProxyNodeManualCreateMutation,
) -> Result<Option<StoredProxyNode>, GatewayError> {
let mut protected = mutation.clone();
let node_id = mutation
.node_id
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned)
.unwrap_or_else(|| uuid::Uuid::new_v4().to_string());
protected.node_id = Some(node_id.clone());
if let Some(password) = mutation.proxy_password.as_deref() {
protected.proxy_password = Some(self.protect_proxy_node_password(&node_id, password)?);
}
self.data
.create_manual_proxy_node(mutation)
.create_manual_proxy_node(&protected)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
@@ -1144,8 +1211,17 @@ impl AppState {
&self,
mutation: &ProxyNodeManualUpdateMutation,
) -> Result<Option<StoredProxyNode>, GatewayError> {
let Some(existing) = self.find_proxy_node(&mutation.node_id).await? else {
return Ok(None);
};
let mut protected = mutation.clone();
protected.node_id = existing.id.clone();
if let Some(password) = mutation.proxy_password.as_deref() {
protected.proxy_password =
Some(self.protect_proxy_node_password(&existing.id, password)?);
}
self.data
.update_manual_proxy_node(mutation)
.update_manual_proxy_node(&protected)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
@@ -1178,6 +1254,9 @@ impl AppState {
&self,
mutation: &ProxyNodeHeartbeatMutation,
) -> Result<Option<StoredProxyNode>, GatewayError> {
crate::state::decrypt_or_migrate_proxy_tunnel_psk(&self.data, &mutation.node_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
self.data
.apply_proxy_node_heartbeat(mutation)
.await
@@ -2015,9 +2094,16 @@ impl AppState {
mut self,
path: impl Into<std::path::PathBuf>,
) -> std::io::Result<Self> {
let encryption_key = self.data.encryption_key().ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"AETHER_GATEWAY_VIDEO_TASK_STORE_PATH requires a configured encryption key",
)
})?;
self.video_tasks = Arc::new(VideoTaskService::with_file_store(
self.video_tasks.truth_source_mode(),
path,
encryption_key,
)?);
Ok(self)
}
@@ -3843,7 +3929,8 @@ mod tests {
use serde_json::json;
use super::{
database_bounded_auth_load_limit, merge_usage_counter_health_snapshots,
admin_usage_aggregate_import_error, database_bounded_auth_load_limit,
merge_usage_counter_health_snapshots,
usage_counter_pending_health_metric_samples_with_timeout,
usage_queue_health_metric_samples_with_timeout, usage_runtime_metric_samples, AppState,
MetricKind, MetricSample, METRIC_SNAPSHOT_TTL,
@@ -3852,6 +3939,23 @@ mod tests {
use crate::cache::SchedulerAffinityTarget;
use crate::data::{GatewayDataConfig, GatewayDataState};
#[test]
fn usage_aggregate_import_invalid_input_is_projected_to_a_safe_client_message() {
let error = admin_usage_aggregate_import_error(
"postgres table stats_user_daily row secret-user contains column password".to_string(),
);
match error {
super::GatewayError::Client { status, message } => {
assert_eq!(status, http::StatusCode::BAD_REQUEST);
assert_eq!(message, "usage aggregate import payload is invalid");
assert!(!message.contains("secret-user"));
assert!(!message.contains("password"));
}
other => panic!("expected client-safe import error, got {other:?}"),
}
}
#[test]
fn auth_load_gate_reserves_half_of_foreground_database_pool() {
assert_eq!(
+47 -27
View File
@@ -13,6 +13,7 @@ use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
use aether_data_contracts::repository::quota::StoredProviderQuotaSnapshot;
use aether_data_contracts::DataLayerError;
use aether_model_fetch::{
aggregate_models_for_cache, build_antigravity_load_code_assist_plan,
fetch_models_from_transports, merge_upstream_metadata, model_fetch_interval_minutes,
@@ -25,7 +26,7 @@ use tracing::{debug, warn};
use super::{AppState, GatewayError};
use crate::clock::current_unix_secs;
use crate::model_fetch::{CodexCatalogRuntime, ModelFetchRuntimeState};
use crate::model_fetch::{safe_model_fetch_error, CodexCatalogRuntime, ModelFetchRuntimeState};
use crate::provider_transport::{GatewayProviderTransportSnapshot, LocalResolvedOAuthRequestAuth};
use crate::request_candidate_runtime::{
RequestCandidateRuntimeCapabilityReader, RequestCandidateRuntimeReader,
@@ -36,6 +37,44 @@ use crate::{execution_runtime, provider_transport};
const MODEL_FETCH_RESPONSE_BODY_LIMIT_BYTES: usize = 8 * 1024 * 1024;
#[async_trait]
impl provider_transport::ProviderTransportSnapshotSource for AppState {
fn encryption_key(&self) -> Option<&str> {
AppState::encryption_key(self)
}
async fn list_provider_catalog_providers_by_ids(
&self,
ids: &[String],
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
self.read_provider_catalog_providers_by_ids(ids)
.await
.map_err(provider_transport_snapshot_data_error)
}
async fn list_provider_catalog_endpoints_by_ids(
&self,
ids: &[String],
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
self.read_provider_catalog_endpoints_by_ids(ids)
.await
.map_err(provider_transport_snapshot_data_error)
}
async fn list_provider_catalog_keys_by_ids(
&self,
ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
self.read_provider_catalog_keys_by_ids(ids)
.await
.map_err(provider_transport_snapshot_data_error)
}
}
fn provider_transport_snapshot_data_error(error: GatewayError) -> DataLayerError {
DataLayerError::UnexpectedValue(error.into_message())
}
impl AppState {
pub(crate) async fn hydrate_antigravity_project_metadata_for_transport(
&self,
@@ -159,7 +198,7 @@ impl AppState {
provider_id = %transport.provider.id,
endpoint_id = %transport.endpoint.id,
key_id = %transport.key.id,
error = %err,
error = %safe_model_fetch_error(&err),
"gemini_cli project metadata hydration failed"
);
return None;
@@ -324,31 +363,12 @@ impl CodexCatalogRuntime for AppState {
return Ok(Some(scope));
}
let decrypted_auth_config = match key.encrypted_auth_config.as_deref() {
Some(ciphertext) => Some(
crate::handlers::shared::decrypt_catalog_secret_with_fallbacks(
self.encryption_key(),
ciphertext,
)
.ok_or_else(|| {
"Codex catalog auth config could not be verified for credential fencing"
.to_string()
})?,
),
None => None,
};
let decrypted_api_key = match key.encrypted_api_key.as_deref() {
Some(ciphertext) => Some(
crate::handlers::shared::decrypt_catalog_secret_with_fallbacks(
self.encryption_key(),
ciphertext,
)
.ok_or_else(|| {
"Codex catalog API key could not be verified for credential fencing".to_string()
})?,
),
None => None,
};
let decrypted_auth_config = self
.decrypt_provider_catalog_key_auth_config(&key)
.map_err(GatewayError::into_message)?;
let decrypted_api_key = self
.decrypt_provider_catalog_key_api_key(&key)
.map_err(GatewayError::into_message)?;
Ok(
crate::model_fetch::codex_catalog_credential_scope_from_stored_key(
+9 -1
View File
@@ -6,6 +6,8 @@ mod app;
mod bootstrap_admin;
mod cache;
mod catalog;
mod catalog_credentials;
mod catalog_proxy;
mod core;
mod cors;
mod integrations;
@@ -23,7 +25,8 @@ pub(crate) use self::admin_types::{
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
AdminPaymentCallbackRecord, AdminSecurityBlacklistEntry, AdminWalletPaymentOrderRecord,
AdminWalletRefundRecord, AdminWalletTransactionRecord, BillingPlanRecord,
BillingPlanWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput,
BillingPlanWriteInput, PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord,
PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate,
UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord,
};
pub use self::app::AppState;
@@ -42,10 +45,15 @@ pub(crate) use self::oauth::{
provider_transport_context_allows_credential_rotation, AgentIdentityAuthConfigFence,
CodexRuntimeOAuthObservation, ProviderTransportCredentialFence,
};
pub(crate) use self::proxy::{
decrypt_or_migrate_proxy_tunnel_psk, decrypt_or_migrate_proxy_tunnel_psk_binding,
unavailable_proxy_snapshot,
};
pub(crate) use self::types::{
AdminWalletMutationOutcome, GatewayAdminPaymentCallbackView, GatewayUserPreferenceView,
GatewayUserSessionView, LocalExecutionRuntimeMissDiagnostic, LocalMutationOutcome,
LocalProviderDeleteTaskState,
};
pub(crate) use self::video::VideoTaskRouteAccess;
use super::provider_transport::provider_transport_snapshot_looks_refreshed;
pub(crate) use super::provider_transport::ProviderTransportSnapshotCacheKey;
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -234,6 +234,23 @@ impl AppState {
record: aether_data::repository::auth::CreateUserApiKeyRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
#[cfg(test)]
{
// Unit-test AppState instances keep users and API keys in separate in-memory
// repositories. Bridge them with the authoritative user record only; an unknown,
// inactive, or deleted user must never be synthesized from the key request.
let Some(user) = self.find_user_auth_by_id(&record.user_id).await? else {
return Ok(None);
};
if !user.is_active || user.is_deleted {
return Ok(None);
}
self.data
.synchronize_user_api_key_owner_for_tests(&user)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
}
let api_key = self
.data
.create_user_api_key(record)
@@ -277,6 +294,32 @@ impl AppState {
Ok(api_key)
}
pub(crate) async fn compare_and_swap_api_key_ciphertext(
&self,
mutation: &aether_data::repository::auth::CompareAndSwapAuthApiKeyCiphertext,
) -> Result<bool, GatewayError> {
self.data
.compare_and_swap_api_key_ciphertext(mutation)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn update_user_api_key_basic_if_unlocked(
&self,
record: aether_data::repository::auth::UpdateUserApiKeyBasicRecord,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
let api_key = self
.data
.update_user_api_key_basic_if_unlocked(record)
.await
.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(
&self,
record: aether_data::repository::auth::UpdateStandaloneApiKeyBasicRecord,
@@ -293,6 +336,22 @@ impl AppState {
Ok(api_key)
}
pub(crate) async fn restore_api_key_if_matches(
&self,
expected: &aether_data::repository::auth::StoredAuthApiKeyExportRecord,
restored: &aether_data::repository::auth::StoredAuthApiKeyExportRecord,
) -> Result<bool, GatewayError> {
let restored = self
.data
.restore_api_key_if_matches(expected, restored)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if restored {
self.invalidate_auth_context_cache();
}
Ok(restored)
}
pub(crate) async fn set_user_api_key_active(
&self,
user_id: &str,
@@ -311,6 +370,24 @@ impl AppState {
Ok(api_key)
}
pub(crate) async fn set_user_api_key_active_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
is_active: bool,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
let api_key = self
.data
.set_user_api_key_active_if_unlocked(user_id, api_key_id, is_active)
.await
.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(
&self,
api_key_id: &str,
@@ -363,6 +440,24 @@ impl AppState {
Ok(api_key)
}
pub(crate) async fn set_user_api_key_allowed_providers_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
allowed_providers: Option<Vec<String>>,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
let api_key = self
.data
.set_user_api_key_allowed_providers_if_unlocked(user_id, api_key_id, allowed_providers)
.await
.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(
&self,
user_id: &str,
@@ -381,6 +476,28 @@ impl AppState {
Ok(api_key)
}
pub(crate) async fn set_user_api_key_force_capabilities_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
force_capabilities: Option<serde_json::Value>,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
let api_key = self
.data
.set_user_api_key_force_capabilities_if_unlocked(
user_id,
api_key_id,
force_capabilities,
)
.await
.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_feature_settings(
&self,
user_id: &str,
@@ -399,6 +516,24 @@ impl AppState {
Ok(api_key)
}
pub(crate) async fn set_user_api_key_feature_settings_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
feature_settings: Option<serde_json::Value>,
) -> Result<Option<aether_data::repository::auth::StoredAuthApiKeyExportRecord>, GatewayError>
{
let api_key = self
.data
.set_user_api_key_feature_settings_if_unlocked(user_id, api_key_id, feature_settings)
.await
.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_api_key_usage_totals(
&self,
api_key_id: &str,
@@ -451,6 +586,22 @@ impl AppState {
Ok(deleted)
}
pub(crate) async fn delete_user_api_key_if_unlocked(
&self,
user_id: &str,
api_key_id: &str,
) -> Result<bool, GatewayError> {
let deleted = self
.data
.delete_user_api_key_if_unlocked(user_id, api_key_id)
.await
.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,
@@ -466,3 +617,190 @@ impl AppState {
Ok(deleted)
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use aether_data::repository::auth::{
AuthApiKeyLookupKey, AuthApiKeyReadRepository, CreateUserApiKeyRecord,
InMemoryAuthApiKeySnapshotRepository, StoredAuthApiKeySnapshot,
};
use aether_data::repository::users::StoredUserAuthRecord;
use crate::data::GatewayDataState;
use crate::AppState;
fn authoritative_user(
user_id: &str,
is_active: bool,
is_deleted: bool,
) -> StoredUserAuthRecord {
StoredUserAuthRecord::new(
user_id.to_string(),
Some(format!("{user_id}@example.com")),
true,
format!("owner-{user_id}"),
Some("server-managed-password-hash".to_string()),
"admin".to_string(),
"oauth".to_string(),
Some(serde_json::json!(["openai"])),
Some(serde_json::json!(["openai:chat"])),
Some(serde_json::json!(["gpt-5"])),
is_active,
is_deleted,
None,
None,
)
.expect("authoritative user should build")
.with_security_version(41)
.expect("security version should be valid")
}
fn create_record(user_id: &str, api_key_id: &str) -> CreateUserApiKeyRecord {
CreateUserApiKeyRecord {
user_id: user_id.to_string(),
api_key_id: api_key_id.to_string(),
key_hash: format!("hash-{api_key_id}"),
key_encrypted: Some(format!("encrypted-{api_key_id}")),
name: Some("first key".to_string()),
allowed_providers: Some(vec!["anthropic".to_string()]),
allowed_api_formats: None,
allowed_models: None,
ip_rules: None,
rate_limit: 0,
concurrent_limit: None,
force_capabilities: None,
feature_settings: None,
is_active: true,
expires_at_unix_secs: None,
auto_delete_on_expiry: false,
total_requests: 0,
total_tokens: 0,
total_cost_usd: 0.0,
}
}
fn state_with_users<I>(
repository: Arc<InMemoryAuthApiKeySnapshotRepository>,
users: I,
) -> AppState
where
I: IntoIterator<Item = StoredUserAuthRecord>,
{
AppState::new()
.expect("gateway state should build")
.with_data_state_for_tests(GatewayDataState::with_auth_api_key_repository_for_tests(
repository,
))
.with_auth_users_for_tests(users)
}
#[tokio::test]
async fn test_gateway_first_user_key_requires_authoritative_active_owner() {
let unknown_repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default());
let unknown_state = state_with_users(
Arc::clone(&unknown_repository),
Vec::<StoredUserAuthRecord>::new(),
);
assert!(unknown_state
.create_user_api_key(create_record("missing-user", "missing-key"))
.await
.expect("unknown owner creation should resolve")
.is_none());
assert!(unknown_repository
.find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("missing-key"))
.await
.expect("unknown key lookup should resolve")
.is_none());
for (user_id, is_active, is_deleted) in [
("inactive-user", false, false),
("deleted-user", true, true),
] {
let repository = Arc::new(InMemoryAuthApiKeySnapshotRepository::default());
let state = state_with_users(
Arc::clone(&repository),
[authoritative_user(user_id, is_active, is_deleted)],
);
let api_key_id = format!("{user_id}-key");
assert!(state
.create_user_api_key(create_record(user_id, &api_key_id))
.await
.expect("ineligible owner creation should resolve")
.is_none());
assert!(repository
.find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId(&api_key_id))
.await
.expect("ineligible key lookup should resolve")
.is_none());
}
}
#[tokio::test]
async fn test_gateway_first_user_key_syncs_owner_without_mutating_authority() {
let stale_owner = StoredAuthApiKeySnapshot::new(
"active-user".to_string(),
"request-derived-owner".to_string(),
None,
"user".to_string(),
"local".to_string(),
false,
false,
None,
None,
None,
"ignored-owner-fixture".to_string(),
None,
true,
false,
false,
None,
None,
None,
None,
None,
None,
)
.expect("stale owner fixture should build");
let repository = Arc::new(
InMemoryAuthApiKeySnapshotRepository::default().with_owner_snapshots([stale_owner]),
);
let authoritative = authoritative_user("active-user", true, false);
let state = state_with_users(Arc::clone(&repository), [authoritative.clone()]);
state
.create_user_api_key(create_record("active-user", "active-key"))
.await
.expect("active owner creation should resolve")
.expect("active authoritative owner should allow its first key");
let snapshot = repository
.find_api_key_snapshot(AuthApiKeyLookupKey::ApiKeyId("active-key"))
.await
.expect("created key lookup should resolve")
.expect("created key should exist");
assert_eq!(snapshot.user_role, "admin");
assert!(snapshot.user_is_active);
assert!(!snapshot.user_is_deleted);
assert_eq!(
snapshot.user_allowed_providers,
Some(vec!["openai".to_string()])
);
assert_eq!(
snapshot.api_key_allowed_providers,
Some(vec!["anthropic".to_string()])
);
let unchanged = state
.find_user_auth_by_id("active-user")
.await
.expect("authoritative owner lookup should resolve")
.expect("authoritative owner should remain");
assert_eq!(unchanged, authoritative);
assert_eq!(unchanged.role, authoritative.role);
assert_eq!(unchanged.is_active, authoritative.is_active);
assert_eq!(unchanged.is_deleted, authoritative.is_deleted);
assert_eq!(unchanged.security_version, 41);
}
}
@@ -122,13 +122,44 @@ impl AppState {
{
let session = session.into();
#[cfg(test)]
if let Some(store) = self.auth_session_store.as_ref() {
if let (Some(user_store), Some(session_store)) = (
self.auth_user_store.as_ref(),
self.auth_session_store.as_ref(),
) {
let existing = {
user_store
.lock()
.expect("auth user store should lock")
.get(&session.user_id)
.cloned()
};
let existing = match existing {
Some(user) => Some(user),
None => self
.data
.find_user_auth_by_id(&session.user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?,
};
let Some(existing) = existing else {
return Ok(None);
};
let mut users = user_store.lock().expect("auth user store should lock");
let user = users.entry(session.user_id.clone()).or_insert(existing);
if !user.is_active
|| user.is_deleted
|| user.security_version != session.security_version
{
return Ok(None);
}
let now = session
.created_at
.or(session.updated_at)
.or(session.last_seen_at)
.unwrap_or_else(chrono::Utc::now);
let mut guard = store.lock().expect("auth session store should lock");
let mut guard = session_store
.lock()
.expect("auth session store should lock");
for existing in guard.values_mut() {
if existing.user_id == session.user_id
&& existing.client_device_id == session.client_device_id
@@ -155,12 +186,96 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn create_user_session_if_password_matches<T>(
&self,
session: T,
expected_password_hash: &str,
) -> Result<Option<GatewayUserSessionView>, GatewayError>
where
T: Into<GatewayUserSessionView>,
{
let session = session.into();
#[cfg(test)]
if self.auth_session_store.is_some() && self.auth_user_store.is_some() {
let existing = {
self.auth_user_store
.as_ref()
.expect("checked auth user store")
.lock()
.expect("auth user store should lock")
.get(&session.user_id)
.cloned()
};
let existing = match existing {
Some(user) => Some(user),
None => self
.data
.find_user_auth_by_id(&session.user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?,
};
let Some(existing) = existing else {
return Ok(None);
};
let mut users = self
.auth_user_store
.as_ref()
.expect("checked auth user store")
.lock()
.expect("auth user store should lock");
let user = users.entry(session.user_id.clone()).or_insert(existing);
if user.password_hash.as_deref() != Some(expected_password_hash)
|| !user.auth_source.eq_ignore_ascii_case("local")
|| !user.is_active
|| user.is_deleted
|| user.security_version != session.security_version
{
return Ok(None);
}
let now = session
.created_at
.or(session.updated_at)
.or(session.last_seen_at)
.unwrap_or_else(chrono::Utc::now);
user.last_login_at = Some(now);
let mut sessions = self
.auth_session_store
.as_ref()
.expect("checked auth session store")
.lock()
.expect("auth session store should lock");
for existing in sessions.values_mut() {
if existing.user_id == session.user_id
&& existing.client_device_id == session.client_device_id
&& !existing.is_revoked()
&& !existing.is_expired(now)
{
existing.revoked_at = Some(now);
existing.revoke_reason = Some("replaced_by_new_login".to_string());
existing.updated_at = Some(now);
}
}
sessions.insert(
format!("{}:{}", session.user_id, session.id),
session.clone().into(),
);
return Ok(Some(session));
}
let raw_session: crate::data::state::StoredUserSessionRecord = session.into();
self.data
.create_user_session_if_password_matches(&raw_session, expected_password_hash)
.await
.map(|value| value.map(Into::into))
.map_err(|err| GatewayError::Internal(err.to_string()))
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn rotate_user_session_refresh_token(
&self,
user_id: &str,
session_id: &str,
previous_refresh_token_hash: &str,
expected_refresh_token_hash: &str,
next_refresh_token_hash: &str,
rotated_at: chrono::DateTime<chrono::Utc>,
expires_at: chrono::DateTime<chrono::Utc>,
@@ -171,8 +286,12 @@ impl AppState {
if let Some(store) = self.auth_session_store.as_ref() {
let key = format!("{user_id}:{session_id}");
let mut guard = store.lock().expect("auth session store should lock");
if let Some(session) = guard.get_mut(&key) {
session.prev_refresh_token_hash = Some(previous_refresh_token_hash.to_string());
if let Some(session) = guard.get_mut(&key).filter(|session| {
session.refresh_token_hash == expected_refresh_token_hash
&& !session.is_revoked()
&& !session.is_expired(rotated_at)
}) {
session.prev_refresh_token_hash = Some(expected_refresh_token_hash.to_string());
session.refresh_token_hash = next_refresh_token_hash.to_string();
session.rotated_at = Some(rotated_at);
session.expires_at = Some(expires_at);
@@ -193,7 +312,7 @@ impl AppState {
.rotate_user_session_refresh_token(
user_id,
session_id,
previous_refresh_token_hash,
expected_refresh_token_hash,
next_refresh_token_hash,
rotated_at,
expires_at,
@@ -288,6 +288,22 @@ impl AppState {
Ok(group)
}
pub(crate) async fn restore_user_group_if_matches(
&self,
expected: &aether_data::repository::users::StoredUserGroup,
restored: &aether_data::repository::users::StoredUserGroup,
) -> Result<bool, GatewayError> {
let restored = self
.data
.restore_user_group_if_matches(expected, restored)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if restored {
self.invalidate_auth_context_cache();
}
Ok(restored)
}
pub(crate) async fn delete_user_group(&self, group_id: &str) -> Result<bool, GatewayError> {
let deleted = self
.data
@@ -369,6 +385,23 @@ impl AppState {
Ok(groups)
}
pub(crate) async fn restore_user_groups_if_matches(
&self,
user_id: &str,
expected_group_ids: &[String],
restored_group_ids: &[String],
) -> Result<bool, GatewayError> {
let restored = self
.data
.restore_user_groups_if_matches(user_id, expected_group_ids, restored_group_ids)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if restored {
self.invalidate_auth_context_cache();
}
Ok(restored)
}
pub(crate) async fn add_user_to_group(
&self,
group_id: &str,
@@ -434,7 +467,9 @@ impl AppState {
pub(crate) async fn update_local_auth_user_profile(
&self,
user_id: &str,
email_present: bool,
email: Option<String>,
email_verified: Option<bool>,
username: Option<String>,
) -> Result<Option<aether_data::repository::users::StoredUserAuthRecord>, GatewayError> {
#[cfg(test)]
@@ -457,8 +492,11 @@ impl AppState {
let Some(mut user) = existing else {
return Ok(None);
};
if let Some(email) = email {
user.email = Some(email);
if email_present {
user.email = email;
}
if let Some(email_verified) = email_verified {
user.email_verified = email_verified;
}
if let Some(username) = username {
user.username = username;
@@ -473,7 +511,7 @@ impl AppState {
let user = self
.data
.update_local_auth_user_profile(user_id, email, username)
.update_local_auth_user_profile(user_id, email_present, email, email_verified, username)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if user.is_some() {
@@ -482,6 +520,159 @@ impl AppState {
Ok(user)
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn restore_local_auth_user_state_if_matches(
&self,
expected_auth: &aether_data::repository::users::StoredUserAuthRecord,
restored_auth: &aether_data::repository::users::StoredUserAuthRecord,
expected_export: &aether_data::repository::users::StoredUserExportRow,
restored_export: &aether_data::repository::users::StoredUserExportRow,
expected_model_capability_settings: Option<&serde_json::Value>,
restored_model_capability_settings: Option<serde_json::Value>,
expected_feature_settings: Option<&serde_json::Value>,
restored_feature_settings: Option<serde_json::Value>,
) -> Result<bool, GatewayError> {
#[cfg(test)]
if let Some(store) = self.auth_user_store.as_ref() {
// The gateway test harness uses an auth overlay for users while most auxiliary data
// remains in the repository. Keep the compare-and-write atomic for that overlay too;
// otherwise import rollback tests would silently exercise a different, unconditional
// path than production.
let current_feature = self
.data
.read_user_feature_settings(&expected_auth.id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let mut users = store.lock().expect("auth user store should lock");
if users.contains_key(&expected_auth.id) {
if expected_auth.id != restored_auth.id
|| expected_export.id != expected_auth.id
|| restored_export.id != restored_auth.id
{
return Ok(false);
}
let Some(current) = users.get(&expected_auth.id).cloned() else {
return Ok(false);
};
if !current.matches_restore_state(expected_auth) {
return Ok(false);
}
let current_model =
self.auth_user_model_capability_store
.as_ref()
.and_then(|settings| {
settings
.lock()
.expect("auth user model capability store should lock")
.get(&expected_auth.id)
.cloned()
});
if current_model.as_ref() != expected_model_capability_settings {
return Ok(false);
}
// Feature settings have no separate test overlay. When the backing repository has
// a row, still honor the snapshot comparison; an absent row is represented by
// `None`, which is the normal overlay case.
if current_feature.as_ref() != expected_feature_settings {
return Ok(false);
}
let security_state_changed = current.role != restored_auth.role
|| current.is_active != restored_auth.is_active;
let removes_active_admin = current.role.eq_ignore_ascii_case("admin")
&& current.is_active
&& !current.is_deleted
&& (!restored_auth.role.eq_ignore_ascii_case("admin")
|| !restored_auth.is_active);
if removes_active_admin
&& users
.values()
.filter(|user| {
user.role.eq_ignore_ascii_case("admin")
&& user.is_active
&& !user.is_deleted
})
.count()
<= 1
{
return Err(GatewayError::LastActiveAdminUpdateDenied);
}
let mut updated = restored_auth.clone();
// Server-managed credentials and timestamps are deliberately not part of this
// aggregate restore. The password has its own nullable CAS operation.
updated.password_hash = current.password_hash;
updated.security_version = current.security_version;
updated.created_at = current.created_at;
updated.last_login_at = current.last_login_at;
updated.auth_source = current.auth_source;
updated.is_deleted = current.is_deleted;
if security_state_changed {
updated.security_version =
updated.security_version.checked_add(1).ok_or_else(|| {
GatewayError::Internal("users.security_version overflow".to_string())
})?;
}
users.insert(updated.id.clone(), updated.clone());
drop(users);
if let Some(settings) = self.auth_user_model_capability_store.as_ref() {
let mut guard = settings
.lock()
.expect("auth user model capability store should lock");
match restored_model_capability_settings {
Some(value) => {
guard.insert(updated.id.clone(), value);
}
None => {
guard.remove(&updated.id);
}
}
}
if security_state_changed {
if let Some(sessions) = self.auth_session_store.as_ref() {
let now = chrono::Utc::now();
for session in sessions
.lock()
.expect("auth session store should lock")
.values_mut()
.filter(|session| {
session.user_id == updated.id && session.revoked_at.is_none()
})
{
session.revoked_at = Some(now);
session.revoke_reason = Some("user_security_state_changed".to_string());
session.updated_at = Some(now);
}
}
}
self.invalidate_auth_context_cache();
return Ok(true);
}
}
let restored = self
.data
.restore_local_auth_user_state_if_matches(
expected_auth,
restored_auth,
expected_export,
restored_export,
expected_model_capability_settings,
restored_model_capability_settings,
expected_feature_settings,
restored_feature_settings,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if restored {
self.invalidate_auth_context_cache();
}
Ok(restored)
}
pub(crate) async fn update_local_auth_user_password_hash(
&self,
user_id: &str,
@@ -509,6 +700,9 @@ impl AppState {
return Ok(None);
};
user.password_hash = Some(password_hash);
user.security_version = user.security_version.checked_add(1).ok_or_else(|| {
GatewayError::Internal("users.security_version overflow".to_string())
})?;
store
.lock()
.expect("auth user store should lock")
@@ -522,6 +716,193 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn restore_local_auth_user_password_hash_if_matches(
&self,
user_id: &str,
expected_password_hash: Option<&str>,
password_hash: Option<String>,
updated_at: chrono::DateTime<chrono::Utc>,
) -> Result<bool, GatewayError> {
#[cfg(test)]
if let Some(store) = self.auth_user_store.as_ref() {
let existing = {
store
.lock()
.expect("auth user store should lock")
.get(user_id)
.cloned()
};
let existing = match existing {
Some(user) => Some(user),
None => self
.data
.find_user_auth_by_id(user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?,
};
let Some(mut user) = existing else {
return Ok(false);
};
if user.password_hash.as_deref() != expected_password_hash {
return Ok(false);
}
user.password_hash = password_hash;
user.security_version = user.security_version.checked_add(1).ok_or_else(|| {
GatewayError::Internal("users.security_version overflow".to_string())
})?;
store
.lock()
.expect("auth user store should lock")
.insert(user.id.clone(), user);
self.invalidate_auth_context_cache();
return Ok(true);
}
self.data
.restore_local_auth_user_password_hash_if_matches(
user_id,
expected_password_hash,
password_hash,
updated_at,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn reset_local_auth_user_password_and_revoke_sessions(
&self,
user_id: &str,
password_hash: String,
changed_at: chrono::DateTime<chrono::Utc>,
) -> Result<bool, GatewayError> {
#[cfg(test)]
if let (Some(user_store), Some(session_store)) = (
self.auth_user_store.as_ref(),
self.auth_session_store.as_ref(),
) {
let existing = {
user_store
.lock()
.expect("auth user store should lock")
.get(user_id)
.cloned()
};
let existing = match existing {
Some(user) => Some(user),
None => self
.data
.find_user_auth_by_id(user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?,
};
let Some(existing) = existing.filter(|user| !user.is_deleted) else {
return Ok(false);
};
let mut users = user_store.lock().expect("auth user store should lock");
let mut sessions = session_store
.lock()
.expect("auth session store should lock");
let user = users.entry(user_id.to_string()).or_insert(existing);
user.password_hash = Some(password_hash);
user.security_version = user.security_version.checked_add(1).ok_or_else(|| {
GatewayError::Internal("users.security_version overflow".to_string())
})?;
for session in sessions
.values_mut()
.filter(|session| session.user_id == user_id && !session.is_revoked())
{
session.revoked_at = Some(changed_at);
session.revoke_reason = Some("admin_password_reset".to_string());
session.updated_at = Some(changed_at);
}
self.invalidate_auth_context_cache();
return Ok(true);
}
let reset = self
.data
.reset_local_auth_user_password_and_revoke_sessions(user_id, password_hash, changed_at)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
if reset {
self.invalidate_auth_context_cache();
}
Ok(reset)
}
pub(crate) async fn change_local_auth_password_and_revoke_sessions(
&self,
user_id: &str,
current_session_id: &str,
expected_password_hash: Option<&str>,
next_password_hash: String,
changed_at: chrono::DateTime<chrono::Utc>,
) -> Result<bool, GatewayError> {
#[cfg(test)]
if let (Some(user_store), Some(session_store)) = (
self.auth_user_store.as_ref(),
self.auth_session_store.as_ref(),
) {
let existing = {
user_store
.lock()
.expect("auth user store should lock")
.get(user_id)
.cloned()
};
let existing = match existing {
Some(user) => Some(user),
None => self
.data
.find_user_auth_by_id(user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?,
};
let Some(existing) = existing else {
return Ok(false);
};
let mut users = user_store.lock().expect("auth user store should lock");
let mut sessions = session_store
.lock()
.expect("auth session store should lock");
let user = users.entry(user_id.to_string()).or_insert(existing);
if user.password_hash.as_deref() != expected_password_hash {
return Ok(false);
}
let current_key = format!("{user_id}:{current_session_id}");
if !sessions
.get(&current_key)
.is_some_and(|session| !session.is_revoked() && !session.is_expired(changed_at))
{
return Ok(false);
}
user.password_hash = Some(next_password_hash);
user.security_version = user.security_version.checked_add(1).ok_or_else(|| {
GatewayError::Internal("users.security_version overflow".to_string())
})?;
for session in sessions
.values_mut()
.filter(|session| session.user_id == user_id && !session.is_revoked())
{
session.revoked_at = Some(changed_at);
session.revoke_reason = Some("password_changed".to_string());
session.updated_at = Some(changed_at);
}
return Ok(true);
}
self.data
.change_local_auth_password_and_revoke_sessions(
user_id,
current_session_id,
expected_password_hash,
next_password_hash,
changed_at,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn create_local_auth_user(
&self,
email: Option<String>,
@@ -642,10 +1023,18 @@ impl AppState {
) -> Result<Option<aether_data::repository::users::StoredUserAuthRecord>, GatewayError> {
#[cfg(test)]
if let Some(store) = self.auth_user_store.as_ref() {
let mut guard = store.lock().expect("auth user store should lock");
let Some(user) = guard.get_mut(user_id) else {
let mut users = store.lock().expect("auth user store should lock");
let Some(user) = users.get_mut(user_id) else {
return Ok(None);
};
let security_state_changed = role
.as_deref()
.is_some_and(|next_role| !user.role.eq_ignore_ascii_case(next_role))
|| is_active.is_some_and(|next_active| user.is_active != next_active);
let mut sessions = self
.auth_session_store
.as_ref()
.map(|sessions| sessions.lock().expect("auth session store should lock"));
if let Some(role) = role {
user.role = role;
}
@@ -661,9 +1050,26 @@ impl AppState {
if let Some(is_active) = is_active {
user.is_active = is_active;
}
if security_state_changed {
user.security_version = user.security_version.checked_add(1).ok_or_else(|| {
GatewayError::Internal("users.security_version overflow".to_string())
})?;
let revoked_at = chrono::Utc::now();
if let Some(sessions) = sessions.as_mut() {
for session in sessions
.values_mut()
.filter(|session| session.user_id == user_id && !session.is_revoked())
{
session.revoked_at = Some(revoked_at);
session.revoke_reason = Some("user_security_state_changed".to_string());
session.updated_at = Some(revoked_at);
}
}
}
let _ = (rate_limit_present, rate_limit);
let user = user.clone();
drop(guard);
drop(sessions);
drop(users);
self.invalidate_auth_context_cache();
return Ok(Some(user));
}
@@ -684,7 +1090,13 @@ impl AppState {
is_active,
)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
.map_err(|err| {
if aether_data::repository::users::is_last_active_admin_update_denied(&err) {
GatewayError::LastActiveAdminUpdateDenied
} else {
GatewayError::Internal(err.to_string())
}
})?;
if user.is_some() {
self.invalidate_auth_context_cache();
}
@@ -788,10 +1200,16 @@ impl AppState {
self.data
.delete_local_auth_user(user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
.map_err(|err| {
if aether_data::repository::users::is_last_active_admin_delete_denied(&err) {
GatewayError::LastActiveAdminDeleteDenied
} else {
GatewayError::Internal(err.to_string())
}
})
}
pub(crate) async fn register_local_auth_user(
pub(crate) async fn register_local_auth_user_with_wallet_outcome(
&self,
email: Option<String>,
email_verified: bool,
@@ -803,6 +1221,7 @@ impl AppState {
Option<(
aether_data::repository::users::StoredUserAuthRecord,
aether_data::repository::wallet::StoredWalletSnapshot,
bool,
)>,
GatewayError,
> {
@@ -868,11 +1287,17 @@ impl AppState {
.lock()
.expect("auth wallet store should lock")
.insert(wallet.id.clone(), wallet.clone());
return Ok(Some((user, wallet)));
super::user_provisioning::record_test_initial_gift_transaction(
self,
&wallet,
&user.id,
"用户初始赠款",
);
return Ok(Some((user, wallet, true)));
}
self.data
.register_local_auth_user(
.register_local_auth_user_with_wallet_outcome(
email,
email_verified,
username,
@@ -883,6 +1308,34 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn register_local_auth_user(
&self,
email: Option<String>,
email_verified: bool,
username: String,
password_hash: String,
initial_gift_usd: f64,
unlimited: bool,
) -> Result<
Option<(
aether_data::repository::users::StoredUserAuthRecord,
aether_data::repository::wallet::StoredWalletSnapshot,
)>,
GatewayError,
> {
Ok(self
.register_local_auth_user_with_wallet_outcome(
email,
email_verified,
username,
password_hash,
initial_gift_usd,
unlimited,
)
.await?
.map(|(user, wallet, _created)| (user, wallet)))
}
}
fn normalized_user_group_ids(group_ids: &[String]) -> BTreeSet<String> {
@@ -944,6 +1397,7 @@ mod tests {
local_rejection: None,
allowed_models: Some(vec!["gpt-4.1".to_string()]),
ip_rules: None,
verified_api_key_hash: None,
}
}
File diff suppressed because it is too large Load Diff
@@ -2,10 +2,12 @@ use super::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, AppState,
BillingPlanRecord, BillingPlanWriteInput, GatewayError, LocalMutationOutcome,
PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, UserDailyQuotaAvailabilityRecord,
UserPlanEntitlementRecord,
PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput,
PaymentGatewaySecretCasUpdate, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord,
};
const PAYMENT_GATEWAY_SECRET_MIGRATION_MAX_ATTEMPTS: usize = 8;
fn data_error(err: impl ToString) -> GatewayError {
GatewayError::Internal(err.to_string())
}
@@ -446,9 +448,73 @@ impl AppState {
&self,
provider: &str,
) -> Result<Option<PaymentGatewayConfigRecord>, GatewayError> {
self.data
.find_payment_gateway_config(provider)
let provider = provider.trim().to_ascii_lowercase();
let mut record = self
.data
.find_payment_gateway_config(&provider)
.await
.map_err(data_error)?;
for _ in 0..PAYMENT_GATEWAY_SECRET_MIGRATION_MAX_ATTEMPTS {
let Some(mut current) = record else {
return Ok(None);
};
let Some(observed) = current.merchant_key_encrypted.as_deref() else {
return Ok(Some(current));
};
let projection = crate::handlers::shared::open_payment_gateway_secret(
self,
&crate::handlers::shared::PaymentGatewaySecretBinding::from_record(&current)
.map_err(|detail| {
GatewayError::Internal(format!(
"payment gateway secret binding is invalid for {}: {detail}",
current.provider
))
})?,
observed,
)
.map_err(|detail| {
GatewayError::Internal(format!(
"payment gateway secret integrity check failed for {}: {detail}",
current.provider
))
})?;
if !projection.migration_required {
return Ok(Some(current));
}
let update = PaymentGatewaySecretCasUpdate {
provider: current.provider.clone(),
expected_merchant_key_encrypted: observed.to_string(),
merchant_key_encrypted: projection.protected.clone(),
};
if self
.data
.compare_and_swap_payment_gateway_secret(&update)
.await
.map_err(data_error)?
{
current.merchant_key_encrypted = Some(projection.protected);
return Ok(Some(current));
}
record = self
.data
.find_payment_gateway_config_strong(&provider)
.await
.map_err(data_error)?;
}
Err(GatewayError::Internal(format!(
"payment gateway secret migration did not converge for {provider}"
)))
}
pub(crate) async fn compare_and_swap_payment_gateway_config(
&self,
input: &PaymentGatewayConfigCasWriteInput,
) -> Result<LocalMutationOutcome<PaymentGatewayConfigRecord>, GatewayError> {
self.data
.compare_and_swap_payment_gateway_config(input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
}
@@ -608,13 +674,20 @@ impl AppState {
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use aether_crypto::{encrypt_python_fernet_plaintext, DEVELOPMENT_ENCRYPTION_KEY};
use aether_data::repository::billing::InMemoryBillingReadRepository;
use aether_data_contracts::repository::billing::{
BillingReadRepository, PaymentGatewayConfigWriteInput,
};
use serde_json::json;
use super::{
AdminBillingCollectorWriteInput, AdminBillingRuleWriteInput, AppState, LocalMutationOutcome,
};
use crate::data::GatewayDataState;
const CACHE_KEY: &str = "billing-mutation-test";
@@ -736,4 +809,59 @@ mod tests {
.get(&CACHE_KEY.to_string(), Duration::from_secs(60),)
.is_some());
}
#[tokio::test]
async fn gateway_lookup_lazily_migrates_legacy_secret_without_touching_other_fields() {
let repository = Arc::new(InMemoryBillingReadRepository::default());
let legacy = encrypt_python_fernet_plaintext(
DEVELOPMENT_ENCRYPTION_KEY,
r#"{"secret_key":"legacy-value"}"#,
)
.expect("legacy secret should encrypt");
repository
.upsert_payment_gateway_config(&PaymentGatewayConfigWriteInput {
provider: "stripe".to_string(),
enabled: true,
endpoint_url: "https://api.stripe.com".to_string(),
callback_base_url: Some("https://example.com".to_string()),
merchant_id: "merchant".to_string(),
merchant_key_encrypted: Some(legacy),
preserve_existing_secret: false,
pay_currency: "USD".to_string(),
usd_exchange_rate: 1.0,
min_recharge_usd: 1.0,
channels_json: json!({"channels": []}),
})
.await
.expect("gateway seed should succeed");
let before = repository
.find_payment_gateway_config("stripe")
.await
.expect("gateway lookup should succeed")
.expect("gateway should exist");
let data = GatewayDataState::with_billing_reader_for_tests(repository.clone())
.with_encryption_key_for_tests(DEVELOPMENT_ENCRYPTION_KEY);
let state = AppState::new()
.expect("app state should build")
.with_data_state_for_tests(data);
let migrated = state
.find_payment_gateway_config("stripe")
.await
.expect("gateway migration should succeed")
.expect("gateway should exist");
assert!(migrated
.merchant_key_encrypted
.as_deref()
.is_some_and(|value| value.starts_with("aether-payment-gateway-secret-v3:")));
let stored = repository
.find_payment_gateway_config("stripe")
.await
.expect("gateway lookup should succeed")
.expect("gateway should exist");
assert_eq!(stored, migrated);
assert_eq!(stored.updated_at_unix_secs, before.updated_at_unix_secs);
assert_eq!(stored.endpoint_url, before.endpoint_url);
assert_eq!(stored.channels_json, before.channels_json);
}
}
@@ -74,7 +74,7 @@ impl AppState {
let effective_status = if order.status == "pending"
&& order
.expires_at_unix_secs
.is_some_and(|value| value < now_unix_secs)
.is_some_and(|value| value <= now_unix_secs)
{
"expired"
} else {
@@ -2,8 +2,8 @@ use super::super::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, AppState,
BillingPlanRecord, BillingPlanWriteInput, GatewayError, LocalMutationOutcome,
PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, UserDailyQuotaAvailabilityRecord,
UserPlanEntitlementRecord,
PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput,
PaymentGatewaySecretCasUpdate, UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord,
};
mod admin;
@@ -111,8 +111,9 @@ impl AppState {
pub(crate) async fn upsert_request_candidate(
&self,
candidate: candidates::UpsertRequestCandidateRecord,
mut candidate: candidates::UpsertRequestCandidateRecord,
) -> Result<Option<candidates::StoredRequestCandidate>, GatewayError> {
candidate.sanitize_for_persistence();
if let Some(queue) = self.request_candidate_queue.as_ref() {
let stored = stored_request_candidate_from_upsert(&candidate)?;
queue
@@ -136,8 +137,9 @@ impl AppState {
/// the async queue is enabled.
pub(crate) async fn enqueue_request_candidate_status(
&self,
candidate: candidates::UpsertRequestCandidateRecord,
mut candidate: candidates::UpsertRequestCandidateRecord,
) -> Result<Option<()>, GatewayError> {
candidate.sanitize_for_persistence();
if let Some(queue) = self.request_candidate_queue.as_ref() {
queue
.enqueue_or_fallback(candidate)
@@ -158,8 +160,9 @@ impl AppState {
/// when the queue is disabled or closed.
pub(crate) fn try_enqueue_request_candidate_status(
&self,
candidate: candidates::UpsertRequestCandidateRecord,
mut candidate: candidates::UpsertRequestCandidateRecord,
) -> Result<(), candidates::UpsertRequestCandidateRecord> {
candidate.sanitize_for_persistence();
let Some(queue) = self.request_candidate_queue.as_ref() else {
return Err(candidate);
};
@@ -14,6 +14,19 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn upsert_gemini_file_mapping_if_owner_matches(
&self,
record: aether_data::repository::gemini_file_mappings::UpsertGeminiFileMappingRecord,
) -> Result<
Option<aether_data::repository::gemini_file_mappings::StoredGeminiFileMapping>,
GatewayError,
> {
self.data
.upsert_gemini_file_mapping_if_owner_matches(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn list_gemini_file_mappings(
&self,
query: &aether_data::repository::gemini_file_mappings::GeminiFileMappingListQuery,
@@ -27,6 +40,50 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_gemini_file_mapping_by_file_name(
&self,
file_name: &str,
) -> Result<
Option<aether_data::repository::gemini_file_mappings::StoredGeminiFileMapping>,
GatewayError,
> {
self.data
.find_gemini_file_mapping_by_file_name(file_name)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_active_gemini_file_mapping_for_user(
&self,
file_name: &str,
user_id: &str,
now_unix_secs: u64,
) -> Result<
Option<aether_data::repository::gemini_file_mappings::StoredGeminiFileMapping>,
GatewayError,
> {
self.data
.find_active_gemini_file_mapping_for_user(file_name, user_id, now_unix_secs)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_active_gemini_file_mapping_for_owner(
&self,
file_name: &str,
key_id: &str,
user_id: &str,
now_unix_secs: u64,
) -> Result<
Option<aether_data::repository::gemini_file_mappings::StoredGeminiFileMapping>,
GatewayError,
> {
self.data
.find_active_gemini_file_mapping_for_owner(file_name, key_id, user_id, now_unix_secs)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn summarize_gemini_file_mappings(
&self,
now_unix_secs: u64,
@@ -48,6 +105,29 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn delete_gemini_file_mapping_by_file_name_for_user(
&self,
file_name: &str,
user_id: &str,
) -> Result<bool, GatewayError> {
self.data
.delete_gemini_file_mapping_by_file_name_for_user(file_name, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn delete_gemini_file_mapping_by_file_name_for_owner(
&self,
file_name: &str,
key_id: &str,
user_id: &str,
) -> Result<bool, GatewayError> {
self.data
.delete_gemini_file_mapping_by_file_name_for_owner(file_name, key_id, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn delete_gemini_file_mapping_by_id(
&self,
mapping_id: &str,
@@ -1,5 +1,6 @@
use std::collections::BTreeMap;
use aether_admin::observability::usage::admin_usage_safe_metadata_value;
use aether_data::repository::audit::AuditLogListQuery;
use chrono::{DateTime, Utc};
use serde_json::{json, Value};
@@ -43,8 +44,8 @@ impl AppState {
"description": record.description,
"ip_address": record.ip_address,
"status_code": record.status_code,
"error_message": record.error_message,
"metadata": record.metadata,
"error_message": record.error_message.as_ref().map(|_| "audit_event_failed"),
"metadata": sanitize_admin_audit_metadata(record.metadata.as_ref()),
"created_at": record.created_at_rfc3339(),
})
})
@@ -72,7 +73,7 @@ impl AppState {
"user_id": record.user_id,
"description": record.description,
"ip_address": record.ip_address,
"metadata": record.metadata,
"metadata": sanitize_admin_audit_metadata(record.metadata.as_ref()),
"created_at": record.created_at_rfc3339(),
})
})
@@ -134,3 +135,41 @@ impl AppState {
fn cutoff_unix_secs(cutoff_time: DateTime<Utc>) -> u64 {
cutoff_time.timestamp().max(0) as u64
}
fn sanitize_admin_audit_metadata(metadata: Option<&Value>) -> Value {
metadata
.map(admin_usage_safe_metadata_value)
.unwrap_or(Value::Null)
}
#[cfg(test)]
mod tests {
use super::sanitize_admin_audit_metadata;
use serde_json::json;
#[test]
fn admin_audit_metadata_drops_credentials_and_url_components() {
let metadata = sanitize_admin_audit_metadata(Some(&json!({
"category": "security",
"authorization": "Bearer audit-secret",
"nested": {
"refresh_token": "refresh-secret",
"endpoint_url": "https://user:[email protected]/v1?token=query-secret#fragment",
"safe_count": 2
}
})));
assert_eq!(metadata["category"], "security");
assert!(metadata.get("authorization").is_none());
assert!(metadata["nested"].get("refresh_token").is_none());
assert_eq!(
metadata["nested"]["endpoint_url"],
"https://example.test/v1"
);
assert_eq!(metadata["nested"]["safe_count"], 2);
let encoded = metadata.to_string();
for secret in ["audit-secret", "refresh-secret", "password", "query-secret"] {
assert!(!encoded.contains(secret), "leaked {secret}");
}
}
}
@@ -142,7 +142,7 @@ impl AppState {
}
if order
.expires_at_unix_secs
.is_some_and(|value| value < chrono::Utc::now().timestamp().max(0) as u64)
.is_some_and(|value| value <= chrono::Utc::now().timestamp().max(0) as u64)
{
return Ok(AdminWalletMutationOutcome::Invalid(
"payment order expired".to_string(),
@@ -4,13 +4,39 @@ use crate::data::state::{
};
use crate::{AppState, GatewayError};
use axum::http::StatusCode;
use tracing::warn;
const REFERRAL_INVALID_INPUT_FALLBACK: &str = "返利请求无效";
fn safe_referral_invalid_input_detail(detail: &str) -> &'static str {
// These messages are deliberate domain-level validation responses. Any
// future adapter/storage detail must stay server-side instead of becoming
// an oracle for database state or schema information.
match detail {
"邀请码无效" => "邀请码无效",
"不能使用自己的邀请码注册" => "不能使用自己的邀请码注册",
"仅失败返利可以补发" => "仅失败返利可以补发",
"返利金额无效,无法补发" => "返利金额无效,无法补发",
_ => REFERRAL_INVALID_INPUT_FALLBACK,
}
}
fn referral_data_error(err: aether_data::DataLayerError) -> GatewayError {
match err {
aether_data::DataLayerError::InvalidInput(detail) => GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: detail,
},
aether_data::DataLayerError::InvalidInput(detail) => {
let safe_detail = safe_referral_invalid_input_detail(&detail);
if safe_detail == REFERRAL_INVALID_INPUT_FALLBACK {
warn!(
event_name = "referral_invalid_input_hidden",
error_length = detail.len(),
"referral data-layer validation detail hidden from client"
);
}
GatewayError::Client {
status: StatusCode::BAD_REQUEST,
message: safe_detail.to_string(),
}
}
other => GatewayError::Internal(other.to_string()),
}
}
@@ -51,6 +77,13 @@ fn config_f64(value: Option<&serde_json::Value>, default: f64) -> f64 {
}
}
fn config_percent(value: Option<&serde_json::Value>) -> f64 {
let value = config_f64(value, 0.0);
(value.is_finite() && value > 0.0 && value <= 100.0)
.then_some(value)
.unwrap_or(0.0)
}
impl AppState {
pub(crate) fn has_referral_data_backend(&self) -> bool {
self.data.has_referral_data_backend()
@@ -93,7 +126,7 @@ impl AppState {
config_string(headcount_trigger.as_ref()).unwrap_or_else(|| "registration".to_string());
Ok(Some(ReferralRewardConfig {
percent_enabled: matches!(mode.as_str(), "percent" | "both"),
percent_rate: config_f64(percent.as_ref(), 0.0),
percent_rate: config_percent(percent.as_ref()),
headcount_enabled: matches!(mode.as_str(), "headcount" | "both"),
headcount_amount_usd: config_f64(headcount_amount.as_ref(), 0.0),
headcount_trigger,
@@ -238,4 +271,43 @@ impl AppState {
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn reconcile_referral_rewards_once(
&self,
) -> Result<crate::data::state::ReferralReconciliationSummary, GatewayError> {
let config = self.referral_reward_config().await?;
self.data
.reconcile_referral_rewards_once(config)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
}
#[cfg(test)]
mod tests {
use super::{referral_data_error, REFERRAL_INVALID_INPUT_FALLBACK};
#[test]
fn referral_invalid_input_projection_allowlists_domain_messages() {
let known = super::referral_data_error(aether_data::DataLayerError::InvalidInput(
"邀请码无效".to_string(),
));
match known {
crate::GatewayError::Client { message, .. } => assert_eq!(message, "邀请码无效"),
other => panic!("expected client error, got {other:?}"),
}
let secret = "database table referral_rewards row reward-secret has invalid wallet";
let unknown = referral_data_error(aether_data::DataLayerError::InvalidInput(
secret.to_string(),
));
match unknown {
crate::GatewayError::Client { message, .. } => {
assert_eq!(message, REFERRAL_INVALID_INPUT_FALLBACK);
assert!(!message.contains("reward-secret"));
assert!(!message.contains("referral_rewards"));
}
other => panic!("expected client error, got {other:?}"),
}
}
}
@@ -1,10 +1,12 @@
use aether_data::repository::wallet::{
AdjustWalletBalanceInput, CompleteAdminWalletRefundInput, CreateManualWalletRechargeInput,
CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput,
AdjustWalletBalanceInput, CompareAndSwapPaymentOrderStripeClientSecretInput,
CompleteAdminWalletRefundInput, CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput,
CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput,
CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput,
CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput, FailAdminWalletRefundInput,
ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome,
WalletMutationOutcome,
FailWalletRechargeCheckoutInput, ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput,
ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput,
UpdateAdminWalletRefundGatewayInput, UpdateWalletRechargeCheckoutInput, WalletMutationOutcome,
};
use crate::{AppState, GatewayError};
@@ -20,6 +22,55 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn update_wallet_recharge_checkout(
&self,
input: UpdateWalletRechargeCheckoutInput,
) -> Result<
Option<WalletMutationOutcome<aether_data::repository::wallet::StoredAdminPaymentOrder>>,
GatewayError,
> {
self.data
.update_wallet_recharge_checkout(input)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn compare_and_swap_payment_order_stripe_client_secret(
&self,
input: CompareAndSwapPaymentOrderStripeClientSecretInput,
) -> Result<Option<bool>, GatewayError> {
self.data
.compare_and_swap_payment_order_stripe_client_secret(input)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn fail_wallet_recharge_checkout(
&self,
input: FailWalletRechargeCheckoutInput,
) -> Result<
Option<WalletMutationOutcome<aether_data::repository::wallet::StoredAdminPaymentOrder>>,
GatewayError,
> {
self.data
.fail_wallet_recharge_checkout(input)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn reclaim_wallet_recharge_checkout(
&self,
input: ReclaimWalletRechargeCheckoutInput,
) -> Result<
Option<WalletMutationOutcome<aether_data::repository::wallet::StoredAdminPaymentOrder>>,
GatewayError,
> {
self.data
.reclaim_wallet_recharge_checkout(input)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn create_plan_purchase_order(
&self,
input: CreatePlanPurchaseOrderInput,
@@ -134,6 +185,19 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn update_admin_wallet_refund_gateway(
&self,
input: UpdateAdminWalletRefundGatewayInput,
) -> Result<
Option<WalletMutationOutcome<aether_data::repository::wallet::StoredAdminWalletRefund>>,
GatewayError,
> {
self.data
.update_admin_wallet_refund_gateway(input)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn fail_admin_wallet_refund(
&self,
input: FailAdminWalletRefundInput,
@@ -191,6 +191,18 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_wallet_recharge_order_by_order_no(
&self,
user_id: &str,
order_no: &str,
) -> Result<Option<aether_data::repository::wallet::StoredAdminPaymentOrder>, GatewayError>
{
self.data
.find_wallet_recharge_order_by_order_no(user_id, order_no)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_pending_plan_purchase_order_by_user_id(
&self,
user_id: &str,
@@ -203,6 +215,28 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_payment_order_by_order_no(
&self,
order_no: &str,
) -> Result<Option<aether_data::repository::wallet::StoredAdminPaymentOrder>, GatewayError>
{
self.data
.find_payment_order_by_order_no(order_no)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_payment_order_by_id(
&self,
order_id: &str,
) -> Result<Option<aether_data::repository::wallet::StoredAdminPaymentOrder>, GatewayError>
{
self.data
.find_admin_payment_order(order_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_wallet_refund(
&self,
wallet_id: &str,
@@ -2,6 +2,9 @@ use super::{
AdminWalletMutationOutcome, AdminWalletRefundRecord, AdminWalletTransactionRecord, AppState,
GatewayError,
};
use aether_data::repository::wallet::{
payment_order_refund_amounts_are_consistent, wallet_refund_proof_is_success,
};
impl AppState {
pub(crate) async fn admin_process_wallet_refund(
@@ -39,6 +42,11 @@ impl AppState {
else {
return Ok(AdminWalletMutationOutcome::NotFound);
};
if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 {
return Ok(AdminWalletMutationOutcome::Invalid(
"refund amount must be finite and greater than zero".to_string(),
));
}
if !matches!(refund.status.as_str(), "approved" | "pending_approval") {
return Ok(AdminWalletMutationOutcome::Invalid(
"refund status is not approvable".to_string(),
@@ -49,8 +57,26 @@ impl AppState {
let mut updated_wallet = wallet.clone();
let before_recharge = updated_wallet.balance;
let before_gift = updated_wallet.gift_balance;
let before_total_refunded = updated_wallet.total_refunded;
let before_total = before_recharge + before_gift;
let after_recharge = before_recharge - amount_usd;
let after_total = after_recharge + before_gift;
let after_total_refunded = before_total_refunded + amount_usd;
if !before_recharge.is_finite()
|| before_recharge < 0.0
|| !before_gift.is_finite()
|| before_gift < 0.0
|| !before_total_refunded.is_finite()
|| before_total_refunded < 0.0
|| !before_total.is_finite()
|| !after_recharge.is_finite()
|| !after_total.is_finite()
|| !after_total_refunded.is_finite()
{
return Ok(AdminWalletMutationOutcome::Invalid(
"wallet balance is invalid".to_string(),
));
}
if after_recharge < 0.0 {
return Ok(AdminWalletMutationOutcome::Invalid(
"refund amount exceeds refundable recharge balance".to_string(),
@@ -72,20 +98,41 @@ impl AppState {
"payment order not found".to_string(),
));
};
if amount_usd > order.refundable_amount_usd {
if order.wallet_id != wallet_id || order.status != "credited" {
return Ok(AdminWalletMutationOutcome::Invalid(
"refund amount exceeds refundable amount".to_string(),
"payment order is not refundable for this wallet".to_string(),
));
}
let order_amount = order.amount_usd;
let refunded_before = order.refunded_amount_usd;
let refundable_before = order.refundable_amount_usd;
let refunded_after = refunded_before + amount_usd;
let refundable_after = refundable_before - amount_usd;
if !payment_order_refund_amounts_are_consistent(
order_amount,
refunded_before,
refundable_before,
) || !refunded_after.is_finite()
|| !refundable_after.is_finite()
|| amount_usd > refundable_before
|| refunded_after < 0.0
|| refunded_after > order_amount
|| refundable_after < 0.0
|| refundable_after > order_amount
{
return Ok(AdminWalletMutationOutcome::Invalid(
"payment order refund amounts are invalid".to_string(),
));
}
let mut order = order;
order.refunded_amount_usd += amount_usd;
order.refundable_amount_usd -= amount_usd;
order.refunded_amount_usd = refunded_after;
order.refundable_amount_usd = refundable_after;
updated_order = Some(order);
}
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
updated_wallet.balance = after_recharge;
updated_wallet.total_refunded = (updated_wallet.total_refunded + amount_usd).max(0.0);
updated_wallet.total_refunded = after_total_refunded;
updated_wallet.updated_at_unix_secs = now_unix_secs;
let transaction = AdminWalletTransactionRecord {
@@ -95,7 +142,7 @@ impl AppState {
reason_code: "refund_out".to_string(),
amount: -amount_usd,
balance_before: before_total,
balance_after: after_recharge + before_gift,
balance_after: after_total,
recharge_balance_before: before_recharge,
recharge_balance_after: after_recharge,
gift_balance_before: before_gift,
@@ -130,6 +177,12 @@ impl AppState {
.expect("admin wallet payment order store should lock")
.insert(updated_order.id.clone(), updated_order);
}
if let Some(transaction_store) = self.admin_wallet_transaction_store.as_ref() {
transaction_store
.lock()
.expect("admin wallet transaction store should lock")
.insert(transaction.id.clone(), transaction.clone());
}
self.invalidate_auth_context_cache();
return Ok(AdminWalletMutationOutcome::Applied((
@@ -187,6 +240,23 @@ impl AppState {
else {
return Ok(AdminWalletMutationOutcome::NotFound);
};
if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 {
return Ok(AdminWalletMutationOutcome::Invalid(
"refund amount must be finite and greater than zero".to_string(),
));
}
if let (Some(existing_id), Some(incoming_id)) =
(refund.gateway_refund_id.as_deref(), gateway_refund_id)
{
if existing_id != incoming_id {
return Ok(AdminWalletMutationOutcome::Invalid(
"gateway refund identifier conflicts with existing evidence".to_string(),
));
}
}
if refund.status == "succeeded" {
return Ok(AdminWalletMutationOutcome::Applied(refund));
}
if refund.status != "processing" {
return Ok(AdminWalletMutationOutcome::Invalid(
"refund status must be processing before completion".to_string(),
@@ -195,9 +265,22 @@ impl AppState {
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
let mut updated_refund = refund;
updated_refund.status = "succeeded".to_string();
updated_refund.gateway_refund_id = gateway_refund_id.map(ToOwned::to_owned);
updated_refund.payout_reference = payout_reference.map(ToOwned::to_owned);
updated_refund.payout_proof = payout_proof;
updated_refund.gateway_refund_id = gateway_refund_id
.map(ToOwned::to_owned)
.or_else(|| updated_refund.gateway_refund_id.clone());
updated_refund.payout_reference = payout_reference
.map(ToOwned::to_owned)
.or_else(|| updated_refund.payout_reference.clone());
// Keep the durable provider response on ordinary retries. A
// terminal success proof is the only completion payload allowed
// to upgrade an earlier processing proof.
if updated_refund.payout_proof.is_none()
|| payout_proof
.as_ref()
.is_some_and(wallet_refund_proof_is_success)
{
updated_refund.payout_proof = payout_proof;
}
updated_refund.completed_at_unix_secs = Some(now_unix_secs);
updated_refund.updated_at_unix_secs = now_unix_secs;
refund_store
@@ -269,6 +352,12 @@ impl AppState {
return Ok(AdminWalletMutationOutcome::NotFound);
};
if !refund.amount_usd.is_finite() || refund.amount_usd <= 0.0 {
return Ok(AdminWalletMutationOutcome::Invalid(
"refund amount must be finite and greater than zero".to_string(),
));
}
let now_unix_secs = chrono::Utc::now().timestamp().max(0) as u64;
if matches!(refund.status.as_str(), "pending_approval" | "approved") {
let mut updated_refund = refund;
@@ -292,15 +381,95 @@ impl AppState {
)));
}
// The in-memory implementation mirrors the database contract:
// only an explicitly offline payout without external evidence may
// release its reservation.
if refund.gateway_refund_id.is_some()
|| refund.payout_proof.is_some()
|| !refund
.refund_mode
.trim()
.eq_ignore_ascii_case("offline_payout")
{
return Ok(AdminWalletMutationOutcome::Invalid(
"cannot fail refund while gateway settlement is processing".to_string(),
));
}
let amount_usd = refund.amount_usd;
let before_recharge = wallet.balance;
let before_gift = wallet.gift_balance;
let before_total_refunded = wallet.total_refunded;
let before_total = before_recharge + before_gift;
let after_recharge = before_recharge + amount_usd;
let after_total = after_recharge + before_gift;
let after_total_refunded = before_total_refunded - amount_usd;
if !before_recharge.is_finite()
|| before_recharge < 0.0
|| !before_gift.is_finite()
|| before_gift < 0.0
|| !before_total_refunded.is_finite()
|| before_total_refunded < 0.0
|| before_total_refunded < amount_usd
|| !before_total.is_finite()
|| !after_recharge.is_finite()
|| !after_total.is_finite()
|| !after_total_refunded.is_finite()
|| after_total_refunded < 0.0
{
return Ok(AdminWalletMutationOutcome::Invalid(
"wallet balance is invalid for refund recovery".to_string(),
));
}
let mut updated_order = None;
if let Some(payment_order_id) = refund.payment_order_id.clone() {
let Some(order_store) = self.admin_wallet_payment_order_store.as_ref() else {
return Ok(AdminWalletMutationOutcome::Unavailable);
};
let Some(order) = order_store
.lock()
.expect("admin wallet payment order store should lock")
.get(&payment_order_id)
.cloned()
else {
return Ok(AdminWalletMutationOutcome::Invalid(
"payment order not found".to_string(),
));
};
if order.wallet_id != wallet_id || order.status != "credited" {
return Ok(AdminWalletMutationOutcome::Invalid(
"payment order is not refundable for this wallet".to_string(),
));
}
let refunded_before = order.refunded_amount_usd;
let refundable_before = order.refundable_amount_usd;
let refunded_after = refunded_before - amount_usd;
let refundable_after = refundable_before + amount_usd;
if !payment_order_refund_amounts_are_consistent(
order.amount_usd,
refunded_before,
refundable_before,
) || !refunded_before.is_finite()
|| refunded_before < amount_usd
|| !refunded_after.is_finite()
|| refunded_after < 0.0
|| refundable_after < 0.0
|| refundable_after > order.amount_usd
{
return Ok(AdminWalletMutationOutcome::Invalid(
"payment order refund amounts are invalid".to_string(),
));
}
let mut order = order;
order.refunded_amount_usd = refunded_after;
order.refundable_amount_usd = refundable_after;
updated_order = Some(order);
}
let mut updated_wallet = wallet.clone();
updated_wallet.balance = after_recharge;
updated_wallet.total_refunded = (updated_wallet.total_refunded - amount_usd).max(0.0);
updated_wallet.total_refunded = after_total_refunded;
updated_wallet.updated_at_unix_secs = now_unix_secs;
let transaction = AdminWalletTransactionRecord {
@@ -310,7 +479,7 @@ impl AppState {
reason_code: "refund_revert".to_string(),
amount: amount_usd,
balance_before: before_total,
balance_after: after_recharge + before_gift,
balance_after: after_total,
recharge_balance_before: before_recharge,
recharge_balance_after: after_recharge,
gift_balance_before: before_gift,
@@ -322,23 +491,13 @@ impl AppState {
created_at_unix_ms: now_unix_secs,
};
if let Some(payment_order_id) = refund.payment_order_id.clone() {
let Some(order_store) = self.admin_wallet_payment_order_store.as_ref() else {
return Ok(AdminWalletMutationOutcome::Unavailable);
};
let maybe_order = order_store
if let Some(updated_order) = updated_order {
self.admin_wallet_payment_order_store
.as_ref()
.expect("admin wallet payment order store should exist")
.lock()
.expect("admin wallet payment order store should lock")
.get(&payment_order_id)
.cloned();
if let Some(mut order) = maybe_order {
order.refunded_amount_usd -= amount_usd;
order.refundable_amount_usd += amount_usd;
order_store
.lock()
.expect("admin wallet payment order store should lock")
.insert(order.id.clone(), order);
}
.insert(updated_order.id.clone(), updated_order);
}
let mut updated_refund = refund;
@@ -354,6 +513,12 @@ impl AppState {
.lock()
.expect("admin wallet refund store should lock")
.insert(updated_refund.id.clone(), updated_refund.clone());
if let Some(transaction_store) = self.admin_wallet_transaction_store.as_ref() {
transaction_store
.lock()
.expect("admin wallet transaction store should lock")
.insert(transaction.id.clone(), transaction.clone());
}
self.invalidate_auth_context_cache();
return Ok(AdminWalletMutationOutcome::Applied((
@@ -446,3 +611,103 @@ fn stored_admin_wallet_transaction_to_gateway(
created_at_unix_ms: transaction.created_at_unix_ms.unwrap_or_default(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn refund_with_proof(proof: serde_json::Value) -> AdminWalletRefundRecord {
AdminWalletRefundRecord {
id: "refund-proof-lifecycle".to_string(),
refund_no: "rf-proof-lifecycle".to_string(),
wallet_id: "wallet-proof-lifecycle".to_string(),
user_id: Some("user-proof-lifecycle".to_string()),
payment_order_id: None,
source_type: "manual".to_string(),
source_id: None,
refund_mode: "original_channel".to_string(),
amount_usd: 4.0,
status: "processing".to_string(),
reason: Some("proof lifecycle regression".to_string()),
failure_reason: None,
gateway_refund_id: Some("gateway-proof-lifecycle".to_string()),
payout_method: Some("wxpay".to_string()),
payout_reference: None,
payout_proof: Some(proof),
requested_by: Some("user-proof-lifecycle".to_string()),
approved_by: Some("admin-proof-lifecycle".to_string()),
processed_by: Some("admin-proof-lifecycle".to_string()),
created_at_unix_ms: 1_710_000_000,
updated_at_unix_secs: 1_710_000_000,
processed_at_unix_secs: Some(1_710_000_010),
completed_at_unix_secs: None,
}
}
#[tokio::test]
async fn completion_preserves_processing_proof_on_non_terminal_retry() {
let processing_proof = json!({
"gateway": "wxpay",
"id": "gateway-proof-lifecycle",
"status": "processing"
});
let state = AppState::new()
.expect("gateway state should build")
.with_admin_wallet_refunds_for_tests([refund_with_proof(processing_proof.clone())]);
let outcome = state
.admin_complete_wallet_refund(
"wallet-proof-lifecycle",
"refund-proof-lifecycle",
Some("gateway-proof-lifecycle"),
None,
Some(json!({
"gateway": "wxpay",
"id": "gateway-proof-lifecycle",
"status": "pending",
"attempt": 2
})),
)
.await
.expect("completion should resolve");
let AdminWalletMutationOutcome::Applied(refund) = outcome else {
panic!("completion should apply");
};
assert_eq!(refund.status, "succeeded");
assert_eq!(refund.payout_proof, Some(processing_proof));
}
#[tokio::test]
async fn completion_allows_terminal_success_proof_to_upgrade_processing_evidence() {
let state = AppState::new()
.expect("gateway state should build")
.with_admin_wallet_refunds_for_tests([refund_with_proof(json!({
"gateway": "wxpay",
"id": "gateway-proof-lifecycle",
"status": "processing"
}))]);
let success_proof = json!({
"gateway": "wxpay",
"id": "gateway-proof-lifecycle",
"status": "succeeded",
"processed_at": "2026-08-29T00:00:00Z"
});
let outcome = state
.admin_complete_wallet_refund(
"wallet-proof-lifecycle",
"refund-proof-lifecycle",
Some("gateway-proof-lifecycle"),
None,
Some(success_proof.clone()),
)
.await
.expect("completion should resolve");
let AdminWalletMutationOutcome::Applied(refund) = outcome else {
panic!("completion should apply");
};
assert_eq!(refund.status, "succeeded");
assert_eq!(refund.payout_proof, Some(success_proof));
}
}
+168 -42
View File
@@ -9,13 +9,49 @@ use aether_data_contracts::repository::usage::{UsageReadRepository, UsageReposit
use aether_data_contracts::repository::video_tasks::{
VideoTaskReadRepository, VideoTaskRepository,
};
use hmac::Mac;
use serde_json::json;
use sha2::{Digest, Sha256};
use super::{AppState, FrontdoorRuntimeGuardConfig, GatewayDataState};
use crate::{provider_transport, usage};
fn auth_email_storage_key_digest_for_tests(domain: &str, parts: &[&str]) -> String {
let secret = std::env::var("JWT_SECRET_KEY")
.ok()
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
.unwrap_or_else(|| "aether-rust-test-jwt-secret-32-bytes-minimum".to_string());
let mut mac = hmac::Hmac::<Sha256>::new_from_slice(secret.as_bytes())
.expect("HMAC should accept the test auth key");
mac.update(b"aether-auth-email-storage-v1\0");
mac.update(domain.as_bytes());
for part in parts {
mac.update(b"\0");
mac.update(part.as_bytes());
}
mac.finalize()
.into_bytes()
.iter()
.map(|byte| format!("{byte:02x}"))
.collect()
}
#[cfg(test)]
impl AppState {
pub(crate) fn with_internal_gateway_auth_secret_for_tests(mut self, secret: &str) -> Self {
self.internal_gateway_auth = Arc::new(
crate::internal_gateway_auth::InternalGatewayAuthConfig::with_secret_for_tests(secret),
);
self
}
pub(crate) fn without_internal_gateway_for_tests(mut self) -> Self {
self.internal_gateway_auth =
Arc::new(crate::internal_gateway_auth::InternalGatewayAuthConfig::disabled_for_tests());
self
}
pub(crate) fn with_data_state_for_tests(mut self, data_state: GatewayDataState) -> Self {
// Request-execution fixtures provide candidate and provider data but
// bypass the production startup bootstrap that creates the enabled
@@ -71,6 +107,38 @@ impl AppState {
self
}
pub(crate) fn with_tunnel_identity_and_relay_secret_for_tests(
mut self,
instance_id: &str,
relay_base_url: Option<&str>,
relay_auth_secret: &str,
) -> Self {
self.tunnel = crate::tunnel::EmbeddedTunnelState::with_data_and_directory_for_tests(
Arc::clone(&self.data),
crate::tunnel::TunnelAttachmentDirectory::for_tests(instance_id, relay_base_url, 90),
relay_auth_secret,
);
self
}
pub(crate) fn with_tunnel_identity_runtime_state_and_relay_secret_for_tests(
mut self,
instance_id: &str,
relay_base_url: Option<&str>,
runtime_state: Arc<aether_runtime_state::RuntimeState>,
relay_auth_secret: &str,
) -> Self {
self.tunnel =
crate::tunnel::EmbeddedTunnelState::with_data_identity_runtime_state_and_relay_secret_for_tests(
Arc::clone(&self.data),
instance_id,
relay_base_url,
runtime_state,
relay_auth_secret,
);
self
}
pub(crate) fn with_video_task_data_reader_for_tests(
mut self,
repository: Arc<dyn VideoTaskReadRepository>,
@@ -239,16 +307,22 @@ impl AppState {
nonce: &str,
payload: serde_json::Value,
) -> Self {
let key = aether_data::repository::provider_oauth::provider_oauth_state_storage_key(nonce);
let plaintext = payload.to_string();
let purpose = format!("provider-oauth-state:{key}");
let sealed =
crate::handlers::shared::seal_runtime_secret_payload(&self, &purpose, &plaintext)
.expect("test provider OAuth state should seal");
let store = self
.provider_oauth_state_store
.get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new())));
store
.lock()
.expect("provider oauth state store should lock")
.insert(format!("provider_oauth_state:{nonce}"), payload.to_string());
.insert(key.clone(), plaintext);
self.runtime_state.kv_set_local_nowait(
&format!("provider_oauth_state:{nonce}"),
payload.to_string(),
&key,
sealed,
Some(Duration::from_secs(
aether_data::repository::provider_oauth::PROVIDER_OAUTH_STATE_TTL_SECS,
)),
@@ -259,23 +333,43 @@ impl AppState {
pub(crate) fn with_provider_oauth_device_session_entry_for_tests(
mut self,
session_id: &str,
payload: serde_json::Value,
mut payload: serde_json::Value,
) -> Self {
if let Some(payload) = payload.as_object_mut() {
payload
.entry("session_id".to_string())
.or_insert_with(|| json!(session_id));
payload
.entry("initiated_by_user_id".to_string())
.or_insert_with(|| json!("admin-user-123"));
payload
.entry("initiated_by_session_id".to_string())
.or_insert_with(|| json!("session-123"));
payload
.entry("initiated_by_management_token_id".to_string())
.or_insert_with(|| json!("management-token-123"));
}
let key =
aether_data::repository::provider_oauth::provider_oauth_device_session_storage_key(
session_id,
);
let plaintext = payload.to_string();
let purpose =
aether_data::repository::provider_oauth::provider_oauth_device_session_secret_purpose(
session_id,
);
let sealed =
crate::handlers::shared::seal_runtime_secret_payload(&self, &purpose, &plaintext)
.expect("test provider OAuth device session should seal");
let store = self
.provider_oauth_device_session_store
.get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new())));
store
.lock()
.expect("provider oauth device session store should lock")
.insert(
format!("device_auth_session:{session_id}"),
payload.to_string(),
);
self.runtime_state.kv_set_local_nowait(
&format!("device_auth_session:{session_id}"),
payload.to_string(),
Some(Duration::from_secs(3600)),
);
.insert(key.clone(), plaintext);
self.runtime_state
.kv_set_local_nowait(&key, sealed, Some(Duration::from_secs(3600)));
self
}
@@ -284,19 +378,26 @@ impl AppState {
task_id: &str,
payload: serde_json::Value,
) -> Self {
let key =
aether_data::repository::provider_oauth::provider_oauth_batch_task_storage_key(task_id);
let plaintext = payload.to_string();
let purpose =
aether_data::repository::provider_oauth::provider_oauth_batch_task_secret_purpose(
task_id,
);
let sealed =
crate::handlers::shared::seal_runtime_secret_payload(&self, &purpose, &plaintext)
.expect("test provider OAuth batch task should seal");
let store = self
.provider_oauth_batch_task_store
.get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new())));
store
.lock()
.expect("provider oauth batch task store should lock")
.insert(
format!("provider_oauth_batch_task:{task_id}"),
payload.to_string(),
);
.insert(key.clone(), plaintext);
self.runtime_state.kv_set_local_nowait(
&format!("provider_oauth_batch_task:{task_id}"),
payload.to_string(),
&key,
sealed,
Some(Duration::from_secs(
aether_data::repository::provider_oauth::PROVIDER_OAUTH_BATCH_TASK_TTL_SECS,
)),
@@ -346,6 +447,11 @@ impl AppState {
self
}
pub(crate) fn without_auth_session_store_for_tests(mut self) -> Self {
self.auth_session_store = None;
self
}
pub(crate) fn without_auth_user_model_capability_store_for_tests(mut self) -> Self {
self.auth_user_model_capability_store = None;
self
@@ -601,47 +707,67 @@ impl AppState {
mut self,
email: &str,
code: &str,
verification_token: &str,
created_at: chrono::DateTime<chrono::Utc>,
) -> Self {
let code_hash = format!(
"{:x}",
Sha256::digest(
format!(
"aether-email-verification\0{}\0{}",
verification_token.trim(),
code.trim()
)
.as_bytes()
)
);
let verification_token_hash =
format!("{:x}", Sha256::digest(verification_token.trim().as_bytes()));
let normalized_email = email.trim().to_ascii_lowercase();
let key = format!(
"email:verification:{}",
auth_email_storage_key_digest_for_tests("pending", &[normalized_email.as_str()])
);
let value = json!({
"code_hash": code_hash,
"created_at": created_at.to_rfc3339(),
"verification_token_hash": verification_token_hash,
})
.to_string();
let store = self
.auth_email_verification_store
.get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new())));
store
.lock()
.expect("auth email verification store should lock")
.insert(
format!("email:verification:{}", email.trim().to_ascii_lowercase()),
json!({
"code": code,
"created_at": created_at.to_rfc3339(),
})
.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)),
);
.insert(key.clone(), value.clone());
self.runtime_state
.kv_set_local_nowait(&key, value, Some(Duration::from_secs(600)));
self
}
pub(crate) fn with_auth_email_verified_for_tests(mut self, email: &str) -> Self {
pub(crate) fn with_auth_email_verified_for_tests(
mut self,
email: &str,
verification_token: &str,
) -> Self {
let normalized_email = email.trim().to_ascii_lowercase();
let key = format!(
"email:verified:{}",
auth_email_storage_key_digest_for_tests(
"registration-proof",
&[normalized_email.as_str(), verification_token.trim()]
)
);
let store = self
.auth_email_verification_store
.get_or_insert_with(|| Arc::new(StdMutex::new(HashMap::new())));
store
.lock()
.expect("auth email verification store should lock")
.insert(
format!("email:verified:{}", email.trim().to_ascii_lowercase()),
"verified".to_string(),
);
.insert(key.clone(), "verified".to_string());
self.runtime_state.kv_set_local_nowait(
&format!("email:verified:{}", email.trim().to_ascii_lowercase()),
&key,
"verified".to_string(),
Some(Duration::from_secs(3600)),
);
+86 -3
View File
@@ -61,7 +61,7 @@ pub(crate) enum AdminWalletMutationOutcome<T> {
Unavailable,
}
#[derive(Debug, Clone, PartialEq)]
#[derive(Clone, PartialEq)]
pub(crate) struct GatewayUserSessionView {
pub(crate) id: String,
pub(crate) user_id: String,
@@ -78,6 +78,7 @@ pub(crate) struct GatewayUserSessionView {
pub(crate) user_agent: Option<String>,
pub(crate) created_at: Option<chrono::DateTime<chrono::Utc>>,
pub(crate) updated_at: Option<chrono::DateTime<chrono::Utc>>,
pub(crate) security_version: i64,
}
impl GatewayUserSessionView {
@@ -131,9 +132,18 @@ impl GatewayUserSessionView {
user_agent,
created_at,
updated_at,
security_version: 0,
})
}
pub(crate) fn with_security_version(mut self, security_version: i64) -> Result<Self, String> {
if security_version < 0 {
return Err("user_sessions.security_version is negative".to_string());
}
self.security_version = security_version;
Ok(self)
}
pub(crate) fn hash_refresh_token(token: &str) -> String {
use sha2::Digest;
@@ -157,8 +167,10 @@ impl GatewayUserSessionView {
let Some(rotated_at) = self.rotated_at else {
return (false, false);
};
let age = now.signed_duration_since(rotated_at);
if prev_hash == &token_hash
&& now.signed_duration_since(rotated_at).num_seconds() <= Self::REFRESH_GRACE_SECONDS
&& age >= chrono::Duration::zero()
&& age <= chrono::Duration::seconds(Self::REFRESH_GRACE_SECONDS)
{
return (true, true);
}
@@ -183,6 +195,48 @@ impl GatewayUserSessionView {
}
}
#[cfg(test)]
mod gateway_user_session_view_tests {
use super::GatewayUserSessionView;
use chrono::{Duration, Utc};
fn session_with_rotation(rotated_at: chrono::DateTime<Utc>) -> GatewayUserSessionView {
GatewayUserSessionView::new(
"session-1".to_string(),
"user-1".to_string(),
"device-1".to_string(),
None,
GatewayUserSessionView::hash_refresh_token("current-token"),
Some(GatewayUserSessionView::hash_refresh_token("previous-token")),
Some(rotated_at),
None,
None,
None,
None,
None,
None,
None,
None,
)
.expect("session should build")
}
#[test]
fn refresh_grace_rejects_future_rotation_timestamps() {
let now = Utc::now();
assert_eq!(
session_with_rotation(now - Duration::seconds(1))
.verify_refresh_token("previous-token", now),
(true, true)
);
assert_eq!(
session_with_rotation(now + Duration::milliseconds(1))
.verify_refresh_token("previous-token", now),
(false, false)
);
}
}
impl From<crate::data::state::StoredUserSessionRecord> for GatewayUserSessionView {
fn from(value: crate::data::state::StoredUserSessionRecord) -> Self {
Self {
@@ -201,6 +255,7 @@ impl From<crate::data::state::StoredUserSessionRecord> for GatewayUserSessionVie
user_agent: value.user_agent,
created_at: value.created_at,
updated_at: value.updated_at,
security_version: value.security_version,
}
}
}
@@ -229,6 +284,7 @@ impl From<GatewayUserSessionView> for crate::data::state::StoredUserSessionRecor
user_agent: value.user_agent,
created_at: value.created_at,
updated_at: value.updated_at,
security_version: value.security_version,
}
}
}
@@ -308,7 +364,7 @@ impl From<GatewayUserPreferenceView> for crate::data::state::StoredUserPreferenc
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub(crate) struct GatewayAdminPaymentCallbackView {
pub(crate) id: String,
pub(crate) payment_order_id: Option<String>,
@@ -325,6 +381,33 @@ pub(crate) struct GatewayAdminPaymentCallbackView {
pub(crate) processed_at_unix_secs: Option<u64>,
}
impl std::fmt::Debug for GatewayAdminPaymentCallbackView {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("GatewayAdminPaymentCallbackView")
.field("id", &self.id)
.field("payment_order_id", &self.payment_order_id)
.field("payment_method", &self.payment_method)
.field("callback_key", &"[REDACTED]")
.field("order_no", &self.order_no)
.field("gateway_order_id", &self.gateway_order_id)
.field(
"payload_hash",
&self.payload_hash.as_ref().map(|_| "[REDACTED]"),
)
.field("signature_valid", &self.signature_valid)
.field("status", &self.status)
.field("payload", &self.payload.as_ref().map(|_| "[REDACTED]"))
.field(
"error_message",
&self.error_message.as_ref().map(|_| "[REDACTED]"),
)
.field("created_at_unix_ms", &self.created_at_unix_ms)
.field("processed_at_unix_secs", &self.processed_at_unix_secs)
.finish()
}
}
impl From<super::AdminPaymentCallbackRecord> for GatewayAdminPaymentCallbackView {
fn from(value: super::AdminPaymentCallbackRecord) -> Self {
Self {
+124 -1
View File
@@ -6,6 +6,13 @@ use aether_data_contracts::repository::video_tasks::{
VideoTaskQueryFilter, VideoTaskStatusCount,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum VideoTaskRouteAccess {
Allowed,
NotFound,
Denied,
}
impl AppState {
pub(crate) async fn read_data_backed_video_task_response(
&self,
@@ -18,6 +25,18 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn read_data_backed_video_task_response_for_user(
&self,
route_family: Option<&str>,
request_path: &str,
user_id: &str,
) -> Result<Option<video_tasks::LocalVideoTaskReadResponse>, GatewayError> {
self.data
.read_video_task_response_for_user(route_family, request_path, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_video_task_by_id(
&self,
task_id: &str,
@@ -38,12 +57,72 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_video_task_by_id_for_user(
&self,
task_id: &str,
user_id: &str,
) -> Result<Option<StoredVideoTask>, GatewayError> {
self.data
.find_video_task_for_user(VideoTaskLookupKey::Id(task_id), user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn find_video_task_by_short_id_for_user(
&self,
short_id: &str,
user_id: &str,
) -> Result<Option<StoredVideoTask>, GatewayError> {
self.data
.find_video_task_for_user(VideoTaskLookupKey::ShortId(short_id), user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(crate) async fn upsert_video_task_snapshot(
&self,
snapshot: &video_tasks::LocalVideoTaskSnapshot,
) -> Result<Option<StoredVideoTask>, GatewayError> {
let mut record = snapshot.to_upsert_record();
// Reconstructed snapshots intentionally omit sensitive/request-only fields. Preserve the
// persisted row's immutable identity and request-shape scalars before writing lifecycle
// changes back, so the repository can continue enforcing immutable-field integrity.
let existing_by_id = self
.data
.find_video_task(VideoTaskLookupKey::Id(record.id.as_str()))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?;
let existing = if existing_by_id.is_some() {
existing_by_id
} else if let Some(short_id) = record.short_id.as_deref() {
self.data
.find_video_task(VideoTaskLookupKey::ShortId(short_id))
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
} else {
None
};
if let Some(existing) = existing {
record.id = existing.id;
record.short_id = existing.short_id;
record.request_id = existing.request_id;
record.user_id = existing.user_id;
record.api_key_id = existing.api_key_id;
record.external_task_id = existing.external_task_id;
record.provider_id = existing.provider_id;
record.endpoint_id = existing.endpoint_id;
record.key_id = existing.key_id;
record.client_api_format = existing.client_api_format;
record.provider_api_format = existing.provider_api_format;
record.format_converted = existing.format_converted;
record.model = existing.model;
record.duration_seconds = existing.duration_seconds;
record.resolution = existing.resolution;
record.aspect_ratio = existing.aspect_ratio;
record.size = existing.size;
}
self.data
.upsert_video_task(snapshot.to_upsert_record())
.upsert_video_task(record)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))
}
@@ -77,6 +156,50 @@ impl AppState {
Ok(true)
}
pub(crate) async fn hydrate_video_task_for_route_for_user(
&self,
route_family: Option<&str>,
request_path: &str,
user_id: &str,
) -> Result<VideoTaskRouteAccess, GatewayError> {
let user_id = user_id.trim();
if user_id.is_empty() {
return Ok(VideoTaskRouteAccess::Denied);
}
let Some(lookup) =
video_tasks::resolve_video_task_hydration_lookup_key(route_family, request_path)
else {
return Ok(VideoTaskRouteAccess::NotFound);
};
if let Some(task) = self
.data
.find_video_task_for_user(lookup, user_id)
.await
.map_err(|err| GatewayError::Internal(err.to_string()))?
{
if !self.video_tasks.hydrate_from_stored_task(&task) {
if let Some(snapshot) = self.reconstruct_video_task_snapshot(&task).await? {
self.video_tasks.record_snapshot(snapshot);
}
}
return Ok(VideoTaskRouteAccess::Allowed);
}
Ok(
match self
.video_tasks
.snapshot_for_route(route_family, request_path)
{
Some(snapshot) if snapshot.belongs_to_user(user_id) => {
VideoTaskRouteAccess::Allowed
}
Some(_) => VideoTaskRouteAccess::Denied,
None => VideoTaskRouteAccess::NotFound,
},
)
}
pub(crate) async fn reconstruct_video_task_snapshot(
&self,
task: &StoredVideoTask,