fix(codex): fence concurrent quota updates

This commit is contained in:
elky
2026-08-14 09:28:07 +08:00
parent f3a12c1008
commit 5b0c763086
65 changed files with 13009 additions and 738 deletions
@@ -8,8 +8,8 @@ use sqlx::{
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyAdminCasUpdate,
ProviderCatalogKeyHealthStateUpdate, ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery,
ProviderCatalogKeyOAuthCredentialCasDelete, ProviderCatalogKeyOAuthRuntimeStateCasUpdate,
ProviderCatalogKeyRuntimeMetadataUpdate, ProviderCatalogKeyStatusSnapshotUpdate,
ProviderCatalogReadRepository, ProviderCatalogUpstreamMetadataNamespaceUpdate,
@@ -936,6 +936,84 @@ WHERE id = ?
self.reload_key(&key.id, "updated").await
}
pub async fn compare_and_update_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, DataLayerError> {
validate_admin_key_cas_update(update)?;
let key = &update.key;
let updated_at = key.updated_at_unix_secs.unwrap_or_else(current_unix_secs) as i64;
let rotation_json = update
.codex_rotation
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys Codex rotation is not serializable: {err}"
))
})?;
let mut builder = QueryBuilder::<MySql>::new("UPDATE provider_api_keys SET provider_id = ");
push_admin_key_assignments(&mut builder, key, updated_at)?;
if update.reset_oauth_runtime {
builder.push(", oauth_invalid_at = NULL, oauth_invalid_reason = NULL, error_count = 0");
}
if let Some(rotation_json) = rotation_json.as_deref() {
builder
.push(", upstream_metadata = JSON_SET(COALESCE(upstream_metadata, JSON_OBJECT()), '$.codex', CAST(")
.push_bind(rotation_json)
.push(" AS JSON))");
}
if update.codex_rotation.is_some() || update.reset_oauth_runtime {
builder.push(", status_snapshot = ");
match (update.codex_rotation.is_some(), update.reset_oauth_runtime) {
(true, true) => builder.push("JSON_SET(COALESCE(status_snapshot, JSON_OBJECT()), '$.quota', CAST('null' AS JSON), '$.oauth', CAST('null' AS JSON))"),
(true, false) => builder.push("JSON_SET(COALESCE(status_snapshot, JSON_OBJECT()), '$.quota', CAST('null' AS JSON))"),
(false, true) => builder.push("JSON_SET(COALESCE(status_snapshot, JSON_OBJECT()), '$.oauth', CAST('null' AS JSON))"),
(false, false) => unreachable!(),
};
}
builder
.push(" WHERE id = ")
.push_bind(&key.id)
.push(" AND BINARY api_key <=> BINARY ")
.push_bind(update.expected_credential.encrypted_api_key.as_deref())
.push(" AND BINARY auth_config <=> BINARY ")
.push_bind(update.expected_encrypted_auth_config.as_deref())
.push(" AND BINARY auth_type = BINARY ")
.push_bind(&update.expected_credential.auth_type)
.push(" AND BINARY provider_id = BINARY ")
.push_bind(&update.expected_credential.provider_id)
.push(
" AND EXISTS (SELECT 1 FROM providers WHERE BINARY providers.id = BINARY provider_api_keys.provider_id AND BINARY providers.provider_type = BINARY ",
)
.push_bind(&update.expected_credential.provider_type)
.push(")");
if update.codex_rotation.is_some() {
builder
.push(" AND JSON_TYPE(COALESCE(upstream_metadata, JSON_OBJECT())) = 'OBJECT'")
.push(" AND NOT (BINARY api_key <=> BINARY ")
.push_bind(key.encrypted_api_key.as_deref())
.push(" AND BINARY auth_config <=> BINARY ")
.push_bind(key.encrypted_auth_config.as_deref())
.push(" AND BINARY auth_type = BINARY ")
.push_bind(&key.auth_type)
.push(" AND BINARY provider_id = BINARY ")
.push_bind(&key.provider_id)
.push(")");
}
if update.codex_rotation.is_some() || update.reset_oauth_runtime {
builder.push(" AND JSON_TYPE(COALESCE(status_snapshot, JSON_OBJECT())) = 'OBJECT'");
}
let rows_affected = builder
.build()
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
pub async fn update_keys(
&self,
keys: &[StoredProviderCatalogKey],
@@ -1007,6 +1085,10 @@ WHERE id = ?
|| expected.auth_type.trim().is_empty()
|| expected.provider_id.trim().is_empty()
|| expected.provider_type.trim().is_empty()
|| delete
.expected_upstream_metadata_namespace
.as_ref()
.is_some_and(|expected| expected.namespace.trim().is_empty())
{
return Err(DataLayerError::InvalidInput(
"provider catalog OAuth credential CAS delete contains empty fields".to_string(),
@@ -1030,6 +1112,33 @@ WHERE id = ?
)
.push_bind(&expected.provider_type)
.push(")");
if let Some(expected) = delete.expected_upstream_metadata_namespace.as_ref() {
let namespace_path = format!(
"$.{}",
serde_json::to_string(&expected.namespace).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata namespace is not serializable: {err}"
))
})?
);
let expected_value = expected
.expected_value
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata expected value is not serializable: {err}"
))
})?;
builder
.push(" AND JSON_TYPE(COALESCE(NULLIF(upstream_metadata, ''), '{}')) = 'OBJECT'")
.push(" AND JSON_EXTRACT(COALESCE(NULLIF(upstream_metadata, ''), '{}'), ")
.push_bind(namespace_path)
.push(") <=> CAST(")
.push_bind(expected_value)
.push(" AS JSON)");
}
let rows_affected = builder
.build()
.execute(&self.pool)
@@ -1339,6 +1448,25 @@ WHERE id = ?
"provider catalog OAuth credential fence must not contain empty fields".to_string(),
));
}
if update
.expected_upstream_metadata_namespace
.as_ref()
.is_some_and(|expected| expected.namespace.trim().is_empty())
{
return Err(DataLayerError::InvalidInput(
"provider catalog OAuth runtime metadata namespace must not be empty".to_string(),
));
}
if update
.upstream_metadata_namespace_to_remove
.as_deref()
.is_some_and(|namespace| namespace.trim().is_empty())
{
return Err(DataLayerError::InvalidInput(
"provider catalog OAuth runtime metadata namespace to remove must not be empty"
.to_string(),
));
}
if !update.status_snapshot_patch.is_object() {
return Err(DataLayerError::InvalidInput(
"provider catalog status snapshot patch must be an object".to_string(),
@@ -1353,6 +1481,22 @@ WHERE id = ?
"provider catalog upstream metadata patch must be an object".to_string(),
));
}
if update
.upstream_metadata_namespace_to_remove
.as_ref()
.is_some_and(|namespace| {
update
.upstream_metadata_patch
.as_ref()
.and_then(serde_json::Value::as_object)
.is_some_and(|patch| patch.contains_key(namespace))
})
{
return Err(DataLayerError::InvalidInput(
"provider catalog OAuth runtime metadata namespace cannot be patched and removed in the same update"
.to_string(),
));
}
let mut builder =
QueryBuilder::<MySql>::new("UPDATE provider_api_keys SET oauth_invalid_at = ");
builder
@@ -1375,9 +1519,29 @@ WHERE id = ?
"provider_api_keys.expires_at",
)?);
}
if let Some(metadata_patch) = update.upstream_metadata_patch.as_ref() {
if update.upstream_metadata_patch.is_some()
|| update.upstream_metadata_namespace_to_remove.is_some()
{
builder.push(", upstream_metadata = ");
push_upstream_metadata_shallow_patch(&mut builder, metadata_patch)?;
if update.upstream_metadata_namespace_to_remove.is_some() {
builder.push("JSON_REMOVE(");
}
if let Some(metadata_patch) = update.upstream_metadata_patch.as_ref() {
push_upstream_metadata_shallow_patch(&mut builder, metadata_patch)?;
} else {
builder.push("COALESCE(NULLIF(upstream_metadata, ''), '{}')");
}
if let Some(namespace) = update.upstream_metadata_namespace_to_remove.as_ref() {
let namespace_path = format!(
"$.{}",
serde_json::to_string(namespace).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata namespace is not serializable: {err}"
))
})?
);
builder.push(", ").push_bind(namespace_path).push(")");
}
}
builder.push(", status_snapshot = ");
push_status_snapshot_shallow_patch(&mut builder, &update.status_snapshot_patch)?;
@@ -1411,6 +1575,39 @@ WHERE id = ?
.push_bind(&expected.provider_type)
.push(")");
}
if update.expected_upstream_metadata_namespace.is_some()
|| update.upstream_metadata_patch.is_some()
|| update.upstream_metadata_namespace_to_remove.is_some()
{
builder
.push(" AND JSON_TYPE(COALESCE(NULLIF(upstream_metadata, ''), '{}')) = 'OBJECT'");
}
if let Some(expected) = update.expected_upstream_metadata_namespace.as_ref() {
let namespace_path = format!(
"$.{}",
serde_json::to_string(&expected.namespace).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata namespace is not serializable: {err}"
))
})?
);
let expected_value = expected
.expected_value
.as_ref()
.map(serde_json::to_string)
.transpose()
.map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata expected value is not serializable: {err}"
))
})?;
builder
.push(" AND JSON_EXTRACT(COALESCE(NULLIF(upstream_metadata, ''), '{}'), ")
.push_bind(namespace_path)
.push(") <=> CAST(")
.push_bind(expected_value)
.push(" AS JSON)");
}
let rows_affected = builder
.build()
.execute(&self.pool)
@@ -1618,6 +1815,7 @@ WHERE id = ?
)
.push(" WHERE id = ")
.push_bind(&update.key_id)
.push(" AND JSON_TYPE(COALESCE(NULLIF(upstream_metadata, ''), '{}')) = 'OBJECT'")
.push(" AND JSON_EXTRACT(COALESCE(NULLIF(upstream_metadata, ''), '{}'), ")
.push_bind(namespace_path)
.push(") <=> CAST(")
@@ -1897,6 +2095,13 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
Self::update_key(self, key).await
}
async fn compare_and_update_key_admin_state(
&self,
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_update_key_admin_state(self, update).await
}
async fn update_keys(
&self,
keys: &[StoredProviderCatalogKey],
@@ -2426,6 +2631,10 @@ SET
model_exclude_patterns = ?,
updated_at = ?
WHERE id = ?
AND BINARY provider_id = BINARY ?
AND BINARY auth_type = BINARY ?
AND BINARY api_key <=> BINARY ?
AND BINARY auth_config <=> BINARY ?
"#
}
@@ -2500,7 +2709,152 @@ fn key_update_query(
"provider_api_keys.model_exclude_patterns",
)?)
.bind(updated_at)
.bind(&key.id))
.bind(&key.id)
.bind(&key.provider_id)
.bind(&key.auth_type)
.bind(&key.encrypted_api_key)
.bind(&key.encrypted_auth_config))
}
fn push_admin_key_assignments<'args>(
builder: &mut QueryBuilder<'args, MySql>,
key: &'args StoredProviderCatalogKey,
updated_at: i64,
) -> Result<(), DataLayerError> {
builder
.push_bind(&key.provider_id)
.push(", name = ")
.push_bind(&key.name)
.push(", api_key = ")
.push_bind(&key.encrypted_api_key)
.push(", auth_type = ")
.push_bind(&key.auth_type)
.push(", capabilities = ")
.push_bind(optional_json_to_string(
&key.capabilities,
"provider_api_keys.capabilities",
)?)
.push(", is_active = ")
.push_bind(key.is_active)
.push(", api_formats = ")
.push_bind(optional_json_to_string(
&key.api_formats,
"provider_api_keys.api_formats",
)?)
.push(", auth_type_by_format = ")
.push_bind(optional_json_to_string(
&key.auth_type_by_format,
"provider_api_keys.auth_type_by_format",
)?)
.push(", allow_auth_channel_mismatch_formats = ")
.push_bind(optional_json_to_string(
&key.allow_auth_channel_mismatch_formats,
"provider_api_keys.allow_auth_channel_mismatch_formats",
)?)
.push(", auth_config = ")
.push_bind(&key.encrypted_auth_config)
.push(", note = ")
.push_bind(&key.note)
.push(", internal_priority = ")
.push_bind(key.internal_priority)
.push(", rate_multipliers = ")
.push_bind(optional_json_to_string(
&key.rate_multipliers,
"provider_api_keys.rate_multipliers",
)?)
.push(", global_priority_by_format = ")
.push_bind(optional_json_to_string(
&key.global_priority_by_format,
"provider_api_keys.global_priority_by_format",
)?)
.push(", allowed_models = ")
.push_bind(optional_json_to_string(
&key.allowed_models,
"provider_api_keys.allowed_models",
)?)
.push(", expires_at = ")
.push_bind(optional_i64_from_u64(
key.expires_at_unix_secs,
"provider_api_keys.expires_at",
)?)
.push(", cache_ttl_minutes = ")
.push_bind(key.cache_ttl_minutes)
.push(", max_probe_interval_minutes = ")
.push_bind(key.max_probe_interval_minutes)
.push(", proxy = ")
.push_bind(optional_json_to_string(
&key.proxy,
"provider_api_keys.proxy",
)?)
.push(", fingerprint = ")
.push_bind(optional_json_to_string(
&key.fingerprint,
"provider_api_keys.fingerprint",
)?)
.push(", rpm_limit = ")
.push_bind(optional_i64_from_u32(key.rpm_limit))
.push(", concurrent_limit = ")
.push_bind(key.concurrent_limit)
.push(", auto_fetch_models = ")
.push_bind(key.auto_fetch_models)
.push(", locked_models = ")
.push_bind(optional_json_to_string(
&key.locked_models,
"provider_api_keys.locked_models",
)?)
.push(", model_include_patterns = ")
.push_bind(optional_json_to_string(
&key.model_include_patterns,
"provider_api_keys.model_include_patterns",
)?)
.push(", model_exclude_patterns = ")
.push_bind(optional_json_to_string(
&key.model_exclude_patterns,
"provider_api_keys.model_exclude_patterns",
)?)
.push(", updated_at = ")
.push_bind(updated_at);
Ok(())
}
fn validate_admin_key_cas_update(
update: &ProviderCatalogKeyAdminCasUpdate,
) -> Result<(), DataLayerError> {
validate_key(&update.key)?;
let expected = &update.expected_credential;
if expected.auth_type.trim().is_empty()
|| expected.provider_id.trim().is_empty()
|| expected.provider_type.trim().is_empty()
|| expected
.encrypted_api_key
.as_deref()
.is_some_and(|value| value.trim().is_empty())
|| update
.expected_encrypted_auth_config
.as_deref()
.is_some_and(|value| value.trim().is_empty())
{
return Err(DataLayerError::InvalidInput(
"provider catalog admin credential fence contains empty fields".to_string(),
));
}
let Some(rotation) = update.codex_rotation.as_ref() else {
return Ok(());
};
let valid_rotation = expected.provider_type.eq_ignore_ascii_case("codex")
&& rotation.as_object().is_some_and(|object| {
object.len() == 1
&& object
.get("credential_generation")
.and_then(serde_json::Value::as_str)
.is_some_and(|generation| !generation.trim().is_empty())
});
if !valid_rotation {
return Err(DataLayerError::InvalidInput(
"provider catalog Codex rotation must contain only credential_generation".to_string(),
));
}
Ok(())
}
fn optional_json_from_string(
@@ -2880,6 +3234,34 @@ mod tests {
] {
assert!(!sql.contains(runtime_assignment));
}
assert!(sql.contains("binary api_key <=> binary ?"));
assert!(sql.contains("binary auth_config <=> binary ?"));
}
#[test]
fn admin_credential_cas_has_atomic_rotation_guards() {
let source = include_str!("provider_catalog.rs");
for predicate in [
"BINARY api_key <=> BINARY ",
"BINARY auth_config <=> BINARY ",
"JSON_TYPE(COALESCE(upstream_metadata, JSON_OBJECT())) = 'OBJECT'",
"JSON_TYPE(COALESCE(status_snapshot, JSON_OBJECT())) = 'OBJECT'",
"JSON_SET(COALESCE(upstream_metadata, JSON_OBJECT()), '$.codex'",
"oauth_invalid_at = NULL, oauth_invalid_reason = NULL, error_count = 0",
] {
assert!(
source.contains(predicate),
"missing admin CAS guard: {predicate}"
);
}
}
#[test]
fn runtime_metadata_cas_requires_an_object_metadata_root() {
let source = include_str!("provider_catalog.rs");
assert!(
source.contains("JSON_TYPE(COALESCE(NULLIF(upstream_metadata, ''), '{}')) = 'OBJECT'")
);
}
#[tokio::test]
async fn empty_id_lists_do_not_connect_to_lazy_pool() {