fix(runtime): avoid JSON precision stream regression

This commit is contained in:
ZheFox
2026-09-03 16:10:30 +08:00
parent 4c6bafe255
commit 89fe9e9f0a
3 changed files with 259 additions and 46 deletions
+1 -1
View File
@@ -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"
@@ -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<bool, DataLayerError> {
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::<bool, _>("metadata_is_object")
.map_postgres_err()?;
let namespace_exists = row
.try_get::<bool, _>("namespace_exists")
.map_postgres_err()?;
let current = row
.try_get::<Option<serde_json::Value>, _>("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<bool, DataLayerError> {
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<StoredProviderCatalogKey, DataLayerError>
#[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]
@@ -17,7 +17,9 @@ pub fn decode_stream_frame_ndjson(line: &[u8]) -> Result<StreamFrame, IoError> {
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);
}
}