mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-06 17:37:47 +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:
@@ -1,5 +1,9 @@
|
||||
use async_trait::async_trait;
|
||||
|
||||
fn redacted_optional_secret<T>(value: &Option<T>) -> Option<&'static str> {
|
||||
value.as_ref().map(|_| "[REDACTED]")
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredAuthApiKeySnapshot {
|
||||
pub user_id: String,
|
||||
@@ -121,7 +125,7 @@ impl StoredAuthApiKeySnapshot {
|
||||
return false;
|
||||
}
|
||||
if let Some(expires_at_unix_secs) = self.api_key_expires_at_unix_secs {
|
||||
if expires_at_unix_secs < now_unix_secs {
|
||||
if expires_at_unix_secs <= now_unix_secs {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
@@ -350,7 +354,7 @@ pub async fn read_resolved_auth_api_key_snapshot_by_user_api_key_ids(
|
||||
.await
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredAuthApiKeyExportRecord {
|
||||
pub user_id: String,
|
||||
pub api_key_id: String,
|
||||
@@ -377,6 +381,22 @@ pub struct StoredAuthApiKeyExportRecord {
|
||||
pub is_standalone: bool,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredAuthApiKeyExportRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredAuthApiKeyExportRecord")
|
||||
.field("user_id", &self.user_id)
|
||||
.field("api_key_id", &self.api_key_id)
|
||||
.field("key_hash", &self.key_hash)
|
||||
.field(
|
||||
"key_encrypted",
|
||||
&redacted_optional_secret(&self.key_encrypted),
|
||||
)
|
||||
.field("is_standalone", &self.is_standalone)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredAuthApiKeyExportRecord {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
@@ -497,7 +517,7 @@ pub struct StandaloneApiKeyExportListQuery {
|
||||
pub is_active: Option<bool>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct CreateUserApiKeyRecord {
|
||||
pub user_id: String,
|
||||
pub api_key_id: String,
|
||||
@@ -511,6 +531,7 @@ pub struct CreateUserApiKeyRecord {
|
||||
pub rate_limit: i32,
|
||||
pub concurrent_limit: Option<i32>,
|
||||
pub force_capabilities: Option<serde_json::Value>,
|
||||
pub feature_settings: Option<serde_json::Value>,
|
||||
pub is_active: bool,
|
||||
pub expires_at_unix_secs: Option<u64>,
|
||||
pub auto_delete_on_expiry: bool,
|
||||
@@ -519,17 +540,61 @@ pub struct CreateUserApiKeyRecord {
|
||||
pub total_cost_usd: f64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
impl std::fmt::Debug for CreateUserApiKeyRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("CreateUserApiKeyRecord")
|
||||
.field("user_id", &self.user_id)
|
||||
.field("api_key_id", &self.api_key_id)
|
||||
.field("key_hash", &self.key_hash)
|
||||
.field(
|
||||
"key_encrypted",
|
||||
&redacted_optional_secret(&self.key_encrypted),
|
||||
)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct UpdateUserApiKeyBasicRecord {
|
||||
pub user_id: String,
|
||||
pub api_key_id: String,
|
||||
pub key_encrypted: Option<String>,
|
||||
/// Whether `key_encrypted` is an explicit replacement, including an explicit `NULL`.
|
||||
/// Ordinary callers should leave this false to retain the existing value.
|
||||
pub key_encrypted_present: bool,
|
||||
pub name: Option<String>,
|
||||
/// Whether `name` is an explicit replacement, including an explicit `NULL`.
|
||||
pub name_present: bool,
|
||||
pub rate_limit: Option<i32>,
|
||||
/// Whether `rate_limit` is an explicit replacement, including an explicit `NULL`.
|
||||
pub rate_limit_present: bool,
|
||||
pub concurrent_limit: Option<i32>,
|
||||
/// Whether `concurrent_limit` is an explicit replacement, including an explicit `NULL`.
|
||||
pub concurrent_limit_present: bool,
|
||||
pub ip_rules: Option<Option<Vec<String>>>,
|
||||
/// `Some(Some(value))` replaces the settings; `Some(None)` clears them; `None` leaves them
|
||||
/// unchanged. Keeping this patch in the basic mutation record lets repositories apply the
|
||||
/// complete user-key update in one atomic write.
|
||||
pub feature_settings: Option<Option<serde_json::Value>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for UpdateUserApiKeyBasicRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("UpdateUserApiKeyBasicRecord")
|
||||
.field("user_id", &self.user_id)
|
||||
.field("api_key_id", &self.api_key_id)
|
||||
.field(
|
||||
"key_encrypted",
|
||||
&redacted_optional_secret(&self.key_encrypted),
|
||||
)
|
||||
.field("key_encrypted_present", &self.key_encrypted_present)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct CreateStandaloneApiKeyRecord {
|
||||
pub user_id: String,
|
||||
pub api_key_id: String,
|
||||
@@ -551,10 +616,35 @@ pub struct CreateStandaloneApiKeyRecord {
|
||||
pub total_cost_usd: f64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
impl std::fmt::Debug for CreateStandaloneApiKeyRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("CreateStandaloneApiKeyRecord")
|
||||
.field("user_id", &self.user_id)
|
||||
.field("api_key_id", &self.api_key_id)
|
||||
.field("key_hash", &self.key_hash)
|
||||
.field(
|
||||
"key_encrypted",
|
||||
&redacted_optional_secret(&self.key_encrypted),
|
||||
)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct UpdateStandaloneApiKeyBasicRecord {
|
||||
pub api_key_id: String,
|
||||
pub key_encrypted: Option<String>,
|
||||
/// Whether `key_encrypted` is an explicit replacement, including an explicit `NULL`.
|
||||
/// Ordinary callers should leave this false to retain the existing value.
|
||||
pub key_encrypted_present: bool,
|
||||
pub name: Option<String>,
|
||||
/// Whether `name` is an explicit replacement, including an explicit `NULL`.
|
||||
/// Ordinary callers should leave this false to retain the existing value.
|
||||
pub name_present: bool,
|
||||
/// `Some(Some(value))` sets the capability map; `Some(None)` clears it; `None` leaves it
|
||||
/// unchanged.
|
||||
pub force_capabilities: Option<Option<serde_json::Value>>,
|
||||
pub rate_limit_present: bool,
|
||||
pub rate_limit: Option<i32>,
|
||||
pub concurrent_limit_present: bool,
|
||||
@@ -569,6 +659,48 @@ pub struct UpdateStandaloneApiKeyBasicRecord {
|
||||
pub auto_delete_on_expiry: bool,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for UpdateStandaloneApiKeyBasicRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("UpdateStandaloneApiKeyBasicRecord")
|
||||
.field("api_key_id", &self.api_key_id)
|
||||
.field(
|
||||
"key_encrypted",
|
||||
&redacted_optional_secret(&self.key_encrypted),
|
||||
)
|
||||
.field("key_encrypted_present", &self.key_encrypted_present)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
/// Replace only the recoverable API-key ciphertext when the complete immutable identity and the
|
||||
/// exact ciphertext observed by the caller still match. This is intentionally narrower than the
|
||||
/// ordinary admin update records so lazy envelope migration cannot overwrite a concurrent secret
|
||||
/// restore or move ciphertext between owners, scopes, hashes, or key IDs.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub struct CompareAndSwapAuthApiKeyCiphertext {
|
||||
pub user_id: String,
|
||||
pub api_key_id: String,
|
||||
pub key_hash: String,
|
||||
pub is_standalone: bool,
|
||||
pub expected_key_encrypted: String,
|
||||
pub key_encrypted: String,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for CompareAndSwapAuthApiKeyCiphertext {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("CompareAndSwapAuthApiKeyCiphertext")
|
||||
.field("user_id", &self.user_id)
|
||||
.field("api_key_id", &self.api_key_id)
|
||||
.field("key_hash", &self.key_hash)
|
||||
.field("is_standalone", &self.is_standalone)
|
||||
.field("expected_key_encrypted", &"[REDACTED]")
|
||||
.field("key_encrypted", &"[REDACTED]")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum AuthApiKeyLookupKey<'a> {
|
||||
KeyHash(&'a str),
|
||||
@@ -650,6 +782,20 @@ pub trait AuthApiKeyReadRepository: Send + Sync {
|
||||
pub trait AuthApiKeyWriteRepository: Send + Sync {
|
||||
async fn touch_last_used_at(&self, api_key_id: &str) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
/// Synchronize an authoritative user snapshot for repositories used by gateway tests.
|
||||
///
|
||||
/// Production database repositories validate the API-key owner in the same transaction as
|
||||
/// key creation and deliberately keep the default no-op implementation. Test repositories
|
||||
/// may override this hook, but must derive owner state exclusively from `StoredUserAuthRecord`
|
||||
/// rather than from an API-key mutation request.
|
||||
async fn synchronize_user_api_key_owner_for_tests(
|
||||
&self,
|
||||
user: &crate::repository::users::StoredUserAuthRecord,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
let _ = user;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn create_user_api_key(
|
||||
&self,
|
||||
record: CreateUserApiKeyRecord,
|
||||
@@ -660,16 +806,53 @@ pub trait AuthApiKeyWriteRepository: Send + Sync {
|
||||
record: CreateStandaloneApiKeyRecord,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError>;
|
||||
|
||||
async fn compare_and_swap_api_key_ciphertext(
|
||||
&self,
|
||||
mutation: &CompareAndSwapAuthApiKeyCiphertext,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
let _ = mutation;
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic API-key ciphertext migration is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn update_user_api_key_basic(
|
||||
&self,
|
||||
record: UpdateUserApiKeyBasicRecord,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError>;
|
||||
|
||||
/// Apply a user-owned key update only while the key remains unlocked.
|
||||
/// Implementations must test ownership, non-standalone status, and
|
||||
/// `is_locked = false` in the same atomic write as the mutation.
|
||||
async fn update_user_api_key_basic_if_unlocked(
|
||||
&self,
|
||||
record: UpdateUserApiKeyBasicRecord,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError> {
|
||||
let _ = record;
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic unlocked user API key update is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn update_standalone_api_key_basic(
|
||||
&self,
|
||||
record: UpdateStandaloneApiKeyBasicRecord,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError>;
|
||||
|
||||
/// Restore an API-key export row only when its complete exported post-state still matches
|
||||
/// `expected`. Implementations must perform the compare and update atomically so a failed
|
||||
/// import cannot overwrite a concurrent administrator change.
|
||||
async fn restore_api_key_if_matches(
|
||||
&self,
|
||||
expected: &StoredAuthApiKeyExportRecord,
|
||||
restored: &StoredAuthApiKeyExportRecord,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
let _ = (expected, restored);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic API key restore is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn set_user_api_key_active(
|
||||
&self,
|
||||
user_id: &str,
|
||||
@@ -677,6 +860,18 @@ pub trait AuthApiKeyWriteRepository: Send + Sync {
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError>;
|
||||
|
||||
async fn set_user_api_key_active_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError> {
|
||||
let _ = (user_id, api_key_id, is_active);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic unlocked user API key status update is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn set_standalone_api_key_active(
|
||||
&self,
|
||||
api_key_id: &str,
|
||||
@@ -697,6 +892,18 @@ pub trait AuthApiKeyWriteRepository: Send + Sync {
|
||||
allowed_providers: Option<Vec<String>>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError>;
|
||||
|
||||
async fn set_user_api_key_allowed_providers_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
allowed_providers: Option<Vec<String>>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError> {
|
||||
let _ = (user_id, api_key_id, allowed_providers);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic unlocked user API key provider update is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn set_user_api_key_force_capabilities(
|
||||
&self,
|
||||
user_id: &str,
|
||||
@@ -704,6 +911,18 @@ pub trait AuthApiKeyWriteRepository: Send + Sync {
|
||||
force_capabilities: Option<serde_json::Value>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError>;
|
||||
|
||||
async fn set_user_api_key_force_capabilities_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
force_capabilities: Option<serde_json::Value>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError> {
|
||||
let _ = (user_id, api_key_id, force_capabilities);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic unlocked user API key capability update is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn set_user_api_key_feature_settings(
|
||||
&self,
|
||||
user_id: &str,
|
||||
@@ -711,6 +930,18 @@ pub trait AuthApiKeyWriteRepository: Send + Sync {
|
||||
feature_settings: Option<serde_json::Value>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError>;
|
||||
|
||||
async fn set_user_api_key_feature_settings_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
feature_settings: Option<serde_json::Value>,
|
||||
) -> Result<Option<StoredAuthApiKeyExportRecord>, crate::DataLayerError> {
|
||||
let _ = (user_id, api_key_id, feature_settings);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic unlocked user API key feature update is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn set_api_key_usage_totals(
|
||||
&self,
|
||||
api_key_id: &str,
|
||||
@@ -725,6 +956,17 @@ pub trait AuthApiKeyWriteRepository: Send + Sync {
|
||||
api_key_id: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn delete_user_api_key_if_unlocked(
|
||||
&self,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
let _ = (user_id, api_key_id);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic unlocked user API key deletion is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn delete_standalone_api_key(
|
||||
&self,
|
||||
api_key_id: &str,
|
||||
@@ -762,7 +1004,9 @@ fn parse_string_list_value(
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, crate::DataLayerError> {
|
||||
match value {
|
||||
serde_json::Value::Null => Ok(None),
|
||||
serde_json::Value::Null => Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains JSON null; use SQL NULL for an unset policy"
|
||||
))),
|
||||
serde_json::Value::Array(array) => parse_string_list_array(array, field_name).map(Some),
|
||||
serde_json::Value::String(raw) => parse_embedded_string_list(raw, field_name),
|
||||
_ => Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
@@ -776,8 +1020,15 @@ fn parse_embedded_string_list(
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, crate::DataLayerError> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
|
||||
return Ok(None);
|
||||
if raw.is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains an empty string"
|
||||
)));
|
||||
}
|
||||
if raw.eq_ignore_ascii_case("null") {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains stringified JSON null; use SQL NULL for an unset policy"
|
||||
)));
|
||||
}
|
||||
|
||||
if let Ok(decoded) = serde_json::from_str::<serde_json::Value>(raw) {
|
||||
@@ -799,9 +1050,12 @@ fn parse_string_list_array(
|
||||
)));
|
||||
};
|
||||
let item = item.trim();
|
||||
if !item.is_empty() {
|
||||
items.push(item.to_string());
|
||||
if item.is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains an empty item"
|
||||
)));
|
||||
}
|
||||
items.push(item.to_string());
|
||||
}
|
||||
Ok(items)
|
||||
}
|
||||
@@ -826,10 +1080,80 @@ mod tests {
|
||||
use super::{
|
||||
read_resolved_auth_api_key_snapshot_by_key_hash,
|
||||
read_resolved_auth_api_key_snapshot_by_user_api_key_ids, AuthApiKeyLookupKey,
|
||||
ResolvedAuthApiKeySnapshot, ResolvedAuthApiKeySnapshotReader, StoredAuthApiKeyExportRecord,
|
||||
StoredAuthApiKeySnapshot,
|
||||
CompareAndSwapAuthApiKeyCiphertext, ResolvedAuthApiKeySnapshot,
|
||||
ResolvedAuthApiKeySnapshotReader, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn api_key_record_debug_output_redacts_recoverable_ciphertext() {
|
||||
let ciphertext = "debug-secret-api-key-ciphertext";
|
||||
let replacement = "debug-secret-api-key-replacement";
|
||||
let record = StoredAuthApiKeyExportRecord::new(
|
||||
"user-1".to_string(),
|
||||
"key-1".to_string(),
|
||||
"key-hash".to_string(),
|
||||
Some(ciphertext.to_string()),
|
||||
Some("test key".to_string()),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
None,
|
||||
false,
|
||||
0,
|
||||
0,
|
||||
0.0,
|
||||
false,
|
||||
)
|
||||
.expect("API key export record should build");
|
||||
let mutation = CompareAndSwapAuthApiKeyCiphertext {
|
||||
user_id: "user-1".to_string(),
|
||||
api_key_id: "key-1".to_string(),
|
||||
key_hash: "key-hash".to_string(),
|
||||
is_standalone: false,
|
||||
expected_key_encrypted: ciphertext.to_string(),
|
||||
key_encrypted: replacement.to_string(),
|
||||
};
|
||||
|
||||
for rendered in [format!("{record:?}"), format!("{mutation:?}")] {
|
||||
assert!(!rendered.contains(ciphertext));
|
||||
assert!(!rendered.contains(replacement));
|
||||
assert!(rendered.contains("[REDACTED]"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stored_security_lists_distinguish_sql_null_from_malformed_json_null() {
|
||||
assert_eq!(
|
||||
super::parse_string_list(None, "api_keys.allowed_providers")
|
||||
.expect("SQL NULL should remain an unset policy"),
|
||||
None
|
||||
);
|
||||
assert!(super::parse_string_list(
|
||||
Some(serde_json::Value::Null),
|
||||
"api_keys.allowed_providers"
|
||||
)
|
||||
.is_err());
|
||||
assert!(super::parse_string_list(
|
||||
Some(serde_json::json!("null")),
|
||||
"api_keys.allowed_providers"
|
||||
)
|
||||
.is_err());
|
||||
assert!(super::parse_string_list(
|
||||
Some(serde_json::json!([" "])),
|
||||
"api_keys.allowed_providers"
|
||||
)
|
||||
.is_err());
|
||||
assert_eq!(
|
||||
super::parse_string_list(Some(serde_json::json!([])), "api_keys.allowed_providers")
|
||||
.expect("an intentional empty policy should preserve its existing semantics"),
|
||||
Some(Vec::new())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn api_format_policy_intersection_preserves_companion_scope() {
|
||||
assert_eq!(
|
||||
@@ -988,6 +1312,8 @@ mod tests {
|
||||
)
|
||||
.expect("snapshot should build");
|
||||
|
||||
assert!(snapshot.is_currently_usable(99));
|
||||
assert!(!snapshot.is_currently_usable(100));
|
||||
assert!(!snapshot.is_currently_usable(101));
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
use async_trait::async_trait;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
fn redacted_optional_secret<T>(value: &Option<T>) -> Option<&'static str> {
|
||||
value.as_ref().map(|_| "[REDACTED]")
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredOAuthProviderModuleConfig {
|
||||
pub provider_type: String,
|
||||
pub display_name: String,
|
||||
@@ -9,6 +13,22 @@ pub struct StoredOAuthProviderModuleConfig {
|
||||
pub redirect_uri: String,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredOAuthProviderModuleConfig {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredOAuthProviderModuleConfig")
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("display_name", &self.display_name)
|
||||
.field("client_id", &self.client_id)
|
||||
.field(
|
||||
"client_secret_encrypted",
|
||||
&redacted_optional_secret(&self.client_secret_encrypted),
|
||||
)
|
||||
.field("redirect_uri", &self.redirect_uri)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredOAuthProviderModuleConfig {
|
||||
pub fn new(
|
||||
provider_type: String,
|
||||
@@ -37,7 +57,7 @@ impl StoredOAuthProviderModuleConfig {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredLdapModuleConfig {
|
||||
pub server_url: String,
|
||||
pub bind_dn: String,
|
||||
@@ -53,6 +73,57 @@ pub struct StoredLdapModuleConfig {
|
||||
pub connect_timeout: Option<i32>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredLdapModuleConfig {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredLdapModuleConfig")
|
||||
.field("server_url", &self.server_url)
|
||||
.field("bind_dn", &self.bind_dn)
|
||||
.field(
|
||||
"bind_password_encrypted",
|
||||
&redacted_optional_secret(&self.bind_password_encrypted),
|
||||
)
|
||||
.field("base_dn", &self.base_dn)
|
||||
.field("user_search_filter", &self.user_search_filter)
|
||||
.field("username_attr", &self.username_attr)
|
||||
.field("email_attr", &self.email_attr)
|
||||
.field("display_name_attr", &self.display_name_attr)
|
||||
.field("is_enabled", &self.is_enabled)
|
||||
.field("is_exclusive", &self.is_exclusive)
|
||||
.field("use_starttls", &self.use_starttls)
|
||||
.field("connect_timeout", &self.connect_timeout)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Explicit mutation semantics for the LDAP bind password.
|
||||
///
|
||||
/// The password is deliberately kept separate from [`StoredLdapModuleConfig`] updates so a
|
||||
/// caller that only changes non-secret fields cannot accidentally write a stale ciphertext back
|
||||
/// to storage.
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub enum LdapBindPasswordUpdate {
|
||||
Preserve,
|
||||
Set(String),
|
||||
Clear,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for LdapBindPasswordUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Preserve => formatter.write_str("Preserve"),
|
||||
Self::Set(_) => formatter.write_str("Set([REDACTED])"),
|
||||
Self::Clear => formatter.write_str("Clear"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub enum CompareAndSwapLdapConfigResult {
|
||||
Applied(StoredLdapModuleConfig),
|
||||
Conflict,
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait AuthModuleReadRepository: Send + Sync {
|
||||
async fn list_enabled_oauth_providers(
|
||||
@@ -66,8 +137,80 @@ pub trait AuthModuleReadRepository: Send + Sync {
|
||||
|
||||
#[async_trait]
|
||||
pub trait AuthModuleWriteRepository: Send + Sync {
|
||||
async fn upsert_ldap_config(
|
||||
/// Atomically create or replace the singleton LDAP configuration.
|
||||
///
|
||||
/// `expected` is the complete snapshot observed by the caller. `None` means the caller
|
||||
/// expects the singleton not to exist. Implementations must compare every persisted config
|
||||
/// field, including the encrypted password, before applying the replacement. The password
|
||||
/// field in `replacement` is never authoritative; only `bind_password_update` controls the
|
||||
/// stored secret.
|
||||
async fn compare_and_swap_ldap_config(
|
||||
&self,
|
||||
config: &StoredLdapModuleConfig,
|
||||
) -> Result<Option<StoredLdapModuleConfig>, crate::DataLayerError>;
|
||||
expected: Option<&StoredLdapModuleConfig>,
|
||||
replacement: &StoredLdapModuleConfig,
|
||||
bind_password_update: &LdapBindPasswordUpdate,
|
||||
) -> Result<CompareAndSwapLdapConfigResult, crate::DataLayerError>;
|
||||
|
||||
/// Delete the singleton LDAP configuration only when every persisted field still matches the
|
||||
/// supplied snapshot. This is used by aggregate-import compensation and must not remove a
|
||||
/// configuration that another operation changed after it was created.
|
||||
async fn delete_ldap_config_if_matches(
|
||||
&self,
|
||||
expected: &StoredLdapModuleConfig,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn compare_and_swap_ldap_bind_password(
|
||||
&self,
|
||||
_expected: &str,
|
||||
_replacement: &str,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
Err(crate::DataLayerError::InvalidConfiguration(
|
||||
"LDAP bind password compare-and-swap is not supported by this repository".to_string(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{LdapBindPasswordUpdate, StoredLdapModuleConfig, StoredOAuthProviderModuleConfig};
|
||||
|
||||
#[test]
|
||||
fn auth_module_debug_output_redacts_encrypted_secrets() {
|
||||
let oauth_secret = "debug-secret-oauth-ciphertext";
|
||||
let oauth = StoredOAuthProviderModuleConfig::new(
|
||||
"linuxdo".to_string(),
|
||||
"Linux.do".to_string(),
|
||||
"client-id".to_string(),
|
||||
Some(oauth_secret.to_string()),
|
||||
"https://example.com/callback".to_string(),
|
||||
)
|
||||
.expect("OAuth module config should build");
|
||||
let ldap_secret = "debug-secret-ldap-ciphertext";
|
||||
let ldap = StoredLdapModuleConfig {
|
||||
server_url: "ldaps://ldap.example.com".to_string(),
|
||||
bind_dn: "cn=admin,dc=example,dc=com".to_string(),
|
||||
bind_password_encrypted: Some(ldap_secret.to_string()),
|
||||
base_dn: "dc=example,dc=com".to_string(),
|
||||
user_search_filter: None,
|
||||
username_attr: None,
|
||||
email_attr: None,
|
||||
display_name_attr: None,
|
||||
is_enabled: true,
|
||||
is_exclusive: false,
|
||||
use_starttls: false,
|
||||
connect_timeout: Some(5),
|
||||
};
|
||||
|
||||
for (rendered, secret) in [
|
||||
(format!("{oauth:?}"), oauth_secret),
|
||||
(format!("{ldap:?}"), ldap_secret),
|
||||
(
|
||||
format!("{:?}", LdapBindPasswordUpdate::Set(ldap_secret.to_string())),
|
||||
ldap_secret,
|
||||
),
|
||||
] {
|
||||
assert!(!rendered.contains(secret));
|
||||
assert!(rendered.contains("[REDACTED]"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,6 +3,45 @@ use std::collections::BTreeMap;
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
const BACKGROUND_TASK_DEFAULT_ERROR_CODE: &str = "background_task_failed";
|
||||
const BACKGROUND_TASK_UNCLASSIFIED_EVENT: &str = "unclassified_event";
|
||||
const MAX_BACKGROUND_TASK_METADATA_FIELDS: usize = 48;
|
||||
const SAFE_BACKGROUND_TASK_METADATA_FIELDS: &[&str] = &[
|
||||
"automatic_deletions",
|
||||
"bytes",
|
||||
"compression",
|
||||
"created_count",
|
||||
"deleted_endpoints",
|
||||
"deleted_keys",
|
||||
"encryption",
|
||||
"error_code",
|
||||
"export_version",
|
||||
"exported_at",
|
||||
"failed",
|
||||
"import_kind",
|
||||
"legacy_encrypted_copies_created",
|
||||
"legacy_encrypted_copies_verified",
|
||||
"legacy_plaintext_objects_deleted",
|
||||
"legacy_plaintext_objects_retained",
|
||||
"object_cleanup_mode",
|
||||
"partition",
|
||||
"provider_type",
|
||||
"replaced_count",
|
||||
"retention_cleanup_candidates",
|
||||
"scheduled_slot",
|
||||
"scope",
|
||||
"sha256",
|
||||
"stage",
|
||||
"status",
|
||||
"success",
|
||||
"total",
|
||||
"total_endpoints",
|
||||
"total_keys",
|
||||
"trigger",
|
||||
"versioned_storage_cleanup_notice",
|
||||
"versioned_storage_cleanup_required",
|
||||
];
|
||||
|
||||
#[derive(
|
||||
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
|
||||
)]
|
||||
@@ -101,6 +140,18 @@ pub struct StoredBackgroundTaskRun {
|
||||
pub updated_at_unix_secs: u64,
|
||||
}
|
||||
|
||||
impl StoredBackgroundTaskRun {
|
||||
pub fn sanitize_persisted_data(&mut self) {
|
||||
self.owner_instance = None;
|
||||
self.created_by = sanitize_background_task_actor(self.created_by.take());
|
||||
self.progress_message = None;
|
||||
self.payload_json = sanitize_background_task_metadata(self.payload_json.take());
|
||||
self.result_json = sanitize_background_task_metadata(self.result_json.take());
|
||||
self.error_message =
|
||||
sanitize_background_task_error_code(self.status, self.error_message.take());
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct UpsertBackgroundTaskRun {
|
||||
pub id: String,
|
||||
@@ -125,6 +176,16 @@ pub struct UpsertBackgroundTaskRun {
|
||||
}
|
||||
|
||||
impl UpsertBackgroundTaskRun {
|
||||
pub fn sanitize_for_persistence(&mut self) {
|
||||
self.owner_instance = None;
|
||||
self.created_by = sanitize_background_task_actor(self.created_by.take());
|
||||
self.progress_message = None;
|
||||
self.payload_json = sanitize_background_task_metadata(self.payload_json.take());
|
||||
self.result_json = sanitize_background_task_metadata(self.result_json.take());
|
||||
self.error_message =
|
||||
sanitize_background_task_error_code(self.status, self.error_message.take());
|
||||
}
|
||||
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
if self.id.trim().is_empty()
|
||||
|| self.task_key.trim().is_empty()
|
||||
@@ -143,7 +204,8 @@ impl UpsertBackgroundTaskRun {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn into_stored(self) -> StoredBackgroundTaskRun {
|
||||
pub fn into_stored(mut self) -> StoredBackgroundTaskRun {
|
||||
self.sanitize_for_persistence();
|
||||
StoredBackgroundTaskRun {
|
||||
id: self.id,
|
||||
task_key: self.task_key,
|
||||
@@ -178,6 +240,14 @@ pub struct StoredBackgroundTaskEvent {
|
||||
pub created_at_unix_secs: u64,
|
||||
}
|
||||
|
||||
impl StoredBackgroundTaskEvent {
|
||||
pub fn sanitize_persisted_data(&mut self) {
|
||||
self.event_type = sanitize_background_task_event_type(&self.event_type);
|
||||
self.message = self.event_type.clone();
|
||||
self.payload_json = sanitize_background_task_metadata(self.payload_json.take());
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub struct UpsertBackgroundTaskEvent {
|
||||
pub id: String,
|
||||
@@ -189,6 +259,12 @@ pub struct UpsertBackgroundTaskEvent {
|
||||
}
|
||||
|
||||
impl UpsertBackgroundTaskEvent {
|
||||
pub fn sanitize_for_persistence(&mut self) {
|
||||
self.event_type = sanitize_background_task_event_type(&self.event_type);
|
||||
self.message = self.event_type.clone();
|
||||
self.payload_json = sanitize_background_task_metadata(self.payload_json.take());
|
||||
}
|
||||
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
if self.id.trim().is_empty()
|
||||
|| self.run_id.trim().is_empty()
|
||||
@@ -202,7 +278,8 @@ impl UpsertBackgroundTaskEvent {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn into_stored(self) -> StoredBackgroundTaskEvent {
|
||||
pub fn into_stored(mut self) -> StoredBackgroundTaskEvent {
|
||||
self.sanitize_for_persistence();
|
||||
StoredBackgroundTaskEvent {
|
||||
id: self.id,
|
||||
run_id: self.run_id,
|
||||
@@ -214,6 +291,199 @@ impl UpsertBackgroundTaskEvent {
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_background_task_error_code(
|
||||
status: BackgroundTaskStatus,
|
||||
value: Option<String>,
|
||||
) -> Option<String> {
|
||||
if status != BackgroundTaskStatus::Failed {
|
||||
return None;
|
||||
}
|
||||
let value = value?.trim().to_ascii_lowercase();
|
||||
let code = match value.as_str() {
|
||||
"background_task_failed"
|
||||
| "background_task_panicked"
|
||||
| "provider_delete_failed"
|
||||
| "provider_oauth_batch_import_failed"
|
||||
| "s3_backup_failed"
|
||||
| "s3_backup_slot_record_failed" => value,
|
||||
_ => BACKGROUND_TASK_DEFAULT_ERROR_CODE.to_string(),
|
||||
};
|
||||
Some(code)
|
||||
}
|
||||
|
||||
fn sanitize_background_task_event_type(value: &str) -> String {
|
||||
let value = value.trim().to_ascii_lowercase();
|
||||
match value.as_str() {
|
||||
"cancel_requested" | "failed" | "queued" | "running" | "skipped" | "succeeded"
|
||||
| "worker_boot" => value,
|
||||
_ => BACKGROUND_TASK_UNCLASSIFIED_EVENT.to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_background_task_actor(value: Option<String>) -> Option<String> {
|
||||
let value = value?.trim().to_ascii_lowercase();
|
||||
matches!(value.as_str(), "admin" | "scheduler" | "system").then_some(value)
|
||||
}
|
||||
|
||||
fn sanitize_background_task_metadata(value: Option<Value>) -> Option<Value> {
|
||||
let Value::Object(object) = value? else {
|
||||
return None;
|
||||
};
|
||||
let mut sanitized = serde_json::Map::new();
|
||||
for (key, value) in object.into_iter().take(MAX_BACKGROUND_TASK_METADATA_FIELDS) {
|
||||
let normalized_key = key.trim().to_ascii_lowercase();
|
||||
if !SAFE_BACKGROUND_TASK_METADATA_FIELDS.contains(&normalized_key.as_str()) {
|
||||
continue;
|
||||
}
|
||||
let Some(value) = sanitize_background_task_metadata_value(&normalized_key, value) else {
|
||||
continue;
|
||||
};
|
||||
sanitized.insert(normalized_key, value);
|
||||
}
|
||||
(!sanitized.is_empty()).then_some(Value::Object(sanitized))
|
||||
}
|
||||
|
||||
fn sanitize_background_task_metadata_value(key: &str, value: Value) -> Option<Value> {
|
||||
match key {
|
||||
"automatic_deletions"
|
||||
| "bytes"
|
||||
| "created_count"
|
||||
| "deleted_endpoints"
|
||||
| "deleted_keys"
|
||||
| "failed"
|
||||
| "legacy_encrypted_copies_created"
|
||||
| "legacy_encrypted_copies_verified"
|
||||
| "legacy_plaintext_objects_deleted"
|
||||
| "legacy_plaintext_objects_retained"
|
||||
| "partition"
|
||||
| "replaced_count"
|
||||
| "retention_cleanup_candidates"
|
||||
| "success"
|
||||
| "total"
|
||||
| "total_endpoints"
|
||||
| "total_keys" => value.is_u64().then_some(value),
|
||||
"versioned_storage_cleanup_required" => value.is_boolean().then_some(value),
|
||||
"error_code" => value
|
||||
.as_str()
|
||||
.and_then(sanitize_background_task_metadata_error_code)
|
||||
.map(Value::String),
|
||||
"scope" => sanitize_background_task_metadata_enum(value, &["config", "data", "users"]),
|
||||
"compression" => sanitize_background_task_metadata_enum(value, &["zstd"]),
|
||||
"encryption" => sanitize_background_task_metadata_enum(value, &["aes-256-gcm-v2"]),
|
||||
"trigger" => sanitize_background_task_metadata_enum(value, &["manual", "scheduled"]),
|
||||
"import_kind" => sanitize_background_task_metadata_enum(
|
||||
value,
|
||||
&["agent_identity", "cookie_authorize", "oauth_batch"],
|
||||
),
|
||||
"provider_type" => sanitize_background_task_metadata_enum(
|
||||
value,
|
||||
&[
|
||||
"antigravity",
|
||||
"chatgpt_web",
|
||||
"claude_code",
|
||||
"codex",
|
||||
"gemini_cli",
|
||||
"kiro",
|
||||
"windsurf",
|
||||
],
|
||||
),
|
||||
"status" => sanitize_background_task_metadata_enum(
|
||||
value,
|
||||
&[
|
||||
"cancelled",
|
||||
"completed",
|
||||
"failed",
|
||||
"pending",
|
||||
"processing",
|
||||
"queued",
|
||||
"running",
|
||||
"skipped",
|
||||
"succeeded",
|
||||
],
|
||||
),
|
||||
"stage" => sanitize_background_task_metadata_enum(
|
||||
value,
|
||||
&[
|
||||
"completed",
|
||||
"deleting_endpoints",
|
||||
"deleting_keys",
|
||||
"deleting_models",
|
||||
"deleting_provider",
|
||||
"disabling",
|
||||
"failed",
|
||||
"preparing",
|
||||
"queued",
|
||||
"skipped",
|
||||
],
|
||||
),
|
||||
"object_cleanup_mode" => sanitize_background_task_metadata_enum(
|
||||
value,
|
||||
&["legacy_plaintext_deleted_after_verified_encryption"],
|
||||
),
|
||||
"sha256" => sanitize_background_task_metadata_hex(value, 64, 64),
|
||||
"export_version" => sanitize_background_task_metadata_version(value),
|
||||
"exported_at" => sanitize_background_task_metadata_rfc3339(value),
|
||||
"scheduled_slot" => sanitize_background_task_metadata_scheduled_slot(value),
|
||||
"versioned_storage_cleanup_notice" => sanitize_background_task_metadata_enum(
|
||||
value,
|
||||
&["legacy_plaintext_versions_require_external_cleanup"],
|
||||
),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_background_task_metadata_enum(value: Value, allowed: &[&str]) -> Option<Value> {
|
||||
let value = value.as_str()?.trim().to_ascii_lowercase();
|
||||
allowed
|
||||
.contains(&value.as_str())
|
||||
.then_some(Value::String(value))
|
||||
}
|
||||
|
||||
fn sanitize_background_task_metadata_hex(value: Value, min: usize, max: usize) -> Option<Value> {
|
||||
let value = value.as_str()?.trim().to_ascii_lowercase();
|
||||
((min..=max).contains(&value.len()) && value.bytes().all(|byte| byte.is_ascii_hexdigit()))
|
||||
.then_some(Value::String(value))
|
||||
}
|
||||
|
||||
fn sanitize_background_task_metadata_version(value: Value) -> Option<Value> {
|
||||
let value = value.as_str()?.trim();
|
||||
(!value.is_empty()
|
||||
&& value.len() <= 32
|
||||
&& value
|
||||
.bytes()
|
||||
.all(|byte| byte.is_ascii_digit() || matches!(byte, b'.' | b'-' | b'_')))
|
||||
.then_some(Value::String(value.to_string()))
|
||||
}
|
||||
|
||||
fn sanitize_background_task_metadata_rfc3339(value: Value) -> Option<Value> {
|
||||
let value = value.as_str()?.trim();
|
||||
(value.len() <= 64 && chrono::DateTime::parse_from_rfc3339(value).is_ok())
|
||||
.then_some(Value::String(value.to_string()))
|
||||
}
|
||||
|
||||
fn sanitize_background_task_metadata_scheduled_slot(value: Value) -> Option<Value> {
|
||||
let value = value.as_str()?.trim();
|
||||
let (unit, timestamp) = value.split_once(':')?;
|
||||
(matches!(unit, "hours" | "days" | "weeks" | "months")
|
||||
&& value.len() <= 80
|
||||
&& chrono::DateTime::parse_from_rfc3339(timestamp).is_ok())
|
||||
.then_some(Value::String(value.to_string()))
|
||||
}
|
||||
|
||||
fn sanitize_background_task_metadata_error_code(value: &str) -> Option<String> {
|
||||
let value = value.trim().to_ascii_lowercase();
|
||||
Some(match value.as_str() {
|
||||
"background_task_failed"
|
||||
| "background_task_panicked"
|
||||
| "provider_delete_failed"
|
||||
| "provider_oauth_batch_import_failed"
|
||||
| "s3_backup_failed"
|
||||
| "s3_backup_slot_record_failed" => value,
|
||||
_ if value.is_empty() => return None,
|
||||
_ => BACKGROUND_TASK_DEFAULT_ERROR_CODE.to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
|
||||
pub struct BackgroundTaskListQuery {
|
||||
pub task_key_substring: Option<String>,
|
||||
@@ -288,3 +558,152 @@ impl<T> BackgroundTaskRepository for T where
|
||||
T: BackgroundTaskReadRepository + BackgroundTaskWriteRepository + Send + Sync
|
||||
{
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
fn background_task_run(status: BackgroundTaskStatus) -> UpsertBackgroundTaskRun {
|
||||
UpsertBackgroundTaskRun {
|
||||
id: "run-1".to_string(),
|
||||
task_key: "security-review".to_string(),
|
||||
kind: BackgroundTaskKind::OnDemand,
|
||||
trigger: "manual".to_string(),
|
||||
status,
|
||||
attempt: 1,
|
||||
max_attempts: 3,
|
||||
owner_instance: None,
|
||||
progress_percent: 50,
|
||||
progress_message: None,
|
||||
payload_json: None,
|
||||
result_json: None,
|
||||
error_message: None,
|
||||
cancel_requested: false,
|
||||
created_by: None,
|
||||
created_at_unix_secs: 1,
|
||||
started_at_unix_secs: Some(2),
|
||||
finished_at_unix_secs: None,
|
||||
updated_at_unix_secs: 3,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn run_sanitization_removes_sensitive_and_nested_metadata() {
|
||||
let mut run = background_task_run(BackgroundTaskStatus::Running);
|
||||
run.owner_instance = Some("gateway-a".to_string());
|
||||
run.created_by = Some("[email protected]".to_string());
|
||||
run.progress_message = Some("Bearer secret-access-token".to_string());
|
||||
run.payload_json = Some(json!({
|
||||
"provider_id": "provider-1",
|
||||
"gateway_instance_id": "gateway-a",
|
||||
"bucket": "safe-bucket",
|
||||
"partition": 7,
|
||||
"success": 1,
|
||||
"access_token": "secret-access-token",
|
||||
"authorization": "Bearer secret-access-token",
|
||||
"password": "secret-password",
|
||||
"error": "upstream detail containing secret",
|
||||
"detail": "private diagnostic",
|
||||
"provider_id ": "Bearer secret-access-token",
|
||||
"gateway_instance_id ": "gateway-a; Authorization: secret",
|
||||
"bucket ": "secret/bucket",
|
||||
"nested": {"refresh_token": "secret-refresh-token"},
|
||||
"unknown": "must not be persisted"
|
||||
}));
|
||||
|
||||
let stored = run.into_stored();
|
||||
|
||||
assert_eq!(stored.progress_message, None);
|
||||
assert_eq!(stored.owner_instance, None);
|
||||
assert_eq!(stored.created_by, None);
|
||||
assert_eq!(
|
||||
stored.payload_json,
|
||||
Some(json!({
|
||||
"partition": 7,
|
||||
"success": 1
|
||||
}))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn run_sanitization_classifies_errors_and_clears_non_failure_errors() {
|
||||
let mut failed = background_task_run(BackgroundTaskStatus::Failed);
|
||||
failed.error_message = Some("upstream response included a credential".to_string());
|
||||
failed.result_json = Some(json!({
|
||||
"error_code": "raw-provider-error",
|
||||
"failed": 2,
|
||||
"token": "secret"
|
||||
}));
|
||||
|
||||
let failed = failed.into_stored();
|
||||
assert_eq!(
|
||||
failed.error_message.as_deref(),
|
||||
Some(BACKGROUND_TASK_DEFAULT_ERROR_CODE)
|
||||
);
|
||||
assert_eq!(
|
||||
failed.result_json,
|
||||
Some(json!({
|
||||
"error_code": BACKGROUND_TASK_DEFAULT_ERROR_CODE,
|
||||
"failed": 2
|
||||
}))
|
||||
);
|
||||
|
||||
let mut succeeded = background_task_run(BackgroundTaskStatus::Succeeded);
|
||||
succeeded.error_message = Some("provider_delete_failed".to_string());
|
||||
assert_eq!(succeeded.into_stored().error_message, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn historical_run_sanitization_applies_the_same_read_boundary() {
|
||||
let mut stored = background_task_run(BackgroundTaskStatus::Failed).into_stored();
|
||||
stored.progress_message = Some("legacy diagnostic with token".to_string());
|
||||
stored.payload_json = Some(json!({
|
||||
"scope": "data",
|
||||
"refresh_token": "legacy-secret",
|
||||
"nested": {"password": "legacy-password"}
|
||||
}));
|
||||
stored.error_message = Some("legacy upstream error: legacy-secret".to_string());
|
||||
|
||||
stored.sanitize_persisted_data();
|
||||
|
||||
assert_eq!(stored.progress_message, None);
|
||||
assert_eq!(stored.payload_json, Some(json!({"scope": "data"})));
|
||||
assert_eq!(
|
||||
stored.error_message.as_deref(),
|
||||
Some(BACKGROUND_TASK_DEFAULT_ERROR_CODE)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn event_sanitization_canonicalizes_type_message_and_payload() {
|
||||
let event = UpsertBackgroundTaskEvent {
|
||||
id: "event-1".to_string(),
|
||||
run_id: "run-1".to_string(),
|
||||
event_type: "provider returned secret-token".to_string(),
|
||||
message: "Authorization: Bearer secret-token".to_string(),
|
||||
payload_json: Some(json!({
|
||||
"stage": "finalize",
|
||||
"bytes": 42,
|
||||
"error_code": "upstream said secret-token",
|
||||
"error": "secret-token",
|
||||
"detail": "credential detail",
|
||||
"token": "secret-token",
|
||||
"nested": {"authorization": "Bearer secret-token"}
|
||||
})),
|
||||
created_at_unix_secs: 4,
|
||||
}
|
||||
.into_stored();
|
||||
|
||||
assert_eq!(event.event_type, BACKGROUND_TASK_UNCLASSIFIED_EVENT);
|
||||
assert_eq!(event.message, BACKGROUND_TASK_UNCLASSIFIED_EVENT);
|
||||
assert_eq!(
|
||||
event.payload_json,
|
||||
Some(json!({
|
||||
"bytes": 42,
|
||||
"error_code": BACKGROUND_TASK_DEFAULT_ERROR_CODE
|
||||
}))
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,9 +1,27 @@
|
||||
mod replacement;
|
||||
mod types;
|
||||
mod usage_policy;
|
||||
|
||||
pub use replacement::{
|
||||
entitlements_have_replacement_selector, entitlements_should_replace_existing,
|
||||
validate_entitlement_replacement_groups, EntitlementReplacementGroupValidationError,
|
||||
ENTITLEMENT_REPLACEMENT_GROUP_FIELD, MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH,
|
||||
};
|
||||
pub use types::{
|
||||
checked_plan_duration_days, checked_plan_duration_days_from_snapshot,
|
||||
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
|
||||
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
|
||||
BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository, PaymentGatewayConfigRecord,
|
||||
PaymentGatewayConfigWriteInput, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
|
||||
BillingPlanRecord, BillingPlanWriteInput, BillingReadRepository,
|
||||
PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigRecord, PaymentGatewayConfigWriteInput,
|
||||
PaymentGatewaySecretCasUpdate, StoredBillingModelContext, UserDailyQuotaAvailabilityRecord,
|
||||
UserPlanEntitlementRecord,
|
||||
};
|
||||
pub use usage_policy::{
|
||||
nonnegative_usd_to_usage_policy_cost_units, parse_usage_policy_entitlements,
|
||||
usd_to_usage_policy_cost_units, UsagePolicyEnforcement, UsagePolicyEntitlement,
|
||||
UsagePolicyEntitlementType, UsagePolicyMetric, UsagePolicyParseError, UsagePolicyRule,
|
||||
UsagePolicyValidationError, UsagePolicyWindow, MAX_USAGE_POLICY_ENTITLEMENTS,
|
||||
MAX_USAGE_POLICY_EXACT_INTEGER, MAX_USAGE_POLICY_ROLLING_WINDOW_SECONDS,
|
||||
MAX_USAGE_POLICY_RULES, MAX_USAGE_POLICY_TEXT_LENGTH, MAX_USAGE_POLICY_TOTAL_RULES,
|
||||
USAGE_POLICY_COST_UNITS_PER_USD, USAGE_POLICY_ENTITLEMENT_TYPE,
|
||||
};
|
||||
|
||||
@@ -0,0 +1,200 @@
|
||||
use std::collections::HashSet;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
pub const ENTITLEMENT_REPLACEMENT_GROUP_FIELD: &str = "replacement_group";
|
||||
pub const MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH: usize = 128;
|
||||
|
||||
const LEGACY_REPLACEMENT_ENTITLEMENT_TYPES: [&str; 2] = ["daily_quota", "membership_group"];
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, thiserror::Error)]
|
||||
pub enum EntitlementReplacementGroupValidationError {
|
||||
#[error("entitlements must be an array")]
|
||||
EntitlementsMustBeArray,
|
||||
#[error("entitlements[{index}].replacement_group must be a string")]
|
||||
InvalidType { index: usize },
|
||||
#[error("entitlements[{index}].replacement_group must not be empty")]
|
||||
Empty { index: usize },
|
||||
#[error("entitlements[{index}].replacement_group exceeds maximum length {max_len}")]
|
||||
TooLong { index: usize, max_len: usize },
|
||||
}
|
||||
|
||||
pub fn validate_entitlement_replacement_groups(
|
||||
entitlements: &Value,
|
||||
) -> Result<(), EntitlementReplacementGroupValidationError> {
|
||||
let items = entitlements
|
||||
.as_array()
|
||||
.ok_or(EntitlementReplacementGroupValidationError::EntitlementsMustBeArray)?;
|
||||
|
||||
for (index, item) in items.iter().enumerate() {
|
||||
let Some(group) = item.get(ENTITLEMENT_REPLACEMENT_GROUP_FIELD) else {
|
||||
continue;
|
||||
};
|
||||
let group = group
|
||||
.as_str()
|
||||
.ok_or(EntitlementReplacementGroupValidationError::InvalidType { index })?
|
||||
.trim();
|
||||
if group.is_empty() {
|
||||
return Err(EntitlementReplacementGroupValidationError::Empty { index });
|
||||
}
|
||||
if group.chars().count() > MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH {
|
||||
return Err(EntitlementReplacementGroupValidationError::TooLong {
|
||||
index,
|
||||
max_len: MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH,
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn entitlements_have_replacement_selector(entitlements: &Value) -> bool {
|
||||
let Some(items) = entitlements.as_array() else {
|
||||
return false;
|
||||
};
|
||||
|
||||
items.iter().any(|item| {
|
||||
let entitlement_type = item.get("type").and_then(Value::as_str);
|
||||
LEGACY_REPLACEMENT_ENTITLEMENT_TYPES.contains(&entitlement_type.unwrap_or_default())
|
||||
|| replacement_group(item).is_some()
|
||||
})
|
||||
}
|
||||
|
||||
pub fn entitlements_should_replace_existing(incoming: &Value, existing: &Value) -> bool {
|
||||
let (Some(incoming), Some(existing)) = (incoming.as_array(), existing.as_array()) else {
|
||||
return false;
|
||||
};
|
||||
|
||||
if LEGACY_REPLACEMENT_ENTITLEMENT_TYPES.iter().any(|kind| {
|
||||
entitlement_items_have_type(incoming, kind) && entitlement_items_have_type(existing, kind)
|
||||
}) {
|
||||
return true;
|
||||
}
|
||||
|
||||
let incoming_groups = incoming
|
||||
.iter()
|
||||
.filter_map(replacement_group)
|
||||
.collect::<HashSet<_>>();
|
||||
!incoming_groups.is_empty()
|
||||
&& existing
|
||||
.iter()
|
||||
.filter_map(replacement_group)
|
||||
.any(|group| incoming_groups.contains(group))
|
||||
}
|
||||
|
||||
fn entitlement_items_have_type(items: &[Value], entitlement_type: &str) -> bool {
|
||||
items
|
||||
.iter()
|
||||
.any(|item| item.get("type").and_then(Value::as_str) == Some(entitlement_type))
|
||||
}
|
||||
|
||||
fn replacement_group(item: &Value) -> Option<&str> {
|
||||
item.get(ENTITLEMENT_REPLACEMENT_GROUP_FIELD)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|group| !group.is_empty())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn legacy_daily_quota_and_membership_groups_remain_mutually_exclusive() {
|
||||
assert!(entitlements_should_replace_existing(
|
||||
&json!([{"type": "daily_quota", "daily_quota_usd": 20}]),
|
||||
&json!([
|
||||
{"type": "daily_quota", "daily_quota_usd": 10},
|
||||
{"type": "usage_policy", "rules": []}
|
||||
]),
|
||||
));
|
||||
assert!(entitlements_should_replace_existing(
|
||||
&json!([{"type": "membership_group", "grant_user_groups": ["pro"]}]),
|
||||
&json!([{"type": "membership_group", "grant_user_groups": ["basic"]}]),
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_policies_stack_by_default() {
|
||||
let incoming = json!([{"type": "usage_policy", "policy_id": "weekly", "rules": []}]);
|
||||
let existing = json!([{"type": "usage_policy", "policy_id": "five-hour", "rules": []}]);
|
||||
|
||||
assert!(!entitlements_have_replacement_selector(&incoming));
|
||||
assert!(!entitlements_should_replace_existing(&incoming, &existing));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn matching_explicit_groups_replace_the_whole_package() {
|
||||
let incoming = json!([{
|
||||
"type": "usage_policy",
|
||||
"replacement_group": "pro-tier",
|
||||
"rules": []
|
||||
}]);
|
||||
let existing = json!([
|
||||
{"type": "wallet_credit", "amount_usd": 10},
|
||||
{
|
||||
"type": "usage_policy",
|
||||
"replacement_group": "pro-tier",
|
||||
"rules": []
|
||||
}
|
||||
]);
|
||||
|
||||
assert!(entitlements_have_replacement_selector(&incoming));
|
||||
assert!(entitlements_should_replace_existing(&incoming, &existing));
|
||||
assert!(!entitlements_should_replace_existing(
|
||||
&incoming,
|
||||
&json!([{
|
||||
"type": "usage_policy",
|
||||
"replacement_group": "team-tier",
|
||||
"rules": []
|
||||
}]),
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_groups_can_span_entitlement_types_and_ignore_outer_whitespace() {
|
||||
assert!(entitlements_should_replace_existing(
|
||||
&json!([{
|
||||
"type": "usage_policy",
|
||||
"replacement_group": " traffic-tier ",
|
||||
"rules": []
|
||||
}]),
|
||||
&json!([{
|
||||
"type": "wallet_credit",
|
||||
"replacement_group": "traffic-tier",
|
||||
"amount_usd": 10
|
||||
}]),
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validates_explicit_group_shape_and_bounds() {
|
||||
assert!(validate_entitlement_replacement_groups(&json!([{
|
||||
"type": "usage_policy",
|
||||
"replacement_group": "pro-tier"
|
||||
}]))
|
||||
.is_ok());
|
||||
assert_eq!(
|
||||
validate_entitlement_replacement_groups(&json!([{
|
||||
"type": "usage_policy",
|
||||
"replacement_group": " "
|
||||
}])),
|
||||
Err(EntitlementReplacementGroupValidationError::Empty { index: 0 })
|
||||
);
|
||||
assert_eq!(
|
||||
validate_entitlement_replacement_groups(&json!([{
|
||||
"type": "usage_policy",
|
||||
"replacement_group": 42
|
||||
}])),
|
||||
Err(EntitlementReplacementGroupValidationError::InvalidType { index: 0 })
|
||||
);
|
||||
assert!(matches!(
|
||||
validate_entitlement_replacement_groups(&json!([{
|
||||
"type": "usage_policy",
|
||||
"replacement_group": "x".repeat(MAX_ENTITLEMENT_REPLACEMENT_GROUP_LENGTH + 1)
|
||||
}])),
|
||||
Err(EntitlementReplacementGroupValidationError::TooLong { index: 0, .. })
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -150,7 +150,7 @@ pub enum AdminBillingMutationOutcome<T> {
|
||||
Unavailable,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct PaymentGatewayConfigRecord {
|
||||
pub provider: String,
|
||||
pub enabled: bool,
|
||||
@@ -166,7 +166,30 @@ pub struct PaymentGatewayConfigRecord {
|
||||
pub updated_at_unix_secs: u64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for PaymentGatewayConfigRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("PaymentGatewayConfigRecord")
|
||||
.field("provider", &self.provider)
|
||||
.field("enabled", &self.enabled)
|
||||
.field("endpoint_url", &"[REDACTED]")
|
||||
.field(
|
||||
"callback_base_url",
|
||||
&self.callback_base_url.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("merchant_id", &self.merchant_id)
|
||||
.field(
|
||||
"merchant_key_encrypted",
|
||||
&self.merchant_key_encrypted.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("channels_json", &"[REDACTED]")
|
||||
.field("created_at_unix_secs", &self.created_at_unix_secs)
|
||||
.field("updated_at_unix_secs", &self.updated_at_unix_secs)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct PaymentGatewayConfigWriteInput {
|
||||
pub provider: String,
|
||||
pub enabled: bool,
|
||||
@@ -181,6 +204,113 @@ pub struct PaymentGatewayConfigWriteInput {
|
||||
pub channels_json: Value,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for PaymentGatewayConfigWriteInput {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("PaymentGatewayConfigWriteInput")
|
||||
.field("provider", &self.provider)
|
||||
.field("enabled", &self.enabled)
|
||||
.field("endpoint_url", &"[REDACTED]")
|
||||
.field(
|
||||
"callback_base_url",
|
||||
&self.callback_base_url.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("merchant_id", &self.merchant_id)
|
||||
.field(
|
||||
"merchant_key_encrypted",
|
||||
&self.merchant_key_encrypted.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("preserve_existing_secret", &self.preserve_existing_secret)
|
||||
.field("channels_json", &"[REDACTED]")
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct PaymentGatewaySecretCasUpdate {
|
||||
pub provider: String,
|
||||
pub expected_merchant_key_encrypted: String,
|
||||
pub merchant_key_encrypted: String,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for PaymentGatewaySecretCasUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("PaymentGatewaySecretCasUpdate")
|
||||
.field("provider", &self.provider)
|
||||
.field("expected_merchant_key_encrypted", &"[REDACTED]")
|
||||
.field("merchant_key_encrypted", &"[REDACTED]")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct PaymentGatewayConfigCasWriteInput {
|
||||
pub input: PaymentGatewayConfigWriteInput,
|
||||
pub expected_existing: bool,
|
||||
pub expected_merchant_key_encrypted: Option<String>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for PaymentGatewayConfigCasWriteInput {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("PaymentGatewayConfigCasWriteInput")
|
||||
.field("input", &self.input)
|
||||
.field("expected_existing", &self.expected_existing)
|
||||
.field(
|
||||
"expected_merchant_key_encrypted",
|
||||
&self
|
||||
.expected_merchant_key_encrypted
|
||||
.as_ref()
|
||||
.map(|_| "[REDACTED]"),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod payment_gateway_debug_tests {
|
||||
use super::{PaymentGatewayConfigCasWriteInput, PaymentGatewayConfigWriteInput};
|
||||
|
||||
#[test]
|
||||
fn payment_gateway_config_debug_output_redacts_credential_material() {
|
||||
let input = PaymentGatewayConfigCasWriteInput {
|
||||
input: PaymentGatewayConfigWriteInput {
|
||||
provider: "stripe".to_string(),
|
||||
enabled: true,
|
||||
endpoint_url: "https://endpoint.example/?key=endpoint-canary".to_string(),
|
||||
callback_base_url: Some(
|
||||
"https://callback.example/?token=callback-canary".to_string(),
|
||||
),
|
||||
merchant_id: "merchant".to_string(),
|
||||
merchant_key_encrypted: Some("merchant-key-canary".to_string()),
|
||||
preserve_existing_secret: false,
|
||||
pay_currency: "USD".to_string(),
|
||||
usd_exchange_rate: 1.0,
|
||||
min_recharge_usd: 1.0,
|
||||
channels_json: serde_json::json!({"secret": "channels-canary"}),
|
||||
},
|
||||
expected_existing: true,
|
||||
expected_merchant_key_encrypted: Some("expected-merchant-key-canary".to_string()),
|
||||
};
|
||||
|
||||
let debug = format!("{input:?}");
|
||||
assert!(debug.contains("[REDACTED]"));
|
||||
for secret in [
|
||||
"endpoint-canary",
|
||||
"callback-canary",
|
||||
"merchant-key-canary",
|
||||
"channels-canary",
|
||||
"expected-merchant-key-canary",
|
||||
] {
|
||||
assert!(
|
||||
!debug.contains(secret),
|
||||
"debug output leaked {secret}: {debug}"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct BillingPlanRecord {
|
||||
pub id: String,
|
||||
@@ -214,6 +344,43 @@ pub struct BillingPlanWriteInput {
|
||||
pub entitlements_json: Value,
|
||||
}
|
||||
|
||||
/// Convert a plan duration into whole days without allowing integer or
|
||||
/// `chrono::TimeDelta` overflow. Plan snapshots are persisted and may later be
|
||||
/// fulfilled by any database adapter, so the accepted range must be portable
|
||||
/// across all of them.
|
||||
pub fn checked_plan_duration_days(duration_unit: &str, duration_value: i64) -> Result<i64, String> {
|
||||
if duration_value <= 0 {
|
||||
return Err("plan duration_value must be positive".to_string());
|
||||
}
|
||||
let days = match duration_unit.trim() {
|
||||
"day" | "custom" => Some(duration_value),
|
||||
"month" => duration_value.checked_mul(30),
|
||||
"year" => duration_value.checked_mul(365),
|
||||
_ => return Err("plan duration_unit is invalid".to_string()),
|
||||
}
|
||||
.ok_or_else(|| "plan duration exceeds the supported range".to_string())?;
|
||||
chrono::TimeDelta::try_days(days)
|
||||
.ok_or_else(|| "plan duration exceeds the supported range".to_string())?;
|
||||
Ok(days)
|
||||
}
|
||||
|
||||
/// Read a persisted plan snapshot using the historical month/one defaults,
|
||||
/// while rejecting malformed or unrepresentable explicit values.
|
||||
pub fn checked_plan_duration_days_from_snapshot(snapshot: &Value) -> Result<i64, String> {
|
||||
let duration_unit = match snapshot.get("duration_unit") {
|
||||
None => "month",
|
||||
Some(Value::String(value)) => value.as_str(),
|
||||
Some(_) => return Err("product_snapshot.duration_unit is invalid".to_string()),
|
||||
};
|
||||
let duration_value = match snapshot.get("duration_value") {
|
||||
None => 1,
|
||||
Some(value) => value
|
||||
.as_i64()
|
||||
.ok_or_else(|| "product_snapshot.duration_value must be an integer".to_string())?,
|
||||
};
|
||||
checked_plan_duration_days(duration_unit, duration_value)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct UserPlanEntitlementRecord {
|
||||
pub id: String,
|
||||
@@ -369,6 +536,37 @@ pub trait BillingReadRepository: Send + Sync {
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
/// Re-read a gateway configuration from the authoritative backing store.
|
||||
/// Implementations without a read cache may delegate to the normal read.
|
||||
async fn find_payment_gateway_config_strong(
|
||||
&self,
|
||||
provider: &str,
|
||||
) -> Result<Option<PaymentGatewayConfigRecord>, crate::DataLayerError> {
|
||||
self.find_payment_gateway_config(provider).await
|
||||
}
|
||||
|
||||
/// Replace only the encrypted merchant secret when the exact previously
|
||||
/// observed ciphertext is still stored. Timestamps and all other fields
|
||||
/// must remain unchanged.
|
||||
async fn compare_and_swap_payment_gateway_secret(
|
||||
&self,
|
||||
update: &PaymentGatewaySecretCasUpdate,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
let _ = update;
|
||||
Ok(false)
|
||||
}
|
||||
|
||||
/// Create a configuration only when absent, or update it only when the
|
||||
/// exact nullable merchant-secret fence still matches.
|
||||
async fn compare_and_swap_payment_gateway_config(
|
||||
&self,
|
||||
input: &PaymentGatewayConfigCasWriteInput,
|
||||
) -> Result<AdminBillingMutationOutcome<PaymentGatewayConfigRecord>, crate::DataLayerError>
|
||||
{
|
||||
let _ = input;
|
||||
Ok(AdminBillingMutationOutcome::Unavailable)
|
||||
}
|
||||
|
||||
async fn upsert_payment_gateway_config(
|
||||
&self,
|
||||
input: &PaymentGatewayConfigWriteInput,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -2,9 +2,12 @@ mod types;
|
||||
|
||||
pub use types::{
|
||||
build_decision_trace, derive_request_candidate_final_status,
|
||||
request_candidate_lifecycle_would_regress, DecisionTrace, DecisionTraceCandidate,
|
||||
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateFinalStatus,
|
||||
RequestCandidateReadRepository, RequestCandidateRepository, RequestCandidateStatus,
|
||||
RequestCandidateTrace, RequestCandidateWriteRepository, StoredRequestCandidate,
|
||||
UpsertRequestCandidateRecord,
|
||||
request_candidate_lifecycle_would_regress, sanitize_request_candidate_api_formats,
|
||||
sanitize_request_candidate_error_type, sanitize_request_candidate_extra_data,
|
||||
sanitize_request_candidate_required_capabilities, sanitize_request_candidate_skip_reason,
|
||||
DecisionTrace, DecisionTraceCandidate, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
||||
RequestCandidateFinalStatus, RequestCandidateReadRepository, RequestCandidateRepository,
|
||||
RequestCandidateStatus, RequestCandidateTrace, RequestCandidateWriteRepository,
|
||||
StoredRequestCandidate, UpsertRequestCandidateRecord, REQUEST_CANDIDATE_ERROR_TYPES,
|
||||
REQUEST_CANDIDATE_ERROR_TYPE_ALIASES, REQUEST_CANDIDATE_SKIP_REASONS,
|
||||
};
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,7 +1,12 @@
|
||||
use async_trait::async_trait;
|
||||
|
||||
pub const GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS: usize = 512;
|
||||
pub const GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS: usize = 512;
|
||||
pub const GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS: usize = 255;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Default)]
|
||||
pub struct GeminiFileMappingListQuery {
|
||||
pub user_id: Option<String>,
|
||||
pub include_expired: bool,
|
||||
pub search: Option<String>,
|
||||
pub offset: usize,
|
||||
@@ -103,6 +108,25 @@ impl UpsertGeminiFileMappingRecord {
|
||||
"gemini_file_mappings.file_name is empty".to_string(),
|
||||
));
|
||||
}
|
||||
validate_text_length(
|
||||
&self.file_name,
|
||||
"file_name",
|
||||
GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS,
|
||||
)?;
|
||||
if let Some(display_name) = self.display_name.as_deref() {
|
||||
validate_text_length(
|
||||
display_name,
|
||||
"display_name",
|
||||
GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS,
|
||||
)?;
|
||||
}
|
||||
if let Some(mime_type) = self.mime_type.as_deref() {
|
||||
validate_text_length(
|
||||
mime_type,
|
||||
"mime_type",
|
||||
GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS,
|
||||
)?;
|
||||
}
|
||||
if self.key_id.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"gemini_file_mappings.key_id is empty".to_string(),
|
||||
@@ -117,6 +141,19 @@ impl UpsertGeminiFileMappingRecord {
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_text_length(
|
||||
value: &str,
|
||||
field: &str,
|
||||
max_chars: usize,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
if value.chars().nth(max_chars).is_some() {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"gemini_file_mappings.{field} exceeds maximum length {max_chars}"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait GeminiFileMappingReadRepository: Send + Sync {
|
||||
async fn find_by_file_name(
|
||||
@@ -124,6 +161,29 @@ pub trait GeminiFileMappingReadRepository: Send + Sync {
|
||||
file_name: &str,
|
||||
) -> Result<Option<StoredGeminiFileMapping>, crate::DataLayerError>;
|
||||
|
||||
/// Return an unexpired mapping only when it belongs to `user_id`.
|
||||
///
|
||||
/// Repository implementations must apply the file name, user and expiry
|
||||
/// predicates in the same read operation. Public callers must not emulate
|
||||
/// this with an unrestricted lookup followed by an ownership check.
|
||||
async fn find_active_by_file_name_for_user(
|
||||
&self,
|
||||
file_name: &str,
|
||||
user_id: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<Option<StoredGeminiFileMapping>, crate::DataLayerError>;
|
||||
|
||||
/// Return an unexpired mapping only when both user and provider key match.
|
||||
/// This is the routing lookup used before forwarding file object requests
|
||||
/// to an upstream provider credential.
|
||||
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>, crate::DataLayerError>;
|
||||
|
||||
async fn list_mappings(
|
||||
&self,
|
||||
query: &GeminiFileMappingListQuery,
|
||||
@@ -141,8 +201,32 @@ pub trait GeminiFileMappingWriteRepository: Send + Sync {
|
||||
&self,
|
||||
record: UpsertGeminiFileMappingRecord,
|
||||
) -> Result<StoredGeminiFileMapping, crate::DataLayerError>;
|
||||
|
||||
/// Insert a new mapping or refresh it only when the persisted owner is the
|
||||
/// same provider key and user. The ownership check and write must be one
|
||||
/// atomic repository operation so callers cannot be bypassed with a
|
||||
/// check-then-write race.
|
||||
async fn upsert_if_owner_matches(
|
||||
&self,
|
||||
record: UpsertGeminiFileMappingRecord,
|
||||
) -> Result<Option<StoredGeminiFileMapping>, crate::DataLayerError>;
|
||||
|
||||
async fn delete_by_file_name(&self, file_name: &str) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn delete_by_file_name_for_user(
|
||||
&self,
|
||||
file_name: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
/// Delete only when both persisted ownership dimensions still match.
|
||||
async fn delete_by_file_name_for_owner(
|
||||
&self,
|
||||
file_name: &str,
|
||||
key_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn delete_by_id(
|
||||
&self,
|
||||
mapping_id: &str,
|
||||
@@ -163,3 +247,64 @@ impl<T> GeminiFileMappingRepository for T where
|
||||
T: GeminiFileMappingReadRepository + GeminiFileMappingWriteRepository
|
||||
{
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
UpsertGeminiFileMappingRecord, GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS,
|
||||
GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS, GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
|
||||
fn record() -> UpsertGeminiFileMappingRecord {
|
||||
UpsertGeminiFileMappingRecord {
|
||||
id: "mapping-1".to_string(),
|
||||
file_name: "files/example".to_string(),
|
||||
key_id: "key-1".to_string(),
|
||||
user_id: Some("user-1".to_string()),
|
||||
display_name: Some("example".to_string()),
|
||||
mime_type: Some("application/octet-stream".to_string()),
|
||||
source_hash: None,
|
||||
expires_at_unix_secs: 1,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mapping_metadata_accepts_schema_limits() {
|
||||
let mut record = record();
|
||||
record.file_name = "f".repeat(GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS);
|
||||
record.display_name = Some("d".repeat(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS));
|
||||
record.mime_type = Some("m".repeat(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS));
|
||||
|
||||
record.validate().expect("schema limits should validate");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mapping_metadata_rejects_values_beyond_schema_limits() {
|
||||
for (field, record) in [
|
||||
("file_name", {
|
||||
let mut record = record();
|
||||
record.file_name = "f".repeat(GEMINI_FILE_MAPPING_MAX_FILE_NAME_CHARS + 1);
|
||||
record
|
||||
}),
|
||||
("display_name", {
|
||||
let mut record = record();
|
||||
record.display_name =
|
||||
Some("d".repeat(GEMINI_FILE_MAPPING_MAX_DISPLAY_NAME_CHARS + 1));
|
||||
record
|
||||
}),
|
||||
("mime_type", {
|
||||
let mut record = record();
|
||||
record.mime_type = Some("m".repeat(GEMINI_FILE_MAPPING_MAX_MIME_TYPE_CHARS + 1));
|
||||
record
|
||||
}),
|
||||
] {
|
||||
let error = record
|
||||
.validate()
|
||||
.expect_err("oversized mapping metadata should fail");
|
||||
assert!(
|
||||
matches!(error, DataLayerError::InvalidInput(message) if message.contains(field))
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -163,7 +163,7 @@ pub struct ManagementTokenListQuery {
|
||||
pub limit: usize,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct CreateManagementTokenRecord {
|
||||
pub id: String,
|
||||
pub user_id: String,
|
||||
@@ -178,6 +178,46 @@ pub struct CreateManagementTokenRecord {
|
||||
pub is_active: bool,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for CreateManagementTokenRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("CreateManagementTokenRecord")
|
||||
.field("id", &self.id)
|
||||
.field("user_id", &self.user_id)
|
||||
.field("token_hash", &"[REDACTED]")
|
||||
.field(
|
||||
"token_prefix",
|
||||
&self.token_prefix.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("name", &self.name)
|
||||
.field("expires_at_unix_secs", &self.expires_at_unix_secs)
|
||||
.field("is_active", &self.is_active)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_management_token_unix_secs_storage_range(
|
||||
value: u64,
|
||||
field_name: &str,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
if i64::try_from(value).is_err() {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"{field_name} exceeds supported storage range"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_optional_management_token_unix_secs_storage_range(
|
||||
value: Option<u64>,
|
||||
field_name: &str,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
match value {
|
||||
Some(value) => validate_management_token_unix_secs_storage_range(value, field_name),
|
||||
None => Ok(()),
|
||||
}
|
||||
}
|
||||
|
||||
impl CreateManagementTokenRecord {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
if self.id.trim().is_empty() {
|
||||
@@ -223,6 +263,10 @@ impl CreateManagementTokenRecord {
|
||||
}
|
||||
}
|
||||
validate_management_token_permissions(self.permissions.as_ref())?;
|
||||
validate_optional_management_token_unix_secs_storage_range(
|
||||
self.expires_at_unix_secs,
|
||||
"expires_at_unix_secs",
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -248,6 +292,21 @@ impl UpdateManagementTokenRecord {
|
||||
"token_id is required".to_string(),
|
||||
));
|
||||
}
|
||||
if self.clear_description && self.description.is_some() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"description and clear_description are mutually exclusive".to_string(),
|
||||
));
|
||||
}
|
||||
if self.clear_allowed_ips && self.allowed_ips.is_some() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"allowed_ips and clear_allowed_ips are mutually exclusive".to_string(),
|
||||
));
|
||||
}
|
||||
if self.clear_expires_at && self.expires_at_unix_secs.is_some() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"expires_at_unix_secs and clear_expires_at are mutually exclusive".to_string(),
|
||||
));
|
||||
}
|
||||
if let Some(name) = &self.name {
|
||||
if name.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
@@ -273,6 +332,10 @@ impl UpdateManagementTokenRecord {
|
||||
}
|
||||
}
|
||||
validate_management_token_permissions(self.permissions.as_ref())?;
|
||||
validate_optional_management_token_unix_secs_storage_range(
|
||||
self.expires_at_unix_secs,
|
||||
"expires_at_unix_secs",
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -301,13 +364,127 @@ fn validate_management_token_permissions(
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct RegenerateManagementTokenSecret {
|
||||
pub token_id: String,
|
||||
pub token_hash: String,
|
||||
pub token_prefix: Option<String>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for RegenerateManagementTokenSecret {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("RegenerateManagementTokenSecret")
|
||||
.field("token_id", &self.token_id)
|
||||
.field("token_hash", &"[REDACTED]")
|
||||
.field(
|
||||
"token_prefix",
|
||||
&self.token_prefix.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ActivateManagementTokenIfMatches {
|
||||
pub expected_token: StoredManagementToken,
|
||||
pub token_hash: String,
|
||||
pub expected_user_security_version: i64,
|
||||
pub now_unix_secs: u64,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ActivateManagementTokenIfMatches {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ActivateManagementTokenIfMatches")
|
||||
.field("expected_token", &self.expected_token)
|
||||
.field("token_hash", &"[REDACTED]")
|
||||
.field(
|
||||
"expected_user_security_version",
|
||||
&self.expected_user_security_version,
|
||||
)
|
||||
.field("now_unix_secs", &self.now_unix_secs)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl ActivateManagementTokenIfMatches {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
if self.expected_token.id.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"token_id is required".to_string(),
|
||||
));
|
||||
}
|
||||
if self.expected_token.user_id.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"user_id is required".to_string(),
|
||||
));
|
||||
}
|
||||
if self.expected_token.name.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"token name is required".to_string(),
|
||||
));
|
||||
}
|
||||
if self.expected_token.is_active {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"activation snapshot must be inactive".to_string(),
|
||||
));
|
||||
}
|
||||
if self.token_hash.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"token_hash is required".to_string(),
|
||||
));
|
||||
}
|
||||
if self.expected_user_security_version < 0 {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"expected_user_security_version must not be negative".to_string(),
|
||||
));
|
||||
}
|
||||
validate_management_token_unix_secs_storage_range(self.now_unix_secs, "now_unix_secs")?;
|
||||
validate_optional_management_token_unix_secs_storage_range(
|
||||
self.expected_token.expires_at_unix_secs,
|
||||
"expires_at_unix_secs",
|
||||
)?;
|
||||
validate_optional_management_token_unix_secs_storage_range(
|
||||
self.expected_token.last_used_at_unix_secs,
|
||||
"last_used_at_unix_secs",
|
||||
)?;
|
||||
validate_optional_management_token_unix_secs_storage_range(
|
||||
self.expected_token.created_at_unix_ms,
|
||||
"created_at_unix_ms",
|
||||
)?;
|
||||
validate_optional_management_token_unix_secs_storage_range(
|
||||
self.expected_token.updated_at_unix_secs,
|
||||
"updated_at_unix_secs",
|
||||
)?;
|
||||
validate_management_token_unix_secs_storage_range(
|
||||
self.expected_token.usage_count,
|
||||
"usage_count",
|
||||
)?;
|
||||
if let Some(allowed_ips) = &self.expected_token.allowed_ips {
|
||||
let Some(items) = allowed_ips.as_array() else {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"IP 限制规则必须是数组".to_string(),
|
||||
));
|
||||
};
|
||||
if items.is_empty() || items.iter().any(|value| value.as_str().is_none()) {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"IP 限制规则只能是非空字符串数组".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
validate_management_token_permissions(self.expected_token.permissions.as_ref())
|
||||
}
|
||||
|
||||
pub fn matches_locked_token_snapshot(
|
||||
&self,
|
||||
current: &StoredManagementToken,
|
||||
current_token_hash: &str,
|
||||
) -> bool {
|
||||
current_token_hash == self.token_hash && current == &self.expected_token
|
||||
}
|
||||
}
|
||||
|
||||
impl RegenerateManagementTokenSecret {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
if self.token_id.trim().is_empty() {
|
||||
@@ -360,22 +537,277 @@ pub trait ManagementTokenWriteRepository: Send + Sync {
|
||||
record: &UpdateManagementTokenRecord,
|
||||
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
|
||||
|
||||
/// Update a token only while it still belongs to `user_id`.
|
||||
///
|
||||
/// Self-service callers must use this owner-scoped mutation instead of
|
||||
/// relying on a preceding read-side ownership check.
|
||||
async fn update_management_token_for_user(
|
||||
&self,
|
||||
record: &UpdateManagementTokenRecord,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
|
||||
|
||||
async fn delete_management_token(&self, token_id: &str) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn delete_management_token_for_user(
|
||||
&self,
|
||||
token_id: &str,
|
||||
user_id: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn set_management_token_active(
|
||||
&self,
|
||||
token_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
|
||||
|
||||
async fn set_management_token_active_for_user(
|
||||
&self,
|
||||
token_id: &str,
|
||||
user_id: &str,
|
||||
is_active: bool,
|
||||
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
|
||||
|
||||
/// Atomically activate an inactive one-time-install token only while all
|
||||
/// security-relevant fields still match the session that issued it.
|
||||
async fn activate_management_token_if_matches(
|
||||
&self,
|
||||
mutation: &ActivateManagementTokenIfMatches,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
/// Delete an unclaimed one-time-install token only while it still matches
|
||||
/// the session that created it and remains inactive.
|
||||
async fn delete_inactive_management_token_if_matches(
|
||||
&self,
|
||||
mutation: &ActivateManagementTokenIfMatches,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn regenerate_management_token_secret(
|
||||
&self,
|
||||
mutation: &RegenerateManagementTokenSecret,
|
||||
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
|
||||
|
||||
async fn regenerate_management_token_secret_for_user(
|
||||
&self,
|
||||
mutation: &RegenerateManagementTokenSecret,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
|
||||
|
||||
async fn record_management_token_usage(
|
||||
&self,
|
||||
token_id: &str,
|
||||
last_used_ip: Option<&str>,
|
||||
) -> Result<Option<StoredManagementToken>, crate::DataLayerError>;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
ActivateManagementTokenIfMatches, CreateManagementTokenRecord,
|
||||
RegenerateManagementTokenSecret, StoredManagementTokenUserSummary,
|
||||
UpdateManagementTokenRecord,
|
||||
};
|
||||
|
||||
fn activation() -> ActivateManagementTokenIfMatches {
|
||||
ActivateManagementTokenIfMatches {
|
||||
expected_token: super::StoredManagementToken::new(
|
||||
"token-1".to_string(),
|
||||
"user-1".to_string(),
|
||||
"install token".to_string(),
|
||||
)
|
||||
.expect("token should build")
|
||||
.with_display_fields(None, Some("ae_install".to_string()), None)
|
||||
.with_permissions(Some(serde_json::json!(["admin:proxy_nodes:write"])))
|
||||
.with_runtime_fields(Some(2_000_000_000), None, None, 0, false)
|
||||
.with_timestamps(Some(1_800_000_000), Some(1_800_000_000)),
|
||||
token_hash: "hash-1".to_string(),
|
||||
expected_user_security_version: 7,
|
||||
now_unix_secs: 1_900_000_000,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn activation_rejects_timestamps_outside_sql_storage_range() {
|
||||
let mut mutation = activation();
|
||||
mutation.expected_token.expires_at_unix_secs = Some(i64::MAX as u64 + 1);
|
||||
assert!(mutation.validate().is_err());
|
||||
|
||||
let mut mutation = activation();
|
||||
mutation.now_unix_secs = i64::MAX as u64 + 1;
|
||||
assert!(mutation.validate().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn activation_requires_non_negative_user_security_version() {
|
||||
let mut mutation = activation();
|
||||
mutation.expected_user_security_version = -1;
|
||||
assert!(mutation.validate().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn activation_matches_the_complete_canonical_token_snapshot() {
|
||||
let mutation = activation();
|
||||
assert!(
|
||||
mutation.matches_locked_token_snapshot(&mutation.expected_token, &mutation.token_hash,)
|
||||
);
|
||||
|
||||
let mut changed = mutation.expected_token.clone();
|
||||
changed.description = Some("changed after session creation".to_string());
|
||||
assert!(!mutation.matches_locked_token_snapshot(&changed, &mutation.token_hash));
|
||||
|
||||
let mut changed = mutation.expected_token.clone();
|
||||
changed.updated_at_unix_secs = changed.updated_at_unix_secs.map(|value| value + 1);
|
||||
assert!(!mutation.matches_locked_token_snapshot(&changed, &mutation.token_hash));
|
||||
|
||||
assert!(
|
||||
!mutation.matches_locked_token_snapshot(&mutation.expected_token, "different-hash",)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn management_token_secret_debug_output_is_redacted() {
|
||||
let activation = ActivateManagementTokenIfMatches {
|
||||
token_hash: "activation-token-hash-canary".to_string(),
|
||||
..activation()
|
||||
};
|
||||
let activation_debug = format!("{activation:?}");
|
||||
assert!(activation_debug.contains("[REDACTED]"));
|
||||
assert!(!activation_debug.contains("activation-token-hash-canary"));
|
||||
|
||||
let regenerate = RegenerateManagementTokenSecret {
|
||||
token_id: "token-1".to_string(),
|
||||
token_hash: "regenerate-token-hash-canary".to_string(),
|
||||
token_prefix: Some("regenerate-token-prefix-canary".to_string()),
|
||||
};
|
||||
let regenerate_debug = format!("{regenerate:?}");
|
||||
assert!(regenerate_debug.contains("[REDACTED]"));
|
||||
assert!(!regenerate_debug.contains("regenerate-token-hash-canary"));
|
||||
assert!(!regenerate_debug.contains("regenerate-token-prefix-canary"));
|
||||
|
||||
let user = StoredManagementTokenUserSummary::new(
|
||||
"user-1".to_string(),
|
||||
None,
|
||||
"admin".to_string(),
|
||||
"admin".to_string(),
|
||||
)
|
||||
.expect("user summary should build");
|
||||
let create = CreateManagementTokenRecord {
|
||||
id: "token-1".to_string(),
|
||||
user_id: user.id.clone(),
|
||||
user,
|
||||
token_hash: "create-token-hash-canary".to_string(),
|
||||
token_prefix: Some("create-token-prefix-canary".to_string()),
|
||||
name: "token".to_string(),
|
||||
description: None,
|
||||
allowed_ips: None,
|
||||
permissions: None,
|
||||
expires_at_unix_secs: None,
|
||||
is_active: true,
|
||||
};
|
||||
let create_debug = format!("{create:?}");
|
||||
assert!(create_debug.contains("[REDACTED]"));
|
||||
assert!(!create_debug.contains("create-token-hash-canary"));
|
||||
assert!(!create_debug.contains("create-token-prefix-canary"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn activation_rejects_json_null_without_conflating_it_with_sql_null() {
|
||||
let mutation = activation();
|
||||
|
||||
let mut json_null_allowed_ips = mutation.expected_token.clone();
|
||||
json_null_allowed_ips.allowed_ips = Some(serde_json::Value::Null);
|
||||
assert!(
|
||||
!mutation.matches_locked_token_snapshot(&json_null_allowed_ips, &mutation.token_hash)
|
||||
);
|
||||
let mut invalid = mutation.clone();
|
||||
invalid.expected_token = json_null_allowed_ips;
|
||||
assert!(invalid.validate().is_err());
|
||||
|
||||
let mut json_null_permissions = mutation.expected_token.clone();
|
||||
json_null_permissions.permissions = Some(serde_json::Value::Null);
|
||||
assert!(
|
||||
!mutation.matches_locked_token_snapshot(&json_null_permissions, &mutation.token_hash)
|
||||
);
|
||||
let mut invalid = mutation;
|
||||
invalid.expected_token = json_null_permissions;
|
||||
assert!(invalid.validate().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn create_and_update_reject_expiry_outside_sql_storage_range() {
|
||||
let unsupported_expiry = i64::MAX as u64 + 1;
|
||||
let user = StoredManagementTokenUserSummary::new(
|
||||
"user-1".to_string(),
|
||||
None,
|
||||
"admin".to_string(),
|
||||
"admin".to_string(),
|
||||
)
|
||||
.expect("user summary should build");
|
||||
let create = CreateManagementTokenRecord {
|
||||
id: "token-1".to_string(),
|
||||
user_id: user.id.clone(),
|
||||
user,
|
||||
token_hash: "hash-1".to_string(),
|
||||
token_prefix: Some("ae_1234".to_string()),
|
||||
name: "token".to_string(),
|
||||
description: None,
|
||||
allowed_ips: None,
|
||||
permissions: Some(serde_json::json!(["admin:proxy_nodes:write"])),
|
||||
expires_at_unix_secs: Some(unsupported_expiry),
|
||||
is_active: false,
|
||||
};
|
||||
assert!(create.validate().is_err());
|
||||
|
||||
let update = UpdateManagementTokenRecord {
|
||||
token_id: "token-1".to_string(),
|
||||
name: None,
|
||||
description: None,
|
||||
clear_description: false,
|
||||
allowed_ips: None,
|
||||
clear_allowed_ips: false,
|
||||
permissions: None,
|
||||
expires_at_unix_secs: Some(unsupported_expiry),
|
||||
clear_expires_at: false,
|
||||
is_active: None,
|
||||
};
|
||||
assert!(update.validate().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn management_token_update_rejects_ambiguous_set_and_clear_operations() {
|
||||
let base = UpdateManagementTokenRecord {
|
||||
token_id: "token-1".to_string(),
|
||||
name: None,
|
||||
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!(UpdateManagementTokenRecord {
|
||||
description: Some("description".to_string()),
|
||||
clear_description: true,
|
||||
..base.clone()
|
||||
}
|
||||
.validate()
|
||||
.is_err());
|
||||
assert!(UpdateManagementTokenRecord {
|
||||
allowed_ips: Some(serde_json::json!(["127.0.0.1"])),
|
||||
clear_allowed_ips: true,
|
||||
..base.clone()
|
||||
}
|
||||
.validate()
|
||||
.is_err());
|
||||
assert!(UpdateManagementTokenRecord {
|
||||
expires_at_unix_secs: Some(100),
|
||||
clear_expires_at: true,
|
||||
..base
|
||||
}
|
||||
.validate()
|
||||
.is_err());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,213 @@
|
||||
use async_trait::async_trait;
|
||||
use std::net::IpAddr;
|
||||
use url::{Host, Url};
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub fn validate_oauth_redirect_uri(value: &str) -> Result<(), String> {
|
||||
let parsed =
|
||||
Url::parse(value).map_err(|_| "redirect_uri must be an absolute URL".to_string())?;
|
||||
let Some(host) = parsed.host() else {
|
||||
return Err("redirect_uri must be an absolute URL".to_string());
|
||||
};
|
||||
let is_loopback = match host {
|
||||
Host::Domain(domain) => domain.eq_ignore_ascii_case("localhost"),
|
||||
Host::Ipv4(address) => address.is_loopback(),
|
||||
Host::Ipv6(address) => address.is_loopback(),
|
||||
};
|
||||
if parsed.scheme() != "https" && !(parsed.scheme() == "http" && is_loopback) {
|
||||
return Err(
|
||||
"redirect_uri must use https, except for localhost or loopback IPs".to_string(),
|
||||
);
|
||||
}
|
||||
if !parsed.username().is_empty() || parsed.password().is_some() {
|
||||
return Err("redirect_uri must not contain URL credentials".to_string());
|
||||
}
|
||||
if parsed.fragment().is_some() {
|
||||
return Err("redirect_uri must not contain a fragment".to_string());
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn validate_oauth_frontend_callback_url(value: &str) -> Result<(), String> {
|
||||
let parsed = Url::parse(value)
|
||||
.map_err(|_| "frontend_callback_url must be an absolute URL".to_string())?;
|
||||
let Some(host) = parsed.host() else {
|
||||
return Err("frontend_callback_url must be an absolute URL".to_string());
|
||||
};
|
||||
let is_loopback = match host {
|
||||
Host::Domain(domain) => domain.eq_ignore_ascii_case("localhost"),
|
||||
Host::Ipv4(address) => address.is_loopback(),
|
||||
Host::Ipv6(address) => address.is_loopback(),
|
||||
};
|
||||
if parsed.scheme() != "https" && !(parsed.scheme() == "http" && is_loopback) {
|
||||
return Err(
|
||||
"frontend_callback_url must use https, except for localhost or loopback IPs"
|
||||
.to_string(),
|
||||
);
|
||||
}
|
||||
if !parsed.username().is_empty() || parsed.password().is_some() {
|
||||
return Err("frontend_callback_url must not contain URL credentials".to_string());
|
||||
}
|
||||
if parsed.query().is_some() || parsed.fragment().is_some() {
|
||||
return Err("frontend_callback_url must not contain a query or fragment".to_string());
|
||||
}
|
||||
if !parsed
|
||||
.path()
|
||||
.trim_end_matches('/')
|
||||
.ends_with("/auth/callback")
|
||||
{
|
||||
return Err("frontend_callback_url path must end with /auth/callback".to_string());
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn validate_oauth_provider_endpoint_config(
|
||||
provider_type: &str,
|
||||
authorization_url_override: Option<&str>,
|
||||
token_url_override: Option<&str>,
|
||||
userinfo_url_override: Option<&str>,
|
||||
extra_config: Option<&serde_json::Value>,
|
||||
) -> Result<(), String> {
|
||||
let provider_type = provider_type.trim().to_ascii_lowercase();
|
||||
let mut provider_chars = provider_type.chars();
|
||||
if !(3..=64).contains(&provider_type.len())
|
||||
|| !provider_chars
|
||||
.next()
|
||||
.is_some_and(|character| character.is_ascii_lowercase())
|
||||
|| !provider_chars.all(|character| {
|
||||
character.is_ascii_lowercase()
|
||||
|| character.is_ascii_digit()
|
||||
|| matches!(character, '_' | '-')
|
||||
})
|
||||
{
|
||||
return Err("provider_type contains invalid characters".to_string());
|
||||
}
|
||||
let is_custom = provider_type == "custom_oidc"
|
||||
|| provider_type.starts_with("custom_oidc_")
|
||||
|| provider_type.starts_with("custom_")
|
||||
|| provider_type.starts_with("oidc_");
|
||||
let allowed_domains = if provider_type == "linuxdo" {
|
||||
vec![
|
||||
"linux.do".to_string(),
|
||||
"connect.linux.do".to_string(),
|
||||
"connect.linuxdo.org".to_string(),
|
||||
]
|
||||
} else if is_custom {
|
||||
oauth_custom_allowed_domains(extra_config)?
|
||||
} else {
|
||||
return Err("unsupported identity OAuth provider_type".to_string());
|
||||
};
|
||||
|
||||
for (field, value) in [
|
||||
("authorization_url_override", authorization_url_override),
|
||||
("token_url_override", token_url_override),
|
||||
("userinfo_url_override", userinfo_url_override),
|
||||
] {
|
||||
let value = value.map(str::trim).filter(|value| !value.is_empty());
|
||||
if is_custom && value.is_none() {
|
||||
return Err(format!("custom OIDC providers must configure {field}"));
|
||||
}
|
||||
if let Some(value) = value {
|
||||
validate_oauth_endpoint_url(field, value, &allowed_domains)?;
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn oauth_custom_allowed_domains(
|
||||
extra_config: Option<&serde_json::Value>,
|
||||
) -> Result<Vec<String>, String> {
|
||||
let values = extra_config
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|object| {
|
||||
object
|
||||
.get("allowed_domains")
|
||||
.or_else(|| object.get("oauth_allowed_domains"))
|
||||
})
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.ok_or_else(|| {
|
||||
"custom OIDC providers must configure extra_config.allowed_domains".to_string()
|
||||
})?;
|
||||
let mut domains = Vec::with_capacity(values.len());
|
||||
for value in values {
|
||||
let domain = value
|
||||
.as_str()
|
||||
.map(str::trim)
|
||||
.map(|value| value.trim_end_matches('.'))
|
||||
.filter(|value| !value.is_empty())
|
||||
.ok_or_else(|| "OAuth allowed_domains must contain only host names".to_string())?;
|
||||
if domain.contains('/')
|
||||
|| domain.contains('\\')
|
||||
|| domain.contains('@')
|
||||
|| domain.contains(':')
|
||||
|| domain.contains(char::is_whitespace)
|
||||
|| domain.parse::<IpAddr>().is_ok()
|
||||
{
|
||||
return Err(
|
||||
"OAuth allowed_domains must contain DNS host names, not IP literals".to_string(),
|
||||
);
|
||||
}
|
||||
domains.push(domain.to_ascii_lowercase());
|
||||
}
|
||||
if domains.is_empty() {
|
||||
return Err("custom OIDC providers must configure allowed domains".to_string());
|
||||
}
|
||||
Ok(domains)
|
||||
}
|
||||
|
||||
fn validate_oauth_endpoint_url(
|
||||
field: &str,
|
||||
value: &str,
|
||||
allowed_domains: &[String],
|
||||
) -> Result<(), String> {
|
||||
let parsed = Url::parse(value).map_err(|_| format!("{field} must be an absolute URL"))?;
|
||||
if parsed.scheme() != "https" || parsed.host_str().is_none() {
|
||||
return Err(format!("{field} must be an absolute https URL"));
|
||||
}
|
||||
if matches!(parsed.host(), Some(Host::Ipv4(_)) | Some(Host::Ipv6(_))) {
|
||||
return Err(format!(
|
||||
"{field} must use a DNS host name, not an IP literal"
|
||||
));
|
||||
}
|
||||
if !parsed.username().is_empty() || parsed.password().is_some() {
|
||||
return Err(format!("{field} must not contain URL credentials"));
|
||||
}
|
||||
if parsed.fragment().is_some() {
|
||||
return Err(format!("{field} must not contain a fragment"));
|
||||
}
|
||||
if field == "authorization_url_override" {
|
||||
for (name, _) in parsed.query_pairs() {
|
||||
if matches!(
|
||||
name.to_ascii_lowercase().as_str(),
|
||||
"response_type"
|
||||
| "client_id"
|
||||
| "redirect_uri"
|
||||
| "state"
|
||||
| "scope"
|
||||
| "code_challenge"
|
||||
| "code_challenge_method"
|
||||
) {
|
||||
return Err(format!(
|
||||
"{field} must not predefine OAuth authorization parameters"
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
if !allowed_domains.is_empty() {
|
||||
let host = parsed
|
||||
.host_str()
|
||||
.map(|value| value.trim_end_matches('.').to_ascii_lowercase())
|
||||
.unwrap_or_default();
|
||||
if !allowed_domains
|
||||
.iter()
|
||||
.any(|domain| host == *domain || host.ends_with(&format!(".{domain}")))
|
||||
{
|
||||
return Err(format!("{field} host is not in the provider allowlist"));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredOAuthProviderConfig {
|
||||
pub provider_type: String,
|
||||
pub display_name: String,
|
||||
@@ -20,6 +227,51 @@ pub struct StoredOAuthProviderConfig {
|
||||
pub updated_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredOAuthProviderConfig {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredOAuthProviderConfig")
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("display_name", &self.display_name)
|
||||
.field("client_id", &self.client_id)
|
||||
.field(
|
||||
"client_secret_encrypted",
|
||||
&self.client_secret_encrypted.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"authorization_url_override",
|
||||
&self
|
||||
.authorization_url_override
|
||||
.as_ref()
|
||||
.map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"token_url_override",
|
||||
&self.token_url_override.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"userinfo_url_override",
|
||||
&self.userinfo_url_override.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("scopes", &self.scopes)
|
||||
.field("redirect_uri", &self.redirect_uri)
|
||||
.field("frontend_callback_url", &self.frontend_callback_url)
|
||||
.field(
|
||||
"attribute_mapping",
|
||||
&self.attribute_mapping.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"extra_config",
|
||||
&self.extra_config.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("icon_url", &self.icon_url)
|
||||
.field("is_enabled", &self.is_enabled)
|
||||
.field("created_at_unix_ms", &self.created_at_unix_ms)
|
||||
.field("updated_at_unix_secs", &self.updated_at_unix_secs)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredOAuthProviderConfig {
|
||||
pub fn new(
|
||||
provider_type: String,
|
||||
@@ -110,7 +362,7 @@ impl StoredOAuthProviderConfig {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize, Default)]
|
||||
#[derive(Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize, Default)]
|
||||
pub enum EncryptedSecretUpdate {
|
||||
#[default]
|
||||
Preserve,
|
||||
@@ -118,6 +370,16 @@ pub enum EncryptedSecretUpdate {
|
||||
Set(String),
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for EncryptedSecretUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self {
|
||||
Self::Preserve => formatter.write_str("Preserve"),
|
||||
Self::Clear => formatter.write_str("Clear"),
|
||||
Self::Set(_) => formatter.write_str("Set([REDACTED])"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl EncryptedSecretUpdate {
|
||||
pub fn mode_name(&self) -> &'static str {
|
||||
match self {
|
||||
@@ -135,7 +397,7 @@ impl EncryptedSecretUpdate {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct UpsertOAuthProviderConfigRecord {
|
||||
pub provider_type: String,
|
||||
pub display_name: String,
|
||||
@@ -153,6 +415,52 @@ pub struct UpsertOAuthProviderConfigRecord {
|
||||
pub is_enabled: bool,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for UpsertOAuthProviderConfigRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("UpsertOAuthProviderConfigRecord")
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("display_name", &self.display_name)
|
||||
.field("client_id", &self.client_id)
|
||||
.field("client_secret_encrypted", &self.client_secret_encrypted)
|
||||
.field(
|
||||
"authorization_url_override",
|
||||
&self
|
||||
.authorization_url_override
|
||||
.as_ref()
|
||||
.map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"token_url_override",
|
||||
&self.token_url_override.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"userinfo_url_override",
|
||||
&self.userinfo_url_override.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("scopes", &self.scopes)
|
||||
.field("redirect_uri", &"[REDACTED]")
|
||||
.field("frontend_callback_url", &"[REDACTED]")
|
||||
.field(
|
||||
"attribute_mapping",
|
||||
&self.attribute_mapping.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"extra_config",
|
||||
&self.extra_config.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("icon_url", &self.icon_url)
|
||||
.field("is_enabled", &self.is_enabled)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum UpsertOAuthProviderConfigOutcome {
|
||||
Upserted(StoredOAuthProviderConfig),
|
||||
DisableRequiresConfirmation { affected_count: usize },
|
||||
}
|
||||
|
||||
impl UpsertOAuthProviderConfigRecord {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
if self.provider_type.trim().is_empty() {
|
||||
@@ -175,11 +483,23 @@ impl UpsertOAuthProviderConfigRecord {
|
||||
"redirect_uri is required".to_string(),
|
||||
));
|
||||
}
|
||||
validate_oauth_redirect_uri(self.redirect_uri.trim())
|
||||
.map_err(crate::DataLayerError::InvalidInput)?;
|
||||
if self.frontend_callback_url.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"frontend_callback_url is required".to_string(),
|
||||
));
|
||||
}
|
||||
validate_oauth_frontend_callback_url(self.frontend_callback_url.trim())
|
||||
.map_err(crate::DataLayerError::InvalidInput)?;
|
||||
validate_oauth_provider_endpoint_config(
|
||||
&self.provider_type,
|
||||
self.authorization_url_override.as_deref(),
|
||||
self.token_url_override.as_deref(),
|
||||
self.userinfo_url_override.as_deref(),
|
||||
self.extra_config.as_ref(),
|
||||
)
|
||||
.map_err(crate::DataLayerError::InvalidInput)?;
|
||||
if let Some(scopes) = &self.scopes {
|
||||
for scope in scopes {
|
||||
if scope.trim().is_empty() {
|
||||
@@ -193,6 +513,186 @@ impl UpsertOAuthProviderConfigRecord {
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
validate_oauth_frontend_callback_url, validate_oauth_provider_endpoint_config,
|
||||
validate_oauth_redirect_uri, EncryptedSecretUpdate, StoredOAuthProviderConfig,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn oauth_provider_debug_output_redacts_encrypted_client_secrets() {
|
||||
let secret = "debug-secret-oauth-provider-ciphertext";
|
||||
let provider = StoredOAuthProviderConfig::new(
|
||||
"linuxdo".to_string(),
|
||||
"Linux.do".to_string(),
|
||||
"client-id".to_string(),
|
||||
"https://gateway.example/api/oauth/linuxdo/callback".to_string(),
|
||||
"https://frontend.example/auth/callback".to_string(),
|
||||
)
|
||||
.expect("provider should build")
|
||||
.with_config_fields(
|
||||
Some(secret.to_string()),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
);
|
||||
|
||||
for rendered in [
|
||||
format!("{provider:?}"),
|
||||
format!("{:?}", EncryptedSecretUpdate::Set(secret.to_string())),
|
||||
] {
|
||||
assert!(!rendered.contains(secret));
|
||||
assert!(rendered.contains("[REDACTED]"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_redirect_uri_requires_absolute_http_url_without_credentials() {
|
||||
assert!(
|
||||
validate_oauth_redirect_uri("https://gateway.example/api/oauth/custom/callback")
|
||||
.is_ok()
|
||||
);
|
||||
assert!(
|
||||
validate_oauth_redirect_uri("http://localhost:8080/api/oauth/custom/callback").is_ok()
|
||||
);
|
||||
for value in [
|
||||
"http://gateway.example/api/oauth/custom/callback",
|
||||
"/api/oauth/custom/callback",
|
||||
"javascript:alert(1)",
|
||||
"https://user:[email protected]/api/oauth/custom/callback",
|
||||
] {
|
||||
assert!(
|
||||
validate_oauth_redirect_uri(value).is_err(),
|
||||
"accepted {value}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_provider_endpoints_require_https_and_the_configured_domain() {
|
||||
let extra = serde_json::json!({"allowed_domains": ["idp.example"]});
|
||||
assert!(validate_oauth_provider_endpoint_config(
|
||||
"custom_oidc_work",
|
||||
Some("https://idp.example/oauth/authorize"),
|
||||
Some("https://idp.example/oauth/token"),
|
||||
Some("https://accounts.idp.example/oauth/userinfo"),
|
||||
Some(&extra),
|
||||
)
|
||||
.is_ok());
|
||||
assert!(validate_oauth_provider_endpoint_config(
|
||||
"custom_oidc_work",
|
||||
Some("https://idp.example/oauth/authorize"),
|
||||
Some("https://attacker.example/oauth/token"),
|
||||
Some("https://idp.example/oauth/userinfo"),
|
||||
Some(&extra),
|
||||
)
|
||||
.is_err());
|
||||
assert!(validate_oauth_provider_endpoint_config(
|
||||
"custom_oidc_work",
|
||||
Some("https://127.0.0.1/oauth/authorize"),
|
||||
Some("https://idp.example/oauth/token"),
|
||||
Some("https://idp.example/oauth/userinfo"),
|
||||
Some(&serde_json::json!({"allowed_domains": ["127.0.0.1"]})),
|
||||
)
|
||||
.is_err());
|
||||
assert!(validate_oauth_provider_endpoint_config(
|
||||
"custom_oidc_work",
|
||||
Some("https://idp.example/oauth/authorize?client_id=attacker"),
|
||||
Some("https://idp.example/oauth/token"),
|
||||
Some("https://idp.example/oauth/userinfo"),
|
||||
Some(&extra),
|
||||
)
|
||||
.is_err());
|
||||
assert!(validate_oauth_provider_endpoint_config(
|
||||
"custom_oidc_work",
|
||||
Some("https://idp.example/oauth/authorize"),
|
||||
Some("http://idp.example/oauth/token"),
|
||||
Some("https://idp.example/oauth/userinfo"),
|
||||
Some(&extra),
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_provider_endpoints_reject_ip_literals_and_predefined_authorization_parameters() {
|
||||
for host in ["127.0.0.1", "[::1]"] {
|
||||
let extra = serde_json::json!({"allowed_domains": [host]});
|
||||
assert!(validate_oauth_provider_endpoint_config(
|
||||
"custom_oidc_work",
|
||||
Some(&format!("https://{host}/oauth/authorize")),
|
||||
Some(&format!("https://{host}/oauth/token")),
|
||||
Some(&format!("https://{host}/oauth/userinfo")),
|
||||
Some(&extra),
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
let extra = serde_json::json!({"allowed_domains": ["idp.example"]});
|
||||
for name in [
|
||||
"response_type",
|
||||
"client_id",
|
||||
"redirect_uri",
|
||||
"state",
|
||||
"scope",
|
||||
"code_challenge",
|
||||
"code_challenge_method",
|
||||
] {
|
||||
assert!(validate_oauth_provider_endpoint_config(
|
||||
"custom_oidc_work",
|
||||
Some(&format!(
|
||||
"https://idp.example/oauth/authorize?{name}=attacker"
|
||||
)),
|
||||
Some("https://idp.example/oauth/token?tenant=workforce"),
|
||||
Some("https://idp.example/oauth/userinfo?schema=current"),
|
||||
Some(&extra),
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
assert!(validate_oauth_provider_endpoint_config(
|
||||
"custom_oidc_work",
|
||||
Some("https://idp.example/oauth/authorize?tenant=workforce"),
|
||||
Some("https://idp.example/oauth/token?tenant=workforce"),
|
||||
Some("https://idp.example/oauth/userinfo?schema=current"),
|
||||
Some(&extra),
|
||||
)
|
||||
.is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oauth_frontend_callback_rejects_token_exfiltration_targets() {
|
||||
for value in [
|
||||
"https://frontend.example/auth/callback",
|
||||
"http://localhost:5173/auth/callback",
|
||||
"http://127.0.0.1:5173/auth/callback",
|
||||
"http://[::1]:5173/auth/callback",
|
||||
] {
|
||||
assert!(
|
||||
validate_oauth_frontend_callback_url(value).is_ok(),
|
||||
"rejected {value}"
|
||||
);
|
||||
}
|
||||
for value in [
|
||||
"http://attacker.example/auth/callback",
|
||||
"https://user:[email protected]/auth/callback",
|
||||
"https://frontend.example/auth/callback?next=https://attacker.example",
|
||||
"https://frontend.example/auth/callback#access_token=stolen",
|
||||
"https://frontend.example/not-the-callback",
|
||||
] {
|
||||
assert!(
|
||||
validate_oauth_frontend_callback_url(value).is_err(),
|
||||
"accepted {value}"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait OAuthProviderReadRepository: Send + Sync {
|
||||
async fn list_oauth_provider_configs(
|
||||
@@ -216,11 +716,42 @@ pub trait OAuthProviderWriteRepository: Send + Sync {
|
||||
async fn upsert_oauth_provider_config(
|
||||
&self,
|
||||
record: &UpsertOAuthProviderConfigRecord,
|
||||
) -> Result<StoredOAuthProviderConfig, crate::DataLayerError>;
|
||||
) -> Result<StoredOAuthProviderConfig, crate::DataLayerError> {
|
||||
match self
|
||||
.upsert_oauth_provider_config_guarded(record, false, false, 0)
|
||||
.await?
|
||||
{
|
||||
UpsertOAuthProviderConfigOutcome::Upserted(provider) => Ok(provider),
|
||||
UpsertOAuthProviderConfigOutcome::DisableRequiresConfirmation { affected_count } => {
|
||||
Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"disabling OAuth provider requires confirmation for {affected_count} affected users"
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn delete_oauth_provider_config(
|
||||
async fn upsert_oauth_provider_config_guarded(
|
||||
&self,
|
||||
record: &UpsertOAuthProviderConfigRecord,
|
||||
ldap_exclusive: bool,
|
||||
force_disable: bool,
|
||||
locked_users_snapshot: usize,
|
||||
) -> Result<UpsertOAuthProviderConfigOutcome, crate::DataLayerError>;
|
||||
|
||||
/// Replace only the stored client secret when the provider and exact previously observed
|
||||
/// ciphertext still match. Implementations must not modify `updated_at` or any non-secret
|
||||
/// provider field; this is used by lazy record-bound ciphertext migration.
|
||||
async fn compare_and_swap_oauth_provider_client_secret(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
expected: &str,
|
||||
replacement: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn delete_oauth_provider_config_if_unlinked(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
has_links_snapshot: bool,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
}
|
||||
|
||||
|
||||
@@ -4,11 +4,12 @@ mod types;
|
||||
pub use snapshot::ProviderCatalogSnapshot;
|
||||
pub use types::{
|
||||
ProviderCatalogKeyAdaptiveState, ProviderCatalogKeyAdaptiveStateUpdate,
|
||||
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyHealthStateUpdate,
|
||||
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyAdminCasUpdate, ProviderCatalogKeyCredentialsCasUpdate,
|
||||
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthCredentialFence,
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate, ProviderCatalogKeyRuntimeMetadataUpdate,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository,
|
||||
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogProviderConfigCasUpdate,
|
||||
ProviderCatalogProxyCasUpdate, ProviderCatalogReadRepository,
|
||||
ProviderCatalogUpstreamMetadataNamespaceExpectation,
|
||||
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
|
||||
@@ -1,11 +1,27 @@
|
||||
use async_trait::async_trait;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
const REDACTED_DEBUG_VALUE: &str = "[REDACTED]";
|
||||
|
||||
fn redacted_debug_option<T>(value: &Option<T>) -> Option<&'static str> {
|
||||
value.as_ref().map(|_| REDACTED_DEBUG_VALUE)
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogUpstreamMetadataNamespaceUpdate {
|
||||
pub namespace: String,
|
||||
pub value: serde_json::Value,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderCatalogUpstreamMetadataNamespaceUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogUpstreamMetadataNamespaceUpdate")
|
||||
.field("namespace", &self.namespace)
|
||||
.field("value", &REDACTED_DEBUG_VALUE)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyAdaptiveState {
|
||||
pub learned_rpm_limit: Option<u32>,
|
||||
@@ -28,7 +44,7 @@ impl ProviderCatalogKeyAdaptiveState {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyAdaptiveStateUpdate {
|
||||
pub key_id: String,
|
||||
/// Optional auth_config fence for request-owned adaptive feedback.
|
||||
@@ -41,7 +57,24 @@ pub struct ProviderCatalogKeyAdaptiveStateUpdate {
|
||||
pub updated_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
impl std::fmt::Debug for ProviderCatalogKeyAdaptiveStateUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogKeyAdaptiveStateUpdate")
|
||||
.field("key_id", &self.key_id)
|
||||
.field(
|
||||
"expected_encrypted_auth_config",
|
||||
&redacted_debug_option(&self.expected_encrypted_auth_config),
|
||||
)
|
||||
.field("expected", &self.expected)
|
||||
.field("next", &self.next)
|
||||
.field("status_snapshot_patch", &REDACTED_DEBUG_VALUE)
|
||||
.field("updated_at_unix_secs", &self.updated_at_unix_secs)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyRuntimeMetadataUpdate {
|
||||
pub key_id: String,
|
||||
pub namespace: String,
|
||||
@@ -59,7 +92,24 @@ pub struct ProviderCatalogKeyRuntimeMetadataUpdate {
|
||||
pub updated_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
impl std::fmt::Debug for ProviderCatalogKeyRuntimeMetadataUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogKeyRuntimeMetadataUpdate")
|
||||
.field("key_id", &self.key_id)
|
||||
.field("namespace", &self.namespace)
|
||||
.field(
|
||||
"expected_upstream_metadata_value",
|
||||
&redacted_debug_option(&self.expected_upstream_metadata_value),
|
||||
)
|
||||
.field("upstream_metadata_value", &REDACTED_DEBUG_VALUE)
|
||||
.field("status_snapshot_patch", &REDACTED_DEBUG_VALUE)
|
||||
.field("updated_at_unix_secs", &self.updated_at_unix_secs)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyStatusSnapshotUpdate {
|
||||
pub key_id: String,
|
||||
/// Top-level status fields owned by the caller.
|
||||
@@ -67,10 +117,21 @@ pub struct ProviderCatalogKeyStatusSnapshotUpdate {
|
||||
pub updated_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderCatalogKeyStatusSnapshotUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogKeyStatusSnapshotUpdate")
|
||||
.field("key_id", &self.key_id)
|
||||
.field("status_snapshot_patch", &REDACTED_DEBUG_VALUE)
|
||||
.field("updated_at_unix_secs", &self.updated_at_unix_secs)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Credential context observed before an OAuth refresh started. Repositories
|
||||
/// compare every field atomically with the runtime-state update so an
|
||||
/// administrator replacement cannot be overwritten by an older refresh.
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyOAuthCredentialFence {
|
||||
/// Exact nullable ciphertext stored in `provider_api_keys.api_key`.
|
||||
pub encrypted_api_key: Option<String>,
|
||||
@@ -79,10 +140,25 @@ pub struct ProviderCatalogKeyOAuthCredentialFence {
|
||||
pub provider_type: String,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderCatalogKeyOAuthCredentialFence {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogKeyOAuthCredentialFence")
|
||||
.field(
|
||||
"encrypted_api_key",
|
||||
&redacted_debug_option(&self.encrypted_api_key),
|
||||
)
|
||||
.field("auth_type", &self.auth_type)
|
||||
.field("provider_id", &self.provider_id)
|
||||
.field("provider_type", &self.provider_type)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Administrator-owned key replacement fenced by the exact credential state
|
||||
/// observed while the edit was prepared. This prevents an older admin request
|
||||
/// from restoring credentials that a concurrent request already replaced.
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyAdminCasUpdate {
|
||||
pub expected_encrypted_auth_config: Option<String>,
|
||||
pub expected_credential: ProviderCatalogKeyOAuthCredentialFence,
|
||||
@@ -102,9 +178,28 @@ pub struct ProviderCatalogKeyAdminCasUpdate {
|
||||
pub reset_oauth_runtime: bool,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderCatalogKeyAdminCasUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogKeyAdminCasUpdate")
|
||||
.field(
|
||||
"expected_encrypted_auth_config",
|
||||
&redacted_debug_option(&self.expected_encrypted_auth_config),
|
||||
)
|
||||
.field("expected_credential", &self.expected_credential)
|
||||
.field("key", &self.key)
|
||||
.field(
|
||||
"codex_rotation",
|
||||
&redacted_debug_option(&self.codex_rotation),
|
||||
)
|
||||
.field("reset_oauth_runtime", &self.reset_oauth_runtime)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Atomic key deletion fenced by the exact OAuth credential generation that
|
||||
/// produced the terminal failure and, when supplied, one metadata namespace.
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyOAuthCredentialCasDelete {
|
||||
pub key_id: String,
|
||||
pub expected_encrypted_auth_config: Option<String>,
|
||||
@@ -114,22 +209,53 @@ pub struct ProviderCatalogKeyOAuthCredentialCasDelete {
|
||||
Option<ProviderCatalogUpstreamMetadataNamespaceExpectation>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderCatalogKeyOAuthCredentialCasDelete {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogKeyOAuthCredentialCasDelete")
|
||||
.field("key_id", &self.key_id)
|
||||
.field(
|
||||
"expected_encrypted_auth_config",
|
||||
&redacted_debug_option(&self.expected_encrypted_auth_config),
|
||||
)
|
||||
.field("expected_credential", &self.expected_credential)
|
||||
.field(
|
||||
"expected_upstream_metadata_namespace",
|
||||
&self.expected_upstream_metadata_namespace,
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Optional single-namespace metadata fence for an OAuth runtime CAS.
|
||||
///
|
||||
/// The outer option on the owning update controls whether the namespace is
|
||||
/// compared. Within an expectation, `None` requires the namespace to be absent.
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogUpstreamMetadataNamespaceExpectation {
|
||||
pub namespace: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub expected_value: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderCatalogUpstreamMetadataNamespaceExpectation {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogUpstreamMetadataNamespaceExpectation")
|
||||
.field("namespace", &self.namespace)
|
||||
.field(
|
||||
"expected_value",
|
||||
&redacted_debug_option(&self.expected_value),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Agent/runtime-owned OAuth state update fenced by the exact encrypted
|
||||
/// auth_config and, when supplied, credential context observed before the
|
||||
/// refresh started. Repositories must update only these fields and return
|
||||
/// `false` when an expected value changed.
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
pub key_id: String,
|
||||
pub expected_encrypted_auth_config: Option<String>,
|
||||
@@ -162,7 +288,53 @@ pub struct ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
pub updated_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
impl std::fmt::Debug for ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogKeyOAuthRuntimeStateCasUpdate")
|
||||
.field("key_id", &self.key_id)
|
||||
.field(
|
||||
"expected_encrypted_auth_config",
|
||||
&redacted_debug_option(&self.expected_encrypted_auth_config),
|
||||
)
|
||||
.field("expected_credential", &self.expected_credential)
|
||||
.field(
|
||||
"expected_upstream_metadata_namespace",
|
||||
&self.expected_upstream_metadata_namespace,
|
||||
)
|
||||
.field("encrypted_auth_config", &REDACTED_DEBUG_VALUE)
|
||||
.field(
|
||||
"encrypted_api_key_update",
|
||||
&redacted_debug_option(&self.encrypted_api_key_update),
|
||||
)
|
||||
.field(
|
||||
"expires_at_unix_secs_update",
|
||||
&self.expires_at_unix_secs_update,
|
||||
)
|
||||
.field(
|
||||
"oauth_invalid_at_unix_secs",
|
||||
&self.oauth_invalid_at_unix_secs,
|
||||
)
|
||||
.field(
|
||||
"oauth_invalid_reason",
|
||||
&redacted_debug_option(&self.oauth_invalid_reason),
|
||||
)
|
||||
.field(
|
||||
"upstream_metadata_patch",
|
||||
&redacted_debug_option(&self.upstream_metadata_patch),
|
||||
)
|
||||
.field(
|
||||
"upstream_metadata_namespace_to_remove",
|
||||
&self.upstream_metadata_namespace_to_remove,
|
||||
)
|
||||
.field("status_snapshot_patch", &REDACTED_DEBUG_VALUE)
|
||||
.field("reset_error_count", &self.reset_error_count)
|
||||
.field("updated_at_unix_secs", &self.updated_at_unix_secs)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyHealthStateUpdate {
|
||||
pub key_id: String,
|
||||
/// Optional auth_config fence for lifecycle-owned health recovery.
|
||||
@@ -174,7 +346,108 @@ pub struct ProviderCatalogKeyHealthStateUpdate {
|
||||
pub circuit_breaker_by_format: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
impl std::fmt::Debug for ProviderCatalogKeyHealthStateUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogKeyHealthStateUpdate")
|
||||
.field("key_id", &self.key_id)
|
||||
.field(
|
||||
"expected_encrypted_auth_config",
|
||||
&redacted_debug_option(&self.expected_encrypted_auth_config),
|
||||
)
|
||||
.field("expected_health_by_format", &self.expected_health_by_format)
|
||||
.field(
|
||||
"expected_circuit_breaker_by_format",
|
||||
&self.expected_circuit_breaker_by_format,
|
||||
)
|
||||
.field("health_by_format", &self.health_by_format)
|
||||
.field("circuit_breaker_by_format", &self.circuit_breaker_by_format)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogProviderConfigCasUpdate {
|
||||
pub provider_id: String,
|
||||
pub expected_config: Option<serde_json::Value>,
|
||||
pub config: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderCatalogProviderConfigCasUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogProviderConfigCasUpdate")
|
||||
.field("provider_id", &self.provider_id)
|
||||
.field(
|
||||
"expected_config",
|
||||
&redacted_debug_option(&self.expected_config),
|
||||
)
|
||||
.field("config", &redacted_debug_option(&self.config))
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogProxyCasUpdate {
|
||||
pub record_id: String,
|
||||
pub expected_proxy: Option<serde_json::Value>,
|
||||
pub proxy: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderCatalogProxyCasUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogProxyCasUpdate")
|
||||
.field("record_id", &self.record_id)
|
||||
.field(
|
||||
"expected_proxy",
|
||||
&redacted_debug_option(&self.expected_proxy),
|
||||
)
|
||||
.field("proxy", &redacted_debug_option(&self.proxy))
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
/// Secret-only migration fenced by the complete catalog-key credential
|
||||
/// identity. The provider fence prevents a legacy credential from being
|
||||
/// re-encrypted for an obsolete provider after a concurrent key move.
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProviderCatalogKeyCredentialsCasUpdate {
|
||||
pub key_id: String,
|
||||
pub expected_provider_id: String,
|
||||
pub expected_encrypted_api_key: Option<String>,
|
||||
pub expected_encrypted_auth_config: Option<String>,
|
||||
pub encrypted_api_key: Option<String>,
|
||||
pub encrypted_auth_config: Option<String>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProviderCatalogKeyCredentialsCasUpdate {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProviderCatalogKeyCredentialsCasUpdate")
|
||||
.field("key_id", &self.key_id)
|
||||
.field("expected_provider_id", &self.expected_provider_id)
|
||||
.field(
|
||||
"expected_encrypted_api_key",
|
||||
&redacted_debug_option(&self.expected_encrypted_api_key),
|
||||
)
|
||||
.field(
|
||||
"expected_encrypted_auth_config",
|
||||
&redacted_debug_option(&self.expected_encrypted_auth_config),
|
||||
)
|
||||
.field(
|
||||
"encrypted_api_key",
|
||||
&redacted_debug_option(&self.encrypted_api_key),
|
||||
)
|
||||
.field(
|
||||
"encrypted_auth_config",
|
||||
&redacted_debug_option(&self.encrypted_auth_config),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredProviderCatalogProvider {
|
||||
pub id: String,
|
||||
pub name: String,
|
||||
@@ -201,6 +474,23 @@ pub struct StoredProviderCatalogProvider {
|
||||
pub updated_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredProviderCatalogProvider {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredProviderCatalogProvider")
|
||||
.field("id", &self.id)
|
||||
.field("name", &self.name)
|
||||
.field("provider_type", &self.provider_type)
|
||||
.field("billing_type", &self.billing_type)
|
||||
.field("is_active", &self.is_active)
|
||||
.field("proxy", &redacted_debug_option(&self.proxy))
|
||||
.field("config", &redacted_debug_option(&self.config))
|
||||
.field("created_at_unix_ms", &self.created_at_unix_ms)
|
||||
.field("updated_at_unix_secs", &self.updated_at_unix_secs)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredProviderCatalogProvider {
|
||||
pub fn new(
|
||||
id: String,
|
||||
@@ -311,7 +601,7 @@ impl StoredProviderCatalogProvider {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredProviderCatalogEndpoint {
|
||||
pub id: String,
|
||||
pub provider_id: String,
|
||||
@@ -332,6 +622,27 @@ pub struct StoredProviderCatalogEndpoint {
|
||||
pub updated_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredProviderCatalogEndpoint {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredProviderCatalogEndpoint")
|
||||
.field("id", &self.id)
|
||||
.field("provider_id", &self.provider_id)
|
||||
.field("api_format", &self.api_format)
|
||||
.field("api_family", &self.api_family)
|
||||
.field("endpoint_kind", &self.endpoint_kind)
|
||||
.field("is_active", &self.is_active)
|
||||
.field("base_url", &REDACTED_DEBUG_VALUE)
|
||||
.field("header_rules", &redacted_debug_option(&self.header_rules))
|
||||
.field("body_rules", &redacted_debug_option(&self.body_rules))
|
||||
.field("config", &redacted_debug_option(&self.config))
|
||||
.field("proxy", &redacted_debug_option(&self.proxy))
|
||||
.field("created_at_unix_ms", &self.created_at_unix_ms)
|
||||
.field("updated_at_unix_secs", &self.updated_at_unix_secs)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredProviderCatalogEndpoint {
|
||||
pub fn new(
|
||||
id: String,
|
||||
@@ -413,7 +724,7 @@ impl StoredProviderCatalogEndpoint {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredProviderCatalogKey {
|
||||
pub id: String,
|
||||
pub provider_id: String,
|
||||
@@ -470,7 +781,44 @@ pub struct StoredProviderCatalogKey {
|
||||
pub circuit_breaker_by_format: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
impl std::fmt::Debug for StoredProviderCatalogKey {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredProviderCatalogKey")
|
||||
.field("id", &self.id)
|
||||
.field("provider_id", &self.provider_id)
|
||||
.field("name", &self.name)
|
||||
.field("auth_type", &self.auth_type)
|
||||
.field("is_active", &self.is_active)
|
||||
.field(
|
||||
"encrypted_api_key",
|
||||
&redacted_debug_option(&self.encrypted_api_key),
|
||||
)
|
||||
.field(
|
||||
"encrypted_auth_config",
|
||||
&redacted_debug_option(&self.encrypted_auth_config),
|
||||
)
|
||||
.field("proxy", &redacted_debug_option(&self.proxy))
|
||||
.field("fingerprint", &redacted_debug_option(&self.fingerprint))
|
||||
.field(
|
||||
"upstream_metadata",
|
||||
&redacted_debug_option(&self.upstream_metadata),
|
||||
)
|
||||
.field(
|
||||
"oauth_invalid_reason",
|
||||
&redacted_debug_option(&self.oauth_invalid_reason),
|
||||
)
|
||||
.field(
|
||||
"status_snapshot",
|
||||
&redacted_debug_option(&self.status_snapshot),
|
||||
)
|
||||
.field("created_at_unix_ms", &self.created_at_unix_ms)
|
||||
.field("updated_at_unix_secs", &self.updated_at_unix_secs)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq)]
|
||||
pub struct StoredProviderCatalogKeyMaintenanceSummary {
|
||||
pub id: String,
|
||||
pub provider_id: String,
|
||||
@@ -478,6 +826,21 @@ pub struct StoredProviderCatalogKeyMaintenanceSummary {
|
||||
pub upstream_metadata: Option<serde_json::Value>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredProviderCatalogKeyMaintenanceSummary {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredProviderCatalogKeyMaintenanceSummary")
|
||||
.field("id", &self.id)
|
||||
.field("provider_id", &self.provider_id)
|
||||
.field("is_active", &self.is_active)
|
||||
.field(
|
||||
"upstream_metadata",
|
||||
&redacted_debug_option(&self.upstream_metadata),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredProviderCatalogKey {
|
||||
pub fn new(
|
||||
id: String,
|
||||
@@ -590,6 +953,18 @@ impl StoredProviderCatalogKey {
|
||||
Ok(self)
|
||||
}
|
||||
|
||||
pub fn with_auth_channel_policy_fields(
|
||||
mut self,
|
||||
auth_type_by_format: Option<serde_json::Value>,
|
||||
allow_auth_channel_mismatch_formats: Option<serde_json::Value>,
|
||||
) -> Result<Self, crate::DataLayerError> {
|
||||
validate_auth_type_by_format(auth_type_by_format.as_ref())?;
|
||||
validate_auth_channel_mismatch_formats(allow_auth_channel_mismatch_formats.as_ref())?;
|
||||
self.auth_type_by_format = auth_type_by_format;
|
||||
self.allow_auth_channel_mismatch_formats = allow_auth_channel_mismatch_formats;
|
||||
Ok(self)
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn with_rate_limit_fields(
|
||||
mut self,
|
||||
@@ -646,6 +1021,58 @@ impl StoredProviderCatalogKey {
|
||||
}
|
||||
}
|
||||
|
||||
fn validate_auth_type_by_format(
|
||||
value: Option<&serde_json::Value>,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
let Some(value) = value else {
|
||||
return Ok(());
|
||||
};
|
||||
let Some(entries) = value.as_object() else {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"provider_api_keys.auth_type_by_format must be a JSON object".to_string(),
|
||||
));
|
||||
};
|
||||
for (api_format, auth_type) in entries {
|
||||
let valid_api_format = !api_format.trim().is_empty();
|
||||
let valid_auth_type = auth_type.as_str().is_some_and(|value| {
|
||||
matches!(
|
||||
value.trim().to_ascii_lowercase().as_str(),
|
||||
"api_key" | "bearer"
|
||||
)
|
||||
});
|
||||
if !valid_api_format || !valid_auth_type {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"provider_api_keys.auth_type_by_format contains an invalid entry".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_auth_channel_mismatch_formats(
|
||||
value: Option<&serde_json::Value>,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
let Some(value) = value else {
|
||||
return Ok(());
|
||||
};
|
||||
let Some(items) = value.as_array() else {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"provider_api_keys.allow_auth_channel_mismatch_formats must be a JSON array"
|
||||
.to_string(),
|
||||
));
|
||||
};
|
||||
if items
|
||||
.iter()
|
||||
.any(|item| item.as_str().is_none_or(|value| value.trim().is_empty()))
|
||||
{
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"provider_api_keys.allow_auth_channel_mismatch_formats contains an invalid entry"
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
impl From<&StoredProviderCatalogKey> for ProviderCatalogKeyAdaptiveState {
|
||||
fn from(key: &StoredProviderCatalogKey) -> Self {
|
||||
Self {
|
||||
@@ -664,7 +1091,60 @@ impl From<&StoredProviderCatalogKey> for ProviderCatalogKeyAdaptiveState {
|
||||
|
||||
#[cfg(test)]
|
||||
mod transport_tests {
|
||||
use super::StoredProviderCatalogKey;
|
||||
use super::{
|
||||
ProviderCatalogKeyCredentialsCasUpdate, ProviderCatalogKeyOAuthCredentialFence,
|
||||
ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
|
||||
ProviderCatalogUpstreamMetadataNamespaceExpectation, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
|
||||
fn assert_debug_redacts<T: std::fmt::Debug>(value: &T, secrets: &[&str]) {
|
||||
let debug = format!("{value:?}");
|
||||
assert!(debug.contains("[REDACTED]"), "debug output: {debug}");
|
||||
for secret in secrets {
|
||||
assert!(
|
||||
!debug.contains(secret),
|
||||
"debug output leaked {secret}: {debug}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
fn sample_key() -> StoredProviderCatalogKey {
|
||||
StoredProviderCatalogKey::new(
|
||||
"key-policy".to_string(),
|
||||
"provider-policy".to_string(),
|
||||
"policy".to_string(),
|
||||
"api_key".to_string(),
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("key should build")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_catalog_key_auth_channel_policy_rejects_malformed_stored_json() {
|
||||
assert!(sample_key()
|
||||
.with_auth_channel_policy_fields(
|
||||
Some(serde_json::json!({"openai:chat": "bearer"})),
|
||||
Some(serde_json::json!([])),
|
||||
)
|
||||
.is_ok());
|
||||
assert!(sample_key()
|
||||
.with_auth_channel_policy_fields(Some(serde_json::Value::Null), None)
|
||||
.is_err());
|
||||
assert!(sample_key()
|
||||
.with_auth_channel_policy_fields(
|
||||
Some(serde_json::json!({"openai:chat": "oauth"})),
|
||||
None,
|
||||
)
|
||||
.is_err());
|
||||
assert!(sample_key()
|
||||
.with_auth_channel_policy_fields(None, Some(serde_json::Value::Null))
|
||||
.is_err());
|
||||
assert!(sample_key()
|
||||
.with_auth_channel_policy_fields(None, Some(serde_json::json!([""])))
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_catalog_key_defaults_concurrent_limit_to_none() {
|
||||
@@ -707,6 +1187,152 @@ mod transport_tests {
|
||||
assert_eq!(key.rpm_limit, Some(120));
|
||||
assert_eq!(key.concurrent_limit, Some(3));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_catalog_debug_output_redacts_credentials_and_transport_metadata() {
|
||||
let mut key = sample_key();
|
||||
key.encrypted_api_key = Some("catalog-api-key-ciphertext-canary".to_string());
|
||||
key.encrypted_auth_config = Some("catalog-auth-config-ciphertext-canary".to_string());
|
||||
key.proxy = Some(serde_json::json!({"password": "catalog-proxy-canary"}));
|
||||
key.fingerprint = Some(serde_json::json!({"device": "catalog-device-canary"}));
|
||||
key.upstream_metadata = Some(serde_json::json!({"token": "catalog-metadata-canary"}));
|
||||
key.oauth_invalid_reason = Some("catalog-oauth-reason-canary".to_string());
|
||||
key.status_snapshot = Some(serde_json::json!({"raw": "catalog-status-canary"}));
|
||||
assert_debug_redacts(
|
||||
&key,
|
||||
&[
|
||||
"catalog-api-key-ciphertext-canary",
|
||||
"catalog-auth-config-ciphertext-canary",
|
||||
"catalog-proxy-canary",
|
||||
"catalog-device-canary",
|
||||
"catalog-metadata-canary",
|
||||
"catalog-oauth-reason-canary",
|
||||
"catalog-status-canary",
|
||||
],
|
||||
);
|
||||
|
||||
let provider = StoredProviderCatalogProvider::new(
|
||||
"provider-debug".to_string(),
|
||||
"debug".to_string(),
|
||||
None,
|
||||
"openai".to_string(),
|
||||
)
|
||||
.expect("provider should build")
|
||||
.with_transport_fields(
|
||||
true,
|
||||
false,
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
Some(serde_json::json!({"password": "provider-proxy-canary"})),
|
||||
None,
|
||||
None,
|
||||
Some(serde_json::json!({"secret": "provider-config-canary"})),
|
||||
);
|
||||
assert_debug_redacts(
|
||||
&provider,
|
||||
&["provider-proxy-canary", "provider-config-canary"],
|
||||
);
|
||||
|
||||
let endpoint = StoredProviderCatalogEndpoint::new(
|
||||
"endpoint-debug".to_string(),
|
||||
"provider-debug".to_string(),
|
||||
"openai:chat".to_string(),
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("endpoint should build")
|
||||
.with_transport_fields(
|
||||
"https://endpoint-user:[email protected]/endpoint-token-canary"
|
||||
.to_string(),
|
||||
Some(serde_json::json!({"Authorization": "endpoint-header-canary"})),
|
||||
Some(serde_json::json!({"credential": "endpoint-body-canary"})),
|
||||
None,
|
||||
None,
|
||||
Some(serde_json::json!({"secret": "endpoint-config-canary"})),
|
||||
None,
|
||||
Some(serde_json::json!({"password": "endpoint-proxy-canary"})),
|
||||
)
|
||||
.expect("endpoint should accept transport fields");
|
||||
assert_debug_redacts(
|
||||
&endpoint,
|
||||
&[
|
||||
"endpoint-password-canary",
|
||||
"endpoint-token-canary",
|
||||
"endpoint-header-canary",
|
||||
"endpoint-body-canary",
|
||||
"endpoint-config-canary",
|
||||
"endpoint-proxy-canary",
|
||||
],
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_catalog_cas_debug_output_redacts_credential_fences() {
|
||||
let fence = ProviderCatalogKeyOAuthCredentialFence {
|
||||
encrypted_api_key: Some("fence-api-key-canary".to_string()),
|
||||
auth_type: "oauth".to_string(),
|
||||
provider_id: "provider-debug".to_string(),
|
||||
provider_type: "codex".to_string(),
|
||||
};
|
||||
assert_debug_redacts(&fence, &["fence-api-key-canary"]);
|
||||
|
||||
let credentials = ProviderCatalogKeyCredentialsCasUpdate {
|
||||
key_id: "key-debug".to_string(),
|
||||
expected_provider_id: "provider-debug".to_string(),
|
||||
expected_encrypted_api_key: Some("expected-api-key-canary".to_string()),
|
||||
expected_encrypted_auth_config: Some("expected-auth-config-canary".to_string()),
|
||||
encrypted_api_key: Some("replacement-api-key-canary".to_string()),
|
||||
encrypted_auth_config: Some("replacement-auth-config-canary".to_string()),
|
||||
};
|
||||
assert_debug_redacts(
|
||||
&credentials,
|
||||
&[
|
||||
"expected-api-key-canary",
|
||||
"expected-auth-config-canary",
|
||||
"replacement-api-key-canary",
|
||||
"replacement-auth-config-canary",
|
||||
],
|
||||
);
|
||||
|
||||
let runtime = ProviderCatalogKeyOAuthRuntimeStateCasUpdate {
|
||||
key_id: "key-debug".to_string(),
|
||||
expected_encrypted_auth_config: Some("runtime-expected-auth-canary".to_string()),
|
||||
expected_credential: Some(fence),
|
||||
expected_upstream_metadata_namespace: Some(
|
||||
ProviderCatalogUpstreamMetadataNamespaceExpectation {
|
||||
namespace: "oauth".to_string(),
|
||||
expected_value: Some(serde_json::json!({"token": "runtime-fence-canary"})),
|
||||
},
|
||||
),
|
||||
encrypted_auth_config: "runtime-auth-config-canary".to_string(),
|
||||
encrypted_api_key_update: Some("runtime-api-key-canary".to_string()),
|
||||
expires_at_unix_secs_update: Some(Some(123)),
|
||||
oauth_invalid_at_unix_secs: Some(124),
|
||||
oauth_invalid_reason: Some("runtime-provider-error-canary".to_string()),
|
||||
upstream_metadata_patch: Some(serde_json::json!({
|
||||
"refresh_token": "runtime-metadata-canary"
|
||||
})),
|
||||
upstream_metadata_namespace_to_remove: None,
|
||||
status_snapshot_patch: serde_json::json!({"raw": "runtime-status-canary"}),
|
||||
reset_error_count: true,
|
||||
updated_at_unix_secs: Some(125),
|
||||
};
|
||||
assert_debug_redacts(
|
||||
&runtime,
|
||||
&[
|
||||
"runtime-expected-auth-canary",
|
||||
"fence-api-key-canary",
|
||||
"runtime-fence-canary",
|
||||
"runtime-auth-config-canary",
|
||||
"runtime-api-key-canary",
|
||||
"runtime-provider-error-canary",
|
||||
"runtime-metadata-canary",
|
||||
"runtime-status-canary",
|
||||
],
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Default)]
|
||||
@@ -848,6 +1474,26 @@ pub trait ProviderCatalogWriteRepository: Send + Sync {
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
) -> Result<StoredProviderCatalogProvider, crate::DataLayerError>;
|
||||
|
||||
async fn compare_and_swap_provider_config(
|
||||
&self,
|
||||
_update: &ProviderCatalogProviderConfigCasUpdate,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
Err(crate::DataLayerError::InvalidConfiguration(
|
||||
"provider catalog config compare-and-swap is not supported by this repository"
|
||||
.to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn compare_and_swap_provider_proxy(
|
||||
&self,
|
||||
_update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
Err(crate::DataLayerError::InvalidConfiguration(
|
||||
"provider catalog provider proxy compare-and-swap is not supported by this repository"
|
||||
.to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn delete_provider(&self, provider_id: &str) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn cleanup_deleted_provider_refs(
|
||||
@@ -868,6 +1514,16 @@ pub trait ProviderCatalogWriteRepository: Send + Sync {
|
||||
endpoint: &StoredProviderCatalogEndpoint,
|
||||
) -> Result<StoredProviderCatalogEndpoint, crate::DataLayerError>;
|
||||
|
||||
async fn compare_and_swap_endpoint_proxy(
|
||||
&self,
|
||||
_update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
Err(crate::DataLayerError::InvalidConfiguration(
|
||||
"provider catalog endpoint proxy compare-and-swap is not supported by this repository"
|
||||
.to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn delete_endpoint(&self, endpoint_id: &str) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn create_key(
|
||||
@@ -880,6 +1536,26 @@ pub trait ProviderCatalogWriteRepository: Send + Sync {
|
||||
key: &StoredProviderCatalogKey,
|
||||
) -> Result<StoredProviderCatalogKey, crate::DataLayerError>;
|
||||
|
||||
async fn compare_and_swap_key_proxy(
|
||||
&self,
|
||||
_update: &ProviderCatalogProxyCasUpdate,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
Err(crate::DataLayerError::InvalidConfiguration(
|
||||
"provider catalog key proxy compare-and-swap is not supported by this repository"
|
||||
.to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn compare_and_swap_key_credentials(
|
||||
&self,
|
||||
_update: &ProviderCatalogKeyCredentialsCasUpdate,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
Err(crate::DataLayerError::InvalidConfiguration(
|
||||
"provider catalog key credential compare-and-swap is not supported by this repository"
|
||||
.to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
/// Compare-and-swap administrator-owned key configuration. Credential
|
||||
/// rotation, Codex namespace replacement, and quota invalidation must be
|
||||
/// committed atomically with the configuration update.
|
||||
@@ -947,20 +1623,11 @@ pub trait ProviderCatalogWriteRepository: Send + Sync {
|
||||
key_id: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
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, crate::DataLayerError>;
|
||||
|
||||
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, crate::DataLayerError>;
|
||||
|
||||
|
||||
@@ -1,9 +1,15 @@
|
||||
use aether_contracts::tunnel_security::TUNNEL_SECURITY_NON_TLS_REQUIRED;
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
const PROXY_NODE_BOUND_TUNNEL_SECRET_PREFIX: &str =
|
||||
"aether-proxy-node-secret-v2:aether-runtime-secret-v1:";
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredProxyNode {
|
||||
pub id: String,
|
||||
#[serde(default = "new_proxy_node_tunnel_generation")]
|
||||
pub tunnel_generation: String,
|
||||
pub name: String,
|
||||
pub ip: String,
|
||||
pub port: i32,
|
||||
@@ -34,6 +40,42 @@ pub struct StoredProxyNode {
|
||||
pub updated_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredProxyNode {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredProxyNode")
|
||||
.field("id", &self.id)
|
||||
.field("tunnel_generation", &self.tunnel_generation)
|
||||
.field("name", &self.name)
|
||||
.field("ip", &self.ip)
|
||||
.field("port", &self.port)
|
||||
.field("region", &self.region)
|
||||
.field("is_manual", &self.is_manual)
|
||||
.field("proxy_url", &self.proxy_url.as_ref().map(|_| "[REDACTED]"))
|
||||
.field("proxy_username", &self.proxy_username)
|
||||
.field(
|
||||
"proxy_password",
|
||||
&self.proxy_password.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("status", &self.status)
|
||||
.field("tunnel_mode", &self.tunnel_mode)
|
||||
.field("tunnel_connected", &self.tunnel_connected)
|
||||
.field(
|
||||
"proxy_metadata",
|
||||
&self.proxy_metadata.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"hardware_info",
|
||||
&self.hardware_info.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field(
|
||||
"remote_config",
|
||||
&self.remote_config.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredProxyNode {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
@@ -76,6 +118,7 @@ impl StoredProxyNode {
|
||||
|
||||
Ok(Self {
|
||||
id,
|
||||
tunnel_generation: new_proxy_node_tunnel_generation(),
|
||||
name,
|
||||
ip,
|
||||
port,
|
||||
@@ -147,11 +190,22 @@ impl StoredProxyNode {
|
||||
self.proxy_password = proxy_password;
|
||||
self
|
||||
}
|
||||
|
||||
pub fn with_tunnel_generation(mut self, tunnel_generation: String) -> Self {
|
||||
self.tunnel_generation = tunnel_generation;
|
||||
self
|
||||
}
|
||||
}
|
||||
|
||||
pub fn new_proxy_node_tunnel_generation() -> String {
|
||||
uuid::Uuid::new_v4().to_string()
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProxyNodeHeartbeatMutation {
|
||||
pub node_id: String,
|
||||
#[serde(default)]
|
||||
pub expected_tunnel_generation: Option<String>,
|
||||
pub heartbeat_interval: Option<i32>,
|
||||
pub active_connections: Option<i32>,
|
||||
pub total_requests_delta: Option<i64>,
|
||||
@@ -166,6 +220,11 @@ pub struct ProxyNodeHeartbeatMutation {
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProxyNodeTrafficMutation {
|
||||
pub node_id: String,
|
||||
/// Incarnation fence captured when the request plan selected this node.
|
||||
/// Missing fences are rejected by the gateway path so a stale plan cannot
|
||||
/// update a node recreated under the same id.
|
||||
#[serde(default)]
|
||||
pub expected_tunnel_generation: Option<String>,
|
||||
pub total_requests_delta: i64,
|
||||
pub failed_requests_delta: i64,
|
||||
pub dns_failures_delta: i64,
|
||||
@@ -174,6 +233,11 @@ pub struct ProxyNodeTrafficMutation {
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProxyNodeRegistrationMutation {
|
||||
/// The stable identity selected by the caller before any secret is
|
||||
/// protected. Re-registration of an existing endpoint must use its
|
||||
/// existing id; repositories reject attempts to replace it.
|
||||
#[serde(default)]
|
||||
pub node_id: Option<String>,
|
||||
pub name: String,
|
||||
pub ip: String,
|
||||
pub port: i32,
|
||||
@@ -190,8 +254,13 @@ pub struct ProxyNodeRegistrationMutation {
|
||||
pub tunnel_mode: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProxyNodeManualCreateMutation {
|
||||
/// Optional caller-selected id used to bind credentials before the row is
|
||||
/// inserted. Repositories generate one only for legacy callers that do
|
||||
/// not provide it.
|
||||
#[serde(default)]
|
||||
pub node_id: Option<String>,
|
||||
pub name: String,
|
||||
pub ip: String,
|
||||
pub port: i32,
|
||||
@@ -202,7 +271,27 @@ pub struct ProxyNodeManualCreateMutation {
|
||||
pub registered_by: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
impl std::fmt::Debug for ProxyNodeManualCreateMutation {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProxyNodeManualCreateMutation")
|
||||
.field("node_id", &self.node_id)
|
||||
.field("name", &self.name)
|
||||
.field("ip", &self.ip)
|
||||
.field("port", &self.port)
|
||||
.field("region", &self.region)
|
||||
.field("proxy_url", &"[REDACTED]")
|
||||
.field("proxy_username", &self.proxy_username)
|
||||
.field(
|
||||
"proxy_password",
|
||||
&self.proxy_password.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.field("registered_by", &self.registered_by)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProxyNodeManualUpdateMutation {
|
||||
pub node_id: String,
|
||||
pub name: Option<String>,
|
||||
@@ -214,9 +303,30 @@ pub struct ProxyNodeManualUpdateMutation {
|
||||
pub proxy_password: Option<String>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for ProxyNodeManualUpdateMutation {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("ProxyNodeManualUpdateMutation")
|
||||
.field("node_id", &self.node_id)
|
||||
.field("name", &self.name)
|
||||
.field("ip", &self.ip)
|
||||
.field("port", &self.port)
|
||||
.field("region", &self.region)
|
||||
.field("proxy_url", &self.proxy_url.as_ref().map(|_| "[REDACTED]"))
|
||||
.field("proxy_username", &self.proxy_username)
|
||||
.field(
|
||||
"proxy_password",
|
||||
&self.proxy_password.as_ref().map(|_| "[REDACTED]"),
|
||||
)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProxyNodeTunnelStatusMutation {
|
||||
pub node_id: String,
|
||||
#[serde(default)]
|
||||
pub expected_tunnel_generation: Option<String>,
|
||||
pub connected: bool,
|
||||
pub conn_count: i32,
|
||||
pub detail: Option<String>,
|
||||
@@ -226,6 +336,8 @@ pub struct ProxyNodeTunnelStatusMutation {
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProxyNodeRemoteConfigMutation {
|
||||
pub node_id: String,
|
||||
#[serde(default)]
|
||||
pub expected_tunnel_generation: Option<String>,
|
||||
pub node_name: Option<String>,
|
||||
pub allowed_ports: Option<Vec<u16>>,
|
||||
pub log_level: Option<String>,
|
||||
@@ -463,6 +575,31 @@ pub fn normalize_proxy_metadata(
|
||||
}
|
||||
}
|
||||
|
||||
pub fn normalize_heartbeat_proxy_metadata(
|
||||
previous_proxy_metadata: Option<&Value>,
|
||||
proxy_metadata: Option<&Value>,
|
||||
proxy_version: Option<&str>,
|
||||
) -> Option<Value> {
|
||||
let Some(Value::Object(mut normalized)) =
|
||||
normalize_proxy_metadata(proxy_metadata, proxy_version)
|
||||
else {
|
||||
return None;
|
||||
};
|
||||
|
||||
// Tunnel security is control-plane state. A heartbeat may refresh runtime
|
||||
// metadata, but it must never introduce or replace this trusted field.
|
||||
normalized.remove("tunnel_security");
|
||||
let merged = preserve_proxy_metadata_tunnel_security(
|
||||
previous_proxy_metadata,
|
||||
Some(Value::Object(normalized)),
|
||||
);
|
||||
merged.filter(|value| {
|
||||
value
|
||||
.as_object()
|
||||
.is_some_and(|metadata| !metadata.is_empty())
|
||||
})
|
||||
}
|
||||
|
||||
pub fn preserve_proxy_metadata_tunnel_security(
|
||||
previous_proxy_metadata: Option<&Value>,
|
||||
next_proxy_metadata: Option<Value>,
|
||||
@@ -477,9 +614,7 @@ pub fn preserve_proxy_metadata_tunnel_security(
|
||||
|
||||
match next_proxy_metadata {
|
||||
Some(Value::Object(mut metadata)) => {
|
||||
metadata
|
||||
.entry("tunnel_security".to_string())
|
||||
.or_insert(tunnel_security);
|
||||
metadata.insert("tunnel_security".to_string(), tunnel_security);
|
||||
Some(Value::Object(metadata))
|
||||
}
|
||||
Some(value) => Some(value),
|
||||
@@ -491,6 +626,70 @@ pub fn preserve_proxy_metadata_tunnel_security(
|
||||
}
|
||||
}
|
||||
|
||||
/// Merge metadata received during a trusted registration/re-registration.
|
||||
///
|
||||
/// Registration is the control-plane path that may rotate a tunnel PSK. A
|
||||
/// registration payload that omits `tunnel_security` is therefore a partial
|
||||
/// metadata refresh and must not clear the previously trusted security state.
|
||||
/// Only a non-empty, gateway-bound v2 ciphertext proves that the registration
|
||||
/// passed through the gateway credential-binding path, and that ciphertext is
|
||||
/// accepted only with the required non-TLS security mode. Mode-only,
|
||||
/// plaintext, malformed, null, scalar, empty, and disabled security values are
|
||||
/// treated as omission and cannot clear a previously trusted object.
|
||||
pub fn merge_proxy_metadata_for_registration(
|
||||
previous_proxy_metadata: Option<&Value>,
|
||||
next_proxy_metadata: Option<Value>,
|
||||
) -> Option<Value> {
|
||||
let Some(next_proxy_metadata) = next_proxy_metadata else {
|
||||
return previous_proxy_metadata.cloned();
|
||||
};
|
||||
let incoming_security_is_explicit =
|
||||
proxy_metadata_has_explicit_tunnel_security(Some(&next_proxy_metadata));
|
||||
|
||||
let Value::Object(mut metadata) = next_proxy_metadata else {
|
||||
// `normalize_proxy_metadata` normally prevents this branch. Keep a
|
||||
// malformed replacement from erasing trusted control-plane state.
|
||||
return previous_proxy_metadata.cloned();
|
||||
};
|
||||
|
||||
if incoming_security_is_explicit {
|
||||
return Some(Value::Object(metadata));
|
||||
}
|
||||
|
||||
// Null, scalar, and empty security objects are not valid replacements.
|
||||
// Remove them before restoring the previous trusted object so malformed
|
||||
// input cannot mask or downgrade the registered security policy.
|
||||
metadata.remove("tunnel_security");
|
||||
if let Some(previous_security) = previous_proxy_metadata
|
||||
.and_then(|value| value.get("tunnel_security"))
|
||||
.filter(|value| value.is_object())
|
||||
.cloned()
|
||||
{
|
||||
metadata.insert("tunnel_security".to_string(), previous_security);
|
||||
}
|
||||
|
||||
(!metadata.is_empty()).then_some(Value::Object(metadata))
|
||||
}
|
||||
|
||||
/// Return whether metadata contains a complete gateway-bound tunnel security
|
||||
/// replacement that a trusted registration may persist.
|
||||
pub fn proxy_metadata_has_explicit_tunnel_security(proxy_metadata: Option<&Value>) -> bool {
|
||||
proxy_metadata
|
||||
.and_then(Value::as_object)
|
||||
.and_then(|metadata| metadata.get("tunnel_security"))
|
||||
.and_then(Value::as_object)
|
||||
.is_some_and(|security| {
|
||||
security.get("mode").and_then(Value::as_str) == Some(TUNNEL_SECURITY_NON_TLS_REQUIRED)
|
||||
&& security
|
||||
.get("encryption_key_encrypted")
|
||||
.and_then(Value::as_str)
|
||||
.and_then(|encrypted| {
|
||||
encrypted.strip_prefix(PROXY_NODE_BOUND_TUNNEL_SECRET_PREFIX)
|
||||
})
|
||||
.is_some_and(|ciphertext| !ciphertext.is_empty())
|
||||
})
|
||||
}
|
||||
|
||||
fn extract_tunnel_metrics_counters(
|
||||
proxy_metadata: Option<&Value>,
|
||||
) -> Option<TunnelMetricsCounters> {
|
||||
@@ -712,6 +911,20 @@ pub trait ProxyNodeReadRepository: Send + Sync {
|
||||
pub trait ProxyNodeWriteRepository: Send + Sync {
|
||||
async fn reset_stale_tunnel_statuses(&self) -> Result<usize, crate::DataLayerError>;
|
||||
|
||||
async fn compare_and_set_proxy_password(
|
||||
&self,
|
||||
node_id: &str,
|
||||
expected: &str,
|
||||
replacement: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn compare_and_set_proxy_metadata(
|
||||
&self,
|
||||
node_id: &str,
|
||||
expected: &serde_json::Value,
|
||||
replacement: &serde_json::Value,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn create_manual_node(
|
||||
&self,
|
||||
mutation: &ProxyNodeManualCreateMutation,
|
||||
@@ -779,12 +992,77 @@ mod tests {
|
||||
|
||||
use super::{
|
||||
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
|
||||
merge_proxy_metadata_for_registration, normalize_heartbeat_proxy_metadata,
|
||||
normalize_proxy_node_scheduling_state, preserve_proxy_metadata_tunnel_security,
|
||||
proxy_node_accepts_new_tunnels, proxy_reported_version,
|
||||
reconcile_remote_config_after_heartbeat, remote_config_scheduling_state,
|
||||
remote_config_upgrade_target, ProxyNodeMetricsStep, StoredProxyNode,
|
||||
remote_config_upgrade_target, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
|
||||
ProxyNodeMetricsStep, StoredProxyNode,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn proxy_node_debug_output_redacts_credentials_and_untrusted_metadata() {
|
||||
let password = "debug-secret-proxy-password";
|
||||
let proxy_url = "https://user:[email protected]";
|
||||
let metadata_secret = "debug-secret-proxy-metadata";
|
||||
let mut stored = StoredProxyNode::new(
|
||||
"node-1".to_string(),
|
||||
"node".to_string(),
|
||||
"127.0.0.1".to_string(),
|
||||
8080,
|
||||
true,
|
||||
"online".to_string(),
|
||||
30,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
false,
|
||||
false,
|
||||
1,
|
||||
)
|
||||
.expect("proxy node should build")
|
||||
.with_manual_proxy_fields(
|
||||
Some(proxy_url.to_string()),
|
||||
Some("user".to_string()),
|
||||
Some(password.to_string()),
|
||||
);
|
||||
stored.proxy_metadata = Some(json!({"secret": metadata_secret}));
|
||||
let create = ProxyNodeManualCreateMutation {
|
||||
node_id: Some("node-1".to_string()),
|
||||
name: "node".to_string(),
|
||||
ip: "127.0.0.1".to_string(),
|
||||
port: 8080,
|
||||
region: None,
|
||||
proxy_url: proxy_url.to_string(),
|
||||
proxy_username: Some("user".to_string()),
|
||||
proxy_password: Some(password.to_string()),
|
||||
registered_by: None,
|
||||
};
|
||||
let update = ProxyNodeManualUpdateMutation {
|
||||
node_id: "node-1".to_string(),
|
||||
name: None,
|
||||
ip: None,
|
||||
port: None,
|
||||
region: None,
|
||||
proxy_url: Some(proxy_url.to_string()),
|
||||
proxy_username: Some("user".to_string()),
|
||||
proxy_password: Some(password.to_string()),
|
||||
};
|
||||
|
||||
for rendered in [
|
||||
format!("{stored:?}"),
|
||||
format!("{create:?}"),
|
||||
format!("{update:?}"),
|
||||
] {
|
||||
for secret in [password, proxy_url, metadata_secret] {
|
||||
assert!(!rendered.contains(secret), "Debug output leaked {secret}");
|
||||
}
|
||||
assert!(rendered.contains("[REDACTED]"));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalizes_reported_versions_and_clears_completed_upgrade_targets() {
|
||||
let remote_config = json!({
|
||||
@@ -985,7 +1263,11 @@ mod tests {
|
||||
});
|
||||
let next = json!({
|
||||
"version": "1.0.1",
|
||||
"tunnel_metrics": {"connect_successes": 1}
|
||||
"tunnel_metrics": {"connect_successes": 1},
|
||||
"tunnel_security": {
|
||||
"mode": "disabled",
|
||||
"encryption_key": "attacker-controlled"
|
||||
}
|
||||
});
|
||||
|
||||
let merged = preserve_proxy_metadata_tunnel_security(Some(&previous), Some(next))
|
||||
@@ -1008,6 +1290,202 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registration_metadata_preserves_omitted_tunnel_security() {
|
||||
let previous = 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"
|
||||
}
|
||||
});
|
||||
let next = json!({
|
||||
"version": "1.1.0",
|
||||
"tunnel_metrics": {"connect_successes": 2}
|
||||
});
|
||||
|
||||
let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next))
|
||||
.expect("registration metadata should remain present");
|
||||
assert_eq!(merged.get("version"), Some(&json!("1.1.0")));
|
||||
assert_eq!(
|
||||
merged.pointer("/tunnel_security/encryption_key_encrypted"),
|
||||
Some(&json!(
|
||||
"aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
|
||||
))
|
||||
);
|
||||
assert_eq!(
|
||||
merged.pointer("/tunnel_metrics/connect_successes"),
|
||||
Some(&json!(2))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registration_metadata_accepts_explicit_tunnel_security_rotation() {
|
||||
let previous = json!({
|
||||
"tunnel_security": {
|
||||
"mode": "non_tls_required",
|
||||
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
|
||||
}
|
||||
});
|
||||
let next = json!({
|
||||
"tunnel_security": {
|
||||
"mode": "non_tls_required",
|
||||
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new"
|
||||
}
|
||||
});
|
||||
|
||||
let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next))
|
||||
.expect("rotated registration metadata should remain present");
|
||||
assert_eq!(
|
||||
merged.pointer("/tunnel_security/encryption_key_encrypted"),
|
||||
Some(&json!(
|
||||
"aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-new"
|
||||
))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registration_metadata_rejects_invalid_security_replacement() {
|
||||
let previous = json!({
|
||||
"tunnel_security": {
|
||||
"mode": "non_tls_required",
|
||||
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
|
||||
}
|
||||
});
|
||||
let next = json!({
|
||||
"version": "1.2.0",
|
||||
"tunnel_security": null
|
||||
});
|
||||
|
||||
let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next))
|
||||
.expect("previous security should be retained");
|
||||
assert_eq!(merged.get("version"), Some(&json!("1.2.0")));
|
||||
assert_eq!(
|
||||
merged.pointer("/tunnel_security/encryption_key_encrypted"),
|
||||
Some(&json!(
|
||||
"aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
|
||||
))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registration_metadata_rejects_mode_only_security_downgrade() {
|
||||
let previous = json!({
|
||||
"tunnel_security": {
|
||||
"mode": "non_tls_required",
|
||||
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
|
||||
}
|
||||
});
|
||||
let next = json!({
|
||||
"version": "1.3.0",
|
||||
"tunnel_security": {"mode": "disabled"}
|
||||
});
|
||||
|
||||
let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next))
|
||||
.expect("mode-only security must not replace the registered credential");
|
||||
assert_eq!(merged.get("version"), Some(&json!("1.3.0")));
|
||||
assert_eq!(
|
||||
merged.pointer("/tunnel_security/encryption_key_encrypted"),
|
||||
Some(&json!(
|
||||
"aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
|
||||
))
|
||||
);
|
||||
assert_eq!(
|
||||
merged.pointer("/tunnel_security/mode"),
|
||||
Some(&json!("non_tls_required"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn registration_metadata_rejects_bound_ciphertext_with_disabled_mode() {
|
||||
let previous = json!({
|
||||
"tunnel_security": {
|
||||
"mode": "non_tls_required",
|
||||
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
|
||||
}
|
||||
});
|
||||
let next = json!({
|
||||
"version": "1.3.1",
|
||||
"tunnel_security": {
|
||||
"mode": "disabled",
|
||||
"encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-attacker"
|
||||
}
|
||||
});
|
||||
|
||||
let merged = merge_proxy_metadata_for_registration(Some(&previous), Some(next))
|
||||
.expect("disabled security must not replace the registered credential");
|
||||
assert_eq!(merged.get("version"), Some(&json!("1.3.1")));
|
||||
assert_eq!(
|
||||
merged.pointer("/tunnel_security/encryption_key_encrypted"),
|
||||
Some(&json!(
|
||||
"aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-old"
|
||||
))
|
||||
);
|
||||
assert_eq!(
|
||||
merged.pointer("/tunnel_security/mode"),
|
||||
Some(&json!("non_tls_required"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn new_registration_drops_invalid_tunnel_security_fields() {
|
||||
for invalid_security in [
|
||||
json!(null),
|
||||
json!("disabled"),
|
||||
json!({}),
|
||||
json!({"mode": "disabled"}),
|
||||
json!({"mode": "disabled", "encryption_key_encrypted": "aether-proxy-node-secret-v2:aether-runtime-secret-v1:sealed-attacker"}),
|
||||
json!({"mode": "non_tls_required", "encryption_key_encrypted": ""}),
|
||||
json!({"mode": "non_tls_required", "encryption_key_encrypted": "not-gateway-bound"}),
|
||||
] {
|
||||
let next = json!({
|
||||
"version": "1.4.0",
|
||||
"tunnel_security": invalid_security
|
||||
});
|
||||
let merged = merge_proxy_metadata_for_registration(None, Some(next))
|
||||
.expect("valid non-security metadata should remain");
|
||||
assert_eq!(merged.get("version"), Some(&json!("1.4.0")));
|
||||
assert!(merged.get("tunnel_security").is_none());
|
||||
}
|
||||
|
||||
assert_eq!(
|
||||
merge_proxy_metadata_for_registration(
|
||||
None,
|
||||
Some(json!({"tunnel_security": {"mode": "disabled"}})),
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn heartbeat_metadata_cannot_introduce_tunnel_security() {
|
||||
let injected = json!({
|
||||
"tunnel_security": {
|
||||
"mode": "disabled",
|
||||
"encryption_key": "attacker-controlled"
|
||||
}
|
||||
});
|
||||
assert_eq!(
|
||||
normalize_heartbeat_proxy_metadata(None, Some(&injected), None),
|
||||
None
|
||||
);
|
||||
|
||||
let injected_with_runtime_metadata = json!({
|
||||
"version": "1.2.3",
|
||||
"arch": "arm64",
|
||||
"tunnel_security": {
|
||||
"mode": "disabled",
|
||||
"encryption_key": "attacker-controlled"
|
||||
}
|
||||
});
|
||||
let normalized =
|
||||
normalize_heartbeat_proxy_metadata(None, Some(&injected_with_runtime_metadata), None)
|
||||
.expect("runtime metadata should remain present");
|
||||
assert_eq!(normalized.get("version"), Some(&json!("1.2.3")));
|
||||
assert_eq!(normalized.get("arch"), Some(&json!("arm64")));
|
||||
assert!(normalized.get("tunnel_security").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn maps_timestamps_to_metric_buckets() {
|
||||
assert_eq!(
|
||||
|
||||
@@ -2,6 +2,11 @@ mod types;
|
||||
|
||||
pub use types::{
|
||||
finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd,
|
||||
settlement_billing_status_for_usage_status, SettlementRepository, SettlementWriteRepository,
|
||||
StoredUsageSettlement, UsageSettlementInput, WalletDebitPlan, SETTLEMENT_EPSILON_USD,
|
||||
settlement_billing_status_for_usage_status, validate_wallet_settlement_values,
|
||||
ReconcileUsagePolicyCostInput, ReleaseUsagePolicyRequestAdmissionInput,
|
||||
ReserveUsagePolicyCostInput, ReserveUsagePolicyCostOutcome, ReserveUsagePolicyRequestInput,
|
||||
ReserveUsagePolicyRequestOutcome, SettlementRepository, SettlementWriteRepository,
|
||||
StoredUsagePolicyCostReservation, StoredUsagePolicyRequestAdmission, StoredUsageSettlement,
|
||||
UsagePolicyCostReservationState, UsagePolicyCostWindow, UsagePolicyRequestAdmissionState,
|
||||
UsagePolicyRequestWindow, UsageSettlementInput, WalletDebitPlan, SETTLEMENT_EPSILON_USD,
|
||||
};
|
||||
|
||||
@@ -1,5 +1,339 @@
|
||||
use async_trait::async_trait;
|
||||
|
||||
use crate::repository::billing::MAX_USAGE_POLICY_TOTAL_RULES;
|
||||
|
||||
const MAX_USAGE_POLICY_LEDGER_ID_BYTES: usize = 128;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct UsagePolicyRequestWindow {
|
||||
pub starts_at_unix_secs: u64,
|
||||
pub ends_at_unix_secs: u64,
|
||||
pub limit_requests: u64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ReserveUsagePolicyRequestInput {
|
||||
pub request_id: String,
|
||||
pub subject_id: String,
|
||||
pub event_token: String,
|
||||
pub admitted_at_unix_secs: u64,
|
||||
/// Exclusive timestamp after which this admission cannot affect any supplied window and its
|
||||
/// idempotency tombstone may be deleted safely.
|
||||
pub retain_until_unix_secs: u64,
|
||||
pub windows: Vec<UsagePolicyRequestWindow>,
|
||||
}
|
||||
|
||||
impl ReserveUsagePolicyRequestInput {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
validate_bounded_id(&self.request_id, "usage policy request_id")?;
|
||||
validate_bounded_id(&self.subject_id, "usage policy subject_id")?;
|
||||
validate_bounded_id(&self.event_token, "usage policy event_token")?;
|
||||
validate_database_u64(self.admitted_at_unix_secs, "usage policy admitted_at")?;
|
||||
validate_database_u64(self.retain_until_unix_secs, "usage policy retain_until")?;
|
||||
if self.windows.is_empty() || self.windows.len() > MAX_USAGE_POLICY_TOTAL_RULES {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"usage policy request admission requires 1 to {MAX_USAGE_POLICY_TOTAL_RULES} windows"
|
||||
)));
|
||||
}
|
||||
for (index, window) in self.windows.iter().enumerate() {
|
||||
validate_database_u64(
|
||||
window.starts_at_unix_secs,
|
||||
&format!("usage policy request window {index} start"),
|
||||
)?;
|
||||
validate_database_u64(
|
||||
window.ends_at_unix_secs,
|
||||
&format!("usage policy request window {index} end"),
|
||||
)?;
|
||||
validate_database_u64(
|
||||
window.limit_requests,
|
||||
&format!("usage policy request window {index} limit"),
|
||||
)?;
|
||||
if window.limit_requests == 0
|
||||
|| window.starts_at_unix_secs >= window.ends_at_unix_secs
|
||||
|| self.admitted_at_unix_secs < window.starts_at_unix_secs
|
||||
|| self.admitted_at_unix_secs >= window.ends_at_unix_secs
|
||||
{
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"usage policy request window {index} does not contain admission or has invalid bounds"
|
||||
)));
|
||||
}
|
||||
if self.retain_until_unix_secs < window.ends_at_unix_secs {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"usage policy request retain_until precedes window {index} end"
|
||||
)));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum UsagePolicyRequestAdmissionState {
|
||||
Active,
|
||||
Released,
|
||||
}
|
||||
|
||||
impl UsagePolicyRequestAdmissionState {
|
||||
pub const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Active => "active",
|
||||
Self::Released => "released",
|
||||
}
|
||||
}
|
||||
|
||||
pub fn parse(value: &str) -> Option<Self> {
|
||||
match value {
|
||||
"active" => Some(Self::Active),
|
||||
"released" => Some(Self::Released),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
#[serde(tag = "status", rename_all = "snake_case")]
|
||||
pub enum ReserveUsagePolicyRequestOutcome {
|
||||
Allowed,
|
||||
Rejected {
|
||||
window_index: usize,
|
||||
limit_requests: u64,
|
||||
used_requests: u64,
|
||||
},
|
||||
AlreadyReleased,
|
||||
Conflict,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ReleaseUsagePolicyRequestAdmissionInput {
|
||||
pub request_id: String,
|
||||
pub subject_id: String,
|
||||
pub event_token: String,
|
||||
pub released_at_unix_secs: u64,
|
||||
}
|
||||
|
||||
impl ReleaseUsagePolicyRequestAdmissionInput {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
validate_bounded_id(&self.request_id, "usage policy request_id")?;
|
||||
validate_bounded_id(&self.subject_id, "usage policy subject_id")?;
|
||||
validate_bounded_id(&self.event_token, "usage policy event_token")?;
|
||||
validate_database_u64(self.released_at_unix_secs, "usage policy released_at")
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredUsagePolicyRequestAdmission {
|
||||
pub request_id: String,
|
||||
pub subject_id: String,
|
||||
pub event_token: String,
|
||||
pub admitted_at_unix_secs: u64,
|
||||
pub retain_until_unix_secs: u64,
|
||||
pub state: UsagePolicyRequestAdmissionState,
|
||||
pub released_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct UsagePolicyCostWindow {
|
||||
pub window_id: String,
|
||||
pub starts_at_unix_secs: u64,
|
||||
pub ends_at_unix_secs: u64,
|
||||
pub limit_cost_units: u64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ReserveUsagePolicyCostInput {
|
||||
pub request_id: String,
|
||||
pub subject_id: String,
|
||||
pub reservation_token: String,
|
||||
pub admitted_at_unix_secs: u64,
|
||||
pub reserved_cost_units: u64,
|
||||
pub reservation_expires_at_unix_secs: u64,
|
||||
/// Exclusive timestamp after which this reservation can no longer affect any future window
|
||||
/// and its idempotency tombstone may be deleted safely.
|
||||
pub retain_until_unix_secs: u64,
|
||||
pub windows: Vec<UsagePolicyCostWindow>,
|
||||
}
|
||||
|
||||
impl ReserveUsagePolicyCostInput {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
validate_bounded_id(&self.request_id, "usage policy request_id")?;
|
||||
validate_bounded_id(&self.subject_id, "usage policy subject_id")?;
|
||||
validate_bounded_id(&self.reservation_token, "usage policy reservation_token")?;
|
||||
validate_cost_units(self.reserved_cost_units, "reserved_cost_units")?;
|
||||
if self.reservation_expires_at_unix_secs <= self.admitted_at_unix_secs {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage policy reservation must expire after admission".to_string(),
|
||||
));
|
||||
}
|
||||
if self.retain_until_unix_secs < self.reservation_expires_at_unix_secs {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage policy retain_until must not precede reservation expiry".to_string(),
|
||||
));
|
||||
}
|
||||
if self.windows.is_empty() || self.windows.len() > MAX_USAGE_POLICY_TOTAL_RULES {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"usage policy reservation requires 1 to {MAX_USAGE_POLICY_TOTAL_RULES} windows"
|
||||
)));
|
||||
}
|
||||
for (index, window) in self.windows.iter().enumerate() {
|
||||
validate_non_empty_id(
|
||||
&window.window_id,
|
||||
&format!("usage policy window {index} id"),
|
||||
)?;
|
||||
validate_cost_units(
|
||||
window.limit_cost_units,
|
||||
&format!("usage policy window {index} limit_cost_units"),
|
||||
)?;
|
||||
if window.limit_cost_units == 0
|
||||
|| window.starts_at_unix_secs >= window.ends_at_unix_secs
|
||||
|| self.admitted_at_unix_secs < window.starts_at_unix_secs
|
||||
|| self.admitted_at_unix_secs >= window.ends_at_unix_secs
|
||||
{
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"usage policy window {index} does not contain admission or has invalid bounds"
|
||||
)));
|
||||
}
|
||||
if self.windows[..index]
|
||||
.iter()
|
||||
.any(|previous| previous.window_id == window.window_id)
|
||||
{
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"usage policy window {index} duplicates window_id {}",
|
||||
window.window_id
|
||||
)));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum UsagePolicyCostReservationState {
|
||||
Reserved,
|
||||
Finalized,
|
||||
Released,
|
||||
}
|
||||
|
||||
impl UsagePolicyCostReservationState {
|
||||
pub const fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::Reserved => "reserved",
|
||||
Self::Finalized => "finalized",
|
||||
Self::Released => "released",
|
||||
}
|
||||
}
|
||||
|
||||
pub fn parse(value: &str) -> Option<Self> {
|
||||
match value {
|
||||
"reserved" => Some(Self::Reserved),
|
||||
"finalized" => Some(Self::Finalized),
|
||||
"released" => Some(Self::Released),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
#[serde(tag = "status", rename_all = "snake_case")]
|
||||
pub enum ReserveUsagePolicyCostOutcome {
|
||||
Allowed {
|
||||
reserved_cost_units: u64,
|
||||
additional_reserved_cost_units: u64,
|
||||
},
|
||||
Rejected {
|
||||
window_index: usize,
|
||||
limit_cost_units: u64,
|
||||
used_cost_units: u64,
|
||||
},
|
||||
AlreadyTerminal {
|
||||
state: UsagePolicyCostReservationState,
|
||||
},
|
||||
Conflict,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ReconcileUsagePolicyCostInput {
|
||||
pub request_id: String,
|
||||
pub subject_id: String,
|
||||
pub reservation_token: String,
|
||||
pub actual_cost_units: u64,
|
||||
pub terminal_state: UsagePolicyCostReservationState,
|
||||
pub finalized_at_unix_secs: u64,
|
||||
}
|
||||
|
||||
impl ReconcileUsagePolicyCostInput {
|
||||
pub fn validate(&self) -> Result<(), crate::DataLayerError> {
|
||||
validate_bounded_id(&self.request_id, "usage policy request_id")?;
|
||||
validate_bounded_id(&self.subject_id, "usage policy subject_id")?;
|
||||
validate_bounded_id(&self.reservation_token, "usage policy reservation_token")?;
|
||||
validate_cost_units(self.actual_cost_units, "actual_cost_units")?;
|
||||
match self.terminal_state {
|
||||
UsagePolicyCostReservationState::Reserved => Err(crate::DataLayerError::InvalidInput(
|
||||
"usage policy reconciliation requires a terminal state".to_string(),
|
||||
)),
|
||||
UsagePolicyCostReservationState::Released if self.actual_cost_units != 0 => {
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"released usage policy reservations must have zero actual cost".to_string(),
|
||||
))
|
||||
}
|
||||
UsagePolicyCostReservationState::Finalized
|
||||
| UsagePolicyCostReservationState::Released => Ok(()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredUsagePolicyCostReservation {
|
||||
pub request_id: String,
|
||||
pub subject_id: String,
|
||||
pub reservation_token: String,
|
||||
pub admitted_at_unix_secs: u64,
|
||||
pub reserved_cost_units: u64,
|
||||
pub actual_cost_units: Option<u64>,
|
||||
pub state: UsagePolicyCostReservationState,
|
||||
pub reservation_expires_at_unix_secs: u64,
|
||||
pub retain_until_unix_secs: u64,
|
||||
pub finalized_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
fn validate_non_empty_id(value: &str, field: &str) -> Result<(), crate::DataLayerError> {
|
||||
if value.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"{field} must not be empty"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_bounded_id(value: &str, field: &str) -> Result<(), crate::DataLayerError> {
|
||||
validate_non_empty_id(value, field)?;
|
||||
if value.len() > MAX_USAGE_POLICY_LEDGER_ID_BYTES {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"{field} exceeds {MAX_USAGE_POLICY_LEDGER_ID_BYTES} bytes"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_database_u64(value: u64, field: &str) -> Result<(), crate::DataLayerError> {
|
||||
if value > i64::MAX as u64 {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"{field} exceeds the database integer range"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn validate_cost_units(value: u64, field: &str) -> Result<(), crate::DataLayerError> {
|
||||
if value > i64::MAX as u64 {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"{field} exceeds the database integer range"
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct UsageSettlementInput {
|
||||
pub request_id: String,
|
||||
@@ -53,6 +387,44 @@ pub struct StoredUsageSettlement {
|
||||
|
||||
#[async_trait]
|
||||
pub trait SettlementWriteRepository: Send + Sync {
|
||||
async fn reserve_usage_policy_request(
|
||||
&self,
|
||||
input: ReserveUsagePolicyRequestInput,
|
||||
) -> Result<ReserveUsagePolicyRequestOutcome, crate::DataLayerError>;
|
||||
|
||||
async fn release_usage_policy_request_admission(
|
||||
&self,
|
||||
input: ReleaseUsagePolicyRequestAdmissionInput,
|
||||
) -> Result<Option<StoredUsagePolicyRequestAdmission>, crate::DataLayerError>;
|
||||
|
||||
async fn cleanup_usage_policy_request_admissions(
|
||||
&self,
|
||||
now_unix_secs: u64,
|
||||
batch_size: usize,
|
||||
) -> Result<usize, crate::DataLayerError> {
|
||||
let _ = (now_unix_secs, batch_size);
|
||||
Ok(0)
|
||||
}
|
||||
|
||||
async fn reserve_usage_policy_cost(
|
||||
&self,
|
||||
input: ReserveUsagePolicyCostInput,
|
||||
) -> Result<ReserveUsagePolicyCostOutcome, crate::DataLayerError>;
|
||||
|
||||
async fn reconcile_usage_policy_cost(
|
||||
&self,
|
||||
input: ReconcileUsagePolicyCostInput,
|
||||
) -> Result<Option<StoredUsagePolicyCostReservation>, crate::DataLayerError>;
|
||||
|
||||
async fn cleanup_usage_policy_cost_reservations(
|
||||
&self,
|
||||
now_unix_secs: u64,
|
||||
batch_size: usize,
|
||||
) -> Result<usize, crate::DataLayerError> {
|
||||
let _ = (now_unix_secs, batch_size);
|
||||
Ok(0)
|
||||
}
|
||||
|
||||
async fn settle_usage(
|
||||
&self,
|
||||
input: UsageSettlementInput,
|
||||
@@ -85,6 +457,35 @@ pub fn finite_wallet_available_usd(recharge_balance: f64, gift_balance: f64) ->
|
||||
recharge_balance.max(0.0) + gift_balance.max(0.0)
|
||||
}
|
||||
|
||||
/// Reject corrupted or overflowing financial values before usage settlement mutates a wallet.
|
||||
/// Recharge balances may legitimately be negative after an admitted request settles, but gift
|
||||
/// balances and cumulative consumption may not be negative. Derived totals must remain finite so
|
||||
/// persisted `NaN`/infinity values cannot turn a finite wallet into an implicit unlimited wallet.
|
||||
pub fn validate_wallet_settlement_values(
|
||||
recharge_balance: f64,
|
||||
gift_balance: f64,
|
||||
total_consumed: f64,
|
||||
additional_consumed: f64,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
let balance_total = recharge_balance + gift_balance;
|
||||
let consumed_after = total_consumed + additional_consumed;
|
||||
if !recharge_balance.is_finite()
|
||||
|| !gift_balance.is_finite()
|
||||
|| gift_balance < 0.0
|
||||
|| !balance_total.is_finite()
|
||||
|| !total_consumed.is_finite()
|
||||
|| total_consumed < 0.0
|
||||
|| !additional_consumed.is_finite()
|
||||
|| additional_consumed < 0.0
|
||||
|| !consumed_after.is_finite()
|
||||
{
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"wallet financial state is invalid for usage settlement".to_string(),
|
||||
));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn plan_finite_wallet_debit(
|
||||
recharge_balance: f64,
|
||||
gift_balance: f64,
|
||||
@@ -115,7 +516,12 @@ pub fn settlement_billable_cost_usd(input: &UsageSettlementInput) -> f64 {
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::UsageSettlementInput;
|
||||
use super::{
|
||||
validate_wallet_settlement_values, ReconcileUsagePolicyCostInput,
|
||||
ReserveUsagePolicyCostInput, ReserveUsagePolicyRequestInput,
|
||||
UsagePolicyCostReservationState, UsagePolicyCostWindow, UsagePolicyRequestWindow,
|
||||
UsageSettlementInput,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_settlement_input() {
|
||||
@@ -133,4 +539,101 @@ mod tests {
|
||||
};
|
||||
assert!(input.validate().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn wallet_settlement_values_reject_corruption_and_overflow() {
|
||||
assert!(validate_wallet_settlement_values(-3.0, 0.0, 12.0, 1.0).is_ok());
|
||||
for (recharge, gift, consumed, additional) in [
|
||||
(f64::NAN, 0.0, 0.0, 1.0),
|
||||
(f64::INFINITY, 0.0, 0.0, 1.0),
|
||||
(0.0, f64::NAN, 0.0, 1.0),
|
||||
(0.0, -0.01, 0.0, 1.0),
|
||||
(0.0, 0.0, f64::INFINITY, 1.0),
|
||||
(0.0, 0.0, -0.01, 1.0),
|
||||
(0.0, 0.0, 0.0, f64::NAN),
|
||||
(f64::MAX, f64::MAX, 0.0, 1.0),
|
||||
(0.0, 0.0, f64::MAX, f64::MAX),
|
||||
] {
|
||||
assert!(
|
||||
validate_wallet_settlement_values(recharge, gift, consumed, additional).is_err()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validates_request_admission_window_and_retention_bounds() {
|
||||
let valid = ReserveUsagePolicyRequestInput {
|
||||
request_id: "request-1".to_string(),
|
||||
subject_id: "user-1".to_string(),
|
||||
event_token: "event-1".to_string(),
|
||||
admitted_at_unix_secs: 100,
|
||||
retain_until_unix_secs: 200,
|
||||
windows: vec![UsagePolicyRequestWindow {
|
||||
starts_at_unix_secs: 50,
|
||||
ends_at_unix_secs: 200,
|
||||
limit_requests: 10,
|
||||
}],
|
||||
};
|
||||
assert!(valid.validate().is_ok());
|
||||
|
||||
let mut invalid_retention = valid.clone();
|
||||
invalid_retention.retain_until_unix_secs = 199;
|
||||
assert!(invalid_retention.validate().is_err());
|
||||
|
||||
let mut invalid_window = valid;
|
||||
invalid_window.windows[0].starts_at_unix_secs = 101;
|
||||
assert!(invalid_window.validate().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_policy_cost_ids_match_sql_column_bounds() {
|
||||
let bounded = "x".repeat(128);
|
||||
let reserve = ReserveUsagePolicyCostInput {
|
||||
request_id: bounded.clone(),
|
||||
subject_id: bounded.clone(),
|
||||
reservation_token: bounded.clone(),
|
||||
admitted_at_unix_secs: 100,
|
||||
reserved_cost_units: 1,
|
||||
reservation_expires_at_unix_secs: 150,
|
||||
retain_until_unix_secs: 200,
|
||||
windows: vec![UsagePolicyCostWindow {
|
||||
window_id: "window-1".to_string(),
|
||||
starts_at_unix_secs: 50,
|
||||
ends_at_unix_secs: 200,
|
||||
limit_cost_units: 10,
|
||||
}],
|
||||
};
|
||||
assert!(reserve.validate().is_ok());
|
||||
|
||||
for field in ["request_id", "subject_id", "reservation_token"] {
|
||||
let mut too_long = reserve.clone();
|
||||
match field {
|
||||
"request_id" => too_long.request_id.push('x'),
|
||||
"subject_id" => too_long.subject_id.push('x'),
|
||||
"reservation_token" => too_long.reservation_token.push('x'),
|
||||
_ => unreachable!(),
|
||||
}
|
||||
assert!(too_long.validate().is_err(), "{field} must be bounded");
|
||||
}
|
||||
|
||||
let reconcile = ReconcileUsagePolicyCostInput {
|
||||
request_id: bounded.clone(),
|
||||
subject_id: bounded.clone(),
|
||||
reservation_token: bounded,
|
||||
actual_cost_units: 1,
|
||||
terminal_state: UsagePolicyCostReservationState::Finalized,
|
||||
finalized_at_unix_secs: 200,
|
||||
};
|
||||
assert!(reconcile.validate().is_ok());
|
||||
for field in ["request_id", "subject_id", "reservation_token"] {
|
||||
let mut too_long = reconcile.clone();
|
||||
match field {
|
||||
"request_id" => too_long.request_id.push('x'),
|
||||
"subject_id" => too_long.subject_id.push('x'),
|
||||
"reservation_token" => too_long.reservation_token.push('x'),
|
||||
_ => unreachable!(),
|
||||
}
|
||||
assert!(too_long.validate().is_err(), "{field} must be bounded");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,58 @@
|
||||
use std::io::Read;
|
||||
|
||||
use crate::DataLayerError;
|
||||
|
||||
/// Hard ceiling for usage JSON after decompression.
|
||||
///
|
||||
/// Usage bodies may contain large model responses, so this stays aligned with the gateway's
|
||||
/// largest routinely buffered response while still bounding gzip expansion from stored data.
|
||||
pub const MAX_DECOMPRESSED_USAGE_JSON_BYTES: usize = 64 * 1024 * 1024;
|
||||
|
||||
pub fn read_decompressed_usage_json(reader: impl Read) -> Result<Vec<u8>, DataLayerError> {
|
||||
read_decompressed_usage_json_with_limit(reader, MAX_DECOMPRESSED_USAGE_JSON_BYTES)
|
||||
}
|
||||
|
||||
fn read_decompressed_usage_json_with_limit(
|
||||
reader: impl Read,
|
||||
limit_bytes: usize,
|
||||
) -> Result<Vec<u8>, DataLayerError> {
|
||||
let read_limit = u64::try_from(limit_bytes)
|
||||
.unwrap_or(u64::MAX)
|
||||
.saturating_add(1);
|
||||
let mut limited = reader.take(read_limit);
|
||||
let mut decoded = Vec::new();
|
||||
limited.read_to_end(&mut decoded).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("failed to decompress usage json: {err}"))
|
||||
})?;
|
||||
if decoded.len() > limit_bytes {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"decompressed usage json exceeds {limit_bytes} bytes"
|
||||
)));
|
||||
}
|
||||
Ok(decoded)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::io::Cursor;
|
||||
|
||||
use super::read_decompressed_usage_json_with_limit;
|
||||
|
||||
#[test]
|
||||
fn decompressed_usage_json_reader_accepts_exact_limit() {
|
||||
let decoded = read_decompressed_usage_json_with_limit(Cursor::new(b"1234"), 4)
|
||||
.expect("payload at the hard limit should decode");
|
||||
|
||||
assert_eq!(decoded, b"1234");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decompressed_usage_json_reader_rejects_limit_plus_one() {
|
||||
let error = read_decompressed_usage_json_with_limit(Cursor::new(b"12345"), 4)
|
||||
.expect_err("payload over the hard limit should fail");
|
||||
|
||||
assert!(error
|
||||
.to_string()
|
||||
.contains("decompressed usage json exceeds 4 bytes"));
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,9 +1,13 @@
|
||||
mod compression;
|
||||
mod metadata_policy;
|
||||
mod policy;
|
||||
mod types;
|
||||
|
||||
pub use compression::{read_decompressed_usage_json, MAX_DECOMPRESSED_USAGE_JSON_BYTES};
|
||||
pub use metadata_policy::*;
|
||||
pub use policy::*;
|
||||
pub use types::{
|
||||
extract_provider_actual_service_tier_from_response,
|
||||
canonical_usage_body_ref_for, extract_provider_actual_service_tier_from_response,
|
||||
extract_provider_cache_ttl_minutes_from_metadata, extract_provider_reasoning_effort_from_body,
|
||||
extract_provider_service_tier_from_body, normalize_provider_service_tier, parse_usage_body_ref,
|
||||
resolve_provider_cache_ttl_minutes, resolve_provider_service_tier_from_request_capture,
|
||||
@@ -34,10 +38,11 @@ pub use types::{
|
||||
UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery,
|
||||
UsageReadRepository, UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
|
||||
UsageTimeSeriesQuery, UsageWriteRepository, LIVE_SESSION_METADATA_KEY,
|
||||
PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY, PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY,
|
||||
PROVIDER_REASONING_EFFORT_METADATA_KEY, PROVIDER_SERVICE_TIER_METADATA_KEY,
|
||||
REALTIME_SESSION_METADATA_KEY, REQUESTED_REASONING_EFFORT_METADATA_KEY,
|
||||
ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY, ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY,
|
||||
USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY,
|
||||
WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY,
|
||||
PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY, PROVIDER_ACTUAL_SERVICE_TIER_METADATA_KEY,
|
||||
PROVIDER_CACHE_TTL_MINUTES_METADATA_KEY, PROVIDER_REASONING_EFFORT_METADATA_KEY,
|
||||
PROVIDER_SERVICE_TIER_METADATA_KEY, REALTIME_SESSION_METADATA_KEY,
|
||||
REQUESTED_REASONING_EFFORT_METADATA_KEY, ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY,
|
||||
ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY, USAGE_AVAILABLE_METADATA_KEY,
|
||||
USAGE_PRICING_AVAILABLE_METADATA_KEY, WEBSOCKET_MODE_METADATA_KEY,
|
||||
WEBSOCKET_TRANSPORT_METADATA_KEY,
|
||||
};
|
||||
|
||||
@@ -1,4 +1,14 @@
|
||||
use super::{StoredRequestUsageAudit, UpsertUsageRecord};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
const MAX_USAGE_CANDIDATE_INDEX: u64 = i32::MAX as u64;
|
||||
const MAX_USAGE_CANDIDATE_ID_LEN: usize = 128;
|
||||
const MAX_USAGE_KEY_NAME_LEN: usize = 255;
|
||||
const MAX_USAGE_PLANNER_KIND_LEN: usize = 64;
|
||||
const MAX_USAGE_ROUTE_FAMILY_LEN: usize = 80;
|
||||
const MAX_USAGE_ROUTE_KIND_LEN: usize = 80;
|
||||
const MAX_USAGE_EXECUTION_PATH_LEN: usize = 80;
|
||||
const MAX_USAGE_RUNTIME_MISS_REASON_LEN: usize = 120;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Default)]
|
||||
pub struct ApiKeyUsageContribution {
|
||||
@@ -201,12 +211,289 @@ pub fn usage_can_recover_terminal_failure(
|
||||
&& incoming_usage_can_recover_terminal_failure(incoming_status, incoming_billing_status)
|
||||
}
|
||||
|
||||
/// Decide whether an incoming lifecycle event may replace an existing usage revision.
|
||||
///
|
||||
/// `updated_at_unix_secs` is authoritative. `finalized_at_unix_secs` breaks ties when writers
|
||||
/// observe multiple transitions in the same second. Equal pending revisions may still progress to
|
||||
/// streaming or terminal states, while every terminal replay requires a strictly newer revision.
|
||||
/// The explicit void-failure recovery remains available at an equal revision, but never for an
|
||||
/// older event.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn usage_lifecycle_update_allowed(
|
||||
existing_status: &str,
|
||||
existing_billing_status: &str,
|
||||
existing_updated_at_unix_secs: u64,
|
||||
existing_finalized_at_unix_secs: Option<u64>,
|
||||
incoming_status: &str,
|
||||
incoming_billing_status: &str,
|
||||
incoming_updated_at_unix_secs: u64,
|
||||
incoming_finalized_at_unix_secs: Option<u64>,
|
||||
) -> bool {
|
||||
let existing_revision = (
|
||||
existing_updated_at_unix_secs,
|
||||
existing_finalized_at_unix_secs.unwrap_or_default(),
|
||||
);
|
||||
let incoming_revision = (
|
||||
incoming_updated_at_unix_secs,
|
||||
incoming_finalized_at_unix_secs.unwrap_or_default(),
|
||||
);
|
||||
if incoming_revision < existing_revision {
|
||||
return false;
|
||||
}
|
||||
|
||||
let can_recover = usage_can_recover_terminal_failure(
|
||||
existing_status,
|
||||
existing_billing_status,
|
||||
incoming_status,
|
||||
incoming_billing_status,
|
||||
);
|
||||
let existing_is_terminal = matches!(existing_status, "completed" | "failed" | "cancelled");
|
||||
let incoming_is_terminal = matches!(incoming_status, "completed" | "failed" | "cancelled");
|
||||
if existing_is_terminal && !incoming_is_terminal {
|
||||
return false;
|
||||
}
|
||||
if existing_status == "streaming" && incoming_status == "pending" {
|
||||
return false;
|
||||
}
|
||||
if incoming_revision == existing_revision && existing_is_terminal && incoming_is_terminal {
|
||||
return can_recover;
|
||||
}
|
||||
|
||||
true
|
||||
}
|
||||
|
||||
pub fn strip_deprecated_usage_display_fields(mut usage: UpsertUsageRecord) -> UpsertUsageRecord {
|
||||
usage.username = None;
|
||||
usage.api_key_name = None;
|
||||
usage
|
||||
}
|
||||
|
||||
pub fn sanitize_usage_for_persistence(mut usage: UpsertUsageRecord) -> UpsertUsageRecord {
|
||||
usage = strip_deprecated_usage_display_fields(usage);
|
||||
sanitize_usage_routing_fields(&mut usage, None);
|
||||
usage.error_message = None;
|
||||
usage.error_category = sanitize_usage_error_category(usage.error_category);
|
||||
if usage.error_category.is_none() && usage.status == "failed" {
|
||||
usage.error_category = usage
|
||||
.status_code
|
||||
.map(usage_error_category_for_status_code)
|
||||
.map(str::to_string);
|
||||
}
|
||||
usage.request_metadata = super::sanitize_usage_request_metadata(usage.request_metadata);
|
||||
usage.request_headers = None;
|
||||
usage.request_body = None;
|
||||
usage.request_body_ref = None;
|
||||
usage.request_body_state = None;
|
||||
usage.provider_request_headers = None;
|
||||
usage.provider_request_body = None;
|
||||
usage.provider_request_body_ref = None;
|
||||
usage.provider_request_body_state = None;
|
||||
usage.response_headers = None;
|
||||
usage.response_body = None;
|
||||
usage.response_body_ref = None;
|
||||
usage.response_body_state = None;
|
||||
usage.client_response_headers = None;
|
||||
usage.client_response_body = None;
|
||||
usage.client_response_body_ref = None;
|
||||
usage.client_response_body_state = None;
|
||||
usage
|
||||
}
|
||||
|
||||
/// Project an event onto the non-content controls accepted by auxiliary usage storage.
|
||||
///
|
||||
/// Explicit `none` states are retained only as tombstones for removing historical captures.
|
||||
/// Every header, body, reference, and non-clear capture state is discarded.
|
||||
pub fn sanitize_usage_capture_controls_for_persistence(
|
||||
mut usage: UpsertUsageRecord,
|
||||
) -> UpsertUsageRecord {
|
||||
// Routing facts are allowed in the transient event metadata for compatibility with older
|
||||
// writers. Project only the known scalar fields into typed slots before the general metadata
|
||||
// sanitizer drops unknown keys. This keeps snapshots useful without re-persisting arbitrary
|
||||
// metadata (or any body/header material).
|
||||
let metadata = usage
|
||||
.request_metadata
|
||||
.as_ref()
|
||||
.and_then(Value::as_object)
|
||||
.cloned();
|
||||
sanitize_usage_routing_fields(&mut usage, metadata.as_ref());
|
||||
let clear_request_body = usage.request_body_state == Some(super::UsageBodyCaptureState::None);
|
||||
let clear_provider_request_body =
|
||||
usage.provider_request_body_state == Some(super::UsageBodyCaptureState::None);
|
||||
let clear_response_body = usage.response_body_state == Some(super::UsageBodyCaptureState::None);
|
||||
let clear_client_response_body =
|
||||
usage.client_response_body_state == Some(super::UsageBodyCaptureState::None);
|
||||
|
||||
let mut usage = sanitize_usage_for_persistence(usage);
|
||||
usage.request_body_state = clear_request_body.then_some(super::UsageBodyCaptureState::None);
|
||||
usage.provider_request_body_state =
|
||||
clear_provider_request_body.then_some(super::UsageBodyCaptureState::None);
|
||||
usage.response_body_state = clear_response_body.then_some(super::UsageBodyCaptureState::None);
|
||||
usage.client_response_body_state =
|
||||
clear_client_response_body.then_some(super::UsageBodyCaptureState::None);
|
||||
usage
|
||||
}
|
||||
|
||||
fn sanitize_usage_routing_fields(
|
||||
usage: &mut UpsertUsageRecord,
|
||||
metadata: Option<&Map<String, Value>>,
|
||||
) {
|
||||
usage.candidate_id = sanitize_usage_routing_string_with_metadata(
|
||||
usage.candidate_id.take(),
|
||||
metadata,
|
||||
"candidate_id",
|
||||
MAX_USAGE_CANDIDATE_ID_LEN,
|
||||
false,
|
||||
);
|
||||
usage.candidate_index =
|
||||
sanitize_usage_routing_index_with_metadata(usage.candidate_index.take(), metadata);
|
||||
usage.key_name = sanitize_usage_routing_string_with_metadata(
|
||||
usage.key_name.take(),
|
||||
metadata,
|
||||
"key_name",
|
||||
MAX_USAGE_KEY_NAME_LEN,
|
||||
true,
|
||||
);
|
||||
usage.planner_kind = sanitize_usage_routing_string_with_metadata(
|
||||
usage.planner_kind.take(),
|
||||
metadata,
|
||||
"planner_kind",
|
||||
MAX_USAGE_PLANNER_KIND_LEN,
|
||||
false,
|
||||
);
|
||||
usage.route_family = sanitize_usage_routing_string_with_metadata(
|
||||
usage.route_family.take(),
|
||||
metadata,
|
||||
"route_family",
|
||||
MAX_USAGE_ROUTE_FAMILY_LEN,
|
||||
false,
|
||||
);
|
||||
usage.route_kind = sanitize_usage_routing_string_with_metadata(
|
||||
usage.route_kind.take(),
|
||||
metadata,
|
||||
"route_kind",
|
||||
MAX_USAGE_ROUTE_KIND_LEN,
|
||||
false,
|
||||
);
|
||||
usage.execution_path = sanitize_usage_routing_string_with_metadata(
|
||||
usage.execution_path.take(),
|
||||
metadata,
|
||||
"execution_path",
|
||||
MAX_USAGE_EXECUTION_PATH_LEN,
|
||||
false,
|
||||
);
|
||||
usage.local_execution_runtime_miss_reason = sanitize_usage_routing_string_with_metadata(
|
||||
usage.local_execution_runtime_miss_reason.take(),
|
||||
metadata,
|
||||
"local_execution_runtime_miss_reason",
|
||||
MAX_USAGE_RUNTIME_MISS_REASON_LEN,
|
||||
false,
|
||||
);
|
||||
}
|
||||
|
||||
fn sanitize_usage_routing_string_with_metadata(
|
||||
typed: Option<String>,
|
||||
metadata: Option<&Map<String, Value>>,
|
||||
key: &str,
|
||||
max_len: usize,
|
||||
allow_spaces: bool,
|
||||
) -> Option<String> {
|
||||
match typed {
|
||||
Some(value) => sanitize_usage_routing_string(Some(value), max_len, allow_spaces),
|
||||
None => metadata_routing_string(metadata, key, max_len, allow_spaces),
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_usage_routing_index_with_metadata(
|
||||
typed: Option<u64>,
|
||||
metadata: Option<&Map<String, Value>>,
|
||||
) -> Option<u64> {
|
||||
match typed {
|
||||
Some(value) => sanitize_usage_routing_index(Some(value)),
|
||||
None => metadata
|
||||
.and_then(|object| object.get("candidate_index"))
|
||||
.and_then(|value| {
|
||||
value
|
||||
.as_u64()
|
||||
.or_else(|| value.as_i64().and_then(|value| u64::try_from(value).ok()))
|
||||
.filter(|value| *value <= MAX_USAGE_CANDIDATE_INDEX)
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
fn metadata_routing_string(
|
||||
metadata: Option<&Map<String, Value>>,
|
||||
key: &str,
|
||||
max_len: usize,
|
||||
allow_spaces: bool,
|
||||
) -> Option<String> {
|
||||
metadata
|
||||
.and_then(|object| object.get(key))
|
||||
.and_then(Value::as_str)
|
||||
.and_then(|value| {
|
||||
sanitize_usage_routing_string(Some(value.to_string()), max_len, allow_spaces)
|
||||
})
|
||||
}
|
||||
|
||||
fn sanitize_usage_routing_index(value: Option<u64>) -> Option<u64> {
|
||||
value.filter(|value| *value <= MAX_USAGE_CANDIDATE_INDEX)
|
||||
}
|
||||
|
||||
fn sanitize_usage_routing_string(
|
||||
value: Option<String>,
|
||||
max_len: usize,
|
||||
allow_spaces: bool,
|
||||
) -> Option<String> {
|
||||
let value = value?;
|
||||
let value = value.trim();
|
||||
if value.is_empty() || value.len() > max_len {
|
||||
return None;
|
||||
}
|
||||
if !value.bytes().all(|byte| {
|
||||
byte.is_ascii_alphanumeric() || b"._:/@+-".contains(&byte) || (allow_spaces && byte == b' ')
|
||||
}) {
|
||||
return None;
|
||||
}
|
||||
Some(value.to_string())
|
||||
}
|
||||
|
||||
pub(crate) fn sanitize_usage_error_category(value: Option<String>) -> Option<String> {
|
||||
let value = value?.trim().to_ascii_lowercase();
|
||||
let category = match value.as_str() {
|
||||
"auth"
|
||||
| "cancelled"
|
||||
| "client_error"
|
||||
| "http_error"
|
||||
| "non_success_status"
|
||||
| "provider_error"
|
||||
| "rate_limit"
|
||||
| "redirect"
|
||||
| "server_error"
|
||||
| "stream_missing_terminal_event"
|
||||
| "stream_terminal_error"
|
||||
| "upstream_error" => value,
|
||||
"" => return None,
|
||||
_ => "other_error".to_string(),
|
||||
};
|
||||
Some(category)
|
||||
}
|
||||
|
||||
/// Return the bounded error category represented by an HTTP status code.
|
||||
///
|
||||
/// Stale-request cleanup may have only a candidate status code available. Do
|
||||
/// not persist provider-supplied diagnostic text in that case; derive one of
|
||||
/// the same fixed categories used by the usage writer instead.
|
||||
pub fn usage_error_category_for_status_code(status_code: u16) -> &'static str {
|
||||
if status_code >= 500 {
|
||||
"server_error"
|
||||
} else if status_code >= 400 {
|
||||
"client_error"
|
||||
} else if status_code >= 300 {
|
||||
"redirect"
|
||||
} else {
|
||||
"non_success_status"
|
||||
}
|
||||
}
|
||||
|
||||
pub fn provider_api_key_usage_is_success(
|
||||
status: &str,
|
||||
status_code: Option<u16>,
|
||||
@@ -339,6 +626,361 @@ fn newer_last_used_at(before: Option<u64>, after: Option<u64>) -> Option<u64> {
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
sanitize_usage_capture_controls_for_persistence, sanitize_usage_error_category,
|
||||
sanitize_usage_for_persistence, usage_error_category_for_status_code,
|
||||
usage_lifecycle_update_allowed,
|
||||
};
|
||||
use crate::repository::usage::{UpsertUsageRecord, UsageBodyCaptureState};
|
||||
|
||||
fn usage_with_http_capture() -> UpsertUsageRecord {
|
||||
UpsertUsageRecord {
|
||||
request_id: "req-sensitive-capture".to_string(),
|
||||
user_id: Some("user-1".to_string()),
|
||||
api_key_id: Some("key-1".to_string()),
|
||||
username: Some("alice".to_string()),
|
||||
api_key_name: Some("primary".to_string()),
|
||||
provider_name: "OpenAI".to_string(),
|
||||
model: "gpt-5".to_string(),
|
||||
target_model: None,
|
||||
provider_id: Some("provider-1".to_string()),
|
||||
provider_endpoint_id: None,
|
||||
provider_api_key_id: None,
|
||||
request_type: Some("chat".to_string()),
|
||||
api_format: Some("openai:chat".to_string()),
|
||||
api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
endpoint_api_format: Some("openai:chat".to_string()),
|
||||
provider_api_family: Some("openai".to_string()),
|
||||
provider_endpoint_kind: Some("chat".to_string()),
|
||||
has_format_conversion: Some(false),
|
||||
is_stream: Some(false),
|
||||
input_tokens: Some(1),
|
||||
output_tokens: Some(2),
|
||||
total_tokens: Some(3),
|
||||
cache_creation_input_tokens: None,
|
||||
cache_creation_ephemeral_5m_input_tokens: None,
|
||||
cache_creation_ephemeral_1h_input_tokens: None,
|
||||
cache_read_input_tokens: None,
|
||||
cache_creation_cost_usd: None,
|
||||
cache_read_cost_usd: None,
|
||||
output_price_per_1m: None,
|
||||
total_cost_usd: Some(0.01),
|
||||
actual_total_cost_usd: Some(0.01),
|
||||
status_code: Some(200),
|
||||
error_message: Some("Bearer secret".to_string()),
|
||||
error_category: Some("provider_error".to_string()),
|
||||
response_time_ms: Some(10),
|
||||
first_byte_time_ms: Some(5),
|
||||
status: "completed".to_string(),
|
||||
billing_status: "settled".to_string(),
|
||||
request_headers: Some(json!({"authorization": "Bearer secret"})),
|
||||
request_body: Some(json!({"prompt": "private"})),
|
||||
request_body_ref: Some("usage://request/body".to_string()),
|
||||
request_body_state: Some(UsageBodyCaptureState::Reference),
|
||||
provider_request_headers: Some(json!({"x-api-key": "secret"})),
|
||||
provider_request_body: Some(json!({"prompt": "private"})),
|
||||
provider_request_body_ref: Some("usage://provider/body".to_string()),
|
||||
provider_request_body_state: Some(UsageBodyCaptureState::Inline),
|
||||
response_headers: Some(json!({"set-cookie": "secret"})),
|
||||
response_body: Some(json!({"output": "private"})),
|
||||
response_body_ref: Some("usage://response/body".to_string()),
|
||||
response_body_state: Some(UsageBodyCaptureState::Truncated),
|
||||
client_response_headers: Some(json!({"x-private": "secret"})),
|
||||
client_response_body: Some(json!({"output": "private"})),
|
||||
client_response_body_ref: Some("usage://client/body".to_string()),
|
||||
client_response_body_state: Some(UsageBodyCaptureState::Disabled),
|
||||
candidate_id: Some("candidate-1".to_string()),
|
||||
candidate_index: Some(0),
|
||||
key_name: None,
|
||||
planner_kind: None,
|
||||
route_family: None,
|
||||
route_kind: None,
|
||||
execution_path: None,
|
||||
local_execution_runtime_miss_reason: None,
|
||||
request_metadata: Some(json!({"client_ip": "203.0.113.8"})),
|
||||
finalized_at_unix_secs: Some(2),
|
||||
created_at_unix_ms: Some(1_000),
|
||||
updated_at_unix_secs: 2,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn usage_error_categories_are_bounded_to_controlled_values() {
|
||||
assert_eq!(
|
||||
sanitize_usage_error_category(Some(" Server_Error ".to_string())).as_deref(),
|
||||
Some("server_error")
|
||||
);
|
||||
assert_eq!(
|
||||
sanitize_usage_error_category(Some("Authorization: Bearer secret".to_string()))
|
||||
.as_deref(),
|
||||
Some("other_error")
|
||||
);
|
||||
assert_eq!(sanitize_usage_error_category(Some(" ".to_string())), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn status_codes_map_to_bounded_error_categories() {
|
||||
assert_eq!(usage_error_category_for_status_code(599), "server_error");
|
||||
assert_eq!(usage_error_category_for_status_code(400), "client_error");
|
||||
assert_eq!(usage_error_category_for_status_code(302), "redirect");
|
||||
assert_eq!(
|
||||
usage_error_category_for_status_code(200),
|
||||
"non_success_status"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn persistence_derives_missing_failed_category_from_status_code() {
|
||||
let mut usage = usage_with_http_capture();
|
||||
usage.status = "failed".to_string();
|
||||
usage.status_code = Some(429);
|
||||
usage.error_category = None;
|
||||
|
||||
let usage = sanitize_usage_for_persistence(usage);
|
||||
assert_eq!(usage.error_category.as_deref(), Some("client_error"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn persistence_boundary_drops_all_http_capture_material() {
|
||||
let usage = sanitize_usage_for_persistence(usage_with_http_capture());
|
||||
|
||||
assert!(usage.username.is_none());
|
||||
assert!(usage.api_key_name.is_none());
|
||||
assert!(usage.error_message.is_none());
|
||||
assert!(usage.request_headers.is_none());
|
||||
assert!(usage.request_body.is_none());
|
||||
assert!(usage.request_body_ref.is_none());
|
||||
assert!(usage.request_body_state.is_none());
|
||||
assert!(usage.provider_request_headers.is_none());
|
||||
assert!(usage.provider_request_body.is_none());
|
||||
assert!(usage.provider_request_body_ref.is_none());
|
||||
assert!(usage.provider_request_body_state.is_none());
|
||||
assert!(usage.response_headers.is_none());
|
||||
assert!(usage.response_body.is_none());
|
||||
assert!(usage.response_body_ref.is_none());
|
||||
assert!(usage.response_body_state.is_none());
|
||||
assert!(usage.client_response_headers.is_none());
|
||||
assert!(usage.client_response_body.is_none());
|
||||
assert!(usage.client_response_body_ref.is_none());
|
||||
assert!(usage.client_response_body_state.is_none());
|
||||
assert_eq!(usage.error_category.as_deref(), Some("provider_error"));
|
||||
assert_eq!(usage.candidate_id.as_deref(), Some("candidate-1"));
|
||||
assert_eq!(
|
||||
usage.request_metadata,
|
||||
Some(json!({"client_ip": "203.0.113.8"}))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auxiliary_capture_projection_keeps_only_explicit_clear_tombstones() {
|
||||
let mut input = usage_with_http_capture();
|
||||
input.request_body_state = Some(UsageBodyCaptureState::None);
|
||||
input.response_body_state = Some(UsageBodyCaptureState::Disabled);
|
||||
|
||||
let usage = sanitize_usage_capture_controls_for_persistence(input);
|
||||
|
||||
assert!(usage.request_headers.is_none());
|
||||
assert!(usage.request_body.is_none());
|
||||
assert!(usage.request_body_ref.is_none());
|
||||
assert_eq!(usage.request_body_state, Some(UsageBodyCaptureState::None));
|
||||
assert!(usage.provider_request_body_state.is_none());
|
||||
assert!(usage.response_body_state.is_none());
|
||||
assert!(usage.client_response_body_state.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn auxiliary_projection_promotes_only_bounded_routing_metadata() {
|
||||
let mut input = usage_with_http_capture();
|
||||
input.candidate_id = None;
|
||||
input.candidate_index = None;
|
||||
input.key_name = None;
|
||||
input.planner_kind = None;
|
||||
input.route_family = None;
|
||||
input.route_kind = None;
|
||||
input.execution_path = None;
|
||||
input.local_execution_runtime_miss_reason = None;
|
||||
input.request_metadata = Some(json!({
|
||||
"trace_id": "trace",
|
||||
"candidate_id": "candidate-from-metadata",
|
||||
"candidate_index": 7,
|
||||
"key_name": "primary key",
|
||||
"planner_kind": "fallback",
|
||||
"route_family": "chat",
|
||||
"route_kind": "remote",
|
||||
"execution_path": "execution_runtime_stream",
|
||||
"local_execution_runtime_miss_reason": "runtime_busy",
|
||||
"authorization": "Bearer should-not-persist",
|
||||
}));
|
||||
|
||||
let usage = sanitize_usage_capture_controls_for_persistence(input);
|
||||
|
||||
assert_eq!(
|
||||
usage.candidate_id.as_deref(),
|
||||
Some("candidate-from-metadata")
|
||||
);
|
||||
assert_eq!(usage.candidate_index, Some(7));
|
||||
assert_eq!(usage.key_name.as_deref(), Some("primary key"));
|
||||
assert_eq!(usage.planner_kind.as_deref(), Some("fallback"));
|
||||
assert_eq!(usage.route_family.as_deref(), Some("chat"));
|
||||
assert_eq!(usage.route_kind.as_deref(), Some("remote"));
|
||||
assert_eq!(
|
||||
usage.execution_path.as_deref(),
|
||||
Some("execution_runtime_stream")
|
||||
);
|
||||
assert_eq!(
|
||||
usage.local_execution_runtime_miss_reason.as_deref(),
|
||||
Some("runtime_busy")
|
||||
);
|
||||
assert_eq!(usage.request_metadata, Some(json!({"trace_id": "trace"})));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn routing_projection_rejects_unbounded_or_control_character_values() {
|
||||
let mut input = usage_with_http_capture();
|
||||
input.candidate_id = Some("candidate\nforged".to_string());
|
||||
input.candidate_index = Some(u64::MAX);
|
||||
input.key_name = Some("key\0name".to_string());
|
||||
input.planner_kind = Some("p".repeat(65));
|
||||
input.route_family = Some("route\tname".to_string());
|
||||
input.route_kind = Some("route-kind".to_string());
|
||||
input.execution_path = Some("execution-path".to_string());
|
||||
input.local_execution_runtime_miss_reason = Some("m".repeat(121));
|
||||
|
||||
let usage = sanitize_usage_for_persistence(input);
|
||||
|
||||
assert!(usage.candidate_id.is_none());
|
||||
assert!(usage.candidate_index.is_none());
|
||||
assert!(usage.key_name.is_none());
|
||||
assert!(usage.planner_kind.is_none());
|
||||
assert!(usage.route_family.is_none());
|
||||
assert_eq!(usage.route_kind.as_deref(), Some("route-kind"));
|
||||
assert_eq!(usage.execution_path.as_deref(), Some("execution-path"));
|
||||
assert!(usage.local_execution_runtime_miss_reason.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn invalid_typed_routing_values_do_not_fall_back_to_metadata() {
|
||||
let mut input = usage_with_http_capture();
|
||||
input.candidate_id = Some("candidate\nforged".to_string());
|
||||
input.candidate_index = Some(u64::MAX);
|
||||
input.key_name = Some("key\0name".to_string());
|
||||
input.planner_kind = Some("p".repeat(65));
|
||||
input.route_family = Some("route\tname".to_string());
|
||||
input.route_kind = Some("route\nkind".to_string());
|
||||
input.execution_path = Some("execution\npath".to_string());
|
||||
input.local_execution_runtime_miss_reason = Some("m".repeat(121));
|
||||
input.request_metadata = Some(json!({
|
||||
"candidate_id": "metadata-candidate",
|
||||
"candidate_index": 7,
|
||||
"key_name": "metadata key",
|
||||
"planner_kind": "metadata-planner",
|
||||
"route_family": "metadata-family",
|
||||
"route_kind": "metadata-kind",
|
||||
"execution_path": "metadata-path",
|
||||
"local_execution_runtime_miss_reason": "metadata-reason",
|
||||
}));
|
||||
|
||||
let usage = sanitize_usage_capture_controls_for_persistence(input);
|
||||
|
||||
assert!(usage.candidate_id.is_none());
|
||||
assert!(usage.candidate_index.is_none());
|
||||
assert!(usage.key_name.is_none());
|
||||
assert!(usage.planner_kind.is_none());
|
||||
assert!(usage.route_family.is_none());
|
||||
assert!(usage.route_kind.is_none());
|
||||
assert!(usage.execution_path.is_none());
|
||||
assert!(usage.local_execution_runtime_miss_reason.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn lifecycle_order_rejects_stale_and_equal_conflicting_terminal_events() {
|
||||
assert!(!usage_lifecycle_update_allowed(
|
||||
"completed",
|
||||
"pending",
|
||||
20,
|
||||
Some(20),
|
||||
"failed",
|
||||
"void",
|
||||
19,
|
||||
Some(19),
|
||||
));
|
||||
assert!(!usage_lifecycle_update_allowed(
|
||||
"completed",
|
||||
"pending",
|
||||
20,
|
||||
Some(20),
|
||||
"failed",
|
||||
"void",
|
||||
20,
|
||||
Some(20),
|
||||
));
|
||||
assert!(!usage_lifecycle_update_allowed(
|
||||
"completed",
|
||||
"pending",
|
||||
20,
|
||||
Some(20),
|
||||
"completed",
|
||||
"settled",
|
||||
20,
|
||||
Some(20),
|
||||
));
|
||||
assert!(usage_lifecycle_update_allowed(
|
||||
"completed",
|
||||
"pending",
|
||||
20,
|
||||
Some(20),
|
||||
"failed",
|
||||
"void",
|
||||
21,
|
||||
Some(21),
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn lifecycle_order_allows_same_second_progress_and_fresh_void_recovery() {
|
||||
assert!(usage_lifecycle_update_allowed(
|
||||
"pending",
|
||||
"pending",
|
||||
20,
|
||||
None,
|
||||
"streaming",
|
||||
"pending",
|
||||
20,
|
||||
None,
|
||||
));
|
||||
assert!(usage_lifecycle_update_allowed(
|
||||
"streaming",
|
||||
"pending",
|
||||
20,
|
||||
None,
|
||||
"completed",
|
||||
"pending",
|
||||
20,
|
||||
Some(20),
|
||||
));
|
||||
assert!(usage_lifecycle_update_allowed(
|
||||
"failed",
|
||||
"void",
|
||||
20,
|
||||
Some(20),
|
||||
"completed",
|
||||
"pending",
|
||||
20,
|
||||
Some(20),
|
||||
));
|
||||
assert!(!usage_lifecycle_update_allowed(
|
||||
"failed",
|
||||
"void",
|
||||
20,
|
||||
Some(21),
|
||||
"completed",
|
||||
"pending",
|
||||
20,
|
||||
Some(20),
|
||||
));
|
||||
}
|
||||
|
||||
use super::{api_key_usage_contribution, provider_api_key_usage_contribution};
|
||||
use crate::repository::usage::StoredRequestUsageAudit;
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ pub const ROUTING_CANDIDATE_SKIP_REASON_METADATA_KEY: &str = "routing_candidate_
|
||||
pub const ROUTING_FAILURE_DIAGNOSTIC_METADATA_KEY: &str = "routing_failure_diagnostic";
|
||||
pub const WEBSOCKET_MODE_METADATA_KEY: &str = "websocket_mode";
|
||||
pub const WEBSOCKET_TRANSPORT_METADATA_KEY: &str = "websocket_transport";
|
||||
pub const PLAN_USAGE_RESERVATION_DEFERRED_METADATA_KEY: &str = "plan_usage_reservation_deferred";
|
||||
/// Whether token/cost usage is authoritative for this audit row.
|
||||
///
|
||||
/// The field is absent for legacy and normally-metered requests. An explicit
|
||||
@@ -386,7 +387,7 @@ impl StoredRequestUsageAudit {
|
||||
total_cost_usd: f64,
|
||||
actual_total_cost_usd: f64,
|
||||
status_code: Option<i32>,
|
||||
error_message: Option<String>,
|
||||
_error_message: Option<String>,
|
||||
error_category: Option<String>,
|
||||
response_time_ms: Option<i32>,
|
||||
first_byte_time_ms: Option<i32>,
|
||||
@@ -421,14 +422,14 @@ impl StoredRequestUsageAudit {
|
||||
"usage.billing_status is empty".to_string(),
|
||||
));
|
||||
}
|
||||
if !total_cost_usd.is_finite() {
|
||||
if !total_cost_usd.is_finite() || total_cost_usd < 0.0 {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"usage.total_cost_usd is not finite".to_string(),
|
||||
"usage.total_cost_usd must be finite and non-negative".to_string(),
|
||||
));
|
||||
}
|
||||
if !actual_total_cost_usd.is_finite() {
|
||||
if !actual_total_cost_usd.is_finite() || actual_total_cost_usd < 0.0 {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"usage.actual_total_cost_usd is not finite".to_string(),
|
||||
"usage.actual_total_cost_usd must be finite and non-negative".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
@@ -468,8 +469,8 @@ impl StoredRequestUsageAudit {
|
||||
total_cost_usd,
|
||||
actual_total_cost_usd,
|
||||
status_code: parse_u16(status_code, "usage.status_code")?,
|
||||
error_message,
|
||||
error_category,
|
||||
error_message: None,
|
||||
error_category: super::policy::sanitize_usage_error_category(error_category),
|
||||
response_time_ms: parse_optional_u64(response_time_ms, "usage.response_time_ms")?,
|
||||
first_byte_time_ms: parse_optional_u64(first_byte_time_ms, "usage.first_byte_time_ms")?,
|
||||
status,
|
||||
@@ -1721,6 +1722,16 @@ pub fn parse_usage_body_ref(body_ref: &str) -> Option<(String, UsageBodyField)>
|
||||
))
|
||||
}
|
||||
|
||||
pub fn canonical_usage_body_ref_for(
|
||||
body_ref: &str,
|
||||
expected_request_id: &str,
|
||||
expected_field: UsageBodyField,
|
||||
) -> Option<String> {
|
||||
parse_usage_body_ref(body_ref)
|
||||
.filter(|(request_id, field)| request_id == expected_request_id && *field == expected_field)
|
||||
.map(|(request_id, field)| usage_body_ref(&request_id, field))
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait UsageReadRepository: Send + Sync {
|
||||
async fn find_by_id(
|
||||
@@ -2043,48 +2054,58 @@ impl UpsertUsageRecord {
|
||||
"usage upsert model cannot be empty".to_string(),
|
||||
));
|
||||
}
|
||||
if self.status.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert status cannot be empty".to_string(),
|
||||
));
|
||||
if !matches!(
|
||||
self.status.as_str(),
|
||||
"pending" | "streaming" | "completed" | "failed" | "cancelled"
|
||||
) {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"invalid usage upsert status: {}",
|
||||
self.status
|
||||
)));
|
||||
}
|
||||
if self.billing_status.trim().is_empty() {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert billing_status cannot be empty".to_string(),
|
||||
));
|
||||
if !matches!(
|
||||
self.billing_status.as_str(),
|
||||
"pending" | "settled" | "void" | "insufficient_quota"
|
||||
) {
|
||||
return Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"invalid usage upsert billing_status: {}",
|
||||
self.billing_status
|
||||
)));
|
||||
}
|
||||
if let Some(value) = self.total_cost_usd {
|
||||
if !value.is_finite() {
|
||||
if !value.is_finite() || value < 0.0 {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert total_cost_usd must be finite".to_string(),
|
||||
"usage upsert total_cost_usd must be finite and non-negative".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(value) = self.cache_creation_cost_usd {
|
||||
if !value.is_finite() {
|
||||
if !value.is_finite() || value < 0.0 {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert cache_creation_cost_usd must be finite".to_string(),
|
||||
"usage upsert cache_creation_cost_usd must be finite and non-negative"
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(value) = self.cache_read_cost_usd {
|
||||
if !value.is_finite() {
|
||||
if !value.is_finite() || value < 0.0 {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert cache_read_cost_usd must be finite".to_string(),
|
||||
"usage upsert cache_read_cost_usd must be finite and non-negative".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(value) = self.output_price_per_1m {
|
||||
if !value.is_finite() {
|
||||
if !value.is_finite() || value < 0.0 {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert output_price_per_1m must be finite".to_string(),
|
||||
"usage upsert output_price_per_1m must be finite and non-negative".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
if let Some(value) = self.actual_total_cost_usd {
|
||||
if !value.is_finite() {
|
||||
if !value.is_finite() || value < 0.0 {
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"usage upsert actual_total_cost_usd must be finite".to_string(),
|
||||
"usage upsert actual_total_cost_usd must be finite and non-negative"
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
@@ -2271,9 +2292,14 @@ pub struct UsageCounterPendingHealthSnapshot {
|
||||
pub pending_by_kind: std::collections::BTreeMap<String, u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct ProxyNodeCounterDelta {
|
||||
pub node_id: String,
|
||||
/// Incarnation fence captured when the request plan selected this node.
|
||||
/// Counter writes must never silently rebind to a different incarnation
|
||||
/// that reused the same node id.
|
||||
#[serde(default)]
|
||||
pub expected_tunnel_generation: Option<String>,
|
||||
pub total_requests_delta: i64,
|
||||
pub failed_requests_delta: i64,
|
||||
pub dns_failures_delta: i64,
|
||||
@@ -2332,6 +2358,8 @@ pub struct UsageCleanupSummary {
|
||||
pub header_cleaned: usize,
|
||||
pub keys_cleaned: usize,
|
||||
pub records_deleted: usize,
|
||||
pub cost_reservations_deleted: usize,
|
||||
pub request_admissions_deleted: usize,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
@@ -2452,12 +2480,13 @@ fn parse_timestamp(value: i64, field_name: &str) -> Result<u64, crate::DataLayer
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
extract_provider_actual_service_tier_from_response,
|
||||
canonical_usage_body_ref_for, extract_provider_actual_service_tier_from_response,
|
||||
extract_provider_service_tier_from_body, resolve_provider_cache_ttl_minutes,
|
||||
StoredRequestUsageAudit, UpsertUsageRecord, UsageBodyCaptureState, UsageBodyCaptureStorage,
|
||||
UsageBodyField, UsageProviderPerformanceQuery, REALTIME_SESSION_METADATA_KEY,
|
||||
USAGE_AVAILABLE_METADATA_KEY, USAGE_PRICING_AVAILABLE_METADATA_KEY,
|
||||
WEBSOCKET_MODE_METADATA_KEY, WEBSOCKET_TRANSPORT_METADATA_KEY,
|
||||
usage_body_ref, StoredRequestUsageAudit, UpsertUsageRecord, UsageBodyCaptureState,
|
||||
UsageBodyCaptureStorage, UsageBodyField, UsageProviderPerformanceQuery,
|
||||
REALTIME_SESSION_METADATA_KEY, USAGE_AVAILABLE_METADATA_KEY,
|
||||
USAGE_PRICING_AVAILABLE_METADATA_KEY, WEBSOCKET_MODE_METADATA_KEY,
|
||||
WEBSOCKET_TRANSPORT_METADATA_KEY,
|
||||
};
|
||||
use serde_json::{json, Value};
|
||||
|
||||
@@ -2503,6 +2532,38 @@ mod tests {
|
||||
.expect("usage should build")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn canonical_body_ref_requires_matching_request_and_field() {
|
||||
assert_eq!(
|
||||
canonical_usage_body_ref_for(
|
||||
" usage://request/req-1/request_body ",
|
||||
"req-1",
|
||||
UsageBodyField::RequestBody,
|
||||
),
|
||||
Some(usage_body_ref("req-1", UsageBodyField::RequestBody))
|
||||
);
|
||||
assert_eq!(
|
||||
canonical_usage_body_ref_for(
|
||||
"usage://request/req-2/request_body",
|
||||
"req-1",
|
||||
UsageBodyField::RequestBody,
|
||||
),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
canonical_usage_body_ref_for(
|
||||
"usage://request/req-1/response_body",
|
||||
"req-1",
|
||||
UsageBodyField::RequestBody,
|
||||
),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
canonical_usage_body_ref_for("blob://opaque", "req-1", UsageBodyField::RequestBody),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_performance_query_defaults_timeline_on_for_legacy_payloads() {
|
||||
let mut payload = json!({
|
||||
@@ -2618,7 +2679,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn rejects_invalid_upsert_payload() {
|
||||
let record = UpsertUsageRecord {
|
||||
let mut record = UpsertUsageRecord {
|
||||
request_id: "".to_string(),
|
||||
user_id: None,
|
||||
api_key_id: None,
|
||||
@@ -2689,6 +2750,42 @@ mod tests {
|
||||
};
|
||||
|
||||
assert!(record.validate().is_err());
|
||||
|
||||
record.request_id = "req-1".to_string();
|
||||
assert!(record.validate().is_ok());
|
||||
|
||||
for invalid_status in ["", " completed ", "success", "COMPLETED"] {
|
||||
record.status = invalid_status.to_string();
|
||||
assert!(
|
||||
record.validate().is_err(),
|
||||
"accepted status {invalid_status:?}"
|
||||
);
|
||||
}
|
||||
|
||||
record.status = "completed".to_string();
|
||||
for invalid_billing_status in ["", " settled ", "paid", "SETTLED"] {
|
||||
record.billing_status = invalid_billing_status.to_string();
|
||||
assert!(
|
||||
record.validate().is_err(),
|
||||
"accepted billing status {invalid_billing_status:?}"
|
||||
);
|
||||
}
|
||||
|
||||
record.billing_status = "pending".to_string();
|
||||
record.total_cost_usd = Some(-0.01);
|
||||
assert!(record.validate().is_err());
|
||||
record.total_cost_usd = None;
|
||||
record.actual_total_cost_usd = Some(-0.01);
|
||||
assert!(record.validate().is_err());
|
||||
record.actual_total_cost_usd = None;
|
||||
record.cache_creation_cost_usd = Some(-0.01);
|
||||
assert!(record.validate().is_err());
|
||||
record.cache_creation_cost_usd = None;
|
||||
record.cache_read_cost_usd = Some(-0.01);
|
||||
assert!(record.validate().is_err());
|
||||
record.cache_read_cost_usd = None;
|
||||
record.output_price_per_1m = Some(-0.01);
|
||||
assert!(record.validate().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -2,6 +2,21 @@ use async_trait::async_trait;
|
||||
use chrono::{DateTime, Utc};
|
||||
use serde_json::Value;
|
||||
|
||||
fn redacted_optional_secret<T>(value: &Option<T>) -> Option<&'static str> {
|
||||
value.as_ref().map(|_| "[REDACTED]")
|
||||
}
|
||||
|
||||
pub const LAST_ACTIVE_ADMIN_UPDATE_DENIED: &str = "last_active_admin_update_denied";
|
||||
pub const LAST_ACTIVE_ADMIN_DELETE_DENIED: &str = "last_active_admin_delete_denied";
|
||||
|
||||
pub fn is_last_active_admin_update_denied(error: &crate::DataLayerError) -> bool {
|
||||
matches!(error, crate::DataLayerError::InvalidInput(message) if message == LAST_ACTIVE_ADMIN_UPDATE_DENIED)
|
||||
}
|
||||
|
||||
pub fn is_last_active_admin_delete_denied(error: &crate::DataLayerError) -> bool {
|
||||
matches!(error, crate::DataLayerError::InvalidInput(message) if message == LAST_ACTIVE_ADMIN_DELETE_DENIED)
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredUserSummary {
|
||||
pub id: String,
|
||||
@@ -47,7 +62,7 @@ impl StoredUserSummary {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredUserAuthRecord {
|
||||
pub id: String,
|
||||
pub email: Option<String>,
|
||||
@@ -64,10 +79,41 @@ pub struct StoredUserAuthRecord {
|
||||
pub allowed_models_mode: String,
|
||||
pub is_active: bool,
|
||||
pub is_deleted: bool,
|
||||
#[serde(default)]
|
||||
pub security_version: i64,
|
||||
pub created_at: Option<DateTime<Utc>>,
|
||||
pub last_login_at: Option<DateTime<Utc>>,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredUserAuthRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredUserAuthRecord")
|
||||
.field("id", &self.id)
|
||||
.field("email", &self.email)
|
||||
.field("email_verified", &self.email_verified)
|
||||
.field("username", &self.username)
|
||||
.field(
|
||||
"password_hash",
|
||||
&redacted_optional_secret(&self.password_hash),
|
||||
)
|
||||
.field("role", &self.role)
|
||||
.field("auth_source", &self.auth_source)
|
||||
.field("allowed_providers", &self.allowed_providers)
|
||||
.field("allowed_providers_mode", &self.allowed_providers_mode)
|
||||
.field("allowed_api_formats", &self.allowed_api_formats)
|
||||
.field("allowed_api_formats_mode", &self.allowed_api_formats_mode)
|
||||
.field("allowed_models", &self.allowed_models)
|
||||
.field("allowed_models_mode", &self.allowed_models_mode)
|
||||
.field("is_active", &self.is_active)
|
||||
.field("is_deleted", &self.is_deleted)
|
||||
.field("security_version", &self.security_version)
|
||||
.field("created_at", &self.created_at)
|
||||
.field("last_login_at", &self.last_login_at)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredUserAuthRecord {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
@@ -126,6 +172,7 @@ impl StoredUserAuthRecord {
|
||||
allowed_models_mode: "unrestricted".to_string(),
|
||||
is_active,
|
||||
is_deleted,
|
||||
security_version: 0,
|
||||
created_at,
|
||||
last_login_at,
|
||||
})
|
||||
@@ -149,6 +196,40 @@ impl StoredUserAuthRecord {
|
||||
Ok(self)
|
||||
}
|
||||
|
||||
pub fn with_security_version(
|
||||
mut self,
|
||||
security_version: i64,
|
||||
) -> Result<Self, crate::DataLayerError> {
|
||||
if security_version < 0 {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"users.security_version is negative".to_string(),
|
||||
));
|
||||
}
|
||||
self.security_version = security_version;
|
||||
Ok(self)
|
||||
}
|
||||
|
||||
/// Compare the user fields that an aggregate import is allowed to restore.
|
||||
/// Passwords, security versions, and timestamps are intentionally excluded:
|
||||
/// password restoration has its own nullable CAS operation and the remaining
|
||||
/// fields are server-managed concurrency markers.
|
||||
pub fn matches_restore_state(&self, expected: &Self) -> bool {
|
||||
self.id == expected.id
|
||||
&& self.email == expected.email
|
||||
&& self.email_verified == expected.email_verified
|
||||
&& self.username == expected.username
|
||||
&& self.role == expected.role
|
||||
&& self.auth_source == expected.auth_source
|
||||
&& self.allowed_providers == expected.allowed_providers
|
||||
&& self.allowed_providers_mode == expected.allowed_providers_mode
|
||||
&& self.allowed_api_formats == expected.allowed_api_formats
|
||||
&& self.allowed_api_formats_mode == expected.allowed_api_formats_mode
|
||||
&& self.allowed_models == expected.allowed_models
|
||||
&& self.allowed_models_mode == expected.allowed_models_mode
|
||||
&& self.is_active == expected.is_active
|
||||
&& self.is_deleted == expected.is_deleted
|
||||
}
|
||||
|
||||
fn with_legacy_policy_modes(mut self) -> Self {
|
||||
self.allowed_providers_mode = legacy_list_policy_mode(&self.allowed_providers);
|
||||
self.allowed_api_formats_mode = legacy_list_policy_mode(&self.allowed_api_formats);
|
||||
@@ -174,6 +255,123 @@ pub struct LdapAuthUserProvisioningOutcome {
|
||||
pub created: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum DeleteUserOAuthLinkOutcome {
|
||||
Deleted,
|
||||
NotFound,
|
||||
LastOAuthBinding,
|
||||
LastLoginMethod,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
pub enum BindUserOAuthLinkOutcome {
|
||||
Bound,
|
||||
IdentityAlreadyBoundToUser,
|
||||
IdentityBoundToAnotherUser,
|
||||
UserAlreadyLinkedProvider,
|
||||
UserNotFound,
|
||||
SessionUnavailable,
|
||||
ProviderNotFound,
|
||||
ProviderDisabled,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
pub struct BindUserOAuthLinkSessionExpectation {
|
||||
pub session_id: String,
|
||||
pub client_device_id: String,
|
||||
pub security_version: i64,
|
||||
pub checked_at: DateTime<Utc>,
|
||||
}
|
||||
|
||||
impl BindUserOAuthLinkSessionExpectation {
|
||||
pub fn new(
|
||||
session_id: impl Into<String>,
|
||||
client_device_id: impl Into<String>,
|
||||
security_version: i64,
|
||||
checked_at: DateTime<Utc>,
|
||||
) -> Result<Self, crate::DataLayerError> {
|
||||
let session_id = session_id.into();
|
||||
let client_device_id = client_device_id.into();
|
||||
if session_id.trim().is_empty()
|
||||
|| session_id != session_id.trim()
|
||||
|| client_device_id.trim().is_empty()
|
||||
|| client_device_id != client_device_id.trim()
|
||||
|| security_version < 0
|
||||
{
|
||||
return Err(crate::DataLayerError::InvalidInput(
|
||||
"OAuth link session expectation is invalid".to_string(),
|
||||
));
|
||||
}
|
||||
Ok(Self {
|
||||
session_id,
|
||||
client_device_id,
|
||||
security_version,
|
||||
checked_at,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq)]
|
||||
pub enum ResolveOAuthLinkedUserOutcome {
|
||||
Linked(StoredUserAuthRecord),
|
||||
NotLinked,
|
||||
ProviderUnavailable,
|
||||
}
|
||||
|
||||
pub fn last_oauth_unbind_denial(
|
||||
auth_source: &str,
|
||||
password_hash: Option<&str>,
|
||||
local_password_login_allowed: bool,
|
||||
) -> Option<DeleteUserOAuthLinkOutcome> {
|
||||
if auth_source.eq_ignore_ascii_case("oauth") {
|
||||
return Some(DeleteUserOAuthLinkOutcome::LastOAuthBinding);
|
||||
}
|
||||
if !auth_source.eq_ignore_ascii_case("local")
|
||||
|| !local_password_login_allowed
|
||||
|| !password_hash.is_some_and(is_valid_bcrypt_hash)
|
||||
{
|
||||
return Some(DeleteUserOAuthLinkOutcome::LastLoginMethod);
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
pub fn is_valid_bcrypt_hash(value: &str) -> bool {
|
||||
use base64::Engine as _;
|
||||
use std::str::FromStr;
|
||||
|
||||
let bytes = value.as_bytes();
|
||||
if value.len() != 60
|
||||
|| !matches!(value.get(0..4), Some("$2a$") | Some("$2b$") | Some("$2y$"))
|
||||
|| !bytes.get(4).is_some_and(u8::is_ascii_digit)
|
||||
|| !bytes.get(5).is_some_and(u8::is_ascii_digit)
|
||||
|| bytes.get(6) != Some(&b'$')
|
||||
{
|
||||
return false;
|
||||
}
|
||||
let Ok(parts) = bcrypt::HashParts::from_str(value) else {
|
||||
return false;
|
||||
};
|
||||
if !(4..=31).contains(&parts.get_cost()) {
|
||||
return false;
|
||||
}
|
||||
let Some(payload) = value.get(7..) else {
|
||||
return false;
|
||||
};
|
||||
if !payload.bytes().all(is_bcrypt_base64_byte) {
|
||||
return false;
|
||||
}
|
||||
bcrypt::BASE_64
|
||||
.decode(&payload[..22])
|
||||
.is_ok_and(|salt| salt.len() == 16)
|
||||
&& bcrypt::BASE_64
|
||||
.decode(&payload[22..])
|
||||
.is_ok_and(|hash| hash.len() == 23)
|
||||
}
|
||||
|
||||
fn is_bcrypt_base64_byte(value: u8) -> bool {
|
||||
value.is_ascii_alphanumeric() || matches!(value, b'.' | b'/')
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredUserOAuthLinkSummary {
|
||||
pub provider_type: String,
|
||||
@@ -218,7 +416,7 @@ impl StoredUserOAuthLinkSummary {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredUserExportRow {
|
||||
pub id: String,
|
||||
pub email: Option<String>,
|
||||
@@ -240,6 +438,35 @@ pub struct StoredUserExportRow {
|
||||
pub is_active: bool,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredUserExportRow {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredUserExportRow")
|
||||
.field("id", &self.id)
|
||||
.field("email", &self.email)
|
||||
.field("email_verified", &self.email_verified)
|
||||
.field("username", &self.username)
|
||||
.field(
|
||||
"password_hash",
|
||||
&redacted_optional_secret(&self.password_hash),
|
||||
)
|
||||
.field("role", &self.role)
|
||||
.field("auth_source", &self.auth_source)
|
||||
.field("allowed_providers", &self.allowed_providers)
|
||||
.field("allowed_providers_mode", &self.allowed_providers_mode)
|
||||
.field("allowed_api_formats", &self.allowed_api_formats)
|
||||
.field("allowed_api_formats_mode", &self.allowed_api_formats_mode)
|
||||
.field("allowed_models", &self.allowed_models)
|
||||
.field("allowed_models_mode", &self.allowed_models_mode)
|
||||
.field("rate_limit", &self.rate_limit)
|
||||
.field("rate_limit_mode", &self.rate_limit_mode)
|
||||
.field("model_capability_settings", &self.model_capability_settings)
|
||||
.field("feature_settings", &self.feature_settings)
|
||||
.field("is_active", &self.is_active)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredUserExportRow {
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
@@ -324,6 +551,25 @@ impl StoredUserExportRow {
|
||||
Ok(self)
|
||||
}
|
||||
|
||||
/// Compare the identity and policy fields that an aggregate import may
|
||||
/// restore. Passwords, rate/settings payloads, and server timestamps are
|
||||
/// checked separately by the rollback operation.
|
||||
pub fn matches_restore_state(&self, expected: &Self) -> bool {
|
||||
self.id == expected.id
|
||||
&& self.email == expected.email
|
||||
&& self.email_verified == expected.email_verified
|
||||
&& self.username == expected.username
|
||||
&& self.role == expected.role
|
||||
&& self.auth_source == expected.auth_source
|
||||
&& self.allowed_providers == expected.allowed_providers
|
||||
&& self.allowed_providers_mode == expected.allowed_providers_mode
|
||||
&& self.allowed_api_formats == expected.allowed_api_formats
|
||||
&& self.allowed_api_formats_mode == expected.allowed_api_formats_mode
|
||||
&& self.allowed_models == expected.allowed_models
|
||||
&& self.allowed_models_mode == expected.allowed_models_mode
|
||||
&& self.is_active == expected.is_active
|
||||
}
|
||||
|
||||
pub fn with_feature_settings(mut self, feature_settings: Option<Value>) -> Self {
|
||||
self.feature_settings = normalize_optional_json(feature_settings);
|
||||
self
|
||||
@@ -342,7 +588,7 @@ impl StoredUserExportRow {
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
#[derive(Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredUserSessionRecord {
|
||||
pub id: String,
|
||||
pub user_id: String,
|
||||
@@ -359,6 +605,35 @@ pub struct StoredUserSessionRecord {
|
||||
pub user_agent: Option<String>,
|
||||
pub created_at: Option<DateTime<Utc>>,
|
||||
pub updated_at: Option<DateTime<Utc>>,
|
||||
#[serde(skip)]
|
||||
pub security_version: i64,
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for StoredUserSessionRecord {
|
||||
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
formatter
|
||||
.debug_struct("StoredUserSessionRecord")
|
||||
.field("id", &self.id)
|
||||
.field("user_id", &self.user_id)
|
||||
.field("client_device_id", &self.client_device_id)
|
||||
.field("device_label", &self.device_label)
|
||||
.field("refresh_token_hash", &"[REDACTED]")
|
||||
.field(
|
||||
"prev_refresh_token_hash",
|
||||
&redacted_optional_secret(&self.prev_refresh_token_hash),
|
||||
)
|
||||
.field("rotated_at", &self.rotated_at)
|
||||
.field("last_seen_at", &self.last_seen_at)
|
||||
.field("expires_at", &self.expires_at)
|
||||
.field("revoked_at", &self.revoked_at)
|
||||
.field("revoke_reason", &self.revoke_reason)
|
||||
.field("ip_address", &self.ip_address)
|
||||
.field("user_agent", &self.user_agent)
|
||||
.field("created_at", &self.created_at)
|
||||
.field("updated_at", &self.updated_at)
|
||||
.field("security_version", &self.security_version)
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
impl StoredUserSessionRecord {
|
||||
@@ -420,9 +695,23 @@ impl StoredUserSessionRecord {
|
||||
user_agent,
|
||||
created_at,
|
||||
updated_at,
|
||||
security_version: 0,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn with_security_version(
|
||||
mut self,
|
||||
security_version: i64,
|
||||
) -> Result<Self, crate::DataLayerError> {
|
||||
if security_version < 0 {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(
|
||||
"user_sessions.security_version is negative".to_string(),
|
||||
));
|
||||
}
|
||||
self.security_version = security_version;
|
||||
Ok(self)
|
||||
}
|
||||
|
||||
pub fn hash_refresh_token(token: &str) -> String {
|
||||
use sha2::Digest;
|
||||
|
||||
@@ -442,8 +731,10 @@ impl StoredUserSessionRecord {
|
||||
let Some(rotated_at) = self.rotated_at else {
|
||||
return (false, false);
|
||||
};
|
||||
let age = now.signed_duration_since(rotated_at);
|
||||
if prev_hash == &token_hash
|
||||
&& now.signed_duration_since(rotated_at).num_seconds() <= Self::REFRESH_GRACE_SECONDS
|
||||
&& age >= chrono::Duration::zero()
|
||||
&& age <= chrono::Duration::seconds(Self::REFRESH_GRACE_SECONDS)
|
||||
{
|
||||
return (true, true);
|
||||
}
|
||||
@@ -744,6 +1035,22 @@ pub trait UserReadRepository: Send + Sync {
|
||||
record: UpsertUserGroupRecord,
|
||||
) -> Result<Option<StoredUserGroup>, crate::DataLayerError>;
|
||||
|
||||
/// Restore an existing user group only when its complete stored snapshot
|
||||
/// still equals `expected`. The compare and replacement must be atomic so
|
||||
/// an import rollback cannot overwrite a concurrent administrator update.
|
||||
/// `false` means the row is missing, the identities differ, or the current
|
||||
/// snapshot no longer matches the expected post-import state.
|
||||
async fn restore_user_group_if_matches(
|
||||
&self,
|
||||
expected: &StoredUserGroup,
|
||||
restored: &StoredUserGroup,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
let _ = (expected, restored);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic user group restore is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn delete_user_group(&self, group_id: &str) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn list_user_group_members(
|
||||
@@ -773,6 +1080,20 @@ pub trait UserReadRepository: Send + Sync {
|
||||
group_ids: &[String],
|
||||
) -> Result<Vec<StoredUserGroup>, crate::DataLayerError>;
|
||||
|
||||
/// Restore a user's group memberships only when the current set still
|
||||
/// equals the post-import set. The compare and replacement are atomic.
|
||||
async fn restore_user_groups_if_matches(
|
||||
&self,
|
||||
user_id: &str,
|
||||
expected_group_ids: &[String],
|
||||
restored_group_ids: &[String],
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
let _ = (user_id, expected_group_ids, restored_group_ids);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic user group restore is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn add_user_to_group(
|
||||
&self,
|
||||
group_id: &str,
|
||||
@@ -824,6 +1145,19 @@ pub trait UserReadRepository: Send + Sync {
|
||||
provider_user_id: &str,
|
||||
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn resolve_enabled_oauth_linked_user(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
provider_user_id: &str,
|
||||
provider_username: Option<&str>,
|
||||
provider_email: Option<&str>,
|
||||
extra_data: Option<Value>,
|
||||
verified_email: Option<&str>,
|
||||
touched_at: DateTime<Utc>,
|
||||
provider_enabled_snapshot: bool,
|
||||
) -> Result<ResolveOAuthLinkedUserOutcome, crate::DataLayerError>;
|
||||
|
||||
async fn touch_oauth_link(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
@@ -837,6 +1171,7 @@ pub trait UserReadRepository: Send + Sync {
|
||||
async fn create_oauth_auth_user(
|
||||
&self,
|
||||
email: Option<String>,
|
||||
email_verified: bool,
|
||||
username: String,
|
||||
created_at: DateTime<Utc>,
|
||||
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
|
||||
@@ -855,8 +1190,22 @@ pub trait UserReadRepository: Send + Sync {
|
||||
|
||||
async fn count_user_oauth_links(&self, user_id: &str) -> Result<u64, crate::DataLayerError>;
|
||||
|
||||
async fn has_oauth_links_for_provider(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn count_locked_users_if_oauth_provider_disabled(
|
||||
&self,
|
||||
_provider_type: &str,
|
||||
_enabled_provider_types_snapshot: &[String],
|
||||
_ldap_exclusive: bool,
|
||||
) -> Result<usize, crate::DataLayerError> {
|
||||
Ok(0)
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn upsert_user_oauth_link(
|
||||
async fn bind_user_oauth_link(
|
||||
&self,
|
||||
user_id: &str,
|
||||
provider_type: &str,
|
||||
@@ -865,13 +1214,51 @@ pub trait UserReadRepository: Send + Sync {
|
||||
provider_email: Option<&str>,
|
||||
extra_data: Option<Value>,
|
||||
linked_at: DateTime<Utc>,
|
||||
) -> Result<(), crate::DataLayerError>;
|
||||
) -> Result<BindUserOAuthLinkOutcome, crate::DataLayerError> {
|
||||
self.bind_user_oauth_link_if_provider_enabled(
|
||||
user_id,
|
||||
provider_type,
|
||||
provider_user_id,
|
||||
provider_username,
|
||||
provider_email,
|
||||
extra_data,
|
||||
linked_at,
|
||||
true,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn bind_user_oauth_link_if_provider_enabled(
|
||||
&self,
|
||||
user_id: &str,
|
||||
provider_type: &str,
|
||||
provider_user_id: &str,
|
||||
provider_username: Option<&str>,
|
||||
provider_email: Option<&str>,
|
||||
extra_data: Option<Value>,
|
||||
linked_at: DateTime<Utc>,
|
||||
provider_enabled_snapshot: bool,
|
||||
session_expectation: Option<&BindUserOAuthLinkSessionExpectation>,
|
||||
) -> Result<BindUserOAuthLinkOutcome, crate::DataLayerError>;
|
||||
|
||||
/// Marks an email as verified only if the user's current normalized email still matches.
|
||||
async fn upgrade_oauth_email_verification_if_matches(
|
||||
&self,
|
||||
user_id: &str,
|
||||
verified_email: &str,
|
||||
verified_at: DateTime<Utc>,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
/// Atomically deletes the requested link only when doing so leaves a valid login method.
|
||||
async fn delete_user_oauth_link(
|
||||
&self,
|
||||
user_id: &str,
|
||||
provider_type: &str,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
local_password_login_allowed: bool,
|
||||
enabled_provider_types_snapshot: &[String],
|
||||
) -> Result<DeleteUserOAuthLinkOutcome, crate::DataLayerError>;
|
||||
|
||||
async fn get_or_create_ldap_auth_user(
|
||||
&self,
|
||||
@@ -891,10 +1278,44 @@ pub trait UserReadRepository: Send + Sync {
|
||||
async fn update_local_auth_user_profile(
|
||||
&self,
|
||||
user_id: &str,
|
||||
email_present: bool,
|
||||
email: Option<String>,
|
||||
email_verified: Option<bool>,
|
||||
username: Option<String>,
|
||||
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
|
||||
|
||||
/// Restore all user fields represented by an aggregate import checkpoint in
|
||||
/// one compare-and-write transaction. The password hash is deliberately
|
||||
/// excluded and must be restored through the nullable password CAS method.
|
||||
/// Implementations must not write anything when the current state differs
|
||||
/// from `expected_auth` (or the exported rate/settings fields).
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn restore_local_auth_user_state_if_matches(
|
||||
&self,
|
||||
expected_auth: &StoredUserAuthRecord,
|
||||
restored_auth: &StoredUserAuthRecord,
|
||||
expected_export: &StoredUserExportRow,
|
||||
restored_export: &StoredUserExportRow,
|
||||
expected_model_capability_settings: Option<&Value>,
|
||||
restored_model_capability_settings: Option<Value>,
|
||||
expected_feature_settings: Option<&Value>,
|
||||
restored_feature_settings: Option<Value>,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
let _ = (
|
||||
expected_auth,
|
||||
restored_auth,
|
||||
expected_export,
|
||||
restored_export,
|
||||
expected_model_capability_settings,
|
||||
restored_model_capability_settings,
|
||||
expected_feature_settings,
|
||||
restored_feature_settings,
|
||||
);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic user state restore is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn update_local_auth_user_password_hash(
|
||||
&self,
|
||||
user_id: &str,
|
||||
@@ -902,6 +1323,37 @@ pub trait UserReadRepository: Send + Sync {
|
||||
updated_at: DateTime<Utc>,
|
||||
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
|
||||
|
||||
/// Replace a user's password hash, including an explicit `NULL`, only when the current hash
|
||||
/// still equals `expected_password_hash`. The compare and write must be atomic.
|
||||
async fn restore_local_auth_user_password_hash_if_matches(
|
||||
&self,
|
||||
user_id: &str,
|
||||
expected_password_hash: Option<&str>,
|
||||
password_hash: Option<String>,
|
||||
updated_at: DateTime<Utc>,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
let _ = (user_id, expected_password_hash, password_hash, updated_at);
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic nullable password restore is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn reset_local_auth_user_password_and_revoke_sessions(
|
||||
&self,
|
||||
user_id: &str,
|
||||
password_hash: String,
|
||||
changed_at: DateTime<Utc>,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
async fn change_local_auth_password_and_revoke_sessions(
|
||||
&self,
|
||||
user_id: &str,
|
||||
current_session_id: &str,
|
||||
expected_password_hash: Option<&str>,
|
||||
next_password_hash: String,
|
||||
changed_at: DateTime<Utc>,
|
||||
) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
async fn update_local_auth_user_admin_fields(
|
||||
&self,
|
||||
@@ -963,6 +1415,23 @@ pub trait UserReadRepository: Send + Sync {
|
||||
|
||||
async fn delete_local_auth_user(&self, user_id: &str) -> Result<bool, crate::DataLayerError>;
|
||||
|
||||
/// Delete a local-auth user only when no wallet currently belongs to the user.
|
||||
///
|
||||
/// Authentication provisioning compensation uses this guard after removing a
|
||||
/// wallet it can prove ownership of. Implementations must evaluate the
|
||||
/// wallet absence check in the same database transaction as the user delete;
|
||||
/// a default implementation is deliberately fail-closed for repositories
|
||||
/// that cannot provide that atomicity.
|
||||
async fn delete_local_auth_user_if_wallet_absent(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<bool, crate::DataLayerError> {
|
||||
let _ = user_id;
|
||||
Err(crate::DataLayerError::InvalidInput(
|
||||
"atomic user deletion without a wallet is not available".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
async fn read_user_preferences(
|
||||
&self,
|
||||
user_id: &str,
|
||||
@@ -989,6 +1458,12 @@ pub trait UserReadRepository: Send + Sync {
|
||||
session: &StoredUserSessionRecord,
|
||||
) -> Result<Option<StoredUserSessionRecord>, crate::DataLayerError>;
|
||||
|
||||
async fn create_user_session_if_password_matches(
|
||||
&self,
|
||||
session: &StoredUserSessionRecord,
|
||||
expected_password_hash: &str,
|
||||
) -> Result<Option<StoredUserSessionRecord>, crate::DataLayerError>;
|
||||
|
||||
async fn touch_user_session(
|
||||
&self,
|
||||
user_id: &str,
|
||||
@@ -1011,7 +1486,7 @@ pub trait UserReadRepository: Send + Sync {
|
||||
&self,
|
||||
user_id: &str,
|
||||
session_id: &str,
|
||||
previous_refresh_token_hash: &str,
|
||||
expected_refresh_token_hash: &str,
|
||||
next_refresh_token_hash: &str,
|
||||
rotated_at: DateTime<Utc>,
|
||||
expires_at: DateTime<Utc>,
|
||||
@@ -1104,7 +1579,9 @@ fn parse_string_list_value(
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, crate::DataLayerError> {
|
||||
match value {
|
||||
Value::Null => Ok(None),
|
||||
Value::Null => Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains JSON null; use SQL NULL for an unset policy"
|
||||
))),
|
||||
Value::Array(array) => parse_string_list_array(array, field_name).map(Some),
|
||||
Value::String(raw) => parse_embedded_string_list(raw, field_name),
|
||||
_ => Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
@@ -1118,8 +1595,15 @@ fn parse_embedded_string_list(
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, crate::DataLayerError> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
|
||||
return Ok(None);
|
||||
if raw.is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains an empty string"
|
||||
)));
|
||||
}
|
||||
if raw.eq_ignore_ascii_case("null") {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains stringified JSON null; use SQL NULL for an unset policy"
|
||||
)));
|
||||
}
|
||||
|
||||
if let Ok(decoded) = serde_json::from_str::<Value>(raw) {
|
||||
@@ -1141,9 +1625,12 @@ fn parse_string_list_array(
|
||||
)));
|
||||
};
|
||||
let item = item.trim();
|
||||
if !item.is_empty() {
|
||||
items.push(item.to_string());
|
||||
if item.is_empty() {
|
||||
return Err(crate::DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains an empty item"
|
||||
)));
|
||||
}
|
||||
items.push(item.to_string());
|
||||
}
|
||||
Ok(items)
|
||||
}
|
||||
@@ -1154,10 +1641,54 @@ mod tests {
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{
|
||||
legacy_list_policy_mode, StoredUserAuthRecord, StoredUserExportRow,
|
||||
is_valid_bcrypt_hash, last_oauth_unbind_denial, legacy_list_policy_mode,
|
||||
DeleteUserOAuthLinkOutcome, StoredUserAuthRecord, StoredUserExportRow,
|
||||
StoredUserPreferenceRecord, StoredUserSessionRecord,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn classifies_last_oauth_unbind_by_remaining_login_method() {
|
||||
let valid_hash =
|
||||
bcrypt::hash("Secret123!", bcrypt::DEFAULT_COST).expect("bcrypt fixture should hash");
|
||||
|
||||
assert_eq!(
|
||||
last_oauth_unbind_denial("oauth", None, false),
|
||||
Some(DeleteUserOAuthLinkOutcome::LastOAuthBinding)
|
||||
);
|
||||
assert_eq!(
|
||||
last_oauth_unbind_denial("local", None, true),
|
||||
Some(DeleteUserOAuthLinkOutcome::LastLoginMethod)
|
||||
);
|
||||
assert_eq!(
|
||||
last_oauth_unbind_denial("local", Some("not-a-password-hash"), true),
|
||||
Some(DeleteUserOAuthLinkOutcome::LastLoginMethod)
|
||||
);
|
||||
assert_eq!(
|
||||
last_oauth_unbind_denial("local", Some(&valid_hash), false),
|
||||
Some(DeleteUserOAuthLinkOutcome::LastLoginMethod)
|
||||
);
|
||||
assert_eq!(
|
||||
last_oauth_unbind_denial("local", Some(&valid_hash), true),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn validates_complete_bcrypt_encoding_and_cost() {
|
||||
let valid_hash =
|
||||
bcrypt::hash("Secret123!", bcrypt::DEFAULT_COST).expect("bcrypt fixture should hash");
|
||||
assert!(is_valid_bcrypt_hash(&valid_hash));
|
||||
|
||||
for invalid in [
|
||||
format!("$2b$99${}", &valid_hash[7..]),
|
||||
format!("$2x$12${}", &valid_hash[7..]),
|
||||
format!("$2b$12${}", "!".repeat(53)),
|
||||
valid_hash[..59].to_string(),
|
||||
] {
|
||||
assert!(!is_valid_bcrypt_hash(&invalid), "accepted {invalid}");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_user_export_row_with_allowed_lists() {
|
||||
let row = StoredUserExportRow::new(
|
||||
@@ -1203,7 +1734,7 @@ mod tests {
|
||||
"user".to_string(),
|
||||
"local".to_string(),
|
||||
Some(serde_json::json!("[\"openai\"]")),
|
||||
Some(serde_json::json!("null")),
|
||||
None,
|
||||
Some(serde_json::json!("gpt-4.1")),
|
||||
None,
|
||||
Some(Value::Null),
|
||||
@@ -1217,6 +1748,29 @@ mod tests {
|
||||
assert_eq!(row.model_capability_settings, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn stored_user_security_lists_distinguish_sql_null_from_malformed_json_null() {
|
||||
assert_eq!(
|
||||
super::parse_string_list(None, "users.allowed_providers")
|
||||
.expect("SQL NULL should remain an unset policy"),
|
||||
None
|
||||
);
|
||||
assert!(
|
||||
super::parse_string_list(Some(serde_json::Value::Null), "users.allowed_providers")
|
||||
.is_err()
|
||||
);
|
||||
assert!(super::parse_string_list(
|
||||
Some(serde_json::json!("null")),
|
||||
"users.allowed_providers"
|
||||
)
|
||||
.is_err());
|
||||
assert!(super::parse_string_list(
|
||||
Some(serde_json::json!([" "])),
|
||||
"users.allowed_providers"
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn rejects_object_allowed_providers_for_user_export_row() {
|
||||
let result = StoredUserExportRow::new(
|
||||
@@ -1266,6 +1820,73 @@ mod tests {
|
||||
assert_eq!(row.allowed_models, Some(vec!["gpt-4.1".to_string()]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn user_record_debug_output_redacts_password_and_refresh_hashes() {
|
||||
let password_hash = "debug-secret-password-hash";
|
||||
let auth = StoredUserAuthRecord::new(
|
||||
"user-debug".to_string(),
|
||||
Some("[email protected]".to_string()),
|
||||
true,
|
||||
"debug-user".to_string(),
|
||||
Some(password_hash.to_string()),
|
||||
"user".to_string(),
|
||||
"local".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
false,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("auth record should build");
|
||||
let export = StoredUserExportRow::new(
|
||||
"user-debug".to_string(),
|
||||
Some("[email protected]".to_string()),
|
||||
true,
|
||||
"debug-user".to_string(),
|
||||
Some(password_hash.to_string()),
|
||||
"user".to_string(),
|
||||
"local".to_string(),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
true,
|
||||
)
|
||||
.expect("export record should build");
|
||||
let current_hash = "debug-secret-current-refresh-hash";
|
||||
let previous_hash = "debug-secret-previous-refresh-hash";
|
||||
let session = StoredUserSessionRecord::new(
|
||||
"session-debug".to_string(),
|
||||
"user-debug".to_string(),
|
||||
"device-debug".to_string(),
|
||||
None,
|
||||
current_hash.to_string(),
|
||||
Some(previous_hash.to_string()),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("session should build");
|
||||
|
||||
for rendered in [format!("{auth:?}"), format!("{export:?}")] {
|
||||
assert!(!rendered.contains(password_hash));
|
||||
assert!(rendered.contains("[REDACTED]"));
|
||||
}
|
||||
let rendered = format!("{session:?}");
|
||||
assert!(!rendered.contains(current_hash));
|
||||
assert!(!rendered.contains(previous_hash));
|
||||
assert!(rendered.contains("[REDACTED]"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn legacy_policy_mode_treats_empty_lists_as_unrestricted() {
|
||||
assert_eq!(legacy_list_policy_mode(&None), "unrestricted");
|
||||
@@ -1306,6 +1927,29 @@ mod tests {
|
||||
session.verify_refresh_token("current-token", now),
|
||||
(true, false)
|
||||
);
|
||||
|
||||
let future_rotation = StoredUserSessionRecord::new(
|
||||
"session-2".to_string(),
|
||||
"user-1".to_string(),
|
||||
"device-1".to_string(),
|
||||
None,
|
||||
StoredUserSessionRecord::hash_refresh_token("current-token"),
|
||||
Some(StoredUserSessionRecord::hash_refresh_token("prev-token")),
|
||||
Some(now + Duration::seconds(1)),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.expect("session should build");
|
||||
assert_eq!(
|
||||
future_rotation.verify_refresh_token("prev-token", now),
|
||||
(false, false)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
||||
@@ -1,6 +1,8 @@
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
|
||||
const SAFE_VIDEO_URL_QUERY_KEYS: &[(&str, &str)] = &[("alt", "media")];
|
||||
|
||||
#[derive(
|
||||
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
|
||||
)]
|
||||
@@ -84,6 +86,13 @@ pub struct StoredVideoTask {
|
||||
}
|
||||
|
||||
impl StoredVideoTask {
|
||||
pub fn effective_api_format(&self) -> Option<&str> {
|
||||
effective_video_task_api_format(
|
||||
self.client_api_format.as_deref(),
|
||||
self.provider_api_format.as_deref(),
|
||||
)
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn new(
|
||||
id: String,
|
||||
@@ -168,7 +177,7 @@ impl StoredVideoTask {
|
||||
None => None,
|
||||
};
|
||||
|
||||
Ok(Self {
|
||||
let mut task = Self {
|
||||
id,
|
||||
short_id,
|
||||
request_id,
|
||||
@@ -206,7 +215,84 @@ impl StoredVideoTask {
|
||||
error_message,
|
||||
video_url,
|
||||
request_metadata,
|
||||
})
|
||||
};
|
||||
task.sanitize_persisted_diagnostics();
|
||||
Ok(task)
|
||||
}
|
||||
|
||||
fn sanitize_persisted_diagnostics(&mut self) {
|
||||
self.prompt = None;
|
||||
self.original_request_body = None;
|
||||
self.progress_message = None;
|
||||
self.error_code = sanitize_video_task_error_code(self.error_code.take());
|
||||
self.error_message = None;
|
||||
self.video_url = sanitize_video_task_url(
|
||||
self.client_api_format.as_deref(),
|
||||
self.provider_api_format.as_deref(),
|
||||
self.video_url.take(),
|
||||
);
|
||||
self.request_metadata = None;
|
||||
}
|
||||
|
||||
/// Verifies that an update still refers to the task identity persisted for `id`.
|
||||
///
|
||||
/// These fields select the owner, upstream target, request shape, or immutable
|
||||
/// creation identity of a video task. Repository implementations must reject an
|
||||
/// upsert that changes any of them instead of treating possession of `id` as
|
||||
/// permission to replace the row.
|
||||
///
|
||||
/// `created_at_unix_ms` is deliberately not compared because some snapshot
|
||||
/// projections recompute it. Repositories preserve the already stored creation
|
||||
/// time while accepting an otherwise matching lifecycle update.
|
||||
pub fn ensure_immutable_identity_matches(
|
||||
&self,
|
||||
incoming: &UpsertVideoTask,
|
||||
) -> Result<(), crate::DataLayerError> {
|
||||
let mismatched_field = if self.id != incoming.id {
|
||||
Some("id")
|
||||
} else if self.short_id != incoming.short_id {
|
||||
Some("short_id")
|
||||
} else if self.request_id != incoming.request_id {
|
||||
Some("request_id")
|
||||
} else if self.user_id != incoming.user_id {
|
||||
Some("user_id")
|
||||
} else if self.api_key_id != incoming.api_key_id {
|
||||
Some("api_key_id")
|
||||
} else if self.external_task_id != incoming.external_task_id {
|
||||
Some("external_task_id")
|
||||
} else if self.provider_id != incoming.provider_id {
|
||||
Some("provider_id")
|
||||
} else if self.endpoint_id != incoming.endpoint_id {
|
||||
Some("endpoint_id")
|
||||
} else if self.key_id != incoming.key_id {
|
||||
Some("key_id")
|
||||
} else if self.client_api_format != incoming.client_api_format {
|
||||
Some("client_api_format")
|
||||
} else if self.provider_api_format != incoming.provider_api_format {
|
||||
Some("provider_api_format")
|
||||
} else if self.format_converted != incoming.format_converted {
|
||||
Some("format_converted")
|
||||
} else if self.model != incoming.model {
|
||||
Some("model")
|
||||
} else if self.duration_seconds != incoming.duration_seconds {
|
||||
Some("duration_seconds")
|
||||
} else if self.resolution != incoming.resolution {
|
||||
Some("resolution")
|
||||
} else if self.aspect_ratio != incoming.aspect_ratio {
|
||||
Some("aspect_ratio")
|
||||
} else if self.size != incoming.size {
|
||||
Some("size")
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
match mismatched_field {
|
||||
Some(field) => Err(crate::DataLayerError::InvalidInput(format!(
|
||||
"video task {} conflicts with persisted immutable field {field}",
|
||||
incoming.id
|
||||
))),
|
||||
None => Ok(()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -252,7 +338,24 @@ pub struct UpsertVideoTask {
|
||||
}
|
||||
|
||||
impl UpsertVideoTask {
|
||||
pub fn into_stored(self) -> StoredVideoTask {
|
||||
pub fn sanitize_for_persistence(&mut self) {
|
||||
self.username = None;
|
||||
self.api_key_name = None;
|
||||
self.prompt = None;
|
||||
self.original_request_body = None;
|
||||
self.progress_message = None;
|
||||
self.error_code = sanitize_video_task_error_code(self.error_code.take());
|
||||
self.error_message = None;
|
||||
self.video_url = sanitize_video_task_url(
|
||||
self.client_api_format.as_deref(),
|
||||
self.provider_api_format.as_deref(),
|
||||
self.video_url.take(),
|
||||
);
|
||||
self.request_metadata = None;
|
||||
}
|
||||
|
||||
pub fn into_stored(mut self) -> StoredVideoTask {
|
||||
self.sanitize_for_persistence();
|
||||
StoredVideoTask {
|
||||
id: self.id,
|
||||
short_id: self.short_id,
|
||||
@@ -295,6 +398,75 @@ impl UpsertVideoTask {
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_video_task_error_code(value: Option<String>) -> Option<String> {
|
||||
let value = value?.trim().to_ascii_lowercase();
|
||||
if value.is_empty() {
|
||||
return None;
|
||||
}
|
||||
Some(match value.as_str() {
|
||||
"authentication_error"
|
||||
| "cancelled"
|
||||
| "content_policy_violation"
|
||||
| "expired"
|
||||
| "invalid_request"
|
||||
| "not_found"
|
||||
| "permission_denied"
|
||||
| "poll_permanent_error"
|
||||
| "poll_timeout"
|
||||
| "provider_error"
|
||||
| "rate_limit_exceeded"
|
||||
| "server_error"
|
||||
| "unknown" => value,
|
||||
_ => "provider_error".to_string(),
|
||||
})
|
||||
}
|
||||
|
||||
fn sanitize_video_task_url(
|
||||
client_api_format: Option<&str>,
|
||||
provider_api_format: Option<&str>,
|
||||
value: Option<String>,
|
||||
) -> Option<String> {
|
||||
if effective_video_task_api_format(client_api_format, provider_api_format)
|
||||
!= Some("gemini:video")
|
||||
{
|
||||
return None;
|
||||
}
|
||||
let mut url = url::Url::parse(value?.trim()).ok()?;
|
||||
if !matches!(url.scheme(), "http" | "https")
|
||||
|| url.host_str().is_none()
|
||||
|| !url.username().is_empty()
|
||||
|| url.password().is_some()
|
||||
{
|
||||
return None;
|
||||
}
|
||||
|
||||
let query = url
|
||||
.query_pairs()
|
||||
.filter(|(key, value)| SAFE_VIDEO_URL_QUERY_KEYS.contains(&(key.as_ref(), value.as_ref())))
|
||||
.map(|(key, value)| (key.into_owned(), value.into_owned()))
|
||||
.collect::<Vec<_>>();
|
||||
url.set_query(None);
|
||||
if !query.is_empty() {
|
||||
url.query_pairs_mut().extend_pairs(query);
|
||||
}
|
||||
url.set_fragment(None);
|
||||
Some(url.into())
|
||||
}
|
||||
|
||||
fn effective_video_task_api_format<'a>(
|
||||
client_api_format: Option<&'a str>,
|
||||
provider_api_format: Option<&'a str>,
|
||||
) -> Option<&'a str> {
|
||||
provider_api_format
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.or_else(|| {
|
||||
client_api_format
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
})
|
||||
}
|
||||
|
||||
impl From<StoredVideoTask> for UpsertVideoTask {
|
||||
fn from(task: StoredVideoTask) -> Self {
|
||||
Self {
|
||||
@@ -376,6 +548,15 @@ pub trait VideoTaskReadRepository: Send + Sync {
|
||||
key: VideoTaskLookupKey<'_>,
|
||||
) -> Result<Option<StoredVideoTask>, crate::DataLayerError>;
|
||||
|
||||
/// Resolve a public task identifier only when the persisted task belongs
|
||||
/// to `user_id`. The lookup key and owner predicate must be evaluated by
|
||||
/// one repository operation.
|
||||
async fn find_for_user(
|
||||
&self,
|
||||
key: VideoTaskLookupKey<'_>,
|
||||
user_id: &str,
|
||||
) -> Result<Option<StoredVideoTask>, crate::DataLayerError>;
|
||||
|
||||
async fn list_active(
|
||||
&self,
|
||||
limit: usize,
|
||||
@@ -468,7 +649,7 @@ fn coerce_optional_unix_secs(
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{StoredVideoTask, VideoTaskStatus};
|
||||
use super::{StoredVideoTask, UpsertVideoTask, VideoTaskStatus};
|
||||
|
||||
#[allow(clippy::type_complexity)]
|
||||
fn base_new_args() -> (
|
||||
@@ -615,4 +796,203 @@ mod tests {
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn immutable_identity_validation_rejects_every_protected_field() {
|
||||
let task = UpsertVideoTask {
|
||||
id: "task-1".to_string(),
|
||||
short_id: Some("short-1".to_string()),
|
||||
request_id: "request-1".to_string(),
|
||||
user_id: Some("user-1".to_string()),
|
||||
api_key_id: Some("api-key-1".to_string()),
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
external_task_id: Some("external-1".to_string()),
|
||||
provider_id: Some("provider-1".to_string()),
|
||||
endpoint_id: Some("endpoint-1".to_string()),
|
||||
key_id: Some("key-1".to_string()),
|
||||
client_api_format: Some("openai:video".to_string()),
|
||||
provider_api_format: Some("gemini:video".to_string()),
|
||||
format_converted: true,
|
||||
model: Some("video-model".to_string()),
|
||||
prompt: None,
|
||||
original_request_body: None,
|
||||
duration_seconds: Some(4),
|
||||
resolution: Some("720p".to_string()),
|
||||
aspect_ratio: Some("16:9".to_string()),
|
||||
size: Some("1280x720".to_string()),
|
||||
status: VideoTaskStatus::Submitted,
|
||||
progress_percent: 0,
|
||||
progress_message: None,
|
||||
retry_count: 0,
|
||||
poll_interval_seconds: 10,
|
||||
next_poll_at_unix_secs: Some(10),
|
||||
poll_count: 0,
|
||||
max_poll_count: 360,
|
||||
created_at_unix_ms: 1,
|
||||
submitted_at_unix_secs: Some(1),
|
||||
completed_at_unix_secs: None,
|
||||
updated_at_unix_secs: 1,
|
||||
error_code: None,
|
||||
error_message: None,
|
||||
video_url: None,
|
||||
request_metadata: None,
|
||||
};
|
||||
let stored = task.clone().into_stored();
|
||||
|
||||
let same_identity_update = UpsertVideoTask {
|
||||
status: VideoTaskStatus::Completed,
|
||||
progress_percent: 100,
|
||||
created_at_unix_ms: 2,
|
||||
completed_at_unix_secs: Some(2),
|
||||
updated_at_unix_secs: 2,
|
||||
..task.clone()
|
||||
};
|
||||
stored
|
||||
.ensure_immutable_identity_matches(&same_identity_update)
|
||||
.expect("mutable state changes should keep the same identity");
|
||||
|
||||
macro_rules! assert_identity_conflict {
|
||||
($field:ident, $value:expr) => {{
|
||||
let mut conflicting = task.clone();
|
||||
conflicting.$field = $value;
|
||||
let error = stored
|
||||
.ensure_immutable_identity_matches(&conflicting)
|
||||
.expect_err(concat!(stringify!($field), " should be immutable"));
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains(concat!("immutable field ", stringify!($field))),
|
||||
"unexpected error for {}: {error}",
|
||||
stringify!($field)
|
||||
);
|
||||
}};
|
||||
}
|
||||
|
||||
assert_identity_conflict!(id, "task-2".to_string());
|
||||
assert_identity_conflict!(short_id, Some("short-2".to_string()));
|
||||
assert_identity_conflict!(request_id, "request-2".to_string());
|
||||
assert_identity_conflict!(user_id, Some("user-2".to_string()));
|
||||
assert_identity_conflict!(api_key_id, Some("api-key-2".to_string()));
|
||||
assert_identity_conflict!(external_task_id, Some("external-2".to_string()));
|
||||
assert_identity_conflict!(provider_id, Some("provider-2".to_string()));
|
||||
assert_identity_conflict!(endpoint_id, Some("endpoint-2".to_string()));
|
||||
assert_identity_conflict!(key_id, Some("key-2".to_string()));
|
||||
assert_identity_conflict!(client_api_format, Some("gemini:video".to_string()));
|
||||
assert_identity_conflict!(provider_api_format, Some("openai:video".to_string()));
|
||||
assert_identity_conflict!(format_converted, false);
|
||||
assert_identity_conflict!(model, Some("other-model".to_string()));
|
||||
assert_identity_conflict!(duration_seconds, Some(8));
|
||||
assert_identity_conflict!(resolution, Some("1080p".to_string()));
|
||||
assert_identity_conflict!(aspect_ratio, Some("9:16".to_string()));
|
||||
assert_identity_conflict!(size, Some("1920x1080".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn upsert_sanitization_drops_sensitive_diagnostics() {
|
||||
let mut task = UpsertVideoTask {
|
||||
id: "task-1".to_string(),
|
||||
short_id: None,
|
||||
request_id: "request-1".to_string(),
|
||||
user_id: Some("user-1".to_string()),
|
||||
api_key_id: Some("api-key-1".to_string()),
|
||||
username: Some("private-user-name".to_string()),
|
||||
api_key_name: Some("private-key-name".to_string()),
|
||||
external_task_id: Some("upstream-1".to_string()),
|
||||
provider_id: Some("provider-1".to_string()),
|
||||
endpoint_id: Some("endpoint-1".to_string()),
|
||||
key_id: Some("key-1".to_string()),
|
||||
client_api_format: Some("openai:video".to_string()),
|
||||
provider_api_format: Some("openai:video".to_string()),
|
||||
format_converted: false,
|
||||
model: Some("video-model".to_string()),
|
||||
prompt: Some("prompt".to_string()),
|
||||
original_request_body: Some(serde_json::json!({
|
||||
"prompt": "private prompt",
|
||||
"api_key": "secret"
|
||||
})),
|
||||
duration_seconds: Some(4),
|
||||
resolution: Some("720p".to_string()),
|
||||
aspect_ratio: Some("16:9".to_string()),
|
||||
size: Some("1280x720".to_string()),
|
||||
status: VideoTaskStatus::Failed,
|
||||
progress_percent: 100,
|
||||
progress_message: Some("provider response: secret".to_string()),
|
||||
retry_count: 1,
|
||||
poll_interval_seconds: 10,
|
||||
next_poll_at_unix_secs: None,
|
||||
poll_count: 2,
|
||||
max_poll_count: 360,
|
||||
created_at_unix_ms: 1,
|
||||
submitted_at_unix_secs: Some(1),
|
||||
completed_at_unix_secs: Some(2),
|
||||
updated_at_unix_secs: 2,
|
||||
error_code: Some("secret provider code".to_string()),
|
||||
error_message: Some("Authorization: Bearer secret".to_string()),
|
||||
video_url: Some("https://cdn.example.test/video.mp4?token=secret".to_string()),
|
||||
request_metadata: Some(serde_json::json!({
|
||||
"rust_local_snapshot": {"transport": {"headers": {"authorization": "secret"}}},
|
||||
"poll_raw_response": {"error": "secret"}
|
||||
})),
|
||||
};
|
||||
|
||||
task.sanitize_for_persistence();
|
||||
|
||||
assert_eq!(task.user_id.as_deref(), Some("user-1"));
|
||||
assert_eq!(task.api_key_id.as_deref(), Some("api-key-1"));
|
||||
assert_eq!(task.username, None);
|
||||
assert_eq!(task.api_key_name, None);
|
||||
assert_eq!(task.original_request_body, None);
|
||||
assert_eq!(task.progress_message, None);
|
||||
assert_eq!(task.error_message, None);
|
||||
assert_eq!(task.error_code.as_deref(), Some("provider_error"));
|
||||
assert_eq!(task.video_url, None);
|
||||
assert_eq!(task.request_metadata, None);
|
||||
assert_eq!(task.prompt, None);
|
||||
assert_eq!(task.duration_seconds, Some(4));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn upsert_sanitization_keeps_only_noncredential_video_urls() {
|
||||
let mut args = base_new_args();
|
||||
args.12 = Some("gemini:video".to_string());
|
||||
args.15 = Some("private prompt".to_string());
|
||||
args.35 = Some(
|
||||
"https://cdn.example.test/video.mp4?key=secret&alt=media&signature=private#fragment"
|
||||
.to_string(),
|
||||
);
|
||||
let task = StoredVideoTask::new(
|
||||
args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7, args.8, args.9,
|
||||
args.10, args.11, args.12, args.13, args.14, args.15, args.16, args.17, args.18,
|
||||
args.19, args.20, args.21, args.22, args.23, args.24, args.25, args.26, args.27,
|
||||
args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36,
|
||||
)
|
||||
.expect("stored task should build");
|
||||
assert_eq!(task.prompt, None);
|
||||
assert_eq!(
|
||||
task.video_url.as_deref(),
|
||||
Some("https://cdn.example.test/video.mp4?alt=media")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn upsert_sanitization_uses_client_format_when_legacy_provider_format_is_blank() {
|
||||
let mut args = base_new_args();
|
||||
args.11 = Some("gemini:video".to_string());
|
||||
args.12 = Some(" ".to_string());
|
||||
args.35 =
|
||||
Some("https://cdn.example.test/video.mp4?key=secret&alt=media#fragment".to_string());
|
||||
let task = StoredVideoTask::new(
|
||||
args.0, args.1, args.2, args.3, args.4, args.5, args.6, args.7, args.8, args.9,
|
||||
args.10, args.11, args.12, args.13, args.14, args.15, args.16, args.17, args.18,
|
||||
args.19, args.20, args.21, args.22, args.23, args.24, args.25, args.26, args.27,
|
||||
args.28, args.29, args.30, args.31, args.32, args.33, args.34, args.35, args.36,
|
||||
)
|
||||
.expect("legacy stored task should build");
|
||||
|
||||
assert_eq!(
|
||||
task.video_url.as_deref(),
|
||||
Some("https://cdn.example.test/video.mp4?alt=media")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -19,6 +19,11 @@ pub struct WalletReadSeed {
|
||||
pub payment_callbacks: Vec<StoredAdminPaymentCallback>,
|
||||
pub wallet_transactions: Vec<StoredAdminWalletTransaction>,
|
||||
pub refunds: Vec<StoredAdminWalletRefund>,
|
||||
/// `(user_id, idempotency_key, refund_id)` entries for mutable in-memory
|
||||
/// repositories. Refund read records intentionally do not expose their
|
||||
/// idempotency keys, so callers must seed this private write index
|
||||
/// explicitly when replay behavior matters.
|
||||
pub refund_idempotency: Vec<(String, String, String)>,
|
||||
pub redeem_batches: Vec<StoredAdminRedeemCodeBatch>,
|
||||
pub redeem_codes: Vec<StoredAdminRedeemCode>,
|
||||
}
|
||||
@@ -248,7 +253,7 @@ impl WalletReadSnapshot {
|
||||
let effective = if order.status == "pending"
|
||||
&& order
|
||||
.expires_at_unix_secs
|
||||
.is_some_and(|value| value < now_unix_secs)
|
||||
.is_some_and(|value| value <= now_unix_secs)
|
||||
{
|
||||
"expired"
|
||||
} else {
|
||||
@@ -282,6 +287,14 @@ impl WalletReadSnapshot {
|
||||
.payment_orders
|
||||
.values()
|
||||
.filter(|order| order.user_id.as_deref() == Some(user_id))
|
||||
.filter(|order| {
|
||||
order
|
||||
.gateway_response
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("order_kind"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
!= Some("plan_purchase")
|
||||
})
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
items.sort_by_key(|item| std::cmp::Reverse(item.created_at_unix_ms));
|
||||
@@ -317,7 +330,15 @@ impl WalletReadSnapshot {
|
||||
) -> Option<StoredAdminPaymentOrder> {
|
||||
self.payment_orders
|
||||
.get(order_id)
|
||||
.filter(|order| order.user_id.as_deref() == Some(user_id))
|
||||
.filter(|order| {
|
||||
order.user_id.as_deref() == Some(user_id)
|
||||
&& order
|
||||
.gateway_response
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("order_kind"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
!= Some("plan_purchase")
|
||||
})
|
||||
.cloned()
|
||||
}
|
||||
|
||||
@@ -441,3 +462,51 @@ fn page<T, P>(items: Vec<T>, offset: usize, limit: usize, build: impl Fn(Vec<T>,
|
||||
let total = items.len() as u64;
|
||||
build(items.into_iter().skip(offset).take(limit).collect(), total)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
AdminPaymentOrderListQuery, StoredAdminPaymentOrder, WalletReadSeed, WalletReadSnapshot,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn payment_orders_expire_at_the_exact_boundary() {
|
||||
let order = StoredAdminPaymentOrder {
|
||||
id: "order-boundary".to_string(),
|
||||
order_no: "order-boundary".to_string(),
|
||||
wallet_id: "wallet-boundary".to_string(),
|
||||
user_id: Some("user-boundary".to_string()),
|
||||
amount_usd: 1.0,
|
||||
pay_amount: None,
|
||||
pay_currency: None,
|
||||
exchange_rate: None,
|
||||
refunded_amount_usd: 0.0,
|
||||
refundable_amount_usd: 0.0,
|
||||
payment_method: "epay".to_string(),
|
||||
payment_provider: Some("epay".to_string()),
|
||||
order_kind: "wallet_recharge".to_string(),
|
||||
gateway_order_id: None,
|
||||
gateway_response: None,
|
||||
status: "pending".to_string(),
|
||||
created_at_unix_ms: 1,
|
||||
paid_at_unix_secs: None,
|
||||
credited_at_unix_secs: None,
|
||||
expires_at_unix_secs: Some(100),
|
||||
};
|
||||
let snapshot = WalletReadSnapshot::new(WalletReadSeed {
|
||||
payment_orders: vec![order],
|
||||
..WalletReadSeed::default()
|
||||
});
|
||||
let page = snapshot.list_admin_payment_orders(
|
||||
&AdminPaymentOrderListQuery {
|
||||
status: Some("expired".to_string()),
|
||||
payment_method: None,
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
},
|
||||
100,
|
||||
);
|
||||
assert_eq!(page.total, 1);
|
||||
assert_eq!(page.items[0].id, "order-boundary");
|
||||
}
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user