mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-06 17:37:47 +08:00
fix: restore security hardening compatibility and validation
Restore authorized rule reveal, explicit full HTTP capture and retention, video task business fields, and valid payment URLs. Add opt-in credential preservation for trusted recovery, fix frontend type contracts and async races, and eliminate PostgreSQL test fixture resource leaks. Document audit coverage and successful fmt and CI-scoped Clippy checks.
This commit is contained in:
@@ -1,13 +1,8 @@
|
||||
use std::{
|
||||
path::{Path, PathBuf},
|
||||
process::{Child, Command, Stdio},
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use sqlx::{query, query_scalar, Connection, PgConnection, PgPool};
|
||||
use sqlx::{query, query_scalar, PgPool};
|
||||
|
||||
use super::{pending_backfills, pending_backfills_from_applied, run_backfills, AppliedBackfill};
|
||||
use crate::lifecycle::migrate::prepare_database_for_startup;
|
||||
use crate::lifecycle::postgres_test_support::ManagedPostgresServer;
|
||||
|
||||
const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_VERSION: i64 = 20260517012000;
|
||||
const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_SQL: &str =
|
||||
@@ -84,194 +79,6 @@ fn corrected_legacy_backfill_is_not_requeued_after_application() {
|
||||
assert!(!pending_versions.contains(&LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_VERSION));
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ManagedPostgresServer {
|
||||
child: Option<Child>,
|
||||
workdir: PathBuf,
|
||||
database_url: String,
|
||||
}
|
||||
|
||||
impl ManagedPostgresServer {
|
||||
async fn try_start() -> Result<Option<Self>, Box<dyn std::error::Error>> {
|
||||
let initdb_bin = std::env::var("AETHER_INITDB_BIN")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or_else(|| "initdb".to_string());
|
||||
let postgres_bin = std::env::var("AETHER_POSTGRES_BIN")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or_else(|| "postgres".to_string());
|
||||
|
||||
if !command_exists(&initdb_bin) || !command_exists(&postgres_bin) {
|
||||
eprintln!(
|
||||
"skipping postgres backfill test because required binaries are unavailable: initdb={}, postgres={}",
|
||||
initdb_bin, postgres_bin
|
||||
);
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
match Self::start(initdb_bin, postgres_bin).await {
|
||||
Ok(server) => Ok(Some(server)),
|
||||
Err(err) if postgres_local_startup_unavailable(err.to_string().as_str()) => {
|
||||
eprintln!(
|
||||
"skipping postgres backfill test because local postgres could not start in this environment: {err}"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
Err(err) => Err(err),
|
||||
}
|
||||
}
|
||||
|
||||
async fn start(
|
||||
initdb_bin: String,
|
||||
postgres_bin: String,
|
||||
) -> Result<Self, Box<dyn std::error::Error>> {
|
||||
let port = reserve_local_port()?;
|
||||
let workdir = std::env::temp_dir().join(format!(
|
||||
"aether-backfill-tests-{}-{}",
|
||||
std::process::id(),
|
||||
port
|
||||
));
|
||||
let data_dir = workdir.join("data");
|
||||
std::fs::create_dir_all(&workdir)?;
|
||||
|
||||
let init_output = Command::new(&initdb_bin)
|
||||
.arg("-D")
|
||||
.arg(&data_dir)
|
||||
.arg("-U")
|
||||
.arg("aether")
|
||||
.arg("--auth=trust")
|
||||
.arg("--encoding=UTF8")
|
||||
.arg("--no-instructions")
|
||||
.output()?;
|
||||
if !init_output.status.success() {
|
||||
return Err(std::io::Error::other(format!(
|
||||
"initdb failed: {}",
|
||||
String::from_utf8_lossy(&init_output.stderr)
|
||||
))
|
||||
.into());
|
||||
}
|
||||
|
||||
let database_url = format!("postgres://[email protected]:{port}/postgres");
|
||||
let log_path = workdir.join("postgres.log");
|
||||
let stdout = std::fs::File::create(&log_path)?;
|
||||
let stderr = stdout.try_clone()?;
|
||||
let mut child = Command::new(&postgres_bin)
|
||||
.arg("-D")
|
||||
.arg(&data_dir)
|
||||
.arg("-h")
|
||||
.arg("127.0.0.1")
|
||||
.arg("-p")
|
||||
.arg(port.to_string())
|
||||
.arg("-F")
|
||||
.arg("-c")
|
||||
.arg("fsync=off")
|
||||
.arg("-c")
|
||||
.arg("synchronous_commit=off")
|
||||
.arg("-c")
|
||||
.arg("full_page_writes=off")
|
||||
.arg("-c")
|
||||
.arg("shared_buffers=8MB")
|
||||
.arg("-c")
|
||||
.arg("max_connections=8")
|
||||
.arg("-c")
|
||||
.arg("dynamic_shared_memory_type=mmap")
|
||||
.arg("-c")
|
||||
.arg("autovacuum=off")
|
||||
.stdout(Stdio::from(stdout))
|
||||
.stderr(Stdio::from(stderr))
|
||||
.spawn()?;
|
||||
|
||||
if let Err(err) = wait_for_postgres(&database_url).await {
|
||||
let _ = child.kill();
|
||||
let _ = child.wait();
|
||||
return Err(err);
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
child: Some(child),
|
||||
workdir,
|
||||
database_url,
|
||||
})
|
||||
}
|
||||
|
||||
fn database_url(&self) -> &str {
|
||||
&self.database_url
|
||||
}
|
||||
|
||||
fn stop(&mut self) {
|
||||
if let Some(mut child) = self.child.take() {
|
||||
let _ = child.kill();
|
||||
let _ = child.wait();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for ManagedPostgresServer {
|
||||
fn drop(&mut self) {
|
||||
self.stop();
|
||||
let _ = std::fs::remove_dir_all(&self.workdir);
|
||||
}
|
||||
}
|
||||
|
||||
fn command_exists(bin: &str) -> bool {
|
||||
if bin.contains(std::path::MAIN_SEPARATOR) {
|
||||
return Path::new(bin).exists();
|
||||
}
|
||||
|
||||
let Some(paths) = std::env::var_os("PATH") else {
|
||||
return false;
|
||||
};
|
||||
|
||||
std::env::split_paths(&paths).any(|path| path.join(bin).exists())
|
||||
}
|
||||
|
||||
fn reserve_local_port() -> Result<u16, std::io::Error> {
|
||||
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
|
||||
let port = listener.local_addr()?.port();
|
||||
drop(listener);
|
||||
Ok(port)
|
||||
}
|
||||
|
||||
fn postgres_shared_memory_unavailable(message: &str) -> bool {
|
||||
let message = message.to_ascii_lowercase();
|
||||
message.contains("shared memory")
|
||||
&& (message.contains("could not create shared memory segment")
|
||||
|| message.contains("shmget")
|
||||
|| message.contains("no space left on device"))
|
||||
}
|
||||
|
||||
fn postgres_local_startup_unavailable(message: &str) -> bool {
|
||||
let message = message.to_ascii_lowercase();
|
||||
postgres_shared_memory_unavailable(&message)
|
||||
|| (message.contains("timed out waiting for local postgres")
|
||||
&& (message.contains("connection refused")
|
||||
|| message.contains("os error 61")
|
||||
|| message.contains("os error 111")))
|
||||
}
|
||||
|
||||
async fn wait_for_postgres(database_url: &str) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let deadline = Instant::now() + Duration::from_secs(10);
|
||||
loop {
|
||||
match PgConnection::connect(database_url).await {
|
||||
Ok(connection) => {
|
||||
connection.close().await?;
|
||||
return Ok(());
|
||||
}
|
||||
Err(_) if Instant::now() < deadline => {
|
||||
tokio::time::sleep(Duration::from_millis(50)).await
|
||||
}
|
||||
Err(err) => {
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::TimedOut,
|
||||
format!("timed out waiting for local postgres: {err}"),
|
||||
)
|
||||
.into())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn run_backfills_rebuilds_stats_and_records_execution() {
|
||||
let Some(server) = ManagedPostgresServer::try_start()
|
||||
|
||||
@@ -538,9 +538,15 @@ fn imported_payload_ids(
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
|
||||
pub struct DataImportOptions {
|
||||
pub preserve_credentials: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
|
||||
pub struct DataCopyOptions {
|
||||
pub omit_request_body_details: bool,
|
||||
pub preserve_credentials: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
@@ -572,6 +578,25 @@ fn set_supported_import_value(
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_import_credential_policy(
|
||||
table_name: &str,
|
||||
object: &mut serde_json::Map<String, Value>,
|
||||
target_has_column: impl Fn(&str) -> bool,
|
||||
options: DataImportOptions,
|
||||
) {
|
||||
let normalized_table = table_name
|
||||
.rsplit('.')
|
||||
.next()
|
||||
.unwrap_or(table_name)
|
||||
.trim_matches(|character| matches!(character, '"' | '`'));
|
||||
if options.preserve_credentials
|
||||
&& matches!(normalized_table, "users" | "api_keys" | "management_tokens")
|
||||
{
|
||||
return;
|
||||
}
|
||||
deactivate_imported_credentials(table_name, object, target_has_column);
|
||||
}
|
||||
|
||||
fn deactivate_imported_credentials(
|
||||
table_name: &str,
|
||||
object: &mut serde_json::Map<String, Value>,
|
||||
@@ -1085,6 +1110,14 @@ pub async fn export_database_jsonl(
|
||||
pub async fn import_database_jsonl(
|
||||
database: SqlDatabaseConfig,
|
||||
input: &str,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
import_database_jsonl_with_options(database, input, DataImportOptions::default()).await
|
||||
}
|
||||
|
||||
pub async fn import_database_jsonl_with_options(
|
||||
database: SqlDatabaseConfig,
|
||||
input: &str,
|
||||
options: DataImportOptions,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
match database.driver {
|
||||
#[cfg(feature = "postgres")]
|
||||
@@ -1092,7 +1125,7 @@ pub async fn import_database_jsonl(
|
||||
let pool =
|
||||
crate::driver::postgres::PostgresPoolFactory::new(database.to_postgres_config()?)?
|
||||
.connect_lazy()?;
|
||||
import_postgres_jsonl(&pool, input).await
|
||||
postgres::import_postgres_jsonl_with_options(&pool, input, options).await
|
||||
}
|
||||
#[cfg(not(feature = "postgres"))]
|
||||
DatabaseDriver::Postgres => Err(DataLayerError::InvalidInput(
|
||||
@@ -1113,7 +1146,14 @@ pub async fn copy_database_records(
|
||||
if options.omit_request_body_details {
|
||||
omit_request_body_details_from_records(&mut records);
|
||||
}
|
||||
import_database_jsonl(target, &encode_jsonl(&records)?).await
|
||||
import_database_jsonl_with_options(
|
||||
target,
|
||||
&encode_jsonl(&records)?,
|
||||
DataImportOptions {
|
||||
preserve_credentials: options.preserve_credentials,
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn omit_request_body_details_from_records(records: &mut Vec<DataExportRecord>) {
|
||||
|
||||
@@ -58,14 +58,30 @@ pub async fn export_postgres_jsonl(
|
||||
pub async fn import_postgres_jsonl(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
input: &str,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
import_postgres_jsonl_with_options(pool, input, DataImportOptions::default()).await
|
||||
}
|
||||
|
||||
pub(super) async fn import_postgres_jsonl_with_options(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
input: &str,
|
||||
options: DataImportOptions,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let plan = build_import_plan(input)?;
|
||||
import_postgres_plan(pool, &plan).await
|
||||
import_postgres_plan_with_options(pool, &plan, options).await
|
||||
}
|
||||
|
||||
pub async fn import_postgres_plan(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
plan: &DataImportPlan,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
import_postgres_plan_with_options(pool, plan, DataImportOptions::default()).await
|
||||
}
|
||||
|
||||
async fn import_postgres_plan_with_options(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
plan: &DataImportPlan,
|
||||
options: DataImportOptions,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let identity_scope = IdentityImportScope::from_plan(plan)?;
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
@@ -75,7 +91,7 @@ pub async fn import_postgres_plan(
|
||||
for domain in &plan.manifest.domains {
|
||||
if *domain == ExportDomain::Auxiliary {
|
||||
for row in plan.rows(*domain) {
|
||||
import_postgres_auxiliary_row(&mut tx, row, &mut column_cache).await?;
|
||||
import_postgres_auxiliary_row(&mut tx, row, &mut column_cache, options).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
continue;
|
||||
@@ -110,6 +126,7 @@ pub async fn import_postgres_plan(
|
||||
*domain,
|
||||
row,
|
||||
&target_columns,
|
||||
options,
|
||||
)
|
||||
.await?;
|
||||
imported = imported.saturating_add(1);
|
||||
@@ -587,11 +604,15 @@ async fn import_postgres_row(
|
||||
domain: ExportDomain,
|
||||
row: &ExportRow,
|
||||
target_columns: &PostgresImportColumns,
|
||||
options: DataImportOptions,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let mut object = normalize_postgres_import_payload(table_name, domain, row, target_columns)?;
|
||||
deactivate_imported_credentials(table_name, &mut object, |column_name| {
|
||||
target_columns.contains_key(column_name)
|
||||
});
|
||||
apply_import_credential_policy(
|
||||
table_name,
|
||||
&mut object,
|
||||
|column_name| target_columns.contains_key(column_name),
|
||||
options,
|
||||
);
|
||||
|
||||
let columns = object.keys().map(String::as_str).collect::<Vec<_>>();
|
||||
let column_sql = columns
|
||||
@@ -839,6 +860,7 @@ async fn import_postgres_billing_row(
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
DataImportOptions::default(),
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -847,6 +869,7 @@ async fn import_postgres_auxiliary_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, PostgresImportColumns>,
|
||||
options: DataImportOptions,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (table_name, payload) = domain_payload_table(row, "auxiliary", None)?;
|
||||
let table = auxiliary_table(&table_name)?;
|
||||
@@ -862,6 +885,7 @@ async fn import_postgres_auxiliary_row(
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
options,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -897,6 +921,7 @@ async fn import_postgres_wallet_row(
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
DataImportOptions::default(),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
@@ -3,11 +3,12 @@ use std::collections::{BTreeMap, BTreeSet};
|
||||
use serde_json::{json, Value};
|
||||
|
||||
use super::{
|
||||
build_import_plan, deactivate_imported_credentials, decode_jsonl, decode_jsonl_with_limits,
|
||||
encode_jsonl, export_postgres_core_jsonl, normalize_imported_binary,
|
||||
normalize_imported_integer_timestamp, normalize_postgres_import_payload,
|
||||
postgres_bytea_json_value, postgres_core_export_domains, DataExportManifest, DataExportRecord,
|
||||
ExportDomain, ExportRow, PostgresImportColumn,
|
||||
apply_import_credential_policy, build_import_plan, deactivate_imported_credentials,
|
||||
decode_jsonl, decode_jsonl_with_limits, encode_jsonl, export_postgres_core_jsonl,
|
||||
normalize_imported_binary, normalize_imported_integer_timestamp,
|
||||
normalize_postgres_import_payload, postgres_bytea_json_value, postgres_core_export_domains,
|
||||
DataExportManifest, DataExportRecord, DataImportOptions, ExportDomain, ExportRow,
|
||||
PostgresImportColumn,
|
||||
};
|
||||
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
use crate::lifecycle::migrate::run_migrations as run_postgres_migrations;
|
||||
@@ -416,6 +417,159 @@ fn imported_proxy_nodes_receive_a_new_offline_tunnel_generation() {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trusted_import_preserves_stable_credentials_only_when_explicitly_requested() {
|
||||
assert!(!DataImportOptions::default().preserve_credentials);
|
||||
for (table, payload) in [
|
||||
("users", json!({"password_hash": "$2b$12$trusted-hash"})),
|
||||
(
|
||||
"public.\"api_keys\"",
|
||||
json!({
|
||||
"key_hash": "trusted-key-hash", "key_encrypted": "trusted-ciphertext",
|
||||
"status": "active", "is_active": true, "is_locked": false,
|
||||
}),
|
||||
),
|
||||
(
|
||||
"management_tokens",
|
||||
json!({"token_hash": "trusted-token-hash", "is_active": true}),
|
||||
),
|
||||
] {
|
||||
for preserve_credentials in [false, true] {
|
||||
let mut object = payload.as_object().unwrap().clone();
|
||||
apply_import_credential_policy(
|
||||
table,
|
||||
&mut object,
|
||||
|_| true,
|
||||
DataImportOptions {
|
||||
preserve_credentials,
|
||||
},
|
||||
);
|
||||
if preserve_credentials {
|
||||
assert_eq!(&object, payload.as_object().unwrap());
|
||||
} else {
|
||||
assert_ne!(&object, payload.as_object().unwrap());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn trusted_import_still_revokes_imported_sessions_and_live_tunnels() {
|
||||
let options = DataImportOptions {
|
||||
preserve_credentials: true,
|
||||
};
|
||||
let mut session = json!({
|
||||
"refresh_token_hash": "old-session", "prev_refresh_token_hash": "older-session",
|
||||
"revoked_at": null, "revoke_reason": null,
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
.clone();
|
||||
apply_import_credential_policy("public.user_sessions", &mut session, |_| true, options);
|
||||
assert_ne!(session["refresh_token_hash"], json!("old-session"));
|
||||
assert_eq!(session["prev_refresh_token_hash"], Value::Null);
|
||||
assert_eq!(
|
||||
session["revoke_reason"],
|
||||
json!("imported_credentials_revoked")
|
||||
);
|
||||
|
||||
let mut node = json!({
|
||||
"tunnel_generation": "old-generation", "tunnel_connected": true,
|
||||
"status": "online", "active_connections": 10,
|
||||
})
|
||||
.as_object()
|
||||
.unwrap()
|
||||
.clone();
|
||||
apply_import_credential_policy("proxy_nodes", &mut node, |_| true, options);
|
||||
assert_ne!(node["tunnel_generation"], json!("old-generation"));
|
||||
assert_eq!(node["tunnel_connected"], json!(false));
|
||||
assert_eq!(node["status"], json!("offline"));
|
||||
assert_eq!(node["active_connections"], json!(0));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
#[ignore = "requires AETHER_TEST_POSTGRES_URL and PostgreSQL migrations"]
|
||||
async fn live_import_credential_policy_round_trips_through_postgres() {
|
||||
let pool = PostgresPoolFactory::new(PostgresPoolConfig {
|
||||
database_url: std::env::var("AETHER_TEST_POSTGRES_URL").unwrap(),
|
||||
..Default::default()
|
||||
})
|
||||
.unwrap()
|
||||
.connect_lazy()
|
||||
.unwrap();
|
||||
run_postgres_migrations(&pool).await.unwrap();
|
||||
for preserve_credentials in [false, true] {
|
||||
let user_id = uuid::Uuid::new_v4().to_string();
|
||||
let key_id = uuid::Uuid::new_v4().to_string();
|
||||
let password_hash = "$2b$12$trusted-import-hash";
|
||||
let key_hash = format!("trusted-{key_id}");
|
||||
sqlx::query("INSERT INTO users (id, username, password_hash, auth_source, email_verified) VALUES ($1, $2, $3, 'local', FALSE)")
|
||||
.bind(&user_id).bind(format!("import-{}", &user_id[..8])).bind(password_hash)
|
||||
.execute(&pool).await.unwrap();
|
||||
sqlx::query("INSERT INTO api_keys (id, user_id, key_hash, key_encrypted, name) VALUES ($1, $2, $3, 'trusted-ciphertext', 'Import probe')")
|
||||
.bind(&key_id).bind(&user_id).bind(&key_hash).execute(&pool).await.unwrap();
|
||||
let user: Value = sqlx::query_scalar("SELECT to_jsonb(users) FROM users WHERE id = $1")
|
||||
.bind(&user_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
let key: Value =
|
||||
sqlx::query_scalar("SELECT to_jsonb(api_keys) FROM api_keys WHERE id = $1")
|
||||
.bind(&key_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
let input = encode_jsonl(&[
|
||||
DataExportRecord::manifest(DataExportManifest::new(
|
||||
1_788_739_200,
|
||||
Some(DatabaseDriver::Postgres),
|
||||
vec![ExportDomain::Users, ExportDomain::ApiKeys],
|
||||
)),
|
||||
DataExportRecord::row(ExportDomain::Users, &user_id, user),
|
||||
DataExportRecord::row(ExportDomain::ApiKeys, &key_id, key),
|
||||
])
|
||||
.unwrap();
|
||||
super::postgres::import_postgres_jsonl_with_options(
|
||||
&pool,
|
||||
&input,
|
||||
DataImportOptions {
|
||||
preserve_credentials,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let imported_password: String =
|
||||
sqlx::query_scalar("SELECT password_hash FROM users WHERE id = $1")
|
||||
.bind(&user_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
let imported_key: (String, Option<String>, bool) =
|
||||
sqlx::query_as("SELECT key_hash, key_encrypted, is_active FROM api_keys WHERE id = $1")
|
||||
.bind(&key_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(imported_password == password_hash, preserve_credentials);
|
||||
assert_eq!(imported_key.0 == key_hash, preserve_credentials);
|
||||
assert_eq!(
|
||||
imported_key.1.as_deref(),
|
||||
preserve_credentials.then_some("trusted-ciphertext")
|
||||
);
|
||||
assert_eq!(imported_key.2, preserve_credentials);
|
||||
sqlx::query("DELETE FROM api_keys WHERE id = $1")
|
||||
.bind(&key_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
sqlx::query("DELETE FROM users WHERE id = $1")
|
||||
.bind(&user_id)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
fn postgres_column(data_type: &str, udt_name: &str) -> PostgresImportColumn {
|
||||
PostgresImportColumn {
|
||||
data_type: data_type.to_ascii_lowercase(),
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
use std::borrow::Cow;
|
||||
use std::collections::BTreeSet;
|
||||
use std::fs;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::process::{Child, Command, Stdio};
|
||||
use std::time::{Duration, Instant};
|
||||
use std::path::PathBuf;
|
||||
|
||||
use sqlx::{
|
||||
migrate::{AppliedMigration, Migrate},
|
||||
@@ -18,6 +16,8 @@ use aether_data_contracts::repository::{
|
||||
},
|
||||
};
|
||||
|
||||
use crate::lifecycle::postgres_test_support::ManagedPostgresServer;
|
||||
|
||||
use super::{
|
||||
postgres::{all_up_migrations, pending_migrations_from_applied, POSTGRES_MIGRATOR},
|
||||
prepare_database_for_startup,
|
||||
@@ -27,146 +27,6 @@ use crate::lifecycle::bootstrap::postgres::{
|
||||
EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION, EMPTY_DATABASE_SNAPSHOT_SQL,
|
||||
};
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ManagedPostgresServer {
|
||||
child: Option<Child>,
|
||||
workdir: PathBuf,
|
||||
database_url: String,
|
||||
}
|
||||
|
||||
impl ManagedPostgresServer {
|
||||
async fn try_start() -> Result<Option<Self>, Box<dyn std::error::Error>> {
|
||||
let required = local_postgres_tests_required();
|
||||
let initdb_bin = std::env::var("AETHER_INITDB_BIN")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or_else(|| "initdb".to_string());
|
||||
let postgres_bin = std::env::var("AETHER_POSTGRES_BIN")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or_else(|| "postgres".to_string());
|
||||
|
||||
if !command_exists(&initdb_bin) || !command_exists(&postgres_bin) {
|
||||
let message = format!(
|
||||
"required postgres integration test binaries are unavailable: initdb={initdb_bin}, postgres={postgres_bin}"
|
||||
);
|
||||
if required {
|
||||
return Err(std::io::Error::new(std::io::ErrorKind::NotFound, message).into());
|
||||
}
|
||||
eprintln!("skipping postgres integration test because {message}");
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
match Self::start(initdb_bin, postgres_bin).await {
|
||||
Ok(server) => Ok(Some(server)),
|
||||
Err(err)
|
||||
if !required && postgres_local_startup_unavailable(err.to_string().as_str()) =>
|
||||
{
|
||||
eprintln!(
|
||||
"skipping postgres integration test because local postgres could not start in this environment: {err}"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
Err(err) => Err(err),
|
||||
}
|
||||
}
|
||||
|
||||
async fn start(
|
||||
initdb_bin: String,
|
||||
postgres_bin: String,
|
||||
) -> Result<Self, Box<dyn std::error::Error>> {
|
||||
let port = reserve_local_port()?;
|
||||
let workdir = std::env::temp_dir().join(format!(
|
||||
"aether-migrate-tests-{}-{}",
|
||||
std::process::id(),
|
||||
port
|
||||
));
|
||||
let data_dir = workdir.join("data");
|
||||
std::fs::create_dir_all(&workdir)?;
|
||||
|
||||
let init_output = Command::new(&initdb_bin)
|
||||
.arg("-D")
|
||||
.arg(&data_dir)
|
||||
.arg("-U")
|
||||
.arg("aether")
|
||||
.arg("--auth=trust")
|
||||
.arg("--encoding=UTF8")
|
||||
.arg("--no-instructions")
|
||||
.output()?;
|
||||
if !init_output.status.success() {
|
||||
return Err(std::io::Error::other(format!(
|
||||
"initdb failed: {}",
|
||||
String::from_utf8_lossy(&init_output.stderr)
|
||||
))
|
||||
.into());
|
||||
}
|
||||
|
||||
let database_url = format!("postgres://[email protected]:{port}/postgres");
|
||||
let log_path = workdir.join("postgres.log");
|
||||
let stdout = std::fs::File::create(&log_path)?;
|
||||
let stderr = stdout.try_clone()?;
|
||||
let mut child = Command::new(&postgres_bin)
|
||||
.arg("-D")
|
||||
.arg(&data_dir)
|
||||
.arg("-h")
|
||||
.arg("127.0.0.1")
|
||||
.arg("-p")
|
||||
.arg(port.to_string())
|
||||
.arg("-k")
|
||||
.arg(&workdir)
|
||||
.arg("-F")
|
||||
.arg("-c")
|
||||
.arg("fsync=off")
|
||||
.arg("-c")
|
||||
.arg("synchronous_commit=off")
|
||||
.arg("-c")
|
||||
.arg("full_page_writes=off")
|
||||
.arg("-c")
|
||||
.arg("shared_buffers=8MB")
|
||||
.arg("-c")
|
||||
.arg("max_connections=8")
|
||||
.arg("-c")
|
||||
.arg("dynamic_shared_memory_type=mmap")
|
||||
.arg("-c")
|
||||
.arg("autovacuum=off")
|
||||
.stdout(Stdio::from(stdout))
|
||||
.stderr(Stdio::from(stderr))
|
||||
.spawn()?;
|
||||
|
||||
if let Err(err) = wait_for_postgres(&database_url).await {
|
||||
let _ = child.kill();
|
||||
let exit_status = child
|
||||
.wait()
|
||||
.map(|status| status.to_string())
|
||||
.unwrap_or_else(|wait_err| format!("unavailable ({wait_err})"));
|
||||
let logs = fs::read_to_string(&log_path)
|
||||
.unwrap_or_else(|read_err| format!("<failed to read postgres log: {read_err}>"));
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::TimedOut,
|
||||
format!("{err}; postgres exit status: {exit_status}; logs:\n{logs}"),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
child: Some(child),
|
||||
workdir,
|
||||
database_url,
|
||||
})
|
||||
}
|
||||
|
||||
fn database_url(&self) -> &str {
|
||||
&self.database_url
|
||||
}
|
||||
|
||||
fn stop(&mut self) {
|
||||
if let Some(mut child) = self.child.take() {
|
||||
let _ = child.kill();
|
||||
let _ = child.wait();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// A clean PostgreSQL database is bootstrapped from the schema snapshot first;
|
||||
/// migrations after the privacy/security frontier are intentionally left
|
||||
/// pending so their data-preserving changes still execute. Exercise the same
|
||||
@@ -191,83 +51,6 @@ async fn prepare_and_apply_clean_postgres_database(pool: &PgPool) {
|
||||
);
|
||||
}
|
||||
|
||||
fn local_postgres_tests_required() -> bool {
|
||||
// CI can opt into failing when the isolated local PostgreSQL fixture is unavailable.
|
||||
std::env::var("AETHER_REQUIRE_LOCAL_POSTGRES_TESTS")
|
||||
.ok()
|
||||
.is_some_and(|value| {
|
||||
matches!(
|
||||
value.trim().to_ascii_lowercase().as_str(),
|
||||
"1" | "true" | "yes" | "on"
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
impl Drop for ManagedPostgresServer {
|
||||
fn drop(&mut self) {
|
||||
self.stop();
|
||||
let _ = std::fs::remove_dir_all(&self.workdir);
|
||||
}
|
||||
}
|
||||
|
||||
fn command_exists(bin: &str) -> bool {
|
||||
if bin.contains(std::path::MAIN_SEPARATOR) {
|
||||
return Path::new(bin).exists();
|
||||
}
|
||||
|
||||
let Some(paths) = std::env::var_os("PATH") else {
|
||||
return false;
|
||||
};
|
||||
|
||||
std::env::split_paths(&paths).any(|path| path.join(bin).exists())
|
||||
}
|
||||
|
||||
fn reserve_local_port() -> Result<u16, std::io::Error> {
|
||||
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
|
||||
let port = listener.local_addr()?.port();
|
||||
drop(listener);
|
||||
Ok(port)
|
||||
}
|
||||
|
||||
fn postgres_shared_memory_unavailable(message: &str) -> bool {
|
||||
let message = message.to_ascii_lowercase();
|
||||
message.contains("shared memory")
|
||||
&& (message.contains("could not create shared memory segment")
|
||||
|| message.contains("shmget")
|
||||
|| message.contains("no space left on device"))
|
||||
}
|
||||
|
||||
fn postgres_local_startup_unavailable(message: &str) -> bool {
|
||||
let message = message.to_ascii_lowercase();
|
||||
postgres_shared_memory_unavailable(&message)
|
||||
|| (message.contains("timed out waiting for local postgres")
|
||||
&& (message.contains("connection refused")
|
||||
|| message.contains("os error 61")
|
||||
|| message.contains("os error 111")))
|
||||
}
|
||||
|
||||
async fn wait_for_postgres(database_url: &str) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let deadline = Instant::now() + Duration::from_secs(10);
|
||||
loop {
|
||||
match PgConnection::connect(database_url).await {
|
||||
Ok(connection) => {
|
||||
connection.close().await?;
|
||||
return Ok(());
|
||||
}
|
||||
Err(_) if Instant::now() < deadline => {
|
||||
tokio::time::sleep(Duration::from_millis(50)).await
|
||||
}
|
||||
Err(err) => {
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::TimedOut,
|
||||
format!("timed out waiting for local postgres: {err}"),
|
||||
)
|
||||
.into())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn table_exists(pool: &PgPool, table_name: &str) -> Result<bool, sqlx::Error> {
|
||||
query_scalar::<_, bool>("SELECT to_regclass($1) IS NOT NULL")
|
||||
.bind(format!("public.{table_name}"))
|
||||
|
||||
@@ -8,3 +8,5 @@ pub mod backfill;
|
||||
pub(crate) mod bootstrap;
|
||||
pub mod export;
|
||||
pub mod migrate;
|
||||
#[cfg(all(test, feature = "postgres"))]
|
||||
mod postgres_test_support;
|
||||
|
||||
@@ -0,0 +1,276 @@
|
||||
use std::{
|
||||
path::{Path, PathBuf},
|
||||
process::{Child, Command, Stdio},
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use sqlx::{Connection, PgConnection};
|
||||
|
||||
#[derive(Debug)]
|
||||
pub(super) struct ManagedPostgresServer {
|
||||
child: Option<Child>,
|
||||
pg_ctl_bin: PathBuf,
|
||||
workdir: PathBuf,
|
||||
data_dir: PathBuf,
|
||||
database_url: String,
|
||||
}
|
||||
|
||||
impl ManagedPostgresServer {
|
||||
pub(super) async fn try_start() -> Result<Option<Self>, Box<dyn std::error::Error>> {
|
||||
let required = local_postgres_tests_required();
|
||||
let initdb_bin = configured_binary("AETHER_INITDB_BIN", "initdb");
|
||||
let postgres_bin = configured_binary("AETHER_POSTGRES_BIN", "postgres");
|
||||
let pg_ctl_bin = std::env::var("AETHER_PG_CTL_BIN")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.map(PathBuf::from)
|
||||
.unwrap_or_else(|| {
|
||||
PathBuf::from(&postgres_bin).with_file_name(if cfg!(windows) {
|
||||
"pg_ctl.exe"
|
||||
} else {
|
||||
"pg_ctl"
|
||||
})
|
||||
});
|
||||
|
||||
if !command_exists(Path::new(&initdb_bin))
|
||||
|| !command_exists(Path::new(&postgres_bin))
|
||||
|| !command_exists(&pg_ctl_bin)
|
||||
{
|
||||
let message = format!(
|
||||
"required postgres integration test binaries are unavailable: initdb={initdb_bin}, postgres={postgres_bin}, pg_ctl={}",
|
||||
pg_ctl_bin.display()
|
||||
);
|
||||
if required {
|
||||
return Err(std::io::Error::new(std::io::ErrorKind::NotFound, message).into());
|
||||
}
|
||||
eprintln!("skipping postgres integration test because {message}");
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
match Self::start(initdb_bin, postgres_bin, pg_ctl_bin).await {
|
||||
Ok(server) => Ok(Some(server)),
|
||||
Err(error) if !required && postgres_local_startup_unavailable(&error.to_string()) => {
|
||||
eprintln!(
|
||||
"skipping postgres integration test because local postgres could not start in this environment: {error}"
|
||||
);
|
||||
Ok(None)
|
||||
}
|
||||
Err(error) => Err(error),
|
||||
}
|
||||
}
|
||||
|
||||
async fn start(
|
||||
initdb_bin: String,
|
||||
postgres_bin: String,
|
||||
pg_ctl_bin: PathBuf,
|
||||
) -> Result<Self, Box<dyn std::error::Error>> {
|
||||
let listener = std::net::TcpListener::bind("127.0.0.1:0")?;
|
||||
let port = listener.local_addr()?.port();
|
||||
drop(listener);
|
||||
let workdir = std::env::temp_dir().join(format!(
|
||||
"aether-lifecycle-tests-{}-{port}",
|
||||
std::process::id()
|
||||
));
|
||||
std::fs::create_dir(&workdir)?;
|
||||
let mut server = Self {
|
||||
child: None,
|
||||
pg_ctl_bin,
|
||||
data_dir: workdir.join("data"),
|
||||
workdir,
|
||||
database_url: format!("postgres://[email protected]:{port}/postgres"),
|
||||
};
|
||||
|
||||
let init_output = Command::new(&initdb_bin)
|
||||
.arg("-D")
|
||||
.arg(&server.data_dir)
|
||||
.args([
|
||||
"-U",
|
||||
"aether",
|
||||
"--auth=trust",
|
||||
"--encoding=UTF8",
|
||||
"--no-instructions",
|
||||
])
|
||||
.output()?;
|
||||
if !init_output.status.success() {
|
||||
return Err(std::io::Error::other(format!(
|
||||
"initdb failed: {}",
|
||||
String::from_utf8_lossy(&init_output.stderr)
|
||||
))
|
||||
.into());
|
||||
}
|
||||
|
||||
let log_path = server.workdir.join("postgres.log");
|
||||
let stdout = std::fs::File::create(&log_path)?;
|
||||
let stderr = stdout.try_clone()?;
|
||||
server.child = Some(
|
||||
Command::new(&postgres_bin)
|
||||
.arg("-D")
|
||||
.arg(&server.data_dir)
|
||||
.args(["-h", "127.0.0.1", "-p"])
|
||||
.arg(port.to_string())
|
||||
.arg("-F")
|
||||
.args(["-c", "unix_socket_directories="])
|
||||
.args(["-c", "fsync=off"])
|
||||
.args(["-c", "synchronous_commit=off"])
|
||||
.args(["-c", "full_page_writes=off"])
|
||||
.args(["-c", "shared_buffers=8MB"])
|
||||
.args(["-c", "max_connections=8"])
|
||||
.args(["-c", "dynamic_shared_memory_type=mmap"])
|
||||
.args(["-c", "autovacuum=off"])
|
||||
.stdout(Stdio::from(stdout))
|
||||
.stderr(Stdio::from(stderr))
|
||||
.spawn()?,
|
||||
);
|
||||
|
||||
if let Err(error) = wait_for_postgres(&server.database_url).await {
|
||||
let logs = std::fs::read_to_string(&log_path).unwrap_or_default();
|
||||
server.stop()?;
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::TimedOut,
|
||||
format!("{error}; postgres logs:\n{logs}"),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
Ok(server)
|
||||
}
|
||||
|
||||
pub(super) fn database_url(&self) -> &str {
|
||||
&self.database_url
|
||||
}
|
||||
|
||||
fn stop(&mut self) -> Result<(), std::io::Error> {
|
||||
let Some(child) = self.child.as_mut() else {
|
||||
return Ok(());
|
||||
};
|
||||
if child.try_wait()?.is_some() {
|
||||
self.child = None;
|
||||
return Ok(());
|
||||
}
|
||||
let output = Command::new(&self.pg_ctl_bin)
|
||||
.arg("-D")
|
||||
.arg(&self.data_dir)
|
||||
.args(["stop", "-m", "fast", "-w", "-t", "10"])
|
||||
.output()?;
|
||||
if !output.status.success() && child.try_wait()?.is_none() {
|
||||
return Err(std::io::Error::other(format!(
|
||||
"pg_ctl stop failed for {}: {}{}",
|
||||
self.data_dir.display(),
|
||||
String::from_utf8_lossy(&output.stdout),
|
||||
String::from_utf8_lossy(&output.stderr),
|
||||
)));
|
||||
}
|
||||
child.wait()?;
|
||||
self.child = None;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for ManagedPostgresServer {
|
||||
fn drop(&mut self) {
|
||||
match self.stop() {
|
||||
Ok(()) => {
|
||||
let _ = std::fs::remove_dir_all(&self.workdir);
|
||||
}
|
||||
Err(error) => {
|
||||
eprintln!(
|
||||
"failed to stop managed postgres; preserving {}: {error}",
|
||||
self.workdir.display(),
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn configured_binary(variable: &str, default: &str) -> String {
|
||||
std::env::var(variable)
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.unwrap_or_else(|| default.to_string())
|
||||
}
|
||||
|
||||
fn command_exists(binary: &Path) -> bool {
|
||||
if binary.is_absolute() || binary.components().count() > 1 {
|
||||
return binary.is_file();
|
||||
}
|
||||
std::env::var_os("PATH")
|
||||
.is_some_and(|paths| std::env::split_paths(&paths).any(|path| path.join(binary).is_file()))
|
||||
}
|
||||
|
||||
fn local_postgres_tests_required() -> bool {
|
||||
std::env::var("AETHER_REQUIRE_LOCAL_POSTGRES_TESTS")
|
||||
.ok()
|
||||
.is_some_and(|value| {
|
||||
matches!(
|
||||
value.trim().to_ascii_lowercase().as_str(),
|
||||
"1" | "true" | "yes" | "on"
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn postgres_local_startup_unavailable(message: &str) -> bool {
|
||||
let message = message.to_ascii_lowercase();
|
||||
(message.contains("shared memory")
|
||||
&& (message.contains("could not create shared memory segment")
|
||||
|| message.contains("shmget")
|
||||
|| message.contains("no space left on device")))
|
||||
|| (message.contains("timed out waiting for local postgres")
|
||||
&& (message.contains("connection refused")
|
||||
|| message.contains("os error 61")
|
||||
|| message.contains("os error 111")))
|
||||
}
|
||||
|
||||
async fn wait_for_postgres(database_url: &str) -> Result<(), Box<dyn std::error::Error>> {
|
||||
let deadline = Instant::now() + Duration::from_secs(10);
|
||||
loop {
|
||||
match PgConnection::connect(database_url).await {
|
||||
Ok(connection) => {
|
||||
connection.close().await?;
|
||||
return Ok(());
|
||||
}
|
||||
Err(_) if Instant::now() < deadline => {
|
||||
tokio::time::sleep(Duration::from_millis(50)).await;
|
||||
}
|
||||
Err(error) => {
|
||||
return Err(std::io::Error::new(
|
||||
std::io::ErrorKind::TimedOut,
|
||||
format!("timed out waiting for local postgres: {error}"),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn managed_postgres_stops_cleanly_with_open_connections() {
|
||||
let Some(mut server) = ManagedPostgresServer::try_start().await.unwrap() else {
|
||||
return;
|
||||
};
|
||||
let connection = PgConnection::connect(server.database_url()).await.unwrap();
|
||||
let workdir = server.workdir.clone();
|
||||
assert!(server.data_dir.join("postmaster.pid").exists());
|
||||
server.stop().unwrap();
|
||||
assert!(server.child.is_none());
|
||||
assert!(!server.data_dir.join("postmaster.pid").exists());
|
||||
server.stop().unwrap();
|
||||
drop(connection);
|
||||
drop(server);
|
||||
assert!(!workdir.exists());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn failed_postgres_stop_retains_ownership_for_retry() {
|
||||
let Some(mut server) = ManagedPostgresServer::try_start().await.unwrap() else {
|
||||
return;
|
||||
};
|
||||
let pg_ctl_bin = server.pg_ctl_bin.clone();
|
||||
let workdir = server.workdir.clone();
|
||||
server.pg_ctl_bin = workdir.join("missing-pg-ctl");
|
||||
assert!(server.stop().is_err());
|
||||
assert!(server.child.as_mut().unwrap().try_wait().unwrap().is_none());
|
||||
assert!(server.data_dir.exists());
|
||||
server.pg_ctl_bin = pg_ctl_bin;
|
||||
server.stop().unwrap();
|
||||
drop(server);
|
||||
assert!(!workdir.exists());
|
||||
}
|
||||
@@ -2752,6 +2752,40 @@ fn hydrate_client_family(item: &mut StoredRequestUsageAudit) {
|
||||
}
|
||||
}
|
||||
|
||||
fn merge_usage_body_capture(
|
||||
incoming_body: Option<Value>,
|
||||
incoming_ref: Option<String>,
|
||||
incoming_state: Option<UsageBodyCaptureState>,
|
||||
existing: Option<&StoredRequestUsageAudit>,
|
||||
field: UsageBodyField,
|
||||
) -> (Option<Value>, Option<String>, Option<UsageBodyCaptureState>) {
|
||||
if matches!(
|
||||
incoming_state,
|
||||
Some(
|
||||
UsageBodyCaptureState::None
|
||||
| UsageBodyCaptureState::Disabled
|
||||
| UsageBodyCaptureState::Unavailable
|
||||
)
|
||||
) {
|
||||
return (None, None, incoming_state);
|
||||
}
|
||||
if incoming_body.is_some() {
|
||||
return (
|
||||
incoming_body,
|
||||
None,
|
||||
incoming_state.or(Some(UsageBodyCaptureState::Inline)),
|
||||
);
|
||||
}
|
||||
if incoming_ref.is_some() {
|
||||
return (None, incoming_ref, Some(UsageBodyCaptureState::Reference));
|
||||
}
|
||||
(
|
||||
existing.and_then(|item| item.body_value(field).cloned()),
|
||||
existing.and_then(|item| item.body_ref(field).map(ToOwned::to_owned)),
|
||||
incoming_state.or_else(|| existing.and_then(|item| item.body_state(field))),
|
||||
)
|
||||
}
|
||||
|
||||
fn request_body_capture_replaces_derived_facts(
|
||||
request_body: Option<&Value>,
|
||||
request_body_state: Option<UsageBodyCaptureState>,
|
||||
@@ -2883,36 +2917,47 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
|
||||
return Ok(existing.clone());
|
||||
}
|
||||
}
|
||||
let capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage);
|
||||
let mut capture_usage = sanitize_usage_capture_controls_for_persistence(capture_usage);
|
||||
if let Some(existing) = by_request_id.get_mut(&usage.request_id) {
|
||||
existing.request_headers = None;
|
||||
existing.request_body = None;
|
||||
existing.request_body_ref = None;
|
||||
existing.request_body_state = None;
|
||||
existing.provider_request_headers = None;
|
||||
existing.provider_request_body = None;
|
||||
existing.provider_request_body_ref = None;
|
||||
existing.provider_request_body_state = None;
|
||||
existing.response_headers = None;
|
||||
existing.response_body = None;
|
||||
existing.response_body_ref = None;
|
||||
existing.response_body_state = None;
|
||||
existing.client_response_headers = None;
|
||||
existing.client_response_body = None;
|
||||
existing.client_response_body_ref = None;
|
||||
existing.client_response_body_state = None;
|
||||
existing.request_metadata =
|
||||
sanitize_usage_request_metadata(existing.request_metadata.take());
|
||||
}
|
||||
{
|
||||
let mut detached_bodies = self.detached_bodies.write().expect("usage repository lock");
|
||||
for field in [
|
||||
UsageBodyField::RequestBody,
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
UsageBodyField::ResponseBody,
|
||||
UsageBodyField::ClientResponseBody,
|
||||
for (field, state, body) in [
|
||||
(
|
||||
UsageBodyField::RequestBody,
|
||||
capture_usage.request_body_state,
|
||||
&capture_usage.request_body,
|
||||
),
|
||||
(
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
capture_usage.provider_request_body_state,
|
||||
&capture_usage.provider_request_body,
|
||||
),
|
||||
(
|
||||
UsageBodyField::ResponseBody,
|
||||
capture_usage.response_body_state,
|
||||
&capture_usage.response_body,
|
||||
),
|
||||
(
|
||||
UsageBodyField::ClientResponseBody,
|
||||
capture_usage.client_response_body_state,
|
||||
&capture_usage.client_response_body,
|
||||
),
|
||||
] {
|
||||
detached_bodies.remove(&usage_body_ref(&usage.request_id, field));
|
||||
if body.is_some()
|
||||
|| matches!(
|
||||
state,
|
||||
Some(
|
||||
UsageBodyCaptureState::None
|
||||
| UsageBodyCaptureState::Disabled
|
||||
| UsageBodyCaptureState::Unavailable
|
||||
)
|
||||
)
|
||||
{
|
||||
detached_bodies.remove(&usage_body_ref(&usage.request_id, field));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2977,10 +3022,36 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
|
||||
}
|
||||
});
|
||||
let request_metadata = sanitize_memory_request_metadata(request_metadata);
|
||||
let request_body_ref = None;
|
||||
let provider_request_body_ref = None;
|
||||
let response_body_ref = None;
|
||||
let client_response_body_ref = None;
|
||||
let (request_body, request_body_ref, request_body_state) = merge_usage_body_capture(
|
||||
capture_usage.request_body.take(),
|
||||
capture_usage.request_body_ref.take(),
|
||||
capture_usage.request_body_state,
|
||||
existing.as_ref(),
|
||||
UsageBodyField::RequestBody,
|
||||
);
|
||||
let (provider_request_body, provider_request_body_ref, provider_request_body_state) =
|
||||
merge_usage_body_capture(
|
||||
capture_usage.provider_request_body.take(),
|
||||
capture_usage.provider_request_body_ref.take(),
|
||||
capture_usage.provider_request_body_state,
|
||||
existing.as_ref(),
|
||||
UsageBodyField::ProviderRequestBody,
|
||||
);
|
||||
let (response_body, response_body_ref, response_body_state) = merge_usage_body_capture(
|
||||
capture_usage.response_body.take(),
|
||||
capture_usage.response_body_ref.take(),
|
||||
capture_usage.response_body_state,
|
||||
existing.as_ref(),
|
||||
UsageBodyField::ResponseBody,
|
||||
);
|
||||
let (client_response_body, client_response_body_ref, client_response_body_state) =
|
||||
merge_usage_body_capture(
|
||||
capture_usage.client_response_body.take(),
|
||||
capture_usage.client_response_body_ref.take(),
|
||||
capture_usage.client_response_body_state,
|
||||
existing.as_ref(),
|
||||
UsageBodyField::ClientResponseBody,
|
||||
);
|
||||
let stored = StoredRequestUsageAudit {
|
||||
id: existing
|
||||
.as_ref()
|
||||
@@ -3089,22 +3160,38 @@ impl UsageWriteRepository for InMemoryUsageReadRepository {
|
||||
),
|
||||
status: usage.status,
|
||||
billing_status: usage.billing_status,
|
||||
request_headers: None,
|
||||
request_body: None,
|
||||
request_headers: capture_usage.request_headers.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|item| item.request_headers.clone())
|
||||
}),
|
||||
request_body,
|
||||
request_body_ref,
|
||||
request_body_state: capture_usage.request_body_state,
|
||||
provider_request_headers: None,
|
||||
provider_request_body: None,
|
||||
request_body_state,
|
||||
provider_request_headers: capture_usage.provider_request_headers.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|item| item.provider_request_headers.clone())
|
||||
}),
|
||||
provider_request_body,
|
||||
provider_request_body_ref,
|
||||
provider_request_body_state: capture_usage.provider_request_body_state,
|
||||
response_headers: None,
|
||||
response_body: None,
|
||||
provider_request_body_state,
|
||||
response_headers: capture_usage.response_headers.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|item| item.response_headers.clone())
|
||||
}),
|
||||
response_body,
|
||||
response_body_ref,
|
||||
response_body_state: capture_usage.response_body_state,
|
||||
client_response_headers: None,
|
||||
client_response_body: None,
|
||||
response_body_state,
|
||||
client_response_headers: capture_usage.client_response_headers.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|item| item.client_response_headers.clone())
|
||||
}),
|
||||
client_response_body,
|
||||
client_response_body_ref,
|
||||
client_response_body_state: capture_usage.client_response_body_state,
|
||||
client_response_body_state,
|
||||
candidate_id: if replace_routing_snapshot {
|
||||
capture_usage.candidate_id
|
||||
} else {
|
||||
|
||||
@@ -136,6 +136,76 @@ fn sample_upsert_usage_record(request_id: &str) -> UpsertUsageRecord {
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upsert_preserves_full_http_captures_across_lifecycle_updates() {
|
||||
let repository = InMemoryUsageReadRepository::default();
|
||||
let mut pending = sample_upsert_usage_record("req-full-capture");
|
||||
pending.request_headers =
|
||||
Some(json!({"content-type": "application/json", "authorization": "Bearer private"}));
|
||||
pending.request_body =
|
||||
Some(json!({"messages": [{"role": "user", "content": "original request"}]}));
|
||||
pending.provider_request_body = Some(json!({"input": "provider request"}));
|
||||
pending.request_body_state = Some(UsageBodyCaptureState::Inline);
|
||||
pending.provider_request_body_state = Some(UsageBodyCaptureState::Inline);
|
||||
let stored_pending = repository.upsert(pending.clone()).await.unwrap();
|
||||
assert_eq!(stored_pending.request_body, pending.request_body);
|
||||
assert_eq!(
|
||||
stored_pending.provider_request_body,
|
||||
pending.provider_request_body
|
||||
);
|
||||
assert_eq!(
|
||||
stored_pending.request_headers,
|
||||
Some(json!({"content-type": "application/json", "authorization": "[redacted]"}))
|
||||
);
|
||||
|
||||
let mut streaming = sample_upsert_usage_record(&pending.request_id);
|
||||
streaming.status = "streaming".to_string();
|
||||
streaming.updated_at_unix_secs += 1;
|
||||
let stored_streaming = repository.upsert(streaming).await.unwrap();
|
||||
assert_eq!(stored_streaming.request_body, pending.request_body);
|
||||
assert_eq!(
|
||||
stored_streaming.provider_request_body,
|
||||
pending.provider_request_body
|
||||
);
|
||||
|
||||
let mut terminal = sample_upsert_usage_record(&pending.request_id);
|
||||
terminal.status = "completed".to_string();
|
||||
terminal.updated_at_unix_secs += 2;
|
||||
terminal.finalized_at_unix_secs = Some(terminal.updated_at_unix_secs);
|
||||
terminal.response_headers =
|
||||
Some(json!({"content-type": "text/event-stream", "set-cookie": "private"}));
|
||||
terminal.response_body = Some(json!("data: upstream response\n\ndata: [DONE]\n\n"));
|
||||
terminal.client_response_body =
|
||||
Some(json!({"choices": [{"message": {"content": "client response"}}]}));
|
||||
let stored_terminal = repository.upsert(terminal.clone()).await.unwrap();
|
||||
assert_eq!(stored_terminal.request_body, pending.request_body);
|
||||
assert_eq!(
|
||||
stored_terminal.provider_request_body,
|
||||
pending.provider_request_body
|
||||
);
|
||||
assert_eq!(stored_terminal.response_body, terminal.response_body);
|
||||
assert_eq!(
|
||||
stored_terminal.client_response_body,
|
||||
terminal.client_response_body
|
||||
);
|
||||
assert_eq!(
|
||||
stored_terminal.response_headers,
|
||||
Some(json!({"content-type": "text/event-stream", "set-cookie": "[redacted]"}))
|
||||
);
|
||||
|
||||
let found = repository
|
||||
.find_by_request_id(&pending.request_id)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(found.request_body, pending.request_body);
|
||||
assert_eq!(found.response_body, terminal.response_body);
|
||||
assert_eq!(
|
||||
repository.upsert(pending).await.unwrap().response_body,
|
||||
terminal.response_body
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn upsert_uses_typed_provider_capture_as_the_fast_fact_snapshot() {
|
||||
for (name, state, incoming_tier, expected_tier) in [
|
||||
|
||||
Reference in New Issue
Block a user