mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
Add multi-database data layer
Introduce aether-data-schema and driver-specific schema generation for Postgres, MySQL, and SQLite. Split data backends, lifecycle, repositories, and gateway runtime integration across database drivers. Verified with cargo fmt --all --check, cargo clippy --workspace --all-targets -- -D warnings, and cargo test --workspace.
This commit is contained in:
@@ -1,5 +1,7 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod mysql;
|
||||
mod postgres;
|
||||
mod sqlite;
|
||||
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::candidate_selection::{
|
||||
@@ -8,4 +10,6 @@ pub(crate) use aether_data_contracts::repository::candidate_selection::{
|
||||
StoredProviderModelMapping, StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
pub use memory::InMemoryMinimalCandidateSelectionReadRepository;
|
||||
pub use sql::SqlxMinimalCandidateSelectionReadRepository;
|
||||
pub use mysql::MysqlMinimalCandidateSelectionReadRepository;
|
||||
pub use postgres::SqlxMinimalCandidateSelectionReadRepository;
|
||||
pub use sqlite::SqliteMinimalCandidateSelectionReadRepository;
|
||||
|
||||
618
crates/aether-data/src/repository/candidate_selection/mysql.rs
Normal file
618
crates/aether-data/src/repository/candidate_selection/mysql.rs
Normal file
@@ -0,0 +1,618 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use super::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use crate::driver::mysql::MysqlPool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
|
||||
const CANDIDATE_SELECTION_COLUMNS: &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,
|
||||
p.config AS provider_config,
|
||||
pe.id AS endpoint_id,
|
||||
COALESCE(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.auth_config AS key_auth_config,
|
||||
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,
|
||||
gm.config AS global_model_config,
|
||||
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 = 1
|
||||
AND pe.is_active = 1
|
||||
AND pak.is_active = 1
|
||||
AND m.is_active = 1
|
||||
AND m.is_available = 1
|
||||
AND gm.is_active = 1
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlMinimalCandidateSelectionReadRepository {
|
||||
pool: MysqlPool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct CandidateSelectionRow {
|
||||
row: StoredMinimalCandidateSelectionRow,
|
||||
provider_pool_enabled: bool,
|
||||
key_auth_config: Option<String>,
|
||||
}
|
||||
|
||||
impl MysqlMinimalCandidateSelectionReadRepository {
|
||||
pub fn new(pool: MysqlPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn load_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<CandidateSelectionRow>, DataLayerError> {
|
||||
let canonical_api_format = normalize_api_format(api_format);
|
||||
let storage_aliases = api_format_aliases(&canonical_api_format);
|
||||
let match_aliases = sql_match_aliases(&storage_aliases);
|
||||
|
||||
let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_SELECTION_COLUMNS);
|
||||
builder.push(" AND LOWER(pe.api_format) IN (");
|
||||
{
|
||||
let mut separated = builder.separated(", ");
|
||||
for alias in &match_aliases {
|
||||
separated.push_bind(alias);
|
||||
}
|
||||
}
|
||||
builder.push(")");
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
let mut items = rows
|
||||
.iter()
|
||||
.map(map_candidate_selection_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
items.retain(|item| {
|
||||
api_format_matches(&item.row.endpoint_api_format, &canonical_api_format)
|
||||
&& item.row.key_supports_api_format(&canonical_api_format)
|
||||
&& key_auth_channel_matches(item, &canonical_api_format)
|
||||
});
|
||||
Ok(items)
|
||||
}
|
||||
|
||||
async fn selected_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self.load_rows_for_api_format(api_format).await?;
|
||||
Ok(sort_rows(select_pool_rows(rows), true))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionReadRepository for MysqlMinimalCandidateSelectionReadRepository {
|
||||
async fn list_for_exact_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.selected_rows_for_api_format(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> {
|
||||
Ok(sort_rows(
|
||||
self.selected_rows_for_api_format(api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| row.global_model_name == global_model_name)
|
||||
.collect(),
|
||||
false,
|
||||
))
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: api_format.to_string(),
|
||||
requested_model_name: requested_model_name.to_string(),
|
||||
offset: 0,
|
||||
limit: u32::MAX,
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self
|
||||
.selected_rows_for_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row_matches_requested_model(row, &query.requested_model_name, &query.api_format)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
Ok(sort_rows(rows, true)
|
||||
.into_iter()
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self
|
||||
.load_rows_for_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.map(|item| item.row)
|
||||
.filter(|row| {
|
||||
row.provider_id == query.provider_id
|
||||
&& row.endpoint_id == query.endpoint_id
|
||||
&& row.model_id == query.model_id
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let mut rows = sort_pool_key_rows(rows);
|
||||
Ok(rows
|
||||
.drain(..)
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.collect())
|
||||
}
|
||||
}
|
||||
|
||||
fn select_pool_rows(rows: Vec<CandidateSelectionRow>) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
let mut selected = Vec::new();
|
||||
let mut pool_rows =
|
||||
BTreeMap::<(String, String, String), StoredMinimalCandidateSelectionRow>::new();
|
||||
for item in rows {
|
||||
if !item.provider_pool_enabled {
|
||||
selected.push(item.row);
|
||||
continue;
|
||||
}
|
||||
let key = (
|
||||
item.row.provider_id.clone(),
|
||||
item.row.endpoint_id.clone(),
|
||||
item.row.model_id.clone(),
|
||||
);
|
||||
match pool_rows.get(&key) {
|
||||
Some(existing)
|
||||
if (existing.key_internal_priority, existing.key_id.as_str())
|
||||
<= (item.row.key_internal_priority, item.row.key_id.as_str()) => {}
|
||||
_ => {
|
||||
pool_rows.insert(key, item.row);
|
||||
}
|
||||
}
|
||||
}
|
||||
selected.extend(pool_rows.into_values());
|
||||
dedupe_candidate_selection_rows(selected)
|
||||
}
|
||||
|
||||
fn sort_rows(
|
||||
mut rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
include_global_model: bool,
|
||||
) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
rows.sort_by(|left, right| {
|
||||
if include_global_model {
|
||||
let ordering = left.global_model_name.cmp(&right.global_model_name);
|
||||
if !ordering.is_eq() {
|
||||
return ordering;
|
||||
}
|
||||
}
|
||||
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))
|
||||
});
|
||||
rows
|
||||
}
|
||||
|
||||
fn sort_pool_key_rows(
|
||||
mut rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
rows.sort_by(|left, right| {
|
||||
left.key_internal_priority
|
||||
.cmp(&right.key_internal_priority)
|
||||
.then(left.key_id.cmp(&right.key_id))
|
||||
});
|
||||
rows
|
||||
}
|
||||
|
||||
fn row_matches_requested_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row.global_model_name == requested_model_name
|
||||
|| row.model_provider_model_name == requested_model_name
|
||||
|| row
|
||||
.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings.iter().any(|mapping| {
|
||||
mapping.api_formats.as_ref().is_none_or(|formats| {
|
||||
formats
|
||||
.iter()
|
||||
.any(|value| api_format_matches(value, api_format))
|
||||
}) && mapping.name == requested_model_name
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn key_auth_channel_matches(row: &CandidateSelectionRow, api_format: &str) -> bool {
|
||||
let provider_type = row.row.provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = row.row.key_auth_type.trim().to_ascii_lowercase();
|
||||
let api_format = normalize_api_format(api_format);
|
||||
match provider_type.as_str() {
|
||||
"codex" => {
|
||||
auth_type == "oauth"
|
||||
&& matches!(
|
||||
api_format.as_str(),
|
||||
"openai:responses" | "openai:responses:compact" | "openai:image"
|
||||
)
|
||||
}
|
||||
"claude_code" => auth_type == "oauth" && api_format == "claude:messages",
|
||||
"kiro" => {
|
||||
api_format == "claude:messages"
|
||||
&& (auth_type == "oauth"
|
||||
|| (auth_type == "bearer"
|
||||
&& row
|
||||
.key_auth_config
|
||||
.as_deref()
|
||||
.is_some_and(|value| !value.trim().is_empty())))
|
||||
}
|
||||
"gemini_cli" | "antigravity" => {
|
||||
auth_type == "oauth" && api_format == "gemini:generate_content"
|
||||
}
|
||||
"vertex_ai" => {
|
||||
(auth_type == "api_key" && api_format == "gemini:generate_content")
|
||||
|| (matches!(auth_type.as_str(), "service_account" | "vertex_ai")
|
||||
&& matches!(
|
||||
api_format.as_str(),
|
||||
"claude:messages" | "gemini:generate_content"
|
||||
))
|
||||
}
|
||||
_ => auth_type != "oauth",
|
||||
}
|
||||
}
|
||||
|
||||
fn dedupe_candidate_selection_rows(
|
||||
rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
let mut seen = BTreeSet::new();
|
||||
rows.into_iter()
|
||||
.filter(|row| {
|
||||
seen.insert((
|
||||
row.endpoint_id.clone(),
|
||||
row.key_id.clone(),
|
||||
row.model_id.clone(),
|
||||
))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn map_candidate_selection_row(row: &MySqlRow) -> Result<CandidateSelectionRow, DataLayerError> {
|
||||
let provider_config = parse_json(row.try_get("provider_config").ok().flatten())?;
|
||||
let global_model_config = parse_json(row.try_get("global_model_config").ok().flatten())?;
|
||||
let provider_pool_enabled = json_object_field_present(&provider_config, "pool_advanced");
|
||||
let global_model_mappings = global_model_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("model_mappings").cloned());
|
||||
let global_model_supports_streaming = global_model_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("streaming"))
|
||||
.and_then(json_bool);
|
||||
Ok(CandidateSelectionRow {
|
||||
row: StoredMinimalCandidateSelectionRow {
|
||||
provider_id: row.try_get("provider_id").map_sql_err()?,
|
||||
provider_name: row.try_get("provider_name").map_sql_err()?,
|
||||
provider_type: row.try_get("provider_type").map_sql_err()?,
|
||||
provider_priority: row.try_get("provider_priority").map_sql_err()?,
|
||||
provider_is_active: row.try_get("provider_is_active").map_sql_err()?,
|
||||
endpoint_id: row.try_get("endpoint_id").map_sql_err()?,
|
||||
endpoint_api_format: row.try_get("endpoint_api_format").map_sql_err()?,
|
||||
endpoint_api_family: row.try_get("endpoint_api_family").map_sql_err()?,
|
||||
endpoint_kind: row.try_get("endpoint_kind").map_sql_err()?,
|
||||
endpoint_is_active: row.try_get("endpoint_is_active").map_sql_err()?,
|
||||
key_id: row.try_get("key_id").map_sql_err()?,
|
||||
key_name: row.try_get("key_name").map_sql_err()?,
|
||||
key_auth_type: row.try_get("key_auth_type").map_sql_err()?,
|
||||
key_is_active: row.try_get("key_is_active").map_sql_err()?,
|
||||
key_api_formats: parse_string_list(
|
||||
parse_json(row.try_get("key_api_formats").ok().flatten())?,
|
||||
"provider_api_keys.api_formats",
|
||||
)?,
|
||||
key_allowed_models: parse_string_list(
|
||||
parse_json(row.try_get("key_allowed_models").ok().flatten())?,
|
||||
"provider_api_keys.allowed_models",
|
||||
)?,
|
||||
key_capabilities: parse_json(row.try_get("key_capabilities").ok().flatten())?,
|
||||
key_internal_priority: row.try_get("key_internal_priority").map_sql_err()?,
|
||||
key_global_priority_by_format: parse_json(
|
||||
row.try_get("key_global_priority_by_format").ok().flatten(),
|
||||
)?,
|
||||
model_id: row.try_get("model_id").map_sql_err()?,
|
||||
global_model_id: row.try_get("global_model_id").map_sql_err()?,
|
||||
global_model_name: row.try_get("global_model_name").map_sql_err()?,
|
||||
global_model_mappings: parse_string_list(
|
||||
global_model_mappings,
|
||||
"global_models.config.model_mappings",
|
||||
)?,
|
||||
global_model_supports_streaming,
|
||||
model_provider_model_name: row.try_get("model_provider_model_name").map_sql_err()?,
|
||||
model_provider_model_mappings: parse_provider_model_mappings(parse_json(
|
||||
row.try_get("model_provider_model_mappings").ok().flatten(),
|
||||
)?)?,
|
||||
model_supports_streaming: row.try_get("model_supports_streaming").map_sql_err()?,
|
||||
model_is_active: row.try_get("model_is_active").map_sql_err()?,
|
||||
model_is_available: row.try_get("model_is_available").map_sql_err()?,
|
||||
},
|
||||
provider_pool_enabled,
|
||||
key_auth_config: row.try_get("key_auth_config").map_sql_err()?,
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_json(value: Option<String>) -> Result<Option<serde_json::Value>, DataLayerError> {
|
||||
value
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.map(|value| {
|
||||
serde_json::from_str(&value).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"candidate selection JSON field is invalid: {err}"
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn json_object_field_present(value: &Option<serde_json::Value>, field: &str) -> bool {
|
||||
value
|
||||
.as_ref()
|
||||
.and_then(|value| value.get(field))
|
||||
.is_some_and(|value| !value.is_null())
|
||||
}
|
||||
|
||||
fn json_bool(value: &serde_json::Value) -> Option<bool> {
|
||||
value.as_bool().or_else(|| {
|
||||
value
|
||||
.as_str()
|
||||
.and_then(|value| value.trim().parse::<bool>().ok())
|
||||
})
|
||||
}
|
||||
|
||||
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);
|
||||
};
|
||||
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_lenient(&object)
|
||||
.map(|mapping| mapping.map(|value| vec![value])),
|
||||
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(Some(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) if !raw.trim().is_empty() => {
|
||||
mappings.push(StoredProviderModelMapping {
|
||||
name: raw.trim().to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
});
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if mappings.is_empty() {
|
||||
Ok(None)
|
||||
} else {
|
||||
Ok(Some(mappings))
|
||||
}
|
||||
}
|
||||
|
||||
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",
|
||||
)?
|
||||
.map(|formats| {
|
||||
formats
|
||||
.into_iter()
|
||||
.map(|value| normalize_api_format(&value))
|
||||
.collect()
|
||||
});
|
||||
|
||||
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,
|
||||
}))
|
||||
}
|
||||
|
||||
fn api_format_aliases(api_format: &str) -> Vec<String> {
|
||||
aether_ai_formats::api_format_storage_aliases(api_format)
|
||||
}
|
||||
|
||||
fn normalize_api_format(api_format: &str) -> String {
|
||||
aether_ai_formats::normalize_api_format_alias(api_format)
|
||||
}
|
||||
|
||||
fn api_format_matches(left: &str, right: &str) -> bool {
|
||||
aether_ai_formats::api_format_alias_matches(left, right)
|
||||
}
|
||||
|
||||
fn sql_match_aliases(api_formats: &[String]) -> Vec<String> {
|
||||
api_formats
|
||||
.iter()
|
||||
.map(|value| value.trim().to_ascii_lowercase())
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::MysqlMinimalCandidateSelectionReadRepository;
|
||||
|
||||
#[tokio::test]
|
||||
async fn repository_builds_from_lazy_pool() {
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new().connect_lazy_with(
|
||||
"mysql://user:pass@localhost:3306/aether"
|
||||
.parse()
|
||||
.expect("mysql options should parse"),
|
||||
);
|
||||
|
||||
let _repository = MysqlMinimalCandidateSelectionReadRepository::new(pool);
|
||||
}
|
||||
}
|
||||
@@ -1008,7 +1008,7 @@ mod tests {
|
||||
parse_provider_model_mappings, parse_string_list, requested_model_selection_page_sql,
|
||||
requested_model_selection_sql, SqlxMinimalCandidateSelectionReadRepository,
|
||||
};
|
||||
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
use crate::repository::candidate_selection::StoredProviderModelMapping;
|
||||
|
||||
#[tokio::test]
|
||||
711
crates/aether-data/src/repository/candidate_selection/sqlite.rs
Normal file
711
crates/aether-data/src/repository/candidate_selection/sqlite.rs
Normal file
@@ -0,0 +1,711 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
|
||||
|
||||
use super::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
|
||||
const CANDIDATE_SELECTION_COLUMNS: &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,
|
||||
p.config AS provider_config,
|
||||
pe.id AS endpoint_id,
|
||||
COALESCE(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.auth_config AS key_auth_config,
|
||||
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,
|
||||
gm.config AS global_model_config,
|
||||
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 = 1
|
||||
AND pe.is_active = 1
|
||||
AND pak.is_active = 1
|
||||
AND m.is_active = 1
|
||||
AND m.is_available = 1
|
||||
AND gm.is_active = 1
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqliteMinimalCandidateSelectionReadRepository {
|
||||
pool: SqlitePool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
struct CandidateSelectionRow {
|
||||
row: StoredMinimalCandidateSelectionRow,
|
||||
provider_pool_enabled: bool,
|
||||
key_auth_config: Option<String>,
|
||||
}
|
||||
|
||||
impl SqliteMinimalCandidateSelectionReadRepository {
|
||||
pub fn new(pool: SqlitePool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn load_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<CandidateSelectionRow>, DataLayerError> {
|
||||
let canonical_api_format = normalize_api_format(api_format);
|
||||
let storage_aliases = api_format_aliases(&canonical_api_format);
|
||||
let match_aliases = sql_match_aliases(&storage_aliases);
|
||||
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_SELECTION_COLUMNS);
|
||||
builder.push(" AND LOWER(pe.api_format) IN (");
|
||||
{
|
||||
let mut separated = builder.separated(", ");
|
||||
for alias in &match_aliases {
|
||||
separated.push_bind(alias);
|
||||
}
|
||||
}
|
||||
builder.push(")");
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
let mut items = rows
|
||||
.iter()
|
||||
.map(map_candidate_selection_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
items.retain(|item| {
|
||||
api_format_matches(&item.row.endpoint_api_format, &canonical_api_format)
|
||||
&& item.row.key_supports_api_format(&canonical_api_format)
|
||||
&& key_auth_channel_matches(item, &canonical_api_format)
|
||||
});
|
||||
Ok(items)
|
||||
}
|
||||
|
||||
async fn selected_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self.load_rows_for_api_format(api_format).await?;
|
||||
Ok(sort_rows(select_pool_rows(rows), true))
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelectionReadRepository {
|
||||
async fn list_for_exact_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.selected_rows_for_api_format(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> {
|
||||
Ok(sort_rows(
|
||||
self.selected_rows_for_api_format(api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| row.global_model_name == global_model_name)
|
||||
.collect(),
|
||||
false,
|
||||
))
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
&self,
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: api_format.to_string(),
|
||||
requested_model_name: requested_model_name.to_string(),
|
||||
offset: 0,
|
||||
limit: u32::MAX,
|
||||
},
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model_page(
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self
|
||||
.selected_rows_for_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row_matches_requested_model(row, &query.requested_model_name, &query.api_format)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
Ok(sort_rows(rows, true)
|
||||
.into_iter()
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self
|
||||
.load_rows_for_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.map(|item| item.row)
|
||||
.filter(|row| {
|
||||
row.provider_id == query.provider_id
|
||||
&& row.endpoint_id == query.endpoint_id
|
||||
&& row.model_id == query.model_id
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let mut rows = sort_pool_key_rows(rows);
|
||||
Ok(rows
|
||||
.drain(..)
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.collect())
|
||||
}
|
||||
}
|
||||
|
||||
fn select_pool_rows(rows: Vec<CandidateSelectionRow>) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
let mut selected = Vec::new();
|
||||
let mut pool_rows =
|
||||
BTreeMap::<(String, String, String), StoredMinimalCandidateSelectionRow>::new();
|
||||
for item in rows {
|
||||
if !item.provider_pool_enabled {
|
||||
selected.push(item.row);
|
||||
continue;
|
||||
}
|
||||
let key = (
|
||||
item.row.provider_id.clone(),
|
||||
item.row.endpoint_id.clone(),
|
||||
item.row.model_id.clone(),
|
||||
);
|
||||
match pool_rows.get(&key) {
|
||||
Some(existing)
|
||||
if (existing.key_internal_priority, existing.key_id.as_str())
|
||||
<= (item.row.key_internal_priority, item.row.key_id.as_str()) => {}
|
||||
_ => {
|
||||
pool_rows.insert(key, item.row);
|
||||
}
|
||||
}
|
||||
}
|
||||
selected.extend(pool_rows.into_values());
|
||||
dedupe_candidate_selection_rows(selected)
|
||||
}
|
||||
|
||||
fn sort_rows(
|
||||
mut rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
include_global_model: bool,
|
||||
) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
rows.sort_by(|left, right| {
|
||||
if include_global_model {
|
||||
let ordering = left.global_model_name.cmp(&right.global_model_name);
|
||||
if !ordering.is_eq() {
|
||||
return ordering;
|
||||
}
|
||||
}
|
||||
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))
|
||||
});
|
||||
rows
|
||||
}
|
||||
|
||||
fn sort_pool_key_rows(
|
||||
mut rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
rows.sort_by(|left, right| {
|
||||
left.key_internal_priority
|
||||
.cmp(&right.key_internal_priority)
|
||||
.then(left.key_id.cmp(&right.key_id))
|
||||
});
|
||||
rows
|
||||
}
|
||||
|
||||
fn row_matches_requested_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row.global_model_name == requested_model_name
|
||||
|| row.model_provider_model_name == requested_model_name
|
||||
|| row
|
||||
.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings.iter().any(|mapping| {
|
||||
mapping.api_formats.as_ref().is_none_or(|formats| {
|
||||
formats
|
||||
.iter()
|
||||
.any(|value| api_format_matches(value, api_format))
|
||||
}) && mapping.name == requested_model_name
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn key_auth_channel_matches(row: &CandidateSelectionRow, api_format: &str) -> bool {
|
||||
let provider_type = row.row.provider_type.trim().to_ascii_lowercase();
|
||||
let auth_type = row.row.key_auth_type.trim().to_ascii_lowercase();
|
||||
let api_format = normalize_api_format(api_format);
|
||||
match provider_type.as_str() {
|
||||
"codex" => {
|
||||
auth_type == "oauth"
|
||||
&& matches!(
|
||||
api_format.as_str(),
|
||||
"openai:responses" | "openai:responses:compact" | "openai:image"
|
||||
)
|
||||
}
|
||||
"claude_code" => auth_type == "oauth" && api_format == "claude:messages",
|
||||
"kiro" => {
|
||||
api_format == "claude:messages"
|
||||
&& (auth_type == "oauth"
|
||||
|| (auth_type == "bearer"
|
||||
&& row
|
||||
.key_auth_config
|
||||
.as_deref()
|
||||
.is_some_and(|value| !value.trim().is_empty())))
|
||||
}
|
||||
"gemini_cli" | "antigravity" => {
|
||||
auth_type == "oauth" && api_format == "gemini:generate_content"
|
||||
}
|
||||
"vertex_ai" => {
|
||||
(auth_type == "api_key" && api_format == "gemini:generate_content")
|
||||
|| (matches!(auth_type.as_str(), "service_account" | "vertex_ai")
|
||||
&& matches!(
|
||||
api_format.as_str(),
|
||||
"claude:messages" | "gemini:generate_content"
|
||||
))
|
||||
}
|
||||
_ => auth_type != "oauth",
|
||||
}
|
||||
}
|
||||
|
||||
fn dedupe_candidate_selection_rows(
|
||||
rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
let mut seen = BTreeSet::new();
|
||||
rows.into_iter()
|
||||
.filter(|row| {
|
||||
seen.insert((
|
||||
row.endpoint_id.clone(),
|
||||
row.key_id.clone(),
|
||||
row.model_id.clone(),
|
||||
))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn map_candidate_selection_row(row: &SqliteRow) -> Result<CandidateSelectionRow, DataLayerError> {
|
||||
let provider_config = parse_json(row.try_get("provider_config").ok().flatten())?;
|
||||
let global_model_config = parse_json(row.try_get("global_model_config").ok().flatten())?;
|
||||
let provider_pool_enabled = json_object_field_present(&provider_config, "pool_advanced");
|
||||
let global_model_mappings = global_model_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("model_mappings").cloned());
|
||||
let global_model_supports_streaming = global_model_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("streaming"))
|
||||
.and_then(json_bool);
|
||||
Ok(CandidateSelectionRow {
|
||||
row: StoredMinimalCandidateSelectionRow {
|
||||
provider_id: row.try_get("provider_id").map_sql_err()?,
|
||||
provider_name: row.try_get("provider_name").map_sql_err()?,
|
||||
provider_type: row.try_get("provider_type").map_sql_err()?,
|
||||
provider_priority: row.try_get("provider_priority").map_sql_err()?,
|
||||
provider_is_active: row.try_get("provider_is_active").map_sql_err()?,
|
||||
endpoint_id: row.try_get("endpoint_id").map_sql_err()?,
|
||||
endpoint_api_format: row.try_get("endpoint_api_format").map_sql_err()?,
|
||||
endpoint_api_family: row.try_get("endpoint_api_family").map_sql_err()?,
|
||||
endpoint_kind: row.try_get("endpoint_kind").map_sql_err()?,
|
||||
endpoint_is_active: row.try_get("endpoint_is_active").map_sql_err()?,
|
||||
key_id: row.try_get("key_id").map_sql_err()?,
|
||||
key_name: row.try_get("key_name").map_sql_err()?,
|
||||
key_auth_type: row.try_get("key_auth_type").map_sql_err()?,
|
||||
key_is_active: row.try_get("key_is_active").map_sql_err()?,
|
||||
key_api_formats: parse_string_list(
|
||||
parse_json(row.try_get("key_api_formats").ok().flatten())?,
|
||||
"provider_api_keys.api_formats",
|
||||
)?,
|
||||
key_allowed_models: parse_string_list(
|
||||
parse_json(row.try_get("key_allowed_models").ok().flatten())?,
|
||||
"provider_api_keys.allowed_models",
|
||||
)?,
|
||||
key_capabilities: parse_json(row.try_get("key_capabilities").ok().flatten())?,
|
||||
key_internal_priority: row.try_get("key_internal_priority").map_sql_err()?,
|
||||
key_global_priority_by_format: parse_json(
|
||||
row.try_get("key_global_priority_by_format").ok().flatten(),
|
||||
)?,
|
||||
model_id: row.try_get("model_id").map_sql_err()?,
|
||||
global_model_id: row.try_get("global_model_id").map_sql_err()?,
|
||||
global_model_name: row.try_get("global_model_name").map_sql_err()?,
|
||||
global_model_mappings: parse_string_list(
|
||||
global_model_mappings,
|
||||
"global_models.config.model_mappings",
|
||||
)?,
|
||||
global_model_supports_streaming,
|
||||
model_provider_model_name: row.try_get("model_provider_model_name").map_sql_err()?,
|
||||
model_provider_model_mappings: parse_provider_model_mappings(parse_json(
|
||||
row.try_get("model_provider_model_mappings").ok().flatten(),
|
||||
)?)?,
|
||||
model_supports_streaming: row.try_get("model_supports_streaming").map_sql_err()?,
|
||||
model_is_active: row.try_get("model_is_active").map_sql_err()?,
|
||||
model_is_available: row.try_get("model_is_available").map_sql_err()?,
|
||||
},
|
||||
provider_pool_enabled,
|
||||
key_auth_config: row.try_get("key_auth_config").map_sql_err()?,
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_json(value: Option<String>) -> Result<Option<serde_json::Value>, DataLayerError> {
|
||||
value
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
.map(|value| {
|
||||
serde_json::from_str(&value).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"candidate selection JSON field is invalid: {err}"
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn json_object_field_present(value: &Option<serde_json::Value>, field: &str) -> bool {
|
||||
value
|
||||
.as_ref()
|
||||
.and_then(|value| value.get(field))
|
||||
.is_some_and(|value| !value.is_null())
|
||||
}
|
||||
|
||||
fn json_bool(value: &serde_json::Value) -> Option<bool> {
|
||||
value.as_bool().or_else(|| {
|
||||
value
|
||||
.as_str()
|
||||
.and_then(|value| value.trim().parse::<bool>().ok())
|
||||
})
|
||||
}
|
||||
|
||||
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);
|
||||
};
|
||||
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_lenient(&object)
|
||||
.map(|mapping| mapping.map(|value| vec![value])),
|
||||
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(Some(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) if !raw.trim().is_empty() => {
|
||||
mappings.push(StoredProviderModelMapping {
|
||||
name: raw.trim().to_string(),
|
||||
priority: 1,
|
||||
api_formats: None,
|
||||
});
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
if mappings.is_empty() {
|
||||
Ok(None)
|
||||
} else {
|
||||
Ok(Some(mappings))
|
||||
}
|
||||
}
|
||||
|
||||
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",
|
||||
)?
|
||||
.map(|formats| {
|
||||
formats
|
||||
.into_iter()
|
||||
.map(|value| normalize_api_format(&value))
|
||||
.collect()
|
||||
});
|
||||
|
||||
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,
|
||||
}))
|
||||
}
|
||||
|
||||
fn api_format_aliases(api_format: &str) -> Vec<String> {
|
||||
aether_ai_formats::api_format_storage_aliases(api_format)
|
||||
}
|
||||
|
||||
fn normalize_api_format(api_format: &str) -> String {
|
||||
aether_ai_formats::normalize_api_format_alias(api_format)
|
||||
}
|
||||
|
||||
fn api_format_matches(left: &str, right: &str) -> bool {
|
||||
aether_ai_formats::api_format_alias_matches(left, right)
|
||||
}
|
||||
|
||||
fn sql_match_aliases(api_formats: &[String]) -> Vec<String> {
|
||||
api_formats
|
||||
.iter()
|
||||
.map(|value| value.trim().to_ascii_lowercase())
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::SqliteMinimalCandidateSelectionReadRepository;
|
||||
use crate::lifecycle::migrate::run_sqlite_migrations;
|
||||
use crate::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, StoredPoolKeyCandidateRowsQuery,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_repository_reads_candidate_selection_rows() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
seed_candidate_selection(&pool).await;
|
||||
|
||||
let repository = SqliteMinimalCandidateSelectionReadRepository::new(pool);
|
||||
let rows = repository
|
||||
.list_for_exact_api_format("openai:chat")
|
||||
.await
|
||||
.expect("candidate rows should load");
|
||||
assert_eq!(
|
||||
rows.iter()
|
||||
.map(|row| row.key_id.as_str())
|
||||
.collect::<Vec<_>>(),
|
||||
vec!["key-1"]
|
||||
);
|
||||
assert_eq!(
|
||||
rows[0].global_model_mappings,
|
||||
Some(vec!["alias-global".to_string()])
|
||||
);
|
||||
assert_eq!(rows[0].global_model_supports_streaming, Some(true));
|
||||
|
||||
let requested = repository
|
||||
.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
requested_model_name: "alias-provider".to_string(),
|
||||
offset: 0,
|
||||
limit: 10,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.expect("requested model rows should load");
|
||||
assert_eq!(requested.len(), 1);
|
||||
|
||||
let pool_keys = repository
|
||||
.list_pool_key_rows_for_group(&StoredPoolKeyCandidateRowsQuery {
|
||||
api_format: "openai:chat".to_string(),
|
||||
provider_id: "provider-1".to_string(),
|
||||
endpoint_id: "endpoint-1".to_string(),
|
||||
model_id: "model-1".to_string(),
|
||||
selected_provider_model_name: "provider-model".to_string(),
|
||||
offset: 1,
|
||||
limit: 1,
|
||||
})
|
||||
.await
|
||||
.expect("pool keys should load");
|
||||
assert_eq!(pool_keys.len(), 1);
|
||||
assert_eq!(pool_keys[0].key_id, "key-2");
|
||||
}
|
||||
|
||||
async fn seed_candidate_selection(pool: &sqlx::SqlitePool) {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO providers (
|
||||
id, name, provider_type, provider_priority, config, is_active, created_at, updated_at
|
||||
)
|
||||
VALUES ('provider-1', 'Provider One', 'custom', 10, '{"pool_advanced":{}}', 1, 1, 1);
|
||||
|
||||
INSERT INTO provider_endpoints (
|
||||
id, provider_id, name, base_url, api_format, is_active, created_at, updated_at
|
||||
)
|
||||
VALUES ('endpoint-1', 'provider-1', 'Endpoint One', 'https://example.test', 'openai:chat', 1, 1, 1);
|
||||
|
||||
INSERT INTO provider_api_keys (
|
||||
id, provider_id, name, auth_type, api_formats, internal_priority, is_active, created_at, updated_at
|
||||
)
|
||||
VALUES
|
||||
('key-1', 'provider-1', 'Key One', 'api_key', '["openai:chat"]', 10, 1, 1, 1),
|
||||
('key-2', 'provider-1', 'Key Two', 'api_key', '["openai:chat"]', 20, 1, 1, 1);
|
||||
|
||||
INSERT INTO global_models (
|
||||
id, name, config, is_active, created_at, updated_at
|
||||
)
|
||||
VALUES ('global-1', 'gpt-5', '{"model_mappings":["alias-global"],"streaming":true}', 1, 1, 1);
|
||||
|
||||
INSERT INTO models (
|
||||
id, provider_id, global_model_id, provider_model_name, provider_model_mappings,
|
||||
supports_streaming, is_active, is_available, created_at, updated_at
|
||||
)
|
||||
VALUES (
|
||||
'model-1', 'provider-1', 'global-1', 'provider-model',
|
||||
'[{"name":"alias-provider","api_formats":["openai:chat"],"priority":1}]',
|
||||
1, 1, 1, 1, 1
|
||||
);
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("candidate selection rows should seed");
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user