mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-09 10:57:03 +08:00
chore: resolve pr 377 checks
This commit is contained in:
@@ -1,5 +1,9 @@
|
||||
mod memory;
|
||||
mod sql;
|
||||
mod mysql;
|
||||
mod postgres;
|
||||
mod sqlite;
|
||||
|
||||
use serde_json::Value;
|
||||
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::global_models::{
|
||||
@@ -11,4 +15,98 @@ pub(crate) use aether_data_contracts::repository::global_models::{
|
||||
StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
};
|
||||
pub use memory::InMemoryGlobalModelReadRepository;
|
||||
pub use sql::SqlxGlobalModelReadRepository;
|
||||
pub use mysql::MysqlGlobalModelReadRepository;
|
||||
pub use postgres::SqlxGlobalModelReadRepository;
|
||||
pub use sqlite::SqliteGlobalModelReadRepository;
|
||||
|
||||
const EMBEDDING_CAPABILITY: &str = "embedding";
|
||||
const EMBEDDING_API_FORMATS: &[&str] = &[
|
||||
"openai:embedding",
|
||||
"gemini:embedding",
|
||||
"jina:embedding",
|
||||
"doubao:embedding",
|
||||
"/v1/embeddings",
|
||||
"/jina/v1/embeddings",
|
||||
];
|
||||
|
||||
pub(super) fn metadata_supports_embedding(
|
||||
supported_capabilities: Option<&Value>,
|
||||
global_config: Option<&Value>,
|
||||
model_config: Option<&Value>,
|
||||
) -> Option<bool> {
|
||||
Some(
|
||||
supported_capabilities.is_some_and(value_contains_embedding_capability)
|
||||
|| global_config.is_some_and(value_contains_embedding_metadata)
|
||||
|| model_config.is_some_and(value_contains_embedding_metadata),
|
||||
)
|
||||
}
|
||||
|
||||
fn value_contains_embedding_capability(value: &Value) -> bool {
|
||||
match value {
|
||||
Value::String(value) => value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY),
|
||||
Value::Array(values) => values.iter().any(value_contains_embedding_capability),
|
||||
Value::Object(object) => {
|
||||
object
|
||||
.get(EMBEDDING_CAPABILITY)
|
||||
.and_then(Value::as_bool)
|
||||
.unwrap_or(false)
|
||||
|| [
|
||||
"capability",
|
||||
"model_type",
|
||||
"type",
|
||||
"task_type",
|
||||
"request_type",
|
||||
]
|
||||
.iter()
|
||||
.any(|key| {
|
||||
object
|
||||
.get(*key)
|
||||
.is_some_and(value_contains_embedding_capability)
|
||||
})
|
||||
|| ["capabilities", "supported_capabilities"]
|
||||
.iter()
|
||||
.any(|key| {
|
||||
object
|
||||
.get(*key)
|
||||
.is_some_and(value_contains_embedding_capability)
|
||||
})
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn value_contains_embedding_metadata(value: &Value) -> bool {
|
||||
match value {
|
||||
Value::String(value) => {
|
||||
value.trim().eq_ignore_ascii_case(EMBEDDING_CAPABILITY)
|
||||
|| is_known_embedding_api_format(value)
|
||||
}
|
||||
Value::Array(values) => values.iter().any(value_contains_embedding_metadata),
|
||||
Value::Object(object) => {
|
||||
value_contains_embedding_capability(value)
|
||||
|| ["api_format", "client_api_format", "provider_api_format"]
|
||||
.iter()
|
||||
.any(|key| {
|
||||
object
|
||||
.get(*key)
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(is_known_embedding_api_format)
|
||||
})
|
||||
|| ["api_formats", "client_api_formats", "provider_api_formats"]
|
||||
.iter()
|
||||
.any(|key| {
|
||||
object
|
||||
.get(*key)
|
||||
.is_some_and(value_contains_embedding_metadata)
|
||||
})
|
||||
}
|
||||
_ => false,
|
||||
}
|
||||
}
|
||||
|
||||
fn is_known_embedding_api_format(value: &str) -> bool {
|
||||
let normalized = value.trim().to_ascii_lowercase();
|
||||
EMBEDDING_API_FORMATS
|
||||
.iter()
|
||||
.any(|format| normalized == *format || normalized.ends_with(*format))
|
||||
}
|
||||
|
||||
@@ -0,0 +1,897 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
|
||||
use super::{
|
||||
metadata_supports_embedding, AdminGlobalModelListQuery, AdminProviderModelListQuery,
|
||||
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelWriteRepository,
|
||||
InMemoryGlobalModelReadRepository, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
|
||||
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
|
||||
StoredAdminProviderModel, StoredProviderActiveGlobalModel, StoredProviderModelStats,
|
||||
StoredPublicCatalogModel, StoredPublicGlobalModel, StoredPublicGlobalModelPage,
|
||||
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
};
|
||||
use crate::driver::mysql::MysqlPool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlGlobalModelReadRepository {
|
||||
pool: MysqlPool,
|
||||
}
|
||||
|
||||
impl MysqlGlobalModelReadRepository {
|
||||
pub fn new(pool: MysqlPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn load_memory(&self) -> Result<InMemoryGlobalModelReadRepository, DataLayerError> {
|
||||
Ok(
|
||||
InMemoryGlobalModelReadRepository::seed(self.load_public_global_models().await?)
|
||||
.with_admin_global_models(self.load_admin_global_models().await?)
|
||||
.with_admin_provider_models(self.load_admin_provider_models().await?)
|
||||
.with_public_catalog_models(self.load_public_catalog_models().await?)
|
||||
.with_provider_model_stats(self.load_provider_model_stats().await?)
|
||||
.with_active_global_model_refs(self.load_active_global_model_refs().await?),
|
||||
)
|
||||
}
|
||||
|
||||
async fn load_public_global_models(
|
||||
&self,
|
||||
) -> Result<Vec<StoredPublicGlobalModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT id, name, display_name, is_active, default_price_per_request,
|
||||
default_tiered_pricing, supported_capabilities, config, usage_count
|
||||
FROM global_models
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_public_global_model_row).collect()
|
||||
}
|
||||
|
||||
async fn load_admin_global_models(
|
||||
&self,
|
||||
) -> Result<Vec<StoredAdminGlobalModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id,
|
||||
name,
|
||||
COALESCE(NULLIF(display_name, ''), name) AS display_name,
|
||||
is_active,
|
||||
default_price_per_request,
|
||||
default_tiered_pricing,
|
||||
supported_capabilities,
|
||||
config,
|
||||
usage_count,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM global_models
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_admin_global_model_row).collect()
|
||||
}
|
||||
|
||||
async fn load_admin_provider_models(
|
||||
&self,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
m.id,
|
||||
m.provider_id,
|
||||
m.global_model_id,
|
||||
m.provider_model_name,
|
||||
m.provider_model_mappings,
|
||||
m.price_per_request,
|
||||
m.tiered_pricing,
|
||||
m.supports_vision,
|
||||
m.supports_function_calling,
|
||||
m.supports_streaming,
|
||||
m.supports_extended_thinking,
|
||||
m.supports_image_generation,
|
||||
m.is_active,
|
||||
m.is_available,
|
||||
m.config,
|
||||
m.created_at AS created_at_unix_ms,
|
||||
m.updated_at AS updated_at_unix_secs,
|
||||
gm.name AS global_model_name,
|
||||
gm.display_name AS global_model_display_name,
|
||||
gm.default_price_per_request AS global_model_default_price_per_request,
|
||||
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
|
||||
gm.supported_capabilities AS global_model_supported_capabilities,
|
||||
gm.config AS global_model_config
|
||||
FROM models m
|
||||
LEFT JOIN global_models gm ON gm.id = m.global_model_id
|
||||
WHERE m.global_model_id IS NOT NULL
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_admin_provider_model_row).collect()
|
||||
}
|
||||
|
||||
async fn load_public_catalog_models(
|
||||
&self,
|
||||
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
m.id,
|
||||
m.provider_id,
|
||||
p.name AS provider_name,
|
||||
p.is_active AS provider_is_active,
|
||||
m.provider_model_name,
|
||||
COALESCE(gm.name, m.provider_model_name) AS name,
|
||||
COALESCE(NULLIF(gm.display_name, ''), m.provider_model_name) AS display_name,
|
||||
gm.config AS global_model_config,
|
||||
gm.supported_capabilities AS global_model_supported_capabilities,
|
||||
m.config AS model_config,
|
||||
m.tiered_pricing,
|
||||
gm.default_tiered_pricing,
|
||||
m.supports_vision,
|
||||
m.supports_function_calling,
|
||||
m.supports_streaming,
|
||||
m.is_active,
|
||||
gm.is_active AS global_model_is_active
|
||||
FROM models m
|
||||
JOIN providers p ON p.id = m.provider_id
|
||||
LEFT JOIN global_models gm ON gm.id = m.global_model_id
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_public_catalog_model_row).collect()
|
||||
}
|
||||
|
||||
async fn load_provider_model_stats(
|
||||
&self,
|
||||
) -> Result<Vec<StoredProviderModelStats>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
provider_id,
|
||||
COUNT(id) AS total_models,
|
||||
SUM(CASE WHEN is_active = 1 THEN 1 ELSE 0 END) AS active_models
|
||||
FROM models
|
||||
GROUP BY provider_id
|
||||
ORDER BY provider_id ASC
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_provider_model_stats_row).collect()
|
||||
}
|
||||
|
||||
async fn load_active_global_model_refs(
|
||||
&self,
|
||||
) -> Result<Vec<StoredProviderActiveGlobalModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT DISTINCT provider_id, global_model_id
|
||||
FROM models
|
||||
WHERE is_active = 1
|
||||
AND global_model_id IS NOT NULL
|
||||
ORDER BY provider_id ASC, global_model_id ASC
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_active_global_model_row).collect()
|
||||
}
|
||||
|
||||
pub async fn create_admin_provider_model(
|
||||
&self,
|
||||
record: &UpsertAdminProviderModelRecord,
|
||||
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
|
||||
let now = current_unix_secs();
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO models (
|
||||
id,
|
||||
provider_id,
|
||||
global_model_id,
|
||||
provider_model_name,
|
||||
provider_model_mappings,
|
||||
price_per_request,
|
||||
tiered_pricing,
|
||||
supports_vision,
|
||||
supports_function_calling,
|
||||
supports_streaming,
|
||||
supports_extended_thinking,
|
||||
supports_image_generation,
|
||||
is_active,
|
||||
is_available,
|
||||
config,
|
||||
created_at,
|
||||
updated_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&record.id)
|
||||
.bind(&record.provider_id)
|
||||
.bind(&record.global_model_id)
|
||||
.bind(&record.provider_model_name)
|
||||
.bind(optional_json_to_string(
|
||||
&record.provider_model_mappings,
|
||||
"models.provider_model_mappings",
|
||||
)?)
|
||||
.bind(record.price_per_request)
|
||||
.bind(optional_json_to_string(
|
||||
&record.tiered_pricing,
|
||||
"models.tiered_pricing",
|
||||
)?)
|
||||
.bind(record.supports_vision)
|
||||
.bind(record.supports_function_calling)
|
||||
.bind(record.supports_streaming)
|
||||
.bind(record.supports_extended_thinking)
|
||||
.bind(record.supports_image_generation)
|
||||
.bind(record.is_active)
|
||||
.bind(record.is_available)
|
||||
.bind(optional_json_to_string(&record.config, "models.config")?)
|
||||
.bind(now as i64)
|
||||
.bind(now as i64)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
self.get_admin_provider_model(&record.provider_id, &record.id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn update_admin_provider_model(
|
||||
&self,
|
||||
record: &UpsertAdminProviderModelRecord,
|
||||
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
|
||||
let now = current_unix_secs();
|
||||
let updated = sqlx::query(
|
||||
r#"
|
||||
UPDATE models
|
||||
SET
|
||||
global_model_id = ?,
|
||||
provider_model_name = ?,
|
||||
provider_model_mappings = ?,
|
||||
price_per_request = ?,
|
||||
tiered_pricing = ?,
|
||||
supports_vision = ?,
|
||||
supports_function_calling = ?,
|
||||
supports_streaming = ?,
|
||||
supports_extended_thinking = ?,
|
||||
supports_image_generation = ?,
|
||||
is_active = ?,
|
||||
is_available = ?,
|
||||
config = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
AND provider_id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(&record.global_model_id)
|
||||
.bind(&record.provider_model_name)
|
||||
.bind(optional_json_to_string(
|
||||
&record.provider_model_mappings,
|
||||
"models.provider_model_mappings",
|
||||
)?)
|
||||
.bind(record.price_per_request)
|
||||
.bind(optional_json_to_string(
|
||||
&record.tiered_pricing,
|
||||
"models.tiered_pricing",
|
||||
)?)
|
||||
.bind(record.supports_vision)
|
||||
.bind(record.supports_function_calling)
|
||||
.bind(record.supports_streaming)
|
||||
.bind(record.supports_extended_thinking)
|
||||
.bind(record.supports_image_generation)
|
||||
.bind(record.is_active)
|
||||
.bind(record.is_available)
|
||||
.bind(optional_json_to_string(&record.config, "models.config")?)
|
||||
.bind(now as i64)
|
||||
.bind(&record.id)
|
||||
.bind(&record.provider_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
if updated.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
self.get_admin_provider_model(&record.provider_id, &record.id)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn delete_admin_provider_model(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
model_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let deleted = sqlx::query(
|
||||
r#"
|
||||
DELETE FROM models
|
||||
WHERE provider_id = ?
|
||||
AND id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(provider_id)
|
||||
.bind(model_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
Ok(deleted.rows_affected() > 0)
|
||||
}
|
||||
|
||||
pub async fn create_admin_global_model(
|
||||
&self,
|
||||
record: &CreateAdminGlobalModelRecord,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
let now = current_unix_secs();
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO global_models (
|
||||
id,
|
||||
name,
|
||||
display_name,
|
||||
is_active,
|
||||
default_price_per_request,
|
||||
default_tiered_pricing,
|
||||
supported_capabilities,
|
||||
usage_count,
|
||||
config,
|
||||
created_at,
|
||||
updated_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, 0, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&record.id)
|
||||
.bind(&record.name)
|
||||
.bind(&record.display_name)
|
||||
.bind(record.is_active)
|
||||
.bind(record.default_price_per_request)
|
||||
.bind(optional_json_to_string(
|
||||
&record.default_tiered_pricing,
|
||||
"global_models.default_tiered_pricing",
|
||||
)?)
|
||||
.bind(optional_json_to_string(
|
||||
&record.supported_capabilities,
|
||||
"global_models.supported_capabilities",
|
||||
)?)
|
||||
.bind(optional_json_to_string(
|
||||
&record.config,
|
||||
"global_models.config",
|
||||
)?)
|
||||
.bind(now as i64)
|
||||
.bind(now as i64)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
self.get_admin_global_model_by_id(&record.id).await
|
||||
}
|
||||
|
||||
pub async fn update_admin_global_model(
|
||||
&self,
|
||||
record: &UpdateAdminGlobalModelRecord,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
let now = current_unix_secs();
|
||||
let updated = sqlx::query(
|
||||
r#"
|
||||
UPDATE global_models
|
||||
SET
|
||||
display_name = ?,
|
||||
is_active = ?,
|
||||
default_price_per_request = ?,
|
||||
default_tiered_pricing = ?,
|
||||
supported_capabilities = ?,
|
||||
config = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(&record.display_name)
|
||||
.bind(record.is_active)
|
||||
.bind(record.default_price_per_request)
|
||||
.bind(optional_json_to_string(
|
||||
&record.default_tiered_pricing,
|
||||
"global_models.default_tiered_pricing",
|
||||
)?)
|
||||
.bind(optional_json_to_string(
|
||||
&record.supported_capabilities,
|
||||
"global_models.supported_capabilities",
|
||||
)?)
|
||||
.bind(optional_json_to_string(
|
||||
&record.config,
|
||||
"global_models.config",
|
||||
)?)
|
||||
.bind(now as i64)
|
||||
.bind(&record.id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
if updated.rows_affected() == 0 {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
self.get_admin_global_model_by_id(&record.id).await
|
||||
}
|
||||
|
||||
pub async fn delete_admin_global_model(
|
||||
&self,
|
||||
global_model_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
DELETE FROM models
|
||||
WHERE global_model_id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(global_model_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let deleted = sqlx::query(
|
||||
r#"
|
||||
DELETE FROM global_models
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(global_model_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
tx.commit().await.map_sql_err()?;
|
||||
|
||||
Ok(deleted.rows_affected() > 0)
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl GlobalModelReadRepository for MysqlGlobalModelReadRepository {
|
||||
async fn list_public_models(
|
||||
&self,
|
||||
query: &PublicGlobalModelQuery,
|
||||
) -> Result<StoredPublicGlobalModelPage, DataLayerError> {
|
||||
self.load_memory().await?.list_public_models(query).await
|
||||
}
|
||||
|
||||
async fn get_public_model_by_name(
|
||||
&self,
|
||||
model_name: &str,
|
||||
) -> Result<Option<StoredPublicGlobalModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.get_public_model_by_name(model_name)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_public_catalog_models(
|
||||
&self,
|
||||
query: &PublicCatalogModelListQuery,
|
||||
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_public_catalog_models(query)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn search_public_catalog_models(
|
||||
&self,
|
||||
query: &PublicCatalogModelSearchQuery,
|
||||
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.search_public_catalog_models(query)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_admin_global_models(
|
||||
&self,
|
||||
query: &AdminGlobalModelListQuery,
|
||||
) -> Result<StoredAdminGlobalModelPage, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_admin_global_models(query)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_admin_provider_models(
|
||||
&self,
|
||||
query: &AdminProviderModelListQuery,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_admin_provider_models(query)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_admin_provider_available_source_models(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_admin_provider_available_source_models(provider_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_admin_provider_model(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
model_id: &str,
|
||||
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.get_admin_provider_model(provider_id, model_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_admin_global_model_by_id(
|
||||
&self,
|
||||
global_model_id: &str,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.get_admin_global_model_by_id(global_model_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_admin_global_model_by_name(
|
||||
&self,
|
||||
model_name: &str,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.get_admin_global_model_by_name(model_name)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_admin_provider_models_by_global_model_id(
|
||||
&self,
|
||||
global_model_id: &str,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_admin_provider_models_by_global_model_id(global_model_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_provider_model_stats(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderModelStats>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_provider_model_stats(provider_ids)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_active_global_model_ids_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderActiveGlobalModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_active_global_model_ids_by_provider_ids(provider_ids)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl GlobalModelWriteRepository for MysqlGlobalModelReadRepository {
|
||||
async fn create_admin_provider_model(
|
||||
&self,
|
||||
record: &UpsertAdminProviderModelRecord,
|
||||
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
|
||||
Self::create_admin_provider_model(self, record).await
|
||||
}
|
||||
|
||||
async fn update_admin_provider_model(
|
||||
&self,
|
||||
record: &UpsertAdminProviderModelRecord,
|
||||
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
|
||||
Self::update_admin_provider_model(self, record).await
|
||||
}
|
||||
|
||||
async fn delete_admin_provider_model(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
model_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Self::delete_admin_provider_model(self, provider_id, model_id).await
|
||||
}
|
||||
|
||||
async fn create_admin_global_model(
|
||||
&self,
|
||||
record: &CreateAdminGlobalModelRecord,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
Self::create_admin_global_model(self, record).await
|
||||
}
|
||||
|
||||
async fn update_admin_global_model(
|
||||
&self,
|
||||
record: &UpdateAdminGlobalModelRecord,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
Self::update_admin_global_model(self, record).await
|
||||
}
|
||||
|
||||
async fn delete_admin_global_model(
|
||||
&self,
|
||||
global_model_id: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
Self::delete_admin_global_model(self, global_model_id).await
|
||||
}
|
||||
}
|
||||
|
||||
fn current_unix_secs() -> u64 {
|
||||
chrono::Utc::now().timestamp().max(0) as u64
|
||||
}
|
||||
|
||||
fn optional_json_to_string(
|
||||
value: &Option<serde_json::Value>,
|
||||
field_name: &str,
|
||||
) -> Result<Option<String>, DataLayerError> {
|
||||
value
|
||||
.as_ref()
|
||||
.map(|value| {
|
||||
serde_json::to_string(value).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains unserializable JSON: {err}"
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn optional_json_from_string(
|
||||
value: Option<String>,
|
||||
field_name: &str,
|
||||
) -> Result<Option<serde_json::Value>, DataLayerError> {
|
||||
value
|
||||
.map(|value| {
|
||||
serde_json::from_str(&value).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"{field_name} contains invalid JSON: {err}"
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn optional_u64(value: Option<i64>, field_name: &str) -> Result<Option<u64>, DataLayerError> {
|
||||
value
|
||||
.map(|value| {
|
||||
u64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!("invalid {field_name}: {value}"))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn first_tier_price(value: Option<&serde_json::Value>, key: &str) -> Option<f64> {
|
||||
value
|
||||
.and_then(|value| value.get("tiers"))
|
||||
.and_then(serde_json::Value::as_array)
|
||||
.and_then(|tiers| tiers.first())
|
||||
.and_then(|tier| tier.get(key))
|
||||
.and_then(serde_json::Value::as_f64)
|
||||
}
|
||||
|
||||
fn map_public_global_model_row(row: &MySqlRow) -> Result<StoredPublicGlobalModel, DataLayerError> {
|
||||
StoredPublicGlobalModel::new(
|
||||
row.try_get("id").map_sql_err()?,
|
||||
row.try_get("name").map_sql_err()?,
|
||||
row.try_get("display_name").map_sql_err()?,
|
||||
row.try_get("is_active").map_sql_err()?,
|
||||
row.try_get("default_price_per_request").map_sql_err()?,
|
||||
optional_json_from_string(
|
||||
row.try_get("default_tiered_pricing").map_sql_err()?,
|
||||
"global_models.default_tiered_pricing",
|
||||
)?,
|
||||
optional_json_from_string(
|
||||
row.try_get("supported_capabilities").map_sql_err()?,
|
||||
"global_models.supported_capabilities",
|
||||
)?,
|
||||
optional_json_from_string(row.try_get("config").map_sql_err()?, "global_models.config")?,
|
||||
row.try_get::<i64, _>("usage_count").map_sql_err()?.max(0) as u64,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_admin_global_model_row(row: &MySqlRow) -> Result<StoredAdminGlobalModel, DataLayerError> {
|
||||
StoredAdminGlobalModel::new(
|
||||
row.try_get("id").map_sql_err()?,
|
||||
row.try_get("name").map_sql_err()?,
|
||||
row.try_get("display_name").map_sql_err()?,
|
||||
row.try_get("is_active").map_sql_err()?,
|
||||
row.try_get("default_price_per_request").map_sql_err()?,
|
||||
optional_json_from_string(
|
||||
row.try_get("default_tiered_pricing").map_sql_err()?,
|
||||
"global_models.default_tiered_pricing",
|
||||
)?,
|
||||
optional_json_from_string(
|
||||
row.try_get("supported_capabilities").map_sql_err()?,
|
||||
"global_models.supported_capabilities",
|
||||
)?,
|
||||
optional_json_from_string(row.try_get("config").map_sql_err()?, "global_models.config")?,
|
||||
0,
|
||||
0,
|
||||
row.try_get::<i64, _>("usage_count").map_sql_err()?.max(0) as u64,
|
||||
optional_u64(
|
||||
row.try_get("created_at_unix_ms").map_sql_err()?,
|
||||
"global_models.created_at",
|
||||
)?,
|
||||
optional_u64(
|
||||
row.try_get("updated_at_unix_secs").map_sql_err()?,
|
||||
"global_models.updated_at",
|
||||
)?,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_admin_provider_model_row(
|
||||
row: &MySqlRow,
|
||||
) -> Result<StoredAdminProviderModel, DataLayerError> {
|
||||
StoredAdminProviderModel::new(
|
||||
row.try_get("id").map_sql_err()?,
|
||||
row.try_get("provider_id").map_sql_err()?,
|
||||
row.try_get("global_model_id").map_sql_err()?,
|
||||
row.try_get("provider_model_name").map_sql_err()?,
|
||||
optional_json_from_string(
|
||||
row.try_get("provider_model_mappings").map_sql_err()?,
|
||||
"models.provider_model_mappings",
|
||||
)?,
|
||||
row.try_get("price_per_request").map_sql_err()?,
|
||||
optional_json_from_string(
|
||||
row.try_get("tiered_pricing").map_sql_err()?,
|
||||
"models.tiered_pricing",
|
||||
)?,
|
||||
row.try_get("supports_vision").map_sql_err()?,
|
||||
row.try_get("supports_function_calling").map_sql_err()?,
|
||||
row.try_get("supports_streaming").map_sql_err()?,
|
||||
row.try_get("supports_extended_thinking").map_sql_err()?,
|
||||
row.try_get("supports_image_generation").map_sql_err()?,
|
||||
row.try_get("is_active").map_sql_err()?,
|
||||
row.try_get("is_available").map_sql_err()?,
|
||||
optional_json_from_string(row.try_get("config").map_sql_err()?, "models.config")?,
|
||||
optional_u64(
|
||||
row.try_get("created_at_unix_ms").map_sql_err()?,
|
||||
"models.created_at",
|
||||
)?,
|
||||
optional_u64(
|
||||
row.try_get("updated_at_unix_secs").map_sql_err()?,
|
||||
"models.updated_at",
|
||||
)?,
|
||||
row.try_get("global_model_name").map_sql_err()?,
|
||||
row.try_get("global_model_display_name").map_sql_err()?,
|
||||
row.try_get("global_model_default_price_per_request")
|
||||
.map_sql_err()?,
|
||||
optional_json_from_string(
|
||||
row.try_get("global_model_default_tiered_pricing")
|
||||
.map_sql_err()?,
|
||||
"global_models.default_tiered_pricing",
|
||||
)?,
|
||||
optional_json_from_string(
|
||||
row.try_get("global_model_supported_capabilities")
|
||||
.map_sql_err()?,
|
||||
"global_models.supported_capabilities",
|
||||
)?,
|
||||
optional_json_from_string(
|
||||
row.try_get("global_model_config").map_sql_err()?,
|
||||
"global_models.config",
|
||||
)?,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_public_catalog_model_row(
|
||||
row: &MySqlRow,
|
||||
) -> Result<StoredPublicCatalogModel, DataLayerError> {
|
||||
let global_model_config = optional_json_from_string(
|
||||
row.try_get("global_model_config").map_sql_err()?,
|
||||
"global_models.config",
|
||||
)?;
|
||||
let global_model_supported_capabilities = optional_json_from_string(
|
||||
row.try_get("global_model_supported_capabilities")
|
||||
.map_sql_err()?,
|
||||
"global_models.supported_capabilities",
|
||||
)?;
|
||||
let model_config =
|
||||
optional_json_from_string(row.try_get("model_config").map_sql_err()?, "models.config")?;
|
||||
let tiered_pricing = optional_json_from_string(
|
||||
row.try_get("tiered_pricing").map_sql_err()?,
|
||||
"models.tiered_pricing",
|
||||
)?;
|
||||
let default_tiered_pricing = optional_json_from_string(
|
||||
row.try_get("default_tiered_pricing").map_sql_err()?,
|
||||
"global_models.default_tiered_pricing",
|
||||
)?;
|
||||
let pricing = tiered_pricing.as_ref().or(default_tiered_pricing.as_ref());
|
||||
let global_model_is_active = row
|
||||
.try_get::<Option<bool>, _>("global_model_is_active")
|
||||
.map_sql_err()?
|
||||
.unwrap_or(true);
|
||||
let model_is_active: bool = row.try_get("is_active").map_sql_err()?;
|
||||
let provider_is_active: bool = row.try_get("provider_is_active").map_sql_err()?;
|
||||
|
||||
StoredPublicCatalogModel::new(
|
||||
row.try_get("id").map_sql_err()?,
|
||||
row.try_get("provider_id").map_sql_err()?,
|
||||
row.try_get("provider_name").map_sql_err()?,
|
||||
row.try_get("provider_model_name").map_sql_err()?,
|
||||
row.try_get("name").map_sql_err()?,
|
||||
row.try_get("display_name").map_sql_err()?,
|
||||
global_model_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("description"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToString::to_string),
|
||||
global_model_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("icon_url"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(ToString::to_string),
|
||||
first_tier_price(pricing, "input_price_per_1m"),
|
||||
first_tier_price(pricing, "output_price_per_1m"),
|
||||
first_tier_price(pricing, "cache_creation_price_per_1m"),
|
||||
first_tier_price(pricing, "cache_read_price_per_1m"),
|
||||
row.try_get("supports_vision").map_sql_err()?,
|
||||
row.try_get("supports_function_calling").map_sql_err()?,
|
||||
row.try_get("supports_streaming").map_sql_err()?,
|
||||
metadata_supports_embedding(
|
||||
global_model_supported_capabilities.as_ref(),
|
||||
global_model_config.as_ref(),
|
||||
model_config.as_ref(),
|
||||
),
|
||||
model_is_active && provider_is_active && global_model_is_active,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_provider_model_stats_row(
|
||||
row: &MySqlRow,
|
||||
) -> Result<StoredProviderModelStats, DataLayerError> {
|
||||
StoredProviderModelStats::new(
|
||||
row.try_get("provider_id").map_sql_err()?,
|
||||
row.try_get("total_models").map_sql_err()?,
|
||||
row.try_get::<Option<i64>, _>("active_models")
|
||||
.map_sql_err()?
|
||||
.unwrap_or(0),
|
||||
)
|
||||
}
|
||||
|
||||
fn map_active_global_model_row(
|
||||
row: &MySqlRow,
|
||||
) -> Result<StoredProviderActiveGlobalModel, DataLayerError> {
|
||||
StoredProviderActiveGlobalModel::new(
|
||||
row.try_get("provider_id").map_sql_err()?,
|
||||
row.try_get("global_model_id").map_sql_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::MysqlGlobalModelReadRepository;
|
||||
|
||||
#[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 = MysqlGlobalModelReadRepository::new(pool);
|
||||
}
|
||||
}
|
||||
+11
-21
@@ -405,7 +405,7 @@ SELECT
|
||||
gm.config,
|
||||
COALESCE(gm_stats.provider_count, 0) AS provider_count,
|
||||
COALESCE(gm_stats.active_provider_count, 0) AS active_provider_count,
|
||||
COALESCE(usage_stats.usage_count, gm.usage_count, 0)::bigint AS usage_count,
|
||||
COALESCE(gm.usage_count, 0)::bigint AS usage_count,
|
||||
EXTRACT(EPOCH FROM gm.created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM gm.updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM global_models gm
|
||||
@@ -423,12 +423,6 @@ LEFT JOIN (
|
||||
JOIN providers p ON p.id = m.provider_id
|
||||
GROUP BY m.global_model_id
|
||||
) gm_stats ON gm_stats.global_model_id = gm.id
|
||||
LEFT JOIN LATERAL (
|
||||
SELECT NULLIF(COUNT(*), 0)::bigint AS usage_count
|
||||
FROM usage_billing_facts AS usage
|
||||
WHERE usage.model = gm.name
|
||||
AND usage.status NOT IN ('pending', 'streaming')
|
||||
) usage_stats ON TRUE
|
||||
WHERE gm.id = $1
|
||||
LIMIT 1
|
||||
"#,
|
||||
@@ -458,7 +452,7 @@ SELECT
|
||||
gm.config,
|
||||
COALESCE(gm_stats.provider_count, 0) AS provider_count,
|
||||
COALESCE(gm_stats.active_provider_count, 0) AS active_provider_count,
|
||||
COALESCE(usage_stats.usage_count, gm.usage_count, 0)::bigint AS usage_count,
|
||||
COALESCE(gm.usage_count, 0)::bigint AS usage_count,
|
||||
EXTRACT(EPOCH FROM gm.created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM gm.updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM global_models gm
|
||||
@@ -476,12 +470,6 @@ LEFT JOIN (
|
||||
JOIN providers p ON p.id = m.provider_id
|
||||
GROUP BY m.global_model_id
|
||||
) gm_stats ON gm_stats.global_model_id = gm.id
|
||||
LEFT JOIN LATERAL (
|
||||
SELECT NULLIF(COUNT(*), 0)::bigint AS usage_count
|
||||
FROM usage_billing_facts AS usage
|
||||
WHERE usage.model = gm.name
|
||||
AND usage.status NOT IN ('pending', 'streaming')
|
||||
) usage_stats ON TRUE
|
||||
WHERE gm.name = $1
|
||||
LIMIT 1
|
||||
"#,
|
||||
@@ -1221,7 +1209,7 @@ mod tests {
|
||||
SqlxGlobalModelReadRepository, LIST_ADMIN_GLOBAL_MODELS_PREFIX,
|
||||
LIST_ADMIN_PROVIDER_MODELS_PREFIX,
|
||||
};
|
||||
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
|
||||
const ADMIN_PROVIDER_MODEL_REQUIRED_COLUMNS: &[&str] = &[
|
||||
"global_model_default_tiered_pricing",
|
||||
@@ -1243,13 +1231,13 @@ mod tests {
|
||||
assert_admin_provider_model_projection_has_required_columns(
|
||||
LIST_ADMIN_PROVIDER_MODELS_PREFIX,
|
||||
);
|
||||
assert_admin_provider_model_projection_has_required_columns(include_str!("sql.rs"));
|
||||
assert_admin_provider_model_projection_has_required_columns(include_str!("postgres.rs"));
|
||||
let supported_capabilities_projection = format!(
|
||||
"{} AS {}",
|
||||
"gm.supported_capabilities", "global_model_supported_capabilities"
|
||||
);
|
||||
assert_eq!(
|
||||
include_str!("sql.rs")
|
||||
include_str!("postgres.rs")
|
||||
.matches(&supported_capabilities_projection)
|
||||
.count(),
|
||||
4
|
||||
@@ -1258,7 +1246,7 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn admin_global_model_sql_projects_usage_count_from_billing_facts() {
|
||||
let source = include_str!("sql.rs");
|
||||
let source = include_str!("postgres.rs");
|
||||
|
||||
assert!(
|
||||
!LIST_ADMIN_GLOBAL_MODELS_PREFIX.contains("usage_billing_facts"),
|
||||
@@ -1292,9 +1280,11 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn global_model_usage_count_read_model_has_backfill_and_delta_maintenance() {
|
||||
let backfill_sql =
|
||||
include_str!("../../../backfills/20260505120000_rebuild_global_model_usage_count.sql");
|
||||
let delta_sql = include_str!("../usage/sql/queries/apply_global_model_usage_delta_sql.sql");
|
||||
let backfill_sql = include_str!(
|
||||
"../../../backfills/postgres/20260505120000_rebuild_global_model_usage_count.sql"
|
||||
);
|
||||
let delta_sql =
|
||||
include_str!("../usage/postgres/queries/apply_global_model_usage_delta_sql.sql");
|
||||
|
||||
assert!(
|
||||
backfill_sql.contains("UPDATE global_models AS gm"),
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user