mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 09:20:22 +08:00
Merge remote-tracking branch 'upstream/main'
# Conflicts: # apps/aether-gateway/src/handlers/admin/provider/endpoints_admin/payloads.rs # apps/aether-gateway/src/handlers/admin/provider/endpoints_admin/reads.rs # apps/aether-gateway/src/handlers/admin/provider/endpoints_admin/update.rs # apps/aether-gateway/src/tests/control/admin/endpoints/routes.rs # frontend/src/features/models/components/GlobalModelFormDialog.vue # frontend/src/features/providers/components/ProviderModelFormDialog.vue # frontend/src/features/providers/components/provider-tabs/__tests__/model-test-request.spec.ts # frontend/src/features/providers/components/provider-tabs/model-test-request.ts
This commit is contained in:
@@ -1,5 +1,6 @@
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use futures_util::TryStreamExt;
|
||||
use serde_json::Value;
|
||||
use sqlx::{Column, Row, TypeInfo, ValueRef};
|
||||
|
||||
@@ -130,6 +131,41 @@ pub struct ExportRow {
|
||||
pub payload: Value,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
|
||||
pub struct DataCopyOptions {
|
||||
pub omit_request_body_details: bool,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct SqliteCopyColumn {
|
||||
name: String,
|
||||
declared_type: String,
|
||||
not_null: bool,
|
||||
has_default: bool,
|
||||
primary_key_position: i64,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum SqliteCopyAffinity {
|
||||
Integer,
|
||||
Real,
|
||||
Text,
|
||||
Blob,
|
||||
Numeric,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct SchemaCopyColumn {
|
||||
sqlite: SqliteCopyColumn,
|
||||
postgres: PostgresImportColumn,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct SchemaCopyTable {
|
||||
table_name: String,
|
||||
columns: Vec<SchemaCopyColumn>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct PostgresImportColumn {
|
||||
data_type: String,
|
||||
@@ -141,6 +177,20 @@ struct PostgresImportColumn {
|
||||
type PostgresImportColumns = BTreeMap<String, PostgresImportColumn>;
|
||||
type ImportColumnNames = BTreeSet<String>;
|
||||
|
||||
const USAGE_REQUEST_BODY_DETAIL_COLUMNS: &[&str] = &[
|
||||
"request_body",
|
||||
"response_body",
|
||||
"provider_request_body",
|
||||
"client_response_body",
|
||||
"request_body_compressed",
|
||||
"response_body_compressed",
|
||||
"provider_request_body_compressed",
|
||||
"client_response_body_compressed",
|
||||
];
|
||||
|
||||
const REQUEST_BODY_DETAIL_TABLES: &[&str] = &["usage_body_blobs", "usage_http_audits"];
|
||||
const LIFECYCLE_TABLES: &[&str] = &["_sqlx_migrations", "schema_backfills"];
|
||||
|
||||
pub fn encode_jsonl(records: &[DataExportRecord]) -> Result<String, DataLayerError> {
|
||||
validate_export_records(records)?;
|
||||
|
||||
@@ -344,6 +394,602 @@ pub async fn import_database_jsonl(
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn copy_database_records(
|
||||
source: SqlDatabaseConfig,
|
||||
target: SqlDatabaseConfig,
|
||||
domains: Vec<ExportDomain>,
|
||||
created_at_unix_secs: u64,
|
||||
options: DataCopyOptions,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
if domains.is_empty()
|
||||
&& source.driver == DatabaseDriver::Postgres
|
||||
&& target.driver == DatabaseDriver::Sqlite
|
||||
{
|
||||
return copy_postgres_to_sqlite_from_target_schema(source, target, options).await;
|
||||
}
|
||||
|
||||
let mut records =
|
||||
decode_jsonl(&export_database_jsonl(source, domains, created_at_unix_secs).await?)?;
|
||||
if options.omit_request_body_details {
|
||||
omit_request_body_details_from_records(&mut records);
|
||||
}
|
||||
import_database_jsonl(target, &encode_jsonl(&records)?).await
|
||||
}
|
||||
|
||||
fn omit_request_body_details_from_records(records: &mut [DataExportRecord]) {
|
||||
for record in records {
|
||||
let DataExportRecord::Row {
|
||||
domain: ExportDomain::Usage,
|
||||
payload,
|
||||
..
|
||||
} = record
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
|
||||
if let Some(object) = payload.as_object_mut() {
|
||||
for column_name in USAGE_REQUEST_BODY_DETAIL_COLUMNS {
|
||||
object.remove(*column_name);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn copy_postgres_to_sqlite_from_target_schema(
|
||||
source: SqlDatabaseConfig,
|
||||
mut target: SqlDatabaseConfig,
|
||||
options: DataCopyOptions,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
target.pool.min_connections = 1;
|
||||
target.pool.max_connections = 1;
|
||||
|
||||
let postgres_pool =
|
||||
crate::driver::postgres::PostgresPoolFactory::new(source.to_postgres_config()?)?
|
||||
.connect_lazy()?;
|
||||
let sqlite_pool = crate::driver::sqlite::SqlitePoolFactory::new(target)?.connect_lazy()?;
|
||||
|
||||
let source_tables = load_postgres_public_table_names(&postgres_pool).await?;
|
||||
let target_tables = load_sqlite_copy_table_names(&sqlite_pool).await?;
|
||||
|
||||
ensure_no_nonempty_source_tables_outside_target_schema(
|
||||
&postgres_pool,
|
||||
&source_tables,
|
||||
&target_tables,
|
||||
options,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let mut imported = 0usize;
|
||||
sqlx::raw_sql("PRAGMA foreign_keys = OFF")
|
||||
.execute(&sqlite_pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
for table_name in target_tables {
|
||||
if copy_table_is_lifecycle(&table_name)
|
||||
|| copy_table_is_sqlite_internal(&table_name)
|
||||
|| !source_tables.contains(&table_name)
|
||||
|| (options.omit_request_body_details && copy_table_is_request_body_detail(&table_name))
|
||||
{
|
||||
continue;
|
||||
}
|
||||
|
||||
let table_plan = build_postgres_sqlite_copy_table_plan(
|
||||
&postgres_pool,
|
||||
&sqlite_pool,
|
||||
&table_name,
|
||||
options,
|
||||
)
|
||||
.await?;
|
||||
if table_plan.columns.is_empty() {
|
||||
continue;
|
||||
}
|
||||
imported = imported.saturating_add(
|
||||
copy_postgres_sqlite_table(&postgres_pool, &sqlite_pool, &table_plan).await?,
|
||||
);
|
||||
}
|
||||
|
||||
sqlx::raw_sql("PRAGMA foreign_keys = ON")
|
||||
.execute(&sqlite_pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
ensure_sqlite_foreign_key_check_passes(&sqlite_pool).await?;
|
||||
Ok(imported)
|
||||
}
|
||||
|
||||
async fn ensure_no_nonempty_source_tables_outside_target_schema(
|
||||
postgres_pool: &crate::driver::postgres::PostgresPool,
|
||||
source_tables: &BTreeSet<String>,
|
||||
target_tables: &BTreeSet<String>,
|
||||
options: DataCopyOptions,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let mut missing = Vec::new();
|
||||
for table_name in source_tables {
|
||||
if copy_table_is_lifecycle(table_name)
|
||||
|| (options.omit_request_body_details && copy_table_is_request_body_detail(table_name))
|
||||
|| target_tables.contains(table_name)
|
||||
{
|
||||
continue;
|
||||
}
|
||||
if postgres_public_table_has_rows(postgres_pool, table_name).await? {
|
||||
missing.push(table_name.clone());
|
||||
}
|
||||
}
|
||||
|
||||
if !missing.is_empty() {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"source Postgres has non-empty public tables that do not exist in the target SQLite schema: {}",
|
||||
missing.join(", ")
|
||||
)));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn build_postgres_sqlite_copy_table_plan(
|
||||
postgres_pool: &crate::driver::postgres::PostgresPool,
|
||||
sqlite_pool: &crate::driver::sqlite::SqlitePool,
|
||||
table_name: &str,
|
||||
options: DataCopyOptions,
|
||||
) -> Result<SchemaCopyTable, DataLayerError> {
|
||||
let sqlite_columns = load_sqlite_copy_columns(sqlite_pool, table_name).await?;
|
||||
let postgres_columns =
|
||||
load_postgres_import_columns(postgres_pool, &format!("public.{table_name}")).await?;
|
||||
let source_has_rows = postgres_public_table_has_rows(postgres_pool, table_name).await?;
|
||||
let mut columns = Vec::new();
|
||||
|
||||
for sqlite_column in sqlite_columns {
|
||||
if options.omit_request_body_details
|
||||
&& table_name == "usage"
|
||||
&& USAGE_REQUEST_BODY_DETAIL_COLUMNS.contains(&sqlite_column.name.as_str())
|
||||
{
|
||||
continue;
|
||||
}
|
||||
|
||||
if let Some(postgres_column) = postgres_columns.get(&sqlite_column.name) {
|
||||
columns.push(SchemaCopyColumn {
|
||||
sqlite: sqlite_column,
|
||||
postgres: postgres_column.clone(),
|
||||
});
|
||||
continue;
|
||||
}
|
||||
|
||||
if source_has_rows && sqlite_copy_column_is_required(&sqlite_column) {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"target SQLite table '{table_name}' has required column '{}' that does not exist in source Postgres",
|
||||
sqlite_column.name
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
if source_has_rows && columns.is_empty() {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"source Postgres table '{table_name}' has rows, but none of its columns exist in target SQLite"
|
||||
)));
|
||||
}
|
||||
|
||||
Ok(SchemaCopyTable {
|
||||
table_name: table_name.to_string(),
|
||||
columns,
|
||||
})
|
||||
}
|
||||
|
||||
async fn copy_postgres_sqlite_table(
|
||||
postgres_pool: &crate::driver::postgres::PostgresPool,
|
||||
sqlite_pool: &crate::driver::sqlite::SqlitePool,
|
||||
table: &SchemaCopyTable,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let source_sql = postgres_schema_copy_select_sql(table)?;
|
||||
let target_sql = sqlite_schema_copy_insert_sql(table)?;
|
||||
let mut rows = sqlx::query(&source_sql).fetch(postgres_pool);
|
||||
let mut imported = 0usize;
|
||||
|
||||
while let Some(row) = rows.try_next().await.map_sql_err()? {
|
||||
let payload = row.try_get::<Value, _>("payload").map_sql_err()?;
|
||||
let object = payload.as_object().ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"postgres copy row for table '{}' did not produce a JSON object",
|
||||
table.table_name
|
||||
))
|
||||
})?;
|
||||
let mut query = sqlx::query(&target_sql);
|
||||
for column in &table.columns {
|
||||
let value = object.get(&column.sqlite.name).ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"postgres copy row for table '{}' is missing column '{}'",
|
||||
table.table_name, column.sqlite.name
|
||||
))
|
||||
})?;
|
||||
query = bind_sqlite_copy_value(query, value, &column.sqlite)?;
|
||||
}
|
||||
query.execute(sqlite_pool).await.map_sql_err()?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
|
||||
Ok(imported)
|
||||
}
|
||||
|
||||
fn postgres_schema_copy_select_sql(table: &SchemaCopyTable) -> Result<String, DataLayerError> {
|
||||
let table_sql = format!(
|
||||
"public.{}",
|
||||
postgres_quote_identifier(table.table_name.as_str())?
|
||||
);
|
||||
let mut payload_parts = Vec::new();
|
||||
for column in &table.columns {
|
||||
if let Some(expr) = postgres_schema_copy_override_expr(column)? {
|
||||
payload_parts.push(sql_string_literal(&column.sqlite.name));
|
||||
payload_parts.push(expr);
|
||||
}
|
||||
}
|
||||
let payload_sql = if payload_parts.is_empty() {
|
||||
"to_jsonb(t)".to_string()
|
||||
} else {
|
||||
format!(
|
||||
"to_jsonb(t) || jsonb_build_object({})",
|
||||
payload_parts.join(", ")
|
||||
)
|
||||
};
|
||||
|
||||
let order_by = table
|
||||
.columns
|
||||
.iter()
|
||||
.filter(|column| column.sqlite.primary_key_position > 0)
|
||||
.map(|column| {
|
||||
postgres_quote_identifier(&column.sqlite.name).map(|quoted| format!("t.{quoted} ASC"))
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
let order_sql = if order_by.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
format!(" ORDER BY {}", order_by.join(", "))
|
||||
};
|
||||
|
||||
Ok(format!(
|
||||
"SELECT {payload_sql} AS payload FROM {table_sql} AS t{order_sql}"
|
||||
))
|
||||
}
|
||||
|
||||
fn postgres_schema_copy_override_expr(
|
||||
column: &SchemaCopyColumn,
|
||||
) -> Result<Option<String>, DataLayerError> {
|
||||
let column_sql = format!("t.{}", postgres_quote_identifier(&column.sqlite.name)?);
|
||||
let affinity = sqlite_copy_affinity(&column.sqlite);
|
||||
|
||||
if affinity == SqliteCopyAffinity::Blob && is_postgres_bytea_column(&column.postgres) {
|
||||
return Ok(Some(format!(
|
||||
"CASE WHEN {column_sql} IS NULL THEN NULL ELSE encode({column_sql}, 'hex') END"
|
||||
)));
|
||||
}
|
||||
|
||||
if affinity == SqliteCopyAffinity::Integer && is_postgres_boolean_column(&column.postgres) {
|
||||
return Ok(Some(format!(
|
||||
"CASE WHEN {column_sql} IS NULL THEN NULL WHEN {column_sql} THEN 1 ELSE 0 END"
|
||||
)));
|
||||
}
|
||||
|
||||
if affinity == SqliteCopyAffinity::Integer
|
||||
&& (is_postgres_timestamp_column(&column.postgres)
|
||||
|| is_postgres_date_column(&column.postgres))
|
||||
{
|
||||
let timestamp_sql = if is_postgres_date_column(&column.postgres) {
|
||||
format!("{column_sql}::timestamp")
|
||||
} else {
|
||||
column_sql.clone()
|
||||
};
|
||||
let multiplier = if sqlite_copy_column_stores_unix_millis(&column.sqlite.name) {
|
||||
" * 1000"
|
||||
} else {
|
||||
""
|
||||
};
|
||||
return Ok(Some(format!(
|
||||
"CASE WHEN {column_sql} IS NULL THEN NULL ELSE FLOOR(EXTRACT(EPOCH FROM {timestamp_sql}){multiplier})::bigint END"
|
||||
)));
|
||||
}
|
||||
|
||||
Ok(None)
|
||||
}
|
||||
|
||||
fn sqlite_schema_copy_insert_sql(table: &SchemaCopyTable) -> Result<String, DataLayerError> {
|
||||
let table_sql = sqlite_quote_identifier(&table.table_name)?;
|
||||
let column_sql = table
|
||||
.columns
|
||||
.iter()
|
||||
.map(|column| sqlite_quote_identifier(&column.sqlite.name))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let placeholder_sql = vec!["?"; table.columns.len()].join(", ");
|
||||
Ok(format!(
|
||||
"INSERT OR REPLACE INTO {table_sql} ({column_sql}) VALUES ({placeholder_sql})"
|
||||
))
|
||||
}
|
||||
|
||||
async fn load_postgres_public_table_names(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
) -> Result<BTreeSet<String>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT table_name
|
||||
FROM information_schema.tables
|
||||
WHERE table_schema = 'public'
|
||||
AND table_type = 'BASE TABLE'
|
||||
ORDER BY table_name
|
||||
"#,
|
||||
)
|
||||
.fetch_all(pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut tables = BTreeSet::new();
|
||||
for row in rows {
|
||||
tables.insert(row.try_get::<String, _>("table_name").map_sql_err()?);
|
||||
}
|
||||
Ok(tables)
|
||||
}
|
||||
|
||||
async fn load_sqlite_copy_table_names(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
) -> Result<BTreeSet<String>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT name
|
||||
FROM sqlite_schema
|
||||
WHERE type = 'table'
|
||||
AND name NOT LIKE 'sqlite_%'
|
||||
ORDER BY name
|
||||
"#,
|
||||
)
|
||||
.fetch_all(pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut tables = BTreeSet::new();
|
||||
for row in rows {
|
||||
let table_name = row.try_get::<String, _>("name").map_sql_err()?;
|
||||
if !copy_table_is_lifecycle(&table_name) && !copy_table_is_sqlite_internal(&table_name) {
|
||||
tables.insert(table_name);
|
||||
}
|
||||
}
|
||||
Ok(tables)
|
||||
}
|
||||
|
||||
async fn load_sqlite_copy_columns(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
table_name: &str,
|
||||
) -> Result<Vec<SqliteCopyColumn>, DataLayerError> {
|
||||
let table_sql = sqlite_quote_identifier(table_name)?;
|
||||
let rows = sqlx::query(&format!("PRAGMA table_info({table_sql})"))
|
||||
.fetch_all(pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut columns = Vec::new();
|
||||
for row in rows {
|
||||
columns.push(SqliteCopyColumn {
|
||||
name: row.try_get::<String, _>("name").map_sql_err()?,
|
||||
declared_type: row
|
||||
.try_get::<Option<String>, _>("type")
|
||||
.map_sql_err()?
|
||||
.unwrap_or_default(),
|
||||
not_null: row.try_get::<i64, _>("notnull").map_sql_err()? != 0,
|
||||
has_default: row
|
||||
.try_get::<Option<String>, _>("dflt_value")
|
||||
.map_sql_err()?
|
||||
.is_some(),
|
||||
primary_key_position: row.try_get::<i64, _>("pk").map_sql_err()?,
|
||||
});
|
||||
}
|
||||
|
||||
if columns.is_empty() {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"target SQLite table '{table_name}' has no visible columns"
|
||||
)));
|
||||
}
|
||||
Ok(columns)
|
||||
}
|
||||
|
||||
async fn postgres_public_table_has_rows(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
table_name: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let table_sql = format!("public.{}", postgres_quote_identifier(table_name)?);
|
||||
sqlx::query_scalar::<_, bool>(&format!(
|
||||
"SELECT EXISTS (SELECT 1 FROM {table_sql} LIMIT 1)"
|
||||
))
|
||||
.fetch_one(pool)
|
||||
.await
|
||||
.map_sql_err()
|
||||
}
|
||||
|
||||
async fn ensure_sqlite_foreign_key_check_passes(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let rows = sqlx::query("PRAGMA foreign_key_check")
|
||||
.fetch_all(pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if rows.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let mut violations = Vec::new();
|
||||
for row in rows.iter().take(10) {
|
||||
let table = row
|
||||
.try_get::<Option<String>, _>("table")
|
||||
.map_sql_err()?
|
||||
.unwrap_or_else(|| "<unknown>".to_string());
|
||||
let rowid = row.try_get::<Option<i64>, _>("rowid").map_sql_err()?;
|
||||
let parent = row
|
||||
.try_get::<Option<String>, _>("parent")
|
||||
.map_sql_err()?
|
||||
.unwrap_or_else(|| "<unknown>".to_string());
|
||||
violations.push(format!("{table} rowid={rowid:?} parent={parent}"));
|
||||
}
|
||||
Err(DataLayerError::InvalidInput(format!(
|
||||
"target SQLite foreign key check failed after copy: {}",
|
||||
violations.join("; ")
|
||||
)))
|
||||
}
|
||||
|
||||
fn copy_table_is_lifecycle(table_name: &str) -> bool {
|
||||
LIFECYCLE_TABLES.contains(&table_name)
|
||||
}
|
||||
|
||||
fn copy_table_is_sqlite_internal(table_name: &str) -> bool {
|
||||
table_name.starts_with("sqlite_")
|
||||
}
|
||||
|
||||
fn copy_table_is_request_body_detail(table_name: &str) -> bool {
|
||||
REQUEST_BODY_DETAIL_TABLES.contains(&table_name)
|
||||
}
|
||||
|
||||
fn sqlite_copy_column_is_required(column: &SqliteCopyColumn) -> bool {
|
||||
(column.not_null || column.primary_key_position > 0) && !column.has_default
|
||||
}
|
||||
|
||||
fn sqlite_copy_column_stores_unix_millis(column_name: &str) -> bool {
|
||||
column_name.ends_with("_unix_ms")
|
||||
}
|
||||
|
||||
fn sqlite_copy_affinity(column: &SqliteCopyColumn) -> SqliteCopyAffinity {
|
||||
let declared_type = column.declared_type.to_ascii_uppercase();
|
||||
if declared_type.contains("INT") {
|
||||
SqliteCopyAffinity::Integer
|
||||
} else if declared_type.contains("CHAR")
|
||||
|| declared_type.contains("CLOB")
|
||||
|| declared_type.contains("TEXT")
|
||||
{
|
||||
SqliteCopyAffinity::Text
|
||||
} else if declared_type.contains("BLOB") || declared_type.trim().is_empty() {
|
||||
SqliteCopyAffinity::Blob
|
||||
} else if declared_type.contains("REAL")
|
||||
|| declared_type.contains("FLOA")
|
||||
|| declared_type.contains("DOUB")
|
||||
{
|
||||
SqliteCopyAffinity::Real
|
||||
} else {
|
||||
SqliteCopyAffinity::Numeric
|
||||
}
|
||||
}
|
||||
|
||||
fn is_postgres_bytea_column(column: &PostgresImportColumn) -> bool {
|
||||
column.data_type == "bytea" || column.udt_name == "bytea"
|
||||
}
|
||||
|
||||
fn is_postgres_date_column(column: &PostgresImportColumn) -> bool {
|
||||
column.data_type == "date" || column.udt_name == "date"
|
||||
}
|
||||
|
||||
fn bind_sqlite_copy_value<'q>(
|
||||
query: sqlx::query::Query<'q, sqlx::Sqlite, sqlx::sqlite::SqliteArguments<'q>>,
|
||||
value: &'q Value,
|
||||
column: &SqliteCopyColumn,
|
||||
) -> Result<sqlx::query::Query<'q, sqlx::Sqlite, sqlx::sqlite::SqliteArguments<'q>>, DataLayerError>
|
||||
{
|
||||
Ok(match sqlite_copy_affinity(column) {
|
||||
SqliteCopyAffinity::Integer => match value {
|
||||
Value::Null => query.bind(Option::<i64>::None),
|
||||
Value::Bool(value) => query.bind(i64::from(*value)),
|
||||
Value::Number(number) => {
|
||||
let value = number
|
||||
.as_i64()
|
||||
.or_else(|| number.as_u64().and_then(|value| i64::try_from(value).ok()))
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{}' expected integer, got {number}",
|
||||
column.name
|
||||
))
|
||||
})?;
|
||||
query.bind(value)
|
||||
}
|
||||
Value::String(value) => query.bind(value.parse::<i64>().map_err(|err| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{}' expected integer string: {err}",
|
||||
column.name
|
||||
))
|
||||
})?),
|
||||
Value::Array(_) | Value::Object(_) => {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{}' expected integer-compatible value",
|
||||
column.name
|
||||
)));
|
||||
}
|
||||
},
|
||||
SqliteCopyAffinity::Real => match value {
|
||||
Value::Null => query.bind(Option::<f64>::None),
|
||||
Value::Number(number) => query.bind(number.as_f64().ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{}' expected finite real value",
|
||||
column.name
|
||||
))
|
||||
})?),
|
||||
Value::String(value) => query.bind(value.parse::<f64>().map_err(|err| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{}' expected real string: {err}",
|
||||
column.name
|
||||
))
|
||||
})?),
|
||||
Value::Bool(value) => query.bind(if *value { 1.0 } else { 0.0 }),
|
||||
Value::Array(_) | Value::Object(_) => {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{}' expected real-compatible value",
|
||||
column.name
|
||||
)));
|
||||
}
|
||||
},
|
||||
SqliteCopyAffinity::Blob => match value {
|
||||
Value::Null => query.bind(Option::<Vec<u8>>::None),
|
||||
Value::String(value) => query.bind(hex_decode(value, &column.name)?),
|
||||
Value::Array(values) => {
|
||||
let mut bytes = Vec::with_capacity(values.len());
|
||||
for value in values {
|
||||
let Some(byte) = value.as_u64().and_then(|value| u8::try_from(value).ok())
|
||||
else {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{}' contains non-byte array value",
|
||||
column.name
|
||||
)));
|
||||
};
|
||||
bytes.push(byte);
|
||||
}
|
||||
query.bind(bytes)
|
||||
}
|
||||
Value::Bool(_) | Value::Number(_) | Value::Object(_) => {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{}' expected blob-compatible value",
|
||||
column.name
|
||||
)));
|
||||
}
|
||||
},
|
||||
SqliteCopyAffinity::Text | SqliteCopyAffinity::Numeric => {
|
||||
bind_sqlite_json_value(query, value)?
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn sql_string_literal(value: &str) -> String {
|
||||
format!("'{}'", value.replace('\'', "''"))
|
||||
}
|
||||
|
||||
fn hex_decode(value: &str, column_name: &str) -> Result<Vec<u8>, DataLayerError> {
|
||||
let value = value.trim();
|
||||
if !value.len().is_multiple_of(2) {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{column_name}' has odd-length hex data"
|
||||
)));
|
||||
}
|
||||
|
||||
let mut bytes = Vec::with_capacity(value.len() / 2);
|
||||
for index in (0..value.len()).step_by(2) {
|
||||
let byte = u8::from_str_radix(&value[index..index + 2], 16).map_err(|err| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"sqlite copy column '{column_name}' has invalid hex data at byte {}: {err}",
|
||||
index / 2
|
||||
))
|
||||
})?;
|
||||
bytes.push(byte);
|
||||
}
|
||||
Ok(bytes)
|
||||
}
|
||||
|
||||
pub async fn export_sqlite_core_jsonl(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
created_at_unix_secs: u64,
|
||||
|
||||
@@ -1,15 +1,15 @@
|
||||
use async_trait::async_trait;
|
||||
use chrono::{TimeZone, Utc};
|
||||
use futures_util::TryStreamExt;
|
||||
use sqlx::{postgres::PgRow, PgPool, Row};
|
||||
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
|
||||
|
||||
use super::types::{
|
||||
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
|
||||
CreateAnnouncementRecord, StoredAnnouncement, StoredAnnouncementPage, UpdateAnnouncementRecord,
|
||||
};
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
use aether_data_query::{push_eq, push_limit, push_limit_offset, WhereClause};
|
||||
|
||||
const FIND_ANNOUNCEMENT_BY_ID_SQL: &str = r#"
|
||||
const ANNOUNCEMENT_SELECT: &str = r#"
|
||||
SELECT
|
||||
a.id,
|
||||
a.title,
|
||||
@@ -26,63 +26,6 @@ SELECT
|
||||
EXTRACT(EPOCH FROM a.updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM announcements a
|
||||
LEFT JOIN users u ON u.id = a.author_id
|
||||
WHERE a.id = $1
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const LIST_ANNOUNCEMENTS_SQL: &str = r#"
|
||||
SELECT
|
||||
a.id,
|
||||
a.title,
|
||||
a.content,
|
||||
a.type,
|
||||
a.priority,
|
||||
a.is_active,
|
||||
a.is_pinned,
|
||||
a.author_id,
|
||||
u.username AS author_username,
|
||||
EXTRACT(EPOCH FROM a.start_time)::bigint AS start_time_unix_secs,
|
||||
EXTRACT(EPOCH FROM a.end_time)::bigint AS end_time_unix_secs,
|
||||
EXTRACT(EPOCH FROM a.created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM a.updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM announcements a
|
||||
LEFT JOIN users u ON u.id = a.author_id
|
||||
WHERE (
|
||||
NOT $1 OR (
|
||||
a.is_active = TRUE
|
||||
AND (a.start_time IS NULL OR a.start_time <= TO_TIMESTAMP($2::double precision))
|
||||
AND (a.end_time IS NULL OR a.end_time >= TO_TIMESTAMP($2::double precision))
|
||||
)
|
||||
)
|
||||
ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC
|
||||
OFFSET $3
|
||||
LIMIT $4
|
||||
"#;
|
||||
|
||||
const COUNT_ANNOUNCEMENTS_SQL: &str = r#"
|
||||
SELECT COUNT(a.id) AS total
|
||||
FROM announcements a
|
||||
WHERE (
|
||||
NOT $1 OR (
|
||||
a.is_active = TRUE
|
||||
AND (a.start_time IS NULL OR a.start_time <= TO_TIMESTAMP($2::double precision))
|
||||
AND (a.end_time IS NULL OR a.end_time >= TO_TIMESTAMP($2::double precision))
|
||||
)
|
||||
)
|
||||
"#;
|
||||
|
||||
const COUNT_UNREAD_ACTIVE_ANNOUNCEMENTS_SQL: &str = r#"
|
||||
SELECT COUNT(a.id) AS total
|
||||
FROM announcements a
|
||||
WHERE a.is_active = TRUE
|
||||
AND (a.start_time IS NULL OR a.start_time <= TO_TIMESTAMP($2::double precision))
|
||||
AND (a.end_time IS NULL OR a.end_time >= TO_TIMESTAMP($2::double precision))
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM announcement_reads r
|
||||
WHERE r.user_id = $1
|
||||
AND r.announcement_id = a.id
|
||||
)
|
||||
"#;
|
||||
|
||||
const CREATE_ANNOUNCEMENT_SQL: &str = r#"
|
||||
@@ -193,6 +136,25 @@ impl SqlxAnnouncementReadRepository {
|
||||
pub fn new(pool: PgPool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
fn apply_active_filter(
|
||||
builder: &mut QueryBuilder<'_, Postgres>,
|
||||
where_clause: &mut WhereClause,
|
||||
active_only: bool,
|
||||
now_unix_secs: u64,
|
||||
) {
|
||||
if !active_only {
|
||||
return;
|
||||
}
|
||||
|
||||
where_clause.push_next(builder);
|
||||
builder
|
||||
.push("a.is_active = TRUE AND (a.start_time IS NULL OR a.start_time <= TO_TIMESTAMP(")
|
||||
.push_bind(now_unix_secs as f64)
|
||||
.push("::double precision)) AND (a.end_time IS NULL OR a.end_time >= TO_TIMESTAMP(")
|
||||
.push_bind(now_unix_secs as f64)
|
||||
.push("::double precision))");
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -201,8 +163,17 @@ impl AnnouncementReadRepository for SqlxAnnouncementReadRepository {
|
||||
&self,
|
||||
announcement_id: &str,
|
||||
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
|
||||
let row = sqlx::query(FIND_ANNOUNCEMENT_BY_ID_SQL)
|
||||
.bind(announcement_id)
|
||||
let mut builder = QueryBuilder::<Postgres>::new(ANNOUNCEMENT_SELECT);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"a.id",
|
||||
announcement_id.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
@@ -214,27 +185,42 @@ impl AnnouncementReadRepository for SqlxAnnouncementReadRepository {
|
||||
query: &AnnouncementListQuery,
|
||||
) -> Result<StoredAnnouncementPage, DataLayerError> {
|
||||
let now_unix_secs = query.now_unix_secs.unwrap_or_else(current_unix_secs);
|
||||
let total_row = sqlx::query(COUNT_ANNOUNCEMENTS_SQL)
|
||||
.bind(query.active_only)
|
||||
.bind(now_unix_secs as f64)
|
||||
let mut count_builder =
|
||||
QueryBuilder::<Postgres>::new("SELECT COUNT(a.id) AS total FROM announcements a");
|
||||
let mut count_where = WhereClause::new();
|
||||
Self::apply_active_filter(
|
||||
&mut count_builder,
|
||||
&mut count_where,
|
||||
query.active_only,
|
||||
now_unix_secs,
|
||||
);
|
||||
let total = count_builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = total_row
|
||||
.try_get::<i64, _>("total")
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64;
|
||||
|
||||
let mut rows = sqlx::query(LIST_ANNOUNCEMENTS_SQL)
|
||||
.bind(query.active_only)
|
||||
.bind(now_unix_secs as f64)
|
||||
.bind(query.offset as i64)
|
||||
.bind(query.limit as i64)
|
||||
.fetch(&self.pool);
|
||||
let mut items = Vec::new();
|
||||
while let Some(row) = rows.try_next().await.map_postgres_err()? {
|
||||
items.push(map_announcement_row(&row)?);
|
||||
}
|
||||
let mut list_builder = QueryBuilder::<Postgres>::new(ANNOUNCEMENT_SELECT);
|
||||
let mut list_where = WhereClause::new();
|
||||
Self::apply_active_filter(
|
||||
&mut list_builder,
|
||||
&mut list_where,
|
||||
query.active_only,
|
||||
now_unix_secs,
|
||||
);
|
||||
list_builder
|
||||
.push(" ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC");
|
||||
push_limit_offset(&mut list_builder, query.limit as i64, query.offset as i64);
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_announcement_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
|
||||
Ok(StoredAnnouncementPage { items, total })
|
||||
}
|
||||
@@ -244,13 +230,22 @@ impl AnnouncementReadRepository for SqlxAnnouncementReadRepository {
|
||||
user_id: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<u64, DataLayerError> {
|
||||
let row = sqlx::query(COUNT_UNREAD_ACTIVE_ANNOUNCEMENTS_SQL)
|
||||
.bind(user_id)
|
||||
.bind(now_unix_secs as f64)
|
||||
let mut builder =
|
||||
QueryBuilder::<Postgres>::new("SELECT COUNT(a.id) AS total FROM announcements a");
|
||||
let mut where_clause = WhereClause::new();
|
||||
Self::apply_active_filter(&mut builder, &mut where_clause, true, now_unix_secs);
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push("NOT EXISTS (SELECT 1 FROM announcement_reads r WHERE r.user_id = ")
|
||||
.push_bind(user_id.to_string())
|
||||
.push(" AND r.announcement_id = a.id)");
|
||||
let total = builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
Ok(row.try_get::<i64, _>("total").map_postgres_err()?.max(0) as u64)
|
||||
.map_postgres_err()?
|
||||
.max(0) as u64;
|
||||
Ok(total)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{sqlite::SqliteRow, Row};
|
||||
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
|
||||
|
||||
use super::types::{
|
||||
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
|
||||
@@ -8,6 +8,7 @@ use super::types::{
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_limit, push_limit_offset, WhereClause};
|
||||
|
||||
const ANNOUNCEMENT_SELECT: &str = r#"
|
||||
SELECT
|
||||
@@ -44,6 +45,27 @@ impl SqliteAnnouncementRepository {
|
||||
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
|
||||
self.find_by_id(announcement_id).await
|
||||
}
|
||||
|
||||
fn apply_active_filter(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
where_clause: &mut WhereClause,
|
||||
active_only: bool,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<(), DataLayerError> {
|
||||
if !active_only {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let now = i64_from_u64(now_unix_secs, "announcements.now")?;
|
||||
where_clause.push_next(builder);
|
||||
builder
|
||||
.push("a.is_active = 1 AND (a.start_time IS NULL OR a.start_time <= ")
|
||||
.push_bind(now)
|
||||
.push(") AND (a.end_time IS NULL OR a.end_time >= ")
|
||||
.push_bind(now)
|
||||
.push(")");
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -52,8 +74,17 @@ impl AnnouncementReadRepository for SqliteAnnouncementRepository {
|
||||
&self,
|
||||
announcement_id: &str,
|
||||
) -> Result<Option<StoredAnnouncement>, DataLayerError> {
|
||||
let row = sqlx::query(&format!("{ANNOUNCEMENT_SELECT} WHERE a.id = ? LIMIT 1"))
|
||||
.bind(announcement_id)
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(ANNOUNCEMENT_SELECT);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"a.id",
|
||||
announcement_id.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
@@ -65,49 +96,38 @@ impl AnnouncementReadRepository for SqliteAnnouncementRepository {
|
||||
query: &AnnouncementListQuery,
|
||||
) -> Result<StoredAnnouncementPage, DataLayerError> {
|
||||
let now_unix_secs = query.now_unix_secs.unwrap_or_else(current_unix_secs);
|
||||
let total_row = sqlx::query(
|
||||
r#"
|
||||
SELECT COUNT(a.id) AS total
|
||||
FROM announcements a
|
||||
WHERE (
|
||||
NOT ? OR (
|
||||
a.is_active = 1
|
||||
AND (a.start_time IS NULL OR a.start_time <= ?)
|
||||
AND (a.end_time IS NULL OR a.end_time >= ?)
|
||||
)
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.bind(query.active_only)
|
||||
.bind(now_unix_secs as i64)
|
||||
.bind(now_unix_secs as i64)
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let total = total_row.try_get::<i64, _>("total").map_sql_err()?.max(0) as u64;
|
||||
let mut count_builder =
|
||||
QueryBuilder::<Sqlite>::new("SELECT COUNT(a.id) AS total FROM announcements a");
|
||||
let mut count_where = WhereClause::new();
|
||||
Self::apply_active_filter(
|
||||
&mut count_builder,
|
||||
&mut count_where,
|
||||
query.active_only,
|
||||
now_unix_secs,
|
||||
)?;
|
||||
let total = count_builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.max(0) as u64;
|
||||
|
||||
let rows = sqlx::query(&format!(
|
||||
r#"
|
||||
{ANNOUNCEMENT_SELECT}
|
||||
WHERE (
|
||||
NOT ? OR (
|
||||
a.is_active = 1
|
||||
AND (a.start_time IS NULL OR a.start_time <= ?)
|
||||
AND (a.end_time IS NULL OR a.end_time >= ?)
|
||||
)
|
||||
)
|
||||
ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC
|
||||
LIMIT ? OFFSET ?
|
||||
"#
|
||||
))
|
||||
.bind(query.active_only)
|
||||
.bind(now_unix_secs as i64)
|
||||
.bind(now_unix_secs as i64)
|
||||
.bind(query.limit as i64)
|
||||
.bind(query.offset as i64)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let mut list_builder = QueryBuilder::<Sqlite>::new(ANNOUNCEMENT_SELECT);
|
||||
let mut list_where = WhereClause::new();
|
||||
Self::apply_active_filter(
|
||||
&mut list_builder,
|
||||
&mut list_where,
|
||||
query.active_only,
|
||||
now_unix_secs,
|
||||
)?;
|
||||
list_builder
|
||||
.push(" ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC");
|
||||
push_limit_offset(&mut list_builder, query.limit as i64, query.offset as i64);
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_announcement_row)
|
||||
@@ -121,28 +141,22 @@ LIMIT ? OFFSET ?
|
||||
user_id: &str,
|
||||
now_unix_secs: u64,
|
||||
) -> Result<u64, DataLayerError> {
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT COUNT(a.id) AS total
|
||||
FROM announcements a
|
||||
WHERE a.is_active = 1
|
||||
AND (a.start_time IS NULL OR a.start_time <= ?)
|
||||
AND (a.end_time IS NULL OR a.end_time >= ?)
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM announcement_reads r
|
||||
WHERE r.user_id = ?
|
||||
AND r.announcement_id = a.id
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.bind(now_unix_secs as i64)
|
||||
.bind(now_unix_secs as i64)
|
||||
.bind(user_id)
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(row.try_get::<i64, _>("total").map_sql_err()?.max(0) as u64)
|
||||
let mut builder =
|
||||
QueryBuilder::<Sqlite>::new("SELECT COUNT(a.id) AS total FROM announcements a");
|
||||
let mut where_clause = WhereClause::new();
|
||||
Self::apply_active_filter(&mut builder, &mut where_clause, true, now_unix_secs)?;
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push("NOT EXISTS (SELECT 1 FROM announcement_reads r WHERE r.user_id = ")
|
||||
.push_bind(user_id.to_string())
|
||||
.push(" AND r.announcement_id = a.id)");
|
||||
let total = builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.max(0) as u64;
|
||||
Ok(total)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,14 +1,14 @@
|
||||
use async_trait::async_trait;
|
||||
use futures_util::{stream::TryStream, TryStreamExt};
|
||||
use sqlx::{postgres::PgRow, PgPool, Row};
|
||||
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
|
||||
|
||||
use super::types::{
|
||||
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
|
||||
StoredOAuthProviderModuleConfig,
|
||||
};
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
use aether_data_query::{push_eq, push_limit, WhereClause};
|
||||
|
||||
const LIST_ENABLED_OAUTH_PROVIDERS_SQL: &str = r#"
|
||||
const OAUTH_PROVIDER_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
provider_type,
|
||||
display_name,
|
||||
@@ -16,11 +16,9 @@ SELECT
|
||||
client_secret_encrypted,
|
||||
redirect_uri
|
||||
FROM oauth_providers
|
||||
WHERE is_enabled = TRUE
|
||||
ORDER BY provider_type ASC
|
||||
"#;
|
||||
|
||||
const GET_LDAP_CONFIG_SQL: &str = r#"
|
||||
const LDAP_CONFIG_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
server_url,
|
||||
bind_dn,
|
||||
@@ -35,8 +33,6 @@ SELECT
|
||||
use_starttls,
|
||||
connect_timeout
|
||||
FROM ldap_configs
|
||||
ORDER BY id ASC
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const UPDATE_LDAP_CONFIG_SQL: &str = r#"
|
||||
@@ -146,18 +142,27 @@ impl SqlxAuthModuleRepository {
|
||||
}
|
||||
}
|
||||
|
||||
async fn collect_query_rows<T, S>(
|
||||
mut rows: S,
|
||||
map_row: fn(&PgRow) -> Result<T, DataLayerError>,
|
||||
) -> Result<Vec<T>, DataLayerError>
|
||||
where
|
||||
S: TryStream<Ok = PgRow, Error = sqlx::Error> + Unpin,
|
||||
{
|
||||
let mut items = Vec::new();
|
||||
while let Some(row) = rows.try_next().await.map_postgres_err()? {
|
||||
items.push(map_row(&row)?);
|
||||
}
|
||||
Ok(items)
|
||||
async fn list_enabled_oauth_providers(
|
||||
pool: &PgPool,
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(OAUTH_PROVIDER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "is_enabled", true);
|
||||
builder.push(" ORDER BY provider_type ASC");
|
||||
let rows = builder.build().fetch_all(pool).await.map_postgres_err()?;
|
||||
rows.iter().map(map_oauth_row).collect()
|
||||
}
|
||||
|
||||
async fn get_ldap_config(pool: &PgPool) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(LDAP_CONFIG_COLUMNS);
|
||||
builder.push(" ORDER BY id ASC");
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
@@ -165,19 +170,11 @@ impl AuthModuleReadRepository for SqlxAuthModuleReadRepository {
|
||||
async fn list_enabled_oauth_providers(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
collect_query_rows(
|
||||
sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL).fetch(&self.pool),
|
||||
map_oauth_row,
|
||||
)
|
||||
.await
|
||||
list_enabled_oauth_providers(&self.pool).await
|
||||
}
|
||||
|
||||
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
get_ldap_config(&self.pool).await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -186,19 +183,11 @@ impl AuthModuleReadRepository for SqlxAuthModuleRepository {
|
||||
async fn list_enabled_oauth_providers(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
collect_query_rows(
|
||||
sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL).fetch(&self.pool),
|
||||
map_oauth_row,
|
||||
)
|
||||
.await
|
||||
list_enabled_oauth_providers(&self.pool).await
|
||||
}
|
||||
|
||||
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
get_ldap_config(&self.pool).await
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{sqlite::SqliteRow, Row};
|
||||
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
|
||||
|
||||
use super::types::{
|
||||
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
|
||||
@@ -8,8 +8,9 @@ use super::types::{
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_limit, WhereClause};
|
||||
|
||||
const LIST_ENABLED_OAUTH_PROVIDERS_SQL: &str = r#"
|
||||
const OAUTH_PROVIDER_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
provider_type,
|
||||
display_name,
|
||||
@@ -17,11 +18,9 @@ SELECT
|
||||
client_secret_encrypted,
|
||||
redirect_uri
|
||||
FROM oauth_providers
|
||||
WHERE is_enabled = 1
|
||||
ORDER BY provider_type ASC
|
||||
"#;
|
||||
|
||||
const GET_LDAP_CONFIG_SQL: &str = r#"
|
||||
const LDAP_CONFIG_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
server_url,
|
||||
bind_dn,
|
||||
@@ -36,8 +35,6 @@ SELECT
|
||||
use_starttls,
|
||||
connect_timeout
|
||||
FROM ldap_configs
|
||||
ORDER BY id ASC
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
@@ -62,24 +59,37 @@ impl SqliteAuthModuleRepository {
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_enabled_oauth_providers(
|
||||
pool: &SqlitePool,
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(OAUTH_PROVIDER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "is_enabled", true);
|
||||
builder.push(" ORDER BY provider_type ASC");
|
||||
let rows = builder.build().fetch_all(pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_oauth_row).collect()
|
||||
}
|
||||
|
||||
async fn get_ldap_config(
|
||||
pool: &SqlitePool,
|
||||
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(LDAP_CONFIG_COLUMNS);
|
||||
builder.push(" ORDER BY id ASC");
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder.build().fetch_optional(pool).await.map_sql_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl AuthModuleReadRepository for SqliteAuthModuleReadRepository {
|
||||
async fn list_enabled_oauth_providers(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_oauth_row).collect()
|
||||
list_enabled_oauth_providers(&self.pool).await
|
||||
}
|
||||
|
||||
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
get_ldap_config(&self.pool).await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -88,19 +98,11 @@ impl AuthModuleReadRepository for SqliteAuthModuleRepository {
|
||||
async fn list_enabled_oauth_providers(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_oauth_row).collect()
|
||||
list_enabled_oauth_providers(&self.pool).await
|
||||
}
|
||||
|
||||
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
|
||||
let row = sqlx::query(GET_LDAP_CONFIG_SQL)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
row.as_ref().map(map_ldap_row).transpose()
|
||||
get_ldap_config(&self.pool).await
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -9,6 +9,9 @@ use super::{
|
||||
};
|
||||
use crate::error::SqlxResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{
|
||||
push_ci_contains, push_eq, push_limit, push_limit_offset, SqlDialect, WhereClause,
|
||||
};
|
||||
|
||||
const RUN_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
@@ -60,35 +63,33 @@ impl SqlxBackgroundTaskRepository {
|
||||
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;
|
||||
}
|
||||
let mut where_clause = if include_where {
|
||||
WhereClause::with_existing_clause()
|
||||
} else {
|
||||
WhereClause::new()
|
||||
};
|
||||
|
||||
if let Some(kind) = query.kind {
|
||||
push_where(builder);
|
||||
builder.push("kind = ").push_bind(kind.as_database());
|
||||
push_eq(builder, &mut where_clause, "kind", kind.as_database());
|
||||
}
|
||||
if let Some(status) = query.status {
|
||||
push_where(builder);
|
||||
builder.push("status = ").push_bind(status.as_database());
|
||||
push_eq(builder, &mut where_clause, "status", status.as_database());
|
||||
}
|
||||
if let Some(trigger) = query.trigger.as_deref() {
|
||||
push_where(builder);
|
||||
builder
|
||||
.push("\"trigger\" = ")
|
||||
.push_bind(trigger.to_string());
|
||||
push_eq(
|
||||
builder,
|
||||
&mut where_clause,
|
||||
&SqlDialect::Postgres.quote_ident("trigger"),
|
||||
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()));
|
||||
push_ci_contains(
|
||||
builder,
|
||||
&mut where_clause,
|
||||
SqlDialect::Postgres,
|
||||
"task_key",
|
||||
task_key_substring,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -99,8 +100,12 @@ impl BackgroundTaskReadRepository for SqlxBackgroundTaskRepository {
|
||||
&self,
|
||||
run_id: &str,
|
||||
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
|
||||
let row = sqlx::query(&format!("{RUN_COLUMNS} WHERE id = $1 LIMIT 1"))
|
||||
.bind(run_id)
|
||||
let mut builder = QueryBuilder::<Postgres>::new(RUN_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "id", run_id.to_string());
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
@@ -124,12 +129,12 @@ impl BackgroundTaskReadRepository for SqlxBackgroundTaskRepository {
|
||||
|
||||
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")?);
|
||||
builder.push(" ORDER BY created_at_unix_secs DESC, updated_at_unix_secs DESC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64_from_usize(limit, "background task run limit")?,
|
||||
i64_from_usize(query.offset, "background task run offset")?,
|
||||
);
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
@@ -153,15 +158,25 @@ impl BackgroundTaskReadRepository for SqlxBackgroundTaskRepository {
|
||||
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()?;
|
||||
let mut builder = QueryBuilder::<Postgres>::new(EVENT_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"run_id",
|
||||
run_id.to_string(),
|
||||
);
|
||||
builder.push(" ORDER BY created_at_unix_secs ASC, id ASC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64_from_usize(limit, "background task event limit")?,
|
||||
i64_from_usize(offset, "background task event offset")?,
|
||||
);
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_event_row).collect()
|
||||
}
|
||||
|
||||
|
||||
@@ -10,6 +10,9 @@ use super::{
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{
|
||||
push_ci_contains, push_eq, push_limit, push_limit_offset, SqlDialect, WhereClause,
|
||||
};
|
||||
|
||||
const RUN_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
@@ -57,46 +60,29 @@ impl SqliteBackgroundTaskRepository {
|
||||
}
|
||||
|
||||
fn apply_run_filter(builder: &mut QueryBuilder<'_, Sqlite>, query: &BackgroundTaskListQuery) {
|
||||
let mut has_where = false;
|
||||
let mut where_clause = WhereClause::new();
|
||||
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());
|
||||
push_eq(builder, &mut where_clause, "kind", 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());
|
||||
push_eq(builder, &mut where_clause, "status", 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());
|
||||
push_eq(
|
||||
builder,
|
||||
&mut where_clause,
|
||||
&SqlDialect::Sqlite.quote_ident("trigger"),
|
||||
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()
|
||||
));
|
||||
push_ci_contains(
|
||||
builder,
|
||||
&mut where_clause,
|
||||
SqlDialect::Sqlite,
|
||||
"task_key",
|
||||
task_key_substring,
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -107,8 +93,12 @@ impl BackgroundTaskReadRepository for SqliteBackgroundTaskRepository {
|
||||
&self,
|
||||
run_id: &str,
|
||||
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
|
||||
let row = sqlx::query(&format!("{RUN_COLUMNS} WHERE id = ? LIMIT 1"))
|
||||
.bind(run_id)
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(RUN_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "id", run_id.to_string());
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
@@ -131,12 +121,12 @@ impl BackgroundTaskReadRepository for SqliteBackgroundTaskRepository {
|
||||
|
||||
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")?);
|
||||
builder.push(" ORDER BY created_at_unix_secs DESC, updated_at_unix_secs DESC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64_from_usize(limit, "run limit")?,
|
||||
i64_from_usize(query.offset, "run offset")?,
|
||||
);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
@@ -155,15 +145,21 @@ impl BackgroundTaskReadRepository for SqliteBackgroundTaskRepository {
|
||||
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()?;
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(EVENT_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"run_id",
|
||||
run_id.to_string(),
|
||||
);
|
||||
builder.push(" ORDER BY created_at_unix_secs ASC, id ASC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64_from_usize(limit, "event limit")?,
|
||||
i64_from_usize(offset, "event offset")?,
|
||||
);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_event_row).collect()
|
||||
}
|
||||
|
||||
|
||||
@@ -961,7 +961,7 @@ ORDER BY expires_at ASC, created_at ASC, id ASC
|
||||
allow_wallet_overage &= grant.allow_wallet_overage;
|
||||
let used = sqlx::query_scalar::<_, f64>(
|
||||
r#"
|
||||
SELECT COALESCE(SUM(amount_usd), 0)
|
||||
SELECT CAST(COALESCE(SUM(amount_usd), 0) AS REAL)
|
||||
FROM entitlement_usage_ledgers
|
||||
WHERE user_entitlement_id = ?
|
||||
AND usage_date = ?
|
||||
|
||||
@@ -45,7 +45,15 @@ SELECT
|
||||
m.provider_model_mappings AS model_provider_model_mappings,
|
||||
m.supports_streaming AS model_supports_streaming,
|
||||
m.is_active AS model_is_active,
|
||||
m.is_available AS model_is_available
|
||||
m.is_available AS model_is_available,
|
||||
CASE
|
||||
WHEN json_valid(p.config) THEN
|
||||
CASE
|
||||
WHEN json_type(p.config, '$.pool_advanced') IS NOT NULL THEN 1
|
||||
ELSE 0
|
||||
END
|
||||
ELSE 0
|
||||
END AS provider_pool_enabled
|
||||
FROM providers p
|
||||
INNER JOIN provider_endpoints pe ON pe.provider_id = p.id
|
||||
INNER JOIN provider_api_keys pak ON pak.provider_id = p.id
|
||||
@@ -67,52 +75,139 @@ pub struct SqliteMinimalCandidateSelectionReadRepository {
|
||||
#[derive(Debug, Clone)]
|
||||
struct CandidateSelectionRow {
|
||||
row: StoredMinimalCandidateSelectionRow,
|
||||
provider_pool_enabled: bool,
|
||||
key_auth_config: Option<String>,
|
||||
key_last_used_at_unix_secs: Option<u64>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
enum SelectedRowsOrder {
|
||||
WithGlobalModel,
|
||||
WithoutGlobalModel,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
enum SelectedRowsFilter<'a> {
|
||||
None,
|
||||
GlobalModel(&'a str),
|
||||
RequestedModel(&'a str),
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy)]
|
||||
struct SqlPage {
|
||||
limit: i64,
|
||||
offset: i64,
|
||||
}
|
||||
|
||||
impl SqliteMinimalCandidateSelectionReadRepository {
|
||||
pub fn new(pool: SqlitePool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn load_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<CandidateSelectionRow>, DataLayerError> {
|
||||
let canonical_api_format = normalize_api_format(api_format);
|
||||
let storage_aliases = api_format_aliases(&canonical_api_format);
|
||||
let match_aliases = sql_match_aliases(&storage_aliases);
|
||||
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_SELECTION_COLUMNS);
|
||||
builder.push(" AND LOWER(pe.api_format) IN (");
|
||||
{
|
||||
let mut separated = builder.separated(", ");
|
||||
for alias in &match_aliases {
|
||||
separated.push_bind(alias);
|
||||
}
|
||||
}
|
||||
builder.push(")");
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
let mut items = rows
|
||||
.iter()
|
||||
.map(map_candidate_selection_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
items.retain(|item| {
|
||||
api_format_matches(&item.row.endpoint_api_format, &canonical_api_format)
|
||||
&& item.row.key_supports_api_format(&canonical_api_format)
|
||||
&& key_auth_channel_matches(item, &canonical_api_format)
|
||||
});
|
||||
Ok(items)
|
||||
}
|
||||
|
||||
async fn selected_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self.load_rows_for_api_format(api_format).await?;
|
||||
Ok(sort_rows(select_pool_rows(rows), true))
|
||||
self.load_selected_rows_for_api_format(
|
||||
api_format,
|
||||
SelectedRowsFilter::None,
|
||||
SelectedRowsOrder::WithGlobalModel,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn load_selected_rows_for_api_format(
|
||||
&self,
|
||||
api_format: &str,
|
||||
filter: SelectedRowsFilter<'_>,
|
||||
order: SelectedRowsOrder,
|
||||
page: Option<SqlPage>,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let canonical_api_format = normalize_api_format(api_format);
|
||||
let storage_aliases = api_format_aliases(&canonical_api_format);
|
||||
let match_aliases = sql_match_aliases(&storage_aliases);
|
||||
let mut rows = Vec::new();
|
||||
|
||||
for storage_api_format in storage_aliases {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new("WITH candidate_rows AS (");
|
||||
builder.push(CANDIDATE_SELECTION_COLUMNS);
|
||||
push_candidate_sql_filters(&mut builder, &storage_api_format, &match_aliases);
|
||||
match filter {
|
||||
SelectedRowsFilter::None => {}
|
||||
SelectedRowsFilter::GlobalModel(global_model_name) => {
|
||||
builder.push(" AND gm.name = ");
|
||||
builder.push_bind(global_model_name);
|
||||
}
|
||||
SelectedRowsFilter::RequestedModel(requested_model_name) => {
|
||||
push_requested_model_sql_filter(
|
||||
&mut builder,
|
||||
requested_model_name,
|
||||
&match_aliases,
|
||||
);
|
||||
}
|
||||
}
|
||||
builder.push(
|
||||
r#"
|
||||
),
|
||||
pool_rows AS (
|
||||
SELECT candidate.*
|
||||
FROM candidate_rows candidate
|
||||
WHERE candidate.provider_pool_enabled = 1
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM candidate_rows other
|
||||
WHERE other.provider_pool_enabled = 1
|
||||
AND other.provider_id = candidate.provider_id
|
||||
AND other.endpoint_id = candidate.endpoint_id
|
||||
AND other.model_id = candidate.model_id
|
||||
AND (
|
||||
other.key_internal_priority < candidate.key_internal_priority
|
||||
OR (
|
||||
other.key_internal_priority = candidate.key_internal_priority
|
||||
AND other.key_id < candidate.key_id
|
||||
)
|
||||
)
|
||||
)
|
||||
),
|
||||
selected_rows AS (
|
||||
SELECT * FROM candidate_rows WHERE provider_pool_enabled = 0
|
||||
UNION ALL
|
||||
SELECT * FROM pool_rows
|
||||
)
|
||||
SELECT * FROM selected_rows
|
||||
"#,
|
||||
);
|
||||
push_selected_rows_order(&mut builder, order);
|
||||
if let Some(page) = page {
|
||||
builder.push(" LIMIT ");
|
||||
builder.push_bind(page.limit);
|
||||
builder.push(" OFFSET ");
|
||||
builder.push_bind(page.offset);
|
||||
}
|
||||
|
||||
let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
let mut items = query_rows
|
||||
.iter()
|
||||
.map(map_candidate_selection_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
items.retain(|item| {
|
||||
api_format_matches(&item.row.endpoint_api_format, &canonical_api_format)
|
||||
&& item.row.key_supports_api_format(&canonical_api_format)
|
||||
&& key_auth_channel_matches(item, &canonical_api_format)
|
||||
});
|
||||
rows.extend(items.into_iter().map(|item| item.row));
|
||||
}
|
||||
|
||||
let rows = match filter {
|
||||
SelectedRowsFilter::RequestedModel(requested_model_name) => rows
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row_matches_requested_model(row, requested_model_name, &canonical_api_format)
|
||||
})
|
||||
.collect(),
|
||||
_ => rows,
|
||||
};
|
||||
Ok(dedupe_candidate_selection_rows(rows))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -130,14 +225,13 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
|
||||
api_format: &str,
|
||||
global_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
Ok(sort_rows(
|
||||
self.selected_rows_for_api_format(api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| row.global_model_name == global_model_name)
|
||||
.collect(),
|
||||
false,
|
||||
))
|
||||
self.load_selected_rows_for_api_format(
|
||||
api_format,
|
||||
SelectedRowsFilter::GlobalModel(global_model_name),
|
||||
SelectedRowsOrder::WithoutGlobalModel,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_for_exact_api_format_and_requested_model(
|
||||
@@ -145,13 +239,11 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
|
||||
api_format: &str,
|
||||
requested_model_name: &str,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
self.list_for_exact_api_format_and_requested_model_page(
|
||||
&StoredRequestedModelCandidateRowsQuery {
|
||||
api_format: api_format.to_string(),
|
||||
requested_model_name: requested_model_name.to_string(),
|
||||
offset: 0,
|
||||
limit: u32::MAX,
|
||||
},
|
||||
self.load_selected_rows_for_api_format(
|
||||
api_format,
|
||||
SelectedRowsFilter::RequestedModel(requested_model_name),
|
||||
SelectedRowsOrder::WithGlobalModel,
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
@@ -160,42 +252,74 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
|
||||
&self,
|
||||
query: &StoredRequestedModelCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self
|
||||
.selected_rows_for_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row_matches_requested_model(row, &query.requested_model_name, &query.api_format)
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
Ok(sort_rows(rows, true)
|
||||
.into_iter()
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.collect())
|
||||
self.load_selected_rows_for_api_format(
|
||||
&query.api_format,
|
||||
SelectedRowsFilter::RequestedModel(&query.requested_model_name),
|
||||
SelectedRowsOrder::WithGlobalModel,
|
||||
Some(SqlPage {
|
||||
limit: i64::from(query.limit.max(1)),
|
||||
offset: i64::from(query.offset),
|
||||
}),
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group(
|
||||
&self,
|
||||
query: &StoredPoolKeyCandidateRowsQuery,
|
||||
) -> Result<Vec<StoredMinimalCandidateSelectionRow>, DataLayerError> {
|
||||
let rows = self
|
||||
.load_rows_for_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row.row.provider_id == query.provider_id
|
||||
&& row.row.endpoint_id == query.endpoint_id
|
||||
&& row.row.model_id == query.model_id
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let mut rows = sort_pool_key_rows(rows, &query.order);
|
||||
Ok(rows
|
||||
.drain(..)
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.map(|item| item.row)
|
||||
.collect())
|
||||
let canonical_api_format = normalize_api_format(&query.api_format);
|
||||
let storage_aliases = api_format_aliases(&canonical_api_format);
|
||||
let match_aliases = sql_match_aliases(&storage_aliases);
|
||||
let mut rows = Vec::<CandidateSelectionRow>::new();
|
||||
let page_in_sql = !matches!(query.order, StoredPoolKeyCandidateOrder::LoadBalance { .. });
|
||||
|
||||
for storage_api_format in storage_aliases {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_SELECTION_COLUMNS);
|
||||
push_candidate_sql_filters(&mut builder, &storage_api_format, &match_aliases);
|
||||
builder.push(" AND p.id = ");
|
||||
builder.push_bind(&query.provider_id);
|
||||
builder.push(" AND pe.id = ");
|
||||
builder.push_bind(&query.endpoint_id);
|
||||
builder.push(" AND m.id = ");
|
||||
builder.push_bind(&query.model_id);
|
||||
if page_in_sql {
|
||||
push_pool_key_order(&mut builder, &query.order);
|
||||
builder.push(" LIMIT ");
|
||||
builder.push_bind(i64::from(query.limit.max(1)));
|
||||
builder.push(" OFFSET ");
|
||||
builder.push_bind(i64::from(query.offset));
|
||||
} else {
|
||||
builder.push(" ORDER BY pak.id ASC");
|
||||
}
|
||||
|
||||
let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
let mut items = query_rows
|
||||
.iter()
|
||||
.map(map_candidate_selection_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
items.retain(|item| {
|
||||
api_format_matches(&item.row.endpoint_api_format, &canonical_api_format)
|
||||
&& item.row.key_supports_api_format(&canonical_api_format)
|
||||
&& key_auth_channel_matches(item, &canonical_api_format)
|
||||
});
|
||||
rows.extend(items);
|
||||
}
|
||||
|
||||
if page_in_sql {
|
||||
Ok(dedupe_candidate_selection_rows(
|
||||
rows.into_iter().map(|item| item.row).collect(),
|
||||
))
|
||||
} else {
|
||||
Ok(dedupe_candidate_selection_rows(
|
||||
sort_pool_key_rows(rows, &query.order)
|
||||
.into_iter()
|
||||
.skip(query.offset as usize)
|
||||
.take(query.limit as usize)
|
||||
.map(|item| item.row)
|
||||
.collect(),
|
||||
))
|
||||
}
|
||||
}
|
||||
|
||||
async fn list_pool_key_rows_for_group_key_ids(
|
||||
@@ -211,75 +335,330 @@ impl MinimalCandidateSelectionReadRepository for SqliteMinimalCandidateSelection
|
||||
.enumerate()
|
||||
.map(|(index, key_id)| (key_id.as_str(), index))
|
||||
.collect::<BTreeMap<_, _>>();
|
||||
let mut rows = self
|
||||
.load_rows_for_api_format(&query.api_format)
|
||||
.await?
|
||||
.into_iter()
|
||||
.filter(|row| {
|
||||
row.row.provider_id == query.provider_id
|
||||
&& row.row.endpoint_id == query.endpoint_id
|
||||
&& row.row.model_id == query.model_id
|
||||
&& key_order.contains_key(row.row.key_id.as_str())
|
||||
})
|
||||
.map(|item| item.row)
|
||||
.collect::<Vec<_>>();
|
||||
let canonical_api_format = normalize_api_format(&query.api_format);
|
||||
let storage_aliases = api_format_aliases(&canonical_api_format);
|
||||
let match_aliases = sql_match_aliases(&storage_aliases);
|
||||
let mut rows = Vec::new();
|
||||
|
||||
for storage_api_format in storage_aliases {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_SELECTION_COLUMNS);
|
||||
push_candidate_sql_filters(&mut builder, &storage_api_format, &match_aliases);
|
||||
builder.push(" AND p.id = ");
|
||||
builder.push_bind(&query.provider_id);
|
||||
builder.push(" AND pe.id = ");
|
||||
builder.push_bind(&query.endpoint_id);
|
||||
builder.push(" AND m.id = ");
|
||||
builder.push_bind(&query.model_id);
|
||||
builder.push(" AND pak.id IN (");
|
||||
{
|
||||
let mut separated = builder.separated(", ");
|
||||
for key_id in &query.key_ids {
|
||||
separated.push_bind(key_id);
|
||||
}
|
||||
}
|
||||
builder.push(")");
|
||||
builder.push(" ORDER BY CASE pak.id");
|
||||
for (index, key_id) in query.key_ids.iter().enumerate() {
|
||||
builder.push(" WHEN ");
|
||||
builder.push_bind(key_id);
|
||||
builder.push(" THEN ");
|
||||
builder.push_bind(i64::try_from(index).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue("key id order index overflowed".to_string())
|
||||
})?);
|
||||
}
|
||||
builder.push(" ELSE ");
|
||||
builder.push_bind(i64::try_from(query.key_ids.len()).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue("key id order length overflowed".to_string())
|
||||
})?);
|
||||
builder.push(" END ASC, pak.id ASC");
|
||||
|
||||
let query_rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
let mut items = query_rows
|
||||
.iter()
|
||||
.map(map_candidate_selection_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
items.retain(|item| {
|
||||
api_format_matches(&item.row.endpoint_api_format, &canonical_api_format)
|
||||
&& item.row.key_supports_api_format(&canonical_api_format)
|
||||
&& key_auth_channel_matches(item, &canonical_api_format)
|
||||
});
|
||||
rows.extend(items.into_iter().map(|item| item.row));
|
||||
}
|
||||
|
||||
let mut rows = dedupe_candidate_selection_rows(rows);
|
||||
rows.sort_by(|left, right| {
|
||||
key_order
|
||||
.get(left.key_id.as_str())
|
||||
.cmp(&key_order.get(right.key_id.as_str()))
|
||||
.then(left.key_id.cmp(&right.key_id))
|
||||
});
|
||||
Ok(dedupe_candidate_selection_rows(rows))
|
||||
Ok(rows)
|
||||
}
|
||||
}
|
||||
|
||||
fn select_pool_rows(rows: Vec<CandidateSelectionRow>) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
let mut selected = Vec::new();
|
||||
let mut pool_rows =
|
||||
BTreeMap::<(String, String, String), StoredMinimalCandidateSelectionRow>::new();
|
||||
for item in rows {
|
||||
if !item.provider_pool_enabled {
|
||||
selected.push(item.row);
|
||||
continue;
|
||||
}
|
||||
let key = (
|
||||
item.row.provider_id.clone(),
|
||||
item.row.endpoint_id.clone(),
|
||||
item.row.model_id.clone(),
|
||||
);
|
||||
match pool_rows.get(&key) {
|
||||
Some(existing)
|
||||
if (existing.key_internal_priority, existing.key_id.as_str())
|
||||
<= (item.row.key_internal_priority, item.row.key_id.as_str()) => {}
|
||||
_ => {
|
||||
pool_rows.insert(key, item.row);
|
||||
}
|
||||
}
|
||||
}
|
||||
selected.extend(pool_rows.into_values());
|
||||
dedupe_candidate_selection_rows(selected)
|
||||
fn push_candidate_sql_filters(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
storage_api_format: &str,
|
||||
match_aliases: &[String],
|
||||
) {
|
||||
builder.push(" AND LOWER(COALESCE(pe.api_format, '')) = ");
|
||||
builder.push_bind(storage_api_format.trim().to_ascii_lowercase());
|
||||
push_key_api_format_sql_filter(builder, match_aliases);
|
||||
push_key_auth_channel_sql_filter(builder, storage_api_format);
|
||||
}
|
||||
|
||||
fn sort_rows(
|
||||
mut rows: Vec<StoredMinimalCandidateSelectionRow>,
|
||||
include_global_model: bool,
|
||||
) -> Vec<StoredMinimalCandidateSelectionRow> {
|
||||
rows.sort_by(|left, right| {
|
||||
if include_global_model {
|
||||
let ordering = left.global_model_name.cmp(&right.global_model_name);
|
||||
if !ordering.is_eq() {
|
||||
return ordering;
|
||||
}
|
||||
fn push_key_api_format_sql_filter(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
match_aliases: &[String],
|
||||
) {
|
||||
builder.push(
|
||||
r#"
|
||||
AND (
|
||||
pak.api_formats IS NULL
|
||||
OR TRIM(pak.api_formats) = ''
|
||||
OR CASE
|
||||
WHEN json_valid(pak.api_formats) THEN
|
||||
(
|
||||
(
|
||||
json_type(pak.api_formats) = 'array'
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM json_each(pak.api_formats) AS fmt
|
||||
WHERE LOWER(TRIM(CAST(fmt.value AS TEXT))) IN (
|
||||
"#,
|
||||
);
|
||||
push_bind_list(builder, match_aliases);
|
||||
builder.push(
|
||||
r#"
|
||||
)
|
||||
)
|
||||
)
|
||||
OR (
|
||||
json_type(pak.api_formats) = 'text'
|
||||
AND LOWER(TRIM(CAST(json_extract(pak.api_formats, '$') AS TEXT))) IN (
|
||||
"#,
|
||||
);
|
||||
push_bind_list(builder, match_aliases);
|
||||
builder.push(
|
||||
r#"
|
||||
)
|
||||
)
|
||||
OR (
|
||||
json_type(pak.api_formats) = 'text'
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM json_each(
|
||||
CASE
|
||||
WHEN json_valid(CAST(json_extract(pak.api_formats, '$') AS TEXT))
|
||||
THEN CAST(json_extract(pak.api_formats, '$') AS TEXT)
|
||||
ELSE '[]'
|
||||
END
|
||||
) AS fmt
|
||||
WHERE LOWER(TRIM(CAST(fmt.value AS TEXT))) IN (
|
||||
"#,
|
||||
);
|
||||
push_bind_list(builder, match_aliases);
|
||||
builder.push(
|
||||
r#"
|
||||
)
|
||||
)
|
||||
)
|
||||
)
|
||||
ELSE 0
|
||||
END
|
||||
OR LOWER(TRIM(pak.api_formats)) IN (
|
||||
"#,
|
||||
);
|
||||
push_bind_list(builder, match_aliases);
|
||||
builder.push(
|
||||
r#"
|
||||
)
|
||||
)
|
||||
"#,
|
||||
);
|
||||
}
|
||||
|
||||
fn push_key_auth_channel_sql_filter(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
storage_api_format: &str,
|
||||
) {
|
||||
let api_format = normalize_api_format(storage_api_format);
|
||||
builder.push(
|
||||
r#"
|
||||
AND (
|
||||
(
|
||||
LOWER(TRIM(p.provider_type)) = 'codex'
|
||||
AND LOWER(TRIM(pak.auth_type)) = 'oauth'
|
||||
AND "#,
|
||||
);
|
||||
builder.push_bind(api_format.clone());
|
||||
builder.push(
|
||||
r#" IN ('openai:responses', 'openai:responses:compact', 'openai:image')
|
||||
)
|
||||
OR (
|
||||
LOWER(TRIM(p.provider_type)) = 'chatgpt_web'
|
||||
AND LOWER(TRIM(pak.auth_type)) IN ('oauth', 'bearer')
|
||||
AND "#,
|
||||
);
|
||||
builder.push_bind(api_format.clone());
|
||||
builder.push(
|
||||
r#" = 'openai:image'
|
||||
)
|
||||
OR (
|
||||
LOWER(TRIM(p.provider_type)) = 'claude_code'
|
||||
AND LOWER(TRIM(pak.auth_type)) = 'oauth'
|
||||
AND "#,
|
||||
);
|
||||
builder.push_bind(api_format.clone());
|
||||
builder.push(
|
||||
r#" = 'claude:messages'
|
||||
)
|
||||
OR (
|
||||
LOWER(TRIM(p.provider_type)) = 'kiro'
|
||||
AND "#,
|
||||
);
|
||||
builder.push_bind(api_format.clone());
|
||||
builder.push(
|
||||
r#" = 'claude:messages'
|
||||
AND (
|
||||
LOWER(TRIM(pak.auth_type)) = 'oauth'
|
||||
OR (
|
||||
LOWER(TRIM(pak.auth_type)) = 'bearer'
|
||||
AND pak.auth_config IS NOT NULL
|
||||
AND TRIM(pak.auth_config) <> ''
|
||||
)
|
||||
)
|
||||
)
|
||||
OR (
|
||||
LOWER(TRIM(p.provider_type)) IN ('gemini_cli', 'antigravity')
|
||||
AND LOWER(TRIM(pak.auth_type)) = 'oauth'
|
||||
AND "#,
|
||||
);
|
||||
builder.push_bind(api_format.clone());
|
||||
builder.push(
|
||||
r#" = 'gemini:generate_content'
|
||||
)
|
||||
OR (
|
||||
LOWER(TRIM(p.provider_type)) = 'vertex_ai'
|
||||
AND (
|
||||
(
|
||||
LOWER(TRIM(pak.auth_type)) = 'api_key'
|
||||
AND "#,
|
||||
);
|
||||
builder.push_bind(api_format.clone());
|
||||
builder.push(
|
||||
r#" = 'gemini:generate_content'
|
||||
)
|
||||
OR (
|
||||
LOWER(TRIM(pak.auth_type)) IN ('service_account', 'vertex_ai')
|
||||
AND "#,
|
||||
);
|
||||
builder.push_bind(api_format.clone());
|
||||
builder.push(
|
||||
r#" IN ('claude:messages', 'gemini:generate_content')
|
||||
)
|
||||
)
|
||||
)
|
||||
OR (
|
||||
LOWER(TRIM(p.provider_type)) NOT IN (
|
||||
'chatgpt_web',
|
||||
'claude_code',
|
||||
'codex',
|
||||
'gemini_cli',
|
||||
'vertex_ai',
|
||||
'antigravity',
|
||||
'kiro'
|
||||
)
|
||||
AND LOWER(TRIM(pak.auth_type)) <> 'oauth'
|
||||
)
|
||||
)
|
||||
"#,
|
||||
);
|
||||
}
|
||||
|
||||
fn push_requested_model_sql_filter(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
requested_model_name: &str,
|
||||
_match_aliases: &[String],
|
||||
) {
|
||||
builder.push(
|
||||
r#"
|
||||
AND (
|
||||
gm.name = "#,
|
||||
);
|
||||
builder.push_bind(requested_model_name.to_string());
|
||||
builder.push(
|
||||
r#"
|
||||
OR m.provider_model_name = "#,
|
||||
);
|
||||
builder.push_bind(requested_model_name.to_string());
|
||||
builder.push(
|
||||
r#"
|
||||
OR (
|
||||
m.provider_model_mappings IS NOT NULL
|
||||
AND m.provider_model_mappings LIKE "#,
|
||||
);
|
||||
builder.push_bind(format!(
|
||||
"%{}%",
|
||||
requested_model_name
|
||||
.replace('\\', "\\\\")
|
||||
.replace('%', "\\%")
|
||||
.replace('_', "\\_")
|
||||
));
|
||||
builder.push(
|
||||
r#"
|
||||
ESCAPE '\'
|
||||
)
|
||||
)
|
||||
"#,
|
||||
);
|
||||
}
|
||||
|
||||
fn push_selected_rows_order(builder: &mut QueryBuilder<'_, Sqlite>, order: SelectedRowsOrder) {
|
||||
builder.push(" ORDER BY ");
|
||||
if matches!(order, SelectedRowsOrder::WithGlobalModel) {
|
||||
builder.push("global_model_name ASC, ");
|
||||
}
|
||||
builder.push(
|
||||
"provider_priority ASC, key_internal_priority ASC, provider_id ASC, endpoint_id ASC, key_id ASC, model_id ASC",
|
||||
);
|
||||
}
|
||||
|
||||
fn push_pool_key_order(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
order: &StoredPoolKeyCandidateOrder,
|
||||
) {
|
||||
match order {
|
||||
StoredPoolKeyCandidateOrder::InternalPriority => {
|
||||
builder.push(" ORDER BY pak.internal_priority ASC, pak.id ASC");
|
||||
}
|
||||
left.provider_priority
|
||||
.cmp(&right.provider_priority)
|
||||
.then(left.key_internal_priority.cmp(&right.key_internal_priority))
|
||||
.then(left.provider_id.cmp(&right.provider_id))
|
||||
.then(left.endpoint_id.cmp(&right.endpoint_id))
|
||||
.then(left.key_id.cmp(&right.key_id))
|
||||
.then(left.model_id.cmp(&right.model_id))
|
||||
});
|
||||
rows
|
||||
StoredPoolKeyCandidateOrder::Lru => {
|
||||
builder.push(
|
||||
" ORDER BY pak.last_used_at IS NOT NULL ASC, pak.last_used_at ASC, pak.internal_priority ASC, pak.id ASC",
|
||||
);
|
||||
}
|
||||
StoredPoolKeyCandidateOrder::CacheAffinity => {
|
||||
builder.push(
|
||||
" ORDER BY pak.last_used_at IS NULL ASC, pak.last_used_at DESC, pak.internal_priority ASC, pak.id ASC",
|
||||
);
|
||||
}
|
||||
StoredPoolKeyCandidateOrder::SingleAccount => {
|
||||
builder.push(
|
||||
" ORDER BY pak.internal_priority ASC, pak.last_used_at IS NULL ASC, pak.last_used_at DESC, pak.id ASC",
|
||||
);
|
||||
}
|
||||
StoredPoolKeyCandidateOrder::LoadBalance { seed } => {
|
||||
let _ = seed;
|
||||
builder.push(" ORDER BY pak.id ASC");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn push_bind_list(builder: &mut QueryBuilder<'_, Sqlite>, values: &[String]) {
|
||||
let mut separated = builder.separated(", ");
|
||||
for value in values {
|
||||
separated.push_bind(value.clone());
|
||||
}
|
||||
}
|
||||
|
||||
fn sort_pool_key_rows(
|
||||
@@ -473,9 +852,8 @@ fn dedupe_candidate_selection_rows(
|
||||
}
|
||||
|
||||
fn map_candidate_selection_row(row: &SqliteRow) -> Result<CandidateSelectionRow, DataLayerError> {
|
||||
let provider_config = parse_json(row.try_get("provider_config").ok().flatten())?;
|
||||
let _provider_config = parse_json(row.try_get("provider_config").ok().flatten())?;
|
||||
let global_model_config = parse_json(row.try_get("global_model_config").ok().flatten())?;
|
||||
let provider_pool_enabled = json_object_field_present(&provider_config, "pool_advanced");
|
||||
let global_model_mappings = global_model_config
|
||||
.as_ref()
|
||||
.and_then(|value| value.get("model_mappings").cloned());
|
||||
@@ -528,7 +906,6 @@ fn map_candidate_selection_row(row: &SqliteRow) -> Result<CandidateSelectionRow,
|
||||
model_is_active: row.try_get("model_is_active").map_sql_err()?,
|
||||
model_is_available: row.try_get("model_is_available").map_sql_err()?,
|
||||
},
|
||||
provider_pool_enabled,
|
||||
key_auth_config: row.try_get("key_auth_config").map_sql_err()?,
|
||||
key_last_used_at_unix_secs: row
|
||||
.try_get::<Option<i64>, _>("key_last_used_at_unix_secs")
|
||||
@@ -550,13 +927,6 @@ fn parse_json(value: Option<String>) -> Result<Option<serde_json::Value>, DataLa
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn json_object_field_present(value: &Option<serde_json::Value>, field: &str) -> bool {
|
||||
value
|
||||
.as_ref()
|
||||
.and_then(|value| value.get(field))
|
||||
.is_some_and(|value| !value.is_null())
|
||||
}
|
||||
|
||||
fn json_bool(value: &serde_json::Value) -> Option<bool> {
|
||||
value.as_bool().or_else(|| {
|
||||
value
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use async_trait::async_trait;
|
||||
use futures_util::{future::BoxFuture, stream::TryStream, TryStreamExt};
|
||||
use sqlx::{postgres::PgRow, PgPool, Row};
|
||||
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
|
||||
use uuid::Uuid;
|
||||
|
||||
use super::{
|
||||
@@ -10,6 +10,7 @@ use super::{
|
||||
};
|
||||
use crate::driver::postgres::PostgresTransactionRunner;
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
use aether_data_query::{push_eq, push_in, push_limit, WhereClause};
|
||||
|
||||
const LIST_BY_REQUEST_ID_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -42,115 +43,6 @@ WHERE request_id = $1
|
||||
ORDER BY candidate_index ASC, retry_index ASC, created_at ASC
|
||||
"#;
|
||||
|
||||
const LIST_RECENT_SQL: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
request_id,
|
||||
user_id,
|
||||
api_key_id,
|
||||
username,
|
||||
api_key_name,
|
||||
candidate_index,
|
||||
retry_index,
|
||||
provider_id,
|
||||
endpoint_id,
|
||||
key_id,
|
||||
status,
|
||||
skip_reason,
|
||||
is_cached,
|
||||
status_code,
|
||||
error_type,
|
||||
error_message,
|
||||
latency_ms,
|
||||
concurrent_requests,
|
||||
extra_data,
|
||||
required_capabilities,
|
||||
CAST(EXTRACT(EPOCH FROM created_at) * 1000 AS BIGINT) AS created_at_unix_ms,
|
||||
CAST(EXTRACT(EPOCH FROM started_at) * 1000 AS BIGINT) AS started_at_unix_ms,
|
||||
CAST(EXTRACT(EPOCH FROM finished_at) * 1000 AS BIGINT) AS finished_at_unix_ms
|
||||
FROM request_candidates
|
||||
ORDER BY created_at DESC
|
||||
LIMIT $1
|
||||
"#;
|
||||
|
||||
const LIST_BY_PROVIDER_ID_SQL: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
request_id,
|
||||
user_id,
|
||||
api_key_id,
|
||||
username,
|
||||
api_key_name,
|
||||
candidate_index,
|
||||
retry_index,
|
||||
provider_id,
|
||||
endpoint_id,
|
||||
key_id,
|
||||
status,
|
||||
skip_reason,
|
||||
is_cached,
|
||||
status_code,
|
||||
error_type,
|
||||
error_message,
|
||||
latency_ms,
|
||||
concurrent_requests,
|
||||
extra_data,
|
||||
required_capabilities,
|
||||
CAST(EXTRACT(EPOCH FROM created_at) * 1000 AS BIGINT) AS created_at_unix_ms,
|
||||
CAST(EXTRACT(EPOCH FROM started_at) * 1000 AS BIGINT) AS started_at_unix_ms,
|
||||
CAST(EXTRACT(EPOCH FROM finished_at) * 1000 AS BIGINT) AS finished_at_unix_ms
|
||||
FROM request_candidates
|
||||
WHERE provider_id = $1
|
||||
ORDER BY created_at DESC
|
||||
LIMIT $2
|
||||
"#;
|
||||
|
||||
const LIST_FINALIZED_BY_ENDPOINT_IDS_SINCE_SQL: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
request_id,
|
||||
user_id,
|
||||
api_key_id,
|
||||
username,
|
||||
api_key_name,
|
||||
candidate_index,
|
||||
retry_index,
|
||||
provider_id,
|
||||
endpoint_id,
|
||||
key_id,
|
||||
status,
|
||||
skip_reason,
|
||||
is_cached,
|
||||
status_code,
|
||||
error_type,
|
||||
error_message,
|
||||
latency_ms,
|
||||
concurrent_requests,
|
||||
extra_data,
|
||||
required_capabilities,
|
||||
CAST(EXTRACT(EPOCH FROM created_at) * 1000 AS BIGINT) AS created_at_unix_ms,
|
||||
CAST(EXTRACT(EPOCH FROM started_at) * 1000 AS BIGINT) AS started_at_unix_ms,
|
||||
CAST(EXTRACT(EPOCH FROM finished_at) * 1000 AS BIGINT) AS finished_at_unix_ms
|
||||
FROM request_candidates
|
||||
WHERE endpoint_id = ANY($1)
|
||||
AND created_at >= TO_TIMESTAMP($2)
|
||||
AND status IN ('success', 'failed', 'skipped')
|
||||
ORDER BY created_at DESC
|
||||
LIMIT $3
|
||||
"#;
|
||||
|
||||
const COUNT_FINALIZED_STATUSES_BY_ENDPOINT_IDS_SINCE_SQL: &str = r#"
|
||||
SELECT
|
||||
endpoint_id,
|
||||
status,
|
||||
COUNT(id) AS count
|
||||
FROM request_candidates
|
||||
WHERE endpoint_id = ANY($1)
|
||||
AND created_at >= TO_TIMESTAMP($2)
|
||||
AND status IN ('success', 'failed', 'skipped')
|
||||
GROUP BY endpoint_id, status
|
||||
"#;
|
||||
|
||||
const AGGREGATE_FINALIZED_TIMELINE_BY_ENDPOINT_IDS_SINCE_SQL: &str = r#"
|
||||
SELECT
|
||||
endpoint_id,
|
||||
@@ -325,13 +217,16 @@ impl SqlxRequestCandidateReadRepository {
|
||||
&self,
|
||||
request_id: &str,
|
||||
) -> Result<Vec<StoredRequestCandidate>, DataLayerError> {
|
||||
collect_query_rows(
|
||||
sqlx::query(LIST_BY_REQUEST_ID_SQL)
|
||||
.bind(request_id)
|
||||
.fetch(&self.pool),
|
||||
map_request_candidate_row,
|
||||
)
|
||||
.await
|
||||
let mut builder = QueryBuilder::<Postgres>::new(candidate_columns());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"request_id",
|
||||
request_id.to_string(),
|
||||
);
|
||||
builder.push(" ORDER BY candidate_index ASC, retry_index ASC, created_at ASC");
|
||||
collect_query_rows(builder.build().fetch(&self.pool), map_request_candidate_row).await
|
||||
}
|
||||
|
||||
pub async fn list_recent(
|
||||
@@ -342,17 +237,17 @@ impl SqlxRequestCandidateReadRepository {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
collect_query_rows(
|
||||
sqlx::query(LIST_RECENT_SQL)
|
||||
.bind(i64::try_from(limit).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid recent request candidate limit: {limit}"
|
||||
))
|
||||
})?)
|
||||
.fetch(&self.pool),
|
||||
map_request_candidate_row,
|
||||
)
|
||||
.await
|
||||
let mut builder = QueryBuilder::<Postgres>::new(candidate_columns());
|
||||
builder.push(" ORDER BY created_at DESC");
|
||||
push_limit(
|
||||
&mut builder,
|
||||
i64::try_from(limit).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid recent request candidate limit: {limit}"
|
||||
))
|
||||
})?,
|
||||
);
|
||||
collect_query_rows(builder.build().fetch(&self.pool), map_request_candidate_row).await
|
||||
}
|
||||
|
||||
pub async fn list_by_provider_id(
|
||||
@@ -364,20 +259,24 @@ impl SqlxRequestCandidateReadRepository {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let limit_value = i64::try_from(limit).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid provider request candidate limit: {limit}"
|
||||
))
|
||||
})?;
|
||||
|
||||
collect_query_rows(
|
||||
sqlx::query(LIST_BY_PROVIDER_ID_SQL)
|
||||
.bind(provider_id)
|
||||
.bind(limit_value)
|
||||
.fetch(&self.pool),
|
||||
map_request_candidate_row,
|
||||
)
|
||||
.await
|
||||
let mut builder = QueryBuilder::<Postgres>::new(candidate_columns());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"provider_id",
|
||||
provider_id.to_string(),
|
||||
);
|
||||
builder.push(" ORDER BY created_at DESC");
|
||||
push_limit(
|
||||
&mut builder,
|
||||
i64::try_from(limit).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid provider request candidate limit: {limit}"
|
||||
))
|
||||
})?,
|
||||
);
|
||||
collect_query_rows(builder.build().fetch(&self.pool), map_request_candidate_row).await
|
||||
}
|
||||
|
||||
pub async fn list_finalized_by_endpoint_ids_since(
|
||||
@@ -390,19 +289,22 @@ impl SqlxRequestCandidateReadRepository {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
collect_query_rows(
|
||||
sqlx::query(LIST_FINALIZED_BY_ENDPOINT_IDS_SINCE_SQL)
|
||||
.bind(endpoint_ids)
|
||||
.bind(since_unix_secs as f64)
|
||||
.bind(i64::try_from(limit).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid finalized request candidate limit: {limit}"
|
||||
))
|
||||
})?)
|
||||
.fetch(&self.pool),
|
||||
map_request_candidate_row,
|
||||
)
|
||||
.await
|
||||
let mut builder = QueryBuilder::<Postgres>::new(candidate_columns());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
|
||||
builder
|
||||
.push(" AND created_at >= TO_TIMESTAMP(")
|
||||
.push_bind(since_unix_secs as f64)
|
||||
.push(") AND status IN ('success', 'failed', 'skipped') ORDER BY created_at DESC");
|
||||
push_limit(
|
||||
&mut builder,
|
||||
i64::try_from(limit).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"invalid finalized request candidate limit: {limit}"
|
||||
))
|
||||
})?,
|
||||
);
|
||||
collect_query_rows(builder.build().fetch(&self.pool), map_request_candidate_row).await
|
||||
}
|
||||
|
||||
pub async fn count_finalized_statuses_by_endpoint_ids_since(
|
||||
@@ -414,29 +316,35 @@ impl SqlxRequestCandidateReadRepository {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let mut rows = sqlx::query(COUNT_FINALIZED_STATUSES_BY_ENDPOINT_IDS_SINCE_SQL)
|
||||
.bind(endpoint_ids)
|
||||
.bind(since_unix_secs as f64)
|
||||
.fetch(&self.pool);
|
||||
let mut counts = Vec::new();
|
||||
while let Some(row) = rows.try_next().await.map_postgres_err()? {
|
||||
let entry = {
|
||||
let status = RequestCandidateStatus::from_database(
|
||||
row_get::<String>(&row, "status")?.as_str(),
|
||||
)?;
|
||||
PublicHealthStatusCount {
|
||||
endpoint_id: row_get(&row, "endpoint_id")?,
|
||||
status,
|
||||
count: u64::try_from(row_get::<i64>(&row, "count")?).map_err(|_| {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(
|
||||
"SELECT endpoint_id, status, COUNT(id) AS count FROM request_candidates",
|
||||
);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
|
||||
builder
|
||||
.push(" AND created_at >= TO_TIMESTAMP(")
|
||||
.push_bind(since_unix_secs as f64)
|
||||
.push(") AND status IN ('success', 'failed', 'skipped') GROUP BY endpoint_id, status");
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter()
|
||||
.map(|row| {
|
||||
Ok(PublicHealthStatusCount {
|
||||
endpoint_id: row_get(row, "endpoint_id")?,
|
||||
status: RequestCandidateStatus::from_database(
|
||||
row_get::<String>(row, "status")?.as_str(),
|
||||
)?,
|
||||
count: u64::try_from(row_get::<i64>(row, "count")?).map_err(|_| {
|
||||
DataLayerError::UnexpectedValue(
|
||||
"public health status count out of range".to_string(),
|
||||
)
|
||||
})?,
|
||||
}
|
||||
};
|
||||
counts.push(entry);
|
||||
}
|
||||
Ok(counts)
|
||||
})
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
pub async fn aggregate_finalized_timeline_by_endpoint_ids_since(
|
||||
@@ -725,6 +633,13 @@ where
|
||||
row.try_get(column).map_postgres_err()
|
||||
}
|
||||
|
||||
fn candidate_columns() -> &'static str {
|
||||
LIST_BY_REQUEST_ID_SQL
|
||||
.split_once("WHERE request_id = $1")
|
||||
.map(|(prefix, _)| prefix)
|
||||
.unwrap_or(LIST_BY_REQUEST_ID_SQL)
|
||||
}
|
||||
|
||||
fn status_to_database(status: RequestCandidateStatus) -> &'static str {
|
||||
match status {
|
||||
RequestCandidateStatus::Available => "available",
|
||||
|
||||
@@ -11,6 +11,7 @@ use super::{
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{push_in, WhereClause};
|
||||
|
||||
const CANDIDATE_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
@@ -132,7 +133,8 @@ impl RequestCandidateReadRepository for SqliteRequestCandidateRepository {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_COLUMNS);
|
||||
push_endpoint_in_clause(&mut builder, endpoint_ids);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
|
||||
builder
|
||||
.push(" AND created_at >= ")
|
||||
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
|
||||
@@ -154,7 +156,8 @@ impl RequestCandidateReadRepository for SqliteRequestCandidateRepository {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(
|
||||
"SELECT endpoint_id, status, COUNT(id) AS count FROM request_candidates",
|
||||
);
|
||||
push_endpoint_in_clause(&mut builder, endpoint_ids);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
|
||||
builder
|
||||
.push(" AND created_at >= ")
|
||||
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
|
||||
@@ -193,7 +196,8 @@ impl RequestCandidateReadRepository for SqliteRequestCandidateRepository {
|
||||
let since_ms = unix_secs_to_ms_i64(since_unix_secs)?;
|
||||
let until_ms = unix_secs_to_ms_i64(until_unix_secs)?;
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(CANDIDATE_COLUMNS);
|
||||
push_endpoint_in_clause(&mut builder, endpoint_ids);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
|
||||
builder
|
||||
.push(" AND created_at >= ")
|
||||
.push_bind(since_ms)
|
||||
@@ -336,20 +340,6 @@ ON CONFLICT(request_id, candidate_index, retry_index) DO UPDATE SET
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn push_endpoint_in_clause<'args>(
|
||||
builder: &mut QueryBuilder<'args, Sqlite>,
|
||||
endpoint_ids: &'args [String],
|
||||
) {
|
||||
builder.push(" WHERE endpoint_id IN (");
|
||||
{
|
||||
let mut separated = builder.separated(", ");
|
||||
for endpoint_id in endpoint_ids {
|
||||
separated.push_bind(endpoint_id);
|
||||
}
|
||||
}
|
||||
builder.push(")");
|
||||
}
|
||||
|
||||
fn merge_candidate(
|
||||
candidate: UpsertRequestCandidateRecord,
|
||||
existing: Option<StoredRequestCandidate>,
|
||||
|
||||
@@ -8,6 +8,7 @@ use super::types::{
|
||||
StoredGeminiFileMappingListPage, UpsertGeminiFileMappingRecord,
|
||||
};
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
use aether_data_query::{push_ci_contains_any, push_limit_offset, SqlDialect, WhereClause};
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqlxGeminiFileMappingRepository {
|
||||
@@ -283,10 +284,10 @@ WHERE expires_at <= TO_TIMESTAMP($1::double precision)
|
||||
}
|
||||
|
||||
fn build_list_count_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, Postgres> {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(
|
||||
"SELECT COUNT(*)::bigint AS total FROM gemini_file_mappings WHERE 1=1",
|
||||
);
|
||||
apply_list_filters(&mut builder, query);
|
||||
let mut builder =
|
||||
QueryBuilder::<Postgres>::new("SELECT COUNT(*)::bigint AS total FROM gemini_file_mappings");
|
||||
let mut where_clause = WhereClause::new();
|
||||
apply_list_filters(&mut builder, &mut where_clause, query);
|
||||
builder
|
||||
}
|
||||
|
||||
@@ -304,23 +305,27 @@ SELECT
|
||||
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM expires_at)::bigint AS expires_at_unix_secs
|
||||
FROM gemini_file_mappings
|
||||
WHERE 1=1
|
||||
"#,
|
||||
);
|
||||
apply_list_filters(&mut builder, query);
|
||||
builder.push(" ORDER BY created_at DESC, file_name ASC LIMIT ");
|
||||
builder.push_bind(i64::try_from(query.limit).unwrap_or(i64::MAX));
|
||||
builder.push(" OFFSET ");
|
||||
builder.push_bind(i64::try_from(query.offset).unwrap_or(i64::MAX));
|
||||
let mut where_clause = WhereClause::new();
|
||||
apply_list_filters(&mut builder, &mut where_clause, query);
|
||||
builder.push(" ORDER BY created_at DESC, file_name ASC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64::try_from(query.limit).unwrap_or(i64::MAX),
|
||||
i64::try_from(query.offset).unwrap_or(i64::MAX),
|
||||
);
|
||||
builder
|
||||
}
|
||||
|
||||
fn apply_list_filters(
|
||||
builder: &mut QueryBuilder<'_, Postgres>,
|
||||
where_clause: &mut WhereClause,
|
||||
query: &GeminiFileMappingListQuery,
|
||||
) {
|
||||
if !query.include_expired {
|
||||
builder.push(" AND expires_at > TO_TIMESTAMP(");
|
||||
where_clause.push_next(builder);
|
||||
builder.push("expires_at > TO_TIMESTAMP(");
|
||||
builder.push_bind(query.now_unix_secs as f64);
|
||||
builder.push("::double precision)");
|
||||
}
|
||||
@@ -330,11 +335,12 @@ fn apply_list_filters(
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let pattern = format!("%{search}%");
|
||||
builder.push(" AND (file_name ILIKE ");
|
||||
builder.push_bind(pattern.clone());
|
||||
builder.push(" OR COALESCE(display_name, '') ILIKE ");
|
||||
builder.push_bind(pattern);
|
||||
builder.push(")");
|
||||
push_ci_contains_any(
|
||||
builder,
|
||||
where_clause,
|
||||
SqlDialect::Postgres,
|
||||
&["file_name", "COALESCE(display_name, '')"],
|
||||
search,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,6 +9,7 @@ use super::types::{
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{push_ci_contains_any, push_limit_offset, SqlDialect, WhereClause};
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqliteGeminiFileMappingRepository {
|
||||
@@ -242,8 +243,9 @@ LIMIT 1
|
||||
|
||||
fn build_list_count_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, Sqlite> {
|
||||
let mut builder =
|
||||
QueryBuilder::<Sqlite>::new("SELECT COUNT(*) AS total FROM gemini_file_mappings WHERE 1=1");
|
||||
apply_list_filters(&mut builder, query);
|
||||
QueryBuilder::<Sqlite>::new("SELECT COUNT(*) AS total FROM gemini_file_mappings");
|
||||
let mut where_clause = WhereClause::new();
|
||||
apply_list_filters(&mut builder, &mut where_clause, query);
|
||||
builder
|
||||
}
|
||||
|
||||
@@ -261,20 +263,27 @@ SELECT
|
||||
created_at AS created_at_unix_ms,
|
||||
expires_at AS expires_at_unix_secs
|
||||
FROM gemini_file_mappings
|
||||
WHERE 1=1
|
||||
"#,
|
||||
);
|
||||
apply_list_filters(&mut builder, query);
|
||||
builder.push(" ORDER BY created_at DESC, file_name ASC LIMIT ");
|
||||
builder.push_bind(i64::try_from(query.limit).unwrap_or(i64::MAX));
|
||||
builder.push(" OFFSET ");
|
||||
builder.push_bind(i64::try_from(query.offset).unwrap_or(i64::MAX));
|
||||
let mut where_clause = WhereClause::new();
|
||||
apply_list_filters(&mut builder, &mut where_clause, query);
|
||||
builder.push(" ORDER BY created_at DESC, file_name ASC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64::try_from(query.limit).unwrap_or(i64::MAX),
|
||||
i64::try_from(query.offset).unwrap_or(i64::MAX),
|
||||
);
|
||||
builder
|
||||
}
|
||||
|
||||
fn apply_list_filters(builder: &mut QueryBuilder<'_, Sqlite>, query: &GeminiFileMappingListQuery) {
|
||||
fn apply_list_filters(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
where_clause: &mut WhereClause,
|
||||
query: &GeminiFileMappingListQuery,
|
||||
) {
|
||||
if !query.include_expired {
|
||||
builder.push(" AND expires_at > ");
|
||||
where_clause.push_next(builder);
|
||||
builder.push("expires_at > ");
|
||||
builder.push_bind(query.now_unix_secs as i64);
|
||||
}
|
||||
if let Some(search) = query
|
||||
@@ -283,12 +292,13 @@ fn apply_list_filters(builder: &mut QueryBuilder<'_, Sqlite>, query: &GeminiFile
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let pattern = format!("%{}%", search.to_ascii_lowercase());
|
||||
builder.push(" AND (LOWER(file_name) LIKE ");
|
||||
builder.push_bind(pattern.clone());
|
||||
builder.push(" OR LOWER(COALESCE(display_name, '')) LIKE ");
|
||||
builder.push_bind(pattern);
|
||||
builder.push(")");
|
||||
push_ci_contains_any(
|
||||
builder,
|
||||
where_clause,
|
||||
SqlDialect::Sqlite,
|
||||
&["file_name", "COALESCE(display_name, '')"],
|
||||
search,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,50 +1,20 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{sqlite::SqliteRow, Row};
|
||||
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
|
||||
|
||||
use super::{
|
||||
metadata_supports_embedding, AdminGlobalModelListQuery, AdminProviderModelListQuery,
|
||||
CreateAdminGlobalModelRecord, GlobalModelReadRepository, GlobalModelWriteRepository,
|
||||
InMemoryGlobalModelReadRepository, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
|
||||
PublicGlobalModelQuery, StoredAdminGlobalModel, StoredAdminGlobalModelPage,
|
||||
StoredAdminProviderModel, StoredProviderActiveGlobalModel, StoredProviderModelStats,
|
||||
StoredPublicCatalogModel, StoredPublicGlobalModel, StoredPublicGlobalModelPage,
|
||||
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
PublicCatalogModelListQuery, PublicCatalogModelSearchQuery, PublicGlobalModelQuery,
|
||||
StoredAdminGlobalModel, StoredAdminGlobalModelPage, StoredAdminProviderModel,
|
||||
StoredProviderActiveGlobalModel, StoredProviderModelStats, StoredPublicCatalogModel,
|
||||
StoredPublicGlobalModel, StoredPublicGlobalModelPage, UpdateAdminGlobalModelRecord,
|
||||
UpsertAdminProviderModelRecord,
|
||||
};
|
||||
use crate::driver::sqlite::{sqlite_optional_real, SqlitePool};
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqliteGlobalModelReadRepository {
|
||||
pool: SqlitePool,
|
||||
}
|
||||
|
||||
impl SqliteGlobalModelReadRepository {
|
||||
pub fn new(pool: SqlitePool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn load_memory(&self) -> Result<InMemoryGlobalModelReadRepository, DataLayerError> {
|
||||
let public_models = self.load_public_global_models().await?;
|
||||
let admin_global_models = self.load_admin_global_models().await?;
|
||||
let admin_provider_models = self.load_admin_provider_models().await?;
|
||||
let public_catalog_models = self.load_public_catalog_models().await?;
|
||||
let provider_model_stats = self.load_provider_model_stats().await?;
|
||||
let active_global_model_refs = self.load_active_global_model_refs().await?;
|
||||
|
||||
Ok(InMemoryGlobalModelReadRepository::seed(public_models)
|
||||
.with_admin_global_models(admin_global_models)
|
||||
.with_admin_provider_models(admin_provider_models)
|
||||
.with_public_catalog_models(public_catalog_models)
|
||||
.with_provider_model_stats(provider_model_stats)
|
||||
.with_active_global_model_refs(active_global_model_refs))
|
||||
}
|
||||
|
||||
async fn load_public_global_models(
|
||||
&self,
|
||||
) -> Result<Vec<StoredPublicGlobalModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
const LIST_PUBLIC_GLOBAL_MODELS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
name,
|
||||
@@ -54,47 +24,73 @@ SELECT
|
||||
default_tiered_pricing,
|
||||
supported_capabilities,
|
||||
config,
|
||||
usage_count
|
||||
0 AS usage_count
|
||||
FROM global_models
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_public_global_model_row).collect()
|
||||
}
|
||||
"#;
|
||||
|
||||
async fn load_admin_global_models(
|
||||
&self,
|
||||
) -> Result<Vec<StoredAdminGlobalModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
const COUNT_PUBLIC_GLOBAL_MODELS_PREFIX: &str = r#"
|
||||
SELECT COUNT(id) AS total
|
||||
FROM global_models
|
||||
"#;
|
||||
|
||||
const LIST_PUBLIC_CATALOG_MODELS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
name,
|
||||
COALESCE(NULLIF(display_name, ''), name) AS display_name,
|
||||
is_active,
|
||||
CAST(default_price_per_request AS REAL) AS default_price_per_request,
|
||||
default_tiered_pricing,
|
||||
supported_capabilities,
|
||||
config,
|
||||
usage_count,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM global_models
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_admin_global_model_row).collect()
|
||||
}
|
||||
m.id,
|
||||
m.provider_id,
|
||||
p.name AS provider_name,
|
||||
p.is_active AS provider_is_active,
|
||||
m.provider_model_name,
|
||||
COALESCE(gm.name, m.provider_model_name) AS name,
|
||||
COALESCE(NULLIF(gm.display_name, ''), m.provider_model_name) AS display_name,
|
||||
gm.config AS global_model_config,
|
||||
gm.supported_capabilities AS global_model_supported_capabilities,
|
||||
m.config AS model_config,
|
||||
m.tiered_pricing,
|
||||
gm.default_tiered_pricing,
|
||||
COALESCE(
|
||||
m.supports_vision,
|
||||
CASE
|
||||
WHEN json_extract(gm.config, '$.vision') IS NULL THEN NULL
|
||||
WHEN LOWER(CAST(json_extract(gm.config, '$.vision') AS TEXT)) IN ('true', '1') THEN 1
|
||||
ELSE 0
|
||||
END,
|
||||
0
|
||||
) AS supports_vision,
|
||||
COALESCE(
|
||||
m.supports_function_calling,
|
||||
CASE
|
||||
WHEN json_extract(gm.config, '$.function_calling') IS NULL THEN NULL
|
||||
WHEN LOWER(CAST(json_extract(gm.config, '$.function_calling') AS TEXT)) IN ('true', '1') THEN 1
|
||||
ELSE 0
|
||||
END,
|
||||
0
|
||||
) AS supports_function_calling,
|
||||
COALESCE(
|
||||
m.supports_streaming,
|
||||
CASE
|
||||
WHEN json_extract(gm.config, '$.streaming') IS NULL THEN NULL
|
||||
WHEN LOWER(CAST(json_extract(gm.config, '$.streaming') AS TEXT)) IN ('true', '1') THEN 1
|
||||
ELSE 0
|
||||
END,
|
||||
1
|
||||
) AS supports_streaming,
|
||||
m.is_active,
|
||||
gm.is_active AS global_model_is_active
|
||||
FROM models m
|
||||
JOIN providers p ON p.id = m.provider_id
|
||||
LEFT JOIN global_models gm ON gm.id = m.global_model_id
|
||||
"#;
|
||||
|
||||
async fn load_admin_provider_models(
|
||||
&self,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
const LIST_PROVIDER_MODEL_STATS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
provider_id,
|
||||
COUNT(id) AS total_models,
|
||||
COALESCE(SUM(CASE WHEN is_active = 1 THEN 1 ELSE 0 END), 0) AS active_models
|
||||
FROM models
|
||||
WHERE provider_id IN (
|
||||
"#;
|
||||
|
||||
const LIST_ADMIN_PROVIDER_MODELS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
m.id,
|
||||
m.provider_id,
|
||||
@@ -109,7 +105,7 @@ SELECT
|
||||
m.supports_extended_thinking,
|
||||
m.supports_image_generation,
|
||||
m.is_active,
|
||||
m.is_available,
|
||||
COALESCE(m.is_available, 1) AS is_available,
|
||||
m.config,
|
||||
m.created_at AS created_at_unix_ms,
|
||||
m.updated_at AS updated_at_unix_secs,
|
||||
@@ -121,85 +117,61 @@ SELECT
|
||||
gm.config AS global_model_config
|
||||
FROM models m
|
||||
LEFT JOIN global_models gm ON gm.id = m.global_model_id
|
||||
WHERE m.global_model_id IS NOT NULL
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_admin_provider_model_row).collect()
|
||||
}
|
||||
"#;
|
||||
|
||||
async fn load_public_catalog_models(
|
||||
&self,
|
||||
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
const LIST_ADMIN_GLOBAL_MODELS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
m.id,
|
||||
m.provider_id,
|
||||
p.name AS provider_name,
|
||||
p.is_active AS provider_is_active,
|
||||
m.provider_model_name,
|
||||
COALESCE(gm.name, m.provider_model_name) AS name,
|
||||
COALESCE(NULLIF(gm.display_name, ''), m.provider_model_name) AS display_name,
|
||||
gm.config AS global_model_config,
|
||||
gm.supported_capabilities AS global_model_supported_capabilities,
|
||||
m.config AS model_config,
|
||||
m.tiered_pricing,
|
||||
gm.id,
|
||||
gm.name,
|
||||
COALESCE(NULLIF(gm.display_name, ''), gm.name) AS display_name,
|
||||
gm.is_active,
|
||||
CAST(gm.default_price_per_request AS REAL) AS default_price_per_request,
|
||||
gm.default_tiered_pricing,
|
||||
m.supports_vision,
|
||||
m.supports_function_calling,
|
||||
m.supports_streaming,
|
||||
m.is_active,
|
||||
gm.is_active AS global_model_is_active
|
||||
FROM models m
|
||||
JOIN providers p ON p.id = m.provider_id
|
||||
LEFT JOIN global_models gm ON gm.id = m.global_model_id
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_public_catalog_model_row).collect()
|
||||
}
|
||||
gm.supported_capabilities,
|
||||
gm.config,
|
||||
COALESCE(gm_stats.provider_count, 0) AS provider_count,
|
||||
COALESCE(gm_stats.active_provider_count, 0) AS active_provider_count,
|
||||
COALESCE(gm.usage_count, 0) AS usage_count,
|
||||
gm.created_at AS created_at_unix_ms,
|
||||
gm.updated_at AS updated_at_unix_secs
|
||||
FROM global_models gm
|
||||
LEFT JOIN (
|
||||
SELECT
|
||||
m.global_model_id,
|
||||
COUNT(DISTINCT m.provider_id) AS provider_count,
|
||||
COUNT(
|
||||
DISTINCT CASE
|
||||
WHEN m.is_active = 1 AND COALESCE(m.is_available, 1) = 1 AND p.is_active = 1 THEN m.provider_id
|
||||
ELSE NULL
|
||||
END
|
||||
) AS active_provider_count
|
||||
FROM models m
|
||||
JOIN providers p ON p.id = m.provider_id
|
||||
GROUP BY m.global_model_id
|
||||
) gm_stats ON gm_stats.global_model_id = gm.id
|
||||
"#;
|
||||
|
||||
async fn load_provider_model_stats(
|
||||
&self,
|
||||
) -> Result<Vec<StoredProviderModelStats>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
const COUNT_ADMIN_GLOBAL_MODELS_PREFIX: &str = r#"
|
||||
SELECT COUNT(id) AS total
|
||||
FROM global_models gm
|
||||
"#;
|
||||
|
||||
const LIST_ACTIVE_GLOBAL_MODEL_IDS_BY_PROVIDER_IDS_PREFIX: &str = r#"
|
||||
SELECT DISTINCT
|
||||
provider_id,
|
||||
COUNT(id) AS total_models,
|
||||
SUM(CASE WHEN is_active = 1 THEN 1 ELSE 0 END) AS active_models
|
||||
global_model_id
|
||||
FROM models
|
||||
GROUP BY provider_id
|
||||
ORDER BY provider_id ASC
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_provider_model_stats_row).collect()
|
||||
}
|
||||
WHERE provider_id IN (
|
||||
"#;
|
||||
|
||||
async fn load_active_global_model_refs(
|
||||
&self,
|
||||
) -> Result<Vec<StoredProviderActiveGlobalModel>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT DISTINCT provider_id, global_model_id
|
||||
FROM models
|
||||
WHERE is_active = 1
|
||||
AND global_model_id IS NOT NULL
|
||||
ORDER BY provider_id ASC, global_model_id ASC
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_active_global_model_row).collect()
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqliteGlobalModelReadRepository {
|
||||
pool: SqlitePool,
|
||||
}
|
||||
|
||||
impl SqliteGlobalModelReadRepository {
|
||||
pub fn new(pool: SqlitePool) -> Self {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
pub async fn create_admin_provider_model(
|
||||
@@ -480,67 +452,172 @@ impl GlobalModelReadRepository for SqliteGlobalModelReadRepository {
|
||||
&self,
|
||||
query: &PublicGlobalModelQuery,
|
||||
) -> Result<StoredPublicGlobalModelPage, DataLayerError> {
|
||||
self.load_memory().await?.list_public_models(query).await
|
||||
let mut count_builder = QueryBuilder::<Sqlite>::new(COUNT_PUBLIC_GLOBAL_MODELS_PREFIX);
|
||||
apply_public_model_filters(&mut count_builder, query);
|
||||
let count_row = count_builder
|
||||
.build()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let total = count_row
|
||||
.try_get::<i64, _>("total")
|
||||
.map(|value| value.max(0) as usize)
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut list_builder = QueryBuilder::<Sqlite>::new(LIST_PUBLIC_GLOBAL_MODELS_PREFIX);
|
||||
apply_public_model_filters(&mut list_builder, query);
|
||||
list_builder
|
||||
.push(" ORDER BY name ASC LIMIT ")
|
||||
.push_bind(query.limit as i64)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(query.offset as i64);
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_public_global_model_row)
|
||||
.collect::<Result<_, _>>()?;
|
||||
|
||||
Ok(StoredPublicGlobalModelPage { items, total })
|
||||
}
|
||||
|
||||
async fn get_public_model_by_name(
|
||||
&self,
|
||||
model_name: &str,
|
||||
) -> Result<Option<StoredPublicGlobalModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.get_public_model_by_name(model_name)
|
||||
.await
|
||||
let row = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id,
|
||||
name,
|
||||
display_name,
|
||||
is_active,
|
||||
CAST(default_price_per_request AS REAL) AS default_price_per_request,
|
||||
default_tiered_pricing,
|
||||
supported_capabilities,
|
||||
config,
|
||||
0 AS usage_count
|
||||
FROM global_models
|
||||
WHERE name = ? AND is_active = 1
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(model_name)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
row.as_ref().map(map_public_global_model_row).transpose()
|
||||
}
|
||||
|
||||
async fn list_public_catalog_models(
|
||||
&self,
|
||||
query: &PublicCatalogModelListQuery,
|
||||
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_public_catalog_models(query)
|
||||
.await
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(LIST_PUBLIC_CATALOG_MODELS_PREFIX);
|
||||
apply_public_catalog_model_filters(&mut builder, query.provider_id.as_deref(), None);
|
||||
builder
|
||||
.push(" ORDER BY p.provider_priority ASC, p.name ASC, COALESCE(gm.name, m.provider_model_name) ASC, m.id ASC LIMIT ")
|
||||
.push_bind(query.limit as i64)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(query.offset as i64);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_public_catalog_model_row).collect()
|
||||
}
|
||||
|
||||
async fn search_public_catalog_models(
|
||||
&self,
|
||||
query: &PublicCatalogModelSearchQuery,
|
||||
) -> Result<Vec<StoredPublicCatalogModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.search_public_catalog_models(query)
|
||||
.await
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(LIST_PUBLIC_CATALOG_MODELS_PREFIX);
|
||||
apply_public_catalog_model_filters(
|
||||
&mut builder,
|
||||
query.provider_id.as_deref(),
|
||||
Some(query.search.as_str()),
|
||||
);
|
||||
builder
|
||||
.push(" ORDER BY p.provider_priority ASC, p.name ASC, COALESCE(gm.name, m.provider_model_name) ASC, m.id ASC LIMIT ")
|
||||
.push_bind(query.limit as i64);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_public_catalog_model_row).collect()
|
||||
}
|
||||
|
||||
async fn list_admin_global_models(
|
||||
&self,
|
||||
query: &AdminGlobalModelListQuery,
|
||||
) -> Result<StoredAdminGlobalModelPage, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_admin_global_models(query)
|
||||
let mut count_builder = QueryBuilder::<Sqlite>::new(COUNT_ADMIN_GLOBAL_MODELS_PREFIX);
|
||||
apply_admin_global_model_filters(&mut count_builder, query);
|
||||
let count_row = count_builder
|
||||
.build()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let total = count_row
|
||||
.try_get::<i64, _>("total")
|
||||
.map(|value| value.max(0) as usize)
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut list_builder = QueryBuilder::<Sqlite>::new(LIST_ADMIN_GLOBAL_MODELS_PREFIX);
|
||||
apply_admin_global_model_filters(&mut list_builder, query);
|
||||
list_builder
|
||||
.push(" ORDER BY name ASC LIMIT ")
|
||||
.push_bind(query.limit as i64)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(query.offset as i64);
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_admin_global_model_row)
|
||||
.collect::<Result<_, _>>()?;
|
||||
Ok(StoredAdminGlobalModelPage { items, total })
|
||||
}
|
||||
|
||||
async fn list_admin_provider_models(
|
||||
&self,
|
||||
query: &AdminProviderModelListQuery,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_admin_provider_models(query)
|
||||
.await
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(LIST_ADMIN_PROVIDER_MODELS_PREFIX);
|
||||
builder
|
||||
.push(" WHERE m.provider_id = ")
|
||||
.push_bind(query.provider_id.trim().to_string());
|
||||
if let Some(is_active) = query.is_active {
|
||||
builder.push(" AND m.is_active = ").push_bind(is_active);
|
||||
}
|
||||
builder
|
||||
.push(" ORDER BY m.created_at DESC, m.id ASC LIMIT ")
|
||||
.push_bind(query.limit as i64)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(query.offset as i64);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_admin_provider_model_row).collect()
|
||||
}
|
||||
|
||||
async fn list_admin_provider_available_source_models(
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_admin_provider_available_source_models(provider_id)
|
||||
.await
|
||||
let rows = sqlx::query(&format!(
|
||||
r#"
|
||||
{LIST_ADMIN_PROVIDER_MODELS_PREFIX}
|
||||
WHERE m.provider_id = ?
|
||||
AND m.is_active = 1
|
||||
AND gm.is_active = 1
|
||||
ORDER BY gm.name ASC, m.created_at DESC, m.id ASC
|
||||
"#
|
||||
))
|
||||
.bind(provider_id)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_admin_provider_model_row).collect()
|
||||
}
|
||||
|
||||
async fn get_admin_provider_model(
|
||||
@@ -548,60 +625,111 @@ impl GlobalModelReadRepository for SqliteGlobalModelReadRepository {
|
||||
provider_id: &str,
|
||||
model_id: &str,
|
||||
) -> Result<Option<StoredAdminProviderModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.get_admin_provider_model(provider_id, model_id)
|
||||
.await
|
||||
let row = sqlx::query(&format!(
|
||||
r#"
|
||||
{LIST_ADMIN_PROVIDER_MODELS_PREFIX}
|
||||
WHERE m.provider_id = ?
|
||||
AND m.id = ?
|
||||
LIMIT 1
|
||||
"#
|
||||
))
|
||||
.bind(provider_id)
|
||||
.bind(model_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
row.as_ref().map(map_admin_provider_model_row).transpose()
|
||||
}
|
||||
|
||||
async fn get_admin_global_model_by_id(
|
||||
&self,
|
||||
global_model_id: &str,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.get_admin_global_model_by_id(global_model_id)
|
||||
.await
|
||||
let row = sqlx::query(&format!(
|
||||
r#"
|
||||
{LIST_ADMIN_GLOBAL_MODELS_PREFIX}
|
||||
WHERE gm.id = ?
|
||||
LIMIT 1
|
||||
"#
|
||||
))
|
||||
.bind(global_model_id)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
row.as_ref().map(map_admin_global_model_row).transpose()
|
||||
}
|
||||
|
||||
async fn get_admin_global_model_by_name(
|
||||
&self,
|
||||
model_name: &str,
|
||||
) -> Result<Option<StoredAdminGlobalModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.get_admin_global_model_by_name(model_name)
|
||||
.await
|
||||
let row = sqlx::query(&format!(
|
||||
r#"
|
||||
{LIST_ADMIN_GLOBAL_MODELS_PREFIX}
|
||||
WHERE gm.name = ?
|
||||
LIMIT 1
|
||||
"#
|
||||
))
|
||||
.bind(model_name)
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
row.as_ref().map(map_admin_global_model_row).transpose()
|
||||
}
|
||||
|
||||
async fn list_admin_provider_models_by_global_model_id(
|
||||
&self,
|
||||
global_model_id: &str,
|
||||
) -> Result<Vec<StoredAdminProviderModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_admin_provider_models_by_global_model_id(global_model_id)
|
||||
.await
|
||||
let rows = sqlx::query(&format!(
|
||||
r#"
|
||||
{LIST_ADMIN_PROVIDER_MODELS_PREFIX}
|
||||
WHERE m.global_model_id = ?
|
||||
ORDER BY m.created_at DESC, m.id ASC
|
||||
"#
|
||||
))
|
||||
.bind(global_model_id)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_admin_provider_model_row).collect()
|
||||
}
|
||||
|
||||
async fn list_provider_model_stats(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderModelStats>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_provider_model_stats(provider_ids)
|
||||
.await
|
||||
if provider_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let mut builder = build_provider_id_list_query(
|
||||
LIST_PROVIDER_MODEL_STATS_PREFIX,
|
||||
provider_ids,
|
||||
")\nGROUP BY provider_id\nORDER BY provider_id ASC",
|
||||
);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_provider_model_stats_row).collect()
|
||||
}
|
||||
|
||||
async fn list_active_global_model_ids_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderActiveGlobalModel>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_active_global_model_ids_by_provider_ids(provider_ids)
|
||||
.await
|
||||
if provider_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let mut builder = build_provider_id_list_query(
|
||||
LIST_ACTIVE_GLOBAL_MODEL_IDS_BY_PROVIDER_IDS_PREFIX,
|
||||
provider_ids,
|
||||
")\nAND is_active = 1\nAND global_model_id IS NOT NULL\nORDER BY provider_id ASC, global_model_id ASC",
|
||||
);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_active_global_model_row).collect()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -705,6 +833,100 @@ fn first_tier_price(value: Option<&serde_json::Value>, key: &str) -> Option<f64>
|
||||
.and_then(serde_json::Value::as_f64)
|
||||
}
|
||||
|
||||
fn apply_public_model_filters(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
query: &PublicGlobalModelQuery,
|
||||
) {
|
||||
builder.push(" WHERE ");
|
||||
match query.is_active {
|
||||
Some(is_active) => {
|
||||
builder.push("is_active = ").push_bind(is_active);
|
||||
}
|
||||
None => {
|
||||
builder.push("is_active = 1");
|
||||
}
|
||||
}
|
||||
|
||||
if let Some(search) = query
|
||||
.search
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let pattern = format!("%{}%", search.to_ascii_lowercase());
|
||||
builder
|
||||
.push(" AND (LOWER(name) LIKE ")
|
||||
.push_bind(pattern.clone())
|
||||
.push(" OR LOWER(display_name) LIKE ")
|
||||
.push_bind(pattern)
|
||||
.push(")");
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_admin_global_model_filters(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
query: &AdminGlobalModelListQuery,
|
||||
) {
|
||||
builder.push(" WHERE 1=1");
|
||||
if let Some(is_active) = query.is_active {
|
||||
builder.push(" AND gm.is_active = ").push_bind(is_active);
|
||||
}
|
||||
if let Some(search) = query
|
||||
.search
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
{
|
||||
let pattern = format!("%{}%", search.to_ascii_lowercase());
|
||||
builder
|
||||
.push(" AND (LOWER(gm.name) LIKE ")
|
||||
.push_bind(pattern.clone())
|
||||
.push(" OR LOWER(gm.display_name) LIKE ")
|
||||
.push_bind(pattern)
|
||||
.push(")");
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_public_catalog_model_filters(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
provider_id: Option<&str>,
|
||||
search: Option<&str>,
|
||||
) {
|
||||
builder.push(" WHERE m.is_active = 1 AND COALESCE(m.is_available, 1) = 1 AND p.is_active = 1 AND COALESCE(gm.is_active, 1) = 1");
|
||||
|
||||
if let Some(provider_id) = provider_id.map(str::trim).filter(|value| !value.is_empty()) {
|
||||
builder
|
||||
.push(" AND m.provider_id = ")
|
||||
.push_bind(provider_id.to_string());
|
||||
}
|
||||
|
||||
if let Some(search) = search.map(str::trim).filter(|value| !value.is_empty()) {
|
||||
let pattern = format!("%{}%", search.to_ascii_lowercase());
|
||||
builder
|
||||
.push(" AND (LOWER(m.provider_model_name) LIKE ")
|
||||
.push_bind(pattern.clone())
|
||||
.push(" OR LOWER(gm.name) LIKE ")
|
||||
.push_bind(pattern.clone())
|
||||
.push(" OR LOWER(gm.display_name) LIKE ")
|
||||
.push_bind(pattern)
|
||||
.push(")");
|
||||
}
|
||||
}
|
||||
|
||||
fn build_provider_id_list_query<'a>(
|
||||
prefix: &'static str,
|
||||
provider_ids: &'a [String],
|
||||
suffix: &'static str,
|
||||
) -> QueryBuilder<'a, Sqlite> {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(prefix);
|
||||
let mut separated = builder.separated(", ");
|
||||
for provider_id in provider_ids {
|
||||
separated.push_bind(provider_id);
|
||||
}
|
||||
separated.push_unseparated(suffix);
|
||||
builder
|
||||
}
|
||||
|
||||
fn map_public_global_model_row(row: &SqliteRow) -> Result<StoredPublicGlobalModel, DataLayerError> {
|
||||
StoredPublicGlobalModel::new(
|
||||
row.try_get("id").map_sql_err()?,
|
||||
@@ -726,6 +948,16 @@ fn map_public_global_model_row(row: &SqliteRow) -> Result<StoredPublicGlobalMode
|
||||
}
|
||||
|
||||
fn map_admin_global_model_row(row: &SqliteRow) -> Result<StoredAdminGlobalModel, DataLayerError> {
|
||||
let provider_count = row
|
||||
.try_get::<i64, _>("provider_count")
|
||||
.map_sql_err()?
|
||||
.max(0) as u64;
|
||||
let active_provider_count = row
|
||||
.try_get::<i64, _>("active_provider_count")
|
||||
.map_sql_err()?
|
||||
.max(0) as u64;
|
||||
let usage_count = row.try_get::<i64, _>("usage_count").map_sql_err()?.max(0) as u64;
|
||||
|
||||
StoredAdminGlobalModel::new(
|
||||
row.try_get("id").map_sql_err()?,
|
||||
row.try_get("name").map_sql_err()?,
|
||||
@@ -741,9 +973,9 @@ fn map_admin_global_model_row(row: &SqliteRow) -> Result<StoredAdminGlobalModel,
|
||||
"global_models.supported_capabilities",
|
||||
)?,
|
||||
optional_json_from_string(row.try_get("config").map_sql_err()?, "global_models.config")?,
|
||||
0,
|
||||
0,
|
||||
row.try_get::<i64, _>("usage_count").map_sql_err()?.max(0) as u64,
|
||||
provider_count,
|
||||
active_provider_count,
|
||||
usage_count,
|
||||
optional_u64(
|
||||
row.try_get("created_at_unix_ms").map_sql_err()?,
|
||||
"global_models.created_at",
|
||||
@@ -898,8 +1130,8 @@ mod tests {
|
||||
use crate::lifecycle::migrate::run_sqlite_migrations;
|
||||
use crate::repository::global_models::{
|
||||
AdminGlobalModelListQuery, AdminProviderModelListQuery, CreateAdminGlobalModelRecord,
|
||||
GlobalModelReadRepository, PublicCatalogModelSearchQuery, PublicGlobalModelQuery,
|
||||
UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
GlobalModelReadRepository, PublicCatalogModelListQuery, PublicCatalogModelSearchQuery,
|
||||
PublicGlobalModelQuery, UpdateAdminGlobalModelRecord, UpsertAdminProviderModelRecord,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
@@ -939,6 +1171,18 @@ mod tests {
|
||||
assert_eq!(catalog.len(), 1);
|
||||
assert_eq!(catalog[0].input_price_per_1m, Some(2.0));
|
||||
|
||||
let catalog_list = repository
|
||||
.list_public_catalog_models(&PublicCatalogModelListQuery {
|
||||
provider_id: None,
|
||||
offset: 0,
|
||||
limit: 10,
|
||||
})
|
||||
.await
|
||||
.expect("catalog list should load");
|
||||
assert_eq!(catalog_list.len(), 2);
|
||||
assert_eq!(catalog_list[0].provider_id, "provider-1");
|
||||
assert_eq!(catalog_list[1].provider_id, "provider-3");
|
||||
|
||||
let admin_globals = repository
|
||||
.list_admin_global_models(&AdminGlobalModelListQuery {
|
||||
offset: 0,
|
||||
@@ -949,7 +1193,8 @@ mod tests {
|
||||
.await
|
||||
.expect("admin globals should load");
|
||||
assert_eq!(admin_globals.total, 1);
|
||||
assert_eq!(admin_globals.items[0].provider_count, 1);
|
||||
assert_eq!(admin_globals.items[0].provider_count, 3);
|
||||
assert_eq!(admin_globals.items[0].active_provider_count, 2);
|
||||
|
||||
let admin_models = repository
|
||||
.list_admin_provider_models(&AdminProviderModelListQuery {
|
||||
@@ -1113,6 +1358,18 @@ mod tests {
|
||||
seed_provider(pool).await;
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO providers (
|
||||
id, name, provider_type, is_active, provider_priority, created_at, updated_at
|
||||
) VALUES
|
||||
('provider-2', 'Inactive Provider', 'custom', 0, 1, 1, 1),
|
||||
('provider-3', 'Alpha Provider', 'custom', 1, 20, 1, 1)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("extra providers should seed");
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO global_models (
|
||||
id, name, display_name, is_active, default_tiered_pricing,
|
||||
supported_capabilities, usage_count, config, created_at, updated_at
|
||||
@@ -1132,10 +1389,20 @@ INSERT INTO models (
|
||||
id, provider_id, global_model_id, provider_model_name, provider_model_mappings,
|
||||
supports_vision, supports_function_calling, supports_streaming, is_active,
|
||||
is_available, created_at, updated_at
|
||||
) VALUES (
|
||||
) VALUES
|
||||
(
|
||||
'model-1', 'provider-1', 'global-1', 'provider-gpt-4.1', '["gpt-4.1"]',
|
||||
1, 1, 1, 1, 1, 4, 5
|
||||
)
|
||||
,
|
||||
(
|
||||
'model-2', 'provider-2', 'global-1', 'inactive-provider-gpt-4.1', '["gpt-4.1"]',
|
||||
1, 1, 1, 1, 1, 6, 7
|
||||
),
|
||||
(
|
||||
'model-3', 'provider-3', 'global-1', 'alpha-provider-gpt-4.1', '["gpt-4.1"]',
|
||||
1, 1, 1, 1, 1, 8, 9
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
@@ -1147,9 +1414,9 @@ INSERT INTO models (
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO providers (
|
||||
id, name, provider_type, is_active, created_at, updated_at
|
||||
id, name, provider_type, is_active, provider_priority, created_at, updated_at
|
||||
) VALUES (
|
||||
'provider-1', 'Provider One', 'custom', 1, 1, 1
|
||||
'provider-1', 'Zulu Provider', 'custom', 1, 10, 1, 1
|
||||
)
|
||||
"#,
|
||||
)
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
use async_trait::async_trait;
|
||||
use futures_util::TryStreamExt;
|
||||
use sqlx::{postgres::PgRow, PgPool, Row};
|
||||
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
|
||||
|
||||
use super::types::{
|
||||
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
|
||||
@@ -9,8 +8,9 @@ use super::types::{
|
||||
UpdateManagementTokenRecord,
|
||||
};
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
use aether_data_query::{push_eq, push_limit, push_limit_offset, push_optional_eq, WhereClause};
|
||||
|
||||
const LIST_MANAGEMENT_TOKENS_SQL: &str = r#"
|
||||
const MANAGEMENT_TOKEN_WITH_USER_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
mt.id,
|
||||
mt.user_id,
|
||||
@@ -32,70 +32,6 @@ SELECT
|
||||
u.role::text AS user_role
|
||||
FROM management_tokens mt
|
||||
JOIN users u ON u.id = mt.user_id
|
||||
WHERE ($1::text IS NULL OR mt.user_id = $1)
|
||||
AND ($2::boolean IS NULL OR mt.is_active = $2)
|
||||
ORDER BY mt.created_at DESC, mt.id DESC
|
||||
OFFSET $3
|
||||
LIMIT $4
|
||||
"#;
|
||||
|
||||
const COUNT_MANAGEMENT_TOKENS_SQL: &str = r#"
|
||||
SELECT COUNT(mt.id) AS total
|
||||
FROM management_tokens mt
|
||||
WHERE ($1::text IS NULL OR mt.user_id = $1)
|
||||
AND ($2::boolean IS NULL OR mt.is_active = $2)
|
||||
"#;
|
||||
|
||||
const GET_MANAGEMENT_TOKEN_WITH_USER_SQL: &str = r#"
|
||||
SELECT
|
||||
mt.id,
|
||||
mt.user_id,
|
||||
mt.name,
|
||||
mt.description,
|
||||
mt.token_prefix,
|
||||
mt.allowed_ips,
|
||||
mt.permissions,
|
||||
EXTRACT(EPOCH FROM mt.expires_at)::bigint AS expires_at_unix_secs,
|
||||
EXTRACT(EPOCH FROM mt.last_used_at)::bigint AS last_used_at_unix_secs,
|
||||
mt.last_used_ip,
|
||||
COALESCE(mt.usage_count, 0) AS usage_count,
|
||||
mt.is_active,
|
||||
EXTRACT(EPOCH FROM mt.created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM mt.updated_at)::bigint AS updated_at_unix_secs,
|
||||
u.id AS user_row_id,
|
||||
u.email AS user_email,
|
||||
u.username AS user_username,
|
||||
u.role::text AS user_role
|
||||
FROM management_tokens mt
|
||||
JOIN users u ON u.id = mt.user_id
|
||||
WHERE mt.id = $1
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const GET_MANAGEMENT_TOKEN_WITH_USER_BY_HASH_SQL: &str = r#"
|
||||
SELECT
|
||||
mt.id,
|
||||
mt.user_id,
|
||||
mt.name,
|
||||
mt.description,
|
||||
mt.token_prefix,
|
||||
mt.allowed_ips,
|
||||
mt.permissions,
|
||||
EXTRACT(EPOCH FROM mt.expires_at)::bigint AS expires_at_unix_secs,
|
||||
EXTRACT(EPOCH FROM mt.last_used_at)::bigint AS last_used_at_unix_secs,
|
||||
mt.last_used_ip,
|
||||
COALESCE(mt.usage_count, 0) AS usage_count,
|
||||
mt.is_active,
|
||||
EXTRACT(EPOCH FROM mt.created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM mt.updated_at)::bigint AS updated_at_unix_secs,
|
||||
u.id AS user_row_id,
|
||||
u.email AS user_email,
|
||||
u.username AS user_username,
|
||||
u.role::text AS user_role
|
||||
FROM management_tokens mt
|
||||
JOIN users u ON u.id = mt.user_id
|
||||
WHERE mt.token_hash = $1
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const DELETE_MANAGEMENT_TOKEN_SQL: &str = r#"
|
||||
@@ -362,24 +298,34 @@ impl ManagementTokenReadRepository for SqlxManagementTokenRepository {
|
||||
&self,
|
||||
query: &ManagementTokenListQuery,
|
||||
) -> Result<StoredManagementTokenListPage, DataLayerError> {
|
||||
let count_row = sqlx::query(COUNT_MANAGEMENT_TOKENS_SQL)
|
||||
.bind(query.user_id.as_deref())
|
||||
.bind(query.is_active)
|
||||
let mut count_builder =
|
||||
QueryBuilder::<Postgres>::new("SELECT COUNT(mt.id) AS total FROM management_tokens mt");
|
||||
let mut count_where = WhereClause::new();
|
||||
apply_management_token_filters(&mut count_builder, &mut count_where, query);
|
||||
let total = count_builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = count_row.try_get::<i64, _>("total").map_postgres_err()?;
|
||||
|
||||
let mut rows = sqlx::query(LIST_MANAGEMENT_TOKENS_SQL)
|
||||
.bind(query.user_id.as_deref())
|
||||
.bind(query.is_active)
|
||||
.bind(i64::try_from(query.offset).unwrap_or(i64::MAX))
|
||||
.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(map_token_with_user_row(&row)?);
|
||||
}
|
||||
let mut list_builder = QueryBuilder::<Postgres>::new(MANAGEMENT_TOKEN_WITH_USER_COLUMNS);
|
||||
let mut list_where = WhereClause::new();
|
||||
apply_management_token_filters(&mut list_builder, &mut list_where, query);
|
||||
list_builder.push(" ORDER BY mt.created_at DESC, mt.id DESC");
|
||||
push_limit_offset(
|
||||
&mut list_builder,
|
||||
i64::try_from(query.limit).unwrap_or(i64::MAX),
|
||||
i64::try_from(query.offset).unwrap_or(i64::MAX),
|
||||
);
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_token_with_user_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
|
||||
Ok(StoredManagementTokenListPage {
|
||||
items,
|
||||
@@ -391,8 +337,17 @@ impl ManagementTokenReadRepository for SqlxManagementTokenRepository {
|
||||
&self,
|
||||
token_id: &str,
|
||||
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
|
||||
let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_SQL)
|
||||
.bind(token_id)
|
||||
let mut builder = QueryBuilder::<Postgres>::new(MANAGEMENT_TOKEN_WITH_USER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"mt.id",
|
||||
token_id.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
@@ -403,8 +358,17 @@ impl ManagementTokenReadRepository for SqlxManagementTokenRepository {
|
||||
&self,
|
||||
token_hash: &str,
|
||||
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
|
||||
let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_BY_HASH_SQL)
|
||||
.bind(token_hash)
|
||||
let mut builder = QueryBuilder::<Postgres>::new(MANAGEMENT_TOKEN_WITH_USER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"mt.token_hash",
|
||||
token_hash.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
@@ -412,6 +376,15 @@ impl ManagementTokenReadRepository for SqlxManagementTokenRepository {
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_management_token_filters<'a>(
|
||||
builder: &mut QueryBuilder<'a, Postgres>,
|
||||
where_clause: &mut WhereClause,
|
||||
query: &'a ManagementTokenListQuery,
|
||||
) {
|
||||
push_optional_eq(builder, where_clause, "mt.user_id", query.user_id.clone());
|
||||
push_optional_eq(builder, where_clause, "mt.is_active", query.is_active);
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ManagementTokenWriteRepository for SqlxManagementTokenRepository {
|
||||
async fn create_management_token(
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{sqlite::SqliteRow, Row, SqlitePool};
|
||||
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite, SqlitePool};
|
||||
|
||||
use super::types::{
|
||||
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
|
||||
@@ -9,6 +9,7 @@ use super::types::{
|
||||
};
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_limit, push_limit_offset, push_optional_eq, WhereClause};
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqliteManagementTokenRepository {
|
||||
@@ -24,8 +25,12 @@ impl SqliteManagementTokenRepository {
|
||||
&self,
|
||||
token_id: &str,
|
||||
) -> Result<Option<StoredManagementToken>, DataLayerError> {
|
||||
let row = sqlx::query(TOKEN_BY_ID_SQL)
|
||||
.bind(token_id)
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(TOKEN_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "id", token_id.to_string());
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
@@ -33,7 +38,7 @@ impl SqliteManagementTokenRepository {
|
||||
}
|
||||
}
|
||||
|
||||
const TOKEN_BY_ID_SQL: &str = r#"
|
||||
const TOKEN_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
user_id,
|
||||
@@ -50,11 +55,9 @@ SELECT
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM management_tokens
|
||||
WHERE id = ?
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const LIST_MANAGEMENT_TOKENS_SQL: &str = r#"
|
||||
const TOKEN_WITH_USER_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
mt.id,
|
||||
mt.user_id,
|
||||
@@ -76,69 +79,6 @@ SELECT
|
||||
u.role AS user_role
|
||||
FROM management_tokens mt
|
||||
JOIN users u ON u.id = mt.user_id
|
||||
WHERE (? IS NULL OR mt.user_id = ?)
|
||||
AND (? IS NULL OR mt.is_active = ?)
|
||||
ORDER BY mt.created_at DESC, mt.id DESC
|
||||
LIMIT ? OFFSET ?
|
||||
"#;
|
||||
|
||||
const COUNT_MANAGEMENT_TOKENS_SQL: &str = r#"
|
||||
SELECT COUNT(mt.id) AS total
|
||||
FROM management_tokens mt
|
||||
WHERE (? IS NULL OR mt.user_id = ?)
|
||||
AND (? IS NULL OR mt.is_active = ?)
|
||||
"#;
|
||||
|
||||
const GET_MANAGEMENT_TOKEN_WITH_USER_SQL: &str = r#"
|
||||
SELECT
|
||||
mt.id,
|
||||
mt.user_id,
|
||||
mt.name,
|
||||
mt.description,
|
||||
mt.token_prefix,
|
||||
mt.allowed_ips,
|
||||
mt.permissions,
|
||||
mt.expires_at AS expires_at_unix_secs,
|
||||
mt.last_used_at AS last_used_at_unix_secs,
|
||||
mt.last_used_ip,
|
||||
COALESCE(mt.usage_count, 0) AS usage_count,
|
||||
mt.is_active,
|
||||
mt.created_at AS created_at_unix_ms,
|
||||
mt.updated_at AS updated_at_unix_secs,
|
||||
u.id AS user_row_id,
|
||||
u.email AS user_email,
|
||||
u.username AS user_username,
|
||||
u.role AS user_role
|
||||
FROM management_tokens mt
|
||||
JOIN users u ON u.id = mt.user_id
|
||||
WHERE mt.id = ?
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const GET_MANAGEMENT_TOKEN_WITH_USER_BY_HASH_SQL: &str = r#"
|
||||
SELECT
|
||||
mt.id,
|
||||
mt.user_id,
|
||||
mt.name,
|
||||
mt.description,
|
||||
mt.token_prefix,
|
||||
mt.allowed_ips,
|
||||
mt.permissions,
|
||||
mt.expires_at AS expires_at_unix_secs,
|
||||
mt.last_used_at AS last_used_at_unix_secs,
|
||||
mt.last_used_ip,
|
||||
COALESCE(mt.usage_count, 0) AS usage_count,
|
||||
mt.is_active,
|
||||
mt.created_at AS created_at_unix_ms,
|
||||
mt.updated_at AS updated_at_unix_secs,
|
||||
u.id AS user_row_id,
|
||||
u.email AS user_email,
|
||||
u.username AS user_username,
|
||||
u.role AS user_role
|
||||
FROM management_tokens mt
|
||||
JOIN users u ON u.id = mt.user_id
|
||||
WHERE mt.token_hash = ?
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
#[async_trait]
|
||||
@@ -147,23 +87,27 @@ impl ManagementTokenReadRepository for SqliteManagementTokenRepository {
|
||||
&self,
|
||||
query: &ManagementTokenListQuery,
|
||||
) -> Result<StoredManagementTokenListPage, DataLayerError> {
|
||||
let count_row = sqlx::query(COUNT_MANAGEMENT_TOKENS_SQL)
|
||||
.bind(query.user_id.as_deref())
|
||||
.bind(query.user_id.as_deref())
|
||||
.bind(query.is_active)
|
||||
.bind(query.is_active)
|
||||
let mut count_builder =
|
||||
QueryBuilder::<Sqlite>::new("SELECT COUNT(mt.id) AS total FROM management_tokens mt");
|
||||
let mut count_where = WhereClause::new();
|
||||
apply_management_token_filters(&mut count_builder, &mut count_where, query);
|
||||
let total = count_builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let total = count_row.try_get::<i64, _>("total").map_sql_err()?;
|
||||
|
||||
let rows = sqlx::query(LIST_MANAGEMENT_TOKENS_SQL)
|
||||
.bind(query.user_id.as_deref())
|
||||
.bind(query.user_id.as_deref())
|
||||
.bind(query.is_active)
|
||||
.bind(query.is_active)
|
||||
.bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
|
||||
.bind(i64::try_from(query.offset).unwrap_or(i64::MAX))
|
||||
let mut list_builder = QueryBuilder::<Sqlite>::new(TOKEN_WITH_USER_COLUMNS);
|
||||
let mut list_where = WhereClause::new();
|
||||
apply_management_token_filters(&mut list_builder, &mut list_where, query);
|
||||
list_builder.push(" ORDER BY mt.created_at DESC, mt.id DESC");
|
||||
push_limit_offset(
|
||||
&mut list_builder,
|
||||
i64::try_from(query.limit).unwrap_or(i64::MAX),
|
||||
i64::try_from(query.offset).unwrap_or(i64::MAX),
|
||||
);
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
@@ -181,8 +125,17 @@ impl ManagementTokenReadRepository for SqliteManagementTokenRepository {
|
||||
&self,
|
||||
token_id: &str,
|
||||
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
|
||||
let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_SQL)
|
||||
.bind(token_id)
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(TOKEN_WITH_USER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"mt.id",
|
||||
token_id.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
@@ -193,8 +146,17 @@ impl ManagementTokenReadRepository for SqliteManagementTokenRepository {
|
||||
&self,
|
||||
token_hash: &str,
|
||||
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
|
||||
let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_BY_HASH_SQL)
|
||||
.bind(token_hash)
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(TOKEN_WITH_USER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"mt.token_hash",
|
||||
token_hash.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
@@ -202,6 +164,15 @@ impl ManagementTokenReadRepository for SqliteManagementTokenRepository {
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_management_token_filters(
|
||||
builder: &mut QueryBuilder<'_, Sqlite>,
|
||||
where_clause: &mut WhereClause,
|
||||
query: &ManagementTokenListQuery,
|
||||
) {
|
||||
push_optional_eq(builder, where_clause, "mt.user_id", query.user_id.clone());
|
||||
push_optional_eq(builder, where_clause, "mt.is_active", query.is_active);
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ManagementTokenWriteRepository for SqliteManagementTokenRepository {
|
||||
async fn create_management_token(
|
||||
|
||||
@@ -1,14 +1,14 @@
|
||||
use async_trait::async_trait;
|
||||
use futures_util::TryStreamExt;
|
||||
use sqlx::{postgres::PgRow, PgPool, Row};
|
||||
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
|
||||
|
||||
use super::types::{
|
||||
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
|
||||
UpsertOAuthProviderConfigRecord,
|
||||
};
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
use aether_data_query::{push_eq, push_limit, WhereClause};
|
||||
|
||||
const LIST_OAUTH_PROVIDER_CONFIGS_SQL: &str = r#"
|
||||
const OAUTH_PROVIDER_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
provider_type,
|
||||
display_name,
|
||||
@@ -26,29 +26,6 @@ SELECT
|
||||
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM oauth_providers
|
||||
ORDER BY provider_type ASC
|
||||
"#;
|
||||
|
||||
const GET_OAUTH_PROVIDER_CONFIG_SQL: &str = r#"
|
||||
SELECT
|
||||
provider_type,
|
||||
display_name,
|
||||
client_id,
|
||||
client_secret_encrypted,
|
||||
authorization_url_override,
|
||||
token_url_override,
|
||||
userinfo_url_override,
|
||||
scopes,
|
||||
redirect_uri,
|
||||
frontend_callback_url,
|
||||
attribute_mapping,
|
||||
extra_config,
|
||||
is_enabled,
|
||||
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM oauth_providers
|
||||
WHERE provider_type = $1
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL: &str = r#"
|
||||
@@ -182,20 +159,31 @@ impl OAuthProviderReadRepository for SqlxOAuthProviderRepository {
|
||||
async fn list_oauth_provider_configs(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
|
||||
let mut rows = sqlx::query(LIST_OAUTH_PROVIDER_CONFIGS_SQL).fetch(&self.pool);
|
||||
let mut items = Vec::new();
|
||||
while let Some(row) = rows.try_next().await.map_postgres_err()? {
|
||||
items.push(map_oauth_provider_row(&row)?);
|
||||
}
|
||||
Ok(items)
|
||||
let mut builder = QueryBuilder::<Postgres>::new(OAUTH_PROVIDER_COLUMNS);
|
||||
builder.push(" ORDER BY provider_type ASC");
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
rows.iter().map(map_oauth_provider_row).collect()
|
||||
}
|
||||
|
||||
async fn get_oauth_provider_config(
|
||||
&self,
|
||||
provider_type: &str,
|
||||
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
|
||||
let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL)
|
||||
.bind(provider_type)
|
||||
let mut builder = QueryBuilder::<Postgres>::new(OAUTH_PROVIDER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"provider_type",
|
||||
provider_type.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{sqlite::SqliteRow, Row};
|
||||
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
|
||||
|
||||
use super::types::{
|
||||
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
|
||||
@@ -8,6 +8,7 @@ use super::types::{
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_limit, WhereClause};
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqliteOAuthProviderRepository {
|
||||
@@ -23,8 +24,17 @@ impl SqliteOAuthProviderRepository {
|
||||
&self,
|
||||
provider_type: &str,
|
||||
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
|
||||
let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL)
|
||||
.bind(provider_type)
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(OAUTH_PROVIDER_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"provider_type",
|
||||
provider_type.to_string(),
|
||||
);
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
@@ -32,7 +42,7 @@ impl SqliteOAuthProviderRepository {
|
||||
}
|
||||
}
|
||||
|
||||
const LIST_OAUTH_PROVIDER_CONFIGS_SQL: &str = r#"
|
||||
const OAUTH_PROVIDER_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
provider_type,
|
||||
display_name,
|
||||
@@ -50,29 +60,6 @@ SELECT
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM oauth_providers
|
||||
ORDER BY provider_type ASC
|
||||
"#;
|
||||
|
||||
const GET_OAUTH_PROVIDER_CONFIG_SQL: &str = r#"
|
||||
SELECT
|
||||
provider_type,
|
||||
display_name,
|
||||
client_id,
|
||||
client_secret_encrypted,
|
||||
authorization_url_override,
|
||||
token_url_override,
|
||||
userinfo_url_override,
|
||||
scopes,
|
||||
redirect_uri,
|
||||
frontend_callback_url,
|
||||
attribute_mapping,
|
||||
extra_config,
|
||||
is_enabled,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM oauth_providers
|
||||
WHERE provider_type = ?
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL: &str = r#"
|
||||
@@ -117,10 +104,9 @@ impl OAuthProviderReadRepository for SqliteOAuthProviderRepository {
|
||||
async fn list_oauth_provider_configs(
|
||||
&self,
|
||||
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
|
||||
let rows = sqlx::query(LIST_OAUTH_PROVIDER_CONFIGS_SQL)
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(OAUTH_PROVIDER_COLUMNS);
|
||||
builder.push(" ORDER BY provider_type ASC");
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_oauth_provider_row).collect()
|
||||
}
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ use super::{
|
||||
use crate::error::SqlxResultExt;
|
||||
use crate::repository::pool_scores::merge_score_reason_patch;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_in, push_limit, push_limit_offset, WhereClause};
|
||||
|
||||
const SCORE_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
@@ -57,25 +58,54 @@ impl PostgresPoolMemberScoreRepository {
|
||||
scope: Option<&PoolScoreScope>,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(identity.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(identity.pool_id.clone())
|
||||
.push(" AND member_kind = ")
|
||||
.push_bind(identity.member_kind.clone())
|
||||
.push(" AND member_id = ")
|
||||
.push_bind(identity.member_id.clone());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_kind",
|
||||
identity.pool_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_id",
|
||||
identity.pool_id.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"member_kind",
|
||||
identity.member_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"member_id",
|
||||
identity.member_id.clone(),
|
||||
);
|
||||
if let Some(scope) = scope {
|
||||
builder
|
||||
.push(" AND capability = ")
|
||||
.push_bind(scope.capability.clone())
|
||||
.push(" AND scope_kind = ")
|
||||
.push_bind(scope.scope_kind.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"capability",
|
||||
scope.capability.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_kind",
|
||||
scope.scope_kind.clone(),
|
||||
);
|
||||
if let Some(scope_id) = &scope.scope_id {
|
||||
builder.push(" AND scope_id = ").push_bind(scope_id.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_id",
|
||||
scope_id.clone(),
|
||||
);
|
||||
} else {
|
||||
builder.push(" AND scope_id IS NULL");
|
||||
where_clause.push_next(&mut builder);
|
||||
builder.push("scope_id IS NULL");
|
||||
}
|
||||
}
|
||||
let rows = builder
|
||||
@@ -94,44 +124,65 @@ impl PoolScoreReadRepository for PostgresPoolMemberScoreRepository {
|
||||
query: &ListRankedPoolMembersQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(query.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(query.pool_id.clone())
|
||||
.push(" AND capability = ")
|
||||
.push_bind(query.capability.clone())
|
||||
.push(" AND scope_kind = ")
|
||||
.push_bind(query.scope_kind.clone());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_kind",
|
||||
query.pool_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_id",
|
||||
query.pool_id.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"capability",
|
||||
query.capability.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_kind",
|
||||
query.scope_kind.clone(),
|
||||
);
|
||||
if let Some(scope_id) = &query.scope_id {
|
||||
builder.push(" AND scope_id = ").push_bind(scope_id.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_id",
|
||||
scope_id.clone(),
|
||||
);
|
||||
} else {
|
||||
builder.push(" AND scope_id IS NULL");
|
||||
where_clause.push_next(&mut builder);
|
||||
builder.push("scope_id IS NULL");
|
||||
}
|
||||
if !query.hard_states.is_empty() {
|
||||
builder.push(" AND hard_state IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for state in &query.hard_states {
|
||||
separated.push_bind(state.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let states = query
|
||||
.hard_states
|
||||
.iter()
|
||||
.map(|state| state.as_database())
|
||||
.collect::<Vec<_>>();
|
||||
push_in(&mut builder, &mut where_clause, "hard_state", &states);
|
||||
}
|
||||
if let Some(statuses) = &query.probe_statuses {
|
||||
if !statuses.is_empty() {
|
||||
builder.push(" AND probe_status IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for status in statuses {
|
||||
separated.push_bind(status.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let statuses = statuses
|
||||
.iter()
|
||||
.map(|status| status.as_database())
|
||||
.collect::<Vec<_>>();
|
||||
push_in(&mut builder, &mut where_clause, "probe_status", &statuses);
|
||||
}
|
||||
}
|
||||
builder
|
||||
.push(" ORDER BY score DESC, last_ranked_at DESC NULLS LAST, member_id ASC, id ASC")
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(query.limit.max(1), "pool score limit")?)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(i64_from_usize(query.offset, "pool score offset")?);
|
||||
builder.push(" ORDER BY score DESC, last_ranked_at DESC NULLS LAST, member_id ASC, id ASC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64_from_usize(query.limit.max(1), "pool score limit")?,
|
||||
i64_from_usize(query.offset, "pool score offset")?,
|
||||
);
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
@@ -145,48 +196,66 @@ impl PoolScoreReadRepository for PostgresPoolMemberScoreRepository {
|
||||
query: &ListPoolMemberScoresQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(query.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(query.pool_id.clone());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_kind",
|
||||
query.pool_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_id",
|
||||
query.pool_id.clone(),
|
||||
);
|
||||
if let Some(capability) = &query.capability {
|
||||
builder
|
||||
.push(" AND capability = ")
|
||||
.push_bind(capability.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"capability",
|
||||
capability.clone(),
|
||||
);
|
||||
}
|
||||
if let Some(scope_kind) = &query.scope_kind {
|
||||
builder
|
||||
.push(" AND scope_kind = ")
|
||||
.push_bind(scope_kind.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_kind",
|
||||
scope_kind.clone(),
|
||||
);
|
||||
}
|
||||
if let Some(scope_id) = &query.scope_id {
|
||||
builder.push(" AND scope_id = ").push_bind(scope_id.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_id",
|
||||
scope_id.clone(),
|
||||
);
|
||||
}
|
||||
if !query.hard_states.is_empty() {
|
||||
builder.push(" AND hard_state IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for state in &query.hard_states {
|
||||
separated.push_bind(state.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let states = query
|
||||
.hard_states
|
||||
.iter()
|
||||
.map(|state| state.as_database())
|
||||
.collect::<Vec<_>>();
|
||||
push_in(&mut builder, &mut where_clause, "hard_state", &states);
|
||||
}
|
||||
if let Some(statuses) = &query.probe_statuses {
|
||||
if !statuses.is_empty() {
|
||||
builder.push(" AND probe_status IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for status in statuses {
|
||||
separated.push_bind(status.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let statuses = statuses
|
||||
.iter()
|
||||
.map(|status| status.as_database())
|
||||
.collect::<Vec<_>>();
|
||||
push_in(&mut builder, &mut where_clause, "probe_status", &statuses);
|
||||
}
|
||||
}
|
||||
builder
|
||||
.push(" ORDER BY score DESC, last_ranked_at DESC NULLS LAST, member_id ASC, id ASC")
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(query.limit.max(1), "pool score limit")?)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(i64_from_usize(query.offset, "pool score offset")?);
|
||||
builder.push(" ORDER BY score DESC, last_ranked_at DESC NULLS LAST, member_id ASC, id ASC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64_from_usize(query.limit.max(1), "pool score limit")?,
|
||||
i64_from_usize(query.offset, "pool score offset")?,
|
||||
);
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
@@ -200,18 +269,30 @@ impl PoolScoreReadRepository for PostgresPoolMemberScoreRepository {
|
||||
query: &ListPoolMemberProbeCandidatesQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(query.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(query.pool_id.clone());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_kind",
|
||||
query.pool_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_id",
|
||||
query.pool_id.clone(),
|
||||
);
|
||||
if let Some(capability) = &query.capability {
|
||||
builder
|
||||
.push(" AND capability = ")
|
||||
.push_bind(capability.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"capability",
|
||||
capability.clone(),
|
||||
);
|
||||
}
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push(" AND hard_state IN ('available','unknown','cooldown','quota_exhausted')")
|
||||
.push("hard_state IN ('available','unknown','cooldown','quota_exhausted')")
|
||||
.push(" AND (probe_status IN ('never','failed','stale')")
|
||||
.push(" OR (probe_status = 'ok' AND (last_probe_success_at IS NULL OR last_probe_success_at <= ")
|
||||
.push_bind(i64_from_u64(
|
||||
@@ -240,12 +321,11 @@ impl PoolScoreReadRepository for PostgresPoolMemberScoreRepository {
|
||||
COALESCE(last_scheduled_at, 0) DESC,
|
||||
member_id ASC
|
||||
"#,
|
||||
)
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(
|
||||
query.limit.max(1),
|
||||
"pool probe candidate limit",
|
||||
)?);
|
||||
);
|
||||
push_limit(
|
||||
&mut builder,
|
||||
i64_from_usize(query.limit.max(1), "pool probe candidate limit")?,
|
||||
);
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
@@ -262,12 +342,8 @@ impl PoolScoreReadRepository for PostgresPoolMemberScoreRepository {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let mut builder = QueryBuilder::<Postgres>::new(SCORE_COLUMNS);
|
||||
builder.push(" WHERE id IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for id in &query.ids {
|
||||
separated.push_bind(id.clone());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(&mut builder, &mut where_clause, "id", &query.ids);
|
||||
let rows = builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
|
||||
@@ -12,6 +12,7 @@ use super::{
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::repository::pool_scores::merge_score_reason_patch;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_in, push_limit, push_limit_offset, WhereClause};
|
||||
|
||||
const SCORE_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
@@ -57,25 +58,54 @@ impl SqlitePoolMemberScoreRepository {
|
||||
scope: Option<&PoolScoreScope>,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(identity.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(identity.pool_id.clone())
|
||||
.push(" AND member_kind = ")
|
||||
.push_bind(identity.member_kind.clone())
|
||||
.push(" AND member_id = ")
|
||||
.push_bind(identity.member_id.clone());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_kind",
|
||||
identity.pool_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_id",
|
||||
identity.pool_id.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"member_kind",
|
||||
identity.member_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"member_id",
|
||||
identity.member_id.clone(),
|
||||
);
|
||||
if let Some(scope) = scope {
|
||||
builder
|
||||
.push(" AND capability = ")
|
||||
.push_bind(scope.capability.clone())
|
||||
.push(" AND scope_kind = ")
|
||||
.push_bind(scope.scope_kind.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"capability",
|
||||
scope.capability.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_kind",
|
||||
scope.scope_kind.clone(),
|
||||
);
|
||||
if let Some(scope_id) = &scope.scope_id {
|
||||
builder.push(" AND scope_id = ").push_bind(scope_id.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_id",
|
||||
scope_id.clone(),
|
||||
);
|
||||
} else {
|
||||
builder.push(" AND scope_id IS NULL");
|
||||
where_clause.push_next(&mut builder);
|
||||
builder.push("scope_id IS NULL");
|
||||
}
|
||||
}
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
@@ -90,44 +120,65 @@ impl PoolScoreReadRepository for SqlitePoolMemberScoreRepository {
|
||||
query: &ListRankedPoolMembersQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(query.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(query.pool_id.clone())
|
||||
.push(" AND capability = ")
|
||||
.push_bind(query.capability.clone())
|
||||
.push(" AND scope_kind = ")
|
||||
.push_bind(query.scope_kind.clone());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_kind",
|
||||
query.pool_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_id",
|
||||
query.pool_id.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"capability",
|
||||
query.capability.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_kind",
|
||||
query.scope_kind.clone(),
|
||||
);
|
||||
if let Some(scope_id) = &query.scope_id {
|
||||
builder.push(" AND scope_id = ").push_bind(scope_id.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_id",
|
||||
scope_id.clone(),
|
||||
);
|
||||
} else {
|
||||
builder.push(" AND scope_id IS NULL");
|
||||
where_clause.push_next(&mut builder);
|
||||
builder.push("scope_id IS NULL");
|
||||
}
|
||||
if !query.hard_states.is_empty() {
|
||||
builder.push(" AND hard_state IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for state in &query.hard_states {
|
||||
separated.push_bind(state.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let states = query
|
||||
.hard_states
|
||||
.iter()
|
||||
.map(|state| state.as_database())
|
||||
.collect::<Vec<_>>();
|
||||
push_in(&mut builder, &mut where_clause, "hard_state", &states);
|
||||
}
|
||||
if let Some(statuses) = &query.probe_statuses {
|
||||
if !statuses.is_empty() {
|
||||
builder.push(" AND probe_status IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for status in statuses {
|
||||
separated.push_bind(status.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let statuses = statuses
|
||||
.iter()
|
||||
.map(|status| status.as_database())
|
||||
.collect::<Vec<_>>();
|
||||
push_in(&mut builder, &mut where_clause, "probe_status", &statuses);
|
||||
}
|
||||
}
|
||||
builder
|
||||
.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC")
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(query.limit.max(1), "pool score limit")?)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(i64_from_usize(query.offset, "pool score offset")?);
|
||||
builder.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64_from_usize(query.limit.max(1), "pool score limit")?,
|
||||
i64_from_usize(query.offset, "pool score offset")?,
|
||||
);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
@@ -137,48 +188,66 @@ impl PoolScoreReadRepository for SqlitePoolMemberScoreRepository {
|
||||
query: &ListPoolMemberScoresQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(query.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(query.pool_id.clone());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_kind",
|
||||
query.pool_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_id",
|
||||
query.pool_id.clone(),
|
||||
);
|
||||
if let Some(capability) = &query.capability {
|
||||
builder
|
||||
.push(" AND capability = ")
|
||||
.push_bind(capability.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"capability",
|
||||
capability.clone(),
|
||||
);
|
||||
}
|
||||
if let Some(scope_kind) = &query.scope_kind {
|
||||
builder
|
||||
.push(" AND scope_kind = ")
|
||||
.push_bind(scope_kind.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_kind",
|
||||
scope_kind.clone(),
|
||||
);
|
||||
}
|
||||
if let Some(scope_id) = &query.scope_id {
|
||||
builder.push(" AND scope_id = ").push_bind(scope_id.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"scope_id",
|
||||
scope_id.clone(),
|
||||
);
|
||||
}
|
||||
if !query.hard_states.is_empty() {
|
||||
builder.push(" AND hard_state IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for state in &query.hard_states {
|
||||
separated.push_bind(state.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let states = query
|
||||
.hard_states
|
||||
.iter()
|
||||
.map(|state| state.as_database())
|
||||
.collect::<Vec<_>>();
|
||||
push_in(&mut builder, &mut where_clause, "hard_state", &states);
|
||||
}
|
||||
if let Some(statuses) = &query.probe_statuses {
|
||||
if !statuses.is_empty() {
|
||||
builder.push(" AND probe_status IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for status in statuses {
|
||||
separated.push_bind(status.as_database());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let statuses = statuses
|
||||
.iter()
|
||||
.map(|status| status.as_database())
|
||||
.collect::<Vec<_>>();
|
||||
push_in(&mut builder, &mut where_clause, "probe_status", &statuses);
|
||||
}
|
||||
}
|
||||
builder
|
||||
.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC")
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(query.limit.max(1), "pool score limit")?)
|
||||
.push(" OFFSET ")
|
||||
.push_bind(i64_from_usize(query.offset, "pool score offset")?);
|
||||
builder.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC");
|
||||
push_limit_offset(
|
||||
&mut builder,
|
||||
i64_from_usize(query.limit.max(1), "pool score limit")?,
|
||||
i64_from_usize(query.offset, "pool score offset")?,
|
||||
);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
@@ -188,18 +257,30 @@ impl PoolScoreReadRepository for SqlitePoolMemberScoreRepository {
|
||||
query: &ListPoolMemberProbeCandidatesQuery,
|
||||
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
|
||||
builder
|
||||
.push(" WHERE pool_kind = ")
|
||||
.push_bind(query.pool_kind.clone())
|
||||
.push(" AND pool_id = ")
|
||||
.push_bind(query.pool_id.clone());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_kind",
|
||||
query.pool_kind.clone(),
|
||||
);
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"pool_id",
|
||||
query.pool_id.clone(),
|
||||
);
|
||||
if let Some(capability) = &query.capability {
|
||||
builder
|
||||
.push(" AND capability = ")
|
||||
.push_bind(capability.clone());
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"capability",
|
||||
capability.clone(),
|
||||
);
|
||||
}
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push(" AND hard_state IN ('available','unknown','cooldown','quota_exhausted')")
|
||||
.push("hard_state IN ('available','unknown','cooldown','quota_exhausted')")
|
||||
.push(" AND (probe_status IN ('never','failed','stale')")
|
||||
.push(" OR (probe_status = 'ok' AND (last_probe_success_at IS NULL OR last_probe_success_at <= ")
|
||||
.push_bind(i64_from_u64(
|
||||
@@ -228,12 +309,11 @@ impl PoolScoreReadRepository for SqlitePoolMemberScoreRepository {
|
||||
COALESCE(last_scheduled_at, 0) DESC,
|
||||
member_id ASC
|
||||
"#,
|
||||
)
|
||||
.push(" LIMIT ")
|
||||
.push_bind(i64_from_usize(
|
||||
query.limit.max(1),
|
||||
"pool probe candidate limit",
|
||||
)?);
|
||||
);
|
||||
push_limit(
|
||||
&mut builder,
|
||||
i64_from_usize(query.limit.max(1), "pool probe candidate limit")?,
|
||||
);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
@@ -246,12 +326,8 @@ impl PoolScoreReadRepository for SqlitePoolMemberScoreRepository {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(SCORE_COLUMNS);
|
||||
builder.push(" WHERE id IN (");
|
||||
let mut separated = builder.separated(", ");
|
||||
for id in &query.ids {
|
||||
separated.push_bind(id.clone());
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(&mut builder, &mut where_clause, "id", &query.ids);
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_score_row).collect()
|
||||
}
|
||||
|
||||
@@ -669,6 +669,11 @@ WHERE id = ?
|
||||
"provider_api_keys.last_probe_increase_at",
|
||||
)?)
|
||||
.bind(optional_i64_from_u32(key.last_rpm_peak))
|
||||
.bind(optional_i64_from_u64(
|
||||
key.last_models_fetch_at_unix_secs,
|
||||
"provider_api_keys.last_models_fetch_at",
|
||||
)?)
|
||||
.bind(&key.last_models_fetch_error)
|
||||
.bind(updated_at)
|
||||
.bind(&key.id)
|
||||
.execute(&self.pool)
|
||||
@@ -1208,6 +1213,8 @@ SET
|
||||
utilization_samples = ?,
|
||||
last_probe_increase_at = ?,
|
||||
last_rpm_peak = ?,
|
||||
last_models_fetch_at = ?,
|
||||
last_models_fetch_error = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
"#
|
||||
|
||||
@@ -11,6 +11,10 @@ use crate::{
|
||||
error::{postgres_error, SqlxResultExt},
|
||||
DataLayerError,
|
||||
};
|
||||
use aether_data_query::{
|
||||
push_ci_contains_any, push_eq, push_in, push_limit_offset, push_optional_eq, SqlDialect,
|
||||
WhereClause,
|
||||
};
|
||||
|
||||
const LIST_PROVIDERS_BY_IDS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
@@ -358,43 +362,14 @@ impl SqlxProviderCatalogReadRepository {
|
||||
&self,
|
||||
active_only: bool,
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
collect_query_rows(
|
||||
sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id,
|
||||
name,
|
||||
description,
|
||||
website,
|
||||
provider_type,
|
||||
CAST(billing_type AS TEXT) AS billing_type,
|
||||
CAST(monthly_quota_usd AS DOUBLE PRECISION) AS monthly_quota_usd,
|
||||
CAST(monthly_used_usd AS DOUBLE PRECISION) AS monthly_used_usd,
|
||||
quota_reset_day,
|
||||
CAST(EXTRACT(EPOCH FROM quota_last_reset_at) AS BIGINT) AS quota_last_reset_at_unix_secs,
|
||||
CAST(EXTRACT(EPOCH FROM quota_expires_at) AS BIGINT) AS quota_expires_at_unix_secs,
|
||||
provider_priority,
|
||||
is_active,
|
||||
keep_priority_on_conversion,
|
||||
enable_format_conversion,
|
||||
concurrent_limit,
|
||||
max_retries,
|
||||
proxy,
|
||||
request_timeout,
|
||||
stream_first_byte_timeout,
|
||||
config,
|
||||
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM providers
|
||||
WHERE ($1::boolean = false OR is_active = true)
|
||||
ORDER BY provider_priority ASC, name ASC
|
||||
"#,
|
||||
)
|
||||
.bind(active_only)
|
||||
.fetch(&self.pool),
|
||||
map_provider_row,
|
||||
)
|
||||
.await
|
||||
let mut builder =
|
||||
QueryBuilder::<Postgres>::new(select_prefix_for_in(LIST_PROVIDERS_BY_IDS_PREFIX));
|
||||
let mut where_clause = WhereClause::new();
|
||||
if active_only {
|
||||
push_eq(&mut builder, &mut where_clause, "is_active", true);
|
||||
}
|
||||
builder.push(" ORDER BY provider_priority ASC, name ASC");
|
||||
collect_query_rows(builder.build().fetch(&self.pool), map_provider_row).await
|
||||
}
|
||||
|
||||
pub async fn list_endpoints_by_ids(
|
||||
@@ -560,12 +535,6 @@ ORDER BY provider_priority ASC, name ASC
|
||||
query.limit
|
||||
))
|
||||
})?;
|
||||
let search_pattern = query
|
||||
.search
|
||||
.as_deref()
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(|value| format!("%{}%", value.to_ascii_lowercase()));
|
||||
let order_by = match query.order {
|
||||
ProviderCatalogKeyListOrder::Name => "internal_priority ASC, name ASC, id ASC",
|
||||
ProviderCatalogKeyListOrder::CreatedAt => {
|
||||
@@ -585,99 +554,25 @@ ORDER BY provider_priority ASC, name ASC
|
||||
}
|
||||
};
|
||||
|
||||
let count_row = sqlx::query(
|
||||
r#"
|
||||
SELECT COUNT(*)::BIGINT AS total
|
||||
FROM provider_api_keys
|
||||
WHERE provider_id = $1
|
||||
AND ($2::TEXT IS NULL OR LOWER(name) LIKE $2 OR LOWER(id) LIKE $2)
|
||||
AND ($3::BOOLEAN IS NULL OR is_active = $3)
|
||||
"#,
|
||||
)
|
||||
.bind(&query.provider_id)
|
||||
.bind(search_pattern.as_deref())
|
||||
.bind(query.is_active)
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
let total = row_get::<i64>(&count_row, "total")?.max(0) as usize;
|
||||
|
||||
let sql = format!(
|
||||
r#"
|
||||
SELECT
|
||||
id,
|
||||
provider_id,
|
||||
name,
|
||||
auth_type,
|
||||
capabilities,
|
||||
is_active,
|
||||
api_formats,
|
||||
auth_type_by_format,
|
||||
allow_auth_channel_mismatch_formats,
|
||||
COALESCE(api_key, encrypted_key) AS api_key,
|
||||
auth_config,
|
||||
note,
|
||||
internal_priority,
|
||||
rate_multipliers,
|
||||
global_priority_by_format,
|
||||
allowed_models,
|
||||
EXTRACT(EPOCH FROM expires_at)::bigint AS expires_at_unix_secs,
|
||||
cache_ttl_minutes,
|
||||
max_probe_interval_minutes,
|
||||
proxy,
|
||||
fingerprint,
|
||||
rpm_limit,
|
||||
concurrent_limit,
|
||||
learned_rpm_limit,
|
||||
concurrent_429_count,
|
||||
rpm_429_count,
|
||||
EXTRACT(EPOCH FROM last_429_at)::bigint AS last_429_at_unix_secs,
|
||||
last_429_type,
|
||||
adjustment_history,
|
||||
utilization_samples,
|
||||
EXTRACT(EPOCH FROM last_probe_increase_at)::bigint AS last_probe_increase_at_unix_secs,
|
||||
last_rpm_peak,
|
||||
request_count,
|
||||
total_tokens,
|
||||
CAST(total_cost_usd AS DOUBLE PRECISION) AS total_cost_usd,
|
||||
success_count,
|
||||
error_count,
|
||||
total_response_time_ms,
|
||||
EXTRACT(EPOCH FROM last_used_at)::bigint AS last_used_at_unix_secs,
|
||||
auto_fetch_models,
|
||||
EXTRACT(EPOCH FROM last_models_fetch_at)::bigint AS last_models_fetch_at_unix_secs,
|
||||
last_models_fetch_error,
|
||||
locked_models,
|
||||
model_include_patterns,
|
||||
model_exclude_patterns,
|
||||
upstream_metadata,
|
||||
EXTRACT(EPOCH FROM oauth_invalid_at)::bigint AS oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason,
|
||||
status_snapshot,
|
||||
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs,
|
||||
health_by_format,
|
||||
circuit_breaker_by_format
|
||||
FROM provider_api_keys
|
||||
WHERE provider_id = $1
|
||||
AND ($2::TEXT IS NULL OR LOWER(name) LIKE $2 OR LOWER(id) LIKE $2)
|
||||
AND ($3::BOOLEAN IS NULL OR is_active = $3)
|
||||
ORDER BY {order_by}
|
||||
OFFSET $4
|
||||
LIMIT $5
|
||||
"#,
|
||||
let mut count_builder = QueryBuilder::<Postgres>::new(
|
||||
"SELECT COUNT(*)::BIGINT AS total FROM provider_api_keys",
|
||||
);
|
||||
let items = collect_query_rows(
|
||||
sqlx::query(&sql)
|
||||
.bind(&query.provider_id)
|
||||
.bind(search_pattern.as_deref())
|
||||
.bind(query.is_active)
|
||||
.bind(offset)
|
||||
.bind(limit)
|
||||
.fetch(&self.pool),
|
||||
map_key_row,
|
||||
)
|
||||
.await?;
|
||||
let mut count_where = WhereClause::new();
|
||||
apply_key_page_filters(&mut count_builder, &mut count_where, query);
|
||||
let total = count_builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
.max(0) as usize;
|
||||
|
||||
let mut list_builder =
|
||||
QueryBuilder::<Postgres>::new(select_prefix_for_in(LIST_KEYS_BY_IDS_PREFIX));
|
||||
let mut list_where = WhereClause::new();
|
||||
apply_key_page_filters(&mut list_builder, &mut list_where, query);
|
||||
list_builder.push(" ORDER BY ").push(order_by);
|
||||
push_limit_offset(&mut list_builder, limit, offset);
|
||||
let items = collect_query_rows(list_builder.build().fetch(&self.pool), map_key_row).await?;
|
||||
|
||||
Ok(StoredProviderCatalogKeyPage { items, total })
|
||||
}
|
||||
@@ -1798,7 +1693,12 @@ SET
|
||||
updated_at = CASE
|
||||
WHEN $38::double precision IS NULL THEN NOW()
|
||||
ELSE TO_TIMESTAMP($38::double precision)
|
||||
END
|
||||
END,
|
||||
last_models_fetch_at = CASE
|
||||
WHEN $42::double precision IS NULL THEN NULL
|
||||
ELSE TO_TIMESTAMP($42::double precision)
|
||||
END,
|
||||
last_models_fetch_error = $43
|
||||
WHERE id = $1
|
||||
"#,
|
||||
)
|
||||
@@ -1846,6 +1746,8 @@ WHERE id = $1
|
||||
.bind(key.expires_at_unix_secs.map(|value| value as f64))
|
||||
.bind(&key.auth_type_by_format)
|
||||
.bind(&key.allow_auth_channel_mismatch_formats)
|
||||
.bind(key.last_models_fetch_at_unix_secs.map(|value| value as f64))
|
||||
.bind(&key.last_models_fetch_error)
|
||||
.execute(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
@@ -2150,16 +2052,56 @@ fn build_list_query<'a>(
|
||||
ids: &'a [String],
|
||||
suffix: &'static str,
|
||||
) -> QueryBuilder<'a, Postgres> {
|
||||
let mut builder = QueryBuilder::<Postgres>::new(prefix);
|
||||
let mut separated = builder.separated(", ");
|
||||
for id in ids {
|
||||
separated.push_bind(id);
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
let mut builder = QueryBuilder::<Postgres>::new(select_prefix_for_in(prefix));
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
in_column_for_prefix(prefix),
|
||||
ids,
|
||||
);
|
||||
builder.push(suffix);
|
||||
builder
|
||||
}
|
||||
|
||||
fn select_prefix_for_in(prefix: &'static str) -> &'static str {
|
||||
prefix
|
||||
.rsplit_once("\nWHERE ")
|
||||
.map(|(select_prefix, _)| select_prefix)
|
||||
.expect("provider catalog IN query prefix must contain WHERE")
|
||||
}
|
||||
|
||||
fn in_column_for_prefix(prefix: &'static str) -> &'static str {
|
||||
prefix
|
||||
.rsplit_once("\nWHERE ")
|
||||
.and_then(|(_, predicate)| predicate.trim().strip_suffix("IN ("))
|
||||
.map(str::trim)
|
||||
.expect("provider catalog IN query prefix must end with IN (")
|
||||
}
|
||||
|
||||
fn apply_key_page_filters<'a>(
|
||||
builder: &mut QueryBuilder<'a, Postgres>,
|
||||
where_clause: &mut WhereClause,
|
||||
query: &'a ProviderCatalogKeyListQuery,
|
||||
) {
|
||||
push_eq(
|
||||
builder,
|
||||
where_clause,
|
||||
"provider_id",
|
||||
query.provider_id.clone(),
|
||||
);
|
||||
if let Some(search) = query.search.as_deref() {
|
||||
push_ci_contains_any(
|
||||
builder,
|
||||
where_clause,
|
||||
SqlDialect::Postgres,
|
||||
&["name", "id"],
|
||||
search,
|
||||
);
|
||||
}
|
||||
push_optional_eq(builder, where_clause, "is_active", query.is_active);
|
||||
}
|
||||
|
||||
fn row_get<T>(row: &PgRow, column: &str) -> Result<T, DataLayerError>
|
||||
where
|
||||
for<'r> T: sqlx::Decode<'r, sqlx::Postgres> + sqlx::Type<sqlx::Postgres>,
|
||||
@@ -2623,8 +2565,9 @@ mod tests {
|
||||
"auth_type_by_format,\n allow_auth_channel_mismatch_formats,\n COALESCE(api_key, encrypted_key) AS api_key",
|
||||
)
|
||||
.count()
|
||||
>= 3
|
||||
>= 2
|
||||
);
|
||||
assert!(source.contains("QueryBuilder::<Postgres>::new(select_prefix_for_in("));
|
||||
assert!(source.contains(".bind(&key.allow_auth_channel_mismatch_formats)"));
|
||||
assert!(source.contains("row.try_get(\"allow_auth_channel_mismatch_formats\").ok()"));
|
||||
}
|
||||
|
||||
@@ -1,15 +1,282 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{sqlite::SqliteRow, Row};
|
||||
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
|
||||
|
||||
use super::{
|
||||
InMemoryProviderCatalogReadRepository, ProviderCatalogKeyListQuery,
|
||||
ProviderCatalogReadRepository, ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats,
|
||||
StoredProviderCatalogProvider,
|
||||
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogReadRepository,
|
||||
ProviderCatalogWriteRepository, StoredProviderCatalogEndpoint, StoredProviderCatalogKey,
|
||||
StoredProviderCatalogKeyPage, StoredProviderCatalogKeyStats, StoredProviderCatalogProvider,
|
||||
};
|
||||
use crate::driver::sqlite::{sqlite_optional_real, SqlitePool};
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{
|
||||
push_ci_contains_any, push_eq, push_in, push_limit_offset, push_optional_eq, SqlDialect,
|
||||
WhereClause,
|
||||
};
|
||||
|
||||
const LIST_PROVIDERS_BY_IDS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
name,
|
||||
description,
|
||||
website,
|
||||
provider_type,
|
||||
billing_type,
|
||||
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,
|
||||
enable_format_conversion,
|
||||
concurrent_limit,
|
||||
max_retries,
|
||||
proxy,
|
||||
request_timeout,
|
||||
stream_first_byte_timeout,
|
||||
config,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM providers
|
||||
WHERE id IN (
|
||||
"#;
|
||||
|
||||
const LIST_ENDPOINTS_BY_IDS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
provider_id,
|
||||
api_format,
|
||||
api_family,
|
||||
endpoint_kind,
|
||||
is_active,
|
||||
health_score,
|
||||
base_url,
|
||||
header_rules,
|
||||
body_rules,
|
||||
max_retries,
|
||||
custom_path,
|
||||
config,
|
||||
format_acceptance_config,
|
||||
proxy,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM provider_endpoints
|
||||
WHERE id IN (
|
||||
"#;
|
||||
|
||||
const LIST_ENDPOINTS_BY_PROVIDER_IDS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
provider_id,
|
||||
api_format,
|
||||
api_family,
|
||||
endpoint_kind,
|
||||
is_active,
|
||||
health_score,
|
||||
base_url,
|
||||
header_rules,
|
||||
body_rules,
|
||||
max_retries,
|
||||
custom_path,
|
||||
config,
|
||||
format_acceptance_config,
|
||||
proxy,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM provider_endpoints
|
||||
WHERE provider_id IN (
|
||||
"#;
|
||||
|
||||
const LIST_KEYS_BY_IDS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
provider_id,
|
||||
name,
|
||||
auth_type,
|
||||
capabilities,
|
||||
is_active,
|
||||
api_formats,
|
||||
auth_type_by_format,
|
||||
allow_auth_channel_mismatch_formats,
|
||||
COALESCE(api_key, encrypted_key) AS api_key,
|
||||
auth_config,
|
||||
note,
|
||||
internal_priority,
|
||||
rate_multipliers,
|
||||
global_priority_by_format,
|
||||
allowed_models,
|
||||
expires_at AS expires_at_unix_secs,
|
||||
cache_ttl_minutes,
|
||||
max_probe_interval_minutes,
|
||||
proxy,
|
||||
fingerprint,
|
||||
rpm_limit,
|
||||
concurrent_limit,
|
||||
learned_rpm_limit,
|
||||
concurrent_429_count,
|
||||
rpm_429_count,
|
||||
last_429_at AS last_429_at_unix_secs,
|
||||
last_429_type,
|
||||
adjustment_history,
|
||||
utilization_samples,
|
||||
last_probe_increase_at AS last_probe_increase_at_unix_secs,
|
||||
last_rpm_peak,
|
||||
request_count,
|
||||
total_tokens,
|
||||
CAST(total_cost_usd AS REAL) AS total_cost_usd,
|
||||
success_count,
|
||||
error_count,
|
||||
total_response_time_ms,
|
||||
last_used_at AS last_used_at_unix_secs,
|
||||
auto_fetch_models,
|
||||
last_models_fetch_at AS last_models_fetch_at_unix_secs,
|
||||
last_models_fetch_error,
|
||||
locked_models,
|
||||
model_include_patterns,
|
||||
model_exclude_patterns,
|
||||
upstream_metadata,
|
||||
oauth_invalid_at AS oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason,
|
||||
status_snapshot,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs,
|
||||
health_by_format,
|
||||
circuit_breaker_by_format
|
||||
FROM provider_api_keys
|
||||
WHERE id IN (
|
||||
"#;
|
||||
|
||||
const LIST_KEYS_BY_PROVIDER_IDS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
provider_id,
|
||||
name,
|
||||
auth_type,
|
||||
capabilities,
|
||||
is_active,
|
||||
api_formats,
|
||||
auth_type_by_format,
|
||||
allow_auth_channel_mismatch_formats,
|
||||
COALESCE(api_key, encrypted_key) AS api_key,
|
||||
auth_config,
|
||||
note,
|
||||
internal_priority,
|
||||
rate_multipliers,
|
||||
global_priority_by_format,
|
||||
allowed_models,
|
||||
expires_at AS expires_at_unix_secs,
|
||||
cache_ttl_minutes,
|
||||
max_probe_interval_minutes,
|
||||
proxy,
|
||||
fingerprint,
|
||||
rpm_limit,
|
||||
concurrent_limit,
|
||||
learned_rpm_limit,
|
||||
concurrent_429_count,
|
||||
rpm_429_count,
|
||||
last_429_at AS last_429_at_unix_secs,
|
||||
last_429_type,
|
||||
adjustment_history,
|
||||
utilization_samples,
|
||||
last_probe_increase_at AS last_probe_increase_at_unix_secs,
|
||||
last_rpm_peak,
|
||||
request_count,
|
||||
total_tokens,
|
||||
CAST(total_cost_usd AS REAL) AS total_cost_usd,
|
||||
success_count,
|
||||
error_count,
|
||||
total_response_time_ms,
|
||||
last_used_at AS last_used_at_unix_secs,
|
||||
auto_fetch_models,
|
||||
last_models_fetch_at AS last_models_fetch_at_unix_secs,
|
||||
last_models_fetch_error,
|
||||
locked_models,
|
||||
model_include_patterns,
|
||||
model_exclude_patterns,
|
||||
upstream_metadata,
|
||||
oauth_invalid_at AS oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason,
|
||||
status_snapshot,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs,
|
||||
health_by_format,
|
||||
circuit_breaker_by_format
|
||||
FROM provider_api_keys
|
||||
WHERE provider_id IN (
|
||||
"#;
|
||||
|
||||
const LIST_KEY_SUMMARIES_BY_PROVIDER_IDS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
provider_id,
|
||||
COALESCE(NULLIF(name, ''), id) AS name,
|
||||
COALESCE(NULLIF(auth_type, ''), 'summary') AS auth_type,
|
||||
NULL AS capabilities,
|
||||
is_active,
|
||||
api_formats,
|
||||
NULL AS auth_type_by_format,
|
||||
NULL AS allow_auth_channel_mismatch_formats,
|
||||
'summary' AS api_key,
|
||||
CASE
|
||||
WHEN auth_config IS NULL THEN NULL
|
||||
ELSE '{}'
|
||||
END AS auth_config,
|
||||
NULL AS note,
|
||||
NULL AS internal_priority,
|
||||
NULL AS rate_multipliers,
|
||||
NULL AS global_priority_by_format,
|
||||
NULL AS allowed_models,
|
||||
NULL AS expires_at_unix_secs,
|
||||
NULL AS cache_ttl_minutes,
|
||||
NULL AS max_probe_interval_minutes,
|
||||
NULL AS proxy,
|
||||
NULL AS fingerprint,
|
||||
NULL AS rpm_limit,
|
||||
NULL AS concurrent_limit,
|
||||
NULL AS learned_rpm_limit,
|
||||
NULL AS concurrent_429_count,
|
||||
NULL AS rpm_429_count,
|
||||
NULL AS last_429_at_unix_secs,
|
||||
NULL AS last_429_type,
|
||||
NULL AS adjustment_history,
|
||||
NULL AS utilization_samples,
|
||||
NULL AS last_probe_increase_at_unix_secs,
|
||||
NULL AS last_rpm_peak,
|
||||
NULL AS request_count,
|
||||
0 AS total_tokens,
|
||||
0.0 AS total_cost_usd,
|
||||
NULL AS success_count,
|
||||
NULL AS error_count,
|
||||
NULL AS total_response_time_ms,
|
||||
NULL AS last_used_at_unix_secs,
|
||||
FALSE AS auto_fetch_models,
|
||||
NULL AS last_models_fetch_at_unix_secs,
|
||||
NULL AS last_models_fetch_error,
|
||||
NULL AS locked_models,
|
||||
NULL AS model_include_patterns,
|
||||
NULL AS model_exclude_patterns,
|
||||
NULL AS upstream_metadata,
|
||||
NULL AS oauth_invalid_at_unix_secs,
|
||||
NULL AS oauth_invalid_reason,
|
||||
NULL AS status_snapshot,
|
||||
NULL AS created_at_unix_ms,
|
||||
NULL AS updated_at_unix_secs,
|
||||
health_by_format,
|
||||
NULL AS circuit_breaker_by_format
|
||||
FROM provider_api_keys
|
||||
WHERE provider_id IN (
|
||||
"#;
|
||||
|
||||
const LIST_KEY_STATS_BY_PROVIDER_IDS_PREFIX: &str = r#"
|
||||
SELECT
|
||||
provider_id,
|
||||
COUNT(*) AS total_keys,
|
||||
SUM(CASE WHEN is_active THEN 1 ELSE 0 END) AS active_keys
|
||||
FROM provider_api_keys
|
||||
WHERE provider_id IN (
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqliteProviderCatalogReadRepository {
|
||||
@@ -21,92 +288,232 @@ impl SqliteProviderCatalogReadRepository {
|
||||
Self { pool }
|
||||
}
|
||||
|
||||
async fn load_memory(&self) -> Result<InMemoryProviderCatalogReadRepository, DataLayerError> {
|
||||
Ok(InMemoryProviderCatalogReadRepository::seed(
|
||||
self.load_providers().await?,
|
||||
self.load_endpoints().await?,
|
||||
self.load_keys().await?,
|
||||
))
|
||||
}
|
||||
pub async fn list_providers_by_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
if provider_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
async fn load_providers(&self) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id, name, description, website, provider_type, billing_type,
|
||||
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,
|
||||
enable_format_conversion, concurrent_limit, max_retries, proxy,
|
||||
request_timeout, stream_first_byte_timeout, config,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM providers
|
||||
"#,
|
||||
let rows = build_list_query(
|
||||
LIST_PROVIDERS_BY_IDS_PREFIX,
|
||||
provider_ids,
|
||||
" ORDER BY name ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_provider_row).collect()
|
||||
}
|
||||
|
||||
async fn load_endpoints(&self) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id, provider_id, api_format, api_family, endpoint_kind, is_active,
|
||||
health_score, base_url, header_rules, body_rules, max_retries,
|
||||
custom_path, config, format_acceptance_config, proxy,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs
|
||||
FROM provider_endpoints
|
||||
WHERE api_format IS NOT NULL
|
||||
"#,
|
||||
pub async fn list_providers(
|
||||
&self,
|
||||
active_only: bool,
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
let mut builder =
|
||||
QueryBuilder::<Sqlite>::new(select_prefix_for_in(LIST_PROVIDERS_BY_IDS_PREFIX));
|
||||
let mut where_clause = WhereClause::new();
|
||||
if active_only {
|
||||
push_eq(&mut builder, &mut where_clause, "is_active", true);
|
||||
}
|
||||
builder.push(" ORDER BY provider_priority ASC, name ASC");
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_provider_row).collect()
|
||||
}
|
||||
|
||||
pub async fn list_endpoints_by_ids(
|
||||
&self,
|
||||
endpoint_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
if endpoint_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let rows = build_list_query(
|
||||
LIST_ENDPOINTS_BY_IDS_PREFIX,
|
||||
endpoint_ids,
|
||||
" ORDER BY api_format ASC, id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_endpoint_row).collect()
|
||||
}
|
||||
|
||||
async fn load_keys(&self) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
id, provider_id, name, auth_type, capabilities, is_active, api_formats,
|
||||
auth_type_by_format, allow_auth_channel_mismatch_formats,
|
||||
COALESCE(api_key, encrypted_key) AS api_key,
|
||||
auth_config, note, internal_priority, rate_multipliers,
|
||||
global_priority_by_format, allowed_models,
|
||||
expires_at AS expires_at_unix_secs,
|
||||
cache_ttl_minutes, max_probe_interval_minutes, proxy, fingerprint,
|
||||
rpm_limit, concurrent_limit, learned_rpm_limit, concurrent_429_count,
|
||||
rpm_429_count, last_429_at AS last_429_at_unix_secs, last_429_type,
|
||||
adjustment_history, utilization_samples,
|
||||
last_probe_increase_at AS last_probe_increase_at_unix_secs,
|
||||
last_rpm_peak, request_count, total_tokens, total_cost_usd,
|
||||
success_count, error_count, total_response_time_ms,
|
||||
last_used_at AS last_used_at_unix_secs, auto_fetch_models,
|
||||
last_models_fetch_at AS last_models_fetch_at_unix_secs,
|
||||
last_models_fetch_error, locked_models, model_include_patterns,
|
||||
model_exclude_patterns, upstream_metadata,
|
||||
oauth_invalid_at AS oauth_invalid_at_unix_secs,
|
||||
oauth_invalid_reason, status_snapshot,
|
||||
created_at AS created_at_unix_ms,
|
||||
updated_at AS updated_at_unix_secs,
|
||||
health_by_format, circuit_breaker_by_format
|
||||
FROM provider_api_keys
|
||||
"#,
|
||||
pub async fn list_endpoints_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
if provider_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let rows = build_list_query(
|
||||
LIST_ENDPOINTS_BY_PROVIDER_IDS_PREFIX,
|
||||
provider_ids,
|
||||
" ORDER BY provider_id ASC, api_format ASC, id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_endpoint_row).collect()
|
||||
}
|
||||
|
||||
pub async fn list_keys_by_ids(
|
||||
&self,
|
||||
key_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
if key_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let rows = build_list_query(
|
||||
LIST_KEYS_BY_IDS_PREFIX,
|
||||
key_ids,
|
||||
" ORDER BY name ASC, id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_key_row).collect()
|
||||
}
|
||||
|
||||
pub async fn list_keys_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
if provider_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let rows = build_list_query(
|
||||
LIST_KEYS_BY_PROVIDER_IDS_PREFIX,
|
||||
provider_ids,
|
||||
" ORDER BY provider_id ASC, name ASC, id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_key_row).collect()
|
||||
}
|
||||
|
||||
pub async fn list_key_summaries_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
if provider_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let rows = build_list_query(
|
||||
LIST_KEY_SUMMARIES_BY_PROVIDER_IDS_PREFIX,
|
||||
provider_ids,
|
||||
" ORDER BY provider_id ASC, id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_key_row).collect()
|
||||
}
|
||||
|
||||
pub async fn list_keys_page(
|
||||
&self,
|
||||
query: &ProviderCatalogKeyListQuery,
|
||||
) -> Result<StoredProviderCatalogKeyPage, DataLayerError> {
|
||||
if query.provider_id.trim().is_empty() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"provider catalog provider_id is empty".to_string(),
|
||||
));
|
||||
}
|
||||
|
||||
let offset = i64::try_from(query.offset).map_err(|_| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"invalid provider catalog key offset: {}",
|
||||
query.offset
|
||||
))
|
||||
})?;
|
||||
let limit = i64::try_from(query.limit).map_err(|_| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"invalid provider catalog key limit: {}",
|
||||
query.limit
|
||||
))
|
||||
})?;
|
||||
let order_by = match query.order {
|
||||
ProviderCatalogKeyListOrder::Name => "internal_priority ASC, name ASC, id ASC",
|
||||
ProviderCatalogKeyListOrder::CreatedAt => {
|
||||
"internal_priority ASC, COALESCE(created_at, 0) ASC, id ASC"
|
||||
}
|
||||
ProviderCatalogKeyListOrder::CreatedAtAsc => {
|
||||
"created_at IS NULL ASC, created_at ASC, name ASC, id ASC"
|
||||
}
|
||||
ProviderCatalogKeyListOrder::CreatedAtDesc => {
|
||||
"created_at IS NULL ASC, created_at DESC, name ASC, id ASC"
|
||||
}
|
||||
ProviderCatalogKeyListOrder::LastUsedAtAsc => {
|
||||
"last_used_at IS NULL ASC, last_used_at ASC, name ASC, id ASC"
|
||||
}
|
||||
ProviderCatalogKeyListOrder::LastUsedAtDesc => {
|
||||
"last_used_at IS NULL ASC, last_used_at DESC, name ASC, id ASC"
|
||||
}
|
||||
};
|
||||
|
||||
let mut count_builder =
|
||||
QueryBuilder::<Sqlite>::new("SELECT COUNT(*) AS total FROM provider_api_keys");
|
||||
let mut count_where = WhereClause::new();
|
||||
apply_key_page_filters(&mut count_builder, &mut count_where, query);
|
||||
let total = count_builder
|
||||
.build_query_scalar::<i64>()
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.max(0) as usize;
|
||||
|
||||
let mut list_builder =
|
||||
QueryBuilder::<Sqlite>::new(select_prefix_for_in(LIST_KEYS_BY_IDS_PREFIX));
|
||||
let mut list_where = WhereClause::new();
|
||||
apply_key_page_filters(&mut list_builder, &mut list_where, query);
|
||||
list_builder.push(" ORDER BY ").push(order_by);
|
||||
push_limit_offset(&mut list_builder, limit, offset);
|
||||
let rows = list_builder
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let items = rows
|
||||
.iter()
|
||||
.map(map_key_row)
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
|
||||
Ok(StoredProviderCatalogKeyPage { items, total })
|
||||
}
|
||||
|
||||
pub async fn list_key_stats_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKeyStats>, DataLayerError> {
|
||||
if provider_ids.is_empty() {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let rows = build_list_query(
|
||||
LIST_KEY_STATS_BY_PROVIDER_IDS_PREFIX,
|
||||
provider_ids,
|
||||
"\nGROUP BY provider_id\nORDER BY provider_id ASC",
|
||||
)
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_key_stats_row).collect()
|
||||
}
|
||||
|
||||
pub async fn create_provider(
|
||||
&self,
|
||||
provider: &StoredProviderCatalogProvider,
|
||||
@@ -671,6 +1078,11 @@ WHERE id = ?
|
||||
"provider_api_keys.last_probe_increase_at",
|
||||
)?)
|
||||
.bind(optional_i64_from_u32(key.last_rpm_peak))
|
||||
.bind(optional_i64_from_u64(
|
||||
key.last_models_fetch_at_unix_secs,
|
||||
"provider_api_keys.last_models_fetch_at",
|
||||
)?)
|
||||
.bind(&key.last_models_fetch_error)
|
||||
.bind(updated_at)
|
||||
.bind(&key.id)
|
||||
.execute(&self.pool)
|
||||
@@ -865,81 +1277,63 @@ impl ProviderCatalogReadRepository for SqliteProviderCatalogReadRepository {
|
||||
&self,
|
||||
active_only: bool,
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
self.load_memory().await?.list_providers(active_only).await
|
||||
Self::list_providers(self, active_only).await
|
||||
}
|
||||
|
||||
async fn list_providers_by_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_providers_by_ids(provider_ids)
|
||||
.await
|
||||
Self::list_providers_by_ids(self, provider_ids).await
|
||||
}
|
||||
|
||||
async fn list_endpoints_by_ids(
|
||||
&self,
|
||||
endpoint_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_endpoints_by_ids(endpoint_ids)
|
||||
.await
|
||||
Self::list_endpoints_by_ids(self, endpoint_ids).await
|
||||
}
|
||||
|
||||
async fn list_endpoints_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogEndpoint>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_endpoints_by_provider_ids(provider_ids)
|
||||
.await
|
||||
Self::list_endpoints_by_provider_ids(self, provider_ids).await
|
||||
}
|
||||
|
||||
async fn list_keys_by_ids(
|
||||
&self,
|
||||
key_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
self.load_memory().await?.list_keys_by_ids(key_ids).await
|
||||
Self::list_keys_by_ids(self, key_ids).await
|
||||
}
|
||||
|
||||
async fn list_keys_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_keys_by_provider_ids(provider_ids)
|
||||
.await
|
||||
Self::list_keys_by_provider_ids(self, provider_ids).await
|
||||
}
|
||||
|
||||
async fn list_key_summaries_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKey>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_key_summaries_by_provider_ids(provider_ids)
|
||||
.await
|
||||
Self::list_key_summaries_by_provider_ids(self, provider_ids).await
|
||||
}
|
||||
|
||||
async fn list_keys_page(
|
||||
&self,
|
||||
query: &ProviderCatalogKeyListQuery,
|
||||
) -> Result<StoredProviderCatalogKeyPage, DataLayerError> {
|
||||
self.load_memory().await?.list_keys_page(query).await
|
||||
Self::list_keys_page(self, query).await
|
||||
}
|
||||
|
||||
async fn list_key_stats_by_provider_ids(
|
||||
&self,
|
||||
provider_ids: &[String],
|
||||
) -> Result<Vec<StoredProviderCatalogKeyStats>, DataLayerError> {
|
||||
self.load_memory()
|
||||
.await?
|
||||
.list_key_stats_by_provider_ids(provider_ids)
|
||||
.await
|
||||
Self::list_key_stats_by_provider_ids(self, provider_ids).await
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1058,6 +1452,61 @@ impl ProviderCatalogWriteRepository for SqliteProviderCatalogReadRepository {
|
||||
}
|
||||
}
|
||||
|
||||
fn build_list_query<'a>(
|
||||
prefix: &'static str,
|
||||
ids: &'a [String],
|
||||
suffix: &'static str,
|
||||
) -> QueryBuilder<'a, Sqlite> {
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(select_prefix_for_in(prefix));
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_in(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
in_column_for_prefix(prefix),
|
||||
ids,
|
||||
);
|
||||
builder.push(suffix);
|
||||
builder
|
||||
}
|
||||
|
||||
fn select_prefix_for_in(prefix: &'static str) -> &'static str {
|
||||
prefix
|
||||
.rsplit_once("\nWHERE ")
|
||||
.map(|(select_prefix, _)| select_prefix)
|
||||
.expect("provider catalog IN query prefix must contain WHERE")
|
||||
}
|
||||
|
||||
fn in_column_for_prefix(prefix: &'static str) -> &'static str {
|
||||
prefix
|
||||
.rsplit_once("\nWHERE ")
|
||||
.and_then(|(_, predicate)| predicate.trim().strip_suffix("IN ("))
|
||||
.map(str::trim)
|
||||
.expect("provider catalog IN query prefix must end with IN (")
|
||||
}
|
||||
|
||||
fn apply_key_page_filters<'a>(
|
||||
builder: &mut QueryBuilder<'a, Sqlite>,
|
||||
where_clause: &mut WhereClause,
|
||||
query: &'a ProviderCatalogKeyListQuery,
|
||||
) {
|
||||
push_eq(
|
||||
builder,
|
||||
where_clause,
|
||||
"provider_id",
|
||||
query.provider_id.clone(),
|
||||
);
|
||||
if let Some(search) = query.search.as_deref() {
|
||||
push_ci_contains_any(
|
||||
builder,
|
||||
where_clause,
|
||||
SqlDialect::Sqlite,
|
||||
&["name", "id"],
|
||||
search,
|
||||
);
|
||||
}
|
||||
push_optional_eq(builder, where_clause, "is_active", query.is_active);
|
||||
}
|
||||
|
||||
fn current_unix_secs() -> u64 {
|
||||
chrono::Utc::now().timestamp().max(0) as u64
|
||||
}
|
||||
@@ -1210,6 +1659,8 @@ SET
|
||||
utilization_samples = ?,
|
||||
last_probe_increase_at = ?,
|
||||
last_rpm_peak = ?,
|
||||
last_models_fetch_at = ?,
|
||||
last_models_fetch_error = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
"#
|
||||
@@ -1346,6 +1797,14 @@ fn map_endpoint_row(row: &SqliteRow) -> Result<StoredProviderCatalogEndpoint, Da
|
||||
)
|
||||
}
|
||||
|
||||
fn map_key_stats_row(row: &SqliteRow) -> Result<StoredProviderCatalogKeyStats, DataLayerError> {
|
||||
StoredProviderCatalogKeyStats::new(
|
||||
row.try_get("provider_id").map_sql_err()?,
|
||||
row.try_get("total_keys").map_sql_err()?,
|
||||
row.try_get("active_keys").map_sql_err()?,
|
||||
)
|
||||
}
|
||||
|
||||
fn map_key_row(row: &SqliteRow) -> Result<StoredProviderCatalogKey, DataLayerError> {
|
||||
let total_cost_usd = sqlite_optional_real(row, "total_cost_usd")?.unwrap_or(0.0);
|
||||
if !total_cost_usd.is_finite() {
|
||||
@@ -1542,8 +2001,8 @@ mod tests {
|
||||
use super::SqliteProviderCatalogReadRepository;
|
||||
use crate::lifecycle::migrate::run_sqlite_migrations;
|
||||
use crate::repository::provider_catalog::{
|
||||
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, ProviderCatalogReadRepository,
|
||||
StoredProviderCatalogEndpoint, StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
ProviderCatalogKeyListOrder, ProviderCatalogKeyListQuery, StoredProviderCatalogEndpoint,
|
||||
StoredProviderCatalogKey, StoredProviderCatalogProvider,
|
||||
};
|
||||
use serde_json::json;
|
||||
|
||||
@@ -1702,7 +2161,7 @@ mod tests {
|
||||
assert_eq!(updated_endpoint.health_score, 0.5);
|
||||
assert!(!updated_endpoint.is_active);
|
||||
|
||||
let key = StoredProviderCatalogKey::new(
|
||||
let mut key = StoredProviderCatalogKey::new(
|
||||
"key-write-1".to_string(),
|
||||
"provider-write-1".to_string(),
|
||||
"Default Key".to_string(),
|
||||
@@ -1740,23 +2199,36 @@ mod tests {
|
||||
Some(json!({"openai:chat":{"score":1}})),
|
||||
Some(json!({"openai:chat":{"open":false}})),
|
||||
);
|
||||
key.last_models_fetch_at_unix_secs = Some(1_730_000_100);
|
||||
key.last_models_fetch_error = Some("stale models fetch error".to_string());
|
||||
let created_key = repository
|
||||
.create_key(&key)
|
||||
.await
|
||||
.expect("key should create");
|
||||
assert_eq!(created_key.concurrent_limit, Some(3));
|
||||
assert_eq!(created_key.total_tokens, 1234);
|
||||
assert_eq!(
|
||||
created_key.last_models_fetch_error.as_deref(),
|
||||
Some("stale models fetch error")
|
||||
);
|
||||
|
||||
let mut updated_key = created_key.clone();
|
||||
updated_key.name = "Updated Key".to_string();
|
||||
updated_key.is_active = false;
|
||||
updated_key.upstream_metadata = Some(json!({"models":["gpt-4.1"]}));
|
||||
updated_key.last_models_fetch_at_unix_secs = Some(1_730_000_200);
|
||||
updated_key.last_models_fetch_error = None;
|
||||
let updated_key = repository
|
||||
.update_key(&updated_key)
|
||||
.await
|
||||
.expect("key should update");
|
||||
assert_eq!(updated_key.name, "Updated Key");
|
||||
assert!(!updated_key.is_active);
|
||||
assert_eq!(
|
||||
updated_key.last_models_fetch_at_unix_secs,
|
||||
Some(1_730_000_200)
|
||||
);
|
||||
assert_eq!(updated_key.last_models_fetch_error, None);
|
||||
|
||||
assert!(repository
|
||||
.update_key_upstream_metadata(
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
use async_trait::async_trait;
|
||||
use futures_util::TryStreamExt;
|
||||
use sha2::{Digest, Sha256};
|
||||
use sqlx::{postgres::PgRow, PgPool, Row};
|
||||
use sqlx::{postgres::PgRow, PgPool, Postgres, QueryBuilder, Row};
|
||||
|
||||
use super::types::{
|
||||
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
|
||||
@@ -17,6 +17,7 @@ use crate::{
|
||||
error::{postgres_error, SqlxResultExt},
|
||||
DataLayerError,
|
||||
};
|
||||
use aether_data_query::{push_eq, push_limit, WhereClause};
|
||||
|
||||
const FIND_PROXY_NODE_SQL: &str = r#"
|
||||
SELECT
|
||||
@@ -54,41 +55,6 @@ WHERE id = $1
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const LIST_PROXY_NODES_SQL: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
name,
|
||||
ip,
|
||||
port,
|
||||
region,
|
||||
is_manual,
|
||||
proxy_url,
|
||||
proxy_username,
|
||||
proxy_password,
|
||||
CAST(status AS TEXT) AS status,
|
||||
registered_by,
|
||||
EXTRACT(EPOCH FROM last_heartbeat_at)::bigint AS last_heartbeat_at_unix_secs,
|
||||
heartbeat_interval,
|
||||
active_connections,
|
||||
total_requests,
|
||||
CAST(avg_latency_ms AS DOUBLE PRECISION) AS avg_latency_ms,
|
||||
failed_requests,
|
||||
dns_failures,
|
||||
stream_errors,
|
||||
proxy_metadata,
|
||||
hardware_info,
|
||||
estimated_max_concurrency,
|
||||
tunnel_mode,
|
||||
tunnel_connected,
|
||||
EXTRACT(EPOCH FROM tunnel_connected_at)::bigint AS tunnel_connected_at_unix_secs,
|
||||
remote_config,
|
||||
config_version,
|
||||
EXTRACT(EPOCH FROM created_at)::bigint AS created_at_unix_ms,
|
||||
EXTRACT(EPOCH FROM updated_at)::bigint AS updated_at_unix_secs
|
||||
FROM proxy_nodes
|
||||
ORDER BY name ASC, id ASC
|
||||
"#;
|
||||
|
||||
const LIST_PROXY_NODE_EVENTS_SQL: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
@@ -103,23 +69,6 @@ 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
|
||||
@@ -870,7 +819,9 @@ impl SqlxProxyNodeRepository {
|
||||
#[async_trait]
|
||||
impl ProxyNodeReadRepository for SqlxProxyNodeRepository {
|
||||
async fn list_proxy_nodes(&self) -> Result<Vec<StoredProxyNode>, DataLayerError> {
|
||||
let mut rows = sqlx::query(LIST_PROXY_NODES_SQL).fetch(&self.pool);
|
||||
let mut builder = QueryBuilder::<Postgres>::new(proxy_node_columns());
|
||||
builder.push(" ORDER BY name ASC, id ASC");
|
||||
let mut rows = builder.build().fetch(&self.pool);
|
||||
let mut items = Vec::new();
|
||||
while let Some(row) = rows.try_next().await.map_postgres_err()? {
|
||||
items.push(Self::row_to_stored(&row)?);
|
||||
@@ -882,8 +833,12 @@ impl ProxyNodeReadRepository for SqlxProxyNodeRepository {
|
||||
&self,
|
||||
node_id: &str,
|
||||
) -> Result<Option<StoredProxyNode>, DataLayerError> {
|
||||
let row = sqlx::query(FIND_PROXY_NODE_SQL)
|
||||
.bind(node_id)
|
||||
let mut builder = QueryBuilder::<Postgres>::new(proxy_node_columns());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "id", node_id.to_string());
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
@@ -895,10 +850,17 @@ impl ProxyNodeReadRepository for SqlxProxyNodeRepository {
|
||||
node_id: &str,
|
||||
limit: usize,
|
||||
) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> {
|
||||
let mut rows = sqlx::query(LIST_PROXY_NODE_EVENTS_SQL)
|
||||
.bind(node_id)
|
||||
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
|
||||
.fetch(&self.pool);
|
||||
let mut builder = QueryBuilder::<Postgres>::new(proxy_node_event_columns());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"node_id",
|
||||
node_id.to_string(),
|
||||
);
|
||||
builder.push(" ORDER BY created_at DESC, id DESC");
|
||||
push_limit(&mut builder, i64::try_from(limit).unwrap_or(i64::MAX));
|
||||
let mut rows = builder.build().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)?);
|
||||
@@ -911,13 +873,38 @@ impl ProxyNodeReadRepository for SqlxProxyNodeRepository {
|
||||
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 builder = QueryBuilder::<Postgres>::new(proxy_node_event_columns());
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"node_id",
|
||||
node_id.to_string(),
|
||||
);
|
||||
if let Some(from_unix_secs) = query.from_unix_secs {
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push("created_at >= TO_TIMESTAMP(")
|
||||
.push_bind(from_unix_secs as f64)
|
||||
.push("::double precision)");
|
||||
}
|
||||
if let Some(to_unix_secs) = query.to_unix_secs {
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push("created_at <= TO_TIMESTAMP(")
|
||||
.push_bind(to_unix_secs as f64)
|
||||
.push("::double precision)");
|
||||
}
|
||||
if let Some(event_type) = query.event_type.as_deref() {
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push("LOWER(CAST(event_type AS TEXT)) = LOWER(")
|
||||
.push_bind(event_type.to_string())
|
||||
.push("::text)");
|
||||
}
|
||||
builder.push(" ORDER BY created_at DESC, id DESC");
|
||||
push_limit(&mut builder, i64::try_from(query.limit).unwrap_or(i64::MAX));
|
||||
let mut rows = builder.build().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)?);
|
||||
@@ -974,6 +961,20 @@ impl ProxyNodeReadRepository for SqlxProxyNodeRepository {
|
||||
}
|
||||
}
|
||||
|
||||
fn proxy_node_columns() -> &'static str {
|
||||
FIND_PROXY_NODE_SQL
|
||||
.split_once("WHERE id = $1")
|
||||
.map(|(prefix, _)| prefix)
|
||||
.unwrap_or(FIND_PROXY_NODE_SQL)
|
||||
}
|
||||
|
||||
fn proxy_node_event_columns() -> &'static str {
|
||||
LIST_PROXY_NODE_EVENTS_SQL
|
||||
.split_once("WHERE node_id = $1")
|
||||
.map(|(prefix, _)| prefix)
|
||||
.unwrap_or(LIST_PROXY_NODE_EVENTS_SQL)
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl ProxyNodeWriteRepository for SqlxProxyNodeRepository {
|
||||
async fn reset_stale_tunnel_statuses(&self) -> Result<usize, DataLayerError> {
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{sqlite::SqliteRow, Row};
|
||||
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
|
||||
|
||||
use super::types::{
|
||||
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
|
||||
@@ -14,6 +14,7 @@ use super::types::{
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
use aether_data_query::{push_eq, push_limit, WhereClause};
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqliteProxyNodeReadRepository {
|
||||
@@ -337,13 +338,23 @@ SELECT
|
||||
FROM proxy_nodes
|
||||
"#;
|
||||
|
||||
const PROXY_NODE_EVENT_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
id,
|
||||
node_id,
|
||||
event_type,
|
||||
detail,
|
||||
event_metadata,
|
||||
created_at AS created_at_unix_ms
|
||||
FROM proxy_node_events
|
||||
"#;
|
||||
|
||||
#[async_trait]
|
||||
impl ProxyNodeReadRepository for SqliteProxyNodeReadRepository {
|
||||
async fn list_proxy_nodes(&self) -> Result<Vec<StoredProxyNode>, DataLayerError> {
|
||||
let rows = sqlx::query(&format!("{PROXY_NODE_COLUMNS} ORDER BY name ASC, id ASC"))
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(PROXY_NODE_COLUMNS);
|
||||
builder.push(" ORDER BY name ASC, id ASC");
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_proxy_node_row).collect()
|
||||
}
|
||||
|
||||
@@ -351,8 +362,12 @@ impl ProxyNodeReadRepository for SqliteProxyNodeReadRepository {
|
||||
&self,
|
||||
node_id: &str,
|
||||
) -> Result<Option<StoredProxyNode>, DataLayerError> {
|
||||
let row = sqlx::query(&format!("{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1"))
|
||||
.bind(node_id)
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(PROXY_NODE_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(&mut builder, &mut where_clause, "id", node_id.to_string());
|
||||
push_limit(&mut builder, 1);
|
||||
let row = builder
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
@@ -364,26 +379,17 @@ impl ProxyNodeReadRepository for SqliteProxyNodeReadRepository {
|
||||
node_id: &str,
|
||||
limit: usize,
|
||||
) -> 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 = ?
|
||||
ORDER BY created_at DESC, id DESC
|
||||
LIMIT ?
|
||||
"#,
|
||||
)
|
||||
.bind(node_id)
|
||||
.bind(i64::try_from(limit).unwrap_or(i64::MAX))
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(PROXY_NODE_EVENT_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"node_id",
|
||||
node_id.to_string(),
|
||||
);
|
||||
builder.push(" ORDER BY created_at DESC, id DESC");
|
||||
push_limit(&mut builder, i64::try_from(limit).unwrap_or(i64::MAX));
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_proxy_node_event_row).collect()
|
||||
}
|
||||
|
||||
@@ -392,51 +398,36 @@ LIMIT ?
|
||||
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()?;
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(PROXY_NODE_EVENT_COLUMNS);
|
||||
let mut where_clause = WhereClause::new();
|
||||
push_eq(
|
||||
&mut builder,
|
||||
&mut where_clause,
|
||||
"node_id",
|
||||
node_id.to_string(),
|
||||
);
|
||||
if let Some(from_unix_secs) = query.from_unix_secs {
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push("created_at >= ")
|
||||
.push_bind(i64::try_from(from_unix_secs).unwrap_or(i64::MAX));
|
||||
}
|
||||
if let Some(to_unix_secs) = query.to_unix_secs {
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push("created_at <= ")
|
||||
.push_bind(i64::try_from(to_unix_secs).unwrap_or(i64::MAX));
|
||||
}
|
||||
if let Some(event_type) = query.event_type.as_deref() {
|
||||
where_clause.push_next(&mut builder);
|
||||
builder
|
||||
.push("LOWER(event_type) = LOWER(")
|
||||
.push_bind(event_type.to_string())
|
||||
.push(")");
|
||||
}
|
||||
builder.push(" ORDER BY created_at DESC, id DESC");
|
||||
push_limit(&mut builder, i64::try_from(query.limit).unwrap_or(i64::MAX));
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
rows.iter().map(map_proxy_node_event_row).collect()
|
||||
}
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@ mod mysql;
|
||||
mod postgres;
|
||||
mod sqlite;
|
||||
|
||||
use aether_data_query::{DialectSql, SelectColumn, SelectQuery};
|
||||
|
||||
#[allow(unused_imports)]
|
||||
pub(crate) use aether_data_contracts::repository::quota::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaRepository, ProviderQuotaWriteRepository,
|
||||
@@ -12,3 +14,35 @@ pub use memory::InMemoryProviderQuotaRepository;
|
||||
pub use mysql::MysqlProviderQuotaRepository;
|
||||
pub use postgres::SqlxProviderQuotaRepository;
|
||||
pub use sqlite::SqliteProviderQuotaRepository;
|
||||
|
||||
fn quota_snapshot_select() -> SelectQuery<'static> {
|
||||
SelectQuery::new("providers").select_columns([
|
||||
SelectColumn::expr("id").alias("provider_id"),
|
||||
SelectColumn::expr(
|
||||
DialectSql::common("billing_type").with_postgres("CAST(billing_type AS TEXT)"),
|
||||
)
|
||||
.alias("billing_type"),
|
||||
SelectColumn::expr(DialectSql::dialect(
|
||||
"CAST(monthly_quota_usd AS DOUBLE PRECISION)",
|
||||
"CAST(monthly_quota_usd AS REAL)",
|
||||
))
|
||||
.alias("monthly_quota_usd"),
|
||||
SelectColumn::expr(DialectSql::dialect(
|
||||
"CAST(COALESCE(monthly_used_usd, 0) AS DOUBLE PRECISION)",
|
||||
"CAST(COALESCE(monthly_used_usd, 0) AS REAL)",
|
||||
))
|
||||
.alias("monthly_used_usd"),
|
||||
SelectColumn::expr("quota_reset_day"),
|
||||
SelectColumn::expr(DialectSql::dialect(
|
||||
"CAST(EXTRACT(EPOCH FROM quota_last_reset_at) AS BIGINT)",
|
||||
"quota_last_reset_at",
|
||||
))
|
||||
.alias("quota_last_reset_at_unix_secs"),
|
||||
SelectColumn::expr(DialectSql::dialect(
|
||||
"CAST(EXTRACT(EPOCH FROM quota_expires_at) AS BIGINT)",
|
||||
"quota_expires_at",
|
||||
))
|
||||
.alias("quota_expires_at_unix_secs"),
|
||||
SelectColumn::expr("is_active"),
|
||||
])
|
||||
}
|
||||
|
||||
@@ -1,40 +1,12 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{PgPool, Row};
|
||||
use sqlx::{PgPool, Postgres, Row};
|
||||
|
||||
use super::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
|
||||
quota_snapshot_select, ProviderQuotaReadRepository, ProviderQuotaWriteRepository,
|
||||
StoredProviderQuotaSnapshot,
|
||||
};
|
||||
use crate::{error::SqlxResultExt, DataLayerError};
|
||||
|
||||
const FIND_BY_PROVIDER_ID_SQL: &str = r#"
|
||||
SELECT
|
||||
id AS provider_id,
|
||||
CAST(billing_type AS TEXT) AS billing_type,
|
||||
CAST(monthly_quota_usd AS DOUBLE PRECISION) AS monthly_quota_usd,
|
||||
CAST(COALESCE(monthly_used_usd, 0) AS DOUBLE PRECISION) AS monthly_used_usd,
|
||||
quota_reset_day,
|
||||
CAST(EXTRACT(EPOCH FROM quota_last_reset_at) AS BIGINT) AS quota_last_reset_at_unix_secs,
|
||||
CAST(EXTRACT(EPOCH FROM quota_expires_at) AS BIGINT) AS quota_expires_at_unix_secs,
|
||||
is_active
|
||||
FROM providers
|
||||
WHERE id = $1
|
||||
LIMIT 1
|
||||
"#;
|
||||
|
||||
const FIND_BY_PROVIDER_IDS_SQL: &str = r#"
|
||||
SELECT
|
||||
id AS provider_id,
|
||||
CAST(billing_type AS TEXT) AS billing_type,
|
||||
CAST(monthly_quota_usd AS DOUBLE PRECISION) AS monthly_quota_usd,
|
||||
CAST(COALESCE(monthly_used_usd, 0) AS DOUBLE PRECISION) AS monthly_used_usd,
|
||||
quota_reset_day,
|
||||
CAST(EXTRACT(EPOCH FROM quota_last_reset_at) AS BIGINT) AS quota_last_reset_at_unix_secs,
|
||||
CAST(EXTRACT(EPOCH FROM quota_expires_at) AS BIGINT) AS quota_expires_at_unix_secs,
|
||||
is_active
|
||||
FROM providers
|
||||
WHERE id = ANY($1::TEXT[])
|
||||
ORDER BY id ASC
|
||||
"#;
|
||||
use aether_data_query::SqlDialect;
|
||||
|
||||
const RESET_DUE_SQL: &str = r#"
|
||||
UPDATE providers
|
||||
@@ -68,8 +40,11 @@ impl ProviderQuotaReadRepository for SqlxProviderQuotaRepository {
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<Option<StoredProviderQuotaSnapshot>, DataLayerError> {
|
||||
let row = sqlx::query(FIND_BY_PROVIDER_ID_SQL)
|
||||
.bind(provider_id)
|
||||
let mut statement = quota_snapshot_select().statement::<Postgres>(SqlDialect::Postgres);
|
||||
statement.where_eq("id", provider_id.to_string()).limit(1);
|
||||
let row = statement
|
||||
.finish()
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?;
|
||||
@@ -84,8 +59,13 @@ impl ProviderQuotaReadRepository for SqlxProviderQuotaRepository {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
sqlx::query(FIND_BY_PROVIDER_IDS_SQL)
|
||||
.bind(provider_ids)
|
||||
let mut statement = quota_snapshot_select().statement::<Postgres>(SqlDialect::Postgres);
|
||||
statement
|
||||
.where_in("id", provider_ids)
|
||||
.order_by_sql("id ASC");
|
||||
statement
|
||||
.finish()
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_postgres_err()?
|
||||
|
||||
@@ -1,25 +1,14 @@
|
||||
use async_trait::async_trait;
|
||||
use sqlx::{sqlite::SqliteRow, QueryBuilder, Row, Sqlite};
|
||||
use sqlx::{sqlite::SqliteRow, Row, Sqlite};
|
||||
|
||||
use super::{
|
||||
ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
|
||||
quota_snapshot_select, ProviderQuotaReadRepository, ProviderQuotaWriteRepository,
|
||||
StoredProviderQuotaSnapshot,
|
||||
};
|
||||
use crate::driver::sqlite::{sqlite_optional_real, sqlite_real, SqlitePool};
|
||||
use crate::error::SqlResultExt;
|
||||
use crate::DataLayerError;
|
||||
|
||||
const QUOTA_COLUMNS: &str = r#"
|
||||
SELECT
|
||||
id AS provider_id,
|
||||
billing_type,
|
||||
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,
|
||||
is_active
|
||||
FROM providers
|
||||
"#;
|
||||
use aether_data_query::SqlDialect;
|
||||
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct SqliteProviderQuotaRepository {
|
||||
@@ -38,8 +27,11 @@ impl ProviderQuotaReadRepository for SqliteProviderQuotaRepository {
|
||||
&self,
|
||||
provider_id: &str,
|
||||
) -> Result<Option<StoredProviderQuotaSnapshot>, DataLayerError> {
|
||||
let row = sqlx::query(&format!("{QUOTA_COLUMNS} WHERE id = ? LIMIT 1"))
|
||||
.bind(provider_id)
|
||||
let mut statement = quota_snapshot_select().statement::<Sqlite>(SqlDialect::Sqlite);
|
||||
statement.where_eq("id", provider_id.to_string()).limit(1);
|
||||
let row = statement
|
||||
.finish()
|
||||
.build()
|
||||
.fetch_optional(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
@@ -54,16 +46,16 @@ impl ProviderQuotaReadRepository for SqliteProviderQuotaRepository {
|
||||
return Ok(Vec::new());
|
||||
}
|
||||
|
||||
let mut builder = QueryBuilder::<Sqlite>::new(QUOTA_COLUMNS);
|
||||
builder.push(" WHERE id IN (");
|
||||
{
|
||||
let mut separated = builder.separated(", ");
|
||||
for provider_id in provider_ids {
|
||||
separated.push_bind(provider_id);
|
||||
}
|
||||
}
|
||||
builder.push(") ORDER BY id ASC");
|
||||
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
|
||||
let mut statement = quota_snapshot_select().statement::<Sqlite>(SqlDialect::Sqlite);
|
||||
statement
|
||||
.where_in("id", provider_ids)
|
||||
.order_by_sql("id ASC");
|
||||
let rows = statement
|
||||
.finish()
|
||||
.build()
|
||||
.fetch_all(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
rows.iter().map(map_row).collect()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -462,6 +462,19 @@ FOR UPDATE
|
||||
}
|
||||
None => Some(0.0),
|
||||
};
|
||||
if let Some(row) = wallet_row.as_ref() {
|
||||
let wallet_id: String = row.try_get("id").map_sql_err()?;
|
||||
let before_recharge: f64 = row.try_get("balance").map_sql_err()?;
|
||||
let before_gift: f64 = row.try_get("gift_balance").map_sql_err()?;
|
||||
let before_total = before_recharge + before_gift;
|
||||
settlement.wallet_id = Some(wallet_id);
|
||||
settlement.wallet_balance_before = Some(before_total);
|
||||
settlement.wallet_balance_after = Some(before_total);
|
||||
settlement.wallet_recharge_balance_before = Some(before_recharge);
|
||||
settlement.wallet_recharge_balance_after = Some(before_recharge);
|
||||
settlement.wallet_gift_balance_before = Some(before_gift);
|
||||
settlement.wallet_gift_balance_after = Some(before_gift);
|
||||
}
|
||||
|
||||
let wallet_debit_cost_usd = if !api_key_is_standalone {
|
||||
if let Some(user_id) = input.user_id.as_deref().filter(|value| !value.is_empty()) {
|
||||
|
||||
@@ -528,6 +528,20 @@ LIMIT 1
|
||||
}
|
||||
None => Some(0.0),
|
||||
};
|
||||
if let Some(row) = wallet_row.as_ref() {
|
||||
let wallet_id: String = row.try_get("id").map_postgres_err()?;
|
||||
let before_recharge: f64 = row.try_get("balance").map_postgres_err()?;
|
||||
let before_gift: f64 =
|
||||
row.try_get("gift_balance").map_postgres_err()?;
|
||||
let before_total = before_recharge + before_gift;
|
||||
settlement.wallet_id = Some(wallet_id);
|
||||
settlement.wallet_balance_before = Some(before_total);
|
||||
settlement.wallet_balance_after = Some(before_total);
|
||||
settlement.wallet_recharge_balance_before = Some(before_recharge);
|
||||
settlement.wallet_recharge_balance_after = Some(before_recharge);
|
||||
settlement.wallet_gift_balance_before = Some(before_gift);
|
||||
settlement.wallet_gift_balance_after = Some(before_gift);
|
||||
}
|
||||
|
||||
let wallet_debit_cost_usd = if !api_key_is_standalone {
|
||||
if let Some(user_id) =
|
||||
|
||||
@@ -271,7 +271,7 @@ ORDER BY expires_at ASC, created_at ASC, id ASC
|
||||
allow_wallet_overage &= grant.allow_wallet_overage;
|
||||
let used = sqlx::query_scalar::<_, f64>(
|
||||
r#"
|
||||
SELECT COALESCE(SUM(amount_usd), 0)
|
||||
SELECT CAST(COALESCE(SUM(amount_usd), 0) AS REAL)
|
||||
FROM entitlement_usage_ledgers
|
||||
WHERE user_entitlement_id = ?
|
||||
AND usage_date = ?
|
||||
@@ -473,6 +473,19 @@ LIMIT 1
|
||||
}
|
||||
None => Some(0.0),
|
||||
};
|
||||
if let Some(row) = wallet_row.as_ref() {
|
||||
let wallet_id: String = row.try_get("id").map_sql_err()?;
|
||||
let before_recharge = sqlite_real(row, "balance")?;
|
||||
let before_gift = sqlite_real(row, "gift_balance")?;
|
||||
let before_total = before_recharge + before_gift;
|
||||
settlement.wallet_id = Some(wallet_id);
|
||||
settlement.wallet_balance_before = Some(before_total);
|
||||
settlement.wallet_balance_after = Some(before_total);
|
||||
settlement.wallet_recharge_balance_before = Some(before_recharge);
|
||||
settlement.wallet_recharge_balance_after = Some(before_recharge);
|
||||
settlement.wallet_gift_balance_before = Some(before_gift);
|
||||
settlement.wallet_gift_balance_after = Some(before_gift);
|
||||
}
|
||||
|
||||
let wallet_debit_cost_usd = if !api_key_is_standalone {
|
||||
if let Some(user_id) = input.user_id.as_deref().filter(|value| !value.is_empty()) {
|
||||
@@ -800,6 +813,58 @@ mod tests {
|
||||
assert_eq!(wallet_total, 12.0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_repository_records_wallet_for_quota_covered_user_usage() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
seed_quota_covered_settlement_rows(&pool).await;
|
||||
|
||||
let repository = SqliteSettlementRepository::new(pool.clone());
|
||||
let settlement = repository
|
||||
.settle_usage(UsageSettlementInput {
|
||||
request_id: "request-quota-covered".to_string(),
|
||||
user_id: Some("user-quota".to_string()),
|
||||
api_key_id: Some("key-quota".to_string()),
|
||||
api_key_is_standalone: false,
|
||||
provider_id: None,
|
||||
status: "completed".to_string(),
|
||||
billing_status: "pending".to_string(),
|
||||
total_cost_usd: 3.0,
|
||||
actual_total_cost_usd: 2.0,
|
||||
finalized_at_unix_secs: Some(1_260),
|
||||
})
|
||||
.await
|
||||
.expect("settlement should run")
|
||||
.expect("usage should exist");
|
||||
|
||||
assert_eq!(settlement.billing_status, "settled");
|
||||
assert_eq!(settlement.wallet_id.as_deref(), Some("wallet-quota"));
|
||||
assert_eq!(settlement.wallet_balance_before, Some(0.0));
|
||||
assert_eq!(settlement.wallet_balance_after, Some(0.0));
|
||||
|
||||
let wallet_total: f64 = sqlx::query_scalar(
|
||||
"SELECT balance + gift_balance FROM wallets WHERE id = 'wallet-quota'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("wallet should load");
|
||||
assert_eq!(wallet_total, 0.0);
|
||||
|
||||
let quota_used: f64 = sqlx::query_scalar(
|
||||
"SELECT CAST(COALESCE(SUM(amount_usd), 0) AS REAL) FROM entitlement_usage_ledgers WHERE request_id = 'request-quota-covered'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("quota ledger should load");
|
||||
assert_eq!(quota_used, 3.0);
|
||||
}
|
||||
|
||||
async fn seed_settlement_rows(pool: &sqlx::SqlitePool) {
|
||||
sqlx::query(
|
||||
r#"
|
||||
@@ -825,4 +890,62 @@ VALUES
|
||||
.await
|
||||
.expect("settlement rows should seed");
|
||||
}
|
||||
|
||||
async fn seed_quota_covered_settlement_rows(pool: &sqlx::SqlitePool) {
|
||||
sqlx::query(
|
||||
r#"
|
||||
INSERT INTO users (
|
||||
id, username, email, role, auth_source, password_hash, is_active,
|
||||
is_deleted, created_at, updated_at
|
||||
) VALUES (
|
||||
'user-quota', 'quota-user', 'quota@example.com', 'user', 'local',
|
||||
'hash', 1, 0, 1, 1
|
||||
);
|
||||
|
||||
INSERT INTO wallets (
|
||||
id, user_id, balance, gift_balance, limit_mode, created_at, updated_at
|
||||
) VALUES (
|
||||
'wallet-quota', 'user-quota', 0.0, 0.0, 'finite', 1, 1
|
||||
);
|
||||
|
||||
INSERT INTO "usage" (
|
||||
request_id, user_id, api_key_id, status, billing_status,
|
||||
total_cost_usd, actual_total_cost_usd
|
||||
) VALUES (
|
||||
'request-quota-covered', 'user-quota', 'key-quota', 'completed',
|
||||
'pending', 3.0, 2.0
|
||||
);
|
||||
|
||||
INSERT INTO billing_plans (
|
||||
id, title, price_amount, price_currency, duration_unit,
|
||||
duration_value, entitlements_json, created_at, updated_at
|
||||
) VALUES (
|
||||
'plan-quota', 'Quota Plan', 0.0, 'USD', 'month', 1,
|
||||
'[{"type":"daily_quota","daily_quota_usd":10.0,"reset_timezone":"Asia/Shanghai","allow_wallet_overage":false}]',
|
||||
1, 1
|
||||
);
|
||||
|
||||
INSERT INTO payment_orders (
|
||||
id, order_no, wallet_id, user_id, amount_usd, refunded_amount_usd,
|
||||
refundable_amount_usd, payment_method, gateway_response, status, created_at
|
||||
) VALUES (
|
||||
'order-quota', 'order-quota', 'wallet-quota', 'user-quota', 0.0, 0.0,
|
||||
0.0, 'admin_manual', '{}', 'credited', 1
|
||||
);
|
||||
|
||||
INSERT INTO user_plan_entitlements (
|
||||
id, user_id, plan_id, payment_order_id, status, starts_at, expires_at,
|
||||
entitlements_snapshot, created_at, updated_at
|
||||
) VALUES (
|
||||
'entitlement-quota', 'user-quota', 'plan-quota', 'order-quota',
|
||||
'active', 1, 9999999999,
|
||||
'[{"type":"daily_quota","daily_quota_usd":10.0,"reset_timezone":"Asia/Shanghai","allow_wallet_overage":false}]',
|
||||
1, 1
|
||||
);
|
||||
"#,
|
||||
)
|
||||
.execute(pool)
|
||||
.await
|
||||
.expect("quota settlement rows should seed");
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1348,8 +1348,41 @@ const REBUILD_PROVIDER_API_KEY_CODEX_WINDOW_USAGE_STATS_SQL: &str =
|
||||
|
||||
const LIST_USAGE_AUDITS_PREFIX: &str = include_str!("queries/list_usage_audits_prefix.sql");
|
||||
const USAGE_RESERVED_PROVIDER_LABELS_FILTER_SQL: &str = " AND BTRIM(COALESCE(\"usage\".provider_name, '')) <> '' AND lower(BTRIM(COALESCE(\"usage\".provider_name, ''))) NOT IN ('unknown', 'unknow', 'pending')";
|
||||
const USAGE_PROVIDER_IDENTITY_FILTER_SQL: &str = " AND BTRIM(COALESCE(\"usage\".provider_id, '')) <> '' AND lower(BTRIM(COALESCE(\"usage\".provider_id, ''))) NOT IN ('unknown', 'unknow', 'pending')";
|
||||
const USAGE_RAW_PROVIDER_GROUP_KEY_SQL: &str = r#"CASE
|
||||
WHEN BTRIM(COALESCE("usage".provider_id, '')) = ''
|
||||
OR lower(BTRIM(COALESCE("usage".provider_id, ''))) IN ('unknown', 'unknow', 'pending')
|
||||
THEN BTRIM("usage".provider_name)
|
||||
ELSE BTRIM("usage".provider_id)
|
||||
END"#;
|
||||
const USAGE_RAW_PROVIDER_DISPLAY_NAME_SQL: &str = r#"CASE
|
||||
WHEN BTRIM(COALESCE("usage".provider_name, '')) = ''
|
||||
OR lower(BTRIM(COALESCE("usage".provider_name, ''))) IN ('unknown', 'unknow', 'pending')
|
||||
THEN NULL
|
||||
ELSE BTRIM("usage".provider_name)
|
||||
END"#;
|
||||
const USAGE_PROVIDER_IDENTITY_JOIN_SQL: &str = r#" LEFT JOIN providers AS provider_by_id
|
||||
ON BTRIM(COALESCE("usage".provider_id, '')) <> ''
|
||||
AND lower(BTRIM(COALESCE("usage".provider_id, ''))) NOT IN ('unknown', 'unknow', 'pending')
|
||||
AND provider_by_id.id = BTRIM("usage".provider_id)"#;
|
||||
const USAGE_RESOLVED_PROVIDER_GROUP_KEY_SQL: &str = r#"COALESCE(
|
||||
provider_by_id.id,
|
||||
BTRIM("usage".provider_id)
|
||||
)"#;
|
||||
const USAGE_RESOLVED_PROVIDER_DISPLAY_NAME_SQL: &str = r#"COALESCE(
|
||||
provider_by_id.name,
|
||||
CASE
|
||||
WHEN BTRIM(COALESCE("usage".provider_name, '')) = ''
|
||||
OR lower(BTRIM(COALESCE("usage".provider_name, ''))) IN ('unknown', 'unknow', 'pending')
|
||||
THEN NULL
|
||||
ELSE BTRIM("usage".provider_name)
|
||||
END
|
||||
)"#;
|
||||
|
||||
struct UsageAuditAggregationSqlFragments {
|
||||
provider_identity_join: &'static str,
|
||||
provider_group_key_expr: &'static str,
|
||||
provider_display_name_expr: &'static str,
|
||||
filtered_extra_where: &'static str,
|
||||
group_key_expr: &'static str,
|
||||
display_name_expr: &'static str,
|
||||
@@ -1365,6 +1398,9 @@ fn usage_audit_aggregation_sql_fragments(
|
||||
) -> UsageAuditAggregationSqlFragments {
|
||||
match group_by {
|
||||
UsageAuditAggregationGroupBy::Model => UsageAuditAggregationSqlFragments {
|
||||
provider_identity_join: "",
|
||||
provider_group_key_expr: USAGE_RAW_PROVIDER_GROUP_KEY_SQL,
|
||||
provider_display_name_expr: USAGE_RAW_PROVIDER_DISPLAY_NAME_SQL,
|
||||
filtered_extra_where: "",
|
||||
group_key_expr: "model",
|
||||
display_name_expr: "NULL::varchar",
|
||||
@@ -1375,6 +1411,9 @@ fn usage_audit_aggregation_sql_fragments(
|
||||
success_count_expr: "NULL::BIGINT",
|
||||
},
|
||||
UsageAuditAggregationGroupBy::Provider => UsageAuditAggregationSqlFragments {
|
||||
provider_identity_join: USAGE_PROVIDER_IDENTITY_JOIN_SQL,
|
||||
provider_group_key_expr: USAGE_RESOLVED_PROVIDER_GROUP_KEY_SQL,
|
||||
provider_display_name_expr: USAGE_RESOLVED_PROVIDER_DISPLAY_NAME_SQL,
|
||||
filtered_extra_where: "",
|
||||
group_key_expr: "provider_group_key",
|
||||
display_name_expr: "provider_display_name",
|
||||
@@ -1385,6 +1424,9 @@ fn usage_audit_aggregation_sql_fragments(
|
||||
success_count_expr: "COALESCE(SUM(success_flag), 0)::BIGINT",
|
||||
},
|
||||
UsageAuditAggregationGroupBy::ApiFormat => UsageAuditAggregationSqlFragments {
|
||||
provider_identity_join: "",
|
||||
provider_group_key_expr: USAGE_RAW_PROVIDER_GROUP_KEY_SQL,
|
||||
provider_display_name_expr: USAGE_RAW_PROVIDER_DISPLAY_NAME_SQL,
|
||||
filtered_extra_where: "",
|
||||
group_key_expr: "api_format_group_key",
|
||||
display_name_expr: "NULL::varchar",
|
||||
@@ -1395,6 +1437,9 @@ fn usage_audit_aggregation_sql_fragments(
|
||||
success_count_expr: "NULL::BIGINT",
|
||||
},
|
||||
UsageAuditAggregationGroupBy::User => UsageAuditAggregationSqlFragments {
|
||||
provider_identity_join: "",
|
||||
provider_group_key_expr: USAGE_RAW_PROVIDER_GROUP_KEY_SQL,
|
||||
provider_display_name_expr: USAGE_RAW_PROVIDER_DISPLAY_NAME_SQL,
|
||||
filtered_extra_where: " AND \"usage\".user_id IS NOT NULL",
|
||||
group_key_expr: "user_id",
|
||||
display_name_expr: "NULL::varchar",
|
||||
@@ -6409,6 +6454,7 @@ WHERE stats_daily_api_key.date >=
|
||||
display_name_expr,
|
||||
avg_response_time_expr,
|
||||
success_count_expr,
|
||||
join_clause,
|
||||
) = match group_by {
|
||||
UsageAuditAggregationGroupBy::Model => (
|
||||
"stats_user_daily_model",
|
||||
@@ -6416,13 +6462,15 @@ WHERE stats_daily_api_key.date >=
|
||||
"NULL::varchar",
|
||||
"NULL::DOUBLE PRECISION",
|
||||
"NULL::BIGINT",
|
||||
"",
|
||||
),
|
||||
UsageAuditAggregationGroupBy::Provider => (
|
||||
"stats_user_daily_provider",
|
||||
"provider_name",
|
||||
"provider_name",
|
||||
"MAX(provider_name)",
|
||||
"CASE WHEN COALESCE(SUM(response_time_samples), 0) > 0 THEN COALESCE(SUM(response_time_sum_ms), 0) / COALESCE(SUM(response_time_samples), 0) ELSE NULL END",
|
||||
"COALESCE(SUM(success_requests), 0)::BIGINT",
|
||||
"",
|
||||
),
|
||||
UsageAuditAggregationGroupBy::ApiFormat => (
|
||||
"stats_user_daily_api_format",
|
||||
@@ -6430,6 +6478,7 @@ WHERE stats_daily_api_key.date >=
|
||||
"NULL::varchar",
|
||||
"CASE WHEN COALESCE(SUM(response_time_samples), 0) > 0 THEN COALESCE(SUM(response_time_sum_ms), 0) / COALESCE(SUM(response_time_samples), 0) ELSE NULL END",
|
||||
"NULL::BIGINT",
|
||||
"",
|
||||
),
|
||||
UsageAuditAggregationGroupBy::User => {
|
||||
return Ok(Vec::new());
|
||||
@@ -6464,6 +6513,7 @@ SELECT
|
||||
{avg_response_time_expr} AS avg_response_time_ms,
|
||||
{success_count_expr} AS success_count
|
||||
FROM {table_name}
|
||||
{join_clause}
|
||||
WHERE date >= $1
|
||||
AND date < $2
|
||||
{provider_extra_where}
|
||||
@@ -6475,6 +6525,7 @@ ORDER BY request_count DESC, group_key ASC
|
||||
avg_response_time_expr = avg_response_time_expr,
|
||||
success_count_expr = success_count_expr,
|
||||
table_name = table_name,
|
||||
join_clause = join_clause,
|
||||
provider_extra_where = provider_extra_where,
|
||||
);
|
||||
|
||||
@@ -6495,9 +6546,9 @@ ORDER BY request_count DESC, group_key ASC
|
||||
) -> Result<Vec<StoredUsageAuditAggregation>, DataLayerError> {
|
||||
let fragments = usage_audit_aggregation_sql_fragments(query.group_by);
|
||||
let provider_extra_where =
|
||||
if matches!(query.group_by, UsageAuditAggregationGroupBy::Provider)
|
||||
|| query.exclude_reserved_provider_labels
|
||||
{
|
||||
if matches!(query.group_by, UsageAuditAggregationGroupBy::Provider) {
|
||||
USAGE_PROVIDER_IDENTITY_FILTER_SQL
|
||||
} else if query.exclude_reserved_provider_labels {
|
||||
USAGE_RESERVED_PROVIDER_LABELS_FILTER_SQL
|
||||
} else {
|
||||
""
|
||||
@@ -6508,18 +6559,8 @@ WITH filtered_usage AS (
|
||||
SELECT
|
||||
"usage".model AS model,
|
||||
"usage".user_id AS user_id,
|
||||
CASE
|
||||
WHEN BTRIM(COALESCE("usage".provider_id, '')) = ''
|
||||
OR lower(BTRIM(COALESCE("usage".provider_id, ''))) IN ('unknown', 'unknow', 'pending')
|
||||
THEN BTRIM("usage".provider_name)
|
||||
ELSE BTRIM("usage".provider_id)
|
||||
END AS provider_group_key,
|
||||
CASE
|
||||
WHEN BTRIM(COALESCE("usage".provider_name, '')) = ''
|
||||
OR lower(BTRIM(COALESCE("usage".provider_name, ''))) IN ('unknown', 'unknow', 'pending')
|
||||
THEN NULL
|
||||
ELSE BTRIM("usage".provider_name)
|
||||
END AS provider_display_name,
|
||||
{provider_group_key_expr} AS provider_group_key,
|
||||
{provider_display_name_expr} AS provider_display_name,
|
||||
COALESCE("usage".api_format, 'unknown') AS api_format_group_key,
|
||||
GREATEST(COALESCE("usage".input_tokens, 0), 0) AS input_tokens,
|
||||
GREATEST(COALESCE("usage".output_tokens, 0), 0) AS output_tokens,
|
||||
@@ -6550,6 +6591,7 @@ WITH filtered_usage AS (
|
||||
ELSE 0
|
||||
END AS success_flag
|
||||
FROM usage_billing_facts AS "usage"
|
||||
{provider_identity_join}
|
||||
WHERE "usage".created_at >= TO_TIMESTAMP($1::double precision)
|
||||
AND "usage".created_at < TO_TIMESTAMP($2::double precision)
|
||||
AND "usage".status NOT IN ('pending', 'streaming')
|
||||
@@ -6649,6 +6691,9 @@ LIMIT $3
|
||||
"#,
|
||||
provider_extra_where = provider_extra_where,
|
||||
filtered_extra_where = fragments.filtered_extra_where,
|
||||
provider_identity_join = fragments.provider_identity_join,
|
||||
provider_group_key_expr = fragments.provider_group_key_expr,
|
||||
provider_display_name_expr = fragments.provider_display_name_expr,
|
||||
group_key_expr = fragments.group_key_expr,
|
||||
display_name_expr = fragments.display_name_expr,
|
||||
secondary_name_expr = fragments.secondary_name_expr,
|
||||
@@ -6683,6 +6728,9 @@ LIMIT $3
|
||||
if matches!(query.group_by, UsageAuditAggregationGroupBy::User) {
|
||||
return self.aggregate_usage_audits_raw(query).await;
|
||||
}
|
||||
if matches!(query.group_by, UsageAuditAggregationGroupBy::Provider) {
|
||||
return self.aggregate_usage_audits_raw(query).await;
|
||||
}
|
||||
if query.exclude_reserved_provider_labels
|
||||
&& !matches!(query.group_by, UsageAuditAggregationGroupBy::Provider)
|
||||
{
|
||||
|
||||
@@ -456,7 +456,17 @@ fn usage_sql_aggregate_usage_audits_supports_daily_model_and_provider_aggregates
|
||||
fn usage_sql_provider_aggregation_excludes_unknown_provider_labels() {
|
||||
let source = include_str!("mod.rs");
|
||||
assert!(source.contains(
|
||||
r#"lower(BTRIM(COALESCE(\"usage\".provider_name, ''))) NOT IN ('unknown', 'unknow', 'pending')"#
|
||||
r#"const USAGE_PROVIDER_IDENTITY_FILTER_SQL: &str = " AND BTRIM(COALESCE(\"usage\".provider_id, '')) <> ''"#
|
||||
));
|
||||
assert!(source.contains("LEFT JOIN providers AS provider_by_id"));
|
||||
assert!(source.contains("provider_by_id.id = BTRIM(\"usage\".provider_id)"));
|
||||
assert!(
|
||||
source.contains("COALESCE(\n provider_by_id.id,\n BTRIM(\"usage\".provider_id)")
|
||||
);
|
||||
assert!(source.contains("COALESCE(\n provider_by_id.name,"));
|
||||
assert!(!source.contains("provider_by_name.name = BTRIM(\"usage\".provider_name)"));
|
||||
assert!(source.contains(
|
||||
"if matches!(query.group_by, UsageAuditAggregationGroupBy::Provider) {\n return self.aggregate_usage_audits_raw(query).await;"
|
||||
));
|
||||
assert!(source.contains("exclude_reserved_provider_labels"));
|
||||
assert!(source.contains(
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user