mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
Merge remote-tracking branch 'origin/main' into codex/pool-key-bulk-management-20260714
# Conflicts: # apps/aether-gateway/src/handlers/admin/request/provider/tasks.rs # frontend/src/api/endpoints/pool.ts
This commit is contained in:
@@ -0,0 +1,365 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
|
||||
use aether_data_contracts::repository::announcements::*;
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
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.requires_ack,
|
||||
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 fn list_required_unread_active_announcements(
|
||||
&self,
|
||||
user_id: &str,
|
||||
now_unix_secs: u64,
|
||||
limit: usize,
|
||||
) -> Result<Vec<StoredAnnouncement>, DataLayerError> {
|
||||
let rows = sqlx::query(&format!(
|
||||
r#"
|
||||
{ANNOUNCEMENT_SELECT}
|
||||
WHERE a.is_active = 1
|
||||
AND a.requires_ack = 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
|
||||
)
|
||||
ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC
|
||||
LIMIT ?
|
||||
"#
|
||||
))
|
||||
.bind(now_unix_secs as i64)
|
||||
.bind(now_unix_secs as i64)
|
||||
.bind(user_id)
|
||||
.bind(limit as i64)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_announcement_row).collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[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,
|
||||
requires_ack, 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(record.requires_ack)
|
||||
.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),
|
||||
requires_ack = COALESCE(?, requires_ack),
|
||||
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(record.requires_ack)
|
||||
.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 mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
sqlx::query("DELETE FROM announcement_reads WHERE announcement_id = ?")
|
||||
.bind(announcement_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let rows_affected = sqlx::query("DELETE FROM announcements WHERE id = ?")
|
||||
.bind(announcement_id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
tx.commit().await.map_sql_err()?;
|
||||
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("requires_ack").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);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,277 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
|
||||
use aether_data_contracts::repository::audit::*;
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlAuditLogReadRepository {
|
||||
pool: MysqlPool,
|
||||
}
|
||||
|
||||
impl MysqlAuditLogReadRepository {
|
||||
pub fn new(pool: MysqlPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl AuditLogReadRepository for MysqlAuditLogReadRepository {
|
||||
async fn list_admin_audit_logs(
|
||||
&self,
|
||||
query: &AuditLogListQuery,
|
||||
) -> Result<StoredAdminAuditLogPage, DataLayerError> {
|
||||
let total = sqlx::query_scalar::<_, i64>(
|
||||
r#"
|
||||
SELECT COUNT(*)
|
||||
FROM audit_logs AS a
|
||||
LEFT JOIN users AS u ON a.user_id = u.id
|
||||
WHERE a.created_at >= ?
|
||||
AND (? IS NULL OR LOWER(u.username) LIKE LOWER(?) ESCAPE '\\')
|
||||
AND (? IS NULL OR a.event_type = ?)
|
||||
"#,
|
||||
)
|
||||
.bind(query.cutoff_unix_secs as i64)
|
||||
.bind(query.username_pattern.as_deref())
|
||||
.bind(query.username_pattern.as_deref())
|
||||
.bind(query.event_type.as_deref())
|
||||
.bind(query.event_type.as_deref())
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
a.id,
|
||||
a.event_type,
|
||||
a.user_id,
|
||||
u.email AS user_email,
|
||||
u.username AS user_username,
|
||||
a.description,
|
||||
a.ip_address,
|
||||
a.status_code,
|
||||
a.error_message,
|
||||
a.event_metadata AS metadata,
|
||||
a.created_at
|
||||
FROM audit_logs AS a
|
||||
LEFT JOIN users AS u ON a.user_id = u.id
|
||||
WHERE a.created_at >= ?
|
||||
AND (? IS NULL OR LOWER(u.username) LIKE LOWER(?) ESCAPE '\\')
|
||||
AND (? IS NULL OR a.event_type = ?)
|
||||
ORDER BY a.created_at DESC
|
||||
LIMIT ? OFFSET ?
|
||||
"#,
|
||||
)
|
||||
.bind(query.cutoff_unix_secs as i64)
|
||||
.bind(query.username_pattern.as_deref())
|
||||
.bind(query.username_pattern.as_deref())
|
||||
.bind(query.event_type.as_deref())
|
||||
.bind(query.event_type.as_deref())
|
||||
.bind(query.limit as i64)
|
||||
.bind(query.offset as i64)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_mysql_admin_audit_log_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
|
||||
Ok(StoredAdminAuditLogPage {
|
||||
items,
|
||||
total: total.max(0) as u64,
|
||||
})
|
||||
}
|
||||
|
||||
async fn list_admin_suspicious_activities(
|
||||
&self,
|
||||
cutoff_unix_secs: u64,
|
||||
) -> Result<Vec<StoredSuspiciousActivity>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT id, event_type, user_id, description, ip_address, event_metadata AS metadata, created_at
|
||||
FROM audit_logs
|
||||
WHERE created_at >= ?
|
||||
AND event_type IN (?, ?, ?, ?)
|
||||
ORDER BY created_at DESC
|
||||
LIMIT 100
|
||||
"#,
|
||||
)
|
||||
.bind(cutoff_unix_secs as i64)
|
||||
.bind(SUSPICIOUS_EVENT_TYPES[0])
|
||||
.bind(SUSPICIOUS_EVENT_TYPES[1])
|
||||
.bind(SUSPICIOUS_EVENT_TYPES[2])
|
||||
.bind(SUSPICIOUS_EVENT_TYPES[3])
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
rows.iter().map(map_mysql_suspicious_activity_row).collect()
|
||||
}
|
||||
|
||||
async fn read_admin_user_behavior_event_counts(
|
||||
&self,
|
||||
user_id: &str,
|
||||
cutoff_unix_secs: u64,
|
||||
) -> Result<std::collections::BTreeMap<String, u64>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT event_type, COUNT(*) AS count
|
||||
FROM audit_logs
|
||||
WHERE user_id = ?
|
||||
AND created_at >= ?
|
||||
GROUP BY event_type
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(cutoff_unix_secs as i64)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
Ok(rows
|
||||
.iter()
|
||||
.filter_map(|row| event_count_from_mysql_row(row).ok())
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn list_user_audit_logs(
|
||||
&self,
|
||||
user_id: &str,
|
||||
query: &AuditLogListQuery,
|
||||
) -> Result<StoredUserAuditLogPage, DataLayerError> {
|
||||
let total = sqlx::query_scalar::<_, i64>(
|
||||
r#"
|
||||
SELECT COUNT(*)
|
||||
FROM audit_logs
|
||||
WHERE user_id = ?
|
||||
AND created_at >= ?
|
||||
AND (? IS NULL OR event_type = ?)
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(query.cutoff_unix_secs as i64)
|
||||
.bind(query.event_type.as_deref())
|
||||
.bind(query.event_type.as_deref())
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT id, event_type, description, ip_address, status_code, created_at
|
||||
FROM audit_logs
|
||||
WHERE user_id = ?
|
||||
AND created_at >= ?
|
||||
AND (? IS NULL OR event_type = ?)
|
||||
ORDER BY created_at DESC
|
||||
LIMIT ? OFFSET ?
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(query.cutoff_unix_secs as i64)
|
||||
.bind(query.event_type.as_deref())
|
||||
.bind(query.event_type.as_deref())
|
||||
.bind(query.limit as i64)
|
||||
.bind(query.offset as i64)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_mysql_user_audit_log_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
|
||||
Ok(StoredUserAuditLogPage {
|
||||
items,
|
||||
total: total.max(0) as u64,
|
||||
})
|
||||
}
|
||||
|
||||
async fn delete_audit_logs_before(
|
||||
&self,
|
||||
cutoff_unix_secs: u64,
|
||||
limit: usize,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let deleted = sqlx::query(
|
||||
r#"
|
||||
DELETE FROM audit_logs
|
||||
WHERE id IN (
|
||||
SELECT id
|
||||
FROM (
|
||||
SELECT id
|
||||
FROM audit_logs
|
||||
WHERE created_at < ?
|
||||
ORDER BY created_at ASC, id ASC
|
||||
LIMIT ?
|
||||
) AS doomed
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.bind(cutoff_unix_secs.min(i64::MAX as u64) as i64)
|
||||
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(usize::try_from(deleted).unwrap_or(usize::MAX))
|
||||
}
|
||||
}
|
||||
|
||||
fn mysql_created_at_unix_secs(row: &MySqlRow) -> Result<u64, DataLayerError> {
|
||||
let value = row.try_get::<i64, _>("created_at").map_sql_err()?;
|
||||
Ok(value.max(0) as u64)
|
||||
}
|
||||
|
||||
fn map_mysql_admin_audit_log_row(row: &MySqlRow) -> Result<StoredAdminAuditLog, DataLayerError> {
|
||||
Ok(StoredAdminAuditLog {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
event_type: row.try_get("event_type").map_sql_err()?,
|
||||
user_id: row.try_get("user_id").map_sql_err()?,
|
||||
user_email: row.try_get("user_email").map_sql_err()?,
|
||||
user_username: row.try_get("user_username").map_sql_err()?,
|
||||
description: row.try_get("description").map_sql_err()?,
|
||||
ip_address: row.try_get("ip_address").map_sql_err()?,
|
||||
status_code: row.try_get("status_code").map_sql_err()?,
|
||||
error_message: row.try_get("error_message").map_sql_err()?,
|
||||
metadata: optional_json_from_text(row.try_get("metadata").map_sql_err()?)?,
|
||||
created_at_unix_secs: mysql_created_at_unix_secs(row)?,
|
||||
})
|
||||
}
|
||||
|
||||
fn map_mysql_suspicious_activity_row(
|
||||
row: &MySqlRow,
|
||||
) -> Result<StoredSuspiciousActivity, DataLayerError> {
|
||||
Ok(StoredSuspiciousActivity {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
event_type: row.try_get("event_type").map_sql_err()?,
|
||||
user_id: row.try_get("user_id").map_sql_err()?,
|
||||
description: row.try_get("description").map_sql_err()?,
|
||||
ip_address: row.try_get("ip_address").map_sql_err()?,
|
||||
metadata: optional_json_from_text(row.try_get("metadata").map_sql_err()?)?,
|
||||
created_at_unix_secs: mysql_created_at_unix_secs(row)?,
|
||||
})
|
||||
}
|
||||
|
||||
fn map_mysql_user_audit_log_row(row: &MySqlRow) -> Result<StoredUserAuditLog, DataLayerError> {
|
||||
Ok(StoredUserAuditLog {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
event_type: row.try_get("event_type").map_sql_err()?,
|
||||
description: row.try_get("description").map_sql_err()?,
|
||||
ip_address: row.try_get("ip_address").map_sql_err()?,
|
||||
status_code: row.try_get("status_code").map_sql_err()?,
|
||||
created_at_unix_secs: mysql_created_at_unix_secs(row)?,
|
||||
})
|
||||
}
|
||||
|
||||
fn event_count_from_mysql_row(row: &MySqlRow) -> Result<(String, u64), DataLayerError> {
|
||||
let event_type = row.try_get("event_type").map_sql_err()?;
|
||||
let count = row.try_get::<i64, _>("count").map_sql_err()?.max(0) as u64;
|
||||
Ok((event_type, count))
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,248 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::auth_modules::*;
|
||||
use aether_data_contracts::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_limit, WhereClause};
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
const OAUTH_PROVIDER_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
provider_type,
|
||||
display_name,
|
||||
client_id,
|
||||
client_secret_encrypted,
|
||||
redirect_uri
|
||||
FROM oauth_providers
|
||||
"#;
|
||||
|
||||
const LDAP_CONFIG_COLUMNS: &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
|
||||
"#;
|
||||
|
||||
#[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 fn list_enabled_oauth_providers(
|
||||
pool: &MysqlPool,
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(OAUTH_PROVIDER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "is_enabled", true);
|
||||
builder.push(" ORDER BY provider_type ASC");
|
||||
let rows = builder.build().fetch_all(pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_oauth_row).collect()
|
||||
}
|
||||
|
||||
async fn get_ldap_config(
|
||||
pool: &MysqlPool,
|
||||
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(LDAP_CONFIG_COLUMNS);
|
||||
builder.push(" ORDER BY id ASC");
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder.build().fetch_optional(pool).await.map_sql_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl AuthModuleReadRepository for MysqlAuthModuleReadRepository {
|
||||
async fn list_enabled_oauth_providers(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
list_enabled_oauth_providers(&self.pool).await
|
||||
}
|
||||
|
||||
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
get_ldap_config(&self.pool).await
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl AuthModuleReadRepository for MysqlAuthModuleRepository {
|
||||
async fn list_enabled_oauth_providers(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
list_enabled_oauth_providers(&self.pool).await
|
||||
}
|
||||
|
||||
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
get_ldap_config(&self.pool).await
|
||||
}
|
||||
}
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,444 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::background_tasks::*;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::{DataLayerError, MysqlPool};
|
||||
|
||||
const RUN_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
task_key,
|
||||
kind,
|
||||
`trigger`,
|
||||
status,
|
||||
attempt,
|
||||
max_attempts,
|
||||
owner_instance,
|
||||
progress_percent,
|
||||
progress_message,
|
||||
payload_json,
|
||||
result_json,
|
||||
error_message,
|
||||
cancel_requested,
|
||||
created_by,
|
||||
created_at_unix_secs,
|
||||
started_at_unix_secs,
|
||||
finished_at_unix_secs,
|
||||
updated_at_unix_secs
|
||||
FROM background_task_runs
|
||||
"#;
|
||||
|
||||
const EVENT_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
run_id,
|
||||
event_type,
|
||||
message,
|
||||
payload_json,
|
||||
created_at_unix_secs
|
||||
FROM background_task_events
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlBackgroundTaskRepository {
|
||||
pool: MysqlPool,
|
||||
}
|
||||
|
||||
impl MysqlBackgroundTaskRepository {
|
||||
pub fn new(pool: MysqlPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
fn apply_run_filter(builder: &mut QueryBuilder<'_, MySql>, query: &BackgroundTaskListQuery) {
|
||||
let mut has_where = false;
|
||||
if let Some(kind) = query.kind {
|
||||
if !has_where {
|
||||
builder.push(" WHERE ");
|
||||
has_where = true;
|
||||
} else {
|
||||
builder.push(" AND ");
|
||||
}
|
||||
builder.push("kind = ").push_bind(kind.as_database());
|
||||
}
|
||||
if let Some(status) = query.status {
|
||||
if !has_where {
|
||||
builder.push(" WHERE ");
|
||||
has_where = true;
|
||||
} else {
|
||||
builder.push(" AND ");
|
||||
}
|
||||
builder.push("status = ").push_bind(status.as_database());
|
||||
}
|
||||
if let Some(trigger) = query.trigger.as_deref() {
|
||||
if !has_where {
|
||||
builder.push(" WHERE ");
|
||||
has_where = true;
|
||||
} else {
|
||||
builder.push(" AND ");
|
||||
}
|
||||
builder.push("`trigger` = ").push_bind(trigger.to_string());
|
||||
}
|
||||
if let Some(task_key_substring) = query.task_key_substring.as_deref() {
|
||||
if !has_where {
|
||||
builder.push(" WHERE ");
|
||||
} else {
|
||||
builder.push(" AND ");
|
||||
}
|
||||
builder.push("LOWER(task_key) LIKE ").push_bind(format!(
|
||||
"%{}%",
|
||||
task_key_substring.trim().to_ascii_lowercase()
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl BackgroundTaskReadRepository for MysqlBackgroundTaskRepository {
|
||||
async fn find_run(
|
||||
&self,
|
||||
run_id: &str,
|
||||
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
|
||||
let row = sqlx::query(&format!("{RUN_COLUMNS} WHERE id = ? LIMIT 1"))
|
||||
.bind(run_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_run_row).transpose()
|
||||
}
|
||||
|
||||
async fn list_runs(
|
||||
&self,
|
||||
query: &BackgroundTaskListQuery,
|
||||
) -> Result<StoredBackgroundTaskRunPage, DataLayerError> {
|
||||
let limit = query.limit.max(1);
|
||||
let mut count_builder =
|
||||
QueryBuilder::<MySql>::new("SELECT COUNT(id) AS total FROM background_task_runs");
|
||||
Self::apply_run_filter(&mut count_builder, query);
|
||||
let total = count_builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut builder = QueryBuilder::<MySql>::new(RUN_COLUMNS);
|
||||
Self::apply_run_filter(&mut builder, query);
|
||||
builder
|
||||
.push(" ORDER BY created_at_unix_secs DESC, updated_at_unix_secs DESC")
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(limit, "run limit")?)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(i64_from_usize(query.offset, "run offset")?);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_run_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
Ok(StoredBackgroundTaskRunPage {
|
||||
items,
|
||||
total: usize::try_from(total).unwrap_or_default(),
|
||||
})
|
||||
}
|
||||
|
||||
async fn list_events(
|
||||
&self,
|
||||
run_id: &str,
|
||||
offset: usize,
|
||||
limit: usize,
|
||||
) -> Result<Vec<StoredBackgroundTaskEvent>, DataLayerError> {
|
||||
let limit = limit.max(1);
|
||||
let rows = sqlx::query(&format!(
|
||||
"{EVENT_COLUMNS} WHERE run_id = ? ORDER BY created_at_unix_secs ASC, id ASC LIMIT ? OFFSET ?"
|
||||
))
|
||||
.bind(run_id)
|
||||
.bind(i64_from_usize(limit, "event limit")?)
|
||||
.bind(i64_from_usize(offset, "event offset")?)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_event_row).collect()
|
||||
}
|
||||
|
||||
async fn summarize_runs(&self) -> Result<BackgroundTaskSummary, DataLayerError> {
|
||||
let total = sqlx::query_scalar::<_, i64>("SELECT COUNT(id) FROM background_task_runs")
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let running_count = sqlx::query_scalar::<_, i64>(
|
||||
"SELECT COUNT(id) FROM background_task_runs WHERE status = 'running'",
|
||||
)
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let status_rows = sqlx::query(
|
||||
"SELECT status, COUNT(id) AS total FROM background_task_runs GROUP BY status",
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let kind_rows =
|
||||
sqlx::query("SELECT kind, COUNT(id) AS total FROM background_task_runs GROUP BY kind")
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut by_status = std::collections::BTreeMap::new();
|
||||
for row in status_rows {
|
||||
let key: String = row.try_get("status").map_sql_err()?;
|
||||
let count: i64 = row.try_get("total").map_sql_err()?;
|
||||
by_status.insert(key, u64::try_from(count).unwrap_or_default());
|
||||
}
|
||||
let mut by_kind = std::collections::BTreeMap::new();
|
||||
for row in kind_rows {
|
||||
let key: String = row.try_get("kind").map_sql_err()?;
|
||||
let count: i64 = row.try_get("total").map_sql_err()?;
|
||||
by_kind.insert(key, u64::try_from(count).unwrap_or_default());
|
||||
}
|
||||
|
||||
Ok(BackgroundTaskSummary {
|
||||
total: u64::try_from(total).unwrap_or_default(),
|
||||
running_count: u64::try_from(running_count).unwrap_or_default(),
|
||||
by_status,
|
||||
by_kind,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl BackgroundTaskWriteRepository for MysqlBackgroundTaskRepository {
|
||||
async fn upsert_run(
|
||||
&self,
|
||||
run: UpsertBackgroundTaskRun,
|
||||
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
|
||||
run.validate()?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO background_task_runs (
|
||||
id,
|
||||
task_key,
|
||||
kind,
|
||||
`trigger`,
|
||||
status,
|
||||
attempt,
|
||||
max_attempts,
|
||||
owner_instance,
|
||||
progress_percent,
|
||||
progress_message,
|
||||
payload_json,
|
||||
result_json,
|
||||
error_message,
|
||||
cancel_requested,
|
||||
created_by,
|
||||
created_at_unix_secs,
|
||||
started_at_unix_secs,
|
||||
finished_at_unix_secs,
|
||||
updated_at_unix_secs
|
||||
) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)
|
||||
ON DUPLICATE KEY UPDATE
|
||||
task_key = VALUES(task_key),
|
||||
kind = VALUES(kind),
|
||||
`trigger` = VALUES(`trigger`),
|
||||
status = VALUES(status),
|
||||
attempt = VALUES(attempt),
|
||||
max_attempts = VALUES(max_attempts),
|
||||
owner_instance = VALUES(owner_instance),
|
||||
progress_percent = VALUES(progress_percent),
|
||||
progress_message = VALUES(progress_message),
|
||||
payload_json = VALUES(payload_json),
|
||||
result_json = VALUES(result_json),
|
||||
error_message = VALUES(error_message),
|
||||
cancel_requested = VALUES(cancel_requested),
|
||||
created_by = VALUES(created_by),
|
||||
created_at_unix_secs = VALUES(created_at_unix_secs),
|
||||
started_at_unix_secs = VALUES(started_at_unix_secs),
|
||||
finished_at_unix_secs = VALUES(finished_at_unix_secs),
|
||||
updated_at_unix_secs = VALUES(updated_at_unix_secs)
|
||||
"#,
|
||||
)
|
||||
.bind(&run.id)
|
||||
.bind(&run.task_key)
|
||||
.bind(run.kind.as_database())
|
||||
.bind(&run.trigger)
|
||||
.bind(run.status.as_database())
|
||||
.bind(i64::from(run.attempt))
|
||||
.bind(i64::from(run.max_attempts))
|
||||
.bind(run.owner_instance.as_deref())
|
||||
.bind(i32::from(run.progress_percent))
|
||||
.bind(run.progress_message.as_deref())
|
||||
.bind(json_to_string(&run.payload_json, "payload_json")?)
|
||||
.bind(json_to_string(&run.result_json, "result_json")?)
|
||||
.bind(run.error_message.as_deref())
|
||||
.bind(run.cancel_requested)
|
||||
.bind(run.created_by.as_deref())
|
||||
.bind(u64_to_i64(
|
||||
run.created_at_unix_secs,
|
||||
"created_at_unix_secs",
|
||||
)?)
|
||||
.bind(run.started_at_unix_secs.map(|value| value as i64))
|
||||
.bind(run.finished_at_unix_secs.map(|value| value as i64))
|
||||
.bind(u64_to_i64(
|
||||
run.updated_at_unix_secs,
|
||||
"updated_at_unix_secs",
|
||||
)?)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
self.find_run(&run.id).await?.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue("background task run missing after upsert".to_string())
|
||||
})
|
||||
}
|
||||
|
||||
async fn request_cancel(
|
||||
&self,
|
||||
run_id: &str,
|
||||
updated_at_unix_secs: u64,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let affected = sqlx::query(
|
||||
"UPDATE background_task_runs SET cancel_requested = TRUE, updated_at_unix_secs = ? WHERE id = ?",
|
||||
)
|
||||
.bind(u64_to_i64(updated_at_unix_secs, "updated_at_unix_secs")?)
|
||||
.bind(run_id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
Ok(affected > 0)
|
||||
}
|
||||
|
||||
async fn upsert_event(
|
||||
&self,
|
||||
event: UpsertBackgroundTaskEvent,
|
||||
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
|
||||
event.validate()?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO background_task_events (
|
||||
id, run_id, event_type, message, payload_json, created_at_unix_secs
|
||||
) VALUES (?, ?, ?, ?, ?, ?)
|
||||
ON DUPLICATE KEY UPDATE
|
||||
run_id = VALUES(run_id),
|
||||
event_type = VALUES(event_type),
|
||||
message = VALUES(message),
|
||||
payload_json = VALUES(payload_json),
|
||||
created_at_unix_secs = VALUES(created_at_unix_secs)
|
||||
"#,
|
||||
)
|
||||
.bind(&event.id)
|
||||
.bind(&event.run_id)
|
||||
.bind(&event.event_type)
|
||||
.bind(&event.message)
|
||||
.bind(json_to_string(&event.payload_json, "payload_json")?)
|
||||
.bind(u64_to_i64(
|
||||
event.created_at_unix_secs,
|
||||
"created_at_unix_secs",
|
||||
)?)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let row = sqlx::query(&format!("{EVENT_COLUMNS} WHERE id = ? LIMIT 1"))
|
||||
.bind(&event.id)
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
map_event_row(&row)
|
||||
}
|
||||
}
|
||||
|
||||
fn map_run_row(row: &MySqlRow) -> Result<StoredBackgroundTaskRun, DataLayerError> {
|
||||
let kind: String = row.try_get("kind").map_sql_err()?;
|
||||
let status: String = row.try_get("status").map_sql_err()?;
|
||||
let attempt: i64 = row.try_get("attempt").map_sql_err()?;
|
||||
let max_attempts: i64 = row.try_get("max_attempts").map_sql_err()?;
|
||||
let progress_percent: i32 = row.try_get("progress_percent").map_sql_err()?;
|
||||
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
|
||||
let started_at_unix_secs: Option<i64> = row.try_get("started_at_unix_secs").map_sql_err()?;
|
||||
let finished_at_unix_secs: Option<i64> = row.try_get("finished_at_unix_secs").map_sql_err()?;
|
||||
let updated_at_unix_secs: i64 = row.try_get("updated_at_unix_secs").map_sql_err()?;
|
||||
|
||||
Ok(StoredBackgroundTaskRun {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
task_key: row.try_get("task_key").map_sql_err()?,
|
||||
kind: BackgroundTaskKind::from_database(&kind)?,
|
||||
trigger: row.try_get("trigger").map_sql_err()?,
|
||||
status: BackgroundTaskStatus::from_database(&status)?,
|
||||
attempt: u32::try_from(attempt).unwrap_or_default(),
|
||||
max_attempts: u32::try_from(max_attempts).unwrap_or_default(),
|
||||
owner_instance: row.try_get("owner_instance").map_sql_err()?,
|
||||
progress_percent: u16::try_from(progress_percent).unwrap_or_default(),
|
||||
progress_message: row.try_get("progress_message").map_sql_err()?,
|
||||
payload_json: parse_optional_json(
|
||||
row.try_get("payload_json").ok().flatten(),
|
||||
"payload_json",
|
||||
)?,
|
||||
result_json: parse_optional_json(row.try_get("result_json").ok().flatten(), "result_json")?,
|
||||
error_message: row.try_get("error_message").map_sql_err()?,
|
||||
cancel_requested: row.try_get("cancel_requested").map_sql_err()?,
|
||||
created_by: row.try_get("created_by").map_sql_err()?,
|
||||
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
|
||||
started_at_unix_secs: started_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
|
||||
finished_at_unix_secs: finished_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
|
||||
updated_at_unix_secs: u64::try_from(updated_at_unix_secs).unwrap_or_default(),
|
||||
})
|
||||
}
|
||||
|
||||
fn map_event_row(row: &MySqlRow) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
|
||||
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
|
||||
Ok(StoredBackgroundTaskEvent {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
run_id: row.try_get("run_id").map_sql_err()?,
|
||||
event_type: row.try_get("event_type").map_sql_err()?,
|
||||
message: row.try_get("message").map_sql_err()?,
|
||||
payload_json: parse_optional_json(
|
||||
row.try_get("payload_json").ok().flatten(),
|
||||
"payload_json",
|
||||
)?,
|
||||
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
|
||||
})
|
||||
}
|
||||
|
||||
fn i64_from_usize(value: usize, label: &str) -> Result<i64, DataLayerError> {
|
||||
i64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
|
||||
})
|
||||
}
|
||||
|
||||
fn u64_to_i64(value: u64, label: &str) -> Result<i64, DataLayerError> {
|
||||
i64::try_from(value).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
|
||||
})
|
||||
}
|
||||
|
||||
fn 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!(
|
||||
"background task {field_name} is unserializable: {err}"
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn parse_optional_json(
|
||||
value: Option<String>,
|
||||
field_name: &str,
|
||||
) -> 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!(
|
||||
"background task {field_name} contains invalid JSON: {err}"
|
||||
))
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,812 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::candidate_selection::{
|
||||
MinimalCandidateSelectionReadRepository, StoredMinimalCandidateSelectionRow,
|
||||
StoredPoolKeyCandidateOrder, StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
StoredPoolKeyCandidateRowsQuery, StoredProviderModelMapping,
|
||||
StoredRequestedModelCandidateRowsQuery,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
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,
|
||||
pak.last_used_at AS key_last_used_at_unix_secs,
|
||||
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>,
|
||||
key_last_used_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
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()
|
||||
.filter(|row| {
|
||||
row.row.provider_id == query.provider_id
|
||||
&& row.row.endpoint_id == query.endpoint_id
|
||||
&& row.row.model_id == query.model_id
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let mut rows = sort_pool_key_rows(rows, &query.order);
|
||||
Ok(rows
|
||||
.drain(..)
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.map(|item| item.row)
|
||||
.collect())
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group_key_ids(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsByKeyIdsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
if query.key_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let key_order = query
|
||||
.key_ids
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(index, key_id)| (key_id.as_str(), index))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let mut rows = self
|
||||
.load_rows_for_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row.row.provider_id == query.provider_id
|
||||
&& row.row.endpoint_id == query.endpoint_id
|
||||
&& row.row.model_id == query.model_id
|
||||
&& key_order.contains_key(row.row.key_id.as_str())
|
||||
})
|
||||
.map(|item| item.row)
|
||||
.collect::<Vec<_>>();
|
||||
rows.sort_by(|left, right| {
|
||||
key_order
|
||||
.get(left.key_id.as_str())
|
||||
.cmp(&key_order.get(right.key_id.as_str()))
|
||||
.then(left.key_id.cmp(&right.key_id))
|
||||
});
|
||||
Ok(dedupe_candidate_selection_rows(rows))
|
||||
}
|
||||
}
|
||||
|
||||
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<CandidateSelectionRow>,
|
||||
order: &StoredPoolKeyCandidateOrder,
|
||||
) -> Vec<CandidateSelectionRow> {
|
||||
rows.sort_by(|left, right| match order {
|
||||
StoredPoolKeyCandidateOrder::InternalPriority => compare_pool_key_internal(left, right),
|
||||
StoredPoolKeyCandidateOrder::Lru => left
|
||||
.key_last_used_at_unix_secs
|
||||
.cmp(&right.key_last_used_at_unix_secs)
|
||||
.then_with(|| compare_pool_key_internal(left, right)),
|
||||
StoredPoolKeyCandidateOrder::CacheAffinity => right
|
||||
.key_last_used_at_unix_secs
|
||||
.cmp(&left.key_last_used_at_unix_secs)
|
||||
.then_with(|| compare_pool_key_internal(left, right)),
|
||||
StoredPoolKeyCandidateOrder::SingleAccount => left
|
||||
.row
|
||||
.key_internal_priority
|
||||
.cmp(&right.row.key_internal_priority)
|
||||
.then_with(|| {
|
||||
right
|
||||
.key_last_used_at_unix_secs
|
||||
.cmp(&left.key_last_used_at_unix_secs)
|
||||
})
|
||||
.then(left.row.key_id.cmp(&right.row.key_id)),
|
||||
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
|
||||
stable_pool_key_hash(seed.as_str(), left.row.key_id.as_str())
|
||||
.cmp(&stable_pool_key_hash(
|
||||
seed.as_str(),
|
||||
right.row.key_id.as_str(),
|
||||
))
|
||||
.then(left.row.key_id.cmp(&right.row.key_id))
|
||||
}
|
||||
});
|
||||
rows
|
||||
}
|
||||
|
||||
fn compare_pool_key_internal(
|
||||
left: &CandidateSelectionRow,
|
||||
right: &CandidateSelectionRow,
|
||||
) -> std::cmp::Ordering {
|
||||
left.row
|
||||
.key_internal_priority
|
||||
.cmp(&right.row.key_internal_priority)
|
||||
.then(left.row.key_id.cmp(&right.row.key_id))
|
||||
}
|
||||
|
||||
fn stable_pool_key_hash(seed: &str, key_id: &str) -> u64 {
|
||||
let mut hash = 0xcbf29ce484222325u64;
|
||||
for byte in seed
|
||||
.as_bytes()
|
||||
.iter()
|
||||
.copied()
|
||||
.chain(std::iter::once(b':'))
|
||||
.chain(key_id.as_bytes().iter().copied())
|
||||
{
|
||||
hash ^= u64::from(byte);
|
||||
hash = hash.wrapping_mul(0x100000001b3);
|
||||
}
|
||||
hash
|
||||
}
|
||||
|
||||
fn row_matches_requested_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
requested_model_name: &str,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
(row_has_available_provider_model(row, api_format)
|
||||
&& row.global_model_name == requested_model_name)
|
||||
|| (row_default_provider_model_name_available(row, api_format)
|
||||
&& row.model_provider_model_name == requested_model_name)
|
||||
|| row
|
||||
.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings.iter().any(|mapping| {
|
||||
mapping_scope_matches(mapping, row, api_format)
|
||||
&& mapping.name == requested_model_name
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn row_has_available_provider_model(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
row_mapping_matches_scope(row, api_format)
|
||||
|| row_default_provider_model_name_available(row, api_format)
|
||||
}
|
||||
|
||||
fn row_default_provider_model_name_available(
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
let Some(mappings) = row.model_provider_model_mappings.as_ref() else {
|
||||
return true;
|
||||
};
|
||||
let mut has_explicit_default_mapping = false;
|
||||
for mapping in mappings {
|
||||
if mapping.name != row.model_provider_model_name {
|
||||
continue;
|
||||
}
|
||||
has_explicit_default_mapping = true;
|
||||
if mapping_scope_matches(mapping, row, api_format) {
|
||||
return true;
|
||||
}
|
||||
}
|
||||
!has_explicit_default_mapping
|
||||
}
|
||||
|
||||
fn row_mapping_matches_scope(row: &StoredMinimalCandidateSelectionRow, api_format: &str) -> bool {
|
||||
row.model_provider_model_mappings
|
||||
.as_ref()
|
||||
.is_some_and(|mappings| {
|
||||
mappings
|
||||
.iter()
|
||||
.any(|mapping| mapping_scope_matches(mapping, row, api_format))
|
||||
})
|
||||
}
|
||||
|
||||
fn mapping_scope_matches(
|
||||
mapping: &StoredProviderModelMapping,
|
||||
row: &StoredMinimalCandidateSelectionRow,
|
||||
api_format: &str,
|
||||
) -> bool {
|
||||
mapping.api_formats.as_ref().is_none_or(|formats| {
|
||||
formats
|
||||
.iter()
|
||||
.any(|value| api_format_scope_covers(value, api_format))
|
||||
}) && mapping.endpoint_ids.as_ref().is_none_or(|endpoint_ids| {
|
||||
endpoint_ids
|
||||
.iter()
|
||||
.any(|endpoint_id| endpoint_id == &row.endpoint_id)
|
||||
})
|
||||
}
|
||||
|
||||
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:search"
|
||||
| "openai:image"
|
||||
)
|
||||
}
|
||||
"chatgpt_web" => {
|
||||
matches!(auth_type.as_str(), "oauth" | "bearer") && api_format == "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"
|
||||
}
|
||||
"grok" => {
|
||||
auth_type == "oauth"
|
||||
&& matches!(
|
||||
api_format.as_str(),
|
||||
"openai:chat" | "openai:responses" | "claude:messages" | "openai:image"
|
||||
)
|
||||
}
|
||||
"windsurf" => {
|
||||
matches!(auth_type.as_str(), "oauth" | "api_key" | "bearer")
|
||||
&& api_format == "openai:chat"
|
||||
}
|
||||
"vertex_ai" => {
|
||||
(auth_type == "api_key"
|
||||
&& matches!(
|
||||
api_format.as_str(),
|
||||
"gemini:generate_content" | "gemini:embedding"
|
||||
))
|
||||
|| (matches!(auth_type.as_str(), "service_account" | "vertex_ai")
|
||||
&& matches!(
|
||||
api_format.as_str(),
|
||||
"claude:messages" | "gemini:generate_content" | "gemini:embedding"
|
||||
))
|
||||
}
|
||||
_ => 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()?,
|
||||
key_last_used_at_unix_secs: row
|
||||
.try_get::<Option<i64>, _>("key_last_used_at_unix_secs")
|
||||
.map_sql_err()?
|
||||
.and_then(|value| u64::try_from(value).ok()),
|
||||
})
|
||||
}
|
||||
|
||||
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,
|
||||
endpoint_ids: None,
|
||||
operations: 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,
|
||||
endpoint_ids: None,
|
||||
operations: 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()
|
||||
});
|
||||
let endpoint_ids = parse_string_list(
|
||||
object.get("endpoint_ids").cloned(),
|
||||
"models.provider_model_mappings.endpoint_ids",
|
||||
)?;
|
||||
let operations = parse_string_list(
|
||||
object.get("operations").cloned(),
|
||||
"models.provider_model_mappings.operations",
|
||||
)?
|
||||
.and_then(normalize_request_operations);
|
||||
|
||||
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,
|
||||
endpoint_ids,
|
||||
operations,
|
||||
}))
|
||||
}
|
||||
|
||||
fn normalize_request_operations(values: Vec<String>) -> Option<Vec<String>> {
|
||||
let operations = values
|
||||
.into_iter()
|
||||
.map(|value| value.trim().to_ascii_lowercase())
|
||||
.filter(|value| !value.is_empty())
|
||||
.collect::<Vec<_>>();
|
||||
(!operations.is_empty()).then_some(operations)
|
||||
}
|
||||
|
||||
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 api_format_scope_covers(allowed: &str, requested: &str) -> bool {
|
||||
aether_ai_formats::api_format_permission_covers(allowed, requested)
|
||||
}
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,808 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::candidates::{
|
||||
request_candidate_lifecycle_would_regress, PublicHealthStatusCount, PublicHealthTimelineBucket,
|
||||
RequestCandidateReadRepository, RequestCandidateStatus, RequestCandidateWriteRepository,
|
||||
StoredRequestCandidate, UpsertRequestCandidateRecord,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
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_attempted_by_request_id(
|
||||
&self,
|
||||
request_id: &str,
|
||||
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
|
||||
let rows = sqlx::query(&format!(
|
||||
"{CANDIDATE_COLUMNS} WHERE request_id = ? \
|
||||
AND (status IN ('streaming', 'success', 'failed', 'cancelled') \
|
||||
OR (status = 'pending' AND started_at IS NOT NULL)) \
|
||||
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 preserve_existing_lifecycle = existing.as_ref().is_some_and(|value| {
|
||||
request_candidate_lifecycle_would_regress(value.status, candidate.status)
|
||||
});
|
||||
let merged_status = if preserve_existing_lifecycle {
|
||||
existing
|
||||
.as_ref()
|
||||
.map(|value| value.status)
|
||||
.unwrap_or(candidate.status)
|
||||
} else {
|
||||
candidate.status
|
||||
};
|
||||
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())),
|
||||
merged_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)),
|
||||
if preserve_existing_lifecycle {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.status_code.map(i32::from))
|
||||
} else {
|
||||
candidate.status_code.map(i32::from).or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.status_code.map(i32::from))
|
||||
})
|
||||
},
|
||||
if preserve_existing_lifecycle {
|
||||
existing.as_ref().and_then(|value| value.error_type.clone())
|
||||
} else {
|
||||
candidate
|
||||
.error_type
|
||||
.or_else(|| existing.as_ref().and_then(|value| value.error_type.clone()))
|
||||
},
|
||||
if preserve_existing_lifecycle {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.error_message.clone())
|
||||
} else {
|
||||
candidate.error_message.or_else(|| {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.error_message.clone())
|
||||
})
|
||||
},
|
||||
if preserve_existing_lifecycle {
|
||||
match existing.as_ref().and_then(|value| value.latency_ms) {
|
||||
Some(value) => Some(to_i32_u64(value)?),
|
||||
None => None,
|
||||
}
|
||||
} else {
|
||||
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()?,
|
||||
if preserve_existing_lifecycle {
|
||||
existing
|
||||
.as_ref()
|
||||
.and_then(|value| value.finished_at_unix_ms)
|
||||
} else {
|
||||
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;
|
||||
use aether_data_contracts::repository::candidates::{
|
||||
RequestCandidateStatus, StoredRequestCandidate, UpsertRequestCandidateRecord,
|
||||
};
|
||||
|
||||
#[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);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merge_candidate_keeps_terminal_status_when_streaming_arrives_late() {
|
||||
let existing = StoredRequestCandidate::new(
|
||||
"candidate-1".to_string(),
|
||||
"request-1".to_string(),
|
||||
Some("user-1".to_string()),
|
||||
Some("key-1".to_string()),
|
||||
None,
|
||||
None,
|
||||
0,
|
||||
0,
|
||||
Some("provider-1".to_string()),
|
||||
Some("endpoint-1".to_string()),
|
||||
Some("provider-key-1".to_string()),
|
||||
RequestCandidateStatus::Success,
|
||||
None,
|
||||
false,
|
||||
Some(200),
|
||||
None,
|
||||
None,
|
||||
Some(123),
|
||||
None,
|
||||
Some(serde_json::json!({"terminal": true})),
|
||||
None,
|
||||
1_000,
|
||||
Some(1_001),
|
||||
Some(1_123),
|
||||
)
|
||||
.expect("existing candidate should build");
|
||||
|
||||
let merged = super::merge_candidate(
|
||||
UpsertRequestCandidateRecord {
|
||||
id: "candidate-late".to_string(),
|
||||
request_id: "request-1".to_string(),
|
||||
user_id: Some("user-1".to_string()),
|
||||
api_key_id: Some("key-1".to_string()),
|
||||
username: None,
|
||||
api_key_name: None,
|
||||
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: RequestCandidateStatus::Streaming,
|
||||
skip_reason: None,
|
||||
is_cached: Some(false),
|
||||
status_code: Some(200),
|
||||
error_type: None,
|
||||
error_message: None,
|
||||
latency_ms: Some(9_999),
|
||||
concurrent_requests: None,
|
||||
extra_data: Some(serde_json::json!({"late": true})),
|
||||
required_capabilities: None,
|
||||
created_at_unix_ms: Some(1_050),
|
||||
started_at_unix_ms: Some(1_051),
|
||||
finished_at_unix_ms: None,
|
||||
},
|
||||
Some(existing),
|
||||
)
|
||||
.expect("candidate should merge");
|
||||
|
||||
assert_eq!(merged.id, "candidate-1");
|
||||
assert_eq!(merged.status, RequestCandidateStatus::Success);
|
||||
assert_eq!(merged.latency_ms, Some(123));
|
||||
assert_eq!(merged.finished_at_unix_ms, Some(1_123));
|
||||
assert_eq!(
|
||||
merged.extra_data,
|
||||
Some(serde_json::json!({"terminal": true, "late": true}))
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
use crate::DataLayerError;
|
||||
|
||||
pub(crate) trait SqlResultExt<T> {
|
||||
fn map_sql_err(self) -> Result<T, DataLayerError>;
|
||||
}
|
||||
|
||||
impl<T> SqlResultExt<T> for Result<T, sqlx::Error> {
|
||||
fn map_sql_err(self) -> Result<T, DataLayerError> {
|
||||
self.map_err(DataLayerError::sql)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,381 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::gemini_file_mappings::{
|
||||
GeminiFileMappingListQuery, GeminiFileMappingMimeTypeCount, GeminiFileMappingReadRepository,
|
||||
GeminiFileMappingStats, GeminiFileMappingWriteRepository, StoredGeminiFileMapping,
|
||||
StoredGeminiFileMappingListPage, UpsertGeminiFileMappingRecord,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
use aether_data_query::{push_ci_contains_any, push_limit_offset, SqlDialect, WhereClause};
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
#[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");
|
||||
let mut where_clause = WhereClause::new();
|
||||
apply_list_filters(&mut builder, &mut where_clause, 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
|
||||
"#,
|
||||
);
|
||||
let mut where_clause = WhereClause::new();
|
||||
apply_list_filters(&mut builder, &mut where_clause, query);
|
||||
builder.push(" ORDER BY created_at DESC, file_name ASC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64::try_from(query.limit).unwrap_or(i64::MAX),
|
||||
i64::try_from(query.offset).unwrap_or(i64::MAX),
|
||||
);
|
||||
builder
|
||||
}
|
||||
|
||||
fn apply_list_filters(
|
||||
builder: &mut QueryBuilder<'_, MySql>,
|
||||
where_clause: &mut WhereClause,
|
||||
query: &GeminiFileMappingListQuery,
|
||||
) {
|
||||
if !query.include_expired {
|
||||
where_clause.push_next(builder);
|
||||
builder.push("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())
|
||||
{
|
||||
push_ci_contains_any(
|
||||
builder,
|
||||
where_clause,
|
||||
SqlDialect::MySql,
|
||||
&["file_name", "COALESCE(display_name, '')"],
|
||||
search,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
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::{build_list_count_query, build_list_rows_query, MysqlGeminiFileMappingRepository};
|
||||
use aether_data_contracts::repository::gemini_file_mappings::GeminiFileMappingListQuery;
|
||||
use sqlx::Execute;
|
||||
|
||||
#[test]
|
||||
fn list_query_uses_shared_mysql_filter_and_pagination_rendering() {
|
||||
let query = GeminiFileMappingListQuery {
|
||||
include_expired: false,
|
||||
search: Some(" Report ".to_string()),
|
||||
offset: 5,
|
||||
limit: 10,
|
||||
now_unix_secs: 123,
|
||||
};
|
||||
|
||||
let mut count = build_list_count_query(&query);
|
||||
let count_sql = count.build().sql().to_string();
|
||||
assert!(count_sql.contains(" WHERE expires_at > ? AND (LOWER(file_name) LIKE ?"));
|
||||
assert!(!count_sql.contains("WHERE 1=1"));
|
||||
|
||||
let mut rows = build_list_rows_query(&query);
|
||||
let rows_sql = rows.build().sql().to_string();
|
||||
assert!(rows_sql.contains("LOWER(COALESCE(display_name, '')) LIKE ?"));
|
||||
assert!(rows_sql.contains(" ORDER BY created_at DESC, file_name ASC LIMIT ? OFFSET ?"));
|
||||
}
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,915 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
|
||||
use aether_data_contracts::repository::global_models::{
|
||||
metadata_supports_embedding, AdminGlobalModelListQuery, AdminProviderModelListQuery,
|
||||
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelSnapshot,
|
||||
GlobalModelWriteRepository, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
|
||||
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
|
||||
StoredAdminProviderModel, StoredProviderActiveGlobalModel, StoredProviderModelStats,
|
||||
StoredPublicCatalogModel, StoredPublicGlobalModel, StoredPublicGlobalModelPage,
|
||||
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlGlobalModelReadRepository {
|
||||
pool: MysqlPool,
|
||||
}
|
||||
|
||||
impl MysqlGlobalModelReadRepository {
|
||||
pub fn new(pool: MysqlPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn load_snapshot(&self) -> Result<GlobalModelSnapshot, DataLayerError> {
|
||||
Ok(
|
||||
GlobalModelSnapshot::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();
|
||||
let usage_count =
|
||||
optional_admin_global_model_usage_count_i64(record.usage_count)?.unwrap_or_default();
|
||||
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 (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.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(usage_count)
|
||||
.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 usage_count = optional_admin_global_model_usage_count_i64(record.usage_count)?;
|
||||
let updated = sqlx::query(
|
||||
r#"
|
||||
UPDATE global_models
|
||||
SET
|
||||
display_name = ?,
|
||||
is_active = ?,
|
||||
default_price_per_request = ?,
|
||||
default_tiered_pricing = ?,
|
||||
supported_capabilities = ?,
|
||||
config = ?,
|
||||
usage_count = COALESCE(?, usage_count),
|
||||
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(usage_count)
|
||||
.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> {
|
||||
Ok(self.load_snapshot().await?.list_public_models(query))
|
||||
}
|
||||
|
||||
async fn get_public_model_by_name(
|
||||
&self,
|
||||
model_name: &str,
|
||||
) -> Result<Option<StoredPublicGlobalModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.get_public_model_by_name(model_name))
|
||||
}
|
||||
|
||||
async fn list_public_catalog_models(
|
||||
&self,
|
||||
query: &PublicCatalogModelListQuery,
|
||||
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.list_public_catalog_models(query))
|
||||
}
|
||||
|
||||
async fn search_public_catalog_models(
|
||||
&self,
|
||||
query: &PublicCatalogModelSearchQuery,
|
||||
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.search_public_catalog_models(query))
|
||||
}
|
||||
|
||||
async fn list_admin_global_models(
|
||||
&self,
|
||||
query: &AdminGlobalModelListQuery,
|
||||
) -> Result<StoredAdminGlobalModelPage, DataLayerError> {
|
||||
Ok(self.load_snapshot().await?.list_admin_global_models(query))
|
||||
}
|
||||
|
||||
async fn list_admin_provider_models(
|
||||
&self,
|
||||
query: &AdminProviderModelListQuery,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.list_admin_provider_models(query))
|
||||
}
|
||||
|
||||
async fn list_admin_provider_available_source_models(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.list_admin_provider_available_source_models(provider_id))
|
||||
}
|
||||
|
||||
async fn get_admin_provider_model(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
model_id: &str,
|
||||
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.get_admin_provider_model(provider_id, model_id))
|
||||
}
|
||||
|
||||
async fn get_admin_global_model_by_id(
|
||||
&self,
|
||||
global_model_id: &str,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.get_admin_global_model_by_id(global_model_id))
|
||||
}
|
||||
|
||||
async fn get_admin_global_model_by_name(
|
||||
&self,
|
||||
model_name: &str,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.get_admin_global_model_by_name(model_name))
|
||||
}
|
||||
|
||||
async fn list_admin_provider_models_by_global_model_id(
|
||||
&self,
|
||||
global_model_id: &str,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.list_admin_provider_models_by_global_model_id(global_model_id))
|
||||
}
|
||||
|
||||
async fn list_provider_model_stats(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderModelStats>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.list_provider_model_stats(provider_ids))
|
||||
}
|
||||
|
||||
async fn list_active_global_model_ids_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderActiveGlobalModel>, DataLayerError> {
|
||||
Ok(self
|
||||
.load_snapshot()
|
||||
.await?
|
||||
.list_active_global_model_ids_by_provider_ids(provider_ids))
|
||||
}
|
||||
}
|
||||
|
||||
#[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()?,
|
||||
)
|
||||
}
|
||||
|
||||
fn optional_admin_global_model_usage_count_i64(
|
||||
value: Option<u64>,
|
||||
) -> Result<Option<i64>, DataLayerError> {
|
||||
value
|
||||
.map(|value| {
|
||||
i64::try_from(value).map_err(|_| {
|
||||
DataLayerError::InvalidInput(
|
||||
"global_models.usage_count exceeds i64 range".to_string(),
|
||||
)
|
||||
})
|
||||
})
|
||||
.transpose()
|
||||
}
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
//! MySQL pool adapter primitives.
|
||||
|
||||
mod announcements;
|
||||
mod audit;
|
||||
mod auth;
|
||||
mod auth_modules;
|
||||
mod background_tasks;
|
||||
mod billing;
|
||||
mod candidate_selection;
|
||||
mod candidates;
|
||||
mod error;
|
||||
mod gemini_file_mappings;
|
||||
mod global_models;
|
||||
mod management_tokens;
|
||||
mod migrations;
|
||||
mod oauth_providers;
|
||||
mod pool;
|
||||
mod pool_scores;
|
||||
mod provider_catalog;
|
||||
mod proxy_nodes;
|
||||
mod quota;
|
||||
mod routing_profiles;
|
||||
mod settlement;
|
||||
mod usage;
|
||||
mod users;
|
||||
mod video_tasks;
|
||||
mod wallet;
|
||||
|
||||
pub use aether_data_contracts::{DataLayerError, DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
|
||||
pub use announcements::MysqlAnnouncementRepository;
|
||||
pub use audit::MysqlAuditLogReadRepository;
|
||||
pub use auth::MysqlAuthApiKeyReadRepository;
|
||||
pub use auth_modules::{MysqlAuthModuleReadRepository, MysqlAuthModuleRepository};
|
||||
pub use background_tasks::MysqlBackgroundTaskRepository;
|
||||
pub use billing::MysqlBillingReadRepository;
|
||||
pub use candidate_selection::MysqlMinimalCandidateSelectionReadRepository;
|
||||
pub use candidates::MysqlRequestCandidateRepository;
|
||||
pub use gemini_file_mappings::MysqlGeminiFileMappingRepository;
|
||||
pub use global_models::MysqlGlobalModelReadRepository;
|
||||
pub use management_tokens::MysqlManagementTokenRepository;
|
||||
pub use migrations::{pending_migrations, prepare_database_for_startup, run_migrations, MIGRATOR};
|
||||
pub use oauth_providers::MysqlOAuthProviderRepository;
|
||||
pub use pool::{MysqlPool, MysqlPoolConfig, MysqlPoolFactory};
|
||||
pub use pool_scores::MysqlPoolMemberScoreRepository;
|
||||
pub use provider_catalog::MysqlProviderCatalogReadRepository;
|
||||
pub use proxy_nodes::MysqlProxyNodeReadRepository;
|
||||
pub use quota::MysqlProviderQuotaRepository;
|
||||
pub use routing_profiles::MysqlRoutingGroupRepository;
|
||||
pub use settlement::MysqlSettlementRepository;
|
||||
pub use usage::{MysqlUsageStorage, MysqlUsageWriteRepository};
|
||||
pub use users::MysqlUserReadRepository;
|
||||
pub use video_tasks::MysqlVideoTaskRepository;
|
||||
pub use wallet::MysqlWalletReadRepository;
|
||||
@@ -0,0 +1,481 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::management_tokens::{
|
||||
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
|
||||
ManagementTokenWriteRepository, RegenerateManagementTokenSecret, StoredManagementToken,
|
||||
StoredManagementTokenListPage, StoredManagementTokenUserSummary, StoredManagementTokenWithUser,
|
||||
UpdateManagementTokenRecord,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_limit, push_limit_offset, push_optional_eq, WhereClause};
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
#[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 mut builder = QueryBuilder::<MySql>::new(TOKEN_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "id", token_id.to_string());
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_token_row).transpose()
|
||||
}
|
||||
}
|
||||
|
||||
const TOKEN_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
user_id,
|
||||
name,
|
||||
description,
|
||||
token_prefix,
|
||||
allowed_ips,
|
||||
permissions,
|
||||
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
|
||||
"#;
|
||||
|
||||
const TOKEN_WITH_USER_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
mt.id,
|
||||
mt.user_id,
|
||||
mt.name,
|
||||
mt.description,
|
||||
mt.token_prefix,
|
||||
mt.allowed_ips,
|
||||
mt.permissions,
|
||||
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
|
||||
"#;
|
||||
|
||||
#[async_trait]
|
||||
impl ManagementTokenReadRepository for MysqlManagementTokenRepository {
|
||||
async fn list_management_tokens(
|
||||
&self,
|
||||
query: &ManagementTokenListQuery,
|
||||
) -> Result<StoredManagementTokenListPage, DataLayerError> {
|
||||
let mut count_builder =
|
||||
QueryBuilder::<MySql>::new("SELECT COUNT(mt.id) AS total FROM management_tokens mt");
|
||||
let mut count_where = WhereClause::new();
|
||||
apply_management_token_filters(&mut count_builder, &mut count_where, query);
|
||||
let total = count_builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut list_builder = QueryBuilder::<MySql>::new(TOKEN_WITH_USER_COLUMNS);
|
||||
let mut list_where = WhereClause::new();
|
||||
apply_management_token_filters(&mut list_builder, &mut list_where, query);
|
||||
list_builder.push(" ORDER BY mt.created_at DESC, mt.id DESC");
|
||||
push_limit_offset(
|
||||
&mut list_builder,
|
||||
i64::try_from(query.limit).unwrap_or(i64::MAX),
|
||||
i64::try_from(query.offset).unwrap_or(i64::MAX),
|
||||
);
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.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 mut builder = QueryBuilder::<MySql>::new(TOKEN_WITH_USER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"mt.id",
|
||||
token_id.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.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 mut builder = QueryBuilder::<MySql>::new(TOKEN_WITH_USER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"mt.token_hash",
|
||||
token_hash.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_token_with_user_row).transpose()
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_management_token_filters(
|
||||
builder: &mut QueryBuilder<'_, MySql>,
|
||||
where_clause: &mut WhereClause,
|
||||
query: &ManagementTokenListQuery,
|
||||
) {
|
||||
push_optional_eq(builder, where_clause, "mt.user_id", query.user_id.clone());
|
||||
push_optional_eq(builder, where_clause, "mt.is_active", query.is_active);
|
||||
}
|
||||
|
||||
#[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,
|
||||
permissions, 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(json_to_string(record.permissions.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(¤t.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 permissions = record.permissions.as_ref().or(current.permissions.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 = ?,
|
||||
permissions = ?,
|
||||
expires_at = ?,
|
||||
is_active = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(name)
|
||||
.bind(description)
|
||||
.bind(json_to_string(allowed_ips)?)
|
||||
.bind(json_to_string(permissions)?)
|
||||
.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_permissions(json_from_string(row.try_get("permissions").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);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
use sqlx::{
|
||||
migrate::{Migrate, MigrateError, Migrator},
|
||||
MySqlPool,
|
||||
};
|
||||
|
||||
use aether_data_contracts::PendingMigrationInfo;
|
||||
|
||||
pub static MIGRATOR: Migrator = sqlx::migrate!("./migrations");
|
||||
|
||||
pub async fn run_migrations(pool: &MySqlPool) -> Result<(), MigrateError> {
|
||||
MIGRATOR.run(pool).await
|
||||
}
|
||||
|
||||
pub async fn pending_migrations(
|
||||
pool: &MySqlPool,
|
||||
) -> Result<Vec<PendingMigrationInfo>, MigrateError> {
|
||||
let mut conn = pool.acquire().await?;
|
||||
let applied_migrations = match conn.list_applied_migrations().await {
|
||||
Ok(applied_migrations) => applied_migrations,
|
||||
Err(err) if is_missing_sqlx_migrations_table_error(&err) => Vec::new(),
|
||||
Err(err) => return Err(err),
|
||||
};
|
||||
Ok(pending_migrations_from_applied(&applied_migrations))
|
||||
}
|
||||
|
||||
pub async fn prepare_database_for_startup(
|
||||
pool: &MySqlPool,
|
||||
) -> Result<Vec<PendingMigrationInfo>, MigrateError> {
|
||||
pending_migrations(pool).await
|
||||
}
|
||||
|
||||
fn is_missing_sqlx_migrations_table_error(err: &MigrateError) -> bool {
|
||||
let message = err.to_string().to_ascii_lowercase();
|
||||
message.contains("_sqlx_migrations")
|
||||
&& (message.contains("no such table")
|
||||
|| message.contains("doesn't exist")
|
||||
|| message.contains("does not exist")
|
||||
|| message.contains("unknown table"))
|
||||
}
|
||||
|
||||
fn pending_migrations_from_applied(
|
||||
applied_migrations: &[sqlx::migrate::AppliedMigration],
|
||||
) -> Vec<PendingMigrationInfo> {
|
||||
let applied_versions = applied_migrations
|
||||
.iter()
|
||||
.map(|migration| migration.version)
|
||||
.collect::<std::collections::HashSet<_>>();
|
||||
MIGRATOR
|
||||
.iter()
|
||||
.filter(|migration| migration.migration_type.is_up_migration())
|
||||
.filter(|migration| !applied_versions.contains(&migration.version))
|
||||
.map(|migration| PendingMigrationInfo {
|
||||
version: migration.version,
|
||||
description: migration.description.to_string(),
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::MIGRATOR;
|
||||
|
||||
#[test]
|
||||
fn embeds_mysql_migration_sources() {
|
||||
let versions = MIGRATOR
|
||||
.iter()
|
||||
.map(|migration| migration.version)
|
||||
.collect::<Vec<_>>();
|
||||
assert!(!versions.is_empty());
|
||||
assert!(versions.windows(2).all(|pair| pair[0] < pair[1]));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,393 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
|
||||
use aether_data_contracts::repository::oauth_providers::{
|
||||
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
|
||||
UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
#[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,
|
||||
icon_url,
|
||||
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,
|
||||
icon_url,
|
||||
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,
|
||||
icon_url,
|
||||
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),
|
||||
icon_url = VALUES(icon_url),
|
||||
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.icon_url.as_deref())
|
||||
.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("icon_url").map_sql_err()?,
|
||||
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);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,93 @@
|
||||
use std::str::FromStr;
|
||||
use std::time::Duration;
|
||||
|
||||
use crate::{DataLayerError, DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
|
||||
use sqlx::mysql::{MySqlConnectOptions, MySqlPoolOptions, MySqlSslMode};
|
||||
use sqlx::MySqlPool as SqlxMysqlPool;
|
||||
|
||||
pub type MysqlPool = SqlxMysqlPool;
|
||||
pub type MysqlPoolConfig = SqlDatabaseConfig;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlPoolFactory {
|
||||
config: MysqlPoolConfig,
|
||||
}
|
||||
|
||||
impl MysqlPoolFactory {
|
||||
pub fn new(config: MysqlPoolConfig) -> Result<Self, DataLayerError> {
|
||||
if config.driver != DatabaseDriver::Mysql {
|
||||
return Err(DataLayerError::InvalidConfiguration(format!(
|
||||
"mysql pool requires mysql driver, got {}",
|
||||
config.driver
|
||||
)));
|
||||
}
|
||||
config.validate()?;
|
||||
Ok(Self { config })
|
||||
}
|
||||
|
||||
pub fn config(&self) -> &MysqlPoolConfig {
|
||||
&self.config
|
||||
}
|
||||
|
||||
pub fn connect_options(&self) -> Result<MySqlConnectOptions, DataLayerError> {
|
||||
let ssl_mode = if self.config.pool.require_ssl {
|
||||
MySqlSslMode::Required
|
||||
} else {
|
||||
MySqlSslMode::Preferred
|
||||
};
|
||||
MySqlConnectOptions::from_str(self.config.url.trim())
|
||||
.map(|options| {
|
||||
options
|
||||
.ssl_mode(ssl_mode)
|
||||
.statement_cache_capacity(self.config.pool.statement_cache_capacity)
|
||||
})
|
||||
.map_err(|err| {
|
||||
DataLayerError::InvalidConfiguration(format!("invalid mysql database url: {err}"))
|
||||
})
|
||||
}
|
||||
|
||||
pub fn connect_lazy(&self) -> Result<MysqlPool, DataLayerError> {
|
||||
let SqlPoolConfig {
|
||||
min_connections,
|
||||
max_connections,
|
||||
acquire_timeout_ms,
|
||||
idle_timeout_ms,
|
||||
max_lifetime_ms,
|
||||
..
|
||||
} = self.config.pool;
|
||||
|
||||
Ok(MySqlPoolOptions::new()
|
||||
.min_connections(min_connections)
|
||||
.max_connections(max_connections)
|
||||
.acquire_timeout(Duration::from_millis(acquire_timeout_ms))
|
||||
.idle_timeout(Duration::from_millis(idle_timeout_ms))
|
||||
.max_lifetime(Duration::from_millis(max_lifetime_ms))
|
||||
.connect_lazy_with(self.connect_options()?))
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::MysqlPoolFactory;
|
||||
use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
|
||||
|
||||
#[tokio::test]
|
||||
async fn factory_builds_lazy_pool_from_valid_config() {
|
||||
let config = SqlDatabaseConfig {
|
||||
driver: DatabaseDriver::Mysql,
|
||||
url: "mysql://user:pass@localhost:3306/aether".to_string(),
|
||||
pool: SqlPoolConfig {
|
||||
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,
|
||||
},
|
||||
};
|
||||
|
||||
let factory = MysqlPoolFactory::new(config).expect("factory should build");
|
||||
let _pool = factory.connect_lazy().expect("lazy pool should build");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,579 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::pool_scores::*;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::{DataLayerError, MysqlPool};
|
||||
|
||||
const SCORE_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
pool_kind,
|
||||
pool_id,
|
||||
member_kind,
|
||||
member_id,
|
||||
capability,
|
||||
scope_kind,
|
||||
scope_id,
|
||||
score,
|
||||
hard_state,
|
||||
score_version,
|
||||
score_reason,
|
||||
last_ranked_at,
|
||||
last_scheduled_at,
|
||||
last_success_at,
|
||||
last_failure_at,
|
||||
failure_count,
|
||||
last_probe_attempt_at,
|
||||
last_probe_success_at,
|
||||
last_probe_failure_at,
|
||||
probe_failure_count,
|
||||
probe_status,
|
||||
updated_at
|
||||
FROM pool_member_scores
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlPoolMemberScoreRepository {
|
||||
pool: MysqlPool,
|
||||
}
|
||||
|
||||
impl MysqlPoolMemberScoreRepository {
|
||||
pub fn new(pool: MysqlPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn find_scores_by_identity(
|
||||
&self,
|
||||
identity: &PoolMemberIdentity,
|
||||
scope: Option<&PoolScoreScope>,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(identity.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(identity.pool_id.clone())
|
||||
.push(" AND member_kind = ")
|
||||
.push_bind(identity.member_kind.clone())
|
||||
.push(" AND member_id = ")
|
||||
.push_bind(identity.member_id.clone());
|
||||
if let Some(scope) = scope {
|
||||
builder
|
||||
.push(" AND capability = ")
|
||||
.push_bind(scope.capability.clone())
|
||||
.push(" AND scope_kind = ")
|
||||
.push_bind(scope.scope_kind.clone());
|
||||
if let Some(scope_id) = &scope.scope_id {
|
||||
builder.push(" AND scope_id = ").push_bind(scope_id.clone());
|
||||
} else {
|
||||
builder.push(" AND scope_id IS NULL");
|
||||
}
|
||||
}
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl PoolScoreReadRepository for MysqlPoolMemberScoreRepository {
|
||||
async fn list_ranked_pool_members(
|
||||
&self,
|
||||
query: &ListRankedPoolMembersQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(query.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(query.pool_id.clone())
|
||||
.push(" AND capability = ")
|
||||
.push_bind(query.capability.clone())
|
||||
.push(" AND scope_kind = ")
|
||||
.push_bind(query.scope_kind.clone());
|
||||
if let Some(scope_id) = &query.scope_id {
|
||||
builder.push(" AND scope_id = ").push_bind(scope_id.clone());
|
||||
} else {
|
||||
builder.push(" AND scope_id IS NULL");
|
||||
}
|
||||
if !query.hard_states.is_empty() {
|
||||
builder.push(" AND hard_state IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for state in &query.hard_states {
|
||||
separated.push_bind(state.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
}
|
||||
if let Some(statuses) = &query.probe_statuses {
|
||||
if !statuses.is_empty() {
|
||||
builder.push(" AND probe_status IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for status in statuses {
|
||||
separated.push_bind(status.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
}
|
||||
}
|
||||
builder
|
||||
.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC")
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(query.limit.max(1), "pool score limit")?)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(i64_from_usize(query.offset, "pool score offset")?);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
|
||||
async fn list_pool_member_scores(
|
||||
&self,
|
||||
query: &ListPoolMemberScoresQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(query.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(query.pool_id.clone());
|
||||
if let Some(capability) = &query.capability {
|
||||
builder
|
||||
.push(" AND capability = ")
|
||||
.push_bind(capability.clone());
|
||||
}
|
||||
if let Some(scope_kind) = &query.scope_kind {
|
||||
builder
|
||||
.push(" AND scope_kind = ")
|
||||
.push_bind(scope_kind.clone());
|
||||
}
|
||||
if let Some(scope_id) = &query.scope_id {
|
||||
builder.push(" AND scope_id = ").push_bind(scope_id.clone());
|
||||
}
|
||||
if !query.hard_states.is_empty() {
|
||||
builder.push(" AND hard_state IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for state in &query.hard_states {
|
||||
separated.push_bind(state.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
}
|
||||
if let Some(statuses) = &query.probe_statuses {
|
||||
if !statuses.is_empty() {
|
||||
builder.push(" AND probe_status IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for status in statuses {
|
||||
separated.push_bind(status.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
}
|
||||
}
|
||||
builder
|
||||
.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC")
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(query.limit.max(1), "pool score limit")?)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(i64_from_usize(query.offset, "pool score offset")?);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
|
||||
async fn list_pool_member_probe_candidates(
|
||||
&self,
|
||||
query: &ListPoolMemberProbeCandidatesQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(query.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(query.pool_id.clone());
|
||||
if let Some(capability) = &query.capability {
|
||||
builder
|
||||
.push(" AND capability = ")
|
||||
.push_bind(capability.clone());
|
||||
}
|
||||
builder
|
||||
.push(" AND hard_state IN ('available','unknown','cooldown','quota_exhausted')")
|
||||
.push(" AND (probe_status IN ('never','failed','stale')")
|
||||
.push(" OR (probe_status = 'ok' AND (last_probe_success_at IS NULL OR last_probe_success_at <= ")
|
||||
.push_bind(i64_from_u64(
|
||||
query.stale_before_unix_secs,
|
||||
"pool probe stale_before_unix_secs",
|
||||
)?)
|
||||
.push("))")
|
||||
.push(" OR (probe_status = 'in_progress' AND (last_probe_attempt_at IS NULL OR last_probe_attempt_at <= ")
|
||||
.push_bind(i64_from_u64(
|
||||
query.stale_before_unix_secs,
|
||||
"pool probe stale_before_unix_secs",
|
||||
)?)
|
||||
.push(")))")
|
||||
.push(
|
||||
r#"
|
||||
ORDER BY
|
||||
CASE
|
||||
WHEN last_scheduled_at IS NOT NULL AND probe_status <> 'ok' THEN 0
|
||||
WHEN hard_state = 'quota_exhausted' THEN 1
|
||||
WHEN hard_state = 'unknown' THEN 2
|
||||
WHEN probe_status = 'stale' THEN 3
|
||||
ELSE 4
|
||||
END ASC,
|
||||
probe_failure_count DESC,
|
||||
COALESCE(last_probe_success_at, 0) ASC,
|
||||
COALESCE(last_scheduled_at, 0) DESC,
|
||||
member_id ASC
|
||||
"#,
|
||||
)
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(
|
||||
query.limit.max(1),
|
||||
"pool probe candidate limit",
|
||||
)?);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
|
||||
async fn get_pool_member_scores_by_ids(
|
||||
&self,
|
||||
query: &GetPoolMemberScoresByIdsQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
if query.ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
|
||||
builder.push(" WHERE id IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for id in &query.ids {
|
||||
separated.push_bind(id.clone());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl PoolMemberScoreWriteRepository for MysqlPoolMemberScoreRepository {
|
||||
async fn upsert_pool_member_score(
|
||||
&self,
|
||||
score: UpsertPoolMemberScore,
|
||||
) -> Result<StoredPoolMemberScore, DataLayerError> {
|
||||
score.validate()?;
|
||||
let stored = score.into_stored();
|
||||
let score_reason = serde_json::to_string(&stored.score_reason)
|
||||
.map_err(|err| DataLayerError::InvalidInput(err.to_string()))?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO pool_member_scores (
|
||||
id, pool_kind, pool_id, member_kind, member_id, capability, scope_kind, scope_id,
|
||||
score, hard_state, score_version, score_reason, last_ranked_at, last_scheduled_at,
|
||||
last_success_at, last_failure_at, failure_count, last_probe_attempt_at,
|
||||
last_probe_success_at, last_probe_failure_at, probe_failure_count, probe_status, updated_at
|
||||
) VALUES (
|
||||
?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?
|
||||
)
|
||||
ON DUPLICATE KEY UPDATE
|
||||
pool_kind = VALUES(pool_kind),
|
||||
pool_id = VALUES(pool_id),
|
||||
member_kind = VALUES(member_kind),
|
||||
member_id = VALUES(member_id),
|
||||
capability = VALUES(capability),
|
||||
scope_kind = VALUES(scope_kind),
|
||||
scope_id = VALUES(scope_id),
|
||||
score = VALUES(score),
|
||||
hard_state = VALUES(hard_state),
|
||||
score_version = VALUES(score_version),
|
||||
score_reason = VALUES(score_reason),
|
||||
last_ranked_at = VALUES(last_ranked_at),
|
||||
last_scheduled_at = COALESCE(VALUES(last_scheduled_at), last_scheduled_at),
|
||||
last_success_at = COALESCE(VALUES(last_success_at), last_success_at),
|
||||
last_failure_at = COALESCE(VALUES(last_failure_at), last_failure_at),
|
||||
failure_count = VALUES(failure_count),
|
||||
last_probe_attempt_at = COALESCE(VALUES(last_probe_attempt_at), last_probe_attempt_at),
|
||||
last_probe_success_at = COALESCE(VALUES(last_probe_success_at), last_probe_success_at),
|
||||
last_probe_failure_at = COALESCE(VALUES(last_probe_failure_at), last_probe_failure_at),
|
||||
probe_failure_count = VALUES(probe_failure_count),
|
||||
probe_status = VALUES(probe_status),
|
||||
updated_at = VALUES(updated_at)
|
||||
"#,
|
||||
)
|
||||
.bind(stored.id.as_str())
|
||||
.bind(stored.pool_kind.as_str())
|
||||
.bind(stored.pool_id.as_str())
|
||||
.bind(stored.member_kind.as_str())
|
||||
.bind(stored.member_id.as_str())
|
||||
.bind(stored.capability.as_str())
|
||||
.bind(stored.scope_kind.as_str())
|
||||
.bind(stored.scope_id.as_deref())
|
||||
.bind(stored.score)
|
||||
.bind(stored.hard_state.as_database())
|
||||
.bind(i64_from_u64(stored.score_version, "pool score version")?)
|
||||
.bind(score_reason)
|
||||
.bind(i64_opt_from_u64(
|
||||
stored.last_ranked_at,
|
||||
"pool score last_ranked_at",
|
||||
)?)
|
||||
.bind(i64_opt_from_u64(
|
||||
stored.last_scheduled_at,
|
||||
"pool score last_scheduled_at",
|
||||
)?)
|
||||
.bind(i64_opt_from_u64(
|
||||
stored.last_success_at,
|
||||
"pool score last_success_at",
|
||||
)?)
|
||||
.bind(i64_opt_from_u64(
|
||||
stored.last_failure_at,
|
||||
"pool score last_failure_at",
|
||||
)?)
|
||||
.bind(i64_from_u64(
|
||||
stored.failure_count,
|
||||
"pool score failure_count",
|
||||
)?)
|
||||
.bind(i64_opt_from_u64(
|
||||
stored.last_probe_attempt_at,
|
||||
"pool score last_probe_attempt_at",
|
||||
)?)
|
||||
.bind(i64_opt_from_u64(
|
||||
stored.last_probe_success_at,
|
||||
"pool score last_probe_success_at",
|
||||
)?)
|
||||
.bind(i64_opt_from_u64(
|
||||
stored.last_probe_failure_at,
|
||||
"pool score last_probe_failure_at",
|
||||
)?)
|
||||
.bind(i64_from_u64(
|
||||
stored.probe_failure_count,
|
||||
"pool score probe_failure_count",
|
||||
)?)
|
||||
.bind(stored.probe_status.as_database())
|
||||
.bind(i64_from_u64(stored.updated_at, "pool score updated_at")?)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(stored)
|
||||
}
|
||||
|
||||
async fn mark_pool_member_probe_in_progress(
|
||||
&self,
|
||||
attempt: PoolMemberProbeAttempt,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let rows = self
|
||||
.find_scores_by_identity(&attempt.identity, attempt.scope.as_ref())
|
||||
.await?;
|
||||
let count = rows.len();
|
||||
for mut row in rows {
|
||||
row.last_probe_attempt_at = Some(attempt.attempted_at);
|
||||
row.probe_status = PoolMemberProbeStatus::InProgress;
|
||||
row.score_reason =
|
||||
merge_score_reason_patch(row.score_reason, attempt.score_reason_patch.clone());
|
||||
row.updated_at = attempt.attempted_at;
|
||||
self.upsert_pool_member_score(upsert_from_stored(row))
|
||||
.await?;
|
||||
}
|
||||
Ok(count)
|
||||
}
|
||||
|
||||
async fn record_pool_member_probe_result(
|
||||
&self,
|
||||
result: PoolMemberProbeResult,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let rows = self
|
||||
.find_scores_by_identity(&result.identity, result.scope.as_ref())
|
||||
.await?;
|
||||
let count = rows.len();
|
||||
for mut row in rows {
|
||||
row.last_probe_attempt_at = Some(result.attempted_at);
|
||||
row.probe_status = result.probe_status;
|
||||
if result.succeeded {
|
||||
row.last_probe_success_at = Some(result.attempted_at);
|
||||
row.probe_failure_count = 0;
|
||||
} else {
|
||||
row.last_probe_failure_at = Some(result.attempted_at);
|
||||
row.probe_failure_count = row.probe_failure_count.saturating_add(1);
|
||||
}
|
||||
if let Some(hard_state) = result.hard_state {
|
||||
row.hard_state = hard_state;
|
||||
}
|
||||
row.score_reason =
|
||||
merge_score_reason_patch(row.score_reason, result.score_reason_patch.clone());
|
||||
row.updated_at = result.attempted_at;
|
||||
self.upsert_pool_member_score(upsert_from_stored(row))
|
||||
.await?;
|
||||
}
|
||||
Ok(count)
|
||||
}
|
||||
|
||||
async fn record_pool_member_schedule_feedback(
|
||||
&self,
|
||||
feedback: PoolMemberScheduleFeedback,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let rows = self
|
||||
.find_scores_by_identity(&feedback.identity, feedback.scope.as_ref())
|
||||
.await?;
|
||||
let count = rows.len();
|
||||
for mut row in rows {
|
||||
row.last_scheduled_at = Some(feedback.scheduled_at);
|
||||
match feedback.succeeded {
|
||||
Some(true) => row.last_success_at = Some(feedback.scheduled_at),
|
||||
Some(false) => {
|
||||
row.last_failure_at = Some(feedback.scheduled_at);
|
||||
row.failure_count = row.failure_count.saturating_add(1);
|
||||
}
|
||||
None => {}
|
||||
}
|
||||
if let Some(hard_state) = feedback.hard_state {
|
||||
row.hard_state = hard_state;
|
||||
}
|
||||
row.score = score_with_delta(row.score, feedback.score_delta);
|
||||
row.score_reason =
|
||||
merge_score_reason_patch(row.score_reason, feedback.score_reason_patch.clone());
|
||||
row.updated_at = feedback.scheduled_at;
|
||||
self.upsert_pool_member_score(upsert_from_stored(row))
|
||||
.await?;
|
||||
}
|
||||
Ok(count)
|
||||
}
|
||||
|
||||
async fn mark_pool_member_hard_state(
|
||||
&self,
|
||||
identity: &PoolMemberIdentity,
|
||||
scope: Option<&PoolScoreScope>,
|
||||
hard_state: PoolMemberHardState,
|
||||
updated_at: u64,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let rows = self.find_scores_by_identity(identity, scope).await?;
|
||||
let count = rows.len();
|
||||
for mut row in rows {
|
||||
row.hard_state = hard_state;
|
||||
row.updated_at = updated_at;
|
||||
self.upsert_pool_member_score(upsert_from_stored(row))
|
||||
.await?;
|
||||
}
|
||||
Ok(count)
|
||||
}
|
||||
|
||||
async fn delete_pool_member_scores_for_member(
|
||||
&self,
|
||||
identity: &PoolMemberIdentity,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let result = sqlx::query(
|
||||
r#"
|
||||
DELETE FROM pool_member_scores
|
||||
WHERE pool_kind = ? AND pool_id = ? AND member_kind = ? AND member_id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(identity.pool_kind.as_str())
|
||||
.bind(identity.pool_id.as_str())
|
||||
.bind(identity.member_kind.as_str())
|
||||
.bind(identity.member_id.as_str())
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(result.rows_affected() as usize)
|
||||
}
|
||||
}
|
||||
|
||||
fn map_score_row(row: &MySqlRow) -> Result<StoredPoolMemberScore, DataLayerError> {
|
||||
let score_reason_raw: String = row.try_get("score_reason").map_sql_err()?;
|
||||
Ok(StoredPoolMemberScore {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
pool_kind: row.try_get("pool_kind").map_sql_err()?,
|
||||
pool_id: row.try_get("pool_id").map_sql_err()?,
|
||||
member_kind: row.try_get("member_kind").map_sql_err()?,
|
||||
member_id: row.try_get("member_id").map_sql_err()?,
|
||||
capability: row.try_get("capability").map_sql_err()?,
|
||||
scope_kind: row.try_get("scope_kind").map_sql_err()?,
|
||||
scope_id: row.try_get("scope_id").map_sql_err()?,
|
||||
score: row.try_get("score").map_sql_err()?,
|
||||
hard_state: PoolMemberHardState::from_database(
|
||||
row.try_get::<String, _>("hard_state")
|
||||
.map_sql_err()?
|
||||
.as_str(),
|
||||
)?,
|
||||
score_version: u64_from_i64(
|
||||
row.try_get("score_version").map_sql_err()?,
|
||||
"pool_member_scores.score_version",
|
||||
)?,
|
||||
score_reason: serde_json::from_str(&score_reason_raw).unwrap_or(serde_json::Value::Null),
|
||||
last_ranked_at: u64_opt_from_i64(
|
||||
row.try_get("last_ranked_at").map_sql_err()?,
|
||||
"pool_member_scores.last_ranked_at",
|
||||
)?,
|
||||
last_scheduled_at: u64_opt_from_i64(
|
||||
row.try_get("last_scheduled_at").map_sql_err()?,
|
||||
"pool_member_scores.last_scheduled_at",
|
||||
)?,
|
||||
last_success_at: u64_opt_from_i64(
|
||||
row.try_get("last_success_at").map_sql_err()?,
|
||||
"pool_member_scores.last_success_at",
|
||||
)?,
|
||||
last_failure_at: u64_opt_from_i64(
|
||||
row.try_get("last_failure_at").map_sql_err()?,
|
||||
"pool_member_scores.last_failure_at",
|
||||
)?,
|
||||
failure_count: u64_from_i64(
|
||||
row.try_get("failure_count").map_sql_err()?,
|
||||
"pool_member_scores.failure_count",
|
||||
)?,
|
||||
last_probe_attempt_at: u64_opt_from_i64(
|
||||
row.try_get("last_probe_attempt_at").map_sql_err()?,
|
||||
"pool_member_scores.last_probe_attempt_at",
|
||||
)?,
|
||||
last_probe_success_at: u64_opt_from_i64(
|
||||
row.try_get("last_probe_success_at").map_sql_err()?,
|
||||
"pool_member_scores.last_probe_success_at",
|
||||
)?,
|
||||
last_probe_failure_at: u64_opt_from_i64(
|
||||
row.try_get("last_probe_failure_at").map_sql_err()?,
|
||||
"pool_member_scores.last_probe_failure_at",
|
||||
)?,
|
||||
probe_failure_count: u64_from_i64(
|
||||
row.try_get("probe_failure_count").map_sql_err()?,
|
||||
"pool_member_scores.probe_failure_count",
|
||||
)?,
|
||||
probe_status: PoolMemberProbeStatus::from_database(
|
||||
row.try_get::<String, _>("probe_status")
|
||||
.map_sql_err()?
|
||||
.as_str(),
|
||||
)?,
|
||||
updated_at: u64_from_i64(
|
||||
row.try_get("updated_at").map_sql_err()?,
|
||||
"pool_member_scores.updated_at",
|
||||
)?,
|
||||
})
|
||||
}
|
||||
|
||||
fn upsert_from_stored(score: StoredPoolMemberScore) -> UpsertPoolMemberScore {
|
||||
UpsertPoolMemberScore {
|
||||
id: score.id,
|
||||
identity: PoolMemberIdentity {
|
||||
pool_kind: score.pool_kind,
|
||||
pool_id: score.pool_id,
|
||||
member_kind: score.member_kind,
|
||||
member_id: score.member_id,
|
||||
},
|
||||
scope: PoolScoreScope {
|
||||
capability: score.capability,
|
||||
scope_kind: score.scope_kind,
|
||||
scope_id: score.scope_id,
|
||||
},
|
||||
score: score.score,
|
||||
hard_state: score.hard_state,
|
||||
score_version: score.score_version,
|
||||
score_reason: score.score_reason,
|
||||
last_ranked_at: score.last_ranked_at,
|
||||
last_scheduled_at: score.last_scheduled_at,
|
||||
last_success_at: score.last_success_at,
|
||||
last_failure_at: score.last_failure_at,
|
||||
failure_count: score.failure_count,
|
||||
last_probe_attempt_at: score.last_probe_attempt_at,
|
||||
last_probe_success_at: score.last_probe_success_at,
|
||||
last_probe_failure_at: score.last_probe_failure_at,
|
||||
probe_failure_count: score.probe_failure_count,
|
||||
probe_status: score.probe_status,
|
||||
updated_at: score.updated_at,
|
||||
}
|
||||
}
|
||||
|
||||
fn i64_from_usize(value: usize, field: &str) -> Result<i64, DataLayerError> {
|
||||
i64::try_from(value)
|
||||
.map_err(|_| DataLayerError::InvalidInput(format!("{field} exceeds signed 64-bit range")))
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,174 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, Row};
|
||||
|
||||
use aether_data_contracts::repository::quota::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
|
||||
};
|
||||
use aether_data_query::{DialectSql, SelectColumn, SelectQuery, SqlDialect};
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::{DataLayerError, MysqlPool};
|
||||
|
||||
fn quota_snapshot_select() -> SelectQuery<'static> {
|
||||
SelectQuery::new("providers").select_columns([
|
||||
SelectColumn::expr("id").alias("provider_id"),
|
||||
SelectColumn::expr(
|
||||
DialectSql::common("billing_type").with_postgres("CAST(billing_type AS TEXT)"),
|
||||
)
|
||||
.alias("billing_type"),
|
||||
SelectColumn::expr(
|
||||
DialectSql::dialect(
|
||||
"CAST(monthly_quota_usd AS DOUBLE PRECISION)",
|
||||
"CAST(monthly_quota_usd AS REAL)",
|
||||
)
|
||||
.with_mysql("CAST(monthly_quota_usd AS DOUBLE)"),
|
||||
)
|
||||
.alias("monthly_quota_usd"),
|
||||
SelectColumn::expr(
|
||||
DialectSql::dialect(
|
||||
"CAST(COALESCE(monthly_used_usd, 0) AS DOUBLE PRECISION)",
|
||||
"CAST(COALESCE(monthly_used_usd, 0) AS REAL)",
|
||||
)
|
||||
.with_mysql("CAST(COALESCE(monthly_used_usd, 0) AS DOUBLE)"),
|
||||
)
|
||||
.alias("monthly_used_usd"),
|
||||
SelectColumn::expr("quota_reset_day"),
|
||||
SelectColumn::expr(
|
||||
DialectSql::dialect(
|
||||
"CAST(EXTRACT(EPOCH FROM quota_last_reset_at) AS BIGINT)",
|
||||
"quota_last_reset_at",
|
||||
)
|
||||
.with_mysql("quota_last_reset_at"),
|
||||
)
|
||||
.alias("quota_last_reset_at_unix_secs"),
|
||||
SelectColumn::expr(
|
||||
DialectSql::dialect(
|
||||
"CAST(EXTRACT(EPOCH FROM quota_expires_at) AS BIGINT)",
|
||||
"quota_expires_at",
|
||||
)
|
||||
.with_mysql("quota_expires_at"),
|
||||
)
|
||||
.alias("quota_expires_at_unix_secs"),
|
||||
SelectColumn::expr("is_active"),
|
||||
])
|
||||
}
|
||||
|
||||
#[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 mut statement = quota_snapshot_select().statement::<MySql>(SqlDialect::MySql);
|
||||
statement.where_eq("id", provider_id.to_string()).limit(1);
|
||||
let row = statement
|
||||
.finish()
|
||||
.build()
|
||||
.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 statement = quota_snapshot_select().statement::<MySql>(SqlDialect::MySql);
|
||||
statement
|
||||
.where_in("id", provider_ids)
|
||||
.order_by_sql("id ASC");
|
||||
let rows = statement
|
||||
.finish()
|
||||
.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::{quota_snapshot_select, MysqlProviderQuotaRepository};
|
||||
use aether_data_query::SqlDialect;
|
||||
|
||||
#[test]
|
||||
fn quota_projection_renders_for_mysql() {
|
||||
let sql = quota_snapshot_select().render(SqlDialect::MySql);
|
||||
|
||||
assert!(sql.contains("id AS `provider_id`"));
|
||||
assert!(sql.contains("CAST(monthly_quota_usd AS DOUBLE) AS `monthly_quota_usd`"));
|
||||
assert!(sql.contains("quota_last_reset_at AS `quota_last_reset_at_unix_secs`"));
|
||||
}
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,414 @@
|
||||
use async_trait::async_trait;
|
||||
use serde_json::Value;
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
|
||||
use aether_data_contracts::repository::routing_profiles::*;
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::pool::MysqlPool;
|
||||
|
||||
const ROUTING_GROUP_SELECT: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
name,
|
||||
description,
|
||||
enabled,
|
||||
is_system_default,
|
||||
config_json,
|
||||
version,
|
||||
created_at,
|
||||
updated_at,
|
||||
published_at
|
||||
FROM routing_groups
|
||||
"#;
|
||||
|
||||
const ROUTING_GROUP_BINDING_SELECT: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
group_id,
|
||||
subject_type,
|
||||
subject_id,
|
||||
is_default,
|
||||
allow_explicit_select,
|
||||
created_at,
|
||||
updated_at
|
||||
FROM routing_group_bindings
|
||||
"#;
|
||||
|
||||
const ROUTING_GROUP_VERSION_SELECT: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
group_id,
|
||||
version,
|
||||
config_json,
|
||||
created_at,
|
||||
created_by
|
||||
FROM routing_group_versions
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct MysqlRoutingGroupRepository {
|
||||
pool: MysqlPool,
|
||||
}
|
||||
|
||||
impl MysqlRoutingGroupRepository {
|
||||
pub fn new(pool: MysqlPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn reload_group(&self, id: &str) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
|
||||
self.find_routing_group(RoutingGroupLookupKey::Id(id)).await
|
||||
}
|
||||
|
||||
async fn find_binding_by_id(
|
||||
&self,
|
||||
id: &str,
|
||||
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
|
||||
let row = sqlx::query(&format!(
|
||||
"{ROUTING_GROUP_BINDING_SELECT} WHERE id = ? LIMIT 1"
|
||||
))
|
||||
.bind(id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_binding_row).transpose()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl RoutingGroupReadRepository for MysqlRoutingGroupRepository {
|
||||
async fn list_routing_groups(&self) -> Result<Vec<StoredRoutingGroup>, DataLayerError> {
|
||||
let rows = sqlx::query(&format!("{ROUTING_GROUP_SELECT} ORDER BY name ASC, id ASC"))
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_group_row).collect()
|
||||
}
|
||||
|
||||
async fn find_routing_group(
|
||||
&self,
|
||||
lookup: RoutingGroupLookupKey<'_>,
|
||||
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
|
||||
let row = match lookup {
|
||||
RoutingGroupLookupKey::Id(id) => sqlx::query(&format!(
|
||||
"{ROUTING_GROUP_SELECT} WHERE id = ? LIMIT 1"
|
||||
))
|
||||
.bind(id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?,
|
||||
RoutingGroupLookupKey::Name(name) => sqlx::query(&format!(
|
||||
"{ROUTING_GROUP_SELECT} WHERE name = ? LIMIT 1"
|
||||
))
|
||||
.bind(name)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?,
|
||||
RoutingGroupLookupKey::SystemDefault => sqlx::query(&format!(
|
||||
"{ROUTING_GROUP_SELECT} WHERE is_system_default = 1 AND enabled = 1 ORDER BY updated_at DESC, id ASC LIMIT 1"
|
||||
))
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?,
|
||||
};
|
||||
row.as_ref().map(map_group_row).transpose()
|
||||
}
|
||||
|
||||
async fn list_routing_group_bindings(
|
||||
&self,
|
||||
query: &RoutingGroupBindingQuery,
|
||||
) -> Result<Vec<StoredRoutingGroupBinding>, DataLayerError> {
|
||||
let rows = sqlx::query(&format!(
|
||||
r#"
|
||||
{ROUTING_GROUP_BINDING_SELECT}
|
||||
WHERE (? IS NULL OR group_id = ?)
|
||||
AND (? IS NULL OR subject_type = ?)
|
||||
AND (? IS NULL OR subject_id = ?)
|
||||
ORDER BY created_at ASC, id ASC
|
||||
"#
|
||||
))
|
||||
.bind(query.group_id.as_deref())
|
||||
.bind(query.group_id.as_deref())
|
||||
.bind(query.subject_type.map(binding_subject_to_database))
|
||||
.bind(query.subject_type.map(binding_subject_to_database))
|
||||
.bind(query.subject_id.as_deref())
|
||||
.bind(query.subject_id.as_deref())
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_binding_row).collect()
|
||||
}
|
||||
|
||||
async fn list_routing_group_versions(
|
||||
&self,
|
||||
group_id: &str,
|
||||
) -> Result<Vec<StoredRoutingGroupVersion>, DataLayerError> {
|
||||
let rows = sqlx::query(&format!(
|
||||
"{ROUTING_GROUP_VERSION_SELECT} WHERE group_id = ? ORDER BY version DESC, created_at DESC, id ASC"
|
||||
))
|
||||
.bind(group_id)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_version_row).collect()
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl RoutingGroupWriteRepository for MysqlRoutingGroupRepository {
|
||||
async fn create_routing_group(
|
||||
&self,
|
||||
record: CreateRoutingGroupRecord,
|
||||
) -> Result<StoredRoutingGroup, DataLayerError> {
|
||||
let group = StoredRoutingGroup::new(record)?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO routing_groups (
|
||||
id, name, description, enabled, is_system_default, config_json,
|
||||
version, created_at, updated_at, published_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&group.id)
|
||||
.bind(&group.name)
|
||||
.bind(&group.description)
|
||||
.bind(group.enabled)
|
||||
.bind(group.is_system_default)
|
||||
.bind(json_to_string(
|
||||
&group.config_json,
|
||||
"routing_groups.config_json",
|
||||
)?)
|
||||
.bind(group.version)
|
||||
.bind(group.created_at)
|
||||
.bind(group.updated_at)
|
||||
.bind(group.published_at)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(group)
|
||||
}
|
||||
|
||||
async fn update_routing_group(
|
||||
&self,
|
||||
id: &str,
|
||||
patch: UpdateRoutingGroupRecord,
|
||||
) -> Result<Option<StoredRoutingGroup>, DataLayerError> {
|
||||
let Some(mut group) = self.reload_group(id).await? else {
|
||||
return Ok(None);
|
||||
};
|
||||
apply_group_patch(&mut group, patch)?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE routing_groups
|
||||
SET name = ?,
|
||||
description = ?,
|
||||
enabled = ?,
|
||||
is_system_default = ?,
|
||||
config_json = ?,
|
||||
version = ?,
|
||||
updated_at = ?,
|
||||
published_at = ?
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(&group.name)
|
||||
.bind(&group.description)
|
||||
.bind(group.enabled)
|
||||
.bind(group.is_system_default)
|
||||
.bind(json_to_string(
|
||||
&group.config_json,
|
||||
"routing_groups.config_json",
|
||||
)?)
|
||||
.bind(group.version)
|
||||
.bind(group.updated_at)
|
||||
.bind(group.published_at)
|
||||
.bind(id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(Some(group))
|
||||
}
|
||||
|
||||
async fn delete_routing_group(&self, id: &str) -> Result<bool, DataLayerError> {
|
||||
let mut tx = self.pool.begin().await.map_sql_err()?;
|
||||
sqlx::query("DELETE FROM routing_group_bindings WHERE group_id = ?")
|
||||
.bind(id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
sqlx::query("DELETE FROM routing_group_versions WHERE group_id = ?")
|
||||
.bind(id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let rows_affected = sqlx::query("DELETE FROM routing_groups WHERE id = ?")
|
||||
.bind(id)
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected();
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(rows_affected > 0)
|
||||
}
|
||||
|
||||
async fn create_routing_group_binding(
|
||||
&self,
|
||||
record: CreateRoutingGroupBindingRecord,
|
||||
) -> Result<StoredRoutingGroupBinding, DataLayerError> {
|
||||
let binding = StoredRoutingGroupBinding::new(record)?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO routing_group_bindings (
|
||||
id, group_id, subject_type, subject_id, is_default,
|
||||
allow_explicit_select, created_at, updated_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&binding.id)
|
||||
.bind(&binding.group_id)
|
||||
.bind(binding_subject_to_database(binding.subject_type))
|
||||
.bind(&binding.subject_id)
|
||||
.bind(binding.is_default)
|
||||
.bind(binding.allow_explicit_select)
|
||||
.bind(binding.created_at)
|
||||
.bind(binding.updated_at)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(binding)
|
||||
}
|
||||
|
||||
async fn delete_routing_group_binding(&self, id: &str) -> Result<bool, DataLayerError> {
|
||||
Ok(
|
||||
sqlx::query("DELETE FROM routing_group_bindings WHERE id = ?")
|
||||
.bind(id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.rows_affected()
|
||||
> 0,
|
||||
)
|
||||
}
|
||||
|
||||
async fn update_routing_group_binding(
|
||||
&self,
|
||||
id: &str,
|
||||
patch: UpdateRoutingGroupBindingRecord,
|
||||
) -> Result<Option<StoredRoutingGroupBinding>, DataLayerError> {
|
||||
let Some(mut binding) = self.find_binding_by_id(id).await? else {
|
||||
return Ok(None);
|
||||
};
|
||||
apply_binding_patch(&mut binding, patch)?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
UPDATE routing_group_bindings
|
||||
SET group_id = ?,
|
||||
subject_type = ?,
|
||||
subject_id = ?,
|
||||
is_default = ?,
|
||||
allow_explicit_select = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
"#,
|
||||
)
|
||||
.bind(&binding.group_id)
|
||||
.bind(binding_subject_to_database(binding.subject_type))
|
||||
.bind(&binding.subject_id)
|
||||
.bind(binding.is_default)
|
||||
.bind(binding.allow_explicit_select)
|
||||
.bind(binding.updated_at)
|
||||
.bind(id)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(Some(binding))
|
||||
}
|
||||
|
||||
async fn create_routing_group_version(
|
||||
&self,
|
||||
record: CreateRoutingGroupVersionRecord,
|
||||
) -> Result<StoredRoutingGroupVersion, DataLayerError> {
|
||||
let version = StoredRoutingGroupVersion::new(record)?;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO routing_group_versions (
|
||||
id, group_id, version, config_json, created_at, created_by
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&version.id)
|
||||
.bind(&version.group_id)
|
||||
.bind(version.version)
|
||||
.bind(json_to_string(
|
||||
&version.config_json,
|
||||
"routing_group_versions.config_json",
|
||||
)?)
|
||||
.bind(version.created_at)
|
||||
.bind(&version.created_by)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(version)
|
||||
}
|
||||
}
|
||||
|
||||
fn map_group_row(row: &MySqlRow) -> Result<StoredRoutingGroup, DataLayerError> {
|
||||
Ok(StoredRoutingGroup {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
name: row.try_get("name").map_sql_err()?,
|
||||
description: row.try_get("description").map_sql_err()?,
|
||||
enabled: row.try_get("enabled").map_sql_err()?,
|
||||
is_system_default: row.try_get("is_system_default").map_sql_err()?,
|
||||
config_json: json_from_string(
|
||||
row.try_get("config_json").map_sql_err()?,
|
||||
"routing_groups.config_json",
|
||||
)?,
|
||||
version: row.try_get("version").map_sql_err()?,
|
||||
created_at: row.try_get("created_at").map_sql_err()?,
|
||||
updated_at: row.try_get("updated_at").map_sql_err()?,
|
||||
published_at: row.try_get("published_at").map_sql_err()?,
|
||||
})
|
||||
}
|
||||
|
||||
fn map_binding_row(row: &MySqlRow) -> Result<StoredRoutingGroupBinding, DataLayerError> {
|
||||
Ok(StoredRoutingGroupBinding {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
group_id: row.try_get("group_id").map_sql_err()?,
|
||||
subject_type: binding_subject_from_database(row.try_get("subject_type").map_sql_err()?)?,
|
||||
subject_id: row.try_get("subject_id").map_sql_err()?,
|
||||
is_default: row.try_get("is_default").map_sql_err()?,
|
||||
allow_explicit_select: row.try_get("allow_explicit_select").map_sql_err()?,
|
||||
created_at: row.try_get("created_at").map_sql_err()?,
|
||||
updated_at: row.try_get("updated_at").map_sql_err()?,
|
||||
})
|
||||
}
|
||||
|
||||
fn map_version_row(row: &MySqlRow) -> Result<StoredRoutingGroupVersion, DataLayerError> {
|
||||
Ok(StoredRoutingGroupVersion {
|
||||
id: row.try_get("id").map_sql_err()?,
|
||||
group_id: row.try_get("group_id").map_sql_err()?,
|
||||
version: row.try_get("version").map_sql_err()?,
|
||||
config_json: json_from_string(
|
||||
row.try_get("config_json").map_sql_err()?,
|
||||
"routing_group_versions.config_json",
|
||||
)?,
|
||||
created_at: row.try_get("created_at").map_sql_err()?,
|
||||
created_by: row.try_get("created_by").map_sql_err()?,
|
||||
})
|
||||
}
|
||||
|
||||
fn json_to_string(value: &Value, field_name: &str) -> Result<String, DataLayerError> {
|
||||
serde_json::to_string(value).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("{field_name} contains unserializable JSON: {err}"))
|
||||
})
|
||||
}
|
||||
|
||||
fn json_from_string(value: String, field_name: &str) -> Result<Value, DataLayerError> {
|
||||
serde_json::from_str(&value).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!("{field_name} contains invalid JSON: {err}"))
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,697 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, Row};
|
||||
|
||||
use aether_data_contracts::repository::settlement::{
|
||||
finite_wallet_available_usd, plan_finite_wallet_debit, settlement_billable_cost_usd,
|
||||
settlement_billing_status_for_usage_status, SettlementWriteRepository, StoredUsageSettlement,
|
||||
UsageSettlementInput, SETTLEMENT_EPSILON_USD,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
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()))
|
||||
}
|
||||
|
||||
#[derive(Debug, Default)]
|
||||
struct DailyQuotaDebitResult {
|
||||
debited_usd: f64,
|
||||
insufficient: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct DailyQuotaGrant {
|
||||
entitlement_id: String,
|
||||
daily_quota_usd: f64,
|
||||
usage_date: String,
|
||||
allow_wallet_overage: bool,
|
||||
}
|
||||
|
||||
fn daily_quota_usage_date(
|
||||
reset_timezone: Option<&str>,
|
||||
now: chrono::DateTime<chrono::Utc>,
|
||||
) -> Result<String, DataLayerError> {
|
||||
let timezone = reset_timezone
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.unwrap_or("Asia/Shanghai")
|
||||
.parse::<chrono_tz::Tz>()
|
||||
.map_err(|err| DataLayerError::InvalidInput(format!("invalid reset_timezone: {err}")))?;
|
||||
Ok(now.with_timezone(&timezone).date_naive().to_string())
|
||||
}
|
||||
|
||||
fn daily_quota_grants_from_entitlement(
|
||||
entitlement_id: &str,
|
||||
entitlements: &serde_json::Value,
|
||||
now: chrono::DateTime<chrono::Utc>,
|
||||
) -> Result<Vec<DailyQuotaGrant>, DataLayerError> {
|
||||
let mut grants = Vec::new();
|
||||
let Some(items) = entitlements.as_array() else {
|
||||
return Ok(grants);
|
||||
};
|
||||
for item in items {
|
||||
if item.get("type").and_then(serde_json::Value::as_str) != Some("daily_quota") {
|
||||
continue;
|
||||
}
|
||||
let daily_quota_usd = item
|
||||
.get("daily_quota_usd")
|
||||
.and_then(serde_json::Value::as_f64)
|
||||
.unwrap_or(0.0);
|
||||
if !daily_quota_usd.is_finite() || daily_quota_usd <= 0.0 {
|
||||
continue;
|
||||
}
|
||||
grants.push(DailyQuotaGrant {
|
||||
entitlement_id: entitlement_id.to_string(),
|
||||
daily_quota_usd,
|
||||
usage_date: daily_quota_usage_date(
|
||||
item.get("reset_timezone")
|
||||
.and_then(serde_json::Value::as_str),
|
||||
now,
|
||||
)?,
|
||||
allow_wallet_overage: item
|
||||
.get("allow_wallet_overage")
|
||||
.and_then(serde_json::Value::as_bool)
|
||||
.unwrap_or(false),
|
||||
});
|
||||
}
|
||||
Ok(grants)
|
||||
}
|
||||
|
||||
async fn consume_daily_quota_mysql(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
user_id: &str,
|
||||
request_id: &str,
|
||||
total_cost_usd: f64,
|
||||
wallet_available_usd: Option<f64>,
|
||||
wallet_can_overdraft: bool,
|
||||
now_unix_secs: i64,
|
||||
) -> Result<DailyQuotaDebitResult, DataLayerError> {
|
||||
if total_cost_usd <= 0.0 {
|
||||
return Ok(DailyQuotaDebitResult::default());
|
||||
}
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT id, entitlements_snapshot
|
||||
FROM user_plan_entitlements
|
||||
WHERE user_id = ?
|
||||
AND status = 'active'
|
||||
AND starts_at <= ?
|
||||
AND expires_at > ?
|
||||
ORDER BY expires_at ASC, created_at ASC, id ASC
|
||||
FOR UPDATE
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.bind(now_unix_secs)
|
||||
.bind(now_unix_secs)
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let now = chrono::Utc::now();
|
||||
let mut grants = Vec::new();
|
||||
for row in rows {
|
||||
let entitlement_id: String = row.try_get("id").map_sql_err()?;
|
||||
let entitlements_raw: String = row.try_get("entitlements_snapshot").map_sql_err()?;
|
||||
let entitlements =
|
||||
serde_json::from_str::<serde_json::Value>(&entitlements_raw).map_err(|err| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"user_plan_entitlements.entitlements_snapshot invalid json: {err}"
|
||||
))
|
||||
})?;
|
||||
grants.extend(daily_quota_grants_from_entitlement(
|
||||
&entitlement_id,
|
||||
&entitlements,
|
||||
now,
|
||||
)?);
|
||||
}
|
||||
if grants.is_empty() {
|
||||
return Ok(DailyQuotaDebitResult::default());
|
||||
}
|
||||
|
||||
let mut grants_with_remaining = Vec::new();
|
||||
let mut total_remaining = 0.0;
|
||||
let mut allow_wallet_overage = true;
|
||||
for grant in grants {
|
||||
allow_wallet_overage &= grant.allow_wallet_overage;
|
||||
let used = sqlx::query_scalar::<_, f64>(
|
||||
r#"
|
||||
SELECT COALESCE(SUM(amount_usd), 0)
|
||||
FROM entitlement_usage_ledgers
|
||||
WHERE user_entitlement_id = ?
|
||||
AND usage_date = ?
|
||||
"#,
|
||||
)
|
||||
.bind(&grant.entitlement_id)
|
||||
.bind(&grant.usage_date)
|
||||
.fetch_one(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let remaining = (grant.daily_quota_usd - used).max(0.0);
|
||||
total_remaining += remaining;
|
||||
grants_with_remaining.push((grant, remaining));
|
||||
}
|
||||
if !allow_wallet_overage && total_remaining + 0.000_000_01 < total_cost_usd {
|
||||
return Ok(DailyQuotaDebitResult {
|
||||
debited_usd: 0.0,
|
||||
insufficient: true,
|
||||
});
|
||||
}
|
||||
if allow_wallet_overage
|
||||
&& !wallet_can_overdraft
|
||||
&& wallet_available_usd.is_some_and(|available| {
|
||||
total_remaining + available + SETTLEMENT_EPSILON_USD < total_cost_usd
|
||||
})
|
||||
{
|
||||
return Ok(DailyQuotaDebitResult {
|
||||
debited_usd: 0.0,
|
||||
insufficient: true,
|
||||
});
|
||||
}
|
||||
|
||||
let mut remaining_cost = total_cost_usd;
|
||||
let mut debited = 0.0;
|
||||
for (grant, balance_before) in grants_with_remaining {
|
||||
if remaining_cost <= 0.000_000_01 || balance_before <= 0.0 {
|
||||
continue;
|
||||
}
|
||||
let amount = remaining_cost.min(balance_before);
|
||||
let balance_after = balance_before - amount;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT IGNORE INTO entitlement_usage_ledgers (
|
||||
id, user_entitlement_id, user_id, request_id, amount_usd,
|
||||
balance_before, balance_after, usage_date, created_at
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
"#,
|
||||
)
|
||||
.bind(uuid::Uuid::new_v4().to_string())
|
||||
.bind(&grant.entitlement_id)
|
||||
.bind(user_id)
|
||||
.bind(request_id)
|
||||
.bind(amount)
|
||||
.bind(balance_before)
|
||||
.bind(balance_after)
|
||||
.bind(&grant.usage_date)
|
||||
.bind(now_unix_secs)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
remaining_cost -= amount;
|
||||
debited += amount;
|
||||
}
|
||||
Ok(DailyQuotaDebitResult {
|
||||
debited_usd: debited,
|
||||
insufficient: false,
|
||||
})
|
||||
}
|
||||
|
||||
#[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 matches!(
|
||||
current_billing_status.as_str(),
|
||||
"settled" | "void" | "insufficient_quota"
|
||||
) {
|
||||
let settlement = settlement_from_row(&usage_row)?;
|
||||
tx.commit().await.map_sql_err()?;
|
||||
return Ok(Some(settlement));
|
||||
}
|
||||
|
||||
let mut final_billing_status =
|
||||
settlement_billing_status_for_usage_status(&input.status).to_string();
|
||||
let mut settlement = StoredUsageSettlement {
|
||||
request_id: input.request_id.clone(),
|
||||
wallet_id: None,
|
||||
billing_status: final_billing_status.clone(),
|
||||
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
|
||||
};
|
||||
|
||||
let wallet_can_overdraft = wallet_row.is_some();
|
||||
let wallet_available_usd = match wallet_row.as_ref() {
|
||||
Some(row) => {
|
||||
let limit_mode: String = row.try_get("limit_mode").map_sql_err()?;
|
||||
if limit_mode.eq_ignore_ascii_case("unlimited") {
|
||||
None
|
||||
} else {
|
||||
Some(finite_wallet_available_usd(
|
||||
row.try_get("balance").map_sql_err()?,
|
||||
row.try_get("gift_balance").map_sql_err()?,
|
||||
))
|
||||
}
|
||||
}
|
||||
None => Some(0.0),
|
||||
};
|
||||
if let Some(row) = wallet_row.as_ref() {
|
||||
let wallet_id: String = row.try_get("id").map_sql_err()?;
|
||||
let before_recharge: f64 = row.try_get("balance").map_sql_err()?;
|
||||
let before_gift: f64 = row.try_get("gift_balance").map_sql_err()?;
|
||||
let before_total = before_recharge + before_gift;
|
||||
settlement.wallet_id = Some(wallet_id);
|
||||
settlement.wallet_balance_before = Some(before_total);
|
||||
settlement.wallet_balance_after = Some(before_total);
|
||||
settlement.wallet_recharge_balance_before = Some(before_recharge);
|
||||
settlement.wallet_recharge_balance_after = Some(before_recharge);
|
||||
settlement.wallet_gift_balance_before = Some(before_gift);
|
||||
settlement.wallet_gift_balance_after = Some(before_gift);
|
||||
}
|
||||
|
||||
let billable_cost_usd = settlement_billable_cost_usd(&input);
|
||||
let wallet_debit_cost_usd = if !api_key_is_standalone {
|
||||
if let Some(user_id) = input.user_id.as_deref().filter(|value| !value.is_empty()) {
|
||||
let quota = consume_daily_quota_mysql(
|
||||
&mut tx,
|
||||
user_id,
|
||||
&input.request_id,
|
||||
billable_cost_usd,
|
||||
wallet_available_usd,
|
||||
wallet_can_overdraft,
|
||||
updated_at,
|
||||
)
|
||||
.await?;
|
||||
if quota.insufficient {
|
||||
final_billing_status = "insufficient_quota".to_string();
|
||||
settlement.billing_status = final_billing_status.clone();
|
||||
0.0
|
||||
} else {
|
||||
(billable_cost_usd - quota.debited_usd).max(0.0)
|
||||
}
|
||||
} else {
|
||||
billable_cost_usd
|
||||
}
|
||||
} else {
|
||||
billable_cost_usd
|
||||
};
|
||||
if final_billing_status != "settled" {
|
||||
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()?;
|
||||
return Ok(Some(settlement));
|
||||
}
|
||||
|
||||
if wallet_debit_cost_usd > SETTLEMENT_EPSILON_USD {
|
||||
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 debit_plan = plan_finite_wallet_debit(
|
||||
before_recharge,
|
||||
before_gift,
|
||||
wallet_debit_cost_usd,
|
||||
);
|
||||
(after_recharge, after_gift) =
|
||||
debit_plan.after_balances(before_recharge, before_gift);
|
||||
}
|
||||
if final_billing_status == "settled" {
|
||||
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(wallet_debit_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);
|
||||
} else {
|
||||
final_billing_status = "insufficient_quota".to_string();
|
||||
settlement.billing_status = final_billing_status.clone();
|
||||
}
|
||||
}
|
||||
|
||||
if final_billing_status != "settled" {
|
||||
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()?;
|
||||
return Ok(Some(settlement));
|
||||
}
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,368 @@
|
||||
use super::{MysqlUsageStorage, MysqlUsageWriteRepository};
|
||||
use crate::run_migrations;
|
||||
use aether_data_contracts::repository::usage::{UpsertUsageRecord, UsageWriteRepository};
|
||||
|
||||
#[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 = MysqlUsageWriteRepository::new(pool);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_usage_daily_heatmap_reads_imported_daily_aggregates() {
|
||||
let source = include_str!("../usage.rs");
|
||||
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("AS SIGNED) AS total_tokens"));
|
||||
assert!(source.contains("CAST(COUNT(*) AS SIGNED) AS requests"));
|
||||
assert!(source.contains("summaries.entry(item.date.clone()).or_insert(item)"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_usage_totals_by_user_ids_reads_imported_user_daily_aggregates() {
|
||||
let source = include_str!("../usage.rs");
|
||||
assert!(source.contains("async fn summarize_usage_totals_by_user_ids"));
|
||||
assert!(source.contains("FROM stats_user_daily"));
|
||||
assert!(source.contains("MAX(`date`) AS latest_date"));
|
||||
assert!(source.contains("AS SIGNED) AS request_count"));
|
||||
assert!(source.contains("CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS request_count"));
|
||||
assert!(source.contains("requested.cutoff_unix_secs"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_dashboard_reads_imported_daily_aggregates() {
|
||||
let source = include_str!("../usage.rs");
|
||||
assert!(source.contains("summarize_dashboard_usage_from_daily_aggregates"));
|
||||
assert!(source.contains("list_dashboard_daily_breakdown_from_daily_aggregates"));
|
||||
assert!(source.contains("FROM stats_daily"));
|
||||
assert!(source.contains("FROM stats_user_daily"));
|
||||
assert!(source.contains("'aggregate' AS model"));
|
||||
assert!(source.contains("AS SIGNED) AS total_requests"));
|
||||
assert!(source.contains("CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS total_requests"));
|
||||
assert!(source.contains("CAST(COALESCE(SUM(total_requests), 0) AS SIGNED) AS requests"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_usage_upsert_keeps_terminal_state_when_streaming_arrives_late() {
|
||||
assert!(super::UPSERT_USAGE_SQL.contains(
|
||||
"status IN ('completed', 'failed', 'cancelled') AND VALUES(status) IN ('pending', 'streaming')"
|
||||
));
|
||||
assert!(super::UPSERT_USAGE_SQL.contains("input_tokens = CASE"));
|
||||
assert!(super::UPSERT_USAGE_SQL.contains("status_code = CASE"));
|
||||
assert!(super::UPSERT_USAGE_SQL.contains("billing_status = CASE"));
|
||||
assert!(super::UPSERT_USAGE_SQL.contains("finalized_at = CASE"));
|
||||
assert!(super::UPSERT_USAGE_SQL.contains("updated_at_unix_secs = CASE"));
|
||||
assert!(super::UPSERT_USAGE_SQL
|
||||
.contains("WHEN status = 'streaming' AND VALUES(status) = 'pending' THEN status"));
|
||||
assert!(super::UPSERT_USAGE_SQL.contains(
|
||||
"WHEN status = 'streaming' AND VALUES(status) = 'streaming' AND VALUES(status_code) IS NULL THEN status_code"
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_usage_write_repository_upserts_when_url_is_set() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
else {
|
||||
eprintln!("skipping mysql usage write smoke test because AETHER_TEST_MYSQL_URL is unset");
|
||||
return;
|
||||
};
|
||||
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect(&database_url)
|
||||
.await
|
||||
.expect("mysql test pool should connect");
|
||||
run_migrations(&pool)
|
||||
.await
|
||||
.expect("mysql migrations should run");
|
||||
|
||||
let suffix = unique_suffix();
|
||||
let user_id = format!("user-{suffix}");
|
||||
let api_key_id = format!("api-key-{suffix}");
|
||||
let provider_id = format!("provider-{suffix}");
|
||||
let provider_key_id = format!("provider-key-{suffix}");
|
||||
seed_stats_targets(&pool, &user_id, &api_key_id, &provider_id, &provider_key_id).await;
|
||||
|
||||
let repository = MysqlUsageWriteRepository::new(pool.clone());
|
||||
let record = repository
|
||||
.upsert(sample_usage(
|
||||
&format!("request-{suffix}"),
|
||||
&user_id,
|
||||
&api_key_id,
|
||||
&provider_id,
|
||||
&provider_key_id,
|
||||
"completed",
|
||||
"pending",
|
||||
1_000,
|
||||
))
|
||||
.await
|
||||
.expect("usage should upsert");
|
||||
|
||||
assert_eq!(record.api_key_id.as_deref(), Some(api_key_id.as_str()));
|
||||
assert_eq!(
|
||||
record.provider_api_key_id.as_deref(),
|
||||
Some(provider_key_id.as_str())
|
||||
);
|
||||
assert_eq!(record.total_tokens, 7);
|
||||
assert_eq!(
|
||||
record.request_metadata.as_ref().unwrap()["upstream_is_stream"],
|
||||
true
|
||||
);
|
||||
let upstream_is_stream: Option<bool> =
|
||||
sqlx::query_scalar("SELECT upstream_is_stream FROM `usage` WHERE request_id = ?")
|
||||
.bind(format!("request-{suffix}"))
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("usage stream mode should load");
|
||||
assert_eq!(upstream_is_stream, Some(true));
|
||||
|
||||
let stats = sqlx::query_as::<_, (i64, i64, f64, Option<i64>)>(
|
||||
"SELECT total_requests, total_tokens, total_cost_usd, last_used_at FROM api_keys WHERE id = ?",
|
||||
)
|
||||
.bind(&api_key_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("api key stats should load");
|
||||
assert_eq!(stats, (1, 7, 0.5, Some(1_000)));
|
||||
|
||||
let provider_stats = sqlx::query_as::<_, (i64, i64, i64, i64, f64, i64, Option<i64>)>(
|
||||
"SELECT request_count, success_count, error_count, total_tokens, total_cost_usd, total_response_time_ms, last_used_at FROM provider_api_keys WHERE id = ?",
|
||||
)
|
||||
.bind(&provider_key_id)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("provider key stats should load");
|
||||
assert_eq!(provider_stats, (1, 1, 0, 7, 0.5, 42, Some(1_000)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_usage_read_repository_reads_usage_contract_views_when_url_is_set() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
else {
|
||||
eprintln!("skipping mysql usage read smoke test because AETHER_TEST_MYSQL_URL is unset");
|
||||
return;
|
||||
};
|
||||
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect(&database_url)
|
||||
.await
|
||||
.expect("mysql test pool should connect");
|
||||
run_migrations(&pool)
|
||||
.await
|
||||
.expect("mysql migrations should run");
|
||||
|
||||
let suffix = unique_suffix();
|
||||
let user_id = format!("user-read-{suffix}");
|
||||
let api_key_id = format!("api-key-read-{suffix}");
|
||||
let provider_id = format!("provider-read-{suffix}");
|
||||
let provider_key_id = format!("provider-key-read-{suffix}");
|
||||
seed_stats_targets(&pool, &user_id, &api_key_id, &provider_id, &provider_key_id).await;
|
||||
|
||||
let writer = MysqlUsageWriteRepository::new(pool.clone());
|
||||
writer
|
||||
.upsert(sample_usage(
|
||||
&format!("request-read-1-{suffix}"),
|
||||
&user_id,
|
||||
&api_key_id,
|
||||
&provider_id,
|
||||
&provider_key_id,
|
||||
"completed",
|
||||
"settled",
|
||||
1_000,
|
||||
))
|
||||
.await
|
||||
.expect("usage should upsert");
|
||||
writer
|
||||
.upsert(sample_usage(
|
||||
&format!("request-read-2-{suffix}"),
|
||||
&user_id,
|
||||
&api_key_id,
|
||||
&provider_id,
|
||||
&provider_key_id,
|
||||
"failed",
|
||||
"void",
|
||||
1_010,
|
||||
))
|
||||
.await
|
||||
.expect("usage should upsert");
|
||||
|
||||
let reader = MysqlUsageStorage::new(pool);
|
||||
let records = reader
|
||||
.load_usage_records()
|
||||
.await
|
||||
.expect("usage records should load");
|
||||
let loaded = records
|
||||
.iter()
|
||||
.find(|item| item.request_id == format!("request-read-1-{suffix}"))
|
||||
.expect("usage should exist");
|
||||
assert_eq!(loaded.total_tokens, 7);
|
||||
assert_eq!(loaded.billing_status, "settled");
|
||||
assert_eq!(
|
||||
records
|
||||
.iter()
|
||||
.filter(|item| item.user_id.as_deref() == Some(&user_id))
|
||||
.count(),
|
||||
2
|
||||
);
|
||||
}
|
||||
|
||||
async fn seed_stats_targets(
|
||||
pool: &sqlx::MySqlPool,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
provider_id: &str,
|
||||
provider_key_id: &str,
|
||||
) {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO users (id, auth_source, created_at, updated_at)
|
||||
VALUES (?, 'local', 1, 1)
|
||||
"#,
|
||||
)
|
||||
.bind(user_id)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("user should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO api_keys (id, user_id, key_hash, created_at, updated_at)
|
||||
VALUES (?, ?, ?, 1, 1)
|
||||
"#,
|
||||
)
|
||||
.bind(api_key_id)
|
||||
.bind(user_id)
|
||||
.bind(format!("hash-{api_key_id}"))
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("api key should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO providers (id, name, provider_type, created_at, updated_at)
|
||||
VALUES (?, ?, 'openai', 1, 1)
|
||||
"#,
|
||||
)
|
||||
.bind(provider_id)
|
||||
.bind(format!("Provider {provider_id}"))
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("provider should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO provider_api_keys (id, provider_id, name, created_at, updated_at)
|
||||
VALUES (?, ?, ?, 1, 1)
|
||||
"#,
|
||||
)
|
||||
.bind(provider_key_id)
|
||||
.bind(provider_id)
|
||||
.bind(format!("Provider Key {provider_key_id}"))
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("provider key should seed");
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn sample_usage(
|
||||
request_id: &str,
|
||||
user_id: &str,
|
||||
api_key_id: &str,
|
||||
provider_id: &str,
|
||||
provider_key_id: &str,
|
||||
status: &str,
|
||||
billing_status: &str,
|
||||
updated_at: u64,
|
||||
) -> UpsertUsageRecord {
|
||||
UpsertUsageRecord {
|
||||
request_id: request_id.to_string(),
|
||||
user_id: Some(user_id.to_string()),
|
||||
api_key_id: Some(api_key_id.to_string()),
|
||||
username: Some("legacy-user".to_string()),
|
||||
api_key_name: Some("legacy-key".to_string()),
|
||||
provider_name: "Provider One".to_string(),
|
||||
model: "model-1".to_string(),
|
||||
target_model: Some("target-model".to_string()),
|
||||
provider_id: Some(provider_id.to_string()),
|
||||
provider_endpoint_id: Some("endpoint-1".to_string()),
|
||||
provider_api_key_id: Some(provider_key_id.to_string()),
|
||||
request_type: Some("chat".to_string()),
|
||||
api_format: Some("openai".to_string()),
|
||||
api_family: Some("chat".to_string()),
|
||||
endpoint_kind: Some("chat".to_string()),
|
||||
endpoint_api_format: Some("openai".to_string()),
|
||||
provider_api_family: Some("chat".to_string()),
|
||||
provider_endpoint_kind: Some("chat".to_string()),
|
||||
has_format_conversion: Some(true),
|
||||
is_stream: Some(false),
|
||||
input_tokens: Some(2),
|
||||
output_tokens: Some(3),
|
||||
total_tokens: None,
|
||||
cache_creation_input_tokens: None,
|
||||
cache_creation_ephemeral_5m_input_tokens: Some(0),
|
||||
cache_creation_ephemeral_1h_input_tokens: Some(0),
|
||||
cache_read_input_tokens: Some(2),
|
||||
cache_creation_cost_usd: Some(0.0),
|
||||
cache_read_cost_usd: Some(0.1),
|
||||
output_price_per_1m: Some(2.0),
|
||||
total_cost_usd: Some(0.5),
|
||||
actual_total_cost_usd: Some(0.4),
|
||||
status_code: Some(200),
|
||||
error_message: None,
|
||||
error_category: None,
|
||||
response_time_ms: Some(42),
|
||||
first_byte_time_ms: Some(12),
|
||||
status: status.to_string(),
|
||||
billing_status: billing_status.to_string(),
|
||||
request_headers: None,
|
||||
request_body: None,
|
||||
request_body_ref: None,
|
||||
request_body_state: None,
|
||||
provider_request_headers: None,
|
||||
provider_request_body: None,
|
||||
provider_request_body_ref: None,
|
||||
provider_request_body_state: None,
|
||||
response_headers: None,
|
||||
response_body: None,
|
||||
response_body_ref: None,
|
||||
response_body_state: None,
|
||||
client_response_headers: None,
|
||||
client_response_body: None,
|
||||
client_response_body_ref: None,
|
||||
client_response_body_state: None,
|
||||
candidate_id: Some("candidate-1".to_string()),
|
||||
candidate_index: Some(1),
|
||||
key_name: Some("key-one".to_string()),
|
||||
planner_kind: Some("default".to_string()),
|
||||
route_family: Some("chat".to_string()),
|
||||
route_kind: Some("completion".to_string()),
|
||||
execution_path: Some("remote".to_string()),
|
||||
local_execution_runtime_miss_reason: None,
|
||||
request_metadata: Some(serde_json::json!({
|
||||
"trace_id": "trace-1",
|
||||
"upstream_is_stream": true,
|
||||
})),
|
||||
finalized_at_unix_secs: Some(updated_at),
|
||||
created_at_unix_ms: Some(updated_at),
|
||||
updated_at_unix_secs: updated_at,
|
||||
}
|
||||
}
|
||||
|
||||
fn unique_suffix() -> String {
|
||||
let nanos = std::time::SystemTime::now()
|
||||
.duration_since(std::time::UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_nanos();
|
||||
format!("{}-{nanos}", std::process::id())
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,690 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
|
||||
|
||||
use aether_data_contracts::repository::video_tasks::{
|
||||
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,
|
||||
VideoTaskQueryFilter, VideoTaskReadRepository, VideoTaskStatus, VideoTaskStatusCount,
|
||||
VideoTaskWriteRepository,
|
||||
};
|
||||
use aether_data_contracts::DataLayerError;
|
||||
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::MysqlPool;
|
||||
|
||||
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);
|
||||
}
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,267 @@
|
||||
use super::MysqlWalletReadRepository;
|
||||
use crate::run_migrations;
|
||||
use aether_data_contracts::repository::wallet::{
|
||||
AdminPaymentOrderListQuery, AdminRedeemCodeListQuery, AdminWalletListQuery, WalletLookupKey,
|
||||
WalletReadRepository,
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_wallet_read_repository_reads_wallet_contract_views() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
else {
|
||||
eprintln!("skipping mysql wallet read smoke test because AETHER_TEST_MYSQL_URL is unset");
|
||||
return;
|
||||
};
|
||||
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect(&database_url)
|
||||
.await
|
||||
.expect("mysql pool should connect");
|
||||
run_migrations(&pool)
|
||||
.await
|
||||
.expect("mysql migrations should run");
|
||||
cleanup_rows(&pool).await;
|
||||
seed_rows(&pool).await;
|
||||
|
||||
let repository = MysqlWalletReadRepository::new(pool);
|
||||
let wallet = repository
|
||||
.find(WalletLookupKey::UserId("user-1"))
|
||||
.await
|
||||
.expect("wallet find should query")
|
||||
.expect("wallet should exist");
|
||||
assert_eq!(wallet.total_adjusted, 3.0);
|
||||
|
||||
let page = repository
|
||||
.list_admin_wallets(&AdminWalletListQuery {
|
||||
status: Some("active".to_string()),
|
||||
owner_type: Some("user".to_string()),
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
})
|
||||
.await
|
||||
.expect("admin wallets should list");
|
||||
let wallet_item = page
|
||||
.items
|
||||
.iter()
|
||||
.find(|item| item.id == "wallet-1")
|
||||
.expect("seeded wallet should be listed");
|
||||
assert!(page.total >= 1);
|
||||
assert_eq!(wallet_item.total_adjusted, 3.0);
|
||||
|
||||
let orders = repository
|
||||
.list_admin_payment_orders(&AdminPaymentOrderListQuery {
|
||||
status: Some("credited".to_string()),
|
||||
payment_method: Some("redeem_code".to_string()),
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
})
|
||||
.await
|
||||
.expect("payment orders should list");
|
||||
assert_eq!(orders.total, 1);
|
||||
assert_eq!(
|
||||
orders.items[0].gateway_response.as_ref().unwrap()["ok"],
|
||||
true
|
||||
);
|
||||
|
||||
let refunds = repository
|
||||
.list_admin_wallet_refunds("wallet-1", 10, 0)
|
||||
.await
|
||||
.expect("refunds should list");
|
||||
assert_eq!(refunds.total, 1);
|
||||
assert_eq!(
|
||||
refunds.items[0].payout_proof.as_ref().unwrap()["proof"],
|
||||
"ok"
|
||||
);
|
||||
|
||||
let callbacks = repository
|
||||
.list_admin_payment_callbacks(Some("redeem_code"), 10, 0)
|
||||
.await
|
||||
.expect("callbacks should list");
|
||||
assert_eq!(callbacks.total, 1);
|
||||
assert!(callbacks.items[0].signature_valid);
|
||||
|
||||
let codes = repository
|
||||
.list_admin_redeem_codes(&AdminRedeemCodeListQuery {
|
||||
batch_id: "batch-1".to_string(),
|
||||
status: Some("redeemed".to_string()),
|
||||
limit: 10,
|
||||
offset: 0,
|
||||
})
|
||||
.await
|
||||
.expect("redeem codes should list");
|
||||
assert_eq!(codes.total, 1);
|
||||
assert_eq!(codes.items[0].masked_code, "ABCD****WXYZ");
|
||||
|
||||
let today = super::current_billing_date("UTC").expect("UTC should parse");
|
||||
sqlx::query("UPDATE wallet_daily_usage_ledgers SET billing_date = ? WHERE id = 'daily-1'")
|
||||
.bind(today)
|
||||
.execute(repository.pool())
|
||||
.await
|
||||
.expect("daily row should update");
|
||||
let daily = repository
|
||||
.find_wallet_today_usage("wallet-1", "UTC")
|
||||
.await
|
||||
.expect("daily usage should query")
|
||||
.expect("daily usage should exist");
|
||||
assert_eq!(daily.total_requests, 2);
|
||||
}
|
||||
|
||||
impl MysqlWalletReadRepository {
|
||||
fn pool(&self) -> &sqlx::MySqlPool {
|
||||
&self.pool
|
||||
}
|
||||
}
|
||||
|
||||
async fn cleanup_rows(pool: &sqlx::MySqlPool) {
|
||||
for sql in [
|
||||
"DELETE FROM wallet_daily_usage_ledgers WHERE id = 'daily-1'",
|
||||
"DELETE FROM redeem_codes WHERE id = 'code-1'",
|
||||
"DELETE FROM redeem_code_batches WHERE id = 'batch-1'",
|
||||
"DELETE FROM wallet_transactions WHERE id = 'tx-1'",
|
||||
"DELETE FROM refund_requests WHERE id = 'refund-1'",
|
||||
"DELETE FROM payment_callbacks WHERE id = 'callback-1'",
|
||||
"DELETE FROM payment_orders WHERE id = 'order-1'",
|
||||
"DELETE FROM wallets WHERE id = 'wallet-1'",
|
||||
"DELETE FROM users WHERE id = 'user-1'",
|
||||
] {
|
||||
sqlx::query(sql)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("cleanup should succeed");
|
||||
}
|
||||
}
|
||||
|
||||
async fn seed_rows(pool: &sqlx::MySqlPool) {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO users (id, username, email, auth_source, created_at, updated_at)
|
||||
VALUES ('user-1', 'Alice', 'alice@example.com', 'local', 1, 1)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("user should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO wallets (
|
||||
id, user_id, balance, gift_balance, total_recharged, total_consumed,
|
||||
total_refunded, total_adjusted, created_at, updated_at
|
||||
) VALUES (
|
||||
'wallet-1', 'user-1', 10.0, 2.0, 20.0, 4.0, 1.0, 3.0, 1, 2
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("wallet should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO payment_orders (
|
||||
id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd,
|
||||
refundable_amount_usd, payment_method, gateway_response, status, created_at
|
||||
) VALUES (
|
||||
'order-1', 'order-no-1', 'wallet-1', 'user-1', 5.0, 1.0, 4.0,
|
||||
'redeem_code', '{"ok":true}', 'credited', 3
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("payment order should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO payment_callbacks (
|
||||
id, payment_order_id, payment_method, callback_key, order_no,
|
||||
signature_valid, payload, created_at
|
||||
) VALUES (
|
||||
'callback-1', 'order-1', 'redeem_code', 'callback-key-1',
|
||||
'order-no-1', 1, '{"event":"paid"}', 4
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("callback should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO refund_requests (
|
||||
id, refund_no, wallet_id, user_id, payment_order_id, source_type,
|
||||
refund_mode, amount_usd, status, payout_proof, created_at, updated_at
|
||||
) VALUES (
|
||||
'refund-1', 'refund-no-1', 'wallet-1', 'user-1', 'order-1',
|
||||
'payment_order', 'offline_payout', 1.0, 'completed',
|
||||
'{"proof":"ok"}', 5, 6
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("refund should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO wallet_transactions (
|
||||
id, wallet_id, category, reason_code, amount, balance_before,
|
||||
balance_after, recharge_balance_before, recharge_balance_after,
|
||||
gift_balance_before, gift_balance_after, created_at
|
||||
) VALUES (
|
||||
'tx-1', 'wallet-1', 'credit', 'manual_adjustment', 3.0, 7.0, 10.0,
|
||||
5.0, 8.0, 2.0, 2.0, 7
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("transaction should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO redeem_code_batches (
|
||||
id, name, amount_usd, total_count, created_at, updated_at
|
||||
) VALUES (
|
||||
'batch-1', 'Batch One', 5.0, 1, 8, 9
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("redeem batch should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO redeem_codes (
|
||||
id, batch_id, code_hash, code_prefix, code_suffix, status,
|
||||
redeemed_by_user_id, redeemed_wallet_id, redeemed_payment_order_id,
|
||||
redeemed_at, created_at, updated_at
|
||||
) VALUES (
|
||||
'code-1', 'batch-1', 'hash-1', 'ABCD', 'WXYZ', 'redeemed',
|
||||
'user-1', 'wallet-1', 'order-1', 10, 8, 10
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("redeem code should seed");
|
||||
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO wallet_daily_usage_ledgers (
|
||||
id, wallet_id, billing_date, billing_timezone, total_cost_usd,
|
||||
total_requests, input_tokens, output_tokens, cache_creation_tokens,
|
||||
cache_read_tokens, aggregated_at, created_at, updated_at
|
||||
) VALUES (
|
||||
'daily-1', 'wallet-1', '2000-01-01', 'UTC', 1.25, 2, 10, 20, 3, 4, 11, 11, 11
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("daily usage should seed");
|
||||
}
|
||||
Reference in New Issue
Block a user