perf(gateway): scale request hot paths for 20k streams

Shard and singleflight hot-path caches, batch and prioritize candidate and usage lifecycle persistence, and extend database and pressure-test instrumentation for 20k concurrent streams.
This commit is contained in:
elky
2026-07-22 02:11:08 +08:00
parent 7756c0913f
commit fc92c4f431
124 changed files with 36325 additions and 3217 deletions
@@ -303,19 +303,73 @@ ON DUPLICATE KEY UPDATE
provider_id = VALUES(provider_id),
endpoint_id = VALUES(endpoint_id),
key_id = VALUES(key_id),
status = VALUES(status),
status = CASE
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
THEN status
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
THEN status
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
THEN status
ELSE VALUES(status)
END,
skip_reason = VALUES(skip_reason),
is_cached = VALUES(is_cached),
status_code = VALUES(status_code),
error_type = VALUES(error_type),
error_message = VALUES(error_message),
latency_ms = VALUES(latency_ms),
status_code = CASE
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
THEN status_code
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
THEN status_code
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
THEN status_code
ELSE COALESCE(VALUES(status_code), status_code)
END,
error_type = CASE
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
THEN error_type
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
THEN error_type
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
THEN error_type
ELSE COALESCE(VALUES(error_type), error_type)
END,
error_message = CASE
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
THEN error_message
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
THEN error_message
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
THEN error_message
ELSE COALESCE(VALUES(error_message), error_message)
END,
latency_ms = CASE
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
THEN latency_ms
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
THEN latency_ms
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
THEN latency_ms
ELSE COALESCE(VALUES(latency_ms), latency_ms)
END,
concurrent_requests = VALUES(concurrent_requests),
extra_data = VALUES(extra_data),
required_capabilities = VALUES(required_capabilities),
created_at = VALUES(created_at),
started_at = VALUES(started_at),
finished_at = VALUES(finished_at)
finished_at = CASE
WHEN status IN ('success', 'failed', 'cancelled', 'skipped')
AND VALUES(status) IN ('available', 'unused', 'pending', 'streaming')
THEN finished_at
WHEN status = 'pending' AND VALUES(status) IN ('available', 'unused')
THEN finished_at
WHEN status = 'streaming' AND VALUES(status) IN ('available', 'unused', 'pending')
THEN finished_at
ELSE COALESCE(VALUES(finished_at), finished_at)
END
"#,
)
.bind(&candidate.id)
@@ -720,8 +774,10 @@ fn optional_u64_to_i64(value: Option<u64>, name: &str) -> Result<Option<i64>, Da
#[cfg(test)]
mod tests {
use super::MysqlRequestCandidateRepository;
use crate::run_migrations;
use aether_data_contracts::repository::candidates::{
RequestCandidateStatus, StoredRequestCandidate, UpsertRequestCandidateRecord,
RequestCandidateReadRepository, RequestCandidateStatus, StoredRequestCandidate,
UpsertRequestCandidateRecord,
};
#[tokio::test]
@@ -735,6 +791,128 @@ mod tests {
let _repository = MysqlRequestCandidateRepository::new(pool);
}
#[tokio::test]
async fn mysql_atomic_conflict_keeps_candidate_lifecycle_monotonic_when_configured() {
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
.ok()
.filter(|value| !value.trim().is_empty())
else {
eprintln!(
"skipping mysql candidate lifecycle test because AETHER_TEST_MYSQL_URL is unset"
);
return;
};
let pool = sqlx::mysql::MySqlPoolOptions::new()
.max_connections(2)
.connect(&database_url)
.await
.expect("mysql test pool should connect");
run_migrations(&pool)
.await
.expect("mysql migrations should run");
let repository = MysqlRequestCandidateRepository::new(pool.clone());
let request_id = format!("candidate-lifecycle-{}", uuid::Uuid::new_v4());
let terminal = stored_candidate(
&request_id,
"terminal",
0,
RequestCandidateStatus::Success,
Some(123),
Some(2_000_002),
);
super::upsert_merged_candidate(&pool, &terminal)
.await
.expect("terminal candidate should insert");
let stale_streaming = stored_candidate(
&request_id,
"stale-streaming",
0,
RequestCandidateStatus::Streaming,
Some(9_999),
Some(9_999_999),
);
super::upsert_merged_candidate(&pool, &stale_streaming)
.await
.expect("stale streaming conflict should execute");
let streaming = stored_candidate(
&request_id,
"streaming",
1,
RequestCandidateStatus::Streaming,
None,
None,
);
super::upsert_merged_candidate(&pool, &streaming)
.await
.expect("streaming candidate should insert");
let stale_pending = stored_candidate(
&request_id,
"stale-pending",
1,
RequestCandidateStatus::Pending,
None,
None,
);
super::upsert_merged_candidate(&pool, &stale_pending)
.await
.expect("stale pending conflict should execute");
let pending = stored_candidate(
&request_id,
"pending",
2,
RequestCandidateStatus::Pending,
Some(321),
None,
);
super::upsert_merged_candidate(&pool, &pending)
.await
.expect("pending candidate should insert");
let stale_available = stored_candidate(
&request_id,
"stale-available",
2,
RequestCandidateStatus::Available,
Some(9_999),
Some(9_999_999),
);
super::upsert_merged_candidate(&pool, &stale_available)
.await
.expect("stale available conflict should execute");
let candidates = repository
.list_by_request_id(&request_id)
.await
.expect("mysql request candidates should load");
let terminal = candidates
.iter()
.find(|candidate| candidate.candidate_index == 0)
.expect("terminal candidate should remain");
assert_eq!(terminal.status, RequestCandidateStatus::Success);
assert_eq!(terminal.latency_ms, Some(123));
assert_eq!(terminal.finished_at_unix_ms, Some(2_000_002));
let streaming = candidates
.iter()
.find(|candidate| candidate.candidate_index == 1)
.expect("streaming candidate should remain");
assert_eq!(streaming.status, RequestCandidateStatus::Streaming);
let pending = candidates
.iter()
.find(|candidate| candidate.candidate_index == 2)
.expect("pending candidate should remain");
assert_eq!(pending.status, RequestCandidateStatus::Pending);
assert_eq!(pending.latency_ms, Some(321));
assert_eq!(pending.finished_at_unix_ms, None);
sqlx::query("DELETE FROM request_candidates WHERE request_id = ?")
.bind(&request_id)
.execute(&pool)
.await
.expect("mysql candidate test rows should clean up");
}
#[test]
fn merge_candidate_keeps_terminal_status_when_streaming_arrives_late() {
let existing = StoredRequestCandidate::new(
@@ -805,4 +983,41 @@ mod tests {
Some(serde_json::json!({"terminal": true, "late": true}))
);
}
fn stored_candidate(
request_id: &str,
id: &str,
candidate_index: u32,
status: RequestCandidateStatus,
latency_ms: Option<i32>,
finished_at_unix_ms: Option<i64>,
) -> StoredRequestCandidate {
StoredRequestCandidate::new(
id.to_string(),
request_id.to_string(),
Some("user-1".to_string()),
Some("key-1".to_string()),
None,
None,
i32::try_from(candidate_index).expect("candidate index should fit"),
0,
Some("provider-1".to_string()),
Some("endpoint-1".to_string()),
Some("provider-key-1".to_string()),
status,
None,
false,
Some(200),
None,
None,
latency_ms,
None,
None,
None,
2_000_000,
Some(2_000_001),
finished_at_unix_ms,
)
.expect("stored candidate should build")
}
}
@@ -34,6 +34,91 @@ SELECT
FROM pool_member_scores
"#;
const UPSERT_PRESERVING_NULLABLE_TIMESTAMPS_SQL: &str = r#"
INSERT INTO pool_member_scores (
id, pool_kind, pool_id, member_kind, member_id, capability, scope_kind, scope_id,
score, hard_state, score_version, score_reason, last_ranked_at, last_scheduled_at,
last_success_at, last_failure_at, failure_count, last_probe_attempt_at,
last_probe_success_at, last_probe_failure_at, probe_failure_count, probe_status, updated_at
) VALUES (
?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?
)
ON DUPLICATE KEY UPDATE
pool_kind = VALUES(pool_kind),
pool_id = VALUES(pool_id),
member_kind = VALUES(member_kind),
member_id = VALUES(member_id),
capability = VALUES(capability),
scope_kind = VALUES(scope_kind),
scope_id = VALUES(scope_id),
score = VALUES(score),
hard_state = VALUES(hard_state),
score_version = VALUES(score_version),
score_reason = VALUES(score_reason),
last_ranked_at = VALUES(last_ranked_at),
last_scheduled_at = COALESCE(VALUES(last_scheduled_at), last_scheduled_at),
last_success_at = COALESCE(VALUES(last_success_at), last_success_at),
last_failure_at = COALESCE(VALUES(last_failure_at), last_failure_at),
failure_count = VALUES(failure_count),
last_probe_attempt_at = COALESCE(VALUES(last_probe_attempt_at), last_probe_attempt_at),
last_probe_success_at = COALESCE(VALUES(last_probe_success_at), last_probe_success_at),
last_probe_failure_at = COALESCE(VALUES(last_probe_failure_at), last_probe_failure_at),
probe_failure_count = VALUES(probe_failure_count),
probe_status = VALUES(probe_status),
updated_at = VALUES(updated_at)
"#;
const UPSERT_OAUTH_RECOVERY_SQL: &str = r#"
INSERT INTO pool_member_scores (
id, pool_kind, pool_id, member_kind, member_id, capability, scope_kind, scope_id,
score, hard_state, score_version, score_reason, last_ranked_at, last_scheduled_at,
last_success_at, last_failure_at, failure_count, last_probe_attempt_at,
last_probe_success_at, last_probe_failure_at, probe_failure_count, probe_status, updated_at
) VALUES (
?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?
)
ON DUPLICATE KEY UPDATE
pool_kind = VALUES(pool_kind),
pool_id = VALUES(pool_id),
member_kind = VALUES(member_kind),
member_id = VALUES(member_id),
capability = VALUES(capability),
scope_kind = VALUES(scope_kind),
scope_id = VALUES(scope_id),
score = IF(updated_at <= VALUES(updated_at), VALUES(score), score),
hard_state = IF(updated_at <= VALUES(updated_at), VALUES(hard_state), hard_state),
score_version = IF(updated_at <= VALUES(updated_at), VALUES(score_version), score_version),
score_reason = IF(updated_at <= VALUES(updated_at), VALUES(score_reason), score_reason),
last_ranked_at = IF(updated_at <= VALUES(updated_at), VALUES(last_ranked_at), last_ranked_at),
failure_count = IF(
last_failure_at IS NULL OR last_failure_at <= VALUES(updated_at),
VALUES(failure_count), failure_count),
last_failure_at = IF(
last_failure_at IS NULL OR last_failure_at <= VALUES(updated_at),
VALUES(last_failure_at), last_failure_at),
probe_status = IF(
(last_probe_attempt_at IS NOT NULL AND last_probe_attempt_at > VALUES(updated_at))
OR (last_probe_success_at IS NOT NULL AND last_probe_success_at > VALUES(updated_at))
OR (last_probe_failure_at IS NOT NULL AND last_probe_failure_at > VALUES(updated_at)),
probe_status, VALUES(probe_status)),
probe_failure_count = IF(
last_probe_failure_at IS NULL OR last_probe_failure_at <= VALUES(updated_at),
VALUES(probe_failure_count), probe_failure_count),
last_probe_failure_at = IF(
last_probe_failure_at IS NULL OR last_probe_failure_at <= VALUES(updated_at),
VALUES(last_probe_failure_at), last_probe_failure_at),
updated_at = GREATEST(updated_at, VALUES(updated_at))
"#;
fn pool_member_score_upsert_sql(mode: PoolMemberScoreUpsertMode) -> &'static str {
match mode {
PoolMemberScoreUpsertMode::PreserveExistingNullableTimestamps => {
UPSERT_PRESERVING_NULLABLE_TIMESTAMPS_SQL
}
PoolMemberScoreUpsertMode::OAuthRecovery => UPSERT_OAUTH_RECOVERY_SQL,
}
}
#[derive(Debug, Clone)]
pub struct MysqlPoolMemberScoreRepository {
pool: MysqlPool,
@@ -252,102 +337,69 @@ impl PoolScoreReadRepository for MysqlPoolMemberScoreRepository {
#[async_trait]
impl PoolMemberScoreWriteRepository for MysqlPoolMemberScoreRepository {
async fn upsert_pool_member_score(
async fn upsert_pool_member_score_with_mode(
&self,
score: UpsertPoolMemberScore,
mode: PoolMemberScoreUpsertMode,
) -> Result<StoredPoolMemberScore, DataLayerError> {
score.validate()?;
let stored = score.into_stored();
let score_reason = serde_json::to_string(&stored.score_reason)
.map_err(|err| DataLayerError::InvalidInput(err.to_string()))?;
sqlx::query(
r#"
INSERT INTO pool_member_scores (
id, pool_kind, pool_id, member_kind, member_id, capability, scope_kind, scope_id,
score, hard_state, score_version, score_reason, last_ranked_at, last_scheduled_at,
last_success_at, last_failure_at, failure_count, last_probe_attempt_at,
last_probe_success_at, last_probe_failure_at, probe_failure_count, probe_status, updated_at
) VALUES (
?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?
)
ON DUPLICATE KEY UPDATE
pool_kind = VALUES(pool_kind),
pool_id = VALUES(pool_id),
member_kind = VALUES(member_kind),
member_id = VALUES(member_id),
capability = VALUES(capability),
scope_kind = VALUES(scope_kind),
scope_id = VALUES(scope_id),
score = VALUES(score),
hard_state = VALUES(hard_state),
score_version = VALUES(score_version),
score_reason = VALUES(score_reason),
last_ranked_at = VALUES(last_ranked_at),
last_scheduled_at = COALESCE(VALUES(last_scheduled_at), last_scheduled_at),
last_success_at = COALESCE(VALUES(last_success_at), last_success_at),
last_failure_at = COALESCE(VALUES(last_failure_at), last_failure_at),
failure_count = VALUES(failure_count),
last_probe_attempt_at = COALESCE(VALUES(last_probe_attempt_at), last_probe_attempt_at),
last_probe_success_at = COALESCE(VALUES(last_probe_success_at), last_probe_success_at),
last_probe_failure_at = COALESCE(VALUES(last_probe_failure_at), last_probe_failure_at),
probe_failure_count = VALUES(probe_failure_count),
probe_status = VALUES(probe_status),
updated_at = VALUES(updated_at)
"#,
)
.bind(stored.id.as_str())
.bind(stored.pool_kind.as_str())
.bind(stored.pool_id.as_str())
.bind(stored.member_kind.as_str())
.bind(stored.member_id.as_str())
.bind(stored.capability.as_str())
.bind(stored.scope_kind.as_str())
.bind(stored.scope_id.as_deref())
.bind(stored.score)
.bind(stored.hard_state.as_database())
.bind(i64_from_u64(stored.score_version, "pool score version")?)
.bind(score_reason)
.bind(i64_opt_from_u64(
stored.last_ranked_at,
"pool score last_ranked_at",
)?)
.bind(i64_opt_from_u64(
stored.last_scheduled_at,
"pool score last_scheduled_at",
)?)
.bind(i64_opt_from_u64(
stored.last_success_at,
"pool score last_success_at",
)?)
.bind(i64_opt_from_u64(
stored.last_failure_at,
"pool score last_failure_at",
)?)
.bind(i64_from_u64(
stored.failure_count,
"pool score failure_count",
)?)
.bind(i64_opt_from_u64(
stored.last_probe_attempt_at,
"pool score last_probe_attempt_at",
)?)
.bind(i64_opt_from_u64(
stored.last_probe_success_at,
"pool score last_probe_success_at",
)?)
.bind(i64_opt_from_u64(
stored.last_probe_failure_at,
"pool score last_probe_failure_at",
)?)
.bind(i64_from_u64(
stored.probe_failure_count,
"pool score probe_failure_count",
)?)
.bind(stored.probe_status.as_database())
.bind(i64_from_u64(stored.updated_at, "pool score updated_at")?)
.execute(&self.pool)
.await
.map_sql_err()?;
sqlx::query(pool_member_score_upsert_sql(mode))
.bind(stored.id.as_str())
.bind(stored.pool_kind.as_str())
.bind(stored.pool_id.as_str())
.bind(stored.member_kind.as_str())
.bind(stored.member_id.as_str())
.bind(stored.capability.as_str())
.bind(stored.scope_kind.as_str())
.bind(stored.scope_id.as_deref())
.bind(stored.score)
.bind(stored.hard_state.as_database())
.bind(i64_from_u64(stored.score_version, "pool score version")?)
.bind(score_reason)
.bind(i64_opt_from_u64(
stored.last_ranked_at,
"pool score last_ranked_at",
)?)
.bind(i64_opt_from_u64(
stored.last_scheduled_at,
"pool score last_scheduled_at",
)?)
.bind(i64_opt_from_u64(
stored.last_success_at,
"pool score last_success_at",
)?)
.bind(i64_opt_from_u64(
stored.last_failure_at,
"pool score last_failure_at",
)?)
.bind(i64_from_u64(
stored.failure_count,
"pool score failure_count",
)?)
.bind(i64_opt_from_u64(
stored.last_probe_attempt_at,
"pool score last_probe_attempt_at",
)?)
.bind(i64_opt_from_u64(
stored.last_probe_success_at,
"pool score last_probe_success_at",
)?)
.bind(i64_opt_from_u64(
stored.last_probe_failure_at,
"pool score last_probe_failure_at",
)?)
.bind(i64_from_u64(
stored.probe_failure_count,
"pool score probe_failure_count",
)?)
.bind(stored.probe_status.as_database())
.bind(i64_from_u64(stored.updated_at, "pool score updated_at")?)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(stored)
}
@@ -577,3 +629,65 @@ fn i64_from_usize(value: usize, field: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field} exceeds signed 64-bit range")))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mysql_upsert_mode_selects_nullable_timestamp_semantics() {
let preserving = pool_member_score_upsert_sql(
PoolMemberScoreUpsertMode::PreserveExistingNullableTimestamps,
);
let recovering = pool_member_score_upsert_sql(PoolMemberScoreUpsertMode::OAuthRecovery);
for field in [
"last_scheduled_at",
"last_success_at",
"last_failure_at",
"last_probe_attempt_at",
"last_probe_success_at",
"last_probe_failure_at",
] {
assert!(preserving.contains(&format!("{field} = COALESCE(VALUES({field}), {field})")));
}
for field in [
"last_scheduled_at",
"last_success_at",
"last_probe_attempt_at",
"last_probe_success_at",
] {
assert!(!recovering.contains(&format!("\n {field} =")));
}
for field in [
"score",
"hard_state",
"score_version",
"score_reason",
"last_ranked_at",
] {
assert!(recovering.contains(&format!(
"{field} = IF(updated_at <= VALUES(updated_at), VALUES({field}), {field})"
)));
}
assert!(
recovering.contains("last_failure_at IS NULL OR last_failure_at <= VALUES(updated_at)")
);
assert!(recovering.contains("failure_count = IF("));
assert!(recovering.contains(
"last_probe_failure_at IS NULL OR last_probe_failure_at <= VALUES(updated_at)"
));
assert!(recovering.contains("probe_failure_count = IF("));
assert!(recovering.contains("updated_at = GREATEST(updated_at, VALUES(updated_at))"));
assert!(
recovering.find("failure_count =").unwrap()
< recovering.find("last_failure_at =").unwrap()
);
assert!(
recovering.find("probe_status =").unwrap()
< recovering.find("last_probe_failure_at =").unwrap()
);
assert_eq!(preserving.matches('?').count(), 23);
assert_eq!(recovering.matches('?').count(), 23);
}
}
@@ -1,3 +1,5 @@
use std::collections::BTreeMap;
use async_trait::async_trait;
use sqlx::{
mysql::{MySqlArguments, MySqlRow},
@@ -6,7 +8,9 @@ use sqlx::{
};
use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogKeyAdaptiveStateUpdate, ProviderCatalogKeyHealthStateUpdate,
ProviderCatalogKeyListQuery, ProviderCatalogKeyRuntimeMetadataUpdate,
ProviderCatalogKeyStatusSnapshotUpdate, ProviderCatalogReadRepository, ProviderCatalogSnapshot,
ProviderCatalogUpstreamMetadataNamespaceUpdate, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
StoredProviderCatalogKeyMaintenanceSummary, StoredProviderCatalogKeyPage,
@@ -716,7 +720,23 @@ WHERE id = ?
}
}
transaction.commit().await.map_sql_err()?;
Ok(keys.to_vec())
let key_ids = keys.iter().map(|key| key.id.clone()).collect::<Vec<_>>();
let mut reloaded = self
.list_keys_by_ids(&key_ids)
.await?
.into_iter()
.map(|key| (key.id.clone(), key))
.collect::<BTreeMap<_, _>>();
keys.iter()
.map(|key| {
reloaded.remove(&key.id).ok_or_else(|| {
DataLayerError::UnexpectedValue(format!(
"updated provider catalog key {} could not be reloaded",
key.id
))
})
})
.collect()
}
pub async fn delete_key(&self, key_id: &str) -> Result<bool, DataLayerError> {
@@ -967,6 +987,38 @@ WHERE id = ?
Ok(rows_affected > 0)
}
pub async fn update_key_oauth_runtime_state(
&self,
key_id: &str,
oauth_invalid_at_unix_secs: Option<u64>,
oauth_invalid_reason: Option<&str>,
encrypted_auth_config_update: Option<&str>,
updated_at_unix_secs: Option<u64>,
) -> Result<bool, DataLayerError> {
validate_non_empty(key_id, "provider catalog key_id")?;
let rows_affected = sqlx::query(
r#"
UPDATE provider_api_keys
SET oauth_invalid_at = ?, oauth_invalid_reason = ?,
auth_config = COALESCE(?, auth_config), updated_at = ?
WHERE id = ?
"#,
)
.bind(optional_i64_from_u64(
oauth_invalid_at_unix_secs,
"provider_api_keys.oauth_invalid_at",
)?)
.bind(oauth_invalid_reason)
.bind(encrypted_auth_config_update)
.bind(updated_at_unix_secs.unwrap_or_else(current_unix_secs) as i64)
.bind(key_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
pub async fn update_key_health_state(
&self,
key_id: &str,
@@ -1000,6 +1052,248 @@ WHERE id = ?
Ok(rows_affected > 0)
}
pub async fn reset_key_error_count(&self, key_id: &str) -> Result<bool, DataLayerError> {
validate_non_empty(key_id, "provider catalog key_id")?;
let rows_affected = sqlx::query(
r#"
UPDATE provider_api_keys
SET error_count = 0, updated_at = ?
WHERE id = ?
"#,
)
.bind(current_unix_secs() as i64)
.bind(key_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
pub async fn compare_and_update_key_adaptive_state(
&self,
update: &ProviderCatalogKeyAdaptiveStateUpdate,
) -> Result<bool, DataLayerError> {
validate_non_empty(&update.key_id, "provider catalog key_id")?;
let status_snapshot_patch = adaptive_status_snapshot_patch(&update.status_snapshot_patch)?;
let expected = update.expected.canonicalized();
let next = update.next.canonicalized();
let mut builder =
QueryBuilder::<MySql>::new("UPDATE provider_api_keys SET learned_rpm_limit = ");
builder
.push_bind(optional_i64_from_u32(next.learned_rpm_limit))
.push(", rpm_429_count = ")
.push_bind(optional_i64_from_u32(next.rpm_429_count))
.push(", last_429_at = ")
.push_bind(optional_i64_from_u64(
next.last_429_at_unix_secs,
"provider_api_keys.last_429_at",
)?)
.push(", last_429_type = ")
.push_bind(&next.last_429_type)
.push(", adjustment_history = ")
.push_bind(optional_json_to_string(
&next.adjustment_history,
"provider_api_keys.adjustment_history",
)?)
.push(", utilization_samples = ")
.push_bind(optional_json_to_string(
&next.utilization_samples,
"provider_api_keys.utilization_samples",
)?)
.push(", last_probe_increase_at = ")
.push_bind(optional_i64_from_u64(
next.last_probe_increase_at_unix_secs,
"provider_api_keys.last_probe_increase_at",
)?)
.push(", last_rpm_peak = ")
.push_bind(optional_i64_from_u32(next.last_rpm_peak))
.push(", concurrent_429_count = ")
.push_bind(optional_i64_from_u32(next.concurrent_429_count))
.push(", status_snapshot = ");
push_status_snapshot_shallow_patch(&mut builder, &status_snapshot_patch)?;
builder
.push(", updated_at = ")
.push_bind(
update
.updated_at_unix_secs
.unwrap_or_else(current_unix_secs) as i64,
)
.push(" WHERE id = ")
.push_bind(&update.key_id)
.push(" AND learned_rpm_limit <=> ")
.push_bind(optional_i64_from_u32(expected.learned_rpm_limit))
.push(" AND rpm_429_count <=> ")
.push_bind(optional_i64_from_u32(expected.rpm_429_count))
.push(" AND last_429_at <=> ")
.push_bind(optional_i64_from_u64(
expected.last_429_at_unix_secs,
"provider_api_keys.last_429_at",
)?)
.push(" AND last_429_type <=> ")
.push_bind(&expected.last_429_type)
.push(" AND JSON_EXTRACT(adjustment_history, '$') <=> CAST(")
.push_bind(optional_json_to_string(
&expected.adjustment_history,
"provider_api_keys.adjustment_history",
)?)
.push(" AS JSON)")
.push(" AND JSON_EXTRACT(utilization_samples, '$') <=> CAST(")
.push_bind(optional_json_to_string(
&expected.utilization_samples,
"provider_api_keys.utilization_samples",
)?)
.push(" AS JSON)")
.push(" AND last_probe_increase_at <=> ")
.push_bind(optional_i64_from_u64(
expected.last_probe_increase_at_unix_secs,
"provider_api_keys.last_probe_increase_at",
)?)
.push(" AND last_rpm_peak <=> ")
.push_bind(optional_i64_from_u32(expected.last_rpm_peak))
.push(" AND concurrent_429_count <=> ")
.push_bind(optional_i64_from_u32(expected.concurrent_429_count));
let rows_affected = builder
.build()
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
pub async fn update_key_runtime_metadata(
&self,
update: &ProviderCatalogKeyRuntimeMetadataUpdate,
) -> Result<bool, DataLayerError> {
validate_runtime_metadata_update(update)?;
let namespace_path = format!(
"$.{}",
serde_json::to_string(&update.namespace).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata namespace is not serializable: {err}"
))
})?
);
let metadata_value =
serde_json::to_string(&update.upstream_metadata_value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.upstream_metadata value is not serializable: {err}"
))
})?;
let expected_metadata_value = update
.expected_upstream_metadata_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}"
))
})?;
let mut builder = QueryBuilder::<MySql>::new(
"UPDATE provider_api_keys SET upstream_metadata = JSON_SET(\
COALESCE(NULLIF(upstream_metadata, ''), '{}'), ",
);
builder
.push_bind(namespace_path.clone())
.push(", CAST(")
.push_bind(metadata_value)
.push(" AS JSON)), status_snapshot = ");
push_status_snapshot_shallow_patch(&mut builder, &update.status_snapshot_patch)?;
builder
.push(", updated_at = ")
.push_bind(
update
.updated_at_unix_secs
.unwrap_or_else(current_unix_secs) as i64,
)
.push(" WHERE id = ")
.push_bind(&update.key_id)
.push(" AND JSON_EXTRACT(COALESCE(NULLIF(upstream_metadata, ''), '{}'), ")
.push_bind(namespace_path)
.push(") <=> CAST(")
.push_bind(expected_metadata_value)
.push(" AS JSON)");
let rows_affected = builder
.build()
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
pub async fn update_key_status_snapshot(
&self,
update: &ProviderCatalogKeyStatusSnapshotUpdate,
) -> Result<bool, DataLayerError> {
validate_non_empty(&update.key_id, "provider catalog key_id")?;
if !update.status_snapshot_patch.is_object() {
return Err(DataLayerError::InvalidInput(
"provider catalog status snapshot patch must be an object".to_string(),
));
}
let mut builder =
QueryBuilder::<MySql>::new("UPDATE provider_api_keys SET status_snapshot = ");
push_status_snapshot_shallow_patch(&mut builder, &update.status_snapshot_patch)?;
builder
.push(", updated_at = ")
.push_bind(
update
.updated_at_unix_secs
.unwrap_or_else(current_unix_secs) as i64,
)
.push(" WHERE id = ")
.push_bind(&update.key_id);
let rows_affected = builder
.build()
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
pub async fn compare_and_update_key_health_state(
&self,
update: &ProviderCatalogKeyHealthStateUpdate,
) -> Result<bool, DataLayerError> {
validate_non_empty(&update.key_id, "provider catalog key_id")?;
let rows_affected = sqlx::query(
r#"
UPDATE provider_api_keys
SET health_by_format = ?, circuit_breaker_by_format = ?, updated_at = ?
WHERE id = ?
AND JSON_EXTRACT(health_by_format, '$') <=> CAST(? AS JSON)
AND JSON_EXTRACT(circuit_breaker_by_format, '$') <=> CAST(? AS JSON)
"#,
)
.bind(optional_json_to_string(
&update.health_by_format,
"provider_api_keys.health_by_format",
)?)
.bind(optional_json_to_string(
&update.circuit_breaker_by_format,
"provider_api_keys.circuit_breaker_by_format",
)?)
.bind(current_unix_secs() as i64)
.bind(&update.key_id)
.bind(optional_json_to_string(
&update.expected_health_by_format,
"provider_api_keys.health_by_format",
)?)
.bind(optional_json_to_string(
&update.expected_circuit_breaker_by_format,
"provider_api_keys.circuit_breaker_by_format",
)?)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn reload_provider(
&self,
provider_id: &str,
@@ -1303,6 +1597,25 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
.await
}
async fn update_key_oauth_runtime_state(
&self,
key_id: &str,
oauth_invalid_at_unix_secs: Option<u64>,
oauth_invalid_reason: Option<&str>,
encrypted_auth_config_update: Option<&str>,
updated_at_unix_secs: Option<u64>,
) -> Result<bool, DataLayerError> {
Self::update_key_oauth_runtime_state(
self,
key_id,
oauth_invalid_at_unix_secs,
oauth_invalid_reason,
encrypted_auth_config_update,
updated_at_unix_secs,
)
.await
}
async fn update_key_health_state(
&self,
key_id: &str,
@@ -1319,6 +1632,38 @@ impl ProviderCatalogWriteRepository for MysqlProviderCatalogReadRepository {
)
.await
}
async fn reset_key_error_count(&self, key_id: &str) -> Result<bool, DataLayerError> {
Self::reset_key_error_count(self, key_id).await
}
async fn compare_and_update_key_adaptive_state(
&self,
update: &ProviderCatalogKeyAdaptiveStateUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_update_key_adaptive_state(self, update).await
}
async fn update_key_runtime_metadata(
&self,
update: &ProviderCatalogKeyRuntimeMetadataUpdate,
) -> Result<bool, DataLayerError> {
Self::update_key_runtime_metadata(self, update).await
}
async fn update_key_status_snapshot(
&self,
update: &ProviderCatalogKeyStatusSnapshotUpdate,
) -> Result<bool, DataLayerError> {
Self::update_key_status_snapshot(self, update).await
}
async fn compare_and_update_key_health_state(
&self,
update: &ProviderCatalogKeyHealthStateUpdate,
) -> Result<bool, DataLayerError> {
Self::compare_and_update_key_health_state(self, update).await
}
}
fn current_unix_secs() -> u64 {
@@ -1334,6 +1679,87 @@ fn validate_non_empty(value: &str, field_name: &str) -> Result<(), DataLayerErro
Ok(())
}
fn adaptive_status_snapshot_patch(
patch: &serde_json::Value,
) -> Result<serde_json::Value, DataLayerError> {
const OWNED_FIELDS: [&str; 6] = [
"observation_count",
"header_observation_count",
"latest_upstream_limit",
"learning_confidence",
"enforcement_active",
"known_boundary",
];
let object = patch.as_object().ok_or_else(|| {
DataLayerError::InvalidInput(
"provider catalog adaptive status snapshot patch must be an object".to_string(),
)
})?;
Ok(serde_json::Value::Object(
OWNED_FIELDS
.into_iter()
.filter_map(|field| {
object
.get(field)
.cloned()
.map(|value| (field.to_string(), value))
})
.collect(),
))
}
fn validate_runtime_metadata_update(
update: &ProviderCatalogKeyRuntimeMetadataUpdate,
) -> Result<(), DataLayerError> {
validate_non_empty(&update.key_id, "provider catalog key_id")?;
validate_non_empty(
&update.namespace,
"provider catalog runtime metadata namespace",
)?;
if !update.status_snapshot_patch.is_object() {
return Err(DataLayerError::InvalidInput(
"provider catalog runtime status snapshot patch must be an object".to_string(),
));
}
Ok(())
}
fn push_status_snapshot_shallow_patch<'args>(
builder: &mut QueryBuilder<'args, MySql>,
patch: &serde_json::Value,
) -> Result<(), DataLayerError> {
let object = patch.as_object().ok_or_else(|| {
DataLayerError::InvalidInput(
"provider catalog status snapshot patch must be an object".to_string(),
)
})?;
if object.is_empty() {
builder.push("COALESCE(NULLIF(status_snapshot, ''), '{}')");
return Ok(());
}
builder.push("JSON_SET(COALESCE(NULLIF(status_snapshot, ''), '{}')");
for (field, value) in object {
let path = format!(
"$.{}",
serde_json::to_string(field).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.status_snapshot field is not serializable: {err}"
))
})?
);
let value = serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"provider_api_keys.status_snapshot value is not serializable: {err}"
))
})?;
builder.push(", ").push_bind(path).push(", CAST(");
builder.push_bind(value).push(" AS JSON)");
}
builder.push(")");
Ok(())
}
fn validate_provider(provider: &StoredProviderCatalogProvider) -> Result<(), DataLayerError> {
validate_non_empty(&provider.id, "provider catalog provider.id")?;
validate_non_empty(&provider.name, "provider catalog provider.name")?;
@@ -1477,17 +1903,6 @@ SET
fingerprint = ?,
rpm_limit = ?,
concurrent_limit = ?,
learned_rpm_limit = ?,
concurrent_429_count = ?,
rpm_429_count = ?,
last_429_at = ?,
last_429_type = ?,
adjustment_history = ?,
utilization_samples = ?,
last_probe_increase_at = ?,
last_rpm_peak = ?,
last_models_fetch_at = ?,
last_models_fetch_error = ?,
auto_fetch_models = ?,
locked_models = ?,
model_include_patterns = ?,
@@ -1554,32 +1969,6 @@ fn key_update_query(
)?)
.bind(optional_i64_from_u32(key.rpm_limit))
.bind(key.concurrent_limit)
.bind(optional_i64_from_u32(key.learned_rpm_limit))
.bind(optional_i64_from_u32(key.concurrent_429_count).unwrap_or(0))
.bind(optional_i64_from_u32(key.rpm_429_count).unwrap_or(0))
.bind(optional_i64_from_u64(
key.last_429_at_unix_secs,
"provider_api_keys.last_429_at",
)?)
.bind(&key.last_429_type)
.bind(optional_json_to_string(
&key.adjustment_history,
"provider_api_keys.adjustment_history",
)?)
.bind(optional_json_to_string(
&key.utilization_samples,
"provider_api_keys.utilization_samples",
)?)
.bind(optional_i64_from_u64(
key.last_probe_increase_at_unix_secs,
"provider_api_keys.last_probe_increase_at",
)?)
.bind(optional_i64_from_u32(key.last_rpm_peak))
.bind(optional_i64_from_u64(
key.last_models_fetch_at_unix_secs,
"provider_api_keys.last_models_fetch_at",
)?)
.bind(&key.last_models_fetch_error)
.bind(key.auto_fetch_models)
.bind(optional_json_to_string(
&key.locked_models,
@@ -1935,6 +2324,23 @@ mod tests {
StoredProviderCatalogProvider,
};
use serde_json::json;
#[test]
fn ordinary_key_update_does_not_own_adaptive_runtime_fields() {
let sql = super::key_update_sql().to_ascii_lowercase();
for runtime_assignment in [
"learned_rpm_limit =",
"rpm_429_count =",
"last_429_at =",
"last_429_type =",
"adjustment_history =",
"utilization_samples =",
"last_probe_increase_at =",
"last_rpm_peak =",
] {
assert!(!sql.contains(runtime_assignment));
}
}
use sqlx::Execute;
#[tokio::test]
@@ -140,6 +140,14 @@ ORDER BY created_at ASC, id ASC
rows.iter().map(map_binding_row).collect()
}
async fn has_any_routing_group_binding(&self) -> Result<bool, DataLayerError> {
let row = sqlx::query("SELECT 1 FROM routing_group_bindings LIMIT 1")
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
Ok(row.is_some())
}
async fn list_routing_group_versions(
&self,
group_id: &str,