chore: resolve pr 377 checks

This commit is contained in:
fawney19
2026-05-06 02:22:32 +08:00
502 changed files with 89553 additions and 21893 deletions

View File

@@ -1,9 +1,13 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
mod types;
pub use memory::InMemoryAnnouncementReadRepository;
pub use sql::SqlxAnnouncementReadRepository;
pub use mysql::MysqlAnnouncementRepository;
pub use postgres::SqlxAnnouncementReadRepository;
pub use sqlite::SqliteAnnouncementRepository;
pub use types::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,

View File

@@ -0,0 +1,322 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use super::types::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const ANNOUNCEMENT_SELECT: &str = r#"
SELECT
a.id,
a.title,
a.content,
a.`type` AS type,
a.priority,
a.is_active,
a.is_pinned,
a.author_id,
u.username AS author_username,
a.start_time AS start_time_unix_secs,
a.end_time AS end_time_unix_secs,
a.created_at AS created_at_unix_ms,
a.updated_at AS updated_at_unix_secs
FROM announcements a
LEFT JOIN users u ON u.id = a.author_id
"#;
#[derive(Debug, Clone)]
pub struct MysqlAnnouncementRepository {
pool: MysqlPool,
}
impl MysqlAnnouncementRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn reload_by_id(
&self,
announcement_id: &str,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
self.find_by_id(announcement_id).await
}
}
#[async_trait]
impl AnnouncementReadRepository for MysqlAnnouncementRepository {
async fn find_by_id(
&self,
announcement_id: &str,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
let row = sqlx::query(&format!("{ANNOUNCEMENT_SELECT} WHERE a.id = ? LIMIT 1"))
.bind(announcement_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_announcement_row).transpose()
}
async fn list_announcements(
&self,
query: &AnnouncementListQuery,
) -> Result<StoredAnnouncementPage, DataLayerError> {
let now_unix_secs = query.now_unix_secs.unwrap_or_else(current_unix_secs);
let total_row = sqlx::query(
r#"
SELECT COUNT(a.id) AS total
FROM announcements a
WHERE (
NOT ? OR (
a.is_active = 1
AND (a.start_time IS NULL OR a.start_time <= ?)
AND (a.end_time IS NULL OR a.end_time >= ?)
)
)
"#,
)
.bind(query.active_only)
.bind(now_unix_secs as i64)
.bind(now_unix_secs as i64)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let total = total_row.try_get::<i64, _>("total").map_sql_err()?.max(0) as u64;
let rows = sqlx::query(&format!(
r#"
{ANNOUNCEMENT_SELECT}
WHERE (
NOT ? OR (
a.is_active = 1
AND (a.start_time IS NULL OR a.start_time <= ?)
AND (a.end_time IS NULL OR a.end_time >= ?)
)
)
ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC
LIMIT ? OFFSET ?
"#
))
.bind(query.active_only)
.bind(now_unix_secs as i64)
.bind(now_unix_secs as i64)
.bind(query.limit as i64)
.bind(query.offset as i64)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows
.iter()
.map(map_announcement_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(StoredAnnouncementPage { items, total })
}
async fn count_unread_active_announcements(
&self,
user_id: &str,
now_unix_secs: u64,
) -> Result<u64, DataLayerError> {
let row = sqlx::query(
r#"
SELECT COUNT(a.id) AS total
FROM announcements a
WHERE a.is_active = 1
AND (a.start_time IS NULL OR a.start_time <= ?)
AND (a.end_time IS NULL OR a.end_time >= ?)
AND NOT EXISTS (
SELECT 1
FROM announcement_reads r
WHERE r.user_id = ?
AND r.announcement_id = a.id
)
"#,
)
.bind(now_unix_secs as i64)
.bind(now_unix_secs as i64)
.bind(user_id)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
Ok(row.try_get::<i64, _>("total").map_sql_err()?.max(0) as u64)
}
}
#[async_trait]
impl AnnouncementWriteRepository for MysqlAnnouncementRepository {
async fn create_announcement(
&self,
record: CreateAnnouncementRecord,
) -> Result<StoredAnnouncement, DataLayerError> {
record.validate()?;
let id = uuid::Uuid::new_v4().to_string();
let now = current_unix_secs() as i64;
sqlx::query(
r#"
INSERT INTO announcements (
id, title, content, `type`, priority, author_id, is_active, is_pinned,
start_time, end_time, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, 1, ?, ?, ?, ?, ?)
"#,
)
.bind(&id)
.bind(record.title)
.bind(record.content)
.bind(record.kind)
.bind(record.priority)
.bind(record.author_id)
.bind(record.is_pinned)
.bind(optional_i64_from_u64(
record.start_time_unix_secs,
"announcements.start_time",
)?)
.bind(optional_i64_from_u64(
record.end_time_unix_secs,
"announcements.end_time",
)?)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_by_id(&id)
.await?
.ok_or_else(|| DataLayerError::UnexpectedValue("created announcement missing".into()))
}
async fn update_announcement(
&self,
record: UpdateAnnouncementRecord,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
record.validate()?;
let id = record.announcement_id;
sqlx::query(
r#"
UPDATE announcements
SET title = COALESCE(?, title),
content = COALESCE(?, content),
`type` = COALESCE(?, `type`),
priority = COALESCE(?, priority),
is_active = COALESCE(?, is_active),
is_pinned = COALESCE(?, is_pinned),
start_time = COALESCE(?, start_time),
end_time = COALESCE(?, end_time),
updated_at = ?
WHERE id = ?
"#,
)
.bind(record.title)
.bind(record.content)
.bind(record.kind)
.bind(record.priority)
.bind(record.is_active)
.bind(record.is_pinned)
.bind(optional_i64_from_u64(
record.start_time_unix_secs,
"announcements.start_time",
)?)
.bind(optional_i64_from_u64(
record.end_time_unix_secs,
"announcements.end_time",
)?)
.bind(current_unix_secs() as i64)
.bind(&id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_by_id(&id).await
}
async fn delete_announcement(&self, announcement_id: &str) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query("DELETE FROM announcements WHERE id = ?")
.bind(announcement_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn mark_announcement_as_read(
&self,
user_id: &str,
announcement_id: &str,
read_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query(
r#"
INSERT IGNORE INTO announcement_reads (id, user_id, announcement_id, read_at)
VALUES (?, ?, ?, ?)
"#,
)
.bind(uuid::Uuid::new_v4().to_string())
.bind(user_id)
.bind(announcement_id)
.bind(i64_from_u64(
read_at_unix_secs,
"announcement_reads.read_at",
)?)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
}
fn current_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn i64_from_u64(value: u64, field_name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}")))
}
fn optional_i64_from_u64(
value: Option<u64>,
field_name: &str,
) -> Result<Option<i64>, DataLayerError> {
value
.map(|value| i64_from_u64(value, field_name))
.transpose()
}
fn map_announcement_row(row: &MySqlRow) -> Result<StoredAnnouncement, DataLayerError> {
StoredAnnouncement::new(
row.try_get("id").map_sql_err()?,
row.try_get("title").map_sql_err()?,
row.try_get("content").map_sql_err()?,
row.try_get("type").map_sql_err()?,
row.try_get("priority").map_sql_err()?,
row.try_get("is_active").map_sql_err()?,
row.try_get("is_pinned").map_sql_err()?,
row.try_get("author_id").map_sql_err()?,
row.try_get("author_username").map_sql_err()?,
row.try_get("start_time_unix_secs").map_sql_err()?,
row.try_get("end_time_unix_secs").map_sql_err()?,
row.try_get("created_at_unix_ms").map_sql_err()?,
row.try_get("updated_at_unix_secs").map_sql_err()?,
)
}
#[cfg(test)]
mod tests {
use super::MysqlAnnouncementRepository;
#[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 = MysqlAnnouncementRepository::new(pool);
}
}

View File

@@ -357,7 +357,7 @@ fn map_announcement_row(row: &PgRow) -> Result<StoredAnnouncement, DataLayerErro
#[cfg(test)]
mod tests {
use super::SqlxAnnouncementReadRepository;
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
#[tokio::test]
async fn repository_constructs_from_lazy_pool() {

View File

@@ -0,0 +1,421 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row};
use super::types::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const ANNOUNCEMENT_SELECT: &str = r#"
SELECT
a.id,
a.title,
a.content,
a.type,
a.priority,
a.is_active,
a.is_pinned,
a.author_id,
u.username AS author_username,
a.start_time AS start_time_unix_secs,
a.end_time AS end_time_unix_secs,
a.created_at AS created_at_unix_ms,
a.updated_at AS updated_at_unix_secs
FROM announcements a
LEFT JOIN users u ON u.id = a.author_id
"#;
#[derive(Debug, Clone)]
pub struct SqliteAnnouncementRepository {
pool: SqlitePool,
}
impl SqliteAnnouncementRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn reload_by_id(
&self,
announcement_id: &str,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
self.find_by_id(announcement_id).await
}
}
#[async_trait]
impl AnnouncementReadRepository for SqliteAnnouncementRepository {
async fn find_by_id(
&self,
announcement_id: &str,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
let row = sqlx::query(&format!("{ANNOUNCEMENT_SELECT} WHERE a.id = ? LIMIT 1"))
.bind(announcement_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_announcement_row).transpose()
}
async fn list_announcements(
&self,
query: &AnnouncementListQuery,
) -> Result<StoredAnnouncementPage, DataLayerError> {
let now_unix_secs = query.now_unix_secs.unwrap_or_else(current_unix_secs);
let total_row = sqlx::query(
r#"
SELECT COUNT(a.id) AS total
FROM announcements a
WHERE (
NOT ? OR (
a.is_active = 1
AND (a.start_time IS NULL OR a.start_time <= ?)
AND (a.end_time IS NULL OR a.end_time >= ?)
)
)
"#,
)
.bind(query.active_only)
.bind(now_unix_secs as i64)
.bind(now_unix_secs as i64)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let total = total_row.try_get::<i64, _>("total").map_sql_err()?.max(0) as u64;
let rows = sqlx::query(&format!(
r#"
{ANNOUNCEMENT_SELECT}
WHERE (
NOT ? OR (
a.is_active = 1
AND (a.start_time IS NULL OR a.start_time <= ?)
AND (a.end_time IS NULL OR a.end_time >= ?)
)
)
ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC
LIMIT ? OFFSET ?
"#
))
.bind(query.active_only)
.bind(now_unix_secs as i64)
.bind(now_unix_secs as i64)
.bind(query.limit as i64)
.bind(query.offset as i64)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows
.iter()
.map(map_announcement_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(StoredAnnouncementPage { items, total })
}
async fn count_unread_active_announcements(
&self,
user_id: &str,
now_unix_secs: u64,
) -> Result<u64, DataLayerError> {
let row = sqlx::query(
r#"
SELECT COUNT(a.id) AS total
FROM announcements a
WHERE a.is_active = 1
AND (a.start_time IS NULL OR a.start_time <= ?)
AND (a.end_time IS NULL OR a.end_time >= ?)
AND NOT EXISTS (
SELECT 1
FROM announcement_reads r
WHERE r.user_id = ?
AND r.announcement_id = a.id
)
"#,
)
.bind(now_unix_secs as i64)
.bind(now_unix_secs as i64)
.bind(user_id)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
Ok(row.try_get::<i64, _>("total").map_sql_err()?.max(0) as u64)
}
}
#[async_trait]
impl AnnouncementWriteRepository for SqliteAnnouncementRepository {
async fn create_announcement(
&self,
record: CreateAnnouncementRecord,
) -> Result<StoredAnnouncement, DataLayerError> {
record.validate()?;
let id = uuid::Uuid::new_v4().to_string();
let now = current_unix_secs() as i64;
sqlx::query(
r#"
INSERT INTO announcements (
id, title, content, type, priority, author_id, is_active, is_pinned,
start_time, end_time, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, 1, ?, ?, ?, ?, ?)
"#,
)
.bind(&id)
.bind(record.title)
.bind(record.content)
.bind(record.kind)
.bind(record.priority)
.bind(record.author_id)
.bind(record.is_pinned)
.bind(optional_i64_from_u64(
record.start_time_unix_secs,
"announcements.start_time",
)?)
.bind(optional_i64_from_u64(
record.end_time_unix_secs,
"announcements.end_time",
)?)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_by_id(&id)
.await?
.ok_or_else(|| DataLayerError::UnexpectedValue("created announcement missing".into()))
}
async fn update_announcement(
&self,
record: UpdateAnnouncementRecord,
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
record.validate()?;
let id = record.announcement_id;
sqlx::query(
r#"
UPDATE announcements
SET title = COALESCE(?, title),
content = COALESCE(?, content),
type = COALESCE(?, type),
priority = COALESCE(?, priority),
is_active = COALESCE(?, is_active),
is_pinned = COALESCE(?, is_pinned),
start_time = COALESCE(?, start_time),
end_time = COALESCE(?, end_time),
updated_at = ?
WHERE id = ?
"#,
)
.bind(record.title)
.bind(record.content)
.bind(record.kind)
.bind(record.priority)
.bind(record.is_active)
.bind(record.is_pinned)
.bind(optional_i64_from_u64(
record.start_time_unix_secs,
"announcements.start_time",
)?)
.bind(optional_i64_from_u64(
record.end_time_unix_secs,
"announcements.end_time",
)?)
.bind(current_unix_secs() as i64)
.bind(&id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_by_id(&id).await
}
async fn delete_announcement(&self, announcement_id: &str) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query("DELETE FROM announcements WHERE id = ?")
.bind(announcement_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn mark_announcement_as_read(
&self,
user_id: &str,
announcement_id: &str,
read_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query(
r#"
INSERT OR IGNORE INTO announcement_reads (id, user_id, announcement_id, read_at)
VALUES (?, ?, ?, ?)
"#,
)
.bind(uuid::Uuid::new_v4().to_string())
.bind(user_id)
.bind(announcement_id)
.bind(i64_from_u64(
read_at_unix_secs,
"announcement_reads.read_at",
)?)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
}
fn current_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn i64_from_u64(value: u64, field_name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}")))
}
fn optional_i64_from_u64(
value: Option<u64>,
field_name: &str,
) -> Result<Option<i64>, DataLayerError> {
value
.map(|value| i64_from_u64(value, field_name))
.transpose()
}
fn map_announcement_row(row: &SqliteRow) -> Result<StoredAnnouncement, DataLayerError> {
StoredAnnouncement::new(
row.try_get("id").map_sql_err()?,
row.try_get("title").map_sql_err()?,
row.try_get("content").map_sql_err()?,
row.try_get("type").map_sql_err()?,
row.try_get("priority").map_sql_err()?,
row.try_get("is_active").map_sql_err()?,
row.try_get("is_pinned").map_sql_err()?,
row.try_get("author_id").map_sql_err()?,
row.try_get("author_username").map_sql_err()?,
row.try_get("start_time_unix_secs").map_sql_err()?,
row.try_get("end_time_unix_secs").map_sql_err()?,
row.try_get("created_at_unix_ms").map_sql_err()?,
row.try_get("updated_at_unix_secs").map_sql_err()?,
)
}
#[cfg(test)]
mod tests {
use super::SqliteAnnouncementRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::announcements::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, UpdateAnnouncementRecord,
};
#[tokio::test]
async fn sqlite_repository_reads_and_writes_announcements() {
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_announcement_user(&pool).await;
let repository = SqliteAnnouncementRepository::new(pool);
let created = repository
.create_announcement(CreateAnnouncementRecord {
title: "Initial".to_string(),
content: "Body".to_string(),
kind: "info".to_string(),
priority: 10,
is_pinned: true,
author_id: "user-1".to_string(),
start_time_unix_secs: Some(100),
end_time_unix_secs: Some(300),
})
.await
.expect("announcement should create");
assert_eq!(created.author_username, Some("admin".to_string()));
assert!(created.is_active);
let page = repository
.list_announcements(&AnnouncementListQuery {
active_only: true,
offset: 0,
limit: 10,
now_unix_secs: Some(200),
})
.await
.expect("announcements should list");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].id, created.id);
let unread = repository
.count_unread_active_announcements("user-1", 200)
.await
.expect("unread count should load");
assert_eq!(unread, 1);
assert!(repository
.mark_announcement_as_read("user-1", &created.id, 210)
.await
.expect("read marker should insert"));
assert!(!repository
.mark_announcement_as_read("user-1", &created.id, 211)
.await
.expect("duplicate read marker should be ignored"));
assert_eq!(
repository
.count_unread_active_announcements("user-1", 200)
.await
.expect("unread count should reload"),
0
);
let updated = repository
.update_announcement(UpdateAnnouncementRecord {
announcement_id: created.id.clone(),
title: Some("Updated".to_string()),
content: None,
kind: None,
priority: Some(20),
is_active: Some(false),
is_pinned: Some(false),
start_time_unix_secs: None,
end_time_unix_secs: None,
})
.await
.expect("announcement should update")
.expect("announcement should exist");
assert_eq!(updated.title, "Updated");
assert!(!updated.is_active);
assert!(repository
.delete_announcement(&created.id)
.await
.expect("announcement should delete"));
assert!(repository
.find_by_id(&created.id)
.await
.expect("find should run")
.is_none());
}
async fn seed_announcement_user(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
INSERT INTO users (
id, email, username, role, auth_source, email_verified, is_active, is_deleted, created_at, updated_at
)
VALUES ('user-1', 'admin@example.com', 'admin', 'admin', 'local', 1, 1, 0, 1, 1)
"#,
)
.execute(pool)
.await
.expect("user should seed");
}
}

File diff suppressed because it is too large Load Diff

View File

@@ -1,9 +1,13 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
mod types;
pub use memory::InMemoryAuthApiKeySnapshotRepository;
pub use sql::SqlxAuthApiKeySnapshotReadRepository;
pub use mysql::MysqlAuthApiKeyReadRepository;
pub use postgres::SqlxAuthApiKeySnapshotReadRepository;
pub use sqlite::SqliteAuthApiKeyReadRepository;
pub use types::{
read_resolved_auth_api_key_snapshot, read_resolved_auth_api_key_snapshot_by_key_hash,
read_resolved_auth_api_key_snapshot_by_user_api_key_ids, AuthApiKeyExportSummary,

View File

@@ -0,0 +1,933 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::types::{
AuthApiKeyExportSummary, AuthApiKeyLookupKey, AuthApiKeyReadRepository,
AuthApiKeyWriteRepository, CreateStandaloneApiKeyRecord, CreateUserApiKeyRecord,
StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const SNAPSHOT_COLUMNS: &str = r#"
SELECT
users.id AS user_id,
users.username,
users.email,
users.role AS user_role,
users.auth_source AS user_auth_source,
users.is_active AS user_is_active,
users.is_deleted AS user_is_deleted,
users.rate_limit AS user_rate_limit,
users.allowed_providers AS user_allowed_providers,
users.allowed_api_formats AS user_allowed_api_formats,
users.allowed_models AS user_allowed_models,
api_keys.id AS api_key_id,
api_keys.name AS api_key_name,
api_keys.is_active AS api_key_is_active,
api_keys.is_locked AS api_key_is_locked,
api_keys.is_standalone AS api_key_is_standalone,
api_keys.rate_limit AS api_key_rate_limit,
api_keys.concurrent_limit AS api_key_concurrent_limit,
api_keys.expires_at AS api_key_expires_at_unix_secs,
api_keys.allowed_providers AS api_key_allowed_providers,
api_keys.allowed_api_formats AS api_key_allowed_api_formats,
api_keys.allowed_models AS api_key_allowed_models
FROM api_keys
JOIN users ON users.id = api_keys.user_id
"#;
const EXPORT_COLUMNS: &str = r#"
SELECT
api_keys.user_id,
api_keys.id AS api_key_id,
api_keys.key_hash,
api_keys.key_encrypted,
api_keys.name,
api_keys.allowed_providers,
api_keys.allowed_api_formats,
api_keys.allowed_models,
api_keys.rate_limit,
api_keys.concurrent_limit,
api_keys.force_capabilities,
api_keys.is_active,
api_keys.expires_at AS expires_at_unix_secs,
api_keys.auto_delete_on_expiry,
api_keys.total_requests,
COALESCE(api_keys.total_tokens, 0) AS total_tokens,
COALESCE(api_keys.total_cost_usd, 0) AS total_cost_usd,
api_keys.last_used_at AS last_used_at_unix_secs,
api_keys.created_at AS created_at_unix_secs,
api_keys.updated_at AS updated_at_unix_secs,
api_keys.is_standalone
FROM api_keys
"#;
#[derive(Debug, Clone)]
pub struct MysqlAuthApiKeyReadRepository {
pool: MysqlPool,
}
impl MysqlAuthApiKeyReadRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn fetch_snapshot_rows(
&self,
mut builder: QueryBuilder<'_, MySql>,
) -> Result<Vec<StoredAuthApiKeySnapshot>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_auth_api_key_snapshot_row).collect()
}
async fn fetch_export_rows(
&self,
mut builder: QueryBuilder<'_, MySql>,
) -> Result<Vec<StoredAuthApiKeyExportRecord>, DataLayerError> {
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_auth_api_key_export_row).collect()
}
async fn reload_export_by_id(
&self,
api_key_id: &str,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
Ok(self
.list_export_api_keys_by_ids(&[api_key_id.to_string()])
.await?
.into_iter()
.next())
}
async fn create_api_key(
&self,
record: CreateApiKeyInsertRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let now = current_unix_secs();
sqlx::query(
r#"
INSERT INTO api_keys (
id, user_id, key_hash, key_encrypted, name, allowed_providers,
allowed_api_formats, allowed_models, rate_limit, concurrent_limit,
force_capabilities, is_active, expires_at, auto_delete_on_expiry,
total_requests, total_tokens, total_cost_usd, is_standalone,
created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&record.api_key_id)
.bind(&record.user_id)
.bind(&record.key_hash)
.bind(&record.key_encrypted)
.bind(&record.name)
.bind(json_string_from_string_list(
record.allowed_providers.as_ref(),
"api_keys.allowed_providers",
)?)
.bind(json_string_from_string_list(
record.allowed_api_formats.as_ref(),
"api_keys.allowed_api_formats",
)?)
.bind(json_string_from_string_list(
record.allowed_models.as_ref(),
"api_keys.allowed_models",
)?)
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(optional_json_to_string(
&record.force_capabilities,
"api_keys.force_capabilities",
)?)
.bind(record.is_active)
.bind(optional_i64_from_u64(
record.expires_at_unix_secs,
"api_keys.expires_at",
)?)
.bind(record.auto_delete_on_expiry)
.bind(i64_from_u64(
record.total_requests,
"api_keys.total_requests",
)?)
.bind(i64_from_u64(record.total_tokens, "api_keys.total_tokens")?)
.bind(record.total_cost_usd)
.bind(record.is_standalone)
.bind(now as i64)
.bind(now as i64)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_export_by_id(&record.api_key_id).await
}
}
struct CreateApiKeyInsertRecord {
user_id: String,
api_key_id: String,
key_hash: String,
key_encrypted: Option<String>,
name: Option<String>,
allowed_providers: Option<Vec<String>>,
allowed_api_formats: Option<Vec<String>>,
allowed_models: Option<Vec<String>>,
rate_limit: Option<i32>,
concurrent_limit: Option<i32>,
force_capabilities: Option<serde_json::Value>,
is_active: bool,
expires_at_unix_secs: Option<u64>,
auto_delete_on_expiry: bool,
total_requests: u64,
total_tokens: u64,
total_cost_usd: f64,
is_standalone: bool,
}
#[async_trait]
impl AuthApiKeyReadRepository for MysqlAuthApiKeyReadRepository {
async fn find_api_key_snapshot(
&self,
key: AuthApiKeyLookupKey<'_>,
) -> Result<Option<StoredAuthApiKeySnapshot>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(SNAPSHOT_COLUMNS);
match key {
AuthApiKeyLookupKey::KeyHash(key_hash) => {
builder
.push(" WHERE api_keys.key_hash = ")
.push_bind(key_hash);
}
AuthApiKeyLookupKey::ApiKeyId(api_key_id) => {
builder.push(" WHERE api_keys.id = ").push_bind(api_key_id);
}
AuthApiKeyLookupKey::UserApiKeyIds {
user_id,
api_key_id,
} => {
builder
.push(" WHERE api_keys.id = ")
.push_bind(api_key_id)
.push(" AND users.id = ")
.push_bind(user_id);
}
}
builder.push(" LIMIT 1");
Ok(self.fetch_snapshot_rows(builder).await?.into_iter().next())
}
async fn list_api_key_snapshots_by_ids(
&self,
api_key_ids: &[String],
) -> Result<Vec<StoredAuthApiKeySnapshot>, DataLayerError> {
if api_key_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(SNAPSHOT_COLUMNS);
push_in_clause(&mut builder, " WHERE api_keys.id IN (", api_key_ids);
builder.push(" ORDER BY api_keys.id ASC");
self.fetch_snapshot_rows(builder).await
}
async fn list_export_api_keys_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredAuthApiKeyExportRecord>, DataLayerError> {
if user_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(EXPORT_COLUMNS);
push_in_clause(&mut builder, " WHERE api_keys.user_id IN (", user_ids);
builder
.push(" AND api_keys.is_standalone = 0 ORDER BY api_keys.user_id ASC, api_keys.id ASC");
self.fetch_export_rows(builder).await
}
async fn list_export_api_keys_by_ids(
&self,
api_key_ids: &[String],
) -> Result<Vec<StoredAuthApiKeyExportRecord>, DataLayerError> {
if api_key_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(EXPORT_COLUMNS);
push_in_clause(&mut builder, " WHERE api_keys.id IN (", api_key_ids);
builder.push(" ORDER BY api_keys.id ASC");
self.fetch_export_rows(builder).await
}
async fn list_export_api_keys_by_name_search(
&self,
name_search: &str,
) -> Result<Vec<StoredAuthApiKeyExportRecord>, DataLayerError> {
let name_search = name_search.trim();
if name_search.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(EXPORT_COLUMNS);
builder
.push(" WHERE LOWER(COALESCE(api_keys.name, '')) LIKE ")
.push_bind(format!("%{}%", name_search.to_ascii_lowercase()))
.push(" ORDER BY api_keys.id ASC");
self.fetch_export_rows(builder).await
}
async fn list_export_standalone_api_keys_page(
&self,
query: &StandaloneApiKeyExportListQuery,
) -> Result<Vec<StoredAuthApiKeyExportRecord>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(EXPORT_COLUMNS);
builder.push(" WHERE api_keys.is_standalone = 1");
if let Some(is_active) = query.is_active {
builder
.push(" AND api_keys.is_active = ")
.push_bind(is_active);
}
builder
.push(" ORDER BY api_keys.id ASC LIMIT ")
.push_bind(i64::try_from(query.limit).map_err(|_| {
DataLayerError::InvalidInput(format!(
"invalid standalone api key export limit: {}",
query.limit
))
})?)
.push(" OFFSET ")
.push_bind(i64::try_from(query.skip).map_err(|_| {
DataLayerError::InvalidInput(format!(
"invalid standalone api key export skip: {}",
query.skip
))
})?);
self.fetch_export_rows(builder).await
}
async fn count_export_standalone_api_keys(
&self,
is_active: Option<bool>,
) -> Result<u64, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(
"SELECT COUNT(*) AS total FROM api_keys WHERE is_standalone = 1",
);
if let Some(is_active) = is_active {
builder.push(" AND is_active = ").push_bind(is_active);
}
let row = builder.build().fetch_one(&self.pool).await.map_sql_err()?;
Ok(row.try_get::<i64, _>("total").map_sql_err()?.max(0) as u64)
}
async fn summarize_export_api_keys_by_user_ids(
&self,
user_ids: &[String],
now_unix_secs: u64,
) -> Result<AuthApiKeyExportSummary, DataLayerError> {
if user_ids.is_empty() {
return Ok(AuthApiKeyExportSummary::default());
}
let mut builder = QueryBuilder::<MySql>::new(
r#"
SELECT
COUNT(*) AS total,
SUM(CASE WHEN is_active = 1 AND (expires_at IS NULL OR expires_at >=
"#,
);
builder.push_bind(now_unix_secs as i64);
builder.push(
r#") THEN 1 ELSE 0 END) AS active
FROM api_keys
"#,
);
push_in_clause(&mut builder, " WHERE user_id IN (", user_ids);
builder.push(" AND is_standalone = 0");
summarize_row(builder.build().fetch_one(&self.pool).await.map_sql_err()?)
}
async fn summarize_export_non_standalone_api_keys(
&self,
now_unix_secs: u64,
) -> Result<AuthApiKeyExportSummary, DataLayerError> {
summarize_api_keys(&self.pool, false, now_unix_secs).await
}
async fn summarize_export_standalone_api_keys(
&self,
now_unix_secs: u64,
) -> Result<AuthApiKeyExportSummary, DataLayerError> {
summarize_api_keys(&self.pool, true, now_unix_secs).await
}
async fn find_export_standalone_api_key_by_id(
&self,
api_key_id: &str,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(EXPORT_COLUMNS);
builder
.push(" WHERE api_keys.is_standalone = 1 AND api_keys.id = ")
.push_bind(api_key_id)
.push(" LIMIT 1");
Ok(self.fetch_export_rows(builder).await?.into_iter().next())
}
async fn list_export_standalone_api_keys(
&self,
) -> Result<Vec<StoredAuthApiKeyExportRecord>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(EXPORT_COLUMNS);
builder.push(" WHERE api_keys.is_standalone = 1 ORDER BY api_keys.id ASC");
self.fetch_export_rows(builder).await
}
}
#[async_trait]
impl AuthApiKeyWriteRepository for MysqlAuthApiKeyReadRepository {
async fn touch_last_used_at(&self, api_key_id: &str) -> Result<bool, DataLayerError> {
let now = current_unix_secs() as i64;
let rows_affected = sqlx::query(
r#"
UPDATE api_keys
SET last_used_at = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(now)
.bind(now)
.bind(api_key_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn create_user_api_key(
&self,
record: CreateUserApiKeyRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.create_api_key(CreateApiKeyInsertRecord {
user_id: record.user_id,
api_key_id: record.api_key_id,
key_hash: record.key_hash,
key_encrypted: record.key_encrypted,
name: record.name,
allowed_providers: record.allowed_providers,
allowed_api_formats: record.allowed_api_formats,
allowed_models: record.allowed_models,
rate_limit: Some(record.rate_limit),
concurrent_limit: record.concurrent_limit,
force_capabilities: record.force_capabilities,
is_active: record.is_active,
expires_at_unix_secs: record.expires_at_unix_secs,
auto_delete_on_expiry: record.auto_delete_on_expiry,
total_requests: record.total_requests,
total_tokens: record.total_tokens,
total_cost_usd: record.total_cost_usd,
is_standalone: false,
})
.await
}
async fn create_standalone_api_key(
&self,
record: CreateStandaloneApiKeyRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.create_api_key(CreateApiKeyInsertRecord {
user_id: record.user_id,
api_key_id: record.api_key_id,
key_hash: record.key_hash,
key_encrypted: record.key_encrypted,
name: record.name,
allowed_providers: record.allowed_providers,
allowed_api_formats: record.allowed_api_formats,
allowed_models: record.allowed_models,
rate_limit: record.rate_limit,
concurrent_limit: record.concurrent_limit,
force_capabilities: record.force_capabilities,
is_active: record.is_active,
expires_at_unix_secs: record.expires_at_unix_secs,
auto_delete_on_expiry: record.auto_delete_on_expiry,
total_requests: record.total_requests,
total_tokens: record.total_tokens,
total_cost_usd: record.total_cost_usd,
is_standalone: true,
})
.await
}
async fn update_user_api_key_basic(
&self,
record: UpdateUserApiKeyBasicRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let now = current_unix_secs() as i64;
sqlx::query(
r#"
UPDATE api_keys
SET name = COALESCE(?, name),
rate_limit = COALESCE(?, rate_limit),
concurrent_limit = COALESCE(?, concurrent_limit),
updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
"#,
)
.bind(record.name.as_deref())
.bind(record.rate_limit)
.bind(record.concurrent_limit)
.bind(now)
.bind(&record.api_key_id)
.bind(&record.user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_export_by_id(&record.api_key_id).await
}
async fn update_standalone_api_key_basic(
&self,
record: UpdateStandaloneApiKeyBasicRecord,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let now = current_unix_secs() as i64;
sqlx::query(
r#"
UPDATE api_keys
SET name = COALESCE(?, name),
rate_limit = CASE WHEN ? THEN ? ELSE rate_limit END,
concurrent_limit = CASE WHEN ? THEN ? ELSE concurrent_limit END,
allowed_providers = CASE WHEN ? THEN ? ELSE allowed_providers END,
allowed_api_formats = CASE WHEN ? THEN ? ELSE allowed_api_formats END,
allowed_models = CASE WHEN ? THEN ? ELSE allowed_models END,
expires_at = CASE WHEN ? THEN ? ELSE expires_at END,
auto_delete_on_expiry = CASE WHEN ? THEN ? ELSE auto_delete_on_expiry END,
updated_at = ?
WHERE id = ?
AND is_standalone = 1
"#,
)
.bind(record.name.as_deref())
.bind(record.rate_limit_present)
.bind(record.rate_limit)
.bind(record.concurrent_limit_present)
.bind(record.concurrent_limit)
.bind(record.allowed_providers.is_some())
.bind(json_string_from_nested_string_list(
&record.allowed_providers,
"api_keys.allowed_providers",
)?)
.bind(record.allowed_api_formats.is_some())
.bind(json_string_from_nested_string_list(
&record.allowed_api_formats,
"api_keys.allowed_api_formats",
)?)
.bind(record.allowed_models.is_some())
.bind(json_string_from_nested_string_list(
&record.allowed_models,
"api_keys.allowed_models",
)?)
.bind(record.expires_at_present)
.bind(optional_i64_from_u64(
record.expires_at_unix_secs,
"api_keys.expires_at",
)?)
.bind(record.auto_delete_on_expiry_present)
.bind(record.auto_delete_on_expiry)
.bind(now)
.bind(&record.api_key_id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_export_by_id(&record.api_key_id).await
}
async fn set_user_api_key_active(
&self,
user_id: &str,
api_key_id: &str,
is_active: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.set_active(api_key_id, Some(user_id), is_active, false)
.await
}
async fn set_standalone_api_key_active(
&self,
api_key_id: &str,
is_active: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
self.set_active(api_key_id, None, is_active, true).await
}
async fn set_user_api_key_locked(
&self,
user_id: &str,
api_key_id: &str,
is_locked: bool,
) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query(
r#"
UPDATE api_keys
SET is_locked = ?, updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
"#,
)
.bind(is_locked)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn set_user_api_key_allowed_providers(
&self,
user_id: &str,
api_key_id: &str,
allowed_providers: Option<Vec<String>>,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
sqlx::query(
r#"
UPDATE api_keys
SET allowed_providers = ?, updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
"#,
)
.bind(json_string_from_string_list(
allowed_providers.as_ref(),
"api_keys.allowed_providers",
)?)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_export_by_id(api_key_id).await
}
async fn set_user_api_key_force_capabilities(
&self,
user_id: &str,
api_key_id: &str,
force_capabilities: Option<serde_json::Value>,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
sqlx::query(
r#"
UPDATE api_keys
SET force_capabilities = ?, updated_at = ?
WHERE id = ?
AND user_id = ?
AND is_standalone = 0
"#,
)
.bind(optional_json_to_string(
&force_capabilities,
"api_keys.force_capabilities",
)?)
.bind(current_unix_secs() as i64)
.bind(api_key_id)
.bind(user_id)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_export_by_id(api_key_id).await
}
async fn delete_user_api_key(
&self,
user_id: &str,
api_key_id: &str,
) -> Result<bool, DataLayerError> {
self.delete_api_key(api_key_id, Some(user_id), false).await
}
async fn delete_standalone_api_key(&self, api_key_id: &str) -> Result<bool, DataLayerError> {
self.delete_api_key(api_key_id, None, true).await
}
}
impl MysqlAuthApiKeyReadRepository {
async fn set_active(
&self,
api_key_id: &str,
user_id: Option<&str>,
is_active: bool,
is_standalone: bool,
) -> Result<Option<StoredAuthApiKeyExportRecord>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new("UPDATE api_keys SET is_active = ");
builder
.push_bind(is_active)
.push(", updated_at = ")
.push_bind(current_unix_secs() as i64)
.push(" WHERE id = ")
.push_bind(api_key_id)
.push(" AND is_standalone = ")
.push_bind(is_standalone);
if let Some(user_id) = user_id {
builder.push(" AND user_id = ").push_bind(user_id);
}
builder.build().execute(&self.pool).await.map_sql_err()?;
self.reload_export_by_id(api_key_id).await
}
async fn delete_api_key(
&self,
api_key_id: &str,
user_id: Option<&str>,
is_standalone: bool,
) -> Result<bool, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new("DELETE FROM api_keys WHERE id = ");
builder
.push_bind(api_key_id)
.push(" AND is_standalone = ")
.push_bind(is_standalone);
if let Some(user_id) = user_id {
builder.push(" AND user_id = ").push_bind(user_id);
}
let rows_affected = builder
.build()
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
}
fn push_in_clause<'args>(
builder: &mut QueryBuilder<'args, MySql>,
prefix: &str,
values: &'args [String],
) {
builder.push(prefix);
{
let mut separated = builder.separated(", ");
for value in values {
separated.push_bind(value);
}
}
builder.push(")");
}
async fn summarize_api_keys(
pool: &MysqlPool,
is_standalone: bool,
now_unix_secs: u64,
) -> Result<AuthApiKeyExportSummary, DataLayerError> {
let row = sqlx::query(
r#"
SELECT
COUNT(*) AS total,
SUM(CASE WHEN is_active = 1 AND (expires_at IS NULL OR expires_at >= ?) THEN 1 ELSE 0 END) AS active
FROM api_keys
WHERE is_standalone = ?
"#,
)
.bind(now_unix_secs as i64)
.bind(is_standalone)
.fetch_one(pool)
.await
.map_sql_err()?;
summarize_row(row)
}
fn summarize_row(row: MySqlRow) -> Result<AuthApiKeyExportSummary, DataLayerError> {
Ok(AuthApiKeyExportSummary {
total: row.try_get::<i64, _>("total").map_sql_err()?.max(0) as u64,
active: row
.try_get::<Option<i64>, _>("active")
.map_sql_err()?
.unwrap_or(0)
.max(0) as u64,
})
}
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 current_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn i64_from_u64(value: u64, field_name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}")))
}
fn optional_i64_from_u64(
value: Option<u64>,
field_name: &str,
) -> Result<Option<i64>, DataLayerError> {
value
.map(|value| i64_from_u64(value, field_name))
.transpose()
}
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 json_string_from_string_list(
value: Option<&Vec<String>>,
field_name: &str,
) -> Result<Option<String>, DataLayerError> {
value
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"{field_name} contains unserializable string list: {err}"
))
})
})
.transpose()
}
fn json_string_from_nested_string_list(
value: &Option<Option<Vec<String>>>,
field_name: &str,
) -> Result<Option<String>, DataLayerError> {
match value {
Some(Some(values)) => json_string_from_string_list(Some(values), field_name),
Some(None) | None => Ok(None),
}
}
fn map_auth_api_key_snapshot_row(
row: &MySqlRow,
) -> Result<StoredAuthApiKeySnapshot, DataLayerError> {
let snapshot = StoredAuthApiKeySnapshot::new(
row.try_get("user_id").map_sql_err()?,
row.try_get("username").map_sql_err()?,
row.try_get("email").map_sql_err()?,
row.try_get("user_role").map_sql_err()?,
row.try_get("user_auth_source").map_sql_err()?,
row.try_get("user_is_active").map_sql_err()?,
row.try_get("user_is_deleted").map_sql_err()?,
optional_json_from_string(
row.try_get("user_allowed_providers").map_sql_err()?,
"users.allowed_providers",
)?,
optional_json_from_string(
row.try_get("user_allowed_api_formats").map_sql_err()?,
"users.allowed_api_formats",
)?,
optional_json_from_string(
row.try_get("user_allowed_models").map_sql_err()?,
"users.allowed_models",
)?,
row.try_get("api_key_id").map_sql_err()?,
row.try_get("api_key_name").map_sql_err()?,
row.try_get("api_key_is_active").map_sql_err()?,
row.try_get("api_key_is_locked").map_sql_err()?,
row.try_get("api_key_is_standalone").map_sql_err()?,
row.try_get("api_key_rate_limit").map_sql_err()?,
row.try_get("api_key_concurrent_limit").map_sql_err()?,
row.try_get("api_key_expires_at_unix_secs").map_sql_err()?,
optional_json_from_string(
row.try_get("api_key_allowed_providers").map_sql_err()?,
"api_keys.allowed_providers",
)?,
optional_json_from_string(
row.try_get("api_key_allowed_api_formats").map_sql_err()?,
"api_keys.allowed_api_formats",
)?,
optional_json_from_string(
row.try_get("api_key_allowed_models").map_sql_err()?,
"api_keys.allowed_models",
)?,
)?;
Ok(snapshot.with_user_rate_limit(row.try_get("user_rate_limit").map_sql_err()?))
}
fn map_auth_api_key_export_row(
row: &MySqlRow,
) -> Result<StoredAuthApiKeyExportRecord, DataLayerError> {
StoredAuthApiKeyExportRecord::new(
row.try_get("user_id").map_sql_err()?,
row.try_get("api_key_id").map_sql_err()?,
row.try_get("key_hash").map_sql_err()?,
row.try_get("key_encrypted").map_sql_err()?,
row.try_get("name").map_sql_err()?,
optional_json_from_string(
row.try_get("allowed_providers").map_sql_err()?,
"api_keys.allowed_providers",
)?,
optional_json_from_string(
row.try_get("allowed_api_formats").map_sql_err()?,
"api_keys.allowed_api_formats",
)?,
optional_json_from_string(
row.try_get("allowed_models").map_sql_err()?,
"api_keys.allowed_models",
)?,
row.try_get("rate_limit").map_sql_err()?,
row.try_get("concurrent_limit").map_sql_err()?,
optional_json_from_string(
row.try_get("force_capabilities").map_sql_err()?,
"api_keys.force_capabilities",
)?,
row.try_get("is_active").map_sql_err()?,
row.try_get("expires_at_unix_secs").map_sql_err()?,
row.try_get("auto_delete_on_expiry").map_sql_err()?,
row.try_get("total_requests").map_sql_err()?,
row.try_get("total_tokens").map_sql_err()?,
row.try_get("total_cost_usd").map_sql_err()?,
row.try_get("is_standalone").map_sql_err()?,
)
.and_then(|record| {
record.with_activity_timestamps(
row.try_get("last_used_at_unix_secs").map_sql_err()?,
row.try_get("created_at_unix_secs").map_sql_err()?,
row.try_get("updated_at_unix_secs").map_sql_err()?,
)
})
}
#[cfg(test)]
mod tests {
use super::MysqlAuthApiKeyReadRepository;
#[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 = MysqlAuthApiKeyReadRepository::new(pool);
}
}

View File

@@ -1451,7 +1451,7 @@ fn map_auth_api_key_export_row(
#[cfg(test)]
mod tests {
use super::{SqlxAuthApiKeySnapshotReadRepository, UPDATE_STANDALONE_API_KEY_BASIC_SQL};
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
#[test]
fn update_standalone_api_key_basic_sql_casts_json_case_values() {

File diff suppressed because it is too large Load Diff

View File

@@ -1,9 +1,13 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
mod types;
pub use memory::InMemoryAuthModuleReadRepository;
pub use sql::{SqlxAuthModuleReadRepository, SqlxAuthModuleRepository};
pub use mysql::{MysqlAuthModuleReadRepository, MysqlAuthModuleRepository};
pub use postgres::{SqlxAuthModuleReadRepository, SqlxAuthModuleRepository};
pub use sqlite::{SqliteAuthModuleReadRepository, SqliteAuthModuleRepository};
pub use types::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
StoredOAuthProviderModuleConfig,

View File

@@ -0,0 +1,248 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use super::types::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
StoredOAuthProviderModuleConfig,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const LIST_ENABLED_OAUTH_PROVIDERS_SQL: &str = r#"
SELECT
provider_type,
display_name,
client_id,
client_secret_encrypted,
redirect_uri
FROM oauth_providers
WHERE is_enabled = 1
ORDER BY provider_type ASC
"#;
const GET_LDAP_CONFIG_SQL: &str = r#"
SELECT
server_url,
bind_dn,
bind_password_encrypted,
base_dn,
user_search_filter,
username_attr,
email_attr,
display_name_attr,
is_enabled,
is_exclusive,
use_starttls,
connect_timeout
FROM ldap_configs
ORDER BY id ASC
LIMIT 1
"#;
#[derive(Debug, Clone)]
pub struct MysqlAuthModuleReadRepository {
pool: MysqlPool,
}
impl MysqlAuthModuleReadRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
}
#[derive(Debug, Clone)]
pub struct MysqlAuthModuleRepository {
pool: MysqlPool,
}
impl MysqlAuthModuleRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
}
#[async_trait]
impl AuthModuleReadRepository for MysqlAuthModuleReadRepository {
async fn list_enabled_oauth_providers(
&self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_oauth_row).collect()
}
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_ldap_row).transpose()
}
}
#[async_trait]
impl AuthModuleReadRepository for MysqlAuthModuleRepository {
async fn list_enabled_oauth_providers(
&self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_oauth_row).collect()
}
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_ldap_row).transpose()
}
}
#[async_trait]
impl AuthModuleWriteRepository for MysqlAuthModuleRepository {
async fn upsert_ldap_config(
&self,
config: &StoredLdapModuleConfig,
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let now = now_unix_secs();
let updated = sqlx::query(
r#"
UPDATE ldap_configs
SET
server_url = ?,
bind_dn = ?,
bind_password_encrypted = ?,
base_dn = ?,
user_search_filter = ?,
username_attr = ?,
email_attr = ?,
display_name_attr = ?,
is_enabled = ?,
is_exclusive = ?,
use_starttls = ?,
connect_timeout = ?,
updated_at = ?
WHERE id = (
SELECT id FROM (
SELECT id
FROM ldap_configs
ORDER BY id ASC
LIMIT 1
) selected_ldap_config
)
"#,
)
.bind(&config.server_url)
.bind(&config.bind_dn)
.bind(config.bind_password_encrypted.as_deref())
.bind(&config.base_dn)
.bind(config.user_search_filter.as_deref())
.bind(config.username_attr.as_deref())
.bind(config.email_attr.as_deref())
.bind(config.display_name_attr.as_deref())
.bind(config.is_enabled)
.bind(config.is_exclusive)
.bind(config.use_starttls)
.bind(config.connect_timeout)
.bind(now as i64)
.execute(&self.pool)
.await
.map_sql_err()?;
if updated.rows_affected() == 0 {
sqlx::query(
r#"
INSERT INTO ldap_configs (
server_url,
bind_dn,
bind_password_encrypted,
base_dn,
user_search_filter,
username_attr,
email_attr,
display_name_attr,
is_enabled,
is_exclusive,
use_starttls,
connect_timeout,
created_at,
updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&config.server_url)
.bind(&config.bind_dn)
.bind(config.bind_password_encrypted.as_deref())
.bind(&config.base_dn)
.bind(config.user_search_filter.as_deref())
.bind(config.username_attr.as_deref())
.bind(config.email_attr.as_deref())
.bind(config.display_name_attr.as_deref())
.bind(config.is_enabled)
.bind(config.is_exclusive)
.bind(config.use_starttls)
.bind(config.connect_timeout)
.bind(now as i64)
.bind(now as i64)
.execute(&self.pool)
.await
.map_sql_err()?;
}
self.get_ldap_config().await
}
}
fn now_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn map_oauth_row(row: &MySqlRow) -> Result<StoredOAuthProviderModuleConfig, DataLayerError> {
StoredOAuthProviderModuleConfig::new(
row.try_get("provider_type").map_sql_err()?,
row.try_get("display_name").map_sql_err()?,
row.try_get("client_id").map_sql_err()?,
row.try_get("client_secret_encrypted").map_sql_err()?,
row.try_get("redirect_uri").map_sql_err()?,
)
}
fn map_ldap_row(row: &MySqlRow) -> Result<StoredLdapModuleConfig, DataLayerError> {
Ok(StoredLdapModuleConfig {
server_url: row.try_get("server_url").map_sql_err()?,
bind_dn: row.try_get("bind_dn").map_sql_err()?,
bind_password_encrypted: row.try_get("bind_password_encrypted").map_sql_err()?,
base_dn: row.try_get("base_dn").map_sql_err()?,
user_search_filter: row.try_get("user_search_filter").map_sql_err()?,
username_attr: row.try_get("username_attr").map_sql_err()?,
email_attr: row.try_get("email_attr").map_sql_err()?,
display_name_attr: row.try_get("display_name_attr").map_sql_err()?,
is_enabled: row.try_get("is_enabled").map_sql_err()?,
is_exclusive: row.try_get("is_exclusive").map_sql_err()?,
use_starttls: row.try_get("use_starttls").map_sql_err()?,
connect_timeout: row.try_get("connect_timeout").map_sql_err()?,
})
}
#[cfg(test)]
mod tests {
use super::{MysqlAuthModuleReadRepository, MysqlAuthModuleRepository};
#[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 = MysqlAuthModuleReadRepository::new(pool.clone());
let _writable_repository = MysqlAuthModuleRepository::new(pool);
}
}

View File

@@ -278,7 +278,7 @@ fn map_ldap_row(row: &PgRow) -> Result<StoredLdapModuleConfig, DataLayerError> {
#[cfg(test)]
mod tests {
use super::{SqlxAuthModuleReadRepository, SqlxAuthModuleRepository};
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
#[tokio::test]
async fn repository_constructs_from_lazy_pool() {

View File

@@ -0,0 +1,304 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row};
use super::types::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
StoredOAuthProviderModuleConfig,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const LIST_ENABLED_OAUTH_PROVIDERS_SQL: &str = r#"
SELECT
provider_type,
display_name,
client_id,
client_secret_encrypted,
redirect_uri
FROM oauth_providers
WHERE is_enabled = 1
ORDER BY provider_type ASC
"#;
const GET_LDAP_CONFIG_SQL: &str = r#"
SELECT
server_url,
bind_dn,
bind_password_encrypted,
base_dn,
user_search_filter,
username_attr,
email_attr,
display_name_attr,
is_enabled,
is_exclusive,
use_starttls,
connect_timeout
FROM ldap_configs
ORDER BY id ASC
LIMIT 1
"#;
#[derive(Debug, Clone)]
pub struct SqliteAuthModuleReadRepository {
pool: SqlitePool,
}
impl SqliteAuthModuleReadRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
}
#[derive(Debug, Clone)]
pub struct SqliteAuthModuleRepository {
pool: SqlitePool,
}
impl SqliteAuthModuleRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
}
#[async_trait]
impl AuthModuleReadRepository for SqliteAuthModuleReadRepository {
async fn list_enabled_oauth_providers(
&self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_oauth_row).collect()
}
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_ldap_row).transpose()
}
}
#[async_trait]
impl AuthModuleReadRepository for SqliteAuthModuleRepository {
async fn list_enabled_oauth_providers(
&self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_oauth_row).collect()
}
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_ldap_row).transpose()
}
}
#[async_trait]
impl AuthModuleWriteRepository for SqliteAuthModuleRepository {
async fn upsert_ldap_config(
&self,
config: &StoredLdapModuleConfig,
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let now = now_unix_secs();
let updated = sqlx::query(
r#"
UPDATE ldap_configs
SET
server_url = ?,
bind_dn = ?,
bind_password_encrypted = ?,
base_dn = ?,
user_search_filter = ?,
username_attr = ?,
email_attr = ?,
display_name_attr = ?,
is_enabled = ?,
is_exclusive = ?,
use_starttls = ?,
connect_timeout = ?,
updated_at = ?
WHERE id = (
SELECT id
FROM ldap_configs
ORDER BY id ASC
LIMIT 1
)
"#,
)
.bind(&config.server_url)
.bind(&config.bind_dn)
.bind(config.bind_password_encrypted.as_deref())
.bind(&config.base_dn)
.bind(config.user_search_filter.as_deref())
.bind(config.username_attr.as_deref())
.bind(config.email_attr.as_deref())
.bind(config.display_name_attr.as_deref())
.bind(config.is_enabled)
.bind(config.is_exclusive)
.bind(config.use_starttls)
.bind(config.connect_timeout)
.bind(now as i64)
.execute(&self.pool)
.await
.map_sql_err()?;
if updated.rows_affected() == 0 {
sqlx::query(
r#"
INSERT INTO ldap_configs (
server_url,
bind_dn,
bind_password_encrypted,
base_dn,
user_search_filter,
username_attr,
email_attr,
display_name_attr,
is_enabled,
is_exclusive,
use_starttls,
connect_timeout,
created_at,
updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&config.server_url)
.bind(&config.bind_dn)
.bind(config.bind_password_encrypted.as_deref())
.bind(&config.base_dn)
.bind(config.user_search_filter.as_deref())
.bind(config.username_attr.as_deref())
.bind(config.email_attr.as_deref())
.bind(config.display_name_attr.as_deref())
.bind(config.is_enabled)
.bind(config.is_exclusive)
.bind(config.use_starttls)
.bind(config.connect_timeout)
.bind(now as i64)
.bind(now as i64)
.execute(&self.pool)
.await
.map_sql_err()?;
}
self.get_ldap_config().await
}
}
fn now_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn map_oauth_row(row: &SqliteRow) -> Result<StoredOAuthProviderModuleConfig, DataLayerError> {
StoredOAuthProviderModuleConfig::new(
row.try_get("provider_type").map_sql_err()?,
row.try_get("display_name").map_sql_err()?,
row.try_get("client_id").map_sql_err()?,
row.try_get("client_secret_encrypted").map_sql_err()?,
row.try_get("redirect_uri").map_sql_err()?,
)
}
fn map_ldap_row(row: &SqliteRow) -> Result<StoredLdapModuleConfig, DataLayerError> {
Ok(StoredLdapModuleConfig {
server_url: row.try_get("server_url").map_sql_err()?,
bind_dn: row.try_get("bind_dn").map_sql_err()?,
bind_password_encrypted: row.try_get("bind_password_encrypted").map_sql_err()?,
base_dn: row.try_get("base_dn").map_sql_err()?,
user_search_filter: row.try_get("user_search_filter").map_sql_err()?,
username_attr: row.try_get("username_attr").map_sql_err()?,
email_attr: row.try_get("email_attr").map_sql_err()?,
display_name_attr: row.try_get("display_name_attr").map_sql_err()?,
is_enabled: row.try_get("is_enabled").map_sql_err()?,
is_exclusive: row.try_get("is_exclusive").map_sql_err()?,
use_starttls: row.try_get("use_starttls").map_sql_err()?,
connect_timeout: row.try_get("connect_timeout").map_sql_err()?,
})
}
#[cfg(test)]
mod tests {
use super::SqliteAuthModuleRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
};
#[tokio::test]
async fn sqlite_repository_reads_and_writes_auth_module_configs() {
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");
sqlx::query(
r#"
INSERT INTO oauth_providers (
provider_type, display_name, client_id, redirect_uri, frontend_callback_url,
is_enabled, created_at, updated_at
) VALUES
('github', 'GitHub', 'github-client', 'https://github.example.com/callback',
'https://frontend.example.com/callback', 1, 1, 1),
('disabled', 'Disabled', 'disabled-client', 'https://disabled.example.com/callback',
'https://frontend.example.com/callback', 0, 1, 1)
"#,
)
.execute(&pool)
.await
.expect("oauth providers should seed");
let repository = SqliteAuthModuleRepository::new(pool);
let oauth = repository
.list_enabled_oauth_providers()
.await
.expect("oauth providers should load");
assert_eq!(oauth.len(), 1);
assert_eq!(oauth[0].provider_type, "github");
let ldap = StoredLdapModuleConfig {
server_url: "ldaps://ldap.example.com".to_string(),
bind_dn: "cn=admin,dc=example,dc=com".to_string(),
bind_password_encrypted: Some("encrypted-password".to_string()),
base_dn: "dc=example,dc=com".to_string(),
user_search_filter: Some("(uid={username})".to_string()),
username_attr: Some("uid".to_string()),
email_attr: Some("mail".to_string()),
display_name_attr: Some("displayName".to_string()),
is_enabled: true,
is_exclusive: false,
use_starttls: true,
connect_timeout: Some(10),
};
let stored = repository
.upsert_ldap_config(&ldap)
.await
.expect("ldap should upsert")
.expect("ldap should be returned");
assert_eq!(stored.server_url, "ldaps://ldap.example.com");
let updated = repository
.upsert_ldap_config(&StoredLdapModuleConfig {
server_url: "ldap://ldap.example.com".to_string(),
..ldap
})
.await
.expect("ldap should update")
.expect("ldap should be returned");
assert_eq!(updated.server_url, "ldap://ldap.example.com");
}
}

View File

@@ -1,11 +1,15 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::billing::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingPresetApplyResult,
AdminBillingRuleRecord, AdminBillingRuleWriteInput, BillingReadRepository,
StoredBillingModelContext,
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
BillingReadRepository, StoredBillingModelContext,
};
pub use memory::InMemoryBillingReadRepository;
pub use sql::SqlxBillingReadRepository;
pub use mysql::MysqlBillingReadRepository;
pub use postgres::SqlxBillingReadRepository;
pub use sqlite::SqliteBillingReadRepository;

View File

@@ -0,0 +1,870 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use super::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
BillingReadRepository, StoredBillingModelContext,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const MODEL_CONTEXT_COLUMNS: &str = r#"
SELECT
p.id AS provider_id,
p.billing_type AS provider_billing_type,
pak.id AS provider_api_key_id,
pak.rate_multipliers AS provider_api_key_rate_multipliers,
pak.cache_ttl_minutes AS provider_api_key_cache_ttl_minutes,
gm.id AS global_model_id,
gm.name AS global_model_name,
gm.config AS global_model_config,
gm.default_price_per_request AS default_price_per_request,
gm.default_tiered_pricing AS default_tiered_pricing,
m.id AS model_id,
m.provider_model_name AS model_provider_model_name,
m.config AS model_config,
m.price_per_request AS model_price_per_request,
m.tiered_pricing AS model_tiered_pricing,
m.provider_model_mappings AS provider_model_mappings,
m.is_available AS model_is_available,
m.created_at AS model_created_at
FROM providers p
"#;
#[derive(Debug, Clone)]
pub struct MysqlBillingReadRepository {
pool: MysqlPool,
}
impl MysqlBillingReadRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
}
#[async_trait]
impl BillingReadRepository for MysqlBillingReadRepository {
async fn find_model_context(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
let rows = sqlx::query(&format!(
r#"
{MODEL_CONTEXT_COLUMNS}
INNER JOIN global_models gm
ON gm.is_active = 1
LEFT JOIN models m
ON m.global_model_id = gm.id
AND m.provider_id = p.id
AND m.is_active = 1
LEFT JOIN provider_api_keys pak
ON pak.id = ?
AND pak.provider_id = p.id
WHERE p.id = ?
AND (
gm.name = ?
OR m.provider_model_name = ?
OR m.provider_model_mappings IS NOT NULL
)
"#
))
.bind(provider_api_key_id)
.bind(provider_id)
.bind(global_model_name)
.bind(global_model_name)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter()
.filter_map(|row| match_rank(row, global_model_name).transpose())
.collect::<Result<Vec<_>, _>>()?
.into_iter()
.min_by_key(|candidate| {
(
candidate.rank,
!candidate.is_available,
candidate.pricing_rank,
candidate.created_at,
)
})
.map(|candidate| candidate.context)
.transpose()
}
async fn find_model_context_by_model_id(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
model_id: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
let row = sqlx::query(&format!(
r#"
{MODEL_CONTEXT_COLUMNS}
INNER JOIN models m
ON m.id = ?
AND m.provider_id = p.id
AND m.is_active = 1
INNER JOIN global_models gm
ON gm.id = m.global_model_id
AND gm.is_active = 1
LEFT JOIN provider_api_keys pak
ON pak.id = ?
AND pak.provider_id = p.id
WHERE p.id = ?
LIMIT 1
"#
))
.bind(model_id)
.bind(provider_api_key_id)
.bind(provider_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_row).transpose()
}
async fn admin_billing_enabled_default_value_exists(
&self,
api_format: &str,
task_type: &str,
dimension_name: &str,
existing_id: Option<&str>,
) -> Result<Option<bool>, DataLayerError> {
let row = sqlx::query(
r#"
SELECT COUNT(*) AS total
FROM dimension_collectors
WHERE api_format = ?
AND task_type = ?
AND dimension_name = ?
AND is_enabled = 1
AND default_value IS NOT NULL
AND (? IS NULL OR id <> ?)
"#,
)
.bind(api_format)
.bind(task_type)
.bind(dimension_name)
.bind(existing_id)
.bind(existing_id)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
Ok(Some(read_count_mysql(&row)? > 0))
}
async fn create_admin_billing_rule(
&self,
input: &AdminBillingRuleWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingRuleRecord>, DataLayerError> {
let id = uuid::Uuid::new_v4().to_string();
let now = current_unix_secs_i64();
let result = sqlx::query(
r#"
INSERT INTO billing_rules (
id, name, task_type, global_model_id, model_id, expression, variables,
dimension_mappings, is_enabled, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id)
.bind(&input.name)
.bind(&input.task_type)
.bind(input.global_model_id.as_deref())
.bind(input.model_id.as_deref())
.bind(&input.expression)
.bind(json_to_string(&input.variables)?)
.bind(json_to_string(&input.dimension_mappings)?)
.bind(input.is_enabled)
.bind(now)
.bind(now)
.execute(&self.pool)
.await;
if let Err(err) = result {
return Ok(AdminBillingMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)));
}
match find_admin_billing_rule_mysql(&self.pool, &id).await? {
Some(record) => Ok(AdminBillingMutationOutcome::Applied(record)),
None => Err(DataLayerError::UnexpectedValue(
"created billing rule missing".to_string(),
)),
}
}
async fn list_admin_billing_rules(
&self,
task_type: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<Option<(Vec<AdminBillingRuleRecord>, u64)>, DataLayerError> {
let total_row = sqlx::query(
r#"
SELECT COUNT(*) AS total
FROM billing_rules
WHERE (? IS NULL OR task_type = ?)
AND (? IS NULL OR is_enabled = ?)
"#,
)
.bind(task_type)
.bind(task_type)
.bind(is_enabled)
.bind(is_enabled)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let total = read_count_mysql(&total_row)?;
let offset = u64::from(page.saturating_sub(1) * page_size);
let rows = sqlx::query(
r#"
SELECT
id, name, task_type, global_model_id, model_id, expression, variables,
dimension_mappings, is_enabled, created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM billing_rules
WHERE (? IS NULL OR task_type = ?)
AND (? IS NULL OR is_enabled = ?)
ORDER BY updated_at DESC, id DESC
LIMIT ? OFFSET ?
"#,
)
.bind(task_type)
.bind(task_type)
.bind(is_enabled)
.bind(is_enabled)
.bind(i64::from(page_size))
.bind(
i64::try_from(offset)
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?,
)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows
.iter()
.map(map_admin_billing_rule_mysql)
.collect::<Result<Vec<_>, _>>()?;
Ok(Some((items, total)))
}
async fn find_admin_billing_rule(
&self,
rule_id: &str,
) -> Result<Option<AdminBillingRuleRecord>, DataLayerError> {
find_admin_billing_rule_mysql(&self.pool, rule_id).await
}
async fn update_admin_billing_rule(
&self,
rule_id: &str,
input: &AdminBillingRuleWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingRuleRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE billing_rules
SET name = ?,
task_type = ?,
global_model_id = ?,
model_id = ?,
expression = ?,
variables = ?,
dimension_mappings = ?,
is_enabled = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(&input.name)
.bind(&input.task_type)
.bind(input.global_model_id.as_deref())
.bind(input.model_id.as_deref())
.bind(&input.expression)
.bind(json_to_string(&input.variables)?)
.bind(json_to_string(&input.dimension_mappings)?)
.bind(input.is_enabled)
.bind(current_unix_secs_i64())
.bind(rule_id)
.execute(&self.pool)
.await;
let affected = match result {
Ok(result) => result.rows_affected(),
Err(err) => {
return Ok(AdminBillingMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
};
if affected == 0 {
return Ok(AdminBillingMutationOutcome::NotFound);
}
match find_admin_billing_rule_mysql(&self.pool, rule_id).await? {
Some(record) => Ok(AdminBillingMutationOutcome::Applied(record)),
None => Ok(AdminBillingMutationOutcome::NotFound),
}
}
async fn create_admin_billing_collector(
&self,
input: &AdminBillingCollectorWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingCollectorRecord>, DataLayerError> {
let id = uuid::Uuid::new_v4().to_string();
let now = current_unix_secs_i64();
let result = sqlx::query(
r#"
INSERT INTO dimension_collectors (
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&id)
.bind(&input.api_format)
.bind(&input.task_type)
.bind(&input.dimension_name)
.bind(&input.source_type)
.bind(input.source_path.as_deref())
.bind(&input.value_type)
.bind(input.transform_expression.as_deref())
.bind(input.default_value.as_deref())
.bind(input.priority)
.bind(input.is_enabled)
.bind(now)
.bind(now)
.execute(&self.pool)
.await;
if let Err(err) = result {
return Ok(AdminBillingMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)));
}
match find_admin_billing_collector_mysql(&self.pool, &id).await? {
Some(record) => Ok(AdminBillingMutationOutcome::Applied(record)),
None => Err(DataLayerError::UnexpectedValue(
"created billing collector missing".to_string(),
)),
}
}
async fn list_admin_billing_collectors(
&self,
api_format: Option<&str>,
task_type: Option<&str>,
dimension_name: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<Option<(Vec<AdminBillingCollectorRecord>, u64)>, DataLayerError> {
let total_row = sqlx::query(
r#"
SELECT COUNT(*) AS total
FROM dimension_collectors
WHERE (? IS NULL OR api_format = ?)
AND (? IS NULL OR task_type = ?)
AND (? IS NULL OR dimension_name = ?)
AND (? IS NULL OR is_enabled = ?)
"#,
)
.bind(api_format)
.bind(api_format)
.bind(task_type)
.bind(task_type)
.bind(dimension_name)
.bind(dimension_name)
.bind(is_enabled)
.bind(is_enabled)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let total = read_count_mysql(&total_row)?;
let offset = u64::from(page.saturating_sub(1) * page_size);
let rows = sqlx::query(
r#"
SELECT
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled,
created_at AS created_at_unix_ms, updated_at AS updated_at_unix_secs
FROM dimension_collectors
WHERE (? IS NULL OR api_format = ?)
AND (? IS NULL OR task_type = ?)
AND (? IS NULL OR dimension_name = ?)
AND (? IS NULL OR is_enabled = ?)
ORDER BY updated_at DESC, priority DESC, id ASC
LIMIT ? OFFSET ?
"#,
)
.bind(api_format)
.bind(api_format)
.bind(task_type)
.bind(task_type)
.bind(dimension_name)
.bind(dimension_name)
.bind(is_enabled)
.bind(is_enabled)
.bind(i64::from(page_size))
.bind(
i64::try_from(offset)
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?,
)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows
.iter()
.map(map_admin_billing_collector_mysql)
.collect::<Result<Vec<_>, _>>()?;
Ok(Some((items, total)))
}
async fn find_admin_billing_collector(
&self,
collector_id: &str,
) -> Result<Option<AdminBillingCollectorRecord>, DataLayerError> {
find_admin_billing_collector_mysql(&self.pool, collector_id).await
}
async fn update_admin_billing_collector(
&self,
collector_id: &str,
input: &AdminBillingCollectorWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingCollectorRecord>, DataLayerError> {
let result = sqlx::query(
r#"
UPDATE dimension_collectors
SET api_format = ?,
task_type = ?,
dimension_name = ?,
source_type = ?,
source_path = ?,
value_type = ?,
transform_expression = ?,
default_value = ?,
priority = ?,
is_enabled = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(&input.api_format)
.bind(&input.task_type)
.bind(&input.dimension_name)
.bind(&input.source_type)
.bind(input.source_path.as_deref())
.bind(&input.value_type)
.bind(input.transform_expression.as_deref())
.bind(input.default_value.as_deref())
.bind(input.priority)
.bind(input.is_enabled)
.bind(current_unix_secs_i64())
.bind(collector_id)
.execute(&self.pool)
.await;
let affected = match result {
Ok(result) => result.rows_affected(),
Err(err) => {
return Ok(AdminBillingMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
};
if affected == 0 {
return Ok(AdminBillingMutationOutcome::NotFound);
}
match find_admin_billing_collector_mysql(&self.pool, collector_id).await? {
Some(record) => Ok(AdminBillingMutationOutcome::Applied(record)),
None => Ok(AdminBillingMutationOutcome::NotFound),
}
}
async fn apply_admin_billing_preset(
&self,
preset: &str,
mode: &str,
collectors: &[AdminBillingCollectorWriteInput],
) -> Result<AdminBillingMutationOutcome<AdminBillingPresetApplyResult>, DataLayerError> {
let mut created = 0_u64;
let mut updated = 0_u64;
let mut skipped = 0_u64;
let mut errors = Vec::new();
for collector in collectors {
let existing_id = match sqlx::query_scalar::<_, String>(
r#"
SELECT id
FROM dimension_collectors
WHERE api_format = ?
AND task_type = ?
AND dimension_name = ?
AND priority = ?
AND is_enabled = 1
LIMIT 1
"#,
)
.bind(&collector.api_format)
.bind(&collector.task_type)
.bind(&collector.dimension_name)
.bind(collector.priority)
.fetch_optional(&self.pool)
.await
{
Ok(value) => value,
Err(err) => {
errors.push(format!(
"Failed to query collector: api_format={} task_type={} dim={}: {}",
collector.api_format, collector.task_type, collector.dimension_name, err
));
continue;
}
};
if let Some(existing_id) = existing_id {
if mode == "overwrite" {
match sqlx::query(
r#"
UPDATE dimension_collectors
SET source_type = ?,
source_path = ?,
value_type = ?,
transform_expression = ?,
default_value = ?,
is_enabled = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(&collector.source_type)
.bind(collector.source_path.as_deref())
.bind(&collector.value_type)
.bind(collector.transform_expression.as_deref())
.bind(collector.default_value.as_deref())
.bind(collector.is_enabled)
.bind(current_unix_secs_i64())
.bind(&existing_id)
.execute(&self.pool)
.await
{
Ok(_) => updated += 1,
Err(err) => errors.push(format!(
"Failed to update collector {}: {}",
existing_id, err
)),
}
} else {
skipped += 1;
}
continue;
}
let id = uuid::Uuid::new_v4().to_string();
let now = current_unix_secs_i64();
match sqlx::query(
r#"
INSERT INTO dimension_collectors (
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled, created_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(id)
.bind(&collector.api_format)
.bind(&collector.task_type)
.bind(&collector.dimension_name)
.bind(&collector.source_type)
.bind(collector.source_path.as_deref())
.bind(&collector.value_type)
.bind(collector.transform_expression.as_deref())
.bind(collector.default_value.as_deref())
.bind(collector.priority)
.bind(collector.is_enabled)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
{
Ok(_) => created += 1,
Err(err) => errors.push(format!(
"Failed to create collector: api_format={} task_type={} dim={}: {}",
collector.api_format, collector.task_type, collector.dimension_name, err
)),
}
}
Ok(AdminBillingMutationOutcome::Applied(
AdminBillingPresetApplyResult {
preset: preset.to_string(),
mode: mode.to_string(),
created,
updated,
skipped,
errors,
},
))
}
}
struct RankedContext {
rank: u8,
is_available: bool,
pricing_rank: u8,
created_at: i64,
context: Result<StoredBillingModelContext, DataLayerError>,
}
fn match_rank(
row: &MySqlRow,
requested_model: &str,
) -> Result<Option<RankedContext>, DataLayerError> {
let provider_model_name: Option<String> =
row.try_get("model_provider_model_name").map_sql_err()?;
let global_model_name: String = row.try_get("global_model_name").map_sql_err()?;
let mappings: Option<String> = row.try_get("provider_model_mappings").ok().flatten();
let rank = if provider_model_name.as_deref() == Some(requested_model) {
0
} else if mappings
.as_deref()
.is_some_and(|mappings| provider_model_mappings_match(mappings, requested_model))
{
1
} else if global_model_name == requested_model {
2
} else {
return Ok(None);
};
let has_model_price = row
.try_get::<Option<f64>, _>("model_price_per_request")
.map_sql_err()?
.is_some()
|| row
.try_get::<Option<String>, _>("model_tiered_pricing")
.ok()
.flatten()
.is_some();
let has_default_price = row
.try_get::<Option<f64>, _>("default_price_per_request")
.map_sql_err()?
.is_some()
|| row
.try_get::<Option<String>, _>("default_tiered_pricing")
.ok()
.flatten()
.is_some();
let pricing_rank = if has_model_price {
0
} else if has_default_price {
1
} else {
2
};
Ok(Some(RankedContext {
rank,
is_available: row
.try_get::<Option<bool>, _>("model_is_available")
.map_sql_err()?
.unwrap_or(false),
pricing_rank,
created_at: row
.try_get::<Option<i64>, _>("model_created_at")
.map_sql_err()?
.unwrap_or(i64::MAX),
context: map_row(row),
}))
}
fn provider_model_mappings_match(raw: &str, requested_model: &str) -> bool {
let Ok(value) = serde_json::from_str::<serde_json::Value>(raw) else {
return raw == requested_model;
};
json_mapping_matches(&value, requested_model)
}
fn json_mapping_matches(value: &serde_json::Value, requested_model: &str) -> bool {
match value {
serde_json::Value::String(value) => value == requested_model,
serde_json::Value::Array(values) => values
.iter()
.any(|value| json_mapping_matches(value, requested_model)),
serde_json::Value::Object(map) => map
.get("name")
.is_some_and(|value| json_mapping_matches(value, requested_model)),
_ => false,
}
}
fn map_row(row: &MySqlRow) -> Result<StoredBillingModelContext, DataLayerError> {
StoredBillingModelContext::new(
row.try_get("provider_id").map_sql_err()?,
row.try_get("provider_billing_type").map_sql_err()?,
row.try_get("provider_api_key_id").map_sql_err()?,
parse_json(
row.try_get("provider_api_key_rate_multipliers")
.ok()
.flatten(),
)?,
row.try_get::<Option<i64>, _>("provider_api_key_cache_ttl_minutes")
.map_sql_err()?,
row.try_get("global_model_id").map_sql_err()?,
row.try_get("global_model_name").map_sql_err()?,
parse_json(row.try_get("global_model_config").ok().flatten())?,
row.try_get("default_price_per_request").map_sql_err()?,
parse_json(row.try_get("default_tiered_pricing").ok().flatten())?,
row.try_get("model_id").map_sql_err()?,
row.try_get("model_provider_model_name").map_sql_err()?,
parse_json(row.try_get("model_config").ok().flatten())?,
row.try_get("model_price_per_request").map_sql_err()?,
parse_json(row.try_get("model_tiered_pricing").ok().flatten())?,
)
}
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!("billing JSON field is invalid: {err}"))
})
})
.transpose()
}
fn current_unix_secs_i64() -> i64 {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs() as i64
}
fn json_to_string(value: &serde_json::Value) -> Result<String, DataLayerError> {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("billing JSON encode failed: {err}"))
})
}
fn read_count_mysql(row: &MySqlRow) -> Result<u64, DataLayerError> {
Ok(row.try_get::<i64, _>("total").map_sql_err()?.max(0) as u64)
}
async fn find_admin_billing_rule_mysql(
pool: &MysqlPool,
rule_id: &str,
) -> Result<Option<AdminBillingRuleRecord>, DataLayerError> {
let row = sqlx::query(
r#"
SELECT
id, name, task_type, global_model_id, model_id, expression, variables,
dimension_mappings, is_enabled, created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM billing_rules
WHERE id = ?
"#,
)
.bind(rule_id)
.fetch_optional(pool)
.await
.map_sql_err()?;
row.as_ref().map(map_admin_billing_rule_mysql).transpose()
}
fn map_admin_billing_rule_mysql(row: &MySqlRow) -> Result<AdminBillingRuleRecord, DataLayerError> {
Ok(AdminBillingRuleRecord {
id: row.try_get("id").map_sql_err()?,
name: row.try_get("name").map_sql_err()?,
task_type: row.try_get("task_type").map_sql_err()?,
global_model_id: row.try_get("global_model_id").map_sql_err()?,
model_id: row.try_get("model_id").map_sql_err()?,
expression: row.try_get("expression").map_sql_err()?,
variables: parse_required_json(row.try_get("variables").map_sql_err()?)?,
dimension_mappings: parse_required_json(row.try_get("dimension_mappings").map_sql_err()?)?,
is_enabled: row.try_get("is_enabled").map_sql_err()?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_sql_err()?
.max(0) as u64,
updated_at_unix_secs: row
.try_get::<i64, _>("updated_at_unix_secs")
.map_sql_err()?
.max(0) as u64,
})
}
async fn find_admin_billing_collector_mysql(
pool: &MysqlPool,
collector_id: &str,
) -> Result<Option<AdminBillingCollectorRecord>, DataLayerError> {
let row = sqlx::query(
r#"
SELECT
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled,
created_at AS created_at_unix_ms, updated_at AS updated_at_unix_secs
FROM dimension_collectors
WHERE id = ?
"#,
)
.bind(collector_id)
.fetch_optional(pool)
.await
.map_sql_err()?;
row.as_ref()
.map(map_admin_billing_collector_mysql)
.transpose()
}
fn map_admin_billing_collector_mysql(
row: &MySqlRow,
) -> Result<AdminBillingCollectorRecord, DataLayerError> {
Ok(AdminBillingCollectorRecord {
id: row.try_get("id").map_sql_err()?,
api_format: row.try_get("api_format").map_sql_err()?,
task_type: row.try_get("task_type").map_sql_err()?,
dimension_name: row.try_get("dimension_name").map_sql_err()?,
source_type: row.try_get("source_type").map_sql_err()?,
source_path: row.try_get("source_path").map_sql_err()?,
value_type: row.try_get("value_type").map_sql_err()?,
transform_expression: row.try_get("transform_expression").map_sql_err()?,
default_value: row.try_get("default_value").map_sql_err()?,
priority: row.try_get("priority").map_sql_err()?,
is_enabled: row.try_get("is_enabled").map_sql_err()?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_sql_err()?
.max(0) as u64,
updated_at_unix_secs: row
.try_get::<i64, _>("updated_at_unix_secs")
.map_sql_err()?
.max(0) as u64,
})
}
fn parse_required_json(raw: String) -> Result<serde_json::Value, DataLayerError> {
serde_json::from_str(&raw).map_err(|err| {
DataLayerError::UnexpectedValue(format!("billing JSON field is invalid: {err}"))
})
}
#[cfg(test)]
mod tests {
use super::MysqlBillingReadRepository;
#[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 = MysqlBillingReadRepository::new(pool);
}
}

View File

@@ -0,0 +1,799 @@
use async_trait::async_trait;
use sqlx::{PgPool, Row};
use super::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
BillingReadRepository, StoredBillingModelContext,
};
use crate::{error::SqlxResultExt, DataLayerError};
const FIND_MODEL_CONTEXT_SQL: &str = r#"
SELECT
p.id AS provider_id,
CAST(p.billing_type AS TEXT) AS provider_billing_type,
pak.id AS provider_api_key_id,
pak.rate_multipliers AS provider_api_key_rate_multipliers,
pak.cache_ttl_minutes AS provider_api_key_cache_ttl_minutes,
gm.id AS global_model_id,
gm.name AS global_model_name,
gm.config AS global_model_config,
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS default_price_per_request,
gm.default_tiered_pricing AS default_tiered_pricing,
m.id AS model_id,
m.provider_model_name AS model_provider_model_name,
m.config AS model_config,
CAST(m.price_per_request AS DOUBLE PRECISION) AS model_price_per_request,
m.tiered_pricing AS model_tiered_pricing
FROM providers p
INNER JOIN global_models gm
ON gm.is_active = TRUE
LEFT JOIN models m
ON m.global_model_id = gm.id
AND m.provider_id = p.id
AND m.is_active = TRUE
LEFT JOIN provider_api_keys pak
ON pak.id = $3
AND pak.provider_id = p.id
WHERE p.id = $1
AND (
gm.name = $2
OR m.provider_model_name = $2
OR (
m.provider_model_mappings IS NOT NULL
AND (
m.provider_model_mappings @> jsonb_build_array(jsonb_build_object('name', $2::TEXT))
OR m.provider_model_mappings @> jsonb_build_array(to_jsonb($2::TEXT))
OR m.provider_model_mappings @> jsonb_build_object('name', $2::TEXT)
OR m.provider_model_mappings = to_jsonb($2::TEXT)
)
)
)
ORDER BY
CASE
WHEN m.provider_model_name = $2 THEN 0
WHEN m.provider_model_mappings IS NOT NULL
AND (
m.provider_model_mappings @> jsonb_build_array(jsonb_build_object('name', $2::TEXT))
OR m.provider_model_mappings @> jsonb_build_array(to_jsonb($2::TEXT))
OR m.provider_model_mappings @> jsonb_build_object('name', $2::TEXT)
OR m.provider_model_mappings = to_jsonb($2::TEXT)
) THEN 1
WHEN gm.name = $2 THEN 2
ELSE 3
END ASC,
COALESCE(m.is_available, FALSE) DESC,
CASE
WHEN m.tiered_pricing IS NOT NULL OR m.price_per_request IS NOT NULL THEN 0
WHEN gm.default_tiered_pricing IS NOT NULL OR gm.default_price_per_request IS NOT NULL THEN 1
ELSE 2
END ASC,
m.created_at ASC
LIMIT 1
"#;
const FIND_MODEL_CONTEXT_BY_MODEL_ID_SQL: &str = r#"
SELECT
p.id AS provider_id,
CAST(p.billing_type AS TEXT) AS provider_billing_type,
pak.id AS provider_api_key_id,
pak.rate_multipliers AS provider_api_key_rate_multipliers,
pak.cache_ttl_minutes AS provider_api_key_cache_ttl_minutes,
gm.id AS global_model_id,
gm.name AS global_model_name,
gm.config AS global_model_config,
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS default_price_per_request,
gm.default_tiered_pricing AS default_tiered_pricing,
m.id AS model_id,
m.provider_model_name AS model_provider_model_name,
m.config AS model_config,
CAST(m.price_per_request AS DOUBLE PRECISION) AS model_price_per_request,
m.tiered_pricing AS model_tiered_pricing
FROM providers p
INNER JOIN models m
ON m.id = $2
AND m.provider_id = p.id
AND m.is_active = TRUE
INNER JOIN global_models gm
ON gm.id = m.global_model_id
AND gm.is_active = TRUE
LEFT JOIN provider_api_keys pak
ON pak.id = $3
AND pak.provider_id = p.id
WHERE p.id = $1
LIMIT 1
"#;
#[derive(Debug, Clone)]
pub struct SqlxBillingReadRepository {
pool: PgPool,
}
impl SqlxBillingReadRepository {
pub fn new(pool: PgPool) -> Self {
Self { pool }
}
pub async fn find_model_context(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
let row = sqlx::query(FIND_MODEL_CONTEXT_SQL)
.bind(provider_id)
.bind(global_model_name)
.bind(provider_api_key_id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_row).transpose()
}
pub async fn find_model_context_by_model_id(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
model_id: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
let row = sqlx::query(FIND_MODEL_CONTEXT_BY_MODEL_ID_SQL)
.bind(provider_id)
.bind(model_id)
.bind(provider_api_key_id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_row).transpose()
}
}
#[async_trait]
impl BillingReadRepository for SqlxBillingReadRepository {
async fn find_model_context(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
Self::find_model_context(self, provider_id, provider_api_key_id, global_model_name).await
}
async fn find_model_context_by_model_id(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
model_id: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
Self::find_model_context_by_model_id(self, provider_id, provider_api_key_id, model_id).await
}
async fn admin_billing_enabled_default_value_exists(
&self,
api_format: &str,
task_type: &str,
dimension_name: &str,
existing_id: Option<&str>,
) -> Result<Option<bool>, DataLayerError> {
let exists = sqlx::query_scalar::<_, bool>(
r#"
SELECT EXISTS(
SELECT 1
FROM dimension_collectors
WHERE api_format = $1
AND task_type = $2
AND dimension_name = $3
AND is_enabled = TRUE
AND default_value IS NOT NULL
AND ($4::TEXT IS NULL OR id <> $4)
)
"#,
)
.bind(api_format)
.bind(task_type)
.bind(dimension_name)
.bind(existing_id)
.fetch_one(&self.pool)
.await
.map_postgres_err()?;
Ok(Some(exists))
}
async fn create_admin_billing_rule(
&self,
input: &AdminBillingRuleWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingRuleRecord>, DataLayerError> {
let rule_id = uuid::Uuid::new_v4().to_string();
let row = match sqlx::query(
r#"
INSERT INTO billing_rules (
id, name, task_type, global_model_id, model_id, expression, variables,
dimension_mappings, is_enabled, created_at, updated_at
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, NOW(), NOW())
RETURNING
id, name, task_type, global_model_id, model_id, expression, variables,
dimension_mappings, is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
"#,
)
.bind(&rule_id)
.bind(&input.name)
.bind(&input.task_type)
.bind(input.global_model_id.as_deref())
.bind(input.model_id.as_deref())
.bind(&input.expression)
.bind(&input.variables)
.bind(&input.dimension_mappings)
.bind(input.is_enabled)
.fetch_one(&self.pool)
.await
{
Ok(row) => row,
Err(sqlx::Error::Database(err)) => {
return Ok(AdminBillingMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
Err(err) => return Err(DataLayerError::postgres(err)),
};
Ok(AdminBillingMutationOutcome::Applied(
map_admin_billing_rule_row(&row)?,
))
}
async fn list_admin_billing_rules(
&self,
task_type: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<Option<(Vec<AdminBillingRuleRecord>, u64)>, DataLayerError> {
let total = read_count(
sqlx::query(
r#"
SELECT COUNT(*) AS total
FROM billing_rules
WHERE ($1::TEXT IS NULL OR task_type = $1)
AND ($2::BOOL IS NULL OR is_enabled = $2)
"#,
)
.bind(task_type)
.bind(is_enabled)
.fetch_one(&self.pool)
.await
.map_postgres_err()?,
)?;
let offset = u64::from(page.saturating_sub(1) * page_size);
let rows = sqlx::query(
r#"
SELECT
id, name, task_type, global_model_id, model_id, expression, variables,
dimension_mappings, is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
FROM billing_rules
WHERE ($1::TEXT IS NULL OR task_type = $1)
AND ($2::BOOL IS NULL OR is_enabled = $2)
ORDER BY updated_at DESC
OFFSET $3
LIMIT $4
"#,
)
.bind(task_type)
.bind(is_enabled)
.bind(
i64::try_from(offset)
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?,
)
.bind(i64::from(page_size))
.fetch_all(&self.pool)
.await
.map_postgres_err()?;
let items = rows
.iter()
.map(map_admin_billing_rule_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(Some((items, total)))
}
async fn find_admin_billing_rule(
&self,
rule_id: &str,
) -> Result<Option<AdminBillingRuleRecord>, DataLayerError> {
let row = sqlx::query(
r#"
SELECT
id, name, task_type, global_model_id, model_id, expression, variables,
dimension_mappings, is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
FROM billing_rules
WHERE id = $1
"#,
)
.bind(rule_id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_admin_billing_rule_row).transpose()
}
async fn update_admin_billing_rule(
&self,
rule_id: &str,
input: &AdminBillingRuleWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingRuleRecord>, DataLayerError> {
let row = match sqlx::query(
r#"
UPDATE billing_rules
SET
name = $2,
task_type = $3,
global_model_id = $4,
model_id = $5,
expression = $6,
variables = $7,
dimension_mappings = $8,
is_enabled = $9,
updated_at = NOW()
WHERE id = $1
RETURNING
id, name, task_type, global_model_id, model_id, expression, variables,
dimension_mappings, is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
"#,
)
.bind(rule_id)
.bind(&input.name)
.bind(&input.task_type)
.bind(input.global_model_id.as_deref())
.bind(input.model_id.as_deref())
.bind(&input.expression)
.bind(&input.variables)
.bind(&input.dimension_mappings)
.bind(input.is_enabled)
.fetch_optional(&self.pool)
.await
{
Ok(row) => row,
Err(sqlx::Error::Database(err)) => {
return Ok(AdminBillingMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
Err(err) => return Err(DataLayerError::postgres(err)),
};
match row {
Some(row) => Ok(AdminBillingMutationOutcome::Applied(
map_admin_billing_rule_row(&row)?,
)),
None => Ok(AdminBillingMutationOutcome::NotFound),
}
}
async fn create_admin_billing_collector(
&self,
input: &AdminBillingCollectorWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingCollectorRecord>, DataLayerError> {
let collector_id = uuid::Uuid::new_v4().to_string();
let row = match sqlx::query(
r#"
INSERT INTO dimension_collectors (
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled, created_at, updated_at
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, NOW(), NOW())
RETURNING
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
"#,
)
.bind(&collector_id)
.bind(&input.api_format)
.bind(&input.task_type)
.bind(&input.dimension_name)
.bind(&input.source_type)
.bind(input.source_path.as_deref())
.bind(&input.value_type)
.bind(input.transform_expression.as_deref())
.bind(input.default_value.as_deref())
.bind(input.priority)
.bind(input.is_enabled)
.fetch_one(&self.pool)
.await
{
Ok(row) => row,
Err(sqlx::Error::Database(err)) => {
return Ok(AdminBillingMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
Err(err) => return Err(DataLayerError::postgres(err)),
};
Ok(AdminBillingMutationOutcome::Applied(
map_admin_billing_collector_row(&row)?,
))
}
async fn list_admin_billing_collectors(
&self,
api_format: Option<&str>,
task_type: Option<&str>,
dimension_name: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<Option<(Vec<AdminBillingCollectorRecord>, u64)>, DataLayerError> {
let total = read_count(
sqlx::query(
r#"
SELECT COUNT(*) AS total
FROM dimension_collectors
WHERE ($1::TEXT IS NULL OR api_format = $1)
AND ($2::TEXT IS NULL OR task_type = $2)
AND ($3::TEXT IS NULL OR dimension_name = $3)
AND ($4::BOOL IS NULL OR is_enabled = $4)
"#,
)
.bind(api_format)
.bind(task_type)
.bind(dimension_name)
.bind(is_enabled)
.fetch_one(&self.pool)
.await
.map_postgres_err()?,
)?;
let offset = u64::from(page.saturating_sub(1) * page_size);
let rows = sqlx::query(
r#"
SELECT
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
FROM dimension_collectors
WHERE ($1::TEXT IS NULL OR api_format = $1)
AND ($2::TEXT IS NULL OR task_type = $2)
AND ($3::TEXT IS NULL OR dimension_name = $3)
AND ($4::BOOL IS NULL OR is_enabled = $4)
ORDER BY updated_at DESC, priority DESC, id ASC
OFFSET $5
LIMIT $6
"#,
)
.bind(api_format)
.bind(task_type)
.bind(dimension_name)
.bind(is_enabled)
.bind(
i64::try_from(offset)
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?,
)
.bind(i64::from(page_size))
.fetch_all(&self.pool)
.await
.map_postgres_err()?;
let items = rows
.iter()
.map(map_admin_billing_collector_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(Some((items, total)))
}
async fn find_admin_billing_collector(
&self,
collector_id: &str,
) -> Result<Option<AdminBillingCollectorRecord>, DataLayerError> {
let row = sqlx::query(
r#"
SELECT
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
FROM dimension_collectors
WHERE id = $1
"#,
)
.bind(collector_id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref()
.map(map_admin_billing_collector_row)
.transpose()
}
async fn update_admin_billing_collector(
&self,
collector_id: &str,
input: &AdminBillingCollectorWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingCollectorRecord>, DataLayerError> {
let row = match sqlx::query(
r#"
UPDATE dimension_collectors
SET
api_format = $2,
task_type = $3,
dimension_name = $4,
source_type = $5,
source_path = $6,
value_type = $7,
transform_expression = $8,
default_value = $9,
priority = $10,
is_enabled = $11,
updated_at = NOW()
WHERE id = $1
RETURNING
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
"#,
)
.bind(collector_id)
.bind(&input.api_format)
.bind(&input.task_type)
.bind(&input.dimension_name)
.bind(&input.source_type)
.bind(input.source_path.as_deref())
.bind(&input.value_type)
.bind(input.transform_expression.as_deref())
.bind(input.default_value.as_deref())
.bind(input.priority)
.bind(input.is_enabled)
.fetch_optional(&self.pool)
.await
{
Ok(row) => row,
Err(sqlx::Error::Database(err)) => {
return Ok(AdminBillingMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
Err(err) => return Err(DataLayerError::postgres(err)),
};
match row {
Some(row) => Ok(AdminBillingMutationOutcome::Applied(
map_admin_billing_collector_row(&row)?,
)),
None => Ok(AdminBillingMutationOutcome::NotFound),
}
}
async fn apply_admin_billing_preset(
&self,
preset: &str,
mode: &str,
collectors: &[AdminBillingCollectorWriteInput],
) -> Result<AdminBillingMutationOutcome<AdminBillingPresetApplyResult>, DataLayerError> {
let mut created = 0_u64;
let mut updated = 0_u64;
let mut skipped = 0_u64;
let mut errors = Vec::new();
for collector in collectors {
let existing_id = match sqlx::query_scalar::<_, String>(
r#"
SELECT id
FROM dimension_collectors
WHERE api_format = $1
AND task_type = $2
AND dimension_name = $3
AND priority = $4
AND is_enabled = TRUE
LIMIT 1
"#,
)
.bind(&collector.api_format)
.bind(&collector.task_type)
.bind(&collector.dimension_name)
.bind(collector.priority)
.fetch_optional(&self.pool)
.await
{
Ok(value) => value,
Err(err) => {
errors.push(format!(
"Failed to query collector: api_format={} task_type={} dim={}: {}",
collector.api_format, collector.task_type, collector.dimension_name, err
));
continue;
}
};
if let Some(existing_id) = existing_id {
if mode == "overwrite" {
match sqlx::query(
r#"
UPDATE dimension_collectors
SET
source_type = $2,
source_path = $3,
value_type = $4,
transform_expression = $5,
default_value = $6,
is_enabled = $7,
updated_at = NOW()
WHERE id = $1
"#,
)
.bind(&existing_id)
.bind(&collector.source_type)
.bind(collector.source_path.as_deref())
.bind(&collector.value_type)
.bind(collector.transform_expression.as_deref())
.bind(collector.default_value.as_deref())
.bind(collector.is_enabled)
.execute(&self.pool)
.await
{
Ok(_) => updated += 1,
Err(err) => errors.push(format!(
"Failed to update collector {}: {}",
existing_id, err
)),
}
} else {
skipped += 1;
}
continue;
}
match sqlx::query(
r#"
INSERT INTO dimension_collectors (
id, api_format, task_type, dimension_name, source_type, source_path, value_type,
transform_expression, default_value, priority, is_enabled, created_at, updated_at
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, NOW(), NOW())
"#,
)
.bind(uuid::Uuid::new_v4().to_string())
.bind(&collector.api_format)
.bind(&collector.task_type)
.bind(&collector.dimension_name)
.bind(&collector.source_type)
.bind(collector.source_path.as_deref())
.bind(&collector.value_type)
.bind(collector.transform_expression.as_deref())
.bind(collector.default_value.as_deref())
.bind(collector.priority)
.bind(collector.is_enabled)
.execute(&self.pool)
.await
{
Ok(_) => created += 1,
Err(err) => errors.push(format!(
"Failed to create collector: api_format={} task_type={} dim={}: {}",
collector.api_format, collector.task_type, collector.dimension_name, err
)),
}
}
Ok(AdminBillingMutationOutcome::Applied(
AdminBillingPresetApplyResult {
preset: preset.to_string(),
mode: mode.to_string(),
created,
updated,
skipped,
errors,
},
))
}
}
fn map_row(row: &sqlx::postgres::PgRow) -> Result<StoredBillingModelContext, DataLayerError> {
StoredBillingModelContext::new(
row.try_get("provider_id").map_postgres_err()?,
row.try_get("provider_billing_type").map_postgres_err()?,
row.try_get("provider_api_key_id").map_postgres_err()?,
row.try_get("provider_api_key_rate_multipliers")
.map_postgres_err()?,
row.try_get::<Option<i32>, _>("provider_api_key_cache_ttl_minutes")
.map_postgres_err()?
.map(i64::from),
row.try_get("global_model_id").map_postgres_err()?,
row.try_get("global_model_name").map_postgres_err()?,
row.try_get("global_model_config").map_postgres_err()?,
row.try_get("default_price_per_request")
.map_postgres_err()?,
row.try_get("default_tiered_pricing").map_postgres_err()?,
row.try_get("model_id").map_postgres_err()?,
row.try_get("model_provider_model_name")
.map_postgres_err()?,
row.try_get("model_config").map_postgres_err()?,
row.try_get("model_price_per_request").map_postgres_err()?,
row.try_get("model_tiered_pricing").map_postgres_err()?,
)
}
fn read_count(row: sqlx::postgres::PgRow) -> Result<u64, DataLayerError> {
Ok(row.try_get::<i64, _>("total").map_postgres_err()?.max(0) as u64)
}
fn map_admin_billing_rule_row(
row: &sqlx::postgres::PgRow,
) -> Result<AdminBillingRuleRecord, DataLayerError> {
Ok(AdminBillingRuleRecord {
id: row.try_get("id").map_postgres_err()?,
name: row.try_get("name").map_postgres_err()?,
task_type: row.try_get("task_type").map_postgres_err()?,
global_model_id: row.try_get("global_model_id").map_postgres_err()?,
model_id: row.try_get("model_id").map_postgres_err()?,
expression: row.try_get("expression").map_postgres_err()?,
variables: row
.try_get::<Option<serde_json::Value>, _>("variables")
.map_postgres_err()?
.unwrap_or_else(|| serde_json::json!({})),
dimension_mappings: row
.try_get::<Option<serde_json::Value>, _>("dimension_mappings")
.map_postgres_err()?
.unwrap_or_else(|| serde_json::json!({})),
is_enabled: row.try_get("is_enabled").map_postgres_err()?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_postgres_err()?
.max(0) as u64,
updated_at_unix_secs: row
.try_get::<i64, _>("updated_at_unix_secs")
.map_postgres_err()?
.max(0) as u64,
})
}
fn map_admin_billing_collector_row(
row: &sqlx::postgres::PgRow,
) -> Result<AdminBillingCollectorRecord, DataLayerError> {
Ok(AdminBillingCollectorRecord {
id: row.try_get("id").map_postgres_err()?,
api_format: row.try_get("api_format").map_postgres_err()?,
task_type: row.try_get("task_type").map_postgres_err()?,
dimension_name: row.try_get("dimension_name").map_postgres_err()?,
source_type: row.try_get("source_type").map_postgres_err()?,
source_path: row.try_get("source_path").map_postgres_err()?,
value_type: row.try_get("value_type").map_postgres_err()?,
transform_expression: row.try_get("transform_expression").map_postgres_err()?,
default_value: row.try_get("default_value").map_postgres_err()?,
priority: row.try_get("priority").map_postgres_err()?,
is_enabled: row.try_get("is_enabled").map_postgres_err()?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_postgres_err()?
.max(0) as u64,
updated_at_unix_secs: row
.try_get::<i64, _>("updated_at_unix_secs")
.map_postgres_err()?
.max(0) as u64,
})
}
#[cfg(test)]
mod tests {
use super::SqlxBillingReadRepository;
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
#[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 = SqlxBillingReadRepository::new(pool);
}
}

View File

@@ -1,214 +0,0 @@
use async_trait::async_trait;
use sqlx::{PgPool, Row};
use super::{BillingReadRepository, StoredBillingModelContext};
use crate::{error::SqlxResultExt, DataLayerError};
const FIND_MODEL_CONTEXT_SQL: &str = r#"
SELECT
p.id AS provider_id,
CAST(p.billing_type AS TEXT) AS provider_billing_type,
pak.id AS provider_api_key_id,
pak.rate_multipliers AS provider_api_key_rate_multipliers,
pak.cache_ttl_minutes AS provider_api_key_cache_ttl_minutes,
gm.id AS global_model_id,
gm.name AS global_model_name,
gm.config AS global_model_config,
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS default_price_per_request,
gm.default_tiered_pricing AS default_tiered_pricing,
m.id AS model_id,
m.provider_model_name AS model_provider_model_name,
m.config AS model_config,
CAST(m.price_per_request AS DOUBLE PRECISION) AS model_price_per_request,
m.tiered_pricing AS model_tiered_pricing
FROM providers p
INNER JOIN global_models gm
ON gm.is_active = TRUE
LEFT JOIN models m
ON m.global_model_id = gm.id
AND m.provider_id = p.id
AND m.is_active = TRUE
LEFT JOIN provider_api_keys pak
ON pak.id = $3
AND pak.provider_id = p.id
WHERE p.id = $1
AND (
gm.name = $2
OR m.provider_model_name = $2
OR (
m.provider_model_mappings IS NOT NULL
AND (
m.provider_model_mappings @> jsonb_build_array(jsonb_build_object('name', $2::TEXT))
OR m.provider_model_mappings @> jsonb_build_array(to_jsonb($2::TEXT))
OR m.provider_model_mappings @> jsonb_build_object('name', $2::TEXT)
OR m.provider_model_mappings = to_jsonb($2::TEXT)
)
)
)
ORDER BY
CASE
WHEN m.provider_model_name = $2 THEN 0
WHEN m.provider_model_mappings IS NOT NULL
AND (
m.provider_model_mappings @> jsonb_build_array(jsonb_build_object('name', $2::TEXT))
OR m.provider_model_mappings @> jsonb_build_array(to_jsonb($2::TEXT))
OR m.provider_model_mappings @> jsonb_build_object('name', $2::TEXT)
OR m.provider_model_mappings = to_jsonb($2::TEXT)
) THEN 1
WHEN gm.name = $2 THEN 2
ELSE 3
END ASC,
COALESCE(m.is_available, FALSE) DESC,
CASE
WHEN m.tiered_pricing IS NOT NULL OR m.price_per_request IS NOT NULL THEN 0
WHEN gm.default_tiered_pricing IS NOT NULL OR gm.default_price_per_request IS NOT NULL THEN 1
ELSE 2
END ASC,
m.created_at ASC
LIMIT 1
"#;
const FIND_MODEL_CONTEXT_BY_MODEL_ID_SQL: &str = r#"
SELECT
p.id AS provider_id,
CAST(p.billing_type AS TEXT) AS provider_billing_type,
pak.id AS provider_api_key_id,
pak.rate_multipliers AS provider_api_key_rate_multipliers,
pak.cache_ttl_minutes AS provider_api_key_cache_ttl_minutes,
gm.id AS global_model_id,
gm.name AS global_model_name,
gm.config AS global_model_config,
CAST(gm.default_price_per_request AS DOUBLE PRECISION) AS default_price_per_request,
gm.default_tiered_pricing AS default_tiered_pricing,
m.id AS model_id,
m.provider_model_name AS model_provider_model_name,
m.config AS model_config,
CAST(m.price_per_request AS DOUBLE PRECISION) AS model_price_per_request,
m.tiered_pricing AS model_tiered_pricing
FROM providers p
INNER JOIN models m
ON m.id = $2
AND m.provider_id = p.id
AND m.is_active = TRUE
INNER JOIN global_models gm
ON gm.id = m.global_model_id
AND gm.is_active = TRUE
LEFT JOIN provider_api_keys pak
ON pak.id = $3
AND pak.provider_id = p.id
WHERE p.id = $1
LIMIT 1
"#;
#[derive(Debug, Clone)]
pub struct SqlxBillingReadRepository {
pool: PgPool,
}
impl SqlxBillingReadRepository {
pub fn new(pool: PgPool) -> Self {
Self { pool }
}
pub async fn find_model_context(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
let row = sqlx::query(FIND_MODEL_CONTEXT_SQL)
.bind(provider_id)
.bind(global_model_name)
.bind(provider_api_key_id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_row).transpose()
}
pub async fn find_model_context_by_model_id(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
model_id: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
let row = sqlx::query(FIND_MODEL_CONTEXT_BY_MODEL_ID_SQL)
.bind(provider_id)
.bind(model_id)
.bind(provider_api_key_id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_row).transpose()
}
}
#[async_trait]
impl BillingReadRepository for SqlxBillingReadRepository {
async fn find_model_context(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
global_model_name: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
Self::find_model_context(self, provider_id, provider_api_key_id, global_model_name).await
}
async fn find_model_context_by_model_id(
&self,
provider_id: &str,
provider_api_key_id: Option<&str>,
model_id: &str,
) -> Result<Option<StoredBillingModelContext>, DataLayerError> {
Self::find_model_context_by_model_id(self, provider_id, provider_api_key_id, model_id).await
}
}
fn map_row(row: &sqlx::postgres::PgRow) -> Result<StoredBillingModelContext, DataLayerError> {
StoredBillingModelContext::new(
row.try_get("provider_id").map_postgres_err()?,
row.try_get("provider_billing_type").map_postgres_err()?,
row.try_get("provider_api_key_id").map_postgres_err()?,
row.try_get("provider_api_key_rate_multipliers")
.map_postgres_err()?,
row.try_get::<Option<i32>, _>("provider_api_key_cache_ttl_minutes")
.map_postgres_err()?
.map(i64::from),
row.try_get("global_model_id").map_postgres_err()?,
row.try_get("global_model_name").map_postgres_err()?,
row.try_get("global_model_config").map_postgres_err()?,
row.try_get("default_price_per_request")
.map_postgres_err()?,
row.try_get("default_tiered_pricing").map_postgres_err()?,
row.try_get("model_id").map_postgres_err()?,
row.try_get("model_provider_model_name")
.map_postgres_err()?,
row.try_get("model_config").map_postgres_err()?,
row.try_get("model_price_per_request").map_postgres_err()?,
row.try_get("model_tiered_pricing").map_postgres_err()?,
)
}
#[cfg(test)]
mod tests {
use super::SqlxBillingReadRepository;
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
#[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 = SqlxBillingReadRepository::new(pool);
}
}

File diff suppressed because it is too large Load Diff

View File

@@ -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;

View 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);
}
}

View File

@@ -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]

View 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");
}
}

View File

@@ -1,5 +1,7 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::candidates::{
@@ -10,4 +12,6 @@ pub(crate) use aether_data_contracts::repository::candidates::{
StoredRequestCandidate, UpsertRequestCandidateRecord,
};
pub use memory::InMemoryRequestCandidateRepository;
pub use sql::SqlxRequestCandidateReadRepository;
pub use mysql::MysqlRequestCandidateRepository;
pub use postgres::SqlxRequestCandidateReadRepository;
pub use sqlite::SqliteRequestCandidateRepository;

View File

@@ -0,0 +1,678 @@
use std::collections::{BTreeMap, BTreeSet};
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::{
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository,
RequestCandidateStatus, RequestCandidateWriteRepository, StoredRequestCandidate,
UpsertRequestCandidateRecord,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const CANDIDATE_COLUMNS: &str = r#"
SELECT
id,
request_id,
user_id,
api_key_id,
username,
api_key_name,
candidate_index,
retry_index,
provider_id,
endpoint_id,
key_id,
status,
skip_reason,
is_cached,
status_code,
error_type,
error_message,
latency_ms,
concurrent_requests,
extra_data,
required_capabilities,
created_at AS created_at_unix_ms,
started_at AS started_at_unix_ms,
finished_at AS finished_at_unix_ms
FROM request_candidates
"#;
#[derive(Debug, Clone)]
pub struct MysqlRequestCandidateRepository {
pool: MysqlPool,
}
impl MysqlRequestCandidateRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn find_by_unique(
&self,
request_id: &str,
candidate_index: u32,
retry_index: u32,
) -> Result<Option<StoredRequestCandidate>, DataLayerError> {
let row = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE request_id = ? AND candidate_index = ? AND retry_index = ? LIMIT 1"
))
.bind(request_id)
.bind(to_i32(candidate_index)?)
.bind(to_i32(retry_index)?)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_candidate_row).transpose()
}
}
#[async_trait]
impl RequestCandidateReadRepository for MysqlRequestCandidateRepository {
async fn list_by_request_id(
&self,
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE request_id = ? ORDER BY candidate_index ASC, retry_index ASC, created_at ASC"
))
.bind(request_id)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_recent(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} ORDER BY created_at DESC LIMIT ?"
))
.bind(limit_i64(limit, "recent request candidate limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_by_provider_id(
&self,
provider_id: &str,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE provider_id = ? ORDER BY created_at DESC LIMIT ?"
))
.bind(provider_id)
.bind(limit_i64(limit, "provider request candidate limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_finalized_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if endpoint_ids.is_empty() || limit == 0 {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_COLUMNS);
push_endpoint_in_clause(&mut builder, endpoint_ids);
builder
.push(" AND created_at >= ")
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
.push(" AND status IN ('success', 'failed', 'skipped')")
.push(" ORDER BY created_at DESC LIMIT ")
.push_bind(limit_i64(limit, "finalized request candidate limit")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn count_finalized_statuses_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
) -> Result<Vec<PublicHealthStatusCount>, DataLayerError> {
if endpoint_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(
"SELECT endpoint_id, status, COUNT(id) AS count FROM request_candidates",
);
push_endpoint_in_clause(&mut builder, endpoint_ids);
builder
.push(" AND created_at >= ")
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
.push(" AND status IN ('success', 'failed', 'skipped')")
.push(" GROUP BY endpoint_id, status");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter()
.map(|row| {
Ok(PublicHealthStatusCount {
endpoint_id: row.try_get("endpoint_id").map_sql_err()?,
status: RequestCandidateStatus::from_database(
row.try_get::<String, _>("status").map_sql_err()?.as_str(),
)?,
count: u64::try_from(row.try_get::<i64, _>("count").map_sql_err()?).map_err(
|_| {
DataLayerError::UnexpectedValue(
"public health status count out of range".to_string(),
)
},
)?,
})
})
.collect()
}
async fn aggregate_finalized_timeline_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
until_unix_secs: u64,
segments: u32,
) -> Result<Vec<PublicHealthTimelineBucket>, DataLayerError> {
if endpoint_ids.is_empty() || segments == 0 || until_unix_secs < since_unix_secs {
return Ok(Vec::new());
}
let since_ms = unix_secs_to_ms_i64(since_unix_secs)?;
let until_ms = unix_secs_to_ms_i64(until_unix_secs)?;
let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_COLUMNS);
push_endpoint_in_clause(&mut builder, endpoint_ids);
builder
.push(" AND created_at >= ")
.push_bind(since_ms)
.push(" AND created_at <= ")
.push_bind(until_ms)
.push(" AND status IN ('success', 'failed', 'skipped')");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
aggregate_timeline(
rows.iter()
.map(map_candidate_row)
.collect::<Result<Vec<_>, _>>()?,
since_unix_secs,
until_unix_secs,
segments,
)
}
}
#[async_trait]
impl RequestCandidateWriteRepository for MysqlRequestCandidateRepository {
async fn upsert(
&self,
candidate: UpsertRequestCandidateRecord,
) -> Result<StoredRequestCandidate, DataLayerError> {
candidate.validate()?;
let existing = self
.find_by_unique(
&candidate.request_id,
candidate.candidate_index,
candidate.retry_index,
)
.await?;
let merged = merge_candidate(candidate, existing)?;
upsert_merged_candidate(&self.pool, &merged).await?;
Ok(merged)
}
async fn delete_created_before(
&self,
created_before_unix_secs: u64,
limit: usize,
) -> Result<usize, DataLayerError> {
if limit == 0 {
return Ok(0);
}
let rows_affected = sqlx::query(
r#"
DELETE FROM request_candidates
WHERE id IN (
SELECT id
FROM (
SELECT id
FROM request_candidates
WHERE created_at < ?
ORDER BY created_at ASC, id ASC
LIMIT ?
) AS old_request_candidates
)
"#,
)
.bind(unix_secs_to_ms_i64(created_before_unix_secs)?)
.bind(limit_i64(limit, "request candidate delete limit")?)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or_default())
}
}
async fn upsert_merged_candidate(
pool: &MysqlPool,
candidate: &StoredRequestCandidate,
) -> Result<(), DataLayerError> {
sqlx::query(
r#"
INSERT INTO request_candidates (
id, request_id, user_id, api_key_id, username, api_key_name,
candidate_index, retry_index, provider_id, endpoint_id, key_id, status,
skip_reason, is_cached, status_code, error_type, error_message, latency_ms,
concurrent_requests, extra_data, required_capabilities, created_at, started_at, finished_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
user_id = VALUES(user_id),
api_key_id = VALUES(api_key_id),
username = VALUES(username),
api_key_name = VALUES(api_key_name),
provider_id = VALUES(provider_id),
endpoint_id = VALUES(endpoint_id),
key_id = VALUES(key_id),
status = VALUES(status),
skip_reason = VALUES(skip_reason),
is_cached = VALUES(is_cached),
status_code = VALUES(status_code),
error_type = VALUES(error_type),
error_message = VALUES(error_message),
latency_ms = VALUES(latency_ms),
concurrent_requests = VALUES(concurrent_requests),
extra_data = VALUES(extra_data),
required_capabilities = VALUES(required_capabilities),
created_at = VALUES(created_at),
started_at = VALUES(started_at),
finished_at = VALUES(finished_at)
"#,
)
.bind(&candidate.id)
.bind(&candidate.request_id)
.bind(&candidate.user_id)
.bind(&candidate.api_key_id)
.bind(&candidate.username)
.bind(&candidate.api_key_name)
.bind(to_i32(candidate.candidate_index)?)
.bind(to_i32(candidate.retry_index)?)
.bind(&candidate.provider_id)
.bind(&candidate.endpoint_id)
.bind(&candidate.key_id)
.bind(status_to_database(candidate.status))
.bind(&candidate.skip_reason)
.bind(candidate.is_cached)
.bind(candidate.status_code.map(i32::from))
.bind(&candidate.error_type)
.bind(&candidate.error_message)
.bind(candidate.latency_ms.map(to_i32_u64).transpose()?)
.bind(candidate.concurrent_requests.map(to_i32).transpose()?)
.bind(json_to_string(&candidate.extra_data)?)
.bind(json_to_string(&candidate.required_capabilities)?)
.bind(u64_to_i64(
candidate.created_at_unix_ms,
"request candidate created_at",
)?)
.bind(optional_u64_to_i64(
candidate.started_at_unix_ms,
"request candidate started_at",
)?)
.bind(optional_u64_to_i64(
candidate.finished_at_unix_ms,
"request candidate finished_at",
)?)
.execute(pool)
.await
.map_sql_err()?;
Ok(())
}
fn push_endpoint_in_clause<'args>(
builder: &mut QueryBuilder<'args, MySql>,
endpoint_ids: &'args [String],
) {
builder.push(" WHERE endpoint_id IN (");
{
let mut separated = builder.separated(", ");
for endpoint_id in endpoint_ids {
separated.push_bind(endpoint_id);
}
}
builder.push(")");
}
fn merge_candidate(
candidate: UpsertRequestCandidateRecord,
existing: Option<StoredRequestCandidate>,
) -> Result<StoredRequestCandidate, DataLayerError> {
let created_at_unix_ms = candidate
.created_at_unix_ms
.filter(|value| *value > 1000)
.or_else(|| {
existing
.as_ref()
.map(|value| value.created_at_unix_ms)
.filter(|value| *value > 1000)
})
.or(candidate.started_at_unix_ms)
.or(candidate.finished_at_unix_ms)
.unwrap_or_else(current_unix_ms);
let id = existing
.as_ref()
.map(|value| value.id.clone())
.unwrap_or(candidate.id);
let extra_data = merge_json_objects(
existing.as_ref().and_then(|value| value.extra_data.clone()),
candidate.extra_data,
);
StoredRequestCandidate::new(
id,
candidate.request_id,
candidate
.user_id
.or_else(|| existing.as_ref().and_then(|value| value.user_id.clone())),
candidate
.api_key_id
.or_else(|| existing.as_ref().and_then(|value| value.api_key_id.clone())),
candidate
.username
.or_else(|| existing.as_ref().and_then(|value| value.username.clone())),
candidate.api_key_name.or_else(|| {
existing
.as_ref()
.and_then(|value| value.api_key_name.clone())
}),
to_i32(candidate.candidate_index)?,
to_i32(candidate.retry_index)?,
candidate.provider_id.or_else(|| {
existing
.as_ref()
.and_then(|value| value.provider_id.clone())
}),
candidate.endpoint_id.or_else(|| {
existing
.as_ref()
.and_then(|value| value.endpoint_id.clone())
}),
candidate
.key_id
.or_else(|| existing.as_ref().and_then(|value| value.key_id.clone())),
candidate.status,
candidate.skip_reason.or_else(|| {
existing
.as_ref()
.and_then(|value| value.skip_reason.clone())
}),
candidate
.is_cached
.unwrap_or_else(|| existing.as_ref().is_some_and(|value| value.is_cached)),
candidate.status_code.map(i32::from).or_else(|| {
existing
.as_ref()
.and_then(|value| value.status_code.map(i32::from))
}),
candidate
.error_type
.or_else(|| existing.as_ref().and_then(|value| value.error_type.clone())),
candidate.error_message.or_else(|| {
existing
.as_ref()
.and_then(|value| value.error_message.clone())
}),
candidate.latency_ms.map(to_i32_u64).transpose()?.or(
match existing.as_ref().and_then(|value| value.latency_ms) {
Some(value) => Some(to_i32_u64(value)?),
None => None,
},
),
candidate.concurrent_requests.map(to_i32).transpose()?.or(
match existing
.as_ref()
.and_then(|value| value.concurrent_requests)
{
Some(value) => Some(to_i32(value)?),
None => None,
},
),
extra_data,
candidate.required_capabilities.or_else(|| {
existing
.as_ref()
.and_then(|value| value.required_capabilities.clone())
}),
u64_to_i64(created_at_unix_ms, "request candidate created_at")?,
candidate
.started_at_unix_ms
.or_else(|| existing.as_ref().and_then(|value| value.started_at_unix_ms))
.map(|value| u64_to_i64(value, "request candidate started_at"))
.transpose()?,
candidate
.finished_at_unix_ms
.or_else(|| {
existing
.as_ref()
.and_then(|value| value.finished_at_unix_ms)
})
.map(|value| u64_to_i64(value, "request candidate finished_at"))
.transpose()?,
)
}
fn aggregate_timeline(
candidates: Vec<StoredRequestCandidate>,
since_unix_secs: u64,
until_unix_secs: u64,
segments: u32,
) -> Result<Vec<PublicHealthTimelineBucket>, DataLayerError> {
let endpoint_ids = candidates
.iter()
.filter_map(|candidate| candidate.endpoint_id.clone())
.collect::<BTreeSet<_>>();
let span_ms = until_unix_secs
.saturating_sub(since_unix_secs)
.saturating_mul(1000)
.max(1);
let since_ms = since_unix_secs.saturating_mul(1000);
let mut buckets = BTreeMap::<(String, u32), PublicHealthTimelineBucket>::new();
for candidate in candidates {
let Some(endpoint_id) = candidate.endpoint_id.clone() else {
continue;
};
let offset = candidate.created_at_unix_ms.saturating_sub(since_ms);
let segment_idx = ((offset.saturating_mul(u64::from(segments))) / span_ms)
.min(u64::from(segments.saturating_sub(1))) as u32;
let bucket = buckets.entry((endpoint_id.clone(), segment_idx)).or_insert(
PublicHealthTimelineBucket {
endpoint_id,
segment_idx,
total_count: 0,
success_count: 0,
failed_count: 0,
min_created_at_unix_ms: Some(candidate.created_at_unix_ms),
max_created_at_unix_ms: Some(candidate.created_at_unix_ms),
},
);
bucket.total_count += 1;
if candidate.status == RequestCandidateStatus::Success {
bucket.success_count += 1;
}
if candidate.status == RequestCandidateStatus::Failed {
bucket.failed_count += 1;
}
bucket.min_created_at_unix_ms = bucket
.min_created_at_unix_ms
.map(|value| value.min(candidate.created_at_unix_ms));
bucket.max_created_at_unix_ms = bucket
.max_created_at_unix_ms
.map(|value| value.max(candidate.created_at_unix_ms));
}
for endpoint_id in endpoint_ids {
for segment_idx in 0..segments {
buckets.entry((endpoint_id.clone(), segment_idx)).or_insert(
PublicHealthTimelineBucket {
endpoint_id: endpoint_id.clone(),
segment_idx,
total_count: 0,
success_count: 0,
failed_count: 0,
min_created_at_unix_ms: None,
max_created_at_unix_ms: None,
},
);
}
}
Ok(buckets.into_values().collect())
}
fn map_candidate_row(row: &MySqlRow) -> Result<StoredRequestCandidate, DataLayerError> {
StoredRequestCandidate::new(
row.try_get("id").map_sql_err()?,
row.try_get("request_id").map_sql_err()?,
row.try_get("user_id").map_sql_err()?,
row.try_get("api_key_id").map_sql_err()?,
row.try_get("username").map_sql_err()?,
row.try_get("api_key_name").map_sql_err()?,
row.try_get("candidate_index").map_sql_err()?,
row.try_get("retry_index").map_sql_err()?,
row.try_get("provider_id").map_sql_err()?,
row.try_get("endpoint_id").map_sql_err()?,
row.try_get("key_id").map_sql_err()?,
RequestCandidateStatus::from_database(
row.try_get::<String, _>("status").map_sql_err()?.as_str(),
)?,
row.try_get("skip_reason").map_sql_err()?,
row.try_get("is_cached").map_sql_err()?,
row.try_get("status_code").map_sql_err()?,
row.try_get("error_type").map_sql_err()?,
row.try_get("error_message").map_sql_err()?,
row.try_get("latency_ms").map_sql_err()?,
row.try_get("concurrent_requests").map_sql_err()?,
parse_json(row.try_get("extra_data").ok().flatten())?,
parse_json(row.try_get("required_capabilities").ok().flatten())?,
row.try_get("created_at_unix_ms").map_sql_err()?,
row.try_get("started_at_unix_ms").map_sql_err()?,
row.try_get("finished_at_unix_ms").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!(
"request_candidates JSON field is invalid: {err}"
))
})
})
.transpose()
}
fn json_to_string(value: &Option<serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.as_ref()
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"request_candidates JSON field is unserializable: {err}"
))
})
})
.transpose()
}
fn merge_json_objects(
existing: Option<serde_json::Value>,
overlay: Option<serde_json::Value>,
) -> Option<serde_json::Value> {
match (existing, overlay) {
(
Some(serde_json::Value::Object(mut existing_object)),
Some(serde_json::Value::Object(overlay_object)),
) => {
existing_object.extend(overlay_object);
Some(serde_json::Value::Object(existing_object))
}
(_existing, Some(overlay)) => Some(overlay),
(existing, None) => existing,
}
}
fn status_to_database(status: RequestCandidateStatus) -> &'static str {
match status {
RequestCandidateStatus::Available => "available",
RequestCandidateStatus::Unused => "unused",
RequestCandidateStatus::Pending => "pending",
RequestCandidateStatus::Streaming => "streaming",
RequestCandidateStatus::Success => "success",
RequestCandidateStatus::Failed => "failed",
RequestCandidateStatus::Cancelled => "cancelled",
RequestCandidateStatus::Skipped => "skipped",
}
}
fn current_unix_ms() -> u64 {
chrono::Utc::now().timestamp_millis().max(0) as u64
}
fn unix_secs_to_ms_i64(value: u64) -> Result<i64, DataLayerError> {
let value = value.checked_mul(1000).ok_or_else(|| {
DataLayerError::UnexpectedValue("request candidate timestamp overflow".to_string())
})?;
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue("request candidate timestamp overflow".to_string())
})
}
fn limit_i64(value: usize, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::UnexpectedValue(format!("invalid {name}: {value}")))
}
fn to_i32(value: u32) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("request candidate value out of range: {value}"))
})
}
fn to_i32_u64(value: u64) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("request candidate value out of range: {value}"))
})
}
fn u64_to_i64(value: u64, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| DataLayerError::UnexpectedValue(format!("{name} overflow")))
}
fn optional_u64_to_i64(value: Option<u64>, name: &str) -> Result<Option<i64>, DataLayerError> {
value.map(|value| u64_to_i64(value, name)).transpose()
}
#[cfg(test)]
mod tests {
use super::MysqlRequestCandidateRepository;
#[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 = MysqlRequestCandidateRepository::new(pool);
}
}

View File

@@ -8,7 +8,7 @@ use super::{
RequestCandidateStatus, RequestCandidateWriteRepository, StoredRequestCandidate,
UpsertRequestCandidateRecord,
};
use crate::postgres::PostgresTransactionRunner;
use crate::driver::postgres::PostgresTransactionRunner;
use crate::{error::SqlxResultExt, DataLayerError};
const LIST_BY_REQUEST_ID_SQL: &str = r#"
@@ -753,7 +753,7 @@ fn to_i32_u64(value: u64) -> Result<i32, DataLayerError> {
#[cfg(test)]
mod tests {
use super::{SqlxRequestCandidateReadRepository, UPSERT_SQL};
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
#[test]
fn upsert_sql_does_not_default_missing_or_epoch_created_at_to_epoch() {

View File

@@ -0,0 +1,777 @@
use std::collections::{BTreeMap, BTreeSet};
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::{
PublicHealthStatusCount, PublicHealthTimelineBucket, RequestCandidateReadRepository,
RequestCandidateStatus, RequestCandidateWriteRepository, StoredRequestCandidate,
UpsertRequestCandidateRecord,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const CANDIDATE_COLUMNS: &str = r#"
SELECT
id,
request_id,
user_id,
api_key_id,
username,
api_key_name,
candidate_index,
retry_index,
provider_id,
endpoint_id,
key_id,
status,
skip_reason,
is_cached,
status_code,
error_type,
error_message,
latency_ms,
concurrent_requests,
extra_data,
required_capabilities,
created_at AS created_at_unix_ms,
started_at AS started_at_unix_ms,
finished_at AS finished_at_unix_ms
FROM request_candidates
"#;
#[derive(Debug, Clone)]
pub struct SqliteRequestCandidateRepository {
pool: SqlitePool,
}
impl SqliteRequestCandidateRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn find_by_unique(
&self,
request_id: &str,
candidate_index: u32,
retry_index: u32,
) -> Result<Option<StoredRequestCandidate>, DataLayerError> {
let row = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE request_id = ? AND candidate_index = ? AND retry_index = ? LIMIT 1"
))
.bind(request_id)
.bind(to_i32(candidate_index)?)
.bind(to_i32(retry_index)?)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_candidate_row).transpose()
}
}
#[async_trait]
impl RequestCandidateReadRepository for SqliteRequestCandidateRepository {
async fn list_by_request_id(
&self,
request_id: &str,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE request_id = ? ORDER BY candidate_index ASC, retry_index ASC, created_at ASC"
))
.bind(request_id)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_recent(
&self,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} ORDER BY created_at DESC LIMIT ?"
))
.bind(limit_i64(limit, "recent request candidate limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_by_provider_id(
&self,
provider_id: &str,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{CANDIDATE_COLUMNS} WHERE provider_id = ? ORDER BY created_at DESC LIMIT ?"
))
.bind(provider_id)
.bind(limit_i64(limit, "provider request candidate limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn list_finalized_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
if endpoint_ids.is_empty() || limit == 0 {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_COLUMNS);
push_endpoint_in_clause(&mut builder, endpoint_ids);
builder
.push(" AND created_at >= ")
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
.push(" AND status IN ('success', 'failed', 'skipped')")
.push(" ORDER BY created_at DESC LIMIT ")
.push_bind(limit_i64(limit, "finalized request candidate limit")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_candidate_row).collect()
}
async fn count_finalized_statuses_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
) -> Result<Vec<PublicHealthStatusCount>, DataLayerError> {
if endpoint_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(
"SELECT endpoint_id, status, COUNT(id) AS count FROM request_candidates",
);
push_endpoint_in_clause(&mut builder, endpoint_ids);
builder
.push(" AND created_at >= ")
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
.push(" AND status IN ('success', 'failed', 'skipped')")
.push(" GROUP BY endpoint_id, status");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter()
.map(|row| {
Ok(PublicHealthStatusCount {
endpoint_id: row.try_get("endpoint_id").map_sql_err()?,
status: RequestCandidateStatus::from_database(
row.try_get::<String, _>("status").map_sql_err()?.as_str(),
)?,
count: u64::try_from(row.try_get::<i64, _>("count").map_sql_err()?).map_err(
|_| {
DataLayerError::UnexpectedValue(
"public health status count out of range".to_string(),
)
},
)?,
})
})
.collect()
}
async fn aggregate_finalized_timeline_by_endpoint_ids_since(
&self,
endpoint_ids: &[String],
since_unix_secs: u64,
until_unix_secs: u64,
segments: u32,
) -> Result<Vec<PublicHealthTimelineBucket>, DataLayerError> {
if endpoint_ids.is_empty() || segments == 0 || until_unix_secs < since_unix_secs {
return Ok(Vec::new());
}
let since_ms = unix_secs_to_ms_i64(since_unix_secs)?;
let until_ms = unix_secs_to_ms_i64(until_unix_secs)?;
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_COLUMNS);
push_endpoint_in_clause(&mut builder, endpoint_ids);
builder
.push(" AND created_at >= ")
.push_bind(since_ms)
.push(" AND created_at <= ")
.push_bind(until_ms)
.push(" AND status IN ('success', 'failed', 'skipped')");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
aggregate_timeline(
rows.iter()
.map(map_candidate_row)
.collect::<Result<Vec<_>, _>>()?,
since_unix_secs,
until_unix_secs,
segments,
)
}
}
#[async_trait]
impl RequestCandidateWriteRepository for SqliteRequestCandidateRepository {
async fn upsert(
&self,
candidate: UpsertRequestCandidateRecord,
) -> Result<StoredRequestCandidate, DataLayerError> {
candidate.validate()?;
let existing = self
.find_by_unique(
&candidate.request_id,
candidate.candidate_index,
candidate.retry_index,
)
.await?;
let merged = merge_candidate(candidate, existing)?;
upsert_merged_candidate(&self.pool, &merged).await?;
Ok(merged)
}
async fn delete_created_before(
&self,
created_before_unix_secs: u64,
limit: usize,
) -> Result<usize, DataLayerError> {
if limit == 0 {
return Ok(0);
}
let rows_affected = sqlx::query(
r#"
DELETE FROM request_candidates
WHERE id IN (
SELECT id
FROM request_candidates
WHERE created_at < ?
ORDER BY created_at ASC, id ASC
LIMIT ?
)
"#,
)
.bind(unix_secs_to_ms_i64(created_before_unix_secs)?)
.bind(limit_i64(limit, "request candidate delete limit")?)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or_default())
}
}
async fn upsert_merged_candidate(
pool: &SqlitePool,
candidate: &StoredRequestCandidate,
) -> Result<(), DataLayerError> {
sqlx::query(
r#"
INSERT INTO request_candidates (
id, request_id, user_id, api_key_id, username, api_key_name,
candidate_index, retry_index, provider_id, endpoint_id, key_id, status,
skip_reason, is_cached, status_code, error_type, error_message, latency_ms,
concurrent_requests, extra_data, required_capabilities, created_at, started_at, finished_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(request_id, candidate_index, retry_index) DO UPDATE SET
user_id = excluded.user_id,
api_key_id = excluded.api_key_id,
username = excluded.username,
api_key_name = excluded.api_key_name,
provider_id = excluded.provider_id,
endpoint_id = excluded.endpoint_id,
key_id = excluded.key_id,
status = excluded.status,
skip_reason = excluded.skip_reason,
is_cached = excluded.is_cached,
status_code = excluded.status_code,
error_type = excluded.error_type,
error_message = excluded.error_message,
latency_ms = excluded.latency_ms,
concurrent_requests = excluded.concurrent_requests,
extra_data = excluded.extra_data,
required_capabilities = excluded.required_capabilities,
created_at = excluded.created_at,
started_at = excluded.started_at,
finished_at = excluded.finished_at
"#,
)
.bind(&candidate.id)
.bind(&candidate.request_id)
.bind(&candidate.user_id)
.bind(&candidate.api_key_id)
.bind(&candidate.username)
.bind(&candidate.api_key_name)
.bind(to_i32(candidate.candidate_index)?)
.bind(to_i32(candidate.retry_index)?)
.bind(&candidate.provider_id)
.bind(&candidate.endpoint_id)
.bind(&candidate.key_id)
.bind(status_to_database(candidate.status))
.bind(&candidate.skip_reason)
.bind(candidate.is_cached)
.bind(candidate.status_code.map(i32::from))
.bind(&candidate.error_type)
.bind(&candidate.error_message)
.bind(candidate.latency_ms.map(to_i32_u64).transpose()?)
.bind(candidate.concurrent_requests.map(to_i32).transpose()?)
.bind(json_to_string(&candidate.extra_data)?)
.bind(json_to_string(&candidate.required_capabilities)?)
.bind(u64_to_i64(
candidate.created_at_unix_ms,
"request candidate created_at",
)?)
.bind(optional_u64_to_i64(
candidate.started_at_unix_ms,
"request candidate started_at",
)?)
.bind(optional_u64_to_i64(
candidate.finished_at_unix_ms,
"request candidate finished_at",
)?)
.execute(pool)
.await
.map_sql_err()?;
Ok(())
}
fn push_endpoint_in_clause<'args>(
builder: &mut QueryBuilder<'args, Sqlite>,
endpoint_ids: &'args [String],
) {
builder.push(" WHERE endpoint_id IN (");
{
let mut separated = builder.separated(", ");
for endpoint_id in endpoint_ids {
separated.push_bind(endpoint_id);
}
}
builder.push(")");
}
fn merge_candidate(
candidate: UpsertRequestCandidateRecord,
existing: Option<StoredRequestCandidate>,
) -> Result<StoredRequestCandidate, DataLayerError> {
let created_at_unix_ms = candidate
.created_at_unix_ms
.filter(|value| *value > 1000)
.or_else(|| {
existing
.as_ref()
.map(|value| value.created_at_unix_ms)
.filter(|value| *value > 1000)
})
.or(candidate.started_at_unix_ms)
.or(candidate.finished_at_unix_ms)
.unwrap_or_else(current_unix_ms);
let id = existing
.as_ref()
.map(|value| value.id.clone())
.unwrap_or(candidate.id);
let extra_data = merge_json_objects(
existing.as_ref().and_then(|value| value.extra_data.clone()),
candidate.extra_data,
);
StoredRequestCandidate::new(
id,
candidate.request_id,
candidate
.user_id
.or_else(|| existing.as_ref().and_then(|value| value.user_id.clone())),
candidate
.api_key_id
.or_else(|| existing.as_ref().and_then(|value| value.api_key_id.clone())),
candidate
.username
.or_else(|| existing.as_ref().and_then(|value| value.username.clone())),
candidate.api_key_name.or_else(|| {
existing
.as_ref()
.and_then(|value| value.api_key_name.clone())
}),
to_i32(candidate.candidate_index)?,
to_i32(candidate.retry_index)?,
candidate.provider_id.or_else(|| {
existing
.as_ref()
.and_then(|value| value.provider_id.clone())
}),
candidate.endpoint_id.or_else(|| {
existing
.as_ref()
.and_then(|value| value.endpoint_id.clone())
}),
candidate
.key_id
.or_else(|| existing.as_ref().and_then(|value| value.key_id.clone())),
candidate.status,
candidate.skip_reason.or_else(|| {
existing
.as_ref()
.and_then(|value| value.skip_reason.clone())
}),
candidate
.is_cached
.unwrap_or_else(|| existing.as_ref().is_some_and(|value| value.is_cached)),
candidate.status_code.map(i32::from).or_else(|| {
existing
.as_ref()
.and_then(|value| value.status_code.map(i32::from))
}),
candidate
.error_type
.or_else(|| existing.as_ref().and_then(|value| value.error_type.clone())),
candidate.error_message.or_else(|| {
existing
.as_ref()
.and_then(|value| value.error_message.clone())
}),
candidate.latency_ms.map(to_i32_u64).transpose()?.or(
match existing.as_ref().and_then(|value| value.latency_ms) {
Some(value) => Some(to_i32_u64(value)?),
None => None,
},
),
candidate.concurrent_requests.map(to_i32).transpose()?.or(
match existing
.as_ref()
.and_then(|value| value.concurrent_requests)
{
Some(value) => Some(to_i32(value)?),
None => None,
},
),
extra_data,
candidate.required_capabilities.or_else(|| {
existing
.as_ref()
.and_then(|value| value.required_capabilities.clone())
}),
u64_to_i64(created_at_unix_ms, "request candidate created_at")?,
candidate
.started_at_unix_ms
.or_else(|| existing.as_ref().and_then(|value| value.started_at_unix_ms))
.map(|value| u64_to_i64(value, "request candidate started_at"))
.transpose()?,
candidate
.finished_at_unix_ms
.or_else(|| {
existing
.as_ref()
.and_then(|value| value.finished_at_unix_ms)
})
.map(|value| u64_to_i64(value, "request candidate finished_at"))
.transpose()?,
)
}
fn aggregate_timeline(
candidates: Vec<StoredRequestCandidate>,
since_unix_secs: u64,
until_unix_secs: u64,
segments: u32,
) -> Result<Vec<PublicHealthTimelineBucket>, DataLayerError> {
let endpoint_ids = candidates
.iter()
.filter_map(|candidate| candidate.endpoint_id.clone())
.collect::<BTreeSet<_>>();
let span_ms = until_unix_secs
.saturating_sub(since_unix_secs)
.saturating_mul(1000)
.max(1);
let since_ms = since_unix_secs.saturating_mul(1000);
let mut buckets = BTreeMap::<(String, u32), PublicHealthTimelineBucket>::new();
for candidate in candidates {
let Some(endpoint_id) = candidate.endpoint_id.clone() else {
continue;
};
let offset = candidate.created_at_unix_ms.saturating_sub(since_ms);
let segment_idx = ((offset.saturating_mul(u64::from(segments))) / span_ms)
.min(u64::from(segments.saturating_sub(1))) as u32;
let bucket = buckets.entry((endpoint_id.clone(), segment_idx)).or_insert(
PublicHealthTimelineBucket {
endpoint_id,
segment_idx,
total_count: 0,
success_count: 0,
failed_count: 0,
min_created_at_unix_ms: Some(candidate.created_at_unix_ms),
max_created_at_unix_ms: Some(candidate.created_at_unix_ms),
},
);
bucket.total_count += 1;
if candidate.status == RequestCandidateStatus::Success {
bucket.success_count += 1;
}
if candidate.status == RequestCandidateStatus::Failed {
bucket.failed_count += 1;
}
bucket.min_created_at_unix_ms = bucket
.min_created_at_unix_ms
.map(|value| value.min(candidate.created_at_unix_ms));
bucket.max_created_at_unix_ms = bucket
.max_created_at_unix_ms
.map(|value| value.max(candidate.created_at_unix_ms));
}
for endpoint_id in endpoint_ids {
for segment_idx in 0..segments {
buckets.entry((endpoint_id.clone(), segment_idx)).or_insert(
PublicHealthTimelineBucket {
endpoint_id: endpoint_id.clone(),
segment_idx,
total_count: 0,
success_count: 0,
failed_count: 0,
min_created_at_unix_ms: None,
max_created_at_unix_ms: None,
},
);
}
}
Ok(buckets.into_values().collect())
}
fn map_candidate_row(row: &SqliteRow) -> Result<StoredRequestCandidate, DataLayerError> {
StoredRequestCandidate::new(
row.try_get("id").map_sql_err()?,
row.try_get("request_id").map_sql_err()?,
row.try_get("user_id").map_sql_err()?,
row.try_get("api_key_id").map_sql_err()?,
row.try_get("username").map_sql_err()?,
row.try_get("api_key_name").map_sql_err()?,
row.try_get("candidate_index").map_sql_err()?,
row.try_get("retry_index").map_sql_err()?,
row.try_get("provider_id").map_sql_err()?,
row.try_get("endpoint_id").map_sql_err()?,
row.try_get("key_id").map_sql_err()?,
RequestCandidateStatus::from_database(
row.try_get::<String, _>("status").map_sql_err()?.as_str(),
)?,
row.try_get("skip_reason").map_sql_err()?,
row.try_get("is_cached").map_sql_err()?,
row.try_get("status_code").map_sql_err()?,
row.try_get("error_type").map_sql_err()?,
row.try_get("error_message").map_sql_err()?,
row.try_get("latency_ms").map_sql_err()?,
row.try_get("concurrent_requests").map_sql_err()?,
parse_json(row.try_get("extra_data").ok().flatten())?,
parse_json(row.try_get("required_capabilities").ok().flatten())?,
row.try_get("created_at_unix_ms").map_sql_err()?,
row.try_get("started_at_unix_ms").map_sql_err()?,
row.try_get("finished_at_unix_ms").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!(
"request_candidates JSON field is invalid: {err}"
))
})
})
.transpose()
}
fn json_to_string(value: &Option<serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.as_ref()
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"request_candidates JSON field is unserializable: {err}"
))
})
})
.transpose()
}
fn merge_json_objects(
existing: Option<serde_json::Value>,
overlay: Option<serde_json::Value>,
) -> Option<serde_json::Value> {
match (existing, overlay) {
(
Some(serde_json::Value::Object(mut existing_object)),
Some(serde_json::Value::Object(overlay_object)),
) => {
existing_object.extend(overlay_object);
Some(serde_json::Value::Object(existing_object))
}
(_existing, Some(overlay)) => Some(overlay),
(existing, None) => existing,
}
}
fn status_to_database(status: RequestCandidateStatus) -> &'static str {
match status {
RequestCandidateStatus::Available => "available",
RequestCandidateStatus::Unused => "unused",
RequestCandidateStatus::Pending => "pending",
RequestCandidateStatus::Streaming => "streaming",
RequestCandidateStatus::Success => "success",
RequestCandidateStatus::Failed => "failed",
RequestCandidateStatus::Cancelled => "cancelled",
RequestCandidateStatus::Skipped => "skipped",
}
}
fn current_unix_ms() -> u64 {
chrono::Utc::now().timestamp_millis().max(0) as u64
}
fn unix_secs_to_ms_i64(value: u64) -> Result<i64, DataLayerError> {
let value = value.checked_mul(1000).ok_or_else(|| {
DataLayerError::UnexpectedValue("request candidate timestamp overflow".to_string())
})?;
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue("request candidate timestamp overflow".to_string())
})
}
fn limit_i64(value: usize, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::UnexpectedValue(format!("invalid {name}: {value}")))
}
fn to_i32(value: u32) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("request candidate value out of range: {value}"))
})
}
fn to_i32_u64(value: u64) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("request candidate value out of range: {value}"))
})
}
fn u64_to_i64(value: u64, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| DataLayerError::UnexpectedValue(format!("{name} overflow")))
}
fn optional_u64_to_i64(value: Option<u64>, name: &str) -> Result<Option<i64>, DataLayerError> {
value.map(|value| u64_to_i64(value, name)).transpose()
}
#[cfg(test)]
mod tests {
use super::SqliteRequestCandidateRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::candidates::{
RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository,
UpsertRequestCandidateRecord,
};
use serde_json::json;
#[tokio::test]
async fn sqlite_repository_writes_and_reads_request_candidates() {
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");
let repository = SqliteRequestCandidateRepository::new(pool);
let created = repository
.upsert(sample_upsert(
"candidate-1",
RequestCandidateStatus::Pending,
Some(json!({"a": 1})),
1_000_000,
))
.await
.expect("candidate should insert");
assert_eq!(created.request_id, "request-1");
let updated = repository
.upsert(sample_upsert(
"candidate-replacement",
RequestCandidateStatus::Success,
Some(json!({"b": 2})),
1_000_500,
))
.await
.expect("candidate should update");
assert_eq!(updated.id, "candidate-1");
assert_eq!(updated.extra_data, Some(json!({"a": 1, "b": 2})));
assert_eq!(
repository
.list_by_request_id("request-1")
.await
.expect("request list should load")
.len(),
1
);
assert_eq!(
repository
.count_finalized_statuses_by_endpoint_ids_since(&["endpoint-1".to_string()], 900)
.await
.expect("status counts should load")[0]
.count,
1
);
assert_eq!(
repository
.aggregate_finalized_timeline_by_endpoint_ids_since(
&["endpoint-1".to_string()],
900,
1200,
3,
)
.await
.expect("timeline should load")
.len(),
3
);
assert_eq!(
repository
.delete_created_before(2_000, 10)
.await
.expect("old candidates should delete"),
1
);
}
fn sample_upsert(
id: &str,
status: RequestCandidateStatus,
extra_data: Option<serde_json::Value>,
created_at_unix_ms: u64,
) -> UpsertRequestCandidateRecord {
UpsertRequestCandidateRecord {
id: id.to_string(),
request_id: "request-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: Some("key-1".to_string()),
username: Some("user".to_string()),
api_key_name: Some("Key".to_string()),
candidate_index: 0,
retry_index: 0,
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("provider-key-1".to_string()),
status,
skip_reason: None,
is_cached: Some(false),
status_code: Some(200),
error_type: None,
error_message: None,
latency_ms: Some(123),
concurrent_requests: Some(2),
extra_data,
required_capabilities: Some(json!({"streaming": true})),
created_at_unix_ms: Some(created_at_unix_ms),
started_at_unix_ms: Some(created_at_unix_ms + 1),
finished_at_unix_ms: Some(created_at_unix_ms + 2),
}
}
}

View File

@@ -1,9 +1,13 @@
pub mod memory;
pub mod sql;
pub mod mysql;
pub mod postgres;
pub mod sqlite;
pub mod types;
pub use memory::InMemoryGeminiFileMappingRepository;
pub use sql::SqlxGeminiFileMappingRepository;
pub use mysql::MysqlGeminiFileMappingRepository;
pub use postgres::SqlxGeminiFileMappingRepository;
pub use sqlite::SqliteGeminiFileMappingRepository;
pub use types::{
GeminiFileMappingListQuery, GeminiFileMappingMimeTypeCount, GeminiFileMappingReadRepository,
GeminiFileMappingRepository, GeminiFileMappingStats, GeminiFileMappingWriteRepository,

View File

@@ -0,0 +1,347 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::types::{
GeminiFileMappingListQuery, GeminiFileMappingMimeTypeCount, GeminiFileMappingReadRepository,
GeminiFileMappingStats, GeminiFileMappingWriteRepository, StoredGeminiFileMapping,
StoredGeminiFileMappingListPage, UpsertGeminiFileMappingRecord,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct MysqlGeminiFileMappingRepository {
pool: MysqlPool,
}
impl MysqlGeminiFileMappingRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn reload_by_file_name(
&self,
file_name: &str,
) -> Result<StoredGeminiFileMapping, DataLayerError> {
self.find_by_file_name(file_name).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("gemini file mapping missing after write".to_string())
})
}
}
#[async_trait]
impl GeminiFileMappingReadRepository for MysqlGeminiFileMappingRepository {
async fn find_by_file_name(
&self,
file_name: &str,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
let row = sqlx::query(
r#"
SELECT
id,
file_name,
key_id,
user_id,
display_name,
mime_type,
source_hash,
created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs
FROM gemini_file_mappings
WHERE file_name = ?
LIMIT 1
"#,
)
.bind(file_name)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_row).transpose()
}
async fn list_mappings(
&self,
query: &GeminiFileMappingListQuery,
) -> Result<StoredGeminiFileMappingListPage, DataLayerError> {
let total = build_list_count_query(query)
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let rows = build_list_rows_query(query)
.build()
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows.iter().map(map_row).collect::<Result<Vec<_>, _>>()?;
Ok(StoredGeminiFileMappingListPage {
items,
total: usize::try_from(total).unwrap_or_default(),
})
}
async fn summarize_mappings(
&self,
now_unix_secs: u64,
) -> Result<GeminiFileMappingStats, DataLayerError> {
let totals = sqlx::query(
r#"
SELECT
COUNT(*) AS total_mappings,
SUM(CASE WHEN expires_at > ? THEN 1 ELSE 0 END) AS active_mappings
FROM gemini_file_mappings
"#,
)
.bind(now_unix_secs as i64)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let total_mappings =
usize::try_from(totals.try_get::<i64, _>("total_mappings").map_sql_err()?)
.unwrap_or_default();
let active_mappings = usize::try_from(
totals
.try_get::<Option<i64>, _>("active_mappings")
.map_sql_err()?
.unwrap_or(0),
)
.unwrap_or_default();
let by_mime_type_rows = sqlx::query(
r#"
SELECT
COALESCE(NULLIF(TRIM(mime_type), ''), 'unknown') AS mime_type,
COUNT(*) AS count
FROM gemini_file_mappings
WHERE expires_at > ?
GROUP BY COALESCE(NULLIF(TRIM(mime_type), ''), 'unknown')
ORDER BY mime_type ASC
"#,
)
.bind(now_unix_secs as i64)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let by_mime_type = by_mime_type_rows
.iter()
.map(|row| {
Ok(GeminiFileMappingMimeTypeCount {
mime_type: row.try_get("mime_type").map_sql_err()?,
count: usize::try_from(row.try_get::<i64, _>("count").map_sql_err()?)
.unwrap_or_default(),
})
})
.collect::<Result<Vec<_>, DataLayerError>>()?;
Ok(GeminiFileMappingStats {
total_mappings,
active_mappings,
expired_mappings: total_mappings.saturating_sub(active_mappings),
by_mime_type,
})
}
}
#[async_trait]
impl GeminiFileMappingWriteRepository for MysqlGeminiFileMappingRepository {
async fn upsert(
&self,
record: UpsertGeminiFileMappingRecord,
) -> Result<StoredGeminiFileMapping, DataLayerError> {
record.validate()?;
sqlx::query(
r#"
INSERT INTO gemini_file_mappings (
id, file_name, key_id, user_id, display_name, mime_type, source_hash,
created_at, expires_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
key_id = VALUES(key_id),
user_id = VALUES(user_id),
display_name = VALUES(display_name),
mime_type = VALUES(mime_type),
source_hash = VALUES(source_hash),
expires_at = VALUES(expires_at)
"#,
)
.bind(&record.id)
.bind(&record.file_name)
.bind(&record.key_id)
.bind(&record.user_id)
.bind(&record.display_name)
.bind(&record.mime_type)
.bind(&record.source_hash)
.bind(current_unix_secs() as i64)
.bind(i64_from_u64(
record.expires_at_unix_secs,
"gemini_file_mappings.expires_at",
)?)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_by_file_name(&record.file_name).await
}
async fn delete_by_file_name(&self, file_name: &str) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query("DELETE FROM gemini_file_mappings WHERE file_name = ?")
.bind(file_name)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn delete_by_id(
&self,
mapping_id: &str,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
let existing = sqlx::query(
r#"
SELECT
id,
file_name,
key_id,
user_id,
display_name,
mime_type,
source_hash,
created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs
FROM gemini_file_mappings
WHERE id = ?
LIMIT 1
"#,
)
.bind(mapping_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
let Some(existing) = existing else {
return Ok(None);
};
sqlx::query("DELETE FROM gemini_file_mappings WHERE id = ?")
.bind(mapping_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(Some(map_row(&existing)?))
}
async fn delete_expired_before(&self, now_unix_secs: u64) -> Result<usize, DataLayerError> {
let rows_affected = sqlx::query("DELETE FROM gemini_file_mappings WHERE expires_at <= ?")
.bind(now_unix_secs as i64)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or_default())
}
}
fn build_list_count_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, MySql> {
let mut builder =
QueryBuilder::<MySql>::new("SELECT COUNT(*) AS total FROM gemini_file_mappings WHERE 1=1");
apply_list_filters(&mut builder, query);
builder
}
fn build_list_rows_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, MySql> {
let mut builder = QueryBuilder::<MySql>::new(
r#"
SELECT
id,
file_name,
key_id,
user_id,
display_name,
mime_type,
source_hash,
created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs
FROM gemini_file_mappings
WHERE 1=1
"#,
);
apply_list_filters(&mut builder, query);
builder.push(" ORDER BY created_at DESC, file_name ASC LIMIT ");
builder.push_bind(i64::try_from(query.limit).unwrap_or(i64::MAX));
builder.push(" OFFSET ");
builder.push_bind(i64::try_from(query.offset).unwrap_or(i64::MAX));
builder
}
fn apply_list_filters(builder: &mut QueryBuilder<'_, MySql>, query: &GeminiFileMappingListQuery) {
if !query.include_expired {
builder.push(" AND expires_at > ");
builder.push_bind(query.now_unix_secs as i64);
}
if let Some(search) = query
.search
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
let pattern = format!("%{}%", search.to_ascii_lowercase());
builder.push(" AND (LOWER(file_name) LIKE ");
builder.push_bind(pattern.clone());
builder.push(" OR LOWER(COALESCE(display_name, '')) LIKE ");
builder.push_bind(pattern);
builder.push(")");
}
}
fn current_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn i64_from_u64(value: u64, field_name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}")))
}
fn map_row(row: &MySqlRow) -> Result<StoredGeminiFileMapping, DataLayerError> {
Ok(StoredGeminiFileMapping {
id: row.try_get("id").map_sql_err()?,
file_name: row.try_get("file_name").map_sql_err()?,
key_id: row.try_get("key_id").map_sql_err()?,
user_id: row.try_get("user_id").ok().flatten(),
display_name: row.try_get("display_name").ok().flatten(),
mime_type: row.try_get("mime_type").ok().flatten(),
source_hash: row.try_get("source_hash").ok().flatten(),
created_at_unix_ms: u64::try_from(
row.try_get::<i64, _>("created_at_unix_ms").map_sql_err()?,
)
.map_err(|_| {
DataLayerError::UnexpectedValue(
"gemini_file_mappings.created_at is invalid".to_string(),
)
})?,
expires_at_unix_secs: u64::try_from(
row.try_get::<i64, _>("expires_at_unix_secs")
.map_sql_err()?,
)
.map_err(|_| {
DataLayerError::UnexpectedValue(
"gemini_file_mappings.expires_at is invalid".to_string(),
)
})?,
})
}
#[cfg(test)]
mod tests {
use super::MysqlGeminiFileMappingRepository;
#[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 = MysqlGeminiFileMappingRepository::new(pool);
}
}

View File

@@ -0,0 +1,420 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::types::{
GeminiFileMappingListQuery, GeminiFileMappingMimeTypeCount, GeminiFileMappingReadRepository,
GeminiFileMappingStats, GeminiFileMappingWriteRepository, StoredGeminiFileMapping,
StoredGeminiFileMappingListPage, UpsertGeminiFileMappingRecord,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct SqliteGeminiFileMappingRepository {
pool: SqlitePool,
}
impl SqliteGeminiFileMappingRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn reload_by_file_name(
&self,
file_name: &str,
) -> Result<StoredGeminiFileMapping, DataLayerError> {
self.find_by_file_name(file_name).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("gemini file mapping missing after write".to_string())
})
}
}
#[async_trait]
impl GeminiFileMappingReadRepository for SqliteGeminiFileMappingRepository {
async fn find_by_file_name(
&self,
file_name: &str,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
let row = sqlx::query(
r#"
SELECT
id,
file_name,
key_id,
user_id,
display_name,
mime_type,
source_hash,
created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs
FROM gemini_file_mappings
WHERE file_name = ?
LIMIT 1
"#,
)
.bind(file_name)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_row).transpose()
}
async fn list_mappings(
&self,
query: &GeminiFileMappingListQuery,
) -> Result<StoredGeminiFileMappingListPage, DataLayerError> {
let total = build_list_count_query(query)
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let rows = build_list_rows_query(query)
.build()
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let items = rows.iter().map(map_row).collect::<Result<Vec<_>, _>>()?;
Ok(StoredGeminiFileMappingListPage {
items,
total: usize::try_from(total).unwrap_or_default(),
})
}
async fn summarize_mappings(
&self,
now_unix_secs: u64,
) -> Result<GeminiFileMappingStats, DataLayerError> {
let totals = sqlx::query(
r#"
SELECT
COUNT(*) AS total_mappings,
SUM(CASE WHEN expires_at > ? THEN 1 ELSE 0 END) AS active_mappings
FROM gemini_file_mappings
"#,
)
.bind(now_unix_secs as i64)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let total_mappings =
usize::try_from(totals.try_get::<i64, _>("total_mappings").map_sql_err()?)
.unwrap_or_default();
let active_mappings = usize::try_from(
totals
.try_get::<Option<i64>, _>("active_mappings")
.map_sql_err()?
.unwrap_or(0),
)
.unwrap_or_default();
let by_mime_type_rows = sqlx::query(
r#"
SELECT
COALESCE(NULLIF(TRIM(mime_type), ''), 'unknown') AS mime_type,
COUNT(*) AS count
FROM gemini_file_mappings
WHERE expires_at > ?
GROUP BY COALESCE(NULLIF(TRIM(mime_type), ''), 'unknown')
ORDER BY mime_type ASC
"#,
)
.bind(now_unix_secs as i64)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let by_mime_type = by_mime_type_rows
.iter()
.map(|row| {
Ok(GeminiFileMappingMimeTypeCount {
mime_type: row.try_get("mime_type").map_sql_err()?,
count: usize::try_from(row.try_get::<i64, _>("count").map_sql_err()?)
.unwrap_or_default(),
})
})
.collect::<Result<Vec<_>, DataLayerError>>()?;
Ok(GeminiFileMappingStats {
total_mappings,
active_mappings,
expired_mappings: total_mappings.saturating_sub(active_mappings),
by_mime_type,
})
}
}
#[async_trait]
impl GeminiFileMappingWriteRepository for SqliteGeminiFileMappingRepository {
async fn upsert(
&self,
record: UpsertGeminiFileMappingRecord,
) -> Result<StoredGeminiFileMapping, DataLayerError> {
record.validate()?;
sqlx::query(
r#"
INSERT INTO gemini_file_mappings (
id, file_name, key_id, user_id, display_name, mime_type, source_hash,
created_at, expires_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(file_name) DO UPDATE SET
key_id = excluded.key_id,
user_id = excluded.user_id,
display_name = excluded.display_name,
mime_type = excluded.mime_type,
source_hash = excluded.source_hash,
expires_at = excluded.expires_at
"#,
)
.bind(&record.id)
.bind(&record.file_name)
.bind(&record.key_id)
.bind(&record.user_id)
.bind(&record.display_name)
.bind(&record.mime_type)
.bind(&record.source_hash)
.bind(current_unix_secs() as i64)
.bind(i64_from_u64(
record.expires_at_unix_secs,
"gemini_file_mappings.expires_at",
)?)
.execute(&self.pool)
.await
.map_sql_err()?;
self.reload_by_file_name(&record.file_name).await
}
async fn delete_by_file_name(&self, file_name: &str) -> Result<bool, DataLayerError> {
let rows_affected = sqlx::query("DELETE FROM gemini_file_mappings WHERE file_name = ?")
.bind(file_name)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(rows_affected > 0)
}
async fn delete_by_id(
&self,
mapping_id: &str,
) -> Result<Option<StoredGeminiFileMapping>, DataLayerError> {
let existing = sqlx::query(
r#"
SELECT
id,
file_name,
key_id,
user_id,
display_name,
mime_type,
source_hash,
created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs
FROM gemini_file_mappings
WHERE id = ?
LIMIT 1
"#,
)
.bind(mapping_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
let Some(existing) = existing else {
return Ok(None);
};
sqlx::query("DELETE FROM gemini_file_mappings WHERE id = ?")
.bind(mapping_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(Some(map_row(&existing)?))
}
async fn delete_expired_before(&self, now_unix_secs: u64) -> Result<usize, DataLayerError> {
let rows_affected = sqlx::query("DELETE FROM gemini_file_mappings WHERE expires_at <= ?")
.bind(now_unix_secs as i64)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or_default())
}
}
fn build_list_count_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, Sqlite> {
let mut builder =
QueryBuilder::<Sqlite>::new("SELECT COUNT(*) AS total FROM gemini_file_mappings WHERE 1=1");
apply_list_filters(&mut builder, query);
builder
}
fn build_list_rows_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, Sqlite> {
let mut builder = QueryBuilder::<Sqlite>::new(
r#"
SELECT
id,
file_name,
key_id,
user_id,
display_name,
mime_type,
source_hash,
created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs
FROM gemini_file_mappings
WHERE 1=1
"#,
);
apply_list_filters(&mut builder, query);
builder.push(" ORDER BY created_at DESC, file_name ASC LIMIT ");
builder.push_bind(i64::try_from(query.limit).unwrap_or(i64::MAX));
builder.push(" OFFSET ");
builder.push_bind(i64::try_from(query.offset).unwrap_or(i64::MAX));
builder
}
fn apply_list_filters(builder: &mut QueryBuilder<'_, Sqlite>, query: &GeminiFileMappingListQuery) {
if !query.include_expired {
builder.push(" AND expires_at > ");
builder.push_bind(query.now_unix_secs as i64);
}
if let Some(search) = query
.search
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
{
let pattern = format!("%{}%", search.to_ascii_lowercase());
builder.push(" AND (LOWER(file_name) LIKE ");
builder.push_bind(pattern.clone());
builder.push(" OR LOWER(COALESCE(display_name, '')) LIKE ");
builder.push_bind(pattern);
builder.push(")");
}
}
fn current_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn i64_from_u64(value: u64, field_name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}")))
}
fn map_row(row: &SqliteRow) -> Result<StoredGeminiFileMapping, DataLayerError> {
Ok(StoredGeminiFileMapping {
id: row.try_get("id").map_sql_err()?,
file_name: row.try_get("file_name").map_sql_err()?,
key_id: row.try_get("key_id").map_sql_err()?,
user_id: row.try_get("user_id").ok().flatten(),
display_name: row.try_get("display_name").ok().flatten(),
mime_type: row.try_get("mime_type").ok().flatten(),
source_hash: row.try_get("source_hash").ok().flatten(),
created_at_unix_ms: u64::try_from(
row.try_get::<i64, _>("created_at_unix_ms").map_sql_err()?,
)
.map_err(|_| {
DataLayerError::UnexpectedValue(
"gemini_file_mappings.created_at is invalid".to_string(),
)
})?,
expires_at_unix_secs: u64::try_from(
row.try_get::<i64, _>("expires_at_unix_secs")
.map_sql_err()?,
)
.map_err(|_| {
DataLayerError::UnexpectedValue(
"gemini_file_mappings.expires_at is invalid".to_string(),
)
})?,
})
}
#[cfg(test)]
mod tests {
use super::SqliteGeminiFileMappingRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::gemini_file_mappings::{
GeminiFileMappingListQuery, GeminiFileMappingReadRepository,
GeminiFileMappingWriteRepository, UpsertGeminiFileMappingRecord,
};
#[tokio::test]
async fn sqlite_repository_round_trips_gemini_file_mappings() {
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");
let repository = SqliteGeminiFileMappingRepository::new(pool);
let created = repository
.upsert(UpsertGeminiFileMappingRecord {
id: "mapping-1".to_string(),
file_name: "files/example.png".to_string(),
key_id: "key-1".to_string(),
user_id: Some("user-1".to_string()),
display_name: Some("Example".to_string()),
mime_type: Some("image/png".to_string()),
source_hash: Some("hash-1".to_string()),
expires_at_unix_secs: 300,
})
.await
.expect("mapping should upsert");
assert_eq!(created.id, "mapping-1");
assert_eq!(created.mime_type, Some("image/png".to_string()));
let updated = repository
.upsert(UpsertGeminiFileMappingRecord {
id: "mapping-replacement".to_string(),
file_name: "files/example.png".to_string(),
key_id: "key-2".to_string(),
user_id: Some("user-2".to_string()),
display_name: Some("Updated".to_string()),
mime_type: Some("image/jpeg".to_string()),
source_hash: Some("hash-2".to_string()),
expires_at_unix_secs: 500,
})
.await
.expect("mapping should update");
assert_eq!(updated.id, "mapping-1");
assert_eq!(updated.key_id, "key-2");
let page = repository
.list_mappings(&GeminiFileMappingListQuery {
include_expired: false,
search: Some("updated".to_string()),
offset: 0,
limit: 10,
now_unix_secs: 400,
})
.await
.expect("mappings should list");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].file_name, "files/example.png");
let stats = repository
.summarize_mappings(400)
.await
.expect("stats should load");
assert_eq!(stats.total_mappings, 1);
assert_eq!(stats.active_mappings, 1);
assert_eq!(stats.by_mime_type[0].mime_type, "image/jpeg");
assert_eq!(
repository
.delete_expired_before(600)
.await
.expect("expired mappings should delete"),
1
);
assert!(repository
.find_by_file_name("files/example.png")
.await
.expect("find should run")
.is_none());
}
}

View File

@@ -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))
}

View File

@@ -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);
}
}

View File

@@ -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

View File

@@ -1,9 +1,13 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
mod types;
pub use memory::InMemoryManagementTokenRepository;
pub use sql::SqlxManagementTokenRepository;
pub use mysql::MysqlManagementTokenRepository;
pub use postgres::SqlxManagementTokenRepository;
pub use sqlite::SqliteManagementTokenRepository;
pub use types::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,

View File

@@ -0,0 +1,500 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use super::types::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
UpdateManagementTokenRecord,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct MysqlManagementTokenRepository {
pool: MysqlPool,
}
impl MysqlManagementTokenRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn get_token(
&self,
token_id: &str,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let row = sqlx::query(TOKEN_BY_ID_SQL)
.bind(token_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_token_row).transpose()
}
}
const TOKEN_BY_ID_SQL: &str = r#"
SELECT
id,
user_id,
name,
description,
token_prefix,
allowed_ips,
expires_at AS expires_at_unix_secs,
last_used_at AS last_used_at_unix_secs,
last_used_ip,
COALESCE(usage_count, 0) AS usage_count,
is_active,
created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM management_tokens
WHERE id = ?
LIMIT 1
"#;
const LIST_MANAGEMENT_TOKENS_SQL: &str = r#"
SELECT
mt.id,
mt.user_id,
mt.name,
mt.description,
mt.token_prefix,
mt.allowed_ips,
mt.expires_at AS expires_at_unix_secs,
mt.last_used_at AS last_used_at_unix_secs,
mt.last_used_ip,
COALESCE(mt.usage_count, 0) AS usage_count,
mt.is_active,
mt.created_at AS created_at_unix_ms,
mt.updated_at AS updated_at_unix_secs,
u.id AS user_row_id,
u.email AS user_email,
u.username AS user_username,
u.role AS user_role
FROM management_tokens mt
JOIN users u ON u.id = mt.user_id
WHERE (? IS NULL OR mt.user_id = ?)
AND (? IS NULL OR mt.is_active = ?)
ORDER BY mt.created_at DESC, mt.id DESC
LIMIT ? OFFSET ?
"#;
const COUNT_MANAGEMENT_TOKENS_SQL: &str = r#"
SELECT COUNT(mt.id) AS total
FROM management_tokens mt
WHERE (? IS NULL OR mt.user_id = ?)
AND (? IS NULL OR mt.is_active = ?)
"#;
const GET_MANAGEMENT_TOKEN_WITH_USER_SQL: &str = r#"
SELECT
mt.id,
mt.user_id,
mt.name,
mt.description,
mt.token_prefix,
mt.allowed_ips,
mt.expires_at AS expires_at_unix_secs,
mt.last_used_at AS last_used_at_unix_secs,
mt.last_used_ip,
COALESCE(mt.usage_count, 0) AS usage_count,
mt.is_active,
mt.created_at AS created_at_unix_ms,
mt.updated_at AS updated_at_unix_secs,
u.id AS user_row_id,
u.email AS user_email,
u.username AS user_username,
u.role AS user_role
FROM management_tokens mt
JOIN users u ON u.id = mt.user_id
WHERE mt.id = ?
LIMIT 1
"#;
const GET_MANAGEMENT_TOKEN_WITH_USER_BY_HASH_SQL: &str = r#"
SELECT
mt.id,
mt.user_id,
mt.name,
mt.description,
mt.token_prefix,
mt.allowed_ips,
mt.expires_at AS expires_at_unix_secs,
mt.last_used_at AS last_used_at_unix_secs,
mt.last_used_ip,
COALESCE(mt.usage_count, 0) AS usage_count,
mt.is_active,
mt.created_at AS created_at_unix_ms,
mt.updated_at AS updated_at_unix_secs,
u.id AS user_row_id,
u.email AS user_email,
u.username AS user_username,
u.role AS user_role
FROM management_tokens mt
JOIN users u ON u.id = mt.user_id
WHERE mt.token_hash = ?
LIMIT 1
"#;
#[async_trait]
impl ManagementTokenReadRepository for MysqlManagementTokenRepository {
async fn list_management_tokens(
&self,
query: &ManagementTokenListQuery,
) -> Result<StoredManagementTokenListPage, DataLayerError> {
let count_row = sqlx::query(COUNT_MANAGEMENT_TOKENS_SQL)
.bind(query.user_id.as_deref())
.bind(query.user_id.as_deref())
.bind(query.is_active)
.bind(query.is_active)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let total = count_row.try_get::<i64, _>("total").map_sql_err()?;
let rows = sqlx::query(LIST_MANAGEMENT_TOKENS_SQL)
.bind(query.user_id.as_deref())
.bind(query.user_id.as_deref())
.bind(query.is_active)
.bind(query.is_active)
.bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
.bind(i64::try_from(query.offset).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
Ok(StoredManagementTokenListPage {
items: rows
.iter()
.map(map_token_with_user_row)
.collect::<Result<Vec<_>, _>>()?,
total: usize::try_from(total.max(0)).unwrap_or(usize::MAX),
})
}
async fn get_management_token_with_user(
&self,
token_id: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_SQL)
.bind(token_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_token_with_user_row).transpose()
}
async fn get_management_token_with_user_by_hash(
&self,
token_hash: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_BY_HASH_SQL)
.bind(token_hash)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_token_with_user_row).transpose()
}
}
#[async_trait]
impl ManagementTokenWriteRepository for MysqlManagementTokenRepository {
async fn create_management_token(
&self,
record: &CreateManagementTokenRecord,
) -> Result<StoredManagementToken, DataLayerError> {
record.validate()?;
let now = now_unix_secs();
sqlx::query(
r#"
INSERT INTO management_tokens (
id, user_id, token_hash, token_prefix, name, description, allowed_ips,
expires_at, is_active, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&record.id)
.bind(&record.user_id)
.bind(&record.token_hash)
.bind(record.token_prefix.as_deref())
.bind(&record.name)
.bind(record.description.as_deref())
.bind(json_to_string(record.allowed_ips.as_ref())?)
.bind(
record
.expires_at_unix_secs
.and_then(|value| i64::try_from(value).ok()),
)
.bind(record.is_active)
.bind(now as i64)
.bind(now as i64)
.execute(&self.pool)
.await
.map_err(|err| map_mysql_write_error(err, Some(record.name.as_str())))?;
self.get_token(&record.id).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("created management token missing".to_string())
})
}
async fn update_management_token(
&self,
record: &UpdateManagementTokenRecord,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
record.validate()?;
let current = self.get_token(&record.token_id).await?;
let Some(current) = current else {
return Ok(None);
};
let name = record.name.as_deref().unwrap_or(&current.name);
let description = if record.clear_description {
None
} else {
record
.description
.as_deref()
.or(current.description.as_deref())
};
let allowed_ips = if record.clear_allowed_ips {
None
} else {
record.allowed_ips.as_ref().or(current.allowed_ips.as_ref())
};
let expires_at = if record.clear_expires_at {
None
} else {
record.expires_at_unix_secs.or(current.expires_at_unix_secs)
};
let is_active = record.is_active.unwrap_or(current.is_active);
let now = now_unix_secs();
let result = sqlx::query(
r#"
UPDATE management_tokens
SET name = ?,
description = ?,
allowed_ips = ?,
expires_at = ?,
is_active = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(name)
.bind(description)
.bind(json_to_string(allowed_ips)?)
.bind(expires_at.and_then(|value| i64::try_from(value).ok()))
.bind(is_active)
.bind(now as i64)
.bind(&record.token_id)
.execute(&self.pool)
.await
.map_err(|err| map_mysql_write_error(err, record.name.as_deref()))?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(&record.token_id).await
}
async fn delete_management_token(&self, token_id: &str) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM management_tokens WHERE id = ?")
.bind(token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn set_management_token_active(
&self,
token_id: &str,
is_active: bool,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let result =
sqlx::query("UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ?")
.bind(is_active)
.bind(now_unix_secs() as i64)
.bind(token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(token_id).await
}
async fn regenerate_management_token_secret(
&self,
mutation: &RegenerateManagementTokenSecret,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
mutation.validate()?;
let result = sqlx::query(
r#"
UPDATE management_tokens
SET token_hash = ?, token_prefix = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(&mutation.token_hash)
.bind(mutation.token_prefix.as_deref())
.bind(now_unix_secs() as i64)
.bind(&mutation.token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(&mutation.token_id).await
}
async fn record_management_token_usage(
&self,
token_id: &str,
last_used_ip: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let now = now_unix_secs();
let result = sqlx::query(
r#"
UPDATE management_tokens
SET last_used_at = ?,
last_used_ip = ?,
usage_count = COALESCE(usage_count, 0) + 1,
updated_at = ?
WHERE id = ?
"#,
)
.bind(now as i64)
.bind(last_used_ip)
.bind(now as i64)
.bind(token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(token_id).await
}
}
fn now_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
value.and_then(|value| u64::try_from(value).ok())
}
fn json_to_string(value: Option<&serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"invalid management token JSON field: {err}"
))
})
})
.transpose()
}
fn json_from_string(value: Option<String>) -> Result<Option<serde_json::Value>, DataLayerError> {
value
.map(|value| {
serde_json::from_str(&value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"invalid management token JSON field: {err}"
))
})
})
.transpose()
}
fn map_mysql_write_error(err: sqlx::Error, requested_name: Option<&str>) -> DataLayerError {
let message = err.to_string();
if message.contains("uq_management_tokens_user_name")
|| message.contains("management_tokens.user_id, management_tokens.name")
{
return DataLayerError::InvalidInput(
requested_name
.map(|name| format!("已存在名为 '{}' 的 Token", name))
.unwrap_or_else(|| "Management Token 名称已存在".to_string()),
);
}
DataLayerError::sql(err)
}
fn map_token_row(row: &MySqlRow) -> Result<StoredManagementToken, DataLayerError> {
Ok(StoredManagementToken::new(
row.try_get("id").map_sql_err()?,
row.try_get("user_id").map_sql_err()?,
row.try_get("name").map_sql_err()?,
)?
.with_display_fields(
row.try_get("description").map_sql_err()?,
row.try_get("token_prefix").map_sql_err()?,
json_from_string(row.try_get("allowed_ips").map_sql_err()?)?,
)
.with_runtime_fields(
optional_unix_secs(row.try_get("expires_at_unix_secs").map_sql_err()?),
optional_unix_secs(row.try_get("last_used_at_unix_secs").map_sql_err()?),
row.try_get("last_used_ip").map_sql_err()?,
u64::try_from(row.try_get::<i64, _>("usage_count").map_sql_err()?).unwrap_or(0),
row.try_get("is_active").map_sql_err()?,
)
.with_timestamps(
optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
optional_unix_secs(row.try_get("updated_at_unix_secs").map_sql_err()?),
))
}
fn map_user_summary_row(
row: &MySqlRow,
) -> Result<StoredManagementTokenUserSummary, DataLayerError> {
StoredManagementTokenUserSummary::new(
row.try_get("user_row_id").map_sql_err()?,
row.try_get("user_email").map_sql_err()?,
row.try_get("user_username").map_sql_err()?,
row.try_get("user_role").map_sql_err()?,
)
}
fn map_token_with_user_row(
row: &MySqlRow,
) -> Result<StoredManagementTokenWithUser, DataLayerError> {
Ok(StoredManagementTokenWithUser::new(
map_token_row(row)?,
map_user_summary_row(row)?,
))
}
#[cfg(test)]
mod tests {
use super::MysqlManagementTokenRepository;
use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
#[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 = MysqlManagementTokenRepository::new(pool);
}
#[test]
fn mysql_management_token_pool_config_remains_driver_specific() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Mysql,
url: "mysql://user:pass@localhost:3306/aether".to_string(),
pool: SqlPoolConfig::default(),
};
assert_eq!(config.driver, DatabaseDriver::Mysql);
}
}

View File

@@ -490,7 +490,7 @@ fn map_token_with_user_row(row: &PgRow) -> Result<StoredManagementTokenWithUser,
#[cfg(test)]
mod tests {
use super::SqlxManagementTokenRepository;
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
#[tokio::test]
async fn repository_constructs_from_lazy_pool() {

View File

@@ -0,0 +1,601 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row, SqlitePool};
use super::types::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
UpdateManagementTokenRecord,
};
use crate::error::SqlResultExt;
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct SqliteManagementTokenRepository {
pool: SqlitePool,
}
impl SqliteManagementTokenRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn get_token(
&self,
token_id: &str,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let row = sqlx::query(TOKEN_BY_ID_SQL)
.bind(token_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_token_row).transpose()
}
}
const TOKEN_BY_ID_SQL: &str = r#"
SELECT
id,
user_id,
name,
description,
token_prefix,
allowed_ips,
expires_at AS expires_at_unix_secs,
last_used_at AS last_used_at_unix_secs,
last_used_ip,
COALESCE(usage_count, 0) AS usage_count,
is_active,
created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM management_tokens
WHERE id = ?
LIMIT 1
"#;
const LIST_MANAGEMENT_TOKENS_SQL: &str = r#"
SELECT
mt.id,
mt.user_id,
mt.name,
mt.description,
mt.token_prefix,
mt.allowed_ips,
mt.expires_at AS expires_at_unix_secs,
mt.last_used_at AS last_used_at_unix_secs,
mt.last_used_ip,
COALESCE(mt.usage_count, 0) AS usage_count,
mt.is_active,
mt.created_at AS created_at_unix_ms,
mt.updated_at AS updated_at_unix_secs,
u.id AS user_row_id,
u.email AS user_email,
u.username AS user_username,
u.role AS user_role
FROM management_tokens mt
JOIN users u ON u.id = mt.user_id
WHERE (? IS NULL OR mt.user_id = ?)
AND (? IS NULL OR mt.is_active = ?)
ORDER BY mt.created_at DESC, mt.id DESC
LIMIT ? OFFSET ?
"#;
const COUNT_MANAGEMENT_TOKENS_SQL: &str = r#"
SELECT COUNT(mt.id) AS total
FROM management_tokens mt
WHERE (? IS NULL OR mt.user_id = ?)
AND (? IS NULL OR mt.is_active = ?)
"#;
const GET_MANAGEMENT_TOKEN_WITH_USER_SQL: &str = r#"
SELECT
mt.id,
mt.user_id,
mt.name,
mt.description,
mt.token_prefix,
mt.allowed_ips,
mt.expires_at AS expires_at_unix_secs,
mt.last_used_at AS last_used_at_unix_secs,
mt.last_used_ip,
COALESCE(mt.usage_count, 0) AS usage_count,
mt.is_active,
mt.created_at AS created_at_unix_ms,
mt.updated_at AS updated_at_unix_secs,
u.id AS user_row_id,
u.email AS user_email,
u.username AS user_username,
u.role AS user_role
FROM management_tokens mt
JOIN users u ON u.id = mt.user_id
WHERE mt.id = ?
LIMIT 1
"#;
const GET_MANAGEMENT_TOKEN_WITH_USER_BY_HASH_SQL: &str = r#"
SELECT
mt.id,
mt.user_id,
mt.name,
mt.description,
mt.token_prefix,
mt.allowed_ips,
mt.expires_at AS expires_at_unix_secs,
mt.last_used_at AS last_used_at_unix_secs,
mt.last_used_ip,
COALESCE(mt.usage_count, 0) AS usage_count,
mt.is_active,
mt.created_at AS created_at_unix_ms,
mt.updated_at AS updated_at_unix_secs,
u.id AS user_row_id,
u.email AS user_email,
u.username AS user_username,
u.role AS user_role
FROM management_tokens mt
JOIN users u ON u.id = mt.user_id
WHERE mt.token_hash = ?
LIMIT 1
"#;
#[async_trait]
impl ManagementTokenReadRepository for SqliteManagementTokenRepository {
async fn list_management_tokens(
&self,
query: &ManagementTokenListQuery,
) -> Result<StoredManagementTokenListPage, DataLayerError> {
let count_row = sqlx::query(COUNT_MANAGEMENT_TOKENS_SQL)
.bind(query.user_id.as_deref())
.bind(query.user_id.as_deref())
.bind(query.is_active)
.bind(query.is_active)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let total = count_row.try_get::<i64, _>("total").map_sql_err()?;
let rows = sqlx::query(LIST_MANAGEMENT_TOKENS_SQL)
.bind(query.user_id.as_deref())
.bind(query.user_id.as_deref())
.bind(query.is_active)
.bind(query.is_active)
.bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
.bind(i64::try_from(query.offset).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
Ok(StoredManagementTokenListPage {
items: rows
.iter()
.map(map_token_with_user_row)
.collect::<Result<Vec<_>, _>>()?,
total: usize::try_from(total.max(0)).unwrap_or(usize::MAX),
})
}
async fn get_management_token_with_user(
&self,
token_id: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_SQL)
.bind(token_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_token_with_user_row).transpose()
}
async fn get_management_token_with_user_by_hash(
&self,
token_hash: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_BY_HASH_SQL)
.bind(token_hash)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_token_with_user_row).transpose()
}
}
#[async_trait]
impl ManagementTokenWriteRepository for SqliteManagementTokenRepository {
async fn create_management_token(
&self,
record: &CreateManagementTokenRecord,
) -> Result<StoredManagementToken, DataLayerError> {
record.validate()?;
let now = now_unix_secs();
sqlx::query(
r#"
INSERT INTO management_tokens (
id, user_id, token_hash, token_prefix, name, description, allowed_ips,
expires_at, is_active, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(&record.id)
.bind(&record.user_id)
.bind(&record.token_hash)
.bind(record.token_prefix.as_deref())
.bind(&record.name)
.bind(record.description.as_deref())
.bind(json_to_string(record.allowed_ips.as_ref())?)
.bind(
record
.expires_at_unix_secs
.and_then(|value| i64::try_from(value).ok()),
)
.bind(record.is_active)
.bind(now as i64)
.bind(now as i64)
.execute(&self.pool)
.await
.map_err(|err| map_sqlite_write_error(err, Some(record.name.as_str())))?;
self.get_token(&record.id).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("created management token missing".to_string())
})
}
async fn update_management_token(
&self,
record: &UpdateManagementTokenRecord,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
record.validate()?;
let current = self.get_token(&record.token_id).await?;
let Some(current) = current else {
return Ok(None);
};
let name = record.name.as_deref().unwrap_or(&current.name);
let description = if record.clear_description {
None
} else {
record
.description
.as_deref()
.or(current.description.as_deref())
};
let allowed_ips = if record.clear_allowed_ips {
None
} else {
record.allowed_ips.as_ref().or(current.allowed_ips.as_ref())
};
let expires_at = if record.clear_expires_at {
None
} else {
record.expires_at_unix_secs.or(current.expires_at_unix_secs)
};
let is_active = record.is_active.unwrap_or(current.is_active);
let now = now_unix_secs();
let result = sqlx::query(
r#"
UPDATE management_tokens
SET name = ?,
description = ?,
allowed_ips = ?,
expires_at = ?,
is_active = ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(name)
.bind(description)
.bind(json_to_string(allowed_ips)?)
.bind(expires_at.and_then(|value| i64::try_from(value).ok()))
.bind(is_active)
.bind(now as i64)
.bind(&record.token_id)
.execute(&self.pool)
.await
.map_err(|err| map_sqlite_write_error(err, record.name.as_deref()))?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(&record.token_id).await
}
async fn delete_management_token(&self, token_id: &str) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM management_tokens WHERE id = ?")
.bind(token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
async fn set_management_token_active(
&self,
token_id: &str,
is_active: bool,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let result =
sqlx::query("UPDATE management_tokens SET is_active = ?, updated_at = ? WHERE id = ?")
.bind(is_active)
.bind(now_unix_secs() as i64)
.bind(token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(token_id).await
}
async fn regenerate_management_token_secret(
&self,
mutation: &RegenerateManagementTokenSecret,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
mutation.validate()?;
let result = sqlx::query(
r#"
UPDATE management_tokens
SET token_hash = ?, token_prefix = ?, updated_at = ?
WHERE id = ?
"#,
)
.bind(&mutation.token_hash)
.bind(mutation.token_prefix.as_deref())
.bind(now_unix_secs() as i64)
.bind(&mutation.token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(&mutation.token_id).await
}
async fn record_management_token_usage(
&self,
token_id: &str,
last_used_ip: Option<&str>,
) -> Result<Option<StoredManagementToken>, DataLayerError> {
let now = now_unix_secs();
let result = sqlx::query(
r#"
UPDATE management_tokens
SET last_used_at = ?,
last_used_ip = ?,
usage_count = COALESCE(usage_count, 0) + 1,
updated_at = ?
WHERE id = ?
"#,
)
.bind(now as i64)
.bind(last_used_ip)
.bind(now as i64)
.bind(token_id)
.execute(&self.pool)
.await
.map_sql_err()?;
if result.rows_affected() == 0 {
return Ok(None);
}
self.get_token(token_id).await
}
}
fn now_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
value.and_then(|value| u64::try_from(value).ok())
}
fn json_to_string(value: Option<&serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"invalid management token JSON field: {err}"
))
})
})
.transpose()
}
fn json_from_string(value: Option<String>) -> Result<Option<serde_json::Value>, DataLayerError> {
value
.map(|value| {
serde_json::from_str(&value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"invalid management token JSON field: {err}"
))
})
})
.transpose()
}
fn map_sqlite_write_error(err: sqlx::Error, requested_name: Option<&str>) -> DataLayerError {
let message = err.to_string();
if message.contains("management_tokens.user_id, management_tokens.name") {
return DataLayerError::InvalidInput(
requested_name
.map(|name| format!("已存在名为 '{}' 的 Token", name))
.unwrap_or_else(|| "Management Token 名称已存在".to_string()),
);
}
DataLayerError::sql(err)
}
fn map_token_row(row: &SqliteRow) -> Result<StoredManagementToken, DataLayerError> {
Ok(StoredManagementToken::new(
row.try_get("id").map_sql_err()?,
row.try_get("user_id").map_sql_err()?,
row.try_get("name").map_sql_err()?,
)?
.with_display_fields(
row.try_get("description").map_sql_err()?,
row.try_get("token_prefix").map_sql_err()?,
json_from_string(row.try_get("allowed_ips").map_sql_err()?)?,
)
.with_runtime_fields(
optional_unix_secs(row.try_get("expires_at_unix_secs").map_sql_err()?),
optional_unix_secs(row.try_get("last_used_at_unix_secs").map_sql_err()?),
row.try_get("last_used_ip").map_sql_err()?,
u64::try_from(row.try_get::<i64, _>("usage_count").map_sql_err()?).unwrap_or(0),
row.try_get("is_active").map_sql_err()?,
)
.with_timestamps(
optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
optional_unix_secs(row.try_get("updated_at_unix_secs").map_sql_err()?),
))
}
fn map_user_summary_row(
row: &SqliteRow,
) -> Result<StoredManagementTokenUserSummary, DataLayerError> {
StoredManagementTokenUserSummary::new(
row.try_get("user_row_id").map_sql_err()?,
row.try_get("user_email").map_sql_err()?,
row.try_get("user_username").map_sql_err()?,
row.try_get("user_role").map_sql_err()?,
)
}
fn map_token_with_user_row(
row: &SqliteRow,
) -> Result<StoredManagementTokenWithUser, DataLayerError> {
Ok(StoredManagementTokenWithUser::new(
map_token_row(row)?,
map_user_summary_row(row)?,
))
}
#[cfg(test)]
mod tests {
use super::SqliteManagementTokenRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::management_tokens::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
ManagementTokenWriteRepository, RegenerateManagementTokenSecret,
StoredManagementTokenUserSummary, UpdateManagementTokenRecord,
};
#[tokio::test]
async fn sqlite_repository_round_trips_management_tokens() {
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");
sqlx::query(
r#"
INSERT INTO users (id, email, username, role, is_active, created_at, updated_at)
VALUES ('user-1', 'user-1@example.com', 'user-1', 'admin', 1, 1, 1)
"#,
)
.execute(&pool)
.await
.expect("seed user should insert");
let repository = SqliteManagementTokenRepository::new(pool);
let user = StoredManagementTokenUserSummary::new(
"user-1".to_string(),
Some("user-1@example.com".to_string()),
"user-1".to_string(),
"admin".to_string(),
)
.expect("user summary should build");
let created = repository
.create_management_token(&CreateManagementTokenRecord {
id: "token-1".to_string(),
user_id: "user-1".to_string(),
user,
token_hash: "hash-1".to_string(),
token_prefix: Some("ae_1234".to_string()),
name: "primary".to_string(),
description: Some("primary token".to_string()),
allowed_ips: Some(serde_json::json!(["127.0.0.1"])),
expires_at_unix_secs: Some(1_800_000_000),
is_active: true,
})
.await
.expect("token should create");
assert_eq!(created.name, "primary");
let page = repository
.list_management_tokens(&ManagementTokenListQuery {
user_id: Some("user-1".to_string()),
is_active: Some(true),
offset: 0,
limit: 10,
})
.await
.expect("tokens should list");
assert_eq!(page.total, 1);
assert_eq!(page.items[0].token.id, "token-1");
let by_hash = repository
.get_management_token_with_user_by_hash("hash-1")
.await
.expect("hash lookup should succeed")
.expect("token should exist");
assert_eq!(by_hash.user.username, "user-1");
let updated = repository
.update_management_token(&UpdateManagementTokenRecord {
token_id: "token-1".to_string(),
name: Some("renamed".to_string()),
description: None,
clear_description: true,
allowed_ips: Some(serde_json::json!(["10.0.0.1"])),
clear_allowed_ips: false,
expires_at_unix_secs: None,
clear_expires_at: true,
is_active: Some(false),
})
.await
.expect("update should succeed")
.expect("token should exist");
assert_eq!(updated.name, "renamed");
assert!(!updated.is_active);
assert_eq!(updated.description, None);
assert_eq!(updated.expires_at_unix_secs, None);
let toggled = repository
.set_management_token_active("token-1", true)
.await
.expect("toggle should succeed")
.expect("token should exist");
assert!(toggled.is_active);
let regenerated = repository
.regenerate_management_token_secret(&RegenerateManagementTokenSecret {
token_id: "token-1".to_string(),
token_hash: "hash-2".to_string(),
token_prefix: Some("ae_5678".to_string()),
})
.await
.expect("regenerate should succeed")
.expect("token should exist");
assert_eq!(regenerated.token_prefix.as_deref(), Some("ae_5678"));
assert!(repository
.get_management_token_with_user_by_hash("hash-1")
.await
.expect("old hash lookup should succeed")
.is_none());
let used = repository
.record_management_token_usage("token-1", Some("127.0.0.1"))
.await
.expect("usage should record")
.expect("token should exist");
assert_eq!(used.usage_count, 1);
assert_eq!(used.last_used_ip.as_deref(), Some("127.0.0.1"));
assert!(repository
.delete_management_token("token-1")
.await
.expect("delete should succeed"));
}
}

View File

@@ -1,3 +1,9 @@
//! Domain repository implementations.
//!
//! Repository contracts and shared DTOs are re-exported from
//! `aether-data-contracts` where possible. Driver-specific files under each
//! domain translate those contracts to concrete Postgres/MySQL/SQLite SQL.
pub mod announcements;
pub mod audit;
pub mod auth;

View File

@@ -1,9 +1,13 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
mod types;
pub use memory::InMemoryOAuthProviderRepository;
pub use sql::SqlxOAuthProviderRepository;
pub use mysql::MysqlOAuthProviderRepository;
pub use postgres::SqlxOAuthProviderRepository;
pub use sqlite::SqliteOAuthProviderRepository;
pub use types::{
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderRepository,
OAuthProviderWriteRepository, StoredOAuthProviderConfig, UpsertOAuthProviderConfigRecord,

View File

@@ -0,0 +1,386 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use super::types::{
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
UpsertOAuthProviderConfigRecord,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct MysqlOAuthProviderRepository {
pool: MysqlPool,
}
impl MysqlOAuthProviderRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn get_provider(
&self,
provider_type: &str,
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL)
.bind(provider_type)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_oauth_provider_row).transpose()
}
}
const LIST_OAUTH_PROVIDER_CONFIGS_SQL: &str = r#"
SELECT
provider_type,
display_name,
client_id,
client_secret_encrypted,
authorization_url_override,
token_url_override,
userinfo_url_override,
scopes,
redirect_uri,
frontend_callback_url,
attribute_mapping,
extra_config,
is_enabled,
created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM oauth_providers
ORDER BY provider_type ASC
"#;
const GET_OAUTH_PROVIDER_CONFIG_SQL: &str = r#"
SELECT
provider_type,
display_name,
client_id,
client_secret_encrypted,
authorization_url_override,
token_url_override,
userinfo_url_override,
scopes,
redirect_uri,
frontend_callback_url,
attribute_mapping,
extra_config,
is_enabled,
created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM oauth_providers
WHERE provider_type = ?
LIMIT 1
"#;
const COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL: &str = r#"
SELECT COUNT(DISTINCT users.id) AS locked_count
FROM users
JOIN user_oauth_links
ON users.id = user_oauth_links.user_id
WHERE users.is_active = 1
AND users.is_deleted = 0
AND user_oauth_links.provider_type = ?
AND (
(
users.auth_source = 'oauth'
AND NOT EXISTS (
SELECT 1
FROM user_oauth_links other_links
JOIN oauth_providers other_provider
ON other_links.provider_type = other_provider.provider_type
WHERE other_links.user_id = users.id
AND other_links.provider_type <> ?
AND other_provider.is_enabled = 1
)
) OR (
? = 1
AND users.auth_source = 'local'
AND users.role <> 'admin'
AND NOT EXISTS (
SELECT 1
FROM user_oauth_links other_links
JOIN oauth_providers other_provider
ON other_links.provider_type = other_provider.provider_type
WHERE other_links.user_id = users.id
AND other_links.provider_type <> ?
AND other_provider.is_enabled = 1
)
)
)
"#;
#[async_trait]
impl OAuthProviderReadRepository for MysqlOAuthProviderRepository {
async fn list_oauth_provider_configs(
&self,
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
let rows = sqlx::query(LIST_OAUTH_PROVIDER_CONFIGS_SQL)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_oauth_provider_row).collect()
}
async fn get_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
self.get_provider(provider_type).await
}
async fn count_locked_users_if_provider_disabled(
&self,
provider_type: &str,
ldap_exclusive: bool,
) -> Result<usize, DataLayerError> {
let row = sqlx::query(COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL)
.bind(provider_type)
.bind(provider_type)
.bind(ldap_exclusive)
.bind(provider_type)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let locked_count = row.try_get::<i64, _>("locked_count").map_sql_err()?;
usize::try_from(locked_count.max(0)).map_err(|_| {
DataLayerError::UnexpectedValue(
"oauth_providers.locked_user_count overflowed".to_string(),
)
})
}
}
#[async_trait]
impl OAuthProviderWriteRepository for MysqlOAuthProviderRepository {
async fn upsert_oauth_provider_config(
&self,
record: &UpsertOAuthProviderConfigRecord,
) -> Result<StoredOAuthProviderConfig, DataLayerError> {
record.validate()?;
let now = now_unix_secs();
sqlx::query(
r#"
INSERT INTO oauth_providers (
provider_type,
display_name,
client_id,
client_secret_encrypted,
authorization_url_override,
token_url_override,
userinfo_url_override,
scopes,
redirect_uri,
frontend_callback_url,
attribute_mapping,
extra_config,
is_enabled,
created_at,
updated_at
) VALUES (
?, ?, ?,
CASE ? WHEN 'set' THEN ? WHEN 'clear' THEN NULL ELSE NULL END,
?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?
)
ON DUPLICATE KEY UPDATE
display_name = VALUES(display_name),
client_id = VALUES(client_id),
client_secret_encrypted = CASE ?
WHEN 'set' THEN ?
WHEN 'clear' THEN NULL
ELSE client_secret_encrypted
END,
authorization_url_override = VALUES(authorization_url_override),
token_url_override = VALUES(token_url_override),
userinfo_url_override = VALUES(userinfo_url_override),
scopes = VALUES(scopes),
redirect_uri = VALUES(redirect_uri),
frontend_callback_url = VALUES(frontend_callback_url),
attribute_mapping = VALUES(attribute_mapping),
extra_config = VALUES(extra_config),
is_enabled = VALUES(is_enabled),
updated_at = VALUES(updated_at)
"#,
)
.bind(&record.provider_type)
.bind(&record.display_name)
.bind(&record.client_id)
.bind(record.client_secret_encrypted.mode_name())
.bind(record.client_secret_encrypted.value())
.bind(record.authorization_url_override.as_deref())
.bind(record.token_url_override.as_deref())
.bind(record.userinfo_url_override.as_deref())
.bind(scopes_to_json_string(record.scopes.as_ref())?)
.bind(&record.redirect_uri)
.bind(&record.frontend_callback_url)
.bind(json_to_string(record.attribute_mapping.as_ref())?)
.bind(json_to_string(record.extra_config.as_ref())?)
.bind(record.is_enabled)
.bind(now as i64)
.bind(now as i64)
.bind(record.client_secret_encrypted.mode_name())
.bind(record.client_secret_encrypted.value())
.execute(&self.pool)
.await
.map_sql_err()?;
self.get_provider(&record.provider_type)
.await?
.ok_or_else(|| {
DataLayerError::UnexpectedValue("upserted OAuth provider missing".to_string())
})
}
async fn delete_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM oauth_providers WHERE provider_type = ?")
.bind(provider_type)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
}
fn now_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
value.and_then(|value| u64::try_from(value).ok())
}
fn json_to_string(value: Option<&serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("invalid OAuth provider JSON field: {err}"))
})
})
.transpose()
}
fn 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 scopes_to_json_string(scopes: Option<&Vec<String>>) -> Result<Option<String>, DataLayerError> {
let value = scopes.map(|items| {
serde_json::Value::Array(
items
.iter()
.cloned()
.map(serde_json::Value::String)
.collect(),
)
});
json_to_string(value.as_ref())
}
fn parse_scopes(value: Option<String>) -> Result<Option<Vec<String>>, DataLayerError> {
let Some(value) = json_from_string(value, "oauth_providers.scopes")? else {
return Ok(None);
};
parse_scopes_value(&value)
}
fn parse_scopes_value(value: &serde_json::Value) -> Result<Option<Vec<String>>, DataLayerError> {
match value {
serde_json::Value::Null => Ok(None),
serde_json::Value::Array(items) => parse_scopes_array(items).map(Some),
serde_json::Value::String(raw) => parse_embedded_scopes(raw),
_ => Err(DataLayerError::UnexpectedValue(
"oauth_providers.scopes is not a JSON array".to_string(),
)),
}
}
fn parse_embedded_scopes(raw: &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_scopes_value(&decoded);
}
Ok(Some(vec![raw.to_string()]))
}
fn parse_scopes_array(items: &[serde_json::Value]) -> Result<Vec<String>, DataLayerError> {
let mut scopes = Vec::with_capacity(items.len());
for item in items {
let Some(scope) = item.as_str() else {
return Err(DataLayerError::UnexpectedValue(
"oauth_providers.scopes contains non-string value".to_string(),
));
};
let scope = scope.trim();
if !scope.is_empty() {
scopes.push(scope.to_string());
}
}
Ok(scopes)
}
fn map_oauth_provider_row(row: &MySqlRow) -> Result<StoredOAuthProviderConfig, DataLayerError> {
Ok(StoredOAuthProviderConfig::new(
row.try_get("provider_type").map_sql_err()?,
row.try_get("display_name").map_sql_err()?,
row.try_get("client_id").map_sql_err()?,
row.try_get("redirect_uri").map_sql_err()?,
row.try_get("frontend_callback_url").map_sql_err()?,
)?
.with_config_fields(
row.try_get("client_secret_encrypted").map_sql_err()?,
row.try_get("authorization_url_override").map_sql_err()?,
row.try_get("token_url_override").map_sql_err()?,
row.try_get("userinfo_url_override").map_sql_err()?,
parse_scopes(row.try_get("scopes").map_sql_err()?)?,
json_from_string(
row.try_get("attribute_mapping").map_sql_err()?,
"oauth_providers.attribute_mapping",
)?,
json_from_string(
row.try_get("extra_config").map_sql_err()?,
"oauth_providers.extra_config",
)?,
row.try_get("is_enabled").map_sql_err()?,
)
.with_timestamps(
optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
optional_unix_secs(row.try_get("updated_at_unix_secs").map_sql_err()?),
))
}
#[cfg(test)]
mod tests {
use super::MysqlOAuthProviderRepository;
#[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 = MysqlOAuthProviderRepository::new(pool);
}
}

View File

@@ -354,7 +354,7 @@ fn map_oauth_provider_row(row: &PgRow) -> Result<StoredOAuthProviderConfig, Data
mod tests {
use super::{parse_scopes, SqlxOAuthProviderRepository};
use crate::{
postgres::{PostgresPoolConfig, PostgresPoolFactory},
driver::postgres::{PostgresPoolConfig, PostgresPoolFactory},
DataLayerError,
};

View File

@@ -0,0 +1,498 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row};
use super::types::{
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
UpsertOAuthProviderConfigRecord,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct SqliteOAuthProviderRepository {
pool: SqlitePool,
}
impl SqliteOAuthProviderRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn get_provider(
&self,
provider_type: &str,
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL)
.bind(provider_type)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_oauth_provider_row).transpose()
}
}
const LIST_OAUTH_PROVIDER_CONFIGS_SQL: &str = r#"
SELECT
provider_type,
display_name,
client_id,
client_secret_encrypted,
authorization_url_override,
token_url_override,
userinfo_url_override,
scopes,
redirect_uri,
frontend_callback_url,
attribute_mapping,
extra_config,
is_enabled,
created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM oauth_providers
ORDER BY provider_type ASC
"#;
const GET_OAUTH_PROVIDER_CONFIG_SQL: &str = r#"
SELECT
provider_type,
display_name,
client_id,
client_secret_encrypted,
authorization_url_override,
token_url_override,
userinfo_url_override,
scopes,
redirect_uri,
frontend_callback_url,
attribute_mapping,
extra_config,
is_enabled,
created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM oauth_providers
WHERE provider_type = ?
LIMIT 1
"#;
const COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL: &str = r#"
SELECT COUNT(DISTINCT users.id) AS locked_count
FROM users
JOIN user_oauth_links
ON users.id = user_oauth_links.user_id
WHERE users.is_active = 1
AND users.is_deleted = 0
AND user_oauth_links.provider_type = ?
AND (
(
users.auth_source = 'oauth'
AND NOT EXISTS (
SELECT 1
FROM user_oauth_links other_links
JOIN oauth_providers other_provider
ON other_links.provider_type = other_provider.provider_type
WHERE other_links.user_id = users.id
AND other_links.provider_type <> ?
AND other_provider.is_enabled = 1
)
) OR (
? = 1
AND users.auth_source = 'local'
AND users.role <> 'admin'
AND NOT EXISTS (
SELECT 1
FROM user_oauth_links other_links
JOIN oauth_providers other_provider
ON other_links.provider_type = other_provider.provider_type
WHERE other_links.user_id = users.id
AND other_links.provider_type <> ?
AND other_provider.is_enabled = 1
)
)
)
"#;
#[async_trait]
impl OAuthProviderReadRepository for SqliteOAuthProviderRepository {
async fn list_oauth_provider_configs(
&self,
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
let rows = sqlx::query(LIST_OAUTH_PROVIDER_CONFIGS_SQL)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_oauth_provider_row).collect()
}
async fn get_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
self.get_provider(provider_type).await
}
async fn count_locked_users_if_provider_disabled(
&self,
provider_type: &str,
ldap_exclusive: bool,
) -> Result<usize, DataLayerError> {
let row = sqlx::query(COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL)
.bind(provider_type)
.bind(provider_type)
.bind(ldap_exclusive)
.bind(provider_type)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let locked_count = row.try_get::<i64, _>("locked_count").map_sql_err()?;
usize::try_from(locked_count.max(0)).map_err(|_| {
DataLayerError::UnexpectedValue(
"oauth_providers.locked_user_count overflowed".to_string(),
)
})
}
}
#[async_trait]
impl OAuthProviderWriteRepository for SqliteOAuthProviderRepository {
async fn upsert_oauth_provider_config(
&self,
record: &UpsertOAuthProviderConfigRecord,
) -> Result<StoredOAuthProviderConfig, DataLayerError> {
record.validate()?;
let now = now_unix_secs();
sqlx::query(
r#"
INSERT INTO oauth_providers (
provider_type,
display_name,
client_id,
client_secret_encrypted,
authorization_url_override,
token_url_override,
userinfo_url_override,
scopes,
redirect_uri,
frontend_callback_url,
attribute_mapping,
extra_config,
is_enabled,
created_at,
updated_at
) VALUES (
?, ?, ?,
CASE ? WHEN 'set' THEN ? WHEN 'clear' THEN NULL ELSE NULL END,
?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?
)
ON CONFLICT(provider_type) DO UPDATE SET
display_name = excluded.display_name,
client_id = excluded.client_id,
client_secret_encrypted = CASE ?
WHEN 'set' THEN ?
WHEN 'clear' THEN NULL
ELSE oauth_providers.client_secret_encrypted
END,
authorization_url_override = excluded.authorization_url_override,
token_url_override = excluded.token_url_override,
userinfo_url_override = excluded.userinfo_url_override,
scopes = excluded.scopes,
redirect_uri = excluded.redirect_uri,
frontend_callback_url = excluded.frontend_callback_url,
attribute_mapping = excluded.attribute_mapping,
extra_config = excluded.extra_config,
is_enabled = excluded.is_enabled,
updated_at = excluded.updated_at
"#,
)
.bind(&record.provider_type)
.bind(&record.display_name)
.bind(&record.client_id)
.bind(record.client_secret_encrypted.mode_name())
.bind(record.client_secret_encrypted.value())
.bind(record.authorization_url_override.as_deref())
.bind(record.token_url_override.as_deref())
.bind(record.userinfo_url_override.as_deref())
.bind(scopes_to_json_string(record.scopes.as_ref())?)
.bind(&record.redirect_uri)
.bind(&record.frontend_callback_url)
.bind(json_to_string(record.attribute_mapping.as_ref())?)
.bind(json_to_string(record.extra_config.as_ref())?)
.bind(record.is_enabled)
.bind(now as i64)
.bind(now as i64)
.bind(record.client_secret_encrypted.mode_name())
.bind(record.client_secret_encrypted.value())
.execute(&self.pool)
.await
.map_sql_err()?;
self.get_provider(&record.provider_type)
.await?
.ok_or_else(|| {
DataLayerError::UnexpectedValue("upserted OAuth provider missing".to_string())
})
}
async fn delete_oauth_provider_config(
&self,
provider_type: &str,
) -> Result<bool, DataLayerError> {
let result = sqlx::query("DELETE FROM oauth_providers WHERE provider_type = ?")
.bind(provider_type)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() > 0)
}
}
fn now_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
value.and_then(|value| u64::try_from(value).ok())
}
fn json_to_string(value: Option<&serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("invalid OAuth provider JSON field: {err}"))
})
})
.transpose()
}
fn 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 scopes_to_json_string(scopes: Option<&Vec<String>>) -> Result<Option<String>, DataLayerError> {
json_to_string(
scopes
.map(|items| {
serde_json::Value::Array(
items
.iter()
.cloned()
.map(serde_json::Value::String)
.collect(),
)
})
.as_ref(),
)
}
fn parse_scopes(value: Option<String>) -> Result<Option<Vec<String>>, DataLayerError> {
let Some(value) = json_from_string(value, "oauth_providers.scopes")? else {
return Ok(None);
};
parse_scopes_value(&value)
}
fn parse_scopes_value(value: &serde_json::Value) -> Result<Option<Vec<String>>, DataLayerError> {
match value {
serde_json::Value::Null => Ok(None),
serde_json::Value::Array(items) => parse_scopes_array(items).map(Some),
serde_json::Value::String(raw) => parse_embedded_scopes(raw),
_ => Err(DataLayerError::UnexpectedValue(
"oauth_providers.scopes is not a JSON array".to_string(),
)),
}
}
fn parse_embedded_scopes(raw: &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_scopes_value(&decoded);
}
Ok(Some(vec![raw.to_string()]))
}
fn parse_scopes_array(items: &[serde_json::Value]) -> Result<Vec<String>, DataLayerError> {
let mut scopes = Vec::with_capacity(items.len());
for item in items {
let Some(scope) = item.as_str() else {
return Err(DataLayerError::UnexpectedValue(
"oauth_providers.scopes contains non-string value".to_string(),
));
};
let scope = scope.trim();
if !scope.is_empty() {
scopes.push(scope.to_string());
}
}
Ok(scopes)
}
fn map_oauth_provider_row(row: &SqliteRow) -> Result<StoredOAuthProviderConfig, DataLayerError> {
Ok(StoredOAuthProviderConfig::new(
row.try_get("provider_type").map_sql_err()?,
row.try_get("display_name").map_sql_err()?,
row.try_get("client_id").map_sql_err()?,
row.try_get("redirect_uri").map_sql_err()?,
row.try_get("frontend_callback_url").map_sql_err()?,
)?
.with_config_fields(
row.try_get("client_secret_encrypted").map_sql_err()?,
row.try_get("authorization_url_override").map_sql_err()?,
row.try_get("token_url_override").map_sql_err()?,
row.try_get("userinfo_url_override").map_sql_err()?,
parse_scopes(row.try_get("scopes").map_sql_err()?)?,
json_from_string(
row.try_get("attribute_mapping").map_sql_err()?,
"oauth_providers.attribute_mapping",
)?,
json_from_string(
row.try_get("extra_config").map_sql_err()?,
"oauth_providers.extra_config",
)?,
row.try_get("is_enabled").map_sql_err()?,
)
.with_timestamps(
optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
optional_unix_secs(row.try_get("updated_at_unix_secs").map_sql_err()?),
))
}
#[cfg(test)]
mod tests {
use super::SqliteOAuthProviderRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::oauth_providers::{
EncryptedSecretUpdate, OAuthProviderReadRepository, OAuthProviderWriteRepository,
UpsertOAuthProviderConfigRecord,
};
fn sample_upsert(provider_type: &str) -> UpsertOAuthProviderConfigRecord {
UpsertOAuthProviderConfigRecord {
provider_type: provider_type.to_string(),
display_name: format!("{provider_type} display"),
client_id: format!("{provider_type}-client"),
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
authorization_url_override: Some(format!("https://{provider_type}.example.com/auth")),
token_url_override: Some(format!("https://{provider_type}.example.com/token")),
userinfo_url_override: None,
scopes: Some(vec!["openid".to_string(), "profile".to_string()]),
redirect_uri: format!("https://{provider_type}.example.com/redirect"),
frontend_callback_url: "https://frontend.example.com/auth/callback".to_string(),
attribute_mapping: Some(serde_json::json!({"email": "email"})),
extra_config: Some(serde_json::json!({"team": true})),
is_enabled: true,
}
}
#[tokio::test]
async fn sqlite_repository_round_trips_oauth_provider_configs() {
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");
let repository = SqliteOAuthProviderRepository::new(pool.clone());
let created = repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Set("secret-1".to_string()),
..sample_upsert("github")
})
.await
.expect("provider should upsert");
assert_eq!(created.client_secret_encrypted.as_deref(), Some("secret-1"));
assert_eq!(
created.scopes,
Some(vec!["openid".to_string(), "profile".to_string()])
);
let updated = repository
.upsert_oauth_provider_config(&UpsertOAuthProviderConfigRecord {
client_secret_encrypted: EncryptedSecretUpdate::Preserve,
display_name: "GitHub".to_string(),
..sample_upsert("github")
})
.await
.expect("provider should update");
assert_eq!(updated.display_name, "GitHub");
assert_eq!(updated.client_secret_encrypted.as_deref(), Some("secret-1"));
let listed = repository
.list_oauth_provider_configs()
.await
.expect("providers should list");
assert_eq!(listed.len(), 1);
let fetched = repository
.get_oauth_provider_config("github")
.await
.expect("provider should fetch")
.expect("provider should exist");
assert_eq!(
fetched.attribute_mapping,
Some(serde_json::json!({"email": "email"}))
);
sqlx::query(
r#"
INSERT INTO users (
id, email, username, role, auth_source, is_active, is_deleted, created_at, updated_at
) VALUES
('user-oauth', 'oauth@example.com', 'oauth-user', 'user', 'oauth', 1, 0, 1, 1),
('user-local', 'local@example.com', 'local-user', 'user', 'local', 1, 0, 1, 1)
"#,
)
.execute(&pool)
.await
.expect("users should seed");
sqlx::query(
r#"
INSERT INTO user_oauth_links (
id, user_id, provider_type, provider_user_id, linked_at
) VALUES
('link-1', 'user-oauth', 'github', 'gh-1', 1),
('link-2', 'user-local', 'github', 'gh-2', 1)
"#,
)
.execute(&pool)
.await
.expect("oauth links should seed");
assert_eq!(
repository
.count_locked_users_if_provider_disabled("github", false)
.await
.expect("locked users should count"),
1
);
assert_eq!(
repository
.count_locked_users_if_provider_disabled("github", true)
.await
.expect("locked users should count"),
2
);
assert!(repository
.delete_oauth_provider_config("github")
.await
.expect("provider should delete"));
}
}

View File

@@ -1,5 +1,7 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::provider_catalog::{
@@ -8,4 +10,6 @@ pub(crate) use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
};
pub use memory::InMemoryProviderCatalogReadRepository;
pub use sql::SqlxProviderCatalogReadRepository;
pub use mysql::MysqlProviderCatalogReadRepository;
pub use postgres::SqlxProviderCatalogReadRepository;
pub use sqlite::SqliteProviderCatalogReadRepository;

File diff suppressed because it is too large Load Diff

View File

@@ -142,7 +142,7 @@ SELECT
api_formats,
auth_type_by_format,
allow_auth_channel_mismatch_formats,
api_key,
COALESCE(api_key, encrypted_key) AS api_key,
auth_config,
note,
internal_priority,
@@ -201,7 +201,7 @@ SELECT
api_formats,
auth_type_by_format,
allow_auth_channel_mismatch_formats,
api_key,
COALESCE(api_key, encrypted_key) AS api_key,
auth_config,
note,
internal_priority,
@@ -602,7 +602,7 @@ SELECT
api_formats,
auth_type_by_format,
allow_auth_channel_mismatch_formats,
api_key,
COALESCE(api_key, encrypted_key) AS api_key,
auth_config,
note,
internal_priority,
@@ -2526,7 +2526,7 @@ fn map_key_row(row: &PgRow) -> Result<StoredProviderCatalogKey, DataLayerError>
#[cfg(test)]
mod tests {
use super::SqlxProviderCatalogReadRepository;
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
#[tokio::test]
async fn repository_constructs_from_lazy_pool() {
@@ -2568,7 +2568,7 @@ mod tests {
assert!(sql.contains("concurrent_limit"));
}
let source = include_str!("sql.rs");
let source = include_str!("postgres.rs");
assert!(source.contains("concurrent_limit,"));
assert!(source.contains("concurrent_limit = $13"));
assert!(source.contains(".bind(key.concurrent_limit)"));
@@ -2578,7 +2578,7 @@ mod tests {
#[test]
fn provider_api_keys_concurrent_limit_schema_is_nullable_without_default() {
let migration = include_str!(
"../../../migrations/20260502000000_add_provider_key_auth_channel_mismatch_formats.sql"
"../../../migrations/postgres/20260502000000_add_provider_key_auth_channel_mismatch_formats.sql"
);
let concurrent_limit_line = migration
.lines()
@@ -2590,7 +2590,7 @@ mod tests {
"add column if not exists concurrent_limit integer;"
);
let baseline = include_str!("../../../bootstrap/20260413020000_baseline_v2.sql");
let baseline = crate::lifecycle::bootstrap::postgres::EMPTY_DATABASE_SNAPSHOT_SQL;
assert!(baseline.contains("CREATE TABLE IF NOT EXISTS public.provider_api_keys"));
assert!(baseline.contains("concurrent_limit integer,"));
}
@@ -2604,10 +2604,12 @@ mod tests {
assert!(sql.contains("allow_auth_channel_mismatch_formats"));
}
let source = include_str!("sql.rs");
let source = include_str!("postgres.rs");
assert!(
source
.matches("auth_type_by_format,\n allow_auth_channel_mismatch_formats,\n api_key",)
.matches(
"auth_type_by_format,\n allow_auth_channel_mismatch_formats,\n COALESCE(api_key, encrypted_key) AS api_key",
)
.count()
>= 3
);
@@ -2617,7 +2619,7 @@ mod tests {
#[test]
fn provider_api_keys_create_key_insert_placeholders_match_bind_order() {
let source = include_str!("sql.rs");
let source = include_str!("postgres.rs");
assert!(source.contains(
" $24,\n $25,\n $26,\n CASE\n WHEN $27::double precision IS NULL THEN NULL"
));

File diff suppressed because it is too large Load Diff

View File

@@ -1,9 +1,13 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
mod types;
pub use memory::InMemoryProxyNodeRepository;
pub use sql::SqlxProxyNodeRepository;
pub use mysql::MysqlProxyNodeReadRepository;
pub use postgres::SqlxProxyNodeRepository;
pub use sqlite::SqliteProxyNodeReadRepository;
pub use types::{
normalize_proxy_node_scheduling_state, proxy_node_accepts_new_tunnels, proxy_reported_version,
reconcile_remote_config_after_heartbeat, remote_config_scheduling_state,

View File

@@ -0,0 +1,882 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use super::types::{
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeHeartbeatMutation,
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeReadRepository,
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyNode, StoredProxyNodeEvent,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct MysqlProxyNodeReadRepository {
pool: MysqlPool,
}
impl MysqlProxyNodeReadRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn upsert_node(&self, node: &StoredProxyNode) -> Result<(), DataLayerError> {
let now = current_unix_secs();
sqlx::query(
r#"
INSERT INTO proxy_nodes (
id, name, ip, port, region, status, registered_by, last_heartbeat_at,
heartbeat_interval, active_connections, total_requests, avg_latency_ms,
is_manual, proxy_url, proxy_username, proxy_password, created_at,
updated_at, remote_config, config_version, hardware_info,
estimated_max_concurrency, tunnel_mode, tunnel_connected, tunnel_connected_at,
failed_requests, dns_failures, stream_errors, proxy_metadata
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
name = VALUES(name),
ip = VALUES(ip),
port = VALUES(port),
region = VALUES(region),
status = VALUES(status),
registered_by = VALUES(registered_by),
last_heartbeat_at = VALUES(last_heartbeat_at),
heartbeat_interval = VALUES(heartbeat_interval),
active_connections = VALUES(active_connections),
total_requests = VALUES(total_requests),
avg_latency_ms = VALUES(avg_latency_ms),
is_manual = VALUES(is_manual),
proxy_url = VALUES(proxy_url),
proxy_username = VALUES(proxy_username),
proxy_password = VALUES(proxy_password),
updated_at = VALUES(updated_at),
remote_config = VALUES(remote_config),
config_version = VALUES(config_version),
hardware_info = VALUES(hardware_info),
estimated_max_concurrency = VALUES(estimated_max_concurrency),
tunnel_mode = VALUES(tunnel_mode),
tunnel_connected = VALUES(tunnel_connected),
tunnel_connected_at = VALUES(tunnel_connected_at),
failed_requests = VALUES(failed_requests),
dns_failures = VALUES(dns_failures),
stream_errors = VALUES(stream_errors),
proxy_metadata = VALUES(proxy_metadata)
"#,
)
.bind(&node.id)
.bind(&node.name)
.bind(&node.ip)
.bind(node.port)
.bind(&node.region)
.bind(&node.status)
.bind(&node.registered_by)
.bind(optional_i64_from_u64(
node.last_heartbeat_at_unix_secs,
"proxy_nodes.last_heartbeat_at",
)?)
.bind(node.heartbeat_interval)
.bind(node.active_connections)
.bind(node.total_requests)
.bind(node.avg_latency_ms)
.bind(node.is_manual)
.bind(&node.proxy_url)
.bind(&node.proxy_username)
.bind(&node.proxy_password)
.bind(node.created_at_unix_ms.unwrap_or(now) as i64)
.bind(node.updated_at_unix_secs.unwrap_or(now) as i64)
.bind(optional_json_to_string(
&node.remote_config,
"proxy_nodes.remote_config",
)?)
.bind(node.config_version)
.bind(optional_json_to_string(
&node.hardware_info,
"proxy_nodes.hardware_info",
)?)
.bind(node.estimated_max_concurrency)
.bind(node.tunnel_mode)
.bind(node.tunnel_connected)
.bind(optional_i64_from_u64(
node.tunnel_connected_at_unix_secs,
"proxy_nodes.tunnel_connected_at",
)?)
.bind(node.failed_requests)
.bind(node.dns_failures)
.bind(node.stream_errors)
.bind(optional_json_to_string(
&node.proxy_metadata,
"proxy_nodes.proxy_metadata",
)?)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(())
}
async fn find_duplicate_proxy_node(
&self,
ip: &str,
port: i32,
excluding_node_id: Option<&str>,
) -> Result<Option<StoredProxyNode>, DataLayerError> {
let row = if let Some(excluding_node_id) = excluding_node_id {
sqlx::query(&format!(
"{PROXY_NODE_COLUMNS} WHERE ip = ? AND port = ? AND id <> ? LIMIT 1"
))
.bind(ip)
.bind(port)
.bind(excluding_node_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?
} else {
sqlx::query(&format!(
"{PROXY_NODE_COLUMNS} WHERE ip = ? AND port = ? LIMIT 1"
))
.bind(ip)
.bind(port)
.fetch_optional(&self.pool)
.await
.map_sql_err()?
};
row.as_ref().map(map_proxy_node_row).transpose()
}
async fn insert_event(
&self,
node_id: &str,
event_type: &str,
detail: Option<&str>,
created_at_unix_secs: Option<u64>,
) -> Result<(), DataLayerError> {
sqlx::query(
r#"
INSERT INTO proxy_node_events (node_id, event_type, detail, created_at)
VALUES (?, ?, ?, ?)
"#,
)
.bind(node_id)
.bind(event_type)
.bind(detail)
.bind(created_at_unix_secs.unwrap_or_else(current_unix_secs) as i64)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(())
}
fn normalize_remote_config(
mutation: &ProxyNodeRemoteConfigMutation,
existing: Option<&serde_json::Value>,
) -> Option<serde_json::Value> {
let mut config = match existing {
Some(serde_json::Value::Object(map)) => map.clone(),
_ => serde_json::Map::new(),
};
if let Some(node_name) = mutation.node_name.as_ref() {
config.insert(
"node_name".to_string(),
serde_json::Value::String(node_name.clone()),
);
}
if let Some(allowed_ports) = mutation.allowed_ports.as_ref() {
config.insert(
"allowed_ports".to_string(),
serde_json::json!(allowed_ports),
);
}
if let Some(log_level) = mutation.log_level.as_ref() {
config.insert(
"log_level".to_string(),
serde_json::Value::String(log_level.clone()),
);
}
if let Some(heartbeat_interval) = mutation.heartbeat_interval {
config.insert(
"heartbeat_interval".to_string(),
serde_json::json!(heartbeat_interval),
);
}
if let Some(scheduling_state) = mutation.scheduling_state.as_ref() {
match scheduling_state {
Some(state) => {
config.insert(
"scheduling_state".to_string(),
serde_json::Value::String(state.clone()),
);
}
None => {
config.remove("scheduling_state");
}
}
}
if let Some(upgrade_to) = mutation.upgrade_to.as_ref() {
match upgrade_to {
Some(version) => {
config.insert(
"upgrade_to".to_string(),
serde_json::Value::String(version.clone()),
);
}
None => {
config.remove("upgrade_to");
}
}
}
(!config.is_empty()).then_some(serde_json::Value::Object(config))
}
}
const PROXY_NODE_COLUMNS: &str = r#"
SELECT
id,
name,
ip,
port,
region,
is_manual,
proxy_url,
proxy_username,
proxy_password,
status,
registered_by,
last_heartbeat_at AS last_heartbeat_at_unix_secs,
heartbeat_interval,
active_connections,
total_requests,
avg_latency_ms,
failed_requests,
dns_failures,
stream_errors,
proxy_metadata,
hardware_info,
estimated_max_concurrency,
tunnel_mode,
tunnel_connected,
tunnel_connected_at AS tunnel_connected_at_unix_secs,
remote_config,
config_version,
created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs
FROM proxy_nodes
"#;
#[async_trait]
impl ProxyNodeReadRepository for MysqlProxyNodeReadRepository {
async fn list_proxy_nodes(&self) -> Result<Vec<StoredProxyNode>, DataLayerError> {
let rows = sqlx::query(&format!("{PROXY_NODE_COLUMNS} ORDER BY name ASC, id ASC"))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_proxy_node_row).collect()
}
async fn find_proxy_node(
&self,
node_id: &str,
) -> Result<Option<StoredProxyNode>, DataLayerError> {
let row = sqlx::query(&format!("{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(node_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_proxy_node_row).transpose()
}
async fn list_proxy_node_events(
&self,
node_id: &str,
limit: usize,
) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> {
let rows = sqlx::query(
r#"
SELECT
id,
node_id,
event_type,
detail,
created_at AS created_at_unix_ms
FROM proxy_node_events
WHERE node_id = ?
ORDER BY created_at DESC, id DESC
LIMIT ?
"#,
)
.bind(node_id)
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_proxy_node_event_row).collect()
}
}
#[async_trait]
impl ProxyNodeWriteRepository for MysqlProxyNodeReadRepository {
async fn reset_stale_tunnel_statuses(&self) -> Result<usize, DataLayerError> {
let now = current_unix_secs() as i64;
let result = sqlx::query(
r#"
UPDATE proxy_nodes
SET tunnel_connected = 0,
status = 'offline',
active_connections = 0,
tunnel_connected_at = ?,
updated_at = ?
WHERE is_manual = 0
AND tunnel_connected = 1
"#,
)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(result.rows_affected() as usize)
}
async fn create_manual_node(
&self,
mutation: &ProxyNodeManualCreateMutation,
) -> Result<StoredProxyNode, DataLayerError> {
if let Some(existing) = self
.find_duplicate_proxy_node(&mutation.ip, mutation.port, None)
.await?
{
return Err(duplicate_proxy_node_error(&existing));
}
let now = Some(current_unix_secs());
let node = StoredProxyNode::new(
uuid::Uuid::new_v4().to_string(),
mutation.name.clone(),
mutation.ip.clone(),
mutation.port,
true,
"online".to_string(),
0,
0,
0,
0,
0,
0,
false,
false,
0,
)?
.with_manual_proxy_fields(
Some(mutation.proxy_url.clone()),
mutation.proxy_username.clone(),
mutation.proxy_password.clone(),
)
.with_runtime_fields(
mutation.region.clone(),
mutation.registered_by.clone(),
None,
None,
None,
None,
None,
None,
None,
now,
now,
);
self.upsert_node(&node).await?;
Ok(node)
}
async fn update_manual_node(
&self,
mutation: &ProxyNodeManualUpdateMutation,
) -> Result<Option<StoredProxyNode>, DataLayerError> {
let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else {
return Ok(None);
};
if !node.is_manual {
return Err(DataLayerError::InvalidInput(
"只能编辑手动添加的代理节点".to_string(),
));
}
let next_ip = mutation.ip.as_deref().unwrap_or(node.ip.as_str());
let next_port = mutation.port.unwrap_or(node.port);
if let Some(existing) = self
.find_duplicate_proxy_node(next_ip, next_port, Some(&mutation.node_id))
.await?
{
return Err(duplicate_proxy_node_error(&existing));
}
if let Some(name) = mutation.name.as_ref() {
node.name = name.clone();
}
if let Some(ip) = mutation.ip.as_ref() {
node.ip = ip.clone();
}
if let Some(port) = mutation.port {
node.port = port;
}
if let Some(region) = mutation.region.as_ref() {
node.region = Some(region.clone());
}
if let Some(proxy_url) = mutation.proxy_url.as_ref() {
node.proxy_url = Some(proxy_url.clone());
}
if let Some(proxy_username) = mutation.proxy_username.as_ref() {
node.proxy_username = Some(proxy_username.clone());
}
if let Some(proxy_password) = mutation.proxy_password.as_ref() {
node.proxy_password = Some(proxy_password.clone());
}
node.updated_at_unix_secs = Some(current_unix_secs());
self.upsert_node(&node).await?;
Ok(Some(node))
}
async fn register_node(
&self,
mutation: &ProxyNodeRegistrationMutation,
) -> Result<StoredProxyNode, DataLayerError> {
let now = Some(current_unix_secs());
let normalized_proxy_metadata = normalize_proxy_metadata(
mutation.proxy_metadata.as_ref(),
mutation.proxy_version.as_deref(),
);
let existing = sqlx::query(&format!(
"{PROXY_NODE_COLUMNS} WHERE ip = ? AND port = ? AND is_manual = 0 ORDER BY created_at ASC, id ASC LIMIT 1"
))
.bind(&mutation.ip)
.bind(mutation.port)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
let mut node = if let Some(row) = existing.as_ref() {
map_proxy_node_row(row)?
} else {
StoredProxyNode::new(
uuid::Uuid::new_v4().to_string(),
mutation.name.clone(),
mutation.ip.clone(),
mutation.port,
false,
"offline".to_string(),
mutation.heartbeat_interval,
mutation.active_connections.unwrap_or(0),
mutation.total_requests.unwrap_or(0),
0,
0,
0,
mutation.tunnel_mode,
false,
0,
)?
.with_runtime_fields(
mutation.region.clone(),
mutation.registered_by.clone(),
now,
mutation.avg_latency_ms,
normalized_proxy_metadata.clone(),
mutation.hardware_info.clone(),
mutation.estimated_max_concurrency,
None,
None,
now,
now,
)
};
node.name = mutation.name.clone();
node.ip = mutation.ip.clone();
node.port = mutation.port;
node.region = mutation.region.clone();
node.registered_by = mutation.registered_by.clone();
node.last_heartbeat_at_unix_secs = now;
node.heartbeat_interval = mutation.heartbeat_interval;
node.tunnel_mode = mutation.tunnel_mode;
if let Some(active_connections) = mutation.active_connections {
node.active_connections = active_connections;
}
if let Some(total_requests) = mutation.total_requests {
node.total_requests = total_requests;
}
if let Some(avg_latency_ms) = mutation.avg_latency_ms {
node.avg_latency_ms = Some(avg_latency_ms);
}
if let Some(hardware_info) = mutation.hardware_info.as_ref() {
node.hardware_info = Some(hardware_info.clone());
}
if let Some(estimated_max_concurrency) = mutation.estimated_max_concurrency {
node.estimated_max_concurrency = Some(estimated_max_concurrency);
}
if let Some(proxy_metadata) = normalized_proxy_metadata {
node.proxy_metadata = Some(proxy_metadata);
}
if node.created_at_unix_ms.is_none() {
node.created_at_unix_ms = now;
}
node.updated_at_unix_secs = now;
self.upsert_node(&node).await?;
Ok(node)
}
async fn apply_heartbeat(
&self,
mutation: &ProxyNodeHeartbeatMutation,
) -> Result<Option<StoredProxyNode>, DataLayerError> {
let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else {
return Ok(None);
};
if !node.tunnel_mode {
return Err(DataLayerError::InvalidInput(
"non-tunnel mode is no longer supported, please upgrade aether-proxy to use tunnel mode"
.to_string(),
));
}
let now = Some(current_unix_secs());
node.last_heartbeat_at_unix_secs = now;
if node.status != "online" || !node.tunnel_connected {
node.status = "online".to_string();
node.tunnel_connected = true;
node.tunnel_connected_at_unix_secs = now;
node.updated_at_unix_secs = now;
}
if let Some(value) = mutation.heartbeat_interval {
node.heartbeat_interval = value;
}
if let Some(value) = mutation.active_connections {
node.active_connections = value;
}
if let Some(value) = mutation.avg_latency_ms {
node.avg_latency_ms = Some(value);
}
if let Some(value) = normalize_proxy_metadata(
mutation.proxy_metadata.as_ref(),
mutation.proxy_version.as_deref(),
) {
node.proxy_metadata = Some(value);
}
if let Some(value) = mutation.total_requests_delta.filter(|value| *value > 0) {
node.total_requests += value;
}
if let Some(value) = mutation.failed_requests_delta.filter(|value| *value > 0) {
node.failed_requests += value;
}
if let Some(value) = mutation.dns_failures_delta.filter(|value| *value > 0) {
node.dns_failures += value;
}
if let Some(value) = mutation.stream_errors_delta.filter(|value| *value > 0) {
node.stream_errors += value;
}
let reconciled_remote_config = reconcile_remote_config_after_heartbeat(
node.remote_config.as_ref(),
mutation.proxy_version.as_deref(),
);
if reconciled_remote_config != node.remote_config {
node.remote_config = reconciled_remote_config;
node.config_version = node.config_version.saturating_add(1);
node.updated_at_unix_secs = now;
}
self.upsert_node(&node).await?;
Ok(Some(node))
}
async fn record_traffic(
&self,
mutation: &ProxyNodeTrafficMutation,
) -> Result<bool, DataLayerError> {
let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else {
return Ok(false);
};
if !node.is_manual {
return Ok(false);
}
node.total_requests += mutation.total_requests_delta.max(0);
node.failed_requests += mutation.failed_requests_delta.max(0);
node.dns_failures += mutation.dns_failures_delta.max(0);
node.stream_errors += mutation.stream_errors_delta.max(0);
node.updated_at_unix_secs = Some(current_unix_secs());
self.upsert_node(&node).await?;
Ok(true)
}
async fn update_tunnel_status(
&self,
mutation: &ProxyNodeTunnelStatusMutation,
) -> Result<Option<StoredProxyNode>, DataLayerError> {
let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else {
return Ok(None);
};
let event_time = mutation
.observed_at_unix_secs
.unwrap_or_else(current_unix_secs);
let event_type = if mutation.connected {
"connected"
} else {
"disconnected"
};
let event_detail = mutation.detail.clone().unwrap_or_else(|| {
format!(
"[tunnel_node_status] conn_count={}",
i32::max(mutation.conn_count, 0)
)
});
if node
.tunnel_connected_at_unix_secs
.is_some_and(|last_transition| event_time < last_transition)
{
self.insert_event(
&mutation.node_id,
event_type,
Some(&format!("[stale_ignored] {event_detail}")),
Some(current_unix_secs()),
)
.await?;
return Ok(Some(node));
}
node.tunnel_connected = mutation.connected;
node.tunnel_connected_at_unix_secs = Some(event_time);
node.status = if mutation.connected {
"online".to_string()
} else {
"offline".to_string()
};
if !mutation.connected {
node.active_connections = 0;
}
node.updated_at_unix_secs = Some(event_time);
self.upsert_node(&node).await?;
self.insert_event(
&mutation.node_id,
event_type,
Some(&event_detail),
Some(event_time),
)
.await?;
Ok(Some(node))
}
async fn unregister_node(
&self,
node_id: &str,
) -> Result<Option<StoredProxyNode>, DataLayerError> {
let Some(mut node) = self.find_proxy_node(node_id).await? else {
return Ok(None);
};
let now = Some(current_unix_secs());
node.status = "offline".to_string();
node.tunnel_connected = false;
node.active_connections = 0;
node.tunnel_connected_at_unix_secs = now;
node.updated_at_unix_secs = now;
self.upsert_node(&node).await?;
Ok(Some(node))
}
async fn delete_node(&self, node_id: &str) -> Result<Option<StoredProxyNode>, DataLayerError> {
let existing = self.find_proxy_node(node_id).await?;
if existing.is_some() {
sqlx::query("DELETE FROM proxy_node_events WHERE node_id = ?")
.bind(node_id)
.execute(&self.pool)
.await
.map_sql_err()?;
sqlx::query("DELETE FROM proxy_nodes WHERE id = ?")
.bind(node_id)
.execute(&self.pool)
.await
.map_sql_err()?;
}
Ok(existing)
}
async fn update_remote_config(
&self,
mutation: &ProxyNodeRemoteConfigMutation,
) -> Result<Option<StoredProxyNode>, DataLayerError> {
let Some(mut node) = self.find_proxy_node(&mutation.node_id).await? else {
return Ok(None);
};
if node.is_manual {
return Err(DataLayerError::InvalidInput(
"手动节点不支持远程配置下发".to_string(),
));
}
if let Some(node_name) = mutation.node_name.as_ref() {
node.name = node_name.clone();
}
node.remote_config = Self::normalize_remote_config(mutation, node.remote_config.as_ref());
node.config_version = node.config_version.saturating_add(1);
node.updated_at_unix_secs = Some(current_unix_secs());
self.upsert_node(&node).await?;
Ok(Some(node))
}
async fn increment_manual_node_requests(
&self,
node_id: &str,
total_delta: i64,
failed_delta: i64,
latency_ms: Option<i64>,
) -> Result<(), DataLayerError> {
let Some(mut node) = self.find_proxy_node(node_id).await? else {
return Ok(());
};
if !node.is_manual {
return Ok(());
}
if total_delta > 0 {
node.total_requests += total_delta;
}
if failed_delta > 0 {
node.failed_requests += failed_delta;
}
if let Some(ms) = latency_ms {
node.avg_latency_ms = Some(ms as f64);
}
node.updated_at_unix_secs = Some(current_unix_secs());
self.upsert_node(&node).await
}
}
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
value.and_then(|value| u64::try_from(value).ok())
}
fn current_unix_secs() -> u64 {
chrono::Utc::now().timestamp().max(0) as u64
}
fn optional_i64_from_u64(
value: Option<u64>,
field_name: &str,
) -> Result<Option<i64>, DataLayerError> {
value
.map(|value| {
i64::try_from(value).map_err(|_| {
DataLayerError::InvalidInput(format!("{field_name} exceeds i64: {value}"))
})
})
.transpose()
}
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 duplicate_proxy_node_error(node: &StoredProxyNode) -> DataLayerError {
DataLayerError::InvalidInput(format!(
"已存在相同地址的代理节点: {} ({}:{})",
node.name, node.ip, node.port
))
}
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 map_proxy_node_row(row: &MySqlRow) -> Result<StoredProxyNode, DataLayerError> {
Ok(StoredProxyNode::new(
row.try_get("id").map_sql_err()?,
row.try_get("name").map_sql_err()?,
row.try_get("ip").map_sql_err()?,
row.try_get("port").map_sql_err()?,
row.try_get("is_manual").map_sql_err()?,
row.try_get("status").map_sql_err()?,
row.try_get("heartbeat_interval").map_sql_err()?,
row.try_get("active_connections").map_sql_err()?,
row.try_get("total_requests").map_sql_err()?,
row.try_get("failed_requests").map_sql_err()?,
row.try_get("dns_failures").map_sql_err()?,
row.try_get("stream_errors").map_sql_err()?,
row.try_get("tunnel_mode").map_sql_err()?,
row.try_get("tunnel_connected").map_sql_err()?,
row.try_get("config_version").map_sql_err()?,
)?
.with_manual_proxy_fields(
row.try_get("proxy_url").map_sql_err()?,
row.try_get("proxy_username").map_sql_err()?,
row.try_get("proxy_password").map_sql_err()?,
)
.with_runtime_fields(
row.try_get("region").map_sql_err()?,
row.try_get("registered_by").map_sql_err()?,
optional_unix_secs(row.try_get("last_heartbeat_at_unix_secs").map_sql_err()?),
row.try_get("avg_latency_ms").map_sql_err()?,
optional_json_from_string(
row.try_get("proxy_metadata").map_sql_err()?,
"proxy_nodes.proxy_metadata",
)?,
optional_json_from_string(
row.try_get("hardware_info").map_sql_err()?,
"proxy_nodes.hardware_info",
)?,
row.try_get("estimated_max_concurrency").map_sql_err()?,
optional_unix_secs(row.try_get("tunnel_connected_at_unix_secs").map_sql_err()?),
optional_json_from_string(
row.try_get("remote_config").map_sql_err()?,
"proxy_nodes.remote_config",
)?,
optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
optional_unix_secs(row.try_get("updated_at_unix_secs").map_sql_err()?),
))
}
fn map_proxy_node_event_row(row: &MySqlRow) -> Result<StoredProxyNodeEvent, DataLayerError> {
Ok(StoredProxyNodeEvent {
id: row.try_get("id").map_sql_err()?,
node_id: row.try_get("node_id").map_sql_err()?,
event_type: row.try_get("event_type").map_sql_err()?,
detail: row.try_get("detail").map_sql_err()?,
created_at_unix_ms: optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
})
}
#[cfg(test)]
mod tests {
use super::MysqlProxyNodeReadRepository;
#[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 = MysqlProxyNodeReadRepository::new(pool);
}
}

File diff suppressed because it is too large Load Diff

View File

@@ -1,5 +1,7 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::quota::{
@@ -7,4 +9,6 @@ pub(crate) use aether_data_contracts::repository::quota::{
StoredProviderQuotaSnapshot,
};
pub use memory::InMemoryProviderQuotaRepository;
pub use sql::SqlxProviderQuotaRepository;
pub use mysql::MysqlProviderQuotaRepository;
pub use postgres::SqlxProviderQuotaRepository;
pub use sqlite::SqliteProviderQuotaRepository;

View File

@@ -0,0 +1,129 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const QUOTA_COLUMNS: &str = r#"
SELECT
id AS provider_id,
billing_type,
monthly_quota_usd,
COALESCE(monthly_used_usd, 0) AS monthly_used_usd,
quota_reset_day,
quota_last_reset_at AS quota_last_reset_at_unix_secs,
quota_expires_at AS quota_expires_at_unix_secs,
is_active
FROM providers
"#;
#[derive(Debug, Clone)]
pub struct MysqlProviderQuotaRepository {
pool: MysqlPool,
}
impl MysqlProviderQuotaRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
}
#[async_trait]
impl ProviderQuotaReadRepository for MysqlProviderQuotaRepository {
async fn find_by_provider_id(
&self,
provider_id: &str,
) -> Result<Option<StoredProviderQuotaSnapshot>, DataLayerError> {
let row = sqlx::query(&format!("{QUOTA_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(provider_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_row).transpose()
}
async fn find_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderQuotaSnapshot>, DataLayerError> {
if provider_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(QUOTA_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for provider_id in provider_ids {
separated.push_bind(provider_id);
}
}
builder.push(") ORDER BY id ASC");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_row).collect()
}
}
#[async_trait]
impl ProviderQuotaWriteRepository for MysqlProviderQuotaRepository {
async fn reset_due(&self, now_unix_secs: u64) -> Result<usize, DataLayerError> {
let now = i64::try_from(now_unix_secs).map_err(|_| {
DataLayerError::InvalidInput("provider quota reset timestamp overflow".to_string())
})?;
let rows_affected = sqlx::query(
r#"
UPDATE providers
SET monthly_used_usd = 0,
quota_last_reset_at = ?,
updated_at = ?
WHERE billing_type = 'monthly_quota'
AND is_active = 1
AND (
quota_last_reset_at IS NULL
OR (? - quota_last_reset_at) >= (quota_reset_day * 86400)
)
"#,
)
.bind(now)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or_default())
}
}
fn map_row(row: &MySqlRow) -> Result<StoredProviderQuotaSnapshot, DataLayerError> {
StoredProviderQuotaSnapshot::new(
row.try_get("provider_id").map_sql_err()?,
row.try_get("billing_type").map_sql_err()?,
row.try_get("monthly_quota_usd").map_sql_err()?,
row.try_get("monthly_used_usd").map_sql_err()?,
row.try_get("quota_reset_day").map_sql_err()?,
row.try_get("quota_last_reset_at_unix_secs").map_sql_err()?,
row.try_get("quota_expires_at_unix_secs").map_sql_err()?,
row.try_get("is_active").map_sql_err()?,
)
}
#[cfg(test)]
mod tests {
use super::MysqlProviderQuotaRepository;
#[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 = MysqlProviderQuotaRepository::new(pool);
}
}

View File

@@ -127,7 +127,7 @@ fn map_row(row: &sqlx::postgres::PgRow) -> Result<StoredProviderQuotaSnapshot, D
#[cfg(test)]
mod tests {
use super::SqlxProviderQuotaRepository;
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
#[tokio::test]
async fn repository_constructs_from_lazy_pool() {

View File

@@ -0,0 +1,183 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const QUOTA_COLUMNS: &str = r#"
SELECT
id AS provider_id,
billing_type,
monthly_quota_usd,
COALESCE(monthly_used_usd, 0) AS monthly_used_usd,
quota_reset_day,
quota_last_reset_at AS quota_last_reset_at_unix_secs,
quota_expires_at AS quota_expires_at_unix_secs,
is_active
FROM providers
"#;
#[derive(Debug, Clone)]
pub struct SqliteProviderQuotaRepository {
pool: SqlitePool,
}
impl SqliteProviderQuotaRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
}
#[async_trait]
impl ProviderQuotaReadRepository for SqliteProviderQuotaRepository {
async fn find_by_provider_id(
&self,
provider_id: &str,
) -> Result<Option<StoredProviderQuotaSnapshot>, DataLayerError> {
let row = sqlx::query(&format!("{QUOTA_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(provider_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_row).transpose()
}
async fn find_by_provider_ids(
&self,
provider_ids: &[String],
) -> Result<Vec<StoredProviderQuotaSnapshot>, DataLayerError> {
if provider_ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(QUOTA_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for provider_id in provider_ids {
separated.push_bind(provider_id);
}
}
builder.push(") ORDER BY id ASC");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_row).collect()
}
}
#[async_trait]
impl ProviderQuotaWriteRepository for SqliteProviderQuotaRepository {
async fn reset_due(&self, now_unix_secs: u64) -> Result<usize, DataLayerError> {
let now = i64::try_from(now_unix_secs).map_err(|_| {
DataLayerError::InvalidInput("provider quota reset timestamp overflow".to_string())
})?;
let rows_affected = sqlx::query(
r#"
UPDATE providers
SET monthly_used_usd = 0,
quota_last_reset_at = ?,
updated_at = ?
WHERE billing_type = 'monthly_quota'
AND is_active = 1
AND (
quota_last_reset_at IS NULL
OR (? - quota_last_reset_at) >= (quota_reset_day * 86400)
)
"#,
)
.bind(now)
.bind(now)
.bind(now)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or_default())
}
}
fn map_row(row: &SqliteRow) -> Result<StoredProviderQuotaSnapshot, DataLayerError> {
StoredProviderQuotaSnapshot::new(
row.try_get("provider_id").map_sql_err()?,
row.try_get("billing_type").map_sql_err()?,
row.try_get("monthly_quota_usd").map_sql_err()?,
row.try_get("monthly_used_usd").map_sql_err()?,
row.try_get("quota_reset_day").map_sql_err()?,
row.try_get("quota_last_reset_at_unix_secs").map_sql_err()?,
row.try_get("quota_expires_at_unix_secs").map_sql_err()?,
row.try_get("is_active").map_sql_err()?,
)
}
#[cfg(test)]
mod tests {
use super::SqliteProviderQuotaRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::quota::{ProviderQuotaReadRepository, ProviderQuotaWriteRepository};
#[tokio::test]
async fn sqlite_repository_reads_and_resets_provider_quotas() {
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_provider_quotas(&pool).await;
let repository = SqliteProviderQuotaRepository::new(pool);
let quota = repository
.find_by_provider_id("provider-1")
.await
.expect("quota should load")
.expect("quota should exist");
assert_eq!(quota.monthly_used_usd, 5.0);
let quotas = repository
.find_by_provider_ids(&["provider-2".to_string(), "provider-1".to_string()])
.await
.expect("quotas should load");
assert_eq!(
quotas
.iter()
.map(|quota| quota.provider_id.as_str())
.collect::<Vec<_>>(),
vec!["provider-1", "provider-2"]
);
let reset = repository
.reset_due(1_000 + 7 * 24 * 60 * 60)
.await
.expect("quota reset should run");
assert_eq!(reset, 1);
let quota = repository
.find_by_provider_id("provider-1")
.await
.expect("quota should reload")
.expect("quota should exist");
assert_eq!(quota.monthly_used_usd, 0.0);
assert_eq!(quota.quota_last_reset_at_unix_secs, Some(605_800));
}
async fn seed_provider_quotas(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
INSERT INTO providers (
id, name, provider_type, billing_type, monthly_quota_usd, monthly_used_usd,
quota_reset_day, quota_last_reset_at, is_active, created_at, updated_at
)
VALUES
('provider-1', 'Provider One', 'openai', 'monthly_quota', 20.0, 5.0, 7, 1000, 1, 1, 1),
('provider-2', 'Provider Two', 'openai', 'payg', NULL, 1.5, NULL, NULL, 1, 1, 1)
"#,
)
.execute(pool)
.await
.expect("providers should seed");
}
}

View File

@@ -1,9 +1,13 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::settlement::{
SettlementRepository, SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput,
};
pub use memory::InMemorySettlementRepository;
pub use sql::SqlxSettlementRepository;
pub use mysql::MysqlSettlementRepository;
pub use postgres::SqlxSettlementRepository;
pub use sqlite::SqliteSettlementRepository;

View File

@@ -0,0 +1,376 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use super::{SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const FIND_USAGE_FOR_SETTLEMENT_SQL: &str = r#"
SELECT
usage_record.request_id,
COALESCE(usage_settlement_snapshots.wallet_id, usage_record.wallet_id) AS wallet_id,
COALESCE(usage_settlement_snapshots.billing_status, usage_record.billing_status) AS billing_status,
COALESCE(
usage_settlement_snapshots.wallet_balance_before,
usage_record.wallet_balance_before
) AS wallet_balance_before,
COALESCE(
usage_settlement_snapshots.wallet_balance_after,
usage_record.wallet_balance_after
) AS wallet_balance_after,
COALESCE(
usage_settlement_snapshots.wallet_recharge_balance_before,
usage_record.wallet_recharge_balance_before
) AS wallet_recharge_balance_before,
COALESCE(
usage_settlement_snapshots.wallet_recharge_balance_after,
usage_record.wallet_recharge_balance_after
) AS wallet_recharge_balance_after,
COALESCE(
usage_settlement_snapshots.wallet_gift_balance_before,
usage_record.wallet_gift_balance_before
) AS wallet_gift_balance_before,
COALESCE(
usage_settlement_snapshots.wallet_gift_balance_after,
usage_record.wallet_gift_balance_after
) AS wallet_gift_balance_after,
usage_settlement_snapshots.provider_monthly_used_usd AS provider_monthly_used_usd,
usage_record.provider_id,
COALESCE(usage_settlement_snapshots.finalized_at, usage_record.finalized_at) AS finalized_at_unix_secs
FROM `usage` AS usage_record
LEFT JOIN usage_settlement_snapshots
ON usage_settlement_snapshots.request_id = usage_record.request_id
WHERE usage_record.request_id = ?
FOR UPDATE
"#;
const FINALIZE_USAGE_BILLING_SQL: &str = r#"
UPDATE `usage`
SET
billing_status = ?,
finalized_at = COALESCE(finalized_at, ?)
WHERE request_id = ?
"#;
const UPSERT_USAGE_SETTLEMENT_SNAPSHOT_SQL: &str = r#"
INSERT INTO usage_settlement_snapshots (
request_id,
billing_status,
wallet_id,
wallet_balance_before,
wallet_balance_after,
wallet_recharge_balance_before,
wallet_recharge_balance_after,
wallet_gift_balance_before,
wallet_gift_balance_after,
provider_monthly_used_usd,
finalized_at,
created_at,
updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
billing_status = VALUES(billing_status),
wallet_id = COALESCE(VALUES(wallet_id), wallet_id),
wallet_balance_before = COALESCE(VALUES(wallet_balance_before), wallet_balance_before),
wallet_balance_after = COALESCE(VALUES(wallet_balance_after), wallet_balance_after),
wallet_recharge_balance_before = COALESCE(
VALUES(wallet_recharge_balance_before),
wallet_recharge_balance_before
),
wallet_recharge_balance_after = COALESCE(
VALUES(wallet_recharge_balance_after),
wallet_recharge_balance_after
),
wallet_gift_balance_before = COALESCE(VALUES(wallet_gift_balance_before), wallet_gift_balance_before),
wallet_gift_balance_after = COALESCE(VALUES(wallet_gift_balance_after), wallet_gift_balance_after),
provider_monthly_used_usd = COALESCE(VALUES(provider_monthly_used_usd), provider_monthly_used_usd),
finalized_at = COALESCE(VALUES(finalized_at), finalized_at),
updated_at = VALUES(updated_at)
"#;
#[derive(Debug, Clone)]
pub struct MysqlSettlementRepository {
pool: MysqlPool,
}
impl MysqlSettlementRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
}
fn settlement_from_row(row: &MySqlRow) -> Result<StoredUsageSettlement, DataLayerError> {
Ok(StoredUsageSettlement {
request_id: row.try_get("request_id").map_sql_err()?,
wallet_id: row.try_get("wallet_id").map_sql_err()?,
billing_status: row.try_get("billing_status").map_sql_err()?,
wallet_balance_before: row.try_get("wallet_balance_before").map_sql_err()?,
wallet_balance_after: row.try_get("wallet_balance_after").map_sql_err()?,
wallet_recharge_balance_before: row
.try_get("wallet_recharge_balance_before")
.map_sql_err()?,
wallet_recharge_balance_after: row
.try_get("wallet_recharge_balance_after")
.map_sql_err()?,
wallet_gift_balance_before: row.try_get("wallet_gift_balance_before").map_sql_err()?,
wallet_gift_balance_after: row.try_get("wallet_gift_balance_after").map_sql_err()?,
provider_monthly_used_usd: row.try_get("provider_monthly_used_usd").map_sql_err()?,
finalized_at_unix_secs: row
.try_get::<Option<i64>, _>("finalized_at_unix_secs")
.map_sql_err()?
.map(|value| value as u64),
})
}
fn now_unix_secs() -> Result<i64, DataLayerError> {
i64::try_from(
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
)
.map_err(|_| DataLayerError::InvalidInput("timestamp overflow".to_string()))
}
#[async_trait]
impl SettlementWriteRepository for MysqlSettlementRepository {
async fn settle_usage(
&self,
input: UsageSettlementInput,
) -> Result<Option<StoredUsageSettlement>, DataLayerError> {
input.validate()?;
let finalized_at = i64::try_from(
input
.finalized_at_unix_secs
.unwrap_or(now_unix_secs()? as u64),
)
.map_err(|_| DataLayerError::InvalidInput("finalized_at overflow".to_string()))?;
let updated_at = now_unix_secs()?;
let mut tx = self.pool.begin().await.map_sql_err()?;
let row = sqlx::query(FIND_USAGE_FOR_SETTLEMENT_SQL)
.bind(&input.request_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(usage_row) = row else {
tx.commit().await.map_sql_err()?;
return Ok(None);
};
let current_billing_status: String = usage_row.try_get("billing_status").map_sql_err()?;
if current_billing_status == "settled" || current_billing_status == "void" {
let settlement = settlement_from_row(&usage_row)?;
tx.commit().await.map_sql_err()?;
return Ok(Some(settlement));
}
let final_billing_status = if input.status == "completed" {
"settled"
} else {
"void"
};
let mut settlement = StoredUsageSettlement {
request_id: input.request_id.clone(),
wallet_id: None,
billing_status: final_billing_status.to_string(),
wallet_balance_before: None,
wallet_balance_after: None,
wallet_recharge_balance_before: None,
wallet_recharge_balance_after: None,
wallet_gift_balance_before: None,
wallet_gift_balance_after: None,
provider_monthly_used_usd: None,
finalized_at_unix_secs: Some(finalized_at as u64),
};
if final_billing_status == "settled" {
let api_key_id = input
.api_key_id
.as_deref()
.filter(|value| !value.is_empty());
let api_key_is_standalone = if input.api_key_is_standalone {
true
} else if let Some(api_key_id) = api_key_id {
sqlx::query_scalar::<_, bool>(
r#"
SELECT is_standalone
FROM api_keys
WHERE id = ?
LIMIT 1
"#,
)
.bind(api_key_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
.unwrap_or(false)
} else {
false
};
let wallet_row = if let Some(api_key_id) = api_key_id {
sqlx::query(
r#"
SELECT id, balance, gift_balance, limit_mode
FROM wallets
WHERE api_key_id = ?
LIMIT 1
FOR UPDATE
"#,
)
.bind(api_key_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
} else {
None
};
let wallet_row = if wallet_row.is_some() {
wallet_row
} else if !api_key_is_standalone {
if let Some(user_id) = input.user_id.as_deref().filter(|value| !value.is_empty()) {
sqlx::query(
r#"
SELECT id, balance, gift_balance, limit_mode
FROM wallets
WHERE user_id = ?
LIMIT 1
FOR UPDATE
"#,
)
.bind(user_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
} else {
None
}
} else {
None
};
if let Some(wallet_row) = wallet_row {
let wallet_id: String = wallet_row.try_get("id").map_sql_err()?;
let before_recharge: f64 = wallet_row.try_get("balance").map_sql_err()?;
let before_gift: f64 = wallet_row.try_get("gift_balance").map_sql_err()?;
let limit_mode: String = wallet_row.try_get("limit_mode").map_sql_err()?;
let before_total = before_recharge + before_gift;
let mut after_recharge = before_recharge;
let mut after_gift = before_gift;
if !limit_mode.eq_ignore_ascii_case("unlimited") {
let gift_deduction = before_gift.max(0.0).min(input.total_cost_usd);
let recharge_deduction = input.total_cost_usd - gift_deduction;
after_gift = before_gift - gift_deduction;
after_recharge = before_recharge - recharge_deduction;
}
sqlx::query(
r#"
UPDATE wallets
SET
balance = ?,
gift_balance = ?,
total_consumed = COALESCE(total_consumed, 0) + ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(after_recharge)
.bind(after_gift)
.bind(input.total_cost_usd)
.bind(updated_at)
.bind(&wallet_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
settlement.wallet_id = Some(wallet_id);
settlement.wallet_balance_before = Some(before_total);
settlement.wallet_balance_after = Some(after_recharge + after_gift);
settlement.wallet_recharge_balance_before = Some(before_recharge);
settlement.wallet_recharge_balance_after = Some(after_recharge);
settlement.wallet_gift_balance_before = Some(before_gift);
settlement.wallet_gift_balance_after = Some(after_gift);
}
if let Some(provider_id) = input
.provider_id
.as_deref()
.filter(|value| !value.is_empty())
{
sqlx::query(
r#"
UPDATE providers
SET
monthly_used_usd = COALESCE(monthly_used_usd, 0) + ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(input.actual_total_cost_usd)
.bind(updated_at)
.bind(provider_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
settlement.provider_monthly_used_usd = sqlx::query_scalar::<_, Option<f64>>(
"SELECT monthly_used_usd FROM providers WHERE id = ? LIMIT 1",
)
.bind(provider_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
.flatten();
}
}
sqlx::query(UPSERT_USAGE_SETTLEMENT_SNAPSHOT_SQL)
.bind(&settlement.request_id)
.bind(&settlement.billing_status)
.bind(settlement.wallet_id.as_deref())
.bind(settlement.wallet_balance_before)
.bind(settlement.wallet_balance_after)
.bind(settlement.wallet_recharge_balance_before)
.bind(settlement.wallet_recharge_balance_after)
.bind(settlement.wallet_gift_balance_before)
.bind(settlement.wallet_gift_balance_after)
.bind(settlement.provider_monthly_used_usd)
.bind(settlement.finalized_at_unix_secs.map(|value| value as i64))
.bind(updated_at)
.bind(updated_at)
.execute(&mut *tx)
.await
.map_sql_err()?;
sqlx::query(FINALIZE_USAGE_BILLING_SQL)
.bind(final_billing_status)
.bind(finalized_at)
.bind(&input.request_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(Some(settlement))
}
}
#[cfg(test)]
mod tests {
use super::MysqlSettlementRepository;
#[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 = MysqlSettlementRepository::new(pool);
}
}

View File

@@ -2,8 +2,8 @@ use async_trait::async_trait;
use sqlx::{PgPool, Row};
use super::{SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput};
use crate::driver::postgres::PostgresTransactionRunner;
use crate::error::SqlxResultExt;
use crate::postgres::PostgresTransactionRunner;
use crate::DataLayerError;
const FIND_USAGE_FOR_SETTLEMENT_SQL: &str = r#"
@@ -438,13 +438,13 @@ mod tests {
#[test]
fn settlement_sql_no_longer_dual_writes_wallet_snapshots_to_usage_rows() {
let source = include_str!("sql.rs");
let source = include_str!("postgres.rs");
assert!(!source.contains("UPDATE \"usage\"\nSET\n wallet_id = $2"));
}
#[test]
fn settlement_sql_blocks_standalone_key_owner_wallet_fallback() {
let source = include_str!("sql.rs");
let source = include_str!("postgres.rs");
assert!(source.contains("SELECT is_standalone"));
assert!(source.contains("} else if !api_key_is_standalone {"));
}

View File

@@ -0,0 +1,521 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row};
use super::{SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const FIND_USAGE_FOR_SETTLEMENT_SQL: &str = r#"
SELECT
usage_record.request_id,
COALESCE(usage_settlement_snapshots.wallet_id, usage_record.wallet_id) AS wallet_id,
COALESCE(usage_settlement_snapshots.billing_status, usage_record.billing_status) AS billing_status,
COALESCE(
usage_settlement_snapshots.wallet_balance_before,
usage_record.wallet_balance_before
) AS wallet_balance_before,
COALESCE(
usage_settlement_snapshots.wallet_balance_after,
usage_record.wallet_balance_after
) AS wallet_balance_after,
COALESCE(
usage_settlement_snapshots.wallet_recharge_balance_before,
usage_record.wallet_recharge_balance_before
) AS wallet_recharge_balance_before,
COALESCE(
usage_settlement_snapshots.wallet_recharge_balance_after,
usage_record.wallet_recharge_balance_after
) AS wallet_recharge_balance_after,
COALESCE(
usage_settlement_snapshots.wallet_gift_balance_before,
usage_record.wallet_gift_balance_before
) AS wallet_gift_balance_before,
COALESCE(
usage_settlement_snapshots.wallet_gift_balance_after,
usage_record.wallet_gift_balance_after
) AS wallet_gift_balance_after,
usage_settlement_snapshots.provider_monthly_used_usd AS provider_monthly_used_usd,
usage_record.provider_id,
COALESCE(usage_settlement_snapshots.finalized_at, usage_record.finalized_at) AS finalized_at_unix_secs
FROM "usage" AS usage_record
LEFT JOIN usage_settlement_snapshots
ON usage_settlement_snapshots.request_id = usage_record.request_id
WHERE usage_record.request_id = ?
"#;
const FINALIZE_USAGE_BILLING_SQL: &str = r#"
UPDATE "usage"
SET
billing_status = ?,
finalized_at = COALESCE(finalized_at, ?)
WHERE request_id = ?
"#;
const UPSERT_USAGE_SETTLEMENT_SNAPSHOT_SQL: &str = r#"
INSERT INTO usage_settlement_snapshots (
request_id,
billing_status,
wallet_id,
wallet_balance_before,
wallet_balance_after,
wallet_recharge_balance_before,
wallet_recharge_balance_after,
wallet_gift_balance_before,
wallet_gift_balance_after,
provider_monthly_used_usd,
finalized_at,
created_at,
updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT (request_id)
DO UPDATE SET
billing_status = excluded.billing_status,
wallet_id = COALESCE(excluded.wallet_id, usage_settlement_snapshots.wallet_id),
wallet_balance_before = COALESCE(
excluded.wallet_balance_before,
usage_settlement_snapshots.wallet_balance_before
),
wallet_balance_after = COALESCE(
excluded.wallet_balance_after,
usage_settlement_snapshots.wallet_balance_after
),
wallet_recharge_balance_before = COALESCE(
excluded.wallet_recharge_balance_before,
usage_settlement_snapshots.wallet_recharge_balance_before
),
wallet_recharge_balance_after = COALESCE(
excluded.wallet_recharge_balance_after,
usage_settlement_snapshots.wallet_recharge_balance_after
),
wallet_gift_balance_before = COALESCE(
excluded.wallet_gift_balance_before,
usage_settlement_snapshots.wallet_gift_balance_before
),
wallet_gift_balance_after = COALESCE(
excluded.wallet_gift_balance_after,
usage_settlement_snapshots.wallet_gift_balance_after
),
provider_monthly_used_usd = COALESCE(
excluded.provider_monthly_used_usd,
usage_settlement_snapshots.provider_monthly_used_usd
),
finalized_at = COALESCE(excluded.finalized_at, usage_settlement_snapshots.finalized_at),
updated_at = excluded.updated_at
"#;
#[derive(Debug, Clone)]
pub struct SqliteSettlementRepository {
pool: SqlitePool,
}
impl SqliteSettlementRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
}
fn settlement_from_row(row: &SqliteRow) -> Result<StoredUsageSettlement, DataLayerError> {
Ok(StoredUsageSettlement {
request_id: row.try_get("request_id").map_sql_err()?,
wallet_id: row.try_get("wallet_id").map_sql_err()?,
billing_status: row.try_get("billing_status").map_sql_err()?,
wallet_balance_before: row.try_get("wallet_balance_before").map_sql_err()?,
wallet_balance_after: row.try_get("wallet_balance_after").map_sql_err()?,
wallet_recharge_balance_before: row
.try_get("wallet_recharge_balance_before")
.map_sql_err()?,
wallet_recharge_balance_after: row
.try_get("wallet_recharge_balance_after")
.map_sql_err()?,
wallet_gift_balance_before: row.try_get("wallet_gift_balance_before").map_sql_err()?,
wallet_gift_balance_after: row.try_get("wallet_gift_balance_after").map_sql_err()?,
provider_monthly_used_usd: row.try_get("provider_monthly_used_usd").map_sql_err()?,
finalized_at_unix_secs: row
.try_get::<Option<i64>, _>("finalized_at_unix_secs")
.map_sql_err()?
.map(|value| value as u64),
})
}
fn now_unix_secs() -> Result<i64, DataLayerError> {
i64::try_from(
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs(),
)
.map_err(|_| DataLayerError::InvalidInput("timestamp overflow".to_string()))
}
#[async_trait]
impl SettlementWriteRepository for SqliteSettlementRepository {
async fn settle_usage(
&self,
input: UsageSettlementInput,
) -> Result<Option<StoredUsageSettlement>, DataLayerError> {
input.validate()?;
let finalized_at = i64::try_from(
input
.finalized_at_unix_secs
.unwrap_or(now_unix_secs()? as u64),
)
.map_err(|_| DataLayerError::InvalidInput("finalized_at overflow".to_string()))?;
let updated_at = now_unix_secs()?;
let mut tx = self.pool.begin().await.map_sql_err()?;
let row = sqlx::query(FIND_USAGE_FOR_SETTLEMENT_SQL)
.bind(&input.request_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?;
let Some(usage_row) = row else {
tx.commit().await.map_sql_err()?;
return Ok(None);
};
let current_billing_status: String = usage_row.try_get("billing_status").map_sql_err()?;
if current_billing_status == "settled" || current_billing_status == "void" {
let settlement = settlement_from_row(&usage_row)?;
tx.commit().await.map_sql_err()?;
return Ok(Some(settlement));
}
let final_billing_status = if input.status == "completed" {
"settled"
} else {
"void"
};
let mut settlement = StoredUsageSettlement {
request_id: input.request_id.clone(),
wallet_id: None,
billing_status: final_billing_status.to_string(),
wallet_balance_before: None,
wallet_balance_after: None,
wallet_recharge_balance_before: None,
wallet_recharge_balance_after: None,
wallet_gift_balance_before: None,
wallet_gift_balance_after: None,
provider_monthly_used_usd: None,
finalized_at_unix_secs: Some(finalized_at as u64),
};
if final_billing_status == "settled" {
let api_key_id = input
.api_key_id
.as_deref()
.filter(|value| !value.is_empty());
let api_key_is_standalone = if input.api_key_is_standalone {
true
} else if let Some(api_key_id) = api_key_id {
sqlx::query_scalar::<_, bool>(
r#"
SELECT is_standalone
FROM api_keys
WHERE id = ?
LIMIT 1
"#,
)
.bind(api_key_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
.unwrap_or(false)
} else {
false
};
let wallet_row = if let Some(api_key_id) = api_key_id {
sqlx::query(
r#"
SELECT id, balance, gift_balance, limit_mode
FROM wallets
WHERE api_key_id = ?
LIMIT 1
"#,
)
.bind(api_key_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
} else {
None
};
let wallet_row = if wallet_row.is_some() {
wallet_row
} else if !api_key_is_standalone {
if let Some(user_id) = input.user_id.as_deref().filter(|value| !value.is_empty()) {
sqlx::query(
r#"
SELECT id, balance, gift_balance, limit_mode
FROM wallets
WHERE user_id = ?
LIMIT 1
"#,
)
.bind(user_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
} else {
None
}
} else {
None
};
if let Some(wallet_row) = wallet_row {
let wallet_id: String = wallet_row.try_get("id").map_sql_err()?;
let before_recharge: f64 = wallet_row.try_get("balance").map_sql_err()?;
let before_gift: f64 = wallet_row.try_get("gift_balance").map_sql_err()?;
let limit_mode: String = wallet_row.try_get("limit_mode").map_sql_err()?;
let before_total = before_recharge + before_gift;
let mut after_recharge = before_recharge;
let mut after_gift = before_gift;
if !limit_mode.eq_ignore_ascii_case("unlimited") {
let gift_deduction = before_gift.max(0.0).min(input.total_cost_usd);
let recharge_deduction = input.total_cost_usd - gift_deduction;
after_gift = before_gift - gift_deduction;
after_recharge = before_recharge - recharge_deduction;
}
sqlx::query(
r#"
UPDATE wallets
SET
balance = ?,
gift_balance = ?,
total_consumed = COALESCE(total_consumed, 0) + ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(after_recharge)
.bind(after_gift)
.bind(input.total_cost_usd)
.bind(updated_at)
.bind(&wallet_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
settlement.wallet_id = Some(wallet_id);
settlement.wallet_balance_before = Some(before_total);
settlement.wallet_balance_after = Some(after_recharge + after_gift);
settlement.wallet_recharge_balance_before = Some(before_recharge);
settlement.wallet_recharge_balance_after = Some(after_recharge);
settlement.wallet_gift_balance_before = Some(before_gift);
settlement.wallet_gift_balance_after = Some(after_gift);
}
if let Some(provider_id) = input
.provider_id
.as_deref()
.filter(|value| !value.is_empty())
{
sqlx::query(
r#"
UPDATE providers
SET
monthly_used_usd = COALESCE(monthly_used_usd, 0) + ?,
updated_at = ?
WHERE id = ?
"#,
)
.bind(input.actual_total_cost_usd)
.bind(updated_at)
.bind(provider_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
settlement.provider_monthly_used_usd = sqlx::query_scalar::<_, Option<f64>>(
"SELECT monthly_used_usd FROM providers WHERE id = ? LIMIT 1",
)
.bind(provider_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
.flatten();
}
}
sqlx::query(UPSERT_USAGE_SETTLEMENT_SNAPSHOT_SQL)
.bind(&settlement.request_id)
.bind(&settlement.billing_status)
.bind(settlement.wallet_id.as_deref())
.bind(settlement.wallet_balance_before)
.bind(settlement.wallet_balance_after)
.bind(settlement.wallet_recharge_balance_before)
.bind(settlement.wallet_recharge_balance_after)
.bind(settlement.wallet_gift_balance_before)
.bind(settlement.wallet_gift_balance_after)
.bind(settlement.provider_monthly_used_usd)
.bind(settlement.finalized_at_unix_secs.map(|value| value as i64))
.bind(updated_at)
.bind(updated_at)
.execute(&mut *tx)
.await
.map_sql_err()?;
sqlx::query(FINALIZE_USAGE_BILLING_SQL)
.bind(final_billing_status)
.bind(finalized_at)
.bind(&input.request_id)
.execute(&mut *tx)
.await
.map_sql_err()?;
tx.commit().await.map_sql_err()?;
Ok(Some(settlement))
}
}
#[cfg(test)]
mod tests {
use super::SqliteSettlementRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::settlement::{SettlementWriteRepository, UsageSettlementInput};
use sqlx::Row;
#[tokio::test]
async fn sqlite_repository_settles_usage_once() {
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_settlement_rows(&pool).await;
let repository = SqliteSettlementRepository::new(pool.clone());
let settlement = repository
.settle_usage(UsageSettlementInput {
request_id: "request-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: None,
api_key_is_standalone: false,
provider_id: Some("provider-1".to_string()),
status: "completed".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 3.0,
actual_total_cost_usd: 2.0,
finalized_at_unix_secs: Some(1_234),
})
.await
.expect("settlement should run")
.expect("usage should exist");
assert_eq!(settlement.billing_status, "settled");
assert_eq!(settlement.wallet_id.as_deref(), Some("wallet-1"));
assert_eq!(settlement.wallet_balance_before, Some(12.0));
assert_eq!(settlement.wallet_balance_after, Some(9.0));
assert_eq!(settlement.wallet_recharge_balance_after, Some(9.0));
assert_eq!(settlement.wallet_gift_balance_after, Some(0.0));
assert_eq!(settlement.provider_monthly_used_usd, Some(7.0));
let wallet = sqlx::query(
"SELECT balance, gift_balance, total_consumed FROM wallets WHERE id = 'wallet-1'",
)
.fetch_one(&pool)
.await
.expect("wallet should load");
assert_eq!(wallet.try_get::<f64, _>("balance").unwrap(), 9.0);
assert_eq!(wallet.try_get::<f64, _>("gift_balance").unwrap(), 0.0);
assert_eq!(wallet.try_get::<f64, _>("total_consumed").unwrap(), 3.0);
let second = repository
.settle_usage(UsageSettlementInput {
request_id: "request-1".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: None,
api_key_is_standalone: false,
provider_id: Some("provider-1".to_string()),
status: "completed".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 3.0,
actual_total_cost_usd: 2.0,
finalized_at_unix_secs: Some(9_999),
})
.await
.expect("second settlement should run")
.expect("usage should exist");
assert_eq!(second.finalized_at_unix_secs, Some(1_234));
let provider_used: f64 =
sqlx::query_scalar("SELECT monthly_used_usd FROM providers WHERE id = 'provider-1'")
.fetch_one(&pool)
.await
.expect("provider should load");
assert_eq!(provider_used, 7.0);
}
#[tokio::test]
async fn sqlite_repository_voids_failed_usage_without_wallet_mutation() {
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_settlement_rows(&pool).await;
let repository = SqliteSettlementRepository::new(pool.clone());
let settlement = repository
.settle_usage(UsageSettlementInput {
request_id: "request-2".to_string(),
user_id: Some("user-1".to_string()),
api_key_id: None,
api_key_is_standalone: false,
provider_id: Some("provider-1".to_string()),
status: "failed".to_string(),
billing_status: "pending".to_string(),
total_cost_usd: 3.0,
actual_total_cost_usd: 2.0,
finalized_at_unix_secs: Some(1_235),
})
.await
.expect("settlement should run")
.expect("usage should exist");
assert_eq!(settlement.billing_status, "void");
assert_eq!(settlement.wallet_id, None);
let wallet_total: f64 =
sqlx::query_scalar("SELECT balance + gift_balance FROM wallets WHERE id = 'wallet-1'")
.fetch_one(&pool)
.await
.expect("wallet should load");
assert_eq!(wallet_total, 12.0);
}
async fn seed_settlement_rows(pool: &sqlx::SqlitePool) {
sqlx::query(
r#"
INSERT INTO providers (
id, name, provider_type, monthly_used_usd, created_at, updated_at
)
VALUES ('provider-1', 'Provider One', 'openai', 5.0, 1, 1);
INSERT INTO wallets (
id, user_id, balance, gift_balance, limit_mode, created_at, updated_at
)
VALUES ('wallet-1', 'user-1', 10.0, 2.0, 'finite', 1, 1);
INSERT INTO "usage" (
request_id, user_id, provider_id, status, billing_status, total_cost_usd, actual_total_cost_usd
)
VALUES
('request-1', 'user-1', 'provider-1', 'completed', 'pending', 3.0, 2.0),
('request-2', 'user-1', 'provider-1', 'failed', 'pending', 3.0, 2.0);
"#,
)
.execute(pool)
.await
.expect("settlement rows should seed");
}
}

View File

@@ -1,11 +1,360 @@
macro_rules! impl_materialized_usage_read_repository {
($repository:ty) => {
#[async_trait::async_trait]
impl $crate::repository::usage::UsageReadRepository for $repository {
async fn find_by_id(
&self,
id: &str,
) -> Result<
Option<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::find_by_id(&repository, id).await
}
async fn list_by_ids(
&self,
ids: &[String],
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_by_ids(&repository, ids).await
}
async fn find_by_request_id(
&self,
request_id: &str,
) -> Result<
Option<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::find_by_request_id(&repository, request_id).await
}
async fn resolve_body_ref(
&self,
body_ref: &str,
) -> Result<Option<serde_json::Value>, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::resolve_body_ref(&repository, body_ref).await
}
async fn list_usage_audits(
&self,
query: &$crate::repository::usage::UsageAuditListQuery,
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_usage_audits(&repository, query).await
}
async fn count_usage_audits(
&self,
query: &$crate::repository::usage::UsageAuditListQuery,
) -> Result<u64, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::count_usage_audits(&repository, query).await
}
async fn list_usage_audits_by_keyword_search(
&self,
query: &$crate::repository::usage::UsageAuditKeywordSearchQuery,
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_usage_audits_by_keyword_search(&repository, query).await
}
async fn count_usage_audits_by_keyword_search(
&self,
query: &$crate::repository::usage::UsageAuditKeywordSearchQuery,
) -> Result<u64, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::count_usage_audits_by_keyword_search(&repository, query).await
}
async fn aggregate_usage_audits(
&self,
query: &$crate::repository::usage::UsageAuditAggregationQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageAuditAggregation>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::aggregate_usage_audits(&repository, query).await
}
async fn summarize_usage_audits(
&self,
query: &$crate::repository::usage::UsageAuditSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageAuditSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_audits(&repository, query).await
}
async fn summarize_usage_totals_by_user_ids(
&self,
user_ids: &[String],
) -> Result<Vec<$crate::repository::usage::StoredUsageUserTotals>, $crate::DataLayerError>
{
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_totals_by_user_ids(&repository, user_ids).await
}
async fn summarize_usage_cache_hit_summary(
&self,
query: &$crate::repository::usage::UsageCacheHitSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageCacheHitSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_cache_hit_summary(&repository, query).await
}
async fn summarize_usage_settled_cost(
&self,
query: &$crate::repository::usage::UsageSettledCostSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageSettledCostSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_settled_cost(&repository, query).await
}
async fn summarize_usage_cache_affinity_hit_summary(
&self,
query: &$crate::repository::usage::UsageCacheAffinityHitSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageCacheAffinityHitSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_cache_affinity_hit_summary(&repository, query).await
}
async fn list_usage_cache_affinity_intervals(
&self,
query: &$crate::repository::usage::UsageCacheAffinityIntervalQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageCacheAffinityIntervalRow>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_usage_cache_affinity_intervals(&repository, query).await
}
async fn summarize_dashboard_usage(
&self,
query: &$crate::repository::usage::UsageDashboardSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageDashboardSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_dashboard_usage(&repository, query).await
}
async fn list_dashboard_daily_breakdown(
&self,
query: &$crate::repository::usage::UsageDashboardDailyBreakdownQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageDashboardDailyBreakdownRow>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_dashboard_daily_breakdown(&repository, query).await
}
async fn summarize_dashboard_provider_counts(
&self,
query: &$crate::repository::usage::UsageDashboardProviderCountsQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageDashboardProviderCount>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_dashboard_provider_counts(&repository, query).await
}
async fn summarize_usage_breakdown(
&self,
query: &$crate::repository::usage::UsageBreakdownSummaryQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageBreakdownSummaryRow>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_breakdown(&repository, query).await
}
async fn count_monitoring_usage_errors(
&self,
query: &$crate::repository::usage::UsageMonitoringErrorCountQuery,
) -> Result<u64, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::count_monitoring_usage_errors(&repository, query).await
}
async fn list_monitoring_usage_errors(
&self,
query: &$crate::repository::usage::UsageMonitoringErrorListQuery,
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_monitoring_usage_errors(&repository, query).await
}
async fn summarize_usage_error_distribution(
&self,
query: &$crate::repository::usage::UsageErrorDistributionQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageErrorDistributionRow>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_error_distribution(&repository, query).await
}
async fn summarize_usage_performance_percentiles(
&self,
query: &$crate::repository::usage::UsagePerformancePercentilesQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsagePerformancePercentilesRow>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_performance_percentiles(&repository, query).await
}
async fn summarize_usage_provider_performance(
&self,
query: &$crate::repository::usage::UsageProviderPerformanceQuery,
) -> Result<
$crate::repository::usage::StoredUsageProviderPerformance,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_provider_performance(&repository, query).await
}
async fn summarize_usage_cost_savings(
&self,
query: &$crate::repository::usage::UsageCostSavingsSummaryQuery,
) -> Result<
$crate::repository::usage::StoredUsageCostSavingsSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_cost_savings(&repository, query).await
}
async fn summarize_usage_time_series(
&self,
query: &$crate::repository::usage::UsageTimeSeriesQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageTimeSeriesBucket>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_time_series(&repository, query).await
}
async fn summarize_usage_leaderboard(
&self,
query: &$crate::repository::usage::UsageLeaderboardQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageLeaderboardSummary>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_leaderboard(&repository, query).await
}
async fn list_recent_usage_audits(
&self,
user_id: Option<&str>,
limit: usize,
) -> Result<
Vec<$crate::repository::usage::StoredRequestUsageAudit>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::list_recent_usage_audits(&repository, user_id, limit).await
}
async fn summarize_total_tokens_by_api_key_ids(
&self,
api_key_ids: &[String],
) -> Result<std::collections::BTreeMap<String, u64>, $crate::DataLayerError> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_total_tokens_by_api_key_ids(&repository, api_key_ids).await
}
async fn summarize_usage_by_provider_api_key_ids(
&self,
provider_api_key_ids: &[String],
) -> Result<
std::collections::BTreeMap<
String,
$crate::repository::usage::StoredProviderApiKeyUsageSummary,
>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_by_provider_api_key_ids(&repository, provider_api_key_ids).await
}
async fn summarize_provider_usage_since(
&self,
provider_id: &str,
since_unix_secs: u64,
) -> Result<
$crate::repository::usage::StoredProviderUsageSummary,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_provider_usage_since(&repository, provider_id, since_unix_secs).await
}
async fn summarize_usage_daily_heatmap(
&self,
query: &$crate::repository::usage::UsageDailyHeatmapQuery,
) -> Result<
Vec<$crate::repository::usage::StoredUsageDailySummary>,
$crate::DataLayerError,
> {
let repository = self.materialize_read_model().await?;
<$crate::repository::usage::InMemoryUsageReadRepository as $crate::repository::usage::UsageReadRepository>::summarize_usage_daily_heatmap(&repository, query).await
}
}
};
}
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::usage::{
StoredProviderApiKeyUsageSummary, StoredProviderUsageSummary, StoredProviderUsageWindow,
StoredRequestUsageAudit, StoredUsageAuditAggregation, StoredUsageAuditSummary,
StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
PendingUsageCleanupSummary, StoredProviderApiKeyUsageSummary, StoredProviderUsageSummary,
StoredProviderUsageWindow, StoredRequestUsageAudit, StoredUsageAuditAggregation,
StoredUsageAuditSummary, StoredUsageBreakdownSummaryRow, StoredUsageCacheAffinityHitSummary,
StoredUsageCacheAffinityIntervalRow, StoredUsageCacheHitSummary, StoredUsageCostSavingsSummary,
StoredUsageDailySummary, StoredUsageDashboardDailyBreakdownRow,
StoredUsageDashboardProviderCount, StoredUsageDashboardSummary,
@@ -17,16 +366,22 @@ pub(crate) use aether_data_contracts::repository::usage::{
UsageAuditAggregationGroupBy, UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery,
UsageAuditListQuery, UsageAuditSummaryQuery, UsageBreakdownGroupBy, UsageBreakdownSummaryQuery,
UsageCacheAffinityHitSummaryQuery, UsageCacheAffinityIntervalGroupBy,
UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery, UsageCostSavingsSummaryQuery,
UsageDailyHeatmapQuery, UsageDashboardDailyBreakdownQuery, UsageDashboardProviderCountsQuery,
UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery, UsageCleanupSummary,
UsageCleanupWindow, UsageCostSavingsSummaryQuery, UsageDailyHeatmapQuery,
UsageDashboardDailyBreakdownQuery, UsageDashboardProviderCountsQuery,
UsageDashboardSummaryQuery, UsageErrorDistributionQuery, UsageLeaderboardGroupBy,
UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, UsageMonitoringErrorListQuery,
UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageReadRepository,
UsageRepository, UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity,
UsageTimeSeriesQuery, UsageWriteRepository,
};
pub mod cleanup {
pub use super::postgres::cleanup::*;
}
pub use memory::InMemoryUsageReadRepository;
pub use sql::SqlxUsageReadRepository;
pub use mysql::{MysqlUsageReadRepository, MysqlUsageWriteRepository};
pub use postgres::SqlxUsageReadRepository;
pub use sqlite::{SqliteUsageReadRepository, SqliteUsageWriteRepository};
#[derive(Debug, Clone, PartialEq, Default)]
pub(crate) struct ApiKeyUsageContribution {

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -11,12 +11,12 @@ use aether_data_contracts::repository::usage::{
UsageAuditAggregationQuery, UsageAuditKeywordSearchQuery, UsageAuditSummaryQuery,
UsageBodyCaptureState, UsageBodyField, UsageBreakdownGroupBy, UsageBreakdownSummaryQuery,
UsageCacheAffinityHitSummaryQuery, UsageCacheAffinityIntervalGroupBy,
UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery, UsageCostSavingsSummaryQuery,
UsageDashboardDailyBreakdownQuery, UsageDashboardProviderCountsQuery,
UsageDashboardSummaryQuery, UsageErrorDistributionQuery, UsageLeaderboardGroupBy,
UsageLeaderboardQuery, UsageMonitoringErrorCountQuery, UsageMonitoringErrorListQuery,
UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery, UsageSettledCostSummaryQuery,
UsageTimeSeriesGranularity, UsageTimeSeriesQuery,
UsageCacheAffinityIntervalQuery, UsageCacheHitSummaryQuery, UsageCleanupSummary,
UsageCleanupWindow, UsageCostSavingsSummaryQuery, UsageDashboardDailyBreakdownQuery,
UsageDashboardProviderCountsQuery, UsageDashboardSummaryQuery, UsageErrorDistributionQuery,
UsageLeaderboardGroupBy, UsageLeaderboardQuery, UsageMonitoringErrorCountQuery,
UsageMonitoringErrorListQuery, UsagePerformancePercentilesQuery, UsageProviderPerformanceQuery,
UsageSettledCostSummaryQuery, UsageTimeSeriesGranularity, UsageTimeSeriesQuery,
};
use async_trait::async_trait;
use chrono::{DateTime, Utc};
@@ -38,16 +38,19 @@ use super::{
api_key_usage_contribution, incoming_usage_can_recover_terminal_failure,
model_usage_contribution, provider_api_key_usage_contribution,
strip_deprecated_usage_display_fields, ApiKeyUsageDelta, ModelUsageDelta,
ProviderApiKeyUsageDelta, StoredProviderApiKeyUsageSummary, StoredProviderUsageSummary,
StoredRequestUsageAudit, StoredUsageDailySummary, UpsertUsageRecord, UsageAuditListQuery,
UsageDailyHeatmapQuery, UsageReadRepository, UsageWriteRepository,
PendingUsageCleanupSummary, ProviderApiKeyUsageDelta, StoredProviderApiKeyUsageSummary,
StoredProviderUsageSummary, StoredRequestUsageAudit, StoredUsageDailySummary,
UpsertUsageRecord, UsageAuditListQuery, UsageDailyHeatmapQuery, UsageReadRepository,
UsageWriteRepository,
};
use crate::postgres::PostgresTransactionRunner;
use crate::driver::postgres::PostgresTransactionRunner;
use crate::{
error::{postgres_error, SqlxResultExt},
DataLayerError,
};
pub mod cleanup;
// Legacy inline body columns on public.usage are deprecated. Keep the threshold at zero so
// newly captured bodies always spill to usage_body_blobs and resolve through usage_http_audits.
const MAX_INLINE_USAGE_BODY_BYTES: usize = 0;
@@ -557,67 +560,6 @@ fn decode_usage_breakdown_summary_row(
})
}
fn decode_usage_audit_aggregation_row(
row: &PgRow,
) -> Result<StoredUsageAuditAggregation, DataLayerError> {
Ok(StoredUsageAuditAggregation {
group_key: row.try_get::<String, _>("group_key").map_postgres_err()?,
display_name: row
.try_get::<Option<String>, _>("display_name")
.map_postgres_err()?,
secondary_name: row
.try_get::<Option<String>, _>("secondary_name")
.map_postgres_err()?,
request_count: row
.try_get::<i64, _>("request_count")
.map_postgres_err()?
.max(0) as u64,
total_tokens: row
.try_get::<i64, _>("total_tokens")
.map_postgres_err()?
.max(0) as u64,
output_tokens: row
.try_get::<i64, _>("output_tokens")
.map_postgres_err()?
.max(0) as u64,
effective_input_tokens: row
.try_get::<i64, _>("effective_input_tokens")
.map_postgres_err()?
.max(0) as u64,
total_input_context: row
.try_get::<i64, _>("total_input_context")
.map_postgres_err()?
.max(0) as u64,
cache_creation_tokens: row
.try_get::<i64, _>("cache_creation_tokens")
.map_postgres_err()?
.max(0) as u64,
cache_creation_ephemeral_5m_tokens: row
.try_get::<i64, _>("cache_creation_ephemeral_5m_tokens")
.map_postgres_err()?
.max(0) as u64,
cache_creation_ephemeral_1h_tokens: row
.try_get::<i64, _>("cache_creation_ephemeral_1h_tokens")
.map_postgres_err()?
.max(0) as u64,
cache_read_tokens: row
.try_get::<i64, _>("cache_read_tokens")
.map_postgres_err()?
.max(0) as u64,
total_cost_usd: row.try_get::<f64, _>("total_cost_usd").map_postgres_err()?,
actual_total_cost_usd: row
.try_get::<f64, _>("actual_total_cost_usd")
.map_postgres_err()?,
avg_response_time_ms: row
.try_get::<Option<f64>, _>("avg_response_time_ms")
.map_postgres_err()?,
success_count: row
.try_get::<Option<i64>, _>("success_count")
.map_postgres_err()?
.map(|value| value.max(0) as u64),
})
}
fn absorb_usage_audit_summary(target: &mut StoredUsageAuditSummary, row: StoredUsageAuditSummary) {
target.total_requests = target.total_requests.saturating_add(row.total_requests);
target.input_tokens = target.input_tokens.saturating_add(row.input_tokens);
@@ -883,6 +825,67 @@ fn decode_usage_audit_summary_row(row: &PgRow) -> Result<StoredUsageAuditSummary
})
}
fn decode_usage_audit_aggregation_row(
row: &PgRow,
) -> Result<StoredUsageAuditAggregation, DataLayerError> {
Ok(StoredUsageAuditAggregation {
group_key: row.try_get::<String, _>("group_key").map_postgres_err()?,
display_name: row
.try_get::<Option<String>, _>("display_name")
.map_postgres_err()?,
secondary_name: row
.try_get::<Option<String>, _>("secondary_name")
.map_postgres_err()?,
request_count: row
.try_get::<i64, _>("request_count")
.map_postgres_err()?
.max(0) as u64,
total_tokens: row
.try_get::<i64, _>("total_tokens")
.map_postgres_err()?
.max(0) as u64,
output_tokens: row
.try_get::<i64, _>("output_tokens")
.map_postgres_err()?
.max(0) as u64,
effective_input_tokens: row
.try_get::<i64, _>("effective_input_tokens")
.map_postgres_err()?
.max(0) as u64,
total_input_context: row
.try_get::<i64, _>("total_input_context")
.map_postgres_err()?
.max(0) as u64,
cache_creation_tokens: row
.try_get::<i64, _>("cache_creation_tokens")
.map_postgres_err()?
.max(0) as u64,
cache_creation_ephemeral_5m_tokens: row
.try_get::<i64, _>("cache_creation_ephemeral_5m_tokens")
.map_postgres_err()?
.max(0) as u64,
cache_creation_ephemeral_1h_tokens: row
.try_get::<i64, _>("cache_creation_ephemeral_1h_tokens")
.map_postgres_err()?
.max(0) as u64,
cache_read_tokens: row
.try_get::<i64, _>("cache_read_tokens")
.map_postgres_err()?
.max(0) as u64,
total_cost_usd: row.try_get::<f64, _>("total_cost_usd").map_postgres_err()?,
actual_total_cost_usd: row
.try_get::<f64, _>("actual_total_cost_usd")
.map_postgres_err()?,
avg_response_time_ms: row
.try_get::<Option<f64>, _>("avg_response_time_ms")
.map_postgres_err()?,
success_count: row
.try_get::<Option<i64>, _>("success_count")
.map_postgres_err()?
.map(|value| value.max(0) as u64),
})
}
fn decode_usage_error_distribution_row(
row: &PgRow,
) -> Result<StoredUsageErrorDistributionRow, DataLayerError> {
@@ -1211,6 +1214,9 @@ const SUMMARIZE_USAGE_BY_PROVIDER_API_KEY_IDS_SQL: &str =
const APPLY_API_KEY_USAGE_DELTA_SQL: &str =
include_str!("queries/apply_api_key_usage_delta_sql.sql");
const APPLY_GLOBAL_MODEL_USAGE_DELTA_SQL: &str =
include_str!("queries/apply_global_model_usage_delta_sql.sql");
const RESET_API_KEY_USAGE_STATS_SQL: &str =
include_str!("queries/reset_api_key_usage_stats_sql.sql");
@@ -1220,9 +1226,6 @@ const REBUILD_API_KEY_USAGE_STATS_SQL: &str =
const APPLY_PROVIDER_API_KEY_USAGE_DELTA_SQL: &str =
include_str!("queries/apply_provider_api_key_usage_delta_sql.sql");
const APPLY_GLOBAL_MODEL_USAGE_DELTA_SQL: &str =
include_str!("queries/apply_global_model_usage_delta_sql.sql");
const RESET_PROVIDER_API_KEY_USAGE_STATS_SQL: &str =
include_str!("queries/reset_provider_api_key_usage_stats_sql.sql");
@@ -1323,6 +1326,99 @@ const LIST_RECENT_USAGE_AUDITS_PREFIX: &str =
const UPSERT_SQL: &str = include_str!("queries/upsert_sql.sql");
const SELECT_STALE_PENDING_USAGE_BATCH_SQL: &str = r#"
SELECT
usage.request_id,
usage.status,
COALESCE(usage_settlement_snapshots.billing_status, usage.billing_status) AS billing_status
FROM usage
LEFT JOIN usage_settlement_snapshots
ON usage_settlement_snapshots.request_id = usage.request_id
WHERE usage.status IN ('pending', 'streaming')
AND usage.created_at < $1
ORDER BY usage.created_at ASC, usage.request_id ASC
LIMIT $2
FOR UPDATE OF usage SKIP LOCKED
"#;
const SELECT_COMPLETED_PENDING_REQUEST_IDS_SQL: &str = r#"
SELECT DISTINCT request_id
FROM request_candidates
WHERE request_id = ANY($1)
AND (
status = 'streaming'
OR (
status = 'success'
AND COALESCE(extra_data->>'stream_completed', 'false') = 'true'
)
)
"#;
const UPDATE_RECOVERED_STALE_USAGE_SQL: &str = r#"
UPDATE usage
SET status = 'completed',
status_code = 200,
error_message = NULL
WHERE request_id = $1
"#;
const UPDATE_FAILED_STALE_USAGE_SQL: &str = r#"
UPDATE usage
SET status = 'failed',
status_code = 504,
error_message = $2
WHERE request_id = $1
"#;
const UPDATE_FAILED_VOID_STALE_USAGE_SQL: &str = r#"
WITH updated_usage AS (
UPDATE usage
SET status = 'failed',
status_code = 504,
error_message = $2,
billing_status = 'void',
finalized_at = $3,
total_cost_usd = 0,
request_cost_usd = 0,
actual_total_cost_usd = 0,
actual_request_cost_usd = 0
WHERE request_id = $1
RETURNING request_id
)
INSERT INTO usage_settlement_snapshots (
request_id,
billing_status,
finalized_at
)
SELECT request_id, 'void', $3
FROM updated_usage
ON CONFLICT (request_id)
DO UPDATE SET
billing_status = EXCLUDED.billing_status,
finalized_at = COALESCE(
usage_settlement_snapshots.finalized_at,
EXCLUDED.finalized_at
),
updated_at = NOW()
"#;
const UPDATE_RECOVERED_STREAMING_CANDIDATES_SQL: &str = r#"
UPDATE request_candidates
SET status = 'success',
finished_at = $2
WHERE request_id = $1
AND status = 'streaming'
"#;
const UPDATE_FAILED_PENDING_CANDIDATES_SQL: &str = r#"
UPDATE request_candidates
SET status = 'failed',
finished_at = $2,
error_message = ''
WHERE request_id = $1
AND status IN ('pending', 'streaming')
"#;
#[derive(Debug, Clone)]
pub struct SqlxUsageReadRepository {
pool: PgPool,
@@ -1398,7 +1494,7 @@ SELECT
COALESCE(SUM(input_tokens), 0)::BIGINT AS input_tokens,
COALESCE(SUM(effective_input_tokens), 0)::BIGINT AS effective_input_tokens,
COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens,
COALESCE(SUM(effective_input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0)::BIGINT AS total_tokens,
COALESCE(SUM(input_tokens + output_tokens), 0)::BIGINT AS total_tokens,
COALESCE(SUM(cache_creation_tokens), 0)::BIGINT AS cache_creation_tokens,
COALESCE(SUM(cache_read_tokens), 0)::BIGINT AS cache_read_tokens,
COALESCE(SUM(total_input_context), 0)::BIGINT AS total_input_context,
@@ -1407,7 +1503,7 @@ SELECT
COALESCE(SUM(total_cost), 0)::DOUBLE PRECISION AS total_cost_usd,
COALESCE(SUM(actual_total_cost), 0)::DOUBLE PRECISION AS actual_total_cost_usd,
COALESCE(SUM(error_requests), 0)::BIGINT AS error_requests,
COALESCE(SUM(response_time_sum_ms), 0) AS response_time_sum_ms,
COALESCE(SUM(response_time_sum_ms), 0)::DOUBLE PRECISION AS response_time_sum_ms,
COALESCE(SUM(response_time_samples), 0)::BIGINT AS response_time_samples
FROM stats_user_daily
WHERE user_id = $1
@@ -1429,7 +1525,7 @@ SELECT
COALESCE(SUM(input_tokens), 0)::BIGINT AS input_tokens,
COALESCE(SUM(effective_input_tokens), 0)::BIGINT AS effective_input_tokens,
COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens,
COALESCE(SUM(effective_input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0)::BIGINT AS total_tokens,
COALESCE(SUM(input_tokens + output_tokens), 0)::BIGINT AS total_tokens,
COALESCE(SUM(cache_creation_tokens), 0)::BIGINT AS cache_creation_tokens,
COALESCE(SUM(cache_read_tokens), 0)::BIGINT AS cache_read_tokens,
COALESCE(SUM(total_input_context), 0)::BIGINT AS total_input_context,
@@ -1438,7 +1534,7 @@ SELECT
COALESCE(SUM(total_cost), 0)::DOUBLE PRECISION AS total_cost_usd,
COALESCE(SUM(actual_total_cost), 0)::DOUBLE PRECISION AS actual_total_cost_usd,
COALESCE(SUM(error_requests), 0)::BIGINT AS error_requests,
COALESCE(SUM(response_time_sum_ms), 0) AS response_time_sum_ms,
COALESCE(SUM(response_time_sum_ms), 0)::DOUBLE PRECISION AS response_time_sum_ms,
COALESCE(SUM(response_time_samples), 0)::BIGINT AS response_time_samples
FROM stats_daily
WHERE date >= $1
@@ -1486,33 +1582,7 @@ SELECT
END
), 0)::BIGINT AS effective_input_tokens,
COALESCE(SUM(GREATEST(COALESCE("usage".output_tokens, 0), 0)), 0)::BIGINT AS output_tokens,
COALESCE(SUM(
CASE
WHEN GREATEST(COALESCE("usage".input_tokens, 0), 0) <= 0 THEN 0
WHEN GREATEST(COALESCE("usage".cache_read_input_tokens, 0), 0) <= 0
THEN GREATEST(COALESCE("usage".input_tokens, 0), 0)
WHEN split_part(lower(COALESCE(COALESCE("usage".endpoint_api_format, "usage".api_format), '')), ':', 1)
IN ('openai', 'gemini', 'google')
THEN GREATEST(
GREATEST(COALESCE("usage".input_tokens, 0), 0)
- GREATEST(COALESCE("usage".cache_read_input_tokens, 0), 0),
0
)
ELSE GREATEST(COALESCE("usage".input_tokens, 0), 0)
END
+ GREATEST(COALESCE("usage".output_tokens, 0), 0)
+ CASE
WHEN COALESCE("usage".cache_creation_input_tokens, 0) = 0
AND (
COALESCE("usage".cache_creation_input_tokens_5m, 0)
+ COALESCE("usage".cache_creation_input_tokens_1h, 0)
) > 0
THEN COALESCE("usage".cache_creation_input_tokens_5m, 0)
+ COALESCE("usage".cache_creation_input_tokens_1h, 0)
ELSE COALESCE("usage".cache_creation_input_tokens, 0)
END
+ GREATEST(COALESCE("usage".cache_read_input_tokens, 0), 0)
), 0)::BIGINT AS total_tokens,
COALESCE(SUM(GREATEST(COALESCE("usage".total_tokens, 0), 0)), 0)::BIGINT AS total_tokens,
COALESCE(SUM(
CASE
WHEN COALESCE("usage".cache_creation_input_tokens, 0) = 0
@@ -2462,7 +2532,7 @@ SELECT
COALESCE(SUM(actual_total_cost), 0)::DOUBLE PRECISION AS actual_total_cost_usd,
COALESCE(SUM(cache_creation_cost), 0)::DOUBLE PRECISION AS cache_creation_cost_usd,
COALESCE(SUM(cache_read_cost), 0)::DOUBLE PRECISION AS cache_read_cost_usd,
COALESCE(SUM(response_time_sum_ms), 0) AS total_response_time_ms,
COALESCE(SUM(response_time_sum_ms), 0)::DOUBLE PRECISION AS total_response_time_ms,
COALESCE(SUM(error_requests), 0)::BIGINT AS error_requests
FROM stats_user_daily
WHERE user_id = $1
@@ -2494,7 +2564,7 @@ SELECT
COALESCE(SUM(actual_total_cost), 0)::DOUBLE PRECISION AS actual_total_cost_usd,
COALESCE(SUM(cache_creation_cost), 0)::DOUBLE PRECISION AS cache_creation_cost_usd,
COALESCE(SUM(cache_read_cost), 0)::DOUBLE PRECISION AS cache_read_cost_usd,
COALESCE(SUM(response_time_sum_ms), 0) AS total_response_time_ms,
COALESCE(SUM(response_time_sum_ms), 0)::DOUBLE PRECISION AS total_response_time_ms,
COALESCE(SUM(error_requests), 0)::BIGINT AS error_requests
FROM stats_daily
WHERE date >= $1
@@ -3770,33 +3840,7 @@ SELECT
"usage".model AS model,
"usage".provider_name AS provider,
COUNT(*)::BIGINT AS requests,
COALESCE(SUM(
CASE
WHEN GREATEST(COALESCE("usage".input_tokens, 0), 0) <= 0 THEN 0
WHEN GREATEST(COALESCE("usage".cache_read_input_tokens, 0), 0) <= 0
THEN GREATEST(COALESCE("usage".input_tokens, 0), 0)
WHEN split_part(lower(COALESCE(COALESCE("usage".endpoint_api_format, "usage".api_format), '')), ':', 1)
IN ('openai', 'gemini', 'google')
THEN GREATEST(
GREATEST(COALESCE("usage".input_tokens, 0), 0)
- GREATEST(COALESCE("usage".cache_read_input_tokens, 0), 0),
0
)
ELSE GREATEST(COALESCE("usage".input_tokens, 0), 0)
END
+ GREATEST(COALESCE("usage".output_tokens, 0), 0)
+ CASE
WHEN COALESCE("usage".cache_creation_input_tokens, 0) = 0
AND (
COALESCE("usage".cache_creation_input_tokens_5m, 0)
+ COALESCE("usage".cache_creation_input_tokens_1h, 0)
) > 0
THEN COALESCE("usage".cache_creation_input_tokens_5m, 0)
+ COALESCE("usage".cache_creation_input_tokens_1h, 0)
ELSE COALESCE("usage".cache_creation_input_tokens, 0)
END
+ GREATEST(COALESCE("usage".cache_read_input_tokens, 0), 0)
), 0)::BIGINT AS total_tokens,
COALESCE(SUM(GREATEST(COALESCE("usage".total_tokens, 0), 0)), 0)::BIGINT AS total_tokens,
COALESCE(SUM(COALESCE(CAST("usage".total_cost_usd AS DOUBLE PRECISION), 0)), 0)
AS total_cost_usd,
COALESCE(SUM(
@@ -4000,7 +4044,7 @@ SELECT
{group_column} AS group_key,
COALESCE(SUM(total_requests), 0)::BIGINT AS request_count,
COALESCE(SUM(input_tokens), 0)::BIGINT AS input_tokens,
COALESCE(SUM(effective_input_tokens + output_tokens + cache_creation_tokens + cache_read_tokens), 0)::BIGINT AS total_tokens,
COALESCE(SUM(total_tokens), 0)::BIGINT AS total_tokens,
COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens,
COALESCE(SUM(effective_input_tokens), 0)::BIGINT AS effective_input_tokens,
COALESCE(SUM(total_input_context), 0)::BIGINT AS total_input_context,
@@ -4735,17 +4779,17 @@ SELECT
COALESCE(SUM(success_flag), 0)::BIGINT AS success_count,
CASE
WHEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
THEN output_tps_duration_ms
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0) > 0
THEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN output_tokens
ELSE 0
END), 0)::DOUBLE PRECISION * 1000.0 / COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
THEN output_tps_duration_ms
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0)::DOUBLE PRECISION
ELSE NULL
@@ -4822,17 +4866,17 @@ SELECT
COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens,
CASE
WHEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
THEN output_tps_duration_ms
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0) > 0
THEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN output_tokens
ELSE 0
END), 0)::DOUBLE PRECISION * 1000.0 / COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
THEN output_tps_duration_ms
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0)::DOUBLE PRECISION
ELSE NULL
@@ -4854,7 +4898,7 @@ SELECT
ELSE NULL
END AS p90_first_byte_time_ms,
COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN 1
ELSE 0
END), 0)::BIGINT AS tps_sample_count,
@@ -4954,17 +4998,17 @@ SELECT
COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens,
CASE
WHEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
THEN output_tps_duration_ms
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0) > 0
THEN COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN output_tokens
ELSE 0
END), 0)::DOUBLE PRECISION * 1000.0 / COALESCE(SUM(CASE
WHEN success_flag = 1 AND output_tps_duration_ms > 0 AND output_tokens > 0
THEN output_tps_duration_ms
WHEN success_flag = 1 AND response_time_ms > 0 AND output_tokens > 0
THEN response_time_ms
ELSE 0
END), 0)::DOUBLE PRECISION
ELSE NULL
@@ -5040,7 +5084,7 @@ SELECT
AS cache_creation_cost_usd,
COALESCE(SUM(
COALESCE(
CAST("usage".input_price_per_1m AS DOUBLE PRECISION),
CAST("usage".output_price_per_1m AS DOUBLE PRECISION),
0
) * GREATEST(COALESCE("usage".cache_read_input_tokens, 0), 0)::DOUBLE PRECISION / 1000000.0
), 0) AS estimated_full_cost_usd
@@ -6157,6 +6201,11 @@ WHERE date >= $1
GROUP BY {group_column}
ORDER BY request_count DESC, group_key ASC
"#,
group_column = group_column,
display_name_expr = display_name_expr,
avg_response_time_expr = avg_response_time_expr,
success_count_expr = success_count_expr,
table_name = table_name,
);
let mut rows = sqlx::query(&sql)
@@ -6337,7 +6386,62 @@ LIMIT $3
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_postgres_err()? {
items.push(decode_usage_audit_aggregation_row(&row)?);
items.push(StoredUsageAuditAggregation {
group_key: row.try_get::<String, _>("group_key").map_postgres_err()?,
display_name: row
.try_get::<Option<String>, _>("display_name")
.map_postgres_err()?,
secondary_name: row
.try_get::<Option<String>, _>("secondary_name")
.map_postgres_err()?,
request_count: row
.try_get::<i64, _>("request_count")
.map_postgres_err()?
.max(0) as u64,
total_tokens: row
.try_get::<i64, _>("total_tokens")
.map_postgres_err()?
.max(0) as u64,
output_tokens: row
.try_get::<i64, _>("output_tokens")
.map_postgres_err()?
.max(0) as u64,
effective_input_tokens: row
.try_get::<i64, _>("effective_input_tokens")
.map_postgres_err()?
.max(0) as u64,
total_input_context: row
.try_get::<i64, _>("total_input_context")
.map_postgres_err()?
.max(0) as u64,
cache_creation_tokens: row
.try_get::<i64, _>("cache_creation_tokens")
.map_postgres_err()?
.max(0) as u64,
cache_creation_ephemeral_5m_tokens: row
.try_get::<i64, _>("cache_creation_ephemeral_5m_tokens")
.map_postgres_err()?
.max(0) as u64,
cache_creation_ephemeral_1h_tokens: row
.try_get::<i64, _>("cache_creation_ephemeral_1h_tokens")
.map_postgres_err()?
.max(0) as u64,
cache_read_tokens: row
.try_get::<i64, _>("cache_read_tokens")
.map_postgres_err()?
.max(0) as u64,
total_cost_usd: row.try_get::<f64, _>("total_cost_usd").map_postgres_err()?,
actual_total_cost_usd: row
.try_get::<f64, _>("actual_total_cost_usd")
.map_postgres_err()?,
avg_response_time_ms: row
.try_get::<Option<f64>, _>("avg_response_time_ms")
.map_postgres_err()?,
success_count: row
.try_get::<Option<i64>, _>("success_count")
.map_postgres_err()?
.map(|value| value.max(0) as u64),
});
}
Ok(items)
}
@@ -6474,13 +6578,13 @@ WHERE "usage".created_at >= $1
r#"
SELECT
date,
total_requests,
input_tokens,
output_tokens,
cache_creation_tokens,
cache_read_tokens,
total_cost,
COALESCE(actual_total_cost, 0) AS actual_total_cost
total_requests::BIGINT AS total_requests,
input_tokens::BIGINT AS input_tokens,
output_tokens::BIGINT AS output_tokens,
cache_creation_tokens::BIGINT AS cache_creation_tokens,
cache_read_tokens::BIGINT AS cache_read_tokens,
total_cost::DOUBLE PRECISION AS total_cost,
COALESCE(actual_total_cost, 0)::DOUBLE PRECISION AS actual_total_cost
FROM stats_user_daily
WHERE user_id = $1
AND date >= $2
@@ -6497,13 +6601,13 @@ ORDER BY date ASC
r#"
SELECT
date,
total_requests,
input_tokens,
output_tokens,
cache_creation_tokens,
cache_read_tokens,
total_cost,
actual_total_cost
total_requests::BIGINT AS total_requests,
input_tokens::BIGINT AS input_tokens,
output_tokens::BIGINT AS output_tokens,
cache_creation_tokens::BIGINT AS cache_creation_tokens,
cache_read_tokens::BIGINT AS cache_read_tokens,
total_cost::DOUBLE PRECISION AS total_cost,
actual_total_cost::DOUBLE PRECISION AS actual_total_cost
FROM stats_daily
WHERE date >= $1
AND date < $2
@@ -7267,6 +7371,40 @@ ORDER BY "usage".user_id ASC
}
}
let before_model_contribution =
previous_usage.as_ref().and_then(model_usage_contribution);
let after_model_contribution = model_usage_contribution(&stored);
match (
before_model_contribution.as_ref(),
after_model_contribution.as_ref(),
) {
(Some(before), Some(after)) if before.model == after.model => {
let delta = ModelUsageDelta::between(before, after);
apply_global_model_usage_delta_in_tx(tx, before.model.as_str(), &delta)
.await?;
}
_ => {
if let Some(before) = before_model_contribution.as_ref() {
let delta = ModelUsageDelta::removal(before);
apply_global_model_usage_delta_in_tx(
tx,
before.model.as_str(),
&delta,
)
.await?;
}
if let Some(after) = after_model_contribution.as_ref() {
let delta = ModelUsageDelta::addition(after);
apply_global_model_usage_delta_in_tx(
tx,
after.model.as_str(),
&delta,
)
.await?;
}
}
}
let before_provider_contribution = previous_usage
.as_ref()
.and_then(provider_api_key_usage_contribution);
@@ -7305,46 +7443,140 @@ ORDER BY "usage".user_id ASC
}
}
}
let before_model_contribution =
previous_usage.as_ref().and_then(model_usage_contribution);
let after_model_contribution = model_usage_contribution(&stored);
match (
before_model_contribution.as_ref(),
after_model_contribution.as_ref(),
) {
(Some(before), Some(after)) if before.model == after.model => {
let delta = ModelUsageDelta::between(before, after);
apply_global_model_usage_delta_in_tx(tx, before.model.as_str(), &delta)
.await?;
}
_ => {
if let Some(before) = before_model_contribution.as_ref() {
let delta = ModelUsageDelta::removal(before);
apply_global_model_usage_delta_in_tx(
tx,
before.model.as_str(),
&delta,
)
.await?;
}
if let Some(after) = after_model_contribution.as_ref() {
let delta = ModelUsageDelta::addition(after);
apply_global_model_usage_delta_in_tx(
tx,
after.model.as_str(),
&delta,
)
.await?;
}
}
}
Ok(stored)
}) as BoxFuture<'_, Result<StoredRequestUsageAudit, DataLayerError>>
})
.await
}
pub async fn cleanup_stale_pending_requests(
&self,
cutoff_unix_secs: u64,
now_unix_secs: u64,
timeout_minutes: u64,
batch_size: usize,
) -> Result<PendingUsageCleanupSummary, DataLayerError> {
if batch_size == 0 {
return Ok(PendingUsageCleanupSummary::default());
}
let cutoff_timestamp = i64::try_from(cutoff_unix_secs).map_err(|_| {
DataLayerError::InvalidInput(format!(
"invalid stale pending usage cutoff: {cutoff_unix_secs}"
))
})?;
let now_timestamp = i64::try_from(now_unix_secs).map_err(|_| {
DataLayerError::InvalidInput(format!(
"invalid stale pending usage timestamp: {now_unix_secs}"
))
})?;
let cutoff = DateTime::<Utc>::from_timestamp(cutoff_timestamp, 0).ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"invalid stale pending usage cutoff: {cutoff_unix_secs}"
))
})?;
let now = DateTime::<Utc>::from_timestamp(now_timestamp, 0).ok_or_else(|| {
DataLayerError::InvalidInput(format!(
"invalid stale pending usage timestamp: {now_unix_secs}"
))
})?;
let mut summary = PendingUsageCleanupSummary::default();
let batch_size_i64 = i64::try_from(batch_size).map_err(|_| {
DataLayerError::InvalidInput(format!(
"invalid stale pending usage batch size: {batch_size}"
))
})?;
loop {
let mut tx = self.pool.begin().await.map_postgres_err()?;
let stale_rows = sqlx::query(SELECT_STALE_PENDING_USAGE_BATCH_SQL)
.bind(cutoff)
.bind(batch_size_i64)
.fetch_all(&mut *tx)
.await
.map_postgres_err()?;
if stale_rows.is_empty() {
tx.rollback().await.map_postgres_err()?;
break;
}
let stale_rows = stale_rows
.iter()
.map(|row| {
Ok(StalePendingUsageRow {
request_id: row.try_get("request_id").map_postgres_err()?,
status: row.try_get("status").map_postgres_err()?,
billing_status: row.try_get("billing_status").map_postgres_err()?,
})
})
.collect::<Result<Vec<_>, DataLayerError>>()?;
let request_ids = stale_rows
.iter()
.map(|row| row.request_id.clone())
.collect::<Vec<_>>();
let completed_request_ids = if request_ids.is_empty() {
Vec::new()
} else {
sqlx::query(SELECT_COMPLETED_PENDING_REQUEST_IDS_SQL)
.bind(request_ids)
.fetch_all(&mut *tx)
.await
.map_postgres_err()?
.iter()
.map(|row| row.try_get("request_id").map_postgres_err())
.collect::<Result<Vec<String>, DataLayerError>>()?
};
for row in stale_rows {
if completed_request_ids.contains(&row.request_id) {
sqlx::query(UPDATE_RECOVERED_STALE_USAGE_SQL)
.bind(&row.request_id)
.execute(&mut *tx)
.await
.map_postgres_err()?;
sqlx::query(UPDATE_RECOVERED_STREAMING_CANDIDATES_SQL)
.bind(&row.request_id)
.bind(now)
.execute(&mut *tx)
.await
.map_postgres_err()?;
summary.recovered += 1;
continue;
}
let error_message = stale_pending_error_message(&row.status, timeout_minutes);
if row.billing_status == "pending" {
sqlx::query(UPDATE_FAILED_VOID_STALE_USAGE_SQL)
.bind(&row.request_id)
.bind(&error_message)
.bind(now)
.execute(&mut *tx)
.await
.map_postgres_err()?;
} else {
sqlx::query(UPDATE_FAILED_STALE_USAGE_SQL)
.bind(&row.request_id)
.bind(&error_message)
.execute(&mut *tx)
.await
.map_postgres_err()?;
}
sqlx::query(UPDATE_FAILED_PENDING_CANDIDATES_SQL)
.bind(&row.request_id)
.bind(now)
.execute(&mut *tx)
.await
.map_postgres_err()?;
summary.failed += 1;
}
tx.commit().await.map_postgres_err()?;
}
Ok(summary)
}
pub async fn rebuild_api_key_usage_stats(&self) -> Result<u64, DataLayerError> {
self.tx_runner
.run_read_write(|tx| {
@@ -7624,6 +7856,42 @@ impl UsageWriteRepository for SqlxUsageReadRepository {
async fn rebuild_provider_api_key_usage_stats(&self) -> Result<u64, DataLayerError> {
Self::rebuild_provider_api_key_usage_stats(self).await
}
async fn cleanup_stale_pending_requests(
&self,
cutoff_unix_secs: u64,
now_unix_secs: u64,
timeout_minutes: u64,
batch_size: usize,
) -> Result<PendingUsageCleanupSummary, DataLayerError> {
Self::cleanup_stale_pending_requests(
self,
cutoff_unix_secs,
now_unix_secs,
timeout_minutes,
batch_size,
)
.await
}
async fn cleanup_usage(
&self,
window: &UsageCleanupWindow,
batch_size: usize,
auto_delete_expired_keys: bool,
) -> Result<UsageCleanupSummary, DataLayerError> {
Self::cleanup_usage(self, window, batch_size, auto_delete_expired_keys).await
}
}
struct StalePendingUsageRow {
request_id: String,
status: String,
billing_status: String,
}
fn stale_pending_error_message(status: &str, timeout_minutes: u64) -> String {
format!("请求超时: 状态 '{status}' 超过 {timeout_minutes} 分钟未完成")
}
async fn find_usage_by_request_id_in_tx(
@@ -7695,6 +7963,32 @@ async fn apply_api_key_usage_delta_in_tx(
Ok(())
}
async fn apply_global_model_usage_delta_in_tx(
tx: &mut sqlx::Transaction<'_, Postgres>,
model: &str,
delta: &ModelUsageDelta,
) -> Result<(), DataLayerError> {
if model.trim().is_empty() {
return Ok(());
}
if delta.is_noop() {
return Ok(());
}
sqlx::query(APPLY_GLOBAL_MODEL_USAGE_DELTA_SQL)
.bind(model)
.bind(i32::try_from(delta.request_count).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
"global_models.usage_count delta exceeds i32: {}",
delta.request_count
))
})?)
.execute(&mut **tx)
.await
.map_postgres_err()?;
Ok(())
}
async fn apply_provider_api_key_usage_delta_in_tx(
tx: &mut sqlx::Transaction<'_, Postgres>,
key_id: &str,
@@ -7757,33 +8051,6 @@ async fn apply_provider_api_key_usage_delta_in_tx(
Ok(())
}
async fn apply_global_model_usage_delta_in_tx(
tx: &mut sqlx::Transaction<'_, Postgres>,
model: &str,
delta: &ModelUsageDelta,
) -> Result<(), DataLayerError> {
let model = model.trim();
if model.is_empty() {
return Ok(());
}
if delta.is_noop() {
return Ok(());
}
sqlx::query(APPLY_GLOBAL_MODEL_USAGE_DELTA_SQL)
.bind(model)
.bind(i32::try_from(delta.request_count).map_err(|_| {
DataLayerError::UnexpectedValue(format!(
"global_models.usage_count delta exceeds i32: {}",
delta.request_count
))
})?)
.execute(&mut **tx)
.await
.map_postgres_err()?;
Ok(())
}
// Build the usage read model from the split storage layout.
//
// Query projections already prefer the newer audit/snapshot owners and only fall back to

View File

@@ -12,7 +12,7 @@ use super::{
UsageHttpAuditRefs, UsageRoutingSnapshot, UsageSettlementPricingSnapshot,
MAX_INLINE_USAGE_BODY_BYTES,
};
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::repository::usage::UpsertUsageRecord;
use aether_data_contracts::repository::usage::UsageBodyField;
@@ -382,6 +382,8 @@ fn usage_sql_summarize_usage_daily_heatmap_supports_daily_aggregates() {
assert!(source.contains("summarize_usage_daily_heatmap_from_daily_aggregates"));
assert!(source.contains("FROM stats_daily"));
assert!(source.contains("FROM stats_user_daily"));
assert!(source.contains("total_requests::BIGINT AS total_requests"));
assert!(source.contains("total_cost::DOUBLE PRECISION AS total_cost"));
assert!(
source.contains("split_dashboard_daily_aggregate_range(start_utc, end_utc, cutoff_utc)")
);
@@ -458,18 +460,14 @@ fn usage_sql_provider_performance_reads_upstream_stream_from_billing_facts() {
#[test]
fn usage_billing_facts_projects_upstream_stream_mode() {
let migration = include_str!(
"../../../../migrations/20260505130000_project_upstream_stream_in_usage_billing_facts.sql"
"../../../../migrations/postgres/20260505130000_project_upstream_stream_in_usage_billing_facts.sql"
);
let baseline = include_str!("../../../../bootstrap/20260413020000_baseline_v2.sql");
for source in [migration, baseline] {
assert!(source.contains("AS upstream_is_stream"));
assert!(source.contains("COALESCE(usage_rows.upstream_is_stream"));
assert!(source.contains("COALESCE(usage_rows.is_stream, FALSE)"));
}
assert!(migration.contains("AS upstream_is_stream"));
assert!(migration.contains("COALESCE(usage_rows.upstream_is_stream"));
assert!(migration.contains("COALESCE(usage_rows.is_stream, FALSE)"));
assert!(migration.contains("ADD COLUMN IF NOT EXISTS upstream_is_stream boolean"));
assert!(migration.contains("request_metadata->>'upstream_is_stream'"));
assert!(baseline.contains("upstream_is_stream boolean"));
}
#[test]

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -1,10 +1,15 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
mod types;
pub use memory::InMemoryUserReadRepository;
pub use sql::SqlxUserReadRepository;
pub use mysql::MysqlUserReadRepository;
pub use postgres::SqlxUserReadRepository;
pub use sqlite::SqliteUserReadRepository;
pub use types::{
StoredUserAuthRecord, StoredUserExportRow, StoredUserPreferenceRecord, StoredUserSessionRecord,
StoredUserSummary, UserExportListQuery, UserExportSummary, UserReadRepository,
StoredUserAuthRecord, StoredUserExportRow, StoredUserOAuthLinkSummary,
StoredUserPreferenceRecord, StoredUserSessionRecord, StoredUserSummary, UserExportListQuery,
UserExportSummary, UserReadRepository,
};

File diff suppressed because it is too large Load Diff

File diff suppressed because it is too large Load Diff

View File

@@ -1,473 +0,0 @@
use async_trait::async_trait;
use futures_util::TryStreamExt;
use sqlx::{PgPool, Postgres, QueryBuilder, Row};
use super::types::{
StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary, UserExportListQuery,
UserExportSummary, UserReadRepository,
};
use crate::{error::SqlxResultExt, DataLayerError};
const LIST_USERS_BY_IDS_SQL: &str = r#"
SELECT
id,
username,
email,
role::text AS role,
is_active,
is_deleted
FROM users
WHERE id = ANY($1::text[])
ORDER BY id ASC
"#;
const LIST_USERS_BY_USERNAME_SEARCH_SQL: &str = r#"
SELECT
id,
username,
email,
role::text AS role,
is_active,
is_deleted
FROM users
WHERE is_deleted IS FALSE
AND LOWER(username) LIKE $1
ORDER BY id ASC
"#;
const LIST_NON_ADMIN_EXPORT_USERS_SQL: &str = r#"
SELECT
id,
email,
email_verified,
username,
password_hash,
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_api_formats,
allowed_models,
rate_limit,
model_capability_settings,
is_active
FROM users
WHERE is_deleted IS FALSE
AND role::text != 'admin'
ORDER BY id ASC
"#;
const LIST_EXPORT_USERS_SQL: &str = r#"
SELECT
id,
email,
email_verified,
username,
password_hash,
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_api_formats,
allowed_models,
rate_limit,
model_capability_settings,
is_active
FROM users
WHERE is_deleted IS FALSE
ORDER BY id ASC
"#;
const LIST_EXPORT_USERS_PAGE_PREFIX: &str = r#"
SELECT
id,
email,
email_verified,
username,
password_hash,
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_api_formats,
allowed_models,
rate_limit,
model_capability_settings,
is_active
FROM users
WHERE is_deleted IS FALSE
"#;
const SUMMARIZE_EXPORT_USERS_SQL: &str = r#"
SELECT
COUNT(*)::BIGINT AS total,
COUNT(*) FILTER (WHERE is_active = TRUE)::BIGINT AS active
FROM users
WHERE is_deleted IS FALSE
"#;
const FIND_EXPORT_USER_BY_ID_SQL: &str = r#"
SELECT
id,
email,
email_verified,
username,
password_hash,
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_api_formats,
allowed_models,
rate_limit,
model_capability_settings,
is_active
FROM users
WHERE is_deleted IS FALSE
AND id = $1
LIMIT 1
"#;
const FIND_USER_AUTH_BY_ID_SQL: &str = r#"
SELECT
id,
email,
email_verified,
username,
password_hash,
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_api_formats,
allowed_models,
is_active,
is_deleted,
created_at,
last_login_at
FROM users
WHERE id = $1
LIMIT 1
"#;
const LIST_USER_AUTH_BY_IDS_SQL: &str = r#"
SELECT
id,
email,
email_verified,
username,
password_hash,
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_api_formats,
allowed_models,
is_active,
is_deleted,
created_at,
last_login_at
FROM users
WHERE id = ANY($1::text[])
ORDER BY id ASC
"#;
const FIND_USER_AUTH_BY_IDENTIFIER_SQL: &str = r#"
SELECT
id,
email,
email_verified,
username,
password_hash,
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_api_formats,
allowed_models,
is_active,
is_deleted,
created_at,
last_login_at
FROM users
WHERE email = $1 OR username = $1
LIMIT 1
"#;
#[derive(Debug, Clone)]
pub struct SqlxUserReadRepository {
pool: PgPool,
}
impl SqlxUserReadRepository {
pub fn new(pool: PgPool) -> Self {
Self { pool }
}
pub async fn list_users_by_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserSummary>, DataLayerError> {
if user_ids.is_empty() {
return Ok(Vec::new());
}
collect_query_rows(
sqlx::query(LIST_USERS_BY_IDS_SQL)
.bind(user_ids)
.fetch(&self.pool),
map_user_row,
)
.await
}
pub async fn list_users_by_username_search(
&self,
username_search: &str,
) -> Result<Vec<StoredUserSummary>, DataLayerError> {
let username_search = username_search.trim();
if username_search.is_empty() {
return Ok(Vec::new());
}
collect_query_rows(
sqlx::query(LIST_USERS_BY_USERNAME_SEARCH_SQL)
.bind(format!("%{}%", username_search.to_ascii_lowercase()))
.fetch(&self.pool),
map_user_row,
)
.await
}
pub async fn list_non_admin_export_users(
&self,
) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
collect_query_rows(
sqlx::query(LIST_NON_ADMIN_EXPORT_USERS_SQL).fetch(&self.pool),
map_user_export_row,
)
.await
}
pub async fn list_export_users(&self) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
collect_query_rows(
sqlx::query(LIST_EXPORT_USERS_SQL).fetch(&self.pool),
map_user_export_row,
)
.await
}
pub async fn list_export_users_page(
&self,
query: &UserExportListQuery,
) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
let mut builder = QueryBuilder::<Postgres>::new(LIST_EXPORT_USERS_PAGE_PREFIX);
if let Some(role) = query.role.as_deref() {
builder
.push(" AND LOWER(role::text) = ")
.push_bind(role.trim().to_ascii_lowercase());
}
if let Some(is_active) = query.is_active {
builder.push(" AND is_active = ").push_bind(is_active);
}
builder
.push(" ORDER BY id ASC OFFSET ")
.push_bind(i64::try_from(query.skip).map_err(|_| {
DataLayerError::InvalidInput(format!("invalid user export skip: {}", query.skip))
})?)
.push(" LIMIT ")
.push_bind(i64::try_from(query.limit).map_err(|_| {
DataLayerError::InvalidInput(format!("invalid user export limit: {}", query.limit))
})?);
let query = builder.build();
collect_query_rows(query.fetch(&self.pool), map_user_export_row).await
}
pub async fn summarize_export_users(&self) -> Result<UserExportSummary, DataLayerError> {
let row = sqlx::query(SUMMARIZE_EXPORT_USERS_SQL)
.fetch_one(&self.pool)
.await
.map_postgres_err()?;
Ok(UserExportSummary {
total: row.try_get::<i64, _>("total").map_postgres_err()?.max(0) as u64,
active: row.try_get::<i64, _>("active").map_postgres_err()?.max(0) as u64,
})
}
pub async fn find_export_user_by_id(
&self,
user_id: &str,
) -> Result<Option<StoredUserExportRow>, DataLayerError> {
let row = sqlx::query(FIND_EXPORT_USER_BY_ID_SQL)
.bind(user_id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_user_export_row).transpose()
}
pub async fn list_user_auth_by_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserAuthRecord>, DataLayerError> {
if user_ids.is_empty() {
return Ok(Vec::new());
}
collect_query_rows(
sqlx::query(LIST_USER_AUTH_BY_IDS_SQL)
.bind(user_ids)
.fetch(&self.pool),
map_user_auth_row,
)
.await
}
pub async fn find_user_auth_by_id(
&self,
user_id: &str,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let row = sqlx::query(FIND_USER_AUTH_BY_ID_SQL)
.bind(user_id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_user_auth_row).transpose()
}
pub async fn find_user_auth_by_identifier(
&self,
identifier: &str,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
let row = sqlx::query(FIND_USER_AUTH_BY_IDENTIFIER_SQL)
.bind(identifier)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_user_auth_row).transpose()
}
}
fn map_user_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserSummary, DataLayerError> {
StoredUserSummary::new(
row.try_get("id").map_postgres_err()?,
row.try_get("username").map_postgres_err()?,
row.try_get("email").map_postgres_err()?,
row.try_get("role").map_postgres_err()?,
row.try_get("is_active").map_postgres_err()?,
row.try_get("is_deleted").map_postgres_err()?,
)
}
fn map_user_export_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserExportRow, DataLayerError> {
StoredUserExportRow::new(
row.try_get("id").map_postgres_err()?,
row.try_get("email").map_postgres_err()?,
row.try_get("email_verified").map_postgres_err()?,
row.try_get("username").map_postgres_err()?,
row.try_get("password_hash").map_postgres_err()?,
row.try_get("role").map_postgres_err()?,
row.try_get("auth_source").map_postgres_err()?,
row.try_get("allowed_providers").map_postgres_err()?,
row.try_get("allowed_api_formats").map_postgres_err()?,
row.try_get("allowed_models").map_postgres_err()?,
row.try_get("rate_limit").map_postgres_err()?,
row.try_get("model_capability_settings")
.map_postgres_err()?,
row.try_get("is_active").map_postgres_err()?,
)
}
fn map_user_auth_row(row: &sqlx::postgres::PgRow) -> Result<StoredUserAuthRecord, DataLayerError> {
StoredUserAuthRecord::new(
row.try_get("id").map_postgres_err()?,
row.try_get("email").map_postgres_err()?,
row.try_get("email_verified").map_postgres_err()?,
row.try_get("username").map_postgres_err()?,
row.try_get("password_hash").map_postgres_err()?,
row.try_get("role").map_postgres_err()?,
row.try_get("auth_source").map_postgres_err()?,
row.try_get("allowed_providers").map_postgres_err()?,
row.try_get("allowed_api_formats").map_postgres_err()?,
row.try_get("allowed_models").map_postgres_err()?,
row.try_get("is_active").map_postgres_err()?,
row.try_get("is_deleted").map_postgres_err()?,
row.try_get("created_at").map_postgres_err()?,
row.try_get("last_login_at").map_postgres_err()?,
)
}
async fn collect_query_rows<T, S>(
mut rows: S,
mapper: fn(&sqlx::postgres::PgRow) -> Result<T, DataLayerError>,
) -> Result<Vec<T>, DataLayerError>
where
S: futures_util::TryStream<Ok = sqlx::postgres::PgRow, Error = sqlx::Error> + Unpin,
{
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_postgres_err()? {
items.push(mapper(&row)?);
}
Ok(items)
}
#[async_trait]
impl UserReadRepository for SqlxUserReadRepository {
async fn list_users_by_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserSummary>, DataLayerError> {
self.list_users_by_ids(user_ids).await
}
async fn list_users_by_username_search(
&self,
username_search: &str,
) -> Result<Vec<StoredUserSummary>, DataLayerError> {
self.list_users_by_username_search(username_search).await
}
async fn list_non_admin_export_users(
&self,
) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
self.list_non_admin_export_users().await
}
async fn list_export_users(&self) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
self.list_export_users().await
}
async fn list_export_users_page(
&self,
query: &UserExportListQuery,
) -> Result<Vec<StoredUserExportRow>, DataLayerError> {
self.list_export_users_page(query).await
}
async fn summarize_export_users(&self) -> Result<UserExportSummary, DataLayerError> {
self.summarize_export_users().await
}
async fn find_export_user_by_id(
&self,
user_id: &str,
) -> Result<Option<StoredUserExportRow>, DataLayerError> {
self.find_export_user_by_id(user_id).await
}
async fn find_user_auth_by_id(
&self,
user_id: &str,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.find_user_auth_by_id(user_id).await
}
async fn list_user_auth_by_ids(
&self,
user_ids: &[String],
) -> Result<Vec<StoredUserAuthRecord>, DataLayerError> {
self.list_user_auth_by_ids(user_ids).await
}
async fn find_user_auth_by_identifier(
&self,
identifier: &str,
) -> Result<Option<StoredUserAuthRecord>, DataLayerError> {
self.find_user_auth_by_identifier(identifier).await
}
}

File diff suppressed because it is too large Load Diff

View File

@@ -137,6 +137,56 @@ impl StoredUserAuthRecord {
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct LdapAuthUserProvisioningOutcome {
pub user: StoredUserAuthRecord,
pub created: bool,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct StoredUserOAuthLinkSummary {
pub provider_type: String,
pub display_name: String,
pub provider_username: Option<String>,
pub provider_email: Option<String>,
pub linked_at: Option<DateTime<Utc>>,
pub last_login_at: Option<DateTime<Utc>>,
pub provider_enabled: bool,
}
impl StoredUserOAuthLinkSummary {
#[allow(clippy::too_many_arguments)]
pub fn new(
provider_type: String,
display_name: String,
provider_username: Option<String>,
provider_email: Option<String>,
linked_at: Option<DateTime<Utc>>,
last_login_at: Option<DateTime<Utc>>,
provider_enabled: bool,
) -> Result<Self, crate::DataLayerError> {
if provider_type.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"user_oauth_links.provider_type is empty".to_string(),
));
}
if display_name.trim().is_empty() {
return Err(crate::DataLayerError::UnexpectedValue(
"oauth_providers.display_name is empty".to_string(),
));
}
Ok(Self {
provider_type,
display_name,
provider_username,
provider_email,
linked_at,
last_login_at,
provider_enabled,
})
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredUserExportRow {
pub id: String,
@@ -430,6 +480,232 @@ pub trait UserReadRepository: Send + Sync {
&self,
identifier: &str,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn find_user_auth_by_email(
&self,
email: &str,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn find_active_user_auth_by_email_ci(
&self,
email: &str,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn find_user_auth_by_username(
&self,
username: &str,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn list_user_oauth_links(
&self,
user_id: &str,
) -> Result<Vec<StoredUserOAuthLinkSummary>, crate::DataLayerError>;
async fn find_oauth_linked_user(
&self,
provider_type: &str,
provider_user_id: &str,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn touch_oauth_link(
&self,
provider_type: &str,
provider_user_id: &str,
provider_username: Option<&str>,
provider_email: Option<&str>,
extra_data: Option<Value>,
touched_at: DateTime<Utc>,
) -> Result<bool, crate::DataLayerError>;
async fn create_oauth_auth_user(
&self,
email: Option<String>,
username: String,
created_at: DateTime<Utc>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn find_oauth_link_owner(
&self,
provider_type: &str,
provider_user_id: &str,
) -> Result<Option<String>, crate::DataLayerError>;
async fn has_user_oauth_provider_link(
&self,
user_id: &str,
provider_type: &str,
) -> Result<bool, crate::DataLayerError>;
async fn count_user_oauth_links(&self, user_id: &str) -> Result<u64, crate::DataLayerError>;
#[allow(clippy::too_many_arguments)]
async fn upsert_user_oauth_link(
&self,
user_id: &str,
provider_type: &str,
provider_user_id: &str,
provider_username: Option<&str>,
provider_email: Option<&str>,
extra_data: Option<Value>,
linked_at: DateTime<Utc>,
) -> Result<(), crate::DataLayerError>;
async fn delete_user_oauth_link(
&self,
user_id: &str,
provider_type: &str,
) -> Result<bool, crate::DataLayerError>;
async fn get_or_create_ldap_auth_user(
&self,
email: String,
username: String,
ldap_dn: Option<String>,
ldap_username: Option<String>,
logged_in_at: DateTime<Utc>,
) -> Result<Option<LdapAuthUserProvisioningOutcome>, crate::DataLayerError>;
async fn touch_auth_user_last_login(
&self,
user_id: &str,
logged_in_at: DateTime<Utc>,
) -> Result<bool, crate::DataLayerError>;
async fn update_local_auth_user_profile(
&self,
user_id: &str,
email: Option<String>,
username: Option<String>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn update_local_auth_user_password_hash(
&self,
user_id: &str,
password_hash: String,
updated_at: DateTime<Utc>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
#[allow(clippy::too_many_arguments)]
async fn update_local_auth_user_admin_fields(
&self,
user_id: &str,
role: Option<String>,
allowed_providers_present: bool,
allowed_providers: Option<Vec<String>>,
allowed_api_formats_present: bool,
allowed_api_formats: Option<Vec<String>>,
allowed_models_present: bool,
allowed_models: Option<Vec<String>>,
rate_limit: Option<i32>,
is_active: Option<bool>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn update_user_model_capability_settings(
&self,
user_id: &str,
settings: Option<Value>,
) -> Result<Option<Value>, crate::DataLayerError>;
async fn create_local_auth_user(
&self,
email: Option<String>,
email_verified: bool,
username: String,
password_hash: String,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
#[allow(clippy::too_many_arguments)]
async fn create_local_auth_user_with_settings(
&self,
email: Option<String>,
email_verified: bool,
username: String,
password_hash: String,
role: String,
allowed_providers: Option<Vec<String>>,
allowed_api_formats: Option<Vec<String>>,
allowed_models: Option<Vec<String>>,
rate_limit: Option<i32>,
) -> Result<Option<StoredUserAuthRecord>, crate::DataLayerError>;
async fn delete_local_auth_user(&self, user_id: &str) -> Result<bool, crate::DataLayerError>;
async fn read_user_preferences(
&self,
user_id: &str,
) -> Result<Option<StoredUserPreferenceRecord>, crate::DataLayerError>;
async fn write_user_preferences(
&self,
preferences: &StoredUserPreferenceRecord,
) -> Result<Option<StoredUserPreferenceRecord>, crate::DataLayerError>;
async fn find_user_session(
&self,
user_id: &str,
session_id: &str,
) -> Result<Option<StoredUserSessionRecord>, crate::DataLayerError>;
async fn list_user_sessions(
&self,
user_id: &str,
) -> Result<Vec<StoredUserSessionRecord>, crate::DataLayerError>;
async fn create_user_session(
&self,
session: &StoredUserSessionRecord,
) -> Result<Option<StoredUserSessionRecord>, crate::DataLayerError>;
async fn touch_user_session(
&self,
user_id: &str,
session_id: &str,
touched_at: DateTime<Utc>,
ip_address: Option<&str>,
user_agent: Option<&str>,
) -> Result<bool, crate::DataLayerError>;
async fn update_user_session_device_label(
&self,
user_id: &str,
session_id: &str,
device_label: &str,
updated_at: DateTime<Utc>,
) -> Result<bool, crate::DataLayerError>;
#[allow(clippy::too_many_arguments)]
async fn rotate_user_session_refresh_token(
&self,
user_id: &str,
session_id: &str,
previous_refresh_token_hash: &str,
next_refresh_token_hash: &str,
rotated_at: DateTime<Utc>,
expires_at: DateTime<Utc>,
ip_address: Option<&str>,
user_agent: Option<&str>,
) -> Result<bool, crate::DataLayerError>;
async fn revoke_user_session(
&self,
user_id: &str,
session_id: &str,
revoked_at: DateTime<Utc>,
reason: &str,
) -> Result<bool, crate::DataLayerError>;
async fn revoke_all_user_sessions(
&self,
user_id: &str,
revoked_at: DateTime<Utc>,
reason: &str,
) -> Result<u64, crate::DataLayerError>;
async fn count_active_admin_users(&self) -> Result<u64, crate::DataLayerError>;
async fn count_active_local_admin_users_with_valid_password(
&self,
) -> Result<u64, crate::DataLayerError>;
}
fn normalize_optional_json(value: Option<Value>) -> Option<Value> {

View File

@@ -1,5 +1,7 @@
mod memory;
mod sql;
mod mysql;
mod postgres;
mod sqlite;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::video_tasks::{
@@ -8,4 +10,6 @@ pub(crate) use aether_data_contracts::repository::video_tasks::{
VideoTaskStatusCount, VideoTaskWriteRepository,
};
pub use memory::InMemoryVideoTaskRepository;
pub use sql::{SqlxVideoTaskReadRepository, SqlxVideoTaskRepository};
pub use mysql::MysqlVideoTaskRepository;
pub use postgres::{SqlxVideoTaskReadRepository, SqlxVideoTaskRepository};
pub use sqlite::SqliteVideoTaskRepository;

View File

@@ -0,0 +1,689 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::{
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskStatus, VideoTaskStatusCount,
VideoTaskWriteRepository,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const VIDEO_TASK_COLUMNS: &str = r#"
SELECT
id,
short_id,
request_id,
user_id,
api_key_id,
username,
api_key_name,
external_task_id,
provider_id,
endpoint_id,
key_id,
client_api_format,
provider_api_format,
format_converted,
model,
prompt,
original_request_body,
duration_seconds,
resolution,
aspect_ratio,
size,
status,
progress_percent,
progress_message,
retry_count,
poll_interval_seconds,
next_poll_at AS next_poll_at_unix_secs,
poll_count,
max_poll_count,
created_at AS created_at_unix_ms,
submitted_at AS submitted_at_unix_secs,
completed_at AS completed_at_unix_secs,
updated_at AS updated_at_unix_secs,
error_code,
error_message,
video_url,
request_metadata
FROM video_tasks
"#;
#[derive(Debug, Clone)]
pub struct MysqlVideoTaskRepository {
pool: MysqlPool,
}
impl MysqlVideoTaskRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
async fn find_by_id(&self, id: &str) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!("{VIDEO_TASK_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_video_task_row).transpose()
}
async fn find_by_short_id(
&self,
short_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!("{VIDEO_TASK_COLUMNS} WHERE short_id = ? LIMIT 1"))
.bind(short_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_video_task_row).transpose()
}
async fn find_by_user_external(
&self,
user_id: &str,
external_task_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE user_id = ? AND external_task_id = ? LIMIT 1"
))
.bind(user_id)
.bind(external_task_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_video_task_row).transpose()
}
async fn reload_ids(&self, ids: &[String]) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(VIDEO_TASK_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for id in ids {
separated.push_bind(id);
}
}
builder.push(")");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let mut tasks = rows
.iter()
.map(map_video_task_row)
.collect::<Result<Vec<_>, _>>()?;
tasks.sort_by(|left, right| {
left.next_poll_at_unix_secs
.cmp(&right.next_poll_at_unix_secs)
.then_with(|| left.updated_at_unix_secs.cmp(&right.updated_at_unix_secs))
});
Ok(tasks)
}
}
#[async_trait]
impl VideoTaskReadRepository for MysqlVideoTaskRepository {
async fn find(
&self,
key: VideoTaskLookupKey<'_>,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
match key {
VideoTaskLookupKey::Id(id) => self.find_by_id(id).await,
VideoTaskLookupKey::ShortId(short_id) => self.find_by_short_id(short_id).await,
VideoTaskLookupKey::UserExternal {
user_id,
external_task_id,
} => self.find_by_user_external(user_id, external_task_id).await,
}
}
async fn list_active(&self, limit: usize) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE status IN ('pending', 'submitted', 'queued', 'processing') ORDER BY updated_at DESC LIMIT ?"
))
.bind(limit_i64(limit, "active video task limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_video_task_row).collect()
}
async fn list_due(
&self,
now_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE status IN ('submitted', 'queued', 'processing') AND next_poll_at IS NOT NULL AND next_poll_at <= ? AND poll_count < max_poll_count ORDER BY next_poll_at ASC, updated_at ASC LIMIT ?"
))
.bind(u64_to_i64(now_unix_secs, "video task now")?)
.bind(limit_i64(limit, "due video task limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_video_task_row).collect()
}
async fn list_page(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<MySql>::new(VIDEO_TASK_COLUMNS);
push_filter(&mut builder, filter, None);
builder
.push(" ORDER BY created_at DESC, updated_at DESC LIMIT ")
.push_bind(limit_i64(limit, "video task page limit")?)
.push(" OFFSET ")
.push_bind(limit_i64(offset, "video task page offset")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_video_task_row).collect()
}
async fn list_page_summary(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
self.list_page(filter, offset, limit).await
}
async fn count(&self, filter: &VideoTaskQueryFilter) -> Result<u64, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new("SELECT COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
count_query(builder, &self.pool).await
}
async fn count_by_status(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<Vec<VideoTaskStatusCount>, DataLayerError> {
let mut builder =
QueryBuilder::<MySql>::new("SELECT status, COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
builder.push(" GROUP BY status ORDER BY status ASC");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter()
.map(|row| {
Ok(VideoTaskStatusCount {
status: VideoTaskStatus::from_database(
row.try_get::<String, _>("status").map_sql_err()?.as_str(),
)?,
count: count_value(row.try_get("total").map_sql_err()?)?,
})
})
.collect()
}
async fn count_distinct_users(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<u64, DataLayerError> {
let mut builder =
QueryBuilder::<MySql>::new("SELECT COUNT(DISTINCT user_id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
push_clause(&mut builder, "user_id IS NOT NULL");
push_clause(&mut builder, "user_id <> ''");
count_query(builder, &self.pool).await
}
async fn top_models(
&self,
filter: &VideoTaskQueryFilter,
limit: usize,
) -> Result<Vec<VideoTaskModelCount>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut builder =
QueryBuilder::<MySql>::new("SELECT model, COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
push_clause(&mut builder, "model IS NOT NULL");
push_clause(&mut builder, "model <> ''");
builder
.push(" GROUP BY model ORDER BY total DESC, model ASC LIMIT ")
.push_bind(limit_i64(limit, "video task top models limit")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter()
.map(|row| {
Ok(VideoTaskModelCount {
model: row.try_get("model").map_sql_err()?,
count: count_value(row.try_get("total").map_sql_err()?)?,
})
})
.collect()
}
async fn count_created_since(
&self,
filter: &VideoTaskQueryFilter,
created_since_unix_secs: u64,
) -> Result<u64, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new("SELECT COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, Some(created_since_unix_secs));
count_query(builder, &self.pool).await
}
}
#[async_trait]
impl VideoTaskWriteRepository for MysqlVideoTaskRepository {
async fn upsert(&self, task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
let id = task.id.clone();
bind_task(sqlx::query(UPSERT_SQL), task, true, false)?
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_by_id(&id)
.await?
.ok_or_else(|| DataLayerError::UnexpectedValue("upserted video task missing".into()))
}
async fn update_if_active(
&self,
task: UpsertVideoTask,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let id = task.id.clone();
let rows_affected = bind_task(sqlx::query(UPDATE_IF_ACTIVE_SQL), task, false, true)?
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
if rows_affected == 0 {
return Ok(None);
}
self.find_by_id(&id).await
}
async fn claim_due(
&self,
now_unix_secs: u64,
claim_until_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let due = self.list_due(now_unix_secs, limit).await?;
let ids = due.iter().map(|task| task.id.clone()).collect::<Vec<_>>();
for id in &ids {
sqlx::query(
"UPDATE video_tasks SET next_poll_at = ?, updated_at = GREATEST(updated_at, ?) WHERE id = ?",
)
.bind(u64_to_i64(claim_until_unix_secs, "video task claim_until")?)
.bind(u64_to_i64(now_unix_secs, "video task now")?)
.bind(id)
.execute(&self.pool)
.await
.map_sql_err()?;
}
self.reload_ids(&ids).await
}
}
const UPSERT_SQL: &str = r#"
INSERT INTO video_tasks (
id, short_id, request_id, user_id, api_key_id, username, api_key_name,
external_task_id, provider_id, endpoint_id, key_id, client_api_format,
provider_api_format, format_converted, model, prompt, original_request_body,
duration_seconds, resolution, aspect_ratio, size, status, progress_percent,
progress_message, retry_count, poll_interval_seconds, next_poll_at, poll_count,
max_poll_count, video_url, error_code, error_message, request_metadata,
created_at, submitted_at, completed_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
short_id = VALUES(short_id),
request_id = VALUES(request_id),
user_id = VALUES(user_id),
api_key_id = VALUES(api_key_id),
username = VALUES(username),
api_key_name = VALUES(api_key_name),
external_task_id = VALUES(external_task_id),
provider_id = VALUES(provider_id),
endpoint_id = VALUES(endpoint_id),
key_id = VALUES(key_id),
client_api_format = VALUES(client_api_format),
provider_api_format = VALUES(provider_api_format),
format_converted = VALUES(format_converted),
model = VALUES(model),
prompt = VALUES(prompt),
original_request_body = VALUES(original_request_body),
duration_seconds = VALUES(duration_seconds),
resolution = VALUES(resolution),
aspect_ratio = VALUES(aspect_ratio),
size = VALUES(size),
status = VALUES(status),
progress_percent = VALUES(progress_percent),
progress_message = VALUES(progress_message),
retry_count = VALUES(retry_count),
poll_interval_seconds = VALUES(poll_interval_seconds),
next_poll_at = VALUES(next_poll_at),
poll_count = VALUES(poll_count),
max_poll_count = VALUES(max_poll_count),
video_url = VALUES(video_url),
error_code = VALUES(error_code),
error_message = VALUES(error_message),
request_metadata = VALUES(request_metadata),
created_at = VALUES(created_at),
submitted_at = VALUES(submitted_at),
completed_at = VALUES(completed_at),
updated_at = VALUES(updated_at)
"#;
const UPDATE_IF_ACTIVE_SQL: &str = r#"
UPDATE video_tasks SET
short_id = ?,
request_id = ?,
user_id = ?,
api_key_id = ?,
username = ?,
api_key_name = ?,
external_task_id = ?,
provider_id = ?,
endpoint_id = ?,
key_id = ?,
client_api_format = ?,
provider_api_format = ?,
format_converted = ?,
model = ?,
prompt = ?,
original_request_body = ?,
duration_seconds = ?,
resolution = ?,
aspect_ratio = ?,
size = ?,
status = ?,
progress_percent = ?,
progress_message = ?,
retry_count = ?,
poll_interval_seconds = ?,
next_poll_at = ?,
poll_count = ?,
max_poll_count = ?,
video_url = ?,
error_code = ?,
error_message = ?,
request_metadata = ?,
created_at = ?,
submitted_at = ?,
completed_at = ?,
updated_at = ?
WHERE id = ?
AND status IN ('pending', 'submitted', 'queued', 'processing')
"#;
fn bind_task<'q>(
query: sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>,
task: UpsertVideoTask,
include_insert_id: bool,
include_update_id: bool,
) -> Result<sqlx::query::Query<'q, MySql, sqlx::mysql::MySqlArguments>, DataLayerError> {
let original_request_body = json_to_string(&task.original_request_body)?;
let request_metadata = json_to_string(&task.request_metadata)?;
let query = if include_insert_id {
query.bind(task.id.clone())
} else {
query
};
let bound = query
.bind(task.short_id)
.bind(task.request_id)
.bind(task.user_id)
.bind(task.api_key_id)
.bind(task.username)
.bind(task.api_key_name)
.bind(task.external_task_id)
.bind(task.provider_id)
.bind(task.endpoint_id)
.bind(task.key_id)
.bind(task.client_api_format)
.bind(task.provider_api_format)
.bind(task.format_converted)
.bind(task.model)
.bind(task.prompt)
.bind(original_request_body)
.bind(optional_u32_to_i32(
task.duration_seconds,
"video task duration_seconds",
)?)
.bind(task.resolution)
.bind(task.aspect_ratio)
.bind(task.size)
.bind(status_to_database(task.status))
.bind(i32::from(task.progress_percent))
.bind(task.progress_message)
.bind(u32_to_i32(task.retry_count, "video task retry_count")?)
.bind(u32_to_i32(
task.poll_interval_seconds,
"video task poll_interval_seconds",
)?)
.bind(optional_u64_to_i64(
task.next_poll_at_unix_secs,
"video task next_poll_at",
)?)
.bind(u32_to_i32(task.poll_count, "video task poll_count")?)
.bind(u32_to_i32(
task.max_poll_count,
"video task max_poll_count",
)?)
.bind(task.video_url)
.bind(task.error_code)
.bind(task.error_message)
.bind(request_metadata)
.bind(u64_to_i64(
task.created_at_unix_ms,
"video task created_at",
)?)
.bind(optional_u64_to_i64(
task.submitted_at_unix_secs,
"video task submitted_at",
)?)
.bind(optional_u64_to_i64(
task.completed_at_unix_secs,
"video task completed_at",
)?)
.bind(u64_to_i64(
task.updated_at_unix_secs,
"video task updated_at",
)?);
if include_update_id {
Ok(bound.bind(task.id))
} else {
Ok(bound)
}
}
fn push_filter<'args>(
builder: &mut QueryBuilder<'args, MySql>,
filter: &'args VideoTaskQueryFilter,
created_since_unix_secs: Option<u64>,
) {
if let Some(user_id) = filter.user_id.as_deref() {
push_clause(builder, "user_id = ");
builder.push_bind(user_id);
}
if let Some(status) = filter.status {
push_clause(builder, "status = ");
builder.push_bind(status_to_database(status));
}
if let Some(model_substring) = filter.model_substring.as_deref() {
push_clause(builder, "LOWER(model) LIKE ");
builder.push_bind(format!(
"%{}%",
escape_like_pattern(&model_substring.trim().to_ascii_lowercase())
));
builder.push(" ESCAPE '\\'");
}
if let Some(client_api_format) = filter.client_api_format.as_deref() {
push_clause(builder, "client_api_format = ");
builder.push_bind(client_api_format);
}
if let Some(created_since_unix_secs) = created_since_unix_secs {
push_clause(builder, "created_at >= ");
builder.push_bind(created_since_unix_secs as i64);
}
}
fn push_clause<'args>(builder: &mut QueryBuilder<'args, MySql>, clause: &str) {
let sql = builder.sql();
if sql.contains(" WHERE ") || sql.contains("\nWHERE ") {
builder.push(" AND ");
} else {
builder.push(" WHERE ");
}
builder.push(clause);
}
async fn count_query(
mut builder: QueryBuilder<'_, MySql>,
pool: &MysqlPool,
) -> Result<u64, DataLayerError> {
let row = builder.build().fetch_one(pool).await.map_sql_err()?;
count_value(row.try_get("total").map_sql_err()?)
}
fn count_value(value: i64) -> Result<u64, DataLayerError> {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("invalid video task count result: {value}"))
})
}
fn map_video_task_row(row: &MySqlRow) -> Result<StoredVideoTask, DataLayerError> {
StoredVideoTask::new(
row.try_get("id").map_sql_err()?,
row.try_get("short_id").map_sql_err()?,
row.try_get("request_id").map_sql_err()?,
row.try_get("user_id").map_sql_err()?,
row.try_get("api_key_id").map_sql_err()?,
row.try_get("username").map_sql_err()?,
row.try_get("api_key_name").map_sql_err()?,
row.try_get("external_task_id").map_sql_err()?,
row.try_get("provider_id").map_sql_err()?,
row.try_get("endpoint_id").map_sql_err()?,
row.try_get("key_id").map_sql_err()?,
row.try_get("client_api_format").map_sql_err()?,
row.try_get("provider_api_format").map_sql_err()?,
row.try_get("format_converted").map_sql_err()?,
row.try_get("model").map_sql_err()?,
row.try_get("prompt").map_sql_err()?,
parse_json(row.try_get("original_request_body").ok().flatten())?,
row.try_get("duration_seconds").map_sql_err()?,
row.try_get("resolution").map_sql_err()?,
row.try_get("aspect_ratio").map_sql_err()?,
row.try_get("size").map_sql_err()?,
VideoTaskStatus::from_database(row.try_get::<String, _>("status").map_sql_err()?.as_str())?,
row.try_get("progress_percent").map_sql_err()?,
row.try_get("progress_message").map_sql_err()?,
row.try_get("retry_count").map_sql_err()?,
row.try_get("poll_interval_seconds").map_sql_err()?,
row.try_get("next_poll_at_unix_secs").map_sql_err()?,
row.try_get("poll_count").map_sql_err()?,
row.try_get("max_poll_count").map_sql_err()?,
row.try_get("created_at_unix_ms").map_sql_err()?,
row.try_get("submitted_at_unix_secs").map_sql_err()?,
row.try_get("completed_at_unix_secs").map_sql_err()?,
row.try_get("updated_at_unix_secs").map_sql_err()?,
row.try_get("error_code").map_sql_err()?,
row.try_get("error_message").map_sql_err()?,
row.try_get("video_url").map_sql_err()?,
parse_json(row.try_get("request_metadata").ok().flatten())?,
)
}
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!("video task JSON field is invalid: {err}"))
})
})
.transpose()
}
fn json_to_string(value: &Option<serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.as_ref()
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"video task JSON field is unserializable: {err}"
))
})
})
.transpose()
}
fn status_to_database(status: VideoTaskStatus) -> &'static str {
match status {
VideoTaskStatus::Pending => "pending",
VideoTaskStatus::Submitted => "submitted",
VideoTaskStatus::Queued => "queued",
VideoTaskStatus::Processing => "processing",
VideoTaskStatus::Completed => "completed",
VideoTaskStatus::Failed => "failed",
VideoTaskStatus::Cancelled => "cancelled",
VideoTaskStatus::Expired => "expired",
VideoTaskStatus::Deleted => "deleted",
}
}
fn escape_like_pattern(value: &str) -> String {
value
.replace('\\', "\\\\")
.replace('%', "\\%")
.replace('_', "\\_")
}
fn limit_i64(value: usize, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::UnexpectedValue(format!("invalid {name}: {value}")))
}
fn u64_to_i64(value: u64, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| DataLayerError::UnexpectedValue(format!("{name} overflow")))
}
fn optional_u64_to_i64(value: Option<u64>, name: &str) -> Result<Option<i64>, DataLayerError> {
value.map(|value| u64_to_i64(value, name)).transpose()
}
fn u32_to_i32(value: u32, name: &str) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| DataLayerError::UnexpectedValue(format!("{name} overflow")))
}
fn optional_u32_to_i32(value: Option<u32>, name: &str) -> Result<Option<i32>, DataLayerError> {
value.map(|value| u32_to_i32(value, name)).transpose()
}
#[cfg(test)]
mod tests {
use super::MysqlVideoTaskRepository;
#[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 = MysqlVideoTaskRepository::new(pool);
}
}

View File

@@ -1067,7 +1067,7 @@ fn map_video_task_row(row: &PgRow) -> Result<StoredVideoTask, DataLayerError> {
#[cfg(test)]
mod tests {
use super::SqlxVideoTaskRepository;
use crate::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
use crate::repository::video_tasks::{
UpsertVideoTask, VideoTaskLookupKey, VideoTaskQueryFilter, VideoTaskReadRepository,
VideoTaskStatus, VideoTaskWriteRepository,

View File

@@ -0,0 +1,824 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::{
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskStatus, VideoTaskStatusCount,
VideoTaskWriteRepository,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const VIDEO_TASK_COLUMNS: &str = r#"
SELECT
id,
short_id,
request_id,
user_id,
api_key_id,
username,
api_key_name,
external_task_id,
provider_id,
endpoint_id,
key_id,
client_api_format,
provider_api_format,
format_converted,
model,
prompt,
original_request_body,
duration_seconds,
resolution,
aspect_ratio,
size,
status,
progress_percent,
progress_message,
retry_count,
poll_interval_seconds,
next_poll_at AS next_poll_at_unix_secs,
poll_count,
max_poll_count,
created_at AS created_at_unix_ms,
submitted_at AS submitted_at_unix_secs,
completed_at AS completed_at_unix_secs,
updated_at AS updated_at_unix_secs,
error_code,
error_message,
video_url,
request_metadata
FROM video_tasks
"#;
#[derive(Debug, Clone)]
pub struct SqliteVideoTaskRepository {
pool: SqlitePool,
}
impl SqliteVideoTaskRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
async fn find_by_id(&self, id: &str) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!("{VIDEO_TASK_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_video_task_row).transpose()
}
async fn find_by_short_id(
&self,
short_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!("{VIDEO_TASK_COLUMNS} WHERE short_id = ? LIMIT 1"))
.bind(short_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_video_task_row).transpose()
}
async fn find_by_user_external(
&self,
user_id: &str,
external_task_id: &str,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let row = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE user_id = ? AND external_task_id = ? LIMIT 1"
))
.bind(user_id)
.bind(external_task_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_video_task_row).transpose()
}
async fn reload_ids(&self, ids: &[String]) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if ids.is_empty() {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(VIDEO_TASK_COLUMNS);
builder.push(" WHERE id IN (");
{
let mut separated = builder.separated(", ");
for id in ids {
separated.push_bind(id);
}
}
builder.push(")");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let mut tasks = rows
.iter()
.map(map_video_task_row)
.collect::<Result<Vec<_>, _>>()?;
tasks.sort_by(|left, right| {
left.next_poll_at_unix_secs
.cmp(&right.next_poll_at_unix_secs)
.then_with(|| left.updated_at_unix_secs.cmp(&right.updated_at_unix_secs))
});
Ok(tasks)
}
}
#[async_trait]
impl VideoTaskReadRepository for SqliteVideoTaskRepository {
async fn find(
&self,
key: VideoTaskLookupKey<'_>,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
match key {
VideoTaskLookupKey::Id(id) => self.find_by_id(id).await,
VideoTaskLookupKey::ShortId(short_id) => self.find_by_short_id(short_id).await,
VideoTaskLookupKey::UserExternal {
user_id,
external_task_id,
} => self.find_by_user_external(user_id, external_task_id).await,
}
}
async fn list_active(&self, limit: usize) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE status IN ('pending', 'submitted', 'queued', 'processing') ORDER BY updated_at DESC LIMIT ?"
))
.bind(limit_i64(limit, "active video task limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_video_task_row).collect()
}
async fn list_due(
&self,
now_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let rows = sqlx::query(&format!(
"{VIDEO_TASK_COLUMNS} WHERE status IN ('submitted', 'queued', 'processing') AND next_poll_at IS NOT NULL AND next_poll_at <= ? AND poll_count < max_poll_count ORDER BY next_poll_at ASC, updated_at ASC LIMIT ?"
))
.bind(u64_to_i64(now_unix_secs, "video task now")?)
.bind(limit_i64(limit, "due video task limit")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_video_task_row).collect()
}
async fn list_page(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut builder = QueryBuilder::<Sqlite>::new(VIDEO_TASK_COLUMNS);
push_filter(&mut builder, filter, None);
builder
.push(" ORDER BY created_at DESC, updated_at DESC LIMIT ")
.push_bind(limit_i64(limit, "video task page limit")?)
.push(" OFFSET ")
.push_bind(limit_i64(offset, "video task page offset")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_video_task_row).collect()
}
async fn list_page_summary(
&self,
filter: &VideoTaskQueryFilter,
offset: usize,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
self.list_page(filter, offset, limit).await
}
async fn count(&self, filter: &VideoTaskQueryFilter) -> Result<u64, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new("SELECT COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
count_query(builder, &self.pool).await
}
async fn count_by_status(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<Vec<VideoTaskStatusCount>, DataLayerError> {
let mut builder =
QueryBuilder::<Sqlite>::new("SELECT status, COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
builder.push(" GROUP BY status ORDER BY status ASC");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter()
.map(|row| {
Ok(VideoTaskStatusCount {
status: VideoTaskStatus::from_database(
row.try_get::<String, _>("status").map_sql_err()?.as_str(),
)?,
count: count_value(row.try_get("total").map_sql_err()?)?,
})
})
.collect()
}
async fn count_distinct_users(
&self,
filter: &VideoTaskQueryFilter,
) -> Result<u64, DataLayerError> {
let mut builder =
QueryBuilder::<Sqlite>::new("SELECT COUNT(DISTINCT user_id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
push_clause(&mut builder, "user_id IS NOT NULL");
push_clause(&mut builder, "user_id <> ''");
count_query(builder, &self.pool).await
}
async fn top_models(
&self,
filter: &VideoTaskQueryFilter,
limit: usize,
) -> Result<Vec<VideoTaskModelCount>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let mut builder =
QueryBuilder::<Sqlite>::new("SELECT model, COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, None);
push_clause(&mut builder, "model IS NOT NULL");
push_clause(&mut builder, "model <> ''");
builder
.push(" GROUP BY model ORDER BY total DESC, model ASC LIMIT ")
.push_bind(limit_i64(limit, "video task top models limit")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter()
.map(|row| {
Ok(VideoTaskModelCount {
model: row.try_get("model").map_sql_err()?,
count: count_value(row.try_get("total").map_sql_err()?)?,
})
})
.collect()
}
async fn count_created_since(
&self,
filter: &VideoTaskQueryFilter,
created_since_unix_secs: u64,
) -> Result<u64, DataLayerError> {
let mut builder = QueryBuilder::<Sqlite>::new("SELECT COUNT(id) AS total FROM video_tasks");
push_filter(&mut builder, filter, Some(created_since_unix_secs));
count_query(builder, &self.pool).await
}
}
#[async_trait]
impl VideoTaskWriteRepository for SqliteVideoTaskRepository {
async fn upsert(&self, task: UpsertVideoTask) -> Result<StoredVideoTask, DataLayerError> {
let id = task.id.clone();
bind_task(sqlx::query(UPSERT_SQL), task, true, false)?
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_by_id(&id)
.await?
.ok_or_else(|| DataLayerError::UnexpectedValue("upserted video task missing".into()))
}
async fn update_if_active(
&self,
task: UpsertVideoTask,
) -> Result<Option<StoredVideoTask>, DataLayerError> {
let id = task.id.clone();
let rows_affected = bind_task(sqlx::query(UPDATE_IF_ACTIVE_SQL), task, false, true)?
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
if rows_affected == 0 {
return Ok(None);
}
self.find_by_id(&id).await
}
async fn claim_due(
&self,
now_unix_secs: u64,
claim_until_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredVideoTask>, DataLayerError> {
if limit == 0 {
return Ok(Vec::new());
}
let due = self.list_due(now_unix_secs, limit).await?;
let ids = due.iter().map(|task| task.id.clone()).collect::<Vec<_>>();
for id in &ids {
sqlx::query(
"UPDATE video_tasks SET next_poll_at = ?, updated_at = MAX(updated_at, ?) WHERE id = ?",
)
.bind(u64_to_i64(claim_until_unix_secs, "video task claim_until")?)
.bind(u64_to_i64(now_unix_secs, "video task now")?)
.bind(id)
.execute(&self.pool)
.await
.map_sql_err()?;
}
self.reload_ids(&ids).await
}
}
const UPSERT_SQL: &str = r#"
INSERT INTO video_tasks (
id, short_id, request_id, user_id, api_key_id, username, api_key_name,
external_task_id, provider_id, endpoint_id, key_id, client_api_format,
provider_api_format, format_converted, model, prompt, original_request_body,
duration_seconds, resolution, aspect_ratio, size, status, progress_percent,
progress_message, retry_count, poll_interval_seconds, next_poll_at, poll_count,
max_poll_count, video_url, error_code, error_message, request_metadata,
created_at, submitted_at, completed_at, updated_at
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET
short_id = excluded.short_id,
request_id = excluded.request_id,
user_id = excluded.user_id,
api_key_id = excluded.api_key_id,
username = excluded.username,
api_key_name = excluded.api_key_name,
external_task_id = excluded.external_task_id,
provider_id = excluded.provider_id,
endpoint_id = excluded.endpoint_id,
key_id = excluded.key_id,
client_api_format = excluded.client_api_format,
provider_api_format = excluded.provider_api_format,
format_converted = excluded.format_converted,
model = excluded.model,
prompt = excluded.prompt,
original_request_body = excluded.original_request_body,
duration_seconds = excluded.duration_seconds,
resolution = excluded.resolution,
aspect_ratio = excluded.aspect_ratio,
size = excluded.size,
status = excluded.status,
progress_percent = excluded.progress_percent,
progress_message = excluded.progress_message,
retry_count = excluded.retry_count,
poll_interval_seconds = excluded.poll_interval_seconds,
next_poll_at = excluded.next_poll_at,
poll_count = excluded.poll_count,
max_poll_count = excluded.max_poll_count,
video_url = excluded.video_url,
error_code = excluded.error_code,
error_message = excluded.error_message,
request_metadata = excluded.request_metadata,
created_at = excluded.created_at,
submitted_at = excluded.submitted_at,
completed_at = excluded.completed_at,
updated_at = excluded.updated_at
"#;
const UPDATE_IF_ACTIVE_SQL: &str = r#"
UPDATE video_tasks SET
short_id = ?,
request_id = ?,
user_id = ?,
api_key_id = ?,
username = ?,
api_key_name = ?,
external_task_id = ?,
provider_id = ?,
endpoint_id = ?,
key_id = ?,
client_api_format = ?,
provider_api_format = ?,
format_converted = ?,
model = ?,
prompt = ?,
original_request_body = ?,
duration_seconds = ?,
resolution = ?,
aspect_ratio = ?,
size = ?,
status = ?,
progress_percent = ?,
progress_message = ?,
retry_count = ?,
poll_interval_seconds = ?,
next_poll_at = ?,
poll_count = ?,
max_poll_count = ?,
video_url = ?,
error_code = ?,
error_message = ?,
request_metadata = ?,
created_at = ?,
submitted_at = ?,
completed_at = ?,
updated_at = ?
WHERE id = ?
AND status IN ('pending', 'submitted', 'queued', 'processing')
"#;
fn bind_task<'q>(
query: sqlx::query::Query<'q, Sqlite, sqlx::sqlite::SqliteArguments<'q>>,
task: UpsertVideoTask,
include_insert_id: bool,
include_update_id: bool,
) -> Result<sqlx::query::Query<'q, Sqlite, sqlx::sqlite::SqliteArguments<'q>>, DataLayerError> {
let original_request_body = json_to_string(&task.original_request_body)?;
let request_metadata = json_to_string(&task.request_metadata)?;
let query = if include_insert_id {
query.bind(task.id.clone())
} else {
query
};
let bound = query
.bind(task.short_id)
.bind(task.request_id)
.bind(task.user_id)
.bind(task.api_key_id)
.bind(task.username)
.bind(task.api_key_name)
.bind(task.external_task_id)
.bind(task.provider_id)
.bind(task.endpoint_id)
.bind(task.key_id)
.bind(task.client_api_format)
.bind(task.provider_api_format)
.bind(task.format_converted)
.bind(task.model)
.bind(task.prompt)
.bind(original_request_body)
.bind(optional_u32_to_i32(
task.duration_seconds,
"video task duration_seconds",
)?)
.bind(task.resolution)
.bind(task.aspect_ratio)
.bind(task.size)
.bind(status_to_database(task.status))
.bind(i32::from(task.progress_percent))
.bind(task.progress_message)
.bind(u32_to_i32(task.retry_count, "video task retry_count")?)
.bind(u32_to_i32(
task.poll_interval_seconds,
"video task poll_interval_seconds",
)?)
.bind(optional_u64_to_i64(
task.next_poll_at_unix_secs,
"video task next_poll_at",
)?)
.bind(u32_to_i32(task.poll_count, "video task poll_count")?)
.bind(u32_to_i32(
task.max_poll_count,
"video task max_poll_count",
)?)
.bind(task.video_url)
.bind(task.error_code)
.bind(task.error_message)
.bind(request_metadata)
.bind(u64_to_i64(
task.created_at_unix_ms,
"video task created_at",
)?)
.bind(optional_u64_to_i64(
task.submitted_at_unix_secs,
"video task submitted_at",
)?)
.bind(optional_u64_to_i64(
task.completed_at_unix_secs,
"video task completed_at",
)?)
.bind(u64_to_i64(
task.updated_at_unix_secs,
"video task updated_at",
)?);
if include_update_id {
Ok(bound.bind(task.id))
} else {
Ok(bound)
}
}
fn push_filter<'args>(
builder: &mut QueryBuilder<'args, Sqlite>,
filter: &'args VideoTaskQueryFilter,
created_since_unix_secs: Option<u64>,
) {
if let Some(user_id) = filter.user_id.as_deref() {
push_clause(builder, "user_id = ");
builder.push_bind(user_id);
}
if let Some(status) = filter.status {
push_clause(builder, "status = ");
builder.push_bind(status_to_database(status));
}
if let Some(model_substring) = filter.model_substring.as_deref() {
push_clause(builder, "LOWER(model) LIKE ");
builder.push_bind(format!(
"%{}%",
escape_like_pattern(&model_substring.trim().to_ascii_lowercase())
));
builder.push(" ESCAPE '\\'");
}
if let Some(client_api_format) = filter.client_api_format.as_deref() {
push_clause(builder, "client_api_format = ");
builder.push_bind(client_api_format);
}
if let Some(created_since_unix_secs) = created_since_unix_secs {
push_clause(builder, "created_at >= ");
builder.push_bind(created_since_unix_secs as i64);
}
}
fn push_clause<'args>(builder: &mut QueryBuilder<'args, Sqlite>, clause: &str) {
let sql = builder.sql();
if sql.contains(" WHERE ") || sql.contains("\nWHERE ") {
builder.push(" AND ");
} else {
builder.push(" WHERE ");
}
builder.push(clause);
}
async fn count_query(
mut builder: QueryBuilder<'_, Sqlite>,
pool: &SqlitePool,
) -> Result<u64, DataLayerError> {
let row = builder.build().fetch_one(pool).await.map_sql_err()?;
count_value(row.try_get("total").map_sql_err()?)
}
fn count_value(value: i64) -> Result<u64, DataLayerError> {
u64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("invalid video task count result: {value}"))
})
}
fn map_video_task_row(row: &SqliteRow) -> Result<StoredVideoTask, DataLayerError> {
StoredVideoTask::new(
row.try_get("id").map_sql_err()?,
row.try_get("short_id").map_sql_err()?,
row.try_get("request_id").map_sql_err()?,
row.try_get("user_id").map_sql_err()?,
row.try_get("api_key_id").map_sql_err()?,
row.try_get("username").map_sql_err()?,
row.try_get("api_key_name").map_sql_err()?,
row.try_get("external_task_id").map_sql_err()?,
row.try_get("provider_id").map_sql_err()?,
row.try_get("endpoint_id").map_sql_err()?,
row.try_get("key_id").map_sql_err()?,
row.try_get("client_api_format").map_sql_err()?,
row.try_get("provider_api_format").map_sql_err()?,
row.try_get("format_converted").map_sql_err()?,
row.try_get("model").map_sql_err()?,
row.try_get("prompt").map_sql_err()?,
parse_json(row.try_get("original_request_body").ok().flatten())?,
row.try_get("duration_seconds").map_sql_err()?,
row.try_get("resolution").map_sql_err()?,
row.try_get("aspect_ratio").map_sql_err()?,
row.try_get("size").map_sql_err()?,
VideoTaskStatus::from_database(row.try_get::<String, _>("status").map_sql_err()?.as_str())?,
row.try_get("progress_percent").map_sql_err()?,
row.try_get("progress_message").map_sql_err()?,
row.try_get("retry_count").map_sql_err()?,
row.try_get("poll_interval_seconds").map_sql_err()?,
row.try_get("next_poll_at_unix_secs").map_sql_err()?,
row.try_get("poll_count").map_sql_err()?,
row.try_get("max_poll_count").map_sql_err()?,
row.try_get("created_at_unix_ms").map_sql_err()?,
row.try_get("submitted_at_unix_secs").map_sql_err()?,
row.try_get("completed_at_unix_secs").map_sql_err()?,
row.try_get("updated_at_unix_secs").map_sql_err()?,
row.try_get("error_code").map_sql_err()?,
row.try_get("error_message").map_sql_err()?,
row.try_get("video_url").map_sql_err()?,
parse_json(row.try_get("request_metadata").ok().flatten())?,
)
}
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!("video task JSON field is invalid: {err}"))
})
})
.transpose()
}
fn json_to_string(value: &Option<serde_json::Value>) -> Result<Option<String>, DataLayerError> {
value
.as_ref()
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"video task JSON field is unserializable: {err}"
))
})
})
.transpose()
}
fn status_to_database(status: VideoTaskStatus) -> &'static str {
match status {
VideoTaskStatus::Pending => "pending",
VideoTaskStatus::Submitted => "submitted",
VideoTaskStatus::Queued => "queued",
VideoTaskStatus::Processing => "processing",
VideoTaskStatus::Completed => "completed",
VideoTaskStatus::Failed => "failed",
VideoTaskStatus::Cancelled => "cancelled",
VideoTaskStatus::Expired => "expired",
VideoTaskStatus::Deleted => "deleted",
}
}
fn escape_like_pattern(value: &str) -> String {
value
.replace('\\', "\\\\")
.replace('%', "\\%")
.replace('_', "\\_")
}
fn limit_i64(value: usize, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::UnexpectedValue(format!("invalid {name}: {value}")))
}
fn u64_to_i64(value: u64, name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| DataLayerError::UnexpectedValue(format!("{name} overflow")))
}
fn optional_u64_to_i64(value: Option<u64>, name: &str) -> Result<Option<i64>, DataLayerError> {
value.map(|value| u64_to_i64(value, name)).transpose()
}
fn u32_to_i32(value: u32, name: &str) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| DataLayerError::UnexpectedValue(format!("{name} overflow")))
}
fn optional_u32_to_i32(value: Option<u32>, name: &str) -> Result<Option<i32>, DataLayerError> {
value.map(|value| u32_to_i32(value, name)).transpose()
}
#[cfg(test)]
mod tests {
use super::SqliteVideoTaskRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::video_tasks::{
UpsertVideoTask, VideoTaskLookupKey, VideoTaskQueryFilter, VideoTaskReadRepository,
VideoTaskStatus, VideoTaskWriteRepository,
};
#[tokio::test]
async fn sqlite_repository_writes_and_reads_video_tasks() {
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");
let repository = SqliteVideoTaskRepository::new(pool);
repository
.upsert(sample_task("task-1", VideoTaskStatus::Submitted, 100))
.await
.expect("task should insert");
repository
.upsert(UpsertVideoTask {
user_id: Some("user-2".to_string()),
model: Some("veo-3-fast".to_string()),
client_api_format: Some("gemini:video".to_string()),
created_at_unix_ms: 260,
updated_at_unix_secs: 260,
status: VideoTaskStatus::Completed,
..sample_task("task-2", VideoTaskStatus::Completed, 260)
})
.await
.expect("task should insert");
assert!(repository
.find(VideoTaskLookupKey::ShortId("short-task-1"))
.await
.expect("short lookup should load")
.is_some());
assert!(repository
.find(VideoTaskLookupKey::UserExternal {
user_id: "user-1",
external_task_id: "ext-task-1",
})
.await
.expect("user/external lookup should load")
.is_some());
let due = repository
.list_due(100, 10)
.await
.expect("due tasks should load");
assert_eq!(due.len(), 1);
let claimed = repository
.claim_due(100, 130, 10)
.await
.expect("due tasks should claim");
assert_eq!(claimed.len(), 1);
assert_eq!(claimed[0].next_poll_at_unix_secs, Some(130));
let filter = VideoTaskQueryFilter {
user_id: Some("user-2".to_string()),
status: Some(VideoTaskStatus::Completed),
model_substring: Some("veo".to_string()),
client_api_format: Some("gemini:video".to_string()),
};
assert_eq!(
repository.count(&filter).await.expect("count should load"),
1
);
assert_eq!(
repository
.count_by_status(&filter)
.await
.expect("status counts should load")[0]
.count,
1
);
assert_eq!(
repository
.top_models(&filter, 10)
.await
.expect("top models should load")[0]
.model,
"veo-3-fast"
);
let updated = repository
.update_if_active(UpsertVideoTask {
status: VideoTaskStatus::Processing,
progress_percent: 50,
..sample_task("task-1", VideoTaskStatus::Processing, 150)
})
.await
.expect("active task should update")
.expect("active task should exist");
assert_eq!(updated.progress_percent, 50);
}
fn sample_task(
id: &str,
status: VideoTaskStatus,
updated_at_unix_secs: u64,
) -> UpsertVideoTask {
UpsertVideoTask {
id: id.to_string(),
short_id: Some(format!("short-{id}")),
request_id: format!("request-{id}"),
user_id: Some("user-1".to_string()),
api_key_id: Some("api-key-1".to_string()),
username: Some("user".to_string()),
api_key_name: Some("primary".to_string()),
external_task_id: Some(format!("ext-{id}")),
provider_id: Some("provider-1".to_string()),
endpoint_id: Some("endpoint-1".to_string()),
key_id: Some("provider-key-1".to_string()),
client_api_format: Some("openai:video".to_string()),
provider_api_format: Some("openai:video".to_string()),
format_converted: false,
model: Some("sora-2".to_string()),
prompt: Some("hello".to_string()),
original_request_body: Some(serde_json::json!({"prompt": "hello"})),
duration_seconds: Some(4),
resolution: Some("720p".to_string()),
aspect_ratio: Some("16:9".to_string()),
size: Some("1280x720".to_string()),
status,
progress_percent: 0,
progress_message: None,
retry_count: 0,
poll_interval_seconds: 10,
next_poll_at_unix_secs: Some(updated_at_unix_secs),
poll_count: 0,
max_poll_count: 360,
created_at_unix_ms: updated_at_unix_secs.saturating_sub(10),
submitted_at_unix_secs: Some(updated_at_unix_secs.saturating_sub(10)),
completed_at_unix_secs: None,
updated_at_unix_secs,
error_code: None,
error_message: None,
video_url: None,
request_metadata: Some(serde_json::json!({"request": id})),
}
}
}

Some files were not shown because too many files have changed in this diff Show More