Merge remote-tracking branch 'origin/aether-rust-pioneer' into codex/async-cleanup-records

# Conflicts:
#	apps/aether-gateway/src/maintenance/mod.rs
#	apps/aether-gateway/src/maintenance/runtime/runners.rs
This commit is contained in:
fawney19
2026-05-10 00:28:39 +08:00
315 changed files with 19750 additions and 3183 deletions

View File

@@ -13,6 +13,9 @@ use crate::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, MysqlAuthModuleReadRepository,
MysqlAuthModuleRepository,
};
use crate::repository::background_tasks::{
BackgroundTaskReadRepository, BackgroundTaskWriteRepository, MysqlBackgroundTaskRepository,
};
use crate::repository::billing::{BillingReadRepository, MysqlBillingReadRepository};
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, MysqlMinimalCandidateSelectionReadRepository,
@@ -123,6 +126,14 @@ impl MysqlBackend {
Arc::new(MysqlBillingReadRepository::new(self.pool_clone()))
}
pub fn background_task_read_repository(&self) -> Arc<dyn BackgroundTaskReadRepository> {
Arc::new(MysqlBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn background_task_write_repository(&self) -> Arc<dyn BackgroundTaskWriteRepository> {
Arc::new(MysqlBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn request_candidate_read_repository(&self) -> Arc<dyn RequestCandidateReadRepository> {
Arc::new(MysqlRequestCandidateRepository::new(self.pool_clone()))
}

View File

@@ -15,6 +15,9 @@ use crate::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, SqlxAuthModuleReadRepository,
SqlxAuthModuleRepository,
};
use crate::repository::background_tasks::{
BackgroundTaskReadRepository, BackgroundTaskWriteRepository, SqlxBackgroundTaskRepository,
};
use crate::repository::billing::{BillingReadRepository, SqlxBillingReadRepository};
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, SqlxMinimalCandidateSelectionReadRepository,
@@ -118,6 +121,14 @@ impl PostgresBackend {
Arc::new(SqlxBillingReadRepository::new(self.pool_clone()))
}
pub fn background_task_read_repository(&self) -> Arc<dyn BackgroundTaskReadRepository> {
Arc::new(SqlxBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn background_task_write_repository(&self) -> Arc<dyn BackgroundTaskWriteRepository> {
Arc::new(SqlxBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn minimal_candidate_selection_read_repository(
&self,
) -> Arc<dyn MinimalCandidateSelectionReadRepository> {

View File

@@ -6,6 +6,7 @@ use crate::repository::announcements::AnnouncementReadRepository;
use crate::repository::audit::AuditLogReadRepository;
use crate::repository::auth::AuthApiKeyReadRepository;
use crate::repository::auth_modules::AuthModuleReadRepository;
use crate::repository::background_tasks::BackgroundTaskReadRepository;
use crate::repository::billing::BillingReadRepository;
use crate::repository::candidate_selection::MinimalCandidateSelectionReadRepository;
use crate::repository::candidates::RequestCandidateReadRepository;
@@ -27,6 +28,7 @@ pub struct DataReadRepositories {
audit_logs: Option<Arc<dyn AuditLogReadRepository>>,
auth_api_keys: Option<Arc<dyn AuthApiKeyReadRepository>>,
auth_modules: Option<Arc<dyn AuthModuleReadRepository>>,
background_tasks: Option<Arc<dyn BackgroundTaskReadRepository>>,
billing: Option<Arc<dyn BillingReadRepository>>,
gemini_file_mappings: Option<Arc<dyn GeminiFileMappingReadRepository>>,
global_models: Option<Arc<dyn GlobalModelReadRepository>>,
@@ -50,6 +52,7 @@ impl fmt::Debug for DataReadRepositories {
.field("has_announcements", &self.announcements.is_some())
.field("has_audit_logs", &self.audit_logs.is_some())
.field("has_auth_modules", &self.auth_modules.is_some())
.field("has_background_tasks", &self.background_tasks.is_some())
.field("has_billing", &self.billing.is_some())
.field(
"has_gemini_file_mappings",
@@ -97,6 +100,10 @@ impl DataReadRepositories {
.map(PostgresBackend::auth_module_read_repository)
.or_else(|| mysql.map(MysqlBackend::auth_module_read_repository))
.or_else(|| sqlite.map(SqliteBackend::auth_module_read_repository)),
background_tasks: postgres
.map(PostgresBackend::background_task_read_repository)
.or_else(|| mysql.map(MysqlBackend::background_task_read_repository))
.or_else(|| sqlite.map(SqliteBackend::background_task_read_repository)),
billing: postgres
.map(PostgresBackend::billing_read_repository)
.or_else(|| mysql.map(MysqlBackend::billing_read_repository))
@@ -177,6 +184,10 @@ impl DataReadRepositories {
self.auth_modules.clone()
}
pub fn background_tasks(&self) -> Option<Arc<dyn BackgroundTaskReadRepository>> {
self.background_tasks.clone()
}
pub fn billing(&self) -> Option<Arc<dyn BillingReadRepository>> {
self.billing.clone()
}
@@ -240,6 +251,7 @@ impl DataReadRepositories {
|| self.announcements.is_some()
|| self.audit_logs.is_some()
|| self.auth_modules.is_some()
|| self.background_tasks.is_some()
|| self.billing.is_some()
|| self.gemini_file_mappings.is_some()
|| self.global_models.is_some()

View File

@@ -13,6 +13,9 @@ use crate::repository::auth_modules::{
AuthModuleReadRepository, AuthModuleWriteRepository, SqliteAuthModuleReadRepository,
SqliteAuthModuleRepository,
};
use crate::repository::background_tasks::{
BackgroundTaskReadRepository, BackgroundTaskWriteRepository, SqliteBackgroundTaskRepository,
};
use crate::repository::billing::{BillingReadRepository, SqliteBillingReadRepository};
use crate::repository::candidate_selection::{
MinimalCandidateSelectionReadRepository, SqliteMinimalCandidateSelectionReadRepository,
@@ -124,6 +127,14 @@ impl SqliteBackend {
Arc::new(SqliteBillingReadRepository::new(self.pool_clone()))
}
pub fn background_task_read_repository(&self) -> Arc<dyn BackgroundTaskReadRepository> {
Arc::new(SqliteBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn background_task_write_repository(&self) -> Arc<dyn BackgroundTaskWriteRepository> {
Arc::new(SqliteBackgroundTaskRepository::new(self.pool_clone()))
}
pub fn request_candidate_read_repository(&self) -> Arc<dyn RequestCandidateReadRepository> {
Arc::new(SqliteRequestCandidateRepository::new(self.pool_clone()))
}

View File

@@ -3,7 +3,7 @@ use sqlx::Row;
use crate::backend::stats_common::{stats_id, unix_ms, unix_secs, utc_from_unix_secs};
use crate::backend::SqliteBackend;
use crate::driver::sqlite::SqlitePool;
use crate::driver::sqlite::{sqlite_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::{
DataLayerError, StatsDailyAggregationInput, StatsDailyAggregationSummary,
@@ -124,9 +124,9 @@ SELECT
COALESCE(SUM(output_tokens), 0) AS output_tokens,
COALESCE(SUM(cache_creation_input_tokens), 0) AS cache_creation_tokens,
COALESCE(SUM(cache_read_input_tokens), 0) AS cache_read_tokens,
COALESCE(SUM(total_cost_usd), 0.0) AS total_cost,
COALESCE(SUM(actual_total_cost_usd), 0.0) AS actual_total_cost,
COALESCE(AVG(response_time_ms), 0.0) AS avg_response_time_ms
CAST(COALESCE(SUM(total_cost_usd), 0) AS REAL) AS total_cost,
CAST(COALESCE(SUM(actual_total_cost_usd), 0) AS REAL) AS actual_total_cost,
CAST(COALESCE(AVG(response_time_ms), 0) AS REAL) AS avg_response_time_ms
FROM "usage"
WHERE created_at_unix_ms >= ?
AND created_at_unix_ms < ?
@@ -188,12 +188,9 @@ ON CONFLICT (hour_utc) DO UPDATE SET
.map_sql_err()?,
)
.bind(row.try_get::<i64, _>("cache_read_tokens").map_sql_err()?)
.bind(row.try_get::<f64, _>("total_cost").map_sql_err()?)
.bind(row.try_get::<f64, _>("actual_total_cost").map_sql_err()?)
.bind(
row.try_get::<f64, _>("avg_response_time_ms")
.map_sql_err()?,
)
.bind(sqlite_real(&row, "total_cost")?)
.bind(sqlite_real(&row, "actual_total_cost")?)
.bind(sqlite_real(&row, "avg_response_time_ms")?)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
.bind(aggregated_at_unix_secs)
@@ -277,12 +274,9 @@ ON CONFLICT ("date") DO UPDATE SET
.map_sql_err()?,
)
.bind(row.try_get::<i64, _>("cache_read_tokens").map_sql_err()?)
.bind(row.try_get::<f64, _>("total_cost").map_sql_err()?)
.bind(row.try_get::<f64, _>("actual_total_cost").map_sql_err()?)
.bind(
row.try_get::<f64, _>("avg_response_time_ms")
.map_sql_err()?,
)
.bind(sqlite_real(&row, "total_cost")?)
.bind(sqlite_real(&row, "actual_total_cost")?)
.bind(sqlite_real(&row, "avg_response_time_ms")?)
.bind(unique_models)
.bind(unique_providers)
.bind(aggregated_at_unix_secs)

View File

@@ -2,6 +2,7 @@ use sha2::{Digest, Sha256};
use sqlx::Row;
use crate::backend::{MysqlBackend, PostgresBackend, SqliteBackend};
use crate::driver::sqlite::sqlite_real;
use crate::error::{SqlResultExt, SqlxResultExt};
use crate::{DataLayerError, WalletDailyUsageAggregationInput, WalletDailyUsageAggregationResult};
@@ -119,7 +120,7 @@ const SQLITE_SELECT_WALLET_DAILY_USAGE_AGGREGATES_SQL: &str = r#"
SELECT
usage_settlement_snapshots.wallet_id AS wallet_id,
COUNT(*) AS total_requests,
COALESCE(SUM("usage".total_cost_usd), 0) AS total_cost_usd,
CAST(COALESCE(SUM("usage".total_cost_usd), 0) AS REAL) AS total_cost_usd,
COALESCE(SUM("usage".input_tokens), 0) AS input_tokens,
COALESCE(SUM("usage".output_tokens), 0) AS output_tokens,
COALESCE(SUM("usage".cache_creation_input_tokens), 0) AS cache_creation_tokens,
@@ -400,7 +401,7 @@ INSERT INTO wallet_daily_usage_ledgers (
.bind(&wallet_id)
.bind(&input.billing_date)
.bind(&input.billing_timezone)
.bind(row.try_get::<f64, _>("total_cost_usd").map_sql_err()?)
.bind(sqlite_real(&row, "total_cost_usd")?)
.bind(row.try_get::<i64, _>("total_requests").map_sql_err()?)
.bind(row.try_get::<i64, _>("input_tokens").map_sql_err()?)
.bind(row.try_get::<i64, _>("output_tokens").map_sql_err()?)

View File

@@ -5,6 +5,7 @@ use super::{MysqlBackend, PostgresBackend, SqliteBackend};
use crate::repository::announcements::AnnouncementWriteRepository;
use crate::repository::auth::AuthApiKeyWriteRepository;
use crate::repository::auth_modules::AuthModuleWriteRepository;
use crate::repository::background_tasks::BackgroundTaskWriteRepository;
use crate::repository::candidates::RequestCandidateWriteRepository;
use crate::repository::gemini_file_mappings::GeminiFileMappingWriteRepository;
use crate::repository::global_models::GlobalModelWriteRepository;
@@ -23,6 +24,7 @@ pub struct DataWriteRepositories {
announcements: Option<Arc<dyn AnnouncementWriteRepository>>,
auth_api_keys: Option<Arc<dyn AuthApiKeyWriteRepository>>,
auth_modules: Option<Arc<dyn AuthModuleWriteRepository>>,
background_tasks: Option<Arc<dyn BackgroundTaskWriteRepository>>,
request_candidates: Option<Arc<dyn RequestCandidateWriteRepository>>,
gemini_file_mappings: Option<Arc<dyn GeminiFileMappingWriteRepository>>,
global_models: Option<Arc<dyn GlobalModelWriteRepository>>,
@@ -43,6 +45,7 @@ impl fmt::Debug for DataWriteRepositories {
.field("has_announcements", &self.announcements.is_some())
.field("has_auth_api_keys", &self.auth_api_keys.is_some())
.field("has_auth_modules", &self.auth_modules.is_some())
.field("has_background_tasks", &self.background_tasks.is_some())
.field("has_request_candidates", &self.request_candidates.is_some())
.field(
"has_gemini_file_mappings",
@@ -81,6 +84,10 @@ impl DataWriteRepositories {
.map(PostgresBackend::auth_module_write_repository)
.or_else(|| mysql.map(MysqlBackend::auth_module_write_repository))
.or_else(|| sqlite.map(SqliteBackend::auth_module_write_repository)),
background_tasks: postgres
.map(PostgresBackend::background_task_write_repository)
.or_else(|| mysql.map(MysqlBackend::background_task_write_repository))
.or_else(|| sqlite.map(SqliteBackend::background_task_write_repository)),
request_candidates: postgres
.map(PostgresBackend::request_candidate_write_repository)
.or_else(|| mysql.map(MysqlBackend::request_candidate_write_repository))
@@ -149,6 +156,10 @@ impl DataWriteRepositories {
self.auth_modules.clone()
}
pub fn background_tasks(&self) -> Option<Arc<dyn BackgroundTaskWriteRepository>> {
self.background_tasks.clone()
}
pub fn usage(&self) -> Option<Arc<dyn UsageWriteRepository>> {
self.usage.clone()
}
@@ -201,6 +212,7 @@ impl DataWriteRepositories {
self.announcements.is_some()
|| self.auth_api_keys.is_some()
|| self.auth_modules.is_some()
|| self.background_tasks.is_some()
|| self.request_candidates.is_some()
|| self.gemini_file_mappings.is_some()
|| self.global_models.is_some()

View File

@@ -1,3 +1,29 @@
mod pool;
pub use pool::{SqlitePool, SqlitePoolConfig, SqlitePoolFactory};
use crate::DataLayerError;
use sqlx::{sqlite::SqliteRow, Row};
pub(crate) fn sqlite_real(row: &SqliteRow, field: &str) -> Result<f64, DataLayerError> {
match row.try_get::<f64, _>(field) {
Ok(value) => Ok(value),
Err(real_err) => match row.try_get::<i64, _>(field) {
Ok(value) => Ok(value as f64),
Err(_) => Err(DataLayerError::sql(real_err)),
},
}
}
pub(crate) fn sqlite_optional_real(
row: &SqliteRow,
field: &str,
) -> Result<Option<f64>, DataLayerError> {
match row.try_get::<Option<f64>, _>(field) {
Ok(value) => Ok(value),
Err(real_err) => match row.try_get::<Option<i64>, _>(field) {
Ok(value) => Ok(value.map(|value| value as f64)),
Err(_) => Err(DataLayerError::sql(real_err)),
},
}
}

View File

@@ -7,7 +7,7 @@ use tracing::info;
// Generated by build.rs from schema/bootstrap/postgres.
pub(crate) static EMPTY_DATABASE_SNAPSHOT_SQL: &str =
include_str!(concat!(env!("OUT_DIR"), "/empty_database_snapshot.sql"));
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260507120000;
pub(crate) const EMPTY_DATABASE_SNAPSHOT_CUTOFF_VERSION: i64 = 20260509000000;
const PUBLIC_BASE_TABLE_COUNT_SQL: &str = r#"
SELECT COUNT(*)::BIGINT

View File

@@ -293,6 +293,8 @@ fn empty_database_snapshot_covers_current_cutoff_versions() {
20260505130000,
20260507000000,
20260507120000,
20260508000000,
20260509000000,
]
);
}
@@ -510,8 +512,24 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
.map(|migration| migration.version)
.collect::<Vec<_>>();
assert_eq!(mysql_versions, vec![20260403000000, 20260507120000]);
assert_eq!(sqlite_versions, vec![20260403000000, 20260507120000]);
assert_eq!(
mysql_versions,
vec![
20260403000000,
20260507120000,
20260508000000,
20260509000000
]
);
assert_eq!(
sqlite_versions,
vec![
20260403000000,
20260507120000,
20260508000000,
20260509000000
]
);
}
#[test]
@@ -1014,6 +1032,8 @@ fn pending_migrations_from_applied_skips_versions_already_applied() {
20260505130000,
20260507000000,
20260507120000,
20260508000000,
20260509000000,
]
);
}

View File

@@ -7,7 +7,7 @@ use super::types::{
StandaloneApiKeyExportListQuery, StoredAuthApiKeyExportRecord, StoredAuthApiKeySnapshot,
UpdateStandaloneApiKeyBasicRecord, UpdateUserApiKeyBasicRecord,
};
use crate::driver::sqlite::SqlitePool;
use crate::driver::sqlite::{sqlite_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::DataLayerError;
@@ -57,7 +57,7 @@ SELECT
api_keys.auto_delete_on_expiry,
api_keys.total_requests,
COALESCE(api_keys.total_tokens, 0) AS total_tokens,
COALESCE(api_keys.total_cost_usd, 0) AS total_cost_usd,
CAST(COALESCE(api_keys.total_cost_usd, 0) AS REAL) AS total_cost_usd,
api_keys.last_used_at AS last_used_at_unix_secs,
api_keys.created_at AS created_at_unix_secs,
api_keys.updated_at AS updated_at_unix_secs,
@@ -904,7 +904,7 @@ fn map_auth_api_key_export_row(
row.try_get("auto_delete_on_expiry").map_sql_err()?,
row.try_get("total_requests").map_sql_err()?,
row.try_get("total_tokens").map_sql_err()?,
row.try_get("total_cost_usd").map_sql_err()?,
sqlite_real(row, "total_cost_usd")?,
row.try_get("is_standalone").map_sql_err()?,
)
.and_then(|record| {

View File

@@ -0,0 +1,218 @@
use std::collections::{BTreeMap, BTreeSet};
use std::sync::RwLock;
use async_trait::async_trait;
use super::{
BackgroundTaskListQuery, BackgroundTaskReadRepository, BackgroundTaskStatus,
BackgroundTaskSummary, BackgroundTaskWriteRepository, StoredBackgroundTaskEvent,
StoredBackgroundTaskRun, StoredBackgroundTaskRunPage, UpsertBackgroundTaskEvent,
UpsertBackgroundTaskRun,
};
use crate::DataLayerError;
#[derive(Debug, Default)]
struct InMemoryBackgroundTaskIndex {
runs: BTreeMap<String, StoredBackgroundTaskRun>,
events_by_run: BTreeMap<String, Vec<StoredBackgroundTaskEvent>>,
}
#[derive(Debug, Default)]
pub struct InMemoryBackgroundTaskRepository {
index: RwLock<InMemoryBackgroundTaskIndex>,
}
impl InMemoryBackgroundTaskRepository {
fn matches_filter(run: &StoredBackgroundTaskRun, query: &BackgroundTaskListQuery) -> bool {
if let Some(kind) = query.kind {
if run.kind != kind {
return false;
}
}
if let Some(status) = query.status {
if run.status != status {
return false;
}
}
if let Some(trigger) = query.trigger.as_deref() {
if run.trigger != trigger {
return false;
}
}
if let Some(task_key_substring) = query.task_key_substring.as_deref() {
let needle = task_key_substring.to_ascii_lowercase();
if !run.task_key.to_ascii_lowercase().contains(&needle) {
return false;
}
}
true
}
pub fn seed_runs<I>(runs: I) -> Self
where
I: IntoIterator<Item = StoredBackgroundTaskRun>,
{
let mut index = InMemoryBackgroundTaskIndex::default();
for run in runs {
index.runs.insert(run.id.clone(), run);
}
Self {
index: RwLock::new(index),
}
}
}
#[async_trait]
impl BackgroundTaskReadRepository for InMemoryBackgroundTaskRepository {
async fn find_run(
&self,
run_id: &str,
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
Ok(self
.index
.read()
.expect("background task repository lock")
.runs
.get(run_id)
.cloned())
}
async fn list_runs(
&self,
query: &BackgroundTaskListQuery,
) -> Result<StoredBackgroundTaskRunPage, DataLayerError> {
let mut items = self
.index
.read()
.expect("background task repository lock")
.runs
.values()
.filter(|run| Self::matches_filter(run, query))
.cloned()
.collect::<Vec<_>>();
items.sort_by(|left, right| {
right
.created_at_unix_secs
.cmp(&left.created_at_unix_secs)
.then_with(|| right.updated_at_unix_secs.cmp(&left.updated_at_unix_secs))
});
let total = items.len();
let limit = query.limit.max(1);
let items = items
.into_iter()
.skip(query.offset)
.take(limit)
.collect::<Vec<_>>();
Ok(StoredBackgroundTaskRunPage { items, total })
}
async fn list_events(
&self,
run_id: &str,
offset: usize,
limit: usize,
) -> Result<Vec<StoredBackgroundTaskEvent>, DataLayerError> {
let Some(events) = self
.index
.read()
.expect("background task repository lock")
.events_by_run
.get(run_id)
.cloned()
else {
return Ok(Vec::new());
};
let limit = limit.max(1);
Ok(events.into_iter().skip(offset).take(limit).collect())
}
async fn summarize_runs(&self) -> Result<BackgroundTaskSummary, DataLayerError> {
let runs = self
.index
.read()
.expect("background task repository lock")
.runs
.values()
.cloned()
.collect::<Vec<_>>();
let mut by_status = BTreeMap::new();
let mut by_kind = BTreeMap::new();
let mut running_count = 0_u64;
for run in runs {
*by_status
.entry(run.status.as_database().to_string())
.or_insert(0) += 1;
*by_kind
.entry(run.kind.as_database().to_string())
.or_insert(0) += 1;
if run.status == BackgroundTaskStatus::Running {
running_count += 1;
}
}
let total = by_status.values().copied().sum();
Ok(BackgroundTaskSummary {
total,
running_count,
by_status,
by_kind,
})
}
}
#[async_trait]
impl BackgroundTaskWriteRepository for InMemoryBackgroundTaskRepository {
async fn upsert_run(
&self,
run: UpsertBackgroundTaskRun,
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
run.validate()?;
let stored = run.into_stored();
self.index
.write()
.expect("background task repository lock")
.runs
.insert(stored.id.clone(), stored.clone());
Ok(stored)
}
async fn request_cancel(
&self,
run_id: &str,
updated_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let mut guard = self.index.write().expect("background task repository lock");
let Some(run) = guard.runs.get_mut(run_id) else {
return Ok(false);
};
run.cancel_requested = true;
run.updated_at_unix_secs = updated_at_unix_secs;
Ok(true)
}
async fn upsert_event(
&self,
event: UpsertBackgroundTaskEvent,
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
event.validate()?;
let stored = event.into_stored();
let mut guard = self.index.write().expect("background task repository lock");
let entries = guard
.events_by_run
.entry(stored.run_id.clone())
.or_default();
if let Some(position) = entries.iter().position(|value| value.id == stored.id) {
entries[position] = stored.clone();
} else {
entries.push(stored.clone());
}
let mut seen = BTreeSet::new();
entries.retain(|entry| seen.insert(entry.id.clone()));
entries.sort_by(|left, right| {
left.created_at_unix_secs
.cmp(&right.created_at_unix_secs)
.then_with(|| left.id.cmp(&right.id))
});
Ok(stored)
}
}

View File

@@ -0,0 +1,17 @@
mod memory;
mod mysql;
mod postgres;
mod sqlite;
#[allow(unused_imports)]
pub(crate) use aether_data_contracts::repository::background_tasks::{
BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskReadRepository,
BackgroundTaskRepository, BackgroundTaskStatus, BackgroundTaskSummary,
BackgroundTaskWriteRepository, StoredBackgroundTaskEvent, StoredBackgroundTaskRun,
StoredBackgroundTaskRunPage, UpsertBackgroundTaskEvent, UpsertBackgroundTaskRun,
};
pub use memory::InMemoryBackgroundTaskRepository;
pub use mysql::MysqlBackgroundTaskRepository;
pub use postgres::SqlxBackgroundTaskRepository;
pub use sqlite::SqliteBackgroundTaskRepository;

View File

@@ -0,0 +1,449 @@
use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::{
BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskReadRepository,
BackgroundTaskStatus, BackgroundTaskSummary, BackgroundTaskWriteRepository,
StoredBackgroundTaskEvent, StoredBackgroundTaskRun, StoredBackgroundTaskRunPage,
UpsertBackgroundTaskEvent, UpsertBackgroundTaskRun,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const RUN_COLUMNS: &str = r#"
SELECT
id,
task_key,
kind,
`trigger`,
status,
attempt,
max_attempts,
owner_instance,
progress_percent,
progress_message,
payload_json,
result_json,
error_message,
cancel_requested,
created_by,
created_at_unix_secs,
started_at_unix_secs,
finished_at_unix_secs,
updated_at_unix_secs
FROM background_task_runs
"#;
const EVENT_COLUMNS: &str = r#"
SELECT
id,
run_id,
event_type,
message,
payload_json,
created_at_unix_secs
FROM background_task_events
"#;
#[derive(Debug, Clone)]
pub struct MysqlBackgroundTaskRepository {
pool: MysqlPool,
}
impl MysqlBackgroundTaskRepository {
pub fn new(pool: MysqlPool) -> Self {
Self { pool }
}
fn apply_run_filter(builder: &mut QueryBuilder<'_, MySql>, query: &BackgroundTaskListQuery) {
let mut has_where = false;
if let Some(kind) = query.kind {
if !has_where {
builder.push(" WHERE ");
has_where = true;
} else {
builder.push(" AND ");
}
builder.push("kind = ").push_bind(kind.as_database());
}
if let Some(status) = query.status {
if !has_where {
builder.push(" WHERE ");
has_where = true;
} else {
builder.push(" AND ");
}
builder.push("status = ").push_bind(status.as_database());
}
if let Some(trigger) = query.trigger.as_deref() {
if !has_where {
builder.push(" WHERE ");
has_where = true;
} else {
builder.push(" AND ");
}
builder.push("`trigger` = ").push_bind(trigger.to_string());
}
if let Some(task_key_substring) = query.task_key_substring.as_deref() {
if !has_where {
builder.push(" WHERE ");
} else {
builder.push(" AND ");
}
builder.push("LOWER(task_key) LIKE ").push_bind(format!(
"%{}%",
task_key_substring.trim().to_ascii_lowercase()
));
}
}
}
#[async_trait]
impl BackgroundTaskReadRepository for MysqlBackgroundTaskRepository {
async fn find_run(
&self,
run_id: &str,
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
let row = sqlx::query(&format!("{RUN_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(run_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_run_row).transpose()
}
async fn list_runs(
&self,
query: &BackgroundTaskListQuery,
) -> Result<StoredBackgroundTaskRunPage, DataLayerError> {
let limit = query.limit.max(1);
let mut count_builder =
QueryBuilder::<MySql>::new("SELECT COUNT(id) AS total FROM background_task_runs");
Self::apply_run_filter(&mut count_builder, query);
let total = count_builder
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let mut builder = QueryBuilder::<MySql>::new(RUN_COLUMNS);
Self::apply_run_filter(&mut builder, query);
builder
.push(" ORDER BY created_at_unix_secs DESC, updated_at_unix_secs DESC")
.push(" LIMIT ")
.push_bind(i64_from_usize(limit, "run limit")?)
.push(" OFFSET ")
.push_bind(i64_from_usize(query.offset, "run offset")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let items = rows
.iter()
.map(map_run_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(StoredBackgroundTaskRunPage {
items,
total: usize::try_from(total).unwrap_or_default(),
})
}
async fn list_events(
&self,
run_id: &str,
offset: usize,
limit: usize,
) -> Result<Vec<StoredBackgroundTaskEvent>, DataLayerError> {
let limit = limit.max(1);
let rows = sqlx::query(&format!(
"{EVENT_COLUMNS} WHERE run_id = ? ORDER BY created_at_unix_secs ASC, id ASC LIMIT ? OFFSET ?"
))
.bind(run_id)
.bind(i64_from_usize(limit, "event limit")?)
.bind(i64_from_usize(offset, "event offset")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_event_row).collect()
}
async fn summarize_runs(&self) -> Result<BackgroundTaskSummary, DataLayerError> {
let total = sqlx::query_scalar::<_, i64>("SELECT COUNT(id) FROM background_task_runs")
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let running_count = sqlx::query_scalar::<_, i64>(
"SELECT COUNT(id) FROM background_task_runs WHERE status = 'running'",
)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let status_rows = sqlx::query(
"SELECT status, COUNT(id) AS total FROM background_task_runs GROUP BY status",
)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let kind_rows =
sqlx::query("SELECT kind, COUNT(id) AS total FROM background_task_runs GROUP BY kind")
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let mut by_status = std::collections::BTreeMap::new();
for row in status_rows {
let key: String = row.try_get("status").map_sql_err()?;
let count: i64 = row.try_get("total").map_sql_err()?;
by_status.insert(key, u64::try_from(count).unwrap_or_default());
}
let mut by_kind = std::collections::BTreeMap::new();
for row in kind_rows {
let key: String = row.try_get("kind").map_sql_err()?;
let count: i64 = row.try_get("total").map_sql_err()?;
by_kind.insert(key, u64::try_from(count).unwrap_or_default());
}
Ok(BackgroundTaskSummary {
total: u64::try_from(total).unwrap_or_default(),
running_count: u64::try_from(running_count).unwrap_or_default(),
by_status,
by_kind,
})
}
}
#[async_trait]
impl BackgroundTaskWriteRepository for MysqlBackgroundTaskRepository {
async fn upsert_run(
&self,
run: UpsertBackgroundTaskRun,
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
run.validate()?;
sqlx::query(
r#"
INSERT INTO background_task_runs (
id,
task_key,
kind,
`trigger`,
status,
attempt,
max_attempts,
owner_instance,
progress_percent,
progress_message,
payload_json,
result_json,
error_message,
cancel_requested,
created_by,
created_at_unix_secs,
started_at_unix_secs,
finished_at_unix_secs,
updated_at_unix_secs
) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)
ON DUPLICATE KEY UPDATE
task_key = VALUES(task_key),
kind = VALUES(kind),
`trigger` = VALUES(`trigger`),
status = VALUES(status),
attempt = VALUES(attempt),
max_attempts = VALUES(max_attempts),
owner_instance = VALUES(owner_instance),
progress_percent = VALUES(progress_percent),
progress_message = VALUES(progress_message),
payload_json = VALUES(payload_json),
result_json = VALUES(result_json),
error_message = VALUES(error_message),
cancel_requested = VALUES(cancel_requested),
created_by = VALUES(created_by),
created_at_unix_secs = VALUES(created_at_unix_secs),
started_at_unix_secs = VALUES(started_at_unix_secs),
finished_at_unix_secs = VALUES(finished_at_unix_secs),
updated_at_unix_secs = VALUES(updated_at_unix_secs)
"#,
)
.bind(&run.id)
.bind(&run.task_key)
.bind(run.kind.as_database())
.bind(&run.trigger)
.bind(run.status.as_database())
.bind(i64::from(run.attempt))
.bind(i64::from(run.max_attempts))
.bind(run.owner_instance.as_deref())
.bind(i32::from(run.progress_percent))
.bind(run.progress_message.as_deref())
.bind(json_to_string(&run.payload_json, "payload_json")?)
.bind(json_to_string(&run.result_json, "result_json")?)
.bind(run.error_message.as_deref())
.bind(run.cancel_requested)
.bind(run.created_by.as_deref())
.bind(u64_to_i64(
run.created_at_unix_secs,
"created_at_unix_secs",
)?)
.bind(run.started_at_unix_secs.map(|value| value as i64))
.bind(run.finished_at_unix_secs.map(|value| value as i64))
.bind(u64_to_i64(
run.updated_at_unix_secs,
"updated_at_unix_secs",
)?)
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_run(&run.id).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("background task run missing after upsert".to_string())
})
}
async fn request_cancel(
&self,
run_id: &str,
updated_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let affected = sqlx::query(
"UPDATE background_task_runs SET cancel_requested = TRUE, updated_at_unix_secs = ? WHERE id = ?",
)
.bind(u64_to_i64(updated_at_unix_secs, "updated_at_unix_secs")?)
.bind(run_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(affected > 0)
}
async fn upsert_event(
&self,
event: UpsertBackgroundTaskEvent,
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
event.validate()?;
sqlx::query(
r#"
INSERT INTO background_task_events (
id, run_id, event_type, message, payload_json, created_at_unix_secs
) VALUES (?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
run_id = VALUES(run_id),
event_type = VALUES(event_type),
message = VALUES(message),
payload_json = VALUES(payload_json),
created_at_unix_secs = VALUES(created_at_unix_secs)
"#,
)
.bind(&event.id)
.bind(&event.run_id)
.bind(&event.event_type)
.bind(&event.message)
.bind(json_to_string(&event.payload_json, "payload_json")?)
.bind(u64_to_i64(
event.created_at_unix_secs,
"created_at_unix_secs",
)?)
.execute(&self.pool)
.await
.map_sql_err()?;
let row = sqlx::query(&format!("{EVENT_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(&event.id)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
map_event_row(&row)
}
}
fn map_run_row(row: &MySqlRow) -> Result<StoredBackgroundTaskRun, DataLayerError> {
let kind: String = row.try_get("kind").map_sql_err()?;
let status: String = row.try_get("status").map_sql_err()?;
let attempt: i64 = row.try_get("attempt").map_sql_err()?;
let max_attempts: i64 = row.try_get("max_attempts").map_sql_err()?;
let progress_percent: i32 = row.try_get("progress_percent").map_sql_err()?;
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
let started_at_unix_secs: Option<i64> = row.try_get("started_at_unix_secs").map_sql_err()?;
let finished_at_unix_secs: Option<i64> = row.try_get("finished_at_unix_secs").map_sql_err()?;
let updated_at_unix_secs: i64 = row.try_get("updated_at_unix_secs").map_sql_err()?;
Ok(StoredBackgroundTaskRun {
id: row.try_get("id").map_sql_err()?,
task_key: row.try_get("task_key").map_sql_err()?,
kind: BackgroundTaskKind::from_database(&kind)?,
trigger: row.try_get("trigger").map_sql_err()?,
status: BackgroundTaskStatus::from_database(&status)?,
attempt: u32::try_from(attempt).unwrap_or_default(),
max_attempts: u32::try_from(max_attempts).unwrap_or_default(),
owner_instance: row.try_get("owner_instance").map_sql_err()?,
progress_percent: u16::try_from(progress_percent).unwrap_or_default(),
progress_message: row.try_get("progress_message").map_sql_err()?,
payload_json: parse_optional_json(
row.try_get("payload_json").ok().flatten(),
"payload_json",
)?,
result_json: parse_optional_json(row.try_get("result_json").ok().flatten(), "result_json")?,
error_message: row.try_get("error_message").map_sql_err()?,
cancel_requested: row.try_get("cancel_requested").map_sql_err()?,
created_by: row.try_get("created_by").map_sql_err()?,
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
started_at_unix_secs: started_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
finished_at_unix_secs: finished_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
updated_at_unix_secs: u64::try_from(updated_at_unix_secs).unwrap_or_default(),
})
}
fn map_event_row(row: &MySqlRow) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
Ok(StoredBackgroundTaskEvent {
id: row.try_get("id").map_sql_err()?,
run_id: row.try_get("run_id").map_sql_err()?,
event_type: row.try_get("event_type").map_sql_err()?,
message: row.try_get("message").map_sql_err()?,
payload_json: parse_optional_json(
row.try_get("payload_json").ok().flatten(),
"payload_json",
)?,
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
})
}
fn i64_from_usize(value: usize, label: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
})
}
fn u64_to_i64(value: u64, label: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
})
}
fn json_to_string(
value: &Option<serde_json::Value>,
field_name: &str,
) -> Result<Option<String>, DataLayerError> {
value
.as_ref()
.map(|value| {
serde_json::to_string(value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"background task {field_name} is unserializable: {err}"
))
})
})
.transpose()
}
fn parse_optional_json(
value: Option<String>,
field_name: &str,
) -> Result<Option<serde_json::Value>, DataLayerError> {
value
.filter(|value| !value.trim().is_empty())
.map(|value| {
serde_json::from_str(&value).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"background task {field_name} contains invalid JSON: {err}"
))
})
})
.transpose()
}

View File

@@ -0,0 +1,426 @@
use async_trait::async_trait;
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
use super::{
BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskReadRepository,
BackgroundTaskStatus, BackgroundTaskSummary, BackgroundTaskWriteRepository,
StoredBackgroundTaskEvent, StoredBackgroundTaskRun, StoredBackgroundTaskRunPage,
UpsertBackgroundTaskEvent, UpsertBackgroundTaskRun,
};
use crate::error::SqlxResultExt;
use crate::DataLayerError;
const RUN_COLUMNS: &str = r#"
SELECT
id,
task_key,
kind,
"trigger",
status,
attempt,
max_attempts,
owner_instance,
progress_percent,
progress_message,
payload_json,
result_json,
error_message,
cancel_requested,
created_by,
created_at_unix_secs,
started_at_unix_secs,
finished_at_unix_secs,
updated_at_unix_secs
FROM background_task_runs
"#;
const EVENT_COLUMNS: &str = r#"
SELECT
id,
run_id,
event_type,
message,
payload_json,
created_at_unix_secs
FROM background_task_events
"#;
#[derive(Debug, Clone)]
pub struct SqlxBackgroundTaskRepository {
pool: PgPool,
}
impl SqlxBackgroundTaskRepository {
pub fn new(pool: PgPool) -> Self {
Self { pool }
}
fn apply_run_filter(
builder: &mut QueryBuilder<'_, Postgres>,
query: &BackgroundTaskListQuery,
include_where: bool,
) {
let mut has_where = include_where;
let mut push_where = |builder: &mut QueryBuilder<'_, Postgres>| {
if has_where {
builder.push(" AND ");
} else {
builder.push(" WHERE ");
has_where = true;
}
};
if let Some(kind) = query.kind {
push_where(builder);
builder.push("kind = ").push_bind(kind.as_database());
}
if let Some(status) = query.status {
push_where(builder);
builder.push("status = ").push_bind(status.as_database());
}
if let Some(trigger) = query.trigger.as_deref() {
push_where(builder);
builder
.push("\"trigger\" = ")
.push_bind(trigger.to_string());
}
if let Some(task_key_substring) = query.task_key_substring.as_deref() {
push_where(builder);
builder
.push("task_key ILIKE ")
.push_bind(format!("%{}%", task_key_substring.trim()));
}
}
}
#[async_trait]
impl BackgroundTaskReadRepository for SqlxBackgroundTaskRepository {
async fn find_run(
&self,
run_id: &str,
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
let row = sqlx::query(&format!("{RUN_COLUMNS} WHERE id = $1 LIMIT 1"))
.bind(run_id)
.fetch_optional(&self.pool)
.await
.map_postgres_err()?;
row.as_ref().map(map_run_row).transpose()
}
async fn list_runs(
&self,
query: &BackgroundTaskListQuery,
) -> Result<StoredBackgroundTaskRunPage, DataLayerError> {
let limit = query.limit.max(1);
let mut count_builder =
QueryBuilder::<Postgres>::new("SELECT COUNT(id) AS total FROM background_task_runs");
Self::apply_run_filter(&mut count_builder, query, false);
let total = count_builder
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_postgres_err()?;
let mut builder = QueryBuilder::<Postgres>::new(RUN_COLUMNS);
Self::apply_run_filter(&mut builder, query, false);
builder
.push(" ORDER BY created_at_unix_secs DESC, updated_at_unix_secs DESC")
.push(" LIMIT ")
.push_bind(i64_from_usize(limit, "background task run limit")?)
.push(" OFFSET ")
.push_bind(i64_from_usize(query.offset, "background task run offset")?);
let rows = builder
.build()
.fetch_all(&self.pool)
.await
.map_postgres_err()?;
let items = rows
.iter()
.map(map_run_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(StoredBackgroundTaskRunPage {
items,
total: usize::try_from(total).unwrap_or_default(),
})
}
async fn list_events(
&self,
run_id: &str,
offset: usize,
limit: usize,
) -> Result<Vec<StoredBackgroundTaskEvent>, DataLayerError> {
let limit = limit.max(1);
let rows = sqlx::query(&format!(
"{EVENT_COLUMNS} WHERE run_id = $1 ORDER BY created_at_unix_secs ASC, id ASC LIMIT $2 OFFSET $3"
))
.bind(run_id)
.bind(i64_from_usize(limit, "background task event limit")?)
.bind(i64_from_usize(offset, "background task event offset")?)
.fetch_all(&self.pool)
.await
.map_postgres_err()?;
rows.iter().map(map_event_row).collect()
}
async fn summarize_runs(&self) -> Result<BackgroundTaskSummary, DataLayerError> {
let total = sqlx::query_scalar::<_, i64>("SELECT COUNT(id) FROM background_task_runs")
.fetch_one(&self.pool)
.await
.map_postgres_err()?;
let running_count = sqlx::query_scalar::<_, i64>(
"SELECT COUNT(id) FROM background_task_runs WHERE status = 'running'",
)
.fetch_one(&self.pool)
.await
.map_postgres_err()?;
let status_rows = sqlx::query(
"SELECT status, COUNT(id) AS total FROM background_task_runs GROUP BY status",
)
.fetch_all(&self.pool)
.await
.map_postgres_err()?;
let kind_rows =
sqlx::query("SELECT kind, COUNT(id) AS total FROM background_task_runs GROUP BY kind")
.fetch_all(&self.pool)
.await
.map_postgres_err()?;
let mut by_status = std::collections::BTreeMap::new();
for row in status_rows {
let key: String = row.try_get("status").map_postgres_err()?;
let count: i64 = row.try_get("total").map_postgres_err()?;
by_status.insert(key, u64::try_from(count).unwrap_or_default());
}
let mut by_kind = std::collections::BTreeMap::new();
for row in kind_rows {
let key: String = row.try_get("kind").map_postgres_err()?;
let count: i64 = row.try_get("total").map_postgres_err()?;
by_kind.insert(key, u64::try_from(count).unwrap_or_default());
}
Ok(BackgroundTaskSummary {
total: u64::try_from(total).unwrap_or_default(),
running_count: u64::try_from(running_count).unwrap_or_default(),
by_status,
by_kind,
})
}
}
#[async_trait]
impl BackgroundTaskWriteRepository for SqlxBackgroundTaskRepository {
async fn upsert_run(
&self,
run: UpsertBackgroundTaskRun,
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
run.validate()?;
sqlx::query(
r#"
INSERT INTO background_task_runs (
id,
task_key,
kind,
"trigger",
status,
attempt,
max_attempts,
owner_instance,
progress_percent,
progress_message,
payload_json,
result_json,
error_message,
cancel_requested,
created_by,
created_at_unix_secs,
started_at_unix_secs,
finished_at_unix_secs,
updated_at_unix_secs
) VALUES (
$1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19
)
ON CONFLICT(id) DO UPDATE SET
task_key = EXCLUDED.task_key,
kind = EXCLUDED.kind,
"trigger" = EXCLUDED."trigger",
status = EXCLUDED.status,
attempt = EXCLUDED.attempt,
max_attempts = EXCLUDED.max_attempts,
owner_instance = EXCLUDED.owner_instance,
progress_percent = EXCLUDED.progress_percent,
progress_message = EXCLUDED.progress_message,
payload_json = EXCLUDED.payload_json,
result_json = EXCLUDED.result_json,
error_message = EXCLUDED.error_message,
cancel_requested = EXCLUDED.cancel_requested,
created_by = EXCLUDED.created_by,
created_at_unix_secs = EXCLUDED.created_at_unix_secs,
started_at_unix_secs = EXCLUDED.started_at_unix_secs,
finished_at_unix_secs = EXCLUDED.finished_at_unix_secs,
updated_at_unix_secs = EXCLUDED.updated_at_unix_secs
"#,
)
.bind(&run.id)
.bind(&run.task_key)
.bind(run.kind.as_database())
.bind(&run.trigger)
.bind(run.status.as_database())
.bind(u32_to_i32(run.attempt, "attempt")?)
.bind(u32_to_i32(run.max_attempts, "max_attempts")?)
.bind(run.owner_instance.as_deref())
.bind(i32::from(run.progress_percent))
.bind(run.progress_message.as_deref())
.bind(run.payload_json.clone())
.bind(run.result_json.clone())
.bind(run.error_message.as_deref())
.bind(run.cancel_requested)
.bind(run.created_by.as_deref())
.bind(u64_to_i64(
run.created_at_unix_secs,
"created_at_unix_secs",
)?)
.bind(run.started_at_unix_secs.map(|value| value as i64))
.bind(run.finished_at_unix_secs.map(|value| value as i64))
.bind(u64_to_i64(
run.updated_at_unix_secs,
"updated_at_unix_secs",
)?)
.execute(&self.pool)
.await
.map_postgres_err()?;
self.find_run(&run.id).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("background task run missing after upsert".to_string())
})
}
async fn request_cancel(
&self,
run_id: &str,
updated_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let affected = sqlx::query(
"UPDATE background_task_runs SET cancel_requested = TRUE, updated_at_unix_secs = $2 WHERE id = $1",
)
.bind(run_id)
.bind(u64_to_i64(updated_at_unix_secs, "updated_at_unix_secs")?)
.execute(&self.pool)
.await
.map_postgres_err()?
.rows_affected();
Ok(affected > 0)
}
async fn upsert_event(
&self,
event: UpsertBackgroundTaskEvent,
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
event.validate()?;
sqlx::query(
r#"
INSERT INTO background_task_events (
id,
run_id,
event_type,
message,
payload_json,
created_at_unix_secs
) VALUES ($1,$2,$3,$4,$5,$6)
ON CONFLICT(id) DO UPDATE SET
run_id = EXCLUDED.run_id,
event_type = EXCLUDED.event_type,
message = EXCLUDED.message,
payload_json = EXCLUDED.payload_json,
created_at_unix_secs = EXCLUDED.created_at_unix_secs
"#,
)
.bind(&event.id)
.bind(&event.run_id)
.bind(&event.event_type)
.bind(&event.message)
.bind(event.payload_json.clone())
.bind(u64_to_i64(
event.created_at_unix_secs,
"created_at_unix_secs",
)?)
.execute(&self.pool)
.await
.map_postgres_err()?;
let row = sqlx::query(&format!("{EVENT_COLUMNS} WHERE id = $1 LIMIT 1"))
.bind(&event.id)
.fetch_one(&self.pool)
.await
.map_postgres_err()?;
map_event_row(&row)
}
}
fn map_run_row(row: &PgRow) -> Result<StoredBackgroundTaskRun, DataLayerError> {
let kind: String = row.try_get("kind").map_postgres_err()?;
let status: String = row.try_get("status").map_postgres_err()?;
let attempt: i32 = row.try_get("attempt").map_postgres_err()?;
let max_attempts: i32 = row.try_get("max_attempts").map_postgres_err()?;
let progress_percent: i32 = row.try_get("progress_percent").map_postgres_err()?;
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_postgres_err()?;
let started_at_unix_secs: Option<i64> =
row.try_get("started_at_unix_secs").map_postgres_err()?;
let finished_at_unix_secs: Option<i64> =
row.try_get("finished_at_unix_secs").map_postgres_err()?;
let updated_at_unix_secs: i64 = row.try_get("updated_at_unix_secs").map_postgres_err()?;
Ok(StoredBackgroundTaskRun {
id: row.try_get("id").map_postgres_err()?,
task_key: row.try_get("task_key").map_postgres_err()?,
kind: BackgroundTaskKind::from_database(&kind)?,
trigger: row.try_get("trigger").map_postgres_err()?,
status: BackgroundTaskStatus::from_database(&status)?,
attempt: u32::try_from(attempt).unwrap_or_default(),
max_attempts: u32::try_from(max_attempts).unwrap_or_default(),
owner_instance: row.try_get("owner_instance").map_postgres_err()?,
progress_percent: u16::try_from(progress_percent).unwrap_or_default(),
progress_message: row.try_get("progress_message").map_postgres_err()?,
payload_json: row.try_get("payload_json").map_postgres_err()?,
result_json: row.try_get("result_json").map_postgres_err()?,
error_message: row.try_get("error_message").map_postgres_err()?,
cancel_requested: row.try_get("cancel_requested").map_postgres_err()?,
created_by: row.try_get("created_by").map_postgres_err()?,
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
started_at_unix_secs: started_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
finished_at_unix_secs: finished_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
updated_at_unix_secs: u64::try_from(updated_at_unix_secs).unwrap_or_default(),
})
}
fn map_event_row(row: &PgRow) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_postgres_err()?;
Ok(StoredBackgroundTaskEvent {
id: row.try_get("id").map_postgres_err()?,
run_id: row.try_get("run_id").map_postgres_err()?,
event_type: row.try_get("event_type").map_postgres_err()?,
message: row.try_get("message").map_postgres_err()?,
payload_json: row.try_get("payload_json").map_postgres_err()?,
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
})
}
fn i64_from_usize(value: usize, label: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
})
}
fn u64_to_i64(value: u64, label: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
})
}
fn u32_to_i32(value: u32, label: &str) -> Result<i32, DataLayerError> {
i32::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
})
}

View File

@@ -0,0 +1,430 @@
use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::{
BackgroundTaskKind, BackgroundTaskListQuery, BackgroundTaskReadRepository,
BackgroundTaskStatus, BackgroundTaskSummary, BackgroundTaskWriteRepository,
StoredBackgroundTaskEvent, StoredBackgroundTaskRun, StoredBackgroundTaskRunPage,
UpsertBackgroundTaskEvent, UpsertBackgroundTaskRun,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
use crate::DataLayerError;
const RUN_COLUMNS: &str = r#"
SELECT
id,
task_key,
kind,
"trigger",
status,
attempt,
max_attempts,
owner_instance,
progress_percent,
progress_message,
payload_json,
result_json,
error_message,
cancel_requested,
created_by,
created_at_unix_secs,
started_at_unix_secs,
finished_at_unix_secs,
updated_at_unix_secs
FROM background_task_runs
"#;
const EVENT_COLUMNS: &str = r#"
SELECT
id,
run_id,
event_type,
message,
payload_json,
created_at_unix_secs
FROM background_task_events
"#;
#[derive(Debug, Clone)]
pub struct SqliteBackgroundTaskRepository {
pool: SqlitePool,
}
impl SqliteBackgroundTaskRepository {
pub fn new(pool: SqlitePool) -> Self {
Self { pool }
}
fn apply_run_filter(builder: &mut QueryBuilder<'_, Sqlite>, query: &BackgroundTaskListQuery) {
let mut has_where = false;
if let Some(kind) = query.kind {
if !has_where {
builder.push(" WHERE ");
has_where = true;
} else {
builder.push(" AND ");
}
builder.push("kind = ").push_bind(kind.as_database());
}
if let Some(status) = query.status {
if !has_where {
builder.push(" WHERE ");
has_where = true;
} else {
builder.push(" AND ");
}
builder.push("status = ").push_bind(status.as_database());
}
if let Some(trigger) = query.trigger.as_deref() {
if !has_where {
builder.push(" WHERE ");
has_where = true;
} else {
builder.push(" AND ");
}
builder
.push("\"trigger\" = ")
.push_bind(trigger.to_string());
}
if let Some(task_key_substring) = query.task_key_substring.as_deref() {
if !has_where {
builder.push(" WHERE ");
} else {
builder.push(" AND ");
}
builder.push("LOWER(task_key) LIKE ").push_bind(format!(
"%{}%",
task_key_substring.trim().to_ascii_lowercase()
));
}
}
}
#[async_trait]
impl BackgroundTaskReadRepository for SqliteBackgroundTaskRepository {
async fn find_run(
&self,
run_id: &str,
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
let row = sqlx::query(&format!("{RUN_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(run_id)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_run_row).transpose()
}
async fn list_runs(
&self,
query: &BackgroundTaskListQuery,
) -> Result<StoredBackgroundTaskRunPage, DataLayerError> {
let limit = query.limit.max(1);
let mut count_builder =
QueryBuilder::<Sqlite>::new("SELECT COUNT(id) AS total FROM background_task_runs");
Self::apply_run_filter(&mut count_builder, query);
let total = count_builder
.build_query_scalar::<i64>()
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let mut builder = QueryBuilder::<Sqlite>::new(RUN_COLUMNS);
Self::apply_run_filter(&mut builder, query);
builder
.push(" ORDER BY created_at_unix_secs DESC, updated_at_unix_secs DESC")
.push(" LIMIT ")
.push_bind(i64_from_usize(limit, "run limit")?)
.push(" OFFSET ")
.push_bind(i64_from_usize(query.offset, "run offset")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let items = rows
.iter()
.map(map_run_row)
.collect::<Result<Vec<_>, _>>()?;
Ok(StoredBackgroundTaskRunPage {
items,
total: usize::try_from(total).unwrap_or_default(),
})
}
async fn list_events(
&self,
run_id: &str,
offset: usize,
limit: usize,
) -> Result<Vec<StoredBackgroundTaskEvent>, DataLayerError> {
let limit = limit.max(1);
let rows = sqlx::query(&format!(
"{EVENT_COLUMNS} WHERE run_id = ? ORDER BY created_at_unix_secs ASC, id ASC LIMIT ? OFFSET ?"
))
.bind(run_id)
.bind(i64_from_usize(limit, "event limit")?)
.bind(i64_from_usize(offset, "event offset")?)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_event_row).collect()
}
async fn summarize_runs(&self) -> Result<BackgroundTaskSummary, DataLayerError> {
let total = sqlx::query_scalar::<_, i64>("SELECT COUNT(id) FROM background_task_runs")
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let running_count = sqlx::query_scalar::<_, i64>(
"SELECT COUNT(id) FROM background_task_runs WHERE status = 'running'",
)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
let status_rows = sqlx::query(
"SELECT status, COUNT(id) AS total FROM background_task_runs GROUP BY status",
)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let kind_rows =
sqlx::query("SELECT kind, COUNT(id) AS total FROM background_task_runs GROUP BY kind")
.fetch_all(&self.pool)
.await
.map_sql_err()?;
let mut by_status = std::collections::BTreeMap::new();
for row in status_rows {
let key: String = row.try_get("status").map_sql_err()?;
let count: i64 = row.try_get("total").map_sql_err()?;
by_status.insert(key, u64::try_from(count).unwrap_or_default());
}
let mut by_kind = std::collections::BTreeMap::new();
for row in kind_rows {
let key: String = row.try_get("kind").map_sql_err()?;
let count: i64 = row.try_get("total").map_sql_err()?;
by_kind.insert(key, u64::try_from(count).unwrap_or_default());
}
Ok(BackgroundTaskSummary {
total: u64::try_from(total).unwrap_or_default(),
running_count: u64::try_from(running_count).unwrap_or_default(),
by_status,
by_kind,
})
}
}
#[async_trait]
impl BackgroundTaskWriteRepository for SqliteBackgroundTaskRepository {
async fn upsert_run(
&self,
run: UpsertBackgroundTaskRun,
) -> Result<StoredBackgroundTaskRun, DataLayerError> {
run.validate()?;
sqlx::query(
r#"
INSERT INTO background_task_runs (
id,
task_key,
kind,
"trigger",
status,
attempt,
max_attempts,
owner_instance,
progress_percent,
progress_message,
payload_json,
result_json,
error_message,
cancel_requested,
created_by,
created_at_unix_secs,
started_at_unix_secs,
finished_at_unix_secs,
updated_at_unix_secs
) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)
ON CONFLICT(id) DO UPDATE SET
task_key = excluded.task_key,
kind = excluded.kind,
"trigger" = excluded."trigger",
status = excluded.status,
attempt = excluded.attempt,
max_attempts = excluded.max_attempts,
owner_instance = excluded.owner_instance,
progress_percent = excluded.progress_percent,
progress_message = excluded.progress_message,
payload_json = excluded.payload_json,
result_json = excluded.result_json,
error_message = excluded.error_message,
cancel_requested = excluded.cancel_requested,
created_by = excluded.created_by,
created_at_unix_secs = excluded.created_at_unix_secs,
started_at_unix_secs = excluded.started_at_unix_secs,
finished_at_unix_secs = excluded.finished_at_unix_secs,
updated_at_unix_secs = excluded.updated_at_unix_secs
"#,
)
.bind(&run.id)
.bind(&run.task_key)
.bind(run.kind.as_database())
.bind(&run.trigger)
.bind(run.status.as_database())
.bind(i64::from(run.attempt))
.bind(i64::from(run.max_attempts))
.bind(run.owner_instance.as_deref())
.bind(i32::from(run.progress_percent))
.bind(run.progress_message.as_deref())
.bind(run.payload_json.as_ref().map(serde_json::Value::to_string))
.bind(run.result_json.as_ref().map(serde_json::Value::to_string))
.bind(run.error_message.as_deref())
.bind(run.cancel_requested)
.bind(run.created_by.as_deref())
.bind(u64_to_i64(
run.created_at_unix_secs,
"created_at_unix_secs",
)?)
.bind(run.started_at_unix_secs.map(|value| value as i64))
.bind(run.finished_at_unix_secs.map(|value| value as i64))
.bind(u64_to_i64(
run.updated_at_unix_secs,
"updated_at_unix_secs",
)?)
.execute(&self.pool)
.await
.map_sql_err()?;
self.find_run(&run.id).await?.ok_or_else(|| {
DataLayerError::UnexpectedValue("background task run missing after upsert".to_string())
})
}
async fn request_cancel(
&self,
run_id: &str,
updated_at_unix_secs: u64,
) -> Result<bool, DataLayerError> {
let affected = sqlx::query(
"UPDATE background_task_runs SET cancel_requested = 1, updated_at_unix_secs = ? WHERE id = ?",
)
.bind(u64_to_i64(updated_at_unix_secs, "updated_at_unix_secs")?)
.bind(run_id)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected();
Ok(affected > 0)
}
async fn upsert_event(
&self,
event: UpsertBackgroundTaskEvent,
) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
event.validate()?;
sqlx::query(
r#"
INSERT INTO background_task_events (
id, run_id, event_type, message, payload_json, created_at_unix_secs
) VALUES (?, ?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET
run_id = excluded.run_id,
event_type = excluded.event_type,
message = excluded.message,
payload_json = excluded.payload_json,
created_at_unix_secs = excluded.created_at_unix_secs
"#,
)
.bind(&event.id)
.bind(&event.run_id)
.bind(&event.event_type)
.bind(&event.message)
.bind(
event
.payload_json
.as_ref()
.map(serde_json::Value::to_string),
)
.bind(u64_to_i64(
event.created_at_unix_secs,
"created_at_unix_secs",
)?)
.execute(&self.pool)
.await
.map_sql_err()?;
let row = sqlx::query(&format!("{EVENT_COLUMNS} WHERE id = ? LIMIT 1"))
.bind(&event.id)
.fetch_one(&self.pool)
.await
.map_sql_err()?;
map_event_row(&row)
}
}
fn map_run_row(row: &SqliteRow) -> Result<StoredBackgroundTaskRun, DataLayerError> {
let kind: String = row.try_get("kind").map_sql_err()?;
let status: String = row.try_get("status").map_sql_err()?;
let attempt: i64 = row.try_get("attempt").map_sql_err()?;
let max_attempts: i64 = row.try_get("max_attempts").map_sql_err()?;
let progress_percent: i32 = row.try_get("progress_percent").map_sql_err()?;
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
let started_at_unix_secs: Option<i64> = row.try_get("started_at_unix_secs").map_sql_err()?;
let finished_at_unix_secs: Option<i64> = row.try_get("finished_at_unix_secs").map_sql_err()?;
let updated_at_unix_secs: i64 = row.try_get("updated_at_unix_secs").map_sql_err()?;
Ok(StoredBackgroundTaskRun {
id: row.try_get("id").map_sql_err()?,
task_key: row.try_get("task_key").map_sql_err()?,
kind: BackgroundTaskKind::from_database(&kind)?,
trigger: row.try_get("trigger").map_sql_err()?,
status: BackgroundTaskStatus::from_database(&status)?,
attempt: u32::try_from(attempt).unwrap_or_default(),
max_attempts: u32::try_from(max_attempts).unwrap_or_default(),
owner_instance: row.try_get("owner_instance").map_sql_err()?,
progress_percent: u16::try_from(progress_percent).unwrap_or_default(),
progress_message: row.try_get("progress_message").map_sql_err()?,
payload_json: parse_optional_json(row.try_get("payload_json").map_sql_err()?)?,
result_json: parse_optional_json(row.try_get("result_json").map_sql_err()?)?,
error_message: row.try_get("error_message").map_sql_err()?,
cancel_requested: row.try_get("cancel_requested").map_sql_err()?,
created_by: row.try_get("created_by").map_sql_err()?,
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
started_at_unix_secs: started_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
finished_at_unix_secs: finished_at_unix_secs.and_then(|value| u64::try_from(value).ok()),
updated_at_unix_secs: u64::try_from(updated_at_unix_secs).unwrap_or_default(),
})
}
fn map_event_row(row: &SqliteRow) -> Result<StoredBackgroundTaskEvent, DataLayerError> {
let created_at_unix_secs: i64 = row.try_get("created_at_unix_secs").map_sql_err()?;
Ok(StoredBackgroundTaskEvent {
id: row.try_get("id").map_sql_err()?,
run_id: row.try_get("run_id").map_sql_err()?,
event_type: row.try_get("event_type").map_sql_err()?,
message: row.try_get("message").map_sql_err()?,
payload_json: parse_optional_json(row.try_get("payload_json").map_sql_err()?)?,
created_at_unix_secs: u64::try_from(created_at_unix_secs).unwrap_or_default(),
})
}
fn parse_optional_json(value: Option<String>) -> Result<Option<serde_json::Value>, DataLayerError> {
value
.map(|raw| {
serde_json::from_str::<serde_json::Value>(&raw).map_err(|err| {
DataLayerError::UnexpectedValue(format!(
"invalid background task json payload: {err}"
))
})
})
.transpose()
}
fn i64_from_usize(value: usize, label: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
})
}
fn u64_to_i64(value: u64, label: &str) -> Result<i64, DataLayerError> {
i64::try_from(value).map_err(|_| {
DataLayerError::UnexpectedValue(format!("background task {label} overflow: {value}"))
})
}

View File

@@ -6,7 +6,7 @@ use super::{
AdminBillingPresetApplyResult, AdminBillingRuleRecord, AdminBillingRuleWriteInput,
BillingReadRepository, StoredBillingModelContext,
};
use crate::driver::sqlite::SqlitePool;
use crate::driver::sqlite::{sqlite_optional_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::DataLayerError;
@@ -20,12 +20,12 @@ SELECT
gm.id AS global_model_id,
gm.name AS global_model_name,
gm.config AS global_model_config,
gm.default_price_per_request AS default_price_per_request,
CAST(gm.default_price_per_request AS REAL) AS default_price_per_request,
gm.default_tiered_pricing AS default_tiered_pricing,
m.id AS model_id,
m.provider_model_name AS model_provider_model_name,
m.config AS model_config,
m.price_per_request AS model_price_per_request,
CAST(m.price_per_request AS REAL) AS model_price_per_request,
m.tiered_pricing AS model_tiered_pricing,
m.provider_model_mappings AS provider_model_mappings,
m.is_available AS model_is_available,
@@ -641,19 +641,13 @@ fn match_rank(
return Ok(None);
};
let has_model_price = row
.try_get::<Option<f64>, _>("model_price_per_request")
.map_sql_err()?
.is_some()
let has_model_price = sqlite_optional_real(row, "model_price_per_request")?.is_some()
|| row
.try_get::<Option<String>, _>("model_tiered_pricing")
.ok()
.flatten()
.is_some();
let has_default_price = row
.try_get::<Option<f64>, _>("default_price_per_request")
.map_sql_err()?
.is_some()
let has_default_price = sqlite_optional_real(row, "default_price_per_request")?.is_some()
|| row
.try_get::<Option<String>, _>("default_tiered_pricing")
.ok()
@@ -717,12 +711,12 @@ fn map_row(row: &SqliteRow) -> Result<StoredBillingModelContext, DataLayerError>
row.try_get("global_model_id").map_sql_err()?,
row.try_get("global_model_name").map_sql_err()?,
parse_json(row.try_get("global_model_config").ok().flatten())?,
row.try_get("default_price_per_request").map_sql_err()?,
sqlite_optional_real(row, "default_price_per_request")?,
parse_json(row.try_get("default_tiered_pricing").ok().flatten())?,
row.try_get("model_id").map_sql_err()?,
row.try_get("model_provider_model_name").map_sql_err()?,
parse_json(row.try_get("model_config").ok().flatten())?,
row.try_get("model_price_per_request").map_sql_err()?,
sqlite_optional_real(row, "model_price_per_request")?,
parse_json(row.try_get("model_tiered_pricing").ok().flatten())?,
)
}

View File

@@ -10,7 +10,7 @@ use super::{
StoredPublicCatalogModel, StoredPublicGlobalModel, StoredPublicGlobalModelPage,
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
};
use crate::driver::sqlite::SqlitePool;
use crate::driver::sqlite::{sqlite_optional_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::DataLayerError;
@@ -50,7 +50,7 @@ SELECT
name,
display_name,
is_active,
default_price_per_request,
CAST(default_price_per_request AS REAL) AS default_price_per_request,
default_tiered_pricing,
supported_capabilities,
config,
@@ -74,7 +74,7 @@ SELECT
name,
COALESCE(NULLIF(display_name, ''), name) AS display_name,
is_active,
default_price_per_request,
CAST(default_price_per_request AS REAL) AS default_price_per_request,
default_tiered_pricing,
supported_capabilities,
config,
@@ -101,7 +101,7 @@ SELECT
m.global_model_id,
m.provider_model_name,
m.provider_model_mappings,
m.price_per_request,
CAST(m.price_per_request AS REAL) AS price_per_request,
m.tiered_pricing,
m.supports_vision,
m.supports_function_calling,
@@ -115,7 +115,7 @@ SELECT
m.updated_at AS updated_at_unix_secs,
gm.name AS global_model_name,
gm.display_name AS global_model_display_name,
gm.default_price_per_request AS global_model_default_price_per_request,
CAST(gm.default_price_per_request AS REAL) AS global_model_default_price_per_request,
gm.default_tiered_pricing AS global_model_default_tiered_pricing,
gm.supported_capabilities AS global_model_supported_capabilities,
gm.config AS global_model_config
@@ -711,7 +711,7 @@ fn map_public_global_model_row(row: &SqliteRow) -> Result<StoredPublicGlobalMode
row.try_get("name").map_sql_err()?,
row.try_get("display_name").map_sql_err()?,
row.try_get("is_active").map_sql_err()?,
row.try_get("default_price_per_request").map_sql_err()?,
sqlite_optional_real(row, "default_price_per_request")?,
optional_json_from_string(
row.try_get("default_tiered_pricing").map_sql_err()?,
"global_models.default_tiered_pricing",
@@ -731,7 +731,7 @@ fn map_admin_global_model_row(row: &SqliteRow) -> Result<StoredAdminGlobalModel,
row.try_get("name").map_sql_err()?,
row.try_get("display_name").map_sql_err()?,
row.try_get("is_active").map_sql_err()?,
row.try_get("default_price_per_request").map_sql_err()?,
sqlite_optional_real(row, "default_price_per_request")?,
optional_json_from_string(
row.try_get("default_tiered_pricing").map_sql_err()?,
"global_models.default_tiered_pricing",
@@ -767,7 +767,7 @@ fn map_admin_provider_model_row(
row.try_get("provider_model_mappings").map_sql_err()?,
"models.provider_model_mappings",
)?,
row.try_get("price_per_request").map_sql_err()?,
sqlite_optional_real(row, "price_per_request")?,
optional_json_from_string(
row.try_get("tiered_pricing").map_sql_err()?,
"models.tiered_pricing",
@@ -790,8 +790,7 @@ fn map_admin_provider_model_row(
)?,
row.try_get("global_model_name").map_sql_err()?,
row.try_get("global_model_display_name").map_sql_err()?,
row.try_get("global_model_default_price_per_request")
.map_sql_err()?,
sqlite_optional_real(row, "global_model_default_price_per_request")?,
optional_json_from_string(
row.try_get("global_model_default_tiered_pricing")
.map_sql_err()?,

View File

@@ -8,6 +8,7 @@ pub mod announcements;
pub mod audit;
pub mod auth;
pub mod auth_modules;
pub mod background_tasks;
pub mod billing;
pub mod candidate_selection;
pub mod candidates;

View File

@@ -7,7 +7,7 @@ use super::{
StoredProviderCatalogKey, StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats,
StoredProviderCatalogProvider,
};
use crate::driver::sqlite::SqlitePool;
use crate::driver::sqlite::{sqlite_optional_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::DataLayerError;
@@ -34,7 +34,9 @@ impl SqliteProviderCatalogReadRepository {
r#"
SELECT
id, name, description, website, provider_type, billing_type,
monthly_quota_usd, monthly_used_usd, quota_reset_day,
CAST(monthly_quota_usd AS REAL) AS monthly_quota_usd,
CAST(monthly_used_usd AS REAL) AS monthly_used_usd,
quota_reset_day,
quota_last_reset_at AS quota_last_reset_at_unix_secs,
quota_expires_at AS quota_expires_at_unix_secs,
provider_priority, is_active, keep_priority_on_conversion,
@@ -1258,8 +1260,8 @@ fn map_provider_row(row: &SqliteRow) -> Result<StoredProviderCatalogProvider, Da
.with_description(row.try_get("description").map_sql_err()?)
.with_billing_fields(
row.try_get("billing_type").map_sql_err()?,
row.try_get("monthly_quota_usd").map_sql_err()?,
row.try_get("monthly_used_usd").map_sql_err()?,
sqlite_optional_real(row, "monthly_quota_usd")?,
sqlite_optional_real(row, "monthly_used_usd")?,
optional_u64(
row.try_get("quota_reset_day").map_sql_err()?,
"providers.quota_reset_day",
@@ -1316,11 +1318,7 @@ fn map_endpoint_row(row: &SqliteRow) -> Result<StoredProviderCatalogEndpoint, Da
"provider_endpoints.updated_at",
)?,
)
.with_health_score(
row.try_get::<Option<f64>, _>("health_score")
.map_sql_err()?
.unwrap_or(1.0),
)
.with_health_score(sqlite_optional_real(row, "health_score")?.unwrap_or(1.0))
.with_transport_fields(
row.try_get("base_url").map_sql_err()?,
optional_json_from_string(
@@ -1349,10 +1347,7 @@ fn map_endpoint_row(row: &SqliteRow) -> Result<StoredProviderCatalogEndpoint, Da
}
fn map_key_row(row: &SqliteRow) -> Result<StoredProviderCatalogKey, DataLayerError> {
let total_cost_usd = row
.try_get::<Option<f64>, _>("total_cost_usd")
.map_sql_err()?
.unwrap_or(0.0);
let total_cost_usd = sqlite_optional_real(row, "total_cost_usd")?.unwrap_or(0.0);
if !total_cost_usd.is_finite() {
return Err(DataLayerError::UnexpectedValue(
"invalid provider_api_keys.total_cost_usd".to_string(),
@@ -1571,6 +1566,8 @@ mod tests {
.expect("providers should list");
assert_eq!(providers.len(), 1);
assert_eq!(providers[0].provider_priority, 10);
assert_eq!(providers[0].monthly_quota_usd, Some(0.0));
assert_eq!(providers[0].monthly_used_usd, Some(0.0));
let endpoints = repository
.list_endpoints_by_provider_ids(&["provider-1".to_string()])
@@ -1827,11 +1824,12 @@ mod tests {
r#"
INSERT INTO providers (
id, name, description, website, provider_type, provider_priority,
monthly_quota_usd, monthly_used_usd,
is_active, keep_priority_on_conversion, enable_format_conversion,
config, created_at, updated_at
) VALUES (
'provider-1', 'Provider One', 'test provider', 'https://example.com',
'custom', 10, 1, 1, 1, '{"region":"us"}', 1, 2
'custom', 10, 0, 0, 1, 1, 1, '{"region":"us"}', 1, 2
)
"#,
)

View File

@@ -7,10 +7,14 @@ use serde_json::{json, Map, Value};
use uuid::Uuid;
use super::types::{
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeHeartbeatMutation,
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeReadRepository,
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery,
ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository,
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyNode, StoredProxyNodeEvent,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket,
StoredProxyNode, StoredProxyNodeEvent, StoredProxyNodeMetricsBucket, TunnelMetricsSample,
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
};
use crate::DataLayerError;
@@ -18,6 +22,8 @@ use crate::DataLayerError;
pub struct InMemoryProxyNodeRepository {
nodes: RwLock<BTreeMap<String, StoredProxyNode>>,
events: RwLock<Vec<StoredProxyNodeEvent>>,
metrics_1m: RwLock<BTreeMap<(String, u64), StoredProxyNodeMetricsBucket>>,
metrics_1h: RwLock<BTreeMap<(String, u64), StoredProxyNodeMetricsBucket>>,
}
impl InMemoryProxyNodeRepository {
@@ -33,6 +39,8 @@ impl InMemoryProxyNodeRepository {
.collect(),
),
events: RwLock::new(Vec::new()),
metrics_1m: RwLock::new(BTreeMap::new()),
metrics_1h: RwLock::new(BTreeMap::new()),
}
}
@@ -49,6 +57,8 @@ impl InMemoryProxyNodeRepository {
.collect(),
),
events: RwLock::new(events.into_iter().collect()),
metrics_1m: RwLock::new(BTreeMap::new()),
metrics_1h: RwLock::new(BTreeMap::new()),
}
}
@@ -63,6 +73,50 @@ impl InMemoryProxyNodeRepository {
events.iter().map(|event| event.id).max().unwrap_or(0) + 1
}
fn upsert_metrics_bucket(
metrics: &mut BTreeMap<(String, u64), StoredProxyNodeMetricsBucket>,
node_id: &str,
bucket_start_unix_secs: u64,
sample: &TunnelMetricsSample,
) {
let key = (node_id.to_string(), bucket_start_unix_secs);
let bucket = metrics
.entry(key)
.or_insert_with(|| StoredProxyNodeMetricsBucket {
node_id: node_id.to_string(),
bucket_start_unix_secs,
samples: 0,
uptime_samples: 0,
active_connections_sum: 0,
active_connections_max: 0,
heartbeat_rtt_ms_sum: 0,
heartbeat_rtt_ms_max: 0,
connect_errors_delta: 0,
disconnects_delta: 0,
error_events_delta: 0,
ws_in_bytes_delta: 0,
ws_out_bytes_delta: 0,
ws_in_frames_delta: 0,
ws_out_frames_delta: 0,
});
bucket.samples += sample.samples;
bucket.uptime_samples += sample.uptime_samples;
bucket.active_connections_sum += sample.active_connections_sum;
bucket.active_connections_max = bucket
.active_connections_max
.max(sample.active_connections_max);
bucket.heartbeat_rtt_ms_sum += sample.heartbeat_rtt_ms_sum;
bucket.heartbeat_rtt_ms_max = bucket.heartbeat_rtt_ms_max.max(sample.heartbeat_rtt_ms_max);
bucket.connect_errors_delta += sample.connect_errors_delta;
bucket.disconnects_delta += sample.disconnects_delta;
bucket.error_events_delta += sample.error_events_delta;
bucket.ws_in_bytes_delta += sample.ws_in_bytes_delta;
bucket.ws_out_bytes_delta += sample.ws_out_bytes_delta;
bucket.ws_in_frames_delta += sample.ws_in_frames_delta;
bucket.ws_out_frames_delta += sample.ws_out_frames_delta;
}
fn normalize_remote_config(
mutation: &ProxyNodeRemoteConfigMutation,
existing: Option<&Value>,
@@ -154,6 +208,130 @@ impl ProxyNodeReadRepository for InMemoryProxyNodeRepository {
items.truncate(limit);
Ok(items)
}
async fn list_proxy_node_events_filtered(
&self,
node_id: &str,
query: &ProxyNodeEventQuery,
) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> {
let events = self.events.read().expect("proxy node repository lock");
let mut items = events
.iter()
.filter(|event| event.node_id == node_id)
.filter(|event| {
query
.from_unix_secs
.map(|from| event.created_at_unix_ms.unwrap_or(0) >= from)
.unwrap_or(true)
})
.filter(|event| {
query
.to_unix_secs
.map(|to| event.created_at_unix_ms.unwrap_or(u64::MAX) <= to)
.unwrap_or(true)
})
.filter(|event| {
query
.event_type
.as_deref()
.map(|event_type| event.event_type.eq_ignore_ascii_case(event_type))
.unwrap_or(true)
})
.cloned()
.collect::<Vec<_>>();
items.sort_by(|left, right| {
right
.created_at_unix_ms
.unwrap_or(0)
.cmp(&left.created_at_unix_ms.unwrap_or(0))
.then(right.id.cmp(&left.id))
});
items.truncate(query.limit);
Ok(items)
}
async fn list_proxy_node_metrics(
&self,
node_id: &str,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyNodeMetricsBucket>, DataLayerError> {
let metrics = match step {
ProxyNodeMetricsStep::OneMinute => self.metrics_1m.read(),
ProxyNodeMetricsStep::OneHour => self.metrics_1h.read(),
}
.expect("proxy node repository lock");
let mut items = metrics
.values()
.filter(|bucket| bucket.node_id == node_id)
.filter(|bucket| bucket.bucket_start_unix_secs >= from_unix_secs)
.filter(|bucket| bucket.bucket_start_unix_secs <= to_unix_secs)
.cloned()
.collect::<Vec<_>>();
items.sort_by_key(|bucket| bucket.bucket_start_unix_secs);
items.truncate(limit);
Ok(items)
}
async fn list_proxy_fleet_metrics(
&self,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyFleetMetricsBucket>, DataLayerError> {
let metrics = match step {
ProxyNodeMetricsStep::OneMinute => self.metrics_1m.read(),
ProxyNodeMetricsStep::OneHour => self.metrics_1h.read(),
}
.expect("proxy node repository lock");
let mut grouped = BTreeMap::<u64, StoredProxyFleetMetricsBucket>::new();
for bucket in metrics.values() {
if bucket.bucket_start_unix_secs < from_unix_secs
|| bucket.bucket_start_unix_secs > to_unix_secs
{
continue;
}
let item = grouped
.entry(bucket.bucket_start_unix_secs)
.or_insert_with(|| StoredProxyFleetMetricsBucket {
bucket_start_unix_secs: bucket.bucket_start_unix_secs,
samples: 0,
uptime_samples: 0,
active_connections_sum: 0,
active_connections_max: 0,
heartbeat_rtt_ms_sum: 0,
heartbeat_rtt_ms_max: 0,
connect_errors_delta: 0,
disconnects_delta: 0,
error_events_delta: 0,
ws_in_bytes_delta: 0,
ws_out_bytes_delta: 0,
ws_in_frames_delta: 0,
ws_out_frames_delta: 0,
});
item.samples += bucket.samples;
item.uptime_samples += bucket.uptime_samples;
item.active_connections_sum += bucket.active_connections_sum;
item.active_connections_max = item
.active_connections_max
.max(bucket.active_connections_max);
item.heartbeat_rtt_ms_sum += bucket.heartbeat_rtt_ms_sum;
item.heartbeat_rtt_ms_max = item.heartbeat_rtt_ms_max.max(bucket.heartbeat_rtt_ms_max);
item.connect_errors_delta += bucket.connect_errors_delta;
item.disconnects_delta += bucket.disconnects_delta;
item.error_events_delta += bucket.error_events_delta;
item.ws_in_bytes_delta += bucket.ws_in_bytes_delta;
item.ws_out_bytes_delta += bucket.ws_out_bytes_delta;
item.ws_in_frames_delta += bucket.ws_in_frames_delta;
item.ws_out_frames_delta += bucket.ws_out_frames_delta;
}
let mut items = grouped.into_values().collect::<Vec<_>>();
items.truncate(limit);
Ok(items)
}
}
#[async_trait]
@@ -376,65 +554,114 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
&self,
mutation: &ProxyNodeHeartbeatMutation,
) -> Result<Option<StoredProxyNode>, DataLayerError> {
let mut nodes = self.nodes.write().expect("proxy node repository lock");
let Some(node) = nodes.get_mut(&mutation.node_id) else {
return Ok(None);
let (node, sample, now_unix_secs) = {
let mut nodes = self.nodes.write().expect("proxy node repository lock");
let Some(node) = nodes.get_mut(&mutation.node_id) else {
return Ok(None);
};
if !node.tunnel_mode {
return Err(DataLayerError::InvalidInput(
"non-tunnel mode is no longer supported, please upgrade aether-proxy to use tunnel mode"
.to_string(),
));
}
let previous_proxy_metadata = node.proxy_metadata.clone();
let now_unix_secs = Self::now_unix_secs().unwrap_or(0);
let now = Some(now_unix_secs);
node.last_heartbeat_at_unix_secs = now;
if node.status != "online" || !node.tunnel_connected {
node.status = "online".to_string();
node.tunnel_connected = true;
node.tunnel_connected_at_unix_secs = now;
node.updated_at_unix_secs = now;
}
if let Some(value) = mutation.heartbeat_interval {
node.heartbeat_interval = value;
}
if let Some(value) = mutation.active_connections {
node.active_connections = value;
}
if let Some(value) = mutation.avg_latency_ms {
node.avg_latency_ms = Some(value);
}
let normalized_proxy_metadata = normalize_proxy_metadata(
mutation.proxy_metadata.as_ref(),
mutation.proxy_version.as_deref(),
);
if let Some(value) = normalized_proxy_metadata {
node.proxy_metadata = Some(value);
}
if let Some(value) = mutation.total_requests_delta.filter(|value| *value > 0) {
node.total_requests += value;
}
if let Some(value) = mutation.failed_requests_delta.filter(|value| *value > 0) {
node.failed_requests += value;
}
if let Some(value) = mutation.dns_failures_delta.filter(|value| *value > 0) {
node.dns_failures += value;
}
if let Some(value) = mutation.stream_errors_delta.filter(|value| *value > 0) {
node.stream_errors += value;
}
let reconciled_remote_config = reconcile_remote_config_after_heartbeat(
node.remote_config.as_ref(),
mutation.proxy_version.as_deref(),
);
if reconciled_remote_config != node.remote_config {
node.remote_config = reconciled_remote_config;
node.config_version = node.config_version.saturating_add(1);
node.updated_at_unix_secs = now;
}
let sample = build_tunnel_metrics_sample(
previous_proxy_metadata.as_ref(),
node.proxy_metadata.as_ref(),
node.active_connections,
node.tunnel_connected,
);
(node.clone(), sample, now_unix_secs)
};
if !node.tunnel_mode {
return Err(DataLayerError::InvalidInput(
"non-tunnel mode is no longer supported, please upgrade aether-proxy to use tunnel mode"
.to_string(),
));
if let Some(sample) = sample.as_ref() {
Self::upsert_metrics_bucket(
&mut self.metrics_1m.write().expect("proxy node repository lock"),
&node.id,
bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneMinute),
sample,
);
Self::upsert_metrics_bucket(
&mut self.metrics_1h.write().expect("proxy node repository lock"),
&node.id,
bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneHour),
sample,
);
let mut events = self.events.write().expect("proxy node repository lock");
for error in &sample.recent_error_events {
let event_id = Self::next_event_id(&events);
events.push(StoredProxyNodeEvent {
id: event_id,
node_id: node.id.clone(),
event_type: PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR.to_string(),
detail: Some(build_tunnel_error_event_detail(error)),
event_metadata: Some(json!({
"source": "heartbeat",
"category": error.category,
"message": error.message,
"timestamp_unix_secs": error.timestamp_unix_secs,
})),
created_at_unix_ms: Some(if error.timestamp_unix_secs == 0 {
now_unix_secs
} else {
error.timestamp_unix_secs
}),
});
}
}
let now = Self::now_unix_secs();
node.last_heartbeat_at_unix_secs = now;
if node.status != "online" || !node.tunnel_connected {
node.status = "online".to_string();
node.tunnel_connected = true;
node.tunnel_connected_at_unix_secs = now;
node.updated_at_unix_secs = now;
}
if let Some(value) = mutation.heartbeat_interval {
node.heartbeat_interval = value;
}
if let Some(value) = mutation.active_connections {
node.active_connections = value;
}
if let Some(value) = mutation.avg_latency_ms {
node.avg_latency_ms = Some(value);
}
let normalized_proxy_metadata = normalize_proxy_metadata(
mutation.proxy_metadata.as_ref(),
mutation.proxy_version.as_deref(),
);
if let Some(value) = normalized_proxy_metadata {
node.proxy_metadata = Some(value);
}
if let Some(value) = mutation.total_requests_delta.filter(|value| *value > 0) {
node.total_requests += value;
}
if let Some(value) = mutation.failed_requests_delta.filter(|value| *value > 0) {
node.failed_requests += value;
}
if let Some(value) = mutation.dns_failures_delta.filter(|value| *value > 0) {
node.dns_failures += value;
}
if let Some(value) = mutation.stream_errors_delta.filter(|value| *value > 0) {
node.stream_errors += value;
}
let reconciled_remote_config = reconcile_remote_config_after_heartbeat(
node.remote_config.as_ref(),
mutation.proxy_version.as_deref(),
);
if reconciled_remote_config != node.remote_config {
node.remote_config = reconciled_remote_config;
node.config_version = node.config_version.saturating_add(1);
node.updated_at_unix_secs = now;
}
Ok(Some(node.clone()))
Ok(Some(node))
}
async fn record_traffic(
@@ -490,6 +717,7 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
node_id: mutation.node_id.clone(),
event_type: event_type.to_string(),
detail: Some(format!("[stale_ignored] {event_detail}")),
event_metadata: None,
created_at_unix_ms: Self::now_unix_secs(),
});
return Ok(Some(node.clone()));
@@ -513,6 +741,7 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
node_id: mutation.node_id.clone(),
event_type: event_type.to_string(),
detail: Some(event_detail),
event_metadata: None,
created_at_unix_ms: Some(event_time),
});
Ok(Some(node.clone()))
@@ -546,6 +775,14 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
.write()
.expect("proxy node repository lock")
.retain(|event| event.node_id != node_id);
self.metrics_1m
.write()
.expect("proxy node repository lock")
.retain(|(metric_node_id, _), _| metric_node_id != node_id);
self.metrics_1h
.write()
.expect("proxy node repository lock")
.retain(|(metric_node_id, _), _| metric_node_id != node_id);
}
Ok(removed)
}
@@ -598,6 +835,44 @@ impl ProxyNodeWriteRepository for InMemoryProxyNodeRepository {
}
Ok(())
}
async fn cleanup_proxy_node_metrics(
&self,
retain_1m_from_unix_secs: u64,
retain_1h_from_unix_secs: u64,
delete_limit: usize,
) -> Result<ProxyNodeMetricsCleanupSummary, DataLayerError> {
let delete_limit = delete_limit.max(1);
let mut metrics_1m = self.metrics_1m.write().expect("proxy node repository lock");
let expired_1m_keys = metrics_1m
.keys()
.filter(|(_, bucket_start)| *bucket_start < retain_1m_from_unix_secs)
.take(delete_limit)
.cloned()
.collect::<Vec<_>>();
let deleted_1m_rows = expired_1m_keys
.iter()
.filter(|key| metrics_1m.remove(key).is_some())
.count();
drop(metrics_1m);
let mut metrics_1h = self.metrics_1h.write().expect("proxy node repository lock");
let expired_1h_keys = metrics_1h
.keys()
.filter(|(_, bucket_start)| *bucket_start < retain_1h_from_unix_secs)
.take(delete_limit)
.cloned()
.collect::<Vec<_>>();
let deleted_1h_rows = expired_1h_keys
.iter()
.filter(|key| metrics_1h.remove(key).is_some())
.count();
Ok(ProxyNodeMetricsCleanupSummary {
deleted_1m_rows,
deleted_1h_rows,
})
}
}
#[cfg(test)]
@@ -750,6 +1025,7 @@ mod tests {
node_id: "node-1".to_string(),
event_type: "connected".to_string(),
detail: Some("older".to_string()),
event_metadata: None,
created_at_unix_ms: Some(1_710_000_000),
},
StoredProxyNodeEvent {
@@ -757,6 +1033,7 @@ mod tests {
node_id: "node-1".to_string(),
event_type: "disconnected".to_string(),
detail: Some("newer".to_string()),
event_metadata: None,
created_at_unix_ms: Some(1_710_000_100),
},
],

View File

@@ -9,11 +9,14 @@ pub use mysql::MysqlProxyNodeReadRepository;
pub use postgres::SqlxProxyNodeRepository;
pub use sqlite::SqliteProxyNodeReadRepository;
pub use types::{
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
normalize_proxy_node_scheduling_state, proxy_node_accepts_new_tunnels, proxy_reported_version,
reconcile_remote_config_after_heartbeat, remote_config_scheduling_state,
remote_config_upgrade_target, ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation,
ProxyNodeManualUpdateMutation, ProxyNodeReadRepository, ProxyNodeRegistrationMutation,
remote_config_upgrade_target, ProxyNodeEventQuery, ProxyNodeHeartbeatMutation,
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeMetricsCleanupSummary,
ProxyNodeMetricsStep, ProxyNodeReadRepository, ProxyNodeRegistrationMutation,
ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation, ProxyNodeTunnelStatusMutation,
ProxyNodeWriteRepository, StoredProxyNode, StoredProxyNodeEvent,
ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket, StoredProxyNode, StoredProxyNodeEvent,
StoredProxyNodeMetricsBucket, TunnelErrorEventRecord, PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
PROXY_NODE_SCHEDULING_STATE_CORDONED, PROXY_NODE_SCHEDULING_STATE_DRAINING,
};

View File

@@ -2,10 +2,14 @@ use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, Row};
use super::types::{
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeHeartbeatMutation,
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeReadRepository,
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery,
ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository,
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyNode, StoredProxyNodeEvent,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket,
StoredProxyNode, StoredProxyNodeEvent, StoredProxyNodeMetricsBucket,
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
};
use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt;
@@ -148,17 +152,22 @@ ON DUPLICATE KEY UPDATE
node_id: &str,
event_type: &str,
detail: Option<&str>,
event_metadata: Option<&serde_json::Value>,
created_at_unix_secs: Option<u64>,
) -> Result<(), DataLayerError> {
sqlx::query(
r#"
INSERT INTO proxy_node_events (node_id, event_type, detail, created_at)
VALUES (?, ?, ?, ?)
INSERT INTO proxy_node_events (node_id, event_type, detail, event_metadata, created_at)
VALUES (?, ?, ?, ?, ?)
"#,
)
.bind(node_id)
.bind(event_type)
.bind(detail)
.bind(optional_json_to_string(
&event_metadata.cloned(),
"proxy_node_events.event_metadata",
)?)
.bind(created_at_unix_secs.unwrap_or_else(current_unix_secs) as i64)
.execute(&self.pool)
.await
@@ -166,6 +175,70 @@ VALUES (?, ?, ?, ?)
Ok(())
}
async fn upsert_metrics_bucket(
&self,
table: &str,
node_id: &str,
bucket_start: u64,
sample: &super::types::TunnelMetricsSample,
) -> Result<(), DataLayerError> {
sqlx::query(&format!(
r#"
INSERT INTO {table} (
node_id,
bucket_start_unix_secs,
samples,
uptime_samples,
active_connections_sum,
active_connections_max,
heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max,
connect_errors_delta,
disconnects_delta,
error_events_delta,
ws_in_bytes_delta,
ws_out_bytes_delta,
ws_in_frames_delta,
ws_out_frames_delta
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON DUPLICATE KEY UPDATE
samples = samples + VALUES(samples),
uptime_samples = uptime_samples + VALUES(uptime_samples),
active_connections_sum = active_connections_sum + VALUES(active_connections_sum),
active_connections_max = GREATEST(active_connections_max, VALUES(active_connections_max)),
heartbeat_rtt_ms_sum = heartbeat_rtt_ms_sum + VALUES(heartbeat_rtt_ms_sum),
heartbeat_rtt_ms_max = GREATEST(heartbeat_rtt_ms_max, VALUES(heartbeat_rtt_ms_max)),
connect_errors_delta = connect_errors_delta + VALUES(connect_errors_delta),
disconnects_delta = disconnects_delta + VALUES(disconnects_delta),
error_events_delta = error_events_delta + VALUES(error_events_delta),
ws_in_bytes_delta = ws_in_bytes_delta + VALUES(ws_in_bytes_delta),
ws_out_bytes_delta = ws_out_bytes_delta + VALUES(ws_out_bytes_delta),
ws_in_frames_delta = ws_in_frames_delta + VALUES(ws_in_frames_delta),
ws_out_frames_delta = ws_out_frames_delta + VALUES(ws_out_frames_delta)
"#
))
.bind(node_id)
.bind(i64::try_from(bucket_start).unwrap_or(i64::MAX))
.bind(sample.samples)
.bind(sample.uptime_samples)
.bind(sample.active_connections_sum)
.bind(sample.active_connections_max)
.bind(sample.heartbeat_rtt_ms_sum)
.bind(sample.heartbeat_rtt_ms_max)
.bind(sample.connect_errors_delta)
.bind(sample.disconnects_delta)
.bind(sample.error_events_delta)
.bind(sample.ws_in_bytes_delta)
.bind(sample.ws_out_bytes_delta)
.bind(sample.ws_in_frames_delta)
.bind(sample.ws_out_frames_delta)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(())
}
fn normalize_remote_config(
mutation: &ProxyNodeRemoteConfigMutation,
existing: Option<&serde_json::Value>,
@@ -298,6 +371,7 @@ SELECT
node_id,
event_type,
detail,
event_metadata,
created_at AS created_at_unix_ms
FROM proxy_node_events
WHERE node_id = ?
@@ -312,6 +386,152 @@ LIMIT ?
.map_sql_err()?;
rows.iter().map(map_proxy_node_event_row).collect()
}
async fn list_proxy_node_events_filtered(
&self,
node_id: &str,
query: &ProxyNodeEventQuery,
) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> {
let rows = sqlx::query(
r#"
SELECT
id,
node_id,
event_type,
detail,
event_metadata,
created_at AS created_at_unix_ms
FROM proxy_node_events
WHERE node_id = ?
AND (? IS NULL OR created_at >= ?)
AND (? IS NULL OR created_at <= ?)
AND (? IS NULL OR LOWER(event_type) = LOWER(?))
ORDER BY created_at DESC, id DESC
LIMIT ?
"#,
)
.bind(node_id)
.bind(
query
.from_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(
query
.from_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(
query
.to_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(
query
.to_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(query.event_type.as_deref())
.bind(query.event_type.as_deref())
.bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_proxy_node_event_row).collect()
}
async fn list_proxy_node_metrics(
&self,
node_id: &str,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyNodeMetricsBucket>, DataLayerError> {
let table = match step {
ProxyNodeMetricsStep::OneMinute => "proxy_node_metrics_1m",
ProxyNodeMetricsStep::OneHour => "proxy_node_metrics_1h",
};
let rows = sqlx::query(&format!(
r#"
SELECT
node_id,
bucket_start_unix_secs,
samples,
uptime_samples,
active_connections_sum,
active_connections_max,
heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max,
connect_errors_delta,
disconnects_delta,
error_events_delta,
ws_in_bytes_delta,
ws_out_bytes_delta,
ws_in_frames_delta,
ws_out_frames_delta
FROM {table}
WHERE node_id = ?
AND bucket_start_unix_secs >= ?
AND bucket_start_unix_secs <= ?
ORDER BY bucket_start_unix_secs ASC
LIMIT ?
"#
))
.bind(node_id)
.bind(i64::try_from(from_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(to_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_proxy_node_metric_row).collect()
}
async fn list_proxy_fleet_metrics(
&self,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyFleetMetricsBucket>, DataLayerError> {
let table = match step {
ProxyNodeMetricsStep::OneMinute => "proxy_node_metrics_1m",
ProxyNodeMetricsStep::OneHour => "proxy_node_metrics_1h",
};
let rows = sqlx::query(&format!(
r#"
SELECT
bucket_start_unix_secs,
SUM(samples) AS samples,
SUM(uptime_samples) AS uptime_samples,
SUM(active_connections_sum) AS active_connections_sum,
MAX(active_connections_max) AS active_connections_max,
SUM(heartbeat_rtt_ms_sum) AS heartbeat_rtt_ms_sum,
MAX(heartbeat_rtt_ms_max) AS heartbeat_rtt_ms_max,
SUM(connect_errors_delta) AS connect_errors_delta,
SUM(disconnects_delta) AS disconnects_delta,
SUM(error_events_delta) AS error_events_delta,
SUM(ws_in_bytes_delta) AS ws_in_bytes_delta,
SUM(ws_out_bytes_delta) AS ws_out_bytes_delta,
SUM(ws_in_frames_delta) AS ws_in_frames_delta,
SUM(ws_out_frames_delta) AS ws_out_frames_delta
FROM {table}
WHERE bucket_start_unix_secs >= ?
AND bucket_start_unix_secs <= ?
GROUP BY bucket_start_unix_secs
ORDER BY bucket_start_unix_secs ASC
LIMIT ?
"#
))
.bind(i64::try_from(from_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(to_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_proxy_fleet_metric_row).collect()
}
}
#[async_trait]
@@ -540,7 +760,9 @@ WHERE is_manual = 0
));
}
let now = Some(current_unix_secs());
let previous_proxy_metadata = node.proxy_metadata.clone();
let now_unix_secs = current_unix_secs();
let now = Some(now_unix_secs);
node.last_heartbeat_at_unix_secs = now;
if node.status != "online" || !node.tunnel_connected {
node.status = "online".to_string();
@@ -584,7 +806,52 @@ WHERE is_manual = 0
node.config_version = node.config_version.saturating_add(1);
node.updated_at_unix_secs = now;
}
let tunnel_metrics_sample = build_tunnel_metrics_sample(
previous_proxy_metadata.as_ref(),
node.proxy_metadata.as_ref(),
node.active_connections,
node.tunnel_connected,
);
self.upsert_node(&node).await?;
if let Some(sample) = tunnel_metrics_sample.as_ref() {
self.upsert_metrics_bucket(
"proxy_node_metrics_1m",
&node.id,
bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneMinute),
sample,
)
.await?;
self.upsert_metrics_bucket(
"proxy_node_metrics_1h",
&node.id,
bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneHour),
sample,
)
.await?;
for error in &sample.recent_error_events {
let detail = build_tunnel_error_event_detail(error);
let event_metadata = serde_json::json!({
"source": "heartbeat",
"category": error.category,
"message": error.message,
"timestamp_unix_secs": error.timestamp_unix_secs,
});
self.insert_event(
&node.id,
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
Some(detail.as_str()),
Some(&event_metadata),
Some(if error.timestamp_unix_secs == 0 {
now_unix_secs
} else {
error.timestamp_unix_secs
}),
)
.await?;
}
}
Ok(Some(node))
}
@@ -638,6 +905,7 @@ WHERE is_manual = 0
&mutation.node_id,
event_type,
Some(&format!("[stale_ignored] {event_detail}")),
None,
Some(current_unix_secs()),
)
.await?;
@@ -660,6 +928,7 @@ WHERE is_manual = 0
&mutation.node_id,
event_type,
Some(&event_detail),
None,
Some(event_time),
)
.await?;
@@ -691,6 +960,16 @@ WHERE is_manual = 0
.execute(&self.pool)
.await
.map_sql_err()?;
sqlx::query("DELETE FROM proxy_node_metrics_1m WHERE node_id = ?")
.bind(node_id)
.execute(&self.pool)
.await
.map_sql_err()?;
sqlx::query("DELETE FROM proxy_node_metrics_1h WHERE node_id = ?")
.bind(node_id)
.execute(&self.pool)
.await
.map_sql_err()?;
sqlx::query("DELETE FROM proxy_nodes WHERE id = ?")
.bind(node_id)
.execute(&self.pool)
@@ -747,6 +1026,49 @@ WHERE is_manual = 0
node.updated_at_unix_secs = Some(current_unix_secs());
self.upsert_node(&node).await
}
async fn cleanup_proxy_node_metrics(
&self,
retain_1m_from_unix_secs: u64,
retain_1h_from_unix_secs: u64,
delete_limit: usize,
) -> Result<ProxyNodeMetricsCleanupSummary, DataLayerError> {
let delete_limit_i64 = i64::try_from(delete_limit.max(1)).unwrap_or(i64::MAX);
let deleted_1m = sqlx::query(
r#"
DELETE FROM proxy_node_metrics_1m
WHERE bucket_start_unix_secs < ?
ORDER BY bucket_start_unix_secs ASC
LIMIT ?
"#,
)
.bind(i64::try_from(retain_1m_from_unix_secs).unwrap_or(i64::MAX))
.bind(delete_limit_i64)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected() as usize;
let deleted_1h = sqlx::query(
r#"
DELETE FROM proxy_node_metrics_1h
WHERE bucket_start_unix_secs < ?
ORDER BY bucket_start_unix_secs ASC
LIMIT ?
"#,
)
.bind(i64::try_from(retain_1h_from_unix_secs).unwrap_or(i64::MAX))
.bind(delete_limit_i64)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected() as usize;
Ok(ProxyNodeMetricsCleanupSummary {
deleted_1m_rows: deleted_1m,
deleted_1h_rows: deleted_1h,
})
}
}
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
@@ -861,10 +1183,63 @@ fn map_proxy_node_event_row(row: &MySqlRow) -> Result<StoredProxyNodeEvent, Data
node_id: row.try_get("node_id").map_sql_err()?,
event_type: row.try_get("event_type").map_sql_err()?,
detail: row.try_get("detail").map_sql_err()?,
event_metadata: optional_json_from_string(
row.try_get("event_metadata").map_sql_err()?,
"proxy_node_events.event_metadata",
)?,
created_at_unix_ms: optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
})
}
fn map_proxy_node_metric_row(
row: &MySqlRow,
) -> Result<StoredProxyNodeMetricsBucket, DataLayerError> {
Ok(StoredProxyNodeMetricsBucket {
node_id: row.try_get("node_id").map_sql_err()?,
bucket_start_unix_secs: optional_unix_secs(
row.try_get("bucket_start_unix_secs").map_sql_err()?,
)
.unwrap_or_default(),
samples: row.try_get("samples").map_sql_err()?,
uptime_samples: row.try_get("uptime_samples").map_sql_err()?,
active_connections_sum: row.try_get("active_connections_sum").map_sql_err()?,
active_connections_max: row.try_get("active_connections_max").map_sql_err()?,
heartbeat_rtt_ms_sum: row.try_get("heartbeat_rtt_ms_sum").map_sql_err()?,
heartbeat_rtt_ms_max: row.try_get("heartbeat_rtt_ms_max").map_sql_err()?,
connect_errors_delta: row.try_get("connect_errors_delta").map_sql_err()?,
disconnects_delta: row.try_get("disconnects_delta").map_sql_err()?,
error_events_delta: row.try_get("error_events_delta").map_sql_err()?,
ws_in_bytes_delta: row.try_get("ws_in_bytes_delta").map_sql_err()?,
ws_out_bytes_delta: row.try_get("ws_out_bytes_delta").map_sql_err()?,
ws_in_frames_delta: row.try_get("ws_in_frames_delta").map_sql_err()?,
ws_out_frames_delta: row.try_get("ws_out_frames_delta").map_sql_err()?,
})
}
fn map_proxy_fleet_metric_row(
row: &MySqlRow,
) -> Result<StoredProxyFleetMetricsBucket, DataLayerError> {
Ok(StoredProxyFleetMetricsBucket {
bucket_start_unix_secs: optional_unix_secs(
row.try_get("bucket_start_unix_secs").map_sql_err()?,
)
.unwrap_or_default(),
samples: row.try_get("samples").map_sql_err()?,
uptime_samples: row.try_get("uptime_samples").map_sql_err()?,
active_connections_sum: row.try_get("active_connections_sum").map_sql_err()?,
active_connections_max: row.try_get("active_connections_max").map_sql_err()?,
heartbeat_rtt_ms_sum: row.try_get("heartbeat_rtt_ms_sum").map_sql_err()?,
heartbeat_rtt_ms_max: row.try_get("heartbeat_rtt_ms_max").map_sql_err()?,
connect_errors_delta: row.try_get("connect_errors_delta").map_sql_err()?,
disconnects_delta: row.try_get("disconnects_delta").map_sql_err()?,
error_events_delta: row.try_get("error_events_delta").map_sql_err()?,
ws_in_bytes_delta: row.try_get("ws_in_bytes_delta").map_sql_err()?,
ws_out_bytes_delta: row.try_get("ws_out_bytes_delta").map_sql_err()?,
ws_in_frames_delta: row.try_get("ws_in_frames_delta").map_sql_err()?,
ws_out_frames_delta: row.try_get("ws_out_frames_delta").map_sql_err()?,
})
}
#[cfg(test)]
mod tests {
use super::MysqlProxyNodeReadRepository;

View File

@@ -4,10 +4,14 @@ use sha2::{Digest, Sha256};
use sqlx::{postgres::PgRow, PgPool, Row};
use super::types::{
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeHeartbeatMutation,
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeReadRepository,
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery,
ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository,
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyNode, StoredProxyNodeEvent,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket,
StoredProxyNode, StoredProxyNodeEvent, StoredProxyNodeMetricsBucket, TunnelMetricsSample,
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
};
use crate::{
error::{postgres_error, SqlxResultExt},
@@ -91,6 +95,7 @@ SELECT
node_id,
CAST(event_type AS TEXT) AS event_type,
detail,
event_metadata,
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms
FROM proxy_node_events
WHERE node_id = $1
@@ -98,6 +103,23 @@ ORDER BY created_at DESC, id DESC
LIMIT $2
"#;
const LIST_PROXY_NODE_EVENTS_FILTERED_SQL: &str = r#"
SELECT
id,
node_id,
CAST(event_type AS TEXT) AS event_type,
detail,
event_metadata,
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms
FROM proxy_node_events
WHERE node_id = $1
AND ($2::double precision IS NULL OR created_at >= TO_TIMESTAMP($2::double precision))
AND ($3::double precision IS NULL OR created_at <= TO_TIMESTAMP($3::double precision))
AND ($4::text IS NULL OR LOWER(CAST(event_type AS TEXT)) = LOWER($4::text))
ORDER BY created_at DESC, id DESC
LIMIT $5
"#;
const APPLY_HEARTBEAT_SQL: &str = r#"
UPDATE proxy_nodes
SET
@@ -381,6 +403,188 @@ WHERE id = $4
AND is_manual = TRUE
"#;
const INSERT_PROXY_NODE_EVENT_SQL: &str = r#"
INSERT INTO proxy_node_events (node_id, event_type, detail, event_metadata, created_at)
VALUES (
$1,
$2,
$3,
$4::json,
CASE
WHEN $5::double precision IS NULL THEN NOW()
ELSE TO_TIMESTAMP($5::double precision)
END
)
"#;
const UPSERT_PROXY_NODE_METRICS_1M_SQL: &str = r#"
INSERT INTO proxy_node_metrics_1m (
node_id,
bucket_start_unix_secs,
samples,
uptime_samples,
active_connections_sum,
active_connections_max,
heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max,
connect_errors_delta,
disconnects_delta,
error_events_delta,
ws_in_bytes_delta,
ws_out_bytes_delta,
ws_in_frames_delta,
ws_out_frames_delta
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15)
ON CONFLICT (node_id, bucket_start_unix_secs) DO UPDATE SET
samples = proxy_node_metrics_1m.samples + EXCLUDED.samples,
uptime_samples = proxy_node_metrics_1m.uptime_samples + EXCLUDED.uptime_samples,
active_connections_sum = proxy_node_metrics_1m.active_connections_sum + EXCLUDED.active_connections_sum,
active_connections_max = GREATEST(proxy_node_metrics_1m.active_connections_max, EXCLUDED.active_connections_max),
heartbeat_rtt_ms_sum = proxy_node_metrics_1m.heartbeat_rtt_ms_sum + EXCLUDED.heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max = GREATEST(proxy_node_metrics_1m.heartbeat_rtt_ms_max, EXCLUDED.heartbeat_rtt_ms_max),
connect_errors_delta = proxy_node_metrics_1m.connect_errors_delta + EXCLUDED.connect_errors_delta,
disconnects_delta = proxy_node_metrics_1m.disconnects_delta + EXCLUDED.disconnects_delta,
error_events_delta = proxy_node_metrics_1m.error_events_delta + EXCLUDED.error_events_delta,
ws_in_bytes_delta = proxy_node_metrics_1m.ws_in_bytes_delta + EXCLUDED.ws_in_bytes_delta,
ws_out_bytes_delta = proxy_node_metrics_1m.ws_out_bytes_delta + EXCLUDED.ws_out_bytes_delta,
ws_in_frames_delta = proxy_node_metrics_1m.ws_in_frames_delta + EXCLUDED.ws_in_frames_delta,
ws_out_frames_delta = proxy_node_metrics_1m.ws_out_frames_delta + EXCLUDED.ws_out_frames_delta
"#;
const UPSERT_PROXY_NODE_METRICS_1H_SQL: &str = r#"
INSERT INTO proxy_node_metrics_1h (
node_id,
bucket_start_unix_secs,
samples,
uptime_samples,
active_connections_sum,
active_connections_max,
heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max,
connect_errors_delta,
disconnects_delta,
error_events_delta,
ws_in_bytes_delta,
ws_out_bytes_delta,
ws_in_frames_delta,
ws_out_frames_delta
)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15)
ON CONFLICT (node_id, bucket_start_unix_secs) DO UPDATE SET
samples = proxy_node_metrics_1h.samples + EXCLUDED.samples,
uptime_samples = proxy_node_metrics_1h.uptime_samples + EXCLUDED.uptime_samples,
active_connections_sum = proxy_node_metrics_1h.active_connections_sum + EXCLUDED.active_connections_sum,
active_connections_max = GREATEST(proxy_node_metrics_1h.active_connections_max, EXCLUDED.active_connections_max),
heartbeat_rtt_ms_sum = proxy_node_metrics_1h.heartbeat_rtt_ms_sum + EXCLUDED.heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max = GREATEST(proxy_node_metrics_1h.heartbeat_rtt_ms_max, EXCLUDED.heartbeat_rtt_ms_max),
connect_errors_delta = proxy_node_metrics_1h.connect_errors_delta + EXCLUDED.connect_errors_delta,
disconnects_delta = proxy_node_metrics_1h.disconnects_delta + EXCLUDED.disconnects_delta,
error_events_delta = proxy_node_metrics_1h.error_events_delta + EXCLUDED.error_events_delta,
ws_in_bytes_delta = proxy_node_metrics_1h.ws_in_bytes_delta + EXCLUDED.ws_in_bytes_delta,
ws_out_bytes_delta = proxy_node_metrics_1h.ws_out_bytes_delta + EXCLUDED.ws_out_bytes_delta,
ws_in_frames_delta = proxy_node_metrics_1h.ws_in_frames_delta + EXCLUDED.ws_in_frames_delta,
ws_out_frames_delta = proxy_node_metrics_1h.ws_out_frames_delta + EXCLUDED.ws_out_frames_delta
"#;
const LIST_PROXY_NODE_METRICS_1M_SQL: &str = r#"
SELECT
node_id,
bucket_start_unix_secs,
samples,
uptime_samples,
active_connections_sum,
active_connections_max,
heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max,
connect_errors_delta,
disconnects_delta,
error_events_delta,
ws_in_bytes_delta,
ws_out_bytes_delta,
ws_in_frames_delta,
ws_out_frames_delta
FROM proxy_node_metrics_1m
WHERE node_id = $1
AND bucket_start_unix_secs >= $2
AND bucket_start_unix_secs <= $3
ORDER BY bucket_start_unix_secs ASC
LIMIT $4
"#;
const LIST_PROXY_NODE_METRICS_1H_SQL: &str = r#"
SELECT
node_id,
bucket_start_unix_secs,
samples,
uptime_samples,
active_connections_sum,
active_connections_max,
heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max,
connect_errors_delta,
disconnects_delta,
error_events_delta,
ws_in_bytes_delta,
ws_out_bytes_delta,
ws_in_frames_delta,
ws_out_frames_delta
FROM proxy_node_metrics_1h
WHERE node_id = $1
AND bucket_start_unix_secs >= $2
AND bucket_start_unix_secs <= $3
ORDER BY bucket_start_unix_secs ASC
LIMIT $4
"#;
const LIST_PROXY_FLEET_METRICS_1M_SQL: &str = r#"
SELECT
bucket_start_unix_secs,
SUM(samples) AS samples,
SUM(uptime_samples) AS uptime_samples,
SUM(active_connections_sum) AS active_connections_sum,
MAX(active_connections_max) AS active_connections_max,
SUM(heartbeat_rtt_ms_sum) AS heartbeat_rtt_ms_sum,
MAX(heartbeat_rtt_ms_max) AS heartbeat_rtt_ms_max,
SUM(connect_errors_delta) AS connect_errors_delta,
SUM(disconnects_delta) AS disconnects_delta,
SUM(error_events_delta) AS error_events_delta,
SUM(ws_in_bytes_delta) AS ws_in_bytes_delta,
SUM(ws_out_bytes_delta) AS ws_out_bytes_delta,
SUM(ws_in_frames_delta) AS ws_in_frames_delta,
SUM(ws_out_frames_delta) AS ws_out_frames_delta
FROM proxy_node_metrics_1m
WHERE bucket_start_unix_secs >= $1
AND bucket_start_unix_secs <= $2
GROUP BY bucket_start_unix_secs
ORDER BY bucket_start_unix_secs ASC
LIMIT $3
"#;
const LIST_PROXY_FLEET_METRICS_1H_SQL: &str = r#"
SELECT
bucket_start_unix_secs,
SUM(samples) AS samples,
SUM(uptime_samples) AS uptime_samples,
SUM(active_connections_sum) AS active_connections_sum,
MAX(active_connections_max) AS active_connections_max,
SUM(heartbeat_rtt_ms_sum) AS heartbeat_rtt_ms_sum,
MAX(heartbeat_rtt_ms_max) AS heartbeat_rtt_ms_max,
SUM(connect_errors_delta) AS connect_errors_delta,
SUM(disconnects_delta) AS disconnects_delta,
SUM(error_events_delta) AS error_events_delta,
SUM(ws_in_bytes_delta) AS ws_in_bytes_delta,
SUM(ws_out_bytes_delta) AS ws_out_bytes_delta,
SUM(ws_in_frames_delta) AS ws_in_frames_delta,
SUM(ws_out_frames_delta) AS ws_out_frames_delta
FROM proxy_node_metrics_1h
WHERE bucket_start_unix_secs >= $1
AND bucket_start_unix_secs <= $2
GROUP BY bucket_start_unix_secs
ORDER BY bucket_start_unix_secs ASC
LIMIT $3
"#;
#[derive(Debug, Clone)]
pub struct SqlxProxyNodeRepository {
pool: PgPool,
@@ -446,12 +650,111 @@ impl SqlxProxyNodeRepository {
node_id: row.try_get("node_id").map_postgres_err()?,
event_type: row.try_get("event_type").map_postgres_err()?,
detail: row.try_get("detail").map_postgres_err()?,
event_metadata: row.try_get("event_metadata").map_postgres_err()?,
created_at_unix_ms: Self::optional_unix_secs(
row.try_get("created_at_unix_ms").map_postgres_err()?,
),
})
}
fn row_to_node_metric(row: &PgRow) -> Result<StoredProxyNodeMetricsBucket, DataLayerError> {
Ok(StoredProxyNodeMetricsBucket {
node_id: row.try_get("node_id").map_postgres_err()?,
bucket_start_unix_secs: Self::optional_unix_secs(
row.try_get("bucket_start_unix_secs").map_postgres_err()?,
)
.unwrap_or_default(),
samples: row.try_get("samples").map_postgres_err()?,
uptime_samples: row.try_get("uptime_samples").map_postgres_err()?,
active_connections_sum: row.try_get("active_connections_sum").map_postgres_err()?,
active_connections_max: row.try_get("active_connections_max").map_postgres_err()?,
heartbeat_rtt_ms_sum: row.try_get("heartbeat_rtt_ms_sum").map_postgres_err()?,
heartbeat_rtt_ms_max: row.try_get("heartbeat_rtt_ms_max").map_postgres_err()?,
connect_errors_delta: row.try_get("connect_errors_delta").map_postgres_err()?,
disconnects_delta: row.try_get("disconnects_delta").map_postgres_err()?,
error_events_delta: row.try_get("error_events_delta").map_postgres_err()?,
ws_in_bytes_delta: row.try_get("ws_in_bytes_delta").map_postgres_err()?,
ws_out_bytes_delta: row.try_get("ws_out_bytes_delta").map_postgres_err()?,
ws_in_frames_delta: row.try_get("ws_in_frames_delta").map_postgres_err()?,
ws_out_frames_delta: row.try_get("ws_out_frames_delta").map_postgres_err()?,
})
}
fn row_to_fleet_metric(row: &PgRow) -> Result<StoredProxyFleetMetricsBucket, DataLayerError> {
Ok(StoredProxyFleetMetricsBucket {
bucket_start_unix_secs: Self::optional_unix_secs(
row.try_get("bucket_start_unix_secs").map_postgres_err()?,
)
.unwrap_or_default(),
samples: row.try_get("samples").map_postgres_err()?,
uptime_samples: row.try_get("uptime_samples").map_postgres_err()?,
active_connections_sum: row.try_get("active_connections_sum").map_postgres_err()?,
active_connections_max: row.try_get("active_connections_max").map_postgres_err()?,
heartbeat_rtt_ms_sum: row.try_get("heartbeat_rtt_ms_sum").map_postgres_err()?,
heartbeat_rtt_ms_max: row.try_get("heartbeat_rtt_ms_max").map_postgres_err()?,
connect_errors_delta: row.try_get("connect_errors_delta").map_postgres_err()?,
disconnects_delta: row.try_get("disconnects_delta").map_postgres_err()?,
error_events_delta: row.try_get("error_events_delta").map_postgres_err()?,
ws_in_bytes_delta: row.try_get("ws_in_bytes_delta").map_postgres_err()?,
ws_out_bytes_delta: row.try_get("ws_out_bytes_delta").map_postgres_err()?,
ws_in_frames_delta: row.try_get("ws_in_frames_delta").map_postgres_err()?,
ws_out_frames_delta: row.try_get("ws_out_frames_delta").map_postgres_err()?,
})
}
async fn insert_event(
&self,
node_id: &str,
event_type: &str,
detail: Option<&str>,
event_metadata: Option<&serde_json::Value>,
created_at_unix_secs: Option<u64>,
) -> Result<(), DataLayerError> {
sqlx::query(INSERT_PROXY_NODE_EVENT_SQL)
.bind(node_id)
.bind(event_type)
.bind(detail)
.bind(event_metadata)
.bind(created_at_unix_secs.map(|value| value as f64))
.execute(&self.pool)
.await
.map_postgres_err()?;
Ok(())
}
async fn upsert_metrics_bucket(
&self,
step: ProxyNodeMetricsStep,
node_id: &str,
bucket_start: u64,
sample: &TunnelMetricsSample,
) -> Result<(), DataLayerError> {
let sql = match step {
ProxyNodeMetricsStep::OneMinute => UPSERT_PROXY_NODE_METRICS_1M_SQL,
ProxyNodeMetricsStep::OneHour => UPSERT_PROXY_NODE_METRICS_1H_SQL,
};
sqlx::query(sql)
.bind(node_id)
.bind(i64::try_from(bucket_start).unwrap_or(i64::MAX))
.bind(sample.samples)
.bind(sample.uptime_samples)
.bind(sample.active_connections_sum)
.bind(sample.active_connections_max)
.bind(sample.heartbeat_rtt_ms_sum)
.bind(sample.heartbeat_rtt_ms_max)
.bind(sample.connect_errors_delta)
.bind(sample.disconnects_delta)
.bind(sample.error_events_delta)
.bind(sample.ws_in_bytes_delta)
.bind(sample.ws_out_bytes_delta)
.bind(sample.ws_in_frames_delta)
.bind(sample.ws_out_frames_delta)
.execute(&self.pool)
.await
.map_postgres_err()?;
Ok(())
}
fn registration_lock_key(ip: &str, port: i32) -> i64 {
let mut hasher = Sha256::new();
hasher.update(ip.as_bytes());
@@ -602,6 +905,73 @@ impl ProxyNodeReadRepository for SqlxProxyNodeRepository {
}
Ok(items)
}
async fn list_proxy_node_events_filtered(
&self,
node_id: &str,
query: &ProxyNodeEventQuery,
) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> {
let mut rows = sqlx::query(LIST_PROXY_NODE_EVENTS_FILTERED_SQL)
.bind(node_id)
.bind(query.from_unix_secs.map(|value| value as f64))
.bind(query.to_unix_secs.map(|value| value as f64))
.bind(query.event_type.as_deref())
.bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
.fetch(&self.pool);
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_postgres_err()? {
items.push(Self::row_to_event(&row)?);
}
Ok(items)
}
async fn list_proxy_node_metrics(
&self,
node_id: &str,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyNodeMetricsBucket>, DataLayerError> {
let sql = match step {
ProxyNodeMetricsStep::OneMinute => LIST_PROXY_NODE_METRICS_1M_SQL,
ProxyNodeMetricsStep::OneHour => LIST_PROXY_NODE_METRICS_1H_SQL,
};
let mut rows = sqlx::query(sql)
.bind(node_id)
.bind(i64::try_from(from_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(to_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.fetch(&self.pool);
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_postgres_err()? {
items.push(Self::row_to_node_metric(&row)?);
}
Ok(items)
}
async fn list_proxy_fleet_metrics(
&self,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyFleetMetricsBucket>, DataLayerError> {
let sql = match step {
ProxyNodeMetricsStep::OneMinute => LIST_PROXY_FLEET_METRICS_1M_SQL,
ProxyNodeMetricsStep::OneHour => LIST_PROXY_FLEET_METRICS_1H_SQL,
};
let mut rows = sqlx::query(sql)
.bind(i64::try_from(from_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(to_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.fetch(&self.pool);
let mut items = Vec::new();
while let Some(row) = rows.try_next().await.map_postgres_err()? {
items.push(Self::row_to_fleet_metric(&row)?);
}
Ok(items)
}
}
#[async_trait]
@@ -822,6 +1192,54 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository {
let Some(updated) = updated else {
return Ok(None);
};
let now_unix_secs = updated
.last_heartbeat_at_unix_secs
.unwrap_or_else(|| chrono::Utc::now().timestamp().max(0) as u64);
let tunnel_metrics_sample = build_tunnel_metrics_sample(
existing.proxy_metadata.as_ref(),
updated.proxy_metadata.as_ref(),
updated.active_connections,
updated.tunnel_connected,
);
if let Some(sample) = tunnel_metrics_sample.as_ref() {
self.upsert_metrics_bucket(
ProxyNodeMetricsStep::OneMinute,
&updated.id,
bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneMinute),
sample,
)
.await?;
self.upsert_metrics_bucket(
ProxyNodeMetricsStep::OneHour,
&updated.id,
bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneHour),
sample,
)
.await?;
for error in &sample.recent_error_events {
let detail = build_tunnel_error_event_detail(error);
let event_metadata = serde_json::json!({
"source": "heartbeat",
"category": error.category,
"message": error.message,
"timestamp_unix_secs": error.timestamp_unix_secs,
});
self.insert_event(
&updated.id,
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
Some(detail.as_str()),
Some(&event_metadata),
Some(if error.timestamp_unix_secs == 0 {
now_unix_secs
} else {
error.timestamp_unix_secs
}),
)
.await?;
}
}
if reconcile_remote_config_after_heartbeat(
updated.remote_config.as_ref(),
@@ -890,23 +1308,15 @@ impl ProxyNodeWriteRepository for SqlxProxyNodeRepository {
.zip(observed_at_unix_secs)
.is_some_and(|(last_transition, observed_at)| observed_at < last_transition)
{
sqlx::query(
r#"
INSERT INTO proxy_node_events (node_id, event_type, detail, created_at)
VALUES (
$1,
$2,
$3,
NOW()
)
"#,
)
.bind(&mutation.node_id)
.bind(event_type)
.bind(format!("[stale_ignored] {event_detail}"))
.execute(&mut *tx)
.await
.map_postgres_err()?;
sqlx::query(INSERT_PROXY_NODE_EVENT_SQL)
.bind(&mutation.node_id)
.bind(event_type)
.bind(format!("[stale_ignored] {event_detail}"))
.bind(None::<serde_json::Value>)
.bind(None::<f64>)
.execute(&mut *tx)
.await
.map_postgres_err()?;
tx.commit().await.map_err(postgres_error)?;
return self.find_proxy_node(&mutation.node_id).await;
}
@@ -942,27 +1352,15 @@ WHERE id = $1
.await
.map_postgres_err()?;
sqlx::query(
r#"
INSERT INTO proxy_node_events (node_id, event_type, detail, created_at)
VALUES (
$1,
$2,
$3,
CASE
WHEN $4::double precision IS NULL THEN NOW()
ELSE TO_TIMESTAMP($4::double precision)
END
)
"#,
)
.bind(&mutation.node_id)
.bind(event_type)
.bind(event_detail)
.bind(observed_at_unix_secs.map(|value| value as f64))
.execute(&mut *tx)
.await
.map_postgres_err()?;
sqlx::query(INSERT_PROXY_NODE_EVENT_SQL)
.bind(&mutation.node_id)
.bind(event_type)
.bind(event_detail)
.bind(None::<serde_json::Value>)
.bind(observed_at_unix_secs.map(|value| value as f64))
.execute(&mut *tx)
.await
.map_postgres_err()?;
tx.commit().await.map_err(postgres_error)?;
self.find_proxy_node(&mutation.node_id).await
@@ -1045,6 +1443,63 @@ VALUES (
.map_postgres_err()?;
Ok(())
}
async fn cleanup_proxy_node_metrics(
&self,
retain_1m_from_unix_secs: u64,
retain_1h_from_unix_secs: u64,
delete_limit: usize,
) -> Result<ProxyNodeMetricsCleanupSummary, DataLayerError> {
let delete_limit_i64 = i64::try_from(delete_limit.max(1)).unwrap_or(i64::MAX);
let deleted_1m = sqlx::query(
r#"
WITH expired AS (
SELECT node_id, bucket_start_unix_secs
FROM proxy_node_metrics_1m
WHERE bucket_start_unix_secs < $1
ORDER BY bucket_start_unix_secs ASC
LIMIT $2
)
DELETE FROM proxy_node_metrics_1m metrics
USING expired
WHERE metrics.node_id = expired.node_id
AND metrics.bucket_start_unix_secs = expired.bucket_start_unix_secs
"#,
)
.bind(i64::try_from(retain_1m_from_unix_secs).unwrap_or(i64::MAX))
.bind(delete_limit_i64)
.execute(&self.pool)
.await
.map_postgres_err()?
.rows_affected() as usize;
let deleted_1h = sqlx::query(
r#"
WITH expired AS (
SELECT node_id, bucket_start_unix_secs
FROM proxy_node_metrics_1h
WHERE bucket_start_unix_secs < $1
ORDER BY bucket_start_unix_secs ASC
LIMIT $2
)
DELETE FROM proxy_node_metrics_1h metrics
USING expired
WHERE metrics.node_id = expired.node_id
AND metrics.bucket_start_unix_secs = expired.bucket_start_unix_secs
"#,
)
.bind(i64::try_from(retain_1h_from_unix_secs).unwrap_or(i64::MAX))
.bind(delete_limit_i64)
.execute(&self.pool)
.await
.map_postgres_err()?
.rows_affected() as usize;
Ok(ProxyNodeMetricsCleanupSummary {
deleted_1m_rows: deleted_1m,
deleted_1h_rows: deleted_1h,
})
}
}
#[cfg(test)]

View File

@@ -2,10 +2,14 @@ use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row};
use super::types::{
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeHeartbeatMutation,
ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation, ProxyNodeReadRepository,
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
normalize_proxy_metadata, reconcile_remote_config_after_heartbeat, ProxyNodeEventQuery,
ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
ProxyNodeMetricsCleanupSummary, ProxyNodeMetricsStep, ProxyNodeReadRepository,
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyNode, StoredProxyNodeEvent,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository, StoredProxyFleetMetricsBucket,
StoredProxyNode, StoredProxyNodeEvent, StoredProxyNodeMetricsBucket,
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
};
use crate::driver::sqlite::SqlitePool;
use crate::error::SqlResultExt;
@@ -148,17 +152,22 @@ ON CONFLICT(id) DO UPDATE SET
node_id: &str,
event_type: &str,
detail: Option<&str>,
event_metadata: Option<&serde_json::Value>,
created_at_unix_secs: Option<u64>,
) -> Result<(), DataLayerError> {
sqlx::query(
r#"
INSERT INTO proxy_node_events (node_id, event_type, detail, created_at)
VALUES (?, ?, ?, ?)
INSERT INTO proxy_node_events (node_id, event_type, detail, event_metadata, created_at)
VALUES (?, ?, ?, ?, ?)
"#,
)
.bind(node_id)
.bind(event_type)
.bind(detail)
.bind(optional_json_to_string(
&event_metadata.cloned(),
"proxy_node_events.event_metadata",
)?)
.bind(created_at_unix_secs.unwrap_or_else(current_unix_secs) as i64)
.execute(&self.pool)
.await
@@ -166,6 +175,70 @@ VALUES (?, ?, ?, ?)
Ok(())
}
async fn upsert_metrics_bucket(
&self,
table: &str,
node_id: &str,
bucket_start: u64,
sample: &super::types::TunnelMetricsSample,
) -> Result<(), DataLayerError> {
sqlx::query(&format!(
r#"
INSERT INTO {table} (
node_id,
bucket_start_unix_secs,
samples,
uptime_samples,
active_connections_sum,
active_connections_max,
heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max,
connect_errors_delta,
disconnects_delta,
error_events_delta,
ws_in_bytes_delta,
ws_out_bytes_delta,
ws_in_frames_delta,
ws_out_frames_delta
)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(node_id, bucket_start_unix_secs) DO UPDATE SET
samples = {table}.samples + excluded.samples,
uptime_samples = {table}.uptime_samples + excluded.uptime_samples,
active_connections_sum = {table}.active_connections_sum + excluded.active_connections_sum,
active_connections_max = MAX({table}.active_connections_max, excluded.active_connections_max),
heartbeat_rtt_ms_sum = {table}.heartbeat_rtt_ms_sum + excluded.heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max = MAX({table}.heartbeat_rtt_ms_max, excluded.heartbeat_rtt_ms_max),
connect_errors_delta = {table}.connect_errors_delta + excluded.connect_errors_delta,
disconnects_delta = {table}.disconnects_delta + excluded.disconnects_delta,
error_events_delta = {table}.error_events_delta + excluded.error_events_delta,
ws_in_bytes_delta = {table}.ws_in_bytes_delta + excluded.ws_in_bytes_delta,
ws_out_bytes_delta = {table}.ws_out_bytes_delta + excluded.ws_out_bytes_delta,
ws_in_frames_delta = {table}.ws_in_frames_delta + excluded.ws_in_frames_delta,
ws_out_frames_delta = {table}.ws_out_frames_delta + excluded.ws_out_frames_delta
"#
))
.bind(node_id)
.bind(i64::try_from(bucket_start).unwrap_or(i64::MAX))
.bind(sample.samples)
.bind(sample.uptime_samples)
.bind(sample.active_connections_sum)
.bind(sample.active_connections_max)
.bind(sample.heartbeat_rtt_ms_sum)
.bind(sample.heartbeat_rtt_ms_max)
.bind(sample.connect_errors_delta)
.bind(sample.disconnects_delta)
.bind(sample.error_events_delta)
.bind(sample.ws_in_bytes_delta)
.bind(sample.ws_out_bytes_delta)
.bind(sample.ws_in_frames_delta)
.bind(sample.ws_out_frames_delta)
.execute(&self.pool)
.await
.map_sql_err()?;
Ok(())
}
fn normalize_remote_config(
mutation: &ProxyNodeRemoteConfigMutation,
existing: Option<&serde_json::Value>,
@@ -298,6 +371,7 @@ SELECT
node_id,
event_type,
detail,
event_metadata,
created_at AS created_at_unix_ms
FROM proxy_node_events
WHERE node_id = ?
@@ -312,6 +386,152 @@ LIMIT ?
.map_sql_err()?;
rows.iter().map(map_proxy_node_event_row).collect()
}
async fn list_proxy_node_events_filtered(
&self,
node_id: &str,
query: &ProxyNodeEventQuery,
) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> {
let rows = sqlx::query(
r#"
SELECT
id,
node_id,
event_type,
detail,
event_metadata,
created_at AS created_at_unix_ms
FROM proxy_node_events
WHERE node_id = ?
AND (? IS NULL OR created_at >= ?)
AND (? IS NULL OR created_at <= ?)
AND (? IS NULL OR LOWER(event_type) = LOWER(?))
ORDER BY created_at DESC, id DESC
LIMIT ?
"#,
)
.bind(node_id)
.bind(
query
.from_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(
query
.from_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(
query
.to_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(
query
.to_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(query.event_type.as_deref())
.bind(query.event_type.as_deref())
.bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_proxy_node_event_row).collect()
}
async fn list_proxy_node_metrics(
&self,
node_id: &str,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyNodeMetricsBucket>, DataLayerError> {
let table = match step {
ProxyNodeMetricsStep::OneMinute => "proxy_node_metrics_1m",
ProxyNodeMetricsStep::OneHour => "proxy_node_metrics_1h",
};
let rows = sqlx::query(&format!(
r#"
SELECT
node_id,
bucket_start_unix_secs,
samples,
uptime_samples,
active_connections_sum,
active_connections_max,
heartbeat_rtt_ms_sum,
heartbeat_rtt_ms_max,
connect_errors_delta,
disconnects_delta,
error_events_delta,
ws_in_bytes_delta,
ws_out_bytes_delta,
ws_in_frames_delta,
ws_out_frames_delta
FROM {table}
WHERE node_id = ?
AND bucket_start_unix_secs >= ?
AND bucket_start_unix_secs <= ?
ORDER BY bucket_start_unix_secs ASC
LIMIT ?
"#
))
.bind(node_id)
.bind(i64::try_from(from_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(to_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_proxy_node_metric_row).collect()
}
async fn list_proxy_fleet_metrics(
&self,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyFleetMetricsBucket>, DataLayerError> {
let table = match step {
ProxyNodeMetricsStep::OneMinute => "proxy_node_metrics_1m",
ProxyNodeMetricsStep::OneHour => "proxy_node_metrics_1h",
};
let rows = sqlx::query(&format!(
r#"
SELECT
bucket_start_unix_secs,
SUM(samples) AS samples,
SUM(uptime_samples) AS uptime_samples,
SUM(active_connections_sum) AS active_connections_sum,
MAX(active_connections_max) AS active_connections_max,
SUM(heartbeat_rtt_ms_sum) AS heartbeat_rtt_ms_sum,
MAX(heartbeat_rtt_ms_max) AS heartbeat_rtt_ms_max,
SUM(connect_errors_delta) AS connect_errors_delta,
SUM(disconnects_delta) AS disconnects_delta,
SUM(error_events_delta) AS error_events_delta,
SUM(ws_in_bytes_delta) AS ws_in_bytes_delta,
SUM(ws_out_bytes_delta) AS ws_out_bytes_delta,
SUM(ws_in_frames_delta) AS ws_in_frames_delta,
SUM(ws_out_frames_delta) AS ws_out_frames_delta
FROM {table}
WHERE bucket_start_unix_secs >= ?
AND bucket_start_unix_secs <= ?
GROUP BY bucket_start_unix_secs
ORDER BY bucket_start_unix_secs ASC
LIMIT ?
"#
))
.bind(i64::try_from(from_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(to_unix_secs).unwrap_or(i64::MAX))
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_proxy_fleet_metric_row).collect()
}
}
#[async_trait]
@@ -540,7 +760,9 @@ WHERE is_manual = 0
));
}
let now = Some(current_unix_secs());
let previous_proxy_metadata = node.proxy_metadata.clone();
let now_unix_secs = current_unix_secs();
let now = Some(now_unix_secs);
node.last_heartbeat_at_unix_secs = now;
if node.status != "online" || !node.tunnel_connected {
node.status = "online".to_string();
@@ -584,7 +806,55 @@ WHERE is_manual = 0
node.config_version = node.config_version.saturating_add(1);
node.updated_at_unix_secs = now;
}
let tunnel_metrics_sample = build_tunnel_metrics_sample(
previous_proxy_metadata.as_ref(),
node.proxy_metadata.as_ref(),
node.active_connections,
node.tunnel_connected,
);
self.upsert_node(&node).await?;
if let Some(sample) = tunnel_metrics_sample.as_ref() {
self.upsert_metrics_bucket(
"proxy_node_metrics_1m",
&node.id,
bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneMinute),
sample,
)
.await?;
self.upsert_metrics_bucket(
"proxy_node_metrics_1h",
&node.id,
bucket_start_unix_secs(now_unix_secs, ProxyNodeMetricsStep::OneHour),
sample,
)
.await?;
for error in &sample.recent_error_events {
let detail = build_tunnel_error_event_detail(error);
let event_metadata = serde_json::json!({
"source": "heartbeat",
"category": error.category,
"message": error.message,
"timestamp_unix_secs": error.timestamp_unix_secs,
});
self.insert_event(
&node.id,
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
Some(detail.as_str()),
Some(&event_metadata),
Some(if error.timestamp_unix_secs == 0 {
now_unix_secs
} else {
error.timestamp_unix_secs
}),
)
.await?;
}
}
Ok(Some(node))
}
@@ -638,6 +908,7 @@ WHERE is_manual = 0
&mutation.node_id,
event_type,
Some(&format!("[stale_ignored] {event_detail}")),
None,
Some(current_unix_secs()),
)
.await?;
@@ -660,6 +931,7 @@ WHERE is_manual = 0
&mutation.node_id,
event_type,
Some(&event_detail),
None,
Some(event_time),
)
.await?;
@@ -691,6 +963,16 @@ WHERE is_manual = 0
.execute(&self.pool)
.await
.map_sql_err()?;
sqlx::query("DELETE FROM proxy_node_metrics_1m WHERE node_id = ?")
.bind(node_id)
.execute(&self.pool)
.await
.map_sql_err()?;
sqlx::query("DELETE FROM proxy_node_metrics_1h WHERE node_id = ?")
.bind(node_id)
.execute(&self.pool)
.await
.map_sql_err()?;
sqlx::query("DELETE FROM proxy_nodes WHERE id = ?")
.bind(node_id)
.execute(&self.pool)
@@ -747,6 +1029,57 @@ WHERE is_manual = 0
node.updated_at_unix_secs = Some(current_unix_secs());
self.upsert_node(&node).await
}
async fn cleanup_proxy_node_metrics(
&self,
retain_1m_from_unix_secs: u64,
retain_1h_from_unix_secs: u64,
delete_limit: usize,
) -> Result<ProxyNodeMetricsCleanupSummary, DataLayerError> {
let delete_limit_i64 = i64::try_from(delete_limit.max(1)).unwrap_or(i64::MAX);
let deleted_1m = sqlx::query(
r#"
DELETE FROM proxy_node_metrics_1m
WHERE (node_id, bucket_start_unix_secs) IN (
SELECT node_id, bucket_start_unix_secs
FROM proxy_node_metrics_1m
WHERE bucket_start_unix_secs < ?
ORDER BY bucket_start_unix_secs ASC
LIMIT ?
)
"#,
)
.bind(i64::try_from(retain_1m_from_unix_secs).unwrap_or(i64::MAX))
.bind(delete_limit_i64)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected() as usize;
let deleted_1h = sqlx::query(
r#"
DELETE FROM proxy_node_metrics_1h
WHERE (node_id, bucket_start_unix_secs) IN (
SELECT node_id, bucket_start_unix_secs
FROM proxy_node_metrics_1h
WHERE bucket_start_unix_secs < ?
ORDER BY bucket_start_unix_secs ASC
LIMIT ?
)
"#,
)
.bind(i64::try_from(retain_1h_from_unix_secs).unwrap_or(i64::MAX))
.bind(delete_limit_i64)
.execute(&self.pool)
.await
.map_sql_err()?
.rows_affected() as usize;
Ok(ProxyNodeMetricsCleanupSummary {
deleted_1m_rows: deleted_1m,
deleted_1h_rows: deleted_1h,
})
}
}
fn optional_unix_secs(value: Option<i64>) -> Option<u64> {
@@ -861,18 +1194,73 @@ fn map_proxy_node_event_row(row: &SqliteRow) -> Result<StoredProxyNodeEvent, Dat
node_id: row.try_get("node_id").map_sql_err()?,
event_type: row.try_get("event_type").map_sql_err()?,
detail: row.try_get("detail").map_sql_err()?,
event_metadata: optional_json_from_string(
row.try_get("event_metadata").map_sql_err()?,
"proxy_node_events.event_metadata",
)?,
created_at_unix_ms: optional_unix_secs(row.try_get("created_at_unix_ms").map_sql_err()?),
})
}
fn map_proxy_node_metric_row(
row: &SqliteRow,
) -> Result<StoredProxyNodeMetricsBucket, DataLayerError> {
Ok(StoredProxyNodeMetricsBucket {
node_id: row.try_get("node_id").map_sql_err()?,
bucket_start_unix_secs: optional_unix_secs(
row.try_get("bucket_start_unix_secs").map_sql_err()?,
)
.unwrap_or_default(),
samples: row.try_get("samples").map_sql_err()?,
uptime_samples: row.try_get("uptime_samples").map_sql_err()?,
active_connections_sum: row.try_get("active_connections_sum").map_sql_err()?,
active_connections_max: row.try_get("active_connections_max").map_sql_err()?,
heartbeat_rtt_ms_sum: row.try_get("heartbeat_rtt_ms_sum").map_sql_err()?,
heartbeat_rtt_ms_max: row.try_get("heartbeat_rtt_ms_max").map_sql_err()?,
connect_errors_delta: row.try_get("connect_errors_delta").map_sql_err()?,
disconnects_delta: row.try_get("disconnects_delta").map_sql_err()?,
error_events_delta: row.try_get("error_events_delta").map_sql_err()?,
ws_in_bytes_delta: row.try_get("ws_in_bytes_delta").map_sql_err()?,
ws_out_bytes_delta: row.try_get("ws_out_bytes_delta").map_sql_err()?,
ws_in_frames_delta: row.try_get("ws_in_frames_delta").map_sql_err()?,
ws_out_frames_delta: row.try_get("ws_out_frames_delta").map_sql_err()?,
})
}
fn map_proxy_fleet_metric_row(
row: &SqliteRow,
) -> Result<StoredProxyFleetMetricsBucket, DataLayerError> {
Ok(StoredProxyFleetMetricsBucket {
bucket_start_unix_secs: optional_unix_secs(
row.try_get("bucket_start_unix_secs").map_sql_err()?,
)
.unwrap_or_default(),
samples: row.try_get("samples").map_sql_err()?,
uptime_samples: row.try_get("uptime_samples").map_sql_err()?,
active_connections_sum: row.try_get("active_connections_sum").map_sql_err()?,
active_connections_max: row.try_get("active_connections_max").map_sql_err()?,
heartbeat_rtt_ms_sum: row.try_get("heartbeat_rtt_ms_sum").map_sql_err()?,
heartbeat_rtt_ms_max: row.try_get("heartbeat_rtt_ms_max").map_sql_err()?,
connect_errors_delta: row.try_get("connect_errors_delta").map_sql_err()?,
disconnects_delta: row.try_get("disconnects_delta").map_sql_err()?,
error_events_delta: row.try_get("error_events_delta").map_sql_err()?,
ws_in_bytes_delta: row.try_get("ws_in_bytes_delta").map_sql_err()?,
ws_out_bytes_delta: row.try_get("ws_out_bytes_delta").map_sql_err()?,
ws_in_frames_delta: row.try_get("ws_in_frames_delta").map_sql_err()?,
ws_out_frames_delta: row.try_get("ws_out_frames_delta").map_sql_err()?,
})
}
#[cfg(test)]
mod tests {
use super::SqliteProxyNodeReadRepository;
use crate::lifecycle::migrate::run_sqlite_migrations;
use crate::repository::proxy_nodes::{
ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation, ProxyNodeManualUpdateMutation,
ProxyNodeReadRepository, ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation,
ProxyNodeTrafficMutation, ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository,
ProxyNodeEventQuery, ProxyNodeHeartbeatMutation, ProxyNodeManualCreateMutation,
ProxyNodeManualUpdateMutation, ProxyNodeMetricsStep, ProxyNodeReadRepository,
ProxyNodeRegistrationMutation, ProxyNodeRemoteConfigMutation, ProxyNodeTrafficMutation,
ProxyNodeTunnelStatusMutation, ProxyNodeWriteRepository,
PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR,
};
use serde_json::json;
@@ -1142,4 +1530,131 @@ VALUES ('node-1', 'registered', 'ok', 3)
.expect("manual node should delete")
.is_some());
}
#[tokio::test]
async fn sqlite_repository_aggregates_proxy_node_metrics_and_filters_events() {
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.expect("sqlite pool should connect");
run_sqlite_migrations(&pool)
.await
.expect("sqlite migrations should run");
let repository = SqliteProxyNodeReadRepository::new(pool);
let registered = repository
.register_node(&ProxyNodeRegistrationMutation {
name: "tunnel-1".to_string(),
ip: "10.0.0.1".to_string(),
port: 7000,
region: None,
heartbeat_interval: 30,
active_connections: Some(0),
total_requests: Some(0),
avg_latency_ms: None,
hardware_info: None,
estimated_max_concurrency: None,
proxy_metadata: None,
proxy_version: Some("1.0.0".to_string()),
registered_by: None,
tunnel_mode: true,
})
.await
.expect("node should register");
let now = super::current_unix_secs();
repository
.apply_heartbeat(&ProxyNodeHeartbeatMutation {
node_id: registered.id.clone(),
heartbeat_interval: Some(30),
active_connections: Some(5),
total_requests_delta: None,
avg_latency_ms: None,
failed_requests_delta: None,
dns_failures_delta: None,
stream_errors_delta: None,
proxy_metadata: Some(json!({
"tunnel_metrics": {
"connect_errors": 4,
"disconnects": 1,
"error_events_total": 1,
"ws_in_bytes": 100,
"ws_out_bytes": 200,
"ws_in_frames": 3,
"ws_out_frames": 6,
"heartbeat_rtt_last_ms": 33
},
"recent_tunnel_errors": [{
"timestamp_unix_secs": now,
"category": "tcp_connect_timeout",
"message": "timeout"
}]
})),
proxy_version: Some("1.0.0".to_string()),
})
.await
.expect("heartbeat should apply")
.expect("node should exist");
let metrics = repository
.list_proxy_node_metrics(
&registered.id,
ProxyNodeMetricsStep::OneMinute,
now.saturating_sub(120),
now.saturating_add(120),
10,
)
.await
.expect("metrics should list");
assert_eq!(metrics.len(), 1);
assert_eq!(metrics[0].samples, 1);
assert_eq!(metrics[0].uptime_samples, 1);
assert_eq!(metrics[0].active_connections_max, 5);
assert_eq!(metrics[0].heartbeat_rtt_ms_sum, 33);
assert_eq!(metrics[0].connect_errors_delta, 4);
assert_eq!(metrics[0].ws_out_frames_delta, 6);
let fleet = repository
.list_proxy_fleet_metrics(
ProxyNodeMetricsStep::OneMinute,
now.saturating_sub(120),
now.saturating_add(120),
10,
)
.await
.expect("fleet metrics should list");
assert_eq!(fleet.len(), 1);
assert_eq!(fleet[0].samples, 1);
assert_eq!(fleet[0].error_events_delta, 1);
let events = repository
.list_proxy_node_events_filtered(
&registered.id,
&ProxyNodeEventQuery {
limit: 10,
from_unix_secs: Some(now.saturating_sub(120)),
to_unix_secs: Some(now.saturating_add(120)),
event_type: Some(PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR.to_string()),
},
)
.await
.expect("events should list");
assert_eq!(events.len(), 1);
assert_eq!(events[0].event_type, PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR);
assert_eq!(
events[0]
.event_metadata
.as_ref()
.and_then(|value| value.get("category"))
.and_then(serde_json::Value::as_str),
Some("tcp_connect_timeout")
);
let cleanup = repository
.cleanup_proxy_node_metrics(now.saturating_add(1), now.saturating_add(1), 10)
.await
.expect("cleanup should run");
assert_eq!(cleanup.deleted_1m_rows, 1);
assert_eq!(cleanup.deleted_1h_rows, 1);
}
}

View File

@@ -1,4 +1,5 @@
use async_trait::async_trait;
use serde_json::Value;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProxyNode {
@@ -239,9 +240,186 @@ pub struct StoredProxyNodeEvent {
pub node_id: String,
pub event_type: String,
pub detail: Option<String>,
pub event_metadata: Option<serde_json::Value>,
pub created_at_unix_ms: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct ProxyNodeEventQuery {
pub limit: usize,
pub from_unix_secs: Option<u64>,
pub to_unix_secs: Option<u64>,
pub event_type: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub enum ProxyNodeMetricsStep {
OneMinute,
OneHour,
}
impl ProxyNodeMetricsStep {
pub fn bucket_size_secs(self) -> u64 {
match self {
Self::OneMinute => 60,
Self::OneHour => 3_600,
}
}
pub fn as_api_value(self) -> &'static str {
match self {
Self::OneMinute => "1m",
Self::OneHour => "1h",
}
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProxyNodeMetricsBucket {
pub node_id: String,
pub bucket_start_unix_secs: u64,
pub samples: i64,
pub uptime_samples: i64,
pub active_connections_sum: i64,
pub active_connections_max: i64,
pub heartbeat_rtt_ms_sum: i64,
pub heartbeat_rtt_ms_max: i64,
pub connect_errors_delta: i64,
pub disconnects_delta: i64,
pub error_events_delta: i64,
pub ws_in_bytes_delta: i64,
pub ws_out_bytes_delta: i64,
pub ws_in_frames_delta: i64,
pub ws_out_frames_delta: i64,
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
pub struct StoredProxyFleetMetricsBucket {
pub bucket_start_unix_secs: u64,
pub samples: i64,
pub uptime_samples: i64,
pub active_connections_sum: i64,
pub active_connections_max: i64,
pub heartbeat_rtt_ms_sum: i64,
pub heartbeat_rtt_ms_max: i64,
pub connect_errors_delta: i64,
pub disconnects_delta: i64,
pub error_events_delta: i64,
pub ws_in_bytes_delta: i64,
pub ws_out_bytes_delta: i64,
pub ws_in_frames_delta: i64,
pub ws_out_frames_delta: i64,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
pub struct ProxyNodeMetricsCleanupSummary {
pub deleted_1m_rows: usize,
pub deleted_1h_rows: usize,
}
pub const PROXY_NODE_EVENT_TYPE_TUNNEL_ERROR: &str = "tunnel_err";
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct TunnelErrorEventRecord {
pub timestamp_unix_secs: u64,
pub category: String,
pub message: String,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct TunnelMetricsCounters {
pub connect_errors: u64,
pub disconnects: u64,
pub error_events_total: u64,
pub ws_in_bytes: u64,
pub ws_out_bytes: u64,
pub ws_in_frames: u64,
pub ws_out_frames: u64,
pub heartbeat_rtt_last_ms: u64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TunnelMetricsSample {
pub samples: i64,
pub uptime_samples: i64,
pub active_connections_sum: i64,
pub active_connections_max: i64,
pub heartbeat_rtt_ms_sum: i64,
pub heartbeat_rtt_ms_max: i64,
pub connect_errors_delta: i64,
pub disconnects_delta: i64,
pub error_events_delta: i64,
pub ws_in_bytes_delta: i64,
pub ws_out_bytes_delta: i64,
pub ws_in_frames_delta: i64,
pub ws_out_frames_delta: i64,
pub recent_error_events: Vec<TunnelErrorEventRecord>,
}
pub fn bucket_start_unix_secs(timestamp_unix_secs: u64, step: ProxyNodeMetricsStep) -> u64 {
let size = step.bucket_size_secs();
timestamp_unix_secs / size * size
}
pub fn build_tunnel_metrics_sample(
previous_proxy_metadata: Option<&Value>,
current_proxy_metadata: Option<&Value>,
active_connections: i32,
tunnel_connected: bool,
) -> Option<TunnelMetricsSample> {
let current = extract_tunnel_metrics_counters(current_proxy_metadata)?;
let previous = extract_tunnel_metrics_counters(previous_proxy_metadata);
let current_recent_errors = extract_recent_tunnel_errors(current_proxy_metadata);
let connect_errors_delta =
counter_delta_u64(previous.map(|v| v.connect_errors), current.connect_errors);
let disconnects_delta = counter_delta_u64(previous.map(|v| v.disconnects), current.disconnects);
let error_events_delta = counter_delta_u64(
previous.map(|v| v.error_events_total),
current.error_events_total,
);
let ws_in_bytes_delta = counter_delta_u64(previous.map(|v| v.ws_in_bytes), current.ws_in_bytes);
let ws_out_bytes_delta =
counter_delta_u64(previous.map(|v| v.ws_out_bytes), current.ws_out_bytes);
let ws_in_frames_delta =
counter_delta_u64(previous.map(|v| v.ws_in_frames), current.ws_in_frames);
let ws_out_frames_delta =
counter_delta_u64(previous.map(|v| v.ws_out_frames), current.ws_out_frames);
let take_recent = usize::try_from(error_events_delta).unwrap_or(usize::MAX);
let recent_error_events = if take_recent == 0 {
Vec::new()
} else {
let capture = take_recent.min(current_recent_errors.len());
let from = current_recent_errors.len().saturating_sub(capture);
current_recent_errors[from..].to_vec()
};
let active_connections = i64::from(active_connections.max(0));
let heartbeat_rtt_last_ms = i64::try_from(current.heartbeat_rtt_last_ms).unwrap_or(i64::MAX);
Some(TunnelMetricsSample {
samples: 1,
uptime_samples: if tunnel_connected { 1 } else { 0 },
active_connections_sum: active_connections,
active_connections_max: active_connections,
heartbeat_rtt_ms_sum: heartbeat_rtt_last_ms,
heartbeat_rtt_ms_max: heartbeat_rtt_last_ms,
connect_errors_delta: i64::try_from(connect_errors_delta).unwrap_or(i64::MAX),
disconnects_delta: i64::try_from(disconnects_delta).unwrap_or(i64::MAX),
error_events_delta: i64::try_from(error_events_delta).unwrap_or(i64::MAX),
ws_in_bytes_delta: i64::try_from(ws_in_bytes_delta).unwrap_or(i64::MAX),
ws_out_bytes_delta: i64::try_from(ws_out_bytes_delta).unwrap_or(i64::MAX),
ws_in_frames_delta: i64::try_from(ws_in_frames_delta).unwrap_or(i64::MAX),
ws_out_frames_delta: i64::try_from(ws_out_frames_delta).unwrap_or(i64::MAX),
recent_error_events,
})
}
pub fn build_tunnel_error_event_detail(event: &TunnelErrorEventRecord) -> String {
format!("[{}] {}", event.category, event.message)
}
pub fn normalize_proxy_metadata(
proxy_metadata: Option<&serde_json::Value>,
proxy_version: Option<&str>,
@@ -276,6 +454,71 @@ pub fn normalize_proxy_metadata(
}
}
fn extract_tunnel_metrics_counters(
proxy_metadata: Option<&Value>,
) -> Option<TunnelMetricsCounters> {
let tunnel_metrics = proxy_metadata
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("tunnel_metrics"))
.and_then(Value::as_object)?;
Some(TunnelMetricsCounters {
connect_errors: json_u64(tunnel_metrics.get("connect_errors")).unwrap_or(0),
disconnects: json_u64(tunnel_metrics.get("disconnects")).unwrap_or(0),
error_events_total: json_u64(tunnel_metrics.get("error_events_total")).unwrap_or(0),
ws_in_bytes: json_u64(tunnel_metrics.get("ws_in_bytes")).unwrap_or(0),
ws_out_bytes: json_u64(tunnel_metrics.get("ws_out_bytes")).unwrap_or(0),
ws_in_frames: json_u64(tunnel_metrics.get("ws_in_frames")).unwrap_or(0),
ws_out_frames: json_u64(tunnel_metrics.get("ws_out_frames")).unwrap_or(0),
heartbeat_rtt_last_ms: json_u64(tunnel_metrics.get("heartbeat_rtt_last_ms")).unwrap_or(0),
})
}
fn extract_recent_tunnel_errors(proxy_metadata: Option<&Value>) -> Vec<TunnelErrorEventRecord> {
proxy_metadata
.and_then(Value::as_object)
.and_then(|metadata| metadata.get("recent_tunnel_errors"))
.and_then(Value::as_array)
.map(|items| {
items
.iter()
.filter_map(|item| {
let item = item.as_object()?;
Some(TunnelErrorEventRecord {
timestamp_unix_secs: json_u64(item.get("timestamp_unix_secs"))
.unwrap_or_default(),
category: item
.get("category")
.and_then(Value::as_str)
.unwrap_or("unknown")
.to_string(),
message: item
.get("message")
.and_then(Value::as_str)
.unwrap_or("n/a")
.to_string(),
})
})
.collect::<Vec<_>>()
})
.unwrap_or_default()
}
fn json_u64(value: Option<&Value>) -> Option<u64> {
value.and_then(|value| {
value
.as_u64()
.or_else(|| value.as_i64().and_then(|n| (n >= 0).then_some(n as u64)))
})
}
fn counter_delta_u64(previous: Option<u64>, current: u64) -> u64 {
match previous {
Some(previous) if current >= previous => current - previous,
Some(_) | None => current,
}
}
fn normalize_proxy_version_label(value: &str) -> Option<String> {
let trimmed = value.trim();
if trimmed.is_empty() {
@@ -375,6 +618,42 @@ pub trait ProxyNodeReadRepository: Send + Sync {
node_id: &str,
limit: usize,
) -> Result<Vec<StoredProxyNodeEvent>, crate::DataLayerError>;
async fn list_proxy_node_events_filtered(
&self,
node_id: &str,
query: &ProxyNodeEventQuery,
) -> Result<Vec<StoredProxyNodeEvent>, crate::DataLayerError> {
let mut items = self.list_proxy_node_events(node_id, query.limit).await?;
if let Some(from_unix_secs) = query.from_unix_secs {
items.retain(|item| item.created_at_unix_ms.unwrap_or(0) >= from_unix_secs);
}
if let Some(to_unix_secs) = query.to_unix_secs {
items.retain(|item| item.created_at_unix_ms.unwrap_or(u64::MAX) <= to_unix_secs);
}
if let Some(event_type) = query.event_type.as_deref() {
items.retain(|item| item.event_type.eq_ignore_ascii_case(event_type));
}
items.truncate(query.limit);
Ok(items)
}
async fn list_proxy_node_metrics(
&self,
node_id: &str,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyNodeMetricsBucket>, crate::DataLayerError>;
async fn list_proxy_fleet_metrics(
&self,
step: ProxyNodeMetricsStep,
from_unix_secs: u64,
to_unix_secs: u64,
limit: usize,
) -> Result<Vec<StoredProxyFleetMetricsBucket>, crate::DataLayerError>;
}
#[async_trait]
@@ -433,6 +712,13 @@ pub trait ProxyNodeWriteRepository: Send + Sync {
failed_delta: i64,
latency_ms: Option<i64>,
) -> Result<(), crate::DataLayerError>;
async fn cleanup_proxy_node_metrics(
&self,
retain_1m_from_unix_secs: u64,
retain_1h_from_unix_secs: u64,
delete_limit: usize,
) -> Result<ProxyNodeMetricsCleanupSummary, crate::DataLayerError>;
}
#[cfg(test)]
@@ -440,9 +726,10 @@ mod tests {
use serde_json::json;
use super::{
normalize_proxy_node_scheduling_state, proxy_node_accepts_new_tunnels,
proxy_reported_version, reconcile_remote_config_after_heartbeat,
remote_config_scheduling_state, remote_config_upgrade_target, StoredProxyNode,
bucket_start_unix_secs, build_tunnel_metrics_sample, normalize_proxy_node_scheduling_state,
proxy_node_accepts_new_tunnels, proxy_reported_version,
reconcile_remote_config_after_heartbeat, remote_config_scheduling_state,
remote_config_upgrade_target, ProxyNodeMetricsStep, StoredProxyNode,
};
#[test]
@@ -527,4 +814,63 @@ mod tests {
assert!(!proxy_node_accepts_new_tunnels(&node));
}
#[test]
fn builds_tunnel_metrics_sample_with_reset_safe_counter_deltas() {
let previous = json!({
"tunnel_metrics": {
"connect_errors": 10,
"disconnects": 5,
"error_events_total": 7,
"ws_in_bytes": 1_000,
"ws_out_bytes": 2_000,
"ws_in_frames": 10,
"ws_out_frames": 20,
"heartbeat_rtt_last_ms": 30
}
});
let current = json!({
"tunnel_metrics": {
"connect_errors": 12,
"disconnects": 2,
"error_events_total": 9,
"ws_in_bytes": 1_500,
"ws_out_bytes": 100,
"ws_in_frames": 11,
"ws_out_frames": 3,
"heartbeat_rtt_last_ms": 44
},
"recent_tunnel_errors": [
{"timestamp_unix_secs": 100, "category": "older", "message": "old"},
{"timestamp_unix_secs": 101, "category": "newer", "message": "new"}
]
});
let sample = build_tunnel_metrics_sample(Some(&previous), Some(&current), 4, true)
.expect("sample should build");
assert_eq!(sample.samples, 1);
assert_eq!(sample.uptime_samples, 1);
assert_eq!(sample.active_connections_sum, 4);
assert_eq!(sample.heartbeat_rtt_ms_sum, 44);
assert_eq!(sample.connect_errors_delta, 2);
assert_eq!(sample.disconnects_delta, 2);
assert_eq!(sample.error_events_delta, 2);
assert_eq!(sample.ws_in_bytes_delta, 500);
assert_eq!(sample.ws_out_bytes_delta, 100);
assert_eq!(sample.ws_out_frames_delta, 3);
assert_eq!(sample.recent_error_events.len(), 2);
assert_eq!(sample.recent_error_events[0].category, "older");
}
#[test]
fn maps_timestamps_to_metric_buckets() {
assert_eq!(
bucket_start_unix_secs(1_710_000_119, ProxyNodeMetricsStep::OneMinute),
1_710_000_060
);
assert_eq!(
bucket_start_unix_secs(1_710_003_999, ProxyNodeMetricsStep::OneHour),
1_710_003_600
);
}
}

View File

@@ -4,7 +4,7 @@ use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
use super::{
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
};
use crate::driver::sqlite::SqlitePool;
use crate::driver::sqlite::{sqlite_optional_real, sqlite_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::DataLayerError;
@@ -12,8 +12,8 @@ const QUOTA_COLUMNS: &str = r#"
SELECT
id AS provider_id,
billing_type,
monthly_quota_usd,
COALESCE(monthly_used_usd, 0) AS monthly_used_usd,
CAST(monthly_quota_usd AS REAL) AS monthly_quota_usd,
CAST(COALESCE(monthly_used_usd, 0) AS REAL) AS monthly_used_usd,
quota_reset_day,
quota_last_reset_at AS quota_last_reset_at_unix_secs,
quota_expires_at AS quota_expires_at_unix_secs,
@@ -77,7 +77,7 @@ impl ProviderQuotaWriteRepository for SqliteProviderQuotaRepository {
let rows_affected = sqlx::query(
r#"
UPDATE providers
SET monthly_used_usd = 0,
SET monthly_used_usd = 0.0,
quota_last_reset_at = ?,
updated_at = ?
WHERE billing_type = 'monthly_quota'
@@ -103,8 +103,8 @@ fn map_row(row: &SqliteRow) -> Result<StoredProviderQuotaSnapshot, DataLayerErro
StoredProviderQuotaSnapshot::new(
row.try_get("provider_id").map_sql_err()?,
row.try_get("billing_type").map_sql_err()?,
row.try_get("monthly_quota_usd").map_sql_err()?,
row.try_get("monthly_used_usd").map_sql_err()?,
sqlite_optional_real(row, "monthly_quota_usd")?,
sqlite_real(row, "monthly_used_usd")?,
row.try_get("quota_reset_day").map_sql_err()?,
row.try_get("quota_last_reset_at_unix_secs").map_sql_err()?,
row.try_get("quota_expires_at_unix_secs").map_sql_err()?,
@@ -138,6 +138,13 @@ mod tests {
.expect("quota should exist");
assert_eq!(quota.monthly_used_usd, 5.0);
let quota = repository
.find_by_provider_id("provider-null-used")
.await
.expect("quota with null usage should load")
.expect("quota with null usage should exist");
assert_eq!(quota.monthly_used_usd, 0.0);
let quotas = repository
.find_by_provider_ids(&["provider-2".to_string(), "provider-1".to_string()])
.await
@@ -173,7 +180,8 @@ INSERT INTO providers (
)
VALUES
('provider-1', 'Provider One', 'openai', 'monthly_quota', 20.0, 5.0, 7, 1000, 1, 1, 1),
('provider-2', 'Provider Two', 'openai', 'payg', NULL, 1.5, NULL, NULL, 1, 1, 1)
('provider-2', 'Provider Two', 'openai', 'payg', NULL, 1.5, NULL, NULL, 1, 1, 1),
('provider-null-used', 'Provider Null Used', 'openai', 'payg', NULL, NULL, NULL, NULL, 1, 1, 1)
"#,
)
.execute(pool)

View File

@@ -2,7 +2,7 @@ use async_trait::async_trait;
use sqlx::{sqlite::SqliteRow, Row};
use super::{SettlementWriteRepository, StoredUsageSettlement, UsageSettlementInput};
use crate::driver::sqlite::SqlitePool;
use crate::driver::sqlite::{sqlite_optional_real, sqlite_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::DataLayerError;
@@ -35,7 +35,7 @@ SELECT
usage_settlement_snapshots.wallet_gift_balance_after,
usage_record.wallet_gift_balance_after
) AS wallet_gift_balance_after,
usage_settlement_snapshots.provider_monthly_used_usd AS provider_monthly_used_usd,
CAST(usage_settlement_snapshots.provider_monthly_used_usd AS REAL) AS provider_monthly_used_usd,
usage_record.provider_id,
COALESCE(usage_settlement_snapshots.finalized_at, usage_record.finalized_at) AS finalized_at_unix_secs
FROM "usage" AS usage_record
@@ -120,17 +120,16 @@ fn settlement_from_row(row: &SqliteRow) -> Result<StoredUsageSettlement, DataLay
request_id: row.try_get("request_id").map_sql_err()?,
wallet_id: row.try_get("wallet_id").map_sql_err()?,
billing_status: row.try_get("billing_status").map_sql_err()?,
wallet_balance_before: row.try_get("wallet_balance_before").map_sql_err()?,
wallet_balance_after: row.try_get("wallet_balance_after").map_sql_err()?,
wallet_recharge_balance_before: row
.try_get("wallet_recharge_balance_before")
.map_sql_err()?,
wallet_recharge_balance_after: row
.try_get("wallet_recharge_balance_after")
.map_sql_err()?,
wallet_gift_balance_before: row.try_get("wallet_gift_balance_before").map_sql_err()?,
wallet_gift_balance_after: row.try_get("wallet_gift_balance_after").map_sql_err()?,
provider_monthly_used_usd: row.try_get("provider_monthly_used_usd").map_sql_err()?,
wallet_balance_before: sqlite_optional_real(row, "wallet_balance_before")?,
wallet_balance_after: sqlite_optional_real(row, "wallet_balance_after")?,
wallet_recharge_balance_before: sqlite_optional_real(
row,
"wallet_recharge_balance_before",
)?,
wallet_recharge_balance_after: sqlite_optional_real(row, "wallet_recharge_balance_after")?,
wallet_gift_balance_before: sqlite_optional_real(row, "wallet_gift_balance_before")?,
wallet_gift_balance_after: sqlite_optional_real(row, "wallet_gift_balance_after")?,
provider_monthly_used_usd: sqlite_optional_real(row, "provider_monthly_used_usd")?,
finalized_at_unix_secs: row
.try_get::<Option<i64>, _>("finalized_at_unix_secs")
.map_sql_err()?
@@ -268,8 +267,8 @@ LIMIT 1
if let Some(wallet_row) = wallet_row {
let wallet_id: String = wallet_row.try_get("id").map_sql_err()?;
let before_recharge: f64 = wallet_row.try_get("balance").map_sql_err()?;
let before_gift: f64 = wallet_row.try_get("gift_balance").map_sql_err()?;
let before_recharge = sqlite_real(&wallet_row, "balance")?;
let before_gift = sqlite_real(&wallet_row, "gift_balance")?;
let limit_mode: String = wallet_row.try_get("limit_mode").map_sql_err()?;
let before_total = before_recharge + before_gift;
let mut after_recharge = before_recharge;
@@ -318,7 +317,7 @@ WHERE id = ?
r#"
UPDATE providers
SET
monthly_used_usd = COALESCE(monthly_used_usd, 0) + ?,
monthly_used_usd = CAST(COALESCE(monthly_used_usd, 0) AS REAL) + ?,
updated_at = ?
WHERE id = ?
"#,
@@ -330,14 +329,15 @@ WHERE id = ?
.await
.map_sql_err()?;
settlement.provider_monthly_used_usd = sqlx::query_scalar::<_, Option<f64>>(
"SELECT monthly_used_usd FROM providers WHERE id = ? LIMIT 1",
settlement.provider_monthly_used_usd = sqlx::query(
"SELECT CAST(monthly_used_usd AS REAL) AS monthly_used_usd FROM providers WHERE id = ? LIMIT 1",
)
.bind(provider_id)
.fetch_optional(&mut *tx)
.await
.map_sql_err()?
.flatten();
.map(|row| sqlite_real(&row, "monthly_used_usd"))
.transpose()?;
}
}

View File

@@ -9,7 +9,7 @@ use super::{
InMemoryUsageReadRepository, PendingUsageCleanupSummary, StoredRequestUsageAudit,
UpsertUsageRecord, UsageWriteRepository,
};
use crate::driver::sqlite::SqlitePool;
use crate::driver::sqlite::{sqlite_optional_real, sqlite_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::DataLayerError;
@@ -42,11 +42,11 @@ SELECT
cache_creation_ephemeral_5m_input_tokens,
cache_creation_ephemeral_1h_input_tokens,
cache_read_input_tokens,
cache_creation_cost_usd,
cache_read_cost_usd,
output_price_per_1m,
total_cost_usd,
actual_total_cost_usd,
CAST(cache_creation_cost_usd AS REAL) AS cache_creation_cost_usd,
CAST(cache_read_cost_usd AS REAL) AS cache_read_cost_usd,
CAST(output_price_per_1m AS REAL) AS output_price_per_1m,
CAST(total_cost_usd AS REAL) AS total_cost_usd,
CAST(actual_total_cost_usd AS REAL) AS actual_total_cost_usd,
status_code,
error_message,
error_category,
@@ -291,7 +291,7 @@ impl UsageWriteRepository for SqliteUsageWriteRepository {
UPDATE api_keys
SET total_requests = 0,
total_tokens = 0,
total_cost_usd = 0,
total_cost_usd = 0.0,
last_used_at = NULL
"#,
)
@@ -305,7 +305,7 @@ SELECT
api_key_id,
COUNT(*) AS total_requests,
COALESCE(SUM(total_tokens), 0) AS total_tokens,
COALESCE(SUM(total_cost_usd), 0) AS total_cost_usd,
CAST(COALESCE(SUM(total_cost_usd), 0) AS REAL) AS total_cost_usd,
MAX(updated_at_unix_secs) AS last_used_at
FROM "usage"
WHERE api_key_id IS NOT NULL AND api_key_id <> ''
@@ -329,7 +329,7 @@ WHERE id = ?
)
.bind(row.try_get::<i64, _>("total_requests").map_sql_err()?)
.bind(row.try_get::<i64, _>("total_tokens").map_sql_err()?)
.bind(row.try_get::<f64, _>("total_cost_usd").map_sql_err()?)
.bind(sqlite_real(row, "total_cost_usd")?)
.bind(
row.try_get::<Option<i64>, _>("last_used_at")
.map_sql_err()?,
@@ -351,7 +351,7 @@ SET request_count = 0,
success_count = 0,
error_count = 0,
total_tokens = 0,
total_cost_usd = 0,
total_cost_usd = 0.0,
total_response_time_ms = 0,
last_used_at = NULL
"#,
@@ -368,7 +368,7 @@ SELECT
status_code,
error_message,
total_tokens,
total_cost_usd,
CAST(total_cost_usd AS REAL) AS total_cost_usd,
response_time_ms,
updated_at_unix_secs
FROM "usage"
@@ -396,7 +396,7 @@ WHERE provider_api_key_id IS NOT NULL AND provider_api_key_id <> ''
entry.error_count += 1;
}
entry.total_tokens += row.try_get::<i64, _>("total_tokens").map_sql_err()?;
entry.total_cost_usd += row.try_get::<f64, _>("total_cost_usd").map_sql_err()?;
entry.total_cost_usd += sqlite_real(&row, "total_cost_usd")?;
entry.total_response_time_ms += row
.try_get::<Option<i64>, _>("response_time_ms")
.map_sql_err()?
@@ -528,8 +528,8 @@ SET status = 'failed',
error_message = ?,
billing_status = 'void',
finalized_at = ?,
total_cost_usd = 0,
actual_total_cost_usd = 0
total_cost_usd = 0.0,
actual_total_cost_usd = 0.0
WHERE request_id = ?
"#,
)
@@ -818,8 +818,8 @@ fn map_usage_row(row: &SqliteRow) -> Result<StoredRequestUsageAudit, DataLayerEr
row_i32(row, "input_tokens")?,
row_i32(row, "output_tokens")?,
row_i32(row, "total_tokens")?,
row.try_get("total_cost_usd").map_sql_err()?,
row.try_get("actual_total_cost_usd").map_sql_err()?,
sqlite_real(row, "total_cost_usd")?,
sqlite_real(row, "actual_total_cost_usd")?,
row_optional_i32(row, "status_code")?,
row.try_get("error_message").map_sql_err()?,
row.try_get("error_category").map_sql_err()?,
@@ -837,9 +837,10 @@ fn map_usage_row(row: &SqliteRow) -> Result<StoredRequestUsageAudit, DataLayerEr
audit.cache_creation_ephemeral_1h_input_tokens =
row_u64(row, "cache_creation_ephemeral_1h_input_tokens")?;
audit.cache_read_input_tokens = row_u64(row, "cache_read_input_tokens")?;
audit.cache_creation_cost_usd = row.try_get("cache_creation_cost_usd").map_sql_err()?;
audit.cache_read_cost_usd = row.try_get("cache_read_cost_usd").map_sql_err()?;
audit.output_price_per_1m = row.try_get("output_price_per_1m").map_sql_err()?;
audit.cache_creation_cost_usd =
sqlite_optional_real(row, "cache_creation_cost_usd")?.unwrap_or(0.0);
audit.cache_read_cost_usd = sqlite_optional_real(row, "cache_read_cost_usd")?.unwrap_or(0.0);
audit.output_price_per_1m = sqlite_optional_real(row, "output_price_per_1m")?;
audit.request_metadata = row
.try_get::<Option<String>, _>("request_metadata")
.map_sql_err()?

View File

@@ -27,7 +27,7 @@ use super::{
StoredWalletDailyUsageLedgerPage, StoredWalletSnapshot, WalletLookupKey, WalletMutationOutcome,
WalletReadRepository, WalletWriteRepository,
};
use crate::driver::sqlite::SqlitePool;
use crate::driver::sqlite::{sqlite_optional_real, sqlite_real, SqlitePool};
use crate::error::SqlResultExt;
use crate::DataLayerError;
@@ -715,7 +715,7 @@ LIMIT 1
tx.commit().await.map_sql_err()?;
return Ok(CreateWalletRefundRequestOutcome::WalletMissing);
};
let wallet_recharge_balance: f64 = get(&wallet_row, "balance")?;
let wallet_recharge_balance = sqlite_real(&wallet_row, "balance")?;
let wallet_reserved_amount: f64 = sqlx::query_scalar(
r#"
SELECT COALESCE(SUM(amount_usd), 0.0)
@@ -779,7 +779,7 @@ WHERE payment_order_id = ?
.fetch_one(&mut *tx)
.await
.map_sql_err()?;
let refundable_amount: f64 = get(&order_row, "refundable_amount_usd")?;
let refundable_amount = sqlite_real(&order_row, "refundable_amount_usd")?;
if input.amount_usd > (refundable_amount - order_reserved_amount) {
tx.commit().await.map_sql_err()?;
return Ok(
@@ -946,7 +946,7 @@ VALUES (?, NULL, ?, ?, ?, ?, ?, ?, 'received', ?, NULL, ?, NULL)
let order_no: String = get(&order_row, "order_no")?;
let order_wallet_id: String = get(&order_row, "wallet_id")?;
let order_payment_method: String = get(&order_row, "payment_method")?;
let order_amount_usd: f64 = get(&order_row, "amount_usd")?;
let order_amount_usd = sqlite_real(&order_row, "amount_usd")?;
let order_status: String = get(&order_row, "status")?;
let expires_at_unix_secs: Option<i64> = get(&order_row, "expires_at_unix_secs")?;
@@ -1070,8 +1070,8 @@ LIMIT 1
});
}
let before_recharge: f64 = get(&wallet_row, "balance")?;
let before_gift: f64 = get(&wallet_row, "gift_balance")?;
let before_recharge = sqlite_real(&wallet_row, "balance")?;
let before_gift = sqlite_real(&wallet_row, "gift_balance")?;
let before_total = before_recharge + before_gift;
let after_recharge = before_recharge + order_amount_usd;
let after_total = after_recharge + before_gift;
@@ -1174,8 +1174,8 @@ WHERE id = ?
return Ok(None);
};
let before_recharge: f64 = get(&row, "balance")?;
let before_gift: f64 = get(&row, "gift_balance")?;
let before_recharge = sqlite_real(&row, "balance")?;
let before_gift = sqlite_real(&row, "gift_balance")?;
let before_total = before_recharge + before_gift;
let mut after_recharge = before_recharge;
let mut after_gift = before_gift;
@@ -1278,8 +1278,8 @@ VALUES (?, ?, 'adjust', 'adjust_admin', ?, ?, ?, ?, ?, ?, ?, 'admin_action', ?,
return Ok(None);
};
let before_recharge: f64 = get(&wallet_row, "balance")?;
let before_gift: f64 = get(&wallet_row, "gift_balance")?;
let before_recharge = sqlite_real(&wallet_row, "balance")?;
let before_gift = sqlite_real(&wallet_row, "gift_balance")?;
let user_id: Option<String> = get(&wallet_row, "user_id")?;
let order_id = uuid::Uuid::new_v4().to_string();
let gateway_response = json_string(
@@ -1416,8 +1416,8 @@ VALUES (?, ?, 'recharge', ?, ?, ?, ?, ?, ?, ?, ?, 'payment_order', ?, ?, ?, ?)
"wallet not found".to_string(),
));
};
let before_recharge: f64 = get(&wallet_row, "balance")?;
let before_gift: f64 = get(&wallet_row, "gift_balance")?;
let before_recharge = sqlite_real(&wallet_row, "balance")?;
let before_gift = sqlite_real(&wallet_row, "gift_balance")?;
let before_total = before_recharge + before_gift;
let amount_usd = refund.amount_usd;
let after_recharge = before_recharge - amount_usd;
@@ -1437,7 +1437,7 @@ VALUES (?, ?, 'recharge', ?, ?, ?, ?, ?, ?, ?, ?, 'payment_order', ?, ?, ?, ?)
"payment order not found".to_string(),
));
};
let refundable_amount: f64 = get(&order_row, "refundable_amount_usd")?;
let refundable_amount = sqlite_real(&order_row, "refundable_amount_usd")?;
if amount_usd > refundable_amount {
tx.commit().await.map_sql_err()?;
return Ok(WalletMutationOutcome::Invalid(
@@ -1665,8 +1665,8 @@ WHERE id = ? AND wallet_id = ?
));
};
let amount_usd = refund.amount_usd;
let before_recharge: f64 = get(&wallet_row, "balance")?;
let before_gift: f64 = get(&wallet_row, "gift_balance")?;
let before_recharge = sqlite_real(&wallet_row, "balance")?;
let before_gift = sqlite_real(&wallet_row, "gift_balance")?;
let before_total = before_recharge + before_gift;
let after_recharge = before_recharge + amount_usd;
@@ -1919,8 +1919,8 @@ WHERE id = ? AND wallet_id = ?
));
}
let before_recharge: f64 = get(&wallet_row, "balance")?;
let before_gift: f64 = get(&wallet_row, "gift_balance")?;
let before_recharge = sqlite_real(&wallet_row, "balance")?;
let before_gift = sqlite_real(&wallet_row, "gift_balance")?;
let before_total = before_recharge + before_gift;
let after_recharge = before_recharge + order.amount_usd;
sqlx::query(
@@ -2358,7 +2358,7 @@ LIMIT 1
let batch_id: String = get(&code_row, "batch_id")?;
let batch_name: String = get(&code_row, "batch_name")?;
let balance_bucket: String = get(&code_row, "balance_bucket")?;
let amount_usd: f64 = get(&code_row, "amount_usd")?;
let amount_usd = sqlite_real(&code_row, "amount_usd")?;
let credits_recharge_balance = redeem_code_credits_recharge_balance(&balance_bucket);
let wallet_row = sqlite_wallet_by_user_id(&mut tx, &input.user_id).await?;
@@ -2374,7 +2374,10 @@ LIMIT 1
};
let (before_recharge, before_gift) = if let Some(row) = wallet_row.as_ref() {
(get(row, "balance")?, get(row, "gift_balance")?)
(
sqlite_real(row, "balance")?,
sqlite_real(row, "gift_balance")?,
)
} else {
sqlx::query(
r#"
@@ -2383,7 +2386,7 @@ INSERT INTO wallets (
total_recharged, total_consumed, total_refunded, total_adjusted,
created_at, updated_at
)
VALUES (?, ?, 0, 0, 'finite', 'USD', 'active', 0, 0, 0, 0, ?, ?)
VALUES (?, ?, 0.0, 0.0, 'finite', 'USD', 'active', 0.0, 0.0, 0.0, 0.0, ?, ?)
"#,
)
.bind(&wallet_id)
@@ -2557,15 +2560,15 @@ fn map_wallet_row(row: &SqliteRow) -> Result<StoredWalletSnapshot, DataLayerErro
get(row, "id")?,
get(row, "user_id")?,
get(row, "api_key_id")?,
get(row, "balance")?,
get(row, "gift_balance")?,
sqlite_real(row, "balance")?,
sqlite_real(row, "gift_balance")?,
get(row, "limit_mode")?,
get(row, "currency")?,
get(row, "status")?,
get(row, "total_recharged")?,
get(row, "total_consumed")?,
get(row, "total_refunded")?,
get(row, "total_adjusted")?,
sqlite_real(row, "total_recharged")?,
sqlite_real(row, "total_consumed")?,
sqlite_real(row, "total_refunded")?,
sqlite_real(row, "total_adjusted")?,
get(row, "updated_at_unix_secs")?,
)
}
@@ -3159,12 +3162,12 @@ fn map_payment_order_row(row: &SqliteRow) -> Result<StoredAdminPaymentOrder, Dat
order_no: get(row, "order_no")?,
wallet_id: get(row, "wallet_id")?,
user_id: get(row, "user_id")?,
amount_usd: get(row, "amount_usd")?,
pay_amount: get(row, "pay_amount")?,
amount_usd: sqlite_real(row, "amount_usd")?,
pay_amount: sqlite_optional_real(row, "pay_amount")?,
pay_currency: get(row, "pay_currency")?,
exchange_rate: get(row, "exchange_rate")?,
refunded_amount_usd: get(row, "refunded_amount_usd")?,
refundable_amount_usd: get(row, "refundable_amount_usd")?,
exchange_rate: sqlite_optional_real(row, "exchange_rate")?,
refunded_amount_usd: sqlite_real(row, "refunded_amount_usd")?,
refundable_amount_usd: sqlite_real(row, "refundable_amount_usd")?,
payment_method: get(row, "payment_method")?,
gateway_order_id: get(row, "gateway_order_id")?,
gateway_response: optional_json(
@@ -3223,13 +3226,13 @@ fn map_wallet_transaction_row(
wallet_id: get(row, "wallet_id")?,
category: get(row, "category")?,
reason_code: get(row, "reason_code")?,
amount: get(row, "amount")?,
balance_before: get(row, "balance_before")?,
balance_after: get(row, "balance_after")?,
recharge_balance_before: get(row, "recharge_balance_before")?,
recharge_balance_after: get(row, "recharge_balance_after")?,
gift_balance_before: get(row, "gift_balance_before")?,
gift_balance_after: get(row, "gift_balance_after")?,
amount: sqlite_real(row, "amount")?,
balance_before: sqlite_real(row, "balance_before")?,
balance_after: sqlite_real(row, "balance_after")?,
recharge_balance_before: sqlite_real(row, "recharge_balance_before")?,
recharge_balance_after: sqlite_real(row, "recharge_balance_after")?,
gift_balance_before: sqlite_real(row, "gift_balance_before")?,
gift_balance_after: sqlite_real(row, "gift_balance_after")?,
link_type: get(row, "link_type")?,
link_id: get(row, "link_id")?,
operator_id: get(row, "operator_id")?,
@@ -3253,7 +3256,7 @@ fn map_refund_row(row: &SqliteRow) -> Result<StoredAdminWalletRefund, DataLayerE
source_type: get(row, "source_type")?,
source_id: get(row, "source_id")?,
refund_mode: get(row, "refund_mode")?,
amount_usd: get(row, "amount_usd")?,
amount_usd: sqlite_real(row, "amount_usd")?,
status: get(row, "status")?,
reason: get(row, "reason")?,
failure_reason: get(row, "failure_reason")?,
@@ -3287,7 +3290,7 @@ fn map_redeem_batch_row(row: &SqliteRow) -> Result<StoredAdminRedeemCodeBatch, D
Ok(StoredAdminRedeemCodeBatch {
id: get(row, "id")?,
name: get(row, "name")?,
amount_usd: get(row, "amount_usd")?,
amount_usd: sqlite_real(row, "amount_usd")?,
currency: get(row, "currency")?,
balance_bucket: get(row, "balance_bucket")?,
total_count: nonnegative_u64(get(row, "total_count")?, "redeem_code_batches.total_count")?,
@@ -3352,7 +3355,7 @@ fn map_daily_usage_row(row: &SqliteRow) -> Result<StoredWalletDailyUsageLedger,
id: get(row, "id")?,
billing_date: get(row, "billing_date")?,
billing_timezone: get(row, "billing_timezone")?,
total_cost_usd: get(row, "total_cost_usd")?,
total_cost_usd: sqlite_real(row, "total_cost_usd")?,
total_requests: nonnegative_u64(
get(row, "total_requests")?,
"wallet_daily_usage_ledgers.total_requests",