mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
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:
File diff suppressed because it is too large
Load Diff
@@ -4,8 +4,8 @@ pub use aether_data_contracts::repository::auth::{
|
||||
read_resolved_auth_api_key_snapshot, read_resolved_auth_api_key_snapshot_by_key_hash,
|
||||
read_resolved_auth_api_key_snapshot_by_user_api_key_ids, AuthApiKeyExportSummary,
|
||||
AuthApiKeyLookupKey, AuthApiKeyReadRepository, AuthApiKeyWriteRepository, AuthRepository,
|
||||
CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord, ResolvedAuthApiKeySnapshot,
|
||||
ResolvedAuthApiKeySnapshotReader, StandaloneApiKeyExportListQuery,
|
||||
CompareAndSwapAuthApiKeyCiphertext, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord,
|
||||
ResolvedAuthApiKeySnapshot, ResolvedAuthApiKeySnapshotReader, StandaloneApiKeyExportListQuery,
|
||||
StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot, UpdateStandaloneApiKeyBasicRecord,
|
||||
UpdateUserApiKeyBasicRecord,
|
||||
};
|
||||
|
||||
@@ -3,8 +3,8 @@ use std::sync::RwLock;
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::{
|
||||
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
|
||||
StoredOAuthProviderModuleConfig,
|
||||
AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult,
|
||||
LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
|
||||
@@ -49,15 +49,73 @@ impl AuthModuleReadRepository for InMemoryAuthModuleReadRepository {
|
||||
|
||||
#[async_trait]
|
||||
impl AuthModuleWriteRepository for InMemoryAuthModuleReadRepository {
|
||||
async fn upsert_ldap_config(
|
||||
async fn compare_and_swap_ldap_config(
|
||||
&self,
|
||||
config: &StoredLdapModuleConfig,
|
||||
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
self.ldap_config
|
||||
expected: Option<&StoredLdapModuleConfig>,
|
||||
replacement: &StoredLdapModuleConfig,
|
||||
bind_password_update: &LdapBindPasswordUpdate,
|
||||
) -> Result<CompareAndSwapLdapConfigResult, DataLayerError> {
|
||||
let mut config = self
|
||||
.ldap_config
|
||||
.write()
|
||||
.expect("auth module ldap repository lock")
|
||||
.replace(config.clone());
|
||||
Ok(Some(config.clone()))
|
||||
.expect("auth module ldap repository lock");
|
||||
if config.as_ref() != expected {
|
||||
return Ok(CompareAndSwapLdapConfigResult::Conflict);
|
||||
}
|
||||
|
||||
let bind_password_encrypted = match bind_password_update {
|
||||
LdapBindPasswordUpdate::Preserve => expected
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidConfiguration(
|
||||
"LDAP bind password cannot be preserved while creating the singleton"
|
||||
.to_string(),
|
||||
)
|
||||
})?
|
||||
.bind_password_encrypted
|
||||
.clone(),
|
||||
LdapBindPasswordUpdate::Set(ciphertext) => Some(ciphertext.clone()),
|
||||
LdapBindPasswordUpdate::Clear => None,
|
||||
};
|
||||
let persisted = StoredLdapModuleConfig {
|
||||
bind_password_encrypted,
|
||||
..replacement.clone()
|
||||
};
|
||||
*config = Some(persisted.clone());
|
||||
Ok(CompareAndSwapLdapConfigResult::Applied(persisted))
|
||||
}
|
||||
|
||||
async fn delete_ldap_config_if_matches(
|
||||
&self,
|
||||
expected: &StoredLdapModuleConfig,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut config = self
|
||||
.ldap_config
|
||||
.write()
|
||||
.expect("auth module ldap repository lock");
|
||||
if config.as_ref() != Some(expected) {
|
||||
return Ok(false);
|
||||
}
|
||||
config.take();
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn compare_and_swap_ldap_bind_password(
|
||||
&self,
|
||||
expected: &str,
|
||||
replacement: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut config = self
|
||||
.ldap_config
|
||||
.write()
|
||||
.expect("auth module ldap repository lock");
|
||||
let Some(config) = config.as_mut() else {
|
||||
return Ok(false);
|
||||
};
|
||||
if config.bind_password_encrypted.as_deref() != Some(expected) {
|
||||
return Ok(false);
|
||||
}
|
||||
config.bind_password_encrypted = Some(replacement.to_string());
|
||||
Ok(true)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -65,9 +123,27 @@ impl AuthModuleWriteRepository for InMemoryAuthModuleReadRepository {
|
||||
mod tests {
|
||||
use super::InMemoryAuthModuleReadRepository;
|
||||
use crate::repository::auth_modules::{
|
||||
AuthModuleReadRepository, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
|
||||
AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult,
|
||||
LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
|
||||
};
|
||||
|
||||
fn ldap_config() -> StoredLdapModuleConfig {
|
||||
StoredLdapModuleConfig {
|
||||
server_url: "ldaps://ldap.example.com".to_string(),
|
||||
bind_dn: "cn=admin,dc=example,dc=com".to_string(),
|
||||
bind_password_encrypted: Some("encrypted-password".to_string()),
|
||||
base_dn: "dc=example,dc=com".to_string(),
|
||||
user_search_filter: Some("(uid={username})".to_string()),
|
||||
username_attr: Some("uid".to_string()),
|
||||
email_attr: Some("mail".to_string()),
|
||||
display_name_attr: Some("displayName".to_string()),
|
||||
is_enabled: true,
|
||||
is_exclusive: false,
|
||||
use_starttls: true,
|
||||
connect_timeout: Some(10),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn reads_seeded_auth_module_configs() {
|
||||
let repository = InMemoryAuthModuleReadRepository::seed(
|
||||
@@ -79,20 +155,7 @@ mod tests {
|
||||
"https://example.com/callback".to_string(),
|
||||
)
|
||||
.expect("oauth provider should build")],
|
||||
Some(StoredLdapModuleConfig {
|
||||
server_url: "ldaps://ldap.example.com".to_string(),
|
||||
bind_dn: "cn=admin,dc=example,dc=com".to_string(),
|
||||
bind_password_encrypted: Some("encrypted-password".to_string()),
|
||||
base_dn: "dc=example,dc=com".to_string(),
|
||||
user_search_filter: Some("(uid={username})".to_string()),
|
||||
username_attr: Some("uid".to_string()),
|
||||
email_attr: Some("mail".to_string()),
|
||||
display_name_attr: Some("displayName".to_string()),
|
||||
is_enabled: true,
|
||||
is_exclusive: false,
|
||||
use_starttls: true,
|
||||
connect_timeout: Some(10),
|
||||
}),
|
||||
Some(ldap_config()),
|
||||
);
|
||||
|
||||
let oauth = repository
|
||||
@@ -111,4 +174,187 @@ mod tests {
|
||||
"ldaps://ldap.example.com"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ldap_compensation_delete_requires_an_exact_match() {
|
||||
let expected = ldap_config();
|
||||
let repository = InMemoryAuthModuleReadRepository::seed(Vec::new(), Some(expected.clone()));
|
||||
let mismatched = StoredLdapModuleConfig {
|
||||
is_enabled: false,
|
||||
..expected.clone()
|
||||
};
|
||||
|
||||
assert!(!repository
|
||||
.delete_ldap_config_if_matches(&mismatched)
|
||||
.await
|
||||
.expect("mismatched delete should execute"));
|
||||
assert!(repository
|
||||
.delete_ldap_config_if_matches(&expected)
|
||||
.await
|
||||
.expect("matching delete should execute"));
|
||||
assert!(repository
|
||||
.get_ldap_config()
|
||||
.await
|
||||
.expect("LDAP config should remain readable")
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ldap_compare_and_swap_separates_preserve_set_and_clear() {
|
||||
let original = ldap_config();
|
||||
let repository = InMemoryAuthModuleReadRepository::seed(Vec::new(), Some(original.clone()));
|
||||
let replacement = StoredLdapModuleConfig {
|
||||
server_url: "ldap://updated.example.com".to_string(),
|
||||
bind_password_encrypted: Some("stale-ciphertext-must-be-ignored".to_string()),
|
||||
..original.clone()
|
||||
};
|
||||
|
||||
let preserved = repository
|
||||
.compare_and_swap_ldap_config(
|
||||
Some(&original),
|
||||
&replacement,
|
||||
&LdapBindPasswordUpdate::Preserve,
|
||||
)
|
||||
.await
|
||||
.expect("preserve CAS should execute");
|
||||
let CompareAndSwapLdapConfigResult::Applied(preserved) = preserved else {
|
||||
panic!("fresh snapshot should apply");
|
||||
};
|
||||
assert_eq!(
|
||||
preserved.bind_password_encrypted.as_deref(),
|
||||
Some("encrypted-password")
|
||||
);
|
||||
|
||||
let set = repository
|
||||
.compare_and_swap_ldap_config(
|
||||
Some(&preserved),
|
||||
&preserved,
|
||||
&LdapBindPasswordUpdate::Set("rotated-ciphertext".to_string()),
|
||||
)
|
||||
.await
|
||||
.expect("set CAS should execute");
|
||||
let CompareAndSwapLdapConfigResult::Applied(set) = set else {
|
||||
panic!("fresh snapshot should apply");
|
||||
};
|
||||
assert_eq!(
|
||||
set.bind_password_encrypted.as_deref(),
|
||||
Some("rotated-ciphertext")
|
||||
);
|
||||
|
||||
let cleared = repository
|
||||
.compare_and_swap_ldap_config(Some(&set), &set, &LdapBindPasswordUpdate::Clear)
|
||||
.await
|
||||
.expect("clear CAS should execute");
|
||||
let CompareAndSwapLdapConfigResult::Applied(cleared) = cleared else {
|
||||
panic!("fresh snapshot should apply");
|
||||
};
|
||||
assert!(cleared.bind_password_encrypted.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ldap_compare_and_swap_rejects_stale_password_and_config_snapshots() {
|
||||
let original = ldap_config();
|
||||
let repository = InMemoryAuthModuleReadRepository::seed(Vec::new(), Some(original.clone()));
|
||||
assert!(repository
|
||||
.compare_and_swap_ldap_bind_password("encrypted-password", "rotated-ciphertext")
|
||||
.await
|
||||
.expect("password rotation should execute"));
|
||||
|
||||
let stale_password_result = repository
|
||||
.compare_and_swap_ldap_config(
|
||||
Some(&original),
|
||||
&StoredLdapModuleConfig {
|
||||
base_dn: "dc=updated,dc=example".to_string(),
|
||||
..original.clone()
|
||||
},
|
||||
&LdapBindPasswordUpdate::Preserve,
|
||||
)
|
||||
.await
|
||||
.expect("stale password CAS should execute");
|
||||
assert_eq!(
|
||||
stale_password_result,
|
||||
CompareAndSwapLdapConfigResult::Conflict
|
||||
);
|
||||
assert_eq!(
|
||||
repository
|
||||
.get_ldap_config()
|
||||
.await
|
||||
.expect("LDAP config should load")
|
||||
.and_then(|config| config.bind_password_encrypted)
|
||||
.as_deref(),
|
||||
Some("rotated-ciphertext")
|
||||
);
|
||||
|
||||
let current = repository
|
||||
.get_ldap_config()
|
||||
.await
|
||||
.expect("LDAP config should load")
|
||||
.expect("LDAP config should exist");
|
||||
let changed = StoredLdapModuleConfig {
|
||||
is_enabled: false,
|
||||
..current.clone()
|
||||
};
|
||||
let applied = repository
|
||||
.compare_and_swap_ldap_config(
|
||||
Some(¤t),
|
||||
&changed,
|
||||
&LdapBindPasswordUpdate::Preserve,
|
||||
)
|
||||
.await
|
||||
.expect("fresh config CAS should execute");
|
||||
assert!(matches!(
|
||||
applied,
|
||||
CompareAndSwapLdapConfigResult::Applied(_)
|
||||
));
|
||||
let stale_config_result = repository
|
||||
.compare_and_swap_ldap_config(
|
||||
Some(¤t),
|
||||
¤t,
|
||||
&LdapBindPasswordUpdate::Preserve,
|
||||
)
|
||||
.await
|
||||
.expect("stale config CAS should execute");
|
||||
assert_eq!(
|
||||
stale_config_result,
|
||||
CompareAndSwapLdapConfigResult::Conflict
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ldap_compare_and_swap_allows_only_one_initial_create() {
|
||||
let repository = InMemoryAuthModuleReadRepository::default();
|
||||
let replacement = StoredLdapModuleConfig {
|
||||
bind_password_encrypted: None,
|
||||
..ldap_config()
|
||||
};
|
||||
|
||||
let first = repository
|
||||
.compare_and_swap_ldap_config(
|
||||
None,
|
||||
&replacement,
|
||||
&LdapBindPasswordUpdate::Set("first-ciphertext".to_string()),
|
||||
)
|
||||
.await
|
||||
.expect("first create should execute");
|
||||
assert!(matches!(first, CompareAndSwapLdapConfigResult::Applied(_)));
|
||||
|
||||
let second = repository
|
||||
.compare_and_swap_ldap_config(
|
||||
None,
|
||||
&replacement,
|
||||
&LdapBindPasswordUpdate::Set("second-ciphertext".to_string()),
|
||||
)
|
||||
.await
|
||||
.expect("second create should execute");
|
||||
assert_eq!(second, CompareAndSwapLdapConfigResult::Conflict);
|
||||
assert_eq!(
|
||||
repository
|
||||
.get_ldap_config()
|
||||
.await
|
||||
.expect("LDAP config should load")
|
||||
.and_then(|config| config.bind_password_encrypted)
|
||||
.as_deref(),
|
||||
Some("first-ciphertext")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
mod memory;
|
||||
|
||||
pub use aether_data_contracts::repository::auth_modules::{
|
||||
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
|
||||
StoredOAuthProviderModuleConfig,
|
||||
AuthModuleReadRepository, AuthModuleWriteRepository, CompareAndSwapLdapConfigResult,
|
||||
LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig,
|
||||
};
|
||||
#[cfg(feature = "mysql")]
|
||||
pub use aether_data_mysql::{MysqlAuthModuleReadRepository, MysqlAuthModuleRepository};
|
||||
|
||||
@@ -53,7 +53,8 @@ impl InMemoryBackgroundTaskRepository {
|
||||
I: IntoIterator<Item = StoredBackgroundTaskRun>,
|
||||
{
|
||||
let mut index = InMemoryBackgroundTaskIndex::default();
|
||||
for run in runs {
|
||||
for mut run in runs {
|
||||
run.sanitize_persisted_data();
|
||||
index.runs.insert(run.id.clone(), run);
|
||||
}
|
||||
Self {
|
||||
@@ -164,8 +165,9 @@ impl BackgroundTaskReadRepository for InMemoryBackgroundTaskRepository {
|
||||
impl BackgroundTaskWriteRepository for InMemoryBackgroundTaskRepository {
|
||||
async fn upsert_run(
|
||||
&self,
|
||||
run: UpsertBackgroundTaskRun,
|
||||
mut run: UpsertBackgroundTaskRun,
|
||||
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
|
||||
run.sanitize_for_persistence();
|
||||
run.validate()?;
|
||||
let stored = run.into_stored();
|
||||
self.index
|
||||
@@ -192,8 +194,9 @@ impl BackgroundTaskWriteRepository for InMemoryBackgroundTaskRepository {
|
||||
|
||||
async fn upsert_event(
|
||||
&self,
|
||||
event: UpsertBackgroundTaskEvent,
|
||||
mut event: UpsertBackgroundTaskEvent,
|
||||
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
|
||||
event.sanitize_for_persistence();
|
||||
event.validate()?;
|
||||
let stored = event.into_stored();
|
||||
let mut guard = self.index.write().expect("background task repository lock");
|
||||
|
||||
@@ -5,8 +5,9 @@ use async_trait::async_trait;
|
||||
|
||||
use super::{
|
||||
AdminBillingMutationOutcome, BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository,
|
||||
PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput, StoredBillingModelContext,
|
||||
UserDailyQuotaAvailabilityRecord, UserPlanEntitlementRecord,
|
||||
PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput,
|
||||
PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
|
||||
UserPlanEntitlementRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
|
||||
@@ -215,6 +216,76 @@ impl BillingReadRepository for InMemoryBillingReadRepository {
|
||||
.cloned())
|
||||
}
|
||||
|
||||
async fn compare_and_swap_payment_gateway_secret(
|
||||
&self,
|
||||
update: &PaymentGatewaySecretCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let provider = update.provider.trim().to_ascii_lowercase();
|
||||
let mut configs = self
|
||||
.gateway_configs_by_provider
|
||||
.write()
|
||||
.expect("billing repository lock");
|
||||
let Some(record) = configs.get_mut(&provider) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if record.merchant_key_encrypted.as_deref()
|
||||
!= Some(update.expected_merchant_key_encrypted.as_str())
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
record.merchant_key_encrypted = Some(update.merchant_key_encrypted.clone());
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn compare_and_swap_payment_gateway_config(
|
||||
&self,
|
||||
mutation: &PaymentGatewayConfigCasWriteInput,
|
||||
) -> Result<AdminBillingMutationOutcome<PaymentGatewayConfigRecord>, DataLayerError> {
|
||||
let input = &mutation.input;
|
||||
let provider = input.provider.trim().to_ascii_lowercase();
|
||||
let now = current_unix_secs();
|
||||
let mut configs = self
|
||||
.gateway_configs_by_provider
|
||||
.write()
|
||||
.expect("billing repository lock");
|
||||
let existing = configs.get(&provider);
|
||||
if mutation.expected_existing {
|
||||
let Some(existing) = existing else {
|
||||
return Ok(AdminBillingMutationOutcome::NotFound);
|
||||
};
|
||||
if existing.merchant_key_encrypted != mutation.expected_merchant_key_encrypted {
|
||||
return Ok(AdminBillingMutationOutcome::NotFound);
|
||||
}
|
||||
} else if existing.is_some() {
|
||||
return Ok(AdminBillingMutationOutcome::NotFound);
|
||||
}
|
||||
|
||||
let created_at = existing
|
||||
.map(|value| value.created_at_unix_secs)
|
||||
.unwrap_or(now);
|
||||
let merchant_key_encrypted = if input.preserve_existing_secret {
|
||||
existing.and_then(|value| value.merchant_key_encrypted.clone())
|
||||
} else {
|
||||
input.merchant_key_encrypted.clone()
|
||||
};
|
||||
let record = PaymentGatewayConfigRecord {
|
||||
provider: provider.clone(),
|
||||
enabled: input.enabled,
|
||||
endpoint_url: input.endpoint_url.clone(),
|
||||
callback_base_url: input.callback_base_url.clone(),
|
||||
merchant_id: input.merchant_id.clone(),
|
||||
merchant_key_encrypted,
|
||||
pay_currency: input.pay_currency.clone(),
|
||||
usd_exchange_rate: input.usd_exchange_rate,
|
||||
min_recharge_usd: input.min_recharge_usd,
|
||||
channels_json: input.channels_json.clone(),
|
||||
created_at_unix_secs: created_at,
|
||||
updated_at_unix_secs: now,
|
||||
};
|
||||
configs.insert(provider, record.clone());
|
||||
Ok(AdminBillingMutationOutcome::Applied(record))
|
||||
}
|
||||
|
||||
async fn upsert_payment_gateway_config(
|
||||
&self,
|
||||
input: &PaymentGatewayConfigWriteInput,
|
||||
@@ -495,7 +566,10 @@ mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::InMemoryBillingReadRepository;
|
||||
use crate::repository::billing::{BillingReadRepository, StoredBillingModelContext};
|
||||
use crate::repository::billing::{
|
||||
AdminBillingMutationOutcome, BillingReadRepository, PaymentGatewayConfigCasWriteInput,
|
||||
PaymentGatewayConfigWriteInput, PaymentGatewaySecretCasUpdate, StoredBillingModelContext,
|
||||
};
|
||||
|
||||
fn sample_context() -> StoredBillingModelContext {
|
||||
StoredBillingModelContext::new(
|
||||
@@ -603,4 +677,105 @@ mod tests {
|
||||
Some("gpt-5-upstream")
|
||||
);
|
||||
}
|
||||
|
||||
fn gateway_input(secret: Option<&str>) -> PaymentGatewayConfigWriteInput {
|
||||
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: secret.map(ToOwned::to_owned),
|
||||
preserve_existing_secret: false,
|
||||
pay_currency: "USD".to_string(),
|
||||
usd_exchange_rate: 1.0,
|
||||
min_recharge_usd: 1.0,
|
||||
channels_json: json!({"channels": []}),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn payment_gateway_cas_prevents_create_overwrite_and_uses_exact_secret_fence() {
|
||||
let repository = InMemoryBillingReadRepository::default();
|
||||
let create = PaymentGatewayConfigCasWriteInput {
|
||||
input: gateway_input(Some("ciphertext-a")),
|
||||
expected_existing: false,
|
||||
expected_merchant_key_encrypted: None,
|
||||
};
|
||||
assert!(matches!(
|
||||
repository
|
||||
.compare_and_swap_payment_gateway_config(&create)
|
||||
.await
|
||||
.expect("create should succeed"),
|
||||
AdminBillingMutationOutcome::Applied(_)
|
||||
));
|
||||
|
||||
let mut competing_create = create.clone();
|
||||
competing_create.input.merchant_id = "overwritten".to_string();
|
||||
assert_eq!(
|
||||
repository
|
||||
.compare_and_swap_payment_gateway_config(&competing_create)
|
||||
.await
|
||||
.expect("conflicting create should be handled"),
|
||||
AdminBillingMutationOutcome::NotFound
|
||||
);
|
||||
|
||||
let mut stale_update = create.clone();
|
||||
stale_update.expected_existing = true;
|
||||
stale_update.expected_merchant_key_encrypted = Some("ciphertext-stale".to_string());
|
||||
stale_update.input.merchant_id = "stale-update".to_string();
|
||||
assert_eq!(
|
||||
repository
|
||||
.compare_and_swap_payment_gateway_config(&stale_update)
|
||||
.await
|
||||
.expect("stale update should be handled"),
|
||||
AdminBillingMutationOutcome::NotFound
|
||||
);
|
||||
let stored = repository
|
||||
.find_payment_gateway_config("stripe")
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("config should exist");
|
||||
assert_eq!(stored.merchant_id, "merchant");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn payment_gateway_secret_cas_changes_no_other_fields() {
|
||||
let repository = InMemoryBillingReadRepository::default();
|
||||
repository
|
||||
.upsert_payment_gateway_config(&gateway_input(Some("legacy-ciphertext")))
|
||||
.await
|
||||
.expect("seed should succeed");
|
||||
let before = repository
|
||||
.find_payment_gateway_config("stripe")
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("config should exist");
|
||||
|
||||
assert!(!repository
|
||||
.compare_and_swap_payment_gateway_secret(&PaymentGatewaySecretCasUpdate {
|
||||
provider: "stripe".to_string(),
|
||||
expected_merchant_key_encrypted: "wrong-ciphertext".to_string(),
|
||||
merchant_key_encrypted: "v2-ciphertext".to_string(),
|
||||
})
|
||||
.await
|
||||
.expect("stale secret CAS should be handled"));
|
||||
assert!(repository
|
||||
.compare_and_swap_payment_gateway_secret(&PaymentGatewaySecretCasUpdate {
|
||||
provider: "stripe".to_string(),
|
||||
expected_merchant_key_encrypted: "legacy-ciphertext".to_string(),
|
||||
merchant_key_encrypted: "v2-ciphertext".to_string(),
|
||||
})
|
||||
.await
|
||||
.expect("secret CAS should succeed"));
|
||||
|
||||
let mut expected = before.clone();
|
||||
expected.merchant_key_encrypted = Some("v2-ciphertext".to_string());
|
||||
let after = repository
|
||||
.find_payment_gateway_config("stripe")
|
||||
.await
|
||||
.expect("lookup should succeed")
|
||||
.expect("config should exist");
|
||||
assert_eq!(after, expected);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,14 +1,18 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::{
|
||||
request_candidate_lifecycle_would_regress, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
||||
RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository,
|
||||
StoredRequestCandidate, UpsertRequestCandidateRecord,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
use async_trait::async_trait;
|
||||
|
||||
fn sanitize_stored_candidate(mut candidate: StoredRequestCandidate) -> StoredRequestCandidate {
|
||||
candidate.sanitize_sensitive_diagnostics();
|
||||
candidate
|
||||
}
|
||||
|
||||
fn merge_extra_data(
|
||||
existing: Option<serde_json::Value>,
|
||||
@@ -38,7 +42,7 @@ impl InMemoryRequestCandidateRepository {
|
||||
I: IntoIterator<Item = StoredRequestCandidate>,
|
||||
{
|
||||
let mut by_id = BTreeMap::new();
|
||||
for item in items {
|
||||
for item in items.into_iter().map(sanitize_stored_candidate) {
|
||||
by_id.insert(item.id.clone(), item);
|
||||
}
|
||||
Self {
|
||||
@@ -60,6 +64,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
|
||||
.values()
|
||||
.filter(|row| row.request_id == request_id)
|
||||
.cloned()
|
||||
.map(sanitize_stored_candidate)
|
||||
.collect::<Vec<_>>();
|
||||
rows.sort_by(|left, right| {
|
||||
left.candidate_index
|
||||
@@ -84,6 +89,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
|
||||
.expect("request candidate repository lock")
|
||||
.values()
|
||||
.cloned()
|
||||
.map(sanitize_stored_candidate)
|
||||
.collect::<Vec<_>>();
|
||||
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
|
||||
rows.truncate(limit);
|
||||
@@ -106,6 +112,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
|
||||
.values()
|
||||
.filter(|row| row.provider_id.as_deref() == Some(provider_id))
|
||||
.cloned()
|
||||
.map(sanitize_stored_candidate)
|
||||
.collect::<Vec<_>>();
|
||||
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
|
||||
rows.truncate(limit);
|
||||
@@ -141,6 +148,7 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
|
||||
)
|
||||
})
|
||||
.cloned()
|
||||
.map(sanitize_stored_candidate)
|
||||
.collect::<Vec<_>>();
|
||||
rows.sort_by_key(|entry| std::cmp::Reverse(entry.created_at_unix_ms));
|
||||
rows.truncate(limit);
|
||||
@@ -294,8 +302,9 @@ impl RequestCandidateReadRepository for InMemoryRequestCandidateRepository {
|
||||
impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
|
||||
async fn upsert(
|
||||
&self,
|
||||
candidate: UpsertRequestCandidateRecord,
|
||||
mut candidate: UpsertRequestCandidateRecord,
|
||||
) -> Result<StoredRequestCandidate, DataLayerError> {
|
||||
candidate.sanitize_for_persistence();
|
||||
candidate.validate()?;
|
||||
|
||||
let mut by_id = self
|
||||
@@ -309,7 +318,8 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
|
||||
&& row.candidate_index == candidate.candidate_index
|
||||
&& row.retry_index == candidate.retry_index
|
||||
})
|
||||
.cloned();
|
||||
.cloned()
|
||||
.map(sanitize_stored_candidate);
|
||||
|
||||
let preserve_existing_lifecycle = existing.as_ref().is_some_and(|row| {
|
||||
request_candidate_lifecycle_would_regress(row.status, candidate.status)
|
||||
@@ -336,29 +346,36 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
|
||||
.map(|row| row.id.clone())
|
||||
.unwrap_or_else(|| candidate.id.clone()),
|
||||
request_id: candidate.request_id.clone(),
|
||||
user_id: candidate
|
||||
.user_id
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.user_id.clone())),
|
||||
api_key_id: candidate
|
||||
.api_key_id
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.api_key_id.clone())),
|
||||
username: candidate
|
||||
.username
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.username.clone())),
|
||||
api_key_name: candidate
|
||||
.api_key_name
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.api_key_name.clone())),
|
||||
user_id: existing
|
||||
.as_ref()
|
||||
.and_then(|row| row.user_id.clone())
|
||||
.or(candidate.user_id),
|
||||
api_key_id: existing
|
||||
.as_ref()
|
||||
.and_then(|row| row.api_key_id.clone())
|
||||
.or(candidate.api_key_id),
|
||||
username: existing
|
||||
.as_ref()
|
||||
.and_then(|row| row.username.clone())
|
||||
.or(candidate.username),
|
||||
api_key_name: existing
|
||||
.as_ref()
|
||||
.and_then(|row| row.api_key_name.clone())
|
||||
.or(candidate.api_key_name),
|
||||
candidate_index: candidate.candidate_index,
|
||||
retry_index: candidate.retry_index,
|
||||
provider_id: candidate
|
||||
.provider_id
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.provider_id.clone())),
|
||||
endpoint_id: candidate
|
||||
.endpoint_id
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.endpoint_id.clone())),
|
||||
key_id: candidate
|
||||
.key_id
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.key_id.clone())),
|
||||
provider_id: existing
|
||||
.as_ref()
|
||||
.and_then(|row| row.provider_id.clone())
|
||||
.or(candidate.provider_id),
|
||||
endpoint_id: existing
|
||||
.as_ref()
|
||||
.and_then(|row| row.endpoint_id.clone())
|
||||
.or(candidate.endpoint_id),
|
||||
key_id: existing
|
||||
.as_ref()
|
||||
.and_then(|row| row.key_id.clone())
|
||||
.or(candidate.key_id),
|
||||
status: merged_status,
|
||||
skip_reason: candidate
|
||||
.skip_reason
|
||||
@@ -380,13 +397,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
|
||||
.error_type
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.error_type.clone()))
|
||||
},
|
||||
error_message: if preserve_existing_lifecycle {
|
||||
existing.as_ref().and_then(|row| row.error_message.clone())
|
||||
} else {
|
||||
candidate
|
||||
.error_message
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.error_message.clone()))
|
||||
},
|
||||
error_message: None,
|
||||
latency_ms: if preserve_existing_lifecycle {
|
||||
existing.as_ref().and_then(|row| row.latency_ms)
|
||||
} else {
|
||||
@@ -407,9 +418,10 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
|
||||
.and_then(|row| row.required_capabilities.clone())
|
||||
}),
|
||||
created_at_unix_ms,
|
||||
started_at_unix_ms: candidate
|
||||
.started_at_unix_ms
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.started_at_unix_ms)),
|
||||
started_at_unix_ms: existing
|
||||
.as_ref()
|
||||
.and_then(|row| row.started_at_unix_ms)
|
||||
.or(candidate.started_at_unix_ms),
|
||||
finished_at_unix_ms: if preserve_existing_lifecycle {
|
||||
existing.as_ref().and_then(|row| row.finished_at_unix_ms)
|
||||
} else {
|
||||
@@ -418,6 +430,7 @@ impl RequestCandidateWriteRepository for InMemoryRequestCandidateRepository {
|
||||
.or_else(|| existing.as_ref().and_then(|row| row.finished_at_unix_ms))
|
||||
},
|
||||
};
|
||||
let stored = sanitize_stored_candidate(stored);
|
||||
|
||||
by_id.insert(stored.id.clone(), stored.clone());
|
||||
Ok(stored)
|
||||
@@ -531,6 +544,138 @@ mod tests {
|
||||
assert_eq!(rows[1].id, "cand-1");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn seed_and_reads_sanitize_candidates_that_bypass_contract_constructors() {
|
||||
let raw_candidate = StoredRequestCandidate {
|
||||
id: "cand-raw".to_string(),
|
||||
request_id: "req-raw".to_string(),
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
candidate_index: 0,
|
||||
retry_index: 0,
|
||||
provider_id: Some("provider-1".to_string()),
|
||||
endpoint_id: Some("endpoint-1".to_string()),
|
||||
key_id: None,
|
||||
status: RequestCandidateStatus::Failed,
|
||||
skip_reason: Some("secret=/private/path".to_string()),
|
||||
is_cached: false,
|
||||
status_code: Some(500),
|
||||
error_type: Some("token=secret".to_string()),
|
||||
error_message: Some("Bearer secret-token".to_string()),
|
||||
latency_ms: Some(10),
|
||||
concurrent_requests: Some(1),
|
||||
extra_data: Some(json!({
|
||||
"gateway_execution_runtime": true,
|
||||
"request_headers": {"authorization": "Bearer secret-token"},
|
||||
"request_body": {"password": "secret"}
|
||||
})),
|
||||
required_capabilities: Some(json!({
|
||||
"streaming": "true",
|
||||
"internal_capability": "secret"
|
||||
})),
|
||||
created_at_unix_ms: 100,
|
||||
started_at_unix_ms: Some(100),
|
||||
finished_at_unix_ms: Some(110),
|
||||
};
|
||||
let repository = InMemoryRequestCandidateRepository::seed(vec![raw_candidate.clone()]);
|
||||
|
||||
{
|
||||
let stored = repository
|
||||
.by_id
|
||||
.read()
|
||||
.expect("request candidate repository lock");
|
||||
let candidate = stored
|
||||
.get("cand-raw")
|
||||
.expect("seeded candidate should exist");
|
||||
assert!(candidate.error_message.is_none());
|
||||
assert_eq!(candidate.skip_reason.as_deref(), Some("unclassified_skip"));
|
||||
assert_eq!(candidate.error_type.as_deref(), Some("unclassified_error"));
|
||||
assert_eq!(
|
||||
candidate.extra_data,
|
||||
Some(json!({"gateway_execution_runtime": true}))
|
||||
);
|
||||
assert_eq!(
|
||||
candidate.required_capabilities,
|
||||
Some(json!({"streaming": true}))
|
||||
);
|
||||
}
|
||||
|
||||
let mut bypassed_candidate = raw_candidate;
|
||||
bypassed_candidate.id = "cand-bypassed".to_string();
|
||||
bypassed_candidate.request_id = "req-bypassed".to_string();
|
||||
repository
|
||||
.by_id
|
||||
.write()
|
||||
.expect("request candidate repository lock")
|
||||
.insert(bypassed_candidate.id.clone(), bypassed_candidate);
|
||||
|
||||
let rows = repository
|
||||
.list_recent(10)
|
||||
.await
|
||||
.expect("list recent should succeed");
|
||||
let candidate = rows
|
||||
.iter()
|
||||
.find(|candidate| candidate.id == "cand-bypassed")
|
||||
.expect("bypassed candidate should be returned");
|
||||
assert!(candidate.error_message.is_none());
|
||||
assert_eq!(
|
||||
candidate.extra_data,
|
||||
Some(json!({"gateway_execution_runtime": true}))
|
||||
);
|
||||
assert_eq!(
|
||||
candidate.required_capabilities,
|
||||
Some(json!({"streaming": true}))
|
||||
);
|
||||
|
||||
let merged = repository
|
||||
.upsert(UpsertRequestCandidateRecord {
|
||||
id: "cand-merged".to_string(),
|
||||
request_id: "req-bypassed".to_string(),
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
candidate_index: 0,
|
||||
retry_index: 0,
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
status: RequestCandidateStatus::Success,
|
||||
skip_reason: None,
|
||||
is_cached: None,
|
||||
status_code: Some(200),
|
||||
error_type: None,
|
||||
error_message: Some("Bearer new-secret".to_string()),
|
||||
latency_ms: Some(12),
|
||||
concurrent_requests: None,
|
||||
extra_data: Some(json!({
|
||||
"stream_completed": true,
|
||||
"request_body": {"password": "new-secret"}
|
||||
})),
|
||||
required_capabilities: Some(json!({
|
||||
"vision": 1,
|
||||
"internal_capability": "new-secret"
|
||||
})),
|
||||
created_at_unix_ms: Some(100),
|
||||
started_at_unix_ms: Some(100),
|
||||
finished_at_unix_ms: Some(112),
|
||||
})
|
||||
.await
|
||||
.expect("candidate merge should succeed");
|
||||
assert_eq!(merged.id, "cand-bypassed");
|
||||
assert!(merged.error_message.is_none());
|
||||
assert_eq!(
|
||||
merged.extra_data,
|
||||
Some(json!({
|
||||
"gateway_execution_runtime": true,
|
||||
"stream_completed": true
|
||||
}))
|
||||
);
|
||||
assert_eq!(merged.required_capabilities, Some(json!({"vision": true})));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn aggregates_finalized_health_data_by_endpoint_ids() {
|
||||
let repository = InMemoryRequestCandidateRepository::seed(vec![
|
||||
@@ -594,7 +739,7 @@ mod tests {
|
||||
"execution_strategy": "local_cross_format",
|
||||
"provider_name": "primary",
|
||||
})),
|
||||
required_capabilities: None,
|
||||
required_capabilities: Some(json!({"streaming": true})),
|
||||
created_at_unix_ms: Some(100),
|
||||
started_at_unix_ms: None,
|
||||
finished_at_unix_ms: None,
|
||||
@@ -608,15 +753,15 @@ mod tests {
|
||||
.upsert(UpsertRequestCandidateRecord {
|
||||
id: "cand-1-replacement".to_string(),
|
||||
request_id: "req-1".to_string(),
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
user_id: Some("attacker-user".to_string()),
|
||||
api_key_id: Some("attacker-api-key".to_string()),
|
||||
username: Some("mallory".to_string()),
|
||||
api_key_name: Some("attacker-key".to_string()),
|
||||
candidate_index: 0,
|
||||
retry_index: 0,
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
provider_id: Some("attacker-provider".to_string()),
|
||||
endpoint_id: Some("attacker-endpoint".to_string()),
|
||||
key_id: Some("attacker-provider-key".to_string()),
|
||||
status: RequestCandidateStatus::Success,
|
||||
skip_reason: None,
|
||||
is_cached: None,
|
||||
@@ -629,7 +774,7 @@ mod tests {
|
||||
"provider_api_format": "openai:responses",
|
||||
"provider_name": "updated",
|
||||
})),
|
||||
required_capabilities: None,
|
||||
required_capabilities: Some(json!({"vision": true})),
|
||||
created_at_unix_ms: None,
|
||||
started_at_unix_ms: Some(101),
|
||||
finished_at_unix_ms: Some(102),
|
||||
@@ -638,6 +783,14 @@ mod tests {
|
||||
.expect("update should succeed");
|
||||
assert_eq!(updated.id, "cand-1");
|
||||
assert_eq!(updated.status, RequestCandidateStatus::Success);
|
||||
assert_eq!(updated.user_id.as_deref(), Some("user-1"));
|
||||
assert_eq!(updated.api_key_id.as_deref(), Some("api-key-1"));
|
||||
assert!(updated.username.is_none());
|
||||
assert!(updated.api_key_name.is_none());
|
||||
assert_eq!(updated.provider_id.as_deref(), Some("provider-1"));
|
||||
assert_eq!(updated.endpoint_id.as_deref(), Some("endpoint-1"));
|
||||
assert_eq!(updated.key_id.as_deref(), Some("key-1"));
|
||||
assert_eq!(updated.required_capabilities, Some(json!({"vision": true})));
|
||||
assert_eq!(updated.status_code, Some(200));
|
||||
assert_eq!(updated.latency_ms, Some(25));
|
||||
assert_eq!(
|
||||
@@ -659,13 +812,13 @@ mod tests {
|
||||
.extra_data
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("provider_name")),
|
||||
Some(&json!("updated"))
|
||||
None
|
||||
);
|
||||
assert_eq!(updated.started_at_unix_ms, Some(101));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upsert_keeps_terminal_candidate_state_when_streaming_arrives_late() {
|
||||
async fn upsert_keeps_first_terminal_candidate_fact_when_another_terminal_arrives_late() {
|
||||
let existing = StoredRequestCandidate::new(
|
||||
"cand-1".to_string(),
|
||||
"req-1".to_string(),
|
||||
@@ -686,7 +839,7 @@ mod tests {
|
||||
Some("retryable upstream failure".to_string()),
|
||||
Some(45),
|
||||
Some(1),
|
||||
Some(json!({"terminal": true})),
|
||||
Some(json!({"stream_completed": true})),
|
||||
None,
|
||||
100,
|
||||
Some(101),
|
||||
@@ -708,7 +861,7 @@ mod tests {
|
||||
provider_id: None,
|
||||
endpoint_id: None,
|
||||
key_id: None,
|
||||
status: RequestCandidateStatus::Streaming,
|
||||
status: RequestCandidateStatus::Success,
|
||||
skip_reason: None,
|
||||
is_cached: None,
|
||||
status_code: Some(200),
|
||||
@@ -716,7 +869,7 @@ mod tests {
|
||||
error_message: None,
|
||||
latency_ms: Some(9_999),
|
||||
concurrent_requests: Some(2),
|
||||
extra_data: Some(json!({"late": true})),
|
||||
extra_data: Some(json!({"gateway_execution_runtime": true})),
|
||||
required_capabilities: None,
|
||||
created_at_unix_ms: None,
|
||||
started_at_unix_ms: Some(102),
|
||||
@@ -729,16 +882,16 @@ mod tests {
|
||||
assert_eq!(updated.status, RequestCandidateStatus::Failed);
|
||||
assert_eq!(updated.status_code, Some(503));
|
||||
assert_eq!(updated.error_type.as_deref(), Some("upstream_error"));
|
||||
assert_eq!(
|
||||
updated.error_message.as_deref(),
|
||||
Some("retryable upstream failure")
|
||||
);
|
||||
assert!(updated.error_message.is_none());
|
||||
assert_eq!(updated.latency_ms, Some(45));
|
||||
assert_eq!(updated.concurrent_requests, Some(2));
|
||||
assert_eq!(updated.finished_at_unix_ms, Some(145));
|
||||
assert_eq!(
|
||||
updated.extra_data,
|
||||
Some(json!({"terminal": true, "late": true}))
|
||||
Some(json!({
|
||||
"gateway_execution_runtime": true,
|
||||
"stream_completed": true
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -41,6 +41,40 @@ impl GeminiFileMappingReadRepository for InMemoryGeminiFileMappingRepository {
|
||||
Ok(guard.get(file_name).cloned())
|
||||
}
|
||||
|
||||
async fn find_active_by_file_name_for_user(
|
||||
&self,
|
||||
file_name: &str,
|
||||
user_id: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
|
||||
let guard = self.by_file.read().expect("gemini mapping repository lock");
|
||||
Ok(guard
|
||||
.get(file_name)
|
||||
.filter(|mapping| {
|
||||
mapping.user_id.as_deref() == Some(user_id)
|
||||
&& mapping.expires_at_unix_secs > now_unix_secs
|
||||
})
|
||||
.cloned())
|
||||
}
|
||||
|
||||
async fn find_active_by_file_name_for_owner(
|
||||
&self,
|
||||
file_name: &str,
|
||||
key_id: &str,
|
||||
user_id: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
|
||||
let guard = self.by_file.read().expect("gemini mapping repository lock");
|
||||
Ok(guard
|
||||
.get(file_name)
|
||||
.filter(|mapping| {
|
||||
mapping.key_id == key_id
|
||||
&& mapping.user_id.as_deref() == Some(user_id)
|
||||
&& mapping.expires_at_unix_secs > now_unix_secs
|
||||
})
|
||||
.cloned())
|
||||
}
|
||||
|
||||
async fn list_mappings(
|
||||
&self,
|
||||
query: &GeminiFileMappingListQuery,
|
||||
@@ -52,6 +86,12 @@ impl GeminiFileMappingReadRepository for InMemoryGeminiFileMappingRepository {
|
||||
.map(|value| value.to_ascii_lowercase());
|
||||
let mut items = guard
|
||||
.values()
|
||||
.filter(|item| {
|
||||
query
|
||||
.user_id
|
||||
.as_deref()
|
||||
.is_none_or(|user_id| item.user_id.as_deref() == Some(user_id))
|
||||
})
|
||||
.filter(|item| query.include_expired || item.expires_at_unix_secs > query.now_unix_secs)
|
||||
.filter(|item| {
|
||||
search.as_deref().is_none_or(|needle| {
|
||||
@@ -146,6 +186,39 @@ impl GeminiFileMappingWriteRepository for InMemoryGeminiFileMappingRepository {
|
||||
Ok(mapping)
|
||||
}
|
||||
|
||||
async fn upsert_if_owner_matches(
|
||||
&self,
|
||||
record: UpsertGeminiFileMappingRecord,
|
||||
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
|
||||
record.validate()?;
|
||||
let mut guard = self
|
||||
.by_file
|
||||
.write()
|
||||
.expect("gemini mapping repository lock");
|
||||
let (id, created_at_unix_ms) = match guard.get(&record.file_name) {
|
||||
Some(existing)
|
||||
if existing.key_id == record.key_id && existing.user_id == record.user_id =>
|
||||
{
|
||||
(existing.id.clone(), existing.created_at_unix_ms)
|
||||
}
|
||||
Some(_) => return Ok(None),
|
||||
None => (record.id.clone(), current_unix_secs()),
|
||||
};
|
||||
let mapping = StoredGeminiFileMapping {
|
||||
id,
|
||||
file_name: record.file_name.clone(),
|
||||
key_id: record.key_id.clone(),
|
||||
user_id: record.user_id.clone(),
|
||||
display_name: record.display_name.clone(),
|
||||
mime_type: record.mime_type.clone(),
|
||||
source_hash: record.source_hash.clone(),
|
||||
created_at_unix_ms,
|
||||
expires_at_unix_secs: record.expires_at_unix_secs,
|
||||
};
|
||||
guard.insert(record.file_name, mapping.clone());
|
||||
Ok(Some(mapping))
|
||||
}
|
||||
|
||||
async fn delete_by_file_name(&self, file_name: &str) -> Result<bool, DataLayerError> {
|
||||
let mut guard = self
|
||||
.by_file
|
||||
@@ -154,6 +227,44 @@ impl GeminiFileMappingWriteRepository for InMemoryGeminiFileMappingRepository {
|
||||
Ok(guard.remove(file_name).is_some())
|
||||
}
|
||||
|
||||
async fn delete_by_file_name_for_user(
|
||||
&self,
|
||||
file_name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut guard = self
|
||||
.by_file
|
||||
.write()
|
||||
.expect("gemini mapping repository lock");
|
||||
if guard
|
||||
.get(file_name)
|
||||
.and_then(|item| item.user_id.as_deref())
|
||||
!= Some(user_id)
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
Ok(guard.remove(file_name).is_some())
|
||||
}
|
||||
|
||||
async fn delete_by_file_name_for_owner(
|
||||
&self,
|
||||
file_name: &str,
|
||||
key_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut guard = self
|
||||
.by_file
|
||||
.write()
|
||||
.expect("gemini mapping repository lock");
|
||||
let owner_matches = guard
|
||||
.get(file_name)
|
||||
.is_some_and(|item| item.key_id == key_id && item.user_id.as_deref() == Some(user_id));
|
||||
if !owner_matches {
|
||||
return Ok(false);
|
||||
}
|
||||
Ok(guard.remove(file_name).is_some())
|
||||
}
|
||||
|
||||
async fn delete_by_id(
|
||||
&self,
|
||||
mapping_id: &str,
|
||||
@@ -224,6 +335,35 @@ mod tests {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn owner_scoped_reads_bind_user_provider_key_and_expiry() -> Result<(), DataLayerError> {
|
||||
let repo = InMemoryGeminiFileMappingRepository::default();
|
||||
repo.upsert(sample_record("id-owner", "files/owned"))
|
||||
.await?;
|
||||
|
||||
assert!(repo
|
||||
.find_active_by_file_name_for_user("files/owned", "user-1", 100)
|
||||
.await?
|
||||
.is_some());
|
||||
assert!(repo
|
||||
.find_active_by_file_name_for_user("files/owned", "user-2", 100)
|
||||
.await?
|
||||
.is_none());
|
||||
assert!(repo
|
||||
.find_active_by_file_name_for_owner("files/owned", "key-1", "user-1", 100)
|
||||
.await?
|
||||
.is_some());
|
||||
assert!(repo
|
||||
.find_active_by_file_name_for_owner("files/owned", "key-2", "user-1", 100)
|
||||
.await?
|
||||
.is_none());
|
||||
assert!(repo
|
||||
.find_active_by_file_name_for_user("files/owned", "user-1", 4_102_444_800)
|
||||
.await?
|
||||
.is_none());
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn delete_removes_entry() -> Result<(), DataLayerError> {
|
||||
let repo = InMemoryGeminiFileMappingRepository::default();
|
||||
@@ -247,6 +387,34 @@ mod tests {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn owner_checked_upsert_cannot_reassign_existing_mapping() -> Result<(), DataLayerError> {
|
||||
let repo = InMemoryGeminiFileMappingRepository::default();
|
||||
let first = repo.upsert(sample_record("id-1", "files/owned")).await?;
|
||||
|
||||
let mut attacker = sample_record("id-2", "files/owned");
|
||||
attacker.key_id = "key-2".to_string();
|
||||
attacker.user_id = Some("user-2".to_string());
|
||||
assert!(repo.upsert_if_owner_matches(attacker).await?.is_none());
|
||||
|
||||
let unchanged = repo
|
||||
.find_by_file_name("files/owned")
|
||||
.await?
|
||||
.expect("mapping should remain");
|
||||
assert_eq!(unchanged.key_id, "key-1");
|
||||
assert_eq!(unchanged.user_id.as_deref(), Some("user-1"));
|
||||
|
||||
let mut refresh = sample_record("id-3", "files/owned");
|
||||
refresh.display_name = Some("refreshed".to_string());
|
||||
let refreshed = repo
|
||||
.upsert_if_owner_matches(refresh)
|
||||
.await?
|
||||
.expect("same owner should refresh");
|
||||
assert_eq!(refreshed.id, first.id);
|
||||
assert_eq!(refreshed.display_name.as_deref(), Some("refreshed"));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_and_summarize_mappings() -> Result<(), DataLayerError> {
|
||||
let repo = InMemoryGeminiFileMappingRepository::seed(vec![
|
||||
@@ -257,6 +425,7 @@ mod tests {
|
||||
|
||||
let page = repo
|
||||
.list_mappings(&GeminiFileMappingListQuery {
|
||||
user_id: None,
|
||||
include_expired: false,
|
||||
search: Some("ga".to_string()),
|
||||
offset: 0,
|
||||
@@ -279,6 +448,46 @@ mod tests {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn owner_filter_and_delete_do_not_cross_user_boundaries() -> Result<(), DataLayerError> {
|
||||
let mut first = repo_item("id-1", "files/alpha", "image/png", 10, 200);
|
||||
first.user_id = Some("user-1".to_string());
|
||||
let mut second = repo_item("id-2", "files/beta", "image/png", 20, 200);
|
||||
second.user_id = Some("user-2".to_string());
|
||||
let repo = InMemoryGeminiFileMappingRepository::seed([first, second]);
|
||||
|
||||
let page = repo
|
||||
.list_mappings(&GeminiFileMappingListQuery {
|
||||
user_id: Some("user-1".to_string()),
|
||||
include_expired: false,
|
||||
search: None,
|
||||
offset: 0,
|
||||
limit: 10,
|
||||
now_unix_secs: 100,
|
||||
})
|
||||
.await?;
|
||||
assert_eq!(page.total, 1);
|
||||
assert_eq!(page.items[0].file_name, "files/alpha");
|
||||
|
||||
assert!(
|
||||
!repo
|
||||
.delete_by_file_name_for_user("files/alpha", "user-2")
|
||||
.await?
|
||||
);
|
||||
assert!(repo.find_by_file_name("files/alpha").await?.is_some());
|
||||
assert!(
|
||||
!repo
|
||||
.delete_by_file_name_for_owner("files/alpha", "key-2", "user-1")
|
||||
.await?
|
||||
);
|
||||
assert!(repo.find_by_file_name("files/alpha").await?.is_some());
|
||||
assert!(
|
||||
repo.delete_by_file_name_for_user("files/alpha", "user-1")
|
||||
.await?
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn delete_by_id_and_cleanup_expired() -> Result<(), DataLayerError> {
|
||||
let repo = InMemoryGeminiFileMappingRepository::seed(vec![
|
||||
|
||||
@@ -6,9 +6,10 @@ use async_trait::async_trait;
|
||||
|
||||
use crate::DataLayerError;
|
||||
use aether_data_contracts::repository::management_tokens::{
|
||||
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
|
||||
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
|
||||
StoredManagementTokenListPage, StoredManagementTokenWithUser, UpdateManagementTokenRecord,
|
||||
ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery,
|
||||
ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret,
|
||||
StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenWithUser,
|
||||
UpdateManagementTokenRecord,
|
||||
};
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
@@ -49,6 +50,189 @@ impl InMemoryManagementTokenRepository {
|
||||
fn remove_hash_for_token(hashes: &mut BTreeMap<String, String>, token_id: &str) {
|
||||
hashes.retain(|_, existing_token_id| existing_token_id != token_id);
|
||||
}
|
||||
|
||||
fn update_management_token_scoped(
|
||||
&self,
|
||||
record: &UpdateManagementTokenRecord,
|
||||
expected_user_id: Option<&str>,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
record.validate()?;
|
||||
|
||||
let mut items = self
|
||||
.items
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let Some(index) = items.iter().position(|item| {
|
||||
item.token.id == record.token_id
|
||||
&& expected_user_id
|
||||
.map(|user_id| item.token.user_id == user_id)
|
||||
.unwrap_or(true)
|
||||
}) else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
if let Some(name) = &record.name {
|
||||
if items.iter().enumerate().any(|(position, item)| {
|
||||
position != index
|
||||
&& item.token.user_id == items[index].token.user_id
|
||||
&& item.token.name == *name
|
||||
}) {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"已存在名为 '{}' 的 Token",
|
||||
name
|
||||
)));
|
||||
}
|
||||
items[index].token.name = name.clone();
|
||||
}
|
||||
|
||||
if record.clear_description {
|
||||
items[index].token.description = None;
|
||||
} else if let Some(description) = &record.description {
|
||||
items[index].token.description = Some(description.clone());
|
||||
}
|
||||
|
||||
if record.clear_allowed_ips {
|
||||
items[index].token.allowed_ips = None;
|
||||
} else if let Some(allowed_ips) = &record.allowed_ips {
|
||||
items[index].token.allowed_ips = Some(allowed_ips.clone());
|
||||
}
|
||||
|
||||
if let Some(permissions) = &record.permissions {
|
||||
items[index].token.permissions = Some(permissions.clone());
|
||||
}
|
||||
|
||||
if record.clear_expires_at {
|
||||
items[index].token.expires_at_unix_secs = None;
|
||||
} else if let Some(expires_at_unix_secs) = record.expires_at_unix_secs {
|
||||
items[index].token.expires_at_unix_secs = Some(expires_at_unix_secs);
|
||||
}
|
||||
|
||||
if let Some(is_active) = record.is_active {
|
||||
items[index].token.is_active = is_active;
|
||||
}
|
||||
|
||||
items[index].token.updated_at_unix_secs = Self::now_unix_secs();
|
||||
Ok(Some(items[index].token.clone()))
|
||||
}
|
||||
|
||||
fn delete_management_token_scoped(
|
||||
&self,
|
||||
token_id: &str,
|
||||
expected_user_id: Option<&str>,
|
||||
) -> bool {
|
||||
let mut items = self
|
||||
.items
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let mut hashes = self
|
||||
.hashes
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let original_len = items.len();
|
||||
items.retain(|item| {
|
||||
item.token.id != token_id
|
||||
|| expected_user_id
|
||||
.map(|user_id| item.token.user_id != user_id)
|
||||
.unwrap_or(false)
|
||||
});
|
||||
if items.len() != original_len {
|
||||
Self::remove_hash_for_token(&mut hashes, token_id);
|
||||
return true;
|
||||
}
|
||||
false
|
||||
}
|
||||
|
||||
fn set_management_token_active_scoped(
|
||||
&self,
|
||||
token_id: &str,
|
||||
expected_user_id: Option<&str>,
|
||||
is_active: bool,
|
||||
) -> Option<StoredManagementToken> {
|
||||
let mut items = self
|
||||
.items
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let item = items.iter_mut().find(|item| {
|
||||
item.token.id == token_id
|
||||
&& expected_user_id
|
||||
.map(|user_id| item.token.user_id == user_id)
|
||||
.unwrap_or(true)
|
||||
})?;
|
||||
item.token.is_active = is_active;
|
||||
item.token.updated_at_unix_secs = Self::now_unix_secs();
|
||||
Some(item.token.clone())
|
||||
}
|
||||
|
||||
fn activate_management_token_if_matches_inner(
|
||||
&self,
|
||||
mutation: &ActivateManagementTokenIfMatches,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
mutation.validate()?;
|
||||
// The in-memory token store does not own the independently stored user row and cannot
|
||||
// atomically verify role/status/security_version with this mutation. Pretending that the
|
||||
// user summary cached beside the token is authoritative would recreate the TOCTOU, so
|
||||
// one-time install activation is intentionally unavailable on this backend.
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
fn delete_inactive_management_token_if_matches_inner(
|
||||
&self,
|
||||
mutation: &ActivateManagementTokenIfMatches,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
mutation.validate()?;
|
||||
|
||||
let mut items = self
|
||||
.items
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let mut hashes = self
|
||||
.hashes
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
if hashes.get(&mutation.token_hash).map(String::as_str)
|
||||
!= Some(mutation.expected_token.id.as_str())
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
let Some(index) = items.iter().position(|item| {
|
||||
mutation.matches_locked_token_snapshot(&item.token, &mutation.token_hash)
|
||||
}) else {
|
||||
return Ok(false);
|
||||
};
|
||||
items.remove(index);
|
||||
Self::remove_hash_for_token(&mut hashes, &mutation.expected_token.id);
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
fn regenerate_management_token_secret_scoped(
|
||||
&self,
|
||||
mutation: &RegenerateManagementTokenSecret,
|
||||
expected_user_id: Option<&str>,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
mutation.validate()?;
|
||||
|
||||
let mut items = self
|
||||
.items
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let mut hashes = self
|
||||
.hashes
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let Some(item) = items.iter_mut().find(|item| {
|
||||
item.token.id == mutation.token_id
|
||||
&& expected_user_id
|
||||
.map(|user_id| item.token.user_id == user_id)
|
||||
.unwrap_or(true)
|
||||
}) else {
|
||||
return Ok(None);
|
||||
};
|
||||
Self::remove_hash_for_token(&mut hashes, &mutation.token_id);
|
||||
hashes.insert(mutation.token_hash.clone(), mutation.token_id.clone());
|
||||
item.token.token_prefix = mutation.token_prefix.clone();
|
||||
item.token.updated_at_unix_secs = Self::now_unix_secs();
|
||||
Ok(Some(item.token.clone()))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -167,76 +351,27 @@ impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository {
|
||||
&self,
|
||||
record: &UpdateManagementTokenRecord,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
record.validate()?;
|
||||
self.update_management_token_scoped(record, None)
|
||||
}
|
||||
|
||||
let mut items = self
|
||||
.items
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let Some(index) = items
|
||||
.iter()
|
||||
.position(|item| item.token.id == record.token_id)
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
if let Some(name) = &record.name {
|
||||
if items.iter().enumerate().any(|(position, item)| {
|
||||
position != index
|
||||
&& item.token.user_id == items[index].token.user_id
|
||||
&& item.token.name == *name
|
||||
}) {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"已存在名为 '{}' 的 Token",
|
||||
name
|
||||
)));
|
||||
}
|
||||
items[index].token.name = name.clone();
|
||||
}
|
||||
|
||||
if record.clear_description {
|
||||
items[index].token.description = None;
|
||||
} else if let Some(description) = &record.description {
|
||||
items[index].token.description = Some(description.clone());
|
||||
}
|
||||
|
||||
if record.clear_allowed_ips {
|
||||
items[index].token.allowed_ips = None;
|
||||
} else if let Some(allowed_ips) = &record.allowed_ips {
|
||||
items[index].token.allowed_ips = Some(allowed_ips.clone());
|
||||
}
|
||||
|
||||
if let Some(permissions) = &record.permissions {
|
||||
items[index].token.permissions = Some(permissions.clone());
|
||||
}
|
||||
|
||||
if record.clear_expires_at {
|
||||
items[index].token.expires_at_unix_secs = None;
|
||||
} else if let Some(expires_at_unix_secs) = record.expires_at_unix_secs {
|
||||
items[index].token.expires_at_unix_secs = Some(expires_at_unix_secs);
|
||||
}
|
||||
|
||||
if let Some(is_active) = record.is_active {
|
||||
items[index].token.is_active = is_active;
|
||||
}
|
||||
|
||||
items[index].token.updated_at_unix_secs = Self::now_unix_secs();
|
||||
Ok(Some(items[index].token.clone()))
|
||||
async fn update_management_token_for_user(
|
||||
&self,
|
||||
record: &UpdateManagementTokenRecord,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
self.update_management_token_scoped(record, Some(user_id))
|
||||
}
|
||||
|
||||
async fn delete_management_token(&self, token_id: &str) -> Result<bool, DataLayerError> {
|
||||
let mut items = self
|
||||
.items
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let mut hashes = self
|
||||
.hashes
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let original_len = items.len();
|
||||
items.retain(|item| item.token.id != token_id);
|
||||
Self::remove_hash_for_token(&mut hashes, token_id);
|
||||
Ok(items.len() != original_len)
|
||||
Ok(self.delete_management_token_scoped(token_id, None))
|
||||
}
|
||||
|
||||
async fn delete_management_token_for_user(
|
||||
&self,
|
||||
token_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Ok(self.delete_management_token_scoped(token_id, Some(user_id)))
|
||||
}
|
||||
|
||||
async fn set_management_token_active(
|
||||
@@ -244,43 +379,45 @@ impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository {
|
||||
token_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
let mut items = self
|
||||
.items
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let Some(item) = items.iter_mut().find(|item| item.token.id == token_id) else {
|
||||
return Ok(None);
|
||||
};
|
||||
item.token.is_active = is_active;
|
||||
item.token.updated_at_unix_secs = Self::now_unix_secs();
|
||||
Ok(Some(item.token.clone()))
|
||||
Ok(self.set_management_token_active_scoped(token_id, None, is_active))
|
||||
}
|
||||
|
||||
async fn set_management_token_active_for_user(
|
||||
&self,
|
||||
token_id: &str,
|
||||
user_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
Ok(self.set_management_token_active_scoped(token_id, Some(user_id), is_active))
|
||||
}
|
||||
|
||||
async fn activate_management_token_if_matches(
|
||||
&self,
|
||||
mutation: &ActivateManagementTokenIfMatches,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
self.activate_management_token_if_matches_inner(mutation)
|
||||
}
|
||||
|
||||
async fn delete_inactive_management_token_if_matches(
|
||||
&self,
|
||||
mutation: &ActivateManagementTokenIfMatches,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
self.delete_inactive_management_token_if_matches_inner(mutation)
|
||||
}
|
||||
|
||||
async fn regenerate_management_token_secret(
|
||||
&self,
|
||||
mutation: &RegenerateManagementTokenSecret,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
mutation.validate()?;
|
||||
self.regenerate_management_token_secret_scoped(mutation, None)
|
||||
}
|
||||
|
||||
let mut items = self
|
||||
.items
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let mut hashes = self
|
||||
.hashes
|
||||
.write()
|
||||
.expect("management token repository lock");
|
||||
let Some(item) = items
|
||||
.iter_mut()
|
||||
.find(|item| item.token.id == mutation.token_id)
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
Self::remove_hash_for_token(&mut hashes, &mutation.token_id);
|
||||
hashes.insert(mutation.token_hash.clone(), mutation.token_id.clone());
|
||||
item.token.token_prefix = mutation.token_prefix.clone();
|
||||
item.token.updated_at_unix_secs = Self::now_unix_secs();
|
||||
Ok(Some(item.token.clone()))
|
||||
async fn regenerate_management_token_secret_for_user(
|
||||
&self,
|
||||
mutation: &RegenerateManagementTokenSecret,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
self.regenerate_management_token_secret_scoped(mutation, Some(user_id))
|
||||
}
|
||||
|
||||
async fn record_management_token_usage(
|
||||
@@ -307,10 +444,10 @@ impl ManagementTokenWriteRepository for InMemoryManagementTokenRepository {
|
||||
mod tests {
|
||||
use super::InMemoryManagementTokenRepository;
|
||||
use crate::repository::management_tokens::{
|
||||
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
|
||||
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
|
||||
StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
|
||||
UpdateManagementTokenRecord,
|
||||
ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery,
|
||||
ManagementTokenReadRepository, ManagementTokenWriteRepository,
|
||||
RegenerateManagementTokenSecret, StoredManagementToken, StoredManagementTokenUserSummary,
|
||||
StoredManagementTokenWithUser, UpdateManagementTokenRecord,
|
||||
};
|
||||
|
||||
fn sample_token(id: &str, user_id: &str, is_active: bool) -> StoredManagementTokenWithUser {
|
||||
@@ -456,4 +593,108 @@ mod tests {
|
||||
.expect("hash lookup should succeed");
|
||||
assert!(deleted_by_hash.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn owner_scoped_mutations_never_cross_user_boundaries() {
|
||||
let repository = InMemoryManagementTokenRepository::seed_with_hashes(
|
||||
vec![sample_token("token-1", "user-1", true)],
|
||||
vec![("hash-1".to_string(), "token-1".to_string())],
|
||||
);
|
||||
let update = UpdateManagementTokenRecord {
|
||||
token_id: "token-1".to_string(),
|
||||
name: Some("hijacked".to_string()),
|
||||
description: None,
|
||||
clear_description: false,
|
||||
allowed_ips: None,
|
||||
clear_allowed_ips: false,
|
||||
permissions: None,
|
||||
expires_at_unix_secs: None,
|
||||
clear_expires_at: false,
|
||||
is_active: None,
|
||||
};
|
||||
|
||||
assert!(repository
|
||||
.update_management_token_for_user(&update, "user-2")
|
||||
.await
|
||||
.expect("scoped update should execute")
|
||||
.is_none());
|
||||
assert!(repository
|
||||
.set_management_token_active_for_user("token-1", "user-2", false)
|
||||
.await
|
||||
.expect("scoped toggle should execute")
|
||||
.is_none());
|
||||
assert!(repository
|
||||
.regenerate_management_token_secret_for_user(
|
||||
&RegenerateManagementTokenSecret {
|
||||
token_id: "token-1".to_string(),
|
||||
token_hash: "hash-hijacked".to_string(),
|
||||
token_prefix: Some("ae_hijacked".to_string()),
|
||||
},
|
||||
"user-2",
|
||||
)
|
||||
.await
|
||||
.expect("scoped regeneration should execute")
|
||||
.is_none());
|
||||
assert!(!repository
|
||||
.delete_management_token_for_user("token-1", "user-2")
|
||||
.await
|
||||
.expect("scoped delete should execute"));
|
||||
|
||||
let unchanged = repository
|
||||
.get_management_token_with_user_by_hash("hash-1")
|
||||
.await
|
||||
.expect("original hash lookup should succeed")
|
||||
.expect("token should remain");
|
||||
assert_eq!(unchanged.token.name, "token-1");
|
||||
assert!(unchanged.token.is_active);
|
||||
assert!(repository
|
||||
.get_management_token_with_user_by_hash("hash-hijacked")
|
||||
.await
|
||||
.expect("replacement hash lookup should succeed")
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn install_activation_fails_closed_without_atomic_user_state() {
|
||||
let mut pending = sample_token("token-1", "user-1", false);
|
||||
pending.token.allowed_ips = Some(serde_json::json!(["127.0.0.1"]));
|
||||
pending.token.permissions = Some(serde_json::json!(["admin:proxy_nodes:write"]));
|
||||
pending.token.expires_at_unix_secs = Some(1_800_000_000);
|
||||
let expected_token = pending.token.clone();
|
||||
let repository = InMemoryManagementTokenRepository::seed_with_hashes(
|
||||
[pending],
|
||||
[("hash-1".to_string(), "token-1".to_string())],
|
||||
);
|
||||
let expected = ActivateManagementTokenIfMatches {
|
||||
expected_token,
|
||||
token_hash: "hash-1".to_string(),
|
||||
expected_user_security_version: 4,
|
||||
now_unix_secs: 1_700_000_000,
|
||||
};
|
||||
|
||||
let mut mismatched = expected.clone();
|
||||
mismatched.expected_token.permissions =
|
||||
Some(serde_json::json!(["admin:proxy_nodes:admin"]));
|
||||
assert!(!repository
|
||||
.activate_management_token_if_matches(&mismatched)
|
||||
.await
|
||||
.expect("mismatched activation should execute"));
|
||||
assert!(!repository
|
||||
.activate_management_token_if_matches(&expected)
|
||||
.await
|
||||
.expect("memory activation should fail closed"));
|
||||
assert!(
|
||||
!repository
|
||||
.get_management_token_with_user("token-1")
|
||||
.await
|
||||
.expect("token lookup should execute")
|
||||
.expect("token should remain")
|
||||
.token
|
||||
.is_active
|
||||
);
|
||||
assert!(repository
|
||||
.delete_inactive_management_token_if_matches(&expected)
|
||||
.await
|
||||
.expect("exact inactive snapshot cleanup should execute"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,10 +1,10 @@
|
||||
mod memory;
|
||||
|
||||
pub use aether_data_contracts::repository::management_tokens::{
|
||||
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
|
||||
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
|
||||
StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
|
||||
UpdateManagementTokenRecord,
|
||||
ActivateManagementTokenIfMatches, CreateManagementTokenRecord, ManagementTokenListQuery,
|
||||
ManagementTokenReadRepository, ManagementTokenWriteRepository, RegenerateManagementTokenSecret,
|
||||
StoredManagementToken, StoredManagementTokenListPage, StoredManagementTokenUserSummary,
|
||||
StoredManagementTokenWithUser, UpdateManagementTokenRecord,
|
||||
};
|
||||
#[cfg(feature = "mysql")]
|
||||
pub use aether_data_mysql::MysqlManagementTokenRepository;
|
||||
|
||||
@@ -7,7 +7,7 @@ use async_trait::async_trait;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_contracts::repository::oauth_providers::{
|
||||
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
|
||||
StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
|
||||
StoredOAuthProviderConfig, UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
@@ -65,15 +65,30 @@ impl OAuthProviderReadRepository for InMemoryOAuthProviderRepository {
|
||||
|
||||
#[async_trait]
|
||||
impl OAuthProviderWriteRepository for InMemoryOAuthProviderRepository {
|
||||
async fn upsert_oauth_provider_config(
|
||||
async fn upsert_oauth_provider_config_guarded(
|
||||
&self,
|
||||
record: &UpsertOAuthProviderConfigRecord,
|
||||
) -> Result<StoredOAuthProviderConfig, DataLayerError> {
|
||||
_ldap_exclusive: bool,
|
||||
force_disable: bool,
|
||||
locked_users_snapshot: usize,
|
||||
) -> Result<UpsertOAuthProviderConfigOutcome, DataLayerError> {
|
||||
record.validate()?;
|
||||
|
||||
let mut items = self.items.write().expect("oauth provider repository lock");
|
||||
let now = Self::now_unix_secs();
|
||||
let existing = items.get(&record.provider_type).cloned();
|
||||
if !force_disable
|
||||
&& locked_users_snapshot > 0
|
||||
&& existing
|
||||
.as_ref()
|
||||
.is_some_and(|provider| provider.is_enabled && !record.is_enabled)
|
||||
{
|
||||
return Ok(
|
||||
UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation {
|
||||
affected_count: locked_users_snapshot,
|
||||
},
|
||||
);
|
||||
}
|
||||
let created_at = existing
|
||||
.as_ref()
|
||||
.and_then(|item| item.created_at_unix_ms)
|
||||
@@ -106,13 +121,34 @@ impl OAuthProviderWriteRepository for InMemoryOAuthProviderRepository {
|
||||
.with_timestamps(created_at, now);
|
||||
|
||||
items.insert(record.provider_type.clone(), item.clone());
|
||||
Ok(item)
|
||||
Ok(UpsertOAuthProviderConfigOutcome::Upserted(item))
|
||||
}
|
||||
|
||||
async fn delete_oauth_provider_config(
|
||||
async fn compare_and_swap_oauth_provider_client_secret(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
expected: &str,
|
||||
replacement: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut items = self.items.write().expect("oauth provider repository lock");
|
||||
let Some(item) = items.get_mut(provider_type) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if item.client_secret_encrypted.as_deref() != Some(expected) {
|
||||
return Ok(false);
|
||||
}
|
||||
item.client_secret_encrypted = Some(replacement.to_string());
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn delete_oauth_provider_config_if_unlinked(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
has_links_snapshot: bool,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
if has_links_snapshot {
|
||||
return Ok(false);
|
||||
}
|
||||
let mut items = self.items.write().expect("oauth provider repository lock");
|
||||
Ok(items.remove(provider_type).is_some())
|
||||
}
|
||||
@@ -123,7 +159,8 @@ mod tests {
|
||||
use super::InMemoryOAuthProviderRepository;
|
||||
use crate::repository::oauth_providers::{
|
||||
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
|
||||
StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
|
||||
StoredOAuthProviderConfig, UpsertOAuthProviderConfigOutcome,
|
||||
UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
|
||||
fn sample_provider(provider_type: &str) -> StoredOAuthProviderConfig {
|
||||
@@ -138,19 +175,31 @@ mod tests {
|
||||
}
|
||||
|
||||
fn sample_upsert(provider_type: &str) -> UpsertOAuthProviderConfigRecord {
|
||||
let is_custom_oidc = provider_type.starts_with("custom_oidc");
|
||||
let endpoint_host = if is_custom_oidc {
|
||||
"idp.example".to_string()
|
||||
} else {
|
||||
format!("{provider_type}.example.com")
|
||||
};
|
||||
UpsertOAuthProviderConfigRecord {
|
||||
provider_type: provider_type.to_string(),
|
||||
display_name: format!("{provider_type} display"),
|
||||
client_id: format!("{provider_type}-client"),
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
|
||||
authorization_url_override: Some(format!("https://{provider_type}.example.com/auth")),
|
||||
token_url_override: Some(format!("https://{provider_type}.example.com/token")),
|
||||
userinfo_url_override: None,
|
||||
authorization_url_override: Some(format!("https://{endpoint_host}/auth")),
|
||||
token_url_override: Some(format!("https://{endpoint_host}/token")),
|
||||
userinfo_url_override: is_custom_oidc
|
||||
.then(|| format!("https://{endpoint_host}/userinfo")),
|
||||
scopes: Some(vec!["openid".to_string(), "profile".to_string()]),
|
||||
redirect_uri: format!("https://{provider_type}.example.com/redirect"),
|
||||
frontend_callback_url: "https://frontend.example.com/auth/callback".to_string(),
|
||||
attribute_mapping: Some(serde_json::json!({"email": "email"})),
|
||||
extra_config: Some(serde_json::json!({"team": true})),
|
||||
extra_config: is_custom_oidc.then(|| {
|
||||
serde_json::json!({
|
||||
"allowed_domains": [endpoint_host],
|
||||
"team": true,
|
||||
})
|
||||
}),
|
||||
icon_url: None,
|
||||
is_enabled: true,
|
||||
}
|
||||
@@ -171,28 +220,106 @@ mod tests {
|
||||
assert_eq!(listed[0].provider_type, "github");
|
||||
assert_eq!(listed[1].provider_type, "linuxdo");
|
||||
|
||||
let created = repository
|
||||
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()),
|
||||
..sample_upsert("google")
|
||||
})
|
||||
let UpsertOAuthProviderConfigOutcome::Upserted(created) = repository
|
||||
.upsert_oauth_provider_config_guarded(
|
||||
&UpsertOAuthProviderConfigRecord {
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()),
|
||||
..sample_upsert("custom_oidc")
|
||||
},
|
||||
false,
|
||||
false,
|
||||
0,
|
||||
)
|
||||
.await
|
||||
.expect("create should succeed");
|
||||
.expect("create should succeed")
|
||||
else {
|
||||
panic!("create unexpectedly required confirmation");
|
||||
};
|
||||
assert_eq!(created.client_secret_encrypted.as_deref(), Some("secret-1"));
|
||||
|
||||
let updated = repository
|
||||
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Clear,
|
||||
..sample_upsert("google")
|
||||
})
|
||||
let UpsertOAuthProviderConfigOutcome::Upserted(updated) = repository
|
||||
.upsert_oauth_provider_config_guarded(
|
||||
&UpsertOAuthProviderConfigRecord {
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Clear,
|
||||
..sample_upsert("custom_oidc")
|
||||
},
|
||||
false,
|
||||
false,
|
||||
0,
|
||||
)
|
||||
.await
|
||||
.expect("update should succeed");
|
||||
.expect("update should succeed")
|
||||
else {
|
||||
panic!("update unexpectedly required confirmation");
|
||||
};
|
||||
assert!(updated.client_secret_encrypted.is_none());
|
||||
|
||||
let deleted = repository
|
||||
.delete_oauth_provider_config("google")
|
||||
.delete_oauth_provider_config_if_unlinked("custom_oidc", false)
|
||||
.await
|
||||
.expect("delete should succeed");
|
||||
assert!(deleted);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn client_secret_cas_preserves_concurrent_non_secret_fields_and_timestamp() {
|
||||
let repository = InMemoryOAuthProviderRepository::default();
|
||||
repository
|
||||
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Set("legacy-secret".to_string()),
|
||||
..sample_upsert("custom_oidc")
|
||||
})
|
||||
.await
|
||||
.expect("provider should create");
|
||||
|
||||
let concurrent = repository
|
||||
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
|
||||
display_name: "concurrent display update".to_string(),
|
||||
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
|
||||
..sample_upsert("custom_oidc")
|
||||
})
|
||||
.await
|
||||
.expect("non-secret update should persist");
|
||||
assert!(repository
|
||||
.compare_and_swap_oauth_provider_client_secret(
|
||||
"custom_oidc",
|
||||
"legacy-secret",
|
||||
"record-bound-v2",
|
||||
)
|
||||
.await
|
||||
.expect("secret CAS should execute"));
|
||||
|
||||
let migrated = repository
|
||||
.get_oauth_provider_config("custom_oidc")
|
||||
.await
|
||||
.expect("provider should read")
|
||||
.expect("provider should exist");
|
||||
assert_eq!(migrated.display_name, "concurrent display update");
|
||||
assert_eq!(
|
||||
migrated.updated_at_unix_secs,
|
||||
concurrent.updated_at_unix_secs
|
||||
);
|
||||
assert_eq!(
|
||||
migrated.client_secret_encrypted.as_deref(),
|
||||
Some("record-bound-v2")
|
||||
);
|
||||
assert!(!repository
|
||||
.compare_and_swap_oauth_provider_client_secret(
|
||||
"custom_oidc",
|
||||
"legacy-secret",
|
||||
"must-not-win",
|
||||
)
|
||||
.await
|
||||
.expect("stale CAS should execute"));
|
||||
assert_eq!(
|
||||
repository
|
||||
.get_oauth_provider_config("custom_oidc")
|
||||
.await
|
||||
.expect("provider should read")
|
||||
.expect("provider should exist")
|
||||
.client_secret_encrypted
|
||||
.as_deref(),
|
||||
Some("record-bound-v2")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
mod memory;
|
||||
|
||||
pub use aether_data_contracts::repository::oauth_providers::{
|
||||
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderRepository,
|
||||
OAuthProviderWriteRepository, StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,
|
||||
validate_oauth_frontend_callback_url, validate_oauth_provider_endpoint_config,
|
||||
validate_oauth_redirect_uri, EncryptedSecretUpdate, OAuthProviderReadRepository,
|
||||
OAuthProviderRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
|
||||
UpsertOAuthProviderConfigOutcome, UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
#[cfg(feature = "mysql")]
|
||||
pub use aether_data_mysql::MysqlOAuthProviderRepository;
|
||||
|
||||
@@ -7,10 +7,12 @@ use serde_json::{json, Map, Value};
|
||||
|
||||
use super::{
|
||||
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
|
||||
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
|
||||
ProviderCatalogKeyListQuery, ProviderCatalogKeyOAuthCredentialCasDelete,
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
|
||||
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate,
|
||||
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
|
||||
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
|
||||
ProviderCatalogProviderConfigCasUpdate, ProviderCatalogProxyCasUpdate,
|
||||
ProviderCatalogReadRepository, ProviderCatalogSnapshot,
|
||||
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
|
||||
@@ -454,6 +456,44 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
|
||||
Ok(stored.clone())
|
||||
}
|
||||
|
||||
async fn compare_and_swap_provider_config(
|
||||
&self,
|
||||
update: &ProviderCatalogProviderConfigCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut index = self
|
||||
.index
|
||||
.write()
|
||||
.expect("provider catalog repository lock");
|
||||
let Some(provider) = index.providers.get_mut(&update.provider_id) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if provider.config != update.expected_config {
|
||||
return Ok(false);
|
||||
}
|
||||
provider.config = update.config.clone();
|
||||
provider.updated_at_unix_secs = Some(current_unix_secs());
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn compare_and_swap_provider_proxy(
|
||||
&self,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut index = self
|
||||
.index
|
||||
.write()
|
||||
.expect("provider catalog repository lock");
|
||||
let Some(provider) = index.providers.get_mut(&update.record_id) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if provider.proxy != update.expected_proxy {
|
||||
return Ok(false);
|
||||
}
|
||||
provider.proxy = update.proxy.clone();
|
||||
provider.updated_at_unix_secs = Some(current_unix_secs());
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn delete_provider(&self, provider_id: &str) -> Result<bool, DataLayerError> {
|
||||
let mut index = self
|
||||
.index
|
||||
@@ -504,6 +544,25 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
|
||||
Ok(stored.clone())
|
||||
}
|
||||
|
||||
async fn compare_and_swap_endpoint_proxy(
|
||||
&self,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut index = self
|
||||
.index
|
||||
.write()
|
||||
.expect("provider catalog repository lock");
|
||||
let Some(endpoint) = index.endpoints.get_mut(&update.record_id) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if endpoint.proxy != update.expected_proxy {
|
||||
return Ok(false);
|
||||
}
|
||||
endpoint.proxy = update.proxy.clone();
|
||||
endpoint.updated_at_unix_secs = Some(current_unix_secs());
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn delete_endpoint(&self, endpoint_id: &str) -> Result<bool, DataLayerError> {
|
||||
let mut index = self
|
||||
.index
|
||||
@@ -548,6 +607,52 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
|
||||
Ok(stored.clone())
|
||||
}
|
||||
|
||||
async fn compare_and_swap_key_proxy(
|
||||
&self,
|
||||
update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut index = self
|
||||
.index
|
||||
.write()
|
||||
.expect("provider catalog repository lock");
|
||||
let Some(key) = index.keys.get_mut(&update.record_id) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if key.proxy != update.expected_proxy {
|
||||
return Ok(false);
|
||||
}
|
||||
key.proxy = update.proxy.clone();
|
||||
key.updated_at_unix_secs = Some(current_unix_secs());
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn compare_and_swap_key_credentials(
|
||||
&self,
|
||||
update: &ProviderCatalogKeyCredentialsCasUpdate,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
if update.key_id.trim().is_empty() || update.expected_provider_id.trim().is_empty() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"provider catalog key credential CAS requires key_id and provider_id".to_string(),
|
||||
));
|
||||
}
|
||||
let mut index = self
|
||||
.index
|
||||
.write()
|
||||
.expect("provider catalog repository lock");
|
||||
let Some(key) = index.keys.get_mut(&update.key_id) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if key.provider_id != update.expected_provider_id
|
||||
|| key.encrypted_api_key != update.expected_encrypted_api_key
|
||||
|| key.encrypted_auth_config != update.expected_encrypted_auth_config
|
||||
{
|
||||
return Ok(false);
|
||||
}
|
||||
key.encrypted_api_key = update.encrypted_api_key.clone();
|
||||
key.encrypted_auth_config = update.encrypted_auth_config.clone();
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn compare_and_update_key_admin_state(
|
||||
&self,
|
||||
update: &ProviderCatalogKeyAdminCasUpdate,
|
||||
@@ -882,40 +987,11 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn update_key_oauth_credentials(
|
||||
&self,
|
||||
key_id: &str,
|
||||
encrypted_api_key: &str,
|
||||
encrypted_auth_config: Option<&str>,
|
||||
expires_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
if encrypted_api_key.trim().is_empty() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"provider catalog oauth api_key is empty".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let mut index = self
|
||||
.index
|
||||
.write()
|
||||
.expect("provider catalog repository lock");
|
||||
let Some(key) = index.keys.get_mut(key_id) else {
|
||||
return Ok(false);
|
||||
};
|
||||
|
||||
key.encrypted_api_key = Some(encrypted_api_key.to_string());
|
||||
key.encrypted_auth_config = encrypted_auth_config.map(ToOwned::to_owned);
|
||||
key.expires_at_unix_secs = expires_at_unix_secs;
|
||||
key.updated_at_unix_secs = Some(current_unix_secs());
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn update_key_oauth_runtime_state(
|
||||
&self,
|
||||
key_id: &str,
|
||||
oauth_invalid_at_unix_secs: Option<u64>,
|
||||
oauth_invalid_reason: Option<&str>,
|
||||
encrypted_auth_config_update: Option<&str>,
|
||||
updated_at_unix_secs: Option<u64>,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut index = self
|
||||
@@ -928,9 +1004,6 @@ impl ProviderCatalogWriteRepository for InMemoryProviderCatalogReadRepository {
|
||||
|
||||
key.oauth_invalid_at_unix_secs = oauth_invalid_at_unix_secs;
|
||||
key.oauth_invalid_reason = oauth_invalid_reason.map(ToOwned::to_owned);
|
||||
if let Some(encrypted_auth_config) = encrypted_auth_config_update {
|
||||
key.encrypted_auth_config = Some(encrypted_auth_config.to_string());
|
||||
}
|
||||
key.updated_at_unix_secs = Some(updated_at_unix_secs.unwrap_or_else(current_unix_secs));
|
||||
Ok(true)
|
||||
}
|
||||
@@ -1576,15 +1649,15 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn updates_oauth_credentials_for_existing_key() {
|
||||
async fn unfenced_oauth_runtime_state_update_preserves_credentials() {
|
||||
let repository = InMemoryProviderCatalogReadRepository::seed(
|
||||
vec![sample_provider("provider-1")],
|
||||
vec![sample_endpoint("endpoint-1", "provider-1")],
|
||||
vec![sample_key("key-1", "provider-1")
|
||||
.with_transport_fields(
|
||||
None,
|
||||
"ciphertext-placeholder".to_string(),
|
||||
Some("ciphertext-auth-1".to_string()),
|
||||
"ciphertext-api".to_string(),
|
||||
Some("ciphertext-auth".to_string()),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
@@ -1596,29 +1669,26 @@ mod tests {
|
||||
);
|
||||
|
||||
assert!(repository
|
||||
.update_key_oauth_credentials(
|
||||
"key-1",
|
||||
"ciphertext-updated-token",
|
||||
Some("ciphertext-auth-2"),
|
||||
Some(4_102_444_800),
|
||||
)
|
||||
.update_key_oauth_runtime_state("key-1", Some(123), Some("refresh failed"), Some(456),)
|
||||
.await
|
||||
.expect("update should succeed"));
|
||||
.expect("runtime state should update"));
|
||||
|
||||
let stored = repository
|
||||
.list_keys_by_ids(&["key-1".to_string()])
|
||||
.await
|
||||
.expect("keys should read");
|
||||
assert_eq!(stored.len(), 1);
|
||||
.expect("key should read")
|
||||
.pop()
|
||||
.expect("key should exist");
|
||||
assert_eq!(stored.encrypted_api_key.as_deref(), Some("ciphertext-api"));
|
||||
assert_eq!(
|
||||
stored[0].encrypted_api_key.as_deref(),
|
||||
Some("ciphertext-updated-token")
|
||||
stored.encrypted_auth_config.as_deref(),
|
||||
Some("ciphertext-auth")
|
||||
);
|
||||
assert_eq!(stored.oauth_invalid_at_unix_secs, Some(123));
|
||||
assert_eq!(
|
||||
stored[0].encrypted_auth_config.as_deref(),
|
||||
Some("ciphertext-auth-2")
|
||||
stored.oauth_invalid_reason.as_deref(),
|
||||
Some("refresh failed")
|
||||
);
|
||||
assert_eq!(stored[0].expires_at_unix_secs, Some(4_102_444_800));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -3,11 +3,12 @@ mod memory;
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::provider_catalog::{
|
||||
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
|
||||
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
|
||||
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate,
|
||||
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence,
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogProviderConfigCasUpdate,
|
||||
ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
|
||||
ProviderCatalogUpstreamMetadataNamespaceExpectation,
|
||||
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
const KIRO_DEVICE_AUTH_SESSION_PREFIX: &str = "device_auth_session:";
|
||||
const PROVIDER_OAUTH_BATCH_TASK_PREFIX: &str = "provider_oauth_batch_task:";
|
||||
const PROVIDER_OAUTH_STATE_PREFIX: &str = "provider_oauth_state:";
|
||||
@@ -8,7 +10,11 @@ pub const PROVIDER_OAUTH_STATE_TTL_SECS: u64 = 600;
|
||||
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredAdminProviderOAuthDeviceSession {
|
||||
pub session_id: String,
|
||||
pub provider_id: String,
|
||||
pub initiated_by_user_id: String,
|
||||
pub initiated_by_session_id: Option<String>,
|
||||
pub initiated_by_management_token_id: Option<String>,
|
||||
pub region: String,
|
||||
pub client_id: String,
|
||||
pub client_secret: String,
|
||||
@@ -34,28 +40,54 @@ pub struct StoredAdminProviderOAuthDeviceSession {
|
||||
pub error_msg: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredAdminProviderOAuthState {
|
||||
pub nonce: String,
|
||||
pub key_id: String,
|
||||
pub provider_id: String,
|
||||
pub provider_type: String,
|
||||
pub pkce_verifier: Option<String>,
|
||||
#[serde(default)]
|
||||
pub expected_encrypted_auth_config: Option<String>,
|
||||
pub initiated_by_user_id: String,
|
||||
#[serde(default)]
|
||||
pub initiated_by_session_id: Option<String>,
|
||||
#[serde(default)]
|
||||
pub initiated_by_management_token_id: Option<String>,
|
||||
pub created_at: u64,
|
||||
}
|
||||
|
||||
pub fn provider_oauth_device_session_storage_key(session_id: &str) -> String {
|
||||
format!("{KIRO_DEVICE_AUTH_SESSION_PREFIX}{session_id}")
|
||||
}
|
||||
|
||||
pub fn provider_oauth_device_session_secret_purpose(session_id: &str) -> String {
|
||||
let storage_key = provider_oauth_device_session_storage_key(session_id);
|
||||
format!(
|
||||
"provider-oauth-device-session:sha256:{:x}",
|
||||
Sha256::digest(storage_key.as_bytes())
|
||||
)
|
||||
}
|
||||
|
||||
pub fn provider_oauth_state_storage_key(nonce: &str) -> String {
|
||||
format!("{PROVIDER_OAUTH_STATE_PREFIX}{nonce}")
|
||||
format!(
|
||||
"{PROVIDER_OAUTH_STATE_PREFIX}sha256:{:x}",
|
||||
Sha256::digest(nonce.as_bytes())
|
||||
)
|
||||
}
|
||||
|
||||
pub fn provider_oauth_batch_task_storage_key(task_id: &str) -> String {
|
||||
format!("{PROVIDER_OAUTH_BATCH_TASK_PREFIX}{task_id}")
|
||||
}
|
||||
|
||||
pub fn provider_oauth_batch_task_secret_purpose(task_id: &str) -> String {
|
||||
let storage_key = provider_oauth_batch_task_storage_key(task_id);
|
||||
format!(
|
||||
"provider-oauth-batch-task:sha256:{:x}",
|
||||
Sha256::digest(storage_key.as_bytes())
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_provider_oauth_batch_task_status_payload(
|
||||
provider_id: &str,
|
||||
state: &serde_json::Map<String, serde_json::Value>,
|
||||
@@ -142,7 +174,8 @@ pub fn build_provider_oauth_batch_task_status_payload(
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_storage_key,
|
||||
build_provider_oauth_batch_task_status_payload, provider_oauth_batch_task_secret_purpose,
|
||||
provider_oauth_batch_task_storage_key, provider_oauth_device_session_secret_purpose,
|
||||
provider_oauth_device_session_storage_key, provider_oauth_state_storage_key,
|
||||
KIRO_DEVICE_AUTH_SESSION_TTL_BUFFER_SECS, PROVIDER_OAUTH_BATCH_TASK_TTL_SECS,
|
||||
PROVIDER_OAUTH_STATE_TTL_SECS,
|
||||
@@ -155,14 +188,23 @@ mod tests {
|
||||
provider_oauth_device_session_storage_key("session-123"),
|
||||
"device_auth_session:session-123"
|
||||
);
|
||||
assert_eq!(
|
||||
provider_oauth_state_storage_key("nonce-123"),
|
||||
"provider_oauth_state:nonce-123"
|
||||
);
|
||||
let first_purpose = provider_oauth_device_session_secret_purpose("session-123");
|
||||
let second_purpose = provider_oauth_device_session_secret_purpose("session-456");
|
||||
assert!(first_purpose.starts_with("provider-oauth-device-session:sha256:"));
|
||||
assert!(!first_purpose.contains("session-123"));
|
||||
assert_ne!(first_purpose, second_purpose);
|
||||
let state_key = provider_oauth_state_storage_key("nonce-123");
|
||||
assert!(state_key.starts_with("provider_oauth_state:sha256:"));
|
||||
assert!(!state_key.contains("nonce-123"));
|
||||
assert_eq!(
|
||||
provider_oauth_batch_task_storage_key("task-123"),
|
||||
"provider_oauth_batch_task:task-123"
|
||||
);
|
||||
let first_task_purpose = provider_oauth_batch_task_secret_purpose("task-123");
|
||||
let second_task_purpose = provider_oauth_batch_task_secret_purpose("task-456");
|
||||
assert!(first_task_purpose.starts_with("provider-oauth-batch-task:sha256:"));
|
||||
assert!(!first_task_purpose.contains("task-123"));
|
||||
assert_ne!(first_task_purpose, second_task_purpose);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -10,13 +10,14 @@ use super::log_reported_tunnel_error_event;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_contracts::repository::proxy_nodes::{
|
||||
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
|
||||
normalize_proxy_metadata, preserve_proxy_metadata_tunnel_security,
|
||||
reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery, ProxyNodeHeartbeatMutation,
|
||||
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeMetricsCleanupSummary,
|
||||
ProxyNodeMetricsStep, ProxyNodeReadRepository, ProxyNodeRegistrationMutation,
|
||||
ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, ProxyNodeTunnelStatusMutation,
|
||||
ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket, StoredProxyNode, StoredProxyNodeEvent,
|
||||
StoredProxyNodeMetricsBucket, TunnelMetricsSample, PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
|
||||
merge_proxy_metadata_for_registration, normalize_heartbeat_proxy_metadata,
|
||||
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery,
|
||||
ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
|
||||
ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository,
|
||||
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
|
||||
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket,
|
||||
StoredProxyNode, StoredProxyNodeEvent, StoredProxyNodeMetricsBucket, TunnelMetricsSample,
|
||||
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
|
||||
};
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
@@ -357,6 +358,42 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
Ok(updated)
|
||||
}
|
||||
|
||||
async fn compare_and_set_proxy_password(
|
||||
&self,
|
||||
node_id: &str,
|
||||
expected: &str,
|
||||
replacement: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut nodes = self.nodes.write().expect("proxy node repository lock");
|
||||
let Some(node) = nodes.get_mut(node_id) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if node.proxy_password.as_deref() != Some(expected) {
|
||||
return Ok(false);
|
||||
}
|
||||
node.proxy_password = Some(replacement.to_string());
|
||||
node.updated_at_unix_secs = Self::now_unix_secs();
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn compare_and_set_proxy_metadata(
|
||||
&self,
|
||||
node_id: &str,
|
||||
expected: &serde_json::Value,
|
||||
replacement: &serde_json::Value,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut nodes = self.nodes.write().expect("proxy node repository lock");
|
||||
let Some(node) = nodes.get_mut(node_id) else {
|
||||
return Ok(false);
|
||||
};
|
||||
if node.proxy_metadata.as_ref() != Some(expected) {
|
||||
return Ok(false);
|
||||
}
|
||||
node.proxy_metadata = Some(replacement.clone());
|
||||
node.updated_at_unix_secs = Self::now_unix_secs();
|
||||
Ok(true)
|
||||
}
|
||||
|
||||
async fn create_manual_node(
|
||||
&self,
|
||||
mutation: &ProxyNodeManualCreateMutation,
|
||||
@@ -369,9 +406,14 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
return Err(Self::duplicate_proxy_node_error(existing));
|
||||
}
|
||||
|
||||
let node_id = requested_proxy_node_id(mutation.node_id.as_deref())?
|
||||
.unwrap_or_else(|| Uuid::new_v4().to_string());
|
||||
if let Some(existing) = nodes.get(&node_id) {
|
||||
return Err(proxy_node_id_in_use_error(existing));
|
||||
}
|
||||
let now = Self::now_unix_secs();
|
||||
let node = StoredProxyNode::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
node_id,
|
||||
mutation.name.clone(),
|
||||
mutation.ip.clone(),
|
||||
mutation.port,
|
||||
@@ -466,18 +508,35 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
) -> Result<StoredProxyNode, DataLayerError> {
|
||||
let mut nodes = self.nodes.write().expect("proxy node repository lock");
|
||||
let now = Self::now_unix_secs();
|
||||
let normalized_proxy_metadata = normalize_proxy_metadata(
|
||||
mutation.proxy_metadata.as_ref(),
|
||||
mutation.proxy_version.as_deref(),
|
||||
let normalized_proxy_metadata = merge_proxy_metadata_for_registration(
|
||||
None,
|
||||
normalize_proxy_metadata(
|
||||
mutation.proxy_metadata.as_ref(),
|
||||
mutation.proxy_version.as_deref(),
|
||||
),
|
||||
);
|
||||
|
||||
if let Some(existing_id) = nodes
|
||||
.iter()
|
||||
.find(|(_, node)| {
|
||||
.filter(|(_, node)| {
|
||||
!node.is_manual && node.ip == mutation.ip && node.port == mutation.port
|
||||
})
|
||||
.min_by(|(_, left), (_, right)| {
|
||||
left.created_at_unix_ms
|
||||
.unwrap_or(u64::MAX)
|
||||
.cmp(&right.created_at_unix_ms.unwrap_or(u64::MAX))
|
||||
.then(left.id.cmp(&right.id))
|
||||
})
|
||||
.map(|(node_id, _)| node_id.clone())
|
||||
{
|
||||
if let Some(requested_id) = requested_proxy_node_id(mutation.node_id.as_deref())? {
|
||||
if requested_id != existing_id {
|
||||
return Err(proxy_node_registration_identity_error(
|
||||
&requested_id,
|
||||
&existing_id,
|
||||
));
|
||||
}
|
||||
}
|
||||
let node = nodes
|
||||
.get_mut(&existing_id)
|
||||
.expect("existing proxy node should be present");
|
||||
@@ -505,9 +564,10 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
if let Some(estimated_max_concurrency) = mutation.estimated_max_concurrency {
|
||||
node.estimated_max_concurrency = Some(estimated_max_concurrency);
|
||||
}
|
||||
if let Some(proxy_metadata) = normalized_proxy_metadata {
|
||||
node.proxy_metadata = Some(proxy_metadata);
|
||||
}
|
||||
node.proxy_metadata = merge_proxy_metadata_for_registration(
|
||||
node.proxy_metadata.as_ref(),
|
||||
normalized_proxy_metadata,
|
||||
);
|
||||
if node.created_at_unix_ms.is_none() {
|
||||
node.created_at_unix_ms = now;
|
||||
}
|
||||
@@ -515,8 +575,13 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
return Ok(node.clone());
|
||||
}
|
||||
|
||||
let node_id = requested_proxy_node_id(mutation.node_id.as_deref())?
|
||||
.unwrap_or_else(|| Uuid::new_v4().to_string());
|
||||
if let Some(existing) = nodes.get(&node_id) {
|
||||
return Err(proxy_node_id_in_use_error(existing));
|
||||
}
|
||||
let mut node = StoredProxyNode::new(
|
||||
Uuid::new_v4().to_string(),
|
||||
node_id,
|
||||
mutation.name.clone(),
|
||||
mutation.ip.clone(),
|
||||
mutation.port,
|
||||
@@ -555,11 +620,18 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
&self,
|
||||
mutation: &ProxyNodeHeartbeatMutation,
|
||||
) -> Result<Option<StoredProxyNode>, DataLayerError> {
|
||||
let mut nodes = self.nodes.write().expect("proxy node repository lock");
|
||||
let (node, sample, now_unix_secs) = {
|
||||
let mut nodes = self.nodes.write().expect("proxy node repository lock");
|
||||
let Some(node) = nodes.get_mut(&mutation.node_id) else {
|
||||
return Ok(None);
|
||||
};
|
||||
if mutation
|
||||
.expected_tunnel_generation
|
||||
.as_deref()
|
||||
.is_some_and(|expected| expected != node.tunnel_generation)
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
if !node.tunnel_mode {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"non-tunnel mode is no longer supported, please upgrade aether-tunnel to use tunnel mode"
|
||||
@@ -587,14 +659,11 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
if let Some(value) = mutation.avg_latency_ms {
|
||||
node.avg_latency_ms = Some(value);
|
||||
}
|
||||
let normalized_proxy_metadata = normalize_proxy_metadata(
|
||||
let normalized_proxy_metadata = normalize_heartbeat_proxy_metadata(
|
||||
previous_proxy_metadata.as_ref(),
|
||||
mutation.proxy_metadata.as_ref(),
|
||||
mutation.proxy_version.as_deref(),
|
||||
);
|
||||
let normalized_proxy_metadata = preserve_proxy_metadata_tunnel_security(
|
||||
previous_proxy_metadata.as_ref(),
|
||||
normalized_proxy_metadata,
|
||||
);
|
||||
if let Some(value) = normalized_proxy_metadata {
|
||||
node.proxy_metadata = Some(value);
|
||||
}
|
||||
@@ -671,6 +740,7 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
});
|
||||
}
|
||||
}
|
||||
drop(nodes);
|
||||
|
||||
Ok(Some(node))
|
||||
}
|
||||
@@ -686,6 +756,12 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
if !node.is_manual {
|
||||
return Ok(false);
|
||||
}
|
||||
let Some(expected_generation) = mutation.expected_tunnel_generation.as_deref() else {
|
||||
return Ok(false);
|
||||
};
|
||||
if expected_generation != node.tunnel_generation {
|
||||
return Ok(false);
|
||||
}
|
||||
|
||||
node.total_requests += mutation.total_requests_delta.max(0);
|
||||
node.failed_requests += mutation.failed_requests_delta.max(0);
|
||||
@@ -703,6 +779,13 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
let Some(node) = nodes.get_mut(&mutation.node_id) else {
|
||||
return Ok(None);
|
||||
};
|
||||
if mutation
|
||||
.expected_tunnel_generation
|
||||
.as_deref()
|
||||
.is_some_and(|expected| expected != node.tunnel_generation)
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let event_time = mutation
|
||||
.observed_at_unix_secs
|
||||
@@ -776,11 +859,11 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
}
|
||||
|
||||
async fn delete_node(&self, node_id: &str) -> Result<Option<StoredProxyNode>, DataLayerError> {
|
||||
let removed = self
|
||||
.nodes
|
||||
.write()
|
||||
.expect("proxy node repository lock")
|
||||
.remove(node_id);
|
||||
// Keep the parent lock until all child state is removed. Registration
|
||||
// also takes this lock, so the same id cannot be recreated between the
|
||||
// parent delete and cleanup of its events or metrics.
|
||||
let mut nodes = self.nodes.write().expect("proxy node repository lock");
|
||||
let removed = nodes.remove(node_id);
|
||||
if removed.is_some() {
|
||||
self.events
|
||||
.write()
|
||||
@@ -795,6 +878,7 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
.expect("proxy node repository lock")
|
||||
.retain(|(metric_node_id, _), _| metric_node_id != node_id);
|
||||
}
|
||||
drop(nodes);
|
||||
Ok(removed)
|
||||
}
|
||||
|
||||
@@ -806,6 +890,13 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
let Some(node) = nodes.get_mut(&mutation.node_id) else {
|
||||
return Ok(None);
|
||||
};
|
||||
if mutation
|
||||
.expected_tunnel_generation
|
||||
.as_deref()
|
||||
.is_some_and(|expected| expected != node.tunnel_generation)
|
||||
{
|
||||
return Ok(None);
|
||||
}
|
||||
if node.is_manual {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"手动节点不支持远程配置下发".to_string(),
|
||||
@@ -886,6 +977,31 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
|
||||
}
|
||||
}
|
||||
|
||||
fn requested_proxy_node_id(value: Option<&str>) -> Result<Option<String>, DataLayerError> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
if value.is_empty() || value.trim() != value {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"proxy node id must be non-empty and unpadded".to_string(),
|
||||
));
|
||||
}
|
||||
Ok(Some(value.to_string()))
|
||||
}
|
||||
|
||||
fn proxy_node_registration_identity_error(requested_id: &str, existing_id: &str) -> DataLayerError {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"proxy node registration identity changed: requested {requested_id}, existing {existing_id}"
|
||||
))
|
||||
}
|
||||
|
||||
fn proxy_node_id_in_use_error(node: &StoredProxyNode) -> DataLayerError {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"proxy node id is already in use: {} ({}:{})",
|
||||
node.id, node.ip, node.port
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::InMemoryProxyNodeRepository;
|
||||
@@ -894,7 +1010,7 @@ mod tests {
|
||||
ProxyNodeRemoteConfigMutation, ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository,
|
||||
StoredProxyNode, StoredProxyNodeEvent,
|
||||
};
|
||||
use serde_json::json;
|
||||
use serde_json::{json, Value};
|
||||
|
||||
fn sample_node() -> StoredProxyNode {
|
||||
StoredProxyNode::new(
|
||||
@@ -937,6 +1053,7 @@ mod tests {
|
||||
let heartbeat = repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-1".to_string(),
|
||||
expected_tunnel_generation: None,
|
||||
heartbeat_interval: Some(45),
|
||||
active_connections: Some(5),
|
||||
total_requests_delta: Some(8),
|
||||
@@ -944,7 +1061,13 @@ mod tests {
|
||||
failed_requests_delta: Some(2),
|
||||
dns_failures_delta: Some(1),
|
||||
stream_errors_delta: Some(3),
|
||||
proxy_metadata: Some(json!({"arch": "arm64"})),
|
||||
proxy_metadata: Some(json!({
|
||||
"arch": "arm64",
|
||||
"tunnel_security": {
|
||||
"mode": "disabled",
|
||||
"encryption_key": "attacker-controlled"
|
||||
}
|
||||
})),
|
||||
proxy_version: Some("1.2.3".to_string()),
|
||||
})
|
||||
.await
|
||||
@@ -966,10 +1089,16 @@ mod tests {
|
||||
.and_then(|value| value.as_str()),
|
||||
Some("1.2.3")
|
||||
);
|
||||
assert!(heartbeat
|
||||
.proxy_metadata
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("tunnel_security"))
|
||||
.is_none());
|
||||
|
||||
let stale = repository
|
||||
.update_tunnel_status(&ProxyNodeTunnelStatusMutation {
|
||||
node_id: "node-1".to_string(),
|
||||
expected_tunnel_generation: None,
|
||||
connected: false,
|
||||
conn_count: 0,
|
||||
detail: None,
|
||||
@@ -994,6 +1123,7 @@ mod tests {
|
||||
let updated = repository
|
||||
.update_tunnel_status(&ProxyNodeTunnelStatusMutation {
|
||||
node_id: "node-1".to_string(),
|
||||
expected_tunnel_generation: None,
|
||||
connected: false,
|
||||
conn_count: 0,
|
||||
detail: None,
|
||||
@@ -1060,6 +1190,96 @@ mod tests {
|
||||
assert_eq!(events[0].detail.as_deref(), Some("newer"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn delete_cleans_child_state_before_same_id_can_be_reused() {
|
||||
let old_node = sample_node();
|
||||
let old_generation = old_node.tunnel_generation.clone();
|
||||
let repository = InMemoryProxyNodeRepository::seed_with_events(
|
||||
vec![old_node],
|
||||
vec![StoredProxyNodeEvent {
|
||||
id: 1,
|
||||
node_id: "node-1".to_string(),
|
||||
event_type: "connected".to_string(),
|
||||
detail: Some("old incarnation".to_string()),
|
||||
event_metadata: None,
|
||||
created_at_unix_ms: Some(1_710_000_000),
|
||||
}],
|
||||
);
|
||||
|
||||
repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: "node-1".to_string(),
|
||||
expected_tunnel_generation: Some(old_generation.clone()),
|
||||
heartbeat_interval: None,
|
||||
active_connections: Some(1),
|
||||
total_requests_delta: None,
|
||||
avg_latency_ms: None,
|
||||
failed_requests_delta: None,
|
||||
dns_failures_delta: None,
|
||||
stream_errors_delta: None,
|
||||
proxy_metadata: Some(json!({
|
||||
"tunnel_metrics": {
|
||||
"connect_errors": 0,
|
||||
"disconnects": 0,
|
||||
"error_events_total": 0,
|
||||
"ws_in_bytes": 0,
|
||||
"ws_out_bytes": 0,
|
||||
"ws_in_frames": 0,
|
||||
"ws_out_frames": 0,
|
||||
"heartbeat_rtt_last_ms": 1
|
||||
}
|
||||
})),
|
||||
proxy_version: None,
|
||||
})
|
||||
.await
|
||||
.expect("heartbeat should create metric buckets")
|
||||
.expect("old node should exist");
|
||||
|
||||
repository
|
||||
.delete_node("node-1")
|
||||
.await
|
||||
.expect("delete should succeed")
|
||||
.expect("old node should be removed");
|
||||
|
||||
let replacement = repository
|
||||
.register_node(&ProxyNodeRegistrationMutation {
|
||||
node_id: Some("node-1".to_string()),
|
||||
name: "replacement".to_string(),
|
||||
ip: "127.0.0.2".to_string(),
|
||||
port: 7002,
|
||||
region: None,
|
||||
heartbeat_interval: 30,
|
||||
active_connections: None,
|
||||
total_requests: None,
|
||||
avg_latency_ms: None,
|
||||
hardware_info: None,
|
||||
estimated_max_concurrency: None,
|
||||
proxy_metadata: None,
|
||||
proxy_version: None,
|
||||
registered_by: None,
|
||||
tunnel_mode: true,
|
||||
})
|
||||
.await
|
||||
.expect("same id should be reusable after delete");
|
||||
assert_ne!(replacement.tunnel_generation, old_generation);
|
||||
|
||||
assert!(repository
|
||||
.list_proxy_node_events("node-1", 10)
|
||||
.await
|
||||
.expect("events should read")
|
||||
.is_empty());
|
||||
for step in [
|
||||
crate::repository::proxy_nodes::ProxyNodeMetricsStep::OneMinute,
|
||||
crate::repository::proxy_nodes::ProxyNodeMetricsStep::OneHour,
|
||||
] {
|
||||
assert!(repository
|
||||
.list_proxy_node_metrics("node-1", step, 0, u64::MAX, 10)
|
||||
.await
|
||||
.expect("metrics should read")
|
||||
.is_empty());
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resets_stale_tunnel_statuses_without_touching_manual_nodes() {
|
||||
let mut stale_tunnel = sample_node();
|
||||
@@ -1101,12 +1321,127 @@ mod tests {
|
||||
assert_eq!(manual.active_connections, 4);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn registration_rejects_rebinding_existing_endpoint_to_different_node_id() {
|
||||
let repository = InMemoryProxyNodeRepository::default();
|
||||
let mutation = ProxyNodeRegistrationMutation {
|
||||
node_id: Some("stable-node-id".to_string()),
|
||||
name: "stable-node".to_string(),
|
||||
ip: "127.0.0.9".to_string(),
|
||||
port: 7009,
|
||||
region: None,
|
||||
heartbeat_interval: 30,
|
||||
active_connections: None,
|
||||
total_requests: None,
|
||||
avg_latency_ms: None,
|
||||
hardware_info: None,
|
||||
estimated_max_concurrency: None,
|
||||
proxy_metadata: Some(json!({"secret_marker": "first"})),
|
||||
proxy_version: None,
|
||||
registered_by: None,
|
||||
tunnel_mode: true,
|
||||
};
|
||||
let registered = repository
|
||||
.register_node(&mutation)
|
||||
.await
|
||||
.expect("initial registration should succeed");
|
||||
assert_eq!(registered.id, "stable-node-id");
|
||||
|
||||
let mut conflicting = mutation;
|
||||
conflicting.node_id = Some("replacement-node-id".to_string());
|
||||
conflicting.proxy_metadata = Some(json!({"secret_marker": "replacement"}));
|
||||
assert!(repository.register_node(&conflicting).await.is_err());
|
||||
let persisted = repository
|
||||
.find_proxy_node("stable-node-id")
|
||||
.await
|
||||
.expect("stable node should read")
|
||||
.expect("stable node should remain");
|
||||
assert_eq!(
|
||||
persisted
|
||||
.proxy_metadata
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("secret_marker")),
|
||||
Some(&json!("first"))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn registration_preserves_omitted_security_and_allows_rotation() {
|
||||
let repository = InMemoryProxyNodeRepository::default();
|
||||
let first_mutation = ProxyNodeRegistrationMutation {
|
||||
node_id: Some("registration-security-node".to_string()),
|
||||
name: "registration-security-node".to_string(),
|
||||
ip: "127.0.0.70".to_string(),
|
||||
port: 7070,
|
||||
region: None,
|
||||
heartbeat_interval: 30,
|
||||
active_connections: None,
|
||||
total_requests: None,
|
||||
avg_latency_ms: None,
|
||||
hardware_info: None,
|
||||
estimated_max_concurrency: None,
|
||||
proxy_metadata: Some(json!({
|
||||
"version": "1.0.0",
|
||||
"tunnel_security": {
|
||||
"mode": "non_tls_required",
|
||||
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
|
||||
}
|
||||
})),
|
||||
proxy_version: None,
|
||||
registered_by: None,
|
||||
tunnel_mode: true,
|
||||
};
|
||||
let first = repository
|
||||
.register_node(&first_mutation)
|
||||
.await
|
||||
.expect("first registration should succeed");
|
||||
|
||||
let mut refreshed_mutation = first_mutation.clone();
|
||||
refreshed_mutation.name = "registration-security-node-refreshed".to_string();
|
||||
refreshed_mutation.proxy_metadata = Some(json!({"runtime": "refreshed"}));
|
||||
refreshed_mutation.proxy_version = Some("2.0.0".to_string());
|
||||
let refreshed = repository
|
||||
.register_node(&refreshed_mutation)
|
||||
.await
|
||||
.expect("metadata-only re-registration should succeed");
|
||||
assert_eq!(refreshed.id, first.id);
|
||||
assert_eq!(
|
||||
refreshed
|
||||
.proxy_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted"))
|
||||
.and_then(Value::as_str),
|
||||
Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old")
|
||||
);
|
||||
|
||||
let mut rotated_mutation = refreshed_mutation;
|
||||
rotated_mutation.proxy_metadata = Some(json!({
|
||||
"tunnel_security": {
|
||||
"mode": "non_tls_required",
|
||||
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new"
|
||||
}
|
||||
}));
|
||||
let rotated = repository
|
||||
.register_node(&rotated_mutation)
|
||||
.await
|
||||
.expect("explicit security rotation should succeed");
|
||||
assert_eq!(
|
||||
rotated
|
||||
.proxy_metadata
|
||||
.as_ref()
|
||||
.and_then(|metadata| metadata.pointer("/tunnel_security/encryption_key_encrypted"))
|
||||
.and_then(Value::as_str),
|
||||
Some("aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new")
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn registers_updates_config_and_unregisters_nodes() {
|
||||
let repository = InMemoryProxyNodeRepository::default();
|
||||
|
||||
let registered = repository
|
||||
.register_node(&ProxyNodeRegistrationMutation {
|
||||
node_id: None,
|
||||
name: "proxy-01".to_string(),
|
||||
ip: "127.0.0.1".to_string(),
|
||||
port: 0,
|
||||
@@ -1131,6 +1466,7 @@ mod tests {
|
||||
let updated = repository
|
||||
.update_remote_config(&ProxyNodeRemoteConfigMutation {
|
||||
node_id: registered.id.clone(),
|
||||
expected_tunnel_generation: None,
|
||||
node_name: Some("proxy-02".to_string()),
|
||||
allowed_ports: Some(vec![443, 8443]),
|
||||
log_level: Some("info".to_string()),
|
||||
@@ -1160,6 +1496,7 @@ mod tests {
|
||||
let after_upgrade = repository
|
||||
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
|
||||
node_id: registered.id.clone(),
|
||||
expected_tunnel_generation: None,
|
||||
heartbeat_interval: None,
|
||||
active_connections: Some(2),
|
||||
total_requests_delta: Some(1),
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -26,6 +26,7 @@ pub enum AdminSystemUsageAggregateImportMode {
|
||||
Skip,
|
||||
Overwrite,
|
||||
Error,
|
||||
ValidateError,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
|
||||
@@ -4,7 +4,8 @@ use std::sync::RwLock;
|
||||
|
||||
use aether_ai_formats::UPSTREAM_IS_STREAM_KEY;
|
||||
use aether_data_contracts::repository::usage::{
|
||||
parse_usage_body_ref, usage_body_ref, StoredUsageAuditAggregation, StoredUsageAuditSummary,
|
||||
canonical_usage_body_ref_for, parse_usage_body_ref, sanitize_usage_request_metadata,
|
||||
usage_body_ref, StoredUsageAuditAggregation, StoredUsageAuditSummary,
|
||||
StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
|
||||
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
|
||||
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
|
||||
@@ -31,7 +32,8 @@ use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
api_key_usage_contribution, provider_api_key_usage_contribution,
|
||||
strip_deprecated_usage_display_fields, usage_can_recover_terminal_failure,
|
||||
sanitize_usage_capture_controls_for_persistence, sanitize_usage_for_persistence,
|
||||
usage_can_recover_terminal_failure, usage_lifecycle_update_allowed,
|
||||
usage_request_metadata_client_family, ApiKeyUsageContribution, ApiKeyUsageDelta,
|
||||
ProviderApiKeyUsageContribution, ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest,
|
||||
StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary,
|
||||
@@ -159,10 +161,6 @@ fn usage_status_is_finalized(status: &str) -> bool {
|
||||
matches!(status, "completed" | "failed" | "cancelled")
|
||||
}
|
||||
|
||||
fn usage_status_is_lifecycle(status: &str) -> bool {
|
||||
matches!(status, "pending" | "streaming")
|
||||
}
|
||||
|
||||
fn merge_usage_timing(existing: Option<u64>, incoming: Option<u64>) -> Option<u64> {
|
||||
match incoming {
|
||||
Some(0) | None => existing.or(incoming),
|
||||
@@ -1176,18 +1174,19 @@ impl UsageReadRepository for InMemoryUsageReadRepository {
|
||||
}
|
||||
|
||||
async fn resolve_body_ref(&self, body_ref: &str) -> Result<Option<Value>, DataLayerError> {
|
||||
let Some((request_id, field)) = parse_usage_body_ref(body_ref) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let canonical_ref = usage_body_ref(&request_id, field);
|
||||
if let Some(value) = self
|
||||
.detached_bodies
|
||||
.read()
|
||||
.expect("usage repository lock")
|
||||
.get(body_ref)
|
||||
.get(&canonical_ref)
|
||||
.cloned()
|
||||
{
|
||||
return Ok(Some(value));
|
||||
}
|
||||
let Some((request_id, field)) = parse_usage_body_ref(body_ref) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let usage = self
|
||||
.by_request_id
|
||||
.read()
|
||||
@@ -2676,44 +2675,74 @@ fn usage_body_ref_from_metadata(
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|object| object.get(field.as_ref_key()))
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.and_then(parse_usage_body_ref)
|
||||
.filter(|(parsed_request_id, parsed_field)| {
|
||||
parsed_request_id == request_id && *parsed_field == field
|
||||
})
|
||||
.map(|(parsed_request_id, parsed_field)| usage_body_ref(&parsed_request_id, parsed_field))
|
||||
.and_then(|body_ref| canonical_usage_body_ref_for(body_ref, request_id, field))
|
||||
}
|
||||
|
||||
fn sanitize_memory_request_metadata(metadata: Option<Value>) -> Option<Value> {
|
||||
sanitize_usage_request_metadata(metadata)
|
||||
}
|
||||
|
||||
fn hydrate_legacy_body_refs(item: &mut StoredRequestUsageAudit) {
|
||||
if item.request_body_ref.is_none() {
|
||||
item.request_body_ref = usage_body_ref_from_metadata(
|
||||
item.request_metadata.as_ref(),
|
||||
&item.request_id,
|
||||
UsageBodyField::RequestBody,
|
||||
);
|
||||
}
|
||||
if item.provider_request_body_ref.is_none() {
|
||||
item.provider_request_body_ref = usage_body_ref_from_metadata(
|
||||
item.request_metadata.as_ref(),
|
||||
&item.request_id,
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
);
|
||||
}
|
||||
if item.response_body_ref.is_none() {
|
||||
item.response_body_ref = usage_body_ref_from_metadata(
|
||||
item.request_metadata.as_ref(),
|
||||
&item.request_id,
|
||||
UsageBodyField::ResponseBody,
|
||||
);
|
||||
}
|
||||
if item.client_response_body_ref.is_none() {
|
||||
item.client_response_body_ref = usage_body_ref_from_metadata(
|
||||
item.request_metadata.as_ref(),
|
||||
&item.request_id,
|
||||
UsageBodyField::ClientResponseBody,
|
||||
);
|
||||
}
|
||||
item.request_body_ref = item
|
||||
.request_body_ref
|
||||
.as_deref()
|
||||
.and_then(|body_ref| {
|
||||
canonical_usage_body_ref_for(body_ref, &item.request_id, UsageBodyField::RequestBody)
|
||||
})
|
||||
.or_else(|| {
|
||||
usage_body_ref_from_metadata(
|
||||
item.request_metadata.as_ref(),
|
||||
&item.request_id,
|
||||
UsageBodyField::RequestBody,
|
||||
)
|
||||
});
|
||||
item.provider_request_body_ref = item
|
||||
.provider_request_body_ref
|
||||
.as_deref()
|
||||
.and_then(|body_ref| {
|
||||
canonical_usage_body_ref_for(
|
||||
body_ref,
|
||||
&item.request_id,
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
)
|
||||
})
|
||||
.or_else(|| {
|
||||
usage_body_ref_from_metadata(
|
||||
item.request_metadata.as_ref(),
|
||||
&item.request_id,
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
)
|
||||
});
|
||||
item.response_body_ref = item
|
||||
.response_body_ref
|
||||
.as_deref()
|
||||
.and_then(|body_ref| {
|
||||
canonical_usage_body_ref_for(body_ref, &item.request_id, UsageBodyField::ResponseBody)
|
||||
})
|
||||
.or_else(|| {
|
||||
usage_body_ref_from_metadata(
|
||||
item.request_metadata.as_ref(),
|
||||
&item.request_id,
|
||||
UsageBodyField::ResponseBody,
|
||||
)
|
||||
});
|
||||
item.client_response_body_ref = item
|
||||
.client_response_body_ref
|
||||
.as_deref()
|
||||
.and_then(|body_ref| {
|
||||
canonical_usage_body_ref_for(
|
||||
body_ref,
|
||||
&item.request_id,
|
||||
UsageBodyField::ClientResponseBody,
|
||||
)
|
||||
})
|
||||
.or_else(|| {
|
||||
usage_body_ref_from_metadata(
|
||||
item.request_metadata.as_ref(),
|
||||
&item.request_id,
|
||||
UsageBodyField::ClientResponseBody,
|
||||
)
|
||||
});
|
||||
}
|
||||
|
||||
fn hydrate_client_family(item: &mut StoredRequestUsageAudit) {
|
||||
@@ -2723,34 +2752,6 @@ fn hydrate_client_family(item: &mut StoredRequestUsageAudit) {
|
||||
}
|
||||
}
|
||||
|
||||
fn persisted_usage_body_ref(
|
||||
incoming_ref: Option<&str>,
|
||||
incoming_body: Option<&Value>,
|
||||
incoming_state: Option<UsageBodyCaptureState>,
|
||||
_metadata: Option<&Value>,
|
||||
existing: Option<&StoredRequestUsageAudit>,
|
||||
field: UsageBodyField,
|
||||
) -> Option<String> {
|
||||
if incoming_state == Some(UsageBodyCaptureState::None) {
|
||||
return None;
|
||||
}
|
||||
if incoming_body.is_some() {
|
||||
return None;
|
||||
}
|
||||
incoming_ref
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(ToOwned::to_owned)
|
||||
.or_else(|| {
|
||||
existing.and_then(|existing| match field {
|
||||
UsageBodyField::RequestBody => existing.request_body_ref.clone(),
|
||||
UsageBodyField::ProviderRequestBody => existing.provider_request_body_ref.clone(),
|
||||
UsageBodyField::ResponseBody => existing.response_body_ref.clone(),
|
||||
UsageBodyField::ClientResponseBody => existing.client_response_body_ref.clone(),
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn request_body_capture_replaces_derived_facts(
|
||||
request_body: Option<&Value>,
|
||||
request_body_state: Option<UsageBodyCaptureState>,
|
||||
@@ -2852,9 +2853,68 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
|
||||
usage: UpsertUsageRecord,
|
||||
) -> Result<StoredRequestUsageAudit, DataLayerError> {
|
||||
usage.validate()?;
|
||||
let usage = strip_deprecated_usage_display_fields(usage);
|
||||
let capture_usage = usage.clone();
|
||||
let usage = sanitize_usage_for_persistence(usage);
|
||||
let mut by_request_id = self.by_request_id.write().expect("usage repository lock");
|
||||
let existing = by_request_id.get(&usage.request_id).cloned();
|
||||
if let Some(existing) = existing.as_ref() {
|
||||
if !usage_lifecycle_update_allowed(
|
||||
&existing.status,
|
||||
&existing.billing_status,
|
||||
existing.updated_at_unix_secs,
|
||||
existing.finalized_at_unix_secs,
|
||||
&usage.status,
|
||||
&usage.billing_status,
|
||||
usage.updated_at_unix_secs,
|
||||
usage.finalized_at_unix_secs,
|
||||
) {
|
||||
return Ok(existing.clone());
|
||||
}
|
||||
let can_recover = usage_can_recover_terminal_failure(
|
||||
existing.status.as_str(),
|
||||
existing.billing_status.as_str(),
|
||||
usage.status.as_str(),
|
||||
usage.billing_status.as_str(),
|
||||
);
|
||||
let completed_terminal_failure_recovery = existing.billing_status == "void"
|
||||
&& matches!(existing.status.as_str(), "failed" | "cancelled")
|
||||
&& usage.status == "completed";
|
||||
if completed_terminal_failure_recovery && !can_recover {
|
||||
return Ok(existing.clone());
|
||||
}
|
||||
}
|
||||
let capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage);
|
||||
if let Some(existing) = by_request_id.get_mut(&usage.request_id) {
|
||||
existing.request_headers = None;
|
||||
existing.request_body = None;
|
||||
existing.request_body_ref = None;
|
||||
existing.request_body_state = None;
|
||||
existing.provider_request_headers = None;
|
||||
existing.provider_request_body = None;
|
||||
existing.provider_request_body_ref = None;
|
||||
existing.provider_request_body_state = None;
|
||||
existing.response_headers = None;
|
||||
existing.response_body = None;
|
||||
existing.response_body_ref = None;
|
||||
existing.response_body_state = None;
|
||||
existing.client_response_headers = None;
|
||||
existing.client_response_body = None;
|
||||
existing.client_response_body_ref = None;
|
||||
existing.client_response_body_state = None;
|
||||
existing.request_metadata =
|
||||
sanitize_usage_request_metadata(existing.request_metadata.take());
|
||||
}
|
||||
{
|
||||
let mut detached_bodies = self.detached_bodies.write().expect("usage repository lock");
|
||||
for field in [
|
||||
UsageBodyField::RequestBody,
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
UsageBodyField::ResponseBody,
|
||||
UsageBodyField::ClientResponseBody,
|
||||
] {
|
||||
detached_bodies.remove(&usage_body_ref(&usage.request_id, field));
|
||||
}
|
||||
}
|
||||
|
||||
let created_at_unix_ms = by_request_id
|
||||
.get(&usage.request_id)
|
||||
@@ -2871,46 +2931,20 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
|
||||
)
|
||||
})
|
||||
.unwrap_or_default();
|
||||
if existing.as_ref().is_some_and(|existing| {
|
||||
let can_recover = usage_can_recover_terminal_failure(
|
||||
existing.status.as_str(),
|
||||
existing.billing_status.as_str(),
|
||||
usage.status.as_str(),
|
||||
usage.billing_status.as_str(),
|
||||
);
|
||||
let finalized_lifecycle_regression = usage_status_is_finalized(&existing.status)
|
||||
&& usage_status_is_lifecycle(&usage.status);
|
||||
let completed_terminal_failure_recovery = existing.billing_status == "void"
|
||||
&& matches!(existing.status.as_str(), "failed" | "cancelled")
|
||||
&& usage.status == "completed";
|
||||
(finalized_lifecycle_regression || completed_terminal_failure_recovery) && !can_recover
|
||||
}) {
|
||||
return Ok(existing.expect("existing usage should be present").clone());
|
||||
}
|
||||
if existing.as_ref().is_some_and(|existing| {
|
||||
existing.billing_status == "pending"
|
||||
&& existing.status == "streaming"
|
||||
&& usage.status == "pending"
|
||||
}) {
|
||||
return Ok(existing.expect("existing usage should be present").clone());
|
||||
}
|
||||
|
||||
let replace_client_request_body_facts = request_body_capture_replaces_derived_facts(
|
||||
usage.request_body.as_ref(),
|
||||
usage.request_body_state,
|
||||
capture_usage.request_body.as_ref(),
|
||||
capture_usage.request_body_state,
|
||||
);
|
||||
let replace_provider_request_body_facts = request_body_capture_replaces_derived_facts(
|
||||
usage.provider_request_body.as_ref(),
|
||||
usage.provider_request_body_state,
|
||||
capture_usage.provider_request_body.as_ref(),
|
||||
capture_usage.provider_request_body_state,
|
||||
);
|
||||
let clear_request_body = usage.request_body_state == Some(UsageBodyCaptureState::None);
|
||||
let clear_request_body =
|
||||
capture_usage.request_body_state == Some(UsageBodyCaptureState::None);
|
||||
let clear_provider_request_body =
|
||||
usage.provider_request_body_state == Some(UsageBodyCaptureState::None);
|
||||
let clear_response_body = usage.response_body_state == Some(UsageBodyCaptureState::None);
|
||||
let clear_client_response_body =
|
||||
usage.client_response_body_state == Some(UsageBodyCaptureState::None);
|
||||
capture_usage.provider_request_body_state == Some(UsageBodyCaptureState::None);
|
||||
let replace_routing_snapshot = usage_status_is_finalized(&usage.status);
|
||||
let mut incoming_request_metadata = usage.request_metadata.clone();
|
||||
let mut incoming_request_metadata = capture_usage.request_metadata.clone();
|
||||
if incoming_request_metadata.is_some()
|
||||
&& (clear_request_body || clear_provider_request_body)
|
||||
{
|
||||
@@ -2942,61 +2976,11 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
|
||||
.and_then(|existing| existing.request_metadata.clone())
|
||||
}
|
||||
});
|
||||
let request_body_ref = persisted_usage_body_ref(
|
||||
usage.request_body_ref.as_deref(),
|
||||
usage.request_body.as_ref(),
|
||||
usage.request_body_state,
|
||||
request_metadata.as_ref(),
|
||||
existing.as_ref(),
|
||||
UsageBodyField::RequestBody,
|
||||
);
|
||||
let provider_request_body_ref = persisted_usage_body_ref(
|
||||
usage.provider_request_body_ref.as_deref(),
|
||||
usage.provider_request_body.as_ref(),
|
||||
usage.provider_request_body_state,
|
||||
request_metadata.as_ref(),
|
||||
existing.as_ref(),
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
);
|
||||
let response_body_ref = persisted_usage_body_ref(
|
||||
usage.response_body_ref.as_deref(),
|
||||
usage.response_body.as_ref(),
|
||||
usage.response_body_state,
|
||||
request_metadata.as_ref(),
|
||||
existing.as_ref(),
|
||||
UsageBodyField::ResponseBody,
|
||||
);
|
||||
let client_response_body_ref = persisted_usage_body_ref(
|
||||
usage.client_response_body_ref.as_deref(),
|
||||
usage.client_response_body.as_ref(),
|
||||
usage.client_response_body_state,
|
||||
request_metadata.as_ref(),
|
||||
existing.as_ref(),
|
||||
UsageBodyField::ClientResponseBody,
|
||||
);
|
||||
if clear_request_body
|
||||
|| clear_provider_request_body
|
||||
|| clear_response_body
|
||||
|| clear_client_response_body
|
||||
{
|
||||
let mut detached_bodies = self.detached_bodies.write().expect("usage repository lock");
|
||||
for (clear, field) in [
|
||||
(clear_request_body, UsageBodyField::RequestBody),
|
||||
(
|
||||
clear_provider_request_body,
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
),
|
||||
(clear_response_body, UsageBodyField::ResponseBody),
|
||||
(
|
||||
clear_client_response_body,
|
||||
UsageBodyField::ClientResponseBody,
|
||||
),
|
||||
] {
|
||||
if clear {
|
||||
detached_bodies.remove(&usage_body_ref(&usage.request_id, field));
|
||||
}
|
||||
}
|
||||
}
|
||||
let request_metadata = sanitize_memory_request_metadata(request_metadata);
|
||||
let request_body_ref = None;
|
||||
let provider_request_body_ref = None;
|
||||
let response_body_ref = None;
|
||||
let client_response_body_ref = None;
|
||||
let stored = StoredRequestUsageAudit {
|
||||
id: existing
|
||||
.as_ref()
|
||||
@@ -3105,159 +3089,97 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
|
||||
),
|
||||
status: usage.status,
|
||||
billing_status: usage.billing_status,
|
||||
request_headers: usage.request_headers.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.request_headers.clone())
|
||||
}),
|
||||
request_body: if clear_request_body {
|
||||
None
|
||||
} else {
|
||||
usage.request_body.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.request_body.clone())
|
||||
})
|
||||
},
|
||||
request_headers: None,
|
||||
request_body: None,
|
||||
request_body_ref,
|
||||
request_body_state: usage.request_body_state.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.request_body_state)
|
||||
}),
|
||||
provider_request_headers: usage.provider_request_headers.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.provider_request_headers.clone())
|
||||
}),
|
||||
provider_request_body: if clear_provider_request_body {
|
||||
None
|
||||
} else {
|
||||
usage.provider_request_body.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.provider_request_body.clone())
|
||||
})
|
||||
},
|
||||
request_body_state: capture_usage.request_body_state,
|
||||
provider_request_headers: None,
|
||||
provider_request_body: None,
|
||||
provider_request_body_ref,
|
||||
provider_request_body_state: usage.provider_request_body_state.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.provider_request_body_state)
|
||||
}),
|
||||
response_headers: usage.response_headers.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.response_headers.clone())
|
||||
}),
|
||||
response_body: if clear_response_body {
|
||||
None
|
||||
} else {
|
||||
usage.response_body.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.response_body.clone())
|
||||
})
|
||||
},
|
||||
provider_request_body_state: capture_usage.provider_request_body_state,
|
||||
response_headers: None,
|
||||
response_body: None,
|
||||
response_body_ref,
|
||||
response_body_state: usage.response_body_state.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.response_body_state)
|
||||
}),
|
||||
client_response_headers: usage.client_response_headers.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.client_response_headers.clone())
|
||||
}),
|
||||
client_response_body: if clear_client_response_body {
|
||||
None
|
||||
} else {
|
||||
usage.client_response_body.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.client_response_body.clone())
|
||||
})
|
||||
},
|
||||
response_body_state: capture_usage.response_body_state,
|
||||
client_response_headers: None,
|
||||
client_response_body: None,
|
||||
client_response_body_ref,
|
||||
client_response_body_state: usage.client_response_body_state.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.client_response_body_state)
|
||||
}),
|
||||
client_response_body_state: capture_usage.client_response_body_state,
|
||||
candidate_id: if replace_routing_snapshot {
|
||||
usage.candidate_id
|
||||
capture_usage.candidate_id
|
||||
} else {
|
||||
usage.candidate_id.or_else(|| {
|
||||
capture_usage.candidate_id.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.routing_candidate_id().map(ToOwned::to_owned))
|
||||
})
|
||||
},
|
||||
candidate_index: if replace_routing_snapshot {
|
||||
usage.candidate_index
|
||||
capture_usage.candidate_index
|
||||
} else {
|
||||
usage.candidate_index.or_else(|| {
|
||||
capture_usage.candidate_index.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.routing_candidate_index())
|
||||
})
|
||||
},
|
||||
key_name: if replace_routing_snapshot {
|
||||
usage.key_name
|
||||
capture_usage.key_name
|
||||
} else {
|
||||
usage.key_name.or_else(|| {
|
||||
capture_usage.key_name.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.routing_key_name().map(ToOwned::to_owned))
|
||||
})
|
||||
},
|
||||
planner_kind: if replace_routing_snapshot {
|
||||
usage.planner_kind
|
||||
capture_usage.planner_kind
|
||||
} else {
|
||||
usage.planner_kind.or_else(|| {
|
||||
capture_usage.planner_kind.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.routing_planner_kind().map(ToOwned::to_owned))
|
||||
})
|
||||
},
|
||||
route_family: if replace_routing_snapshot {
|
||||
usage.route_family
|
||||
capture_usage.route_family
|
||||
} else {
|
||||
usage.route_family.or_else(|| {
|
||||
capture_usage.route_family.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.routing_route_family().map(ToOwned::to_owned))
|
||||
})
|
||||
},
|
||||
route_kind: if replace_routing_snapshot {
|
||||
usage.route_kind
|
||||
capture_usage.route_kind
|
||||
} else {
|
||||
usage.route_kind.or_else(|| {
|
||||
capture_usage.route_kind.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|existing| existing.routing_route_kind().map(ToOwned::to_owned))
|
||||
})
|
||||
},
|
||||
execution_path: if replace_routing_snapshot {
|
||||
usage.execution_path
|
||||
capture_usage.execution_path
|
||||
} else {
|
||||
usage.execution_path.or_else(|| {
|
||||
capture_usage.execution_path.or_else(|| {
|
||||
existing.as_ref().and_then(|existing| {
|
||||
existing.routing_execution_path().map(ToOwned::to_owned)
|
||||
})
|
||||
})
|
||||
},
|
||||
local_execution_runtime_miss_reason: if replace_routing_snapshot {
|
||||
usage.local_execution_runtime_miss_reason
|
||||
capture_usage.local_execution_runtime_miss_reason
|
||||
} else {
|
||||
usage.local_execution_runtime_miss_reason.or_else(|| {
|
||||
existing.as_ref().and_then(|existing| {
|
||||
existing
|
||||
.routing_local_execution_runtime_miss_reason()
|
||||
.map(ToOwned::to_owned)
|
||||
capture_usage
|
||||
.local_execution_runtime_miss_reason
|
||||
.or_else(|| {
|
||||
existing.as_ref().and_then(|existing| {
|
||||
existing
|
||||
.routing_local_execution_runtime_miss_reason()
|
||||
.map(ToOwned::to_owned)
|
||||
})
|
||||
})
|
||||
})
|
||||
},
|
||||
client_family: usage_request_metadata_client_family(request_metadata.as_ref())
|
||||
.map(ToOwned::to_owned)
|
||||
|
||||
@@ -788,6 +788,64 @@ async fn upsert_allows_completed_recovery_after_void_failure() {
|
||||
assert_eq!(stored.total_tokens, 10);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn stale_terminal_event_cannot_replace_usage_routing_or_counter_contribution() {
|
||||
let auth_api_keys = sample_auth_api_key_repository(&["api-key-1"]);
|
||||
let repository = InMemoryUsageReadRepository::default()
|
||||
.with_auth_api_key_repository(Arc::clone(&auth_api_keys));
|
||||
|
||||
let mut newer = sample_upsert_usage_record("req-stale-terminal");
|
||||
newer.api_key_id = Some("api-key-1".to_string());
|
||||
newer.status = "completed".to_string();
|
||||
newer.status_code = Some(200);
|
||||
newer.total_tokens = Some(5);
|
||||
newer.total_cost_usd = Some(0.5);
|
||||
newer.candidate_id = Some("candidate-new".to_string());
|
||||
newer.route_kind = Some("route-new".to_string());
|
||||
newer.updated_at_unix_secs = 200;
|
||||
newer.finalized_at_unix_secs = Some(200);
|
||||
repository
|
||||
.upsert(newer)
|
||||
.await
|
||||
.expect("newer terminal usage should upsert");
|
||||
|
||||
let mut stale = sample_upsert_usage_record("req-stale-terminal");
|
||||
stale.api_key_id = Some("api-key-1".to_string());
|
||||
stale.status = "failed".to_string();
|
||||
stale.billing_status = "void".to_string();
|
||||
stale.status_code = Some(503);
|
||||
stale.total_tokens = Some(999);
|
||||
stale.total_cost_usd = Some(99.0);
|
||||
stale.candidate_id = Some("candidate-stale".to_string());
|
||||
stale.route_kind = Some("route-stale".to_string());
|
||||
stale.updated_at_unix_secs = 199;
|
||||
stale.finalized_at_unix_secs = Some(199);
|
||||
let stored = repository
|
||||
.upsert(stale)
|
||||
.await
|
||||
.expect("stale terminal usage should be ignored");
|
||||
|
||||
assert_eq!(stored.status, "completed");
|
||||
assert_eq!(stored.billing_status, "pending");
|
||||
assert_eq!(stored.status_code, Some(200));
|
||||
assert_eq!(stored.total_tokens, 5);
|
||||
assert_eq!(stored.total_cost_usd, 0.5);
|
||||
assert_eq!(stored.routing_candidate_id(), Some("candidate-new"));
|
||||
assert_eq!(stored.routing_route_kind(), Some("route-new"));
|
||||
assert_eq!(stored.updated_at_unix_secs, 200);
|
||||
|
||||
let key = auth_api_keys
|
||||
.list_export_api_keys_by_ids(&["api-key-1".to_string()])
|
||||
.await
|
||||
.expect("api key stats should load")
|
||||
.into_iter()
|
||||
.next()
|
||||
.expect("api key should exist");
|
||||
assert_eq!(key.total_requests, 1);
|
||||
assert_eq!(key.total_tokens, 5);
|
||||
assert_eq!(key.total_cost_usd, 0.5);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upsert_rejects_non_authoritative_void_failure_recovery() {
|
||||
let repository = InMemoryUsageReadRepository::default();
|
||||
@@ -1189,6 +1247,23 @@ async fn detached_body_seed_moves_large_payloads_behind_usage_refs() {
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn seed_discards_cross_request_and_cross_field_body_refs() {
|
||||
let mut usage = sample_usage("req-ref-target", 100);
|
||||
usage.request_body_ref = Some("usage://request/req-ref-owner/request_body".to_string());
|
||||
usage.response_body_ref = Some("usage://request/req-ref-target/request_body".to_string());
|
||||
|
||||
let repository = InMemoryUsageReadRepository::seed(vec![usage]);
|
||||
let stored = repository
|
||||
.find_by_request_id("req-ref-target")
|
||||
.await
|
||||
.expect("find should succeed")
|
||||
.expect("usage should exist");
|
||||
|
||||
assert!(stored.request_body_ref.is_none());
|
||||
assert!(stored.response_body_ref.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upsert_writes_usage_record() {
|
||||
let repository = InMemoryUsageReadRepository::default();
|
||||
@@ -1520,12 +1595,7 @@ async fn upsert_does_not_backfill_typed_body_refs_from_request_metadata() {
|
||||
.expect("upsert should succeed");
|
||||
|
||||
assert_eq!(stored.request_body_ref, None);
|
||||
assert_eq!(
|
||||
stored.request_metadata,
|
||||
Some(json!({
|
||||
"request_body_ref": "usage://request/req-upsert-body-ref-metadata/request_body"
|
||||
}))
|
||||
);
|
||||
assert_eq!(stored.request_metadata, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
@@ -6,33 +6,35 @@ mod mysql;
|
||||
pub(crate) use aether_data_contracts::repository::usage::{
|
||||
api_key_usage_contribution, incoming_usage_can_recover_terminal_failure,
|
||||
model_usage_contribution, provider_api_key_usage_contribution, provider_api_key_usage_is_error,
|
||||
provider_api_key_usage_is_success, strip_deprecated_usage_display_fields,
|
||||
usage_can_recover_terminal_failure, usage_request_metadata_client_family, ApiKeyLastUsedDelta,
|
||||
ApiKeyUsageContribution, ApiKeyUsageDelta, ManagementTokenCounterDelta, ModelUsageContribution,
|
||||
ModelUsageDelta, PendingUsageCleanupSummary, ProviderApiKeyUsageContribution,
|
||||
ProviderApiKeyUsageDelta, ProviderApiKeyWindowUsageRequest, ProxyNodeCounterDelta,
|
||||
StoredProviderApiKeyUsageSummary, StoredProviderApiKeyWindowUsageSummary,
|
||||
StoredProviderUsageSummary, StoredProviderUsageWindow, StoredRequestUsageAudit,
|
||||
StoredUsageAuditAggregation, StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow,
|
||||
StoredUsageCacheAffinityHitSummary, StoredUsageCacheAffinityIntervalRow,
|
||||
StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary, StoredUsageDailySummary,
|
||||
StoredUsageDashboardDailyBreakdownRow, StoredUsageDashboardProviderCount,
|
||||
StoredUsageDashboardStatsSummary, StoredUsageDashboardSummary, StoredUsageErrorDistributionRow,
|
||||
StoredUsageLeaderboardSummary, StoredUsagePerformancePercentilesRow,
|
||||
StoredUsageProviderPerformance, StoredUsageProviderPerformanceProviderRow,
|
||||
StoredUsageProviderPerformanceSummary, StoredUsageProviderPerformanceTimelineRow,
|
||||
StoredUsageSettledCostSummary, StoredUsageTimeSeriesBucket, StoredUsageUserTotals,
|
||||
UpsertUsageRecord, UsageAuditAggregationGroupBy, UsageAuditAggregationQuery,
|
||||
UsageAuditKeywordSearchQuery, UsageAuditListQuery, UsageAuditSummaryQuery,
|
||||
UsageBreakdownGroupBy, UsageBreakdownSummaryQuery, UsageCacheAffinityHitSummaryQuery,
|
||||
UsageCacheAffinityIntervalGroupBy, UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery,
|
||||
UsageCleanupPreviewCounts, UsageCleanupSummary, UsageCleanupWindow,
|
||||
UsageCostSavingsSummaryQuery, UsageCounterFlushSummary, UsageCounterHealthSnapshot,
|
||||
UsageCounterPendingHealthSnapshot, UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery,
|
||||
UsageDashboardProviderCountsQuery, UsageDashboardSummaryQuery, UsageErrorDistributionQuery,
|
||||
UsageLeaderboardGroupBy, UsageLeaderboardQuery, UsageMonitoringErrorCountQuery,
|
||||
UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery,
|
||||
UsageReadRepository, UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
|
||||
provider_api_key_usage_is_success, sanitize_usage_capture_controls_for_persistence,
|
||||
sanitize_usage_for_persistence, strip_deprecated_usage_display_fields,
|
||||
usage_can_recover_terminal_failure, usage_lifecycle_update_allowed,
|
||||
usage_request_metadata_client_family, ApiKeyLastUsedDelta, ApiKeyUsageContribution,
|
||||
ApiKeyUsageDelta, ManagementTokenCounterDelta, ModelUsageContribution, ModelUsageDelta,
|
||||
PendingUsageCleanupSummary, ProviderApiKeyUsageContribution, ProviderApiKeyUsageDelta,
|
||||
ProviderApiKeyWindowUsageRequest, ProxyNodeCounterDelta, StoredProviderApiKeyUsageSummary,
|
||||
StoredProviderApiKeyWindowUsageSummary, StoredProviderUsageSummary, StoredProviderUsageWindow,
|
||||
StoredRequestUsageAudit, StoredUsageAuditAggregation, StoredUsageAuditSummary,
|
||||
StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
|
||||
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
|
||||
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow,
|
||||
StoredUsageDashboardProviderCount, StoredUsageDashboardStatsSummary,
|
||||
StoredUsageDashboardSummary, StoredUsageErrorDistributionRow, StoredUsageLeaderboardSummary,
|
||||
StoredUsagePerformancePercentilesRow, StoredUsageProviderPerformance,
|
||||
StoredUsageProviderPerformanceProviderRow, StoredUsageProviderPerformanceSummary,
|
||||
StoredUsageProviderPerformanceTimelineRow, StoredUsageSettledCostSummary,
|
||||
StoredUsageTimeSeriesBucket, StoredUsageUserTotals, UpsertUsageRecord,
|
||||
UsageAuditAggregationGroupBy, UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery,
|
||||
UsageAuditListQuery, UsageAuditSummaryQuery, UsageBreakdownGroupBy, UsageBreakdownSummaryQuery,
|
||||
UsageCacheAffinityHitSummaryQuery, UsageCacheAffinityIntervalGroupBy,
|
||||
UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery, UsageCleanupPreviewCounts,
|
||||
UsageCleanupSummary, UsageCleanupWindow, UsageCostSavingsSummaryQuery,
|
||||
UsageCounterFlushSummary, UsageCounterHealthSnapshot, UsageCounterPendingHealthSnapshot,
|
||||
UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery, UsageDashboardProviderCountsQuery,
|
||||
UsageDashboardSummaryQuery, UsageErrorDistributionQuery, UsageLeaderboardGroupBy,
|
||||
UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, UsageMonitoringErrorListQuery,
|
||||
UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageReadRepository,
|
||||
UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
|
||||
UsageTimeSeriesQuery, UsageWriteRepository,
|
||||
};
|
||||
#[cfg(feature = "postgres")]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,11 +1,15 @@
|
||||
mod memory;
|
||||
|
||||
pub use aether_data_contracts::repository::users::{
|
||||
normalize_user_group_name, LdapAuthUserProvisioningOutcome, StoredUserAuthRecord,
|
||||
is_last_active_admin_delete_denied, is_last_active_admin_update_denied, is_valid_bcrypt_hash,
|
||||
last_oauth_unbind_denial, normalize_user_group_name, BindUserOAuthLinkOutcome,
|
||||
BindUserOAuthLinkSessionExpectation, DeleteUserOAuthLinkOutcome,
|
||||
LdapAuthUserProvisioningOutcome, ResolveOAuthLinkedUserOutcome, StoredUserAuthRecord,
|
||||
StoredUserExportRow, StoredUserGroup, StoredUserGroupMember, StoredUserGroupMembership,
|
||||
StoredUserOAuthLinkSummary, StoredUserPreferenceRecord, StoredUserSessionRecord,
|
||||
StoredUserSummary, UpsertUserGroupRecord, UserExportListQuery, UserExportSortBy,
|
||||
UserExportSortOrder, UserExportSummary, UserReadRepository,
|
||||
UserExportSortOrder, UserExportSummary, UserReadRepository, LAST_ACTIVE_ADMIN_DELETE_DENIED,
|
||||
LAST_ACTIVE_ADMIN_UPDATE_DENIED,
|
||||
};
|
||||
#[cfg(feature = "mysql")]
|
||||
pub use aether_data_mysql::MysqlUserReadRepository;
|
||||
|
||||
@@ -14,6 +14,7 @@ use crate::DataLayerError;
|
||||
struct MemoryVideoTaskIndex {
|
||||
by_id: BTreeMap<String, StoredVideoTask>,
|
||||
short_to_id: BTreeMap<String, String>,
|
||||
request_to_id: BTreeMap<String, String>,
|
||||
user_external_to_id: BTreeMap<(String, String), String>,
|
||||
}
|
||||
|
||||
@@ -28,6 +29,7 @@ impl InMemoryVideoTaskRepository {
|
||||
if let Some(short_id) = previous.short_id {
|
||||
index.short_to_id.remove(&short_id);
|
||||
}
|
||||
index.request_to_id.remove(&previous.request_id);
|
||||
if let (Some(user_id), Some(external_task_id)) =
|
||||
(previous.user_id, previous.external_task_id)
|
||||
{
|
||||
@@ -40,6 +42,9 @@ impl InMemoryVideoTaskRepository {
|
||||
if let Some(short_id) = &task.short_id {
|
||||
index.short_to_id.insert(short_id.clone(), task.id.clone());
|
||||
}
|
||||
index
|
||||
.request_to_id
|
||||
.insert(task.request_id.clone(), task.id.clone());
|
||||
if let (Some(user_id), Some(external_task_id)) = (&task.user_id, &task.external_task_id) {
|
||||
index
|
||||
.user_external_to_id
|
||||
@@ -49,6 +54,35 @@ impl InMemoryVideoTaskRepository {
|
||||
task
|
||||
}
|
||||
|
||||
fn ensure_unique_keys_available(
|
||||
index: &MemoryVideoTaskIndex,
|
||||
task: &UpsertVideoTask,
|
||||
) -> Result<(), DataLayerError> {
|
||||
if let Some(short_id) = task.short_id.as_deref() {
|
||||
if index
|
||||
.short_to_id
|
||||
.get(short_id)
|
||||
.is_some_and(|existing_id| existing_id != &task.id)
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"video task {} conflicts with existing short_id {short_id}",
|
||||
task.id
|
||||
)));
|
||||
}
|
||||
}
|
||||
if index
|
||||
.request_to_id
|
||||
.get(&task.request_id)
|
||||
.is_some_and(|existing_id| existing_id != &task.id)
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"video task {} conflicts with existing request_id {}",
|
||||
task.id, task.request_id
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn matches_filter(task: &StoredVideoTask, filter: &VideoTaskQueryFilter) -> bool {
|
||||
if let Some(user_id) = filter.user_id.as_deref() {
|
||||
if task.user_id.as_deref() != Some(user_id) {
|
||||
@@ -103,6 +137,36 @@ impl VideoTaskReadRepository for InMemoryVideoTaskRepository {
|
||||
})
|
||||
}
|
||||
|
||||
async fn find_for_user(
|
||||
&self,
|
||||
key: VideoTaskLookupKey<'_>,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredVideoTask>, DataLayerError> {
|
||||
let index = self.index.read().expect("video task repository lock");
|
||||
let task = match key {
|
||||
VideoTaskLookupKey::Id(id) => index.by_id.get(id),
|
||||
VideoTaskLookupKey::ShortId(short_id) => index
|
||||
.short_to_id
|
||||
.get(short_id)
|
||||
.and_then(|id| index.by_id.get(id)),
|
||||
VideoTaskLookupKey::UserExternal {
|
||||
user_id: lookup_user_id,
|
||||
external_task_id,
|
||||
} => {
|
||||
if lookup_user_id != user_id {
|
||||
return Ok(None);
|
||||
}
|
||||
index
|
||||
.user_external_to_id
|
||||
.get(&(lookup_user_id.to_string(), external_task_id.to_string()))
|
||||
.and_then(|id| index.by_id.get(id))
|
||||
}
|
||||
};
|
||||
Ok(task
|
||||
.filter(|task| task.user_id.as_deref() == Some(user_id))
|
||||
.cloned())
|
||||
}
|
||||
|
||||
async fn list_active(&self, limit: usize) -> Result<Vec<StoredVideoTask>, DataLayerError> {
|
||||
if limit == 0 {
|
||||
return Ok(Vec::new());
|
||||
@@ -306,22 +370,34 @@ impl VideoTaskReadRepository for InMemoryVideoTaskRepository {
|
||||
|
||||
#[async_trait]
|
||||
impl VideoTaskWriteRepository for InMemoryVideoTaskRepository {
|
||||
async fn upsert(&self, task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
|
||||
async fn upsert(&self, mut task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
|
||||
let mut index = self.index.write().expect("video task repository lock");
|
||||
Self::ensure_unique_keys_available(&index, &task)?;
|
||||
if let Some(existing) = index.by_id.get(&task.id) {
|
||||
existing.ensure_immutable_identity_matches(&task)?;
|
||||
task.created_at_unix_ms = existing.created_at_unix_ms;
|
||||
}
|
||||
Ok(Self::store_locked(&mut index, task.into_stored()))
|
||||
}
|
||||
|
||||
async fn update_if_active(
|
||||
&self,
|
||||
task: UpsertVideoTask,
|
||||
mut task: UpsertVideoTask,
|
||||
) -> Result<Option<StoredVideoTask>, DataLayerError> {
|
||||
let mut index = self.index.write().expect("video task repository lock");
|
||||
if Self::ensure_unique_keys_available(&index, &task).is_err() {
|
||||
return Ok(None);
|
||||
}
|
||||
let Some(existing) = index.by_id.get(&task.id) else {
|
||||
return Ok(None);
|
||||
};
|
||||
if !existing.status.is_active() {
|
||||
return Ok(None);
|
||||
}
|
||||
if existing.ensure_immutable_identity_matches(&task).is_err() {
|
||||
return Ok(None);
|
||||
}
|
||||
task.created_at_unix_ms = existing.created_at_unix_ms;
|
||||
Ok(Some(Self::store_locked(&mut index, task.into_stored())))
|
||||
}
|
||||
|
||||
@@ -455,6 +531,34 @@ mod tests {
|
||||
.is_some());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn owner_scoped_lookup_rejects_foreign_user_for_every_identifier() {
|
||||
let repo = InMemoryVideoTaskRepository::default();
|
||||
repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
|
||||
.await
|
||||
.expect("upsert should succeed");
|
||||
|
||||
for key in [
|
||||
VideoTaskLookupKey::Id("task-1"),
|
||||
VideoTaskLookupKey::ShortId("short-task-1"),
|
||||
VideoTaskLookupKey::UserExternal {
|
||||
user_id: "user-1",
|
||||
external_task_id: "ext-task-1",
|
||||
},
|
||||
] {
|
||||
assert!(repo
|
||||
.find_for_user(key, "user-1")
|
||||
.await
|
||||
.expect("owner lookup should succeed")
|
||||
.is_some());
|
||||
assert!(repo
|
||||
.find_for_user(key, "user-2")
|
||||
.await
|
||||
.expect("foreign lookup should succeed")
|
||||
.is_none());
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn list_active_only_returns_active_tasks_in_descending_update_order() {
|
||||
let repo = InMemoryVideoTaskRepository::default();
|
||||
@@ -478,59 +582,69 @@ mod tests {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upsert_replaces_secondary_indexes() {
|
||||
async fn upsert_rejects_immutable_identity_replacement() {
|
||||
let repo = InMemoryVideoTaskRepository::default();
|
||||
repo.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
|
||||
.await
|
||||
.expect("upsert should succeed");
|
||||
|
||||
repo.upsert(UpsertVideoTask {
|
||||
id: "task-1".to_string(),
|
||||
short_id: Some("short-task-1b".to_string()),
|
||||
request_id: "request-task-1b".to_string(),
|
||||
user_id: Some("user-2".to_string()),
|
||||
api_key_id: Some("api-key-2".to_string()),
|
||||
username: Some("user-2".to_string()),
|
||||
api_key_name: Some("secondary".to_string()),
|
||||
external_task_id: Some("ext-task-1b".to_string()),
|
||||
provider_id: Some("provider-2".to_string()),
|
||||
endpoint_id: Some("endpoint-2".to_string()),
|
||||
key_id: Some("provider-key-2".to_string()),
|
||||
client_api_format: Some("gemini:video".to_string()),
|
||||
provider_api_format: Some("gemini:video".to_string()),
|
||||
format_converted: false,
|
||||
model: Some("veo-3".to_string()),
|
||||
prompt: Some("remix".to_string()),
|
||||
original_request_body: Some(serde_json::json!({"prompt": "remix"})),
|
||||
duration_seconds: Some(8),
|
||||
resolution: Some("1080p".to_string()),
|
||||
aspect_ratio: Some("16:9".to_string()),
|
||||
size: Some("720p".to_string()),
|
||||
status: VideoTaskStatus::Processing,
|
||||
progress_percent: 50,
|
||||
progress_message: Some("processing".to_string()),
|
||||
retry_count: 1,
|
||||
poll_interval_seconds: 10,
|
||||
next_poll_at_unix_secs: Some(200),
|
||||
poll_count: 2,
|
||||
max_poll_count: 360,
|
||||
created_at_unix_ms: 150,
|
||||
submitted_at_unix_secs: Some(150),
|
||||
completed_at_unix_secs: None,
|
||||
updated_at_unix_secs: 200,
|
||||
error_code: None,
|
||||
error_message: None,
|
||||
video_url: None,
|
||||
request_metadata: None,
|
||||
})
|
||||
.await
|
||||
.expect("upsert should succeed");
|
||||
let conflict = repo
|
||||
.upsert(UpsertVideoTask {
|
||||
id: "task-1".to_string(),
|
||||
short_id: Some("short-task-1b".to_string()),
|
||||
request_id: "request-task-1b".to_string(),
|
||||
user_id: Some("user-2".to_string()),
|
||||
api_key_id: Some("api-key-2".to_string()),
|
||||
username: Some("user-2".to_string()),
|
||||
api_key_name: Some("secondary".to_string()),
|
||||
external_task_id: Some("ext-task-1b".to_string()),
|
||||
provider_id: Some("provider-2".to_string()),
|
||||
endpoint_id: Some("endpoint-2".to_string()),
|
||||
key_id: Some("provider-key-2".to_string()),
|
||||
client_api_format: Some("gemini:video".to_string()),
|
||||
provider_api_format: Some("gemini:video".to_string()),
|
||||
format_converted: false,
|
||||
model: Some("veo-3".to_string()),
|
||||
prompt: Some("remix".to_string()),
|
||||
original_request_body: Some(serde_json::json!({"prompt": "remix"})),
|
||||
duration_seconds: Some(8),
|
||||
resolution: Some("1080p".to_string()),
|
||||
aspect_ratio: Some("16:9".to_string()),
|
||||
size: Some("720p".to_string()),
|
||||
status: VideoTaskStatus::Processing,
|
||||
progress_percent: 50,
|
||||
progress_message: Some("processing".to_string()),
|
||||
retry_count: 1,
|
||||
poll_interval_seconds: 10,
|
||||
next_poll_at_unix_secs: Some(200),
|
||||
poll_count: 2,
|
||||
max_poll_count: 360,
|
||||
created_at_unix_ms: 150,
|
||||
submitted_at_unix_secs: Some(150),
|
||||
completed_at_unix_secs: None,
|
||||
updated_at_unix_secs: 200,
|
||||
error_code: None,
|
||||
error_message: None,
|
||||
video_url: None,
|
||||
request_metadata: None,
|
||||
})
|
||||
.await
|
||||
.expect_err("identity replacement should be rejected");
|
||||
assert!(conflict.to_string().contains("immutable field short_id"));
|
||||
|
||||
let stored = repo
|
||||
.find(VideoTaskLookupKey::Id("task-1"))
|
||||
.await
|
||||
.expect("find should succeed")
|
||||
.expect("original task should remain");
|
||||
assert_eq!(stored.request_id, "request-task-1");
|
||||
assert_eq!(stored.user_id.as_deref(), Some("user-1"));
|
||||
assert_eq!(stored.status, VideoTaskStatus::Submitted);
|
||||
assert!(repo
|
||||
.find(VideoTaskLookupKey::ShortId("short-task-1"))
|
||||
.await
|
||||
.expect("find should succeed")
|
||||
.is_none());
|
||||
.is_some());
|
||||
assert!(repo
|
||||
.find(VideoTaskLookupKey::UserExternal {
|
||||
user_id: "user-1",
|
||||
@@ -538,12 +652,117 @@ mod tests {
|
||||
})
|
||||
.await
|
||||
.expect("find should succeed")
|
||||
.is_none());
|
||||
.is_some());
|
||||
assert!(repo
|
||||
.find(VideoTaskLookupKey::ShortId("short-task-1b"))
|
||||
.await
|
||||
.expect("find should succeed")
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upsert_allows_same_identity_status_update() {
|
||||
let repo = InMemoryVideoTaskRepository::default();
|
||||
let task = sample_task("task-1", VideoTaskStatus::Submitted, 100);
|
||||
repo.upsert(task.clone())
|
||||
.await
|
||||
.expect("initial upsert should succeed");
|
||||
|
||||
let updated = repo
|
||||
.upsert(UpsertVideoTask {
|
||||
status: VideoTaskStatus::Processing,
|
||||
progress_percent: 50,
|
||||
poll_count: 2,
|
||||
created_at_unix_ms: 999,
|
||||
updated_at_unix_secs: 200,
|
||||
..task
|
||||
})
|
||||
.await
|
||||
.expect("same identity update should succeed");
|
||||
|
||||
assert_eq!(updated.status, VideoTaskStatus::Processing);
|
||||
assert_eq!(updated.progress_percent, 50);
|
||||
assert_eq!(updated.poll_count, 2);
|
||||
assert_eq!(updated.created_at_unix_ms, 90);
|
||||
assert_eq!(updated.updated_at_unix_secs, 200);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn update_if_active_rejects_identity_conflict_without_modification() {
|
||||
let repo = InMemoryVideoTaskRepository::default();
|
||||
let task = sample_task("task-1", VideoTaskStatus::Submitted, 100);
|
||||
repo.upsert(task.clone())
|
||||
.await
|
||||
.expect("initial upsert should succeed");
|
||||
|
||||
let result = repo
|
||||
.update_if_active(UpsertVideoTask {
|
||||
user_id: Some("attacker".to_string()),
|
||||
status: VideoTaskStatus::Completed,
|
||||
progress_percent: 100,
|
||||
updated_at_unix_secs: 200,
|
||||
..task
|
||||
})
|
||||
.await
|
||||
.expect("guarded update should execute");
|
||||
assert!(result.is_none());
|
||||
|
||||
let stored = repo
|
||||
.find(VideoTaskLookupKey::Id("task-1"))
|
||||
.await
|
||||
.expect("find should succeed")
|
||||
.expect("original task should remain");
|
||||
assert_eq!(stored.user_id.as_deref(), Some("user-1"));
|
||||
assert_eq!(stored.status, VideoTaskStatus::Submitted);
|
||||
assert_eq!(stored.progress_percent, 0);
|
||||
assert_eq!(stored.updated_at_unix_secs, 100);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upsert_rejects_secondary_unique_key_takeover() {
|
||||
let repo = InMemoryVideoTaskRepository::default();
|
||||
let original = sample_task("task-1", VideoTaskStatus::Submitted, 100);
|
||||
repo.upsert(original.clone())
|
||||
.await
|
||||
.expect("initial upsert should succeed");
|
||||
|
||||
let short_id_conflict = repo
|
||||
.upsert(UpsertVideoTask {
|
||||
id: "task-2".to_string(),
|
||||
request_id: "request-task-2".to_string(),
|
||||
..original.clone()
|
||||
})
|
||||
.await
|
||||
.expect_err("a short id must not be reassigned to another task");
|
||||
assert!(short_id_conflict.to_string().contains("existing short_id"));
|
||||
|
||||
let request_id_conflict = repo
|
||||
.upsert(UpsertVideoTask {
|
||||
id: "task-3".to_string(),
|
||||
short_id: Some("short-task-3".to_string()),
|
||||
..original
|
||||
})
|
||||
.await
|
||||
.expect_err("a request id must not be reassigned to another task");
|
||||
assert!(request_id_conflict
|
||||
.to_string()
|
||||
.contains("existing request_id"));
|
||||
|
||||
assert!(repo
|
||||
.find(VideoTaskLookupKey::ShortId("short-task-1"))
|
||||
.await
|
||||
.expect("find should succeed")
|
||||
.is_some());
|
||||
assert!(repo
|
||||
.find(VideoTaskLookupKey::Id("task-2"))
|
||||
.await
|
||||
.expect("find should succeed")
|
||||
.is_none());
|
||||
assert!(repo
|
||||
.find(VideoTaskLookupKey::Id("task-3"))
|
||||
.await
|
||||
.expect("find should succeed")
|
||||
.is_none());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,19 +1,34 @@
|
||||
mod memory;
|
||||
|
||||
pub use aether_data_contracts::repository::wallet::{
|
||||
redeem_code_credits_recharge_balance, redeem_code_payment_method,
|
||||
redeem_code_refundable_amount, AdjustWalletBalanceInput, AdminPaymentCallbackRecord,
|
||||
AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery,
|
||||
AdminWalletLedgerQuery, AdminWalletListQuery, AdminWalletPaymentOrderRecord,
|
||||
AdminWalletRefundRecord, AdminWalletRefundRequestListQuery, AdminWalletTransactionRecord,
|
||||
CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput,
|
||||
CreateAdminRedeemCodeBatchResult, CreateManualWalletRechargeInput,
|
||||
CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome, CreateWalletRechargeOrderInput,
|
||||
CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput,
|
||||
CreateWalletRefundRequestOutcome, CreatedAdminRedeemCodePlaintext,
|
||||
CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput,
|
||||
canonicalize_payment_method, canonicalize_wallet_refund_fields,
|
||||
payment_order_is_uncertain_wallet_checkout_placeholder,
|
||||
payment_order_refund_amounts_are_consistent,
|
||||
payment_order_stripe_client_secret_cas_replacement, project_wallet_gateway_response,
|
||||
project_wallet_recharge_gateway_response, redeem_code_credits_recharge_balance,
|
||||
redeem_code_payment_method, redeem_code_refundable_amount, stored_timestamp_unix_secs,
|
||||
validate_admin_redeem_code_batch_input, validate_payment_order_credit_amounts,
|
||||
validate_plan_purchase_order_input, validate_plan_wallet_credit_entitlements,
|
||||
validate_redeem_wallet_credit, validate_wallet_recharge_order_input,
|
||||
wallet_recharge_checkout_claim_response, wallet_recharge_checkout_claim_token,
|
||||
wallet_recharge_checkout_claimed_at, wallet_recharge_checkout_failed_response,
|
||||
wallet_recharge_checkout_uncertain_response, wallet_recharge_order_created_at_unix_secs,
|
||||
wallet_recharge_order_is_checkout_placeholder,
|
||||
wallet_recharge_order_is_reclaimable_placeholder, wallet_recharge_replay_matches,
|
||||
wallet_recharge_response_is_checkout_placeholder, wallet_refund_proof_is_success,
|
||||
AdjustWalletBalanceInput, AdminPaymentCallbackRecord, AdminPaymentOrderListQuery,
|
||||
AdminRedeemCodeBatchListQuery, AdminRedeemCodeListQuery, AdminWalletLedgerQuery,
|
||||
AdminWalletListQuery, AdminWalletPaymentOrderRecord, AdminWalletRefundRecord,
|
||||
AdminWalletRefundRequestListQuery, AdminWalletTransactionRecord, CanonicalWalletRefundFields,
|
||||
CompareAndSwapPaymentOrderStripeClientSecretInput, CompleteAdminWalletRefundInput,
|
||||
CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult,
|
||||
CreateManualWalletRechargeInput, CreatePlanPurchaseOrderInput, CreatePlanPurchaseOrderOutcome,
|
||||
CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome,
|
||||
CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome,
|
||||
CreatedAdminRedeemCodePlaintext, CreditAdminPaymentOrderInput, DeleteAdminRedeemCodeBatchInput,
|
||||
DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput,
|
||||
ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome,
|
||||
FailWalletRechargeCheckoutInput, InitializeAuthWalletOutcome, ProcessAdminWalletRefundInput,
|
||||
ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome, ReclaimWalletRechargeCheckoutInput,
|
||||
RedeemWalletCodeInput, RedeemWalletCodeOutcome, StoredAdminPaymentCallback,
|
||||
StoredAdminPaymentCallbackPage, StoredAdminPaymentOrder, StoredAdminPaymentOrderPage,
|
||||
StoredAdminRedeemCode, StoredAdminRedeemCodeBatch, StoredAdminRedeemCodeBatchPage,
|
||||
@@ -22,9 +37,10 @@ pub use aether_data_contracts::repository::wallet::{
|
||||
StoredAdminWalletRefundPage, StoredAdminWalletRefundRequestItem,
|
||||
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
|
||||
StoredAdminWalletTransactionPage, StoredWalletDailyUsageLedger,
|
||||
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome,
|
||||
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAdminWalletRefundGatewayInput,
|
||||
UpdateWalletRechargeCheckoutInput, WalletLookupKey, WalletMutationOutcome,
|
||||
WalletReadRepository, WalletReadSeed, WalletReadSnapshot, WalletRepository,
|
||||
WalletWriteRepository,
|
||||
WalletWriteRepository, WALLET_RECHARGE_CHECKOUT_CLAIM_LEASE_SECS,
|
||||
};
|
||||
#[cfg(feature = "mysql")]
|
||||
pub use aether_data_mysql::MysqlWalletReadRepository;
|
||||
|
||||
Reference in New Issue
Block a user