Add multi-database data layer

Introduce aether-data-schema and driver-specific schema generation for Postgres, MySQL, and SQLite.

Split data backends, lifecycle, repositories, and gateway runtime integration across database drivers.

Verified with cargo fmt --all --check, cargo clippy --workspace --all-targets -- -D warnings, and cargo test --workspace.
This commit is contained in:
fawney19
2026-05-05 18:27:36 +08:00
parent 099653f732
commit fce7e959e5
372 changed files with 86217 additions and 21160 deletions

View File

@@ -1,10 +1,11 @@
use aether_data::postgres::PostgresPoolConfig;
use aether_data::redis::RedisClientConfig;
use aether_data::DataLayerConfig;
use aether_data::driver::postgres::PostgresPoolConfig;
use aether_data::driver::redis::RedisClientConfig;
use aether_data::{DataLayerConfig, SqlDatabaseConfig};
use std::fmt;
#[derive(Clone, Default)]
pub struct GatewayDataConfig {
database: Option<SqlDatabaseConfig>,
postgres: Option<PostgresPoolConfig>,
redis: Option<RedisClientConfig>,
encryption_key: Option<String>,
@@ -13,6 +14,7 @@ pub struct GatewayDataConfig {
impl fmt::Debug for GatewayDataConfig {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("GatewayDataConfig")
.field("database", &self.database)
.field("postgres", &self.postgres)
.field("redis", &self.redis)
.field("has_encryption_key", &self.encryption_key.is_some())
@@ -27,12 +29,23 @@ impl GatewayDataConfig {
pub fn from_postgres_config(postgres: PostgresPoolConfig) -> Self {
Self {
database: Some(SqlDatabaseConfig::from_postgres_config(postgres.clone())),
postgres: Some(postgres),
redis: None,
encryption_key: None,
}
}
pub fn from_database_config(database: SqlDatabaseConfig) -> Self {
let postgres = database.to_postgres_config().ok();
Self {
database: Some(database),
postgres,
redis: None,
encryption_key: None,
}
}
pub fn from_postgres_url(database_url: impl Into<String>, require_ssl: bool) -> Self {
let mut postgres = PostgresPoolConfig::default();
postgres.database_url = database_url.into();
@@ -44,6 +57,10 @@ impl GatewayDataConfig {
self.postgres.as_ref()
}
pub fn database(&self) -> Option<&SqlDatabaseConfig> {
self.database.as_ref()
}
pub fn redis(&self) -> Option<&RedisClientConfig> {
self.redis.as_ref()
}
@@ -80,11 +97,12 @@ impl GatewayDataConfig {
}
pub fn is_enabled(&self) -> bool {
self.postgres.is_some() || self.redis.is_some()
self.database.is_some() || self.postgres.is_some() || self.redis.is_some()
}
pub fn to_data_layer_config(&self) -> DataLayerConfig {
DataLayerConfig {
database: self.database.clone(),
postgres: self.postgres.clone(),
redis: self.redis.clone(),
}

File diff suppressed because it is too large Load Diff

View File

@@ -1,5 +1,5 @@
use aether_data::redis::{RedisKvRunner, RedisKvRunnerConfig, RedisLockRunner};
use aether_data::{DataBackends, DataLayerError};
use aether_data::driver::redis::{RedisKvRunner, RedisKvRunnerConfig, RedisLockRunner};
use aether_data::{DataBackends, DataLayerError, DatabaseDriver};
use super::{GatewayDataConfig, GatewayDataState, StoredSystemConfigEntry};
@@ -138,6 +138,36 @@ impl GatewayDataState {
self.backends.is_some()
}
pub(crate) fn has_database_maintenance_backend(&self) -> bool {
self.backends
.as_ref()
.is_some_and(|backends| backends.has_database_maintenance_backend())
}
pub(crate) fn has_database_pool_summary(&self) -> bool {
self.backends
.as_ref()
.is_some_and(|backends| backends.has_database_pool_summary())
}
pub(crate) fn has_wallet_daily_usage_aggregation_backend(&self) -> bool {
self.backends
.as_ref()
.is_some_and(|backends| backends.has_wallet_daily_usage_aggregation_backend())
}
pub(crate) fn has_stats_hourly_aggregation_backend(&self) -> bool {
self.backends
.as_ref()
.is_some_and(|backends| backends.has_stats_hourly_aggregation_backend())
}
pub(crate) fn has_stats_daily_aggregation_backend(&self) -> bool {
self.backends
.as_ref()
.is_some_and(|backends| backends.has_stats_daily_aggregation_backend())
}
pub(crate) fn has_auth_api_key_reader(&self) -> bool {
self.auth_api_key_reader.is_some()
}
@@ -158,6 +188,13 @@ impl GatewayDataState {
self.announcement_writer.is_some()
}
pub(crate) fn has_audit_log_reader(&self) -> bool {
self.backends
.as_ref()
.and_then(|backends| backends.read().audit_logs())
.is_some()
}
pub(crate) fn has_management_token_reader(&self) -> bool {
self.management_token_reader.is_some()
}
@@ -223,8 +260,7 @@ impl GatewayDataState {
|| self
.backends
.as_ref()
.and_then(|backends| backends.postgres())
.is_some()
.is_some_and(|backends| backends.has_system_config_backend())
}
pub(crate) fn oauth_refresh_lock_runner(&self) -> Option<RedisLockRunner> {
@@ -240,15 +276,10 @@ impl GatewayDataState {
.and_then(|backend| backend.kv_runner(RedisKvRunnerConfig::default()).ok())
}
pub(crate) fn postgres_pool(&self) -> Option<aether_data::postgres::PostgresPool> {
pub(crate) fn database_driver(&self) -> Option<DatabaseDriver> {
self.backends
.as_ref()
.and_then(|backends| backends.postgres())
.map(|backend| backend.pool_clone())
}
pub(crate) fn postgres_max_connections(&self) -> Option<u32> {
self.config.postgres().map(|config| config.max_connections)
.and_then(|backends| backends.database_driver())
}
pub(crate) fn has_provider_quota_writer(&self) -> bool {
@@ -307,14 +338,10 @@ impl GatewayDataState {
.get(key)
.map(|entry| entry.value.clone()));
}
match self
.backends
.as_ref()
.and_then(|backends| backends.postgres())
{
Some(backend) => backend.find_system_config_value(key).await,
None => Ok(None),
}
let Some(backends) = self.backends.as_ref() else {
return Ok(None);
};
backends.find_system_config_value(key).await
}
pub(crate) async fn upsert_system_config_value(
@@ -340,14 +367,10 @@ impl GatewayDataState {
.cloned()
.collect());
}
match self
.backends
.as_ref()
.and_then(|backends| backends.postgres())
{
Some(backend) => backend.list_system_config_entries().await,
None => Ok(Vec::new()),
}
let Some(backends) = self.backends.as_ref() else {
return Ok(Vec::new());
};
backends.list_system_config_entries().await
}
pub(crate) async fn upsert_system_config_entry(
@@ -370,23 +393,20 @@ impl GatewayDataState {
values.insert(key.to_string(), entry.clone());
return Ok(entry);
}
match self
.backends
.as_ref()
.and_then(|backends| backends.postgres())
{
Some(backend) => {
backend
.upsert_system_config_entry(key, value, description)
.await
if let Some(backends) = self.backends.as_ref() {
if let Some(entry) = backends
.upsert_system_config_entry(key, value, description)
.await?
{
return Ok(entry);
}
None => Ok(StoredSystemConfigEntry {
key: key.to_string(),
value: value.clone(),
description: description.map(ToOwned::to_owned),
updated_at_unix_secs: Some(current_system_config_updated_at_unix_secs()),
}),
}
Ok(StoredSystemConfigEntry {
key: key.to_string(),
value: value.clone(),
description: description.map(ToOwned::to_owned),
updated_at_unix_secs: Some(current_system_config_updated_at_unix_secs()),
})
}
pub(crate) async fn delete_system_config_value(
@@ -400,25 +420,17 @@ impl GatewayDataState {
.remove(key)
.is_some());
}
match self
.backends
.as_ref()
.and_then(|backends| backends.postgres())
{
Some(backend) => backend.delete_system_config_value(key).await,
None => Ok(false),
}
let Some(backends) = self.backends.as_ref() else {
return Ok(false);
};
backends.delete_system_config_value(key).await
}
pub(crate) async fn read_admin_system_stats(
&self,
) -> Result<super::AdminSystemStats, DataLayerError> {
match self
.backends
.as_ref()
.and_then(|backends| backends.postgres())
{
Some(backend) => backend.read_admin_system_stats().await,
match self.backends.as_ref() {
Some(backends) => backends.read_admin_system_stats().await,
None => Ok(super::AdminSystemStats::default()),
}
}

View File

@@ -1,6 +1,6 @@
use aether_billing::enrich_usage_event_with_billing;
use aether_billing::BillingModelContextLookup;
use aether_data::redis::RedisStreamRunner;
use aether_data::driver::redis::RedisStreamRunner;
use aether_data::repository::audit::RequestAuditReader;
use aether_data::repository::auth::{
AuthApiKeyLookupKey, ResolvedAuthApiKeySnapshotReader, StoredAuthApiKeySnapshot,

View File

@@ -11,12 +11,17 @@ use crate::provider_transport::{
read_provider_transport_snapshot, GatewayProviderTransportSnapshot,
};
use crate::video_tasks::LocalVideoTaskReadResponse;
use aether_data::redis::{RedisKvRunner, RedisKvRunnerConfig, RedisLockRunner, RedisStreamRunner};
use aether_data::driver::redis::{
RedisKvRunner, RedisKvRunnerConfig, RedisLockRunner, RedisStreamRunner,
};
use aether_data::repository::announcements::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,
};
use aether_data::repository::audit::RequestAuditBundle;
use aether_data::repository::audit::{
AuditLogListQuery, RequestAuditBundle, StoredAdminAuditLogPage, StoredSuspiciousActivity,
StoredUserAuditLogPage,
};
use aether_data::repository::auth::{
AuthApiKeyLookupKey, AuthApiKeyReadRepository, AuthApiKeyWriteRepository,
StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
@@ -47,7 +52,8 @@ use aether_data::repository::proxy_nodes::{
};
pub(crate) use aether_data::repository::system::{AdminSystemStats, StoredSystemConfigEntry};
use aether_data::repository::users::{
StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary, UserReadRepository,
StoredUserAuthRecord, StoredUserExportRow, StoredUserOAuthLinkSummary, StoredUserSummary,
UserReadRepository,
};
pub(crate) use aether_data::repository::users::{
StoredUserPreferenceRecord, StoredUserSessionRecord,
@@ -72,8 +78,13 @@ use aether_data::repository::wallet::{
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome,
WalletReadRepository, WalletWriteRepository,
};
use aether_data::{DataBackends, DataLayerError};
use aether_data::{
DataBackends, DataLayerError, DatabaseMaintenanceSummary, WalletDailyUsageAggregationInput,
WalletDailyUsageAggregationResult,
};
use aether_data_contracts::repository::billing::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
BillingReadRepository, StoredBillingModelContext,
};
use aether_data_contracts::repository::candidate_selection::{
@@ -104,8 +115,8 @@ use aether_data_contracts::repository::settlement::{
SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput,
};
use aether_data_contracts::repository::usage::{
StoredProviderUsageSummary, StoredRequestUsageAudit, UpsertUsageRecord, UsageReadRepository,
UsageWriteRepository,
PendingUsageCleanupSummary, StoredProviderUsageSummary, StoredRequestUsageAudit,
UpsertUsageRecord, UsageReadRepository, UsageWriteRepository,
};
use aether_data_contracts::repository::video_tasks::{
StoredVideoTask, UpsertVideoTask, VideoTaskLookupKey, VideoTaskModelCount,

View File

@@ -1,36 +1,139 @@
use super::{
read_decision_trace, read_provider_transport_snapshot, read_request_candidate_trace,
AdjustWalletBalanceInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery,
AdjustWalletBalanceInput, AdminBillingCollectorRecord, AdminBillingCollectorWriteInput,
AdminBillingMutationOutcome, AdminBillingPresetApplyResult, AdminBillingRuleRecord,
AdminBillingRuleWriteInput, AdminPaymentOrderListQuery, AdminRedeemCodeBatchListQuery,
AdminRedeemCodeListQuery, AdminWalletLedgerQuery, AdminWalletListQuery,
AdminWalletRefundRequestListQuery, AnnouncementListQuery, CompleteAdminWalletRefundInput,
CreateAdminRedeemCodeBatchInput, CreateAdminRedeemCodeBatchResult, CreateAnnouncementRecord,
CreateManualWalletRechargeInput, CreateWalletRechargeOrderInput,
CreateWalletRechargeOrderOutcome, CreateWalletRefundRequestInput,
CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput, DataLayerError, DecisionTrace,
DeleteAdminRedeemCodeBatchInput, DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput,
FailAdminWalletRefundInput, GatewayDataState, GatewayProviderTransportSnapshot,
LocalVideoTaskReadResponse, ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput,
ProcessPaymentCallbackOutcome, RedeemWalletCodeInput, RedeemWalletCodeOutcome,
RedisStreamRunner, RequestAuditBundle, RequestCandidateTrace, StoredAdminPaymentCallbackPage,
AdminWalletRefundRequestListQuery, AnnouncementListQuery, AuditLogListQuery,
CompleteAdminWalletRefundInput, CreateAdminRedeemCodeBatchInput,
CreateAdminRedeemCodeBatchResult, CreateAnnouncementRecord, CreateManualWalletRechargeInput,
CreateWalletRechargeOrderInput, CreateWalletRechargeOrderOutcome,
CreateWalletRefundRequestInput, CreateWalletRefundRequestOutcome, CreditAdminPaymentOrderInput,
DataLayerError, DatabaseMaintenanceSummary, DecisionTrace, DeleteAdminRedeemCodeBatchInput,
DisableAdminRedeemCodeBatchInput, DisableAdminRedeemCodeInput, FailAdminWalletRefundInput,
GatewayDataState, GatewayProviderTransportSnapshot, LocalVideoTaskReadResponse,
ProcessAdminWalletRefundInput, ProcessPaymentCallbackInput, ProcessPaymentCallbackOutcome,
RedeemWalletCodeInput, RedeemWalletCodeOutcome, RedisStreamRunner, RequestAuditBundle,
RequestCandidateTrace, StoredAdminAuditLogPage, StoredAdminPaymentCallbackPage,
StoredAdminPaymentOrder, StoredAdminPaymentOrderPage, StoredAdminRedeemCodeBatch,
StoredAdminRedeemCodeBatchPage, StoredAdminRedeemCodePage, StoredAdminWalletLedgerPage,
StoredAdminWalletListPage, StoredAdminWalletRefund, StoredAdminWalletRefundPage,
StoredAdminWalletRefundRequestPage, StoredAdminWalletTransaction,
StoredAdminWalletTransactionPage, StoredAnnouncement, StoredAnnouncementPage,
StoredBillingModelContext, StoredProviderQuotaSnapshot, StoredProviderUsageSummary,
StoredRequestUsageAudit, StoredUsageSettlement, StoredUserAuthRecord, StoredUserExportRow,
StoredUserSummary, StoredVideoTask, StoredWalletDailyUsageLedger,
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, UpdateAnnouncementRecord,
UpsertUsageRecord, UpsertVideoTask, UsageSettlementInput, VideoTaskLookupKey,
VideoTaskModelCount, VideoTaskQueryFilter, VideoTaskStatusCount, WalletLookupKey,
WalletMutationOutcome,
StoredRequestUsageAudit, StoredSuspiciousActivity, StoredUsageSettlement,
StoredUserAuditLogPage, StoredUserAuthRecord, StoredUserExportRow, StoredUserSummary,
StoredVideoTask, StoredWalletDailyUsageLedger, StoredWalletDailyUsageLedgerPage,
StoredWalletSnapshot, UpdateAnnouncementRecord, UpsertUsageRecord, UpsertVideoTask,
UsageSettlementInput, VideoTaskLookupKey, VideoTaskModelCount, VideoTaskQueryFilter,
VideoTaskStatusCount, WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult,
WalletLookupKey, WalletMutationOutcome,
};
use aether_data_contracts::repository::usage::{
StoredUsageDailySummary, UsageAuditListQuery, UsageDailyHeatmapQuery,
PendingUsageCleanupSummary, StoredUsageDailySummary, UsageAuditListQuery, UsageCleanupSummary,
UsageCleanupWindow, UsageDailyHeatmapQuery,
};
use aether_video_tasks_core::read_data_backed_video_task_response;
impl GatewayDataState {
pub(crate) async fn run_database_maintenance(
&self,
table_names: &[&str],
) -> Result<DatabaseMaintenanceSummary, DataLayerError> {
match &self.backends {
Some(backends) => backends.run_database_maintenance(table_names).await,
None => Ok(DatabaseMaintenanceSummary::default()),
}
}
pub(crate) async fn run_database_migrations(
&self,
) -> Result<bool, sqlx::migrate::MigrateError> {
match &self.backends {
Some(backends) => backends.run_database_migrations().await,
None => Ok(false),
}
}
pub(crate) async fn run_database_backfills(&self) -> Result<bool, sqlx::migrate::MigrateError> {
match &self.backends {
Some(backends) => backends.run_database_backfills().await,
None => Ok(false),
}
}
pub(crate) async fn pending_database_migrations(
&self,
) -> Result<
Option<Vec<aether_data::lifecycle::migrate::PendingMigrationInfo>>,
sqlx::migrate::MigrateError,
> {
match &self.backends {
Some(backends) => backends.pending_database_migrations().await,
None => Ok(None),
}
}
pub(crate) async fn prepare_database_for_startup(
&self,
) -> Result<
Option<Vec<aether_data::lifecycle::migrate::PendingMigrationInfo>>,
sqlx::migrate::MigrateError,
> {
match &self.backends {
Some(backends) => backends.prepare_database_for_startup().await,
None => Ok(None),
}
}
pub(crate) async fn pending_database_backfills(
&self,
) -> Result<
Option<Vec<aether_data::lifecycle::backfill::PendingBackfillInfo>>,
sqlx::migrate::MigrateError,
> {
match &self.backends {
Some(backends) => backends.pending_database_backfills().await,
None => Ok(None),
}
}
pub(crate) fn database_pool_summary(&self) -> Option<aether_data::DatabasePoolSummary> {
self.backends
.as_ref()
.and_then(|backends| backends.database_pool_summary())
}
pub(crate) async fn aggregate_wallet_daily_usage(
&self,
input: &WalletDailyUsageAggregationInput,
) -> Result<WalletDailyUsageAggregationResult, DataLayerError> {
match &self.backends {
Some(backends) => backends.aggregate_wallet_daily_usage(input).await,
None => Ok(WalletDailyUsageAggregationResult::default()),
}
}
pub(crate) async fn aggregate_stats_hourly(
&self,
input: &aether_data::StatsHourlyAggregationInput,
) -> Result<Option<aether_data::StatsHourlyAggregationSummary>, DataLayerError> {
match &self.backends {
Some(backends) => backends.aggregate_stats_hourly(input).await,
None => Ok(None),
}
}
pub(crate) async fn aggregate_stats_daily(
&self,
input: &aether_data::StatsDailyAggregationInput,
) -> Result<Option<aether_data::StatsDailyAggregationSummary>, DataLayerError> {
match &self.backends {
Some(backends) => backends.aggregate_stats_daily(input).await,
None => Ok(None),
}
}
pub(crate) async fn list_announcements(
&self,
query: &AnnouncementListQuery,
@@ -51,6 +154,91 @@ impl GatewayDataState {
}
}
pub(crate) async fn list_admin_audit_logs(
&self,
query: &AuditLogListQuery,
) -> Result<StoredAdminAuditLogPage, DataLayerError> {
let Some(repository) = self
.backends
.as_ref()
.and_then(|backends| backends.read().audit_logs())
else {
return Ok(StoredAdminAuditLogPage {
items: Vec::new(),
total: 0,
});
};
repository.list_admin_audit_logs(query).await
}
pub(crate) async fn list_admin_suspicious_activities(
&self,
cutoff_unix_secs: u64,
) -> Result<Vec<StoredSuspiciousActivity>, DataLayerError> {
let Some(repository) = self
.backends
.as_ref()
.and_then(|backends| backends.read().audit_logs())
else {
return Ok(Vec::new());
};
repository
.list_admin_suspicious_activities(cutoff_unix_secs)
.await
}
pub(crate) async fn read_admin_user_behavior_event_counts(
&self,
user_id: &str,
cutoff_unix_secs: u64,
) -> Result<std::collections::BTreeMap<String, u64>, DataLayerError> {
let Some(repository) = self
.backends
.as_ref()
.and_then(|backends| backends.read().audit_logs())
else {
return Ok(std::collections::BTreeMap::new());
};
repository
.read_admin_user_behavior_event_counts(user_id, cutoff_unix_secs)
.await
}
pub(crate) async fn list_user_audit_logs(
&self,
user_id: &str,
query: &AuditLogListQuery,
) -> Result<StoredUserAuditLogPage, DataLayerError> {
let Some(repository) = self
.backends
.as_ref()
.and_then(|backends| backends.read().audit_logs())
else {
return Ok(StoredUserAuditLogPage {
items: Vec::new(),
total: 0,
});
};
repository.list_user_audit_logs(user_id, query).await
}
pub(crate) async fn delete_audit_logs_before(
&self,
cutoff_unix_secs: u64,
limit: usize,
) -> Result<usize, DataLayerError> {
let Some(repository) = self
.backends
.as_ref()
.and_then(|backends| backends.read().audit_logs())
else {
return Ok(0);
};
repository
.delete_audit_logs_before(cutoff_unix_secs, limit)
.await
}
pub(crate) async fn count_unread_active_announcements(
&self,
user_id: &str,
@@ -748,6 +936,44 @@ impl GatewayDataState {
}
}
pub(crate) async fn cleanup_stale_pending_requests(
&self,
cutoff_unix_secs: u64,
now_unix_secs: u64,
timeout_minutes: u64,
batch_size: usize,
) -> Result<PendingUsageCleanupSummary, DataLayerError> {
match &self.usage_writer {
Some(repository) => {
repository
.cleanup_stale_pending_requests(
cutoff_unix_secs,
now_unix_secs,
timeout_minutes,
batch_size,
)
.await
}
None => Ok(PendingUsageCleanupSummary::default()),
}
}
pub(crate) async fn cleanup_usage(
&self,
window: &UsageCleanupWindow,
batch_size: usize,
auto_delete_expired_keys: bool,
) -> Result<UsageCleanupSummary, DataLayerError> {
match &self.usage_writer {
Some(repository) => {
repository
.cleanup_usage(window, batch_size, auto_delete_expired_keys)
.await
}
None => Ok(UsageCleanupSummary::default()),
}
}
pub(crate) async fn find_request_usage_by_request_id(
&self,
request_id: &str,
@@ -1257,6 +1483,153 @@ impl GatewayDataState {
}
}
pub(crate) async fn admin_billing_enabled_default_value_exists(
&self,
api_format: &str,
task_type: &str,
dimension_name: &str,
existing_id: Option<&str>,
) -> Result<Option<bool>, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.admin_billing_enabled_default_value_exists(
api_format,
task_type,
dimension_name,
existing_id,
)
.await
}
None => Ok(None),
}
}
pub(crate) async fn create_admin_billing_rule(
&self,
input: &AdminBillingRuleWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingRuleRecord>, DataLayerError> {
match &self.billing_reader {
Some(repository) => repository.create_admin_billing_rule(input).await,
None => Ok(AdminBillingMutationOutcome::Unavailable),
}
}
pub(crate) async fn list_admin_billing_rules(
&self,
task_type: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<Option<(Vec<AdminBillingRuleRecord>, u64)>, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.list_admin_billing_rules(task_type, is_enabled, page, page_size)
.await
}
None => Ok(None),
}
}
pub(crate) async fn find_admin_billing_rule(
&self,
rule_id: &str,
) -> Result<Option<AdminBillingRuleRecord>, DataLayerError> {
match &self.billing_reader {
Some(repository) => repository.find_admin_billing_rule(rule_id).await,
None => Ok(None),
}
}
pub(crate) async fn update_admin_billing_rule(
&self,
rule_id: &str,
input: &AdminBillingRuleWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingRuleRecord>, DataLayerError> {
match &self.billing_reader {
Some(repository) => repository.update_admin_billing_rule(rule_id, input).await,
None => Ok(AdminBillingMutationOutcome::Unavailable),
}
}
pub(crate) async fn create_admin_billing_collector(
&self,
input: &AdminBillingCollectorWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingCollectorRecord>, DataLayerError> {
match &self.billing_reader {
Some(repository) => repository.create_admin_billing_collector(input).await,
None => Ok(AdminBillingMutationOutcome::Unavailable),
}
}
pub(crate) async fn list_admin_billing_collectors(
&self,
api_format: Option<&str>,
task_type: Option<&str>,
dimension_name: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<Option<(Vec<AdminBillingCollectorRecord>, u64)>, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.list_admin_billing_collectors(
api_format,
task_type,
dimension_name,
is_enabled,
page,
page_size,
)
.await
}
None => Ok(None),
}
}
pub(crate) async fn find_admin_billing_collector(
&self,
collector_id: &str,
) -> Result<Option<AdminBillingCollectorRecord>, DataLayerError> {
match &self.billing_reader {
Some(repository) => repository.find_admin_billing_collector(collector_id).await,
None => Ok(None),
}
}
pub(crate) async fn update_admin_billing_collector(
&self,
collector_id: &str,
input: &AdminBillingCollectorWriteInput,
) -> Result<AdminBillingMutationOutcome<AdminBillingCollectorRecord>, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.update_admin_billing_collector(collector_id, input)
.await
}
None => Ok(AdminBillingMutationOutcome::Unavailable),
}
}
pub(crate) async fn apply_admin_billing_preset(
&self,
preset: &str,
mode: &str,
collectors: &[AdminBillingCollectorWriteInput],
) -> Result<AdminBillingMutationOutcome<AdminBillingPresetApplyResult>, DataLayerError> {
match &self.billing_reader {
Some(repository) => {
repository
.apply_admin_billing_preset(preset, mode, collectors)
.await
}
None => Ok(AdminBillingMutationOutcome::Unavailable),
}
}
pub(crate) async fn read_request_candidate_trace(
&self,
request_id: &str,

View File

@@ -8,8 +8,11 @@ use aether_data::repository::candidate_selection::InMemoryMinimalCandidateSelect
use aether_data::repository::candidates::InMemoryRequestCandidateRepository;
use aether_data::repository::provider_catalog::InMemoryProviderCatalogReadRepository;
use aether_data::repository::usage::InMemoryUsageReadRepository;
use aether_data::repository::users::{
InMemoryUserReadRepository, StoredUserAuthRecord, StoredUserPreferenceRecord,
};
use aether_data::repository::video_tasks::InMemoryVideoTaskRepository;
use aether_data::DataLayerError;
use aether_data::{DataLayerError, DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
use aether_data_contracts::repository::candidate_selection::{
StoredMinimalCandidateSelectionRow, StoredProviderModelMapping,
};
@@ -144,6 +147,176 @@ async fn app_state_wires_gateway_data_state_from_config() {
assert!(state.data.has_video_task_reader());
}
#[tokio::test]
async fn app_state_prepares_sqlite_database_startup() -> Result<(), Box<dyn std::error::Error>> {
let mut pool = SqlPoolConfig::default();
pool.min_connections = 0;
pool.max_connections = 1;
let database = SqlDatabaseConfig::new(DatabaseDriver::Sqlite, "sqlite::memory:", pool)?;
let state =
AppState::new()?.with_data_config(GatewayDataConfig::from_database_config(database))?;
let pending = state
.prepare_database_for_startup()
.await?
.expect("sqlite database should expose migration state");
assert!(
!pending.is_empty(),
"fresh sqlite gateway databases should report pending migrations"
);
assert!(
state.run_database_migrations().await?,
"sqlite gateway database should run migrations"
);
let pending = state
.prepare_database_for_startup()
.await?
.expect("sqlite database should expose migration state");
assert!(
pending.is_empty(),
"sqlite gateway databases should be current after migrations"
);
Ok(())
}
#[tokio::test]
async fn data_state_checks_user_uniqueness_through_user_reader() {
let user = StoredUserAuthRecord::new(
"user-1".to_string(),
Some("alice@example.com".to_string()),
true,
"alice".to_string(),
Some("hash".to_string()),
"user".to_string(),
"local".to_string(),
None,
None,
None,
true,
false,
None,
None,
)
.expect("auth user should build");
let admin = StoredUserAuthRecord::new(
"admin-1".to_string(),
Some("admin@example.com".to_string()),
true,
"admin".to_string(),
Some(format!("$2b$12${}", "a".repeat(53))),
"admin".to_string(),
"local".to_string(),
None,
None,
None,
true,
false,
None,
None,
)
.expect("admin user should build");
let state = GatewayDataState::with_user_reader_for_tests(Arc::new(
InMemoryUserReadRepository::seed_auth_users(vec![user, admin]),
));
assert!(state
.is_other_user_auth_email_taken("alice@example.com", "other-user")
.await
.expect("email uniqueness should check"));
assert!(!state
.is_other_user_auth_email_taken("alice@example.com", "user-1")
.await
.expect("same user email should not be taken"));
assert!(!state
.is_other_user_auth_email_taken("alice", "other-user")
.await
.expect("email lookup should not match username"));
assert!(state
.is_other_user_auth_username_taken("alice", "other-user")
.await
.expect("username uniqueness should check"));
assert_eq!(
state
.count_active_admin_users()
.await
.expect("active admin count should check"),
1
);
assert_eq!(
state
.count_active_local_admin_users_with_valid_password()
.await
.expect("valid local admin count should check"),
1
);
let preferences = StoredUserPreferenceRecord {
user_id: "user-1".to_string(),
avatar_url: Some("https://example.test/avatar.png".to_string()),
bio: Some("hello".to_string()),
default_provider_id: None,
default_provider_name: None,
theme: "dark".to_string(),
language: "en-US".to_string(),
timezone: "UTC".to_string(),
email_notifications: false,
usage_alerts: true,
announcement_notifications: false,
};
assert_eq!(
state
.write_user_preferences(&preferences)
.await
.expect("preferences should write through repository"),
Some(preferences.clone())
);
assert_eq!(
state
.read_user_preferences("user-1")
.await
.expect("preferences should read through repository"),
Some(preferences)
);
}
#[tokio::test]
async fn data_state_finds_active_provider_name_through_catalog_reader() {
let active = StoredProviderCatalogProvider::new(
"provider-1".to_string(),
"Provider One".to_string(),
None,
"openai".to_string(),
)
.expect("provider should build");
let inactive = StoredProviderCatalogProvider::new(
"provider-2".to_string(),
"Provider Two".to_string(),
None,
"openai".to_string(),
)
.expect("provider should build")
.with_transport_fields(false, false, false, None, None, None, None, None, None);
let state = GatewayDataState::with_provider_catalog_reader_for_tests(Arc::new(
InMemoryProviderCatalogReadRepository::seed(vec![active, inactive], Vec::new(), Vec::new()),
));
assert_eq!(
state
.find_active_provider_name("provider-1")
.await
.expect("provider lookup should succeed"),
Some("Provider One".to_string())
);
assert_eq!(
state
.find_active_provider_name("provider-2")
.await
.expect("inactive provider lookup should succeed"),
None
);
}
fn sample_auth_snapshot(api_key_id: &str, user_id: &str) -> StoredAuthApiKeySnapshot {
StoredAuthApiKeySnapshot::new(
user_id.to_string(),

View File

@@ -1,5 +1,5 @@
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339};
use crate::handlers::admin::shared::query_param_value;
use crate::GatewayError;
use axum::{
body::{Body, Bytes},
@@ -9,7 +9,6 @@ use axum::{
};
use regex::Regex;
use serde_json::json;
use sqlx::Row;
const ADMIN_BILLING_DATA_UNAVAILABLE_DETAIL: &str = "Admin billing data unavailable";
@@ -185,20 +184,6 @@ fn admin_billing_validate_safe_expression(expression: &str) -> Result<(), String
Ok(())
}
fn admin_billing_optional_epoch_value(
row: &sqlx::postgres::PgRow,
field: &str,
) -> Result<Option<String>, GatewayError> {
let value = row
.try_get::<Option<i64>, _>(field)
.map_err(|err| GatewayError::Internal(err.to_string()))?;
match value {
None => Ok(None),
Some(value) if value < 0 => Ok(None),
Some(value) => Ok(unix_secs_to_rfc3339(value as u64)),
}
}
pub(crate) async fn maybe_build_local_admin_billing_response(
state: &AdminAppState<'_>,
request_context: &AdminRequestContext<'_>,

View File

@@ -1,7 +1,6 @@
use super::{
build_admin_payment_callback_payload, build_admin_payment_callback_payload_from_record,
build_admin_payments_bad_request_response, parse_admin_payments_limit,
parse_admin_payments_offset,
build_admin_payment_callback_payload_from_record, build_admin_payments_bad_request_response,
parse_admin_payments_limit, parse_admin_payments_offset,
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::query_param_value;

View File

@@ -12,13 +12,13 @@ mod shared;
use self::shared::{
admin_payment_operator_id, admin_payment_order_id_from_detail_path,
admin_payment_order_id_from_suffix_path, build_admin_payment_callback_payload,
build_admin_payment_callback_payload_from_record, build_admin_payment_order_not_found_response,
build_admin_payment_order_payload, build_admin_payment_orders_page_response,
build_admin_payments_backend_unavailable_response, build_admin_payments_bad_request_response,
build_admin_payments_data_unavailable_response, normalize_admin_payment_currency,
normalize_admin_payment_optional_string, normalize_admin_payment_positive_number,
parse_admin_payments_limit, parse_admin_payments_offset, AdminPaymentOrderCreditRequest,
admin_payment_order_id_from_suffix_path, build_admin_payment_callback_payload_from_record,
build_admin_payment_order_not_found_response, build_admin_payment_order_payload,
build_admin_payment_orders_page_response, build_admin_payments_backend_unavailable_response,
build_admin_payments_bad_request_response, build_admin_payments_data_unavailable_response,
normalize_admin_payment_currency, normalize_admin_payment_optional_string,
normalize_admin_payment_positive_number, parse_admin_payments_limit,
parse_admin_payments_offset, AdminPaymentOrderCreditRequest,
};
pub(crate) async fn maybe_build_local_admin_payments_response(

View File

@@ -1,6 +1,6 @@
use crate::handlers::admin::request::AdminRequestContext;
use crate::handlers::admin::shared::{query_param_value, unix_secs_to_rfc3339};
use crate::{GatewayAdminPaymentCallbackView, GatewayError};
use crate::GatewayAdminPaymentCallbackView;
use axum::{
body::Body,
http,
@@ -8,7 +8,6 @@ use axum::{
Json,
};
use serde_json::json;
use sqlx::Row;
const ADMIN_PAYMENTS_DATA_UNAVAILABLE_DETAIL: &str = "Admin payments data unavailable";
@@ -220,34 +219,6 @@ pub(super) fn build_admin_payment_order_payload(
})
}
pub(super) fn build_admin_payment_callback_payload(
row: &sqlx::postgres::PgRow,
) -> Result<serde_json::Value, GatewayError> {
Ok(json!({
"id": row.try_get::<String, _>("id").map_err(|err| GatewayError::Internal(err.to_string()))?,
"payment_order_id": row.try_get::<Option<String>, _>("payment_order_id").map_err(|err| GatewayError::Internal(err.to_string()))?,
"payment_method": row.try_get::<String, _>("payment_method").map_err(|err| GatewayError::Internal(err.to_string()))?,
"callback_key": row.try_get::<String, _>("callback_key").map_err(|err| GatewayError::Internal(err.to_string()))?,
"order_no": row.try_get::<Option<String>, _>("order_no").map_err(|err| GatewayError::Internal(err.to_string()))?,
"gateway_order_id": row.try_get::<Option<String>, _>("gateway_order_id").map_err(|err| GatewayError::Internal(err.to_string()))?,
"payload_hash": row.try_get::<Option<String>, _>("payload_hash").map_err(|err| GatewayError::Internal(err.to_string()))?,
"signature_valid": row.try_get::<bool, _>("signature_valid").map_err(|err| GatewayError::Internal(err.to_string()))?,
"status": row.try_get::<String, _>("status").map_err(|err| GatewayError::Internal(err.to_string()))?,
"payload": row.try_get::<Option<serde_json::Value>, _>("payload").map_err(|err| GatewayError::Internal(err.to_string()))?,
"error_message": row.try_get::<Option<String>, _>("error_message").map_err(|err| GatewayError::Internal(err.to_string()))?,
"created_at": row
.try_get::<Option<i64>, _>("created_at_unix_ms")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339),
"processed_at": row
.try_get::<Option<i64>, _>("processed_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339),
}))
}
pub(super) fn build_admin_payment_callback_payload_from_record(
record: &GatewayAdminPaymentCallbackView,
) -> serde_json::Value {

View File

@@ -61,7 +61,7 @@ pub(in super::super) async fn build_admin_wallet_adjust_response(
));
}
let operator_id = admin_wallet_operator_id(request_context);
let has_postgres = state.has_postgres_pool();
let has_wallet_writer = state.has_wallet_data_writer();
let Some((wallet, transaction)) = state
.admin_adjust_wallet_balance(
&wallet_id,
@@ -72,7 +72,7 @@ pub(in super::super) async fn build_admin_wallet_adjust_response(
)
.await?
else {
return if has_postgres {
return if has_wallet_writer {
Ok(build_admin_wallet_not_found_response())
} else {
Ok(build_admin_wallets_data_unavailable_response())

View File

@@ -61,7 +61,7 @@ pub(in super::super) async fn build_admin_wallet_recharge_response(
));
}
let operator_id = admin_wallet_operator_id(request_context);
let has_postgres = state.has_postgres_pool();
let has_wallet_writer = state.has_wallet_data_writer();
let Some((wallet, payment_order)) = state
.admin_create_manual_wallet_recharge(
&wallet_id,
@@ -72,7 +72,7 @@ pub(in super::super) async fn build_admin_wallet_recharge_response(
)
.await?
else {
return if has_postgres {
return if has_wallet_writer {
Ok(build_admin_wallet_not_found_response())
} else {
Ok(build_admin_wallets_data_unavailable_response())

View File

@@ -1,8 +1,6 @@
use super::requests::ADMIN_WALLETS_API_KEY_GIFT_ADJUST_DETAIL;
use crate::handlers::admin::request::AdminRequestContext;
use crate::handlers::admin::shared::query_param_value;
use crate::GatewayError;
use sqlx::Row;
pub(in super::super) fn admin_wallet_operator_id(
request_context: &AdminRequestContext<'_>,
@@ -198,17 +196,6 @@ pub(in super::super) fn parse_admin_wallets_owner_type_filter(
}
}
pub(in super::super) fn optional_epoch_value(
row: &sqlx::postgres::PgRow,
key: &str,
) -> Result<Option<String>, GatewayError> {
Ok(row
.try_get::<Option<i64>, _>(key)
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(crate::handlers::admin::shared::unix_secs_to_rfc3339))
}
pub(in super::super) fn admin_wallet_build_order_no(now: chrono::DateTime<chrono::Utc>) -> String {
format!(
"po_{}_{}",

View File

@@ -6,7 +6,6 @@ use super::route_filters::{
};
use crate::constants::INTERNAL_GATEWAY_PATH_PREFIXES;
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::query::monitoring as monitoring_query;
use crate::GatewayError;
use aether_admin::observability::monitoring::{
admin_monitoring_bad_request_response, admin_monitoring_user_behavior_user_id_from_path,
@@ -41,33 +40,21 @@ pub(super) async fn build_admin_monitoring_audit_logs_response(
Err(detail) => return Ok(admin_monitoring_bad_request_response(detail)),
};
let Some(pool) = state.postgres_pool() else {
return Ok(build_admin_monitoring_audit_logs_payload_response(
Vec::new(),
0,
limit,
offset,
username,
event_type,
days,
));
};
let cutoff_time = chrono::Utc::now() - chrono::Duration::days(days);
let username_pattern = username
.as_deref()
.map(admin_monitoring_escape_like_pattern)
.map(|value| format!("%{value}%"));
let (items, total) = monitoring_query::list_admin_audit_logs(
&pool,
cutoff_time,
username_pattern.as_deref(),
event_type.as_deref(),
limit,
offset,
)
.await?;
let (items, total) = state
.list_admin_audit_logs(
cutoff_time,
username_pattern.as_deref(),
event_type.as_deref(),
limit,
offset,
)
.await?;
Ok(build_admin_monitoring_audit_logs_payload_response(
items, total, limit, offset, username, event_type, days,
@@ -85,14 +72,8 @@ pub(super) async fn build_admin_monitoring_suspicious_activities_response(
Err(detail) => return Ok(admin_monitoring_bad_request_response(detail)),
};
let Some(pool) = state.postgres_pool() else {
return Ok(
build_admin_monitoring_suspicious_activities_payload_response(Vec::new(), hours),
);
};
let cutoff_time = chrono::Utc::now() - chrono::Duration::hours(hours);
let activities = monitoring_query::list_admin_suspicious_activities(&pool, cutoff_time).await?;
let activities = state.list_admin_suspicious_activities(cutoff_time).await?;
Ok(build_admin_monitoring_suspicious_activities_payload_response(activities, hours))
}
@@ -112,22 +93,11 @@ pub(super) async fn build_admin_monitoring_user_behavior_response(
Err(detail) => return Ok(admin_monitoring_bad_request_response(detail)),
};
let Some(pool) = state.postgres_pool() else {
return Ok(build_admin_monitoring_user_behavior_payload_response(
user_id,
days,
std::collections::BTreeMap::new(),
0,
0,
0,
));
};
let cutoff_time = chrono::Utc::now() - chrono::Duration::days(days);
let event_counts =
monitoring_query::read_admin_user_behavior_event_counts(&pool, &user_id, cutoff_time)
.await?;
let event_counts = state
.read_admin_user_behavior_event_counts(&user_id, cutoff_time)
.await?;
let failed_requests = event_counts
.get("request_failed")

View File

@@ -23,7 +23,7 @@ async fn count_admin_monitoring_cache_affinity_entries(state: &AdminAppState<'_>
}
async fn scan_admin_monitoring_namespaced_keys(
runner: &aether_data::redis::RedisKvRunner,
runner: &aether_data::driver::redis::RedisKvRunner,
pattern: &str,
) -> Result<Vec<String>, GatewayError> {
let mut connection = runner

View File

@@ -1,7 +1,4 @@
use crate::handlers::admin::request::AdminAppState;
use crate::query::usage_heatmap::{
list_usage_heatmap_aggregate_rows, read_stats_daily_cutoff_date,
};
use crate::GatewayError;
use aether_admin::observability::stats::round_to;
use aether_admin::observability::usage::{
@@ -89,49 +86,15 @@ pub(super) async fn build_admin_usage_heatmap_response(
async fn build_admin_heatmap_summaries(
state: &AdminAppState<'_>,
created_from_unix_secs: u64,
start_date: chrono::NaiveDate,
today: chrono::NaiveDate,
_start_date: chrono::NaiveDate,
_today: chrono::NaiveDate,
) -> Result<Vec<StoredUsageDailySummary>, GatewayError> {
let query = UsageDailyHeatmapQuery {
created_from_unix_secs,
user_id: None,
admin_mode: true,
};
let Some(pool) = state.app().postgres_pool() else {
return state.summarize_usage_daily_heatmap(&query).await;
};
let Some(cutoff_date) = read_stats_daily_cutoff_date(&pool).await? else {
return state.summarize_usage_daily_heatmap(&query).await;
};
let cutoff_day = cutoff_date.date_naive().min(today);
let mut summaries =
list_usage_heatmap_aggregate_rows(&pool, start_date, cutoff_day, None).await?;
let raw_start_date = start_date.max(cutoff_day);
if raw_start_date <= today {
let raw_start_of_day = raw_start_date
.and_hms_opt(0, 0, 0)
.expect("heatmap day start should be valid");
let raw_created_from_unix_secs = u64::try_from(
chrono::DateTime::<chrono::Utc>::from_naive_utc_and_offset(
raw_start_of_day,
chrono::Utc,
)
.timestamp(),
)
.unwrap_or_default();
summaries.extend(
state
.summarize_usage_daily_heatmap(&UsageDailyHeatmapQuery {
created_from_unix_secs: raw_created_from_unix_secs,
user_id: None,
admin_mode: true,
})
.await?,
);
}
let mut summaries = state.summarize_usage_daily_heatmap(&query).await?;
summaries.sort_by(|left, right| left.date.cmp(&right.date));
Ok(summaries)
}

View File

@@ -3,7 +3,9 @@ use crate::handlers::admin::provider::shared::support::{
ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_KEYS, ADMIN_PROVIDER_MAPPING_PREVIEW_MAX_MODELS,
};
use crate::handlers::admin::request::AdminAppState;
use crate::handlers::admin::shared::{decrypt_catalog_secret_with_fallbacks, json_string_list};
use crate::handlers::admin::shared::{
decrypt_catalog_secret_with_fallbacks, json_string_list, take_secret_prefix, take_secret_suffix,
};
use crate::handlers::public::matches_model_mapping_for_models;
use crate::{GatewayError, LocalProviderDeleteTaskState};
use aether_data_contracts::repository::global_models::{
@@ -175,14 +177,15 @@ pub(crate) fn mapping_preview_masked_catalog_api_key(
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
.map(|value| {
if value.len() > 8 {
let char_count = value.chars().count();
if char_count > 8 {
format!(
"{}***{}",
&value[..4],
&value[value.len().saturating_sub(4)..]
take_secret_prefix(&value, 4),
take_secret_suffix(&value, 4)
)
} else if value.len() >= 2 {
format!("{}***", &value[..2])
} else if char_count >= 2 {
format!("{}***", take_secret_prefix(&value, 2))
} else {
"***".to_string()
}

View File

@@ -1,4 +1,4 @@
use aether_data::redis::RedisKeyspace;
use aether_data::driver::redis::RedisKeyspace;
pub(super) fn pool_sticky_pattern(keyspace: &RedisKeyspace, provider_id: &str) -> String {
keyspace.key(&format!("ap:{provider_id}:sticky:*"))

View File

@@ -7,7 +7,7 @@ use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolRuntimeState, ADMIN_PROVIDER_POOL_SCAN_BATCH,
};
use crate::GatewayError;
use aether_data::redis::RedisKvRunner;
use aether_data::driver::redis::RedisKvRunner;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};
use tracing::warn;

View File

@@ -5,7 +5,7 @@ use super::keys::{
use crate::handlers::admin::provider::shared::support::{
AdminProviderPoolConfig, AdminProviderPoolUnschedulableRule,
};
use aether_data::redis::RedisKvRunner;
use aether_data::driver::redis::RedisKvRunner;
use regex::Regex;
use std::collections::BTreeMap;
use std::time::{SystemTime, UNIX_EPOCH};

View File

@@ -116,7 +116,7 @@ impl<'a> AdminAppState<'a> {
self.app.mark_provider_key_rpm_reset(key_id, now_unix_secs)
}
pub(crate) fn redis_kv_runner(&self) -> Option<aether_data::redis::RedisKvRunner> {
pub(crate) fn redis_kv_runner(&self) -> Option<aether_data::driver::redis::RedisKvRunner> {
self.app.redis_kv_runner()
}
@@ -128,8 +128,8 @@ impl<'a> AdminAppState<'a> {
self.app.provider_key_rpm_reset_at(key_id, now_unix_secs)
}
pub(crate) fn has_postgres_pool(&self) -> bool {
self.app.postgres_pool().is_some()
pub(crate) fn has_wallet_data_writer(&self) -> bool {
self.app.has_wallet_data_writer()
}
pub(crate) fn mark_admin_monitoring_error_stats_reset(&self, now_unix_secs: u64) {

View File

@@ -12,6 +12,6 @@ pub(crate) use crate::handlers::shared::{
masked_catalog_api_key, normalize_json_array, normalize_json_object, normalize_string_list,
parse_catalog_auth_config_json, provider_catalog_key_supports_format,
provider_key_health_summary, provider_key_status_snapshot_payload, query_param_bool,
query_param_optional_bool, query_param_value, unix_secs_to_rfc3339,
OFFICIAL_EXTERNAL_MODEL_PROVIDERS,
query_param_optional_bool, query_param_value, take_secret_prefix, take_secret_suffix,
unix_secs_to_rfc3339, OFFICIAL_EXTERNAL_MODEL_PROVIDERS,
};

View File

@@ -4,7 +4,6 @@ use super::support::{
};
use crate::handlers::admin::request::{AdminAppState, AdminRequestContext};
use crate::handlers::admin::shared::{query_param_optional_bool, query_param_value};
use crate::query::user_rollups::list_user_usage_totals_from_stats_summary;
use crate::GatewayError;
use axum::{
body::Body,
@@ -43,19 +42,10 @@ pub(in super::super) async fn build_admin_list_users_response(
.iter()
.map(|row| row.id.clone())
.collect::<Vec<_>>();
let usage_totals_future = async {
let Some(pool) = state.app().postgres_pool() else {
return state.summarize_usage_totals_by_user_ids(&user_ids).await;
};
match list_user_usage_totals_from_stats_summary(&pool, &user_ids).await? {
Some(items) => Ok(items),
None => state.summarize_usage_totals_by_user_ids(&user_ids).await,
}
};
let (auth_rows_result, wallet_rows_result, usage_totals_result) = tokio::join!(
state.list_user_auth_by_ids(&user_ids),
state.list_wallet_snapshots_by_user_ids(&user_ids),
usage_totals_future,
state.summarize_usage_totals_by_user_ids(&user_ids),
);
let auth_by_user_id = auth_rows_result?
.into_iter()

View File

@@ -71,7 +71,7 @@ use self::support_test_connection::maybe_build_local_test_connection_response;
use self::support_user_me::maybe_build_local_users_me_response;
use self::support_wallet::{
maybe_build_local_wallet_response, sanitize_wallet_gateway_response,
wallet_normalize_optional_string_field, wallet_payment_order_payload_from_row,
wallet_normalize_optional_string_field,
};
pub(crate) fn build_unhandled_public_support_response(

View File

@@ -18,18 +18,6 @@ use chrono::Datelike;
use serde_json::json;
use std::collections::{BTreeMap, BTreeSet};
use crate::query::dashboard_stats::{
list_admin_dashboard_daily_model_aggregates, list_admin_dashboard_daily_provider_aggregates,
list_admin_dashboard_daily_totals_aggregates, list_admin_dashboard_hourly_model_aggregates,
list_admin_dashboard_hourly_provider_aggregates, list_admin_dashboard_hourly_totals_aggregates,
list_user_dashboard_daily_model_aggregates, list_user_dashboard_daily_totals_aggregates,
list_user_dashboard_hourly_model_aggregates, list_user_dashboard_hourly_totals_aggregates,
read_stats_hourly_cutoff, summarize_dashboard_usage_from_daily_aggregates,
DashboardDailyModelAggregateRow, DashboardDailyProviderAggregateRow,
DashboardDailyTotalsAggregateRow,
};
use crate::query::usage_heatmap::read_stats_daily_cutoff_date;
#[derive(Debug, Clone, Copy)]
struct DashboardDateRange {
start_date: chrono::NaiveDate,
@@ -438,93 +426,6 @@ fn dashboard_range_bounds_unix(range: DashboardDateRange) -> Option<(u64, u64)>
Some((start_utc.max(0) as u64, end_utc.max(0) as u64))
}
fn dashboard_range_bounds_utc(
range: DashboardDateRange,
) -> Option<(chrono::DateTime<chrono::Utc>, chrono::DateTime<chrono::Utc>)> {
let (created_from_unix_secs, created_until_unix_secs) = dashboard_range_bounds_unix(range)?;
let start_utc =
chrono::DateTime::<chrono::Utc>::from_timestamp(created_from_unix_secs as i64, 0)?;
let end_utc =
chrono::DateTime::<chrono::Utc>::from_timestamp(created_until_unix_secs as i64, 0)?;
Some((start_utc, end_utc))
}
fn dashboard_local_day_bounds_utc_exclusive(
range: DashboardDateRange,
) -> Option<(chrono::DateTime<chrono::Utc>, chrono::DateTime<chrono::Utc>)> {
let offset = chrono::Duration::minutes(i64::from(range.tz_offset_minutes));
let start_local = range.start_date.and_hms_opt(0, 0, 0)?;
let end_exclusive_local = range
.end_date
.checked_add_signed(chrono::Duration::days(1))?
.and_hms_opt(0, 0, 0)?;
let start_utc = chrono::DateTime::<chrono::Utc>::from_naive_utc_and_offset(
start_local.checked_sub_signed(offset)?,
chrono::Utc,
);
let end_exclusive_utc = chrono::DateTime::<chrono::Utc>::from_naive_utc_and_offset(
end_exclusive_local.checked_sub_signed(offset)?,
chrono::Utc,
);
Some((start_utc, end_exclusive_utc))
}
fn dashboard_utc_midnight(value: chrono::DateTime<chrono::Utc>) -> chrono::DateTime<chrono::Utc> {
chrono::DateTime::<chrono::Utc>::from_naive_utc_and_offset(
value
.date_naive()
.and_hms_opt(0, 0, 0)
.expect("midnight should be valid"),
chrono::Utc,
)
}
fn dashboard_next_utc_midnight(
value: chrono::DateTime<chrono::Utc>,
) -> chrono::DateTime<chrono::Utc> {
let midnight = dashboard_utc_midnight(value);
if value == midnight {
midnight
} else {
midnight + chrono::Duration::days(1)
}
}
fn dashboard_unix_secs(value: chrono::DateTime<chrono::Utc>) -> u64 {
value.timestamp().max(0) as u64
}
fn dashboard_absorb_dashboard_summary(
target: &mut StoredUsageDashboardSummary,
part: &StoredUsageDashboardSummary,
) {
target.total_requests = target.total_requests.saturating_add(part.total_requests);
target.input_tokens = target.input_tokens.saturating_add(part.input_tokens);
target.effective_input_tokens = target
.effective_input_tokens
.saturating_add(part.effective_input_tokens);
target.output_tokens = target.output_tokens.saturating_add(part.output_tokens);
target.total_tokens = target.total_tokens.saturating_add(part.total_tokens);
target.cache_creation_tokens = target
.cache_creation_tokens
.saturating_add(part.cache_creation_tokens);
target.cache_read_tokens = target
.cache_read_tokens
.saturating_add(part.cache_read_tokens);
target.total_input_context = target
.total_input_context
.saturating_add(part.total_input_context);
target.cache_creation_cost_usd += part.cache_creation_cost_usd;
target.cache_read_cost_usd += part.cache_read_cost_usd;
target.total_cost_usd += part.total_cost_usd;
target.actual_total_cost_usd += part.actual_total_cost_usd;
target.error_requests = target.error_requests.saturating_add(part.error_requests);
target.response_time_sum_ms += part.response_time_sum_ms;
target.response_time_samples = target
.response_time_samples
.saturating_add(part.response_time_samples);
}
async fn dashboard_summary_for_unix_range_raw(
state: &AppState,
created_from_unix_secs: u64,
@@ -607,87 +508,7 @@ async fn dashboard_summary_for_range(
user_id: Option<&str>,
error_context: &str,
) -> Result<StoredUsageDashboardSummary, Response<Body>> {
let Some(pool) = state.postgres_pool() else {
return dashboard_summary_for_range_raw(state, range, user_id, error_context).await;
};
let cutoff_date = match read_stats_daily_cutoff_date(&pool).await {
Ok(value) => value,
Err(err) => {
return Err(build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
))
}
};
let Some(cutoff_date) = cutoff_date else {
return dashboard_summary_for_range_raw(state, range, user_id, error_context).await;
};
let Some((start_utc, end_utc)) = dashboard_range_bounds_utc(range) else {
return Err(build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: invalid time range"),
false,
));
};
let aggregate_start = dashboard_next_utc_midnight(start_utc);
let aggregate_end = dashboard_utc_midnight(end_utc).min(cutoff_date);
let mut summary = StoredUsageDashboardSummary::default();
if aggregate_start < aggregate_end {
let leading_end = aggregate_start.min(end_utc);
if start_utc < leading_end {
let raw = dashboard_summary_for_unix_range_raw(
state,
dashboard_unix_secs(start_utc),
dashboard_unix_secs(leading_end),
user_id,
error_context,
)
.await?;
dashboard_absorb_dashboard_summary(&mut summary, &raw);
}
let aggregate = summarize_dashboard_usage_from_daily_aggregates(
&pool,
aggregate_start,
aggregate_end,
user_id,
)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
dashboard_absorb_dashboard_summary(&mut summary, &aggregate);
if aggregate_end < end_utc {
let raw = dashboard_summary_for_unix_range_raw(
state,
dashboard_unix_secs(aggregate_end),
dashboard_unix_secs(end_utc),
user_id,
error_context,
)
.await?;
dashboard_absorb_dashboard_summary(&mut summary, &raw);
}
} else if start_utc < end_utc {
let raw = dashboard_summary_for_unix_range_raw(
state,
dashboard_unix_secs(start_utc),
dashboard_unix_secs(end_utc),
user_id,
error_context,
)
.await?;
dashboard_absorb_dashboard_summary(&mut summary, &raw);
}
Ok(summary)
dashboard_summary_for_range_raw(state, range, user_id, error_context).await
}
async fn dashboard_daily_breakdown_for_range(
@@ -753,75 +574,6 @@ fn dashboard_apply_daily_breakdown_rows(
}
}
fn dashboard_record_daily_totals_aggregate(
by_date: &mut std::collections::BTreeMap<chrono::NaiveDate, DashboardDailyAggregate>,
row: &DashboardDailyTotalsAggregateRow,
) {
let Ok(date) = chrono::NaiveDate::parse_from_str(&row.date, "%Y-%m-%d") else {
return;
};
let aggregate = by_date.entry(date).or_default();
aggregate.totals.requests = aggregate.totals.requests.saturating_add(row.requests);
aggregate.totals.total_tokens = aggregate
.totals
.total_tokens
.saturating_add(row.total_tokens);
aggregate.totals.total_cost_usd += row.total_cost_usd;
aggregate.totals.response_time_sum_ms += row.response_time_sum_ms;
aggregate.totals.response_time_samples = aggregate
.totals
.response_time_samples
.saturating_add(row.response_time_samples);
}
fn dashboard_record_daily_model_aggregate(
by_date: &mut std::collections::BTreeMap<chrono::NaiveDate, DashboardDailyAggregate>,
model_summary: &mut std::collections::BTreeMap<String, DashboardModelAggregate>,
row: &DashboardDailyModelAggregateRow,
) {
let Ok(date) = chrono::NaiveDate::parse_from_str(&row.date, "%Y-%m-%d") else {
return;
};
let aggregate = by_date.entry(date).or_default();
let model = aggregate.models.entry(row.model.clone()).or_default();
model.requests = model.requests.saturating_add(row.requests);
model.tokens = model.tokens.saturating_add(row.total_tokens);
model.cost += row.total_cost_usd;
model.response_time_sum_ms += row.response_time_sum_ms;
model.response_time_samples = model
.response_time_samples
.saturating_add(row.response_time_samples);
let summary = model_summary.entry(row.model.clone()).or_default();
summary.requests = summary.requests.saturating_add(row.requests);
summary.tokens = summary.tokens.saturating_add(row.total_tokens);
summary.cost += row.total_cost_usd;
summary.response_time_sum_ms += row.response_time_sum_ms;
summary.response_time_samples = summary
.response_time_samples
.saturating_add(row.response_time_samples);
}
fn dashboard_record_daily_provider_aggregate(
by_date: &mut std::collections::BTreeMap<chrono::NaiveDate, DashboardDailyAggregate>,
provider_summary: &mut std::collections::BTreeMap<String, DashboardProviderAggregate>,
row: &DashboardDailyProviderAggregateRow,
) {
let Ok(date) = chrono::NaiveDate::parse_from_str(&row.date, "%Y-%m-%d") else {
return;
};
let aggregate = by_date.entry(date).or_default();
let provider = aggregate.providers.entry(row.provider.clone()).or_default();
provider.requests = provider.requests.saturating_add(row.requests);
provider.tokens = provider.tokens.saturating_add(row.total_tokens);
provider.cost += row.total_cost_usd;
let summary = provider_summary.entry(row.provider.clone()).or_default();
summary.requests = summary.requests.saturating_add(row.requests);
summary.tokens = summary.tokens.saturating_add(row.total_tokens);
summary.cost += row.total_cost_usd;
}
fn dashboard_build_daily_stats_payload(
range: DashboardDateRange,
is_admin: bool,
@@ -972,477 +724,6 @@ fn dashboard_build_daily_stats_payload(
payload
}
fn dashboard_range_supports_hourly_rollup(range: DashboardDateRange) -> bool {
range.tz_offset_minutes % 60 == 0
}
async fn dashboard_admin_hourly_stats_aggregate_payload(
state: &AppState,
range: DashboardDateRange,
error_context: &str,
) -> Result<Option<serde_json::Value>, Response<Body>> {
if !dashboard_range_supports_hourly_rollup(range) {
return Ok(None);
}
let Some(pool) = state.postgres_pool() else {
return Ok(None);
};
let cutoff_utc = match read_stats_hourly_cutoff(&pool).await {
Ok(value) => value,
Err(err) => {
return Err(build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
))
}
};
let Some(cutoff_utc) = cutoff_utc else {
return Ok(None);
};
let Some((range_start_utc, range_end_exclusive_utc)) =
dashboard_local_day_bounds_utc_exclusive(range)
else {
return Ok(None);
};
let aggregate_end_utc = range_end_exclusive_utc.min(cutoff_utc);
if range_start_utc >= aggregate_end_utc {
return Ok(None);
}
let daily_totals = list_admin_dashboard_hourly_totals_aggregates(
&pool,
range_start_utc,
aggregate_end_utc,
range.tz_offset_minutes,
)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let daily_models = list_admin_dashboard_hourly_model_aggregates(
&pool,
range_start_utc,
aggregate_end_utc,
range.tz_offset_minutes,
)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let daily_providers = list_admin_dashboard_hourly_provider_aggregates(
&pool,
range_start_utc,
aggregate_end_utc,
range.tz_offset_minutes,
)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let mut by_date =
std::collections::BTreeMap::<chrono::NaiveDate, DashboardDailyAggregate>::new();
let mut model_summary = std::collections::BTreeMap::<String, DashboardModelAggregate>::new();
let mut provider_summary =
std::collections::BTreeMap::<String, DashboardProviderAggregate>::new();
for row in &daily_totals {
dashboard_record_daily_totals_aggregate(&mut by_date, row);
}
for row in &daily_models {
dashboard_record_daily_model_aggregate(&mut by_date, &mut model_summary, row);
}
for row in &daily_providers {
dashboard_record_daily_provider_aggregate(&mut by_date, &mut provider_summary, row);
}
if aggregate_end_utc < range_end_exclusive_utc {
let raw_rows = dashboard_daily_breakdown_for_unix_range_raw(
state,
dashboard_unix_secs(aggregate_end_utc),
dashboard_unix_secs(range_end_exclusive_utc),
range.tz_offset_minutes,
None,
error_context,
)
.await?;
dashboard_apply_daily_breakdown_rows(
&raw_rows,
&mut by_date,
&mut model_summary,
&mut provider_summary,
);
}
Ok(Some(dashboard_build_daily_stats_payload(
range,
true,
&by_date,
&model_summary,
&provider_summary,
)))
}
async fn dashboard_user_hourly_stats_aggregate_payload(
state: &AppState,
range: DashboardDateRange,
user_id: &str,
error_context: &str,
) -> Result<Option<serde_json::Value>, Response<Body>> {
if !dashboard_range_supports_hourly_rollup(range) {
return Ok(None);
}
let Some(pool) = state.postgres_pool() else {
return Ok(None);
};
let cutoff_utc = match read_stats_hourly_cutoff(&pool).await {
Ok(value) => value,
Err(err) => {
return Err(build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
))
}
};
let Some(cutoff_utc) = cutoff_utc else {
return Ok(None);
};
let Some((range_start_utc, range_end_exclusive_utc)) =
dashboard_local_day_bounds_utc_exclusive(range)
else {
return Ok(None);
};
let aggregate_end_utc = range_end_exclusive_utc.min(cutoff_utc);
if range_start_utc >= aggregate_end_utc {
return Ok(None);
}
let daily_totals = list_user_dashboard_hourly_totals_aggregates(
&pool,
range_start_utc,
aggregate_end_utc,
range.tz_offset_minutes,
user_id,
)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let daily_models = list_user_dashboard_hourly_model_aggregates(
&pool,
range_start_utc,
aggregate_end_utc,
range.tz_offset_minutes,
user_id,
)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let mut by_date =
std::collections::BTreeMap::<chrono::NaiveDate, DashboardDailyAggregate>::new();
let mut model_summary = std::collections::BTreeMap::<String, DashboardModelAggregate>::new();
let mut provider_summary =
std::collections::BTreeMap::<String, DashboardProviderAggregate>::new();
for row in &daily_totals {
dashboard_record_daily_totals_aggregate(&mut by_date, row);
}
for row in &daily_models {
dashboard_record_daily_model_aggregate(&mut by_date, &mut model_summary, row);
}
if aggregate_end_utc < range_end_exclusive_utc {
let raw_rows = dashboard_daily_breakdown_for_unix_range_raw(
state,
dashboard_unix_secs(aggregate_end_utc),
dashboard_unix_secs(range_end_exclusive_utc),
range.tz_offset_minutes,
Some(user_id),
error_context,
)
.await?;
dashboard_apply_daily_breakdown_rows(
&raw_rows,
&mut by_date,
&mut model_summary,
&mut provider_summary,
);
}
Ok(Some(dashboard_build_daily_stats_payload(
range,
false,
&by_date,
&model_summary,
&provider_summary,
)))
}
async fn dashboard_admin_daily_stats_aggregate_payload(
state: &AppState,
range: DashboardDateRange,
error_context: &str,
) -> Result<Option<serde_json::Value>, Response<Body>> {
if range.tz_offset_minutes != 0 {
return dashboard_admin_hourly_stats_aggregate_payload(state, range, error_context).await;
}
let Some(pool) = state.postgres_pool() else {
return Ok(None);
};
let cutoff_date = match read_stats_daily_cutoff_date(&pool).await {
Ok(value) => value,
Err(err) => {
return Err(build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
))
}
};
let Some(cutoff_date) = cutoff_date else {
return Ok(None);
};
let Some(range_end_exclusive) = range.end_date.checked_add_signed(chrono::Duration::days(1))
else {
return Ok(None);
};
let aggregate_end_exclusive = range_end_exclusive.min(cutoff_date.date_naive());
if range.start_date >= aggregate_end_exclusive {
return Ok(None);
}
let aggregate_start_utc = chrono::DateTime::<chrono::Utc>::from_naive_utc_and_offset(
range
.start_date
.and_hms_opt(0, 0, 0)
.expect("midnight should be valid"),
chrono::Utc,
);
let aggregate_end_utc = chrono::DateTime::<chrono::Utc>::from_naive_utc_and_offset(
aggregate_end_exclusive
.and_hms_opt(0, 0, 0)
.expect("midnight should be valid"),
chrono::Utc,
);
let daily_totals =
list_admin_dashboard_daily_totals_aggregates(&pool, aggregate_start_utc, aggregate_end_utc)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let daily_models =
list_admin_dashboard_daily_model_aggregates(&pool, aggregate_start_utc, aggregate_end_utc)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let daily_providers = list_admin_dashboard_daily_provider_aggregates(
&pool,
aggregate_start_utc,
aggregate_end_utc,
)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let mut by_date =
std::collections::BTreeMap::<chrono::NaiveDate, DashboardDailyAggregate>::new();
let mut model_summary = std::collections::BTreeMap::<String, DashboardModelAggregate>::new();
let mut provider_summary =
std::collections::BTreeMap::<String, DashboardProviderAggregate>::new();
for row in &daily_totals {
dashboard_record_daily_totals_aggregate(&mut by_date, row);
}
for row in &daily_models {
dashboard_record_daily_model_aggregate(&mut by_date, &mut model_summary, row);
}
for row in &daily_providers {
dashboard_record_daily_provider_aggregate(&mut by_date, &mut provider_summary, row);
}
if aggregate_end_exclusive <= range.end_date {
let raw_range = DashboardDateRange {
start_date: aggregate_end_exclusive,
end_date: range.end_date,
tz_offset_minutes: range.tz_offset_minutes,
};
let raw_rows =
dashboard_daily_breakdown_for_range(state, raw_range, None, error_context).await?;
dashboard_apply_daily_breakdown_rows(
&raw_rows,
&mut by_date,
&mut model_summary,
&mut provider_summary,
);
}
Ok(Some(dashboard_build_daily_stats_payload(
range,
true,
&by_date,
&model_summary,
&provider_summary,
)))
}
async fn dashboard_user_daily_stats_aggregate_payload(
state: &AppState,
range: DashboardDateRange,
user_id: &str,
error_context: &str,
) -> Result<Option<serde_json::Value>, Response<Body>> {
if range.tz_offset_minutes != 0 {
return dashboard_user_hourly_stats_aggregate_payload(state, range, user_id, error_context)
.await;
}
let Some(pool) = state.postgres_pool() else {
return Ok(None);
};
let cutoff_date = match read_stats_daily_cutoff_date(&pool).await {
Ok(value) => value,
Err(err) => {
return Err(build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
))
}
};
let Some(cutoff_date) = cutoff_date else {
return Ok(None);
};
let Some(range_end_exclusive) = range.end_date.checked_add_signed(chrono::Duration::days(1))
else {
return Ok(None);
};
let aggregate_end_exclusive = range_end_exclusive.min(cutoff_date.date_naive());
if range.start_date >= aggregate_end_exclusive {
return Ok(None);
}
let aggregate_start_utc = chrono::DateTime::<chrono::Utc>::from_naive_utc_and_offset(
range
.start_date
.and_hms_opt(0, 0, 0)
.expect("midnight should be valid"),
chrono::Utc,
);
let aggregate_end_utc = chrono::DateTime::<chrono::Utc>::from_naive_utc_and_offset(
aggregate_end_exclusive
.and_hms_opt(0, 0, 0)
.expect("midnight should be valid"),
chrono::Utc,
);
let daily_totals = list_user_dashboard_daily_totals_aggregates(
&pool,
aggregate_start_utc,
aggregate_end_utc,
user_id,
)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let daily_models = list_user_dashboard_daily_model_aggregates(
&pool,
aggregate_start_utc,
aggregate_end_utc,
user_id,
)
.await
.map_err(|err| {
build_auth_error_response(
http::StatusCode::INTERNAL_SERVER_ERROR,
format!("{error_context}: {err:?}"),
false,
)
})?;
let mut by_date =
std::collections::BTreeMap::<chrono::NaiveDate, DashboardDailyAggregate>::new();
let mut model_summary = std::collections::BTreeMap::<String, DashboardModelAggregate>::new();
let mut provider_summary =
std::collections::BTreeMap::<String, DashboardProviderAggregate>::new();
for row in &daily_totals {
dashboard_record_daily_totals_aggregate(&mut by_date, row);
}
for row in &daily_models {
dashboard_record_daily_model_aggregate(&mut by_date, &mut model_summary, row);
}
if aggregate_end_exclusive <= range.end_date {
let raw_range = DashboardDateRange {
start_date: aggregate_end_exclusive,
end_date: range.end_date,
tz_offset_minutes: range.tz_offset_minutes,
};
let raw_rows =
dashboard_daily_breakdown_for_range(state, raw_range, Some(user_id), error_context)
.await?;
dashboard_apply_daily_breakdown_rows(
&raw_rows,
&mut by_date,
&mut model_summary,
&mut provider_summary,
);
}
Ok(Some(dashboard_build_daily_stats_payload(
range,
false,
&by_date,
&model_summary,
&provider_summary,
)))
}
fn dashboard_usage_totals_from_summary(
summary: &StoredUsageDashboardSummary,
) -> DashboardUsageTotals {
@@ -1888,37 +1169,6 @@ pub(super) async fn handle_dashboard_daily_stats_get(
Err(detail) => return dashboard_bad_request_response(detail),
};
let user_filter = (!is_admin).then_some(auth.user.id.as_str());
if is_admin {
match dashboard_admin_daily_stats_aggregate_payload(
state,
range,
"dashboard daily stats lookup failed",
)
.await
{
Ok(Some(payload)) => {
return dashboard_cached_json_response(state, cache_key, cache_ttl, &payload)
}
Ok(None) => {}
Err(response) => return response,
}
} else {
match dashboard_user_daily_stats_aggregate_payload(
state,
range,
auth.user.id.as_str(),
"dashboard daily stats lookup failed",
)
.await
{
Ok(Some(payload)) => {
return dashboard_cached_json_response(state, cache_key, cache_ttl, &payload)
}
Ok(None) => {}
Err(response) => return response,
}
}
let usage = match dashboard_daily_breakdown_for_range(
state,
range,

View File

@@ -8,7 +8,6 @@ use chrono::Utc;
use serde_json::{json, Value};
use crate::handlers::shared::query_param_value;
use crate::query::monitoring as monitoring_query;
use super::{
build_auth_error_response, resolve_authenticated_local_user, AppState,
@@ -112,27 +111,16 @@ pub(super) async fn handle_user_audit_logs(
}
};
let Some(pool) = state.postgres_pool() else {
return build_user_monitoring_audit_logs_payload(
Vec::new(),
0,
let cutoff_time = Utc::now() - chrono::Duration::days(days);
let (items, total) = match state
.list_user_audit_logs(
&auth.user.id,
cutoff_time,
event_type.as_deref(),
limit,
offset,
event_type,
days,
);
};
let cutoff_time = Utc::now() - chrono::Duration::days(days);
let (items, total) = match monitoring_query::list_user_audit_logs(
&pool,
&auth.user.id,
cutoff_time,
event_type.as_deref(),
limit,
offset,
)
.await
)
.await
{
Ok(value) => value,
Err(err) => {

View File

@@ -4,8 +4,8 @@ pub(super) use super::{build_auth_error_response, AppState, GatewayPublicRequest
#[path = "payment/gateway.rs"]
pub(super) mod payment_gateway;
#[path = "payment/postgres.rs"]
mod payment_postgres;
#[path = "payment/repository.rs"]
mod payment_repository;
#[path = "payment/route.rs"]
mod payment_route;
#[path = "payment/shared.rs"]
@@ -14,7 +14,7 @@ mod payment_shared;
#[path = "payment/test_support.rs"]
mod payment_test_support;
use self::payment_postgres::handle_payment_callback_with_postgres;
use self::payment_repository::handle_payment_callback_with_wallet_repository;
use self::payment_shared::NormalizedPaymentCallbackRequest;
const PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL: &str = "支付回调存储暂不可用";

View File

@@ -11,14 +11,14 @@ use super::{
GatewayPublicRequestContext,
};
pub(super) async fn handle_payment_callback_with_postgres(
pub(super) async fn handle_payment_callback_with_wallet_repository(
state: &AppState,
payment_method: &str,
request_context: &GatewayPublicRequestContext,
payload: &NormalizedPaymentCallbackRequest,
signature_valid: bool,
) -> Response<Body> {
if state.postgres_pool().is_none() {
if !state.has_database_wallet_data_writer() {
return build_payment_callback_storage_unavailable_response();
}
@@ -128,7 +128,7 @@ pub(super) async fn handle_payment_callback_with_postgres(
#[cfg(test)]
mod tests {
use super::{
handle_payment_callback_with_postgres, AppState, NormalizedPaymentCallbackRequest,
handle_payment_callback_with_wallet_repository, AppState, NormalizedPaymentCallbackRequest,
};
use crate::control::GatewayPublicRequestContext;
use crate::handlers::public::support::support_payment::PAYMENT_CALLBACK_STORAGE_UNAVAILABLE_DETAIL;
@@ -137,10 +137,10 @@ mod tests {
use serde_json::json;
#[tokio::test]
async fn payment_callback_postgres_handler_returns_explicit_503_without_pool() {
async fn payment_callback_repository_handler_returns_explicit_503_without_wallet_writer() {
let state = AppState::new().expect("state should build");
let request_context = GatewayPublicRequestContext::from_request_parts(
"trace-payment-callback-postgres-missing",
"trace-payment-callback-wallet-writer-missing",
&Method::POST,
&"/api/payment/callback/alipay"
.parse::<Uri>()
@@ -159,7 +159,7 @@ mod tests {
payload: json!({ "status": "paid" }),
};
let response = handle_payment_callback_with_postgres(
let response = handle_payment_callback_with_wallet_repository(
&state,
"alipay",
&request_context,

View File

@@ -7,7 +7,7 @@ use super::payment_shared::{
};
use super::{
build_auth_error_response, build_payment_callback_storage_unavailable_response,
handle_payment_callback_with_postgres, AppState, GatewayPublicRequestContext,
handle_payment_callback_with_wallet_repository, AppState, GatewayPublicRequestContext,
};
pub(super) async fn maybe_build_local_payment_callback_route_response(
@@ -104,9 +104,9 @@ pub(super) async fn maybe_build_local_payment_callback_route_response(
}
};
if state.postgres_pool().is_some() {
if state.has_database_wallet_data_writer() {
return Some(
handle_payment_callback_with_postgres(
handle_payment_callback_with_wallet_repository(
state,
&payment_method,
request_context,

View File

@@ -18,9 +18,6 @@ use axum::{
use chrono::Utc;
use serde_json::json;
use crate::query::usage_heatmap::{
list_usage_heatmap_aggregate_rows, read_stats_daily_cutoff_date,
};
use crate::GatewayError;
use super::{
@@ -1204,8 +1201,8 @@ pub(super) async fn handle_users_me_usage_heatmap_get(
async fn build_usage_heatmap_summaries(
state: &AppState,
created_from_unix_secs: u64,
start_date: chrono::NaiveDate,
today: chrono::NaiveDate,
_start_date: chrono::NaiveDate,
_today: chrono::NaiveDate,
user_id: Option<&str>,
) -> Result<Vec<StoredUsageDailySummary>, GatewayError> {
let query = aether_data_contracts::repository::usage::UsageDailyHeatmapQuery {
@@ -1213,43 +1210,7 @@ async fn build_usage_heatmap_summaries(
user_id: user_id.map(ToOwned::to_owned),
admin_mode: user_id.is_none(),
};
let Some(pool) = state.postgres_pool() else {
return state.summarize_usage_daily_heatmap(&query).await;
};
let Some(cutoff_date) = read_stats_daily_cutoff_date(&pool).await? else {
return state.summarize_usage_daily_heatmap(&query).await;
};
let cutoff_day = cutoff_date.date_naive().min(today);
let mut summaries =
list_usage_heatmap_aggregate_rows(&pool, start_date, cutoff_day, user_id).await?;
let raw_start_date = start_date.max(cutoff_day);
if raw_start_date <= today {
let raw_start_of_day = raw_start_date
.and_hms_opt(0, 0, 0)
.expect("heatmap day start should be valid");
let raw_created_from_unix_secs = u64::try_from(
chrono::DateTime::<chrono::Utc>::from_naive_utc_and_offset(
raw_start_of_day,
chrono::Utc,
)
.timestamp(),
)
.unwrap_or_default();
summaries.extend(
state
.summarize_usage_daily_heatmap(
&aether_data_contracts::repository::usage::UsageDailyHeatmapQuery {
created_from_unix_secs: raw_created_from_unix_secs,
user_id: user_id.map(ToOwned::to_owned),
admin_mode: user_id.is_none(),
},
)
.await?,
);
}
let mut summaries = state.summarize_usage_daily_heatmap(&query).await?;
summaries.sort_by(|left, right| left.date.cmp(&right.date));
Ok(summaries)
}

View File

@@ -34,13 +34,11 @@ use self::reads::{
parse_wallet_limit, parse_wallet_offset, wallet_fixed_offset, wallet_today_billing_date_string,
wallet_transaction_payload_from_record,
};
pub(crate) use self::recharge::sanitize_wallet_gateway_response;
use self::recharge::{
handle_wallet_create_recharge, handle_wallet_recharge_detail, handle_wallet_recharge_list,
wallet_recharge_detail_path_matches,
};
pub(crate) use self::recharge::{
sanitize_wallet_gateway_response, wallet_payment_order_payload_from_row,
};
use self::redeem::handle_wallet_redeem;
use self::refunds::{
handle_wallet_create_refund, handle_wallet_refund_detail, handle_wallet_refunds_list,

View File

@@ -5,8 +5,8 @@ use super::{
build_auth_error_response, build_auth_json_response, build_wallet_payload,
build_wallet_recharge_storage_unavailable_response, http, parse_wallet_limit,
parse_wallet_offset, resolve_authenticated_local_user, unix_secs_to_rfc3339,
wallet_normalize_optional_string_field, AppState, Body, GatewayError,
GatewayPublicRequestContext, Response, WALLET_SAFE_GATEWAY_RESPONSE_KEYS,
wallet_normalize_optional_string_field, AppState, Body, GatewayPublicRequestContext, Response,
WALLET_SAFE_GATEWAY_RESPONSE_KEYS,
};
#[cfg(test)]
use super::{
@@ -16,7 +16,6 @@ use super::{
use chrono::Utc;
use serde::Deserialize;
use serde_json::json;
use sqlx::Row;
use uuid::Uuid;
#[derive(Debug, Deserialize)]
@@ -152,65 +151,6 @@ fn build_wallet_payment_order_payload(
})
}
pub(crate) fn wallet_payment_order_payload_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<serde_json::Value, GatewayError> {
let created_at = row
.try_get::<Option<i64>, _>("created_at_unix_ms")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339);
let paid_at = row
.try_get::<Option<i64>, _>("paid_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339);
let credited_at = row
.try_get::<Option<i64>, _>("credited_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339);
let expires_at = row
.try_get::<Option<i64>, _>("expires_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339);
Ok(build_wallet_payment_order_payload(
row.try_get::<String, _>("id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<String, _>("order_no")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<String, _>("wallet_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<Option<String>, _>("user_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<f64, _>("amount_usd")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<Option<f64>, _>("pay_amount")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<Option<String>, _>("pay_currency")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<Option<f64>, _>("exchange_rate")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<f64, _>("refunded_amount_usd")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<f64, _>("refundable_amount_usd")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<String, _>("payment_method")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<Option<String>, _>("gateway_order_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<Option<serde_json::Value>, _>("gateway_response")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get::<String, _>("effective_status")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
created_at,
paid_at,
credited_at,
expires_at,
))
}
fn wallet_payment_order_payload_from_record(
record: &aether_data::repository::wallet::StoredAdminPaymentOrder,
) -> serde_json::Value {
@@ -285,7 +225,7 @@ pub(super) async fn handle_wallet_create_recharge(
}
};
if state.postgres_pool().is_none() {
if !state.has_database_wallet_data_writer() {
#[cfg(test)]
{
let Some(wallet) = wallet else {
@@ -485,11 +425,12 @@ pub(super) async fn handle_wallet_recharge_list(
}
};
#[cfg(test)]
let (items, total) = if state.postgres_pool().is_none() && items.is_empty() && total == 0 {
wallet_test_recharge_orders_for_user(&auth.user.id, limit, offset)
} else {
(items, total)
};
let (items, total) =
if !state.has_database_wallet_data_writer() && items.is_empty() && total == 0 {
wallet_test_recharge_orders_for_user(&auth.user.id, limit, offset)
} else {
(items, total)
};
let mut payload = json!({
"items": items,

View File

@@ -2,8 +2,7 @@ use super::{
build_auth_error_response, build_auth_json_response, build_wallet_payload,
build_wallet_refund_storage_unavailable_response, http, parse_wallet_limit,
parse_wallet_offset, resolve_authenticated_local_user, unix_secs_to_rfc3339,
wallet_normalize_optional_string_field, AppState, Body, GatewayError,
GatewayPublicRequestContext, Response,
wallet_normalize_optional_string_field, AppState, Body, GatewayPublicRequestContext, Response,
};
#[cfg(test)]
use super::{
@@ -13,7 +12,6 @@ use super::{
use chrono::Utc;
use serde::Deserialize;
use serde_json::json;
use sqlx::Row;
use uuid::Uuid;
#[derive(Debug, Deserialize)]
@@ -94,51 +92,6 @@ pub(super) fn wallet_refund_detail_path_matches(request_path: &str) -> bool {
wallet_refund_id_from_path(request_path).is_some()
}
fn wallet_refund_payload_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<serde_json::Value, GatewayError> {
let created_at = row
.try_get::<Option<i64>, _>("created_at_unix_ms")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339);
let updated_at = row
.try_get::<Option<i64>, _>("updated_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339);
let processed_at = row
.try_get::<Option<i64>, _>("processed_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339);
let completed_at = row
.try_get::<Option<i64>, _>("completed_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.and_then(|value| u64::try_from(value).ok())
.and_then(unix_secs_to_rfc3339);
Ok(json!({
"id": row.try_get::<String, _>("id").map_err(|err| GatewayError::Internal(err.to_string()))?,
"refund_no": row.try_get::<String, _>("refund_no").map_err(|err| GatewayError::Internal(err.to_string()))?,
"payment_order_id": row.try_get::<Option<String>, _>("payment_order_id").map_err(|err| GatewayError::Internal(err.to_string()))?,
"source_type": row.try_get::<String, _>("source_type").map_err(|err| GatewayError::Internal(err.to_string()))?,
"source_id": row.try_get::<Option<String>, _>("source_id").map_err(|err| GatewayError::Internal(err.to_string()))?,
"refund_mode": row.try_get::<String, _>("refund_mode").map_err(|err| GatewayError::Internal(err.to_string()))?,
"amount_usd": row.try_get::<f64, _>("amount_usd").map_err(|err| GatewayError::Internal(err.to_string()))?,
"status": row.try_get::<String, _>("status").map_err(|err| GatewayError::Internal(err.to_string()))?,
"reason": row.try_get::<Option<String>, _>("reason").map_err(|err| GatewayError::Internal(err.to_string()))?,
"failure_reason": row.try_get::<Option<String>, _>("failure_reason").map_err(|err| GatewayError::Internal(err.to_string()))?,
"gateway_refund_id": row.try_get::<Option<String>, _>("gateway_refund_id").map_err(|err| GatewayError::Internal(err.to_string()))?,
"payout_method": row.try_get::<Option<String>, _>("payout_method").map_err(|err| GatewayError::Internal(err.to_string()))?,
"payout_reference": row.try_get::<Option<String>, _>("payout_reference").map_err(|err| GatewayError::Internal(err.to_string()))?,
"payout_proof": row.try_get::<Option<serde_json::Value>, _>("payout_proof").map_err(|err| GatewayError::Internal(err.to_string()))?,
"created_at": created_at,
"updated_at": updated_at,
"processed_at": processed_at,
"completed_at": completed_at,
}))
}
fn wallet_refund_payload_from_record(
record: &aether_data::repository::wallet::StoredAdminWalletRefund,
) -> serde_json::Value {
@@ -255,18 +208,19 @@ pub(super) async fn handle_wallet_refunds_list(
})
.collect::<Vec<_>>();
#[cfg(test)]
let (items, total) = if state.postgres_pool().is_none() && items.is_empty() && total == 0 {
let all_items = wallet_test_refunds_for_wallet(&wallet.id);
let total = all_items.len() as u64;
let items = all_items
.into_iter()
.skip(offset)
.take(limit)
.collect::<Vec<_>>();
(items, total)
} else {
(items, total)
};
let (items, total) =
if !state.has_database_wallet_data_writer() && items.is_empty() && total == 0 {
let all_items = wallet_test_refunds_for_wallet(&wallet.id);
let total = all_items.len() as u64;
let items = all_items
.into_iter()
.skip(offset)
.take(limit)
.collect::<Vec<_>>();
(items, total)
} else {
(items, total)
};
let mut payload = json!({
"items": items,
@@ -394,7 +348,7 @@ pub(super) async fn handle_wallet_create_refund(
);
};
if state.postgres_pool().is_none() {
if !state.has_database_wallet_data_writer() {
#[cfg(test)]
{
if let Some(idempotency_key) = payload.idempotency_key.as_deref() {

View File

@@ -102,6 +102,29 @@ pub(crate) fn encrypt_catalog_secret_with_fallbacks(
encrypt_python_fernet_plaintext(encryption_key.as_ref(), plaintext).ok()
}
pub(crate) fn take_secret_prefix(value: &str, prefix_chars: usize) -> &str {
let end = value
.char_indices()
.nth(prefix_chars)
.map(|(index, _)| index)
.unwrap_or(value.len());
&value[..end]
}
pub(crate) fn take_secret_suffix(value: &str, suffix_chars: usize) -> &str {
if suffix_chars == 0 {
return &value[value.len()..];
}
let start = value
.char_indices()
.rev()
.nth(suffix_chars - 1)
.map(|(index, _)| index)
.unwrap_or(0);
&value[start..]
}
pub(crate) fn masked_catalog_api_key(state: &AppState, key: &StoredProviderCatalogKey) -> String {
match key.auth_type.trim() {
"service_account" | "vertex_ai" => "[Service Account]".to_string(),
@@ -117,13 +140,13 @@ pub(crate) fn masked_catalog_api_key(state: &AppState, key: &StoredProviderCatal
};
decrypt_catalog_secret_with_fallbacks(state.encryption_key(), ciphertext)
.map(|value| {
if value.len() <= 12 {
if value.chars().count() <= 12 {
format!("{value}***")
} else {
format!(
"{}***{}",
&value[..8],
&value[value.len().saturating_sub(4)..]
take_secret_prefix(&value, 8),
take_secret_suffix(&value, 4)
)
}
})
@@ -1621,6 +1644,39 @@ mod tests {
.expect("key transport should build")
}
#[test]
fn masked_catalog_api_key_handles_unicode_plaintext_without_panicking() {
let state = AppState::new().expect("gateway should build");
let encrypted_api_key =
encrypt_python_fernet_plaintext(DEVELOPMENT_ENCRYPTION_KEY, "测试-密钥-1234567890")
.expect("api key ciphertext should build");
let key = StoredProviderCatalogKey::new(
"key-unicode".to_string(),
"provider-test".to_string(),
"default".to_string(),
"api_key".to_string(),
None,
true,
)
.expect("key should build")
.with_transport_fields(
Some(json!(["openai:chat"])),
encrypted_api_key,
None,
None,
None,
None,
None,
None,
None,
)
.expect("key transport should build");
let masked = masked_catalog_api_key(&state, &key);
assert!(masked.contains("***"));
assert_ne!(masked, "***ERROR***");
}
#[test]
fn provider_key_status_snapshot_payload_backfills_missing_quota_from_upstream_metadata() {
let mut key = sample_catalog_key();

View File

@@ -25,7 +25,7 @@ pub(crate) use self::catalog::{
encrypt_catalog_secret_with_fallbacks, masked_catalog_api_key, parse_catalog_auth_config_json,
provider_catalog_key_supports_format, provider_key_health_summary,
provider_key_status_snapshot_payload, sync_provider_key_oauth_status_snapshot,
sync_provider_key_quota_status_snapshot,
sync_provider_key_quota_status_snapshot, take_secret_prefix, take_secret_suffix,
};
pub(crate) use self::email_templates::{
admin_email_template_definition, admin_email_template_html_key,

View File

@@ -51,7 +51,6 @@ mod oauth;
mod orchestration;
mod provider_key_auth;
pub(crate) use aether_provider_transport as provider_transport;
mod query;
mod rate_limit;
mod request_candidate_runtime;
mod router;

View File

@@ -2,12 +2,15 @@
#[global_allocator]
static GLOBAL: tikv_jemallocator::Jemalloc = tikv_jemallocator::Jemalloc;
use clap::{Args as ClapArgs, Parser, ValueEnum};
use std::path::PathBuf;
use clap::{Args as ClapArgs, Parser, Subcommand, ValueEnum};
use tracing::{debug, info, warn};
use aether_crypto::warm_python_fernet_secret;
use aether_data::postgres::PostgresPoolConfig;
use aether_data::redis::RedisClientConfig;
use aether_data::driver::redis::RedisClientConfig;
use aether_data::lifecycle::export::{export_database_jsonl, import_database_jsonl, ExportDomain};
use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig, DEFAULT_SQLITE_DATABASE_URL};
use aether_gateway::{
attach_static_frontend, build_router_with_state, set_gateway_frontdoor_app_port, AppState,
FrontdoorCorsConfig, FrontdoorUserRpmConfig, GatewayDataConfig, UsageRuntimeConfig,
@@ -50,6 +53,56 @@ impl DeploymentTopologyArg {
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)]
enum DatabaseDriverArg {
Sqlite,
Mysql,
Postgres,
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)]
enum ExportDomainArg {
Users,
ApiKeys,
Providers,
ProviderKeys,
Endpoints,
Models,
GlobalModels,
SystemConfigs,
Wallets,
Usage,
Billing,
}
impl From<ExportDomainArg> for ExportDomain {
fn from(value: ExportDomainArg) -> Self {
match value {
ExportDomainArg::Users => ExportDomain::Users,
ExportDomainArg::ApiKeys => ExportDomain::ApiKeys,
ExportDomainArg::Providers => ExportDomain::Providers,
ExportDomainArg::ProviderKeys => ExportDomain::ProviderKeys,
ExportDomainArg::Endpoints => ExportDomain::Endpoints,
ExportDomainArg::Models => ExportDomain::Models,
ExportDomainArg::GlobalModels => ExportDomain::GlobalModels,
ExportDomainArg::SystemConfigs => ExportDomain::SystemConfigs,
ExportDomainArg::Wallets => ExportDomain::Wallets,
ExportDomainArg::Usage => ExportDomain::Usage,
ExportDomainArg::Billing => ExportDomain::Billing,
}
}
}
impl From<DatabaseDriverArg> for DatabaseDriver {
fn from(value: DatabaseDriverArg) -> Self {
match value {
DatabaseDriverArg::Sqlite => DatabaseDriver::Sqlite,
DatabaseDriverArg::Mysql => DatabaseDriver::Mysql,
DatabaseDriverArg::Postgres => DatabaseDriver::Postgres,
}
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)]
enum NodeRoleArg {
All,
@@ -71,6 +124,21 @@ impl NodeRoleArg {
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)]
enum RuntimeBackendArg {
Redis,
Memory,
}
impl RuntimeBackendArg {
const fn as_str(self) -> &'static str {
match self {
Self::Redis => "redis",
Self::Memory => "memory",
}
}
}
#[derive(Copy, Clone, Debug, Eq, PartialEq, ValueEnum)]
enum GatewayLogFormatArg {
Pretty,
@@ -148,6 +216,12 @@ fn env_var_trimmed(name: &str) -> Option<String> {
#[derive(ClapArgs, Debug, Clone)]
struct GatewayDataArgs {
#[arg(long, env = "AETHER_DATABASE_DRIVER")]
database_driver: Option<DatabaseDriverArg>,
#[arg(long, env = "AETHER_DATABASE_URL")]
database_url: Option<String>,
#[arg(long, env = "AETHER_GATEWAY_DATA_POSTGRES_URL")]
postgres_url: Option<String>,
@@ -211,6 +285,55 @@ struct GatewayDataArgs {
}
impl GatewayDataArgs {
fn effective_database_driver(&self) -> Option<DatabaseDriver> {
self.database_driver.map(Into::into).or_else(|| {
self.database_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(DatabaseDriver::from_database_url)
})
}
fn effective_database_url(&self) -> Option<String> {
let configured_url = self
.database_url
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty())
.map(ToOwned::to_owned);
match (self.effective_database_driver(), configured_url) {
(Some(DatabaseDriver::Sqlite), None) => Some(DEFAULT_SQLITE_DATABASE_URL.to_string()),
(_, Some(url)) => Some(url),
(None, None) => self.effective_postgres_url(),
(Some(DatabaseDriver::Postgres), None) => self.effective_postgres_url(),
(Some(DatabaseDriver::Mysql), None) => None,
}
}
fn effective_sql_database_config(&self) -> Option<SqlDatabaseConfig> {
let url = self.effective_database_url()?;
let driver = self
.effective_database_driver()
.or_else(|| DatabaseDriver::from_database_url(&url))
.unwrap_or(DatabaseDriver::Postgres);
Some(SqlDatabaseConfig {
driver,
url,
pool: SqlPoolConfig {
min_connections: self.postgres_min_connections,
max_connections: self.postgres_max_connections,
acquire_timeout_ms: self.postgres_acquire_timeout_ms,
idle_timeout_ms: self.postgres_idle_timeout_ms,
max_lifetime_ms: self.postgres_max_lifetime_ms,
statement_cache_capacity: self.postgres_statement_cache_capacity,
require_ssl: driver != DatabaseDriver::Sqlite && self.postgres_require_ssl,
},
})
}
fn effective_postgres_url(&self) -> Option<String> {
self.postgres_url
.as_deref()
@@ -270,20 +393,11 @@ impl GatewayDataArgs {
}
fn to_config(&self) -> GatewayDataConfig {
let database_url = self.effective_postgres_url();
let database = self.effective_sql_database_config();
let redis_url = self.effective_redis_url();
let mut config = match database_url.as_deref() {
Some(database_url) => GatewayDataConfig::from_postgres_config(PostgresPoolConfig {
database_url: database_url.to_string(),
min_connections: self.postgres_min_connections,
max_connections: self.postgres_max_connections,
acquire_timeout_ms: self.postgres_acquire_timeout_ms,
idle_timeout_ms: self.postgres_idle_timeout_ms,
max_lifetime_ms: self.postgres_max_lifetime_ms,
statement_cache_capacity: self.postgres_statement_cache_capacity,
require_ssl: self.postgres_require_ssl,
}),
let mut config = match database {
Some(database) => GatewayDataConfig::from_database_config(database),
None => GatewayDataConfig::disabled(),
};
@@ -458,6 +572,35 @@ struct GatewayLoggingArgs {
log_max_files: usize,
}
#[derive(Subcommand, Debug, Clone)]
enum DataCommand {
/// Export persistent SQL data to database-neutral JSONL.
Export(DataExportArgs),
/// Import database-neutral JSONL into the selected SQL database.
Import(DataImportArgs),
}
#[derive(ClapArgs, Debug, Clone)]
struct DataExportArgs {
#[command(flatten)]
data: GatewayDataArgs,
#[arg(long)]
output: PathBuf,
#[arg(long, value_enum, value_delimiter = ',')]
domains: Vec<ExportDomainArg>,
}
#[derive(ClapArgs, Debug, Clone)]
struct DataImportArgs {
#[command(flatten)]
data: GatewayDataArgs,
#[arg(long)]
input: PathBuf,
}
impl GatewayLoggingArgs {
fn apply_to_runtime_config(
&self,
@@ -498,6 +641,9 @@ impl GatewayLoggingArgs {
about = "Phase 3a Rust ingress gateway for Aether"
)]
struct Args {
#[command(subcommand)]
command: Option<DataCommand>,
#[arg(long, env = "APP_PORT", default_value_t = 8084)]
app_port: u16,
@@ -605,6 +751,9 @@ struct Args {
)]
distributed_request_command_timeout_ms: u64,
#[arg(long, env = "AETHER_RUNTIME_BACKEND", value_enum)]
runtime_backend: Option<RuntimeBackendArg>,
#[command(flatten)]
data: GatewayDataArgs,
@@ -622,13 +771,37 @@ struct Args {
}
impl Args {
fn effective_runtime_backend(
&self,
database: Option<&SqlDatabaseConfig>,
data_redis_url: Option<&str>,
) -> RuntimeBackendArg {
if let Some(runtime_backend) = self.runtime_backend {
return runtime_backend;
}
if matches!(self.deployment_topology, DeploymentTopologyArg::MultiNode) {
return RuntimeBackendArg::Redis;
}
if database.is_some_and(|database| database.driver == DatabaseDriver::Sqlite) {
return RuntimeBackendArg::Memory;
}
if data_redis_url.is_some() {
RuntimeBackendArg::Redis
} else {
RuntimeBackendArg::Memory
}
}
fn runtime_config(&self) -> Result<ServiceRuntimeConfig, std::io::Error> {
let default_log_filter =
if self.migrate || self.apply_backfills || self.auto_prepare_database {
"aether_gateway=info,aether_data=info"
} else {
"aether_gateway=info"
};
let default_log_filter = if self.command.is_some()
|| self.migrate
|| self.apply_backfills
|| self.auto_prepare_database
{
"aether_gateway=info,aether_data=info"
} else {
"aether_gateway=info"
};
let config = self
.logging
.apply_to_runtime_config(ServiceRuntimeConfig::new(
@@ -691,15 +864,22 @@ async fn run_healthcheck(
fn validate_deployment_topology(
args: &Args,
data_postgres_url: Option<&str>,
database: Option<&SqlDatabaseConfig>,
data_redis_url: Option<&str>,
runtime_backend: RuntimeBackendArg,
) -> Result<(), std::io::Error> {
if matches!(args.deployment_topology, DeploymentTopologyArg::SingleNode) {
if data_postgres_url.is_none() && data_redis_url.is_none() {
if database.is_none() && data_redis_url.is_none() {
warn!(
"single-node deployment is starting without Postgres or Redis; local-only mode is allowed, but admin/auth/billing persistence will be limited"
"single-node deployment is starting without SQL database or Redis; local-only mode is allowed, but admin/auth/billing persistence will be limited"
);
}
if matches!(runtime_backend, RuntimeBackendArg::Redis) && data_redis_url.is_none() {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"AETHER_RUNTIME_BACKEND=redis requires REDIS_URL or AETHER_GATEWAY_DATA_REDIS_URL",
));
}
return Ok(());
}
@@ -711,8 +891,8 @@ fn validate_deployment_topology(
}
let mut missing = Vec::new();
if data_postgres_url.is_none() {
missing.push("DATABASE_URL or AETHER_GATEWAY_DATA_POSTGRES_URL");
if database.is_none() {
missing.push("AETHER_DATABASE_URL, DATABASE_URL, or AETHER_GATEWAY_DATA_POSTGRES_URL");
}
if data_redis_url.is_none() {
missing.push("REDIS_URL or AETHER_GATEWAY_DATA_REDIS_URL");
@@ -728,6 +908,20 @@ fn validate_deployment_topology(
));
}
if matches!(runtime_backend, RuntimeBackendArg::Memory) {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"AETHER_RUNTIME_BACKEND=memory is only valid for single-node deployment",
));
}
if database.is_some_and(|database| database.driver == DatabaseDriver::Sqlite) {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"AETHER_DATABASE_DRIVER=sqlite is only valid for single-node deployment",
));
}
if args
.video_task_store_path
.as_deref()
@@ -772,6 +966,10 @@ fn main() -> Result<(), Box<dyn std::error::Error>> {
async fn run() -> Result<(), Box<dyn std::error::Error>> {
let args = Args::parse();
if let Some(command) = args.command.as_ref() {
init_service_runtime(args.runtime_config()?)?;
return run_data_command(command).await;
}
if args.migrate {
init_service_runtime(args.runtime_config()?)?;
return run_explicit_migrations(&args).await;
@@ -787,12 +985,16 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
return run_healthcheck(app_port, args.healthcheck_timeout_ms).await;
}
init_service_runtime(args.runtime_config()?)?;
let sql_database_config = args.data.effective_sql_database_config();
let data_postgres_url = args.data.effective_postgres_url();
let data_redis_url = args.data.effective_redis_url();
let runtime_backend =
args.effective_runtime_backend(sql_database_config.as_ref(), data_redis_url.as_deref());
validate_deployment_topology(
&args,
data_postgres_url.as_deref(),
sql_database_config.as_ref(),
data_redis_url.as_deref(),
runtime_backend,
)?;
let data_config = args.data.to_config();
let rate_limit_config = if matches!(args.deployment_topology, DeploymentTopologyArg::MultiNode)
@@ -814,6 +1016,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
environment = %args.frontdoor.environment,
deployment_topology = args.deployment_topology.as_str(),
node_role = args.node_role.as_str(),
runtime_backend = runtime_backend.as_str(),
frontdoor_mode = "compatibility_frontdoor",
log_format = ?args.logging.log_format,
log_destination = args.logging.log_destination.as_str(),
@@ -844,6 +1047,11 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
.as_deref()
.or(data_redis_url.as_deref())
.is_some(),
data_database_configured = sql_database_config.is_some(),
data_database_driver = sql_database_config
.as_ref()
.map(|database| database.driver.as_str())
.unwrap_or("-"),
data_postgres_configured = data_postgres_url.is_some(),
data_redis_configured = data_redis_url.is_some(),
data_has_encryption_key = data_config.encryption_key().is_some(),
@@ -935,7 +1143,7 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
execution_runtime_configured = state.execution_runtime_configured(),
"aether-gateway data layer configured"
);
prepare_postgres_startup_requirements(&state, args.auto_prepare_database).await?;
prepare_database_startup_requirements(&state, args.auto_prepare_database).await?;
let reset_stale_proxy_nodes = state.reset_stale_proxy_node_tunnel_statuses().await?;
if reset_stale_proxy_nodes > 0 {
info!(
@@ -991,11 +1199,88 @@ async fn run() -> Result<(), Box<dyn std::error::Error>> {
Ok(())
}
async fn run_data_command(command: &DataCommand) -> Result<(), Box<dyn std::error::Error>> {
match command {
DataCommand::Export(args) => run_data_export(args).await,
DataCommand::Import(args) => run_data_import(args).await,
}
}
fn required_sql_database_config(
data: &GatewayDataArgs,
) -> Result<SqlDatabaseConfig, Box<dyn std::error::Error>> {
data.effective_sql_database_config().ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"AETHER_DATABASE_DRIVER/AETHER_DATABASE_URL, AETHER_GATEWAY_DATA_POSTGRES_URL, or DATABASE_URL is required",
)
.into()
})
}
fn requested_export_domains(args: &DataExportArgs) -> Vec<ExportDomain> {
args.domains
.iter()
.copied()
.map(Into::into)
.collect::<Vec<_>>()
}
fn current_unix_secs() -> Result<u64, std::time::SystemTimeError> {
Ok(std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)?
.as_secs())
}
async fn run_data_export(args: &DataExportArgs) -> Result<(), Box<dyn std::error::Error>> {
let database = required_sql_database_config(&args.data)?;
let driver = database.driver;
let domains = requested_export_domains(args);
let created_at_unix_secs = current_unix_secs()?;
let encoded = export_database_jsonl(database, domains, created_at_unix_secs).await?;
tokio::fs::write(&args.output, encoded.as_bytes()).await?;
info!(
driver = %driver,
output = %args.output.display(),
bytes = encoded.len(),
"database export complete"
);
println!(
"exported {} bytes from {} to {}",
encoded.len(),
driver,
args.output.display()
);
Ok(())
}
async fn run_data_import(args: &DataImportArgs) -> Result<(), Box<dyn std::error::Error>> {
let database = required_sql_database_config(&args.data)?;
let driver = database.driver;
let input = tokio::fs::read_to_string(&args.input).await?;
let imported = import_database_jsonl(database, &input).await?;
info!(
driver = %driver,
input = %args.input.display(),
imported,
"database import complete"
);
println!(
"imported {} records into {} from {}",
imported,
driver,
args.input.display()
);
Ok(())
}
async fn run_explicit_migrations(args: &Args) -> Result<(), Box<dyn std::error::Error>> {
if args.data.effective_postgres_url().is_none() {
if args.data.effective_sql_database_config().is_none() {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"AETHER_GATEWAY_DATA_POSTGRES_URL or DATABASE_URL is required when running --migrate",
"AETHER_DATABASE_DRIVER/AETHER_DATABASE_URL, AETHER_GATEWAY_DATA_POSTGRES_URL, or DATABASE_URL is required when running --migrate",
)
.into());
}
@@ -1008,7 +1293,7 @@ async fn run_explicit_migrations(args: &Args) -> Result<(), Box<dyn std::error::
let state = AppState::new()?.with_data_config(args.data.to_config())?;
let pending = state
.pending_postgres_migrations()
.pending_database_migrations()
.await?
.unwrap_or_default();
if pending.is_empty() {
@@ -1029,28 +1314,29 @@ async fn run_explicit_migrations(args: &Args) -> Result<(), Box<dyn std::error::
pending_versions = %format_pending_migrations(&pending),
"running database migrations by explicit request..."
);
if state.run_postgres_migrations().await? {
if state.run_database_migrations().await? {
info!("database migrations complete");
}
Ok(())
}
async fn run_explicit_backfills(args: &Args) -> Result<(), Box<dyn std::error::Error>> {
args.data.effective_postgres_url().ok_or_else(|| {
let database = args.data.effective_sql_database_config().ok_or_else(|| {
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"AETHER_GATEWAY_DATA_POSTGRES_URL or DATABASE_URL is required when running --apply-backfills",
"AETHER_DATABASE_DRIVER/AETHER_DATABASE_URL, AETHER_GATEWAY_DATA_POSTGRES_URL, or DATABASE_URL is required when running --apply-backfills",
)
})?;
let state = AppState::new()?.with_data_config(args.data.to_config())?;
ensure_postgres_schema_is_current(&state).await?;
ensure_database_schema_is_current(&state).await?;
let pending = state
.pending_postgres_backfills()
.pending_database_backfills()
.await?
.unwrap_or_default();
if pending.is_empty() {
info!(
driver = %database.driver,
pending_backfills = 0,
"database backfills already up to date"
);
@@ -1067,19 +1353,19 @@ async fn run_explicit_backfills(args: &Args) -> Result<(), Box<dyn std::error::E
pending_versions = %format_pending_backfills(&pending),
"running database backfills by explicit request..."
);
if state.run_postgres_backfills().await? {
if state.run_database_backfills().await? {
info!("database backfills complete");
}
Ok(())
}
async fn prepare_postgres_startup_requirements(
async fn prepare_database_startup_requirements(
state: &AppState,
auto_prepare_database: bool,
) -> Result<(), Box<dyn std::error::Error>> {
if !auto_prepare_database {
ensure_postgres_schema_is_current(state).await?;
ensure_postgres_backfills_are_current(state).await?;
ensure_database_schema_is_current(state).await?;
ensure_database_backfills_are_current(state).await?;
return Ok(());
}
@@ -1087,7 +1373,7 @@ async fn prepare_postgres_startup_requirements(
"auto database preparation enabled; applying pending migrations and backfills before serving traffic"
);
let Some(pending_migrations) = state.prepare_postgres_for_startup().await? else {
let Some(pending_migrations) = state.prepare_database_for_startup().await? else {
return Ok(());
};
if !pending_migrations.is_empty() {
@@ -1101,12 +1387,12 @@ async fn prepare_postgres_startup_requirements(
pending_versions = %format_pending_migrations(&pending_migrations),
"running database migrations during service startup..."
);
if state.run_postgres_migrations().await? {
if state.run_database_migrations().await? {
info!("database migrations complete during service startup");
}
}
let Some(pending_backfills) = state.pending_postgres_backfills().await? else {
let Some(pending_backfills) = state.pending_database_backfills().await? else {
return Ok(());
};
if pending_backfills.is_empty() {
@@ -1123,14 +1409,16 @@ async fn prepare_postgres_startup_requirements(
pending_versions = %format_pending_backfills(&pending_backfills),
"running database backfills during service startup..."
);
if state.run_postgres_backfills().await? {
if state.run_database_backfills().await? {
info!("database backfills complete during service startup");
}
Ok(())
}
fn format_pending_migrations(pending: &[aether_data::migrate::PendingMigrationInfo]) -> String {
fn format_pending_migrations(
pending: &[aether_data::lifecycle::migrate::PendingMigrationInfo],
) -> String {
pending
.iter()
.map(|migration| format!("{} ({})", migration.version, migration.description))
@@ -1138,7 +1426,9 @@ fn format_pending_migrations(pending: &[aether_data::migrate::PendingMigrationIn
.join(", ")
}
fn format_pending_backfills(pending: &[aether_data::backfill::PendingBackfillInfo]) -> String {
fn format_pending_backfills(
pending: &[aether_data::lifecycle::backfill::PendingBackfillInfo],
) -> String {
pending
.iter()
.map(|backfill| format!("{} ({})", backfill.version, backfill.description))
@@ -1146,10 +1436,10 @@ fn format_pending_backfills(pending: &[aether_data::backfill::PendingBackfillInf
.join(", ")
}
async fn ensure_postgres_backfills_are_current(
async fn ensure_database_backfills_are_current(
state: &AppState,
) -> Result<(), Box<dyn std::error::Error>> {
let Some(pending) = state.pending_postgres_backfills().await? else {
let Some(pending) = state.pending_database_backfills().await? else {
return Ok(());
};
if pending.is_empty() {
@@ -1162,10 +1452,10 @@ async fn ensure_postgres_backfills_are_current(
Err(pending_backfills_error(pending.len(), next.version, &next.description).into())
}
async fn ensure_postgres_schema_is_current(
async fn ensure_database_schema_is_current(
state: &AppState,
) -> Result<(), Box<dyn std::error::Error>> {
let Some(pending) = state.prepare_postgres_for_startup().await? else {
let Some(pending) = state.prepare_database_for_startup().await? else {
return Ok(());
};
if pending.is_empty() {
@@ -1207,16 +1497,19 @@ fn pending_backfills_error(
#[cfg(test)]
mod tests {
use super::{
ensure_postgres_backfills_are_current, ensure_postgres_schema_is_current,
ensure_database_backfills_are_current, ensure_database_schema_is_current,
pending_backfills_error, pending_schema_error, resolve_healthcheck_url, Args,
DeploymentTopologyArg, GatewayDataArgs, GatewayFrontdoorArgs, GatewayLogDestinationArg,
GatewayLogFormatArg, GatewayLogRotationArg, GatewayLoggingArgs, GatewayRateLimitArgs,
GatewayUsageArgs, NodeRoleArg, VideoTaskTruthSourceArg,
DatabaseDriverArg, DeploymentTopologyArg, GatewayDataArgs, GatewayFrontdoorArgs,
GatewayLogDestinationArg, GatewayLogFormatArg, GatewayLogRotationArg, GatewayLoggingArgs,
GatewayRateLimitArgs, GatewayUsageArgs, NodeRoleArg, RuntimeBackendArg,
VideoTaskTruthSourceArg,
};
use aether_data::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
use aether_gateway::AppState;
fn test_args() -> Args {
Args {
command: None,
app_port: 8084,
healthcheck: false,
healthcheck_timeout_ms: 3_000,
@@ -1237,7 +1530,10 @@ mod tests {
distributed_request_lease_ttl_ms: 30_000,
distributed_request_renew_interval_ms: 10_000,
distributed_request_command_timeout_ms: 1_000,
runtime_backend: None,
data: GatewayDataArgs {
database_driver: None,
database_url: None,
postgres_url: None,
encryption_key: None,
redis_url: None,
@@ -1337,6 +1633,166 @@ mod tests {
);
}
#[test]
fn sqlite_database_defaults_to_memory_runtime_backend() {
let args = test_args();
let database = SqlDatabaseConfig::new(
DatabaseDriver::Sqlite,
"sqlite://./data/aether.db".to_string(),
SqlPoolConfig::default(),
)
.expect("sqlite config should build");
assert_eq!(
args.effective_runtime_backend(Some(&database), Some("redis://127.0.0.1/0")),
RuntimeBackendArg::Memory
);
}
#[test]
fn redis_url_defaults_to_redis_runtime_backend_for_server_database() {
let args = test_args();
let database = SqlDatabaseConfig::new(
DatabaseDriver::Postgres,
"postgres://postgres:postgres@localhost/aether".to_string(),
SqlPoolConfig::default(),
)
.expect("postgres config should build");
assert_eq!(
args.effective_runtime_backend(Some(&database), Some("redis://127.0.0.1/0")),
RuntimeBackendArg::Redis
);
}
#[test]
fn mysql_database_with_redis_defaults_to_redis_runtime_backend() {
let args = test_args();
let database = SqlDatabaseConfig::new(
DatabaseDriver::Mysql,
"mysql://aether:aether@localhost:3306/aether".to_string(),
SqlPoolConfig::default(),
)
.expect("mysql config should build");
assert_eq!(
args.effective_runtime_backend(Some(&database), Some("redis://127.0.0.1/0")),
RuntimeBackendArg::Redis
);
}
#[test]
fn sqlite_database_allows_explicit_redis_runtime_backend_when_redis_is_configured() {
let mut args = test_args();
args.runtime_backend = Some(RuntimeBackendArg::Redis);
let database = SqlDatabaseConfig::new(
DatabaseDriver::Sqlite,
"sqlite://./data/aether.db".to_string(),
SqlPoolConfig::default(),
)
.expect("sqlite config should build");
assert_eq!(
args.effective_runtime_backend(Some(&database), Some("redis://127.0.0.1/0")),
RuntimeBackendArg::Redis
);
super::validate_deployment_topology(
&args,
Some(&database),
Some("redis://127.0.0.1/0"),
RuntimeBackendArg::Redis,
)
.expect("single-node sqlite should allow explicit redis runtime");
}
#[test]
fn single_node_sqlite_without_redis_allows_memory_runtime_backend() {
let args = test_args();
let database = SqlDatabaseConfig::new(
DatabaseDriver::Sqlite,
"sqlite://./data/aether.db".to_string(),
SqlPoolConfig::default(),
)
.expect("sqlite config should build");
super::validate_deployment_topology(
&args,
Some(&database),
None,
RuntimeBackendArg::Memory,
)
.expect("single-node sqlite memory runtime should be accepted");
}
#[test]
fn multi_node_rejects_memory_runtime_backend() {
let mut args = test_args();
args.deployment_topology = DeploymentTopologyArg::MultiNode;
args.node_role = NodeRoleArg::Frontdoor;
let database = SqlDatabaseConfig::new(
DatabaseDriver::Postgres,
"postgres://postgres:postgres@localhost/aether".to_string(),
SqlPoolConfig::default(),
)
.expect("postgres config should build");
let error = super::validate_deployment_topology(
&args,
Some(&database),
Some("redis://127.0.0.1/0"),
RuntimeBackendArg::Memory,
)
.expect_err("multi-node memory runtime should be rejected");
assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
assert!(error.to_string().contains("AETHER_RUNTIME_BACKEND=memory"));
}
#[test]
fn multi_node_rejects_missing_redis_runtime_backend() {
let mut args = test_args();
args.deployment_topology = DeploymentTopologyArg::MultiNode;
args.node_role = NodeRoleArg::Frontdoor;
let database = SqlDatabaseConfig::new(
DatabaseDriver::Postgres,
"postgres://postgres:postgres@localhost/aether".to_string(),
SqlPoolConfig::default(),
)
.expect("postgres config should build");
let error = super::validate_deployment_topology(
&args,
Some(&database),
None,
RuntimeBackendArg::Redis,
)
.expect_err("multi-node should require redis");
assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
assert!(error.to_string().contains("REDIS_URL"));
}
#[test]
fn multi_node_rejects_sqlite_database_backend() {
let mut args = test_args();
args.deployment_topology = DeploymentTopologyArg::MultiNode;
args.node_role = NodeRoleArg::Frontdoor;
let database = SqlDatabaseConfig::new(
DatabaseDriver::Sqlite,
"sqlite://./data/aether.db".to_string(),
SqlPoolConfig::default(),
)
.expect("sqlite config should build");
let error = super::validate_deployment_topology(
&args,
Some(&database),
Some("redis://127.0.0.1/0"),
RuntimeBackendArg::Redis,
)
.expect_err("multi-node sqlite should be rejected");
assert_eq!(error.kind(), std::io::ErrorKind::InvalidInput);
assert!(error.to_string().contains("AETHER_DATABASE_DRIVER=sqlite"));
}
#[test]
fn pending_schema_error_mentions_explicit_migrate_command() {
let error = pending_schema_error(2, 20260413020000, "squash usage schema split");
@@ -1363,37 +1819,37 @@ mod tests {
}
#[tokio::test]
async fn ensure_postgres_schema_is_current_is_noop_without_postgres_pool() {
async fn ensure_database_schema_is_current_is_noop_without_database_pool() {
let state = AppState::new().expect("state should build");
ensure_postgres_schema_is_current(&state)
ensure_database_schema_is_current(&state)
.await
.expect("disabled data backend should not block startup");
}
#[tokio::test]
async fn ensure_postgres_backfills_are_current_is_noop_without_postgres_pool() {
async fn ensure_database_backfills_are_current_is_noop_without_database_pool() {
let state = AppState::new().expect("state should build");
ensure_postgres_backfills_are_current(&state)
ensure_database_backfills_are_current(&state)
.await
.expect("disabled data backend should not block startup");
}
#[tokio::test]
async fn auto_prepare_database_is_noop_without_postgres_pool() {
async fn auto_prepare_database_is_noop_without_database_pool() {
let state = AppState::new().expect("state should build");
super::prepare_postgres_startup_requirements(&state, true)
super::prepare_database_startup_requirements(&state, true)
.await
.expect("disabled data backend should not block startup");
}
#[tokio::test]
async fn explicit_migrate_requires_postgres_url() {
async fn explicit_migrate_requires_database_url() {
let args = test_args();
let error = super::run_explicit_migrations(&args)
.await
.expect_err("missing postgres URL should fail");
.expect_err("missing database URL should fail");
let message = error.to_string();
assert!(message.contains("AETHER_GATEWAY_DATA_POSTGRES_URL or DATABASE_URL"));
assert!(message.contains("AETHER_DATABASE_DRIVER/AETHER_DATABASE_URL"));
assert!(message.contains("--migrate"));
}
@@ -1404,20 +1860,40 @@ mod tests {
let error = super::run_explicit_migrations(&args)
.await
.expect_err("missing postgres URL should fail before any app port validation");
.expect_err("missing database URL should fail before any app port validation");
let message = error.to_string();
assert!(message.contains("AETHER_GATEWAY_DATA_POSTGRES_URL or DATABASE_URL"));
assert!(message.contains("AETHER_DATABASE_DRIVER/AETHER_DATABASE_URL"));
assert!(!message.contains("APP_PORT"));
}
#[tokio::test]
async fn explicit_backfills_require_postgres_url() {
async fn explicit_backfills_require_database_url() {
let args = test_args();
let error = super::run_explicit_backfills(&args)
.await
.expect_err("missing postgres URL should fail");
.expect_err("missing database URL should fail");
let message = error.to_string();
assert!(message.contains("AETHER_GATEWAY_DATA_POSTGRES_URL or DATABASE_URL"));
assert!(message.contains("AETHER_DATABASE_DRIVER/AETHER_DATABASE_URL"));
assert!(message.contains("--apply-backfills"));
}
#[tokio::test]
async fn explicit_backfills_are_noop_for_sqlite_database() {
let mut args = test_args();
let database_path = std::env::temp_dir().join(format!(
"aether-sqlite-backfill-noop-{}-{}.db",
std::process::id(),
crate::current_unix_secs().expect("clock should be available")
));
args.data.database_driver = Some(DatabaseDriverArg::Sqlite);
args.data.database_url = Some(format!("sqlite://{}", database_path.display()));
super::run_explicit_migrations(&args)
.await
.expect("sqlite migrations should run before backfills");
super::run_explicit_backfills(&args)
.await
.expect("sqlite backfills should be an explicit no-op");
let _ = std::fs::remove_file(database_path);
}
}

File diff suppressed because it is too large Load Diff

View File

@@ -3,26 +3,14 @@ use chrono::{DateTime, Utc};
use crate::data::GatewayDataState;
use super::{
postgres_error, system_config_bool, system_config_u64, system_config_usize,
DELETE_AUDIT_LOGS_BEFORE_SQL,
};
use super::{system_config_bool, system_config_u64, system_config_usize};
pub(crate) async fn cleanup_audit_logs_once(
data: &GatewayDataState,
) -> Result<usize, DataLayerError> {
cleanup_audit_logs_with(data, |cutoff_time, delete_limit| async move {
let Some(pool) = data.postgres_pool() else {
return Ok(0);
};
let deleted = sqlx::query(DELETE_AUDIT_LOGS_BEFORE_SQL)
.bind(cutoff_time)
.bind(i64::try_from(delete_limit).unwrap_or(i64::MAX))
.execute(&pool)
data.delete_audit_logs_before(cutoff_time.timestamp().max(0) as u64, delete_limit)
.await
.map_err(postgres_error)?
.rows_affected();
Ok(usize::try_from(deleted).unwrap_or(usize::MAX))
})
.await
}

View File

@@ -3,7 +3,7 @@ use tracing::warn;
use crate::data::GatewayDataState;
use super::{postgres_error, system_config_bool, DB_MAINTENANCE_TABLES};
use super::{system_config_bool, DB_MAINTENANCE_TABLES};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) struct DbMaintenanceRunSummary {
@@ -14,25 +14,25 @@ pub(super) struct DbMaintenanceRunSummary {
pub(super) async fn perform_db_maintenance_once(
data: &GatewayDataState,
) -> Result<DbMaintenanceRunSummary, DataLayerError> {
let Some(pool) = data.postgres_pool() else {
if !data.has_database_maintenance_backend() {
return Ok(DbMaintenanceRunSummary {
attempted: 0,
succeeded: 0,
});
};
}
run_db_maintenance_with(data, |table_name| {
let pool = pool.clone();
async move {
let statement = format!("VACUUM ANALYZE {table_name}");
sqlx::raw_sql(&statement)
.execute(&pool)
.await
.map_err(postgres_error)?;
Ok(())
}
if !system_config_bool(data, "enable_db_maintenance", true).await? {
return Ok(DbMaintenanceRunSummary {
attempted: 0,
succeeded: 0,
});
}
let summary = data.run_database_maintenance(DB_MAINTENANCE_TABLES).await?;
Ok(DbMaintenanceRunSummary {
attempted: summary.attempted,
succeeded: summary.succeeded,
})
.await
}
pub(super) async fn run_db_maintenance_with<F, Fut>(

View File

@@ -1,19 +1,11 @@
use std::collections::HashSet;
use chrono::Utc;
use futures_util::TryStreamExt;
use sqlx::Row;
use crate::data::GatewayDataState;
use aether_data_contracts::DataLayerError;
use super::{
pending_cleanup_batch_size, pending_cleanup_timeout_minutes, postgres_error,
SELECT_COMPLETED_PENDING_REQUEST_IDS_SQL, SELECT_STALE_PENDING_USAGE_BATCH_SQL,
UPDATE_FAILED_PENDING_CANDIDATES_SQL, UPDATE_FAILED_STALE_USAGE_SQL,
UPDATE_FAILED_VOID_STALE_USAGE_SQL, UPDATE_RECOVERED_STALE_USAGE_SQL,
UPDATE_RECOVERED_STREAMING_CANDIDATES_SQL,
};
use super::{pending_cleanup_batch_size, pending_cleanup_timeout_minutes};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub(crate) struct PendingCleanupSummary {
@@ -47,123 +39,27 @@ pub(super) struct PendingCleanupBatchPlan {
pub(crate) async fn cleanup_stale_pending_requests_once(
data: &GatewayDataState,
) -> Result<PendingCleanupSummary, DataLayerError> {
let Some(pool) = data.postgres_pool() else {
if !data.has_usage_writer() {
return Ok(PendingCleanupSummary::default());
};
}
let timeout_minutes = pending_cleanup_timeout_minutes(data).await?;
let batch_size = pending_cleanup_batch_size(data).await?;
let cutoff_time =
Utc::now() - chrono::Duration::minutes(i64::try_from(timeout_minutes).unwrap_or(i64::MAX));
let active_statuses = vec!["pending", "streaming"];
let mut summary = PendingCleanupSummary::default();
let now_unix_secs = Utc::now().timestamp().max(0) as u64;
let cutoff_unix_secs = now_unix_secs.saturating_sub(timeout_minutes.saturating_mul(60));
let summary = data
.cleanup_stale_pending_requests(
cutoff_unix_secs,
now_unix_secs,
timeout_minutes,
batch_size,
)
.await?;
loop {
let mut tx = pool.begin().await.map_err(postgres_error)?;
let stale_rows = {
let mut stale_rows_stream = sqlx::query(SELECT_STALE_PENDING_USAGE_BATCH_SQL)
.bind(active_statuses.clone())
.bind(cutoff_time)
.bind(i64::try_from(batch_size).unwrap_or(i64::MAX))
.fetch(&mut *tx);
let mut stale_rows = Vec::new();
while let Some(row) = stale_rows_stream.try_next().await.map_err(postgres_error)? {
stale_rows.push(row);
}
stale_rows
};
if stale_rows.is_empty() {
tx.rollback().await.map_err(postgres_error)?;
break;
}
let stale_rows = stale_rows
.into_iter()
.map(|row| {
Ok::<StalePendingUsageRow, DataLayerError>(StalePendingUsageRow {
id: row.try_get::<String, _>("id").map_err(postgres_error)?,
request_id: row
.try_get::<String, _>("request_id")
.map_err(postgres_error)?,
status: row.try_get::<String, _>("status").map_err(postgres_error)?,
billing_status: row
.try_get::<String, _>("billing_status")
.map_err(postgres_error)?,
})
})
.collect::<Result<Vec<_>, DataLayerError>>()?;
let request_ids = stale_rows
.iter()
.map(|row| row.request_id.clone())
.collect::<Vec<_>>();
let completed_request_ids = if request_ids.is_empty() {
HashSet::new()
} else {
{
let mut completed_rows = sqlx::query(SELECT_COMPLETED_PENDING_REQUEST_IDS_SQL)
.bind(request_ids)
.fetch(&mut *tx);
let mut completed_request_ids = HashSet::new();
while let Some(row) = completed_rows.try_next().await.map_err(postgres_error)? {
if let Ok(request_id) = row.try_get::<String, _>("request_id") {
completed_request_ids.insert(request_id);
}
}
completed_request_ids
}
};
let plan = plan_pending_cleanup_batch(stale_rows, &completed_request_ids, timeout_minutes);
let now = Utc::now();
for usage_id in &plan.recovered_usage_ids {
sqlx::query(UPDATE_RECOVERED_STALE_USAGE_SQL)
.bind(usage_id)
.execute(&mut *tx)
.await
.map_err(postgres_error)?;
}
for failed_row in &plan.failed_usage_rows {
if failed_row.should_void_billing {
sqlx::query(UPDATE_FAILED_VOID_STALE_USAGE_SQL)
.bind(&failed_row.id)
.bind(&failed_row.error_message)
.bind(now)
.execute(&mut *tx)
.await
.map_err(postgres_error)?;
} else {
sqlx::query(UPDATE_FAILED_STALE_USAGE_SQL)
.bind(&failed_row.id)
.bind(&failed_row.error_message)
.execute(&mut *tx)
.await
.map_err(postgres_error)?;
}
}
if !plan.recovered_request_ids.is_empty() {
sqlx::query(UPDATE_RECOVERED_STREAMING_CANDIDATES_SQL)
.bind(plan.recovered_request_ids.clone())
.bind(now)
.execute(&mut *tx)
.await
.map_err(postgres_error)?;
}
if !plan.failed_request_ids.is_empty() {
sqlx::query(UPDATE_FAILED_PENDING_CANDIDATES_SQL)
.bind(plan.failed_request_ids.clone())
.bind(now)
.bind(active_statuses.clone())
.execute(&mut *tx)
.await
.map_err(postgres_error)?;
}
tx.commit().await.map_err(postgres_error)?;
summary.failed += plan.failed_usage_rows.len();
summary.recovered += plan.recovered_usage_ids.len();
}
Ok(summary)
Ok(PendingCleanupSummary {
failed: summary.failed,
recovered: summary.recovered,
})
}
pub(super) fn plan_pending_cleanup_batch(

View File

@@ -1,7 +1,7 @@
use std::collections::BTreeMap;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use aether_data::redis::{RedisKvRunner, RedisLockLease, RedisLockRunner};
use aether_data::driver::redis::{RedisKvRunner, RedisLockLease, RedisLockRunner};
use aether_data_contracts::repository::provider_catalog::{
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
};

View File

@@ -12,7 +12,7 @@ use super::{
perform_provider_checkin_once, perform_stats_aggregation_once,
perform_stats_hourly_aggregation_once, perform_usage_cleanup_once,
perform_wallet_daily_usage_aggregation_once, record_proxy_upgrade_traffic_success,
summarize_postgres_pool,
summarize_database_pool,
};
pub(super) async fn run_audit_cleanup_once(data: &GatewayDataState) -> Result<(), DataLayerError> {
@@ -212,20 +212,21 @@ pub(super) async fn run_usage_cleanup_once(data: &GatewayDataState) -> Result<()
}
pub(super) fn run_pool_monitor_once(data: &GatewayDataState) {
let Some(summary) = summarize_postgres_pool(data) else {
let Some(summary) = summarize_database_pool(data) else {
return;
};
info!(
event_name = "postgres_pool_sampled",
event_name = "database_pool_sampled",
log_type = "ops",
worker = "pool_monitor",
driver = %summary.driver,
checked_out = summary.checked_out,
pool_size = summary.pool_size,
idle = summary.idle,
max_connections = summary.max_connections,
usage_rate = summary.usage_rate,
"gateway postgres pool status"
"gateway database pool status"
);
}

View File

@@ -1,662 +1,24 @@
use chrono::{DateTime, Utc};
use sqlx::Row;
use uuid::Uuid;
use aether_data::{DataLayerError, StatsDailyAggregationInput, StatsDailyAggregationSummary};
use chrono::Utc;
use crate::data::GatewayDataState;
use aether_data_contracts::DataLayerError;
use super::{
postgres_error, stats_aggregation_target_day, system_config_bool, PercentileSummary,
StatsAggregationSummary, DELETE_STATS_DAILY_ERRORS_FOR_DATE_SQL, INSERT_STATS_DAILY_ERROR_SQL,
INSERT_STATS_SUMMARY_SQL, SELECT_EXISTING_STATS_SUMMARY_ID_SQL,
SELECT_LATEST_STATS_DAILY_DATE_SQL, SELECT_NEXT_STATS_DAILY_BUCKET_SQL,
SELECT_STATS_DAILY_AGGREGATE_SQL, SELECT_STATS_DAILY_FALLBACK_COUNT_SQL,
SELECT_STATS_DAILY_FIRST_BYTE_PERCENTILES_SQL,
SELECT_STATS_DAILY_RESPONSE_TIME_PERCENTILES_SQL, SELECT_STATS_SUMMARY_ENTITY_COUNTS_SQL,
SELECT_STATS_SUMMARY_TOTALS_SQL, UPDATE_STATS_SUMMARY_SQL, UPSERT_STATS_DAILY_API_KEY_SQL,
UPSERT_STATS_DAILY_COST_SAVINGS_MODEL_PROVIDER_SQL, UPSERT_STATS_DAILY_COST_SAVINGS_MODEL_SQL,
UPSERT_STATS_DAILY_COST_SAVINGS_PROVIDER_SQL, UPSERT_STATS_DAILY_COST_SAVINGS_SQL,
UPSERT_STATS_DAILY_MODEL_PROVIDER_SQL, UPSERT_STATS_DAILY_MODEL_SQL,
UPSERT_STATS_DAILY_PROVIDER_SQL, UPSERT_STATS_DAILY_SQL,
UPSERT_STATS_USER_DAILY_API_FORMAT_SQL,
UPSERT_STATS_USER_DAILY_COST_SAVINGS_MODEL_PROVIDER_SQL,
UPSERT_STATS_USER_DAILY_COST_SAVINGS_MODEL_SQL,
UPSERT_STATS_USER_DAILY_COST_SAVINGS_PROVIDER_SQL, UPSERT_STATS_USER_DAILY_COST_SAVINGS_SQL,
UPSERT_STATS_USER_DAILY_MODEL_PROVIDER_SQL, UPSERT_STATS_USER_DAILY_MODEL_SQL,
UPSERT_STATS_USER_DAILY_PROVIDER_SQL, UPSERT_STATS_USER_DAILY_SQL,
UPSERT_STATS_USER_SUMMARY_SQL,
};
use super::{stats_aggregation_target_day, system_config_bool};
pub(super) async fn perform_stats_aggregation_once(
data: &GatewayDataState,
) -> Result<Option<StatsAggregationSummary>, DataLayerError> {
let Some(pool) = data.postgres_pool() else {
) -> Result<Option<StatsDailyAggregationSummary>, DataLayerError> {
if !data.has_stats_daily_aggregation_backend() {
return Ok(None);
};
}
if !system_config_bool(data, "enable_stats_aggregation", true).await? {
return Ok(None);
}
let now_utc = Utc::now();
let target_day_utc = stats_aggregation_target_day(now_utc);
let Some(day_start_utc) = next_stats_aggregation_day(&pool, target_day_utc)
.await
.map_err(postgres_error)?
else {
return Ok(None);
};
perform_stats_aggregation_for_day(&pool, day_start_utc, now_utc)
.await
.map(Some)
.map_err(postgres_error)
}
async fn next_stats_aggregation_day(
pool: &aether_data::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: &aether_data::postgres::PostgresPool,
day_start_utc: DateTime<Utc>,
now_utc: DateTime<Utc>,
) -> Result<StatsAggregationSummary, 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(StatsAggregationSummary {
day_start_utc,
total_requests,
model_rows,
provider_rows,
api_key_rows,
error_rows,
user_rows,
data.aggregate_stats_daily(&StatsDailyAggregationInput {
target_day_utc: stats_aggregation_target_day(now_utc),
aggregated_at: now_utc,
})
}
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")?),
})
}
fn percentile_ms_to_i64(value: Option<f64>) -> Option<i64> {
value.and_then(|raw| raw.is_finite().then_some(raw.floor() as i64))
}
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(())
.await
}

View File

@@ -1,223 +1,24 @@
use chrono::{DateTime, Utc};
use sqlx::Row;
use uuid::Uuid;
use aether_data::{DataLayerError, StatsHourlyAggregationInput, StatsHourlyAggregationSummary};
use chrono::Utc;
use crate::data::GatewayDataState;
use aether_data_contracts::DataLayerError;
use super::{
stats_hourly_aggregation_target_hour, system_config_bool, SELECT_LATEST_STATS_HOURLY_HOUR_SQL,
SELECT_NEXT_STATS_HOURLY_BUCKET_SQL, SELECT_STATS_HOURLY_AGGREGATE_SQL,
UPSERT_STATS_HOURLY_MODEL_SQL, UPSERT_STATS_HOURLY_PROVIDER_SQL, UPSERT_STATS_HOURLY_SQL,
UPSERT_STATS_HOURLY_USER_MODEL_SQL, UPSERT_STATS_HOURLY_USER_SQL,
};
#[derive(Debug, Clone, PartialEq)]
pub(super) struct StatsHourlyAggregationSummary {
pub(super) hour_utc: DateTime<Utc>,
pub(super) total_requests: i64,
pub(super) user_rows: usize,
pub(super) user_model_rows: usize,
pub(super) model_rows: usize,
pub(super) provider_rows: usize,
}
use super::{stats_hourly_aggregation_target_hour, system_config_bool};
pub(super) async fn perform_stats_hourly_aggregation_once(
data: &GatewayDataState,
) -> Result<Option<StatsHourlyAggregationSummary>, DataLayerError> {
let Some(pool) = data.postgres_pool() else {
if !data.has_stats_hourly_aggregation_backend() {
return Ok(None);
};
}
if !system_config_bool(data, "enable_stats_aggregation", true).await? {
return Ok(None);
}
let now_utc = Utc::now();
let target_hour_utc = stats_hourly_aggregation_target_hour(now_utc);
let Some(hour_utc) = next_stats_hourly_bucket(&pool, target_hour_utc)
.await
.map_err(postgres_error)?
else {
return Ok(None);
};
perform_stats_hourly_aggregation_for_hour(&pool, hour_utc, now_utc)
.await
.map(Some)
.map_err(postgres_error)
}
async fn next_stats_hourly_bucket(
pool: &aether_data::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: &aether_data::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,
data.aggregate_stats_hourly(&StatsHourlyAggregationInput {
target_hour_utc: stats_hourly_aggregation_target_hour(now_utc),
aggregated_at: now_utc,
})
}
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))
}
fn postgres_error(error: sqlx::Error) -> DataLayerError {
DataLayerError::postgres(error)
.await
}

View File

@@ -25,12 +25,10 @@ use super::{
spawn_proxy_upgrade_rollout_worker, spawn_stats_aggregation_worker,
spawn_stats_hourly_aggregation_worker, spawn_usage_cleanup_worker,
spawn_wallet_daily_usage_aggregation_worker, start_proxy_upgrade_rollout,
stats_aggregation_target_day, stats_hourly_aggregation_target_hour, summarize_postgres_pool,
stats_aggregation_target_day, stats_hourly_aggregation_target_hour, summarize_database_pool,
usage_cleanup_settings, usage_cleanup_window, wallet_daily_usage_aggregation_target, AppState,
DbMaintenanceRunSummary, FailedPendingUsageRow, GatewayDataState,
ProxyUpgradeRolloutProbeConfig, StalePendingUsageRow, UsageCleanupSettings,
DELETE_STALE_WALLET_DAILY_USAGE_LEDGERS_SQL, SELECT_STALE_PENDING_USAGE_BATCH_SQL,
UPDATE_FAILED_VOID_STALE_USAGE_SQL, UPSERT_WALLET_DAILY_USAGE_LEDGER_SQL, USAGE_CLEANUP_HOUR,
ProxyUpgradeRolloutProbeConfig, StalePendingUsageRow, UsageCleanupSettings, USAGE_CLEANUP_HOUR,
USAGE_CLEANUP_MINUTE, WALLET_DAILY_USAGE_AGGREGATION_HOUR,
WALLET_DAILY_USAGE_AGGREGATION_MINUTE,
};
@@ -41,12 +39,12 @@ async fn spawn_audit_cleanup_worker_skips_when_postgres_unavailable() {
}
#[tokio::test]
async fn spawn_db_maintenance_worker_skips_when_postgres_unavailable() {
async fn spawn_db_maintenance_worker_skips_when_database_maintenance_unavailable() {
assert!(spawn_db_maintenance_worker(Arc::new(GatewayDataState::disabled())).is_none());
}
#[tokio::test]
async fn spawn_pending_cleanup_worker_skips_when_postgres_unavailable() {
async fn spawn_pending_cleanup_worker_skips_when_usage_writer_unavailable() {
assert!(spawn_pending_cleanup_worker(Arc::new(GatewayDataState::disabled())).is_none());
}
@@ -89,24 +87,6 @@ async fn spawn_pool_quota_probe_worker_skips_when_provider_catalog_unavailable()
assert!(spawn_pool_quota_probe_worker(state).is_none());
}
#[test]
fn wallet_daily_usage_queries_use_settlement_snapshots_for_wallet_identity() {
assert!(UPSERT_WALLET_DAILY_USAGE_LEDGER_SQL.contains("JOIN usage_settlement_snapshots"));
assert!(UPSERT_WALLET_DAILY_USAGE_LEDGER_SQL.contains("usage_settlement_snapshots.wallet_id"));
assert!(DELETE_STALE_WALLET_DAILY_USAGE_LEDGERS_SQL.contains("JOIN usage_settlement_snapshots"));
assert!(DELETE_STALE_WALLET_DAILY_USAGE_LEDGERS_SQL
.contains("usage_settlement_snapshots.wallet_id = ledgers.wallet_id"));
}
#[test]
fn pending_cleanup_queries_use_settlement_snapshots_for_billing_authority() {
assert!(SELECT_STALE_PENDING_USAGE_BATCH_SQL.contains("LEFT JOIN usage_settlement_snapshots"));
assert!(SELECT_STALE_PENDING_USAGE_BATCH_SQL
.contains("COALESCE(usage_settlement_snapshots.billing_status, usage.billing_status)"));
assert!(UPDATE_FAILED_VOID_STALE_USAGE_SQL.contains("INSERT INTO usage_settlement_snapshots"));
assert!(UPDATE_FAILED_VOID_STALE_USAGE_SQL.contains("billing_status = EXCLUDED.billing_status"));
}
fn sample_connected_proxy_node(
node_id: &str,
heartbeat_interval: i32,
@@ -548,24 +528,25 @@ async fn proxy_upgrade_rollout_active_probe_advances_next_wave_after_version_con
}
#[tokio::test]
async fn spawn_stats_aggregation_worker_skips_when_postgres_unavailable() {
async fn spawn_stats_aggregation_worker_skips_when_stats_daily_backend_unavailable() {
assert!(spawn_stats_aggregation_worker(Arc::new(GatewayDataState::disabled())).is_none());
}
#[tokio::test]
async fn spawn_stats_hourly_aggregation_worker_skips_when_postgres_unavailable() {
async fn spawn_stats_hourly_aggregation_worker_skips_when_stats_hourly_backend_unavailable() {
assert!(
spawn_stats_hourly_aggregation_worker(Arc::new(GatewayDataState::disabled())).is_none()
);
}
#[tokio::test]
async fn spawn_usage_cleanup_worker_skips_when_postgres_unavailable() {
async fn spawn_usage_cleanup_worker_skips_when_usage_writer_unavailable() {
assert!(spawn_usage_cleanup_worker(Arc::new(GatewayDataState::disabled())).is_none());
}
#[tokio::test]
async fn spawn_wallet_daily_usage_aggregation_worker_skips_when_postgres_unavailable() {
async fn spawn_wallet_daily_usage_aggregation_worker_skips_when_wallet_daily_usage_backend_unavailable(
) {
assert!(
spawn_wallet_daily_usage_aggregation_worker(Arc::new(GatewayDataState::disabled()))
.is_none()
@@ -780,9 +761,9 @@ fn usage_cleanup_window_uses_non_overlapping_ranges() {
}
#[tokio::test]
async fn summarize_postgres_pool_uses_busy_connections_for_usage_rate() {
async fn summarize_database_pool_uses_busy_connections_for_usage_rate() {
let data = GatewayDataState::from_config(crate::data::GatewayDataConfig::from_postgres_config(
aether_data::postgres::PostgresPoolConfig {
aether_data::driver::postgres::PostgresPoolConfig {
database_url: "postgres://localhost/aether".to_string(),
min_connections: 1,
max_connections: 8,
@@ -795,8 +776,9 @@ async fn summarize_postgres_pool_uses_busy_connections_for_usage_rate() {
))
.expect("gateway data state should build");
let summary = summarize_postgres_pool(&data).expect("pool summary should exist");
let summary = summarize_database_pool(&data).expect("pool summary should exist");
assert_eq!(summary.driver, aether_data::DatabaseDriver::Postgres);
assert_eq!(summary.checked_out, 0);
assert_eq!(summary.pool_size, 0);
assert_eq!(summary.idle, 0);

View File

@@ -1,865 +1,27 @@
use std::io::Write;
use aether_data_contracts::repository::usage::{
parse_usage_body_ref, usage_body_ref, UsageBodyField,
};
use aether_data_contracts::repository::usage::UsageCleanupSummary;
use aether_data_contracts::DataLayerError;
use chrono::{DateTime, Utc};
use flate2::{write::GzEncoder, Compression};
use futures_util::TryStreamExt;
use serde_json::{Map, Value};
use sqlx::Row;
use tracing::warn;
use chrono::Utc;
use crate::data::GatewayDataState;
use super::{
system_config_bool, usage_cleanup_settings, usage_cleanup_window, ExpiredApiKeyRow,
UsageBodyCleanupRow, UsageBodyCompressionRow, UsageCleanupSummary, CLEAR_USAGE_BODY_FIELDS_SQL,
CLEAR_USAGE_HEADER_FIELDS_SQL, CLEAR_USAGE_HTTP_AUDIT_BODY_REFS_SQL,
CLEAR_USAGE_HTTP_AUDIT_HEADERS_SQL, DELETE_EMPTY_USAGE_HTTP_AUDITS_SQL,
DELETE_EXPIRED_API_KEY_SQL, DELETE_OLD_USAGE_RECORDS_SQL, DELETE_USAGE_BODY_BLOBS_SQL,
DISABLE_EXPIRED_API_KEY_SQL, DISABLE_EXPIRED_API_KEY_WALLET_SQL,
EXPIRED_API_KEY_PRE_CLEAN_BATCH_SIZE, NULLIFY_REQUEST_CANDIDATE_API_KEY_BATCH_SQL,
NULLIFY_USAGE_API_KEY_BATCH_SQL, SELECT_EXPIRED_ACTIVE_API_KEYS_SQL,
SELECT_USAGE_BODY_COMPRESSION_BATCH_SQL, SELECT_USAGE_BODY_COMPRESSION_ROW_SQL,
SELECT_USAGE_HEADER_BATCH_SQL, SELECT_USAGE_LEGACY_BODY_REF_METADATA_BATCH_SQL,
SELECT_USAGE_STALE_BODY_BATCH_SQL, UPDATE_USAGE_BODY_COMPRESSION_SQL,
UPDATE_USAGE_REQUEST_METADATA_SQL, UPSERT_USAGE_BODY_BLOB_SQL,
UPSERT_USAGE_HTTP_AUDIT_BODY_REFS_SQL,
};
use super::{system_config_bool, usage_cleanup_settings, usage_cleanup_window};
pub(super) async fn perform_usage_cleanup_once(
data: &GatewayDataState,
) -> Result<UsageCleanupSummary, DataLayerError> {
let Some(pool) = data.postgres_pool() else {
if !data.has_usage_writer() {
return Ok(UsageCleanupSummary::default());
};
}
if !system_config_bool(data, "enable_auto_cleanup", true).await? {
return Ok(UsageCleanupSummary::default());
}
let settings = usage_cleanup_settings(data).await?;
let window = usage_cleanup_window(Utc::now(), settings);
let records_deleted =
delete_old_usage_records(&pool, window.log_cutoff, settings.batch_size).await?;
let header_cleaned = cleanup_usage_header_fields(
&pool,
window.header_cutoff,
data.cleanup_usage(
&window,
settings.batch_size,
Some(window.log_cutoff),
settings.auto_delete_expired_keys,
)
.await?;
let legacy_body_refs_migrated = migrate_legacy_usage_body_ref_metadata(
&pool,
window.detail_cutoff,
settings.batch_size,
Some(window.compressed_cutoff),
)
.await?;
let body_cleaned = cleanup_usage_stale_body_fields(
&pool,
window.compressed_cutoff,
settings.batch_size,
Some(window.log_cutoff),
)
.await?;
let body_externalized = compress_usage_body_fields(
&pool,
window.detail_cutoff,
settings.batch_size,
Some(window.compressed_cutoff),
)
.await?;
let keys_cleaned =
match cleanup_expired_api_keys(&pool, settings.auto_delete_expired_keys).await {
Ok(count) => count,
Err(err) => {
warn!(error = %err, "gateway expired api key cleanup failed");
0
}
};
Ok(UsageCleanupSummary {
body_externalized,
legacy_body_refs_migrated,
body_cleaned,
header_cleaned,
keys_cleaned,
records_deleted,
})
}
async fn migrate_legacy_usage_body_ref_metadata(
pool: &aether_data::postgres::PostgresPool,
cutoff_time: DateTime<Utc>,
batch_size: usize,
newer_than: Option<DateTime<Utc>>,
) -> Result<usize, DataLayerError> {
if matches!(newer_than, Some(value) if value >= cutoff_time) {
warn!(
cutoff_time = %cutoff_time,
newer_than = ?newer_than,
"gateway usage legacy body-ref migration skipped due to invalid window"
);
return Ok(0);
}
let mut total_migrated = 0usize;
loop {
let rows = sqlx::query(SELECT_USAGE_LEGACY_BODY_REF_METADATA_BATCH_SQL)
.bind(cutoff_time)
.bind(newer_than)
.bind(i64::try_from(batch_size).unwrap_or(i64::MAX))
.fetch_all(pool)
.await
.map_err(postgres_error)?
.into_iter()
.map(|row| {
Ok(UsageLegacyBodyRefMetadataRow {
id: row.try_get::<String, _>("id").map_err(postgres_error)?,
request_id: row
.try_get::<String, _>("request_id")
.map_err(postgres_error)?,
request_metadata: row
.try_get::<Option<Value>, _>("request_metadata")
.map_err(postgres_error)?,
})
})
.collect::<Result<Vec<_>, DataLayerError>>()?;
if rows.is_empty() {
break;
}
let mut batch_migrated = 0usize;
for row in rows {
let Some(plan) =
migrate_legacy_body_ref_metadata_plan(&row.request_id, row.request_metadata)
else {
continue;
};
let mut tx = pool.begin().await.map_err(postgres_error)?;
if plan.refs.any_present() {
sqlx::query(UPSERT_USAGE_HTTP_AUDIT_BODY_REFS_SQL)
.bind(&row.request_id)
.bind(plan.refs.request_body_ref.as_deref())
.bind(plan.refs.provider_request_body_ref.as_deref())
.bind(plan.refs.response_body_ref.as_deref())
.bind(plan.refs.client_response_body_ref.as_deref())
.bind("ref_backed")
.execute(&mut *tx)
.await
.map_err(postgres_error)?;
}
let updated = sqlx::query(UPDATE_USAGE_REQUEST_METADATA_SQL)
.bind(&row.id)
.bind(plan.request_metadata)
.execute(&mut *tx)
.await
.map_err(postgres_error)?
.rows_affected();
tx.commit().await.map_err(postgres_error)?;
if updated > 0 {
batch_migrated += 1;
}
}
total_migrated += batch_migrated;
if batch_migrated == 0 || batch_migrated < batch_size {
break;
}
}
Ok(total_migrated)
}
async fn delete_old_usage_records(
pool: &aether_data::postgres::PostgresPool,
cutoff_time: DateTime<Utc>,
batch_size: usize,
) -> Result<usize, DataLayerError> {
let mut total_deleted = 0usize;
loop {
let deleted = sqlx::query(DELETE_OLD_USAGE_RECORDS_SQL)
.bind(cutoff_time)
.bind(i64::try_from(batch_size).unwrap_or(i64::MAX))
.execute(pool)
.await
.map_err(postgres_error)?
.rows_affected();
let deleted = usize::try_from(deleted).unwrap_or(usize::MAX);
total_deleted += deleted;
if deleted < batch_size {
break;
}
}
Ok(total_deleted)
}
async fn cleanup_usage_header_fields(
pool: &aether_data::postgres::PostgresPool,
cutoff_time: DateTime<Utc>,
batch_size: usize,
newer_than: Option<DateTime<Utc>>,
) -> Result<usize, DataLayerError> {
if matches!(newer_than, Some(value) if value >= cutoff_time) {
warn!(
cutoff_time = %cutoff_time,
newer_than = ?newer_than,
"gateway usage header cleanup skipped due to invalid window"
);
return Ok(0);
}
let mut total_cleaned = 0usize;
loop {
let mut stream = sqlx::query(SELECT_USAGE_HEADER_BATCH_SQL)
.bind(cutoff_time)
.bind(newer_than)
.bind(i64::try_from(batch_size).unwrap_or(i64::MAX))
.fetch(pool);
let mut rows = Vec::new();
while let Some(row) = stream.try_next().await.map_err(postgres_error)? {
rows.push(UsageBodyCleanupRow {
id: row.try_get::<String, _>("id").map_err(postgres_error)?,
request_id: row
.try_get::<String, _>("request_id")
.map_err(postgres_error)?,
});
}
if rows.is_empty() {
break;
}
let ids = rows.iter().map(|row| row.id.clone()).collect::<Vec<_>>();
let request_ids = rows
.iter()
.map(|row| row.request_id.clone())
.collect::<Vec<_>>();
let cleaned = sqlx::query(CLEAR_USAGE_HEADER_FIELDS_SQL)
.bind(ids)
.execute(pool)
.await
.map_err(postgres_error)?
.rows_affected();
sqlx::query(CLEAR_USAGE_HTTP_AUDIT_HEADERS_SQL)
.bind(&request_ids)
.execute(pool)
.await
.map_err(postgres_error)?;
sqlx::query(DELETE_EMPTY_USAGE_HTTP_AUDITS_SQL)
.bind(request_ids)
.execute(pool)
.await
.map_err(postgres_error)?;
let cleaned = usize::try_from(cleaned).unwrap_or(usize::MAX);
total_cleaned += cleaned;
if cleaned == 0 || cleaned < batch_size {
break;
}
}
Ok(total_cleaned)
}
async fn cleanup_usage_stale_body_fields(
pool: &aether_data::postgres::PostgresPool,
cutoff_time: DateTime<Utc>,
batch_size: usize,
newer_than: Option<DateTime<Utc>>,
) -> Result<usize, DataLayerError> {
if matches!(newer_than, Some(value) if value >= cutoff_time) {
warn!(
cutoff_time = %cutoff_time,
newer_than = ?newer_than,
"gateway usage body cleanup skipped due to invalid window"
);
return Ok(0);
}
let mut total_cleaned = 0usize;
loop {
let mut stream = sqlx::query(SELECT_USAGE_STALE_BODY_BATCH_SQL)
.bind(cutoff_time)
.bind(newer_than)
.bind(i64::try_from(batch_size).unwrap_or(i64::MAX))
.fetch(pool);
let mut rows = Vec::new();
while let Some(row) = stream.try_next().await.map_err(postgres_error)? {
rows.push(UsageBodyCleanupRow {
id: row.try_get::<String, _>("id").map_err(postgres_error)?,
request_id: row
.try_get::<String, _>("request_id")
.map_err(postgres_error)?,
});
}
if rows.is_empty() {
break;
}
let ids = rows.iter().map(|row| row.id.clone()).collect::<Vec<_>>();
let request_ids = rows
.iter()
.map(|row| row.request_id.clone())
.collect::<Vec<_>>();
let cleaned = sqlx::query(CLEAR_USAGE_BODY_FIELDS_SQL)
.bind(ids)
.execute(pool)
.await
.map_err(postgres_error)?
.rows_affected();
sqlx::query(DELETE_USAGE_BODY_BLOBS_SQL)
.bind(&request_ids)
.execute(pool)
.await
.map_err(postgres_error)?;
sqlx::query(CLEAR_USAGE_HTTP_AUDIT_BODY_REFS_SQL)
.bind(&request_ids)
.execute(pool)
.await
.map_err(postgres_error)?;
sqlx::query(DELETE_EMPTY_USAGE_HTTP_AUDITS_SQL)
.bind(request_ids)
.execute(pool)
.await
.map_err(postgres_error)?;
let cleaned = usize::try_from(cleaned).unwrap_or(usize::MAX);
total_cleaned += cleaned;
if cleaned == 0 || cleaned < batch_size {
break;
}
}
Ok(total_cleaned)
}
async fn compress_usage_body_fields(
pool: &aether_data::postgres::PostgresPool,
cutoff_time: DateTime<Utc>,
batch_size: usize,
newer_than: Option<DateTime<Utc>>,
) -> Result<usize, DataLayerError> {
if matches!(newer_than, Some(value) if value >= cutoff_time) {
warn!(
cutoff_time = %cutoff_time,
newer_than = ?newer_than,
"gateway usage body compression skipped due to invalid window"
);
return Ok(0);
}
let mut total_compressed = 0usize;
let mut no_progress_count = 0usize;
let batch_size = batch_size.clamp(1, 25);
loop {
let mut stream = sqlx::query(SELECT_USAGE_BODY_COMPRESSION_BATCH_SQL)
.bind(cutoff_time)
.bind(newer_than)
.bind(i64::try_from(batch_size).unwrap_or(i64::MAX))
.fetch(pool);
let mut ids = Vec::new();
while let Some(row) = stream.try_next().await.map_err(postgres_error)? {
ids.push(row.try_get::<String, _>("id").map_err(postgres_error)?);
}
if ids.is_empty() {
break;
}
let mut batch_success = 0usize;
for id in ids {
let row = sqlx::query(SELECT_USAGE_BODY_COMPRESSION_ROW_SQL)
.bind(&id)
.fetch_optional(pool)
.await
.map_err(postgres_error)?;
let Some(row) = row else {
continue;
};
let row = UsageBodyCompressionRow {
id: row.try_get::<String, _>("id").map_err(postgres_error)?,
request_id: row
.try_get::<String, _>("request_id")
.map_err(postgres_error)?,
request_body: row
.try_get::<Option<Value>, _>("request_body")
.map_err(postgres_error)?,
request_body_compressed: row
.try_get::<Option<Vec<u8>>, _>("request_body_compressed")
.map_err(postgres_error)?,
response_body: row
.try_get::<Option<Value>, _>("response_body")
.map_err(postgres_error)?,
response_body_compressed: row
.try_get::<Option<Vec<u8>>, _>("response_body_compressed")
.map_err(postgres_error)?,
provider_request_body: row
.try_get::<Option<Value>, _>("provider_request_body")
.map_err(postgres_error)?,
provider_request_body_compressed: row
.try_get::<Option<Vec<u8>>, _>("provider_request_body_compressed")
.map_err(postgres_error)?,
client_response_body: row
.try_get::<Option<Value>, _>("client_response_body")
.map_err(postgres_error)?,
client_response_body_compressed: row
.try_get::<Option<Vec<u8>>, _>("client_response_body_compressed")
.map_err(postgres_error)?,
};
let detached = build_usage_body_externalization(&row)?;
if detached.refs.any_present() {
let mut tx = pool.begin().await.map_err(postgres_error)?;
for blob in &detached.blobs {
sqlx::query(UPSERT_USAGE_BODY_BLOB_SQL)
.bind(&blob.body_ref)
.bind(&row.request_id)
.bind(blob.body_field)
.bind(&blob.payload_gzip)
.execute(&mut *tx)
.await
.map_err(postgres_error)?;
}
sqlx::query(UPSERT_USAGE_HTTP_AUDIT_BODY_REFS_SQL)
.bind(&row.request_id)
.bind(detached.refs.request_body_ref.as_deref())
.bind(detached.refs.provider_request_body_ref.as_deref())
.bind(detached.refs.response_body_ref.as_deref())
.bind(detached.refs.client_response_body_ref.as_deref())
.bind("ref_backed")
.execute(&mut *tx)
.await
.map_err(postgres_error)?;
let updated = sqlx::query(UPDATE_USAGE_BODY_COMPRESSION_SQL)
.bind(&row.id)
.execute(&mut *tx)
.await
.map_err(postgres_error)?
.rows_affected();
tx.commit().await.map_err(postgres_error)?;
if updated > 0 {
batch_success += 1;
}
continue;
}
let updated = sqlx::query(UPDATE_USAGE_BODY_COMPRESSION_SQL)
.bind(&row.id)
.execute(pool)
.await
.map_err(postgres_error)?
.rows_affected();
if updated > 0 {
batch_success += 1;
}
}
if batch_success == 0 {
no_progress_count += 1;
if no_progress_count >= 3 {
warn!(
"gateway usage body compression stopped after repeated zero-progress batches"
);
break;
}
} else {
no_progress_count = 0;
}
total_compressed += batch_success;
}
Ok(total_compressed)
}
fn compress_usage_json_value(value: &Value) -> Result<Vec<u8>, DataLayerError> {
let bytes = serde_json::to_vec(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to serialize usage json for gzip: {err}"))
})?;
let mut encoder = GzEncoder::new(Vec::new(), Compression::new(6));
encoder.write_all(&bytes).map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to gzip usage json: {err}"))
})?;
encoder.finish().map_err(|err| {
DataLayerError::UnexpectedValue(format!("failed to finish gzipped usage json: {err}"))
})
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct UsageDetachedBodyBlobWrite {
body_ref: String,
body_field: &'static str,
payload_gzip: Vec<u8>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
struct UsageDetachedBodyRefs {
request_body_ref: Option<String>,
provider_request_body_ref: Option<String>,
response_body_ref: Option<String>,
client_response_body_ref: Option<String>,
}
#[derive(Debug, Clone, PartialEq)]
struct UsageLegacyBodyRefMetadataRow {
id: String,
request_id: String,
request_metadata: Option<Value>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
struct UsageLegacyBodyRefMigrationPlan {
refs: UsageDetachedBodyRefs,
request_metadata: Option<Value>,
}
impl UsageDetachedBodyRefs {
fn any_present(&self) -> bool {
self.request_body_ref.is_some()
|| self.provider_request_body_ref.is_some()
|| self.response_body_ref.is_some()
|| self.client_response_body_ref.is_some()
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
struct UsageBodyExternalizationPlan {
blobs: Vec<UsageDetachedBodyBlobWrite>,
refs: UsageDetachedBodyRefs,
}
fn migrate_legacy_body_ref_metadata_plan(
request_id: &str,
request_metadata: Option<Value>,
) -> Option<UsageLegacyBodyRefMigrationPlan> {
let mut metadata = match request_metadata {
Some(Value::Object(object)) => object,
_ => return None,
};
let mut refs = UsageDetachedBodyRefs::default();
let mut removed_any = false;
for field in [
UsageBodyField::RequestBody,
UsageBodyField::ProviderRequestBody,
UsageBodyField::ResponseBody,
UsageBodyField::ClientResponseBody,
] {
let key = field.as_ref_key();
let Some(value) = metadata.remove(key) else {
continue;
};
removed_any = true;
let parsed = value
.as_str()
.map(str::trim)
.filter(|value| !value.is_empty())
.and_then(parse_usage_body_ref)
.filter(|(parsed_request_id, parsed_field)| {
parsed_request_id == request_id && *parsed_field == field
})
.map(|(parsed_request_id, parsed_field)| {
usage_body_ref(&parsed_request_id, parsed_field)
});
match field {
UsageBodyField::RequestBody => refs.request_body_ref = parsed,
UsageBodyField::ProviderRequestBody => refs.provider_request_body_ref = parsed,
UsageBodyField::ResponseBody => refs.response_body_ref = parsed,
UsageBodyField::ClientResponseBody => refs.client_response_body_ref = parsed,
}
}
if !removed_any {
return None;
}
Some(UsageLegacyBodyRefMigrationPlan {
refs,
request_metadata: (!metadata.is_empty()).then_some(Value::Object(metadata)),
})
}
fn build_usage_body_externalization(
row: &UsageBodyCompressionRow,
) -> Result<UsageBodyExternalizationPlan, DataLayerError> {
let mut plan = UsageBodyExternalizationPlan::default();
maybe_externalize_usage_body_field(
&mut plan,
&row.request_id,
UsageBodyField::RequestBody,
row.request_body.as_ref(),
row.request_body_compressed.as_deref(),
)?;
maybe_externalize_usage_body_field(
&mut plan,
&row.request_id,
UsageBodyField::ProviderRequestBody,
row.provider_request_body.as_ref(),
row.provider_request_body_compressed.as_deref(),
)?;
maybe_externalize_usage_body_field(
&mut plan,
&row.request_id,
UsageBodyField::ResponseBody,
row.response_body.as_ref(),
row.response_body_compressed.as_deref(),
)?;
maybe_externalize_usage_body_field(
&mut plan,
&row.request_id,
UsageBodyField::ClientResponseBody,
row.client_response_body.as_ref(),
row.client_response_body_compressed.as_deref(),
)?;
Ok(plan)
}
fn maybe_externalize_usage_body_field(
plan: &mut UsageBodyExternalizationPlan,
request_id: &str,
field: UsageBodyField,
inline_body: Option<&Value>,
compressed_body: Option<&[u8]>,
) -> Result<(), DataLayerError> {
let Some(payload_gzip) = (match inline_body {
Some(value) => Some(compress_usage_json_value(value)?),
None => compressed_body.map(|value| value.to_vec()),
}) else {
return Ok(());
};
let body_ref = usage_body_ref(request_id, field);
plan.blobs.push(UsageDetachedBodyBlobWrite {
body_ref: body_ref.clone(),
body_field: field.as_storage_field(),
payload_gzip,
});
match field {
UsageBodyField::RequestBody => plan.refs.request_body_ref = Some(body_ref),
UsageBodyField::ProviderRequestBody => plan.refs.provider_request_body_ref = Some(body_ref),
UsageBodyField::ResponseBody => plan.refs.response_body_ref = Some(body_ref),
UsageBodyField::ClientResponseBody => plan.refs.client_response_body_ref = Some(body_ref),
}
Ok(())
}
async fn cleanup_expired_api_keys(
pool: &aether_data::postgres::PostgresPool,
auto_delete_expired_keys: bool,
) -> Result<usize, DataLayerError> {
let mut expired_keys = sqlx::query(SELECT_EXPIRED_ACTIVE_API_KEYS_SQL).fetch(pool);
let mut cleaned = 0usize;
while let Some(row) = expired_keys.try_next().await.map_err(postgres_error)? {
let api_key_id = row.try_get::<String, _>("id").map_err(postgres_error)?;
let key = ExpiredApiKeyRow {
id: api_key_id.as_str(),
auto_delete_on_expiry: row
.try_get::<Option<bool>, _>("auto_delete_on_expiry")
.map_err(postgres_error)?,
};
let should_delete = key
.auto_delete_on_expiry
.unwrap_or(auto_delete_expired_keys);
if should_delete {
nullify_expired_api_key_usage_refs(pool, key.id).await?;
nullify_expired_api_key_candidate_refs(pool, key.id).await?;
sqlx::query(DISABLE_EXPIRED_API_KEY_WALLET_SQL)
.bind(key.id)
.execute(pool)
.await
.map_err(postgres_error)?;
let deleted = sqlx::query(DELETE_EXPIRED_API_KEY_SQL)
.bind(key.id)
.execute(pool)
.await
.map_err(postgres_error)?
.rows_affected();
if deleted > 0 {
cleaned += 1;
}
} else {
let updated = sqlx::query(DISABLE_EXPIRED_API_KEY_SQL)
.bind(key.id)
.bind(Utc::now())
.execute(pool)
.await
.map_err(postgres_error)?
.rows_affected();
if updated > 0 {
cleaned += 1;
}
}
}
Ok(cleaned)
}
async fn nullify_expired_api_key_usage_refs(
pool: &aether_data::postgres::PostgresPool,
api_key_id: &str,
) -> Result<(), DataLayerError> {
loop {
let updated = sqlx::query(NULLIFY_USAGE_API_KEY_BATCH_SQL)
.bind(api_key_id)
.bind(i64::try_from(EXPIRED_API_KEY_PRE_CLEAN_BATCH_SIZE).unwrap_or(i64::MAX))
.execute(pool)
.await
.map_err(postgres_error)?
.rows_affected();
let updated = usize::try_from(updated).unwrap_or(usize::MAX);
if updated < EXPIRED_API_KEY_PRE_CLEAN_BATCH_SIZE {
break;
}
}
Ok(())
}
async fn nullify_expired_api_key_candidate_refs(
pool: &aether_data::postgres::PostgresPool,
api_key_id: &str,
) -> Result<(), DataLayerError> {
loop {
let updated = sqlx::query(NULLIFY_REQUEST_CANDIDATE_API_KEY_BATCH_SQL)
.bind(api_key_id)
.bind(i64::try_from(EXPIRED_API_KEY_PRE_CLEAN_BATCH_SIZE).unwrap_or(i64::MAX))
.execute(pool)
.await
.map_err(postgres_error)?
.rows_affected();
let updated = usize::try_from(updated).unwrap_or(usize::MAX);
if updated < EXPIRED_API_KEY_PRE_CLEAN_BATCH_SIZE {
break;
}
}
Ok(())
}
fn postgres_error(error: sqlx::Error) -> DataLayerError {
DataLayerError::postgres(error)
}
#[cfg(test)]
mod tests {
use std::io::Read;
use flate2::read::GzDecoder;
use serde_json::json;
use super::{
build_usage_body_externalization, compress_usage_json_value,
migrate_legacy_body_ref_metadata_plan, UsageBodyCompressionRow,
};
fn inflate_json(bytes: &[u8]) -> serde_json::Value {
let mut decoder = GzDecoder::new(bytes);
let mut decoded = Vec::new();
decoder
.read_to_end(&mut decoded)
.expect("gzip should decode");
serde_json::from_slice(&decoded).expect("json should decode")
}
#[test]
fn usage_body_externalization_moves_inline_json_into_ref_backed_blobs() {
let row = UsageBodyCompressionRow {
id: "usage-1".to_string(),
request_id: "req-1".to_string(),
request_body: Some(json!({"hello": "world"})),
request_body_compressed: None,
response_body: None,
response_body_compressed: None,
provider_request_body: Some(json!({"provider": true})),
provider_request_body_compressed: None,
client_response_body: None,
client_response_body_compressed: None,
};
let plan = build_usage_body_externalization(&row).expect("plan should build");
assert_eq!(plan.blobs.len(), 2);
assert_eq!(
plan.refs.request_body_ref.as_deref(),
Some("usage://request/req-1/request_body")
);
assert_eq!(
plan.refs.provider_request_body_ref.as_deref(),
Some("usage://request/req-1/provider_request_body")
);
assert_eq!(
inflate_json(&plan.blobs[0].payload_gzip),
json!({"hello": "world"})
);
assert_eq!(
inflate_json(&plan.blobs[1].payload_gzip),
json!({"provider": true})
);
}
#[test]
fn usage_body_externalization_reuses_existing_compressed_payloads() {
let compressed = compress_usage_json_value(&json!({"legacy": true}))
.expect("compressed payload should build");
let row = UsageBodyCompressionRow {
id: "usage-1".to_string(),
request_id: "req-legacy".to_string(),
request_body: None,
request_body_compressed: Some(compressed.clone()),
response_body: None,
response_body_compressed: None,
provider_request_body: None,
provider_request_body_compressed: None,
client_response_body: None,
client_response_body_compressed: None,
};
let plan = build_usage_body_externalization(&row).expect("plan should build");
assert_eq!(plan.blobs.len(), 1);
assert_eq!(plan.blobs[0].payload_gzip, compressed);
assert_eq!(
plan.refs.request_body_ref.as_deref(),
Some("usage://request/req-legacy/request_body")
);
}
#[test]
fn legacy_body_ref_metadata_migration_moves_matching_refs_and_strips_keys() {
let plan = migrate_legacy_body_ref_metadata_plan(
"req-1",
Some(json!({
"trace_id": "trace-1",
"request_body_ref": "usage://request/req-1/request_body",
"response_body_ref": "usage://request/req-1/response_body"
})),
)
.expect("migration plan should exist");
assert_eq!(
plan.refs.request_body_ref.as_deref(),
Some("usage://request/req-1/request_body")
);
assert_eq!(
plan.refs.response_body_ref.as_deref(),
Some("usage://request/req-1/response_body")
);
assert_eq!(
plan.request_metadata,
Some(json!({
"trace_id": "trace-1"
}))
);
}
#[test]
fn legacy_body_ref_metadata_migration_strips_invalid_and_cross_request_refs() {
let plan = migrate_legacy_body_ref_metadata_plan(
"req-1",
Some(json!({
"request_body_ref": "blob://legacy-request",
"provider_request_body_ref": "usage://request/req-other/provider_request_body",
"candidate_index": 2
})),
)
.expect("migration plan should exist");
assert!(!plan.refs.any_present());
assert_eq!(
plan.request_metadata,
Some(json!({
"candidate_index": 2
}))
);
}
.await
}

View File

@@ -3,10 +3,7 @@ use chrono::{DateTime, Utc};
use crate::data::GatewayDataState;
use aether_data_contracts::DataLayerError;
use super::{
maintenance_timezone, wallet_daily_usage_aggregation_target,
DELETE_STALE_WALLET_DAILY_USAGE_LEDGERS_SQL, UPSERT_WALLET_DAILY_USAGE_LEDGER_SQL,
};
use super::{maintenance_timezone, wallet_daily_usage_aggregation_target};
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct WalletDailyUsageAggregationSummary {
@@ -30,46 +27,29 @@ pub(super) async fn perform_wallet_daily_usage_aggregation_once(
let timezone = maintenance_timezone();
let now_utc = Utc::now();
let target = wallet_daily_usage_aggregation_target(now_utc, timezone);
let Some(pool) = data.postgres_pool() else {
if !data.has_wallet_daily_usage_aggregation_backend() {
return Ok(WalletDailyUsageAggregationSummary {
billing_date: target.billing_date,
billing_timezone: target.billing_timezone,
aggregated_wallets: 0,
deleted_stale_ledgers: 0,
});
};
}
let mut tx = pool.begin().await.map_err(postgres_error)?;
let aggregated_wallets = sqlx::query(UPSERT_WALLET_DAILY_USAGE_LEDGER_SQL)
.bind(target.window_start_utc)
.bind(target.window_end_utc)
.bind(target.billing_date)
.bind(target.billing_timezone.as_str())
.bind(now_utc)
.execute(&mut *tx)
.await
.map_err(postgres_error)?
.rows_affected();
let deleted_stale_ledgers = sqlx::query(DELETE_STALE_WALLET_DAILY_USAGE_LEDGERS_SQL)
.bind(target.billing_date)
.bind(target.billing_timezone.as_str())
.bind(target.window_start_utc)
.bind(target.window_end_utc)
.execute(&mut *tx)
.await
.map_err(postgres_error)?
.rows_affected();
tx.commit().await.map_err(postgres_error)?;
let result = data
.aggregate_wallet_daily_usage(&aether_data::WalletDailyUsageAggregationInput {
billing_date: target.billing_date.to_string(),
billing_timezone: target.billing_timezone.clone(),
window_start_unix_secs: target.window_start_utc.timestamp().max(0) as u64,
window_end_unix_secs: target.window_end_utc.timestamp().max(0) as u64,
aggregated_at_unix_secs: now_utc.timestamp().max(0) as u64,
})
.await?;
Ok(WalletDailyUsageAggregationSummary {
billing_date: target.billing_date,
billing_timezone: target.billing_timezone,
aggregated_wallets: usize::try_from(aggregated_wallets).unwrap_or(usize::MAX),
deleted_stale_ledgers: usize::try_from(deleted_stale_ledgers).unwrap_or(usize::MAX),
aggregated_wallets: result.aggregated_wallets,
deleted_stale_ledgers: result.deleted_stale_ledgers,
})
}
fn postgres_error(error: sqlx::Error) -> DataLayerError {
DataLayerError::postgres(error)
}

View File

@@ -42,7 +42,7 @@ fn log_maintenance_worker_failure(
pub(crate) fn spawn_audit_cleanup_worker(
data: Arc<GatewayDataState>,
) -> Option<tokio::task::JoinHandle<()>> {
if data.postgres_pool().is_none() {
if !data.has_audit_log_reader() {
return None;
}
@@ -65,7 +65,7 @@ pub(crate) fn spawn_audit_cleanup_worker(
pub(crate) fn spawn_db_maintenance_worker(
data: Arc<GatewayDataState>,
) -> Option<tokio::task::JoinHandle<()>> {
if data.postgres_pool().is_none() {
if !data.has_database_maintenance_backend() {
return None;
}
@@ -83,7 +83,7 @@ pub(crate) fn spawn_db_maintenance_worker(
pub(crate) fn spawn_wallet_daily_usage_aggregation_worker(
data: Arc<GatewayDataState>,
) -> Option<tokio::task::JoinHandle<()>> {
if data.postgres_pool().is_none() {
if !data.has_wallet_daily_usage_aggregation_backend() {
return None;
}
@@ -107,7 +107,7 @@ pub(crate) fn spawn_wallet_daily_usage_aggregation_worker(
pub(crate) fn spawn_stats_aggregation_worker(
data: Arc<GatewayDataState>,
) -> Option<tokio::task::JoinHandle<()>> {
if data.postgres_pool().is_none() {
if !data.has_stats_daily_aggregation_backend() {
return None;
}
@@ -137,7 +137,7 @@ pub(crate) fn spawn_stats_aggregation_worker(
pub(crate) fn spawn_usage_cleanup_worker(
data: Arc<GatewayDataState>,
) -> Option<tokio::task::JoinHandle<()>> {
if data.postgres_pool().is_none() {
if !data.has_usage_writer() {
return None;
}
@@ -224,7 +224,7 @@ pub(crate) fn spawn_gemini_file_mapping_cleanup_worker(
pub(crate) fn spawn_pending_cleanup_worker(
data: Arc<GatewayDataState>,
) -> Option<tokio::task::JoinHandle<()>> {
if data.postgres_pool().is_none() {
if !data.has_usage_writer() {
return None;
}
@@ -296,7 +296,7 @@ pub(crate) fn spawn_proxy_upgrade_rollout_worker(
pub(crate) fn spawn_pool_monitor_worker(
data: Arc<GatewayDataState>,
) -> Option<tokio::task::JoinHandle<()>> {
if data.postgres_pool().is_none() {
if !data.has_database_pool_summary() {
return None;
}
@@ -314,7 +314,7 @@ pub(crate) fn spawn_pool_monitor_worker(
pub(crate) fn spawn_stats_hourly_aggregation_worker(
data: Arc<GatewayDataState>,
) -> Option<tokio::task::JoinHandle<()>> {
if data.postgres_pool().is_none() {
if !data.has_stats_hourly_aggregation_backend() {
return None;
}

View File

@@ -1,265 +1,17 @@
use crate::handlers::shared::decrypt_catalog_secret_with_fallbacks;
use crate::{AppState, GatewayError};
use aether_data::repository::oauth_providers::StoredOAuthProviderConfig;
use aether_data::repository::users::StoredUserAuthRecord;
use aether_data::repository::users::{StoredUserAuthRecord, StoredUserOAuthLinkSummary};
use aether_oauth::identity::{IdentityClaims, IdentityOAuthProviderConfig};
use chrono::{DateTime, Utc};
use chrono::Utc;
use serde::Serialize;
use serde_json::{json, Value};
use sqlx::Row;
use uuid::Uuid;
const LINUXDO_AUTHORIZE_URL: &str = "https://connect.linux.do/oauth2/authorize";
const LINUXDO_TOKEN_URL: &str = "https://connect.linux.do/oauth2/token";
const LINUXDO_USERINFO_URL: &str = "https://connect.linux.do/api/user";
const FIND_OAUTH_LINKED_USER_SQL: &str = r#"
SELECT
users.id,
users.email,
users.email_verified,
users.username,
users.password_hash,
users.role::text AS role,
users.auth_source::text AS auth_source,
users.allowed_providers,
users.allowed_api_formats,
users.allowed_models,
users.is_active,
users.is_deleted,
users.created_at,
users.last_login_at
FROM user_oauth_links
JOIN users ON users.id = user_oauth_links.user_id
WHERE user_oauth_links.provider_type = $1
AND user_oauth_links.provider_user_id = $2
LIMIT 1
"#;
const FIND_USER_BY_EMAIL_SQL: &str = r#"
SELECT
id,
email,
email_verified,
username,
password_hash,
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_api_formats,
allowed_models,
is_active,
is_deleted,
created_at,
last_login_at
FROM users
WHERE LOWER(email) = LOWER($1)
AND is_deleted IS FALSE
LIMIT 1
"#;
const CHECK_USERNAME_TAKEN_SQL: &str = r#"
SELECT id
FROM users
WHERE username = $1
LIMIT 1
"#;
const CREATE_OAUTH_USER_SQL: &str = r#"
INSERT INTO users (
id,
email,
email_verified,
username,
password_hash,
role,
auth_source,
is_active,
is_deleted,
created_at,
updated_at,
last_login_at
)
VALUES (
$1,
$2,
TRUE,
$3,
NULL,
'user'::userrole,
'oauth'::authsource,
TRUE,
FALSE,
$4,
$4,
$4
)
RETURNING
id,
email,
email_verified,
username,
password_hash,
role::text AS role,
auth_source::text AS auth_source,
allowed_providers,
allowed_api_formats,
allowed_models,
is_active,
is_deleted,
created_at,
last_login_at
"#;
const UPSERT_OAUTH_LINK_SQL: &str = r#"
INSERT INTO user_oauth_links (
id,
user_id,
provider_type,
provider_user_id,
provider_username,
provider_email,
extra_data,
linked_at,
last_login_at
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $8)
ON CONFLICT (user_id, provider_type) DO UPDATE
SET provider_user_id = EXCLUDED.provider_user_id,
provider_username = EXCLUDED.provider_username,
provider_email = EXCLUDED.provider_email,
extra_data = EXCLUDED.extra_data,
last_login_at = EXCLUDED.last_login_at
"#;
const TOUCH_OAUTH_LINK_SQL: &str = r#"
UPDATE user_oauth_links
SET provider_username = COALESCE($3, provider_username),
provider_email = COALESCE($4, provider_email),
extra_data = COALESCE($5, extra_data),
last_login_at = $6
WHERE provider_type = $1
AND provider_user_id = $2
"#;
const CREATE_AUTH_USER_WALLET_SQL: &str = r#"
INSERT INTO wallets (
id,
user_id,
api_key_id,
balance,
gift_balance,
limit_mode,
currency,
status,
total_recharged,
total_consumed,
total_refunded,
total_adjusted,
created_at,
updated_at
)
VALUES (
$1,
$2,
NULL,
0,
$3,
$4,
'USD',
'active',
0,
0,
0,
$3,
NOW(),
NOW()
)
"#;
const CREATE_AUTH_USER_WALLET_GIFT_TX_SQL: &str = r#"
INSERT INTO wallet_transactions (
id,
wallet_id,
category,
reason_code,
amount,
balance_before,
balance_after,
recharge_balance_before,
recharge_balance_after,
gift_balance_before,
gift_balance_after,
link_type,
link_id,
operator_id,
description,
created_at
)
VALUES (
$1,
$2,
'gift',
'gift_initial',
$3,
0,
$3,
0,
0,
0,
$3,
'system_task',
$4,
NULL,
'用户初始赠款',
NOW()
)
"#;
const LIST_OAUTH_LINKS_SQL: &str = r#"
SELECT
user_oauth_links.provider_type,
oauth_providers.display_name,
user_oauth_links.provider_username,
user_oauth_links.provider_email,
user_oauth_links.linked_at,
user_oauth_links.last_login_at,
oauth_providers.is_enabled AS provider_enabled
FROM user_oauth_links
JOIN oauth_providers
ON oauth_providers.provider_type = user_oauth_links.provider_type
WHERE user_oauth_links.user_id = $1
ORDER BY user_oauth_links.linked_at ASC
"#;
const FIND_OAUTH_LINK_OWNER_SQL: &str = r#"
SELECT user_id
FROM user_oauth_links
WHERE provider_type = $1
AND provider_user_id = $2
LIMIT 1
"#;
const FIND_USER_PROVIDER_LINK_OWNER_SQL: &str = r#"
SELECT user_id
FROM user_oauth_links
WHERE user_id = $1
AND provider_type = $2
LIMIT 1
"#;
const COUNT_USER_OAUTH_LINKS_SQL: &str = r#"
SELECT COUNT(*)::bigint AS link_count
FROM user_oauth_links
WHERE user_id = $1
"#;
const DELETE_USER_OAUTH_LINK_SQL: &str = r#"
DELETE FROM user_oauth_links
WHERE user_id = $1
AND provider_type = $2
"#;
#[derive(Debug, Clone, PartialEq, Eq, Serialize)]
pub(crate) struct IdentityOAuthProviderSummary {
pub(crate) provider_type: String,
@@ -350,15 +102,14 @@ pub(crate) async fn list_identity_oauth_links(
state: &AppState,
user_id: &str,
) -> Result<Vec<IdentityOAuthLinkSummary>, GatewayError> {
let Some(pool) = state.postgres_pool() else {
return Ok(Vec::new());
};
let rows = sqlx::query(LIST_OAUTH_LINKS_SQL)
.bind(user_id)
.fetch_all(&pool)
state
.data
.list_user_oauth_links(user_id)
.await
.map_err(sql_gateway_error)?;
rows.iter().map(map_link_summary_row).collect()
.map_err(data_gateway_error)?
.into_iter()
.map(map_link_summary)
.collect()
}
pub(crate) async fn list_bindable_identity_oauth_providers(
@@ -382,39 +133,36 @@ pub(crate) async fn resolve_identity_oauth_login_user(
state: &AppState,
claims: &IdentityClaims,
) -> Result<StoredUserAuthRecord, IdentityOAuthAccountError> {
let Some(pool) = state.postgres_pool() else {
return Err(IdentityOAuthAccountError::ProviderUnavailable);
};
let now = Utc::now();
if let Some(row) = sqlx::query(FIND_OAUTH_LINKED_USER_SQL)
.bind(&claims.provider_type)
.bind(&claims.subject)
.fetch_optional(&pool)
if let Some(user) = state
.data
.find_oauth_linked_user(&claims.provider_type, &claims.subject)
.await
.map_err(repo_sql_error)?
.map_err(repo_data_error)?
{
sqlx::query(TOUCH_OAUTH_LINK_SQL)
.bind(&claims.provider_type)
.bind(&claims.subject)
.bind(claims.username.as_deref())
.bind(claims.email.as_deref())
.bind(Some(claims.raw.clone()))
.bind(now)
.execute(&pool)
state
.data
.touch_oauth_link(
&claims.provider_type,
&claims.subject,
claims.username.as_deref(),
claims.email.as_deref(),
Some(claims.raw.clone()),
now,
)
.await
.map_err(repo_sql_error)?;
return map_user_auth_row(&row).map_err(repo_data_error);
.map_err(repo_data_error)?;
return Ok(user);
}
let email = normalize_identity_email(claims.email.as_deref());
if let Some(email) = email.as_deref() {
if let Some(row) = sqlx::query(FIND_USER_BY_EMAIL_SQL)
.bind(email)
.fetch_optional(&pool)
if let Some(existing) = state
.data
.find_active_user_auth_by_email_ci(email)
.await
.map_err(repo_sql_error)?
.map_err(repo_data_error)?
{
let existing = map_user_auth_row(&row).map_err(repo_data_error)?;
return Err(match existing.auth_source.to_ascii_lowercase().as_str() {
"local" => IdentityOAuthAccountError::EmailExistsLocal,
"ldap" => IdentityOAuthAccountError::EmailIsLdap,
@@ -442,22 +190,31 @@ pub(crate) async fn resolve_identity_oauth_login_user(
.map(|value| system_config_f64(value, 10.0))
.unwrap_or(10.0);
let mut tx = pool.begin().await.map_err(repo_sql_error)?;
let username = unique_oauth_username(&mut tx, claims).await?;
let user_id = Uuid::new_v4().to_string();
let row = sqlx::query(CREATE_OAUTH_USER_SQL)
.bind(&user_id)
.bind(email.as_deref())
.bind(&username)
.bind(now)
.fetch_one(&mut *tx)
let username = unique_oauth_username(state, claims).await?;
let user = state
.data
.create_oauth_auth_user(email, username, now)
.await
.map_err(repo_sql_error)?;
let user = map_user_auth_row(&row).map_err(repo_data_error)?;
create_initial_wallet_in_tx(&mut tx, &user.id, initial_gift).await?;
upsert_oauth_link_in_tx(&mut tx, &user.id, claims, now).await?;
tx.commit().await.map_err(repo_sql_error)?;
.map_err(repo_data_error)?
.ok_or_else(|| IdentityOAuthAccountError::Storage("oauth user not created".to_string()))?;
match state
.initialize_auth_user_wallet(&user.id, initial_gift, false)
.await
{
Ok(Some(_wallet)) => {}
Ok(None) => {
let _ = state.delete_local_auth_user(&user.id).await;
return Err(IdentityOAuthAccountError::ProviderUnavailable);
}
Err(err) => {
let _ = state.delete_local_auth_user(&user.id).await;
return Err(IdentityOAuthAccountError::Storage(format!("{err:?}")));
}
}
if let Err(err) = upsert_oauth_link(state, &user.id, claims, now).await {
let _ = state.delete_local_auth_user(&user.id).await;
return Err(err);
}
Ok(user)
}
@@ -469,34 +226,25 @@ pub(crate) async fn bind_identity_oauth_to_user(
if user.auth_source.eq_ignore_ascii_case("ldap") {
return Err(IdentityOAuthAccountError::EmailIsLdap);
}
let Some(pool) = state.postgres_pool() else {
return Err(IdentityOAuthAccountError::ProviderUnavailable);
};
if let Some(row) = sqlx::query(FIND_OAUTH_LINK_OWNER_SQL)
.bind(&claims.provider_type)
.bind(&claims.subject)
.fetch_optional(&pool)
if let Some(owner) = state
.data
.find_oauth_link_owner(&claims.provider_type, &claims.subject)
.await
.map_err(repo_sql_error)?
.map_err(repo_data_error)?
{
let owner: String = row.try_get("user_id").map_err(repo_sql_error)?;
if owner != user.id {
return Err(IdentityOAuthAccountError::OAuthAlreadyBound);
}
}
if sqlx::query(FIND_USER_PROVIDER_LINK_OWNER_SQL)
.bind(&user.id)
.bind(&claims.provider_type)
.fetch_optional(&pool)
if state
.data
.has_user_oauth_provider_link(&user.id, &claims.provider_type)
.await
.map_err(repo_sql_error)?
.is_some()
.map_err(repo_data_error)?
{
return Err(IdentityOAuthAccountError::AlreadyBoundProvider);
}
let mut tx = pool.begin().await.map_err(repo_sql_error)?;
upsert_oauth_link_in_tx(&mut tx, &user.id, claims, Utc::now()).await?;
tx.commit().await.map_err(repo_sql_error)?;
upsert_oauth_link(state, &user.id, claims, Utc::now()).await?;
Ok(())
}
@@ -508,28 +256,22 @@ pub(crate) async fn unbind_identity_oauth(
if user.auth_source.eq_ignore_ascii_case("ldap") {
return Err(IdentityOAuthAccountError::EmailIsLdap);
}
let Some(pool) = state.postgres_pool() else {
return Err(IdentityOAuthAccountError::ProviderUnavailable);
};
let row = sqlx::query(COUNT_USER_OAUTH_LINKS_SQL)
.bind(&user.id)
.fetch_one(&pool)
let link_count = state
.data
.count_user_oauth_links(&user.id)
.await
.map_err(repo_sql_error)?;
let link_count: i64 = row.try_get("link_count").map_err(repo_sql_error)?;
.map_err(repo_data_error)?;
if user.auth_source.eq_ignore_ascii_case("oauth") && link_count <= 1 {
return Err(IdentityOAuthAccountError::LastOAuthBinding);
}
if !user.auth_source.eq_ignore_ascii_case("local") && link_count <= 1 {
return Err(IdentityOAuthAccountError::LastLoginMethod);
}
let result = sqlx::query(DELETE_USER_OAUTH_LINK_SQL)
.bind(&user.id)
.bind(provider_type.trim())
.execute(&pool)
state
.data
.delete_user_oauth_link(&user.id, provider_type.trim())
.await
.map_err(repo_sql_error)?;
Ok(result.rows_affected() > 0)
.map_err(repo_data_error)
}
fn stored_provider_config_to_identity_config(
@@ -590,52 +332,22 @@ fn identity_provider_defaults(
}
}
fn map_link_summary_row(
row: &sqlx::postgres::PgRow,
fn map_link_summary(
row: StoredUserOAuthLinkSummary,
) -> Result<IdentityOAuthLinkSummary, GatewayError> {
Ok(IdentityOAuthLinkSummary {
provider_type: row.try_get("provider_type").map_err(sql_gateway_error)?,
display_name: row.try_get("display_name").map_err(sql_gateway_error)?,
provider_username: row
.try_get("provider_username")
.map_err(sql_gateway_error)?,
provider_email: row.try_get("provider_email").map_err(sql_gateway_error)?,
linked_at: row
.try_get::<Option<DateTime<Utc>>, _>("linked_at")
.map_err(sql_gateway_error)?
.map(|value| value.to_rfc3339()),
last_login_at: row
.try_get::<Option<DateTime<Utc>>, _>("last_login_at")
.map_err(sql_gateway_error)?
.map(|value| value.to_rfc3339()),
provider_enabled: row.try_get("provider_enabled").map_err(sql_gateway_error)?,
provider_type: row.provider_type,
display_name: row.display_name,
provider_username: row.provider_username,
provider_email: row.provider_email,
linked_at: row.linked_at.map(|value| value.to_rfc3339()),
last_login_at: row.last_login_at.map(|value| value.to_rfc3339()),
provider_enabled: row.provider_enabled,
})
}
fn map_user_auth_row(
row: &sqlx::postgres::PgRow,
) -> Result<StoredUserAuthRecord, aether_data::DataLayerError> {
StoredUserAuthRecord::new(
row.try_get("id").map_err(data_unexpected)?,
row.try_get("email").map_err(data_unexpected)?,
row.try_get("email_verified").map_err(data_unexpected)?,
row.try_get("username").map_err(data_unexpected)?,
row.try_get("password_hash").map_err(data_unexpected)?,
row.try_get("role").map_err(data_unexpected)?,
row.try_get("auth_source").map_err(data_unexpected)?,
row.try_get("allowed_providers").map_err(data_unexpected)?,
row.try_get("allowed_api_formats")
.map_err(data_unexpected)?,
row.try_get("allowed_models").map_err(data_unexpected)?,
row.try_get("is_active").map_err(data_unexpected)?,
row.try_get("is_deleted").map_err(data_unexpected)?,
row.try_get("created_at").map_err(data_unexpected)?,
row.try_get("last_login_at").map_err(data_unexpected)?,
)
}
async fn unique_oauth_username(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
state: &AppState,
claims: &IdentityClaims,
) -> Result<String, IdentityOAuthAccountError> {
let base = normalize_oauth_username(
@@ -661,11 +373,11 @@ async fn unique_oauth_username(
short_uuid()
)
};
let taken = sqlx::query(CHECK_USERNAME_TAKEN_SQL)
.bind(&candidate)
.fetch_optional(&mut **tx)
let taken = state
.data
.find_user_auth_by_username(&candidate)
.await
.map_err(repo_sql_error)?
.map_err(repo_data_error)?
.is_some();
if !taken {
return Ok(candidate);
@@ -674,53 +386,25 @@ async fn unique_oauth_username(
Ok(format!("oauth_{}", short_uuid()))
}
async fn upsert_oauth_link_in_tx(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
async fn upsert_oauth_link(
state: &AppState,
user_id: &str,
claims: &IdentityClaims,
now: DateTime<Utc>,
now: chrono::DateTime<Utc>,
) -> Result<(), IdentityOAuthAccountError> {
sqlx::query(UPSERT_OAUTH_LINK_SQL)
.bind(Uuid::new_v4().to_string())
.bind(user_id)
.bind(&claims.provider_type)
.bind(&claims.subject)
.bind(claims.username.as_deref())
.bind(claims.email.as_deref())
.bind(Some(claims.raw.clone()))
.bind(now)
.execute(&mut **tx)
state
.data
.upsert_user_oauth_link(
user_id,
&claims.provider_type,
&claims.subject,
claims.username.as_deref(),
claims.email.as_deref(),
Some(claims.raw.clone()),
now,
)
.await
.map_err(repo_sql_error)?;
Ok(())
}
async fn create_initial_wallet_in_tx(
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
user_id: &str,
initial_gift_usd: f64,
) -> Result<(), IdentityOAuthAccountError> {
let gift_amount = initial_gift_usd.max(0.0);
let wallet_id = Uuid::new_v4().to_string();
sqlx::query(CREATE_AUTH_USER_WALLET_SQL)
.bind(&wallet_id)
.bind(user_id)
.bind(gift_amount)
.bind("finite")
.execute(&mut **tx)
.await
.map_err(repo_sql_error)?;
if gift_amount > 0.0 {
sqlx::query(CREATE_AUTH_USER_WALLET_GIFT_TX_SQL)
.bind(Uuid::new_v4().to_string())
.bind(&wallet_id)
.bind(gift_amount)
.bind(user_id)
.execute(&mut **tx)
.await
.map_err(repo_sql_error)?;
}
Ok(())
.map_err(repo_data_error)
}
fn normalize_identity_email(value: Option<&str>) -> Option<String> {
@@ -786,18 +470,10 @@ fn system_config_f64(value: &Value, default: f64) -> f64 {
}
}
fn repo_sql_error(error: sqlx::Error) -> IdentityOAuthAccountError {
IdentityOAuthAccountError::Storage(error.to_string())
}
fn repo_data_error(error: aether_data::DataLayerError) -> IdentityOAuthAccountError {
IdentityOAuthAccountError::Storage(error.to_string())
}
fn sql_gateway_error(error: sqlx::Error) -> GatewayError {
fn data_gateway_error(error: aether_data::DataLayerError) -> GatewayError {
GatewayError::Internal(error.to_string())
}
fn data_unexpected(error: sqlx::Error) -> aether_data::DataLayerError {
aether_data::DataLayerError::UnexpectedValue(error.to_string())
}

View File

@@ -93,7 +93,7 @@ pub(crate) enum LocalExecutionEffect<'a> {
}
struct PoolFeedbackContext {
runner: aether_data::redis::RedisKvRunner,
runner: aether_data::driver::redis::RedisKvRunner,
pool_config: AdminProviderPoolConfig,
sticky_session_token: Option<String>,
}

View File

@@ -1,721 +0,0 @@
use crate::state::AdminBillingPresetApplyResult;
use crate::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingRuleRecord,
AdminBillingRuleWriteInput, GatewayError, LocalMutationOutcome,
};
use aether_data::postgres::PostgresPool;
use futures_util::TryStreamExt;
use sqlx::Row;
fn internal(err: impl ToString) -> GatewayError {
GatewayError::Internal(err.to_string())
}
pub(crate) async fn admin_billing_enabled_default_value_exists(
pool: &PostgresPool,
api_format: &str,
task_type: &str,
dimension_name: &str,
existing_id: Option<&str>,
) -> Result<bool, GatewayError> {
sqlx::query_scalar::<_, bool>(
r#"
SELECT EXISTS(
SELECT 1
FROM dimension_collectors
WHERE api_format = $1
AND task_type = $2
AND dimension_name = $3
AND is_enabled = TRUE
AND default_value IS NOT NULL
AND ($4::TEXT IS NULL OR id <> $4)
)
"#,
)
.bind(api_format)
.bind(task_type)
.bind(dimension_name)
.bind(existing_id)
.fetch_one(pool)
.await
.map_err(internal)
}
pub(crate) async fn create_admin_billing_rule(
pool: &PostgresPool,
input: &AdminBillingRuleWriteInput,
) -> Result<LocalMutationOutcome<AdminBillingRuleRecord>, GatewayError> {
let rule_id = uuid::Uuid::new_v4().to_string();
let row = match sqlx::query(
r#"
INSERT INTO billing_rules (
id,
name,
task_type,
global_model_id,
model_id,
expression,
variables,
dimension_mappings,
is_enabled,
created_at,
updated_at
)
VALUES (
$1,
$2,
$3,
$4,
$5,
$6,
$7,
$8,
$9,
NOW(),
NOW()
)
RETURNING
id,
name,
task_type,
global_model_id,
model_id,
expression,
variables,
dimension_mappings,
is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
"#,
)
.bind(&rule_id)
.bind(&input.name)
.bind(&input.task_type)
.bind(input.global_model_id.as_deref())
.bind(input.model_id.as_deref())
.bind(&input.expression)
.bind(&input.variables)
.bind(&input.dimension_mappings)
.bind(input.is_enabled)
.fetch_one(pool)
.await
{
Ok(row) => row,
Err(sqlx::Error::Database(err)) => {
return Ok(LocalMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
Err(err) => return Err(GatewayError::Internal(err.to_string())),
};
Ok(LocalMutationOutcome::Applied(admin_billing_rule_from_row(
&row,
)?))
}
pub(crate) async fn list_admin_billing_rules(
pool: &PostgresPool,
task_type: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<(Vec<AdminBillingRuleRecord>, u64), GatewayError> {
let total = read_count(
sqlx::query(
r#"
SELECT COUNT(*) AS total
FROM billing_rules
WHERE ($1::TEXT IS NULL OR task_type = $1)
AND ($2::BOOL IS NULL OR is_enabled = $2)
"#,
)
.bind(task_type)
.bind(is_enabled)
.fetch_one(pool)
.await
.map_err(internal)?,
)?;
let offset = u64::from(page.saturating_sub(1) * page_size);
let mut rows = sqlx::query(
r#"
SELECT
id,
name,
task_type,
global_model_id,
model_id,
expression,
variables,
dimension_mappings,
is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
FROM billing_rules
WHERE ($1::TEXT IS NULL OR task_type = $1)
AND ($2::BOOL IS NULL OR is_enabled = $2)
ORDER BY updated_at DESC
OFFSET $3
LIMIT $4
"#,
)
.bind(task_type)
.bind(is_enabled)
.bind(i64::try_from(offset).map_err(|err| GatewayError::Internal(err.to_string()))?)
.bind(i64::from(page_size))
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_err(internal)? {
items.push(admin_billing_rule_from_row(&row)?);
}
Ok((items, total))
}
pub(crate) async fn find_admin_billing_rule(
pool: &PostgresPool,
rule_id: &str,
) -> Result<Option<AdminBillingRuleRecord>, GatewayError> {
let row = sqlx::query(
r#"
SELECT
id,
name,
task_type,
global_model_id,
model_id,
expression,
variables,
dimension_mappings,
is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
FROM billing_rules
WHERE id = $1
"#,
)
.bind(rule_id)
.fetch_optional(pool)
.await
.map_err(internal)?;
row.as_ref().map(admin_billing_rule_from_row).transpose()
}
pub(crate) async fn update_admin_billing_rule(
pool: &PostgresPool,
rule_id: &str,
input: &AdminBillingRuleWriteInput,
) -> Result<LocalMutationOutcome<AdminBillingRuleRecord>, GatewayError> {
let row = match sqlx::query(
r#"
UPDATE billing_rules
SET
name = $2,
task_type = $3,
global_model_id = $4,
model_id = $5,
expression = $6,
variables = $7,
dimension_mappings = $8,
is_enabled = $9,
updated_at = NOW()
WHERE id = $1
RETURNING
id,
name,
task_type,
global_model_id,
model_id,
expression,
variables,
dimension_mappings,
is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
"#,
)
.bind(rule_id)
.bind(&input.name)
.bind(&input.task_type)
.bind(input.global_model_id.as_deref())
.bind(input.model_id.as_deref())
.bind(&input.expression)
.bind(&input.variables)
.bind(&input.dimension_mappings)
.bind(input.is_enabled)
.fetch_optional(pool)
.await
{
Ok(row) => row,
Err(sqlx::Error::Database(err)) => {
return Ok(LocalMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
Err(err) => return Err(GatewayError::Internal(err.to_string())),
};
match row {
Some(row) => Ok(LocalMutationOutcome::Applied(admin_billing_rule_from_row(
&row,
)?)),
None => Ok(LocalMutationOutcome::NotFound),
}
}
pub(crate) async fn create_admin_billing_collector(
pool: &PostgresPool,
input: &AdminBillingCollectorWriteInput,
) -> Result<LocalMutationOutcome<AdminBillingCollectorRecord>, GatewayError> {
let collector_id = uuid::Uuid::new_v4().to_string();
let row = match sqlx::query(
r#"
INSERT INTO dimension_collectors (
id,
api_format,
task_type,
dimension_name,
source_type,
source_path,
value_type,
transform_expression,
default_value,
priority,
is_enabled,
created_at,
updated_at
)
VALUES (
$1,
$2,
$3,
$4,
$5,
$6,
$7,
$8,
$9,
$10,
$11,
NOW(),
NOW()
)
RETURNING
id,
api_format,
task_type,
dimension_name,
source_type,
source_path,
value_type,
transform_expression,
default_value,
priority,
is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
"#,
)
.bind(&collector_id)
.bind(&input.api_format)
.bind(&input.task_type)
.bind(&input.dimension_name)
.bind(&input.source_type)
.bind(input.source_path.as_deref())
.bind(&input.value_type)
.bind(input.transform_expression.as_deref())
.bind(input.default_value.as_deref())
.bind(input.priority)
.bind(input.is_enabled)
.fetch_one(pool)
.await
{
Ok(row) => row,
Err(sqlx::Error::Database(err)) => {
return Ok(LocalMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
Err(err) => return Err(GatewayError::Internal(err.to_string())),
};
Ok(LocalMutationOutcome::Applied(
admin_billing_collector_from_row(&row)?,
))
}
pub(crate) async fn list_admin_billing_collectors(
pool: &PostgresPool,
api_format: Option<&str>,
task_type: Option<&str>,
dimension_name: Option<&str>,
is_enabled: Option<bool>,
page: u32,
page_size: u32,
) -> Result<(Vec<AdminBillingCollectorRecord>, u64), GatewayError> {
let total = read_count(
sqlx::query(
r#"
SELECT COUNT(*) AS total
FROM dimension_collectors
WHERE ($1::TEXT IS NULL OR api_format = $1)
AND ($2::TEXT IS NULL OR task_type = $2)
AND ($3::TEXT IS NULL OR dimension_name = $3)
AND ($4::BOOL IS NULL OR is_enabled = $4)
"#,
)
.bind(api_format)
.bind(task_type)
.bind(dimension_name)
.bind(is_enabled)
.fetch_one(pool)
.await
.map_err(internal)?,
)?;
let offset = u64::from(page.saturating_sub(1) * page_size);
let mut rows = sqlx::query(
r#"
SELECT
id,
api_format,
task_type,
dimension_name,
source_type,
source_path,
value_type,
transform_expression,
default_value,
priority,
is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
FROM dimension_collectors
WHERE ($1::TEXT IS NULL OR api_format = $1)
AND ($2::TEXT IS NULL OR task_type = $2)
AND ($3::TEXT IS NULL OR dimension_name = $3)
AND ($4::BOOL IS NULL OR is_enabled = $4)
ORDER BY updated_at DESC, priority DESC, id ASC
OFFSET $5
LIMIT $6
"#,
)
.bind(api_format)
.bind(task_type)
.bind(dimension_name)
.bind(is_enabled)
.bind(i64::try_from(offset).map_err(|err| GatewayError::Internal(err.to_string()))?)
.bind(i64::from(page_size))
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_err(internal)? {
items.push(admin_billing_collector_from_row(&row)?);
}
Ok((items, total))
}
pub(crate) async fn find_admin_billing_collector(
pool: &PostgresPool,
collector_id: &str,
) -> Result<Option<AdminBillingCollectorRecord>, GatewayError> {
let row = sqlx::query(
r#"
SELECT
id,
api_format,
task_type,
dimension_name,
source_type,
source_path,
value_type,
transform_expression,
default_value,
priority,
is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
FROM dimension_collectors
WHERE id = $1
"#,
)
.bind(collector_id)
.fetch_optional(pool)
.await
.map_err(internal)?;
row.as_ref()
.map(admin_billing_collector_from_row)
.transpose()
}
pub(crate) async fn update_admin_billing_collector(
pool: &PostgresPool,
collector_id: &str,
input: &AdminBillingCollectorWriteInput,
) -> Result<LocalMutationOutcome<AdminBillingCollectorRecord>, GatewayError> {
let row = match sqlx::query(
r#"
UPDATE dimension_collectors
SET
api_format = $2,
task_type = $3,
dimension_name = $4,
source_type = $5,
source_path = $6,
value_type = $7,
transform_expression = $8,
default_value = $9,
priority = $10,
is_enabled = $11,
updated_at = NOW()
WHERE id = $1
RETURNING
id,
api_format,
task_type,
dimension_name,
source_type,
source_path,
value_type,
transform_expression,
default_value,
priority,
is_enabled,
CAST(EXTRACT(EPOCH FROM created_at) AS BIGINT) AS created_at_unix_ms,
CAST(EXTRACT(EPOCH FROM updated_at) AS BIGINT) AS updated_at_unix_secs
"#,
)
.bind(collector_id)
.bind(&input.api_format)
.bind(&input.task_type)
.bind(&input.dimension_name)
.bind(&input.source_type)
.bind(input.source_path.as_deref())
.bind(&input.value_type)
.bind(input.transform_expression.as_deref())
.bind(input.default_value.as_deref())
.bind(input.priority)
.bind(input.is_enabled)
.fetch_optional(pool)
.await
{
Ok(row) => row,
Err(sqlx::Error::Database(err)) => {
return Ok(LocalMutationOutcome::Invalid(format!(
"Integrity error: {err}"
)))
}
Err(err) => return Err(GatewayError::Internal(err.to_string())),
};
match row {
Some(row) => Ok(LocalMutationOutcome::Applied(
admin_billing_collector_from_row(&row)?,
)),
None => Ok(LocalMutationOutcome::NotFound),
}
}
pub(crate) async fn apply_admin_billing_preset(
pool: &PostgresPool,
preset: &str,
mode: &str,
collectors: &[AdminBillingCollectorWriteInput],
) -> Result<LocalMutationOutcome<AdminBillingPresetApplyResult>, GatewayError> {
let mut created = 0_u64;
let mut updated = 0_u64;
let mut skipped = 0_u64;
let mut errors = Vec::new();
for collector in collectors {
let existing_id = match sqlx::query_scalar::<_, String>(
r#"
SELECT id
FROM dimension_collectors
WHERE api_format = $1
AND task_type = $2
AND dimension_name = $3
AND priority = $4
AND is_enabled = TRUE
LIMIT 1
"#,
)
.bind(&collector.api_format)
.bind(&collector.task_type)
.bind(&collector.dimension_name)
.bind(collector.priority)
.fetch_optional(pool)
.await
{
Ok(value) => value,
Err(err) => {
errors.push(format!(
"Failed to query collector: api_format={} task_type={} dim={}: {}",
collector.api_format, collector.task_type, collector.dimension_name, err
));
continue;
}
};
if let Some(existing_id) = existing_id {
if mode == "overwrite" {
match sqlx::query(
r#"
UPDATE dimension_collectors
SET
source_type = $2,
source_path = $3,
value_type = $4,
transform_expression = $5,
default_value = $6,
is_enabled = $7,
updated_at = NOW()
WHERE id = $1
"#,
)
.bind(&existing_id)
.bind(&collector.source_type)
.bind(collector.source_path.as_deref())
.bind(&collector.value_type)
.bind(collector.transform_expression.as_deref())
.bind(collector.default_value.as_deref())
.bind(collector.is_enabled)
.execute(pool)
.await
{
Ok(_) => updated += 1,
Err(err) => errors.push(format!(
"Failed to update collector {}: {}",
existing_id, err
)),
}
} else {
skipped += 1;
}
continue;
}
match sqlx::query(
r#"
INSERT INTO dimension_collectors (
id,
api_format,
task_type,
dimension_name,
source_type,
source_path,
value_type,
transform_expression,
default_value,
priority,
is_enabled,
created_at,
updated_at
)
VALUES (
$1,
$2,
$3,
$4,
$5,
$6,
$7,
$8,
$9,
$10,
$11,
NOW(),
NOW()
)
"#,
)
.bind(uuid::Uuid::new_v4().to_string())
.bind(&collector.api_format)
.bind(&collector.task_type)
.bind(&collector.dimension_name)
.bind(&collector.source_type)
.bind(collector.source_path.as_deref())
.bind(&collector.value_type)
.bind(collector.transform_expression.as_deref())
.bind(collector.default_value.as_deref())
.bind(collector.priority)
.bind(collector.is_enabled)
.execute(pool)
.await
{
Ok(_) => created += 1,
Err(err) => errors.push(format!(
"Failed to create collector: api_format={} task_type={} dim={}: {}",
collector.api_format, collector.task_type, collector.dimension_name, err
)),
}
}
Ok(LocalMutationOutcome::Applied(
AdminBillingPresetApplyResult {
preset: preset.to_string(),
mode: mode.to_string(),
created,
updated,
skipped,
errors,
},
))
}
fn read_count(row: sqlx::postgres::PgRow) -> Result<u64, GatewayError> {
Ok(row.try_get::<i64, _>("total").map_err(internal)?.max(0) as u64)
}
fn admin_billing_rule_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<AdminBillingRuleRecord, GatewayError> {
Ok(AdminBillingRuleRecord {
id: row.try_get("id").map_err(internal)?,
name: row.try_get("name").map_err(internal)?,
task_type: row.try_get("task_type").map_err(internal)?,
global_model_id: row.try_get("global_model_id").map_err(internal)?,
model_id: row.try_get("model_id").map_err(internal)?,
expression: row.try_get("expression").map_err(internal)?,
variables: row
.try_get::<Option<serde_json::Value>, _>("variables")
.map_err(internal)?
.unwrap_or_else(|| serde_json::json!({})),
dimension_mappings: row
.try_get::<Option<serde_json::Value>, _>("dimension_mappings")
.map_err(internal)?
.unwrap_or_else(|| serde_json::json!({})),
is_enabled: row.try_get("is_enabled").map_err(internal)?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_err(internal)?
.max(0) as u64,
updated_at_unix_secs: row
.try_get::<i64, _>("updated_at_unix_secs")
.map_err(internal)?
.max(0) as u64,
})
}
fn admin_billing_collector_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<AdminBillingCollectorRecord, GatewayError> {
Ok(AdminBillingCollectorRecord {
id: row.try_get("id").map_err(internal)?,
api_format: row.try_get("api_format").map_err(internal)?,
task_type: row.try_get("task_type").map_err(internal)?,
dimension_name: row.try_get("dimension_name").map_err(internal)?,
source_type: row.try_get("source_type").map_err(internal)?,
source_path: row.try_get("source_path").map_err(internal)?,
value_type: row.try_get("value_type").map_err(internal)?,
transform_expression: row.try_get("transform_expression").map_err(internal)?,
default_value: row.try_get("default_value").map_err(internal)?,
priority: row.try_get("priority").map_err(internal)?,
is_enabled: row.try_get("is_enabled").map_err(internal)?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_err(internal)?
.max(0) as u64,
updated_at_unix_secs: row
.try_get::<i64, _>("updated_at_unix_secs")
.map_err(internal)?
.max(0) as u64,
})
}

View File

@@ -1,853 +0,0 @@
use aether_data::postgres::PostgresPool;
use aether_data_contracts::repository::usage::StoredUsageDashboardSummary;
use chrono::{DateTime, Utc};
use futures_util::TryStreamExt;
use sqlx::Row;
use crate::GatewayError;
fn internal(err: impl ToString) -> GatewayError {
GatewayError::Internal(err.to_string())
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct DashboardDailyTotalsAggregateRow {
pub(crate) date: String,
pub(crate) requests: u64,
pub(crate) total_tokens: u64,
pub(crate) total_cost_usd: f64,
pub(crate) response_time_sum_ms: f64,
pub(crate) response_time_samples: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct DashboardDailyModelAggregateRow {
pub(crate) date: String,
pub(crate) model: String,
pub(crate) requests: u64,
pub(crate) total_tokens: u64,
pub(crate) total_cost_usd: f64,
pub(crate) response_time_sum_ms: f64,
pub(crate) response_time_samples: u64,
}
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct DashboardDailyProviderAggregateRow {
pub(crate) date: String,
pub(crate) provider: String,
pub(crate) requests: u64,
pub(crate) total_tokens: u64,
pub(crate) total_cost_usd: f64,
}
pub(crate) async fn summarize_dashboard_usage_from_daily_aggregates(
pool: &PostgresPool,
start_day_utc: DateTime<Utc>,
end_day_utc: DateTime<Utc>,
user_id: Option<&str>,
) -> Result<StoredUsageDashboardSummary, GatewayError> {
let row = if let Some(user_id) = user_id {
sqlx::query(
r#"
SELECT
COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests,
COALESCE(SUM(input_tokens), 0)::BIGINT AS input_tokens,
COALESCE(SUM(effective_input_tokens), 0)::BIGINT AS effective_input_tokens,
COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens,
COALESCE(SUM(input_tokens + output_tokens), 0)::BIGINT AS total_tokens,
COALESCE(SUM(cache_creation_tokens), 0)::BIGINT AS cache_creation_tokens,
COALESCE(SUM(cache_read_tokens), 0)::BIGINT AS cache_read_tokens,
COALESCE(SUM(total_input_context), 0)::BIGINT AS total_input_context,
CAST(COALESCE(SUM(cache_creation_cost), 0) AS DOUBLE PRECISION) AS cache_creation_cost_usd,
CAST(COALESCE(SUM(cache_read_cost), 0) AS DOUBLE PRECISION) AS cache_read_cost_usd,
CAST(COALESCE(SUM(total_cost), 0) AS DOUBLE PRECISION) AS total_cost_usd,
CAST(COALESCE(SUM(actual_total_cost), 0) AS DOUBLE PRECISION) AS actual_total_cost_usd,
COALESCE(SUM(error_requests), 0)::BIGINT AS error_requests,
COALESCE(SUM(response_time_sum_ms), 0) AS response_time_sum_ms,
COALESCE(SUM(response_time_samples), 0)::BIGINT AS response_time_samples
FROM stats_user_daily
WHERE user_id = $1
AND date >= $2
AND date < $3
"#,
)
.bind(user_id)
.bind(start_day_utc)
.bind(end_day_utc)
.fetch_one(pool)
.await
.map_err(|err| internal(format!("user daily aggregate summary lookup failed: {err}")))?
} else {
sqlx::query(
r#"
SELECT
COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests,
COALESCE(SUM(input_tokens), 0)::BIGINT AS input_tokens,
COALESCE(SUM(effective_input_tokens), 0)::BIGINT AS effective_input_tokens,
COALESCE(SUM(output_tokens), 0)::BIGINT AS output_tokens,
COALESCE(SUM(input_tokens + output_tokens), 0)::BIGINT AS total_tokens,
COALESCE(SUM(cache_creation_tokens), 0)::BIGINT AS cache_creation_tokens,
COALESCE(SUM(cache_read_tokens), 0)::BIGINT AS cache_read_tokens,
COALESCE(SUM(total_input_context), 0)::BIGINT AS total_input_context,
CAST(COALESCE(SUM(cache_creation_cost), 0) AS DOUBLE PRECISION) AS cache_creation_cost_usd,
CAST(COALESCE(SUM(cache_read_cost), 0) AS DOUBLE PRECISION) AS cache_read_cost_usd,
CAST(COALESCE(SUM(total_cost), 0) AS DOUBLE PRECISION) AS total_cost_usd,
CAST(COALESCE(SUM(actual_total_cost), 0) AS DOUBLE PRECISION) AS actual_total_cost_usd,
COALESCE(SUM(error_requests), 0)::BIGINT AS error_requests,
COALESCE(SUM(response_time_sum_ms), 0) AS response_time_sum_ms,
COALESCE(SUM(response_time_samples), 0)::BIGINT AS response_time_samples
FROM stats_daily
WHERE date >= $1
AND date < $2
"#,
)
.bind(start_day_utc)
.bind(end_day_utc)
.fetch_one(pool)
.await
.map_err(|err| internal(format!("daily aggregate summary lookup failed: {err}")))?
};
Ok(StoredUsageDashboardSummary {
total_requests: row
.try_get::<i64, _>("total_requests")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
input_tokens: row
.try_get::<i64, _>("input_tokens")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
effective_input_tokens: row
.try_get::<i64, _>("effective_input_tokens")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
output_tokens: row
.try_get::<i64, _>("output_tokens")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
total_tokens: row
.try_get::<i64, _>("total_tokens")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
cache_creation_tokens: row
.try_get::<i64, _>("cache_creation_tokens")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
cache_read_tokens: row
.try_get::<i64, _>("cache_read_tokens")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
total_input_context: row
.try_get::<i64, _>("total_input_context")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
cache_creation_cost_usd: row
.try_get::<f64, _>("cache_creation_cost_usd")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?,
cache_read_cost_usd: row
.try_get::<f64, _>("cache_read_cost_usd")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?,
total_cost_usd: row
.try_get::<f64, _>("total_cost_usd")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?,
actual_total_cost_usd: row
.try_get::<f64, _>("actual_total_cost_usd")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?,
error_requests: row
.try_get::<i64, _>("error_requests")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
response_time_sum_ms: row
.try_get::<f64, _>("response_time_sum_ms")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?,
response_time_samples: row
.try_get::<i64, _>("response_time_samples")
.map_err(|err| internal(format!("aggregate summary decode failed: {err}")))?
.max(0) as u64,
})
}
pub(crate) async fn read_stats_hourly_cutoff(
pool: &PostgresPool,
) -> Result<Option<DateTime<Utc>>, GatewayError> {
let row = sqlx::query(
r#"
SELECT MAX(hour_utc) AS latest_hour
FROM stats_hourly
WHERE is_complete IS TRUE
"#,
)
.fetch_one(pool)
.await
.map_err(|err| internal(format!("hourly aggregate cutoff lookup failed: {err}")))?;
let latest_hour = row
.try_get::<Option<DateTime<Utc>>, _>("latest_hour")
.map_err(|err| internal(format!("hourly aggregate cutoff decode failed: {err}")))?;
Ok(latest_hour.map(|value| value + chrono::Duration::hours(1)))
}
pub(crate) async fn list_admin_dashboard_daily_totals_aggregates(
pool: &PostgresPool,
start_day_utc: DateTime<Utc>,
end_day_utc: DateTime<Utc>,
) -> Result<Vec<DashboardDailyTotalsAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
date,
total_requests,
input_tokens,
output_tokens,
COALESCE(total_cost, 0)::DOUBLE PRECISION AS total_cost,
response_time_sum_ms,
response_time_samples
FROM stats_daily
WHERE date >= $1
AND date < $2
ORDER BY date ASC
"#,
)
.bind(start_day_utc)
.bind(end_day_utc)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("daily aggregate totals read failed: {err}")))?
{
let date = row
.try_get::<DateTime<Utc>, _>("date")
.map_err(|err| internal(format!("daily aggregate totals decode failed: {err}")))?;
let input_tokens = row
.try_get::<i64, _>("input_tokens")
.map_err(|err| internal(format!("daily aggregate totals decode failed: {err}")))?;
let output_tokens = row
.try_get::<i64, _>("output_tokens")
.map_err(|err| internal(format!("daily aggregate totals decode failed: {err}")))?;
items.push(DashboardDailyTotalsAggregateRow {
date: date.date_naive().to_string(),
requests: row
.try_get::<i32, _>("total_requests")
.map_err(|err| internal(format!("daily aggregate totals decode failed: {err}")))?
.max(0) as u64,
total_tokens: input_tokens.saturating_add(output_tokens).max(0) as u64,
total_cost_usd: row
.try_get::<f64, _>("total_cost")
.map_err(|err| internal(format!("daily aggregate totals decode failed: {err}")))?,
response_time_sum_ms: row
.try_get::<f64, _>("response_time_sum_ms")
.map_err(|err| internal(format!("daily aggregate totals decode failed: {err}")))?,
response_time_samples: row
.try_get::<i64, _>("response_time_samples")
.map_err(|err| internal(format!("daily aggregate totals decode failed: {err}")))?
.max(0) as u64,
});
}
Ok(items)
}
pub(crate) async fn list_admin_dashboard_hourly_totals_aggregates(
pool: &PostgresPool,
start_utc: DateTime<Utc>,
end_utc: DateTime<Utc>,
tz_offset_minutes: i32,
) -> Result<Vec<DashboardDailyTotalsAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
CAST(DATE(hour_utc + ($3::integer * INTERVAL '1 minute')) AS TEXT) AS date,
COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests,
COALESCE(SUM(input_tokens + output_tokens), 0)::BIGINT AS total_tokens,
CAST(COALESCE(SUM(total_cost), 0) AS DOUBLE PRECISION) AS total_cost,
CAST(COALESCE(SUM(response_time_sum_ms), 0) AS DOUBLE PRECISION) AS response_time_sum_ms,
COALESCE(SUM(response_time_samples), 0)::BIGINT AS response_time_samples
FROM stats_hourly
WHERE hour_utc >= $1
AND hour_utc < $2
GROUP BY date
ORDER BY date ASC
"#,
)
.bind(start_utc)
.bind(end_utc)
.bind(tz_offset_minutes)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("hourly aggregate totals read failed: {err}")))?
{
items.push(DashboardDailyTotalsAggregateRow {
date: row
.try_get::<String, _>("date")
.map_err(|err| internal(format!("hourly aggregate totals decode failed: {err}")))?,
requests: row
.try_get::<i64, _>("total_requests")
.map_err(|err| internal(format!("hourly aggregate totals decode failed: {err}")))?
.max(0) as u64,
total_tokens: row
.try_get::<i64, _>("total_tokens")
.map_err(|err| internal(format!("hourly aggregate totals decode failed: {err}")))?
.max(0) as u64,
total_cost_usd: row
.try_get::<f64, _>("total_cost")
.map_err(|err| internal(format!("hourly aggregate totals decode failed: {err}")))?,
response_time_sum_ms: row
.try_get::<f64, _>("response_time_sum_ms")
.map_err(|err| internal(format!("hourly aggregate totals decode failed: {err}")))?,
response_time_samples: row
.try_get::<i64, _>("response_time_samples")
.map_err(|err| internal(format!("hourly aggregate totals decode failed: {err}")))?
.max(0) as u64,
});
}
Ok(items)
}
pub(crate) async fn list_admin_dashboard_daily_model_aggregates(
pool: &PostgresPool,
start_day_utc: DateTime<Utc>,
end_day_utc: DateTime<Utc>,
) -> Result<Vec<DashboardDailyModelAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
date,
model,
total_requests,
input_tokens,
output_tokens,
COALESCE(total_cost, 0)::DOUBLE PRECISION AS total_cost,
response_time_sum_ms,
response_time_samples
FROM stats_daily_model
WHERE date >= $1
AND date < $2
ORDER BY date ASC, total_cost DESC, model ASC
"#,
)
.bind(start_day_utc)
.bind(end_day_utc)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("daily aggregate model read failed: {err}")))?
{
let date = row
.try_get::<DateTime<Utc>, _>("date")
.map_err(|err| internal(format!("daily aggregate model decode failed: {err}")))?;
let input_tokens = row
.try_get::<i64, _>("input_tokens")
.map_err(|err| internal(format!("daily aggregate model decode failed: {err}")))?;
let output_tokens = row
.try_get::<i64, _>("output_tokens")
.map_err(|err| internal(format!("daily aggregate model decode failed: {err}")))?;
items.push(DashboardDailyModelAggregateRow {
date: date.date_naive().to_string(),
model: row
.try_get::<String, _>("model")
.map_err(|err| internal(format!("daily aggregate model decode failed: {err}")))?,
requests: row
.try_get::<i32, _>("total_requests")
.map_err(|err| internal(format!("daily aggregate model decode failed: {err}")))?
.max(0) as u64,
total_tokens: input_tokens.saturating_add(output_tokens).max(0) as u64,
total_cost_usd: row
.try_get::<f64, _>("total_cost")
.map_err(|err| internal(format!("daily aggregate model decode failed: {err}")))?,
response_time_sum_ms: row
.try_get::<f64, _>("response_time_sum_ms")
.map_err(|err| internal(format!("daily aggregate model decode failed: {err}")))?,
response_time_samples: row
.try_get::<i64, _>("response_time_samples")
.map_err(|err| internal(format!("daily aggregate model decode failed: {err}")))?
.max(0) as u64,
});
}
Ok(items)
}
pub(crate) async fn list_admin_dashboard_hourly_model_aggregates(
pool: &PostgresPool,
start_utc: DateTime<Utc>,
end_utc: DateTime<Utc>,
tz_offset_minutes: i32,
) -> Result<Vec<DashboardDailyModelAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
CAST(DATE(hour_utc + ($3::integer * INTERVAL '1 minute')) AS TEXT) AS date,
model,
COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests,
COALESCE(SUM(input_tokens + output_tokens), 0)::BIGINT AS total_tokens,
CAST(COALESCE(SUM(total_cost), 0) AS DOUBLE PRECISION) AS total_cost,
CAST(COALESCE(SUM(response_time_sum_ms), 0) AS DOUBLE PRECISION) AS response_time_sum_ms,
COALESCE(SUM(response_time_samples), 0)::BIGINT AS response_time_samples
FROM stats_hourly_model
WHERE hour_utc >= $1
AND hour_utc < $2
GROUP BY date, model
ORDER BY date ASC, total_cost DESC, model ASC
"#,
)
.bind(start_utc)
.bind(end_utc)
.bind(tz_offset_minutes)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("hourly aggregate model read failed: {err}")))?
{
items.push(DashboardDailyModelAggregateRow {
date: row
.try_get::<String, _>("date")
.map_err(|err| internal(format!("hourly aggregate model decode failed: {err}")))?,
model: row
.try_get::<String, _>("model")
.map_err(|err| internal(format!("hourly aggregate model decode failed: {err}")))?,
requests: row
.try_get::<i64, _>("total_requests")
.map_err(|err| internal(format!("hourly aggregate model decode failed: {err}")))?
.max(0) as u64,
total_tokens: row
.try_get::<i64, _>("total_tokens")
.map_err(|err| internal(format!("hourly aggregate model decode failed: {err}")))?
.max(0) as u64,
total_cost_usd: row
.try_get::<f64, _>("total_cost")
.map_err(|err| internal(format!("hourly aggregate model decode failed: {err}")))?,
response_time_sum_ms: row
.try_get::<f64, _>("response_time_sum_ms")
.map_err(|err| internal(format!("hourly aggregate model decode failed: {err}")))?,
response_time_samples: row
.try_get::<i64, _>("response_time_samples")
.map_err(|err| internal(format!("hourly aggregate model decode failed: {err}")))?
.max(0) as u64,
});
}
Ok(items)
}
pub(crate) async fn list_admin_dashboard_daily_provider_aggregates(
pool: &PostgresPool,
start_day_utc: DateTime<Utc>,
end_day_utc: DateTime<Utc>,
) -> Result<Vec<DashboardDailyProviderAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
date,
provider_name,
total_requests,
input_tokens,
output_tokens,
COALESCE(total_cost, 0)::DOUBLE PRECISION AS total_cost
FROM stats_daily_provider
WHERE date >= $1
AND date < $2
ORDER BY date ASC, total_cost DESC, provider_name ASC
"#,
)
.bind(start_day_utc)
.bind(end_day_utc)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("daily aggregate provider read failed: {err}")))?
{
let date = row
.try_get::<DateTime<Utc>, _>("date")
.map_err(|err| internal(format!("daily aggregate provider decode failed: {err}")))?;
let input_tokens = row
.try_get::<i64, _>("input_tokens")
.map_err(|err| internal(format!("daily aggregate provider decode failed: {err}")))?;
let output_tokens = row
.try_get::<i64, _>("output_tokens")
.map_err(|err| internal(format!("daily aggregate provider decode failed: {err}")))?;
items.push(DashboardDailyProviderAggregateRow {
date: date.date_naive().to_string(),
provider: row.try_get::<String, _>("provider_name").map_err(|err| {
internal(format!("daily aggregate provider decode failed: {err}"))
})?,
requests: row
.try_get::<i32, _>("total_requests")
.map_err(|err| internal(format!("daily aggregate provider decode failed: {err}")))?
.max(0) as u64,
total_tokens: input_tokens.saturating_add(output_tokens).max(0) as u64,
total_cost_usd: row.try_get::<f64, _>("total_cost").map_err(|err| {
internal(format!("daily aggregate provider decode failed: {err}"))
})?,
});
}
Ok(items)
}
pub(crate) async fn list_admin_dashboard_hourly_provider_aggregates(
pool: &PostgresPool,
start_utc: DateTime<Utc>,
end_utc: DateTime<Utc>,
tz_offset_minutes: i32,
) -> Result<Vec<DashboardDailyProviderAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
CAST(DATE(hour_utc + ($3::integer * INTERVAL '1 minute')) AS TEXT) AS date,
provider_name,
COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests,
COALESCE(SUM(input_tokens + output_tokens), 0)::BIGINT AS total_tokens,
CAST(COALESCE(SUM(total_cost), 0) AS DOUBLE PRECISION) AS total_cost
FROM stats_hourly_provider
WHERE hour_utc >= $1
AND hour_utc < $2
GROUP BY date, provider_name
ORDER BY date ASC, total_cost DESC, provider_name ASC
"#,
)
.bind(start_utc)
.bind(end_utc)
.bind(tz_offset_minutes)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("hourly aggregate provider read failed: {err}")))?
{
items.push(DashboardDailyProviderAggregateRow {
date: row.try_get::<String, _>("date").map_err(|err| {
internal(format!("hourly aggregate provider decode failed: {err}"))
})?,
provider: row.try_get::<String, _>("provider_name").map_err(|err| {
internal(format!("hourly aggregate provider decode failed: {err}"))
})?,
requests: row
.try_get::<i64, _>("total_requests")
.map_err(|err| internal(format!("hourly aggregate provider decode failed: {err}")))?
.max(0) as u64,
total_tokens: row
.try_get::<i64, _>("total_tokens")
.map_err(|err| internal(format!("hourly aggregate provider decode failed: {err}")))?
.max(0) as u64,
total_cost_usd: row.try_get::<f64, _>("total_cost").map_err(|err| {
internal(format!("hourly aggregate provider decode failed: {err}"))
})?,
});
}
Ok(items)
}
pub(crate) async fn list_user_dashboard_daily_totals_aggregates(
pool: &PostgresPool,
start_day_utc: DateTime<Utc>,
end_day_utc: DateTime<Utc>,
user_id: &str,
) -> Result<Vec<DashboardDailyTotalsAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
date,
total_requests,
input_tokens,
output_tokens,
COALESCE(total_cost, 0)::DOUBLE PRECISION AS total_cost,
response_time_sum_ms,
response_time_samples
FROM stats_user_daily
WHERE user_id = $1
AND date >= $2
AND date < $3
ORDER BY date ASC
"#,
)
.bind(user_id)
.bind(start_day_utc)
.bind(end_day_utc)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("user daily aggregate totals read failed: {err}")))?
{
let date = row
.try_get::<DateTime<Utc>, _>("date")
.map_err(|err| internal(format!("user daily aggregate totals decode failed: {err}")))?;
let input_tokens = row
.try_get::<i64, _>("input_tokens")
.map_err(|err| internal(format!("user daily aggregate totals decode failed: {err}")))?;
let output_tokens = row
.try_get::<i64, _>("output_tokens")
.map_err(|err| internal(format!("user daily aggregate totals decode failed: {err}")))?;
items.push(DashboardDailyTotalsAggregateRow {
date: date.date_naive().to_string(),
requests: row
.try_get::<i32, _>("total_requests")
.map_err(|err| {
internal(format!("user daily aggregate totals decode failed: {err}"))
})?
.max(0) as u64,
total_tokens: input_tokens.saturating_add(output_tokens).max(0) as u64,
total_cost_usd: row.try_get::<f64, _>("total_cost").map_err(|err| {
internal(format!("user daily aggregate totals decode failed: {err}"))
})?,
response_time_sum_ms: row
.try_get::<f64, _>("response_time_sum_ms")
.map_err(|err| {
internal(format!("user daily aggregate totals decode failed: {err}"))
})?,
response_time_samples: row
.try_get::<i64, _>("response_time_samples")
.map_err(|err| {
internal(format!("user daily aggregate totals decode failed: {err}"))
})?
.max(0) as u64,
});
}
Ok(items)
}
pub(crate) async fn list_user_dashboard_hourly_totals_aggregates(
pool: &PostgresPool,
start_utc: DateTime<Utc>,
end_utc: DateTime<Utc>,
tz_offset_minutes: i32,
user_id: &str,
) -> Result<Vec<DashboardDailyTotalsAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
CAST(DATE(hour_utc + ($4::integer * INTERVAL '1 minute')) AS TEXT) AS date,
COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests,
COALESCE(SUM(input_tokens + output_tokens), 0)::BIGINT AS total_tokens,
CAST(COALESCE(SUM(total_cost), 0) AS DOUBLE PRECISION) AS total_cost,
CAST(COALESCE(SUM(response_time_sum_ms), 0) AS DOUBLE PRECISION) AS response_time_sum_ms,
COALESCE(SUM(response_time_samples), 0)::BIGINT AS response_time_samples
FROM stats_hourly_user
WHERE user_id = $1
AND hour_utc >= $2
AND hour_utc < $3
GROUP BY date
ORDER BY date ASC
"#,
)
.bind(user_id)
.bind(start_utc)
.bind(end_utc)
.bind(tz_offset_minutes)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("user hourly aggregate totals read failed: {err}")))?
{
items.push(DashboardDailyTotalsAggregateRow {
date: row.try_get::<String, _>("date").map_err(|err| {
internal(format!("user hourly aggregate totals decode failed: {err}"))
})?,
requests: row
.try_get::<i64, _>("total_requests")
.map_err(|err| {
internal(format!("user hourly aggregate totals decode failed: {err}"))
})?
.max(0) as u64,
total_tokens: row
.try_get::<i64, _>("total_tokens")
.map_err(|err| {
internal(format!("user hourly aggregate totals decode failed: {err}"))
})?
.max(0) as u64,
total_cost_usd: row.try_get::<f64, _>("total_cost").map_err(|err| {
internal(format!("user hourly aggregate totals decode failed: {err}"))
})?,
response_time_sum_ms: row
.try_get::<f64, _>("response_time_sum_ms")
.map_err(|err| {
internal(format!("user hourly aggregate totals decode failed: {err}"))
})?,
response_time_samples: row
.try_get::<i64, _>("response_time_samples")
.map_err(|err| {
internal(format!("user hourly aggregate totals decode failed: {err}"))
})?
.max(0) as u64,
});
}
Ok(items)
}
pub(crate) async fn list_user_dashboard_daily_model_aggregates(
pool: &PostgresPool,
start_day_utc: DateTime<Utc>,
end_day_utc: DateTime<Utc>,
user_id: &str,
) -> Result<Vec<DashboardDailyModelAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
date,
model,
total_requests,
input_tokens,
output_tokens,
COALESCE(total_cost, 0)::DOUBLE PRECISION AS total_cost,
response_time_sum_ms,
response_time_samples
FROM stats_user_daily_model
WHERE user_id = $1
AND date >= $2
AND date < $3
ORDER BY date ASC, total_cost DESC, model ASC
"#,
)
.bind(user_id)
.bind(start_day_utc)
.bind(end_day_utc)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("user daily aggregate model read failed: {err}")))?
{
let date = row
.try_get::<DateTime<Utc>, _>("date")
.map_err(|err| internal(format!("user daily aggregate model decode failed: {err}")))?;
let input_tokens = row
.try_get::<i64, _>("input_tokens")
.map_err(|err| internal(format!("user daily aggregate model decode failed: {err}")))?;
let output_tokens = row
.try_get::<i64, _>("output_tokens")
.map_err(|err| internal(format!("user daily aggregate model decode failed: {err}")))?;
items.push(DashboardDailyModelAggregateRow {
date: date.date_naive().to_string(),
model: row.try_get::<String, _>("model").map_err(|err| {
internal(format!("user daily aggregate model decode failed: {err}"))
})?,
requests: row
.try_get::<i32, _>("total_requests")
.map_err(|err| {
internal(format!("user daily aggregate model decode failed: {err}"))
})?
.max(0) as u64,
total_tokens: input_tokens.saturating_add(output_tokens).max(0) as u64,
total_cost_usd: row.try_get::<f64, _>("total_cost").map_err(|err| {
internal(format!("user daily aggregate model decode failed: {err}"))
})?,
response_time_sum_ms: row
.try_get::<f64, _>("response_time_sum_ms")
.map_err(|err| {
internal(format!("user daily aggregate model decode failed: {err}"))
})?,
response_time_samples: row
.try_get::<i64, _>("response_time_samples")
.map_err(|err| {
internal(format!("user daily aggregate model decode failed: {err}"))
})?
.max(0) as u64,
});
}
Ok(items)
}
pub(crate) async fn list_user_dashboard_hourly_model_aggregates(
pool: &PostgresPool,
start_utc: DateTime<Utc>,
end_utc: DateTime<Utc>,
tz_offset_minutes: i32,
user_id: &str,
) -> Result<Vec<DashboardDailyModelAggregateRow>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
CAST(DATE(hour_utc + ($4::integer * INTERVAL '1 minute')) AS TEXT) AS date,
model,
COALESCE(SUM(total_requests), 0)::BIGINT AS total_requests,
COALESCE(SUM(input_tokens + output_tokens), 0)::BIGINT AS total_tokens,
CAST(COALESCE(SUM(total_cost), 0) AS DOUBLE PRECISION) AS total_cost,
CAST(COALESCE(SUM(response_time_sum_ms), 0) AS DOUBLE PRECISION) AS response_time_sum_ms,
COALESCE(SUM(response_time_samples), 0)::BIGINT AS response_time_samples
FROM stats_hourly_user_model
WHERE user_id = $1
AND hour_utc >= $2
AND hour_utc < $3
GROUP BY date, model
ORDER BY date ASC, total_cost DESC, model ASC
"#,
)
.bind(user_id)
.bind(start_utc)
.bind(end_utc)
.bind(tz_offset_minutes)
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("user hourly aggregate model read failed: {err}")))?
{
items.push(DashboardDailyModelAggregateRow {
date: row.try_get::<String, _>("date").map_err(|err| {
internal(format!("user hourly aggregate model decode failed: {err}"))
})?,
model: row.try_get::<String, _>("model").map_err(|err| {
internal(format!("user hourly aggregate model decode failed: {err}"))
})?,
requests: row
.try_get::<i64, _>("total_requests")
.map_err(|err| {
internal(format!("user hourly aggregate model decode failed: {err}"))
})?
.max(0) as u64,
total_tokens: row
.try_get::<i64, _>("total_tokens")
.map_err(|err| {
internal(format!("user hourly aggregate model decode failed: {err}"))
})?
.max(0) as u64,
total_cost_usd: row.try_get::<f64, _>("total_cost").map_err(|err| {
internal(format!("user hourly aggregate model decode failed: {err}"))
})?,
response_time_sum_ms: row
.try_get::<f64, _>("response_time_sum_ms")
.map_err(|err| {
internal(format!("user hourly aggregate model decode failed: {err}"))
})?,
response_time_samples: row
.try_get::<i64, _>("response_time_samples")
.map_err(|err| {
internal(format!("user hourly aggregate model decode failed: {err}"))
})?
.max(0) as u64,
});
}
Ok(items)
}

View File

@@ -1,5 +0,0 @@
pub(crate) mod billing;
pub(crate) mod dashboard_stats;
pub(crate) mod monitoring;
pub(crate) mod usage_heatmap;
pub(crate) mod user_rollups;

View File

@@ -1,266 +0,0 @@
use std::collections::BTreeMap;
use aether_data::postgres::PostgresPool;
use chrono::{DateTime, Utc};
use futures_util::TryStreamExt;
use serde_json::{json, Value};
use sqlx::Row;
use crate::GatewayError;
fn internal(err: impl ToString) -> GatewayError {
GatewayError::Internal(err.to_string())
}
pub(crate) async fn list_admin_audit_logs(
pool: &PostgresPool,
cutoff_time: DateTime<Utc>,
username_pattern: Option<&str>,
event_type: Option<&str>,
limit: usize,
offset: usize,
) -> Result<(Vec<Value>, usize), GatewayError> {
let total = sqlx::query_scalar::<_, i64>(
r#"
SELECT COUNT(*)
FROM audit_logs AS a
LEFT JOIN users AS u ON a.user_id = u.id
WHERE a.created_at >= $1
AND ($2::text IS NULL OR u.username ILIKE $2 ESCAPE '\')
AND ($3::text IS NULL OR a.event_type = $3)
"#,
)
.bind(cutoff_time)
.bind(username_pattern)
.bind(event_type)
.fetch_one(pool)
.await
.map_err(|err| GatewayError::Internal(format!("admin audit logs count failed: {err}")))?;
let mut rows = sqlx::query(
r#"
SELECT
a.id,
a.event_type,
a.user_id,
u.email AS user_email,
u.username AS user_username,
a.description,
a.ip_address,
a.status_code,
a.error_message,
a.event_metadata AS metadata,
a.created_at
FROM audit_logs AS a
LEFT JOIN users AS u ON a.user_id = u.id
WHERE a.created_at >= $1
AND ($2::text IS NULL OR u.username ILIKE $2 ESCAPE '\')
AND ($3::text IS NULL OR a.event_type = $3)
ORDER BY a.created_at DESC
LIMIT $4 OFFSET $5
"#,
)
.bind(cutoff_time)
.bind(username_pattern)
.bind(event_type)
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.bind(i64::try_from(offset).unwrap_or(i64::MAX))
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| GatewayError::Internal(format!("admin audit logs read failed: {err}")))?
{
items.push(admin_audit_log_row_to_json(row));
}
Ok((items, usize::try_from(total.max(0)).unwrap_or(usize::MAX)))
}
pub(crate) async fn list_admin_suspicious_activities(
pool: &PostgresPool,
cutoff_time: DateTime<Utc>,
) -> Result<Vec<Value>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
id,
event_type,
user_id,
description,
ip_address,
event_metadata AS metadata,
created_at
FROM audit_logs
WHERE created_at >= $1
AND event_type = ANY($2)
ORDER BY created_at DESC
LIMIT 100
"#,
)
.bind(cutoff_time)
.bind(vec![
"suspicious_activity",
"unauthorized_access",
"login_failed",
"request_rate_limited",
])
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_err(|err| {
GatewayError::Internal(format!("admin suspicious activities read failed: {err}"))
})? {
items.push(admin_suspicious_row_to_json(row));
}
Ok(items)
}
pub(crate) async fn read_admin_user_behavior_event_counts(
pool: &PostgresPool,
user_id: &str,
cutoff_time: DateTime<Utc>,
) -> Result<BTreeMap<String, u64>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT event_type, COUNT(*)::bigint AS count
FROM audit_logs
WHERE user_id = $1
AND created_at >= $2
GROUP BY event_type
"#,
)
.bind(user_id)
.bind(cutoff_time)
.fetch(pool);
let mut counts = BTreeMap::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| GatewayError::Internal(format!("admin user behavior read failed: {err}")))?
{
let Ok(event_type) = row.try_get::<String, _>("event_type") else {
continue;
};
let count = row
.try_get::<i64, _>("count")
.ok()
.and_then(|value| u64::try_from(value.max(0)).ok())
.unwrap_or(0);
counts.insert(event_type, count);
}
Ok(counts)
}
pub(crate) async fn list_user_audit_logs(
pool: &PostgresPool,
user_id: &str,
cutoff_time: DateTime<Utc>,
event_type: Option<&str>,
limit: usize,
offset: usize,
) -> Result<(Vec<Value>, usize), GatewayError> {
let total = match sqlx::query_scalar::<_, i64>(
r#"
SELECT COUNT(*)
FROM audit_logs
WHERE user_id = $1
AND created_at >= $2
AND ($3::text IS NULL OR event_type = $3)
"#,
)
.bind(user_id)
.bind(cutoff_time)
.bind(event_type)
.fetch_one(pool)
.await
{
Ok(value) => usize::try_from(value.max(0)).unwrap_or(usize::MAX),
Err(err) => {
return Err(GatewayError::Internal(format!(
"user audit logs count failed: {err}"
)))
}
};
let mut rows = sqlx::query(
r#"
SELECT id, event_type, description, ip_address, status_code, created_at
FROM audit_logs
WHERE user_id = $1
AND created_at >= $2
AND ($3::text IS NULL OR event_type = $3)
ORDER BY created_at DESC
LIMIT $4 OFFSET $5
"#,
)
.bind(user_id)
.bind(cutoff_time)
.bind(event_type)
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.bind(i64::try_from(offset).unwrap_or(i64::MAX))
.fetch(pool);
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| GatewayError::Internal(format!("user audit logs read failed: {err}")))?
{
items.push(user_audit_log_row_to_json(row));
}
Ok((items, total))
}
fn admin_audit_log_row_to_json(row: sqlx::postgres::PgRow) -> Value {
let created_at = row
.try_get::<chrono::DateTime<chrono::Utc>, _>("created_at")
.ok()
.map(|value| value.to_rfc3339());
json!({
"id": row.try_get::<String, _>("id").ok(),
"event_type": row.try_get::<String, _>("event_type").ok(),
"user_id": row.try_get::<Option<String>, _>("user_id").ok().flatten(),
"user_email": row.try_get::<Option<String>, _>("user_email").ok().flatten(),
"user_username": row.try_get::<Option<String>, _>("user_username").ok().flatten(),
"description": row.try_get::<Option<String>, _>("description").ok().flatten(),
"ip_address": row.try_get::<Option<String>, _>("ip_address").ok().flatten(),
"status_code": row.try_get::<Option<i32>, _>("status_code").ok().flatten(),
"error_message": row.try_get::<Option<String>, _>("error_message").ok().flatten(),
"metadata": row.try_get::<Option<serde_json::Value>, _>("metadata").ok().flatten(),
"created_at": created_at,
})
}
fn admin_suspicious_row_to_json(row: sqlx::postgres::PgRow) -> Value {
let created_at = row
.try_get::<chrono::DateTime<chrono::Utc>, _>("created_at")
.ok()
.map(|value| value.to_rfc3339());
json!({
"id": row.try_get::<String, _>("id").ok(),
"event_type": row.try_get::<String, _>("event_type").ok(),
"user_id": row.try_get::<Option<String>, _>("user_id").ok().flatten(),
"description": row.try_get::<Option<String>, _>("description").ok().flatten(),
"ip_address": row.try_get::<Option<String>, _>("ip_address").ok().flatten(),
"metadata": row.try_get::<Option<serde_json::Value>, _>("metadata").ok().flatten(),
"created_at": created_at,
})
}
fn user_audit_log_row_to_json(row: sqlx::postgres::PgRow) -> Value {
let created_at = row
.try_get::<chrono::DateTime<chrono::Utc>, _>("created_at")
.ok()
.map(|value| value.to_rfc3339());
json!({
"id": row.try_get::<String, _>("id").ok(),
"event_type": row.try_get::<String, _>("event_type").ok(),
"description": row.try_get::<String, _>("description").ok(),
"ip_address": row.try_get::<Option<String>, _>("ip_address").ok().flatten(),
"status_code": row.try_get::<Option<i32>, _>("status_code").ok().flatten(),
"created_at": created_at,
})
}

View File

@@ -1,176 +0,0 @@
use aether_data::postgres::PostgresPool;
use aether_data_contracts::repository::usage::StoredUsageDailySummary;
use chrono::{DateTime, NaiveDate, Utc};
use futures_util::TryStreamExt;
use sqlx::Row;
use crate::GatewayError;
const USER_HEATMAP_AGGREGATE_SQL: &str = r#"
SELECT
date,
total_requests,
input_tokens,
output_tokens,
cache_creation_tokens,
cache_read_tokens,
COALESCE(total_cost, 0)::DOUBLE PRECISION AS total_cost,
COALESCE(actual_total_cost, 0)::DOUBLE PRECISION AS actual_total_cost
FROM stats_user_daily
WHERE user_id = $1
AND date >= $2
AND date < $3
ORDER BY date ASC
"#;
const GLOBAL_HEATMAP_AGGREGATE_SQL: &str = r#"
SELECT
date,
total_requests,
input_tokens,
output_tokens,
cache_creation_tokens,
cache_read_tokens,
COALESCE(total_cost, 0)::DOUBLE PRECISION AS total_cost,
COALESCE(actual_total_cost, 0)::DOUBLE PRECISION AS actual_total_cost
FROM stats_daily
WHERE date >= $1
AND date < $2
ORDER BY date ASC
"#;
fn internal(err: impl ToString) -> GatewayError {
GatewayError::Internal(err.to_string())
}
pub(crate) async fn read_stats_daily_cutoff_date(
pool: &PostgresPool,
) -> Result<Option<DateTime<Utc>>, GatewayError> {
let row = sqlx::query(
r#"
SELECT cutoff_date
FROM stats_summary
ORDER BY updated_at DESC, created_at DESC
LIMIT 1
"#,
)
.fetch_optional(pool)
.await
.map_err(|err| internal(format!("stats summary cutoff lookup failed: {err}")))?;
let Some(row) = row else {
return Ok(None);
};
row.try_get::<DateTime<Utc>, _>("cutoff_date")
.map(Some)
.map_err(|err| internal(format!("stats summary cutoff decode failed: {err}")))
}
pub(crate) async fn list_usage_heatmap_aggregate_rows(
pool: &PostgresPool,
start_date: NaiveDate,
end_date_exclusive: NaiveDate,
user_id: Option<&str>,
) -> Result<Vec<StoredUsageDailySummary>, GatewayError> {
if start_date >= end_date_exclusive {
return Ok(Vec::new());
}
let start_at = DateTime::<Utc>::from_naive_utc_and_offset(
start_date
.and_hms_opt(0, 0, 0)
.expect("midnight should be valid"),
Utc,
);
let end_at = DateTime::<Utc>::from_naive_utc_and_offset(
end_date_exclusive
.and_hms_opt(0, 0, 0)
.expect("midnight should be valid"),
Utc,
);
let mut rows = if let Some(user_id) = user_id {
sqlx::query(USER_HEATMAP_AGGREGATE_SQL)
.bind(user_id)
.bind(start_at)
.bind(end_at)
.fetch(pool)
} else {
sqlx::query(GLOBAL_HEATMAP_AGGREGATE_SQL)
.bind(start_at)
.bind(end_at)
.fetch(pool)
};
let mut items = Vec::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("aggregate heatmap read failed: {err}")))?
{
let date = row
.try_get::<DateTime<Utc>, _>("date")
.map_err(|err| internal(format!("aggregate heatmap date decode failed: {err}")))?;
let requests = row
.try_get::<i32, _>("total_requests")
.map_err(|err| internal(format!("aggregate heatmap request decode failed: {err}")))?;
let input_tokens = row
.try_get::<i64, _>("input_tokens")
.map_err(|err| internal(format!("aggregate heatmap token decode failed: {err}")))?;
let output_tokens = row
.try_get::<i64, _>("output_tokens")
.map_err(|err| internal(format!("aggregate heatmap token decode failed: {err}")))?;
let cache_creation_tokens = row
.try_get::<i64, _>("cache_creation_tokens")
.map_err(|err| internal(format!("aggregate heatmap token decode failed: {err}")))?;
let cache_read_tokens = row
.try_get::<i64, _>("cache_read_tokens")
.map_err(|err| internal(format!("aggregate heatmap token decode failed: {err}")))?;
let total_cost_usd = row
.try_get::<f64, _>("total_cost")
.map_err(|err| internal(format!("aggregate heatmap cost decode failed: {err}")))?;
let actual_total_cost_usd = row.try_get::<f64, _>("actual_total_cost").map_err(|err| {
internal(format!(
"aggregate heatmap actual cost decode failed: {err}"
))
})?;
items.push(StoredUsageDailySummary {
date: date.date_naive().to_string(),
requests: u64::try_from(requests.max(0)).unwrap_or_default(),
total_tokens: u64::try_from(
input_tokens
.saturating_add(output_tokens)
.saturating_add(cache_creation_tokens)
.saturating_add(cache_read_tokens)
.max(0),
)
.unwrap_or_default(),
total_cost_usd,
actual_total_cost_usd,
});
}
Ok(items)
}
#[cfg(test)]
mod tests {
use super::{GLOBAL_HEATMAP_AGGREGATE_SQL, USER_HEATMAP_AGGREGATE_SQL};
#[test]
fn user_heatmap_query_casts_cost_columns_to_double_precision() {
assert!(USER_HEATMAP_AGGREGATE_SQL
.contains("COALESCE(total_cost, 0)::DOUBLE PRECISION AS total_cost"));
assert!(USER_HEATMAP_AGGREGATE_SQL
.contains("COALESCE(actual_total_cost, 0)::DOUBLE PRECISION AS actual_total_cost"));
}
#[test]
fn global_heatmap_query_casts_cost_columns_to_double_precision() {
assert!(GLOBAL_HEATMAP_AGGREGATE_SQL
.contains("COALESCE(total_cost, 0)::DOUBLE PRECISION AS total_cost"));
assert!(GLOBAL_HEATMAP_AGGREGATE_SQL
.contains("COALESCE(actual_total_cost, 0)::DOUBLE PRECISION AS actual_total_cost"));
}
}

View File

@@ -1,154 +0,0 @@
use std::collections::BTreeMap;
use aether_data::postgres::PostgresPool;
use aether_data_contracts::repository::usage::StoredUsageUserTotals;
use chrono::{DateTime, Utc};
use futures_util::TryStreamExt;
use sqlx::Row;
use crate::query::usage_heatmap::read_stats_daily_cutoff_date;
use crate::GatewayError;
fn internal(err: impl ToString) -> GatewayError {
GatewayError::Internal(err.to_string())
}
pub(crate) async fn list_user_usage_totals_from_stats_summary(
pool: &PostgresPool,
user_ids: &[String],
) -> Result<Option<Vec<StoredUsageUserTotals>>, GatewayError> {
if user_ids.is_empty() {
return Ok(Some(Vec::new()));
}
let Some(cutoff_date) = read_stats_daily_cutoff_date(pool).await? else {
return Ok(None);
};
let mut totals = load_stats_user_summary_rows(pool, user_ids).await?;
absorb_stats_user_summary_tail(pool, cutoff_date, user_ids, &mut totals).await?;
let mut items = user_ids
.iter()
.map(|user_id| {
totals
.remove(user_id)
.unwrap_or_else(|| StoredUsageUserTotals {
user_id: user_id.clone(),
request_count: 0,
total_tokens: 0,
})
})
.collect::<Vec<_>>();
items.sort_by(|left, right| left.user_id.cmp(&right.user_id));
Ok(Some(items))
}
async fn load_stats_user_summary_rows(
pool: &PostgresPool,
user_ids: &[String],
) -> Result<BTreeMap<String, StoredUsageUserTotals>, GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
user_id,
COALESCE(all_time_requests, 0)::BIGINT AS request_count,
COALESCE(
all_time_input_tokens
+ all_time_output_tokens
+ all_time_cache_creation_tokens
+ all_time_cache_read_tokens,
0
)::BIGINT AS total_tokens
FROM stats_user_summary
WHERE user_id = ANY($1::TEXT[])
ORDER BY user_id ASC
"#,
)
.bind(user_ids)
.fetch(pool);
let mut items = BTreeMap::new();
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("stats_user_summary lookup failed: {err}")))?
{
let user_id = row
.try_get::<String, _>("user_id")
.map_err(|err| internal(format!("stats_user_summary decode failed: {err}")))?;
let request_count = row
.try_get::<i64, _>("request_count")
.map_err(|err| internal(format!("stats_user_summary decode failed: {err}")))?
.max(0) as u64;
let total_tokens = row
.try_get::<i64, _>("total_tokens")
.map_err(|err| internal(format!("stats_user_summary decode failed: {err}")))?
.max(0) as u64;
items.insert(
user_id.clone(),
StoredUsageUserTotals {
user_id,
request_count,
total_tokens,
},
);
}
Ok(items)
}
async fn absorb_stats_user_summary_tail(
pool: &PostgresPool,
cutoff_date: DateTime<Utc>,
user_ids: &[String],
totals: &mut BTreeMap<String, StoredUsageUserTotals>,
) -> Result<(), GatewayError> {
let mut rows = sqlx::query(
r#"
SELECT
"usage".user_id,
COUNT(*)::BIGINT AS request_count,
COALESCE(SUM(GREATEST(COALESCE("usage".total_tokens, 0), 0)), 0)::BIGINT AS total_tokens
FROM usage_billing_facts AS "usage"
WHERE "usage".user_id = ANY($1::TEXT[])
AND "usage".created_at >= $2
AND "usage".status NOT IN ('pending', 'streaming')
AND "usage".provider_name NOT IN ('unknown', 'pending')
GROUP BY "usage".user_id
ORDER BY "usage".user_id ASC
"#,
)
.bind(user_ids)
.bind(cutoff_date)
.fetch(pool);
while let Some(row) = rows
.try_next()
.await
.map_err(|err| internal(format!("stats_user_summary tail lookup failed: {err}")))?
{
let user_id = row
.try_get::<String, _>("user_id")
.map_err(|err| internal(format!("stats_user_summary tail decode failed: {err}")))?;
let request_count = row
.try_get::<i64, _>("request_count")
.map_err(|err| internal(format!("stats_user_summary tail decode failed: {err}")))?
.max(0) as u64;
let total_tokens = row
.try_get::<i64, _>("total_tokens")
.map_err(|err| internal(format!("stats_user_summary tail decode failed: {err}")))?
.max(0) as u64;
let entry = totals
.entry(user_id.clone())
.or_insert_with(|| StoredUsageUserTotals {
user_id,
request_count: 0,
total_tokens: 0,
});
entry.request_count = entry.request_count.saturating_add(request_count);
entry.total_tokens = entry.total_tokens.saturating_add(total_tokens);
}
Ok(())
}

View File

@@ -308,7 +308,7 @@ impl FrontdoorUserRpmLimiter {
async fn check_and_consume_redis(
&self,
runner: &aether_data::redis::RedisKvRunner,
runner: &aether_data::driver::redis::RedisKvRunner,
plan: &RpmPlan,
) -> Result<FrontdoorUserRpmOutcome, GatewayError> {
let user_key = runner.keyspace().key(&plan.user_rpm_key);

View File

@@ -4,6 +4,6 @@ pub(crate) use aether_data::repository::wallet::{
AdminWalletTransactionRecord,
};
pub(crate) use aether_data_contracts::repository::billing::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingPresetApplyResult,
AdminBillingRuleRecord, AdminBillingRuleWriteInput,
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
};

View File

@@ -253,52 +253,39 @@ impl AppState {
Ok(self)
}
pub async fn run_postgres_migrations(&self) -> Result<bool, sqlx::migrate::MigrateError> {
let Some(pool) = self.postgres_pool() else {
return Ok(false);
};
aether_data::migrate::run_migrations(&pool).await?;
Ok(true)
pub async fn run_database_migrations(&self) -> Result<bool, sqlx::migrate::MigrateError> {
self.data.run_database_migrations().await
}
pub async fn run_postgres_backfills(&self) -> Result<bool, sqlx::migrate::MigrateError> {
let Some(pool) = self.postgres_pool() else {
return Ok(false);
};
aether_data::backfill::run_backfills(&pool).await?;
Ok(true)
pub async fn run_database_backfills(&self) -> Result<bool, sqlx::migrate::MigrateError> {
self.data.run_database_backfills().await
}
pub async fn pending_postgres_migrations(
pub async fn pending_database_migrations(
&self,
) -> Result<Option<Vec<aether_data::migrate::PendingMigrationInfo>>, sqlx::migrate::MigrateError>
{
let Some(pool) = self.postgres_pool() else {
return Ok(None);
};
Ok(Some(aether_data::migrate::pending_migrations(&pool).await?))
) -> Result<
Option<Vec<aether_data::lifecycle::migrate::PendingMigrationInfo>>,
sqlx::migrate::MigrateError,
> {
self.data.pending_database_migrations().await
}
pub async fn prepare_postgres_for_startup(
pub async fn prepare_database_for_startup(
&self,
) -> Result<Option<Vec<aether_data::migrate::PendingMigrationInfo>>, sqlx::migrate::MigrateError>
{
let Some(pool) = self.postgres_pool() else {
return Ok(None);
};
Ok(Some(
aether_data::migrate::prepare_database_for_startup(&pool).await?,
))
) -> Result<
Option<Vec<aether_data::lifecycle::migrate::PendingMigrationInfo>>,
sqlx::migrate::MigrateError,
> {
self.data.prepare_database_for_startup().await
}
pub async fn pending_postgres_backfills(
pub async fn pending_database_backfills(
&self,
) -> Result<Option<Vec<aether_data::backfill::PendingBackfillInfo>>, sqlx::migrate::MigrateError>
{
let Some(pool) = self.postgres_pool() else {
return Ok(None);
};
Ok(Some(aether_data::backfill::pending_backfills(&pool).await?))
) -> Result<
Option<Vec<aether_data::lifecycle::backfill::PendingBackfillInfo>>,
sqlx::migrate::MigrateError,
> {
self.data.pending_database_backfills().await
}
pub fn with_video_task_poller_config(mut self, interval: Duration, batch_size: usize) -> Self {
@@ -753,14 +740,10 @@ impl AppState {
self.data.has_redis_backend()
}
pub(crate) fn redis_kv_runner(&self) -> Option<aether_data::redis::RedisKvRunner> {
pub(crate) fn redis_kv_runner(&self) -> Option<aether_data::driver::redis::RedisKvRunner> {
self.data.kv_runner()
}
pub(crate) fn postgres_pool(&self) -> Option<aether_data::postgres::PostgresPool> {
self.data.postgres_pool()
}
pub(crate) fn remove_scheduler_affinity_cache_entry(&self, cache_key: &str) -> bool {
self.scheduler_affinity_cache.remove(cache_key).is_some()
}

View File

@@ -18,10 +18,10 @@ mod types;
mod video;
pub(crate) use self::admin_types::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingPresetApplyResult,
AdminBillingRuleRecord, AdminBillingRuleWriteInput, AdminPaymentCallbackRecord,
AdminSecurityBlacklistEntry, AdminWalletPaymentOrderRecord, AdminWalletRefundRecord,
AdminWalletTransactionRecord,
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
AdminPaymentCallbackRecord, AdminSecurityBlacklistEntry, AdminWalletPaymentOrderRecord,
AdminWalletRefundRecord, AdminWalletTransactionRecord,
};
pub use self::app::AppState;
pub(crate) use self::cache::{

View File

@@ -155,6 +155,7 @@ impl AppState {
.map_err(|err| GatewayError::Internal(err.to_string()))
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn rotate_user_session_refresh_token(
&self,
user_id: &str,

View File

@@ -1,9 +1,21 @@
use super::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingPresetApplyResult,
AdminBillingRuleRecord, AdminBillingRuleWriteInput, AppState, GatewayError,
LocalMutationOutcome,
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, AppState,
GatewayError, LocalMutationOutcome,
};
use crate::query::billing as billing_query;
fn data_error(err: impl ToString) -> GatewayError {
GatewayError::Internal(err.to_string())
}
fn local_mutation_outcome<T>(outcome: AdminBillingMutationOutcome<T>) -> LocalMutationOutcome<T> {
match outcome {
AdminBillingMutationOutcome::Applied(value) => LocalMutationOutcome::Applied(value),
AdminBillingMutationOutcome::NotFound => LocalMutationOutcome::NotFound,
AdminBillingMutationOutcome::Invalid(detail) => LocalMutationOutcome::Invalid(detail),
AdminBillingMutationOutcome::Unavailable => LocalMutationOutcome::Unavailable,
}
}
impl AppState {
pub(crate) async fn admin_billing_enabled_default_value_exists(
@@ -30,17 +42,17 @@ impl AppState {
return Ok(exists);
}
let Some(pool) = self.postgres_pool() else {
return Ok(false);
};
billing_query::admin_billing_enabled_default_value_exists(
&pool,
api_format,
task_type,
dimension_name,
existing_id,
)
.await
Ok(self
.data
.admin_billing_enabled_default_value_exists(
api_format,
task_type,
dimension_name,
existing_id,
)
.await
.map_err(data_error)?
.unwrap_or(false))
}
pub(crate) async fn create_admin_billing_rule(
@@ -70,10 +82,11 @@ impl AppState {
return Ok(LocalMutationOutcome::Applied(record));
}
let Some(pool) = self.postgres_pool() else {
return Ok(LocalMutationOutcome::Unavailable);
};
billing_query::create_admin_billing_rule(&pool, input).await
self.data
.create_admin_billing_rule(input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
}
pub(crate) async fn list_admin_billing_rules(
@@ -111,13 +124,10 @@ impl AppState {
return Ok(Some((items, total)));
}
let Some(pool) = self.postgres_pool() else {
return Ok(None);
};
let (items, total) =
billing_query::list_admin_billing_rules(&pool, task_type, is_enabled, page, page_size)
.await?;
Ok(Some((items, total)))
self.data
.list_admin_billing_rules(task_type, is_enabled, page, page_size)
.await
.map_err(data_error)
}
pub(crate) async fn read_admin_billing_rule(
@@ -133,10 +143,10 @@ impl AppState {
.cloned());
}
let Some(pool) = self.postgres_pool() else {
return Ok(None);
};
billing_query::find_admin_billing_rule(&pool, rule_id).await
self.data
.find_admin_billing_rule(rule_id)
.await
.map_err(data_error)
}
pub(crate) async fn update_admin_billing_rule(
@@ -162,10 +172,11 @@ impl AppState {
return Ok(LocalMutationOutcome::Applied(record.clone()));
}
let Some(pool) = self.postgres_pool() else {
return Ok(LocalMutationOutcome::Unavailable);
};
billing_query::update_admin_billing_rule(&pool, rule_id, input).await
self.data
.update_admin_billing_rule(rule_id, input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
}
pub(crate) async fn create_admin_billing_collector(
@@ -197,10 +208,11 @@ impl AppState {
return Ok(LocalMutationOutcome::Applied(record));
}
let Some(pool) = self.postgres_pool() else {
return Ok(LocalMutationOutcome::Unavailable);
};
billing_query::create_admin_billing_collector(&pool, input).await
self.data
.create_admin_billing_collector(input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
}
pub(crate) async fn list_admin_billing_collectors(
@@ -243,20 +255,17 @@ impl AppState {
return Ok(Some((items, total)));
}
let Some(pool) = self.postgres_pool() else {
return Ok(None);
};
let (items, total) = billing_query::list_admin_billing_collectors(
&pool,
api_format,
task_type,
dimension_name,
is_enabled,
page,
page_size,
)
.await?;
Ok(Some((items, total)))
self.data
.list_admin_billing_collectors(
api_format,
task_type,
dimension_name,
is_enabled,
page,
page_size,
)
.await
.map_err(data_error)
}
pub(crate) async fn read_admin_billing_collector(
@@ -272,10 +281,10 @@ impl AppState {
.cloned());
}
let Some(pool) = self.postgres_pool() else {
return Ok(None);
};
billing_query::find_admin_billing_collector(&pool, collector_id).await
self.data
.find_admin_billing_collector(collector_id)
.await
.map_err(data_error)
}
pub(crate) async fn update_admin_billing_collector(
@@ -305,10 +314,11 @@ impl AppState {
return Ok(LocalMutationOutcome::Applied(record.clone()));
}
let Some(pool) = self.postgres_pool() else {
return Ok(LocalMutationOutcome::Unavailable);
};
billing_query::update_admin_billing_collector(&pool, collector_id, input).await
self.data
.update_admin_billing_collector(collector_id, input)
.await
.map(local_mutation_outcome)
.map_err(data_error)
}
pub(crate) async fn apply_admin_billing_preset(
@@ -389,9 +399,10 @@ impl AppState {
));
}
let Some(pool) = self.postgres_pool() else {
return Ok(LocalMutationOutcome::Unavailable);
};
billing_query::apply_admin_billing_preset(&pool, preset, mode, collectors).await
self.data
.apply_admin_billing_preset(preset, mode, collectors)
.await
.map(local_mutation_outcome)
.map_err(data_error)
}
}

View File

@@ -1,7 +1,7 @@
use super::super::{
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingPresetApplyResult,
AdminBillingRuleRecord, AdminBillingRuleWriteInput, AppState, GatewayError,
LocalMutationOutcome,
AdminBillingCollectorRecord, AdminBillingCollectorWriteInput, AdminBillingMutationOutcome,
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput, AppState,
GatewayError, LocalMutationOutcome,
};
mod admin;

View File

@@ -13,6 +13,7 @@ mod auth;
mod billing;
mod candidate_queries;
mod gemini_files;
mod monitoring;
mod payments;
mod security;
mod usage_queries;
@@ -77,13 +78,21 @@ impl AppState {
self.data.has_wallet_writer()
}
pub fn has_database_wallet_data_writer(&self) -> bool {
self.data.has_wallet_writer() && self.data.database_driver().is_some()
}
pub fn has_auth_user_write_capability(&self) -> bool {
#[cfg(test)]
if self.auth_user_store.is_some() {
return true;
}
#[cfg(test)]
if !self.data.has_backends() {
return false;
}
self.postgres_pool().is_some()
self.data.has_user_reader()
}
pub fn has_auth_wallet_write_capability(&self) -> bool {
@@ -92,7 +101,7 @@ impl AppState {
return true;
}
self.postgres_pool().is_some()
self.data.has_wallet_writer()
}
pub fn has_provider_quota_data_writer(&self) -> bool {

View File

@@ -0,0 +1,136 @@
use std::collections::BTreeMap;
use aether_data::repository::audit::AuditLogListQuery;
use chrono::{DateTime, Utc};
use serde_json::{json, Value};
use super::{AppState, GatewayError};
impl AppState {
pub(crate) async fn list_admin_audit_logs(
&self,
cutoff_time: DateTime<Utc>,
username_pattern: Option<&str>,
event_type: Option<&str>,
limit: usize,
offset: usize,
) -> Result<(Vec<Value>, usize), GatewayError> {
let query = AuditLogListQuery {
cutoff_unix_secs: cutoff_unix_secs(cutoff_time),
username_pattern: username_pattern.map(str::to_string),
event_type: event_type.map(str::to_string),
limit,
offset,
};
let page = self
.data
.list_admin_audit_logs(&query)
.await
.map_err(|err| {
GatewayError::Internal(format!("admin audit logs read failed: {err}"))
})?;
let total = usize::try_from(page.total).unwrap_or(usize::MAX);
let items = page
.items
.iter()
.map(|record| {
json!({
"id": record.id,
"event_type": record.event_type,
"user_id": record.user_id,
"user_email": record.user_email,
"user_username": record.user_username,
"description": record.description,
"ip_address": record.ip_address,
"status_code": record.status_code,
"error_message": record.error_message,
"metadata": record.metadata,
"created_at": record.created_at_rfc3339(),
})
})
.collect();
Ok((items, total))
}
pub(crate) async fn list_admin_suspicious_activities(
&self,
cutoff_time: DateTime<Utc>,
) -> Result<Vec<Value>, GatewayError> {
let activities = self
.data
.list_admin_suspicious_activities(cutoff_unix_secs(cutoff_time))
.await
.map_err(|err| {
GatewayError::Internal(format!("admin suspicious activities read failed: {err}"))
})?;
Ok(activities
.iter()
.map(|record| {
json!({
"id": record.id,
"event_type": record.event_type,
"user_id": record.user_id,
"description": record.description,
"ip_address": record.ip_address,
"metadata": record.metadata,
"created_at": record.created_at_rfc3339(),
})
})
.collect())
}
pub(crate) async fn read_admin_user_behavior_event_counts(
&self,
user_id: &str,
cutoff_time: DateTime<Utc>,
) -> Result<BTreeMap<String, u64>, GatewayError> {
self.data
.read_admin_user_behavior_event_counts(user_id, cutoff_unix_secs(cutoff_time))
.await
.map_err(|err| {
GatewayError::Internal(format!("admin user behavior read failed: {err}"))
})
}
pub(crate) async fn list_user_audit_logs(
&self,
user_id: &str,
cutoff_time: DateTime<Utc>,
event_type: Option<&str>,
limit: usize,
offset: usize,
) -> Result<(Vec<Value>, usize), GatewayError> {
let query = AuditLogListQuery {
cutoff_unix_secs: cutoff_unix_secs(cutoff_time),
username_pattern: None,
event_type: event_type.map(str::to_string),
limit,
offset,
};
let page = self
.data
.list_user_audit_logs(user_id, &query)
.await
.map_err(|err| GatewayError::Internal(format!("user audit logs read failed: {err}")))?;
let total = usize::try_from(page.total).unwrap_or(usize::MAX);
let items = page
.items
.iter()
.map(|record| {
json!({
"id": record.id,
"event_type": record.event_type,
"description": record.description,
"ip_address": record.ip_address,
"status_code": record.status_code,
"created_at": record.created_at_rfc3339(),
})
})
.collect();
Ok((items, total))
}
}
fn cutoff_unix_secs(cutoff_time: DateTime<Utc>) -> u64 {
cutoff_time.timestamp().max(0) as u64
}

View File

@@ -1,10 +1,3 @@
use sqlx::Row;
use super::{
AdminBillingCollectorRecord, AdminBillingRuleRecord, AdminWalletPaymentOrderRecord,
AdminWalletRefundRecord, GatewayError,
};
pub(crate) fn admin_wallet_build_order_no(now: chrono::DateTime<chrono::Utc>) -> String {
format!(
"po_{}_{}",
@@ -21,274 +14,3 @@ pub(crate) fn admin_payment_gateway_response_map(
_ => serde_json::Map::new(),
}
}
pub(super) fn admin_wallet_snapshot_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<aether_data::repository::wallet::StoredWalletSnapshot, GatewayError> {
aether_data::repository::wallet::StoredWalletSnapshot::new(
row.try_get("id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("user_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("api_key_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("balance")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("gift_balance")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("limit_mode")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("currency")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("status")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("total_recharged")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("total_consumed")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("total_refunded")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("total_adjusted")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
row.try_get("updated_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
)
.map_err(|err| GatewayError::Internal(err.to_string()))
}
pub(super) fn admin_wallet_payment_order_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<AdminWalletPaymentOrderRecord, GatewayError> {
Ok(AdminWalletPaymentOrderRecord {
id: row
.try_get("id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
order_no: row
.try_get("order_no")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
wallet_id: row
.try_get("wallet_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
user_id: row
.try_get("user_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
amount_usd: row
.try_get("amount_usd")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
pay_amount: row
.try_get("pay_amount")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
pay_currency: row
.try_get("pay_currency")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
exchange_rate: row
.try_get("exchange_rate")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
refunded_amount_usd: row
.try_get("refunded_amount_usd")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
refundable_amount_usd: row
.try_get("refundable_amount_usd")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
payment_method: row
.try_get("payment_method")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
gateway_order_id: row
.try_get("gateway_order_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
status: row
.try_get("status")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
gateway_response: row
.try_get("gateway_response")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.max(0) as u64,
paid_at_unix_secs: row
.try_get::<Option<i64>, _>("paid_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.map(|value| value.max(0) as u64),
credited_at_unix_secs: row
.try_get::<Option<i64>, _>("credited_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.map(|value| value.max(0) as u64),
expires_at_unix_secs: row
.try_get::<Option<i64>, _>("expires_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.map(|value| value.max(0) as u64),
})
}
pub(super) fn admin_wallet_refund_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<AdminWalletRefundRecord, GatewayError> {
Ok(AdminWalletRefundRecord {
id: row
.try_get("id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
refund_no: row
.try_get("refund_no")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
wallet_id: row
.try_get("wallet_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
user_id: row
.try_get("user_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
payment_order_id: row
.try_get("payment_order_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
source_type: row
.try_get("source_type")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
source_id: row
.try_get("source_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
refund_mode: row
.try_get("refund_mode")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
amount_usd: row
.try_get("amount_usd")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
status: row
.try_get("status")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
reason: row
.try_get("reason")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
failure_reason: row
.try_get("failure_reason")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
gateway_refund_id: row
.try_get("gateway_refund_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
payout_method: row
.try_get("payout_method")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
payout_reference: row
.try_get("payout_reference")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
payout_proof: row
.try_get("payout_proof")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
requested_by: row
.try_get("requested_by")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
approved_by: row
.try_get("approved_by")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
processed_by: row
.try_get("processed_by")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.max(0) as u64,
updated_at_unix_secs: row
.try_get::<i64, _>("updated_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.max(0) as u64,
processed_at_unix_secs: row
.try_get::<Option<i64>, _>("processed_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.map(|value| value.max(0) as u64),
completed_at_unix_secs: row
.try_get::<Option<i64>, _>("completed_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.map(|value| value.max(0) as u64),
})
}
pub(super) fn admin_billing_rule_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<AdminBillingRuleRecord, GatewayError> {
Ok(AdminBillingRuleRecord {
id: row
.try_get("id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
name: row
.try_get("name")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
task_type: row
.try_get("task_type")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
global_model_id: row
.try_get("global_model_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
model_id: row
.try_get("model_id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
expression: row
.try_get("expression")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
variables: row
.try_get::<Option<serde_json::Value>, _>("variables")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.unwrap_or_else(|| serde_json::json!({})),
dimension_mappings: row
.try_get::<Option<serde_json::Value>, _>("dimension_mappings")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.unwrap_or_else(|| serde_json::json!({})),
is_enabled: row
.try_get("is_enabled")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.max(0) as u64,
updated_at_unix_secs: row
.try_get::<i64, _>("updated_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.max(0) as u64,
})
}
pub(super) fn admin_billing_collector_from_row(
row: &sqlx::postgres::PgRow,
) -> Result<AdminBillingCollectorRecord, GatewayError> {
Ok(AdminBillingCollectorRecord {
id: row
.try_get("id")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
api_format: row
.try_get("api_format")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
task_type: row
.try_get("task_type")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
dimension_name: row
.try_get("dimension_name")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
source_type: row
.try_get("source_type")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
source_path: row
.try_get("source_path")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
value_type: row
.try_get("value_type")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
transform_expression: row
.try_get("transform_expression")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
default_value: row
.try_get("default_value")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
priority: row
.try_get("priority")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
is_enabled: row
.try_get("is_enabled")
.map_err(|err| GatewayError::Internal(err.to_string()))?,
created_at_unix_ms: row
.try_get::<i64, _>("created_at_unix_ms")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.max(0) as u64,
updated_at_unix_secs: row
.try_get::<i64, _>("updated_at_unix_secs")
.map_err(|err| GatewayError::Internal(err.to_string()))?
.max(0) as u64,
})
}

View File

@@ -119,7 +119,7 @@ fn admin_wrapped_state_owns_billing_capabilities() {
let admin_request =
read_workspace_module_tree("apps/aether-gateway/src/handlers/admin/request/mod.rs");
for pattern in [
"pub(crate) fn has_postgres_pool(&self) -> bool",
"pub(crate) fn has_wallet_data_writer(&self) -> bool",
"pub(crate) async fn list_admin_billing_collectors(",
"pub(crate) async fn read_admin_billing_collector(",
"pub(crate) async fn create_admin_billing_collector(",

View File

@@ -4701,6 +4701,7 @@ fn retired_api_format_occurrences_are_whitelisted() {
"crates/aether-ai-formats/src/protocol/matrix.rs",
"crates/aether-ai-formats/src/protocol/registry.rs",
"crates/aether-data/src/migrate.rs",
"crates/aether-data/src/lifecycle/migrate/tests.rs",
"crates/aether-usage-runtime/src/report.rs",
"frontend/src/api/endpoints/types/__tests__/api-format.spec.ts",
];

View File

@@ -23,6 +23,10 @@ pub(super) fn assert_no_sqlx_queries(root_relative_path: &str) {
let patterns = [
"sqlx::query(",
"sqlx::query_scalar",
"sqlx::postgres::PgRow",
"sqlx::Row",
"PostgresPoolFactory",
"PostgresPool",
"query_scalar::<",
"QueryBuilder<",
];

View File

@@ -1,5 +1,9 @@
use super::*;
fn production_source(source: &str) -> &str {
source.split("#[cfg(test)]").next().unwrap_or(source)
}
#[test]
fn handlers_do_not_inline_sql_queries() {
assert_no_sqlx_queries("src/handlers");
@@ -10,11 +14,157 @@ fn gateway_runtime_does_not_inline_sql_queries() {
assert_no_sqlx_queries("src/state/runtime");
}
#[test]
fn aether_data_bootstrap_snapshot_is_built_from_schema_sources() {
let build_rs = read_workspace_file("crates/aether-data/build.rs");
assert!(
build_rs.contains("schema/bootstrap/postgres/manifest.txt"),
"build.rs should source the bootstrap snapshot from schema/bootstrap/postgres"
);
let compose_schema = read_workspace_file("crates/aether-data/schema/compose_schema.sh");
assert!(
compose_schema.contains("check_bootstrap_sources"),
"compose_schema.sh should still validate bootstrap source fragments"
);
assert!(
!compose_schema.contains("bootstrap/postgres/20260413020000_empty_database_snapshot.sql"),
"compose_schema.sh should not depend on the outer bootstrap artifact anymore"
);
let bootstrap = read_workspace_file("crates/aether-data/src/lifecycle/bootstrap/postgres.rs");
assert!(
bootstrap.contains("include_str!(concat!(env!(\"OUT_DIR\"), \"/empty_database_snapshot.sql\"))"),
"lifecycle/bootstrap/postgres.rs should embed the generated bootstrap snapshot from OUT_DIR"
);
assert!(
!bootstrap
.contains("../../../bootstrap/postgres/20260413020000_empty_database_snapshot.sql"),
"lifecycle/bootstrap/postgres.rs should not read the outer bootstrap artifact directly"
);
let provider_catalog =
read_workspace_file("crates/aether-data/src/repository/provider_catalog/postgres.rs");
assert!(
!provider_catalog.contains("../../../bootstrap/postgres/20260413020000_empty_database_snapshot.sql"),
"provider_catalog tests should use the shared bootstrap snapshot constant instead of the outer bootstrap artifact"
);
}
#[test]
fn aether_data_backend_pool_modules_do_not_own_maintenance_sql() {
for path in [
"crates/aether-data/src/backend/postgres.rs",
"crates/aether-data/src/backend/mysql.rs",
"crates/aether-data/src/backend/sqlite.rs",
] {
let source = read_workspace_file(path);
let production = production_source(&source);
for forbidden in [
"run_table_maintenance(",
"aggregate_wallet_daily_usage(",
"aggregate_stats_hourly(",
"aggregate_stats_daily(",
"find_system_config_value(",
"list_system_config_entries(",
"upsert_system_config_entry(",
"read_admin_system_stats(",
"sqlx::query(",
"sqlx::query_scalar",
"sqlx::raw_sql(",
] {
assert!(
!production.contains(forbidden),
"{path} should stay focused on pool and repository construction instead of owning maintenance SQL via {forbidden}"
);
}
}
let maintenance = read_workspace_file("crates/aether-data/src/backend/maintenance.rs");
for pattern in [
"Self::Postgres(postgres) => postgres.run_table_maintenance(table_names).await",
"Self::Mysql(mysql) => mysql.run_table_maintenance(table_names).await",
"Self::Sqlite(sqlite) => sqlite.run_table_maintenance(table_names).await",
"Self::Postgres(postgres) => postgres.aggregate_wallet_daily_usage(input).await",
"Self::Mysql(mysql) => mysql.aggregate_wallet_daily_usage(input).await",
"Self::Sqlite(sqlite) => sqlite.aggregate_wallet_daily_usage(input).await",
"Self::Postgres(postgres) => postgres.aggregate_stats_hourly(input).await",
"Self::Mysql(mysql) => mysql.aggregate_stats_hourly(input).await",
"Self::Sqlite(sqlite) => sqlite.aggregate_stats_hourly(input).await",
"Self::Postgres(postgres) => postgres.aggregate_stats_daily(input).await",
"Self::Mysql(mysql) => mysql.aggregate_stats_daily(input).await",
"Self::Sqlite(sqlite) => sqlite.aggregate_stats_daily(input).await",
] {
assert!(
maintenance.contains(pattern),
"backend/maintenance.rs should own SQL-driver maintenance dispatch {pattern}"
);
}
}
#[test]
fn testkit_does_not_copy_aether_business_schema_sql() {
let owner_relay_baseline =
read_workspace_file("crates/aether-testkit/src/bin/multi_instance_owner_relay_baseline.rs");
for forbidden in [
"CREATE TYPE proxynodestatus",
"CREATE TABLE IF NOT EXISTS system_configs",
"CREATE TABLE IF NOT EXISTS proxy_nodes",
"CREATE TABLE IF NOT EXISTS proxy_node_events",
"PgConnection::connect",
"sqlx::{Connection, Executor, PgConnection}",
] {
assert!(
!owner_relay_baseline.contains(forbidden),
"owner relay baseline should use aether-data schema bootstrap instead of copying business schema SQL via {forbidden}"
);
}
assert!(
owner_relay_baseline.contains("prepare_aether_postgres_schema(&postgres_url).await?"),
"owner relay baseline should prepare business schema through aether-testkit's aether-data helper"
);
let postgres_testkit = read_workspace_file("crates/aether-testkit/src/postgres.rs");
for required in [
"pub async fn prepare_aether_postgres_schema",
"DataBackends::from_config",
".prepare_database_for_startup()",
".run_database_migrations()",
] {
assert!(
postgres_testkit.contains(required),
"testkit Postgres helper should delegate Aether schema setup to aether-data via {required}"
);
}
}
#[test]
fn gateway_main_keeps_database_export_import_driver_selection_in_data_layer() {
let main_rs = read_workspace_file("apps/aether-gateway/src/main.rs");
for forbidden in [
"PostgresPoolFactory",
"MysqlPoolFactory",
"SqlitePoolFactory",
"to_postgres_config()",
] {
assert!(
!main_rs.contains(forbidden),
"main.rs should delegate database export/import driver selection to aether-data instead of {forbidden}"
);
}
for required in ["export_database_jsonl", "import_database_jsonl"] {
assert!(
main_rs.contains(required),
"main.rs should use aether-data {required}"
);
}
}
#[test]
fn wallet_repository_does_not_reexport_settlement_types() {
let wallet_mod = read_workspace_file("crates/aether-data/src/repository/wallet/mod.rs");
let wallet_types = read_workspace_file("crates/aether-data/src/repository/wallet/types.rs");
let wallet_sql = read_workspace_file("crates/aether-data/src/repository/wallet/sql.rs");
let wallet_sql = read_workspace_file("crates/aether-data/src/repository/wallet/postgres.rs");
let wallet_memory = read_workspace_file("crates/aether-data/src/repository/wallet/memory.rs");
assert!(
@@ -35,7 +185,7 @@ fn wallet_repository_does_not_reexport_settlement_types() {
);
assert!(
!wallet_sql.contains("impl SettlementWriteRepository"),
"wallet/sql.rs should not implement SettlementWriteRepository"
"wallet/postgres.rs should not implement SettlementWriteRepository"
);
assert!(
!wallet_memory.contains("impl SettlementWriteRepository"),
@@ -57,8 +207,9 @@ fn gateway_system_config_types_are_owned_by_aether_data() {
let state_core = read_workspace_file("apps/aether-gateway/src/data/state/core.rs");
for pattern in [
"backend.list_system_config_entries().await",
"upsert_system_config_entry(key, value, description)",
"backends.list_system_config_entries().await",
".upsert_system_config_entry(key, value, description)",
"backends.read_admin_system_stats().await",
"AdminSystemStats::default()",
] {
assert!(
@@ -66,6 +217,17 @@ fn gateway_system_config_types_are_owned_by_aether_data() {
"data/state/core.rs should use shared system DTO path {pattern}"
);
}
let data_backends = read_workspace_file("crates/aether-data/src/backend/maintenance.rs");
for pattern in [
"postgres.list_system_config_entries().await",
"mysql.list_system_config_entries().await",
"sqlite.list_system_config_entries().await",
] {
assert!(
data_backends.contains(pattern),
"aether-data backends should own driver-specific system config dispatch {pattern}"
);
}
for pattern in [
"|(key, value, description, updated_at_unix_secs)|",
"Ok((0, 0, 0, 0))",
@@ -180,6 +342,7 @@ fn gateway_auth_data_layer_does_not_keep_ldap_row_wrapper() {
"fn map_ldap_user_auth_row(",
"Result<Option<StoredLdapAuthUserRow>, DataLayerError>",
"existing.user.",
"map_user_auth_row(row)",
] {
assert!(
!gateway_auth_state.contains(pattern),
@@ -187,13 +350,14 @@ fn gateway_auth_data_layer_does_not_keep_ldap_row_wrapper() {
);
}
let user_sql = read_workspace_file("crates/aether-data/src/repository/users/postgres.rs");
for pattern in [
"Result<Option<StoredUserAuthRecord>, DataLayerError>",
"return map_user_auth_row(row).map(Some);",
] {
assert!(
gateway_auth_state.contains(pattern),
"data/state/auth.rs should use shared user auth record directly via {pattern}"
user_sql.contains(pattern),
"aether-data user repository should use shared user auth record directly via {pattern}"
);
}
}

View File

@@ -909,7 +909,7 @@ async fn gateway_returns_conflict_for_admin_create_user_when_writer_unavailable(
let mut state = AppState::new().expect("gateway should build");
state =
state.with_data_state_for_tests(crate::data::GatewayDataState::with_user_reader_for_tests(
Arc::new(InMemoryUserReadRepository::seed_auth_users(Vec::new())),
Arc::new(InMemoryUserReadRepository::seed_auth_users(Vec::new()).read_only()),
));
state.auth_user_store = None;
state.auth_wallet_store = None;
@@ -1032,9 +1032,9 @@ async fn gateway_returns_conflict_for_admin_update_user_when_writer_unavailable(
}
}));
let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![
sample_admin_user("user-1"),
]));
let user_repository = Arc::new(
InMemoryUserReadRepository::seed_auth_users(vec![sample_admin_user("user-1")]).read_only(),
);
let (upstream_url, upstream_handle) = start_server(upstream).await;
let mut state = AppState::new()

View File

@@ -4215,7 +4215,8 @@ async fn gateway_returns_service_unavailable_for_users_me_detail_update_without_
]),
now + chrono::Duration::hours(1),
);
let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]));
let user_repository =
Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]).read_only());
let (gateway_url, upstream_hits, gateway_handle, upstream_handle) =
start_auth_gateway_with_builder(|| {
let data_state =
@@ -4365,7 +4366,8 @@ async fn gateway_returns_service_unavailable_for_users_me_password_change_withou
]),
now + chrono::Duration::hours(1),
);
let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]));
let user_repository =
Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]).read_only());
let (gateway_url, upstream_hits, gateway_handle, upstream_handle) =
start_auth_gateway_with_builder(|| {
let data_state =
@@ -4731,7 +4733,8 @@ async fn gateway_returns_service_unavailable_for_users_me_preferences_update_wit
]),
now + chrono::Duration::hours(1),
);
let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]));
let user_repository =
Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]).read_only());
let (gateway_url, upstream_hits, gateway_handle, upstream_handle) =
start_auth_gateway_with_builder(|| {
let data_state =
@@ -8210,7 +8213,8 @@ async fn gateway_returns_service_unavailable_for_users_me_model_capabilities_upd
]),
now + chrono::Duration::hours(1),
);
let user_repository = Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]));
let user_repository =
Arc::new(InMemoryUserReadRepository::seed_auth_users(vec![user]).read_only());
let (gateway_url, upstream_hits, gateway_handle, upstream_handle) =
start_auth_gateway_with_builder(|| {
let data_state =