refactor: 大规模模块拆分与代码精简,新增 ai-pipeline/data-contracts 独立 crate

- 新增 aether-ai-pipeline 和 aether-data-contracts crate,将 pipeline 逻辑与数据契约从 gateway 中解耦
- 重构 admin handlers:拆分单体模块为 auth/billing/endpoint/features/model/observability/provider/system 等独立子模块
- 合并 chat/cli 重复代码路径:精简 conversion、finalize、planner 中的 sync/chat/cli 分支
- 重构 scheduler/executor/data 层,引入 facade 模式降低模块间耦合
- 移除冗余的 intent 模块,将 plan_fallback/policy/stream_path/sync_path 迁移至 executor
- 前端适配:调整 admin API 调用和 provider 模型测试对话框
This commit is contained in:
fawney19
2026-04-07 02:50:19 +08:00
parent 763ff03a7b
commit 5d96d6673b
732 changed files with 28593 additions and 20666 deletions

View File

@@ -3,7 +3,7 @@ use std::sync::RwLock;
use async_trait::async_trait;
use super::types::{
use super::{
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,

View File

@@ -1,11 +1,11 @@
mod memory;
mod sql;
mod types;
pub use memory::InMemoryProviderCatalogReadRepository;
pub use sql::SqlxProviderCatalogReadRepository;
pub use types::{
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::provider_catalog::{
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
};
pub use memory::InMemoryProviderCatalogReadRepository;
pub use sql::SqlxProviderCatalogReadRepository;

View File

@@ -1,12 +1,15 @@
use async_trait::async_trait;
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
use super::types::{
use super::{
ProviderCatalogKeyListQuery, ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogKeyPage,
StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
};
use crate::DataLayerError;
use crate::{
error::{postgres_error, SqlxResultExt},
DataLayerError,
};
const LIST_PROVIDERS_BY_IDS_PREFIX: &str = r#"
SELECT
@@ -271,7 +274,8 @@ impl SqlxProviderCatalogReadRepository {
)
.build()
.fetch_all(&self.pool)
.await?;
.await
.map_postgres_err()?;
rows.iter().map(map_provider_row).collect()
}
@@ -312,7 +316,8 @@ ORDER BY provider_priority ASC, name ASC
)
.bind(active_only)
.fetch_all(&self.pool)
.await?;
.await
.map_postgres_err()?;
rows.iter().map(map_provider_row).collect()
}
@@ -334,17 +339,16 @@ ORDER BY provider_priority ASC, name ASC
.await
{
Ok(rows) => rows,
Err(error) if is_missing_endpoint_health_score_column(&error) => {
build_list_query(
LIST_ENDPOINTS_BY_IDS_PREFIX_LEGACY,
endpoint_ids,
" ORDER BY api_format ASC, id ASC",
)
.build()
.fetch_all(&self.pool)
.await?
}
Err(error) => return Err(error.into()),
Err(error) if is_missing_endpoint_health_score_column(&error) => build_list_query(
LIST_ENDPOINTS_BY_IDS_PREFIX_LEGACY,
endpoint_ids,
" ORDER BY api_format ASC, id ASC",
)
.build()
.fetch_all(&self.pool)
.await
.map_postgres_err()?,
Err(error) => return Err(postgres_error(error)),
};
rows.iter().map(map_endpoint_row).collect()
}
@@ -367,17 +371,16 @@ ORDER BY provider_priority ASC, name ASC
.await
{
Ok(rows) => rows,
Err(error) if is_missing_endpoint_health_score_column(&error) => {
build_list_query(
LIST_ENDPOINTS_BY_PROVIDER_IDS_PREFIX_LEGACY,
provider_ids,
" ORDER BY provider_id ASC, api_format ASC, id ASC",
)
.build()
.fetch_all(&self.pool)
.await?
}
Err(error) => return Err(error.into()),
Err(error) if is_missing_endpoint_health_score_column(&error) => build_list_query(
LIST_ENDPOINTS_BY_PROVIDER_IDS_PREFIX_LEGACY,
provider_ids,
" ORDER BY provider_id ASC, api_format ASC, id ASC",
)
.build()
.fetch_all(&self.pool)
.await
.map_postgres_err()?,
Err(error) => return Err(postgres_error(error)),
};
rows.iter().map(map_endpoint_row).collect()
}
@@ -397,7 +400,8 @@ ORDER BY provider_priority ASC, name ASC
)
.build()
.fetch_all(&self.pool)
.await?;
.await
.map_postgres_err()?;
rows.iter().map(map_key_row).collect()
}
@@ -416,7 +420,8 @@ ORDER BY provider_priority ASC, name ASC
)
.build()
.fetch_all(&self.pool)
.await?;
.await
.map_postgres_err()?;
rows.iter().map(map_key_row).collect()
}
@@ -462,8 +467,9 @@ WHERE provider_id = $1
.bind(search_pattern.as_deref())
.bind(query.is_active)
.fetch_one(&self.pool)
.await?;
let total = count_row.try_get::<i64, _>("total")?.max(0) as usize;
.await
.map_postgres_err()?;
let total = row_get::<i64>(&count_row, "total")?.max(0) as usize;
let rows = sqlx::query(
r#"
@@ -530,7 +536,8 @@ LIMIT $5
.bind(offset)
.bind(limit)
.fetch_all(&self.pool)
.await?;
.await
.map_postgres_err()?;
let items = rows
.iter()
.map(map_key_row)
@@ -554,7 +561,8 @@ LIMIT $5
)
.build()
.fetch_all(&self.pool)
.await?;
.await
.map_postgres_err()?;
rows.iter().map(map_key_stats_row).collect()
}
@@ -597,7 +605,8 @@ WHERE id = $1
.bind(encrypted_auth_config)
.bind(expires_at_unix_secs.map(|value| value as f64))
.execute(&self.pool)
.await?
.await
.map_postgres_err()?
.rows_affected();
Ok(rows_affected > 0)
@@ -634,7 +643,7 @@ WHERE id = $1
));
}
let mut tx = self.pool.begin().await?;
let mut tx = self.pool.begin().await.map_postgres_err()?;
if let Some(target_priority) = shift_existing_priorities_from {
sqlx::query(
@@ -647,7 +656,8 @@ WHERE provider_priority IS NOT NULL
)
.bind(target_priority)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
}
sqlx::query(
@@ -752,9 +762,10 @@ INSERT INTO providers (
.bind(provider.created_at_unix_secs.map(|value| value as f64))
.bind(provider.updated_at_unix_secs.map(|value| value as f64))
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
tx.commit().await?;
tx.commit().await.map_err(postgres_error)?;
self.list_providers_by_ids(std::slice::from_ref(&provider.id))
.await?
@@ -871,7 +882,8 @@ WHERE id = $1
.bind(&provider.config)
.bind(provider.updated_at_unix_secs.map(|value| value as f64))
.execute(&self.pool)
.await?
.await
.map_postgres_err()?
.rows_affected();
if rows_affected == 0 {
@@ -908,7 +920,8 @@ WHERE id = $1
)
.bind(provider_id)
.execute(&self.pool)
.await?
.await
.map_postgres_err()?
.rows_affected();
Ok(rows_affected > 0)
@@ -926,26 +939,30 @@ WHERE id = $1
));
}
let mut tx = self.pool.begin().await?;
let mut tx = self.pool.begin().await.map_postgres_err()?;
sqlx::query(
"UPDATE user_preferences SET default_provider_id = NULL WHERE default_provider_id = $1",
)
.bind(provider_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
sqlx::query("UPDATE usage SET provider_id = NULL WHERE provider_id = $1")
.bind(provider_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
sqlx::query("UPDATE video_tasks SET provider_id = NULL WHERE provider_id = $1")
.bind(provider_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
sqlx::query("DELETE FROM request_candidates WHERE provider_id = $1")
.bind(provider_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
for endpoint_id in endpoint_ids {
sqlx::query(
@@ -953,44 +970,52 @@ WHERE id = $1
)
.bind(endpoint_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
sqlx::query("UPDATE video_tasks SET endpoint_id = NULL WHERE endpoint_id = $1")
.bind(endpoint_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
sqlx::query("DELETE FROM request_candidates WHERE endpoint_id = $1")
.bind(endpoint_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
}
for key_id in key_ids {
sqlx::query("DELETE FROM gemini_file_mappings WHERE key_id = $1")
.bind(key_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
sqlx::query(
"UPDATE usage SET provider_api_key_id = NULL WHERE provider_api_key_id = $1",
)
.bind(key_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
sqlx::query("UPDATE video_tasks SET key_id = NULL WHERE key_id = $1")
.bind(key_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
}
sqlx::query("DELETE FROM api_key_provider_mappings WHERE provider_id = $1")
.bind(provider_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
sqlx::query("DELETE FROM provider_usage_tracking WHERE provider_id = $1")
.bind(provider_id)
.execute(&mut *tx)
.await?;
.await
.map_postgres_err()?;
tx.commit().await?;
tx.commit().await.map_err(postgres_error)?;
Ok(())
}
@@ -1016,7 +1041,8 @@ WHERE id = $1
)
.bind(key_id)
.execute(&self.pool)
.await?
.await
.map_postgres_err()?
.rows_affected();
Ok(rows_affected > 0)
@@ -1217,7 +1243,8 @@ INSERT INTO provider_api_keys (
.bind(key.created_at_unix_secs.map(|value| value as f64))
.bind(key.updated_at_unix_secs.map(|value| value as f64))
.execute(&self.pool)
.await?;
.await
.map_postgres_err()?;
self.list_keys_by_ids(std::slice::from_ref(&key.id))
.await?
@@ -1246,7 +1273,7 @@ INSERT INTO provider_api_keys (
));
}
sqlx::query(
match sqlx::query(
r#"
INSERT INTO provider_endpoints (
id,
@@ -1311,7 +1338,77 @@ INSERT INTO provider_endpoints (
.bind(endpoint.created_at_unix_secs.map(|value| value as f64))
.bind(endpoint.updated_at_unix_secs.map(|value| value as f64))
.execute(&self.pool)
.await?;
.await
{
Ok(_) => {}
Err(error) if is_missing_endpoint_health_score_column(&error) => {
sqlx::query(
r#"
INSERT INTO provider_endpoints (
id,
provider_id,
api_format,
api_family,
endpoint_kind,
is_active,
base_url,
header_rules,
body_rules,
max_retries,
custom_path,
config,
format_acceptance_config,
proxy,
created_at,
updated_at
) VALUES (
$1,
$2,
$3,
$4,
$5,
$6,
$7,
$8,
$9,
$10,
$11,
$12,
$13,
$14,
CASE
WHEN $15::double precision IS NULL THEN NOW()
ELSE TO_TIMESTAMP($15::double precision)
END,
CASE
WHEN $16::double precision IS NULL THEN NOW()
ELSE TO_TIMESTAMP($16::double precision)
END
)
"#,
)
.bind(&endpoint.id)
.bind(&endpoint.provider_id)
.bind(&endpoint.api_format)
.bind(&endpoint.api_family)
.bind(&endpoint.endpoint_kind)
.bind(endpoint.is_active)
.bind(&endpoint.base_url)
.bind(&endpoint.header_rules)
.bind(&endpoint.body_rules)
.bind(endpoint.max_retries)
.bind(&endpoint.custom_path)
.bind(&endpoint.config)
.bind(&endpoint.format_acceptance_config)
.bind(&endpoint.proxy)
.bind(endpoint.created_at_unix_secs.map(|value| value as f64))
.bind(endpoint.updated_at_unix_secs.map(|value| value as f64))
.execute(&self.pool)
.await
.map_postgres_err()?;
}
Err(error) => return Err(postgres_error(error)),
}
self.list_endpoints_by_ids(std::slice::from_ref(&endpoint.id))
.await?
@@ -1340,7 +1437,7 @@ INSERT INTO provider_endpoints (
));
}
let rows_affected = sqlx::query(
let rows_affected = match sqlx::query(
r#"
UPDATE provider_endpoints
SET
@@ -1382,8 +1479,54 @@ WHERE id = $1
.bind(&endpoint.proxy)
.bind(endpoint.updated_at_unix_secs.map(|value| value as f64))
.execute(&self.pool)
.await?
.rows_affected();
.await
{
Ok(result) => result.rows_affected(),
Err(error) if is_missing_endpoint_health_score_column(&error) => sqlx::query(
r#"
UPDATE provider_endpoints
SET
provider_id = $2,
api_format = $3,
api_family = $4,
endpoint_kind = $5,
is_active = $6,
base_url = $7,
header_rules = $8,
body_rules = $9,
max_retries = $10,
custom_path = $11,
config = $12,
format_acceptance_config = $13,
proxy = $14,
updated_at = CASE
WHEN $15::double precision IS NULL THEN NOW()
ELSE TO_TIMESTAMP($15::double precision)
END
WHERE id = $1
"#,
)
.bind(&endpoint.id)
.bind(&endpoint.provider_id)
.bind(&endpoint.api_format)
.bind(&endpoint.api_family)
.bind(&endpoint.endpoint_kind)
.bind(endpoint.is_active)
.bind(&endpoint.base_url)
.bind(&endpoint.header_rules)
.bind(&endpoint.body_rules)
.bind(endpoint.max_retries)
.bind(&endpoint.custom_path)
.bind(&endpoint.config)
.bind(&endpoint.format_acceptance_config)
.bind(&endpoint.proxy)
.bind(endpoint.updated_at_unix_secs.map(|value| value as f64))
.execute(&self.pool)
.await
.map_postgres_err()?
.rows_affected(),
Err(error) => return Err(postgres_error(error)),
};
if rows_affected == 0 {
return Err(DataLayerError::UnexpectedValue(format!(
@@ -1419,7 +1562,8 @@ WHERE id = $1
)
.bind(endpoint_id)
.execute(&self.pool)
.await?
.await
.map_postgres_err()?
.rows_affected();
Ok(rows_affected > 0)
@@ -1511,7 +1655,8 @@ WHERE id = $1
.bind(key.is_active)
.bind(key.updated_at_unix_secs.map(|value| value as f64))
.execute(&self.pool)
.await?
.await
.map_postgres_err()?
.rows_affected();
if rows_affected == 0 {
@@ -1548,7 +1693,8 @@ WHERE id = $1
)
.bind(key_id)
.execute(&self.pool)
.await?
.await
.map_postgres_err()?
.rows_affected();
Ok(rows_affected > 0)
@@ -1583,7 +1729,8 @@ WHERE id = $1
.bind(health_by_format)
.bind(circuit_breaker_by_format)
.execute(&self.pool)
.await?
.await
.map_postgres_err()?
.rows_affected();
Ok(rows_affected > 0)
@@ -1769,9 +1916,15 @@ fn build_list_query<'a>(
builder
}
fn row_get<T>(row: &PgRow, column: &str) -> Result<T, DataLayerError>
where
for<'r> T: sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type<sqlx::Postgres>,
{
row.try_get(column).map_postgres_err()
}
fn map_provider_row(row: &PgRow) -> Result<StoredProviderCatalogProvider, DataLayerError> {
let quota_reset_day = row
.try_get::<Option<i32>, _>("quota_reset_day")?
let quota_reset_day = row_get::<Option<i32>>(row, "quota_reset_day")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1780,8 +1933,7 @@ fn map_provider_row(row: &PgRow) -> Result<StoredProviderCatalogProvider, DataLa
})
})
.transpose()?;
let created_at_unix_secs = row
.try_get::<Option<i64>, _>("created_at_unix_secs")?
let created_at_unix_secs = row_get::<Option<i64>>(row, "created_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1790,8 +1942,7 @@ fn map_provider_row(row: &PgRow) -> Result<StoredProviderCatalogProvider, DataLa
})
})
.transpose()?;
let updated_at_unix_secs = row
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
let updated_at_unix_secs = row_get::<Option<i64>>(row, "updated_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1801,40 +1952,37 @@ fn map_provider_row(row: &PgRow) -> Result<StoredProviderCatalogProvider, DataLa
})
.transpose()?;
Ok(StoredProviderCatalogProvider::new(
row.try_get("id")?,
row.try_get("name")?,
row.try_get("website")?,
row.try_get("provider_type")?,
row_get(row, "id")?,
row_get(row, "name")?,
row_get(row, "website")?,
row_get(row, "provider_type")?,
)?
.with_description(row.try_get("description")?)
.with_description(row_get(row, "description")?)
.with_billing_fields(
row.try_get("billing_type")?,
row.try_get("monthly_quota_usd")?,
row.try_get("monthly_used_usd")?,
row_get(row, "billing_type")?,
row_get(row, "monthly_quota_usd")?,
row_get(row, "monthly_used_usd")?,
quota_reset_day,
row.try_get::<Option<i64>, _>("quota_last_reset_at_unix_secs")?
.map(|value| value as u64),
row.try_get::<Option<i64>, _>("quota_expires_at_unix_secs")?
.map(|value| value as u64),
row_get::<Option<i64>>(row, "quota_last_reset_at_unix_secs")?.map(|value| value as u64),
row_get::<Option<i64>>(row, "quota_expires_at_unix_secs")?.map(|value| value as u64),
)
.with_routing_fields(row.try_get("provider_priority")?)
.with_routing_fields(row_get(row, "provider_priority")?)
.with_transport_fields(
row.try_get("is_active")?,
row.try_get("keep_priority_on_conversion")?,
row.try_get("enable_format_conversion")?,
row.try_get("concurrent_limit")?,
row.try_get("max_retries")?,
row.try_get("proxy")?,
row.try_get("request_timeout")?,
row.try_get("stream_first_byte_timeout")?,
row.try_get("config")?,
row_get(row, "is_active")?,
row_get(row, "keep_priority_on_conversion")?,
row_get(row, "enable_format_conversion")?,
row_get(row, "concurrent_limit")?,
row_get(row, "max_retries")?,
row_get(row, "proxy")?,
row_get(row, "request_timeout")?,
row_get(row, "stream_first_byte_timeout")?,
row_get(row, "config")?,
)
.with_timestamps(created_at_unix_secs, updated_at_unix_secs))
}
fn map_endpoint_row(row: &PgRow) -> Result<StoredProviderCatalogEndpoint, DataLayerError> {
let created_at_unix_secs = row
.try_get::<Option<i64>, _>("created_at_unix_secs")?
let created_at_unix_secs = row_get::<Option<i64>>(row, "created_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1843,8 +1991,7 @@ fn map_endpoint_row(row: &PgRow) -> Result<StoredProviderCatalogEndpoint, DataLa
})
})
.transpose()?;
let updated_at_unix_secs = row
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
let updated_at_unix_secs = row_get::<Option<i64>>(row, "updated_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1854,12 +2001,12 @@ fn map_endpoint_row(row: &PgRow) -> Result<StoredProviderCatalogEndpoint, DataLa
})
.transpose()?;
StoredProviderCatalogEndpoint::new(
row.try_get("id")?,
row.try_get("provider_id")?,
row.try_get("api_format")?,
row.try_get("api_family")?,
row.try_get("endpoint_kind")?,
row.try_get("is_active")?,
row_get(row, "id")?,
row_get(row, "provider_id")?,
row_get(row, "api_format")?,
row_get(row, "api_family")?,
row_get(row, "endpoint_kind")?,
row_get(row, "is_active")?,
)?
.with_timestamps(created_at_unix_secs, updated_at_unix_secs)
.with_health_score(
@@ -1869,14 +2016,14 @@ fn map_endpoint_row(row: &PgRow) -> Result<StoredProviderCatalogEndpoint, DataLa
.unwrap_or(1.0),
)
.with_transport_fields(
row.try_get("base_url")?,
row.try_get("header_rules")?,
row.try_get("body_rules")?,
row.try_get("max_retries")?,
row.try_get("custom_path")?,
row.try_get("config")?,
row.try_get("format_acceptance_config")?,
row.try_get("proxy")?,
row_get(row, "base_url")?,
row_get(row, "header_rules")?,
row_get(row, "body_rules")?,
row_get(row, "max_retries")?,
row_get(row, "custom_path")?,
row_get(row, "config")?,
row_get(row, "format_acceptance_config")?,
row_get(row, "proxy")?,
)
}
@@ -1893,15 +2040,14 @@ fn is_missing_endpoint_health_score_column(error: &sqlx::Error) -> bool {
fn map_key_stats_row(row: &PgRow) -> Result<StoredProviderCatalogKeyStats, DataLayerError> {
StoredProviderCatalogKeyStats::new(
row.try_get("provider_id")?,
row.try_get("total_keys")?,
row.try_get("active_keys")?,
row_get(row, "provider_id")?,
row_get(row, "total_keys")?,
row_get(row, "active_keys")?,
)
}
fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError> {
let rpm_limit = row
.try_get::<Option<i32>, _>("rpm_limit")?
let rpm_limit = row_get::<Option<i32>>(row, "rpm_limit")?
.map(|value| {
u32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1910,8 +2056,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let learned_rpm_limit = row
.try_get::<Option<i32>, _>("learned_rpm_limit")?
let learned_rpm_limit = row_get::<Option<i32>>(row, "learned_rpm_limit")?
.map(|value| {
u32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1920,8 +2065,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let concurrent_429_count = row
.try_get::<Option<i32>, _>("concurrent_429_count")?
let concurrent_429_count = row_get::<Option<i32>>(row, "concurrent_429_count")?
.map(|value| {
u32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1930,8 +2074,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let rpm_429_count = row
.try_get::<Option<i32>, _>("rpm_429_count")?
let rpm_429_count = row_get::<Option<i32>>(row, "rpm_429_count")?
.map(|value| {
u32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1940,8 +2083,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let request_count = row
.try_get::<Option<i32>, _>("request_count")?
let request_count = row_get::<Option<i32>>(row, "request_count")?
.map(|value| {
u32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1950,8 +2092,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let success_count = row
.try_get::<Option<i32>, _>("success_count")?
let success_count = row_get::<Option<i32>>(row, "success_count")?
.map(|value| {
u32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1960,8 +2101,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let error_count = row
.try_get::<Option<i32>, _>("error_count")?
let error_count = row_get::<Option<i32>>(row, "error_count")?
.map(|value| {
u32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1970,8 +2110,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let total_response_time_ms = row
.try_get::<Option<i32>, _>("total_response_time_ms")?
let total_response_time_ms = row_get::<Option<i32>>(row, "total_response_time_ms")?
.map(|value| {
u32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -1980,28 +2119,27 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let last_probe_increase_at_unix_secs = row
.try_get::<Option<i64>, _>("last_probe_increase_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
"invalid provider_api_keys.last_probe_increase_at_unix_secs: {value}"
))
let last_probe_increase_at_unix_secs =
row_get::<Option<i64>>(row, "last_probe_increase_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
"invalid provider_api_keys.last_probe_increase_at_unix_secs: {value}"
))
})
})
})
.transpose()?;
let last_models_fetch_at_unix_secs = row
.try_get::<Option<i64>, _>("last_models_fetch_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
"invalid provider_api_keys.last_models_fetch_at_unix_secs: {value}"
))
.transpose()?;
let last_models_fetch_at_unix_secs =
row_get::<Option<i64>>(row, "last_models_fetch_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
"invalid provider_api_keys.last_models_fetch_at_unix_secs: {value}"
))
})
})
})
.transpose()?;
let oauth_invalid_at_unix_secs = row
.try_get::<Option<i64>, _>("oauth_invalid_at_unix_secs")?
.transpose()?;
let oauth_invalid_at_unix_secs = row_get::<Option<i64>>(row, "oauth_invalid_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -2010,8 +2148,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let last_used_at_unix_secs = row
.try_get::<Option<i64>, _>("last_used_at_unix_secs")?
let last_used_at_unix_secs = row_get::<Option<i64>>(row, "last_used_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -2020,8 +2157,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let created_at_unix_secs = row
.try_get::<Option<i64>, _>("created_at_unix_secs")?
let created_at_unix_secs = row_get::<Option<i64>>(row, "created_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -2030,8 +2166,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
})
})
.transpose()?;
let updated_at_unix_secs = row
.try_get::<Option<i64>, _>("updated_at_unix_secs")?
let updated_at_unix_secs = row_get::<Option<i64>>(row, "updated_at_unix_secs")?
.map(|value| {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
@@ -2042,24 +2177,24 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
.transpose()?;
StoredProviderCatalogKey::new(
row.try_get("id")?,
row.try_get("provider_id")?,
row.try_get("name")?,
row.try_get("auth_type")?,
row.try_get("capabilities")?,
row.try_get("is_active")?,
row_get(row, "id")?,
row_get(row, "provider_id")?,
row_get(row, "name")?,
row_get(row, "auth_type")?,
row_get(row, "capabilities")?,
row_get(row, "is_active")?,
)?
.with_transport_fields(
row.try_get("api_formats")?,
row.try_get("api_key")?,
row.try_get("auth_config")?,
row.try_get("rate_multipliers")?,
row.try_get("global_priority_by_format")?,
row.try_get("allowed_models")?,
row.try_get::<Option<i64>, _>("expires_at_unix_secs")?
row_get(row, "api_formats")?,
row_get(row, "api_key")?,
row_get(row, "auth_config")?,
row_get(row, "rate_multipliers")?,
row_get(row, "global_priority_by_format")?,
row_get(row, "allowed_models")?,
row_get::<Option<i64>>(row, "expires_at_unix_secs")?
.and_then(|value| u64::try_from(value).ok()),
row.try_get("proxy")?,
row.try_get("fingerprint")?,
row_get(row, "proxy")?,
row_get(row, "fingerprint")?,
)
.map(|key| {
let mut key = key

View File

@@ -1,738 +0,0 @@
use async_trait::async_trait;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogProvider {
pub id: String,
pub name: String,
pub description: Option<String>,
pub website: Option<String>,
pub provider_type: String,
pub billing_type: Option<String>,
pub monthly_quota_usd: Option<f64>,
pub monthly_used_usd: Option<f64>,
pub quota_reset_day: Option<u64>,
pub quota_last_reset_at_unix_secs: Option<u64>,
pub quota_expires_at_unix_secs: Option<u64>,
pub provider_priority: i32,
pub is_active: bool,
pub keep_priority_on_conversion: bool,
pub enable_format_conversion: bool,
pub concurrent_limit: Option<i32>,
pub max_retries: Option<i32>,
pub proxy: Option<serde_json::Value>,
pub request_timeout_secs: Option<f64>,
pub stream_first_byte_timeout_secs: Option<f64>,
pub config: Option<serde_json::Value>,
pub created_at_unix_secs: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
}
impl StoredProviderCatalogProvider {
pub fn new(
id: String,
name: String,
website: Option<String>,
provider_type: String,
) -> Result<Self, crate::DataLayerError> {
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"providers.name is empty".to_string(),
));
}
if provider_type.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"providers.provider_type is empty".to_string(),
));
}
Ok(Self {
id,
name,
description: None,
website,
provider_type,
billing_type: None,
monthly_quota_usd: None,
monthly_used_usd: None,
quota_reset_day: None,
quota_last_reset_at_unix_secs: None,
quota_expires_at_unix_secs: None,
provider_priority: 0,
is_active: true,
keep_priority_on_conversion: false,
enable_format_conversion: false,
concurrent_limit: None,
max_retries: None,
proxy: None,
request_timeout_secs: None,
stream_first_byte_timeout_secs: None,
config: None,
created_at_unix_secs: None,
updated_at_unix_secs: None,
})
}
#[allow(clippy::too_many_arguments)]
pub fn with_transport_fields(
mut self,
is_active: bool,
keep_priority_on_conversion: bool,
enable_format_conversion: bool,
concurrent_limit: Option<i32>,
max_retries: Option<i32>,
proxy: Option<serde_json::Value>,
request_timeout_secs: Option<f64>,
stream_first_byte_timeout_secs: Option<f64>,
config: Option<serde_json::Value>,
) -> Self {
self.is_active = is_active;
self.keep_priority_on_conversion = keep_priority_on_conversion;
self.enable_format_conversion = enable_format_conversion;
self.concurrent_limit = concurrent_limit;
self.max_retries = max_retries;
self.proxy = proxy;
self.request_timeout_secs = request_timeout_secs;
self.stream_first_byte_timeout_secs = stream_first_byte_timeout_secs;
self.config = config;
self
}
pub fn with_description(mut self, description: Option<String>) -> Self {
self.description = description;
self
}
#[allow(clippy::too_many_arguments)]
pub fn with_billing_fields(
mut self,
billing_type: Option<String>,
monthly_quota_usd: Option<f64>,
monthly_used_usd: Option<f64>,
quota_reset_day: Option<u64>,
quota_last_reset_at_unix_secs: Option<u64>,
quota_expires_at_unix_secs: Option<u64>,
) -> Self {
self.billing_type = billing_type;
self.monthly_quota_usd = monthly_quota_usd;
self.monthly_used_usd = monthly_used_usd;
self.quota_reset_day = quota_reset_day;
self.quota_last_reset_at_unix_secs = quota_last_reset_at_unix_secs;
self.quota_expires_at_unix_secs = quota_expires_at_unix_secs;
self
}
pub fn with_routing_fields(mut self, provider_priority: i32) -> Self {
self.provider_priority = provider_priority;
self
}
pub fn with_timestamps(
mut self,
created_at_unix_secs: Option<u64>,
updated_at_unix_secs: Option<u64>,
) -> Self {
self.created_at_unix_secs = created_at_unix_secs;
self.updated_at_unix_secs = updated_at_unix_secs;
self
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogEndpoint {
pub id: String,
pub provider_id: String,
pub api_format: String,
pub api_family: Option<String>,
pub endpoint_kind: Option<String>,
pub is_active: bool,
pub health_score: f64,
pub base_url: String,
pub header_rules: Option<serde_json::Value>,
pub body_rules: Option<serde_json::Value>,
pub max_retries: Option<i32>,
pub custom_path: Option<String>,
pub config: Option<serde_json::Value>,
pub format_acceptance_config: Option<serde_json::Value>,
pub proxy: Option<serde_json::Value>,
pub created_at_unix_secs: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
}
impl StoredProviderCatalogEndpoint {
pub fn new(
id: String,
provider_id: String,
api_format: String,
api_family: Option<String>,
endpoint_kind: Option<String>,
is_active: bool,
) -> Result<Self, crate::DataLayerError> {
if api_format.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider_endpoints.api_format is empty".to_string(),
));
}
Ok(Self {
id,
provider_id,
api_format,
api_family,
endpoint_kind,
is_active,
health_score: 1.0,
base_url: String::new(),
header_rules: None,
body_rules: None,
max_retries: None,
custom_path: None,
config: None,
format_acceptance_config: None,
proxy: None,
created_at_unix_secs: None,
updated_at_unix_secs: None,
})
}
#[allow(clippy::too_many_arguments)]
pub fn with_transport_fields(
mut self,
base_url: String,
header_rules: Option<serde_json::Value>,
body_rules: Option<serde_json::Value>,
max_retries: Option<i32>,
custom_path: Option<String>,
config: Option<serde_json::Value>,
format_acceptance_config: Option<serde_json::Value>,
proxy: Option<serde_json::Value>,
) -> Result<Self, crate::DataLayerError> {
if base_url.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider_endpoints.base_url is empty".to_string(),
));
}
self.base_url = base_url;
self.header_rules = header_rules;
self.body_rules = body_rules;
self.max_retries = max_retries;
self.custom_path = custom_path;
self.config = config;
self.format_acceptance_config = format_acceptance_config;
self.proxy = proxy;
Ok(self)
}
pub fn with_health_score(mut self, health_score: f64) -> Self {
self.health_score = health_score;
self
}
pub fn with_timestamps(
mut self,
created_at_unix_secs: Option<u64>,
updated_at_unix_secs: Option<u64>,
) -> Self {
self.created_at_unix_secs = created_at_unix_secs;
self.updated_at_unix_secs = updated_at_unix_secs;
self
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogKey {
pub id: String,
pub provider_id: String,
pub name: String,
pub auth_type: String,
pub capabilities: Option<serde_json::Value>,
pub is_active: bool,
pub api_formats: Option<serde_json::Value>,
pub encrypted_api_key: String,
pub encrypted_auth_config: Option<String>,
pub note: Option<String>,
pub internal_priority: i32,
pub rate_multipliers: Option<serde_json::Value>,
pub global_priority_by_format: Option<serde_json::Value>,
pub allowed_models: Option<serde_json::Value>,
pub expires_at_unix_secs: Option<u64>,
pub cache_ttl_minutes: i32,
pub max_probe_interval_minutes: i32,
pub proxy: Option<serde_json::Value>,
pub fingerprint: Option<serde_json::Value>,
pub rpm_limit: Option<u32>,
pub learned_rpm_limit: Option<u32>,
pub concurrent_429_count: Option<u32>,
pub rpm_429_count: Option<u32>,
pub last_429_at_unix_secs: Option<u64>,
pub last_429_type: Option<String>,
pub adjustment_history: Option<serde_json::Value>,
pub utilization_samples: Option<serde_json::Value>,
pub last_probe_increase_at_unix_secs: Option<u64>,
pub request_count: Option<u32>,
pub success_count: Option<u32>,
pub error_count: Option<u32>,
pub total_response_time_ms: Option<u32>,
pub last_used_at_unix_secs: Option<u64>,
pub auto_fetch_models: bool,
pub last_models_fetch_at_unix_secs: Option<u64>,
pub last_models_fetch_error: Option<String>,
pub locked_models: Option<serde_json::Value>,
pub model_include_patterns: Option<serde_json::Value>,
pub model_exclude_patterns: Option<serde_json::Value>,
pub upstream_metadata: Option<serde_json::Value>,
pub oauth_invalid_at_unix_secs: Option<u64>,
pub oauth_invalid_reason: Option<String>,
pub status_snapshot: Option<serde_json::Value>,
pub created_at_unix_secs: Option<u64>,
pub updated_at_unix_secs: Option<u64>,
pub health_by_format: Option<serde_json::Value>,
pub circuit_breaker_by_format: Option<serde_json::Value>,
}
impl StoredProviderCatalogKey {
pub fn new(
id: String,
provider_id: String,
name: String,
auth_type: String,
capabilities: Option<serde_json::Value>,
is_active: bool,
) -> Result<Self, crate::DataLayerError> {
if name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider_api_keys.name is empty".to_string(),
));
}
if auth_type.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider_api_keys.auth_type is empty".to_string(),
));
}
Ok(Self {
id,
provider_id,
name,
auth_type,
capabilities,
is_active,
api_formats: None,
encrypted_api_key: String::new(),
encrypted_auth_config: None,
note: None,
internal_priority: 50,
rate_multipliers: None,
global_priority_by_format: None,
allowed_models: None,
expires_at_unix_secs: None,
cache_ttl_minutes: 5,
max_probe_interval_minutes: 32,
proxy: None,
fingerprint: None,
rpm_limit: None,
learned_rpm_limit: None,
concurrent_429_count: None,
rpm_429_count: None,
last_429_at_unix_secs: None,
last_429_type: None,
adjustment_history: None,
utilization_samples: None,
last_probe_increase_at_unix_secs: None,
request_count: None,
success_count: None,
error_count: None,
total_response_time_ms: None,
last_used_at_unix_secs: None,
auto_fetch_models: false,
last_models_fetch_at_unix_secs: None,
last_models_fetch_error: None,
locked_models: None,
model_include_patterns: None,
model_exclude_patterns: None,
upstream_metadata: None,
oauth_invalid_at_unix_secs: None,
oauth_invalid_reason: None,
status_snapshot: None,
created_at_unix_secs: None,
updated_at_unix_secs: None,
health_by_format: None,
circuit_breaker_by_format: None,
})
}
#[allow(clippy::too_many_arguments)]
pub fn with_transport_fields(
mut self,
api_formats: Option<serde_json::Value>,
encrypted_api_key: String,
encrypted_auth_config: Option<String>,
rate_multipliers: Option<serde_json::Value>,
global_priority_by_format: Option<serde_json::Value>,
allowed_models: Option<serde_json::Value>,
expires_at_unix_secs: Option<u64>,
proxy: Option<serde_json::Value>,
fingerprint: Option<serde_json::Value>,
) -> Result<Self, crate::DataLayerError> {
if encrypted_api_key.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider_api_keys.api_key is empty".to_string(),
));
}
self.api_formats = api_formats;
self.encrypted_api_key = encrypted_api_key;
self.encrypted_auth_config = encrypted_auth_config;
self.rate_multipliers = rate_multipliers;
self.global_priority_by_format = global_priority_by_format;
self.allowed_models = allowed_models;
self.expires_at_unix_secs = expires_at_unix_secs;
self.proxy = proxy;
self.fingerprint = fingerprint;
Ok(self)
}
#[allow(clippy::too_many_arguments)]
pub fn with_rate_limit_fields(
mut self,
rpm_limit: Option<u32>,
learned_rpm_limit: Option<u32>,
concurrent_429_count: Option<u32>,
rpm_429_count: Option<u32>,
last_429_at_unix_secs: Option<u64>,
adjustment_history: Option<serde_json::Value>,
request_count: Option<u32>,
success_count: Option<u32>,
) -> Self {
self.rpm_limit = rpm_limit;
self.learned_rpm_limit = learned_rpm_limit;
self.concurrent_429_count = concurrent_429_count;
self.rpm_429_count = rpm_429_count;
self.last_429_at_unix_secs = last_429_at_unix_secs;
self.adjustment_history = adjustment_history;
self.request_count = request_count;
self.success_count = success_count;
self
}
pub fn with_usage_fields(
mut self,
error_count: Option<u32>,
total_response_time_ms: Option<u32>,
) -> Self {
self.error_count = error_count;
self.total_response_time_ms = total_response_time_ms;
self
}
pub fn with_health_fields(
mut self,
health_by_format: Option<serde_json::Value>,
circuit_breaker_by_format: Option<serde_json::Value>,
) -> Self {
self.health_by_format = health_by_format;
self.circuit_breaker_by_format = circuit_breaker_by_format;
self
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct ProviderCatalogKeyListQuery {
pub provider_id: String,
pub search: Option<String>,
pub is_active: Option<bool>,
pub offset: usize,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogKeyPage {
pub items: Vec<StoredProviderCatalogKey>,
pub total: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredProviderCatalogKeyStats {
pub provider_id: String,
pub total_keys: u64,
pub active_keys: u64,
}
impl StoredProviderCatalogKeyStats {
pub fn new(
provider_id: String,
total_keys: i64,
active_keys: i64,
) -> Result<Self, crate::DataLayerError> {
if provider_id.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"provider key stats provider_id is empty".to_string(),
));
}
if total_keys < 0 || active_keys < 0 {
return Err(crate::DataLayerError::UnexpectedValue(
"provider key stats count is negative".to_string(),
));
}
Ok(Self {
provider_id,
total_keys: total_keys as u64,
active_keys: active_keys as u64,
})
}
}
#[async_trait]
pub trait ProviderCatalogReadRepository: Send + Sync {
async fn list_providers(
&self,
active_only: bool,
) -> Result<Vec<StoredProviderCatalogProvider>, crate::DataLayerError>;
async fn list_providers_by_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogProvider>, crate::DataLayerError>;
async fn list_endpoints_by_ids(
&self,
endpoint_ids: &[String],
) -> Result<Vec<StoredProviderCatalogEndpoint>, crate::DataLayerError>;
async fn list_endpoints_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogEndpoint>, crate::DataLayerError>;
async fn list_keys_by_ids(
&self,
key_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, crate::DataLayerError>;
async fn list_keys_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKey>, crate::DataLayerError>;
async fn list_keys_page(
&self,
query: &ProviderCatalogKeyListQuery,
) -> Result<StoredProviderCatalogKeyPage, crate::DataLayerError>;
async fn list_key_stats_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderCatalogKeyStats>, crate::DataLayerError>;
}
#[async_trait]
pub trait ProviderCatalogWriteRepository: Send + Sync {
async fn create_provider(
&self,
provider: &StoredProviderCatalogProvider,
shift_existing_priorities_from: Option<i32>,
) -> Result<StoredProviderCatalogProvider, crate::DataLayerError>;
async fn update_provider(
&self,
provider: &StoredProviderCatalogProvider,
) -> Result<StoredProviderCatalogProvider, crate::DataLayerError>;
async fn delete_provider(&self, provider_id: &str) -> Result<bool, crate::DataLayerError>;
async fn cleanup_deleted_provider_refs(
&self,
provider_id: &str,
endpoint_ids: &[String],
key_ids: &[String],
) -> Result<(), crate::DataLayerError>;
async fn create_endpoint(
&self,
endpoint: &StoredProviderCatalogEndpoint,
) -> Result<StoredProviderCatalogEndpoint, crate::DataLayerError>;
async fn update_endpoint(
&self,
endpoint: &StoredProviderCatalogEndpoint,
) -> Result<StoredProviderCatalogEndpoint, crate::DataLayerError>;
async fn delete_endpoint(&self, endpoint_id: &str) -> Result<bool, crate::DataLayerError>;
async fn create_key(
&self,
key: &StoredProviderCatalogKey,
) -> Result<StoredProviderCatalogKey, crate::DataLayerError>;
async fn update_key(
&self,
key: &StoredProviderCatalogKey,
) -> Result<StoredProviderCatalogKey, crate::DataLayerError>;
async fn delete_key(&self, key_id: &str) -> Result<bool, crate::DataLayerError>;
async fn clear_key_oauth_invalid_marker(
&self,
key_id: &str,
) -> Result<bool, crate::DataLayerError>;
async fn update_key_oauth_credentials(
&self,
key_id: &str,
encrypted_api_key: &str,
encrypted_auth_config: Option<&str>,
expires_at_unix_secs: Option<u64>,
) -> Result<bool, crate::DataLayerError>;
async fn update_key_health_state(
&self,
key_id: &str,
is_active: bool,
health_by_format: Option<&serde_json::Value>,
circuit_breaker_by_format: Option<&serde_json::Value>,
) -> Result<bool, crate::DataLayerError>;
}
#[cfg(test)]
mod tests {
use super::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};
#[test]
fn rejects_empty_provider_name() {
assert!(StoredProviderCatalogProvider::new(
"provider-1".to_string(),
"".to_string(),
None,
"custom".to_string(),
)
.is_err());
}
#[test]
fn rejects_empty_endpoint_api_format() {
assert!(StoredProviderCatalogEndpoint::new(
"endpoint-1".to_string(),
"provider-1".to_string(),
"".to_string(),
None,
None,
true,
)
.is_err());
}
#[test]
fn rejects_empty_endpoint_base_url() {
let endpoint = StoredProviderCatalogEndpoint::new(
"endpoint-1".to_string(),
"provider-1".to_string(),
"openai:chat".to_string(),
None,
None,
true,
)
.expect("endpoint should build");
assert!(endpoint
.with_transport_fields("".to_string(), None, None, None, None, None, None, None,)
.is_err());
}
#[test]
fn rejects_empty_key_auth_type() {
assert!(StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"default".to_string(),
"".to_string(),
None,
true,
)
.is_err());
}
#[test]
fn rejects_empty_encrypted_api_key() {
let key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"default".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build");
assert!(key
.with_transport_fields(
None,
"".to_string(),
None,
None,
None,
None,
None,
None,
None
)
.is_err());
}
#[test]
fn stores_key_rate_limit_fields() {
let key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"default".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_rate_limit_fields(
Some(100),
Some(80),
Some(2),
Some(3),
Some(1_700_000_000),
Some(serde_json::json!([{"new_limit": 80}])),
Some(120),
Some(110),
);
assert_eq!(key.rpm_limit, Some(100));
assert_eq!(key.learned_rpm_limit, Some(80));
assert_eq!(key.concurrent_429_count, Some(2));
assert_eq!(key.rpm_429_count, Some(3));
assert_eq!(key.last_429_at_unix_secs, Some(1_700_000_000));
assert_eq!(key.request_count, Some(120));
assert_eq!(key.success_count, Some(110));
}
#[test]
fn stores_key_health_fields() {
let key = StoredProviderCatalogKey::new(
"key-1".to_string(),
"provider-1".to_string(),
"default".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_health_fields(
Some(serde_json::json!({"openai:chat": {"health_score": 0.4}})),
Some(serde_json::json!({"openai:chat": {"open": true}})),
);
assert_eq!(
key.health_by_format,
Some(serde_json::json!({"openai:chat": {"health_score": 0.4}}))
);
assert_eq!(
key.circuit_breaker_by_format,
Some(serde_json::json!({"openai:chat": {"open": true}}))
);
}
}