From 89fe9e9f0a75b6a868ec56e4a03f642e17d59bac Mon Sep 17 00:00:00 2001 From: ZheFox <77232781+zhefox@users.noreply.github.com> Date: Thu, 3 Sep 2026 16:10:30 +0800 Subject: [PATCH] fix(runtime): avoid JSON precision stream regression --- Cargo.toml | 2 +- .../adapters/postgres/src/provider_catalog.rs | 281 +++++++++++++++--- .../execution/src/stream/ndjson.rs | 22 +- 3 files changed, 259 insertions(+), 46 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 34bd0926f..ae184e1cb 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -125,7 +125,7 @@ regex = "1" rustls = { version = "0.23", features = ["ring"] } semver = "1" serde = { version = "1", features = ["derive"] } -serde_json = { version = "1", features = ["arbitrary_precision", "preserve_order"] } +serde_json = { version = "1", features = ["preserve_order"] } serde_path_to_error = "0.1" sha2 = "0.10" socket2 = "0.6" diff --git a/crates/aether-data/adapters/postgres/src/provider_catalog.rs b/crates/aether-data/adapters/postgres/src/provider_catalog.rs index 3c68881b5..a8e79e44b 100644 --- a/crates/aether-data/adapters/postgres/src/provider_catalog.rs +++ b/crates/aether-data/adapters/postgres/src/provider_catalog.rs @@ -398,7 +398,18 @@ WHERE id = $1 AND ($6::text IS NULL OR auth_config IS NOT DISTINCT FROM $6) "#; -const KEY_RUNTIME_METADATA_CAS_SQL: &str = r#" +const KEY_RUNTIME_METADATA_NAMESPACE_LOCK_SQL: &str = r#" +SELECT + jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object' + AS metadata_is_object, + COALESCE(upstream_metadata, '{}'::jsonb) ? $2 AS namespace_exists, + COALESCE(upstream_metadata, '{}'::jsonb) -> $2 AS namespace_value +FROM provider_api_keys +WHERE id = $1 +FOR UPDATE +"#; + +const KEY_RUNTIME_METADATA_UPDATE_SQL: &str = r#" UPDATE provider_api_keys SET upstream_metadata = COALESCE(upstream_metadata, '{}'::jsonb) @@ -410,10 +421,58 @@ SET END WHERE id = $1 AND jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object' - AND (COALESCE(upstream_metadata, '{}'::jsonb) -> $2) - IS NOT DISTINCT FROM $6::jsonb "#; +fn runtime_metadata_namespace_matches( + metadata_is_object: bool, + namespace_exists: bool, + current: Option<&serde_json::Value>, + expected: Option<&serde_json::Value>, +) -> bool { + metadata_is_object + && match expected { + Some(expected) => namespace_exists && current == Some(expected), + None => !namespace_exists, + } +} + +async fn lock_runtime_metadata_namespace_matches( + tx: &mut sqlx::Transaction<'_, Postgres>, + key_id: &str, + namespace: &str, + expected: Option<&serde_json::Value>, +) -> Result { + let Some(row) = sqlx::query(KEY_RUNTIME_METADATA_NAMESPACE_LOCK_SQL) + .bind(key_id) + .bind(namespace) + .fetch_optional(&mut **tx) + .await + .map_postgres_err()? + else { + return Ok(false); + }; + let metadata_is_object = row + .try_get::("metadata_is_object") + .map_postgres_err()?; + let namespace_exists = row + .try_get::("namespace_exists") + .map_postgres_err()?; + let current = row + .try_get::, _>("namespace_value") + .map_postgres_err()?; + + // PostgreSQL jsonb retains decimal lexemes that serde_json's default + // Number representation rounds to f64. Re-read and compare while holding + // the row lock instead of binding that rounded value back into a jsonb + // equality predicate, which would report a false CAS conflict. + Ok(runtime_metadata_namespace_matches( + metadata_is_object, + namespace_exists, + current.as_ref(), + expected, + )) +} + fn validate_key_for_update(key: &StoredProviderCatalogKey) -> Result<(), DataLayerError> { if key.id.trim().is_empty() { return Err(DataLayerError::InvalidInput( @@ -1041,6 +1100,20 @@ WHERE id = $1 .to_string(), )); } + let mut tx = self.pool.begin().await.map_postgres_err()?; + if let Some(expected) = update.expected_upstream_metadata_namespace.as_ref() { + let matches = lock_runtime_metadata_namespace_matches( + &mut tx, + &update.key_id, + &expected.namespace, + expected.expected_value.as_ref(), + ) + .await?; + if !matches { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + } let rows_affected = sqlx::query( r#" UPDATE provider_api_keys @@ -1089,14 +1162,6 @@ WHERE id = $1 AND providers.provider_type = $18 ) ) - AND ( - $19::boolean IS FALSE - OR ( - jsonb_typeof(COALESCE(upstream_metadata, '{}'::jsonb)) = 'object' - AND (COALESCE(upstream_metadata, '{}'::jsonb) -> $20) - IS NOT DISTINCT FROM $21::jsonb - ) - ) "#, ) .bind(&update.key_id) @@ -1142,24 +1207,16 @@ WHERE id = $1 .as_ref() .map(|expected| expected.provider_type.as_str()), ) - .bind(update.expected_upstream_metadata_namespace.is_some()) - .bind( - update - .expected_upstream_metadata_namespace - .as_ref() - .map(|expected| expected.namespace.as_str()), - ) - .bind( - update - .expected_upstream_metadata_namespace - .as_ref() - .and_then(|expected| expected.expected_value.as_ref()), - ) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()? .rows_affected(); - Ok(rows_affected > 0) + if rows_affected == 0 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_postgres_err()?; + Ok(true) } pub async fn create_provider( @@ -2702,18 +2759,34 @@ WHERE id = $1 update: &ProviderCatalogKeyRuntimeMetadataUpdate, ) -> Result { validate_runtime_metadata_update(update)?; - let rows_affected = sqlx::query(KEY_RUNTIME_METADATA_CAS_SQL) + let mut tx = self.pool.begin().await.map_postgres_err()?; + if !lock_runtime_metadata_namespace_matches( + &mut tx, + &update.key_id, + &update.namespace, + update.expected_upstream_metadata_value.as_ref(), + ) + .await? + { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + let rows_affected = sqlx::query(KEY_RUNTIME_METADATA_UPDATE_SQL) .bind(&update.key_id) .bind(&update.namespace) .bind(&update.upstream_metadata_value) .bind(&update.status_snapshot_patch) .bind(update.updated_at_unix_secs.map(|value| value as f64)) - .bind(update.expected_upstream_metadata_value.as_ref()) - .execute(&self.pool) + .execute(&mut *tx) .await .map_postgres_err()? .rows_affected(); - Ok(rows_affected > 0) + if rows_affected == 0 { + tx.rollback().await.map_postgres_err()?; + return Ok(false); + } + tx.commit().await.map_postgres_err()?; + Ok(true) } pub async fn update_key_status_snapshot( @@ -3544,6 +3617,9 @@ fn map_key_row(row: &PgRow) -> Result #[cfg(test)] mod tests { + use aether_data_contracts::repository::provider_catalog::ProviderCatalogKeyRuntimeMetadataUpdate; + use serde_json::json; + use super::SqlxProviderCatalogReadRepository; use crate::{PostgresPoolConfig, PostgresPoolFactory}; @@ -3662,24 +3738,141 @@ mod tests { } #[test] - fn runtime_metadata_cas_compares_only_the_requested_namespace() { - let sql = super::KEY_RUNTIME_METADATA_CAS_SQL.to_ascii_lowercase(); - assert!(sql.contains("upstream_metadata, '{}'::jsonb) -> $2")); - assert!(sql.contains("jsonb_typeof(coalesce(upstream_metadata, '{}'::jsonb)) = 'object'")); - assert!(sql.contains("is not distinct from $6::jsonb")); - assert!(sql.contains("status_snapshot::jsonb")); - assert!(!sql.contains("is_active")); + fn runtime_metadata_cas_locks_only_the_requested_namespace() { + let lock_sql = super::KEY_RUNTIME_METADATA_NAMESPACE_LOCK_SQL.to_ascii_lowercase(); + let update_sql = super::KEY_RUNTIME_METADATA_UPDATE_SQL.to_ascii_lowercase(); + + assert!(lock_sql.contains("upstream_metadata, '{}'::jsonb) -> $2")); + assert!(lock_sql.contains("upstream_metadata, '{}'::jsonb) ? $2")); + assert!(lock_sql.contains("for update")); + assert!(update_sql + .contains("jsonb_typeof(coalesce(upstream_metadata, '{}'::jsonb)) = 'object'")); + assert!(update_sql.contains("status_snapshot::jsonb")); + assert!(!update_sql.contains("is_active")); } #[test] - fn runtime_metadata_cas_preserves_postgres_jsonb_numeric_lexemes() { - // PostgreSQL jsonb keeps the exact decimal value. Losing that lexeme while - // decoding into serde_json::Value makes a subsequent namespace CAS compare - // a nearby-but-different number and report a false conflict. - let postgres_json = r#"{"used_percent":9.373679999999995}"#; - let decoded: serde_json::Value = serde_json::from_str(postgres_json).unwrap(); + fn runtime_metadata_namespace_cas_distinguishes_missing_from_json_null() { + assert!(super::runtime_metadata_namespace_matches( + true, false, None, None, + )); + assert!(!super::runtime_metadata_namespace_matches( + true, + true, + Some(&serde_json::Value::Null), + None, + )); + assert!(super::runtime_metadata_namespace_matches( + true, + true, + Some(&serde_json::Value::Null), + Some(&serde_json::Value::Null), + )); + assert!(!super::runtime_metadata_namespace_matches( + false, false, None, None, + )); + } - assert_eq!(serde_json::to_string(&decoded).unwrap(), postgres_json); + #[tokio::test] + #[ignore = "requires AETHER_TEST_DATABASE_URL and PostgreSQL migrations"] + async fn live_runtime_metadata_cas_handles_high_precision_jsonb_numbers() { + let database_url = std::env::var("AETHER_TEST_DATABASE_URL") + .expect("AETHER_TEST_DATABASE_URL must point at the test database"); + let factory = PostgresPoolFactory::new(PostgresPoolConfig { + database_url, + min_connections: 1, + max_connections: 2, + acquire_timeout_ms: 10_000, + idle_timeout_ms: 30_000, + max_lifetime_ms: 60_000, + statement_cache_capacity: 64, + require_ssl: false, + }) + .expect("factory should build"); + let repository = SqlxProviderCatalogReadRepository::new( + factory.connect_lazy().expect("lazy pool should build"), + ); + crate::run_migrations(repository.pool()) + .await + .expect("test database migrations should succeed"); + + let suffix = uuid::Uuid::new_v4().simple().to_string(); + let provider_id = uuid::Uuid::new_v4().to_string(); + let key_id = uuid::Uuid::new_v4().to_string(); + let provider_name = format!("provider-metadata-cas-{suffix}"); + let key_name = format!("key-metadata-cas-{suffix}"); + sqlx::query( + "INSERT INTO providers (id, name, provider_type) VALUES ($1, $2, 'antigravity')", + ) + .bind(&provider_id) + .bind(&provider_name) + .execute(repository.pool()) + .await + .expect("provider fixture should insert"); + sqlx::query( + r#" +INSERT INTO provider_api_keys ( + id, name, provider_id, total_tokens, total_cost_usd, upstream_metadata +) +VALUES ($1, $2, $3, 0, 0, $4::jsonb) +"#, + ) + .bind(&key_id) + .bind(&key_name) + .bind(&provider_id) + .bind(r#"{"antigravity":{"used_percent":0.123456789012345678901234567890}}"#) + .execute(repository.pool()) + .await + .expect("provider key fixture should insert"); + + let observed = sqlx::query_scalar::<_, serde_json::Value>( + "SELECT upstream_metadata -> 'antigravity' FROM provider_api_keys WHERE id = $1", + ) + .bind(&key_id) + .fetch_one(repository.pool()) + .await + .expect("metadata namespace should load"); + assert_ne!( + serde_json::to_string(&observed).expect("metadata should serialize"), + r#"{"used_percent":0.123456789012345678901234567890}"#, + "the fixture must exercise precision loss in serde_json's default number representation", + ); + + let updated = repository + .update_key_runtime_metadata(&ProviderCatalogKeyRuntimeMetadataUpdate { + key_id: key_id.clone(), + namespace: "antigravity".to_string(), + expected_upstream_metadata_value: Some(observed), + upstream_metadata_value: json!({"used_percent": 12.5}), + status_snapshot_patch: json!({"quota": {"used_percent": 12.5}}), + updated_at_unix_secs: Some(1_700_000_000), + }) + .await + .expect("runtime metadata CAS should execute"); + assert!( + updated, + "matching metadata must not report a false CAS conflict" + ); + + let stored = sqlx::query_scalar::<_, serde_json::Value>( + "SELECT upstream_metadata -> 'antigravity' FROM provider_api_keys WHERE id = $1", + ) + .bind(&key_id) + .fetch_one(repository.pool()) + .await + .expect("updated metadata namespace should load"); + assert_eq!(stored, json!({"used_percent": 12.5})); + + sqlx::query("DELETE FROM provider_api_keys WHERE id = $1") + .bind(&key_id) + .execute(repository.pool()) + .await + .expect("provider key fixture should delete"); + sqlx::query("DELETE FROM providers WHERE id = $1") + .bind(&provider_id) + .execute(repository.pool()) + .await + .expect("provider fixture should delete"); } #[test] diff --git a/crates/aether-gateway/execution/src/stream/ndjson.rs b/crates/aether-gateway/execution/src/stream/ndjson.rs index 990eb2411..d7ba00e66 100644 --- a/crates/aether-gateway/execution/src/stream/ndjson.rs +++ b/crates/aether-gateway/execution/src/stream/ndjson.rs @@ -17,7 +17,9 @@ pub fn decode_stream_frame_ndjson(line: &[u8]) -> Result { mod tests { use std::collections::BTreeMap; - use aether_contracts::{StreamFramePayload, StreamFrameType}; + use aether_contracts::{ + ExecutionStreamTerminalSummary, StandardizedUsage, StreamFramePayload, StreamFrameType, + }; use super::{decode_stream_frame_ndjson, encode_stream_frame_ndjson}; @@ -41,4 +43,22 @@ mod tests { decode_stream_frame_ndjson(raw.trim_ascii_end()).expect("frame should decode"); assert_eq!(decoded, frame); } + + #[test] + fn ndjson_round_trip_preserves_terminal_usage_with_fractional_fields() { + let mut usage = StandardizedUsage::new(); + usage.cache_storage_token_hours = 0.125; + let frame = + aether_contracts::StreamFrame::eof_with_summary(Some(ExecutionStreamTerminalSummary { + standardized_usage: Some(usage), + observed_finish: true, + ..ExecutionStreamTerminalSummary::default() + })); + + let raw = encode_stream_frame_ndjson(&frame).expect("frame should encode"); + let decoded = decode_stream_frame_ndjson(raw.trim_ascii_end()) + .expect("terminal usage frame should decode"); + + assert_eq!(decoded, frame); + } }