mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
feat: 扩展 Rust gateway 全功能模块,新增 billing/crypto/wallet crate 及完整数据层
- 新增 aether-billing、aether-crypto、aether-wallet 独立 crate - aether-data 扩展 repository 层:announcements、auth_modules、billing、 candidate_selection、gemini_file_mappings、global_models、management_tokens、 oauth_providers、proxy_nodes、quota、users、wallet 等模块 - aether-gateway 新增 api/auth/billing/control/middleware/scheduler/usage/ video_tasks/hooks/maintenance/model_fetch/provider_transport 等功能模块 - 重构 executor decision 和 gateway state 为模块目录结构 - 新增 gateway router、frontdoor 路由层及对应测试 - Python 侧 API 路由重构,新增 compat/support 模块 - 前端 Logo 组件更新及 Provider 管理页面调整
This commit is contained in:
154
crates/aether-data/src/repository/candidate_selection/memory.rs
Normal file
154
crates/aether-data/src/repository/candidate_selection/memory.rs
Normal file
@@ -0,0 +1,154 @@
|
||||
use std::sync::RwLock;
|
||||
|
||||
use async_trait::async_trait;
|
||||
|
||||
use super::types::{MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow};
|
||||
use crate::DataLayerError;
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
pub struct InMemoryMinimalCandidateSelectionReadRepository {
|
||||
rows: RwLock<Vec<StoredMinimalCandidateSelectionRow>>,
|
||||
}
|
||||
|
||||
impl InMemoryMinimalCandidateSelectionReadRepository {
|
||||
pub fn seed<I>(rows: I) -> Self
|
||||
where
|
||||
I: IntoIterator<Item = StoredMinimalCandidateSelectionRow>,
|
||||
{
|
||||
Self {
|
||||
rows: RwLock::new(rows.into_iter().collect()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionReadRepository for InMemoryMinimalCandidateSelectionReadRepository {
|
||||
async fn list_for_exact_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let api_format = api_format.trim();
|
||||
let mut rows = self
|
||||
.rows
|
||||
.read()
|
||||
.expect("candidate selection repository lock")
|
||||
.iter()
|
||||
.filter(|row| {
|
||||
row.provider_is_active
|
||||
&& row.endpoint_is_active
|
||||
&& row.key_is_active
|
||||
&& row.model_is_active
|
||||
&& row.model_is_available
|
||||
&& row.endpoint_api_format.eq_ignore_ascii_case(api_format)
|
||||
&& row.key_supports_api_format(api_format)
|
||||
})
|
||||
.cloned()
|
||||
.collect::<Vec<_>>();
|
||||
rows.sort_by(|left, right| {
|
||||
left.provider_priority
|
||||
.cmp(&right.provider_priority)
|
||||
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
|
||||
.then(left.provider_id.cmp(&right.provider_id))
|
||||
.then(left.endpoint_id.cmp(&right.endpoint_id))
|
||||
.then(left.key_id.cmp(&right.key_id))
|
||||
.then(left.model_id.cmp(&right.model_id))
|
||||
});
|
||||
Ok(rows)
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_global_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self.list_for_exact_api_format(api_format).await?;
|
||||
Ok(rows
|
||||
.into_iter()
|
||||
.filter(|row| row.global_model_name == global_model_name)
|
||||
.collect())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
use crate::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
};
|
||||
|
||||
fn sample_row(
|
||||
provider_id: &str,
|
||||
api_format: &str,
|
||||
global_model_name: &str,
|
||||
provider_priority: i32,
|
||||
) -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: provider_id.to_string(),
|
||||
provider_name: provider_id.to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
provider_priority,
|
||||
provider_is_active: true,
|
||||
endpoint_id: format!("endpoint-{provider_id}"),
|
||||
endpoint_api_format: api_format.to_string(),
|
||||
endpoint_api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
endpoint_is_active: true,
|
||||
key_id: format!("key-{provider_id}"),
|
||||
key_name: "prod".to_string(),
|
||||
key_auth_type: "api_key".to_string(),
|
||||
key_is_active: true,
|
||||
key_api_formats: Some(vec![api_format.to_string()]),
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 50,
|
||||
key_global_priority_by_format: None,
|
||||
model_id: format!("model-{provider_id}"),
|
||||
global_model_id: "global-model-1".to_string(),
|
||||
global_model_name: global_model_name.to_string(),
|
||||
global_model_mappings: None,
|
||||
global_model_supports_streaming: Some(true),
|
||||
model_provider_model_name: global_model_name.to_string(),
|
||||
model_provider_model_mappings: None,
|
||||
model_supports_streaming: None,
|
||||
model_is_active: true,
|
||||
model_is_available: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn filters_by_exact_api_format_and_global_model() {
|
||||
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
sample_row("provider-2", "openai:chat", "gpt-4.1", 20),
|
||||
sample_row("provider-1", "openai:chat", "gpt-4.1", 10),
|
||||
sample_row("provider-3", "openai:responses", "gpt-4.1", 5),
|
||||
sample_row("provider-4", "openai:chat", "gpt-4.1-mini", 1),
|
||||
]);
|
||||
|
||||
let rows = repository
|
||||
.list_for_exact_api_format_and_global_model("openai:chat", "gpt-4.1")
|
||||
.await
|
||||
.expect("list should succeed");
|
||||
|
||||
assert_eq!(rows.len(), 2);
|
||||
assert_eq!(rows[0].provider_id, "provider-1");
|
||||
assert_eq!(rows[1].provider_id, "provider-2");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn filters_by_exact_api_format_only() {
|
||||
let repository = InMemoryMinimalCandidateSelectionReadRepository::seed(vec![
|
||||
sample_row("provider-2", "openai:chat", "gpt-4.1", 20),
|
||||
sample_row("provider-1", "openai:chat", "gpt-4.1-mini", 10),
|
||||
sample_row("provider-3", "openai:responses", "gpt-4.1", 5),
|
||||
]);
|
||||
|
||||
let rows = repository
|
||||
.list_for_exact_api_format("openai:chat")
|
||||
.await
|
||||
.expect("list should succeed");
|
||||
|
||||
assert_eq!(rows.len(), 2);
|
||||
assert_eq!(rows[0].provider_id, "provider-1");
|
||||
assert_eq!(rows[1].provider_id, "provider-2");
|
||||
}
|
||||
}
|
||||
10
crates/aether-data/src/repository/candidate_selection/mod.rs
Normal file
10
crates/aether-data/src/repository/candidate_selection/mod.rs
Normal file
@@ -0,0 +1,10 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod types;
|
||||
|
||||
pub use memory::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
pub use sql::SqlxMinimalCandidateSelectionReadRepository;
|
||||
pub use types::{
|
||||
MinimalCandidateSelectionReadRepository, MinimalCandidateSelectionRepository,
|
||||
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
|
||||
};
|
||||
548
crates/aether-data/src/repository/candidate_selection/sql.rs
Normal file
548
crates/aether-data/src/repository/candidate_selection/sql.rs
Normal file
@@ -0,0 +1,548 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row};
|
||||
|
||||
use super::types::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredProviderModelMapping,
|
||||
};
|
||||
use crate::DataLayerError;
|
||||
|
||||
const LIST_FOR_EXACT_API_FORMAT_SQL: &str = r#"
|
||||
SELECT
|
||||
p.id AS provider_id,
|
||||
p.name AS provider_name,
|
||||
p.provider_type AS provider_type,
|
||||
p.provider_priority AS provider_priority,
|
||||
p.is_active AS provider_is_active,
|
||||
pe.id AS endpoint_id,
|
||||
pe.api_format AS endpoint_api_format,
|
||||
pe.api_family AS endpoint_api_family,
|
||||
pe.endpoint_kind AS endpoint_kind,
|
||||
pe.is_active AS endpoint_is_active,
|
||||
pak.id AS key_id,
|
||||
pak.name AS key_name,
|
||||
pak.auth_type AS key_auth_type,
|
||||
pak.is_active AS key_is_active,
|
||||
pak.api_formats AS key_api_formats,
|
||||
pak.allowed_models AS key_allowed_models,
|
||||
pak.capabilities AS key_capabilities,
|
||||
pak.internal_priority AS key_internal_priority,
|
||||
pak.global_priority_by_format AS key_global_priority_by_format,
|
||||
m.id AS model_id,
|
||||
m.global_model_id AS global_model_id,
|
||||
gm.name AS global_model_name,
|
||||
CASE
|
||||
WHEN gm.config IS NOT NULL THEN gm.config -> 'model_mappings'
|
||||
ELSE NULL
|
||||
END AS global_model_mappings,
|
||||
CASE
|
||||
WHEN gm.config IS NOT NULL AND gm.config ? 'streaming'
|
||||
THEN (gm.config ->> 'streaming')::BOOLEAN
|
||||
ELSE NULL
|
||||
END AS global_model_supports_streaming,
|
||||
m.provider_model_name AS model_provider_model_name,
|
||||
m.provider_model_mappings AS model_provider_model_mappings,
|
||||
m.supports_streaming AS model_supports_streaming,
|
||||
m.is_active AS model_is_active,
|
||||
m.is_available AS model_is_available
|
||||
FROM providers p
|
||||
INNER JOIN provider_endpoints pe
|
||||
ON pe.provider_id = p.id
|
||||
INNER JOIN provider_api_keys pak
|
||||
ON pak.provider_id = p.id
|
||||
INNER JOIN models m
|
||||
ON m.provider_id = p.id
|
||||
INNER JOIN global_models gm
|
||||
ON gm.id = m.global_model_id
|
||||
WHERE p.is_active = TRUE
|
||||
AND pe.is_active = TRUE
|
||||
AND pak.is_active = TRUE
|
||||
AND m.is_active = TRUE
|
||||
AND m.is_available = TRUE
|
||||
AND gm.is_active = TRUE
|
||||
AND LOWER(pe.api_format) = LOWER($1)
|
||||
AND (
|
||||
pak.api_formats IS NULL
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM json_array_elements_text(pak.api_formats) AS fmt(value)
|
||||
WHERE LOWER(fmt.value) = LOWER($1)
|
||||
)
|
||||
)
|
||||
ORDER BY
|
||||
gm.name ASC,
|
||||
p.provider_priority ASC,
|
||||
pak.internal_priority ASC,
|
||||
p.id ASC,
|
||||
pe.id ASC,
|
||||
pak.id ASC,
|
||||
m.id ASC
|
||||
"#;
|
||||
|
||||
const LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL: &str = r#"
|
||||
SELECT
|
||||
p.id AS provider_id,
|
||||
p.name AS provider_name,
|
||||
p.provider_type AS provider_type,
|
||||
p.provider_priority AS provider_priority,
|
||||
p.is_active AS provider_is_active,
|
||||
pe.id AS endpoint_id,
|
||||
pe.api_format AS endpoint_api_format,
|
||||
pe.api_family AS endpoint_api_family,
|
||||
pe.endpoint_kind AS endpoint_kind,
|
||||
pe.is_active AS endpoint_is_active,
|
||||
pak.id AS key_id,
|
||||
pak.name AS key_name,
|
||||
pak.auth_type AS key_auth_type,
|
||||
pak.is_active AS key_is_active,
|
||||
pak.api_formats AS key_api_formats,
|
||||
pak.allowed_models AS key_allowed_models,
|
||||
pak.capabilities AS key_capabilities,
|
||||
pak.internal_priority AS key_internal_priority,
|
||||
pak.global_priority_by_format AS key_global_priority_by_format,
|
||||
m.id AS model_id,
|
||||
m.global_model_id AS global_model_id,
|
||||
gm.name AS global_model_name,
|
||||
CASE
|
||||
WHEN gm.config IS NOT NULL THEN gm.config -> 'model_mappings'
|
||||
ELSE NULL
|
||||
END AS global_model_mappings,
|
||||
CASE
|
||||
WHEN gm.config IS NOT NULL AND gm.config ? 'streaming'
|
||||
THEN (gm.config ->> 'streaming')::BOOLEAN
|
||||
ELSE NULL
|
||||
END AS global_model_supports_streaming,
|
||||
m.provider_model_name AS model_provider_model_name,
|
||||
m.provider_model_mappings AS model_provider_model_mappings,
|
||||
m.supports_streaming AS model_supports_streaming,
|
||||
m.is_active AS model_is_active,
|
||||
m.is_available AS model_is_available
|
||||
FROM providers p
|
||||
INNER JOIN provider_endpoints pe
|
||||
ON pe.provider_id = p.id
|
||||
INNER JOIN provider_api_keys pak
|
||||
ON pak.provider_id = p.id
|
||||
INNER JOIN models m
|
||||
ON m.provider_id = p.id
|
||||
INNER JOIN global_models gm
|
||||
ON gm.id = m.global_model_id
|
||||
WHERE p.is_active = TRUE
|
||||
AND pe.is_active = TRUE
|
||||
AND pak.is_active = TRUE
|
||||
AND m.is_active = TRUE
|
||||
AND m.is_available = TRUE
|
||||
AND gm.is_active = TRUE
|
||||
AND LOWER(pe.api_format) = LOWER($1)
|
||||
AND gm.name = $2
|
||||
AND (
|
||||
pak.api_formats IS NULL
|
||||
OR EXISTS (
|
||||
SELECT 1
|
||||
FROM json_array_elements_text(pak.api_formats) AS fmt(value)
|
||||
WHERE LOWER(fmt.value) = LOWER($1)
|
||||
)
|
||||
)
|
||||
ORDER BY
|
||||
p.provider_priority ASC,
|
||||
pak.internal_priority ASC,
|
||||
p.id ASC,
|
||||
pe.id ASC,
|
||||
pak.id ASC,
|
||||
m.id ASC
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqlxMinimalCandidateSelectionReadRepository {
|
||||
pool: PgPool,
|
||||
}
|
||||
|
||||
impl SqlxMinimalCandidateSelectionReadRepository {
|
||||
pub fn new(pool: PgPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
pub fn pool(&self) -> &PgPool {
|
||||
&self.pool
|
||||
}
|
||||
|
||||
pub async fn list_for_exact_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_FOR_EXACT_API_FORMAT_SQL)
|
||||
.bind(api_format)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
rows.iter().map(map_candidate_selection_row).collect()
|
||||
}
|
||||
|
||||
pub async fn list_for_exact_api_format_and_global_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_FOR_EXACT_API_FORMAT_AND_GLOBAL_MODEL_SQL)
|
||||
.bind(api_format)
|
||||
.bind(global_model_name)
|
||||
.fetch_all(&self.pool)
|
||||
.await?;
|
||||
rows.iter().map(map_candidate_selection_row).collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionReadRepository for SqlxMinimalCandidateSelectionReadRepository {
|
||||
async fn list_for_exact_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Self::list_for_exact_api_format(self, api_format).await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_global_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Self::list_for_exact_api_format_and_global_model(self, api_format, global_model_name).await
|
||||
}
|
||||
}
|
||||
|
||||
fn map_candidate_selection_row(
|
||||
row: &sqlx::postgres::PgRow,
|
||||
) -> Result<StoredMinimalCandidateSelectionRow, DataLayerError> {
|
||||
Ok(StoredMinimalCandidateSelectionRow {
|
||||
provider_id: row.try_get("provider_id")?,
|
||||
provider_name: row.try_get("provider_name")?,
|
||||
provider_type: row.try_get("provider_type")?,
|
||||
provider_priority: row.try_get("provider_priority")?,
|
||||
provider_is_active: row.try_get("provider_is_active")?,
|
||||
endpoint_id: row.try_get("endpoint_id")?,
|
||||
endpoint_api_format: row.try_get("endpoint_api_format")?,
|
||||
endpoint_api_family: row.try_get("endpoint_api_family")?,
|
||||
endpoint_kind: row.try_get("endpoint_kind")?,
|
||||
endpoint_is_active: row.try_get("endpoint_is_active")?,
|
||||
key_id: row.try_get("key_id")?,
|
||||
key_name: row.try_get("key_name")?,
|
||||
key_auth_type: row.try_get("key_auth_type")?,
|
||||
key_is_active: row.try_get("key_is_active")?,
|
||||
key_api_formats: parse_string_list(
|
||||
row.try_get("key_api_formats")?,
|
||||
"provider_api_keys.api_formats",
|
||||
)?,
|
||||
key_allowed_models: parse_string_list(
|
||||
row.try_get("key_allowed_models")?,
|
||||
"provider_api_keys.allowed_models",
|
||||
)?,
|
||||
key_capabilities: row.try_get("key_capabilities")?,
|
||||
key_internal_priority: row.try_get("key_internal_priority")?,
|
||||
key_global_priority_by_format: row.try_get("key_global_priority_by_format")?,
|
||||
model_id: row.try_get("model_id")?,
|
||||
global_model_id: row.try_get("global_model_id")?,
|
||||
global_model_name: row.try_get("global_model_name")?,
|
||||
global_model_mappings: parse_string_list(
|
||||
row.try_get("global_model_mappings")?,
|
||||
"global_models.config.model_mappings",
|
||||
)?,
|
||||
global_model_supports_streaming: row.try_get("global_model_supports_streaming")?,
|
||||
model_provider_model_name: row.try_get("model_provider_model_name")?,
|
||||
model_provider_model_mappings: parse_provider_model_mappings(
|
||||
row.try_get("model_provider_model_mappings")?,
|
||||
)?,
|
||||
model_supports_streaming: row.try_get("model_supports_streaming")?,
|
||||
model_is_active: row.try_get("model_is_active")?,
|
||||
model_is_available: row.try_get("model_is_available")?,
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_string_list(
|
||||
value: Option<serde_json::Value>,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
parse_string_list_value(&value, field_name)
|
||||
}
|
||||
|
||||
fn parse_string_list_value(
|
||||
value: &serde_json::Value,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
match value {
|
||||
serde_json::Value::Null => Ok(None),
|
||||
serde_json::Value::Array(array) => parse_string_list_array(array, field_name).map(Some),
|
||||
serde_json::Value::String(raw) => parse_embedded_string_list(raw, field_name),
|
||||
_ => Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} is not a JSON array"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_embedded_string_list(
|
||||
raw: &str,
|
||||
field_name: &str,
|
||||
) -> Result<Option<Vec<String>>, DataLayerError> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if let Ok(decoded) = serde_json::from_str::<serde_json::Value>(raw) {
|
||||
return parse_string_list_value(&decoded, field_name);
|
||||
}
|
||||
|
||||
Ok(Some(vec![raw.to_string()]))
|
||||
}
|
||||
|
||||
fn parse_string_list_array(
|
||||
array: &[serde_json::Value],
|
||||
field_name: &str,
|
||||
) -> Result<Vec<String>, DataLayerError> {
|
||||
let mut items = Vec::with_capacity(array.len());
|
||||
for item in array {
|
||||
let Some(item) = item.as_str() else {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains a non-string item"
|
||||
)));
|
||||
};
|
||||
let item = item.trim();
|
||||
if !item.is_empty() {
|
||||
items.push(item.to_string());
|
||||
}
|
||||
}
|
||||
Ok(items)
|
||||
}
|
||||
|
||||
fn parse_provider_model_mappings(
|
||||
value: Option<serde_json::Value>,
|
||||
) -> Result<Option<Vec<StoredProviderModelMapping>>, DataLayerError> {
|
||||
let Some(value) = value else {
|
||||
return Ok(None);
|
||||
};
|
||||
parse_provider_model_mappings_value(&value)
|
||||
}
|
||||
|
||||
fn parse_provider_model_mappings_value(
|
||||
value: &serde_json::Value,
|
||||
) -> Result<Option<Vec<StoredProviderModelMapping>>, DataLayerError> {
|
||||
match value {
|
||||
serde_json::Value::Null => Ok(None),
|
||||
serde_json::Value::Array(array) => parse_provider_model_mappings_array(array),
|
||||
serde_json::Value::Object(object) => {
|
||||
parse_provider_model_mapping_object(object).map(|mapping| Some(vec![mapping]))
|
||||
}
|
||||
serde_json::Value::String(raw) => parse_embedded_provider_model_mappings(raw),
|
||||
_ => Err(DataLayerError::UnexpectedValue(
|
||||
"models.provider_model_mappings is not a JSON array".to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_embedded_provider_model_mappings(
|
||||
raw: &str,
|
||||
) -> Result<Option<Vec<StoredProviderModelMapping>>, DataLayerError> {
|
||||
let raw = raw.trim();
|
||||
if raw.is_empty() || raw.eq_ignore_ascii_case("null") {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
if let Ok(decoded) = serde_json::from_str::<serde_json::Value>(raw) {
|
||||
return parse_provider_model_mappings_value(&decoded);
|
||||
}
|
||||
|
||||
Ok(Some(vec![StoredProviderModelMapping {
|
||||
name: raw.to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
}]))
|
||||
}
|
||||
|
||||
fn parse_provider_model_mappings_array(
|
||||
array: &[serde_json::Value],
|
||||
) -> Result<Option<Vec<StoredProviderModelMapping>>, DataLayerError> {
|
||||
let mut mappings = Vec::with_capacity(array.len());
|
||||
for raw in array {
|
||||
match raw {
|
||||
serde_json::Value::Object(object) => {
|
||||
if let Some(mapping) = parse_provider_model_mapping_object_lenient(object)? {
|
||||
mappings.push(mapping);
|
||||
}
|
||||
}
|
||||
serde_json::Value::String(raw) => {
|
||||
let raw = raw.trim();
|
||||
if !raw.is_empty() {
|
||||
mappings.push(StoredProviderModelMapping {
|
||||
name: raw.to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
});
|
||||
}
|
||||
}
|
||||
serde_json::Value::Null => {}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if mappings.is_empty() {
|
||||
Ok(None)
|
||||
} else {
|
||||
Ok(Some(mappings))
|
||||
}
|
||||
}
|
||||
|
||||
fn parse_provider_model_mapping_object(
|
||||
object: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Result<StoredProviderModelMapping, DataLayerError> {
|
||||
parse_provider_model_mapping_object_lenient(object)?.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"models.provider_model_mappings item is missing a valid name".to_string(),
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_provider_model_mapping_object_lenient(
|
||||
object: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Result<Option<StoredProviderModelMapping>, DataLayerError> {
|
||||
let Some(name) = object
|
||||
.get("name")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
|
||||
let priority = object
|
||||
.get("priority")
|
||||
.and_then(serde_json::Value::as_i64)
|
||||
.unwrap_or(1)
|
||||
.max(1);
|
||||
let api_formats = parse_string_list(
|
||||
object.get("api_formats").cloned(),
|
||||
"models.provider_model_mappings.api_formats",
|
||||
)?;
|
||||
|
||||
Ok(Some(StoredProviderModelMapping {
|
||||
name: name.to_string(),
|
||||
priority: i32::try_from(priority).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid models.provider_model_mappings.priority: {priority}"
|
||||
))
|
||||
})?,
|
||||
api_formats,
|
||||
}))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
parse_provider_model_mappings, parse_string_list,
|
||||
SqlxMinimalCandidateSelectionReadRepository,
|
||||
};
|
||||
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
use crate::repository::candidate_selection::StoredProviderModelMapping;
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_constructs_from_lazy_pool() {
|
||||
let factory = PostgresPoolFactory::new(PostgresPoolConfig {
|
||||
database_url: "postgres://localhost/aether".to_string(),
|
||||
min_connections: 1,
|
||||
max_connections: 4,
|
||||
acquire_timeout_ms: 1_000,
|
||||
idle_timeout_ms: 5_000,
|
||||
max_lifetime_ms: 30_000,
|
||||
statement_cache_capacity: 64,
|
||||
require_ssl: false,
|
||||
})
|
||||
.expect("factory should build");
|
||||
|
||||
let pool = factory.connect_lazy().expect("pool should build");
|
||||
let repository = SqlxMinimalCandidateSelectionReadRepository::new(pool);
|
||||
let _ = repository.pool();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_string_list_accepts_stringified_array() {
|
||||
let parsed = parse_string_list(
|
||||
Some(json!("[\"gpt-5.2\", \"gpt-5\"]")),
|
||||
"provider_api_keys.allowed_models",
|
||||
)
|
||||
.expect("stringified array should parse");
|
||||
|
||||
assert_eq!(
|
||||
parsed,
|
||||
Some(vec!["gpt-5.2".to_string(), "gpt-5".to_string()])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_string_list_accepts_single_string() {
|
||||
let parsed = parse_string_list(Some(json!("gpt-5.2")), "provider_api_keys.allowed_models")
|
||||
.expect("single string should parse");
|
||||
|
||||
assert_eq!(parsed, Some(vec!["gpt-5.2".to_string()]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_provider_model_mappings_accepts_stringified_array() {
|
||||
let parsed = parse_provider_model_mappings(Some(json!(
|
||||
"[{\"name\":\"gpt-5.2\",\"priority\":2,\"api_formats\":[\"openai:chat\"]}]"
|
||||
)))
|
||||
.expect("stringified provider_model_mappings should parse");
|
||||
|
||||
assert_eq!(
|
||||
parsed,
|
||||
Some(vec![StoredProviderModelMapping {
|
||||
name: "gpt-5.2".to_string(),
|
||||
priority: 2,
|
||||
api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
}])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_provider_model_mappings_accepts_single_string_alias() {
|
||||
let parsed = parse_provider_model_mappings(Some(json!("gpt-5.2")))
|
||||
.expect("single-string provider_model_mappings should parse");
|
||||
|
||||
assert_eq!(
|
||||
parsed,
|
||||
Some(vec![StoredProviderModelMapping {
|
||||
name: "gpt-5.2".to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
}])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_provider_model_mappings_skips_invalid_array_items() {
|
||||
let parsed = parse_provider_model_mappings(Some(json!([
|
||||
{"name": "gpt-5.2", "priority": 1},
|
||||
{"priority": 2},
|
||||
3,
|
||||
null,
|
||||
"gpt-5.2-mini"
|
||||
])))
|
||||
.expect("mixed provider_model_mappings should parse");
|
||||
|
||||
assert_eq!(
|
||||
parsed,
|
||||
Some(vec![
|
||||
StoredProviderModelMapping {
|
||||
name: "gpt-5.2".to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
},
|
||||
StoredProviderModelMapping {
|
||||
name: "gpt-5.2-mini".to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
}
|
||||
])
|
||||
);
|
||||
}
|
||||
}
|
||||
144
crates/aether-data/src/repository/candidate_selection/types.rs
Normal file
144
crates/aether-data/src/repository/candidate_selection/types.rs
Normal file
@@ -0,0 +1,144 @@
|
||||
use async_trait::async_trait;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredProviderModelMapping {
|
||||
pub name: String,
|
||||
pub priority: i32,
|
||||
pub api_formats: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct StoredMinimalCandidateSelectionRow {
|
||||
pub provider_id: String,
|
||||
pub provider_name: String,
|
||||
pub provider_type: String,
|
||||
pub provider_priority: i32,
|
||||
pub provider_is_active: bool,
|
||||
pub endpoint_id: String,
|
||||
pub endpoint_api_format: String,
|
||||
pub endpoint_api_family: Option<String>,
|
||||
pub endpoint_kind: Option<String>,
|
||||
pub endpoint_is_active: bool,
|
||||
pub key_id: String,
|
||||
pub key_name: String,
|
||||
pub key_auth_type: String,
|
||||
pub key_is_active: bool,
|
||||
pub key_api_formats: Option<Vec<String>>,
|
||||
pub key_allowed_models: Option<Vec<String>>,
|
||||
pub key_capabilities: Option<serde_json::Value>,
|
||||
pub key_internal_priority: i32,
|
||||
pub key_global_priority_by_format: Option<serde_json::Value>,
|
||||
pub model_id: String,
|
||||
pub global_model_id: String,
|
||||
pub global_model_name: String,
|
||||
pub global_model_mappings: Option<Vec<String>>,
|
||||
pub global_model_supports_streaming: Option<bool>,
|
||||
pub model_provider_model_name: String,
|
||||
pub model_provider_model_mappings: Option<Vec<StoredProviderModelMapping>>,
|
||||
pub model_supports_streaming: Option<bool>,
|
||||
pub model_is_active: bool,
|
||||
pub model_is_available: bool,
|
||||
}
|
||||
|
||||
impl StoredMinimalCandidateSelectionRow {
|
||||
pub fn supports_streaming(&self) -> bool {
|
||||
self.model_supports_streaming
|
||||
.or(self.global_model_supports_streaming)
|
||||
.unwrap_or(true)
|
||||
}
|
||||
|
||||
pub fn key_supports_api_format(&self, api_format: &str) -> bool {
|
||||
let target = api_format.trim();
|
||||
match self.key_api_formats.as_deref() {
|
||||
None => true,
|
||||
Some(formats) => formats
|
||||
.iter()
|
||||
.any(|value| value.eq_ignore_ascii_case(target)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
pub trait MinimalCandidateSelectionReadRepository: Send + Sync {
|
||||
async fn list_for_exact_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
|
||||
|
||||
async fn list_for_exact_api_format_and_global_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, crate::DataLayerError>;
|
||||
}
|
||||
|
||||
pub trait MinimalCandidateSelectionRepository:
|
||||
MinimalCandidateSelectionReadRepository + Send + Sync
|
||||
{
|
||||
}
|
||||
|
||||
impl<T> MinimalCandidateSelectionRepository for T where
|
||||
T: MinimalCandidateSelectionReadRepository + Send + Sync
|
||||
{
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{StoredMinimalCandidateSelectionRow, StoredProviderModelMapping};
|
||||
|
||||
fn sample_row() -> StoredMinimalCandidateSelectionRow {
|
||||
StoredMinimalCandidateSelectionRow {
|
||||
provider_id: "provider-1".to_string(),
|
||||
provider_name: "OpenAI".to_string(),
|
||||
provider_type: "custom".to_string(),
|
||||
provider_priority: 10,
|
||||
provider_is_active: true,
|
||||
endpoint_id: "endpoint-1".to_string(),
|
||||
endpoint_api_format: "openai:chat".to_string(),
|
||||
endpoint_api_family: Some("openai".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
endpoint_is_active: true,
|
||||
key_id: "key-1".to_string(),
|
||||
key_name: "prod".to_string(),
|
||||
key_auth_type: "api_key".to_string(),
|
||||
key_is_active: true,
|
||||
key_api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
key_allowed_models: None,
|
||||
key_capabilities: None,
|
||||
key_internal_priority: 50,
|
||||
key_global_priority_by_format: None,
|
||||
model_id: "model-1".to_string(),
|
||||
global_model_id: "global-model-1".to_string(),
|
||||
global_model_name: "gpt-4.1".to_string(),
|
||||
global_model_mappings: Some(vec!["gpt-4\\.1-.*".to_string()]),
|
||||
global_model_supports_streaming: Some(true),
|
||||
model_provider_model_name: "gpt-4.1-upstream".to_string(),
|
||||
model_provider_model_mappings: Some(vec![StoredProviderModelMapping {
|
||||
name: "gpt-4.1-canary".to_string(),
|
||||
priority: 1,
|
||||
api_formats: Some(vec!["openai:chat".to_string()]),
|
||||
}]),
|
||||
model_supports_streaming: None,
|
||||
model_is_active: true,
|
||||
model_is_available: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn defaults_streaming_support_to_true() {
|
||||
let mut row = sample_row();
|
||||
row.model_supports_streaming = None;
|
||||
row.global_model_supports_streaming = None;
|
||||
|
||||
assert!(row.supports_streaming());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn key_api_formats_none_means_support_all_formats() {
|
||||
let mut row = sample_row();
|
||||
row.key_api_formats = None;
|
||||
|
||||
assert!(row.key_supports_api_format("openai:chat"));
|
||||
assert!(row.key_supports_api_format("openai:responses"));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user