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:
elky
2026-09-07 21:14:27 +08:00
parent a5c3699ae9
commit a90d564931
191 changed files with 6785 additions and 1643 deletions
@@ -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 [