mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 01:10:23 +08:00
Merge remote-tracking branch 'origin/pr-482'
# Conflicts: # .github/workflows/release.yml
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,
|
||||
|
||||
@@ -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,16 +1,279 @@
|
||||
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;
|
||||
|
||||
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 {
|
||||
pool: SqlitePool,
|
||||
@@ -21,92 +284,340 @@ 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> {
|
||||
pub async fn list_providers(
|
||||
&self,
|
||||
active_only: bool,
|
||||
) -> Result<Vec<StoredProviderCatalogProvider>, 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,
|
||||
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 provider_endpoints
|
||||
WHERE api_format IS NOT NULL
|
||||
FROM providers
|
||||
WHERE (? = FALSE OR is_active = TRUE)
|
||||
ORDER BY provider_priority ASC, name ASC
|
||||
"#,
|
||||
)
|
||||
.bind(active_only)
|
||||
.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 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 => {
|
||||
"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 count_row = sqlx::query(
|
||||
r#"
|
||||
SELECT COUNT(*) AS total
|
||||
FROM provider_api_keys
|
||||
WHERE provider_id = ?
|
||||
AND (? IS NULL OR LOWER(name) LIKE ? OR LOWER(id) LIKE ?)
|
||||
AND (? IS NULL OR is_active = ?)
|
||||
"#,
|
||||
)
|
||||
.bind(&query.provider_id)
|
||||
.bind(search_pattern.as_deref())
|
||||
.bind(search_pattern.as_deref())
|
||||
.bind(search_pattern.as_deref())
|
||||
.bind(query.is_active)
|
||||
.bind(query.is_active)
|
||||
.fetch_one(&self.pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let total = count_row.try_get::<i64, _>("total").map_sql_err()?.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,
|
||||
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 = ?
|
||||
AND (? IS NULL OR LOWER(name) LIKE ? OR LOWER(id) LIKE ?)
|
||||
AND (? IS NULL OR is_active = ?)
|
||||
ORDER BY {order_by}
|
||||
LIMIT ?
|
||||
OFFSET ?
|
||||
"#,
|
||||
);
|
||||
let rows = sqlx::query(&sql)
|
||||
.bind(&query.provider_id)
|
||||
.bind(search_pattern.as_deref())
|
||||
.bind(search_pattern.as_deref())
|
||||
.bind(search_pattern.as_deref())
|
||||
.bind(query.is_active)
|
||||
.bind(query.is_active)
|
||||
.bind(limit)
|
||||
.bind(offset)
|
||||
.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,
|
||||
@@ -870,81 +1381,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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1063,6 +1556,21 @@ 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(prefix);
|
||||
let mut separated = builder.separated(", ");
|
||||
for id in ids {
|
||||
separated.push_bind(id);
|
||||
}
|
||||
separated.push_unseparated(")");
|
||||
builder.push(suffix);
|
||||
builder
|
||||
}
|
||||
|
||||
fn current_unix_secs() -> u64 {
|
||||
chrono::Utc::now().timestamp().max(0) as u64
|
||||
}
|
||||
@@ -1353,6 +1861,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() {
|
||||
@@ -1549,8 +2065,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;
|
||||
|
||||
|
||||
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