refactor(workspace): enforce layered crate boundaries

This commit is contained in:
elky
2026-07-15 23:47:19 +08:00
parent a728c090a9
commit 8616fe6ee2
969 changed files with 40187 additions and 27240 deletions
@@ -0,0 +1,81 @@
use std::fmt;
#[cfg(feature = "postgres")]
use super::PostgresBackend;
#[cfg(feature = "postgres")]
use crate::driver::postgres::{PostgresLeaseRunner, PostgresLeaseRunnerConfig};
#[cfg(feature = "postgres")]
use crate::DataLayerError;
#[derive(Clone, Default)]
pub struct DataLeaseBackends {
#[cfg(feature = "postgres")]
postgres: Option<PostgresLeaseRunner>,
}
impl fmt::Debug for DataLeaseBackends {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DataLeaseBackends")
.field("has_postgres", &self.has_any())
.finish()
}
}
impl DataLeaseBackends {
#[cfg(feature = "postgres")]
pub(crate) fn from_postgres(
postgres: Option<&PostgresBackend>,
) -> Result<Self, DataLayerError> {
Ok(Self {
postgres: postgres
.map(|backend| backend.lease_runner(PostgresLeaseRunnerConfig::default()))
.transpose()?,
})
}
#[cfg(feature = "postgres")]
pub fn postgres(&self) -> Option<PostgresLeaseRunner> {
self.postgres.clone()
}
pub fn has_any(&self) -> bool {
cfg!(feature = "postgres") && {
#[cfg(feature = "postgres")]
{
self.postgres.is_some()
}
#[cfg(not(feature = "postgres"))]
{
false
}
}
}
}
#[cfg(all(test, feature = "postgres"))]
mod tests {
use super::DataLeaseBackends;
use crate::backend::PostgresBackend;
use crate::driver::postgres::PostgresPoolConfig;
#[tokio::test]
async fn builds_postgres_lease_runner_from_backend() {
let backend = PostgresBackend::from_config(PostgresPoolConfig {
database_url: "postgres://localhost/aether".to_string(),
min_connections: 1,
max_connections: 4,
acquire_timeout_ms: 1_000,
idle_timeout_ms: 5_000,
max_lifetime_ms: 30_000,
statement_cache_capacity: 64,
require_ssl: false,
})
.expect("postgres backend should build");
let leases =
DataLeaseBackends::from_postgres(Some(&backend)).expect("lease backends should build");
assert!(leases.has_any());
assert!(leases.postgres().is_some());
}
}
@@ -0,0 +1,654 @@
#[cfg(feature = "mysql")]
mod mysql;
#[cfg(feature = "postgres")]
mod postgres;
#[cfg(feature = "sqlite")]
mod sqlite;
use super::{summarize_pool, DataBackends, SqlBackendRef};
use crate::maintenance::{
DatabaseMaintenanceSummary, DatabasePoolSummary, DatabasePostgresActivityGroup,
DatabasePostgresObservabilitySnapshot, StatsDailyAggregationInput,
StatsDailyAggregationSummary, StatsHourlyAggregationInput, StatsHourlyAggregationSummary,
WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult,
};
use crate::repository::system::{
AdminSystemPurgeSummary, AdminSystemPurgeTarget, AdminSystemStats,
AdminSystemUsageAggregateImportMode, AdminSystemUsageAggregateImportSummary,
AdminSystemUsageAggregateSnapshot, StoredSystemConfigEntry,
};
use crate::DataLayerError;
use sqlx::migrate::MigrateError;
async fn warm_pool<DB>(pool: &sqlx::Pool<DB>, min_connections: u32) -> Result<(), DataLayerError>
where
DB: sqlx::Database,
{
let mut connections = Vec::with_capacity(min_connections as usize);
for _ in 0..min_connections {
connections.push(pool.acquire().await.map_err(DataLayerError::sql)?);
}
Ok(())
}
pub(super) fn maintenance_identifier(value: &str) -> Result<&str, DataLayerError> {
let valid = !value.is_empty()
&& value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'_');
if valid {
Ok(value)
} else {
Err(DataLayerError::InvalidInput(format!(
"invalid maintenance table name: {value}"
)))
}
}
impl DataBackends {
pub fn has_database_maintenance_backend(&self) -> bool {
self.sql_backend().is_some()
}
pub fn has_database_pool_summary(&self) -> bool {
self.sql_backend().is_some()
}
/// Establishes the configured minimum number of SQL connections before the service reports
/// ready. Driver pools are built lazily, so relying on request traffic to grow them can make
/// the first concurrency ramp consume nearly every connection in the small cold pool.
pub async fn warm_database_pool(&self) -> Result<(), DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.warm_database_pool().await,
None => Ok(()),
}
}
pub fn has_system_config_backend(&self) -> bool {
self.sql_backend().is_some()
}
pub fn has_wallet_daily_usage_aggregation_backend(&self) -> bool {
self.sql_backend().is_some()
}
pub fn has_stats_hourly_aggregation_backend(&self) -> bool {
self.sql_backend().is_some()
}
pub fn has_stats_daily_aggregation_backend(&self) -> bool {
self.sql_backend().is_some()
}
pub async fn run_database_maintenance(
&self,
table_names: &[&str],
) -> Result<DatabaseMaintenanceSummary, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.run_database_maintenance(table_names).await,
None => Ok(DatabaseMaintenanceSummary::default()),
}
}
pub async fn run_database_migrations(&self) -> Result<bool, MigrateError> {
match self.sql_backend() {
Some(backend) => backend.run_database_migrations().await,
None => Ok(false),
}
}
pub async fn run_database_backfills(&self) -> Result<bool, MigrateError> {
match self.sql_backend() {
Some(backend) => backend.run_database_backfills().await,
None => Ok(false),
}
}
pub async fn pending_database_migrations(
&self,
) -> Result<Option<Vec<crate::lifecycle::migrate::PendingMigrationInfo>>, MigrateError> {
match self.sql_backend() {
Some(backend) => backend.pending_database_migrations().await,
None => Ok(None),
}
}
pub async fn prepare_database_for_startup(
&self,
) -> Result<Option<Vec<crate::lifecycle::migrate::PendingMigrationInfo>>, MigrateError> {
match self.sql_backend() {
Some(backend) => backend.prepare_database_for_startup().await,
None => Ok(None),
}
}
pub async fn pending_database_backfills(
&self,
) -> Result<Option<Vec<crate::lifecycle::backfill::PendingBackfillInfo>>, MigrateError> {
match self.sql_backend() {
Some(backend) => backend.pending_database_backfills().await,
None => Ok(None),
}
}
pub fn database_pool_summary(&self) -> Option<DatabasePoolSummary> {
self.sql_backend().map(SqlBackendRef::database_pool_summary)
}
pub async fn postgres_observability_snapshot(
&self,
) -> Result<Option<DatabasePostgresObservabilitySnapshot>, DataLayerError> {
#[cfg(not(feature = "postgres"))]
return Ok(None);
#[cfg(feature = "postgres")]
match self.postgres() {
Some(postgres) => postgres.postgres_observability_snapshot().await.map(Some),
None => Ok(None),
}
}
pub async fn postgres_activity_groups(
&self,
limit: i64,
) -> Result<Vec<DatabasePostgresActivityGroup>, DataLayerError> {
#[cfg(not(feature = "postgres"))]
let _ = limit;
#[cfg(not(feature = "postgres"))]
return Ok(Vec::new());
#[cfg(feature = "postgres")]
match self.postgres() {
Some(postgres) => postgres.postgres_activity_groups(limit).await,
None => Ok(Vec::new()),
}
}
pub async fn aggregate_wallet_daily_usage(
&self,
input: &WalletDailyUsageAggregationInput,
) -> Result<WalletDailyUsageAggregationResult, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.aggregate_wallet_daily_usage(input).await,
None => Ok(WalletDailyUsageAggregationResult::default()),
}
}
pub async fn aggregate_stats_hourly(
&self,
input: &StatsHourlyAggregationInput,
) -> Result<Option<StatsHourlyAggregationSummary>, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.aggregate_stats_hourly(input).await,
None => Ok(None),
}
}
pub async fn aggregate_stats_daily(
&self,
input: &StatsDailyAggregationInput,
) -> Result<Option<StatsDailyAggregationSummary>, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.aggregate_stats_daily(input).await,
None => Ok(None),
}
}
pub async fn find_system_config_value(
&self,
key: &str,
) -> Result<Option<serde_json::Value>, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.find_system_config_value(key).await,
None => Ok(None),
}
}
pub async fn list_system_config_entries(
&self,
) -> Result<Vec<StoredSystemConfigEntry>, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.list_system_config_entries().await,
None => Ok(Vec::new()),
}
}
pub async fn upsert_system_config_entry(
&self,
key: &str,
value: &serde_json::Value,
description: Option<&str>,
) -> Result<Option<StoredSystemConfigEntry>, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend
.upsert_system_config_entry(key, value, description)
.await
.map(Some),
None => Ok(None),
}
}
pub async fn delete_system_config_value(&self, key: &str) -> Result<bool, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.delete_system_config_value(key).await,
None => Ok(false),
}
}
pub async fn read_admin_system_stats(&self) -> Result<AdminSystemStats, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.read_admin_system_stats().await,
None => Ok(AdminSystemStats::default()),
}
}
pub async fn purge_admin_system_data(
&self,
target: AdminSystemPurgeTarget,
) -> Result<AdminSystemPurgeSummary, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.purge_admin_system_data(target).await,
None => Ok(AdminSystemPurgeSummary::default()),
}
}
pub async fn export_admin_system_usage_aggregates(
&self,
) -> Result<AdminSystemUsageAggregateSnapshot, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.export_admin_system_usage_aggregates().await,
None => Ok(AdminSystemUsageAggregateSnapshot::default()),
}
}
pub async fn import_admin_system_usage_aggregates(
&self,
snapshot: &AdminSystemUsageAggregateSnapshot,
user_id_map: &std::collections::BTreeMap<String, String>,
api_key_id_map: &std::collections::BTreeMap<String, String>,
mode: AdminSystemUsageAggregateImportMode,
) -> Result<AdminSystemUsageAggregateImportSummary, DataLayerError> {
match self.sql_backend() {
Some(backend) => {
backend
.import_admin_system_usage_aggregates(
snapshot,
user_id_map,
api_key_id_map,
mode,
)
.await
}
None => Ok(AdminSystemUsageAggregateImportSummary::default()),
}
}
pub async fn purge_admin_request_bodies_batch(
&self,
batch_size: usize,
) -> Result<AdminSystemPurgeSummary, DataLayerError> {
match self.sql_backend() {
Some(backend) => backend.purge_admin_request_bodies_batch(batch_size).await,
None => Ok(AdminSystemPurgeSummary::default()),
}
}
}
impl<'a> SqlBackendRef<'a> {
async fn warm_database_pool(self) -> Result<(), DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => {
warm_pool(postgres.pool(), postgres.config().min_connections).await
}
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => {
warm_pool(mysql.pool(), mysql.config().pool.min_connections).await
}
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => {
warm_pool(sqlite.pool(), sqlite.config().pool.min_connections).await
}
}
}
async fn run_database_maintenance(
self,
table_names: &[&str],
) -> Result<DatabaseMaintenanceSummary, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.run_table_maintenance(table_names).await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.run_table_maintenance(table_names).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.run_table_maintenance(table_names).await,
}
}
async fn run_database_migrations(self) -> Result<bool, MigrateError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => {
crate::lifecycle::migrate::run_migrations(postgres.pool()).await?;
Ok(true)
}
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => {
crate::lifecycle::migrate::run_mysql_migrations(mysql.pool()).await?;
Ok(true)
}
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => {
crate::lifecycle::migrate::run_sqlite_migrations(sqlite.pool()).await?;
Ok(true)
}
}
}
async fn run_database_backfills(self) -> Result<bool, MigrateError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => {
crate::lifecycle::backfill::run_backfills(postgres.pool()).await?;
Ok(true)
}
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => {
crate::lifecycle::backfill::run_mysql_backfills(mysql.pool()).await?;
Ok(true)
}
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => {
crate::lifecycle::backfill::run_sqlite_backfills(sqlite.pool()).await?;
Ok(true)
}
}
}
async fn pending_database_migrations(
self,
) -> Result<Option<Vec<crate::lifecycle::migrate::PendingMigrationInfo>>, MigrateError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => Ok(Some(
crate::lifecycle::migrate::pending_migrations(postgres.pool()).await?,
)),
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => Ok(Some(
crate::lifecycle::migrate::pending_mysql_migrations(mysql.pool()).await?,
)),
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => Ok(Some(
crate::lifecycle::migrate::pending_sqlite_migrations(sqlite.pool()).await?,
)),
}
}
async fn prepare_database_for_startup(
self,
) -> Result<Option<Vec<crate::lifecycle::migrate::PendingMigrationInfo>>, MigrateError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => Ok(Some(
crate::lifecycle::migrate::prepare_database_for_startup(postgres.pool()).await?,
)),
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => Ok(Some(
crate::lifecycle::migrate::prepare_mysql_database_for_startup(mysql.pool()).await?,
)),
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => Ok(Some(
crate::lifecycle::migrate::prepare_sqlite_database_for_startup(sqlite.pool())
.await?,
)),
}
}
async fn pending_database_backfills(
self,
) -> Result<Option<Vec<crate::lifecycle::backfill::PendingBackfillInfo>>, MigrateError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => Ok(Some(
crate::lifecycle::backfill::pending_backfills(postgres.pool()).await?,
)),
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => Ok(Some(
crate::lifecycle::backfill::pending_mysql_backfills(mysql.pool()).await?,
)),
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => Ok(Some(
crate::lifecycle::backfill::pending_sqlite_backfills(sqlite.pool()).await?,
)),
}
}
fn database_pool_summary(self) -> DatabasePoolSummary {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => summarize_pool(
crate::database::DatabaseDriver::Postgres,
usize::try_from(postgres.pool().size()).unwrap_or(usize::MAX),
postgres.pool().num_idle(),
postgres.config().max_connections,
),
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => summarize_pool(
crate::database::DatabaseDriver::Mysql,
usize::try_from(mysql.pool().size()).unwrap_or(usize::MAX),
mysql.pool().num_idle(),
mysql.config().pool.max_connections,
),
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => summarize_pool(
crate::database::DatabaseDriver::Sqlite,
usize::try_from(sqlite.pool().size()).unwrap_or(usize::MAX),
sqlite.pool().num_idle(),
sqlite.config().pool.max_connections,
),
}
}
async fn aggregate_wallet_daily_usage(
self,
input: &WalletDailyUsageAggregationInput,
) -> Result<WalletDailyUsageAggregationResult, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.aggregate_wallet_daily_usage(input).await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.aggregate_wallet_daily_usage(input).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.aggregate_wallet_daily_usage(input).await,
}
}
async fn aggregate_stats_hourly(
self,
input: &StatsHourlyAggregationInput,
) -> Result<Option<StatsHourlyAggregationSummary>, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.aggregate_stats_hourly(input).await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.aggregate_stats_hourly(input).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.aggregate_stats_hourly(input).await,
}
}
async fn aggregate_stats_daily(
self,
input: &StatsDailyAggregationInput,
) -> Result<Option<StatsDailyAggregationSummary>, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.aggregate_stats_daily(input).await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.aggregate_stats_daily(input).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.aggregate_stats_daily(input).await,
}
}
async fn find_system_config_value(
self,
key: &str,
) -> Result<Option<serde_json::Value>, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.find_system_config_value(key).await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.find_system_config_value(key).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.find_system_config_value(key).await,
}
}
async fn list_system_config_entries(
self,
) -> Result<Vec<StoredSystemConfigEntry>, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.list_system_config_entries().await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.list_system_config_entries().await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.list_system_config_entries().await,
}
}
async fn upsert_system_config_entry(
self,
key: &str,
value: &serde_json::Value,
description: Option<&str>,
) -> Result<StoredSystemConfigEntry, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => {
postgres
.upsert_system_config_entry(key, value, description)
.await
}
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => {
mysql
.upsert_system_config_entry(key, value, description)
.await
}
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => {
sqlite
.upsert_system_config_entry(key, value, description)
.await
}
}
}
async fn delete_system_config_value(self, key: &str) -> Result<bool, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.delete_system_config_value(key).await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.delete_system_config_value(key).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.delete_system_config_value(key).await,
}
}
async fn read_admin_system_stats(self) -> Result<AdminSystemStats, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.read_admin_system_stats().await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.read_admin_system_stats().await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.read_admin_system_stats().await,
}
}
async fn purge_admin_system_data(
self,
target: AdminSystemPurgeTarget,
) -> Result<AdminSystemPurgeSummary, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.purge_admin_system_data(target).await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.purge_admin_system_data(target).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.purge_admin_system_data(target).await,
}
}
async fn export_admin_system_usage_aggregates(
self,
) -> Result<AdminSystemUsageAggregateSnapshot, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.export_admin_system_usage_aggregates().await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.export_admin_system_usage_aggregates().await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.export_admin_system_usage_aggregates().await,
}
}
async fn import_admin_system_usage_aggregates(
self,
snapshot: &AdminSystemUsageAggregateSnapshot,
user_id_map: &std::collections::BTreeMap<String, String>,
api_key_id_map: &std::collections::BTreeMap<String, String>,
mode: AdminSystemUsageAggregateImportMode,
) -> Result<AdminSystemUsageAggregateImportSummary, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => {
postgres
.import_admin_system_usage_aggregates(
snapshot,
user_id_map,
api_key_id_map,
mode,
)
.await
}
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => {
mysql
.import_admin_system_usage_aggregates(
snapshot,
user_id_map,
api_key_id_map,
mode,
)
.await
}
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => {
sqlite
.import_admin_system_usage_aggregates(
snapshot,
user_id_map,
api_key_id_map,
mode,
)
.await
}
}
}
async fn purge_admin_request_bodies_batch(
self,
batch_size: usize,
) -> Result<AdminSystemPurgeSummary, DataLayerError> {
match self {
#[cfg(feature = "postgres")]
Self::Postgres(postgres) => postgres.purge_admin_request_bodies_batch(batch_size).await,
#[cfg(feature = "mysql")]
Self::Mysql(mysql) => mysql.purge_admin_request_bodies_batch(batch_size).await,
#[cfg(feature = "sqlite")]
Self::Sqlite(sqlite) => sqlite.purge_admin_request_bodies_batch(batch_size).await,
}
}
}
@@ -0,0 +1,28 @@
use crate::backend::MysqlBackend;
use crate::error::SqlResultExt;
use crate::{DataLayerError, DatabaseMaintenanceSummary};
use super::maintenance_identifier;
impl MysqlBackend {
pub async fn run_table_maintenance(
&self,
table_names: &[&str],
) -> Result<DatabaseMaintenanceSummary, DataLayerError> {
let mut summary = DatabaseMaintenanceSummary::default();
for table_name in table_names {
let table_name = maintenance_identifier(table_name)?;
summary.attempted += 1;
let statement = format!("ANALYZE TABLE `{table_name}`");
if sqlx::raw_sql(&statement)
.execute(self.pool())
.await
.map_sql_err()
.is_ok()
{
summary.succeeded += 1;
}
}
Ok(summary)
}
}
@@ -0,0 +1,663 @@
use sqlx::Row;
use crate::backend::PostgresBackend;
use crate::error::SqlxResultExt;
use crate::maintenance::{
DatabaseMaintenanceSummary, DatabasePostgresActivityGroup,
DatabasePostgresObservabilitySnapshot,
};
use crate::DataLayerError;
use super::maintenance_identifier;
impl PostgresBackend {
pub async fn run_table_maintenance(
&self,
table_names: &[&str],
) -> Result<DatabaseMaintenanceSummary, DataLayerError> {
let mut summary = DatabaseMaintenanceSummary::default();
for table_name in table_names {
let table_name = maintenance_identifier(table_name)?;
summary.attempted += 1;
let statement = format!("VACUUM ANALYZE \"{table_name}\"");
if sqlx::raw_sql(&statement)
.execute(self.pool())
.await
.map_postgres_err()
.is_ok()
{
summary.succeeded += 1;
}
}
Ok(summary)
}
pub async fn postgres_observability_snapshot(
&self,
) -> Result<DatabasePostgresObservabilitySnapshot, DataLayerError> {
const ACTIVITY_SQL: &str = r#"
SELECT
COUNT(*) FILTER (WHERE state = 'active')::BIGINT AS active_connections,
COUNT(*) FILTER (WHERE state = 'idle')::BIGINT AS idle_connections,
COUNT(*) FILTER (WHERE state = 'idle in transaction')::BIGINT AS idle_in_transaction_connections,
COUNT(*) FILTER (WHERE state = 'active' AND wait_event_type IS NOT NULL)::BIGINT AS waiting_connections,
COUNT(*) FILTER (WHERE state = 'active' AND wait_event_type = 'Lock')::BIGINT AS lock_waiting_connections,
COALESCE(MAX(EXTRACT(EPOCH FROM now() - query_start) * 1000) FILTER (WHERE state = 'active' AND query_start IS NOT NULL), 0)::BIGINT AS oldest_active_query_age_ms,
COALESCE(MAX(EXTRACT(EPOCH FROM now() - xact_start) * 1000) FILTER (WHERE xact_start IS NOT NULL), 0)::BIGINT AS oldest_transaction_age_ms
FROM pg_stat_activity
WHERE datname = current_database()
AND pid <> pg_backend_pid()
AND backend_type = 'client backend'
"#;
const DEADLOCKS_SQL: &str = r#"
SELECT
COALESCE(SUM(deadlocks), 0)::BIGINT AS deadlocks_total,
COALESCE(SUM(blks_read), 0)::BIGINT AS block_read_total,
COALESCE(SUM(blks_hit), 0)::BIGINT AS block_hit_total,
COALESCE(SUM(temp_files), 0)::BIGINT AS temp_files_total,
COALESCE(SUM(temp_bytes), 0)::BIGINT AS temp_bytes_total,
COALESCE(SUM(xact_commit), 0)::BIGINT AS xact_commit_total,
COALESCE(SUM(xact_rollback), 0)::BIGINT AS xact_rollback_total
FROM pg_stat_database
WHERE datname = current_database()
"#;
let activity = sqlx::query(ACTIVITY_SQL)
.fetch_one(self.pool())
.await
.map_postgres_err()?;
let database = sqlx::query(DEADLOCKS_SQL)
.fetch_one(self.pool())
.await
.map_postgres_err()?;
let wal = self.postgres_wal_observability_snapshot().await;
let checkpoint = self.postgres_checkpoint_observability_snapshot().await;
let statements = self.postgres_statement_observability_snapshot().await;
let block_read_total = row_u64(&database, "block_read_total")?;
let block_hit_total = row_u64(&database, "block_hit_total")?;
Ok(DatabasePostgresObservabilitySnapshot {
active_connections: row_u64(&activity, "active_connections")?,
idle_connections: row_u64(&activity, "idle_connections")?,
idle_in_transaction_connections: row_u64(&activity, "idle_in_transaction_connections")?,
waiting_connections: row_u64(&activity, "waiting_connections")?,
lock_waiting_connections: row_u64(&activity, "lock_waiting_connections")?,
oldest_active_query_age_ms: row_u64(&activity, "oldest_active_query_age_ms")?,
oldest_transaction_age_ms: row_u64(&activity, "oldest_transaction_age_ms")?,
deadlocks_total: row_u64(&database, "deadlocks_total")?,
block_read_total,
block_hit_total,
block_cache_hit_rate_basis_points: ratio_to_basis_points(
block_hit_total,
block_read_total.saturating_add(block_hit_total),
),
temp_files_total: row_u64(&database, "temp_files_total")?,
temp_bytes_total: row_u64(&database, "temp_bytes_total")?,
xact_commit_total: row_u64(&database, "xact_commit_total")?,
xact_rollback_total: row_u64(&database, "xact_rollback_total")?,
wal_observability_available: wal.available,
wal_observability_unavailable: wal.unavailable,
wal_records_total: wal.records_total,
wal_fpi_total: wal.fpi_total,
wal_bytes_total: wal.bytes_total,
wal_buffers_full_total: wal.buffers_full_total,
wal_write_total: wal.write_total,
wal_sync_total: wal.sync_total,
wal_write_time_ms_total: wal.write_time_ms_total,
wal_sync_time_ms_total: wal.sync_time_ms_total,
checkpoint_observability_available: checkpoint.available,
checkpoint_observability_unavailable: checkpoint.unavailable,
checkpoints_timed_total: checkpoint.timed_total,
checkpoints_requested_total: checkpoint.requested_total,
checkpoint_write_time_ms_total: checkpoint.write_time_ms_total,
checkpoint_sync_time_ms_total: checkpoint.sync_time_ms_total,
buffers_checkpoint_total: checkpoint.buffers_checkpoint_total,
buffers_backend_total: checkpoint.buffers_backend_total,
statement_observability_available: statements.available,
statement_observability_unavailable: statements.unavailable,
statement_top_calls_total: statements.top_calls_total,
statement_top_exec_time_ms_total: statements.top_exec_time_ms_total,
statement_top_max_mean_exec_time_ms: statements.top_max_mean_exec_time_ms,
statement_top_max_exec_time_ms: statements.top_max_exec_time_ms,
statement_top_shared_blks_read_total: statements.top_shared_blks_read_total,
statement_top_shared_blks_hit_total: statements.top_shared_blks_hit_total,
statement_top_temp_blks_total: statements.top_temp_blks_total,
})
}
pub async fn postgres_activity_groups(
&self,
limit: i64,
) -> Result<Vec<DatabasePostgresActivityGroup>, DataLayerError> {
const ACTIVITY_GROUP_SQL: &str = r#"
WITH normalized_activity AS (
SELECT
COALESCE(NULLIF(state, ''), 'unknown') AS state,
COALESCE(NULLIF(wait_event_type, ''), 'none') AS wait_event_type,
COALESCE(NULLIF(wait_event, ''), 'none') AS wait_event,
LEFT(
regexp_replace(
regexp_replace(
COALESCE(NULLIF(query, ''), '<empty>'),
'\s+',
' ',
'g'
),
'([0-9a-fA-F]{8,}|[0-9]+)',
'?',
'g'
),
160
) AS query_prefix,
COALESCE(EXTRACT(EPOCH FROM now() - query_start) * 1000, 0)::BIGINT AS query_age_ms,
COALESCE(EXTRACT(EPOCH FROM now() - xact_start) * 1000, 0)::BIGINT AS transaction_age_ms
FROM pg_stat_activity
WHERE datname = current_database()
AND pid <> pg_backend_pid()
AND backend_type = 'client backend'
)
SELECT
state,
wait_event_type,
wait_event,
query_prefix,
COUNT(*)::BIGINT AS connections,
COALESCE(MAX(query_age_ms), 0)::BIGINT AS max_query_age_ms,
COALESCE(MAX(transaction_age_ms), 0)::BIGINT AS max_transaction_age_ms
FROM normalized_activity
GROUP BY state, wait_event_type, wait_event, query_prefix
ORDER BY connections DESC, max_transaction_age_ms DESC, max_query_age_ms DESC
LIMIT $1
"#;
let rows = sqlx::query(ACTIVITY_GROUP_SQL)
.bind(limit.clamp(1, 20))
.fetch_all(self.pool())
.await
.map_postgres_err()?;
rows.into_iter()
.map(|row| {
Ok(DatabasePostgresActivityGroup {
state: row.try_get::<String, _>("state").map_postgres_err()?,
wait_event_type: row
.try_get::<String, _>("wait_event_type")
.map_postgres_err()?,
wait_event: row.try_get::<String, _>("wait_event").map_postgres_err()?,
query_prefix: row
.try_get::<String, _>("query_prefix")
.map_postgres_err()?,
connections: row_u64(&row, "connections")?,
max_query_age_ms: row_u64(&row, "max_query_age_ms")?,
max_transaction_age_ms: row_u64(&row, "max_transaction_age_ms")?,
})
})
.collect()
}
async fn postgres_wal_observability_snapshot(&self) -> PostgresWalObservabilitySnapshot {
if !self
.postgres_catalog_relation_has_columns(
"pg_catalog.pg_stat_wal",
&["wal_records", "wal_fpi", "wal_bytes", "wal_buffers_full"],
)
.await
{
return PostgresWalObservabilitySnapshot::default();
}
const WAL_SQL: &str = r#"
SELECT
COALESCE(SUM(wal_records), 0)::BIGINT AS records_total,
COALESCE(SUM(wal_fpi), 0)::BIGINT AS fpi_total,
COALESCE(SUM(wal_bytes), 0)::BIGINT AS bytes_total,
COALESCE(SUM(wal_buffers_full), 0)::BIGINT AS buffers_full_total
FROM pg_stat_wal
"#;
match sqlx::query(WAL_SQL).fetch_one(self.pool()).await {
Ok(row) => {
let io = self.postgres_wal_io_observability_snapshot().await;
PostgresWalObservabilitySnapshot {
available: 1,
records_total: row_u64(&row, "records_total").unwrap_or_default(),
fpi_total: row_u64(&row, "fpi_total").unwrap_or_default(),
bytes_total: row_u64(&row, "bytes_total").unwrap_or_default(),
buffers_full_total: row_u64(&row, "buffers_full_total").unwrap_or_default(),
write_total: io.write_total,
sync_total: io.sync_total,
write_time_ms_total: io.write_time_ms_total,
sync_time_ms_total: io.sync_time_ms_total,
..PostgresWalObservabilitySnapshot::default()
}
}
Err(_) => PostgresWalObservabilitySnapshot {
unavailable: 1,
..PostgresWalObservabilitySnapshot::default()
},
}
}
async fn postgres_wal_io_observability_snapshot(&self) -> PostgresWalIoObservabilitySnapshot {
if self
.postgres_catalog_relation_has_columns(
"pg_catalog.pg_stat_wal",
&["wal_write", "wal_sync", "wal_write_time", "wal_sync_time"],
)
.await
{
return self.postgres_wal_legacy_io_observability_snapshot().await;
}
if !self
.postgres_catalog_relation_has_columns(
"pg_catalog.pg_stat_io",
&["object", "writes", "fsyncs", "write_time", "fsync_time"],
)
.await
{
return PostgresWalIoObservabilitySnapshot::default();
}
const WAL_IO_SQL: &str = r#"
SELECT
COALESCE(SUM(writes), 0)::BIGINT AS write_total,
COALESCE(SUM(fsyncs), 0)::BIGINT AS sync_total,
COALESCE(SUM(write_time), 0)::BIGINT AS write_time_ms_total,
COALESCE(SUM(fsync_time), 0)::BIGINT AS sync_time_ms_total
FROM pg_stat_io
WHERE object = 'wal'
"#;
match sqlx::query(WAL_IO_SQL).fetch_one(self.pool()).await {
Ok(row) => PostgresWalIoObservabilitySnapshot {
write_total: row_u64(&row, "write_total").unwrap_or_default(),
sync_total: row_u64(&row, "sync_total").unwrap_or_default(),
write_time_ms_total: row_u64(&row, "write_time_ms_total").unwrap_or_default(),
sync_time_ms_total: row_u64(&row, "sync_time_ms_total").unwrap_or_default(),
},
Err(_) => PostgresWalIoObservabilitySnapshot::default(),
}
}
async fn postgres_wal_legacy_io_observability_snapshot(
&self,
) -> PostgresWalIoObservabilitySnapshot {
const WAL_IO_SQL: &str = r#"
SELECT
COALESCE(SUM(wal_write), 0)::BIGINT AS write_total,
COALESCE(SUM(wal_sync), 0)::BIGINT AS sync_total,
COALESCE(SUM(wal_write_time), 0)::BIGINT AS write_time_ms_total,
COALESCE(SUM(wal_sync_time), 0)::BIGINT AS sync_time_ms_total
FROM pg_stat_wal
"#;
match sqlx::query(WAL_IO_SQL).fetch_one(self.pool()).await {
Ok(row) => PostgresWalIoObservabilitySnapshot {
write_total: row_u64(&row, "write_total").unwrap_or_default(),
sync_total: row_u64(&row, "sync_total").unwrap_or_default(),
write_time_ms_total: row_u64(&row, "write_time_ms_total").unwrap_or_default(),
sync_time_ms_total: row_u64(&row, "sync_time_ms_total").unwrap_or_default(),
},
Err(_) => PostgresWalIoObservabilitySnapshot::default(),
}
}
async fn postgres_checkpoint_observability_snapshot(
&self,
) -> PostgresCheckpointObservabilitySnapshot {
if self
.postgres_catalog_relation_has_columns(
"pg_catalog.pg_stat_checkpointer",
&[
"num_timed",
"num_requested",
"write_time",
"sync_time",
"buffers_written",
],
)
.await
{
return self
.postgres_checkpoint_observability_snapshot_from_checkpointer()
.await;
}
if !self
.postgres_catalog_relation_has_columns(
"pg_catalog.pg_stat_bgwriter",
&[
"checkpoints_timed",
"checkpoints_req",
"checkpoint_write_time",
"checkpoint_sync_time",
"buffers_checkpoint",
"buffers_backend",
],
)
.await
{
return PostgresCheckpointObservabilitySnapshot::default();
}
const CHECKPOINT_SQL: &str = r#"
SELECT
COALESCE(SUM(checkpoints_timed), 0)::BIGINT AS timed_total,
COALESCE(SUM(checkpoints_req), 0)::BIGINT AS requested_total,
COALESCE(SUM(checkpoint_write_time), 0)::BIGINT AS write_time_ms_total,
COALESCE(SUM(checkpoint_sync_time), 0)::BIGINT AS sync_time_ms_total,
COALESCE(SUM(buffers_checkpoint), 0)::BIGINT AS buffers_checkpoint_total,
COALESCE(SUM(buffers_backend), 0)::BIGINT AS buffers_backend_total
FROM pg_stat_bgwriter
"#;
match sqlx::query(CHECKPOINT_SQL).fetch_one(self.pool()).await {
Ok(row) => PostgresCheckpointObservabilitySnapshot {
available: 1,
timed_total: row_u64(&row, "timed_total").unwrap_or_default(),
requested_total: row_u64(&row, "requested_total").unwrap_or_default(),
write_time_ms_total: row_u64(&row, "write_time_ms_total").unwrap_or_default(),
sync_time_ms_total: row_u64(&row, "sync_time_ms_total").unwrap_or_default(),
buffers_checkpoint_total: row_u64(&row, "buffers_checkpoint_total")
.unwrap_or_default(),
buffers_backend_total: row_u64(&row, "buffers_backend_total").unwrap_or_default(),
..PostgresCheckpointObservabilitySnapshot::default()
},
Err(_) => PostgresCheckpointObservabilitySnapshot {
unavailable: 1,
..PostgresCheckpointObservabilitySnapshot::default()
},
}
}
async fn postgres_checkpoint_observability_snapshot_from_checkpointer(
&self,
) -> PostgresCheckpointObservabilitySnapshot {
const CHECKPOINT_SQL: &str = r#"
SELECT
COALESCE(SUM(num_timed), 0)::BIGINT AS timed_total,
COALESCE(SUM(num_requested), 0)::BIGINT AS requested_total,
COALESCE(SUM(write_time), 0)::BIGINT AS write_time_ms_total,
COALESCE(SUM(sync_time), 0)::BIGINT AS sync_time_ms_total,
COALESCE(SUM(buffers_written), 0)::BIGINT AS buffers_checkpoint_total
FROM pg_stat_checkpointer
"#;
match sqlx::query(CHECKPOINT_SQL).fetch_one(self.pool()).await {
Ok(row) => PostgresCheckpointObservabilitySnapshot {
available: 1,
timed_total: row_u64(&row, "timed_total").unwrap_or_default(),
requested_total: row_u64(&row, "requested_total").unwrap_or_default(),
write_time_ms_total: row_u64(&row, "write_time_ms_total").unwrap_or_default(),
sync_time_ms_total: row_u64(&row, "sync_time_ms_total").unwrap_or_default(),
buffers_checkpoint_total: row_u64(&row, "buffers_checkpoint_total")
.unwrap_or_default(),
..PostgresCheckpointObservabilitySnapshot::default()
},
Err(_) => PostgresCheckpointObservabilitySnapshot {
unavailable: 1,
..PostgresCheckpointObservabilitySnapshot::default()
},
}
}
async fn postgres_statement_observability_snapshot(
&self,
) -> PostgresStatementObservabilitySnapshot {
let extension_installed = sqlx::query_scalar::<_, bool>(
"SELECT EXISTS (SELECT 1 FROM pg_extension WHERE extname = 'pg_stat_statements')",
)
.fetch_one(self.pool())
.await
.unwrap_or(false);
if !extension_installed {
return PostgresStatementObservabilitySnapshot::default();
}
if self
.postgres_catalog_relation_has_columns(
"pg_stat_statements",
&[
"calls",
"total_exec_time",
"mean_exec_time",
"max_exec_time",
"shared_blks_read",
"shared_blks_hit",
"temp_blks_read",
"temp_blks_written",
"dbid",
],
)
.await
{
return self
.postgres_statement_observability_snapshot_with_exec_time()
.await;
}
if self
.postgres_catalog_relation_has_columns(
"pg_stat_statements",
&[
"calls",
"total_time",
"mean_time",
"max_time",
"shared_blks_read",
"shared_blks_hit",
"temp_blks_read",
"temp_blks_written",
"dbid",
],
)
.await
{
return self
.postgres_statement_observability_snapshot_with_total_time()
.await;
}
PostgresStatementObservabilitySnapshot::default()
}
async fn postgres_statement_observability_snapshot_with_exec_time(
&self,
) -> PostgresStatementObservabilitySnapshot {
const STATEMENTS_SQL: &str = r#"
SELECT
COALESCE(SUM(calls), 0)::BIGINT AS top_calls_total,
COALESCE(SUM(total_exec_time), 0)::BIGINT AS top_exec_time_ms_total,
COALESCE(MAX(mean_exec_time), 0)::BIGINT AS top_max_mean_exec_time_ms,
COALESCE(MAX(max_exec_time), 0)::BIGINT AS top_max_exec_time_ms,
COALESCE(SUM(shared_blks_read), 0)::BIGINT AS top_shared_blks_read_total,
COALESCE(SUM(shared_blks_hit), 0)::BIGINT AS top_shared_blks_hit_total,
COALESCE(SUM(temp_blks_read + temp_blks_written), 0)::BIGINT AS top_temp_blks_total
FROM (
SELECT
calls,
total_exec_time,
mean_exec_time,
max_exec_time,
shared_blks_read,
shared_blks_hit,
temp_blks_read,
temp_blks_written
FROM pg_stat_statements
WHERE dbid = (SELECT oid FROM pg_database WHERE datname = current_database())
ORDER BY total_exec_time DESC
LIMIT 20
) top_statements
"#;
match sqlx::query(STATEMENTS_SQL).fetch_one(self.pool()).await {
Ok(row) => PostgresStatementObservabilitySnapshot {
available: 1,
top_calls_total: row_u64(&row, "top_calls_total").unwrap_or_default(),
top_exec_time_ms_total: row_u64(&row, "top_exec_time_ms_total").unwrap_or_default(),
top_max_mean_exec_time_ms: row_u64(&row, "top_max_mean_exec_time_ms")
.unwrap_or_default(),
top_max_exec_time_ms: row_u64(&row, "top_max_exec_time_ms").unwrap_or_default(),
top_shared_blks_read_total: row_u64(&row, "top_shared_blks_read_total")
.unwrap_or_default(),
top_shared_blks_hit_total: row_u64(&row, "top_shared_blks_hit_total")
.unwrap_or_default(),
top_temp_blks_total: row_u64(&row, "top_temp_blks_total").unwrap_or_default(),
..PostgresStatementObservabilitySnapshot::default()
},
Err(_) => PostgresStatementObservabilitySnapshot {
unavailable: 1,
..PostgresStatementObservabilitySnapshot::default()
},
}
}
async fn postgres_statement_observability_snapshot_with_total_time(
&self,
) -> PostgresStatementObservabilitySnapshot {
const STATEMENTS_SQL: &str = r#"
SELECT
COALESCE(SUM(calls), 0)::BIGINT AS top_calls_total,
COALESCE(SUM(total_time), 0)::BIGINT AS top_exec_time_ms_total,
COALESCE(MAX(mean_time), 0)::BIGINT AS top_max_mean_exec_time_ms,
COALESCE(MAX(max_time), 0)::BIGINT AS top_max_exec_time_ms,
COALESCE(SUM(shared_blks_read), 0)::BIGINT AS top_shared_blks_read_total,
COALESCE(SUM(shared_blks_hit), 0)::BIGINT AS top_shared_blks_hit_total,
COALESCE(SUM(temp_blks_read + temp_blks_written), 0)::BIGINT AS top_temp_blks_total
FROM (
SELECT
calls,
total_time,
mean_time,
max_time,
shared_blks_read,
shared_blks_hit,
temp_blks_read,
temp_blks_written
FROM pg_stat_statements
WHERE dbid = (SELECT oid FROM pg_database WHERE datname = current_database())
ORDER BY total_time DESC
LIMIT 20
) top_statements
"#;
match sqlx::query(STATEMENTS_SQL).fetch_one(self.pool()).await {
Ok(row) => PostgresStatementObservabilitySnapshot {
available: 1,
top_calls_total: row_u64(&row, "top_calls_total").unwrap_or_default(),
top_exec_time_ms_total: row_u64(&row, "top_exec_time_ms_total").unwrap_or_default(),
top_max_mean_exec_time_ms: row_u64(&row, "top_max_mean_exec_time_ms")
.unwrap_or_default(),
top_max_exec_time_ms: row_u64(&row, "top_max_exec_time_ms").unwrap_or_default(),
top_shared_blks_read_total: row_u64(&row, "top_shared_blks_read_total")
.unwrap_or_default(),
top_shared_blks_hit_total: row_u64(&row, "top_shared_blks_hit_total")
.unwrap_or_default(),
top_temp_blks_total: row_u64(&row, "top_temp_blks_total").unwrap_or_default(),
..PostgresStatementObservabilitySnapshot::default()
},
Err(_) => PostgresStatementObservabilitySnapshot {
unavailable: 1,
..PostgresStatementObservabilitySnapshot::default()
},
}
}
async fn postgres_catalog_relation_has_columns(
&self,
relation: &str,
columns: &[&str],
) -> bool {
if !self.postgres_catalog_relation_exists(relation).await {
return false;
}
for column in columns {
if !self.postgres_catalog_column_exists(relation, column).await {
return false;
}
}
true
}
async fn postgres_catalog_column_exists(&self, relation: &str, column: &str) -> bool {
sqlx::query_scalar::<_, bool>(
r#"
SELECT EXISTS (
SELECT 1
FROM pg_attribute
WHERE attrelid = to_regclass($1)
AND attname = $2
AND NOT attisdropped
)
"#,
)
.bind(relation)
.bind(column)
.fetch_one(self.pool())
.await
.unwrap_or(false)
}
async fn postgres_catalog_relation_exists(&self, relation: &str) -> bool {
sqlx::query_scalar::<_, Option<String>>("SELECT to_regclass($1)::TEXT")
.bind(relation)
.fetch_one(self.pool())
.await
.ok()
.flatten()
.is_some()
}
}
fn row_u64(row: &sqlx::postgres::PgRow, name: &str) -> Result<u64, DataLayerError> {
row.try_get::<i64, _>(name)
.map(u64_from_i64)
.map_postgres_err()
}
fn u64_from_i64(value: i64) -> u64 {
u64::try_from(value).unwrap_or_default()
}
fn ratio_to_basis_points(value: u64, total: u64) -> u64 {
value.saturating_mul(10_000).checked_div(total).unwrap_or(0)
}
#[derive(Debug, Clone, Copy, Default)]
struct PostgresWalObservabilitySnapshot {
available: u64,
unavailable: u64,
records_total: u64,
fpi_total: u64,
bytes_total: u64,
buffers_full_total: u64,
write_total: u64,
sync_total: u64,
write_time_ms_total: u64,
sync_time_ms_total: u64,
}
#[derive(Debug, Clone, Copy, Default)]
struct PostgresWalIoObservabilitySnapshot {
write_total: u64,
sync_total: u64,
write_time_ms_total: u64,
sync_time_ms_total: u64,
}
#[derive(Debug, Clone, Copy, Default)]
struct PostgresCheckpointObservabilitySnapshot {
available: u64,
unavailable: u64,
timed_total: u64,
requested_total: u64,
write_time_ms_total: u64,
sync_time_ms_total: u64,
buffers_checkpoint_total: u64,
buffers_backend_total: u64,
}
#[derive(Debug, Clone, Copy, Default)]
struct PostgresStatementObservabilitySnapshot {
available: u64,
unavailable: u64,
top_calls_total: u64,
top_exec_time_ms_total: u64,
top_max_mean_exec_time_ms: u64,
top_max_exec_time_ms: u64,
top_shared_blks_read_total: u64,
top_shared_blks_hit_total: u64,
top_temp_blks_total: u64,
}
@@ -0,0 +1,34 @@
use crate::backend::SqliteBackend;
use crate::error::SqlResultExt;
use crate::{DataLayerError, DatabaseMaintenanceSummary};
use super::maintenance_identifier;
impl SqliteBackend {
pub async fn run_table_maintenance(
&self,
table_names: &[&str],
) -> Result<DatabaseMaintenanceSummary, DataLayerError> {
let mut summary = DatabaseMaintenanceSummary::default();
for table_name in table_names {
let table_name = maintenance_identifier(table_name)?;
summary.attempted += 1;
let statement = format!("ANALYZE \"{table_name}\"");
if sqlx::raw_sql(&statement)
.execute(self.pool())
.await
.map_sql_err()
.is_ok()
{
summary.succeeded += 1;
}
}
if summary.succeeded > 0 {
sqlx::raw_sql("PRAGMA optimize")
.execute(self.pool())
.await
.map_sql_err()?;
}
Ok(summary)
}
}
@@ -0,0 +1,478 @@
//! Backend composition layer.
//!
//! `DataBackends` chooses the configured SQL driver, builds low-level pools,
//! instantiates concrete repositories, and exposes app-facing read/write,
//! lease, transaction, and maintenance handles. Request-path repository SQL
//! belongs in the selected `aether-data-*` adapter; backend-owned maintenance
//! SQL lives in focused modules such as `stats`, `wallet`, and `system`.
mod leases;
mod maintenance;
#[cfg(feature = "mysql")]
mod mysql;
#[cfg(feature = "postgres")]
mod postgres;
mod read;
mod referrals;
#[cfg(feature = "sqlite")]
mod sqlite;
mod stats;
#[cfg(any(feature = "mysql", feature = "sqlite"))]
mod stats_common;
mod system;
mod transactions;
mod wallet;
mod write;
use crate::maintenance::DatabasePoolSummary;
pub use leases::DataLeaseBackends;
#[cfg(feature = "mysql")]
pub use mysql::MysqlBackend;
#[cfg(feature = "postgres")]
pub use postgres::PostgresBackend;
pub use read::DataReadRepositories;
pub use referrals::{
ReferralAdminStats, ReferralDataState, ReferralMutationStatus, ReferralRelationshipListQuery,
ReferralRelationshipRecord, ReferralRewardConfig, ReferralRewardListQuery,
ReferralRewardRecord, ReferralUserDashboard,
};
#[cfg(feature = "sqlite")]
pub use sqlite::SqliteBackend;
pub use transactions::DataTransactionBackends;
pub use write::DataWriteRepositories;
use crate::database::DatabaseDriver;
use crate::{DataLayerConfig, DataLayerError};
#[derive(Clone, Copy)]
enum SqlBackendRef<'a> {
#[cfg(feature = "postgres")]
Postgres(&'a PostgresBackend),
#[cfg(feature = "mysql")]
Mysql(&'a MysqlBackend),
#[cfg(feature = "sqlite")]
Sqlite(&'a SqliteBackend),
}
#[derive(Debug, Clone, Default)]
pub struct DataBackends {
config: DataLayerConfig,
#[cfg(feature = "postgres")]
postgres: Option<PostgresBackend>,
#[cfg(feature = "mysql")]
mysql: Option<MysqlBackend>,
#[cfg(feature = "sqlite")]
sqlite: Option<SqliteBackend>,
leases: DataLeaseBackends,
read: DataReadRepositories,
transactions: DataTransactionBackends,
write: DataWriteRepositories,
}
fn summarize_pool(
driver: DatabaseDriver,
pool_size: usize,
idle: usize,
max_connections: u32,
) -> DatabasePoolSummary {
let max_connections = max_connections.max(1);
let checked_out = pool_size.saturating_sub(idle);
let usage_rate = checked_out as f64 / f64::from(max_connections) * 100.0;
DatabasePoolSummary {
driver,
checked_out,
pool_size,
idle,
max_connections,
usage_rate,
}
}
fn ensure_driver_enabled(driver: DatabaseDriver) -> Result<(), DataLayerError> {
match driver {
#[cfg(feature = "postgres")]
DatabaseDriver::Postgres => Ok(()),
#[cfg(not(feature = "postgres"))]
DatabaseDriver::Postgres => Err(DataLayerError::InvalidInput(
"PostgreSQL driver is not enabled for this aether-data build".to_string(),
)),
#[cfg(feature = "mysql")]
DatabaseDriver::Mysql => Ok(()),
#[cfg(not(feature = "mysql"))]
DatabaseDriver::Mysql => Err(DataLayerError::InvalidInput(
"MySQL driver is not enabled for this aether-data build".to_string(),
)),
#[cfg(feature = "sqlite")]
DatabaseDriver::Sqlite => Ok(()),
#[cfg(not(feature = "sqlite"))]
DatabaseDriver::Sqlite => Err(DataLayerError::InvalidInput(
"SQLite driver is not enabled for this aether-data build".to_string(),
)),
}
}
impl DataBackends {
fn sql_backend(&self) -> Option<SqlBackendRef<'_>> {
#[cfg(feature = "postgres")]
if let Some(postgres) = self.postgres.as_ref() {
return Some(SqlBackendRef::Postgres(postgres));
}
#[cfg(feature = "mysql")]
if let Some(mysql) = self.mysql.as_ref() {
return Some(SqlBackendRef::Mysql(mysql));
}
#[cfg(feature = "sqlite")]
if let Some(sqlite) = self.sqlite.as_ref() {
return Some(SqlBackendRef::Sqlite(sqlite));
}
None
}
pub fn from_config(config: DataLayerConfig) -> Result<Self, DataLayerError> {
config.validate()?;
let database = config.effective_database();
if let Some(database) = database.as_ref() {
ensure_driver_enabled(database.driver)?;
}
#[cfg(feature = "postgres")]
let postgres = match database.clone() {
Some(database) if database.driver == DatabaseDriver::Postgres => Some(
PostgresBackend::from_config(database.to_postgres_config()?)?,
),
_ => None,
};
#[cfg(feature = "mysql")]
let mysql = match database.clone() {
Some(database) if database.driver == DatabaseDriver::Mysql => {
Some(MysqlBackend::from_config(database)?)
}
_ => None,
};
#[cfg(feature = "sqlite")]
let sqlite = match database.clone() {
Some(database) if database.driver == DatabaseDriver::Sqlite => {
Some(SqliteBackend::from_config(database)?)
}
_ => None,
};
#[cfg(feature = "postgres")]
let leases = DataLeaseBackends::from_postgres(postgres.as_ref())?;
#[cfg(not(feature = "postgres"))]
let leases = DataLeaseBackends::default();
let read = DataReadRepositories::from_backends(
#[cfg(feature = "postgres")]
postgres.as_ref(),
#[cfg(feature = "mysql")]
mysql.as_ref(),
#[cfg(feature = "sqlite")]
sqlite.as_ref(),
);
#[cfg(feature = "postgres")]
let transactions = DataTransactionBackends::from_postgres(postgres.as_ref());
#[cfg(not(feature = "postgres"))]
let transactions = DataTransactionBackends::default();
let write = DataWriteRepositories::from_backends(
#[cfg(feature = "postgres")]
postgres.as_ref(),
#[cfg(feature = "mysql")]
mysql.as_ref(),
#[cfg(feature = "sqlite")]
sqlite.as_ref(),
);
Ok(Self {
config,
#[cfg(feature = "postgres")]
postgres,
#[cfg(feature = "mysql")]
mysql,
#[cfg(feature = "sqlite")]
sqlite,
leases,
read,
transactions,
write,
})
}
pub fn config(&self) -> &DataLayerConfig {
&self.config
}
#[cfg(feature = "postgres")]
pub fn postgres(&self) -> Option<&PostgresBackend> {
self.postgres.as_ref()
}
pub fn database_driver(&self) -> Option<DatabaseDriver> {
self.config
.effective_database()
.map(|database| database.driver)
}
#[cfg(feature = "mysql")]
pub fn mysql(&self) -> Option<&MysqlBackend> {
self.mysql.as_ref()
}
#[cfg(feature = "sqlite")]
pub fn sqlite(&self) -> Option<&SqliteBackend> {
self.sqlite.as_ref()
}
pub fn read(&self) -> &DataReadRepositories {
&self.read
}
pub fn leases(&self) -> &DataLeaseBackends {
&self.leases
}
pub fn transactions(&self) -> &DataTransactionBackends {
&self.transactions
}
pub fn write(&self) -> &DataWriteRepositories {
&self.write
}
pub fn has_runtime_backends(&self) -> bool {
self.leases.has_any()
|| self.read.has_any()
|| self.transactions.has_any()
|| self.write.has_any()
}
}
#[cfg(test)]
mod tests {
use super::DataBackends;
#[cfg(feature = "postgres")]
use crate::driver::postgres::PostgresPoolConfig;
use crate::{DataLayerConfig, DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
#[test]
#[cfg(not(feature = "mysql"))]
fn rejects_mysql_when_driver_is_not_enabled() {
let error = DataBackends::from_config(DataLayerConfig::from_database(SqlDatabaseConfig {
driver: DatabaseDriver::Mysql,
url: "mysql://user:pass@localhost/aether".to_string(),
pool: SqlPoolConfig::default(),
}))
.expect_err("disabled mysql should fail explicitly");
assert!(error.to_string().contains("MySQL driver is not enabled"));
}
#[test]
#[cfg(not(feature = "sqlite"))]
fn rejects_sqlite_when_driver_is_not_enabled() {
let error = DataBackends::from_config(DataLayerConfig::from_database(SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite://./data/aether.db".to_string(),
pool: SqlPoolConfig::default(),
}))
.expect_err("disabled sqlite should fail explicitly");
assert!(error.to_string().contains("SQLite driver is not enabled"));
}
#[test]
fn builds_empty_backends_from_default_config() {
let backends = DataBackends::from_config(DataLayerConfig::default())
.expect("empty config should be accepted");
assert!(!backends.has_runtime_backends());
#[cfg(feature = "postgres")]
assert!(backends.postgres().is_none());
#[cfg(feature = "mysql")]
assert!(backends.mysql().is_none());
#[cfg(feature = "sqlite")]
assert!(backends.sqlite().is_none());
#[cfg(feature = "postgres")]
assert!(backends.leases().postgres().is_none());
assert!(backends.read().auth_api_keys().is_none());
assert!(backends.read().auth_modules().is_none());
assert!(backends.read().billing().is_none());
assert!(backends.read().gemini_file_mappings().is_none());
assert!(backends.read().global_models().is_none());
assert!(backends.read().management_tokens().is_none());
assert!(backends.read().oauth_providers().is_none());
assert!(backends.read().proxy_nodes().is_none());
assert!(backends.read().minimal_candidate_selection().is_none());
assert!(backends.read().request_candidates().is_none());
assert!(backends.read().provider_catalog().is_none());
assert!(backends.read().usage().is_none());
assert!(backends.read().video_tasks().is_none());
#[cfg(feature = "postgres")]
assert!(backends.transactions().postgres().is_none());
assert!(backends.write().settlement().is_none());
assert!(backends.write().usage().is_none());
}
#[tokio::test]
#[cfg(feature = "postgres")]
async fn builds_postgres_backend_from_config() {
let backends = DataBackends::from_config(DataLayerConfig {
database: None,
postgres: Some(PostgresPoolConfig {
database_url: "postgres://localhost/aether".to_string(),
min_connections: 1,
max_connections: 4,
acquire_timeout_ms: 1_000,
idle_timeout_ms: 5_000,
max_lifetime_ms: 30_000,
statement_cache_capacity: 64,
require_ssl: false,
}),
})
.expect("postgres backend should build");
assert!(backends.has_runtime_backends());
#[cfg(feature = "postgres")]
assert!(backends.postgres().is_some());
#[cfg(feature = "mysql")]
assert!(backends.mysql().is_none());
#[cfg(feature = "sqlite")]
assert!(backends.sqlite().is_none());
#[cfg(feature = "postgres")]
assert!(backends.leases().postgres().is_some());
assert!(backends.read().auth_api_keys().is_some());
assert!(backends.read().auth_modules().is_some());
assert!(backends.read().billing().is_some());
assert!(backends.read().gemini_file_mappings().is_some());
assert!(backends.read().global_models().is_some());
assert!(backends.read().management_tokens().is_some());
assert!(backends.read().minimal_candidate_selection().is_some());
assert!(backends.read().oauth_providers().is_some());
assert!(backends.read().proxy_nodes().is_some());
assert!(backends.read().minimal_candidate_selection().is_some());
assert!(backends.read().request_candidates().is_some());
assert!(backends.read().provider_catalog().is_some());
assert!(backends.read().provider_quotas().is_some());
assert!(backends.read().usage().is_some());
assert!(backends.read().video_tasks().is_some());
assert!(backends.read().wallets().is_some());
assert!(backends.transactions().postgres().is_some());
assert!(backends.write().auth_modules().is_some());
assert!(backends.write().gemini_file_mappings().is_some());
assert!(backends.write().management_tokens().is_some());
assert!(backends.write().oauth_providers().is_some());
assert!(backends.write().proxy_nodes().is_some());
assert!(backends.write().provider_catalog().is_some());
assert!(backends.write().provider_quotas().is_some());
assert!(backends.write().settlement().is_some());
assert!(backends.write().usage().is_some());
assert!(backends.write().wallets().is_some());
assert!(backends.config().effective_database().is_some());
}
#[tokio::test]
#[cfg(feature = "mysql")]
async fn builds_mysql_backend_from_database_config_with_first_core_repository() {
let backends = DataBackends::from_config(DataLayerConfig {
database: Some(SqlDatabaseConfig {
driver: DatabaseDriver::Mysql,
url: "mysql://user:pass@localhost:3306/aether".to_string(),
pool: SqlPoolConfig::default(),
}),
postgres: None,
})
.expect("mysql backend should build");
assert!(backends.has_runtime_backends());
#[cfg(feature = "postgres")]
assert!(backends.postgres().is_none());
#[cfg(feature = "mysql")]
assert!(backends.mysql().is_some());
#[cfg(feature = "sqlite")]
assert!(backends.sqlite().is_none());
assert!(backends.read().has_any());
assert!(backends.read().announcements().is_some());
assert!(backends.read().auth_api_keys().is_some());
assert!(backends.read().auth_modules().is_some());
assert!(backends.read().billing().is_some());
assert!(backends.read().gemini_file_mappings().is_some());
assert!(backends.read().global_models().is_some());
assert!(backends.read().management_tokens().is_some());
assert!(backends.read().minimal_candidate_selection().is_some());
assert!(backends.read().oauth_providers().is_some());
assert!(backends.read().provider_catalog().is_some());
assert!(backends.read().provider_quotas().is_some());
assert!(backends.read().proxy_nodes().is_some());
assert!(backends.read().request_candidates().is_some());
assert!(backends.read().users().is_some());
assert!(backends.read().video_tasks().is_some());
assert!(backends.has_stats_hourly_aggregation_backend());
assert!(backends.has_stats_daily_aggregation_backend());
assert!(backends.write().has_any());
assert!(backends.write().announcements().is_some());
assert!(backends.write().auth_api_keys().is_some());
assert!(backends.write().auth_modules().is_some());
assert!(backends.write().gemini_file_mappings().is_some());
assert!(backends.write().global_models().is_some());
assert!(backends.write().management_tokens().is_some());
assert!(backends.write().oauth_providers().is_some());
assert!(backends.write().proxy_nodes().is_some());
assert!(backends.write().provider_catalog().is_some());
assert!(backends.write().provider_quotas().is_some());
assert!(backends.write().request_candidates().is_some());
assert!(backends.write().video_tasks().is_some());
assert!(backends.write().wallets().is_some());
assert!(backends.config().effective_database().is_some());
}
#[tokio::test]
#[cfg(feature = "sqlite")]
async fn builds_sqlite_backend_from_database_config_with_first_core_repository() {
let backends = DataBackends::from_config(DataLayerConfig {
database: Some(SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite://./data/aether.db".to_string(),
pool: SqlPoolConfig::default(),
}),
postgres: None,
})
.expect("sqlite backend should build");
assert!(backends.has_runtime_backends());
#[cfg(feature = "postgres")]
assert!(backends.postgres().is_none());
#[cfg(feature = "mysql")]
assert!(backends.mysql().is_none());
#[cfg(feature = "sqlite")]
assert!(backends.sqlite().is_some());
assert!(backends.read().has_any());
assert!(backends.read().announcements().is_some());
assert!(backends.read().auth_api_keys().is_some());
assert!(backends.read().auth_modules().is_some());
assert!(backends.read().billing().is_some());
assert!(backends.read().gemini_file_mappings().is_some());
assert!(backends.read().global_models().is_some());
assert!(backends.read().management_tokens().is_some());
assert!(backends.read().oauth_providers().is_some());
assert!(backends.read().provider_catalog().is_some());
assert!(backends.read().provider_quotas().is_some());
assert!(backends.read().proxy_nodes().is_some());
assert!(backends.read().request_candidates().is_some());
assert!(backends.read().users().is_some());
assert!(backends.read().video_tasks().is_some());
assert!(backends.has_stats_hourly_aggregation_backend());
assert!(backends.has_stats_daily_aggregation_backend());
assert!(backends.write().has_any());
assert!(backends.write().announcements().is_some());
assert!(backends.write().auth_api_keys().is_some());
assert!(backends.write().auth_modules().is_some());
assert!(backends.write().gemini_file_mappings().is_some());
assert!(backends.write().global_models().is_some());
assert!(backends.write().management_tokens().is_some());
assert!(backends.write().oauth_providers().is_some());
assert!(backends.write().proxy_nodes().is_some());
assert!(backends.write().provider_catalog().is_some());
assert!(backends.write().provider_quotas().is_some());
assert!(backends.write().request_candidates().is_some());
assert!(backends.write().video_tasks().is_some());
assert!(backends.write().wallets().is_some());
assert!(backends.config().effective_database().is_some());
}
}
@@ -0,0 +1,646 @@
use std::sync::Arc;
use crate::database::SqlDatabaseConfig;
use crate::driver::mysql::{MysqlPool, MysqlPoolFactory};
use crate::repository::announcements::{
AnnouncementReadRepository, AnnouncementWriteRepository, MysqlAnnouncementRepository,
};
use crate::repository::audit::{AuditLogReadRepository, MysqlAuditLogReadRepository};
use crate::repository::auth::{
AuthApiKeyReadRepository, AuthApiKeyWriteRepository, MysqlAuthApiKeyReadRepository,
};
use crate::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, MysqlAuthModuleReadRepository,
MysqlAuthModuleRepository,
};
use crate::repository::background_tasks::{
BackgroundTaskReadRepository, BackgroundTaskWriteRepository, MysqlBackgroundTaskRepository,
};
use crate::repository::billing::{BillingReadRepository, MysqlBillingReadRepository};
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, MysqlMinimalCandidateSelectionReadRepository,
};
use crate::repository::candidates::{
MysqlRequestCandidateRepository, RequestCandidateReadRepository,
RequestCandidateWriteRepository,
};
use crate::repository::gemini_file_mappings::{
GeminiFileMappingReadRepository, GeminiFileMappingWriteRepository,
MysqlGeminiFileMappingRepository,
};
use crate::repository::global_models::{
GlobalModelReadRepository, GlobalModelWriteRepository, MysqlGlobalModelReadRepository,
};
use crate::repository::management_tokens::{
ManagementTokenReadRepository, ManagementTokenWriteRepository, MysqlManagementTokenRepository,
};
use crate::repository::oauth_providers::{
MysqlOAuthProviderRepository, OAuthProviderReadRepository, OAuthProviderWriteRepository,
};
use crate::repository::pool_scores::{
MysqlPoolMemberScoreRepository, PoolMemberScoreWriteRepository, PoolScoreReadRepository,
};
use crate::repository::provider_catalog::{
MysqlProviderCatalogReadRepository, ProviderCatalogReadRepository,
ProviderCatalogWriteRepository,
};
use crate::repository::proxy_nodes::{
MysqlProxyNodeReadRepository, ProxyNodeReadRepository, ProxyNodeWriteRepository,
};
use crate::repository::quota::{
MysqlProviderQuotaRepository, ProviderQuotaReadRepository, ProviderQuotaWriteRepository,
};
use crate::repository::routing_profiles::{
MysqlRoutingGroupRepository, RoutingGroupReadRepository, RoutingGroupWriteRepository,
};
use crate::repository::settlement::{MysqlSettlementRepository, SettlementWriteRepository};
use crate::repository::usage::{
MysqlUsageReadRepository, MysqlUsageWriteRepository, UsageReadRepository, UsageWriteRepository,
};
use crate::repository::users::{MysqlUserReadRepository, UserReadRepository};
use crate::repository::video_tasks::{
MysqlVideoTaskRepository, VideoTaskReadRepository, VideoTaskWriteRepository,
};
use crate::repository::wallet::{
MysqlWalletReadRepository, WalletReadRepository, WalletWriteRepository,
};
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct MysqlBackend {
config: SqlDatabaseConfig,
pool: MysqlPool,
}
impl MysqlBackend {
pub fn from_config(config: SqlDatabaseConfig) -> Result<Self, DataLayerError> {
let factory = MysqlPoolFactory::new(config.clone())?;
let pool = factory.connect_lazy()?;
Ok(Self { config, pool })
}
pub fn config(&self) -> &SqlDatabaseConfig {
&self.config
}
pub fn pool(&self) -> &MysqlPool {
&self.pool
}
pub fn pool_clone(&self) -> MysqlPool {
self.pool.clone()
}
pub fn auth_api_key_read_repository(&self) -> Arc<dyn AuthApiKeyReadRepository> {
Arc::new(MysqlAuthApiKeyReadRepository::new(self.pool_clone()))
}
pub fn announcement_read_repository(&self) -> Arc<dyn AnnouncementReadRepository> {
Arc::new(MysqlAnnouncementRepository::new(self.pool_clone()))
}
pub fn audit_log_read_repository(&self) -> Arc<dyn AuditLogReadRepository> {
Arc::new(MysqlAuditLogReadRepository::new(self.pool_clone()))
}
pub fn announcement_write_repository(&self) -> Arc<dyn AnnouncementWriteRepository> {
Arc::new(MysqlAnnouncementRepository::new(self.pool_clone()))
}
pub fn auth_api_key_write_repository(&self) -> Arc<dyn AuthApiKeyWriteRepository> {
Arc::new(MysqlAuthApiKeyReadRepository::new(self.pool_clone()))
}
pub fn management_token_read_repository(&self) -> Arc<dyn ManagementTokenReadRepository> {
Arc::new(MysqlManagementTokenRepository::new(self.pool_clone()))
}
pub fn management_token_write_repository(&self) -> Arc<dyn ManagementTokenWriteRepository> {
Arc::new(MysqlManagementTokenRepository::new(self.pool_clone()))
}
pub fn auth_module_read_repository(&self) -> Arc<dyn AuthModuleReadRepository> {
Arc::new(MysqlAuthModuleReadRepository::new(self.pool_clone()))
}
pub fn auth_module_write_repository(&self) -> Arc<dyn AuthModuleWriteRepository> {
Arc::new(MysqlAuthModuleRepository::new(self.pool_clone()))
}
pub fn billing_read_repository(&self) -> Arc<dyn BillingReadRepository> {
Arc::new(MysqlBillingReadRepository::new(self.pool_clone()))
}
pub fn background_task_read_repository(&self) -> Arc<dyn BackgroundTaskReadRepository> {
Arc::new(MysqlBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn background_task_write_repository(&self) -> Arc<dyn BackgroundTaskWriteRepository> {
Arc::new(MysqlBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn request_candidate_read_repository(&self) -> Arc<dyn RequestCandidateReadRepository> {
Arc::new(MysqlRequestCandidateRepository::new(self.pool_clone()))
}
pub fn request_candidate_write_repository(&self) -> Arc<dyn RequestCandidateWriteRepository> {
Arc::new(MysqlRequestCandidateRepository::new(self.pool_clone()))
}
pub fn minimal_candidate_selection_read_repository(
&self,
) -> Arc<dyn MinimalCandidateSelectionReadRepository> {
Arc::new(MysqlMinimalCandidateSelectionReadRepository::new(
self.pool_clone(),
))
}
pub fn gemini_file_mapping_read_repository(&self) -> Arc<dyn GeminiFileMappingReadRepository> {
Arc::new(MysqlGeminiFileMappingRepository::new(self.pool_clone()))
}
pub fn gemini_file_mapping_write_repository(
&self,
) -> Arc<dyn GeminiFileMappingWriteRepository> {
Arc::new(MysqlGeminiFileMappingRepository::new(self.pool_clone()))
}
pub fn global_model_read_repository(&self) -> Arc<dyn GlobalModelReadRepository> {
Arc::new(MysqlGlobalModelReadRepository::new(self.pool_clone()))
}
pub fn global_model_write_repository(&self) -> Arc<dyn GlobalModelWriteRepository> {
Arc::new(MysqlGlobalModelReadRepository::new(self.pool_clone()))
}
pub fn oauth_provider_read_repository(&self) -> Arc<dyn OAuthProviderReadRepository> {
Arc::new(MysqlOAuthProviderRepository::new(self.pool_clone()))
}
pub fn oauth_provider_write_repository(&self) -> Arc<dyn OAuthProviderWriteRepository> {
Arc::new(MysqlOAuthProviderRepository::new(self.pool_clone()))
}
pub fn provider_catalog_read_repository(&self) -> Arc<dyn ProviderCatalogReadRepository> {
Arc::new(MysqlProviderCatalogReadRepository::new(self.pool_clone()))
}
pub fn provider_catalog_write_repository(&self) -> Arc<dyn ProviderCatalogWriteRepository> {
Arc::new(MysqlProviderCatalogReadRepository::new(self.pool_clone()))
}
pub fn pool_score_read_repository(&self) -> Arc<dyn PoolScoreReadRepository> {
Arc::new(MysqlPoolMemberScoreRepository::new(self.pool_clone()))
}
pub fn pool_score_write_repository(&self) -> Arc<dyn PoolMemberScoreWriteRepository> {
Arc::new(MysqlPoolMemberScoreRepository::new(self.pool_clone()))
}
pub fn routing_group_read_repository(&self) -> Arc<dyn RoutingGroupReadRepository> {
Arc::new(MysqlRoutingGroupRepository::new(self.pool_clone()))
}
pub fn routing_group_write_repository(&self) -> Arc<dyn RoutingGroupWriteRepository> {
Arc::new(MysqlRoutingGroupRepository::new(self.pool_clone()))
}
pub fn proxy_node_read_repository(&self) -> Arc<dyn ProxyNodeReadRepository> {
Arc::new(MysqlProxyNodeReadRepository::new(self.pool_clone()))
}
pub fn proxy_node_write_repository(&self) -> Arc<dyn ProxyNodeWriteRepository> {
Arc::new(MysqlProxyNodeReadRepository::new(self.pool_clone()))
}
pub fn provider_quota_read_repository(&self) -> Arc<dyn ProviderQuotaReadRepository> {
Arc::new(MysqlProviderQuotaRepository::new(self.pool_clone()))
}
pub fn provider_quota_write_repository(&self) -> Arc<dyn ProviderQuotaWriteRepository> {
Arc::new(MysqlProviderQuotaRepository::new(self.pool_clone()))
}
pub fn settlement_write_repository(&self) -> Arc<dyn SettlementWriteRepository> {
Arc::new(MysqlSettlementRepository::new(self.pool_clone()))
}
pub fn usage_write_repository(&self) -> Arc<dyn UsageWriteRepository> {
Arc::new(MysqlUsageWriteRepository::new(self.pool_clone()))
}
pub fn usage_read_repository(&self) -> Arc<dyn UsageReadRepository> {
Arc::new(MysqlUsageReadRepository::new(self.pool_clone()))
}
pub fn user_read_repository(&self) -> Arc<dyn UserReadRepository> {
Arc::new(MysqlUserReadRepository::new(self.pool_clone()))
}
pub fn video_task_read_repository(&self) -> Arc<dyn VideoTaskReadRepository> {
Arc::new(MysqlVideoTaskRepository::new(self.pool_clone()))
}
pub fn video_task_write_repository(&self) -> Arc<dyn VideoTaskWriteRepository> {
Arc::new(MysqlVideoTaskRepository::new(self.pool_clone()))
}
pub fn wallet_read_repository(&self) -> Arc<dyn WalletReadRepository> {
Arc::new(MysqlWalletReadRepository::new(self.pool_clone()))
}
pub fn wallet_write_repository(&self) -> Arc<dyn WalletWriteRepository> {
Arc::new(MysqlWalletReadRepository::new(self.pool_clone()))
}
}
#[cfg(test)]
mod tests {
use super::MysqlBackend;
use crate::lifecycle::migrate::run_mysql_migrations;
use crate::{
DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, StatsDailyAggregationInput,
StatsHourlyAggregationInput, WalletDailyUsageAggregationInput,
};
#[tokio::test]
async fn backend_retains_config_and_pool() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Mysql,
url: "mysql://user:pass@localhost:3306/aether".to_string(),
pool: SqlPoolConfig::default(),
};
let backend = MysqlBackend::from_config(config.clone()).expect("backend should build");
assert_eq!(backend.config(), &config);
let _pool = backend.pool();
let _pool_clone = backend.pool_clone();
}
#[tokio::test]
async fn mysql_wallet_daily_usage_aggregation_uses_settlement_wallets_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 wallet daily usage aggregation smoke test because AETHER_TEST_MYSQL_URL is unset"
);
return;
};
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Mysql,
url: database_url,
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = MysqlBackend::from_config(config).expect("backend should build");
run_mysql_migrations(backend.pool())
.await
.expect("mysql migrations should run");
let suffix = format!(
"{}-{}",
std::process::id(),
chrono::Utc::now().timestamp_nanos_opt().unwrap_or_default()
);
let wallet_id = format!("wallet-daily-{suffix}");
let stale_wallet_id = format!("wallet-daily-stale-{suffix}");
let timezone = format!("Test/WalletDaily/{suffix}");
let request_one = format!("request-daily-1-{suffix}");
let request_two = format!("request-daily-2-{suffix}");
let request_zero = format!("request-daily-zero-{suffix}");
let request_outside = format!("request-daily-outside-{suffix}");
let stale_ledger_id = format!("stale-ledger-{suffix}");
let unique_offset = chrono::Utc::now()
.timestamp_nanos_opt()
.unwrap_or_default()
.rem_euclid(10_000_000);
let window_start = 4_100_000_000_i64 + unique_offset * 1_000;
let window_end = window_start + 200;
let first_finalized_at = window_start;
let last_finalized_at = window_start + 100;
let zero_finalized_at = window_start + 150;
let outside_finalized_at = window_end;
let seed_created_at = window_start - 100;
let aggregated_at = window_end + 100;
sqlx::query(
r#"
INSERT INTO wallets (id, user_id, balance, gift_balance, limit_mode, created_at, updated_at)
VALUES
(?, ?, 10.0, 2.0, 'finite', 1, 1),
(?, ?, 0.0, 0.0, 'finite', 1, 1)
"#,
)
.bind(&wallet_id)
.bind(format!("user-{wallet_id}"))
.bind(&stale_wallet_id)
.bind(format!("user-{stale_wallet_id}"))
.execute(backend.pool())
.await
.expect("wallets should seed");
sqlx::query(
r#"
INSERT INTO `usage` (
request_id, wallet_id, provider_name, model, status, billing_status,
total_cost_usd, input_tokens, output_tokens, cache_creation_input_tokens,
cache_read_input_tokens, finalized_at, created_at_unix_ms, updated_at_unix_secs
) VALUES
(?, 'wrong-wallet', 'provider', 'model', 'completed', 'pending',
1.25, 10, 20, 3, 4, ?, ?, ?),
(?, NULL, 'provider', 'model', 'completed', 'pending',
2.00, 5, 7, 1, 2, ?, ?, ?),
(?, NULL, 'provider', 'model', 'completed', 'pending',
0.00, 100, 100, 0, 0, ?, ?, ?),
(?, NULL, 'provider', 'model', 'completed', 'pending',
9.00, 50, 50, 0, 0, ?, ?, ?)
"#,
)
.bind(&request_one)
.bind(seed_created_at)
.bind(seed_created_at * 1000)
.bind(seed_created_at)
.bind(&request_two)
.bind(seed_created_at + 1)
.bind((seed_created_at + 1) * 1000)
.bind(seed_created_at + 1)
.bind(&request_zero)
.bind(seed_created_at + 2)
.bind((seed_created_at + 2) * 1000)
.bind(seed_created_at + 2)
.bind(&request_outside)
.bind(seed_created_at + 3)
.bind((seed_created_at + 3) * 1000)
.bind(seed_created_at + 3)
.execute(backend.pool())
.await
.expect("usage should seed");
sqlx::query(
r#"
INSERT INTO usage_settlement_snapshots (
request_id, billing_status, wallet_id, finalized_at, created_at, updated_at
) VALUES
(?, 'settled', ?, ?, ?, ?),
(?, 'settled', ?, ?, ?, ?),
(?, 'settled', ?, ?, ?, ?),
(?, 'settled', ?, ?, ?, ?)
"#,
)
.bind(&request_one)
.bind(&wallet_id)
.bind(first_finalized_at)
.bind(first_finalized_at)
.bind(first_finalized_at)
.bind(&request_two)
.bind(&wallet_id)
.bind(last_finalized_at)
.bind(last_finalized_at)
.bind(last_finalized_at)
.bind(&request_zero)
.bind(&wallet_id)
.bind(zero_finalized_at)
.bind(zero_finalized_at)
.bind(zero_finalized_at)
.bind(&request_outside)
.bind(&wallet_id)
.bind(outside_finalized_at)
.bind(outside_finalized_at)
.bind(outside_finalized_at)
.execute(backend.pool())
.await
.expect("settlement snapshots 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 (?, ?, '2026-05-03', ?, 7.0, 3, 1, 1, 0, 0, ?, ?, ?)
"#,
)
.bind(&stale_ledger_id)
.bind(&stale_wallet_id)
.bind(&timezone)
.bind(seed_created_at)
.bind(seed_created_at)
.bind(seed_created_at)
.execute(backend.pool())
.await
.expect("stale ledger should seed");
let summary = backend
.aggregate_wallet_daily_usage(&WalletDailyUsageAggregationInput {
billing_date: "2026-05-03".to_string(),
billing_timezone: timezone.clone(),
window_start_unix_secs: window_start as u64,
window_end_unix_secs: window_end as u64,
aggregated_at_unix_secs: aggregated_at as u64,
})
.await
.expect("wallet daily usage aggregation should run");
assert_eq!(summary.aggregated_wallets, 1);
assert_eq!(summary.deleted_stale_ledgers, 1);
let ledger = sqlx::query_as::<
_,
(
String,
f64,
i64,
i64,
i64,
i64,
i64,
Option<i64>,
Option<i64>,
i64,
),
>(
r#"
SELECT
wallet_id,
total_cost_usd,
total_requests,
input_tokens,
output_tokens,
cache_creation_tokens,
cache_read_tokens,
first_finalized_at,
last_finalized_at,
aggregated_at
FROM wallet_daily_usage_ledgers
WHERE wallet_id = ?
AND billing_date = '2026-05-03'
AND billing_timezone = ?
"#,
)
.bind(&wallet_id)
.bind(&timezone)
.fetch_one(backend.pool())
.await
.expect("aggregated ledger should load");
assert_eq!(ledger.0, wallet_id);
assert!((ledger.1 - 3.25).abs() < f64::EPSILON);
assert_eq!(ledger.2, 2);
assert_eq!(ledger.3, 15);
assert_eq!(ledger.4, 27);
assert_eq!(ledger.5, 4);
assert_eq!(ledger.6, 6);
assert_eq!(ledger.7, Some(first_finalized_at));
assert_eq!(ledger.8, Some(last_finalized_at));
assert_eq!(ledger.9, aggregated_at);
let stale_count: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM wallet_daily_usage_ledgers WHERE id = ?")
.bind(&stale_ledger_id)
.fetch_one(backend.pool())
.await
.expect("stale ledger count should load");
assert_eq!(stale_count, 0);
}
#[tokio::test]
async fn mysql_stats_aggregation_runs_after_mysql_migrations_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 stats aggregation smoke test because AETHER_TEST_MYSQL_URL is unset"
);
return;
};
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Mysql,
url: database_url,
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = MysqlBackend::from_config(config).expect("backend should build");
run_mysql_migrations(backend.pool())
.await
.expect("mysql migrations should run");
for sql in [
"DELETE FROM stats_daily WHERE `date` = 0",
"DELETE FROM stats_hourly WHERE hour_utc = 3600",
"DELETE FROM usage_settlement_snapshots WHERE request_id LIKE 'request-daily-%' OR request_id LIKE 'stats-%'",
"DELETE FROM `usage` WHERE request_id LIKE 'request-%' OR request_id LIKE 'export-request-%' OR request_id LIKE 'stats-%'",
] {
sqlx::query(sql)
.execute(backend.pool())
.await
.expect("stats smoke cleanup should run");
}
sqlx::query(
r#"
INSERT INTO `usage` (
request_id, user_id, api_key_id, provider_name, model, status, billing_status,
status_code, error_category, input_tokens, output_tokens,
cache_creation_input_tokens, cache_read_input_tokens, total_cost_usd,
actual_total_cost_usd, response_time_ms, created_at_unix_ms, updated_at_unix_secs
) VALUES
('stats-1', 'user-1', 'key-1', 'provider-a', 'model-a', 'completed', 'settled',
200, NULL, 10, 20, 1, 2, 0.30, 0.25, 100, 3600000, 3600),
('stats-2', 'user-2', 'key-2', 'provider-b', 'model-b', 'failed', 'void',
500, 'upstream_error', 5, 7, 0, 1, 0.20, 0.20, 300, 3610000, 3610),
('stats-pending', 'user-3', 'key-3', 'provider-a', 'model-a', 'pending', 'pending',
NULL, NULL, 100, 100, 0, 0, 9.99, 9.99, 50, 3620000, 3620),
('stats-unknown-provider', 'user-4', 'key-4', 'unknown', 'model-a', 'completed', 'settled',
200, NULL, 100, 100, 0, 0, 9.99, 9.99, 50, 3630000, 3630)
"#,
)
.execute(backend.pool())
.await
.expect("usage stats rows should seed");
let target_hour = chrono::DateTime::<chrono::Utc>::from_timestamp(3600, 0)
.expect("target hour should be valid");
let aggregated_at = chrono::DateTime::<chrono::Utc>::from_timestamp(7200, 0)
.expect("aggregation time should be valid");
let hourly = backend
.aggregate_stats_hourly(&StatsHourlyAggregationInput {
target_hour_utc: target_hour,
aggregated_at,
})
.await
.expect("hourly stats aggregation should run")
.expect("hourly bucket should aggregate");
assert_eq!(hourly.hour_utc, target_hour);
assert_eq!(hourly.total_requests, 2);
assert_eq!(hourly.user_rows, 2);
assert_eq!(hourly.user_model_rows, 2);
assert_eq!(hourly.model_rows, 2);
assert_eq!(hourly.provider_rows, 2);
let hourly_row = sqlx::query_as::<_, (i64, i64, i64, i64, f64)>(
r#"
SELECT total_requests, success_requests, error_requests, input_tokens, total_cost
FROM stats_hourly
WHERE hour_utc = 3600
"#,
)
.fetch_one(backend.pool())
.await
.expect("hourly stats row should load");
assert_eq!(hourly_row.0, 2);
assert_eq!(hourly_row.1, 1);
assert_eq!(hourly_row.2, 1);
assert_eq!(hourly_row.3, 15);
assert!((hourly_row.4 - 0.50).abs() < f64::EPSILON);
let second_hourly = backend
.aggregate_stats_hourly(&StatsHourlyAggregationInput {
target_hour_utc: target_hour,
aggregated_at,
})
.await
.expect("second hourly aggregation should run");
assert!(second_hourly.is_none());
let target_day = chrono::DateTime::<chrono::Utc>::from_timestamp(0, 0)
.expect("target day should be valid");
let daily = backend
.aggregate_stats_daily(&StatsDailyAggregationInput {
target_day_utc: target_day,
aggregated_at,
})
.await
.expect("daily stats aggregation should run")
.expect("daily bucket should aggregate");
assert_eq!(daily.day_start_utc, target_day);
assert_eq!(daily.total_requests, 2);
assert_eq!(daily.model_rows, 2);
assert_eq!(daily.provider_rows, 2);
assert_eq!(daily.api_key_rows, 2);
assert_eq!(daily.error_rows, 1);
assert_eq!(daily.user_rows, 2);
let daily_row = sqlx::query_as::<_, (i64, i64, i64, i64)>(
r#"
SELECT total_requests, success_requests, error_requests, unique_models
FROM stats_daily
WHERE `date` = 0
"#,
)
.fetch_one(backend.pool())
.await
.expect("daily stats row should load");
assert_eq!(daily_row, (2, 1, 1, 2));
}
}
@@ -0,0 +1,330 @@
use std::sync::Arc;
use crate::driver::postgres::{
PostgresLeaseRunner, PostgresLeaseRunnerConfig, PostgresPool, PostgresPoolConfig,
PostgresPoolFactory, PostgresTransactionRunner,
};
use crate::repository::announcements::{
AnnouncementReadRepository, AnnouncementWriteRepository, SqlxAnnouncementReadRepository,
};
use crate::repository::audit::{AuditLogReadRepository, PostgresAuditLogReadRepository};
use crate::repository::auth::{
AuthApiKeyReadRepository, AuthApiKeyWriteRepository, SqlxAuthApiKeySnapshotReadRepository,
};
use crate::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, SqlxAuthModuleReadRepository,
SqlxAuthModuleRepository,
};
use crate::repository::background_tasks::{
BackgroundTaskReadRepository, BackgroundTaskWriteRepository, SqlxBackgroundTaskRepository,
};
use crate::repository::billing::{BillingReadRepository, SqlxBillingReadRepository};
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, SqlxMinimalCandidateSelectionReadRepository,
};
use crate::repository::candidates::{
RequestCandidateReadRepository, RequestCandidateWriteRepository,
SqlxRequestCandidateReadRepository,
};
use crate::repository::gemini_file_mappings::{
GeminiFileMappingReadRepository, GeminiFileMappingWriteRepository,
SqlxGeminiFileMappingRepository,
};
use crate::repository::global_models::{
GlobalModelReadRepository, GlobalModelWriteRepository, SqlxGlobalModelReadRepository,
};
use crate::repository::management_tokens::{
ManagementTokenReadRepository, ManagementTokenWriteRepository, SqlxManagementTokenRepository,
};
use crate::repository::oauth_providers::{
OAuthProviderReadRepository, OAuthProviderWriteRepository, SqlxOAuthProviderRepository,
};
use crate::repository::pool_scores::{
PoolMemberScoreWriteRepository, PoolScoreReadRepository, PostgresPoolMemberScoreRepository,
};
use crate::repository::provider_catalog::{
ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
SqlxProviderCatalogReadRepository,
};
use crate::repository::proxy_nodes::{
ProxyNodeReadRepository, ProxyNodeWriteRepository, SqlxProxyNodeRepository,
};
use crate::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, SqlxProviderQuotaRepository,
};
use crate::repository::routing_profiles::{
PostgresRoutingGroupRepository, RoutingGroupReadRepository, RoutingGroupWriteRepository,
};
use crate::repository::settlement::{SettlementWriteRepository, SqlxSettlementRepository};
use crate::repository::usage::{
SqlxUsageReadRepository, UsageReadRepository, UsageWriteRepository,
};
use crate::repository::users::{SqlxUserReadRepository, UserReadRepository};
use crate::repository::video_tasks::{
SqlxVideoTaskReadRepository, SqlxVideoTaskRepository, VideoTaskReadRepository,
VideoTaskWriteRepository,
};
use crate::repository::wallet::{
SqlxWalletRepository, WalletReadRepository, WalletWriteRepository,
};
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct PostgresBackend {
config: PostgresPoolConfig,
pool: PostgresPool,
}
impl PostgresBackend {
pub fn from_config(config: PostgresPoolConfig) -> Result<Self, DataLayerError> {
let factory = PostgresPoolFactory::new(config.clone())?;
let pool = factory.connect_lazy()?;
Ok(Self { config, pool })
}
pub fn config(&self) -> &PostgresPoolConfig {
&self.config
}
pub fn pool(&self) -> &PostgresPool {
&self.pool
}
pub fn pool_clone(&self) -> PostgresPool {
self.pool.clone()
}
pub fn auth_api_key_read_repository(&self) -> Arc<dyn AuthApiKeyReadRepository> {
Arc::new(SqlxAuthApiKeySnapshotReadRepository::new(self.pool_clone()))
}
pub fn announcement_read_repository(&self) -> Arc<dyn AnnouncementReadRepository> {
Arc::new(SqlxAnnouncementReadRepository::new(self.pool_clone()))
}
pub fn audit_log_read_repository(&self) -> Arc<dyn AuditLogReadRepository> {
Arc::new(PostgresAuditLogReadRepository::new(self.pool_clone()))
}
pub fn announcement_write_repository(&self) -> Arc<dyn AnnouncementWriteRepository> {
Arc::new(SqlxAnnouncementReadRepository::new(self.pool_clone()))
}
pub fn auth_api_key_write_repository(&self) -> Arc<dyn AuthApiKeyWriteRepository> {
Arc::new(SqlxAuthApiKeySnapshotReadRepository::new(self.pool_clone()))
}
pub fn auth_module_read_repository(&self) -> Arc<dyn AuthModuleReadRepository> {
Arc::new(SqlxAuthModuleReadRepository::new(self.pool_clone()))
}
pub fn auth_module_write_repository(&self) -> Arc<dyn AuthModuleWriteRepository> {
Arc::new(SqlxAuthModuleRepository::new(self.pool_clone()))
}
pub fn billing_read_repository(&self) -> Arc<dyn BillingReadRepository> {
Arc::new(SqlxBillingReadRepository::new(self.pool_clone()))
}
pub fn background_task_read_repository(&self) -> Arc<dyn BackgroundTaskReadRepository> {
Arc::new(SqlxBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn background_task_write_repository(&self) -> Arc<dyn BackgroundTaskWriteRepository> {
Arc::new(SqlxBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn minimal_candidate_selection_read_repository(
&self,
) -> Arc<dyn MinimalCandidateSelectionReadRepository> {
Arc::new(SqlxMinimalCandidateSelectionReadRepository::new(
self.pool_clone(),
))
}
pub fn request_candidate_read_repository(&self) -> Arc<dyn RequestCandidateReadRepository> {
Arc::new(SqlxRequestCandidateReadRepository::new(self.pool_clone()))
}
pub fn request_candidate_write_repository(&self) -> Arc<dyn RequestCandidateWriteRepository> {
Arc::new(SqlxRequestCandidateReadRepository::new(self.pool_clone()))
}
pub fn gemini_file_mapping_read_repository(&self) -> Arc<dyn GeminiFileMappingReadRepository> {
Arc::new(SqlxGeminiFileMappingRepository::new(self.pool_clone()))
}
pub fn gemini_file_mapping_write_repository(
&self,
) -> Arc<dyn GeminiFileMappingWriteRepository> {
Arc::new(SqlxGeminiFileMappingRepository::new(self.pool_clone()))
}
pub fn global_model_read_repository(&self) -> Arc<dyn GlobalModelReadRepository> {
Arc::new(SqlxGlobalModelReadRepository::new(self.pool_clone()))
}
pub fn global_model_write_repository(&self) -> Arc<dyn GlobalModelWriteRepository> {
Arc::new(SqlxGlobalModelReadRepository::new(self.pool_clone()))
}
pub fn management_token_read_repository(&self) -> Arc<dyn ManagementTokenReadRepository> {
Arc::new(SqlxManagementTokenRepository::new(self.pool_clone()))
}
pub fn management_token_write_repository(&self) -> Arc<dyn ManagementTokenWriteRepository> {
Arc::new(SqlxManagementTokenRepository::new(self.pool_clone()))
}
pub fn oauth_provider_read_repository(&self) -> Arc<dyn OAuthProviderReadRepository> {
Arc::new(SqlxOAuthProviderRepository::new(self.pool_clone()))
}
pub fn oauth_provider_write_repository(&self) -> Arc<dyn OAuthProviderWriteRepository> {
Arc::new(SqlxOAuthProviderRepository::new(self.pool_clone()))
}
pub fn proxy_node_read_repository(&self) -> Arc<dyn ProxyNodeReadRepository> {
Arc::new(SqlxProxyNodeRepository::new(self.pool_clone()))
}
pub fn proxy_node_write_repository(&self) -> Arc<dyn ProxyNodeWriteRepository> {
Arc::new(SqlxProxyNodeRepository::new(self.pool_clone()))
}
pub fn provider_catalog_read_repository(&self) -> Arc<dyn ProviderCatalogReadRepository> {
Arc::new(SqlxProviderCatalogReadRepository::new(self.pool_clone()))
}
pub fn provider_catalog_write_repository(&self) -> Arc<dyn ProviderCatalogWriteRepository> {
Arc::new(SqlxProviderCatalogReadRepository::new(self.pool_clone()))
}
pub fn pool_score_read_repository(&self) -> Arc<dyn PoolScoreReadRepository> {
Arc::new(PostgresPoolMemberScoreRepository::new(self.pool_clone()))
}
pub fn pool_score_write_repository(&self) -> Arc<dyn PoolMemberScoreWriteRepository> {
Arc::new(PostgresPoolMemberScoreRepository::new(self.pool_clone()))
}
pub fn routing_group_read_repository(&self) -> Arc<dyn RoutingGroupReadRepository> {
Arc::new(PostgresRoutingGroupRepository::new(self.pool_clone()))
}
pub fn routing_group_write_repository(&self) -> Arc<dyn RoutingGroupWriteRepository> {
Arc::new(PostgresRoutingGroupRepository::new(self.pool_clone()))
}
pub fn provider_quota_read_repository(&self) -> Arc<dyn ProviderQuotaReadRepository> {
Arc::new(SqlxProviderQuotaRepository::new(self.pool_clone()))
}
pub fn usage_read_repository(&self) -> Arc<dyn UsageReadRepository> {
Arc::new(SqlxUsageReadRepository::new(self.pool_clone()))
}
pub fn user_read_repository(&self) -> Arc<dyn UserReadRepository> {
Arc::new(SqlxUserReadRepository::new(self.pool_clone()))
}
pub fn usage_write_repository(&self) -> Arc<dyn UsageWriteRepository> {
Arc::new(SqlxUsageReadRepository::new(self.pool_clone()))
}
pub fn wallet_read_repository(&self) -> Arc<dyn WalletReadRepository> {
Arc::new(SqlxWalletRepository::new(self.pool_clone()))
}
pub fn wallet_write_repository(&self) -> Arc<dyn WalletWriteRepository> {
Arc::new(SqlxWalletRepository::new(self.pool_clone()))
}
pub fn settlement_write_repository(&self) -> Arc<dyn SettlementWriteRepository> {
Arc::new(SqlxSettlementRepository::new(self.pool_clone()))
}
pub fn video_task_read_repository(&self) -> Arc<dyn VideoTaskReadRepository> {
Arc::new(SqlxVideoTaskReadRepository::new(self.pool_clone()))
}
pub fn video_task_write_repository(&self) -> Arc<dyn VideoTaskWriteRepository> {
Arc::new(SqlxVideoTaskRepository::new(self.pool_clone()))
}
pub fn transaction_runner(&self) -> PostgresTransactionRunner {
PostgresTransactionRunner::new(self.pool_clone())
}
pub fn lease_runner(
&self,
config: PostgresLeaseRunnerConfig,
) -> Result<PostgresLeaseRunner, DataLayerError> {
PostgresLeaseRunner::new(self.transaction_runner(), config)
}
pub fn provider_quota_write_repository(&self) -> Arc<dyn ProviderQuotaWriteRepository> {
Arc::new(SqlxProviderQuotaRepository::new(self.pool_clone()))
}
}
#[cfg(test)]
mod tests {
use super::PostgresBackend;
use crate::driver::postgres::{PostgresLeaseRunnerConfig, PostgresPoolConfig};
#[tokio::test]
async fn backend_retains_config_and_pool() {
let config = PostgresPoolConfig {
database_url: "postgres://localhost/aether".to_string(),
min_connections: 1,
max_connections: 4,
acquire_timeout_ms: 1_000,
idle_timeout_ms: 5_000,
max_lifetime_ms: 30_000,
statement_cache_capacity: 64,
require_ssl: false,
};
let backend =
PostgresBackend::from_config(config.clone()).expect("backend should build lazily");
assert_eq!(backend.config(), &config);
let _pool = backend.pool();
let _pool_clone = backend.pool_clone();
let _auth_api_key_reader = backend.auth_api_key_read_repository();
let _auth_api_key_writer = backend.auth_api_key_write_repository();
let _auth_module_reader = backend.auth_module_read_repository();
let _billing_reader = backend.billing_read_repository();
let _gemini_file_mapping_reader = backend.gemini_file_mapping_read_repository();
let _global_model_reader = backend.global_model_read_repository();
let _global_model_writer = backend.global_model_write_repository();
let _management_token_reader = backend.management_token_read_repository();
let _management_token_writer = backend.management_token_write_repository();
let _oauth_provider_reader = backend.oauth_provider_read_repository();
let _oauth_provider_writer = backend.oauth_provider_write_repository();
let _proxy_node_reader = backend.proxy_node_read_repository();
let _proxy_node_writer = backend.proxy_node_write_repository();
let _minimal_candidate_selection_reader =
backend.minimal_candidate_selection_read_repository();
let _request_candidate_reader = backend.request_candidate_read_repository();
let _request_candidate_writer = backend.request_candidate_write_repository();
let _gemini_file_mapping_writer = backend.gemini_file_mapping_write_repository();
let _provider_catalog_reader = backend.provider_catalog_read_repository();
let _provider_catalog_writer = backend.provider_catalog_write_repository();
let _provider_quota_reader = backend.provider_quota_read_repository();
let _usage_reader = backend.usage_read_repository();
let _usage_writer = backend.usage_write_repository();
let _wallet_reader = backend.wallet_read_repository();
let _wallet_writer = backend.wallet_write_repository();
let _settlement_writer = backend.settlement_write_repository();
let _video_task_reader = backend.video_task_read_repository();
let _video_task_writer = backend.video_task_write_repository();
let _transaction_runner = backend.transaction_runner();
let _lease_runner = backend
.lease_runner(PostgresLeaseRunnerConfig::default())
.expect("lease runner should build");
let _provider_quota_writer = backend.provider_quota_write_repository();
}
}
@@ -0,0 +1,491 @@
use std::fmt;
use std::sync::Arc;
#[cfg(feature = "mysql")]
use super::MysqlBackend;
#[cfg(feature = "postgres")]
use super::PostgresBackend;
#[cfg(feature = "sqlite")]
use super::SqliteBackend;
use crate::repository::announcements::AnnouncementReadRepository;
use crate::repository::audit::AuditLogReadRepository;
use crate::repository::auth::AuthApiKeyReadRepository;
use crate::repository::auth_modules::AuthModuleReadRepository;
use crate::repository::background_tasks::BackgroundTaskReadRepository;
use crate::repository::billing::BillingReadRepository;
use crate::repository::candidate_selection::MinimalCandidateSelectionReadRepository;
use crate::repository::candidates::RequestCandidateReadRepository;
use crate::repository::gemini_file_mappings::GeminiFileMappingReadRepository;
use crate::repository::global_models::GlobalModelReadRepository;
use crate::repository::management_tokens::ManagementTokenReadRepository;
use crate::repository::oauth_providers::OAuthProviderReadRepository;
use crate::repository::pool_scores::PoolScoreReadRepository;
use crate::repository::provider_catalog::ProviderCatalogReadRepository;
use crate::repository::proxy_nodes::ProxyNodeReadRepository;
use crate::repository::quota::ProviderQuotaReadRepository;
use crate::repository::routing_profiles::RoutingGroupReadRepository;
use crate::repository::usage::UsageReadRepository;
use crate::repository::users::UserReadRepository;
use crate::repository::video_tasks::VideoTaskReadRepository;
use crate::repository::wallet::WalletReadRepository;
#[derive(Clone, Default)]
pub struct DataReadRepositories {
announcements: Option<Arc<dyn AnnouncementReadRepository>>,
audit_logs: Option<Arc<dyn AuditLogReadRepository>>,
auth_api_keys: Option<Arc<dyn AuthApiKeyReadRepository>>,
auth_modules: Option<Arc<dyn AuthModuleReadRepository>>,
background_tasks: Option<Arc<dyn BackgroundTaskReadRepository>>,
billing: Option<Arc<dyn BillingReadRepository>>,
gemini_file_mappings: Option<Arc<dyn GeminiFileMappingReadRepository>>,
global_models: Option<Arc<dyn GlobalModelReadRepository>>,
management_tokens: Option<Arc<dyn ManagementTokenReadRepository>>,
oauth_providers: Option<Arc<dyn OAuthProviderReadRepository>>,
pool_scores: Option<Arc<dyn PoolScoreReadRepository>>,
proxy_nodes: Option<Arc<dyn ProxyNodeReadRepository>>,
minimal_candidate_selection: Option<Arc<dyn MinimalCandidateSelectionReadRepository>>,
request_candidates: Option<Arc<dyn RequestCandidateReadRepository>>,
provider_catalog: Option<Arc<dyn ProviderCatalogReadRepository>>,
provider_quotas: Option<Arc<dyn ProviderQuotaReadRepository>>,
routing_groups: Option<Arc<dyn RoutingGroupReadRepository>>,
usage: Option<Arc<dyn UsageReadRepository>>,
users: Option<Arc<dyn UserReadRepository>>,
video_tasks: Option<Arc<dyn VideoTaskReadRepository>>,
wallets: Option<Arc<dyn WalletReadRepository>>,
}
impl fmt::Debug for DataReadRepositories {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DataReadRepositories")
.field("has_auth_api_keys", &self.auth_api_keys.is_some())
.field("has_announcements", &self.announcements.is_some())
.field("has_audit_logs", &self.audit_logs.is_some())
.field("has_auth_modules", &self.auth_modules.is_some())
.field("has_background_tasks", &self.background_tasks.is_some())
.field("has_billing", &self.billing.is_some())
.field(
"has_gemini_file_mappings",
&self.gemini_file_mappings.is_some(),
)
.field("has_global_models", &self.global_models.is_some())
.field("has_management_tokens", &self.management_tokens.is_some())
.field("has_oauth_providers", &self.oauth_providers.is_some())
.field("has_pool_scores", &self.pool_scores.is_some())
.field("has_proxy_nodes", &self.proxy_nodes.is_some())
.field(
"has_minimal_candidate_selection",
&self.minimal_candidate_selection.is_some(),
)
.field("has_request_candidates", &self.request_candidates.is_some())
.field("has_provider_catalog", &self.provider_catalog.is_some())
.field("has_provider_quotas", &self.provider_quotas.is_some())
.field("has_routing_groups", &self.routing_groups.is_some())
.field("has_usage", &self.usage.is_some())
.field("has_users", &self.users.is_some())
.field("has_video_tasks", &self.video_tasks.is_some())
.field("has_wallets", &self.wallets.is_some())
.finish()
}
}
impl DataReadRepositories {
pub(crate) fn from_backends(
#[cfg(feature = "postgres")] postgres: Option<&PostgresBackend>,
#[cfg(feature = "mysql")] mysql: Option<&MysqlBackend>,
#[cfg(feature = "sqlite")] sqlite: Option<&SqliteBackend>,
) -> Self {
let mut repositories = Self::default();
#[cfg(feature = "postgres")]
if let Some(postgres) = postgres {
repositories.install_postgres(postgres);
}
#[cfg(feature = "mysql")]
if let Some(mysql) = mysql {
repositories.install_mysql(mysql);
}
#[cfg(feature = "sqlite")]
if let Some(sqlite) = sqlite {
repositories.install_sqlite(sqlite);
}
repositories
}
#[cfg(feature = "postgres")]
fn install_postgres(&mut self, backend: &PostgresBackend) {
if self.announcements.is_none() {
self.announcements = Some(PostgresBackend::announcement_read_repository(backend));
}
if self.audit_logs.is_none() {
self.audit_logs = Some(PostgresBackend::audit_log_read_repository(backend));
}
if self.auth_api_keys.is_none() {
self.auth_api_keys = Some(PostgresBackend::auth_api_key_read_repository(backend));
}
if self.auth_modules.is_none() {
self.auth_modules = Some(PostgresBackend::auth_module_read_repository(backend));
}
if self.background_tasks.is_none() {
self.background_tasks = Some(PostgresBackend::background_task_read_repository(backend));
}
if self.billing.is_none() {
self.billing = Some(PostgresBackend::billing_read_repository(backend));
}
if self.gemini_file_mappings.is_none() {
self.gemini_file_mappings = Some(PostgresBackend::gemini_file_mapping_read_repository(
backend,
));
}
if self.global_models.is_none() {
self.global_models = Some(PostgresBackend::global_model_read_repository(backend));
}
if self.management_tokens.is_none() {
self.management_tokens =
Some(PostgresBackend::management_token_read_repository(backend));
}
if self.oauth_providers.is_none() {
self.oauth_providers = Some(PostgresBackend::oauth_provider_read_repository(backend));
}
if self.pool_scores.is_none() {
self.pool_scores = Some(PostgresBackend::pool_score_read_repository(backend));
}
if self.proxy_nodes.is_none() {
self.proxy_nodes = Some(PostgresBackend::proxy_node_read_repository(backend));
}
if self.minimal_candidate_selection.is_none() {
self.minimal_candidate_selection =
Some(PostgresBackend::minimal_candidate_selection_read_repository(backend));
}
if self.request_candidates.is_none() {
self.request_candidates =
Some(PostgresBackend::request_candidate_read_repository(backend));
}
if self.provider_catalog.is_none() {
self.provider_catalog =
Some(PostgresBackend::provider_catalog_read_repository(backend));
}
if self.provider_quotas.is_none() {
self.provider_quotas = Some(PostgresBackend::provider_quota_read_repository(backend));
}
if self.routing_groups.is_none() {
self.routing_groups = Some(PostgresBackend::routing_group_read_repository(backend));
}
if self.usage.is_none() {
self.usage = Some(PostgresBackend::usage_read_repository(backend));
}
if self.users.is_none() {
self.users = Some(PostgresBackend::user_read_repository(backend));
}
if self.video_tasks.is_none() {
self.video_tasks = Some(PostgresBackend::video_task_read_repository(backend));
}
if self.wallets.is_none() {
self.wallets = Some(PostgresBackend::wallet_read_repository(backend));
}
}
#[cfg(feature = "mysql")]
fn install_mysql(&mut self, backend: &MysqlBackend) {
if self.announcements.is_none() {
self.announcements = Some(MysqlBackend::announcement_read_repository(backend));
}
if self.audit_logs.is_none() {
self.audit_logs = Some(MysqlBackend::audit_log_read_repository(backend));
}
if self.auth_api_keys.is_none() {
self.auth_api_keys = Some(MysqlBackend::auth_api_key_read_repository(backend));
}
if self.auth_modules.is_none() {
self.auth_modules = Some(MysqlBackend::auth_module_read_repository(backend));
}
if self.background_tasks.is_none() {
self.background_tasks = Some(MysqlBackend::background_task_read_repository(backend));
}
if self.billing.is_none() {
self.billing = Some(MysqlBackend::billing_read_repository(backend));
}
if self.gemini_file_mappings.is_none() {
self.gemini_file_mappings =
Some(MysqlBackend::gemini_file_mapping_read_repository(backend));
}
if self.global_models.is_none() {
self.global_models = Some(MysqlBackend::global_model_read_repository(backend));
}
if self.management_tokens.is_none() {
self.management_tokens = Some(MysqlBackend::management_token_read_repository(backend));
}
if self.oauth_providers.is_none() {
self.oauth_providers = Some(MysqlBackend::oauth_provider_read_repository(backend));
}
if self.pool_scores.is_none() {
self.pool_scores = Some(MysqlBackend::pool_score_read_repository(backend));
}
if self.proxy_nodes.is_none() {
self.proxy_nodes = Some(MysqlBackend::proxy_node_read_repository(backend));
}
if self.minimal_candidate_selection.is_none() {
self.minimal_candidate_selection = Some(
MysqlBackend::minimal_candidate_selection_read_repository(backend),
);
}
if self.request_candidates.is_none() {
self.request_candidates =
Some(MysqlBackend::request_candidate_read_repository(backend));
}
if self.provider_catalog.is_none() {
self.provider_catalog = Some(MysqlBackend::provider_catalog_read_repository(backend));
}
if self.provider_quotas.is_none() {
self.provider_quotas = Some(MysqlBackend::provider_quota_read_repository(backend));
}
if self.routing_groups.is_none() {
self.routing_groups = Some(MysqlBackend::routing_group_read_repository(backend));
}
if self.usage.is_none() {
self.usage = Some(MysqlBackend::usage_read_repository(backend));
}
if self.users.is_none() {
self.users = Some(MysqlBackend::user_read_repository(backend));
}
if self.video_tasks.is_none() {
self.video_tasks = Some(MysqlBackend::video_task_read_repository(backend));
}
if self.wallets.is_none() {
self.wallets = Some(MysqlBackend::wallet_read_repository(backend));
}
}
#[cfg(feature = "sqlite")]
fn install_sqlite(&mut self, backend: &SqliteBackend) {
if self.announcements.is_none() {
self.announcements = Some(SqliteBackend::announcement_read_repository(backend));
}
if self.audit_logs.is_none() {
self.audit_logs = Some(SqliteBackend::audit_log_read_repository(backend));
}
if self.auth_api_keys.is_none() {
self.auth_api_keys = Some(SqliteBackend::auth_api_key_read_repository(backend));
}
if self.auth_modules.is_none() {
self.auth_modules = Some(SqliteBackend::auth_module_read_repository(backend));
}
if self.background_tasks.is_none() {
self.background_tasks = Some(SqliteBackend::background_task_read_repository(backend));
}
if self.billing.is_none() {
self.billing = Some(SqliteBackend::billing_read_repository(backend));
}
if self.gemini_file_mappings.is_none() {
self.gemini_file_mappings =
Some(SqliteBackend::gemini_file_mapping_read_repository(backend));
}
if self.global_models.is_none() {
self.global_models = Some(SqliteBackend::global_model_read_repository(backend));
}
if self.management_tokens.is_none() {
self.management_tokens = Some(SqliteBackend::management_token_read_repository(backend));
}
if self.oauth_providers.is_none() {
self.oauth_providers = Some(SqliteBackend::oauth_provider_read_repository(backend));
}
if self.pool_scores.is_none() {
self.pool_scores = Some(SqliteBackend::pool_score_read_repository(backend));
}
if self.proxy_nodes.is_none() {
self.proxy_nodes = Some(SqliteBackend::proxy_node_read_repository(backend));
}
if self.minimal_candidate_selection.is_none() {
self.minimal_candidate_selection = Some(
SqliteBackend::minimal_candidate_selection_read_repository(backend),
);
}
if self.request_candidates.is_none() {
self.request_candidates =
Some(SqliteBackend::request_candidate_read_repository(backend));
}
if self.provider_catalog.is_none() {
self.provider_catalog = Some(SqliteBackend::provider_catalog_read_repository(backend));
}
if self.provider_quotas.is_none() {
self.provider_quotas = Some(SqliteBackend::provider_quota_read_repository(backend));
}
if self.routing_groups.is_none() {
self.routing_groups = Some(SqliteBackend::routing_group_read_repository(backend));
}
if self.usage.is_none() {
self.usage = Some(SqliteBackend::usage_read_repository(backend));
}
if self.users.is_none() {
self.users = Some(SqliteBackend::user_read_repository(backend));
}
if self.video_tasks.is_none() {
self.video_tasks = Some(SqliteBackend::video_task_read_repository(backend));
}
if self.wallets.is_none() {
self.wallets = Some(SqliteBackend::wallet_read_repository(backend));
}
}
#[cfg(test)]
#[cfg(feature = "postgres")]
pub(crate) fn from_postgres(postgres: Option<&PostgresBackend>) -> Self {
Self::from_backends(
postgres,
#[cfg(feature = "mysql")]
None,
#[cfg(feature = "sqlite")]
None,
)
}
pub fn auth_api_keys(&self) -> Option<Arc<dyn AuthApiKeyReadRepository>> {
self.auth_api_keys.clone()
}
pub fn announcements(&self) -> Option<Arc<dyn AnnouncementReadRepository>> {
self.announcements.clone()
}
pub fn audit_logs(&self) -> Option<Arc<dyn AuditLogReadRepository>> {
self.audit_logs.clone()
}
pub fn auth_modules(&self) -> Option<Arc<dyn AuthModuleReadRepository>> {
self.auth_modules.clone()
}
pub fn background_tasks(&self) -> Option<Arc<dyn BackgroundTaskReadRepository>> {
self.background_tasks.clone()
}
pub fn billing(&self) -> Option<Arc<dyn BillingReadRepository>> {
self.billing.clone()
}
pub fn gemini_file_mappings(&self) -> Option<Arc<dyn GeminiFileMappingReadRepository>> {
self.gemini_file_mappings.clone()
}
pub fn global_models(&self) -> Option<Arc<dyn GlobalModelReadRepository>> {
self.global_models.clone()
}
pub fn management_tokens(&self) -> Option<Arc<dyn ManagementTokenReadRepository>> {
self.management_tokens.clone()
}
pub fn oauth_providers(&self) -> Option<Arc<dyn OAuthProviderReadRepository>> {
self.oauth_providers.clone()
}
pub fn pool_scores(&self) -> Option<Arc<dyn PoolScoreReadRepository>> {
self.pool_scores.clone()
}
pub fn proxy_nodes(&self) -> Option<Arc<dyn ProxyNodeReadRepository>> {
self.proxy_nodes.clone()
}
pub fn minimal_candidate_selection(
&self,
) -> Option<Arc<dyn MinimalCandidateSelectionReadRepository>> {
self.minimal_candidate_selection.clone()
}
pub fn request_candidates(&self) -> Option<Arc<dyn RequestCandidateReadRepository>> {
self.request_candidates.clone()
}
pub fn provider_catalog(&self) -> Option<Arc<dyn ProviderCatalogReadRepository>> {
self.provider_catalog.clone()
}
pub fn provider_quotas(&self) -> Option<Arc<dyn ProviderQuotaReadRepository>> {
self.provider_quotas.clone()
}
pub fn routing_groups(&self) -> Option<Arc<dyn RoutingGroupReadRepository>> {
self.routing_groups.clone()
}
pub fn usage(&self) -> Option<Arc<dyn UsageReadRepository>> {
self.usage.clone()
}
pub fn users(&self) -> Option<Arc<dyn UserReadRepository>> {
self.users.clone()
}
pub fn video_tasks(&self) -> Option<Arc<dyn VideoTaskReadRepository>> {
self.video_tasks.clone()
}
pub fn wallets(&self) -> Option<Arc<dyn WalletReadRepository>> {
self.wallets.clone()
}
pub fn has_any(&self) -> bool {
self.auth_api_keys.is_some()
|| self.announcements.is_some()
|| self.audit_logs.is_some()
|| self.auth_modules.is_some()
|| self.background_tasks.is_some()
|| self.billing.is_some()
|| self.gemini_file_mappings.is_some()
|| self.global_models.is_some()
|| self.management_tokens.is_some()
|| self.oauth_providers.is_some()
|| self.pool_scores.is_some()
|| self.proxy_nodes.is_some()
|| self.minimal_candidate_selection.is_some()
|| self.request_candidates.is_some()
|| self.provider_catalog.is_some()
|| self.provider_quotas.is_some()
|| self.routing_groups.is_some()
|| self.usage.is_some()
|| self.users.is_some()
|| self.video_tasks.is_some()
|| self.wallets.is_some()
}
}
#[cfg(all(test, feature = "postgres"))]
mod tests {
use super::DataReadRepositories;
use crate::backend::PostgresBackend;
use crate::driver::postgres::PostgresPoolConfig;
#[tokio::test]
async fn builds_read_repositories_from_postgres_backend() {
let backend = PostgresBackend::from_config(PostgresPoolConfig {
database_url: "postgres://localhost/aether".to_string(),
min_connections: 1,
max_connections: 4,
acquire_timeout_ms: 1_000,
idle_timeout_ms: 5_000,
max_lifetime_ms: 30_000,
statement_cache_capacity: 64,
require_ssl: false,
})
.expect("postgres backend should build");
let read = DataReadRepositories::from_postgres(Some(&backend));
assert!(read.has_any());
assert!(read.announcements().is_some());
assert!(read.audit_logs().is_some());
assert!(read.auth_api_keys().is_some());
assert!(read.auth_modules().is_some());
assert!(read.billing().is_some());
assert!(read.gemini_file_mappings().is_some());
assert!(read.global_models().is_some());
assert!(read.management_tokens().is_some());
assert!(read.oauth_providers().is_some());
assert!(read.proxy_nodes().is_some());
assert!(read.minimal_candidate_selection().is_some());
assert!(read.request_candidates().is_some());
assert!(read.provider_catalog().is_some());
assert!(read.provider_quotas().is_some());
assert!(read.usage().is_some());
assert!(read.video_tasks().is_some());
assert!(read.wallets().is_some());
}
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,938 @@
use std::sync::Arc;
use crate::database::SqlDatabaseConfig;
use crate::driver::sqlite::{SqlitePool, SqlitePoolFactory};
use crate::repository::announcements::{
AnnouncementReadRepository, AnnouncementWriteRepository, SqliteAnnouncementRepository,
};
use crate::repository::audit::{AuditLogReadRepository, SqliteAuditLogReadRepository};
use crate::repository::auth::{
AuthApiKeyReadRepository, AuthApiKeyWriteRepository, SqliteAuthApiKeyReadRepository,
};
use crate::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, SqliteAuthModuleReadRepository,
SqliteAuthModuleRepository,
};
use crate::repository::background_tasks::{
BackgroundTaskReadRepository, BackgroundTaskWriteRepository, SqliteBackgroundTaskRepository,
};
use crate::repository::billing::{BillingReadRepository, SqliteBillingReadRepository};
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, SqliteMinimalCandidateSelectionReadRepository,
};
use crate::repository::candidates::{
RequestCandidateReadRepository, RequestCandidateWriteRepository,
SqliteRequestCandidateRepository,
};
use crate::repository::gemini_file_mappings::{
GeminiFileMappingReadRepository, GeminiFileMappingWriteRepository,
SqliteGeminiFileMappingRepository,
};
use crate::repository::global_models::{
GlobalModelReadRepository, GlobalModelWriteRepository, SqliteGlobalModelReadRepository,
};
use crate::repository::management_tokens::{
ManagementTokenReadRepository, ManagementTokenWriteRepository, SqliteManagementTokenRepository,
};
use crate::repository::oauth_providers::{
OAuthProviderReadRepository, OAuthProviderWriteRepository, SqliteOAuthProviderRepository,
};
use crate::repository::pool_scores::{
PoolMemberScoreWriteRepository, PoolScoreReadRepository, SqlitePoolMemberScoreRepository,
};
use crate::repository::provider_catalog::{
ProviderCatalogReadRepository, ProviderCatalogWriteRepository,
SqliteProviderCatalogReadRepository,
};
use crate::repository::proxy_nodes::{
ProxyNodeReadRepository, ProxyNodeWriteRepository, SqliteProxyNodeReadRepository,
};
use crate::repository::quota::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, SqliteProviderQuotaRepository,
};
use crate::repository::routing_profiles::{
RoutingGroupReadRepository, RoutingGroupWriteRepository, SqliteRoutingGroupRepository,
};
use crate::repository::settlement::{SettlementWriteRepository, SqliteSettlementRepository};
use crate::repository::usage::{
SqliteUsageReadRepository, SqliteUsageWriteRepository, UsageReadRepository,
UsageWriteRepository,
};
use crate::repository::users::{SqliteUserReadRepository, UserReadRepository};
use crate::repository::video_tasks::{
SqliteVideoTaskRepository, VideoTaskReadRepository, VideoTaskWriteRepository,
};
use crate::repository::wallet::{
SqliteWalletReadRepository, WalletReadRepository, WalletWriteRepository,
};
use crate::DataLayerError;
#[derive(Debug, Clone)]
pub struct SqliteBackend {
config: SqlDatabaseConfig,
pool: SqlitePool,
}
impl SqliteBackend {
pub fn from_config(config: SqlDatabaseConfig) -> Result<Self, DataLayerError> {
let factory = SqlitePoolFactory::new(config.clone())?;
let pool = factory.connect_lazy()?;
Ok(Self { config, pool })
}
pub fn config(&self) -> &SqlDatabaseConfig {
&self.config
}
pub fn pool(&self) -> &SqlitePool {
&self.pool
}
pub fn pool_clone(&self) -> SqlitePool {
self.pool.clone()
}
pub fn auth_api_key_read_repository(&self) -> Arc<dyn AuthApiKeyReadRepository> {
Arc::new(SqliteAuthApiKeyReadRepository::new(self.pool_clone()))
}
pub fn announcement_read_repository(&self) -> Arc<dyn AnnouncementReadRepository> {
Arc::new(SqliteAnnouncementRepository::new(self.pool_clone()))
}
pub fn audit_log_read_repository(&self) -> Arc<dyn AuditLogReadRepository> {
Arc::new(SqliteAuditLogReadRepository::new(self.pool_clone()))
}
pub fn announcement_write_repository(&self) -> Arc<dyn AnnouncementWriteRepository> {
Arc::new(SqliteAnnouncementRepository::new(self.pool_clone()))
}
pub fn auth_api_key_write_repository(&self) -> Arc<dyn AuthApiKeyWriteRepository> {
Arc::new(SqliteAuthApiKeyReadRepository::new(self.pool_clone()))
}
pub fn management_token_read_repository(&self) -> Arc<dyn ManagementTokenReadRepository> {
Arc::new(SqliteManagementTokenRepository::new(self.pool_clone()))
}
pub fn management_token_write_repository(&self) -> Arc<dyn ManagementTokenWriteRepository> {
Arc::new(SqliteManagementTokenRepository::new(self.pool_clone()))
}
pub fn auth_module_read_repository(&self) -> Arc<dyn AuthModuleReadRepository> {
Arc::new(SqliteAuthModuleReadRepository::new(self.pool_clone()))
}
pub fn auth_module_write_repository(&self) -> Arc<dyn AuthModuleWriteRepository> {
Arc::new(SqliteAuthModuleRepository::new(self.pool_clone()))
}
pub fn billing_read_repository(&self) -> Arc<dyn BillingReadRepository> {
Arc::new(SqliteBillingReadRepository::new(self.pool_clone()))
}
pub fn background_task_read_repository(&self) -> Arc<dyn BackgroundTaskReadRepository> {
Arc::new(SqliteBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn background_task_write_repository(&self) -> Arc<dyn BackgroundTaskWriteRepository> {
Arc::new(SqliteBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn request_candidate_read_repository(&self) -> Arc<dyn RequestCandidateReadRepository> {
Arc::new(SqliteRequestCandidateRepository::new(self.pool_clone()))
}
pub fn request_candidate_write_repository(&self) -> Arc<dyn RequestCandidateWriteRepository> {
Arc::new(SqliteRequestCandidateRepository::new(self.pool_clone()))
}
pub fn minimal_candidate_selection_read_repository(
&self,
) -> Arc<dyn MinimalCandidateSelectionReadRepository> {
Arc::new(SqliteMinimalCandidateSelectionReadRepository::new(
self.pool_clone(),
))
}
pub fn gemini_file_mapping_read_repository(&self) -> Arc<dyn GeminiFileMappingReadRepository> {
Arc::new(SqliteGeminiFileMappingRepository::new(self.pool_clone()))
}
pub fn gemini_file_mapping_write_repository(
&self,
) -> Arc<dyn GeminiFileMappingWriteRepository> {
Arc::new(SqliteGeminiFileMappingRepository::new(self.pool_clone()))
}
pub fn global_model_read_repository(&self) -> Arc<dyn GlobalModelReadRepository> {
Arc::new(SqliteGlobalModelReadRepository::new(self.pool_clone()))
}
pub fn global_model_write_repository(&self) -> Arc<dyn GlobalModelWriteRepository> {
Arc::new(SqliteGlobalModelReadRepository::new(self.pool_clone()))
}
pub fn user_read_repository(&self) -> Arc<dyn UserReadRepository> {
Arc::new(SqliteUserReadRepository::new(self.pool_clone()))
}
pub fn video_task_read_repository(&self) -> Arc<dyn VideoTaskReadRepository> {
Arc::new(SqliteVideoTaskRepository::new(self.pool_clone()))
}
pub fn video_task_write_repository(&self) -> Arc<dyn VideoTaskWriteRepository> {
Arc::new(SqliteVideoTaskRepository::new(self.pool_clone()))
}
pub fn oauth_provider_read_repository(&self) -> Arc<dyn OAuthProviderReadRepository> {
Arc::new(SqliteOAuthProviderRepository::new(self.pool_clone()))
}
pub fn oauth_provider_write_repository(&self) -> Arc<dyn OAuthProviderWriteRepository> {
Arc::new(SqliteOAuthProviderRepository::new(self.pool_clone()))
}
pub fn provider_catalog_read_repository(&self) -> Arc<dyn ProviderCatalogReadRepository> {
Arc::new(SqliteProviderCatalogReadRepository::new(self.pool_clone()))
}
pub fn provider_catalog_write_repository(&self) -> Arc<dyn ProviderCatalogWriteRepository> {
Arc::new(SqliteProviderCatalogReadRepository::new(self.pool_clone()))
}
pub fn pool_score_read_repository(&self) -> Arc<dyn PoolScoreReadRepository> {
Arc::new(SqlitePoolMemberScoreRepository::new(self.pool_clone()))
}
pub fn pool_score_write_repository(&self) -> Arc<dyn PoolMemberScoreWriteRepository> {
Arc::new(SqlitePoolMemberScoreRepository::new(self.pool_clone()))
}
pub fn routing_group_read_repository(&self) -> Arc<dyn RoutingGroupReadRepository> {
Arc::new(SqliteRoutingGroupRepository::new(self.pool_clone()))
}
pub fn routing_group_write_repository(&self) -> Arc<dyn RoutingGroupWriteRepository> {
Arc::new(SqliteRoutingGroupRepository::new(self.pool_clone()))
}
pub fn proxy_node_read_repository(&self) -> Arc<dyn ProxyNodeReadRepository> {
Arc::new(SqliteProxyNodeReadRepository::new(self.pool_clone()))
}
pub fn proxy_node_write_repository(&self) -> Arc<dyn ProxyNodeWriteRepository> {
Arc::new(SqliteProxyNodeReadRepository::new(self.pool_clone()))
}
pub fn provider_quota_read_repository(&self) -> Arc<dyn ProviderQuotaReadRepository> {
Arc::new(SqliteProviderQuotaRepository::new(self.pool_clone()))
}
pub fn provider_quota_write_repository(&self) -> Arc<dyn ProviderQuotaWriteRepository> {
Arc::new(SqliteProviderQuotaRepository::new(self.pool_clone()))
}
pub fn settlement_write_repository(&self) -> Arc<dyn SettlementWriteRepository> {
Arc::new(SqliteSettlementRepository::new(self.pool_clone()))
}
pub fn usage_write_repository(&self) -> Arc<dyn UsageWriteRepository> {
Arc::new(SqliteUsageWriteRepository::new(self.pool_clone()))
}
pub fn usage_read_repository(&self) -> Arc<dyn UsageReadRepository> {
Arc::new(SqliteUsageReadRepository::new(self.pool_clone()))
}
pub fn wallet_read_repository(&self) -> Arc<dyn WalletReadRepository> {
Arc::new(SqliteWalletReadRepository::new(self.pool_clone()))
}
pub fn wallet_write_repository(&self) -> Arc<dyn WalletWriteRepository> {
Arc::new(SqliteWalletReadRepository::new(self.pool_clone()))
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use super::SqliteBackend;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::system::{
AdminSystemPurgeTarget, AdminSystemStatsDailyAggregate,
AdminSystemStatsDailyApiKeyAggregate, AdminSystemStatsUserDailyAggregate,
AdminSystemUsageAggregateImportMode, AdminSystemUsageAggregateSnapshot,
};
use crate::{
DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, StatsDailyAggregationInput,
StatsHourlyAggregationInput, WalletDailyUsageAggregationInput,
};
#[tokio::test]
async fn backend_retains_config_and_pool() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite://./data/aether.db".to_string(),
pool: SqlPoolConfig::default(),
};
let backend = SqliteBackend::from_config(config.clone()).expect("backend should build");
assert_eq!(backend.config(), &config);
let _pool = backend.pool();
let _pool_clone = backend.pool_clone();
}
#[tokio::test]
async fn system_config_round_trips_after_sqlite_migrations() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite::memory:".to_string(),
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = SqliteBackend::from_config(config).expect("backend should build");
run_sqlite_migrations(backend.pool())
.await
.expect("sqlite migrations should run");
let value = serde_json::json!({"enabled": true});
let stored = backend
.upsert_system_config_entry("feature.local", &value, Some("local flag"))
.await
.expect("system config should upsert");
assert_eq!(stored.value, value);
assert_eq!(
backend
.find_system_config_value("feature.local")
.await
.expect("system config should read"),
Some(value)
);
assert_eq!(
backend
.list_system_config_entries()
.await
.expect("system config should list")
.len(),
2
);
assert!(backend
.delete_system_config_value("feature.local")
.await
.expect("system config should delete"));
}
#[tokio::test]
async fn table_maintenance_runs_after_sqlite_migrations() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite::memory:".to_string(),
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = SqliteBackend::from_config(config).expect("backend should build");
run_sqlite_migrations(backend.pool())
.await
.expect("sqlite migrations should run");
let summary = backend
.run_table_maintenance(&["usage", "request_candidates", "audit_logs"])
.await
.expect("sqlite table maintenance should run");
assert_eq!(summary.attempted, 3);
assert_eq!(summary.succeeded, 3);
}
#[tokio::test]
async fn admin_system_config_purge_deletes_config_scope_and_preserves_users() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite::memory:".to_string(),
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = SqliteBackend::from_config(config).expect("backend should build");
run_sqlite_migrations(backend.pool())
.await
.expect("sqlite migrations should run");
sqlx::query(
"INSERT INTO users (id, email, username, role, created_at, updated_at) VALUES ('admin-1', '[email protected]', 'admin', 'admin', 1, 1)",
)
.execute(backend.pool())
.await
.expect("user should insert");
sqlx::query(
"INSERT INTO providers (id, name, provider_type, created_at, updated_at) VALUES ('provider-1', 'OpenAI', 'openai', 1, 1)",
)
.execute(backend.pool())
.await
.expect("provider should insert");
sqlx::query(
"INSERT INTO system_configs (id, key, value, created_at, updated_at) VALUES ('config-1', 'site_name', '\"Aether\"', 1, 1)",
)
.execute(backend.pool())
.await
.expect("system config should insert");
let summary = backend
.purge_admin_system_data(AdminSystemPurgeTarget::Config)
.await
.expect("config purge should run");
assert!(summary.total() >= 2);
assert_eq!(sqlite_count(backend.pool(), "system_configs").await, 0);
assert_eq!(sqlite_count(backend.pool(), "providers").await, 0);
assert_eq!(sqlite_count(backend.pool(), "users").await, 1);
}
#[tokio::test]
async fn admin_system_users_purge_deletes_only_non_admin_users_and_keys() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite::memory:".to_string(),
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = SqliteBackend::from_config(config).expect("backend should build");
run_sqlite_migrations(backend.pool())
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO users (id, email, username, role, created_at, updated_at)
VALUES
('admin-1', '[email protected]', 'admin', 'admin', 1, 1),
('user-1', '[email protected]', 'alice', 'user', 1, 1)
"#,
)
.execute(backend.pool())
.await
.expect("users should insert");
sqlx::query(
r#"
INSERT INTO api_keys (id, user_id, key_hash, name, created_at, updated_at, total_requests, total_tokens, total_cost_usd)
VALUES
('admin-key-1', 'admin-1', 'hash-admin', 'admin-key', 1, 1, 5, 50, 0.5),
('user-key-1', 'user-1', 'hash-user', 'user-key', 1, 1, 7, 70, 0.7)
"#,
)
.execute(backend.pool())
.await
.expect("api keys should insert");
sqlx::query(
r#"
INSERT INTO stats_daily_api_key (id, api_key_id, "date", total_requests, created_at, updated_at)
VALUES
('admin-key-stats-1', 'admin-key-1', 1, 5, 1, 1),
('user-key-stats-1', 'user-key-1', 1, 7, 1, 1)
"#,
)
.execute(backend.pool())
.await
.expect("api key stats should insert");
let summary = backend
.purge_admin_system_data(AdminSystemPurgeTarget::Users)
.await
.expect("users purge should run");
assert!(summary.total() >= 2);
assert_eq!(sqlite_count(backend.pool(), "users").await, 1);
assert_eq!(sqlite_count(backend.pool(), "api_keys").await, 1);
assert_eq!(sqlite_count(backend.pool(), "stats_daily_api_key").await, 1);
let admin_exists: i64 =
sqlx::query_scalar("SELECT COUNT(*) FROM users WHERE id = 'admin-1'")
.fetch_one(backend.pool())
.await
.expect("admin count should load");
assert_eq!(admin_exists, 1);
}
#[tokio::test]
async fn admin_system_usage_aggregates_round_trip_after_sqlite_migrations() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite::memory:".to_string(),
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = SqliteBackend::from_config(config).expect("backend should build");
run_sqlite_migrations(backend.pool())
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO users (id, email, username, role, created_at, updated_at)
VALUES ('target-user-1', '[email protected]', 'target', 'user', 1, 1)
"#,
)
.execute(backend.pool())
.await
.expect("target user should insert");
sqlx::query(
r#"
INSERT INTO api_keys (id, user_id, key_hash, name, created_at, updated_at)
VALUES ('target-key-1', 'target-user-1', 'hash-target-key', 'target key', 1, 1)
"#,
)
.execute(backend.pool())
.await
.expect("target key should insert");
let snapshot = AdminSystemUsageAggregateSnapshot {
stats_daily: vec![AdminSystemStatsDailyAggregate {
date_unix_secs: 86_400,
total_requests: 9,
success_requests: 8,
error_requests: 1,
input_tokens: 100,
output_tokens: 200,
cache_creation_tokens: 3,
cache_read_tokens: 4,
total_cost: 1.25,
actual_total_cost: 1.0,
is_complete: true,
aggregated_at_unix_secs: Some(90_000),
}],
stats_user_daily: vec![AdminSystemStatsUserDailyAggregate {
user_id: "source-user-1".to_string(),
username: Some("source".to_string()),
date_unix_secs: 86_400,
total_requests: 5,
success_requests: 5,
error_requests: 0,
input_tokens: 50,
output_tokens: 60,
cache_creation_tokens: 1,
cache_read_tokens: 2,
total_cost: 0.5,
}],
stats_daily_api_key: vec![AdminSystemStatsDailyApiKeyAggregate {
api_key_id: "source-key-1".to_string(),
api_key_name: Some("source key".to_string()),
date_unix_secs: 86_400,
total_requests: 4,
success_requests: 3,
error_requests: 1,
input_tokens: 40,
output_tokens: 30,
cache_creation_tokens: 2,
cache_read_tokens: 1,
total_cost: 0.75,
}],
};
let user_id_map =
BTreeMap::from([("source-user-1".to_string(), "target-user-1".to_string())]);
let api_key_id_map =
BTreeMap::from([("source-key-1".to_string(), "target-key-1".to_string())]);
let summary = backend
.import_admin_system_usage_aggregates(
&snapshot,
&user_id_map,
&api_key_id_map,
AdminSystemUsageAggregateImportMode::Overwrite,
)
.await
.expect("usage aggregates should import");
assert_eq!(summary.stats_daily.created, 1);
assert_eq!(summary.stats_user_daily.created, 1);
assert_eq!(summary.stats_daily_api_key.created, 1);
let exported = backend
.export_admin_system_usage_aggregates()
.await
.expect("usage aggregates should export");
assert_eq!(exported.stats_daily.len(), 1);
assert_eq!(exported.stats_daily[0].total_requests, 9);
assert_eq!(exported.stats_daily[0].actual_total_cost, 1.0);
assert_eq!(exported.stats_user_daily.len(), 1);
assert_eq!(exported.stats_user_daily[0].user_id, "target-user-1");
assert_eq!(exported.stats_user_daily[0].total_requests, 5);
assert_eq!(exported.stats_daily_api_key.len(), 1);
assert_eq!(exported.stats_daily_api_key[0].api_key_id, "target-key-1");
assert_eq!(exported.stats_daily_api_key[0].total_requests, 4);
}
#[tokio::test]
async fn admin_system_request_bodies_purge_clears_inline_usage_body_fields() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite::memory:".to_string(),
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = SqliteBackend::from_config(config).expect("backend should build");
run_sqlite_migrations(backend.pool())
.await
.expect("sqlite migrations should run");
for (column, ty) in [
("request_body", "TEXT"),
("response_body", "TEXT"),
("provider_request_body", "TEXT"),
("client_response_body", "TEXT"),
("request_body_compressed", "BLOB"),
("response_body_compressed", "BLOB"),
("provider_request_body_compressed", "BLOB"),
("client_response_body_compressed", "BLOB"),
] {
sqlx::query(&format!(r#"ALTER TABLE "usage" ADD COLUMN {column} {ty}"#))
.execute(backend.pool())
.await
.expect("legacy body column should be added");
}
sqlx::query(
r#"
INSERT INTO "usage" (
request_id,
provider_name,
model,
request_body,
response_body,
provider_request_body,
client_response_body,
request_body_compressed,
response_body_compressed,
provider_request_body_compressed,
client_response_body_compressed,
created_at_unix_ms
)
VALUES (
'request-1',
'openai',
'gpt-4.1',
'client request',
'provider response',
'provider request',
'client response',
X'01',
X'02',
X'03',
X'04',
1
)
"#,
)
.execute(backend.pool())
.await
.expect("usage row should insert");
let summary = backend
.purge_admin_system_data(AdminSystemPurgeTarget::RequestBodies)
.await
.expect("request body purge should run");
assert_eq!(summary.affected.get("usage_body_fields_cleaned"), Some(&1));
let remaining: i64 = sqlx::query_scalar(
r#"
SELECT COUNT(*)
FROM "usage"
WHERE request_body IS NOT NULL
OR response_body IS NOT NULL
OR provider_request_body IS NOT NULL
OR client_response_body IS NOT NULL
OR request_body_compressed IS NOT NULL
OR response_body_compressed IS NOT NULL
OR provider_request_body_compressed IS NOT NULL
OR client_response_body_compressed IS NOT NULL
"#,
)
.fetch_one(backend.pool())
.await
.expect("remaining body count should load");
assert_eq!(remaining, 0);
}
async fn sqlite_count(pool: &sqlx::SqlitePool, table: &str) -> i64 {
let sql = format!("SELECT COUNT(*) FROM \"{table}\"");
sqlx::query_scalar::<_, i64>(&sql)
.fetch_one(pool)
.await
.expect("count should load")
}
#[tokio::test]
async fn wallet_daily_usage_aggregation_uses_settlement_wallets_after_sqlite_migrations() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite::memory:".to_string(),
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = SqliteBackend::from_config(config).expect("backend should build");
run_sqlite_migrations(backend.pool())
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO wallets (id, user_id, balance, gift_balance, limit_mode, created_at, updated_at)
VALUES
('wallet-1', 'user-1', 10.0, 2.0, 'finite', 1, 1),
('wallet-stale', 'user-stale', 0.0, 0.0, 'finite', 1, 1)
"#,
)
.execute(backend.pool())
.await
.expect("wallets should seed");
sqlx::query(
r#"
INSERT INTO "usage" (
request_id, wallet_id, provider_name, model, status, billing_status,
total_cost_usd, input_tokens, output_tokens, cache_creation_input_tokens,
cache_read_input_tokens, finalized_at, created_at_unix_ms, updated_at_unix_secs
) VALUES
('request-1', 'wrong-wallet', 'provider', 'model', 'completed', 'pending',
1.25, 10, 20, 3, 4, 900, 900000, 900),
('request-2', NULL, 'provider', 'model', 'completed', 'pending',
2.00, 5, 7, 1, 2, 901, 901000, 901),
('request-zero', NULL, 'provider', 'model', 'completed', 'pending',
0.00, 100, 100, 0, 0, 902, 902000, 902),
('request-outside', NULL, 'provider', 'model', 'completed', 'pending',
9.00, 50, 50, 0, 0, 903, 903000, 903)
"#,
)
.execute(backend.pool())
.await
.expect("usage should seed");
sqlx::query(
r#"
INSERT INTO usage_settlement_snapshots (
request_id, billing_status, wallet_id, finalized_at, created_at, updated_at
) VALUES
('request-1', 'settled', 'wallet-1', 1000, 1000, 1000),
('request-2', 'settled', 'wallet-1', 1100, 1100, 1100),
('request-zero', 'settled', 'wallet-1', 1150, 1150, 1150),
('request-outside', 'settled', 'wallet-1', 1200, 1200, 1200)
"#,
)
.execute(backend.pool())
.await
.expect("settlement snapshots 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 (
'stale-ledger', 'wallet-stale', '2026-05-03', 'Asia/Shanghai',
7.0, 3, 1, 1, 0, 0, 999, 999, 999
)
"#,
)
.execute(backend.pool())
.await
.expect("stale ledger should seed");
let summary = backend
.aggregate_wallet_daily_usage(&WalletDailyUsageAggregationInput {
billing_date: "2026-05-03".to_string(),
billing_timezone: "Asia/Shanghai".to_string(),
window_start_unix_secs: 1000,
window_end_unix_secs: 1200,
aggregated_at_unix_secs: 1300,
})
.await
.expect("wallet daily usage aggregation should run");
assert_eq!(summary.aggregated_wallets, 1);
assert_eq!(summary.deleted_stale_ledgers, 1);
let ledger = sqlx::query_as::<
_,
(
String,
String,
f64,
i64,
i64,
i64,
i64,
i64,
Option<i64>,
Option<i64>,
i64,
),
>(
r#"
SELECT
id,
wallet_id,
total_cost_usd,
total_requests,
input_tokens,
output_tokens,
cache_creation_tokens,
cache_read_tokens,
first_finalized_at,
last_finalized_at,
aggregated_at
FROM wallet_daily_usage_ledgers
WHERE billing_date = '2026-05-03'
AND billing_timezone = 'Asia/Shanghai'
"#,
)
.fetch_one(backend.pool())
.await
.expect("aggregated ledger should load");
assert_eq!(ledger.0.len(), 64);
assert_eq!(ledger.1, "wallet-1");
assert!((ledger.2 - 3.25).abs() < f64::EPSILON);
assert_eq!(ledger.3, 2);
assert_eq!(ledger.4, 15);
assert_eq!(ledger.5, 27);
assert_eq!(ledger.6, 4);
assert_eq!(ledger.7, 6);
assert_eq!(ledger.8, Some(1000));
assert_eq!(ledger.9, Some(1100));
assert_eq!(ledger.10, 1300);
let stale_count: i64 = sqlx::query_scalar(
"SELECT COUNT(*) FROM wallet_daily_usage_ledgers WHERE id = 'stale-ledger'",
)
.fetch_one(backend.pool())
.await
.expect("stale ledger count should load");
assert_eq!(stale_count, 0);
}
#[tokio::test]
async fn stats_aggregation_runs_after_sqlite_migrations() {
let config = SqlDatabaseConfig {
driver: DatabaseDriver::Sqlite,
url: "sqlite::memory:".to_string(),
pool: SqlPoolConfig {
max_connections: 1,
..SqlPoolConfig::default()
},
};
let backend = SqliteBackend::from_config(config).expect("backend should build");
run_sqlite_migrations(backend.pool())
.await
.expect("sqlite migrations should run");
sqlx::query(
r#"
INSERT INTO "usage" (
request_id, user_id, api_key_id, provider_name, model, status, billing_status,
status_code, error_category, input_tokens, output_tokens,
cache_creation_input_tokens, cache_read_input_tokens, total_cost_usd,
actual_total_cost_usd, response_time_ms, created_at_unix_ms, updated_at_unix_secs
) VALUES
('stats-1', 'user-1', 'key-1', 'provider-a', 'model-a', 'completed', 'settled',
200, NULL, 10, 20, 1, 2, 0.30, 0.25, 100, 3600000, 3600),
('stats-2', 'user-2', 'key-2', 'provider-b', 'model-b', 'failed', 'void',
500, 'upstream_error', 5, 7, 0, 1, 0.20, 0.20, 300, 3610000, 3610),
('stats-pending', 'user-3', 'key-3', 'provider-a', 'model-a', 'pending', 'pending',
NULL, NULL, 100, 100, 0, 0, 9.99, 9.99, 50, 3620000, 3620),
('stats-unknown-provider', 'user-4', 'key-4', 'unknown', 'model-a', 'completed', 'settled',
200, NULL, 100, 100, 0, 0, 9.99, 9.99, 50, 3630000, 3630)
"#,
)
.execute(backend.pool())
.await
.expect("usage stats rows should seed");
let target_hour = chrono::DateTime::<chrono::Utc>::from_timestamp(3600, 0)
.expect("target hour should be valid");
let aggregated_at = chrono::DateTime::<chrono::Utc>::from_timestamp(7200, 0)
.expect("aggregation time should be valid");
let hourly = backend
.aggregate_stats_hourly(&StatsHourlyAggregationInput {
target_hour_utc: target_hour,
aggregated_at,
})
.await
.expect("hourly stats aggregation should run")
.expect("hourly bucket should aggregate");
assert_eq!(hourly.hour_utc, target_hour);
assert_eq!(hourly.total_requests, 2);
assert_eq!(hourly.user_rows, 2);
assert_eq!(hourly.user_model_rows, 2);
assert_eq!(hourly.model_rows, 2);
assert_eq!(hourly.provider_rows, 2);
let hourly_row = sqlx::query_as::<_, (i64, i64, i64, i64, f64)>(
r#"
SELECT total_requests, success_requests, error_requests, input_tokens, total_cost
FROM stats_hourly
WHERE hour_utc = 3600
"#,
)
.fetch_one(backend.pool())
.await
.expect("hourly stats row should load");
assert_eq!(hourly_row.0, 2);
assert_eq!(hourly_row.1, 1);
assert_eq!(hourly_row.2, 1);
assert_eq!(hourly_row.3, 15);
assert!((hourly_row.4 - 0.50).abs() < f64::EPSILON);
let second_hourly = backend
.aggregate_stats_hourly(&StatsHourlyAggregationInput {
target_hour_utc: target_hour,
aggregated_at,
})
.await
.expect("second hourly aggregation should run");
assert!(second_hourly.is_none());
let target_day = chrono::DateTime::<chrono::Utc>::from_timestamp(0, 0)
.expect("target day should be valid");
let daily = backend
.aggregate_stats_daily(&StatsDailyAggregationInput {
target_day_utc: target_day,
aggregated_at,
})
.await
.expect("daily stats aggregation should run")
.expect("daily bucket should aggregate");
assert_eq!(daily.day_start_utc, target_day);
assert_eq!(daily.total_requests, 2);
assert_eq!(daily.model_rows, 2);
assert_eq!(daily.provider_rows, 2);
assert_eq!(daily.api_key_rows, 2);
assert_eq!(daily.error_rows, 1);
assert_eq!(daily.user_rows, 2);
let daily_row = sqlx::query_as::<_, (i64, i64, i64, i64)>(
r#"
SELECT total_requests, success_requests, error_requests, unique_models
FROM stats_daily
WHERE "date" = 0
"#,
)
.fetch_one(backend.pool())
.await
.expect("daily stats row should load");
assert_eq!(daily_row, (2, 1, 1, 2));
}
}
@@ -0,0 +1,8 @@
#[cfg(feature = "mysql")]
pub(crate) mod mysql;
#[cfg(feature = "postgres")]
pub(crate) mod postgres_daily;
#[cfg(feature = "postgres")]
pub(crate) mod postgres_hourly;
#[cfg(feature = "sqlite")]
pub(crate) mod sqlite;
@@ -0,0 +1,371 @@
use chrono::{DateTime, Utc};
use sqlx::Row;
use crate::backend::stats_common::{stats_id, unix_ms, unix_secs, utc_from_unix_secs};
use crate::backend::MysqlBackend;
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::{
DataLayerError, StatsDailyAggregationInput, StatsDailyAggregationSummary,
StatsHourlyAggregationInput, StatsHourlyAggregationSummary,
};
impl MysqlBackend {
pub async fn aggregate_stats_hourly(
&self,
input: &StatsHourlyAggregationInput,
) -> Result<Option<StatsHourlyAggregationSummary>, DataLayerError> {
let Some(hour_utc_unix_secs) =
next_mysql_stats_hourly_bucket(self.pool(), input.target_hour_utc).await?
else {
return Ok(None);
};
perform_mysql_stats_hourly_aggregation(self.pool(), hour_utc_unix_secs, input.aggregated_at)
.await
.map(Some)
}
pub async fn aggregate_stats_daily(
&self,
input: &StatsDailyAggregationInput,
) -> Result<Option<StatsDailyAggregationSummary>, DataLayerError> {
let Some(day_start_unix_secs) =
next_mysql_stats_daily_bucket(self.pool(), input.target_day_utc).await?
else {
return Ok(None);
};
perform_mysql_stats_daily_aggregation(self.pool(), day_start_unix_secs, input.aggregated_at)
.await
.map(Some)
}
}
async fn next_mysql_stats_hourly_bucket(
pool: &MysqlPool,
target_hour_utc: DateTime<Utc>,
) -> Result<Option<i64>, DataLayerError> {
let latest_hour: Option<i64> =
sqlx::query_scalar("SELECT MAX(hour_utc) FROM stats_hourly WHERE is_complete <> 0")
.fetch_one(pool)
.await
.map_sql_err()?;
let search_from = latest_hour.map(|value| value + 3600).unwrap_or(0);
let search_until = unix_secs(target_hour_utc) + 3600;
if search_from >= search_until {
return Ok(None);
}
let next_bucket: Option<i64> = sqlx::query_scalar(
r#"
SELECT CAST(MIN(FLOOR(created_at_unix_ms / 3600000) * 3600) AS SIGNED)
FROM `usage`
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
"#,
)
.bind(unix_ms(search_from)?)
.bind(unix_ms(search_until)?)
.fetch_one(pool)
.await
.map_sql_err()?;
Ok(next_bucket.filter(|value| *value <= unix_secs(target_hour_utc)))
}
async fn next_mysql_stats_daily_bucket(
pool: &MysqlPool,
target_day_utc: DateTime<Utc>,
) -> Result<Option<i64>, DataLayerError> {
let latest_day: Option<i64> =
sqlx::query_scalar("SELECT MAX(`date`) FROM stats_daily WHERE is_complete <> 0")
.fetch_one(pool)
.await
.map_sql_err()?;
let search_from = latest_day.map(|value| value + 86_400).unwrap_or(0);
let search_until = unix_secs(target_day_utc) + 86_400;
if search_from >= search_until {
return Ok(None);
}
let next_bucket: Option<i64> = sqlx::query_scalar(
r#"
SELECT CAST(MIN(FLOOR(created_at_unix_ms / 86400000) * 86400) AS SIGNED)
FROM `usage`
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
"#,
)
.bind(unix_ms(search_from)?)
.bind(unix_ms(search_until)?)
.fetch_one(pool)
.await
.map_sql_err()?;
Ok(next_bucket.filter(|value| *value <= unix_secs(target_day_utc)))
}
const MYSQL_STATS_AGGREGATE_SQL: &str = r#"
SELECT
CAST(COUNT(*) AS SIGNED) AS total_requests,
CAST(COALESCE(SUM(CASE
WHEN status = 'failed'
OR status_code >= 400
OR (error_category IS NOT NULL AND error_category <> '')
THEN 1 ELSE 0 END), 0) AS SIGNED) AS error_requests,
CAST(COALESCE(SUM(input_tokens), 0) AS SIGNED) AS input_tokens,
CAST(COALESCE(SUM(output_tokens), 0) AS SIGNED) AS output_tokens,
CAST(COALESCE(SUM(cache_creation_input_tokens), 0) AS SIGNED) AS cache_creation_tokens,
CAST(COALESCE(SUM(cache_read_input_tokens), 0) AS SIGNED) AS cache_read_tokens,
CAST(COALESCE(SUM(total_cost_usd), 0.0) AS DOUBLE) AS total_cost,
CAST(COALESCE(SUM(actual_total_cost_usd), 0.0) AS DOUBLE) AS actual_total_cost,
CAST(COALESCE(AVG(response_time_ms), 0.0) AS DOUBLE) AS avg_response_time_ms
FROM `usage`
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
"#;
async fn perform_mysql_stats_hourly_aggregation(
pool: &MysqlPool,
hour_utc_unix_secs: i64,
aggregated_at: DateTime<Utc>,
) -> Result<StatsHourlyAggregationSummary, DataLayerError> {
let start_ms = unix_ms(hour_utc_unix_secs)?;
let end_ms = unix_ms(hour_utc_unix_secs + 3600)?;
let aggregated_at_unix_secs = unix_secs(aggregated_at);
let mut tx = pool.begin().await.map_sql_err()?;
let row = sqlx::query(MYSQL_STATS_AGGREGATE_SQL)
.bind(start_ms)
.bind(end_ms)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let total_requests: i64 = row.try_get("total_requests").map_sql_err()?;
let error_requests: i64 = row.try_get("error_requests").map_sql_err()?;
sqlx::query(
r#"
INSERT INTO stats_hourly (
id, hour_utc, total_requests, success_requests, error_requests,
input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens,
total_cost, actual_total_cost, avg_response_time_ms, is_complete,
aggregated_at, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, TRUE, ?, ?, ?)
ON DUPLICATE KEY UPDATE
total_requests = VALUES(total_requests),
success_requests = VALUES(success_requests),
error_requests = VALUES(error_requests),
input_tokens = VALUES(input_tokens),
output_tokens = VALUES(output_tokens),
cache_creation_tokens = VALUES(cache_creation_tokens),
cache_read_tokens = VALUES(cache_read_tokens),
total_cost = VALUES(total_cost),
actual_total_cost = VALUES(actual_total_cost),
avg_response_time_ms = VALUES(avg_response_time_ms),
is_complete = VALUES(is_complete),
aggregated_at = VALUES(aggregated_at),
updated_at = VALUES(updated_at)
"#,
)
.bind(stats_id(&format!("stats-hourly:{hour_utc_unix_secs}")))
.bind(hour_utc_unix_secs)
.bind(total_requests)
.bind(total_requests.saturating_sub(error_requests))
.bind(error_requests)
.bind(row.try_get::<i64, _>("input_tokens").map_sql_err()?)
.bind(row.try_get::<i64, _>("output_tokens").map_sql_err()?)
.bind(
row.try_get::<i64, _>("cache_creation_tokens")
.map_sql_err()?,
)
.bind(row.try_get::<i64, _>("cache_read_tokens").map_sql_err()?)
.bind(row.try_get::<f64, _>("total_cost").map_sql_err()?)
.bind(row.try_get::<f64, _>("actual_total_cost").map_sql_err()?)
.bind(
row.try_get::<f64, _>("avg_response_time_ms")
.map_sql_err()?,
)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
.execute(&mut *tx)
.await
.map_sql_err()?;
let user_rows = mysql_group_count(&mut tx, "user_id", start_ms, end_ms).await?;
let user_model_rows = mysql_group_count(&mut tx, "user_id, model", start_ms, end_ms).await?;
let model_rows = mysql_group_count(&mut tx, "model", start_ms, end_ms).await?;
let provider_rows = mysql_group_count(&mut tx, "provider_name", start_ms, end_ms).await?;
tx.commit().await.map_sql_err()?;
Ok(StatsHourlyAggregationSummary {
hour_utc: utc_from_unix_secs(hour_utc_unix_secs, "stats_hourly.hour_utc")?,
total_requests,
user_rows,
user_model_rows,
model_rows,
provider_rows,
})
}
async fn perform_mysql_stats_daily_aggregation(
pool: &MysqlPool,
day_start_unix_secs: i64,
aggregated_at: DateTime<Utc>,
) -> Result<StatsDailyAggregationSummary, DataLayerError> {
let start_ms = unix_ms(day_start_unix_secs)?;
let end_ms = unix_ms(day_start_unix_secs + 86_400)?;
let aggregated_at_unix_secs = unix_secs(aggregated_at);
let mut tx = pool.begin().await.map_sql_err()?;
let row = sqlx::query(MYSQL_STATS_AGGREGATE_SQL)
.bind(start_ms)
.bind(end_ms)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let total_requests: i64 = row.try_get("total_requests").map_sql_err()?;
let error_requests: i64 = row.try_get("error_requests").map_sql_err()?;
let unique_models = mysql_group_count(&mut tx, "model", start_ms, end_ms).await? as i64;
let unique_providers =
mysql_group_count(&mut tx, "provider_name", start_ms, end_ms).await? as i64;
sqlx::query(
r#"
INSERT INTO stats_daily (
id, `date`, total_requests, success_requests, error_requests,
input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens,
total_cost, actual_total_cost, avg_response_time_ms, fallback_count,
unique_models, unique_providers, is_complete, aggregated_at, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 0, ?, ?, TRUE, ?, ?, ?)
ON DUPLICATE KEY UPDATE
total_requests = VALUES(total_requests),
success_requests = VALUES(success_requests),
error_requests = VALUES(error_requests),
input_tokens = VALUES(input_tokens),
output_tokens = VALUES(output_tokens),
cache_creation_tokens = VALUES(cache_creation_tokens),
cache_read_tokens = VALUES(cache_read_tokens),
total_cost = VALUES(total_cost),
actual_total_cost = VALUES(actual_total_cost),
avg_response_time_ms = VALUES(avg_response_time_ms),
fallback_count = VALUES(fallback_count),
unique_models = VALUES(unique_models),
unique_providers = VALUES(unique_providers),
is_complete = VALUES(is_complete),
aggregated_at = VALUES(aggregated_at),
updated_at = VALUES(updated_at)
"#,
)
.bind(stats_id(&format!("stats-daily:{day_start_unix_secs}")))
.bind(day_start_unix_secs)
.bind(total_requests)
.bind(total_requests.saturating_sub(error_requests))
.bind(error_requests)
.bind(row.try_get::<i64, _>("input_tokens").map_sql_err()?)
.bind(row.try_get::<i64, _>("output_tokens").map_sql_err()?)
.bind(
row.try_get::<i64, _>("cache_creation_tokens")
.map_sql_err()?,
)
.bind(row.try_get::<i64, _>("cache_read_tokens").map_sql_err()?)
.bind(row.try_get::<f64, _>("total_cost").map_sql_err()?)
.bind(row.try_get::<f64, _>("actual_total_cost").map_sql_err()?)
.bind(
row.try_get::<f64, _>("avg_response_time_ms")
.map_sql_err()?,
)
.bind(unique_models)
.bind(unique_providers)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
.execute(&mut *tx)
.await
.map_sql_err()?;
let model_rows = usize::try_from(unique_models).unwrap_or(usize::MAX);
let provider_rows = usize::try_from(unique_providers).unwrap_or(usize::MAX);
let api_key_rows = mysql_group_count(&mut tx, "api_key_id", start_ms, end_ms).await?;
let error_rows = mysql_error_group_count(&mut tx, start_ms, end_ms).await?;
let user_rows = mysql_group_count(&mut tx, "user_id", start_ms, end_ms).await?;
tx.commit().await.map_sql_err()?;
Ok(StatsDailyAggregationSummary {
day_start_utc: utc_from_unix_secs(day_start_unix_secs, "stats_daily.date")?,
total_requests,
model_rows,
provider_rows,
api_key_rows,
error_rows,
user_rows,
})
}
async fn mysql_group_count(
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
group_columns: &str,
start_ms: i64,
end_ms: i64,
) -> Result<usize, DataLayerError> {
let not_empty = group_columns
.split(',')
.map(str::trim)
.map(|column| format!("{column} IS NOT NULL AND {column} <> ''"))
.collect::<Vec<_>>()
.join(" AND ");
let sql = format!(
r#"
SELECT COUNT(*)
FROM (
SELECT 1
FROM `usage`
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
AND {not_empty}
GROUP BY {group_columns}
) AS grouped
"#
);
let count: i64 = sqlx::query_scalar(&sql)
.bind(start_ms)
.bind(end_ms)
.fetch_one(&mut **tx)
.await
.map_sql_err()?;
Ok(usize::try_from(count.max(0)).unwrap_or(usize::MAX))
}
async fn mysql_error_group_count(
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
start_ms: i64,
end_ms: i64,
) -> Result<usize, DataLayerError> {
let count: i64 = sqlx::query_scalar(
r#"
SELECT COUNT(*)
FROM (
SELECT 1
FROM `usage`
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
AND (
status = 'failed'
OR status_code >= 400
OR (error_category IS NOT NULL AND error_category <> '')
)
GROUP BY COALESCE(NULLIF(error_category, ''), 'unknown_error'), provider_name, model
) AS grouped
"#,
)
.bind(start_ms)
.bind(end_ms)
.fetch_one(&mut **tx)
.await
.map_sql_err()?;
Ok(usize::try_from(count.max(0)).unwrap_or(usize::MAX))
}
@@ -0,0 +1,638 @@
use chrono::{DateTime, Utc};
use sqlx::Row;
use uuid::Uuid;
use crate::backend::PostgresBackend;
use crate::{
error::postgres_error, DataLayerError, StatsDailyAggregationInput, StatsDailyAggregationSummary,
};
mod percentiles;
mod sql;
use self::percentiles::{percentile_ms_to_i64, PercentileSummary};
use self::sql::*;
impl PostgresBackend {
pub async fn aggregate_stats_daily(
&self,
input: &StatsDailyAggregationInput,
) -> Result<Option<StatsDailyAggregationSummary>, DataLayerError> {
let Some(day_start_utc) = next_stats_aggregation_day(self.pool(), input.target_day_utc)
.await
.map_err(postgres_error)?
else {
return Ok(None);
};
perform_stats_aggregation_for_day(self.pool(), day_start_utc, input.aggregated_at)
.await
.map(Some)
.map_err(postgres_error)
}
}
async fn next_stats_aggregation_day(
pool: &crate::driver::postgres::PostgresPool,
target_day_utc: DateTime<Utc>,
) -> Result<Option<DateTime<Utc>>, sqlx::Error> {
let latest_row = sqlx::query(SELECT_LATEST_STATS_DAILY_DATE_SQL)
.fetch_one(pool)
.await?;
let latest_day = latest_row.try_get::<Option<DateTime<Utc>>, _>("latest_date")?;
let search_from = latest_day
.map(|value| value + chrono::Duration::days(1))
.unwrap_or_else(|| {
DateTime::<Utc>::from_timestamp(0, 0).expect("unix epoch should be valid")
});
let search_until = target_day_utc + chrono::Duration::days(1);
if search_from >= search_until {
return Ok(None);
}
let next_row = sqlx::query(SELECT_NEXT_STATS_DAILY_BUCKET_SQL)
.bind(search_from)
.bind(search_until)
.fetch_one(pool)
.await?;
let next_bucket = next_row.try_get::<Option<DateTime<Utc>>, _>("next_bucket")?;
Ok(next_bucket.filter(|value| *value <= target_day_utc))
}
async fn perform_stats_aggregation_for_day(
pool: &crate::driver::postgres::PostgresPool,
day_start_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<StatsDailyAggregationSummary, sqlx::Error> {
let day_end_utc = day_start_utc + chrono::Duration::days(1);
let mut tx = pool.begin().await?;
let aggregate_row = sqlx::query(SELECT_STATS_DAILY_AGGREGATE_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.fetch_one(&mut *tx)
.await?;
let total_requests = aggregate_row.try_get::<i64, _>("total_requests")?;
let error_requests = aggregate_row.try_get::<i64, _>("error_requests")?;
let success_requests = total_requests.saturating_sub(error_requests);
let fallback_count = sqlx::query(SELECT_STATS_DAILY_FALLBACK_COUNT_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(vec!["success", "failed"])
.fetch_one(&mut *tx)
.await?
.try_get::<i64, _>("fallback_count")?;
let response_percentiles = fetch_stats_daily_percentiles(
&mut tx,
SELECT_STATS_DAILY_RESPONSE_TIME_PERCENTILES_SQL,
day_start_utc,
day_end_utc,
)
.await?;
let first_byte_percentiles = fetch_stats_daily_percentiles(
&mut tx,
SELECT_STATS_DAILY_FIRST_BYTE_PERCENTILES_SQL,
day_start_utc,
day_end_utc,
)
.await?;
sqlx::query(UPSERT_STATS_DAILY_SQL)
.bind(Uuid::new_v4().to_string())
.bind(day_start_utc)
.bind(total_requests)
.bind(aggregate_row.try_get::<i64, _>("cache_hit_total_requests")?)
.bind(aggregate_row.try_get::<i64, _>("cache_hit_requests")?)
.bind(aggregate_row.try_get::<i64, _>("completed_total_requests")?)
.bind(aggregate_row.try_get::<i64, _>("completed_cache_hit_requests")?)
.bind(aggregate_row.try_get::<i64, _>("completed_input_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("completed_cache_creation_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("completed_cache_read_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("completed_total_input_context")?)
.bind(aggregate_row.try_get::<f64, _>("completed_cache_creation_cost")?)
.bind(aggregate_row.try_get::<f64, _>("completed_cache_read_cost")?)
.bind(aggregate_row.try_get::<f64, _>("settled_total_cost")?)
.bind(aggregate_row.try_get::<i64, _>("settled_total_requests")?)
.bind(aggregate_row.try_get::<i64, _>("settled_input_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("settled_output_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("settled_cache_creation_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("settled_cache_read_tokens")?)
.bind(aggregate_row.try_get::<Option<i64>, _>("settled_first_finalized_at_unix_secs")?)
.bind(aggregate_row.try_get::<Option<i64>, _>("settled_last_finalized_at_unix_secs")?)
.bind(success_requests)
.bind(error_requests)
.bind(aggregate_row.try_get::<i64, _>("input_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("effective_input_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("output_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("cache_creation_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("cache_creation_ephemeral_5m_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("cache_creation_ephemeral_1h_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("cache_read_tokens")?)
.bind(aggregate_row.try_get::<i64, _>("total_input_context")?)
.bind(aggregate_row.try_get::<f64, _>("total_cost")?)
.bind(aggregate_row.try_get::<f64, _>("actual_total_cost")?)
.bind(aggregate_row.try_get::<f64, _>("input_cost")?)
.bind(aggregate_row.try_get::<f64, _>("output_cost")?)
.bind(aggregate_row.try_get::<f64, _>("cache_creation_cost")?)
.bind(aggregate_row.try_get::<f64, _>("cache_read_cost")?)
.bind(aggregate_row.try_get::<f64, _>("response_time_sum_ms")?)
.bind(aggregate_row.try_get::<i64, _>("response_time_samples")?)
.bind(aggregate_row.try_get::<f64, _>("avg_response_time_ms")?)
.bind(response_percentiles.p50)
.bind(response_percentiles.p90)
.bind(response_percentiles.p99)
.bind(first_byte_percentiles.p50)
.bind(first_byte_percentiles.p90)
.bind(first_byte_percentiles.p99)
.bind(fallback_count)
.bind(aggregate_row.try_get::<i64, _>("unique_models")?)
.bind(aggregate_row.try_get::<i64, _>("unique_providers")?)
.bind(true)
.bind(now_utc)
.bind(now_utc)
.bind(now_utc)
.execute(&mut *tx)
.await?;
let model_rows =
upsert_stats_daily_model_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
let provider_rows =
upsert_stats_daily_provider_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
upsert_stats_daily_model_provider_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
upsert_stats_daily_cost_savings_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
upsert_stats_daily_cost_savings_provider_rows(&mut tx, day_start_utc, day_end_utc, now_utc)
.await?;
upsert_stats_daily_cost_savings_model_rows(&mut tx, day_start_utc, day_end_utc, now_utc)
.await?;
upsert_stats_daily_cost_savings_model_provider_rows(
&mut tx,
day_start_utc,
day_end_utc,
now_utc,
)
.await?;
let api_key_rows =
upsert_stats_daily_api_key_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
let error_rows =
refresh_stats_daily_error_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
let user_rows =
upsert_stats_user_daily_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
upsert_stats_user_daily_model_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
upsert_stats_user_daily_model_provider_rows(&mut tx, day_start_utc, day_end_utc, now_utc)
.await?;
upsert_stats_user_daily_provider_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
upsert_stats_user_daily_cost_savings_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
upsert_stats_user_daily_cost_savings_provider_rows(
&mut tx,
day_start_utc,
day_end_utc,
now_utc,
)
.await?;
upsert_stats_user_daily_cost_savings_model_rows(&mut tx, day_start_utc, day_end_utc, now_utc)
.await?;
upsert_stats_user_daily_cost_savings_model_provider_rows(
&mut tx,
day_start_utc,
day_end_utc,
now_utc,
)
.await?;
upsert_stats_user_daily_api_format_rows(&mut tx, day_start_utc, day_end_utc, now_utc).await?;
refresh_stats_summary_row(&mut tx, day_end_utc, now_utc).await?;
refresh_stats_user_summary_rows(&mut tx, day_end_utc, now_utc).await?;
tx.commit().await?;
Ok(StatsDailyAggregationSummary {
day_start_utc,
total_requests,
model_rows,
provider_rows,
api_key_rows,
error_rows,
user_rows,
})
}
async fn fetch_stats_daily_percentiles(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
sql: &str,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
) -> Result<PercentileSummary, sqlx::Error> {
let row = sqlx::query(sql)
.bind(day_start_utc)
.bind(day_end_utc)
.fetch_one(&mut **tx)
.await?;
let sample_count = row.try_get::<i64, _>("sample_count")?;
if sample_count < 10 {
return Ok(PercentileSummary::default());
}
Ok(PercentileSummary {
p50: percentile_ms_to_i64(row.try_get::<Option<f64>, _>("p50")?),
p90: percentile_ms_to_i64(row.try_get::<Option<f64>, _>("p90")?),
p99: percentile_ms_to_i64(row.try_get::<Option<f64>, _>("p99")?),
})
}
async fn upsert_stats_daily_model_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_DAILY_MODEL_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_daily_provider_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_DAILY_PROVIDER_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_daily_model_provider_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_DAILY_MODEL_PROVIDER_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_daily_cost_savings_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_DAILY_COST_SAVINGS_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_daily_cost_savings_provider_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_DAILY_COST_SAVINGS_PROVIDER_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_daily_cost_savings_model_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_DAILY_COST_SAVINGS_MODEL_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_daily_cost_savings_model_provider_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_DAILY_COST_SAVINGS_MODEL_PROVIDER_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_daily_api_key_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_DAILY_API_KEY_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn refresh_stats_daily_error_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
sqlx::query(DELETE_STATS_DAILY_ERRORS_FOR_DATE_SQL)
.bind(day_start_utc)
.execute(&mut **tx)
.await?;
let rows_affected = sqlx::query(INSERT_STATS_DAILY_ERROR_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_user_daily_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_USER_DAILY_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_user_daily_model_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_USER_DAILY_MODEL_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_user_daily_model_provider_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_USER_DAILY_MODEL_PROVIDER_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_user_daily_provider_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_USER_DAILY_PROVIDER_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_user_daily_cost_savings_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_USER_DAILY_COST_SAVINGS_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_user_daily_cost_savings_provider_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_USER_DAILY_COST_SAVINGS_PROVIDER_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_user_daily_cost_savings_model_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_USER_DAILY_COST_SAVINGS_MODEL_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_user_daily_cost_savings_model_provider_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_USER_DAILY_COST_SAVINGS_MODEL_PROVIDER_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_user_daily_api_format_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
day_start_utc: DateTime<Utc>,
day_end_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_USER_DAILY_API_FORMAT_SQL)
.bind(day_start_utc)
.bind(day_end_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn refresh_stats_summary_row(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
cutoff_date: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<(), sqlx::Error> {
let totals_row = sqlx::query(SELECT_STATS_SUMMARY_TOTALS_SQL)
.bind(cutoff_date)
.fetch_one(&mut **tx)
.await?;
let entity_counts_row = sqlx::query(SELECT_STATS_SUMMARY_ENTITY_COUNTS_SQL)
.fetch_one(&mut **tx)
.await?;
let existing_summary_id = sqlx::query_scalar::<_, String>(SELECT_EXISTING_STATS_SUMMARY_ID_SQL)
.fetch_optional(&mut **tx)
.await?;
let all_time_requests = totals_row.try_get::<i64, _>("all_time_requests")?;
let all_time_success_requests = totals_row.try_get::<i64, _>("all_time_success_requests")?;
let all_time_error_requests = totals_row.try_get::<i64, _>("all_time_error_requests")?;
let all_time_input_tokens = totals_row.try_get::<i64, _>("all_time_input_tokens")?;
let all_time_output_tokens = totals_row.try_get::<i64, _>("all_time_output_tokens")?;
let all_time_cache_creation_tokens =
totals_row.try_get::<i64, _>("all_time_cache_creation_tokens")?;
let all_time_cache_read_tokens = totals_row.try_get::<i64, _>("all_time_cache_read_tokens")?;
let all_time_cost = totals_row.try_get::<f64, _>("all_time_cost")?;
let all_time_actual_cost = totals_row.try_get::<f64, _>("all_time_actual_cost")?;
let total_users = entity_counts_row.try_get::<i64, _>("total_users")?;
let active_users = entity_counts_row.try_get::<i64, _>("active_users")?;
let total_api_keys = entity_counts_row.try_get::<i64, _>("total_api_keys")?;
let active_api_keys = entity_counts_row.try_get::<i64, _>("active_api_keys")?;
if let Some(summary_id) = existing_summary_id {
sqlx::query(UPDATE_STATS_SUMMARY_SQL)
.bind(summary_id)
.bind(cutoff_date)
.bind(all_time_requests)
.bind(all_time_success_requests)
.bind(all_time_error_requests)
.bind(all_time_input_tokens)
.bind(all_time_output_tokens)
.bind(all_time_cache_creation_tokens)
.bind(all_time_cache_read_tokens)
.bind(all_time_cost)
.bind(all_time_actual_cost)
.bind(total_users)
.bind(active_users)
.bind(total_api_keys)
.bind(active_api_keys)
.bind(now_utc)
.execute(&mut **tx)
.await?;
} else {
sqlx::query(INSERT_STATS_SUMMARY_SQL)
.bind(Uuid::new_v4().to_string())
.bind(cutoff_date)
.bind(all_time_requests)
.bind(all_time_success_requests)
.bind(all_time_error_requests)
.bind(all_time_input_tokens)
.bind(all_time_output_tokens)
.bind(all_time_cache_creation_tokens)
.bind(all_time_cache_read_tokens)
.bind(all_time_cost)
.bind(all_time_actual_cost)
.bind(total_users)
.bind(active_users)
.bind(total_api_keys)
.bind(active_api_keys)
.bind(now_utc)
.bind(now_utc)
.execute(&mut **tx)
.await?;
}
Ok(())
}
async fn refresh_stats_user_summary_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
cutoff_date: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<(), sqlx::Error> {
sqlx::query(UPSERT_STATS_USER_SUMMARY_SQL)
.bind(cutoff_date)
.bind(now_utc)
.execute(&mut **tx)
.await?;
Ok(())
}
@@ -0,0 +1,10 @@
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub(super) struct PercentileSummary {
pub(super) p50: Option<i64>,
pub(super) p90: Option<i64>,
pub(super) p99: Option<i64>,
}
pub(super) fn percentile_ms_to_i64(value: Option<f64>) -> Option<i64> {
value.and_then(|raw| raw.is_finite().then_some(raw.floor() as i64))
}
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,203 @@
use chrono::{DateTime, Utc};
use sqlx::Row;
use uuid::Uuid;
use crate::backend::PostgresBackend;
use crate::{
error::postgres_error, DataLayerError, StatsHourlyAggregationInput,
StatsHourlyAggregationSummary,
};
mod sql;
use self::sql::*;
impl PostgresBackend {
pub async fn aggregate_stats_hourly(
&self,
input: &StatsHourlyAggregationInput,
) -> Result<Option<StatsHourlyAggregationSummary>, DataLayerError> {
let Some(hour_utc) = next_stats_hourly_bucket(self.pool(), input.target_hour_utc)
.await
.map_err(postgres_error)?
else {
return Ok(None);
};
perform_stats_hourly_aggregation_for_hour(self.pool(), hour_utc, input.aggregated_at)
.await
.map(Some)
.map_err(postgres_error)
}
}
async fn next_stats_hourly_bucket(
pool: &crate::driver::postgres::PostgresPool,
target_hour_utc: DateTime<Utc>,
) -> Result<Option<DateTime<Utc>>, sqlx::Error> {
let latest_row = sqlx::query(SELECT_LATEST_STATS_HOURLY_HOUR_SQL)
.fetch_one(pool)
.await?;
let latest_hour = latest_row.try_get::<Option<DateTime<Utc>>, _>("latest_hour")?;
let search_from = latest_hour
.map(|value| value + chrono::Duration::hours(1))
.unwrap_or_else(|| {
DateTime::<Utc>::from_timestamp(0, 0).expect("unix epoch should be valid")
});
let search_until = target_hour_utc + chrono::Duration::hours(1);
if search_from >= search_until {
return Ok(None);
}
let next_row = sqlx::query(SELECT_NEXT_STATS_HOURLY_BUCKET_SQL)
.bind(search_from)
.bind(search_until)
.fetch_one(pool)
.await?;
let next_bucket = next_row.try_get::<Option<DateTime<Utc>>, _>("next_bucket")?;
Ok(next_bucket.filter(|value| *value <= target_hour_utc))
}
async fn perform_stats_hourly_aggregation_for_hour(
pool: &crate::driver::postgres::PostgresPool,
hour_utc: DateTime<Utc>,
aggregated_at: DateTime<Utc>,
) -> Result<StatsHourlyAggregationSummary, sqlx::Error> {
let hour_end = hour_utc + chrono::Duration::hours(1);
let mut tx = pool.begin().await?;
let row = sqlx::query(SELECT_STATS_HOURLY_AGGREGATE_SQL)
.bind(hour_utc)
.bind(hour_end)
.fetch_one(&mut *tx)
.await?;
let total_requests = row.try_get::<i64, _>("total_requests")?;
let error_requests = row.try_get::<i64, _>("error_requests")?;
let success_requests = total_requests.saturating_sub(error_requests);
sqlx::query(UPSERT_STATS_HOURLY_SQL)
.bind(Uuid::new_v4().to_string())
.bind(hour_utc)
.bind(total_requests)
.bind(row.try_get::<i64, _>("cache_hit_total_requests")?)
.bind(row.try_get::<i64, _>("cache_hit_requests")?)
.bind(row.try_get::<i64, _>("completed_total_requests")?)
.bind(row.try_get::<i64, _>("completed_cache_hit_requests")?)
.bind(row.try_get::<i64, _>("completed_input_tokens")?)
.bind(row.try_get::<i64, _>("completed_cache_creation_tokens")?)
.bind(row.try_get::<i64, _>("completed_cache_read_tokens")?)
.bind(row.try_get::<i64, _>("completed_total_input_context")?)
.bind(row.try_get::<f64, _>("completed_cache_creation_cost")?)
.bind(row.try_get::<f64, _>("completed_cache_read_cost")?)
.bind(row.try_get::<f64, _>("settled_total_cost")?)
.bind(row.try_get::<i64, _>("settled_total_requests")?)
.bind(row.try_get::<i64, _>("settled_input_tokens")?)
.bind(row.try_get::<i64, _>("settled_output_tokens")?)
.bind(row.try_get::<i64, _>("settled_cache_creation_tokens")?)
.bind(row.try_get::<i64, _>("settled_cache_read_tokens")?)
.bind(row.try_get::<Option<i64>, _>("settled_first_finalized_at_unix_secs")?)
.bind(row.try_get::<Option<i64>, _>("settled_last_finalized_at_unix_secs")?)
.bind(success_requests)
.bind(error_requests)
.bind(row.try_get::<i64, _>("input_tokens")?)
.bind(row.try_get::<i64, _>("output_tokens")?)
.bind(row.try_get::<i64, _>("cache_creation_tokens")?)
.bind(row.try_get::<i64, _>("cache_read_tokens")?)
.bind(row.try_get::<f64, _>("total_cost")?)
.bind(row.try_get::<f64, _>("actual_total_cost")?)
.bind(row.try_get::<f64, _>("response_time_sum_ms")?)
.bind(row.try_get::<i64, _>("response_time_samples")?)
.bind(row.try_get::<f64, _>("avg_response_time_ms")?)
.bind(true)
.bind(aggregated_at)
.bind(aggregated_at)
.bind(aggregated_at)
.execute(&mut *tx)
.await?;
let user_rows =
upsert_stats_hourly_user_rows(&mut tx, hour_utc, hour_end, aggregated_at).await?;
let user_model_rows =
upsert_stats_hourly_user_model_rows(&mut tx, hour_utc, hour_end, aggregated_at).await?;
let model_rows =
upsert_stats_hourly_model_rows(&mut tx, hour_utc, hour_end, aggregated_at).await?;
let provider_rows =
upsert_stats_hourly_provider_rows(&mut tx, hour_utc, hour_end, aggregated_at).await?;
tx.commit().await?;
Ok(StatsHourlyAggregationSummary {
hour_utc,
total_requests,
user_rows,
user_model_rows,
model_rows,
provider_rows,
})
}
async fn upsert_stats_hourly_user_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
hour_utc: DateTime<Utc>,
hour_end: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_HOURLY_USER_SQL)
.bind(hour_utc)
.bind(hour_end)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_hourly_model_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
hour_utc: DateTime<Utc>,
hour_end: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_HOURLY_MODEL_SQL)
.bind(hour_utc)
.bind(hour_end)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_hourly_user_model_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
hour_utc: DateTime<Utc>,
hour_end: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_HOURLY_USER_MODEL_SQL)
.bind(hour_utc)
.bind(hour_end)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
async fn upsert_stats_hourly_provider_rows(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
hour_utc: DateTime<Utc>,
hour_end: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<usize, sqlx::Error> {
let rows_affected = sqlx::query(UPSERT_STATS_HOURLY_PROVIDER_SQL)
.bind(hour_utc)
.bind(hour_end)
.bind(now_utc)
.execute(&mut **tx)
.await?
.rows_affected();
Ok(usize::try_from(rows_affected).unwrap_or(usize::MAX))
}
@@ -0,0 +1,941 @@
pub(super) const SELECT_LATEST_STATS_HOURLY_HOUR_SQL: &str = r#"
SELECT MAX(hour_utc) AS latest_hour
FROM stats_hourly
WHERE is_complete IS TRUE
"#;
pub(super) const SELECT_NEXT_STATS_HOURLY_BUCKET_SQL: &str = r#"
SELECT date_trunc('hour', MIN(created_at)) AS next_bucket
FROM usage_billing_facts AS usage
WHERE created_at >= $1
AND created_at < $2
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
"#;
pub(super) const SELECT_STATS_HOURLY_AGGREGATE_SQL: &str = r#"
SELECT
CAST(COUNT(usage.id) AS BIGINT) AS cache_hit_total_requests,
CAST(
COUNT(usage.id) FILTER (
WHERE GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0) > 0
) AS BIGINT
) AS cache_hit_requests,
CAST(
COUNT(usage.id) FILTER (WHERE usage.status = 'completed') AS BIGINT
) AS completed_total_requests,
CAST(
COUNT(usage.id) FILTER (
WHERE usage.status = 'completed'
AND GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0) > 0
) AS BIGINT
) AS completed_cache_hit_requests,
CAST(
COALESCE(
SUM(GREATEST(COALESCE(usage.input_tokens, 0), 0))
FILTER (WHERE usage.status = 'completed'),
0
) AS BIGINT
) AS completed_input_tokens,
CAST(
COALESCE(
SUM(
CASE
WHEN COALESCE(usage.cache_creation_input_tokens, 0) = 0
AND (
COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
) > 0
THEN COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
ELSE COALESCE(usage.cache_creation_input_tokens, 0)
END
) FILTER (WHERE usage.status = 'completed'),
0
) AS BIGINT
) AS completed_cache_creation_tokens,
CAST(
COALESCE(
SUM(GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0))
FILTER (WHERE usage.status = 'completed'),
0
) AS BIGINT
) AS completed_cache_read_tokens,
CAST(
COALESCE(
SUM(
CASE
WHEN split_part(
lower(
COALESCE(
COALESCE(usage.endpoint_api_format, usage.api_format),
''
)
),
':',
1
) IN ('claude', 'anthropic')
THEN GREATEST(COALESCE(usage.input_tokens, 0), 0)
+ CASE
WHEN COALESCE(usage.cache_creation_input_tokens, 0) = 0
AND (
COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
) > 0
THEN COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
ELSE COALESCE(usage.cache_creation_input_tokens, 0)
END
+ GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0)
WHEN split_part(
lower(
COALESCE(
COALESCE(usage.endpoint_api_format, usage.api_format),
''
)
),
':',
1
) IN ('openai', 'gemini', 'google')
THEN (
CASE
WHEN GREATEST(COALESCE(usage.input_tokens, 0), 0) <= 0
THEN 0
WHEN GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0) <= 0
THEN GREATEST(COALESCE(usage.input_tokens, 0), 0)
ELSE GREATEST(
GREATEST(COALESCE(usage.input_tokens, 0), 0)
- GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0),
0
)
END
) + GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0)
ELSE CASE
WHEN (
CASE
WHEN COALESCE(usage.cache_creation_input_tokens, 0) = 0
AND (
COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
) > 0
THEN COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
ELSE COALESCE(usage.cache_creation_input_tokens, 0)
END
) > 0
THEN GREATEST(COALESCE(usage.input_tokens, 0), 0)
+ (
CASE
WHEN COALESCE(usage.cache_creation_input_tokens, 0) = 0
AND (
COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
) > 0
THEN COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
ELSE COALESCE(usage.cache_creation_input_tokens, 0)
END
)
+ GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0)
ELSE GREATEST(COALESCE(usage.input_tokens, 0), 0)
+ GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0)
END
END
) FILTER (WHERE usage.status = 'completed'),
0
) AS BIGINT
) AS completed_total_input_context,
CAST(
COALESCE(
SUM(
COALESCE(CAST(usage.cache_creation_cost_usd AS DOUBLE PRECISION), 0)
) FILTER (WHERE usage.status = 'completed'),
0
) AS DOUBLE PRECISION
) AS completed_cache_creation_cost,
CAST(
COALESCE(
SUM(COALESCE(CAST(usage.cache_read_cost_usd AS DOUBLE PRECISION), 0))
FILTER (WHERE usage.status = 'completed'),
0
) AS DOUBLE PRECISION
) AS completed_cache_read_cost,
CAST(
COALESCE(
SUM(COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0)) FILTER (
WHERE usage.billing_status = 'settled'
AND COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0) > 0
),
0
) AS DOUBLE PRECISION
) AS settled_total_cost,
CAST(
COUNT(usage.id) FILTER (
WHERE usage.billing_status = 'settled'
AND COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0) > 0
) AS BIGINT
) AS settled_total_requests,
CAST(
COALESCE(
SUM(GREATEST(COALESCE(usage.input_tokens, 0), 0)) FILTER (
WHERE usage.billing_status = 'settled'
AND COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0) > 0
),
0
) AS BIGINT
) AS settled_input_tokens,
CAST(
COALESCE(
SUM(GREATEST(COALESCE(usage.output_tokens, 0), 0)) FILTER (
WHERE usage.billing_status = 'settled'
AND COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0) > 0
),
0
) AS BIGINT
) AS settled_output_tokens,
CAST(
COALESCE(
SUM(GREATEST(COALESCE(usage.cache_creation_input_tokens, 0), 0)) FILTER (
WHERE usage.billing_status = 'settled'
AND COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0) > 0
),
0
) AS BIGINT
) AS settled_cache_creation_tokens,
CAST(
COALESCE(
SUM(GREATEST(COALESCE(usage.cache_read_input_tokens, 0), 0)) FILTER (
WHERE usage.billing_status = 'settled'
AND COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0) > 0
),
0
) AS BIGINT
) AS settled_cache_read_tokens,
MIN(CAST(EXTRACT(EPOCH FROM usage.finalized_at) AS BIGINT)) FILTER (
WHERE usage.billing_status = 'settled'
AND COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0) > 0
) AS settled_first_finalized_at_unix_secs,
MAX(CAST(EXTRACT(EPOCH FROM usage.finalized_at) AS BIGINT)) FILTER (
WHERE usage.billing_status = 'settled'
AND COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0) > 0
) AS settled_last_finalized_at_unix_secs,
CAST(
COUNT(usage.id) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
) AS BIGINT
) AS total_requests,
CAST(COALESCE(
SUM(
CASE
WHEN usage.status_code >= 400
OR lower(COALESCE(usage.status, '')) = 'failed'
OR usage.error_message IS NOT NULL THEN 1
ELSE 0
END
) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS BIGINT) AS error_requests,
CAST(
COALESCE(
SUM(usage.input_tokens) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS BIGINT
) AS input_tokens,
CAST(
COALESCE(
SUM(usage.output_tokens) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS BIGINT
) AS output_tokens,
CAST(COALESCE(
SUM(
CASE
WHEN COALESCE(usage.cache_creation_input_tokens, 0) = 0
AND (
COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
) > 0
THEN COALESCE(usage.cache_creation_input_tokens_5m, 0)
+ COALESCE(usage.cache_creation_input_tokens_1h, 0)
ELSE COALESCE(usage.cache_creation_input_tokens, 0)
END
) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS BIGINT) AS cache_creation_tokens,
CAST(
COALESCE(
SUM(usage.cache_read_input_tokens) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS BIGINT
) AS cache_read_tokens,
CAST(
COALESCE(
SUM(usage.total_cost_usd) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS DOUBLE PRECISION
) AS total_cost,
CAST(
COALESCE(
SUM(usage.actual_total_cost_usd) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS DOUBLE PRECISION
) AS actual_total_cost,
COALESCE(
SUM(
CASE
WHEN usage.response_time_ms IS NOT NULL
THEN GREATEST(COALESCE(usage.response_time_ms, 0), 0)::DOUBLE PRECISION
ELSE 0
END
) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS response_time_sum_ms,
CAST(COALESCE(
SUM(
CASE
WHEN usage.response_time_ms IS NOT NULL THEN 1
ELSE 0
END
) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS BIGINT) AS response_time_samples,
CAST(
COALESCE(
AVG(usage.response_time_ms) FILTER (
WHERE usage.status NOT IN ('pending', 'streaming')
AND usage.provider_name NOT IN ('unknown', 'pending')
),
0
) AS DOUBLE PRECISION
) AS avg_response_time_ms
FROM usage_billing_facts AS usage
WHERE usage.created_at >= $1
AND usage.created_at < $2
"#;
pub(super) const UPSERT_STATS_HOURLY_SQL: &str = r#"
INSERT INTO stats_hourly (
id,
hour_utc,
total_requests,
cache_hit_total_requests,
cache_hit_requests,
completed_total_requests,
completed_cache_hit_requests,
completed_input_tokens,
completed_cache_creation_tokens,
completed_cache_read_tokens,
completed_total_input_context,
completed_cache_creation_cost,
completed_cache_read_cost,
settled_total_cost,
settled_total_requests,
settled_input_tokens,
settled_output_tokens,
settled_cache_creation_tokens,
settled_cache_read_tokens,
settled_first_finalized_at_unix_secs,
settled_last_finalized_at_unix_secs,
success_requests,
error_requests,
input_tokens,
output_tokens,
cache_creation_tokens,
cache_read_tokens,
total_cost,
actual_total_cost,
response_time_sum_ms,
response_time_samples,
avg_response_time_ms,
is_complete,
aggregated_at,
created_at,
updated_at
)
VALUES (
$1, $2, $3, $4, $5, $6, $7, $8,
$9, $10, $11, $12, $13, $14, $15, $16,
$17, $18, $19, $20, $21, $22, $23, $24,
$25, $26, $27, $28, $29, $30, $31, $32,
$33, $34, $35, $36
)
ON CONFLICT (hour_utc)
DO UPDATE SET
total_requests = EXCLUDED.total_requests,
cache_hit_total_requests = EXCLUDED.cache_hit_total_requests,
cache_hit_requests = EXCLUDED.cache_hit_requests,
completed_total_requests = EXCLUDED.completed_total_requests,
completed_cache_hit_requests = EXCLUDED.completed_cache_hit_requests,
completed_input_tokens = EXCLUDED.completed_input_tokens,
completed_cache_creation_tokens = EXCLUDED.completed_cache_creation_tokens,
completed_cache_read_tokens = EXCLUDED.completed_cache_read_tokens,
completed_total_input_context = EXCLUDED.completed_total_input_context,
completed_cache_creation_cost = EXCLUDED.completed_cache_creation_cost,
completed_cache_read_cost = EXCLUDED.completed_cache_read_cost,
settled_total_cost = EXCLUDED.settled_total_cost,
settled_total_requests = EXCLUDED.settled_total_requests,
settled_input_tokens = EXCLUDED.settled_input_tokens,
settled_output_tokens = EXCLUDED.settled_output_tokens,
settled_cache_creation_tokens = EXCLUDED.settled_cache_creation_tokens,
settled_cache_read_tokens = EXCLUDED.settled_cache_read_tokens,
settled_first_finalized_at_unix_secs = EXCLUDED.settled_first_finalized_at_unix_secs,
settled_last_finalized_at_unix_secs = EXCLUDED.settled_last_finalized_at_unix_secs,
success_requests = EXCLUDED.success_requests,
error_requests = EXCLUDED.error_requests,
input_tokens = EXCLUDED.input_tokens,
output_tokens = EXCLUDED.output_tokens,
cache_creation_tokens = EXCLUDED.cache_creation_tokens,
cache_read_tokens = EXCLUDED.cache_read_tokens,
total_cost = EXCLUDED.total_cost,
actual_total_cost = EXCLUDED.actual_total_cost,
response_time_sum_ms = EXCLUDED.response_time_sum_ms,
response_time_samples = EXCLUDED.response_time_samples,
avg_response_time_ms = EXCLUDED.avg_response_time_ms,
is_complete = EXCLUDED.is_complete,
aggregated_at = EXCLUDED.aggregated_at,
updated_at = EXCLUDED.updated_at
"#;
pub(super) const UPSERT_STATS_HOURLY_USER_SQL: &str = r#"
WITH aggregated AS (
SELECT
user_id,
CAST(COUNT(id) AS BIGINT) AS total_requests,
CAST(COALESCE(
SUM(
CASE
WHEN status_code >= 400
OR lower(COALESCE(status, '')) = 'failed'
OR error_message IS NOT NULL THEN 1
ELSE 0
END
),
0
) AS BIGINT) AS error_requests,
CAST(COALESCE(SUM(input_tokens), 0) AS BIGINT) AS input_tokens,
CAST(COALESCE(SUM(output_tokens), 0) AS BIGINT) AS output_tokens,
CAST(COALESCE(
SUM(
CASE
WHEN COALESCE(cache_creation_input_tokens, 0) = 0
AND (
COALESCE(cache_creation_input_tokens_5m, 0)
+ COALESCE(cache_creation_input_tokens_1h, 0)
) > 0
THEN COALESCE(cache_creation_input_tokens_5m, 0)
+ COALESCE(cache_creation_input_tokens_1h, 0)
ELSE COALESCE(cache_creation_input_tokens, 0)
END
),
0
) AS BIGINT) AS cache_creation_tokens,
CAST(COALESCE(SUM(cache_read_input_tokens), 0) AS BIGINT) AS cache_read_tokens,
CAST(COALESCE(SUM(total_cost_usd), 0) AS DOUBLE PRECISION) AS total_cost,
CAST(COALESCE(SUM(actual_total_cost_usd), 0) AS DOUBLE PRECISION) AS actual_total_cost,
CAST(
COALESCE(
SUM(
CASE
WHEN billing_status = 'settled'
AND COALESCE(CAST(total_cost_usd AS DOUBLE PRECISION), 0) > 0
THEN COALESCE(CAST(total_cost_usd AS DOUBLE PRECISION), 0)
ELSE 0
END
),
0
) AS DOUBLE PRECISION
) AS settled_total_cost,
CAST(COALESCE(
SUM(
CASE
WHEN billing_status = 'settled'
AND COALESCE(CAST(total_cost_usd AS DOUBLE PRECISION), 0) > 0
THEN 1
ELSE 0
END
),
0
) AS BIGINT) AS settled_total_requests,
CAST(COALESCE(
SUM(
CASE
WHEN billing_status = 'settled'
AND COALESCE(CAST(total_cost_usd AS DOUBLE PRECISION), 0) > 0
THEN GREATEST(COALESCE(input_tokens, 0), 0)
ELSE 0
END
),
0
) AS BIGINT) AS settled_input_tokens,
CAST(COALESCE(
SUM(
CASE
WHEN billing_status = 'settled'
AND COALESCE(CAST(total_cost_usd AS DOUBLE PRECISION), 0) > 0
THEN GREATEST(COALESCE(output_tokens, 0), 0)
ELSE 0
END
),
0
) AS BIGINT) AS settled_output_tokens,
CAST(COALESCE(
SUM(
CASE
WHEN billing_status = 'settled'
AND COALESCE(CAST(total_cost_usd AS DOUBLE PRECISION), 0) > 0
THEN GREATEST(COALESCE(cache_creation_input_tokens, 0), 0)
ELSE 0
END
),
0
) AS BIGINT) AS settled_cache_creation_tokens,
CAST(COALESCE(
SUM(
CASE
WHEN billing_status = 'settled'
AND COALESCE(CAST(total_cost_usd AS DOUBLE PRECISION), 0) > 0
THEN GREATEST(COALESCE(cache_read_input_tokens, 0), 0)
ELSE 0
END
),
0
) AS BIGINT) AS settled_cache_read_tokens,
MIN(
CASE
WHEN billing_status = 'settled'
AND COALESCE(CAST(total_cost_usd AS DOUBLE PRECISION), 0) > 0
AND finalized_at IS NOT NULL
THEN CAST(EXTRACT(EPOCH FROM finalized_at) AS BIGINT)
ELSE NULL
END
) AS settled_first_finalized_at_unix_secs,
MAX(
CASE
WHEN billing_status = 'settled'
AND COALESCE(CAST(total_cost_usd AS DOUBLE PRECISION), 0) > 0
AND finalized_at IS NOT NULL
THEN CAST(EXTRACT(EPOCH FROM finalized_at) AS BIGINT)
ELSE NULL
END
) AS settled_last_finalized_at_unix_secs,
COALESCE(
SUM(
CASE
WHEN response_time_ms IS NOT NULL
THEN GREATEST(COALESCE(response_time_ms, 0), 0)::DOUBLE PRECISION
ELSE 0
END
),
0
) AS response_time_sum_ms,
CAST(COALESCE(
SUM(
CASE
WHEN response_time_ms IS NOT NULL THEN 1
ELSE 0
END
),
0
) AS BIGINT) AS response_time_samples
FROM usage_billing_facts AS usage
WHERE created_at >= $1
AND created_at < $2
AND user_id IS NOT NULL
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
GROUP BY user_id
)
INSERT INTO stats_hourly_user (
id,
hour_utc,
user_id,
total_requests,
success_requests,
error_requests,
input_tokens,
output_tokens,
cache_creation_tokens,
cache_read_tokens,
total_cost,
actual_total_cost,
settled_total_cost,
settled_total_requests,
settled_input_tokens,
settled_output_tokens,
settled_cache_creation_tokens,
settled_cache_read_tokens,
settled_first_finalized_at_unix_secs,
settled_last_finalized_at_unix_secs,
response_time_sum_ms,
response_time_samples,
created_at,
updated_at
)
SELECT
md5(CONCAT('stats-hourly-user:', aggregated.user_id, ':', CAST($1 AS TEXT))),
$1,
aggregated.user_id,
aggregated.total_requests,
GREATEST(aggregated.total_requests - aggregated.error_requests, 0),
aggregated.error_requests,
aggregated.input_tokens,
aggregated.output_tokens,
aggregated.cache_creation_tokens,
aggregated.cache_read_tokens,
aggregated.total_cost,
aggregated.actual_total_cost,
aggregated.settled_total_cost,
aggregated.settled_total_requests,
aggregated.settled_input_tokens,
aggregated.settled_output_tokens,
aggregated.settled_cache_creation_tokens,
aggregated.settled_cache_read_tokens,
aggregated.settled_first_finalized_at_unix_secs,
aggregated.settled_last_finalized_at_unix_secs,
aggregated.response_time_sum_ms,
aggregated.response_time_samples,
$3,
$3
FROM aggregated
ON CONFLICT (hour_utc, user_id)
DO UPDATE SET
total_requests = EXCLUDED.total_requests,
success_requests = EXCLUDED.success_requests,
error_requests = EXCLUDED.error_requests,
input_tokens = EXCLUDED.input_tokens,
output_tokens = EXCLUDED.output_tokens,
cache_creation_tokens = EXCLUDED.cache_creation_tokens,
cache_read_tokens = EXCLUDED.cache_read_tokens,
total_cost = EXCLUDED.total_cost,
actual_total_cost = EXCLUDED.actual_total_cost,
settled_total_cost = EXCLUDED.settled_total_cost,
settled_total_requests = EXCLUDED.settled_total_requests,
settled_input_tokens = EXCLUDED.settled_input_tokens,
settled_output_tokens = EXCLUDED.settled_output_tokens,
settled_cache_creation_tokens = EXCLUDED.settled_cache_creation_tokens,
settled_cache_read_tokens = EXCLUDED.settled_cache_read_tokens,
settled_first_finalized_at_unix_secs = EXCLUDED.settled_first_finalized_at_unix_secs,
settled_last_finalized_at_unix_secs = EXCLUDED.settled_last_finalized_at_unix_secs,
response_time_sum_ms = EXCLUDED.response_time_sum_ms,
response_time_samples = EXCLUDED.response_time_samples,
updated_at = EXCLUDED.updated_at
"#;
pub(super) const UPSERT_STATS_HOURLY_MODEL_SQL: &str = r#"
WITH aggregated AS (
SELECT
model,
CAST(COUNT(id) AS BIGINT) AS total_requests,
CAST(COALESCE(SUM(input_tokens), 0) AS BIGINT) AS input_tokens,
CAST(COALESCE(SUM(output_tokens), 0) AS BIGINT) AS output_tokens,
CAST(COALESCE(SUM(total_cost_usd), 0) AS DOUBLE PRECISION) AS total_cost,
COALESCE(
SUM(
CASE
WHEN response_time_ms IS NOT NULL
THEN GREATEST(COALESCE(response_time_ms, 0), 0)::DOUBLE PRECISION
ELSE 0
END
),
0
) AS response_time_sum_ms,
CAST(COALESCE(
SUM(
CASE
WHEN response_time_ms IS NOT NULL THEN 1
ELSE 0
END
),
0
) AS BIGINT) AS response_time_samples,
CAST(COALESCE(AVG(response_time_ms), 0) AS DOUBLE PRECISION) AS avg_response_time_ms
FROM usage_billing_facts AS usage
WHERE created_at >= $1
AND created_at < $2
AND model IS NOT NULL
AND model <> ''
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
GROUP BY model
)
INSERT INTO stats_hourly_model (
id,
hour_utc,
model,
total_requests,
input_tokens,
output_tokens,
total_cost,
response_time_sum_ms,
response_time_samples,
avg_response_time_ms,
created_at,
updated_at
)
SELECT
md5(CONCAT('stats-hourly-model:', aggregated.model, ':', CAST($1 AS TEXT))),
$1,
aggregated.model,
aggregated.total_requests,
aggregated.input_tokens,
aggregated.output_tokens,
aggregated.total_cost,
aggregated.response_time_sum_ms,
aggregated.response_time_samples,
aggregated.avg_response_time_ms,
$3,
$3
FROM aggregated
ON CONFLICT (hour_utc, model)
DO UPDATE SET
total_requests = EXCLUDED.total_requests,
input_tokens = EXCLUDED.input_tokens,
output_tokens = EXCLUDED.output_tokens,
total_cost = EXCLUDED.total_cost,
response_time_sum_ms = EXCLUDED.response_time_sum_ms,
response_time_samples = EXCLUDED.response_time_samples,
avg_response_time_ms = EXCLUDED.avg_response_time_ms,
updated_at = EXCLUDED.updated_at
"#;
pub(super) const UPSERT_STATS_HOURLY_USER_MODEL_SQL: &str = r#"
WITH aggregated AS (
SELECT
user_id,
model,
CAST(COUNT(id) AS BIGINT) AS total_requests,
CAST(COALESCE(SUM(input_tokens), 0) AS BIGINT) AS input_tokens,
CAST(COALESCE(SUM(output_tokens), 0) AS BIGINT) AS output_tokens,
CAST(COALESCE(SUM(total_cost_usd), 0) AS DOUBLE PRECISION) AS total_cost,
COALESCE(
SUM(
CASE
WHEN response_time_ms IS NOT NULL
THEN GREATEST(COALESCE(response_time_ms, 0), 0)::DOUBLE PRECISION
ELSE 0
END
),
0
) AS response_time_sum_ms,
CAST(COALESCE(
SUM(
CASE
WHEN response_time_ms IS NOT NULL THEN 1
ELSE 0
END
),
0
) AS BIGINT) AS response_time_samples
FROM usage_billing_facts AS usage
WHERE created_at >= $1
AND created_at < $2
AND user_id IS NOT NULL
AND model IS NOT NULL
AND model <> ''
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
GROUP BY user_id, model
)
INSERT INTO stats_hourly_user_model (
id,
hour_utc,
user_id,
model,
total_requests,
input_tokens,
output_tokens,
total_cost,
response_time_sum_ms,
response_time_samples,
created_at,
updated_at
)
SELECT
md5(CONCAT('stats-hourly-user-model:', aggregated.user_id, ':', aggregated.model, ':', CAST($1 AS TEXT))),
$1,
aggregated.user_id,
aggregated.model,
aggregated.total_requests,
aggregated.input_tokens,
aggregated.output_tokens,
aggregated.total_cost,
aggregated.response_time_sum_ms,
aggregated.response_time_samples,
$3,
$3
FROM aggregated
ON CONFLICT (hour_utc, user_id, model)
DO UPDATE SET
total_requests = EXCLUDED.total_requests,
input_tokens = EXCLUDED.input_tokens,
output_tokens = EXCLUDED.output_tokens,
total_cost = EXCLUDED.total_cost,
response_time_sum_ms = EXCLUDED.response_time_sum_ms,
response_time_samples = EXCLUDED.response_time_samples,
updated_at = EXCLUDED.updated_at
"#;
pub(super) const UPSERT_STATS_HOURLY_PROVIDER_SQL: &str = r#"
WITH aggregated AS (
SELECT
provider_name,
CAST(COUNT(id) AS BIGINT) AS total_requests,
CAST(COALESCE(SUM(input_tokens), 0) AS BIGINT) AS input_tokens,
CAST(COALESCE(SUM(output_tokens), 0) AS BIGINT) AS output_tokens,
CAST(COALESCE(SUM(total_cost_usd), 0) AS DOUBLE PRECISION) AS total_cost
FROM usage_billing_facts AS usage
WHERE created_at >= $1
AND created_at < $2
AND provider_name IS NOT NULL
AND provider_name <> ''
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
GROUP BY provider_name
)
INSERT INTO stats_hourly_provider (
id,
hour_utc,
provider_name,
total_requests,
input_tokens,
output_tokens,
total_cost,
created_at,
updated_at
)
SELECT
md5(CONCAT('stats-hourly-provider:', aggregated.provider_name, ':', CAST($1 AS TEXT))),
$1,
aggregated.provider_name,
aggregated.total_requests,
aggregated.input_tokens,
aggregated.output_tokens,
aggregated.total_cost,
$3,
$3
FROM aggregated
ON CONFLICT (hour_utc, provider_name)
DO UPDATE SET
total_requests = EXCLUDED.total_requests,
input_tokens = EXCLUDED.input_tokens,
output_tokens = EXCLUDED.output_tokens,
total_cost = EXCLUDED.total_cost,
updated_at = EXCLUDED.updated_at
"#;
#[cfg(test)]
mod tests {
use super::SELECT_STATS_HOURLY_AGGREGATE_SQL;
#[test]
fn hourly_aggregate_uses_one_time_bounded_base_table_scan() {
let sql = SELECT_STATS_HOURLY_AGGREGATE_SQL;
assert_eq!(sql.matches("FROM usage_billing_facts AS usage").count(), 1);
assert_eq!(sql.matches("WHERE usage.created_at >= $1").count(), 1);
assert_eq!(sql.matches("AND usage.created_at < $2").count(), 1);
assert!(!sql.contains("MATERIALIZED"));
assert!(!sql.contains("bucket_usage"));
}
#[test]
fn hourly_aggregate_keeps_independent_row_populations() {
let sql = SELECT_STATS_HOURLY_AGGREGATE_SQL;
assert_eq!(sql.matches("usage.status = 'completed'").count(), 8);
assert_eq!(sql.matches("usage.billing_status = 'settled'").count(), 8);
assert_eq!(
sql.matches("COALESCE(CAST(usage.total_cost_usd AS DOUBLE PRECISION), 0) > 0")
.count(),
8
);
assert_eq!(
sql.matches("usage.status NOT IN ('pending', 'streaming')")
.count(),
11
);
assert_eq!(
sql.matches("usage.provider_name NOT IN ('unknown', 'pending')")
.count(),
11
);
let first_aggregate = sql
.trim_start()
.strip_prefix("SELECT\n")
.expect("hourly aggregate should remain a direct SELECT");
assert!(first_aggregate
.starts_with(" CAST(COUNT(usage.id) AS BIGINT) AS cache_hit_total_requests,"));
assert!(
sql.contains("MIN(CAST(EXTRACT(EPOCH FROM usage.finalized_at) AS BIGINT)) FILTER (")
);
assert!(
sql.contains("MAX(CAST(EXTRACT(EPOCH FROM usage.finalized_at) AS BIGINT)) FILTER (")
);
}
#[test]
fn hourly_aggregate_preserves_the_decoder_column_contract() {
let sql = SELECT_STATS_HOURLY_AGGREGATE_SQL;
let aliases = [
"cache_hit_total_requests",
"cache_hit_requests",
"completed_total_requests",
"completed_cache_hit_requests",
"completed_input_tokens",
"completed_cache_creation_tokens",
"completed_cache_read_tokens",
"completed_total_input_context",
"completed_cache_creation_cost",
"completed_cache_read_cost",
"settled_total_cost",
"settled_total_requests",
"settled_input_tokens",
"settled_output_tokens",
"settled_cache_creation_tokens",
"settled_cache_read_tokens",
"settled_first_finalized_at_unix_secs",
"settled_last_finalized_at_unix_secs",
"total_requests",
"error_requests",
"input_tokens",
"output_tokens",
"cache_creation_tokens",
"cache_read_tokens",
"total_cost",
"actual_total_cost",
"response_time_sum_ms",
"response_time_samples",
"avg_response_time_ms",
];
for alias in aliases {
assert_eq!(
sql.matches(&format!(" AS {alias}")).count(),
1,
"hourly aggregate alias {alias} must appear exactly once"
);
}
}
}
@@ -0,0 +1,373 @@
use chrono::{DateTime, Utc};
use sqlx::Row;
use crate::backend::stats_common::{stats_id, unix_ms, unix_secs, utc_from_unix_secs};
use crate::backend::SqliteBackend;
use crate::driver::sqlite::{sqlite_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::{
DataLayerError, StatsDailyAggregationInput, StatsDailyAggregationSummary,
StatsHourlyAggregationInput, StatsHourlyAggregationSummary,
};
impl SqliteBackend {
pub async fn aggregate_stats_hourly(
&self,
input: &StatsHourlyAggregationInput,
) -> Result<Option<StatsHourlyAggregationSummary>, DataLayerError> {
let Some(hour_utc_unix_secs) =
next_sqlite_stats_hourly_bucket(self.pool(), input.target_hour_utc).await?
else {
return Ok(None);
};
perform_sqlite_stats_hourly_aggregation(
self.pool(),
hour_utc_unix_secs,
input.aggregated_at,
)
.await
.map(Some)
}
pub async fn aggregate_stats_daily(
&self,
input: &StatsDailyAggregationInput,
) -> Result<Option<StatsDailyAggregationSummary>, DataLayerError> {
let Some(day_start_unix_secs) =
next_sqlite_stats_daily_bucket(self.pool(), input.target_day_utc).await?
else {
return Ok(None);
};
perform_sqlite_stats_daily_aggregation(
self.pool(),
day_start_unix_secs,
input.aggregated_at,
)
.await
.map(Some)
}
}
async fn next_sqlite_stats_hourly_bucket(
pool: &SqlitePool,
target_hour_utc: DateTime<Utc>,
) -> Result<Option<i64>, DataLayerError> {
let latest_hour: Option<i64> =
sqlx::query_scalar("SELECT MAX(hour_utc) FROM stats_hourly WHERE is_complete <> 0")
.fetch_one(pool)
.await
.map_sql_err()?;
let search_from = latest_hour.map(|value| value + 3600).unwrap_or(0);
let search_until = unix_secs(target_hour_utc) + 3600;
if search_from >= search_until {
return Ok(None);
}
let next_bucket: Option<i64> = sqlx::query_scalar(
r#"
SELECT MIN(CAST(created_at_unix_ms / 3600000 AS INTEGER) * 3600)
FROM "usage"
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
"#,
)
.bind(unix_ms(search_from)?)
.bind(unix_ms(search_until)?)
.fetch_one(pool)
.await
.map_sql_err()?;
Ok(next_bucket.filter(|value| *value <= unix_secs(target_hour_utc)))
}
async fn next_sqlite_stats_daily_bucket(
pool: &SqlitePool,
target_day_utc: DateTime<Utc>,
) -> Result<Option<i64>, DataLayerError> {
let latest_day: Option<i64> =
sqlx::query_scalar(r#"SELECT MAX("date") FROM stats_daily WHERE is_complete <> 0"#)
.fetch_one(pool)
.await
.map_sql_err()?;
let search_from = latest_day.map(|value| value + 86_400).unwrap_or(0);
let search_until = unix_secs(target_day_utc) + 86_400;
if search_from >= search_until {
return Ok(None);
}
let next_bucket: Option<i64> = sqlx::query_scalar(
r#"
SELECT MIN(CAST(created_at_unix_ms / 86400000 AS INTEGER) * 86400)
FROM "usage"
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
"#,
)
.bind(unix_ms(search_from)?)
.bind(unix_ms(search_until)?)
.fetch_one(pool)
.await
.map_sql_err()?;
Ok(next_bucket.filter(|value| *value <= unix_secs(target_day_utc)))
}
const SQLITE_STATS_AGGREGATE_SQL: &str = r#"
SELECT
COUNT(*) AS total_requests,
COALESCE(SUM(CASE
WHEN status = 'failed'
OR status_code >= 400
OR (error_category IS NOT NULL AND error_category <> '')
THEN 1 ELSE 0 END), 0) AS error_requests,
COALESCE(SUM(input_tokens), 0) AS input_tokens,
COALESCE(SUM(output_tokens), 0) AS output_tokens,
COALESCE(SUM(cache_creation_input_tokens), 0) AS cache_creation_tokens,
COALESCE(SUM(cache_read_input_tokens), 0) AS cache_read_tokens,
CAST(COALESCE(SUM(total_cost_usd), 0) AS REAL) AS total_cost,
CAST(COALESCE(SUM(actual_total_cost_usd), 0) AS REAL) AS actual_total_cost,
CAST(COALESCE(AVG(response_time_ms), 0) AS REAL) AS avg_response_time_ms
FROM "usage"
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
"#;
async fn perform_sqlite_stats_hourly_aggregation(
pool: &SqlitePool,
hour_utc_unix_secs: i64,
aggregated_at: DateTime<Utc>,
) -> Result<StatsHourlyAggregationSummary, DataLayerError> {
let start_ms = unix_ms(hour_utc_unix_secs)?;
let end_ms = unix_ms(hour_utc_unix_secs + 3600)?;
let aggregated_at_unix_secs = unix_secs(aggregated_at);
let mut tx = pool.begin().await.map_sql_err()?;
let row = sqlx::query(SQLITE_STATS_AGGREGATE_SQL)
.bind(start_ms)
.bind(end_ms)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let total_requests: i64 = row.try_get("total_requests").map_sql_err()?;
let error_requests: i64 = row.try_get("error_requests").map_sql_err()?;
sqlx::query(
r#"
INSERT INTO stats_hourly (
id, hour_utc, total_requests, success_requests, error_requests,
input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens,
total_cost, actual_total_cost, avg_response_time_ms, is_complete,
aggregated_at, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 1, ?, ?, ?)
ON CONFLICT (hour_utc) DO UPDATE SET
total_requests = excluded.total_requests,
success_requests = excluded.success_requests,
error_requests = excluded.error_requests,
input_tokens = excluded.input_tokens,
output_tokens = excluded.output_tokens,
cache_creation_tokens = excluded.cache_creation_tokens,
cache_read_tokens = excluded.cache_read_tokens,
total_cost = excluded.total_cost,
actual_total_cost = excluded.actual_total_cost,
avg_response_time_ms = excluded.avg_response_time_ms,
is_complete = excluded.is_complete,
aggregated_at = excluded.aggregated_at,
updated_at = excluded.updated_at
"#,
)
.bind(stats_id(&format!("stats-hourly:{hour_utc_unix_secs}")))
.bind(hour_utc_unix_secs)
.bind(total_requests)
.bind(total_requests.saturating_sub(error_requests))
.bind(error_requests)
.bind(row.try_get::<i64, _>("input_tokens").map_sql_err()?)
.bind(row.try_get::<i64, _>("output_tokens").map_sql_err()?)
.bind(
row.try_get::<i64, _>("cache_creation_tokens")
.map_sql_err()?,
)
.bind(row.try_get::<i64, _>("cache_read_tokens").map_sql_err()?)
.bind(sqlite_real(&row, "total_cost")?)
.bind(sqlite_real(&row, "actual_total_cost")?)
.bind(sqlite_real(&row, "avg_response_time_ms")?)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
.execute(&mut *tx)
.await
.map_sql_err()?;
let user_rows = sqlite_group_count(&mut tx, "user_id", start_ms, end_ms).await?;
let user_model_rows = sqlite_group_count(&mut tx, "user_id, model", start_ms, end_ms).await?;
let model_rows = sqlite_group_count(&mut tx, "model", start_ms, end_ms).await?;
let provider_rows = sqlite_group_count(&mut tx, "provider_name", start_ms, end_ms).await?;
tx.commit().await.map_sql_err()?;
Ok(StatsHourlyAggregationSummary {
hour_utc: utc_from_unix_secs(hour_utc_unix_secs, "stats_hourly.hour_utc")?,
total_requests,
user_rows,
user_model_rows,
model_rows,
provider_rows,
})
}
async fn perform_sqlite_stats_daily_aggregation(
pool: &SqlitePool,
day_start_unix_secs: i64,
aggregated_at: DateTime<Utc>,
) -> Result<StatsDailyAggregationSummary, DataLayerError> {
let start_ms = unix_ms(day_start_unix_secs)?;
let end_ms = unix_ms(day_start_unix_secs + 86_400)?;
let aggregated_at_unix_secs = unix_secs(aggregated_at);
let mut tx = pool.begin().await.map_sql_err()?;
let row = sqlx::query(SQLITE_STATS_AGGREGATE_SQL)
.bind(start_ms)
.bind(end_ms)
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let total_requests: i64 = row.try_get("total_requests").map_sql_err()?;
let error_requests: i64 = row.try_get("error_requests").map_sql_err()?;
let unique_models = sqlite_group_count(&mut tx, "model", start_ms, end_ms).await? as i64;
let unique_providers =
sqlite_group_count(&mut tx, "provider_name", start_ms, end_ms).await? as i64;
sqlx::query(
r#"
INSERT INTO stats_daily (
id, "date", total_requests, success_requests, error_requests,
input_tokens, output_tokens, cache_creation_tokens, cache_read_tokens,
total_cost, actual_total_cost, avg_response_time_ms, fallback_count,
unique_models, unique_providers, is_complete, aggregated_at, created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 0, ?, ?, 1, ?, ?, ?)
ON CONFLICT ("date") DO UPDATE SET
total_requests = excluded.total_requests,
success_requests = excluded.success_requests,
error_requests = excluded.error_requests,
input_tokens = excluded.input_tokens,
output_tokens = excluded.output_tokens,
cache_creation_tokens = excluded.cache_creation_tokens,
cache_read_tokens = excluded.cache_read_tokens,
total_cost = excluded.total_cost,
actual_total_cost = excluded.actual_total_cost,
avg_response_time_ms = excluded.avg_response_time_ms,
fallback_count = excluded.fallback_count,
unique_models = excluded.unique_models,
unique_providers = excluded.unique_providers,
is_complete = excluded.is_complete,
aggregated_at = excluded.aggregated_at,
updated_at = excluded.updated_at
"#,
)
.bind(stats_id(&format!("stats-daily:{day_start_unix_secs}")))
.bind(day_start_unix_secs)
.bind(total_requests)
.bind(total_requests.saturating_sub(error_requests))
.bind(error_requests)
.bind(row.try_get::<i64, _>("input_tokens").map_sql_err()?)
.bind(row.try_get::<i64, _>("output_tokens").map_sql_err()?)
.bind(
row.try_get::<i64, _>("cache_creation_tokens")
.map_sql_err()?,
)
.bind(row.try_get::<i64, _>("cache_read_tokens").map_sql_err()?)
.bind(sqlite_real(&row, "total_cost")?)
.bind(sqlite_real(&row, "actual_total_cost")?)
.bind(sqlite_real(&row, "avg_response_time_ms")?)
.bind(unique_models)
.bind(unique_providers)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
.execute(&mut *tx)
.await
.map_sql_err()?;
let model_rows = usize::try_from(unique_models).unwrap_or(usize::MAX);
let provider_rows = usize::try_from(unique_providers).unwrap_or(usize::MAX);
let api_key_rows = sqlite_group_count(&mut tx, "api_key_id", start_ms, end_ms).await?;
let error_rows = sqlite_error_group_count(&mut tx, start_ms, end_ms).await?;
let user_rows = sqlite_group_count(&mut tx, "user_id", start_ms, end_ms).await?;
tx.commit().await.map_sql_err()?;
Ok(StatsDailyAggregationSummary {
day_start_utc: utc_from_unix_secs(day_start_unix_secs, "stats_daily.date")?,
total_requests,
model_rows,
provider_rows,
api_key_rows,
error_rows,
user_rows,
})
}
async fn sqlite_group_count(
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
group_columns: &str,
start_ms: i64,
end_ms: i64,
) -> Result<usize, DataLayerError> {
let not_empty = group_columns
.split(',')
.map(str::trim)
.map(|column| format!("{column} IS NOT NULL AND {column} <> ''"))
.collect::<Vec<_>>()
.join(" AND ");
let sql = format!(
r#"
SELECT COUNT(*)
FROM (
SELECT 1
FROM "usage"
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
AND {not_empty}
GROUP BY {group_columns}
)
"#
);
let count: i64 = sqlx::query_scalar(&sql)
.bind(start_ms)
.bind(end_ms)
.fetch_one(&mut **tx)
.await
.map_sql_err()?;
Ok(usize::try_from(count.max(0)).unwrap_or(usize::MAX))
}
async fn sqlite_error_group_count(
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
start_ms: i64,
end_ms: i64,
) -> Result<usize, DataLayerError> {
let count: i64 = sqlx::query_scalar(
r#"
SELECT COUNT(*)
FROM (
SELECT 1
FROM "usage"
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
AND status NOT IN ('pending', 'streaming')
AND provider_name NOT IN ('unknown', 'pending')
AND (
status = 'failed'
OR status_code >= 400
OR (error_category IS NOT NULL AND error_category <> '')
)
GROUP BY COALESCE(NULLIF(error_category, ''), 'unknown_error'), provider_name, model
)
"#,
)
.bind(start_ms)
.bind(end_ms)
.fetch_one(&mut **tx)
.await
.map_sql_err()?;
Ok(usize::try_from(count.max(0)).unwrap_or(usize::MAX))
}
@@ -0,0 +1,33 @@
use chrono::{DateTime, Utc};
use sha2::{Digest, Sha256};
use crate::DataLayerError;
pub(crate) fn unix_secs(value: DateTime<Utc>) -> i64 {
value.timestamp().max(0)
}
pub(crate) fn unix_ms(value: i64) -> Result<i64, DataLayerError> {
value.checked_mul(1000).ok_or_else(|| {
DataLayerError::InvalidInput(format!("timestamp overflow while converting {value} to ms"))
})
}
pub(crate) fn utc_from_unix_secs(
value: i64,
field_name: &str,
) -> Result<DateTime<Utc>, DataLayerError> {
DateTime::<Utc>::from_timestamp(value, 0).ok_or_else(|| {
DataLayerError::UnexpectedValue(format!("{field_name} contains invalid timestamp {value}"))
})
}
pub(crate) fn stats_id(value: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(value.as_bytes());
hasher
.finalize()
.iter()
.map(|byte| format!("{byte:02x}"))
.collect()
}
@@ -0,0 +1,189 @@
use std::collections::BTreeMap;
#[cfg(feature = "mysql")]
use super::MysqlBackend;
#[cfg(feature = "postgres")]
use super::PostgresBackend;
#[cfg(feature = "sqlite")]
use super::SqliteBackend;
use crate::repository::system::{
AdminSystemPurgeSummary, AdminSystemPurgeTarget, AdminSystemUsageAggregateImportMode,
AdminSystemUsageAggregateImportSummary, AdminSystemUsageAggregateSnapshot,
StoredSystemConfigEntry,
};
use crate::DataLayerError;
#[cfg(feature = "mysql")]
mod mysql;
#[cfg(feature = "postgres")]
mod postgres;
#[cfg(feature = "sqlite")]
mod sqlite;
const ADMIN_CONFIG_PURGE_TABLES: &[&str] = &[
"api_key_provider_mappings",
"gemini_file_mappings",
"provider_usage_tracking",
"billing_rules",
"dimension_collectors",
"models",
"provider_endpoints",
"provider_api_keys",
"providers",
"global_models",
"proxy_node_events",
"proxy_nodes",
"user_oauth_links",
"ldap_configs",
"oauth_providers",
"auth_modules",
"system_configs",
];
const ADMIN_STATS_PURGE_TABLES: &[&str] = &[
"stats_user_daily_cost_savings_model_provider",
"stats_user_daily_cost_savings_model",
"stats_user_daily_cost_savings_provider",
"stats_user_daily_cost_savings",
"stats_daily_cost_savings_model_provider",
"stats_daily_cost_savings_model",
"stats_daily_cost_savings_provider",
"stats_daily_cost_savings",
"stats_user_daily_model_provider",
"stats_daily_model_provider",
"stats_user_daily_api_format",
"stats_user_daily_provider",
"stats_hourly_user_model",
"stats_user_daily_model",
"stats_user_summary",
"stats_hourly_user",
"stats_user_daily",
"stats_daily_api_key",
"stats_daily_error",
"stats_daily_model",
"stats_daily_provider",
"stats_hourly_model",
"stats_hourly_provider",
"stats_summary",
"stats_hourly",
"stats_daily",
];
const ADMIN_USAGE_CHILD_TABLES: &[&str] = &[
"usage_body_blobs",
"usage_http_audits",
"usage_routing_snapshots",
"usage_settlement_snapshots",
];
const USAGE_BODY_FIELD_COLUMNS: &[&str] = &[
"request_body",
"response_body",
"provider_request_body",
"client_response_body",
"request_body_compressed",
"response_body_compressed",
"provider_request_body_compressed",
"client_response_body_compressed",
];
const ADMIN_USER_SCOPED_TABLES: &[&str] = &[
"stats_user_daily_cost_savings_model_provider",
"stats_user_daily_cost_savings_model",
"stats_user_daily_cost_savings_provider",
"stats_user_daily_cost_savings",
"stats_user_daily_model_provider",
"stats_user_daily_api_format",
"stats_user_daily_provider",
"stats_hourly_user_model",
"stats_user_daily_model",
"stats_user_summary",
"stats_hourly_user",
"stats_user_daily",
"user_model_usage_counts",
"announcement_reads",
"management_tokens",
"user_preferences",
"user_sessions",
"user_oauth_links",
];
fn checked_sql_identifier(value: &str) -> Result<&str, DataLayerError> {
if !value.is_empty()
&& value
.bytes()
.all(|byte| byte.is_ascii_alphanumeric() || byte == b'_')
{
Ok(value)
} else {
Err(DataLayerError::InvalidInput(format!(
"invalid SQL identifier: {value}"
)))
}
}
#[cfg(any(feature = "mysql", feature = "sqlite"))]
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 range")))
}
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 u64_from_i64(value: i64) -> u64 {
value.max(0) as u64
}
fn add_aggregate_import_count(
counter: &mut crate::repository::system::AdminSystemUsageAggregateImportCounter,
existed: bool,
) {
if existed {
counter.updated += 1;
} else {
counter.created += 1;
}
}
fn should_skip_imported_aggregate(
exists: bool,
mode: AdminSystemUsageAggregateImportMode,
table: &str,
date_unix_secs: u64,
) -> Result<bool, DataLayerError> {
if !exists {
return Ok(false);
}
match mode {
AdminSystemUsageAggregateImportMode::Skip => Ok(true),
AdminSystemUsageAggregateImportMode::Overwrite => Ok(false),
AdminSystemUsageAggregateImportMode::Error => Err(DataLayerError::InvalidInput(format!(
"{table} aggregate already exists for date_unix_secs={date_unix_secs}"
))),
}
}
#[cfg(any(feature = "mysql", feature = "sqlite"))]
fn serialize_json_value(value: &serde_json::Value) -> Result<String, DataLayerError> {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("invalid system config JSON value: {err}"))
})
}
#[cfg(any(feature = "mysql", feature = "sqlite"))]
fn parse_json_value(value: String) -> Result<serde_json::Value, DataLayerError> {
serde_json::from_str(&value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("invalid system config JSON value: {err}"))
})
}
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,74 @@
use std::fmt;
#[cfg(feature = "postgres")]
use super::PostgresBackend;
#[cfg(feature = "postgres")]
use crate::driver::postgres::PostgresTransactionRunner;
#[derive(Clone, Default)]
pub struct DataTransactionBackends {
#[cfg(feature = "postgres")]
postgres: Option<PostgresTransactionRunner>,
}
impl fmt::Debug for DataTransactionBackends {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DataTransactionBackends")
.field("has_postgres", &self.has_any())
.finish()
}
}
impl DataTransactionBackends {
#[cfg(feature = "postgres")]
pub(crate) fn from_postgres(postgres: Option<&PostgresBackend>) -> Self {
Self {
postgres: postgres.map(PostgresBackend::transaction_runner),
}
}
#[cfg(feature = "postgres")]
pub fn postgres(&self) -> Option<PostgresTransactionRunner> {
self.postgres.clone()
}
pub fn has_any(&self) -> bool {
cfg!(feature = "postgres") && {
#[cfg(feature = "postgres")]
{
self.postgres.is_some()
}
#[cfg(not(feature = "postgres"))]
{
false
}
}
}
}
#[cfg(all(test, feature = "postgres"))]
mod tests {
use super::DataTransactionBackends;
use crate::backend::PostgresBackend;
use crate::driver::postgres::PostgresPoolConfig;
#[tokio::test]
async fn builds_postgres_transaction_runner_from_backend() {
let backend = PostgresBackend::from_config(PostgresPoolConfig {
database_url: "postgres://localhost/aether".to_string(),
min_connections: 1,
max_connections: 4,
acquire_timeout_ms: 1_000,
idle_timeout_ms: 5_000,
max_lifetime_ms: 30_000,
statement_cache_capacity: 64,
require_ssl: false,
})
.expect("postgres backend should build");
let transactions = DataTransactionBackends::from_postgres(Some(&backend));
assert!(transactions.has_any());
assert!(transactions.postgres().is_some());
}
}
@@ -0,0 +1,74 @@
//! Driver-specific wallet usage aggregation adapters.
#[cfg(feature = "mysql")]
mod mysql;
#[cfg(feature = "postgres")]
mod postgres;
#[cfg(feature = "sqlite")]
mod sqlite;
#[cfg(any(feature = "mysql", feature = "sqlite"))]
use sha2::{Digest, Sha256};
use crate::DataLayerError;
pub(super) fn u64_to_i64(value: u64, field_name: &str) -> Result<i64, DataLayerError> {
i64::try_from(value)
.map_err(|_| DataLayerError::InvalidInput(format!("invalid {field_name}: {value}")))
}
#[cfg(feature = "postgres")]
pub(super) fn unix_secs_to_utc(
value: u64,
field_name: &str,
) -> Result<chrono::DateTime<chrono::Utc>, DataLayerError> {
let value = u64_to_i64(value, field_name)?;
chrono::DateTime::<chrono::Utc>::from_timestamp(value, 0)
.ok_or_else(|| DataLayerError::InvalidInput(format!("invalid {field_name}: {value}")))
}
#[cfg(any(feature = "mysql", feature = "sqlite"))]
pub(super) fn wallet_daily_usage_id(
wallet_id: &str,
billing_date: &str,
billing_timezone: &str,
) -> String {
let mut hasher = Sha256::new();
hasher.update(b"wallet-daily-usage:");
hasher.update(wallet_id.as_bytes());
hasher.update(b":");
hasher.update(billing_date.as_bytes());
hasher.update(b":");
hasher.update(billing_timezone.as_bytes());
hasher
.finalize()
.iter()
.map(|byte| format!("{byte:02x}"))
.collect()
}
#[cfg(test)]
mod tests {
use super::u64_to_i64;
#[cfg(any(feature = "mysql", feature = "sqlite"))]
use super::wallet_daily_usage_id;
#[cfg(any(feature = "mysql", feature = "sqlite"))]
#[test]
fn wallet_daily_usage_ids_are_stable_and_partition_specific() {
let first = wallet_daily_usage_id("wallet-1", "2026-07-13", "UTC");
let same = wallet_daily_usage_id("wallet-1", "2026-07-13", "UTC");
let other_day = wallet_daily_usage_id("wallet-1", "2026-07-14", "UTC");
assert_eq!(first, same);
assert_ne!(first, other_day);
assert_eq!(first.len(), 64);
}
#[test]
fn rejects_timestamps_outside_i64_range() {
if usize::BITS >= 64 {
assert!(u64_to_i64(u64::MAX, "window_start").is_err());
}
}
}
@@ -0,0 +1,155 @@
use sqlx::Row;
use crate::backend::MysqlBackend;
use crate::error::SqlResultExt;
use crate::{DataLayerError, WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult};
use super::{u64_to_i64, wallet_daily_usage_id};
const SELECT_WALLET_DAILY_USAGE_AGGREGATES_SQL: &str = r#"
SELECT
usage_settlement_snapshots.wallet_id AS wallet_id,
CAST(COUNT(*) AS SIGNED) AS total_requests,
CAST(COALESCE(SUM(`usage`.total_cost_usd), 0) AS DOUBLE) AS total_cost_usd,
CAST(COALESCE(SUM(`usage`.input_tokens), 0) AS SIGNED) AS input_tokens,
CAST(COALESCE(SUM(`usage`.output_tokens), 0) AS SIGNED) AS output_tokens,
CAST(COALESCE(SUM(`usage`.cache_creation_input_tokens), 0) AS SIGNED) AS cache_creation_tokens,
CAST(COALESCE(SUM(`usage`.cache_read_input_tokens), 0) AS SIGNED) AS cache_read_tokens,
MIN(COALESCE(usage_settlement_snapshots.finalized_at, `usage`.finalized_at)) AS first_finalized_at,
MAX(COALESCE(usage_settlement_snapshots.finalized_at, `usage`.finalized_at)) AS last_finalized_at
FROM `usage`
JOIN usage_settlement_snapshots
ON usage_settlement_snapshots.request_id = `usage`.request_id
WHERE usage_settlement_snapshots.wallet_id IS NOT NULL
AND usage_settlement_snapshots.wallet_id <> ''
AND COALESCE(usage_settlement_snapshots.billing_status, `usage`.billing_status) = 'settled'
AND `usage`.total_cost_usd > 0
AND COALESCE(usage_settlement_snapshots.finalized_at, `usage`.finalized_at) >= ?
AND COALESCE(usage_settlement_snapshots.finalized_at, `usage`.finalized_at) < ?
GROUP BY usage_settlement_snapshots.wallet_id
"#;
impl MysqlBackend {
pub async fn aggregate_wallet_daily_usage(
&self,
input: &WalletDailyUsageAggregationInput,
) -> Result<WalletDailyUsageAggregationResult, DataLayerError> {
let window_start = u64_to_i64(input.window_start_unix_secs, "window_start")?;
let window_end = u64_to_i64(input.window_end_unix_secs, "window_end")?;
let aggregated_at = u64_to_i64(input.aggregated_at_unix_secs, "aggregated_at")?;
let mut tx = self.pool().begin().await.map_sql_err()?;
let rows = sqlx::query(SELECT_WALLET_DAILY_USAGE_AGGREGATES_SQL)
.bind(window_start)
.bind(window_end)
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
let mut aggregated_wallets = 0usize;
for row in rows {
let wallet_id: String = row.try_get("wallet_id").map_sql_err()?;
sqlx::query(
r#"
DELETE FROM wallet_daily_usage_ledgers
WHERE wallet_id = ?
AND billing_date = ?
AND billing_timezone = ?
"#,
)
.bind(&wallet_id)
.bind(&input.billing_date)
.bind(&input.billing_timezone)
.execute(&mut *tx)
.await
.map_sql_err()?;
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,
first_finalized_at,
last_finalized_at,
aggregated_at,
created_at,
updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(wallet_daily_usage_id(
&wallet_id,
&input.billing_date,
&input.billing_timezone,
))
.bind(&wallet_id)
.bind(&input.billing_date)
.bind(&input.billing_timezone)
.bind(row.try_get::<f64, _>("total_cost_usd").map_sql_err()?)
.bind(row.try_get::<i64, _>("total_requests").map_sql_err()?)
.bind(row.try_get::<i64, _>("input_tokens").map_sql_err()?)
.bind(row.try_get::<i64, _>("output_tokens").map_sql_err()?)
.bind(
row.try_get::<i64, _>("cache_creation_tokens")
.map_sql_err()?,
)
.bind(row.try_get::<i64, _>("cache_read_tokens").map_sql_err()?)
.bind(
row.try_get::<Option<i64>, _>("first_finalized_at")
.map_sql_err()?,
)
.bind(
row.try_get::<Option<i64>, _>("last_finalized_at")
.map_sql_err()?,
)
.bind(aggregated_at)
.bind(aggregated_at)
.bind(aggregated_at)
.execute(&mut *tx)
.await
.map_sql_err()?;
aggregated_wallets += 1;
}
let deleted_stale_ledgers = sqlx::query(
r#"
DELETE FROM wallet_daily_usage_ledgers
WHERE billing_date = ?
AND billing_timezone = ?
AND NOT EXISTS (
SELECT 1
FROM `usage`
JOIN usage_settlement_snapshots
ON usage_settlement_snapshots.request_id = `usage`.request_id
WHERE usage_settlement_snapshots.wallet_id = wallet_daily_usage_ledgers.wallet_id
AND COALESCE(usage_settlement_snapshots.billing_status, `usage`.billing_status) = 'settled'
AND `usage`.total_cost_usd > 0
AND COALESCE(usage_settlement_snapshots.finalized_at, `usage`.finalized_at) >= ?
AND COALESCE(usage_settlement_snapshots.finalized_at, `usage`.finalized_at) < ?
)
"#,
)
.bind(&input.billing_date)
.bind(&input.billing_timezone)
.bind(window_start)
.bind(window_end)
.execute(&mut *tx)
.await
.map_sql_err()?
.rows_affected();
tx.commit().await.map_sql_err()?;
Ok(WalletDailyUsageAggregationResult {
aggregated_wallets,
deleted_stale_ledgers: usize::try_from(deleted_stale_ledgers).unwrap_or(usize::MAX),
})
}
}
@@ -0,0 +1,135 @@
use crate::backend::PostgresBackend;
use crate::error::SqlxResultExt;
use crate::{DataLayerError, WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult};
use super::unix_secs_to_utc;
const UPSERT_WALLET_DAILY_USAGE_LEDGER_SQL: &str = r#"
WITH aggregated AS (
SELECT
usage_settlement_snapshots.wallet_id,
COUNT(*) AS total_requests,
CAST(COALESCE(SUM(usage.total_cost_usd), 0) AS DOUBLE PRECISION) AS total_cost_usd,
COALESCE(SUM(usage.input_tokens), 0) AS input_tokens,
COALESCE(SUM(usage.output_tokens), 0) AS output_tokens,
COALESCE(SUM(usage.cache_creation_input_tokens), 0) AS cache_creation_tokens,
COALESCE(SUM(usage.cache_read_input_tokens), 0) AS cache_read_tokens,
MIN(COALESCE(usage_settlement_snapshots.finalized_at, usage.finalized_at)) AS first_finalized_at,
MAX(COALESCE(usage_settlement_snapshots.finalized_at, usage.finalized_at)) AS last_finalized_at
FROM usage_billing_facts AS usage
JOIN usage_settlement_snapshots
ON usage_settlement_snapshots.request_id = usage.request_id
WHERE usage_settlement_snapshots.wallet_id IS NOT NULL
AND COALESCE(usage_settlement_snapshots.billing_status, usage.billing_status) = 'settled'
AND usage.total_cost_usd > 0
AND COALESCE(usage_settlement_snapshots.finalized_at, usage.finalized_at) >= $1
AND COALESCE(usage_settlement_snapshots.finalized_at, usage.finalized_at) < $2
GROUP BY usage_settlement_snapshots.wallet_id
)
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,
first_finalized_at,
last_finalized_at,
aggregated_at,
created_at,
updated_at
)
SELECT
md5(CONCAT('wallet-daily-usage:', aggregated.wallet_id, ':', CAST($3 AS TEXT), ':', $4)),
aggregated.wallet_id,
$3,
$4,
aggregated.total_cost_usd,
aggregated.total_requests,
aggregated.input_tokens,
aggregated.output_tokens,
aggregated.cache_creation_tokens,
aggregated.cache_read_tokens,
aggregated.first_finalized_at,
aggregated.last_finalized_at,
$5,
$5,
$5
FROM aggregated
ON CONFLICT (wallet_id, billing_date, billing_timezone)
DO UPDATE SET
total_cost_usd = EXCLUDED.total_cost_usd,
total_requests = EXCLUDED.total_requests,
input_tokens = EXCLUDED.input_tokens,
output_tokens = EXCLUDED.output_tokens,
cache_creation_tokens = EXCLUDED.cache_creation_tokens,
cache_read_tokens = EXCLUDED.cache_read_tokens,
first_finalized_at = EXCLUDED.first_finalized_at,
last_finalized_at = EXCLUDED.last_finalized_at,
aggregated_at = EXCLUDED.aggregated_at,
updated_at = EXCLUDED.updated_at
"#;
const DELETE_STALE_WALLET_DAILY_USAGE_LEDGERS_SQL: &str = r#"
DELETE FROM wallet_daily_usage_ledgers AS ledgers
WHERE ledgers.billing_date = $1
AND ledgers.billing_timezone = $2
AND NOT EXISTS (
SELECT 1
FROM usage_billing_facts AS usage
JOIN usage_settlement_snapshots
ON usage_settlement_snapshots.request_id = usage.request_id
WHERE usage_settlement_snapshots.wallet_id = ledgers.wallet_id
AND COALESCE(usage_settlement_snapshots.billing_status, usage.billing_status) = 'settled'
AND usage.total_cost_usd > 0
AND COALESCE(usage_settlement_snapshots.finalized_at, usage.finalized_at) >= $3
AND COALESCE(usage_settlement_snapshots.finalized_at, usage.finalized_at) < $4
)
"#;
impl PostgresBackend {
pub async fn aggregate_wallet_daily_usage(
&self,
input: &WalletDailyUsageAggregationInput,
) -> Result<WalletDailyUsageAggregationResult, DataLayerError> {
let billing_date = chrono::NaiveDate::parse_from_str(&input.billing_date, "%Y-%m-%d")
.map_err(|err| {
DataLayerError::InvalidInput(format!("invalid wallet billing_date: {err}"))
})?;
let window_start = unix_secs_to_utc(input.window_start_unix_secs, "window_start")?;
let window_end = unix_secs_to_utc(input.window_end_unix_secs, "window_end")?;
let aggregated_at = unix_secs_to_utc(input.aggregated_at_unix_secs, "aggregated_at")?;
let mut tx = self.pool().begin().await.map_postgres_err()?;
let aggregated_wallets = sqlx::query(UPSERT_WALLET_DAILY_USAGE_LEDGER_SQL)
.bind(window_start)
.bind(window_end)
.bind(billing_date)
.bind(input.billing_timezone.as_str())
.bind(aggregated_at)
.execute(&mut *tx)
.await
.map_postgres_err()?
.rows_affected();
let deleted_stale_ledgers = sqlx::query(DELETE_STALE_WALLET_DAILY_USAGE_LEDGERS_SQL)
.bind(billing_date)
.bind(input.billing_timezone.as_str())
.bind(window_start)
.bind(window_end)
.execute(&mut *tx)
.await
.map_postgres_err()?
.rows_affected();
tx.commit().await.map_postgres_err()?;
Ok(WalletDailyUsageAggregationResult {
aggregated_wallets: usize::try_from(aggregated_wallets).unwrap_or(usize::MAX),
deleted_stale_ledgers: usize::try_from(deleted_stale_ledgers).unwrap_or(usize::MAX),
})
}
}
@@ -0,0 +1,156 @@
use sqlx::Row;
use crate::backend::SqliteBackend;
use crate::driver::sqlite::sqlite_real;
use crate::error::SqlResultExt;
use crate::{DataLayerError, WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult};
use super::{u64_to_i64, wallet_daily_usage_id};
const SELECT_WALLET_DAILY_USAGE_AGGREGATES_SQL: &str = r#"
SELECT
usage_settlement_snapshots.wallet_id AS wallet_id,
COUNT(*) AS total_requests,
CAST(COALESCE(SUM("usage".total_cost_usd), 0) AS REAL) AS total_cost_usd,
COALESCE(SUM("usage".input_tokens), 0) AS input_tokens,
COALESCE(SUM("usage".output_tokens), 0) AS output_tokens,
COALESCE(SUM("usage".cache_creation_input_tokens), 0) AS cache_creation_tokens,
COALESCE(SUM("usage".cache_read_input_tokens), 0) AS cache_read_tokens,
MIN(COALESCE(usage_settlement_snapshots.finalized_at, "usage".finalized_at)) AS first_finalized_at,
MAX(COALESCE(usage_settlement_snapshots.finalized_at, "usage".finalized_at)) AS last_finalized_at
FROM "usage"
JOIN usage_settlement_snapshots
ON usage_settlement_snapshots.request_id = "usage".request_id
WHERE usage_settlement_snapshots.wallet_id IS NOT NULL
AND usage_settlement_snapshots.wallet_id <> ''
AND COALESCE(usage_settlement_snapshots.billing_status, "usage".billing_status) = 'settled'
AND "usage".total_cost_usd > 0
AND COALESCE(usage_settlement_snapshots.finalized_at, "usage".finalized_at) >= ?
AND COALESCE(usage_settlement_snapshots.finalized_at, "usage".finalized_at) < ?
GROUP BY usage_settlement_snapshots.wallet_id
"#;
impl SqliteBackend {
pub async fn aggregate_wallet_daily_usage(
&self,
input: &WalletDailyUsageAggregationInput,
) -> Result<WalletDailyUsageAggregationResult, DataLayerError> {
let window_start = u64_to_i64(input.window_start_unix_secs, "window_start")?;
let window_end = u64_to_i64(input.window_end_unix_secs, "window_end")?;
let aggregated_at = u64_to_i64(input.aggregated_at_unix_secs, "aggregated_at")?;
let mut tx = self.pool().begin().await.map_sql_err()?;
let rows = sqlx::query(SELECT_WALLET_DAILY_USAGE_AGGREGATES_SQL)
.bind(window_start)
.bind(window_end)
.fetch_all(&mut *tx)
.await
.map_sql_err()?;
let mut aggregated_wallets = 0usize;
for row in rows {
let wallet_id: String = row.try_get("wallet_id").map_sql_err()?;
sqlx::query(
r#"
DELETE FROM wallet_daily_usage_ledgers
WHERE wallet_id = ?
AND billing_date = ?
AND billing_timezone = ?
"#,
)
.bind(&wallet_id)
.bind(&input.billing_date)
.bind(&input.billing_timezone)
.execute(&mut *tx)
.await
.map_sql_err()?;
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,
first_finalized_at,
last_finalized_at,
aggregated_at,
created_at,
updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
"#,
)
.bind(wallet_daily_usage_id(
&wallet_id,
&input.billing_date,
&input.billing_timezone,
))
.bind(&wallet_id)
.bind(&input.billing_date)
.bind(&input.billing_timezone)
.bind(sqlite_real(&row, "total_cost_usd")?)
.bind(row.try_get::<i64, _>("total_requests").map_sql_err()?)
.bind(row.try_get::<i64, _>("input_tokens").map_sql_err()?)
.bind(row.try_get::<i64, _>("output_tokens").map_sql_err()?)
.bind(
row.try_get::<i64, _>("cache_creation_tokens")
.map_sql_err()?,
)
.bind(row.try_get::<i64, _>("cache_read_tokens").map_sql_err()?)
.bind(
row.try_get::<Option<i64>, _>("first_finalized_at")
.map_sql_err()?,
)
.bind(
row.try_get::<Option<i64>, _>("last_finalized_at")
.map_sql_err()?,
)
.bind(aggregated_at)
.bind(aggregated_at)
.bind(aggregated_at)
.execute(&mut *tx)
.await
.map_sql_err()?;
aggregated_wallets += 1;
}
let deleted_stale_ledgers = sqlx::query(
r#"
DELETE FROM wallet_daily_usage_ledgers
WHERE billing_date = ?
AND billing_timezone = ?
AND NOT EXISTS (
SELECT 1
FROM "usage"
JOIN usage_settlement_snapshots
ON usage_settlement_snapshots.request_id = "usage".request_id
WHERE usage_settlement_snapshots.wallet_id = wallet_daily_usage_ledgers.wallet_id
AND COALESCE(usage_settlement_snapshots.billing_status, "usage".billing_status) = 'settled'
AND "usage".total_cost_usd > 0
AND COALESCE(usage_settlement_snapshots.finalized_at, "usage".finalized_at) >= ?
AND COALESCE(usage_settlement_snapshots.finalized_at, "usage".finalized_at) < ?
)
"#,
)
.bind(&input.billing_date)
.bind(&input.billing_timezone)
.bind(window_start)
.bind(window_end)
.execute(&mut *tx)
.await
.map_sql_err()?
.rows_affected();
tx.commit().await.map_sql_err()?;
Ok(WalletDailyUsageAggregationResult {
aggregated_wallets,
deleted_stale_ledgers: usize::try_from(deleted_stale_ledgers).unwrap_or(usize::MAX),
})
}
}
@@ -0,0 +1,430 @@
use std::fmt;
use std::sync::Arc;
#[cfg(feature = "mysql")]
use super::MysqlBackend;
#[cfg(feature = "postgres")]
use super::PostgresBackend;
#[cfg(feature = "sqlite")]
use super::SqliteBackend;
use crate::repository::announcements::AnnouncementWriteRepository;
use crate::repository::auth::AuthApiKeyWriteRepository;
use crate::repository::auth_modules::AuthModuleWriteRepository;
use crate::repository::background_tasks::BackgroundTaskWriteRepository;
use crate::repository::candidates::RequestCandidateWriteRepository;
use crate::repository::gemini_file_mappings::GeminiFileMappingWriteRepository;
use crate::repository::global_models::GlobalModelWriteRepository;
use crate::repository::management_tokens::ManagementTokenWriteRepository;
use crate::repository::oauth_providers::OAuthProviderWriteRepository;
use crate::repository::pool_scores::PoolMemberScoreWriteRepository;
use crate::repository::provider_catalog::ProviderCatalogWriteRepository;
use crate::repository::proxy_nodes::ProxyNodeWriteRepository;
use crate::repository::quota::ProviderQuotaWriteRepository;
use crate::repository::routing_profiles::RoutingGroupWriteRepository;
use crate::repository::settlement::SettlementWriteRepository;
use crate::repository::usage::UsageWriteRepository;
use crate::repository::video_tasks::VideoTaskWriteRepository;
use crate::repository::wallet::WalletWriteRepository;
#[derive(Clone, Default)]
pub struct DataWriteRepositories {
announcements: Option<Arc<dyn AnnouncementWriteRepository>>,
auth_api_keys: Option<Arc<dyn AuthApiKeyWriteRepository>>,
auth_modules: Option<Arc<dyn AuthModuleWriteRepository>>,
background_tasks: Option<Arc<dyn BackgroundTaskWriteRepository>>,
request_candidates: Option<Arc<dyn RequestCandidateWriteRepository>>,
gemini_file_mappings: Option<Arc<dyn GeminiFileMappingWriteRepository>>,
global_models: Option<Arc<dyn GlobalModelWriteRepository>>,
management_tokens: Option<Arc<dyn ManagementTokenWriteRepository>>,
oauth_providers: Option<Arc<dyn OAuthProviderWriteRepository>>,
pool_scores: Option<Arc<dyn PoolMemberScoreWriteRepository>>,
proxy_nodes: Option<Arc<dyn ProxyNodeWriteRepository>>,
provider_catalog: Option<Arc<dyn ProviderCatalogWriteRepository>>,
provider_quotas: Option<Arc<dyn ProviderQuotaWriteRepository>>,
routing_groups: Option<Arc<dyn RoutingGroupWriteRepository>>,
settlement: Option<Arc<dyn SettlementWriteRepository>>,
usage: Option<Arc<dyn UsageWriteRepository>>,
video_tasks: Option<Arc<dyn VideoTaskWriteRepository>>,
wallets: Option<Arc<dyn WalletWriteRepository>>,
}
impl fmt::Debug for DataWriteRepositories {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("DataWriteRepositories")
.field("has_announcements", &self.announcements.is_some())
.field("has_auth_api_keys", &self.auth_api_keys.is_some())
.field("has_auth_modules", &self.auth_modules.is_some())
.field("has_background_tasks", &self.background_tasks.is_some())
.field("has_request_candidates", &self.request_candidates.is_some())
.field(
"has_gemini_file_mappings",
&self.gemini_file_mappings.is_some(),
)
.field("has_global_models", &self.global_models.is_some())
.field("has_management_tokens", &self.management_tokens.is_some())
.field("has_oauth_providers", &self.oauth_providers.is_some())
.field("has_pool_scores", &self.pool_scores.is_some())
.field("has_proxy_nodes", &self.proxy_nodes.is_some())
.field("has_provider_catalog", &self.provider_catalog.is_some())
.field("has_provider_quotas", &self.provider_quotas.is_some())
.field("has_routing_groups", &self.routing_groups.is_some())
.field("has_settlement", &self.settlement.is_some())
.field("has_usage", &self.usage.is_some())
.field("has_video_tasks", &self.video_tasks.is_some())
.field("has_wallets", &self.wallets.is_some())
.finish()
}
}
impl DataWriteRepositories {
pub(crate) fn from_backends(
#[cfg(feature = "postgres")] postgres: Option<&PostgresBackend>,
#[cfg(feature = "mysql")] mysql: Option<&MysqlBackend>,
#[cfg(feature = "sqlite")] sqlite: Option<&SqliteBackend>,
) -> Self {
let mut repositories = Self::default();
#[cfg(feature = "postgres")]
if let Some(postgres) = postgres {
repositories.install_postgres(postgres);
}
#[cfg(feature = "mysql")]
if let Some(mysql) = mysql {
repositories.install_mysql(mysql);
}
#[cfg(feature = "sqlite")]
if let Some(sqlite) = sqlite {
repositories.install_sqlite(sqlite);
}
repositories
}
#[cfg(feature = "postgres")]
fn install_postgres(&mut self, backend: &PostgresBackend) {
if self.announcements.is_none() {
self.announcements = Some(PostgresBackend::announcement_write_repository(backend));
}
if self.auth_api_keys.is_none() {
self.auth_api_keys = Some(PostgresBackend::auth_api_key_write_repository(backend));
}
if self.auth_modules.is_none() {
self.auth_modules = Some(PostgresBackend::auth_module_write_repository(backend));
}
if self.background_tasks.is_none() {
self.background_tasks =
Some(PostgresBackend::background_task_write_repository(backend));
}
if self.request_candidates.is_none() {
self.request_candidates =
Some(PostgresBackend::request_candidate_write_repository(backend));
}
if self.gemini_file_mappings.is_none() {
self.gemini_file_mappings = Some(
PostgresBackend::gemini_file_mapping_write_repository(backend),
);
}
if self.global_models.is_none() {
self.global_models = Some(PostgresBackend::global_model_write_repository(backend));
}
if self.management_tokens.is_none() {
self.management_tokens =
Some(PostgresBackend::management_token_write_repository(backend));
}
if self.oauth_providers.is_none() {
self.oauth_providers = Some(PostgresBackend::oauth_provider_write_repository(backend));
}
if self.pool_scores.is_none() {
self.pool_scores = Some(PostgresBackend::pool_score_write_repository(backend));
}
if self.proxy_nodes.is_none() {
self.proxy_nodes = Some(PostgresBackend::proxy_node_write_repository(backend));
}
if self.provider_catalog.is_none() {
self.provider_catalog =
Some(PostgresBackend::provider_catalog_write_repository(backend));
}
if self.provider_quotas.is_none() {
self.provider_quotas = Some(PostgresBackend::provider_quota_write_repository(backend));
}
if self.routing_groups.is_none() {
self.routing_groups = Some(PostgresBackend::routing_group_write_repository(backend));
}
if self.settlement.is_none() {
self.settlement = Some(PostgresBackend::settlement_write_repository(backend));
}
if self.usage.is_none() {
self.usage = Some(PostgresBackend::usage_write_repository(backend));
}
if self.video_tasks.is_none() {
self.video_tasks = Some(PostgresBackend::video_task_write_repository(backend));
}
if self.wallets.is_none() {
self.wallets = Some(PostgresBackend::wallet_write_repository(backend));
}
}
#[cfg(feature = "mysql")]
fn install_mysql(&mut self, backend: &MysqlBackend) {
if self.announcements.is_none() {
self.announcements = Some(MysqlBackend::announcement_write_repository(backend));
}
if self.auth_api_keys.is_none() {
self.auth_api_keys = Some(MysqlBackend::auth_api_key_write_repository(backend));
}
if self.auth_modules.is_none() {
self.auth_modules = Some(MysqlBackend::auth_module_write_repository(backend));
}
if self.background_tasks.is_none() {
self.background_tasks = Some(MysqlBackend::background_task_write_repository(backend));
}
if self.request_candidates.is_none() {
self.request_candidates =
Some(MysqlBackend::request_candidate_write_repository(backend));
}
if self.gemini_file_mappings.is_none() {
self.gemini_file_mappings =
Some(MysqlBackend::gemini_file_mapping_write_repository(backend));
}
if self.global_models.is_none() {
self.global_models = Some(MysqlBackend::global_model_write_repository(backend));
}
if self.management_tokens.is_none() {
self.management_tokens = Some(MysqlBackend::management_token_write_repository(backend));
}
if self.oauth_providers.is_none() {
self.oauth_providers = Some(MysqlBackend::oauth_provider_write_repository(backend));
}
if self.pool_scores.is_none() {
self.pool_scores = Some(MysqlBackend::pool_score_write_repository(backend));
}
if self.proxy_nodes.is_none() {
self.proxy_nodes = Some(MysqlBackend::proxy_node_write_repository(backend));
}
if self.provider_catalog.is_none() {
self.provider_catalog = Some(MysqlBackend::provider_catalog_write_repository(backend));
}
if self.provider_quotas.is_none() {
self.provider_quotas = Some(MysqlBackend::provider_quota_write_repository(backend));
}
if self.routing_groups.is_none() {
self.routing_groups = Some(MysqlBackend::routing_group_write_repository(backend));
}
if self.settlement.is_none() {
self.settlement = Some(MysqlBackend::settlement_write_repository(backend));
}
if self.usage.is_none() {
self.usage = Some(MysqlBackend::usage_write_repository(backend));
}
if self.video_tasks.is_none() {
self.video_tasks = Some(MysqlBackend::video_task_write_repository(backend));
}
if self.wallets.is_none() {
self.wallets = Some(MysqlBackend::wallet_write_repository(backend));
}
}
#[cfg(feature = "sqlite")]
fn install_sqlite(&mut self, backend: &SqliteBackend) {
if self.announcements.is_none() {
self.announcements = Some(SqliteBackend::announcement_write_repository(backend));
}
if self.auth_api_keys.is_none() {
self.auth_api_keys = Some(SqliteBackend::auth_api_key_write_repository(backend));
}
if self.auth_modules.is_none() {
self.auth_modules = Some(SqliteBackend::auth_module_write_repository(backend));
}
if self.background_tasks.is_none() {
self.background_tasks = Some(SqliteBackend::background_task_write_repository(backend));
}
if self.request_candidates.is_none() {
self.request_candidates =
Some(SqliteBackend::request_candidate_write_repository(backend));
}
if self.gemini_file_mappings.is_none() {
self.gemini_file_mappings =
Some(SqliteBackend::gemini_file_mapping_write_repository(backend));
}
if self.global_models.is_none() {
self.global_models = Some(SqliteBackend::global_model_write_repository(backend));
}
if self.management_tokens.is_none() {
self.management_tokens =
Some(SqliteBackend::management_token_write_repository(backend));
}
if self.oauth_providers.is_none() {
self.oauth_providers = Some(SqliteBackend::oauth_provider_write_repository(backend));
}
if self.pool_scores.is_none() {
self.pool_scores = Some(SqliteBackend::pool_score_write_repository(backend));
}
if self.proxy_nodes.is_none() {
self.proxy_nodes = Some(SqliteBackend::proxy_node_write_repository(backend));
}
if self.provider_catalog.is_none() {
self.provider_catalog = Some(SqliteBackend::provider_catalog_write_repository(backend));
}
if self.provider_quotas.is_none() {
self.provider_quotas = Some(SqliteBackend::provider_quota_write_repository(backend));
}
if self.routing_groups.is_none() {
self.routing_groups = Some(SqliteBackend::routing_group_write_repository(backend));
}
if self.settlement.is_none() {
self.settlement = Some(SqliteBackend::settlement_write_repository(backend));
}
if self.usage.is_none() {
self.usage = Some(SqliteBackend::usage_write_repository(backend));
}
if self.video_tasks.is_none() {
self.video_tasks = Some(SqliteBackend::video_task_write_repository(backend));
}
if self.wallets.is_none() {
self.wallets = Some(SqliteBackend::wallet_write_repository(backend));
}
}
#[cfg(test)]
#[cfg(feature = "postgres")]
pub(crate) fn from_postgres(postgres: Option<&PostgresBackend>) -> Self {
Self::from_backends(
postgres,
#[cfg(feature = "mysql")]
None,
#[cfg(feature = "sqlite")]
None,
)
}
pub fn announcements(&self) -> Option<Arc<dyn AnnouncementWriteRepository>> {
self.announcements.clone()
}
pub fn auth_api_keys(&self) -> Option<Arc<dyn AuthApiKeyWriteRepository>> {
self.auth_api_keys.clone()
}
pub fn auth_modules(&self) -> Option<Arc<dyn AuthModuleWriteRepository>> {
self.auth_modules.clone()
}
pub fn background_tasks(&self) -> Option<Arc<dyn BackgroundTaskWriteRepository>> {
self.background_tasks.clone()
}
pub fn usage(&self) -> Option<Arc<dyn UsageWriteRepository>> {
self.usage.clone()
}
pub fn request_candidates(&self) -> Option<Arc<dyn RequestCandidateWriteRepository>> {
self.request_candidates.clone()
}
pub fn gemini_file_mappings(&self) -> Option<Arc<dyn GeminiFileMappingWriteRepository>> {
self.gemini_file_mappings.clone()
}
pub fn global_models(&self) -> Option<Arc<dyn GlobalModelWriteRepository>> {
self.global_models.clone()
}
pub fn management_tokens(&self) -> Option<Arc<dyn ManagementTokenWriteRepository>> {
self.management_tokens.clone()
}
pub fn oauth_providers(&self) -> Option<Arc<dyn OAuthProviderWriteRepository>> {
self.oauth_providers.clone()
}
pub fn pool_scores(&self) -> Option<Arc<dyn PoolMemberScoreWriteRepository>> {
self.pool_scores.clone()
}
pub fn proxy_nodes(&self) -> Option<Arc<dyn ProxyNodeWriteRepository>> {
self.proxy_nodes.clone()
}
pub fn provider_quotas(&self) -> Option<Arc<dyn ProviderQuotaWriteRepository>> {
self.provider_quotas.clone()
}
pub fn routing_groups(&self) -> Option<Arc<dyn RoutingGroupWriteRepository>> {
self.routing_groups.clone()
}
pub fn provider_catalog(&self) -> Option<Arc<dyn ProviderCatalogWriteRepository>> {
self.provider_catalog.clone()
}
pub fn settlement(&self) -> Option<Arc<dyn SettlementWriteRepository>> {
self.settlement.clone()
}
pub fn video_tasks(&self) -> Option<Arc<dyn VideoTaskWriteRepository>> {
self.video_tasks.clone()
}
pub fn wallets(&self) -> Option<Arc<dyn WalletWriteRepository>> {
self.wallets.clone()
}
pub fn has_any(&self) -> bool {
self.announcements.is_some()
|| self.auth_api_keys.is_some()
|| self.auth_modules.is_some()
|| self.background_tasks.is_some()
|| self.request_candidates.is_some()
|| self.gemini_file_mappings.is_some()
|| self.global_models.is_some()
|| self.management_tokens.is_some()
|| self.oauth_providers.is_some()
|| self.pool_scores.is_some()
|| self.proxy_nodes.is_some()
|| self.provider_catalog.is_some()
|| self.provider_quotas.is_some()
|| self.routing_groups.is_some()
|| self.settlement.is_some()
|| self.usage.is_some()
|| self.video_tasks.is_some()
|| self.wallets.is_some()
}
}
#[cfg(all(test, feature = "postgres"))]
mod tests {
use super::DataWriteRepositories;
use crate::backend::PostgresBackend;
use crate::driver::postgres::PostgresPoolConfig;
#[tokio::test]
async fn builds_write_repositories_from_postgres_backend() {
let backend = PostgresBackend::from_config(PostgresPoolConfig {
database_url: "postgres://localhost/aether".to_string(),
min_connections: 1,
max_connections: 4,
acquire_timeout_ms: 1_000,
idle_timeout_ms: 5_000,
max_lifetime_ms: 30_000,
statement_cache_capacity: 64,
require_ssl: false,
})
.expect("postgres backend should build");
let write = DataWriteRepositories::from_postgres(Some(&backend));
assert!(write.has_any());
assert!(write.announcements().is_some());
assert!(write.auth_api_keys().is_some());
assert!(write.auth_modules().is_some());
assert!(write.request_candidates().is_some());
assert!(write.gemini_file_mappings().is_some());
assert!(write.global_models().is_some());
assert!(write.management_tokens().is_some());
assert!(write.oauth_providers().is_some());
assert!(write.proxy_nodes().is_some());
assert!(write.provider_catalog().is_some());
assert!(write.provider_quotas().is_some());
assert!(write.settlement().is_some());
assert!(write.usage().is_some());
assert!(write.video_tasks().is_some());
assert!(write.wallets().is_some());
}
}