mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-09 12:40:20 +08:00
fix(runtime): avoid JSON precision stream regression
This commit is contained in:
+1
-1
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user