refactor(data): limit query abstraction to postgres and sqlite

This commit is contained in:
Kayphoon
2026-05-17 15:40:44 +08:00
parent 77640d51a6
commit f29cca72ba
13 changed files with 498 additions and 563 deletions

View File

@@ -4,14 +4,12 @@ use sqlx::{Database, Encode, QueryBuilder, Type};
pub enum SqlDialect { pub enum SqlDialect {
Postgres, Postgres,
Sqlite, Sqlite,
Mysql,
} }
impl SqlDialect { impl SqlDialect {
pub fn quote_ident(self, ident: &str) -> String { pub fn quote_ident(self, ident: &str) -> String {
let quote = match self { let quote = match self {
Self::Postgres | Self::Sqlite => '"', Self::Postgres | Self::Sqlite => '"',
Self::Mysql => '`',
}; };
let escaped = ident.replace(quote, &format!("{quote}{quote}")); let escaped = ident.replace(quote, &format!("{quote}{quote}"));
format!("{quote}{escaped}{quote}") format!("{quote}{escaped}{quote}")
@@ -31,7 +29,6 @@ pub struct DialectSql<'a> {
common: Option<&'a str>, common: Option<&'a str>,
postgres: Option<&'a str>, postgres: Option<&'a str>,
sqlite: Option<&'a str>, sqlite: Option<&'a str>,
mysql: Option<&'a str>,
} }
impl<'a> DialectSql<'a> { impl<'a> DialectSql<'a> {
@@ -40,16 +37,14 @@ impl<'a> DialectSql<'a> {
common: Some(sql), common: Some(sql),
postgres: None, postgres: None,
sqlite: None, sqlite: None,
mysql: None,
} }
} }
pub const fn dialect(postgres: &'a str, sqlite: &'a str, mysql: &'a str) -> Self { pub const fn dialect(postgres: &'a str, sqlite: &'a str) -> Self {
Self { Self {
common: None, common: None,
postgres: Some(postgres), postgres: Some(postgres),
sqlite: Some(sqlite), sqlite: Some(sqlite),
mysql: Some(mysql),
} }
} }
@@ -63,16 +58,10 @@ impl<'a> DialectSql<'a> {
self self
} }
pub fn with_mysql(mut self, sql: &'a str) -> Self {
self.mysql = Some(sql);
self
}
pub fn sql(self, dialect: SqlDialect) -> &'a str { pub fn sql(self, dialect: SqlDialect) -> &'a str {
match dialect { match dialect {
SqlDialect::Postgres => self.postgres.or(self.common), SqlDialect::Postgres => self.postgres.or(self.common),
SqlDialect::Sqlite => self.sqlite.or(self.common), SqlDialect::Sqlite => self.sqlite.or(self.common),
SqlDialect::Mysql => self.mysql.or(self.common),
} }
.expect("dialect SQL expression is missing for selected dialect") .expect("dialect SQL expression is missing for selected dialect")
} }
@@ -470,7 +459,7 @@ fn push_ci_contains_predicate<'args, DB>(
.push(" ILIKE ") .push(" ILIKE ")
.push_bind(format!("%{trimmed}%")); .push_bind(format!("%{trimmed}%"));
} }
SqlDialect::Sqlite | SqlDialect::Mysql => { SqlDialect::Sqlite => {
builder builder
.push("LOWER(") .push("LOWER(")
.push(column_sql) .push(column_sql)
@@ -522,13 +511,12 @@ pub fn push_order_by<DB>(
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use sqlx::{Execute, MySql, Postgres, QueryBuilder, Sqlite}; use sqlx::{Execute, Postgres, QueryBuilder, Sqlite};
#[test] #[test]
fn quotes_identifiers_by_dialect() { fn quotes_identifiers_by_dialect() {
assert_eq!(SqlDialect::Postgres.quote_ident("trigger"), "\"trigger\""); assert_eq!(SqlDialect::Postgres.quote_ident("trigger"), "\"trigger\"");
assert_eq!(SqlDialect::Sqlite.quote_ident("trigger"), "\"trigger\""); assert_eq!(SqlDialect::Sqlite.quote_ident("trigger"), "\"trigger\"");
assert_eq!(SqlDialect::Mysql.quote_ident("trigger"), "`trigger`");
assert_eq!( assert_eq!(
SqlDialect::Postgres.quote_path(&["usage", "id"]), SqlDialect::Postgres.quote_path(&["usage", "id"]),
"\"usage\".\"id\"" "\"usage\".\"id\""
@@ -571,7 +559,7 @@ mod tests {
} }
#[test] #[test]
fn ci_contains_uses_lower_like_for_sqlite_and_mysql() { fn ci_contains_uses_lower_like_for_sqlite() {
let mut sqlite_builder = QueryBuilder::<Sqlite>::new("SELECT * FROM items"); let mut sqlite_builder = QueryBuilder::<Sqlite>::new("SELECT * FROM items");
let mut sqlite_where = WhereClause::new(); let mut sqlite_where = WhereClause::new();
push_ci_contains( push_ci_contains(
@@ -585,20 +573,6 @@ mod tests {
.build() .build()
.sql() .sql()
.contains(" WHERE LOWER(task_key) LIKE ?")); .contains(" WHERE LOWER(task_key) LIKE ?"));
let mut mysql_builder = QueryBuilder::<MySql>::new("SELECT * FROM items");
let mut mysql_where = WhereClause::new();
push_ci_contains(
&mut mysql_builder,
&mut mysql_where,
SqlDialect::Mysql,
"task_key",
" Fetch ",
);
assert!(mysql_builder
.build()
.sql()
.contains(" WHERE LOWER(task_key) LIKE ?"));
} }
#[test] #[test]
@@ -652,7 +626,6 @@ mod tests {
SelectColumn::expr(DialectSql::dialect( SelectColumn::expr(DialectSql::dialect(
"CAST(monthly_quota_usd AS DOUBLE PRECISION)", "CAST(monthly_quota_usd AS DOUBLE PRECISION)",
"CAST(monthly_quota_usd AS REAL)", "CAST(monthly_quota_usd AS REAL)",
"monthly_quota_usd",
)) ))
.alias("monthly_quota_usd"), .alias("monthly_quota_usd"),
]); ]);
@@ -662,8 +635,8 @@ mod tests {
"SELECT id AS \"provider_id\", CAST(monthly_quota_usd AS DOUBLE PRECISION) AS \"monthly_quota_usd\" FROM providers" "SELECT id AS \"provider_id\", CAST(monthly_quota_usd AS DOUBLE PRECISION) AS \"monthly_quota_usd\" FROM providers"
); );
assert_eq!( assert_eq!(
query.render(SqlDialect::Mysql), query.render(SqlDialect::Sqlite),
"SELECT id AS `provider_id`, monthly_quota_usd AS `monthly_quota_usd` FROM providers" "SELECT id AS \"provider_id\", CAST(monthly_quota_usd AS REAL) AS \"monthly_quota_usd\" FROM providers"
); );
} }

View File

@@ -1,5 +1,5 @@
use async_trait::async_trait; use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use sqlx::{mysql::MySqlRow, Row};
use super::types::{ use super::types::{
AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository, AnnouncementListQuery, AnnouncementReadRepository, AnnouncementWriteRepository,
@@ -8,7 +8,6 @@ use super::types::{
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::{push_eq, push_limit, push_limit_offset, WhereClause};
const ANNOUNCEMENT_SELECT: &str = r#" const ANNOUNCEMENT_SELECT: &str = r#"
SELECT SELECT
@@ -45,27 +44,6 @@ impl MysqlAnnouncementRepository {
) -> Result<Option<StoredAnnouncement>, DataLayerError> { ) -> Result<Option<StoredAnnouncement>, DataLayerError> {
self.find_by_id(announcement_id).await self.find_by_id(announcement_id).await
} }
fn apply_active_filter(
builder: &mut QueryBuilder<'_, MySql>,
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] #[async_trait]
@@ -74,17 +52,8 @@ impl AnnouncementReadRepository for MysqlAnnouncementRepository {
&self, &self,
announcement_id: &str, announcement_id: &str,
) -> Result<Option<StoredAnnouncement>, DataLayerError> { ) -> Result<Option<StoredAnnouncement>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(ANNOUNCEMENT_SELECT); let row = sqlx::query(&format!("{ANNOUNCEMENT_SELECT} WHERE a.id = ? LIMIT 1"))
let mut where_clause = WhereClause::new(); .bind(announcement_id)
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) .fetch_optional(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
@@ -96,38 +65,49 @@ impl AnnouncementReadRepository for MysqlAnnouncementRepository {
query: &AnnouncementListQuery, query: &AnnouncementListQuery,
) -> Result<StoredAnnouncementPage, DataLayerError> { ) -> Result<StoredAnnouncementPage, DataLayerError> {
let now_unix_secs = query.now_unix_secs.unwrap_or_else(current_unix_secs); let now_unix_secs = query.now_unix_secs.unwrap_or_else(current_unix_secs);
let mut count_builder = let total_row = sqlx::query(
QueryBuilder::<MySql>::new("SELECT COUNT(a.id) AS total FROM announcements a"); r#"
let mut count_where = WhereClause::new(); SELECT COUNT(a.id) AS total
Self::apply_active_filter( FROM announcements a
&mut count_builder, WHERE (
&mut count_where, NOT ? OR (
query.active_only, a.is_active = 1
now_unix_secs, AND (a.start_time IS NULL OR a.start_time <= ?)
)?; AND (a.end_time IS NULL OR a.end_time >= ?)
let total = count_builder )
.build_query_scalar::<i64>() )
.fetch_one(&self.pool) "#,
.await )
.map_sql_err()? .bind(query.active_only)
.max(0) as u64; .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 list_builder = QueryBuilder::<MySql>::new(ANNOUNCEMENT_SELECT); let rows = sqlx::query(&format!(
let mut list_where = WhereClause::new(); r#"
Self::apply_active_filter( {ANNOUNCEMENT_SELECT}
&mut list_builder, WHERE (
&mut list_where, NOT ? OR (
query.active_only, a.is_active = 1
now_unix_secs, AND (a.start_time IS NULL OR a.start_time <= ?)
)?; AND (a.end_time IS NULL OR a.end_time >= ?)
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); ORDER BY a.is_pinned DESC, a.priority DESC, a.created_at DESC, a.id ASC
let rows = list_builder LIMIT ? OFFSET ?
.build() "#
.fetch_all(&self.pool) ))
.await .bind(query.active_only)
.map_sql_err()?; .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 items = rows let items = rows
.iter() .iter()
.map(map_announcement_row) .map(map_announcement_row)
@@ -141,22 +121,28 @@ impl AnnouncementReadRepository for MysqlAnnouncementRepository {
user_id: &str, user_id: &str,
now_unix_secs: u64, now_unix_secs: u64,
) -> Result<u64, DataLayerError> { ) -> Result<u64, DataLayerError> {
let mut builder = let row = sqlx::query(
QueryBuilder::<MySql>::new("SELECT COUNT(a.id) AS total FROM announcements a"); r#"
let mut where_clause = WhereClause::new(); SELECT COUNT(a.id) AS total
Self::apply_active_filter(&mut builder, &mut where_clause, true, now_unix_secs)?; FROM announcements a
where_clause.push_next(&mut builder); WHERE a.is_active = 1
builder AND (a.start_time IS NULL OR a.start_time <= ?)
.push("NOT EXISTS (SELECT 1 FROM announcement_reads r WHERE r.user_id = ") AND (a.end_time IS NULL OR a.end_time >= ?)
.push_bind(user_id.to_string()) AND NOT EXISTS (
.push(" AND r.announcement_id = a.id)"); SELECT 1
let total = builder FROM announcement_reads r
.build_query_scalar::<i64>() WHERE r.user_id = ?
.fetch_one(&self.pool) AND r.announcement_id = a.id
.await )
.map_sql_err()? "#,
.max(0) as u64; )
Ok(total) .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)
} }
} }

View File

@@ -1,5 +1,5 @@
use async_trait::async_trait; use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use sqlx::{mysql::MySqlRow, Row};
use super::types::{ use super::types::{
AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig, AuthModuleReadRepository, AuthModuleWriteRepository, StoredLdapModuleConfig,
@@ -8,9 +8,8 @@ use super::types::{
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::{push_eq, push_limit, WhereClause};
const OAUTH_PROVIDER_COLUMNS: &str = r#" const LIST_ENABLED_OAUTH_PROVIDERS_SQL: &str = r#"
SELECT SELECT
provider_type, provider_type,
display_name, display_name,
@@ -18,9 +17,11 @@ SELECT
client_secret_encrypted, client_secret_encrypted,
redirect_uri redirect_uri
FROM oauth_providers FROM oauth_providers
WHERE is_enabled = 1
ORDER BY provider_type ASC
"#; "#;
const LDAP_CONFIG_COLUMNS: &str = r#" const GET_LDAP_CONFIG_SQL: &str = r#"
SELECT SELECT
server_url, server_url,
bind_dn, bind_dn,
@@ -35,6 +36,8 @@ SELECT
use_starttls, use_starttls,
connect_timeout connect_timeout
FROM ldap_configs FROM ldap_configs
ORDER BY id ASC
LIMIT 1
"#; "#;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
@@ -59,37 +62,24 @@ impl MysqlAuthModuleRepository {
} }
} }
async fn list_enabled_oauth_providers(
pool: &MysqlPool,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::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: &MysqlPool,
) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::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] #[async_trait]
impl AuthModuleReadRepository for MysqlAuthModuleReadRepository { impl AuthModuleReadRepository for MysqlAuthModuleReadRepository {
async fn list_enabled_oauth_providers( async fn list_enabled_oauth_providers(
&self, &self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> { ) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
list_enabled_oauth_providers(&self.pool).await let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_oauth_row).collect()
} }
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> { async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
get_ldap_config(&self.pool).await let row = sqlx::query(GET_LDAP_CONFIG_SQL)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_ldap_row).transpose()
} }
} }
@@ -98,11 +88,19 @@ impl AuthModuleReadRepository for MysqlAuthModuleRepository {
async fn list_enabled_oauth_providers( async fn list_enabled_oauth_providers(
&self, &self,
) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> { ) -> Result<Vec<StoredOAuthProviderModuleConfig>, DataLayerError> {
list_enabled_oauth_providers(&self.pool).await let rows = sqlx::query(LIST_ENABLED_OAUTH_PROVIDERS_SQL)
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_oauth_row).collect()
} }
async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> { async fn get_ldap_config(&self) -> Result<Option<StoredLdapModuleConfig>, DataLayerError> {
get_ldap_config(&self.pool).await let row = sqlx::query(GET_LDAP_CONFIG_SQL)
.fetch_optional(&self.pool)
.await
.map_sql_err()?;
row.as_ref().map(map_ldap_row).transpose()
} }
} }

View File

@@ -10,9 +10,6 @@ use super::{
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::{
push_ci_contains, push_eq, push_limit, push_limit_offset, SqlDialect, WhereClause,
};
const RUN_COLUMNS: &str = r#" const RUN_COLUMNS: &str = r#"
SELECT SELECT
@@ -60,29 +57,44 @@ impl MysqlBackgroundTaskRepository {
} }
fn apply_run_filter(builder: &mut QueryBuilder<'_, MySql>, query: &BackgroundTaskListQuery) { fn apply_run_filter(builder: &mut QueryBuilder<'_, MySql>, query: &BackgroundTaskListQuery) {
let mut where_clause = WhereClause::new(); let mut has_where = false;
if let Some(kind) = query.kind { if let Some(kind) = query.kind {
push_eq(builder, &mut where_clause, "kind", kind.as_database()); if !has_where {
builder.push(" WHERE ");
has_where = true;
} else {
builder.push(" AND ");
}
builder.push("kind = ").push_bind(kind.as_database());
} }
if let Some(status) = query.status { if let Some(status) = query.status {
push_eq(builder, &mut where_clause, "status", status.as_database()); if !has_where {
builder.push(" WHERE ");
has_where = true;
} else {
builder.push(" AND ");
}
builder.push("status = ").push_bind(status.as_database());
} }
if let Some(trigger) = query.trigger.as_deref() { if let Some(trigger) = query.trigger.as_deref() {
push_eq( if !has_where {
builder, builder.push(" WHERE ");
&mut where_clause, has_where = true;
&SqlDialect::Mysql.quote_ident("trigger"), } else {
trigger.to_string(), builder.push(" AND ");
); }
builder.push("`trigger` = ").push_bind(trigger.to_string());
} }
if let Some(task_key_substring) = query.task_key_substring.as_deref() { if let Some(task_key_substring) = query.task_key_substring.as_deref() {
push_ci_contains( if !has_where {
builder, builder.push(" WHERE ");
&mut where_clause, } else {
SqlDialect::Mysql, builder.push(" AND ");
"task_key", }
task_key_substring, builder.push("LOWER(task_key) LIKE ").push_bind(format!(
); "%{}%",
task_key_substring.trim().to_ascii_lowercase()
));
} }
} }
} }
@@ -93,12 +105,8 @@ impl BackgroundTaskReadRepository for MysqlBackgroundTaskRepository {
&self, &self,
run_id: &str, run_id: &str,
) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> { ) -> Result<Option<StoredBackgroundTaskRun>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(RUN_COLUMNS); let row = sqlx::query(&format!("{RUN_COLUMNS} WHERE id = ? LIMIT 1"))
let mut where_clause = WhereClause::new(); .bind(run_id)
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) .fetch_optional(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
@@ -121,12 +129,12 @@ impl BackgroundTaskReadRepository for MysqlBackgroundTaskRepository {
let mut builder = QueryBuilder::<MySql>::new(RUN_COLUMNS); let mut builder = QueryBuilder::<MySql>::new(RUN_COLUMNS);
Self::apply_run_filter(&mut builder, query); Self::apply_run_filter(&mut builder, query);
builder.push(" ORDER BY created_at_unix_secs DESC, updated_at_unix_secs DESC"); builder
push_limit_offset( .push(" ORDER BY created_at_unix_secs DESC, updated_at_unix_secs DESC")
&mut builder, .push(" LIMIT ")
i64_from_usize(limit, "run limit")?, .push_bind(i64_from_usize(limit, "run limit")?)
i64_from_usize(query.offset, "run offset")?, .push(" OFFSET ")
); .push_bind(i64_from_usize(query.offset, "run offset")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
let items = rows let items = rows
.iter() .iter()
@@ -145,21 +153,15 @@ impl BackgroundTaskReadRepository for MysqlBackgroundTaskRepository {
limit: usize, limit: usize,
) -> Result<Vec<StoredBackgroundTaskEvent>, DataLayerError> { ) -> Result<Vec<StoredBackgroundTaskEvent>, DataLayerError> {
let limit = limit.max(1); let limit = limit.max(1);
let mut builder = QueryBuilder::<MySql>::new(EVENT_COLUMNS); let rows = sqlx::query(&format!(
let mut where_clause = WhereClause::new(); "{EVENT_COLUMNS} WHERE run_id = ? ORDER BY created_at_unix_secs ASC, id ASC LIMIT ? OFFSET ?"
push_eq( ))
&mut builder, .bind(run_id)
&mut where_clause, .bind(i64_from_usize(limit, "event limit")?)
"run_id", .bind(i64_from_usize(offset, "event offset")?)
run_id.to_string(), .fetch_all(&self.pool)
); .await
builder.push(" ORDER BY created_at_unix_secs ASC, id ASC"); .map_sql_err()?;
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() rows.iter().map(map_event_row).collect()
} }

View File

@@ -11,7 +11,6 @@ use super::{
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::{push_in, WhereClause};
const CANDIDATE_COLUMNS: &str = r#" const CANDIDATE_COLUMNS: &str = r#"
SELECT SELECT
@@ -133,8 +132,7 @@ impl RequestCandidateReadRepository for MysqlRequestCandidateRepository {
return Ok(Vec::new()); return Ok(Vec::new());
} }
let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_COLUMNS); let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_COLUMNS);
let mut where_clause = WhereClause::new(); push_endpoint_in_clause(&mut builder, endpoint_ids);
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
builder builder
.push(" AND created_at >= ") .push(" AND created_at >= ")
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?) .push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
@@ -156,8 +154,7 @@ impl RequestCandidateReadRepository for MysqlRequestCandidateRepository {
let mut builder = QueryBuilder::<MySql>::new( let mut builder = QueryBuilder::<MySql>::new(
"SELECT endpoint_id, status, COUNT(id) AS count FROM request_candidates", "SELECT endpoint_id, status, COUNT(id) AS count FROM request_candidates",
); );
let mut where_clause = WhereClause::new(); push_endpoint_in_clause(&mut builder, endpoint_ids);
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
builder builder
.push(" AND created_at >= ") .push(" AND created_at >= ")
.push_bind(unix_secs_to_ms_i64(since_unix_secs)?) .push_bind(unix_secs_to_ms_i64(since_unix_secs)?)
@@ -196,8 +193,7 @@ impl RequestCandidateReadRepository for MysqlRequestCandidateRepository {
let since_ms = unix_secs_to_ms_i64(since_unix_secs)?; let since_ms = unix_secs_to_ms_i64(since_unix_secs)?;
let until_ms = unix_secs_to_ms_i64(until_unix_secs)?; let until_ms = unix_secs_to_ms_i64(until_unix_secs)?;
let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_COLUMNS); let mut builder = QueryBuilder::<MySql>::new(CANDIDATE_COLUMNS);
let mut where_clause = WhereClause::new(); push_endpoint_in_clause(&mut builder, endpoint_ids);
push_in(&mut builder, &mut where_clause, "endpoint_id", endpoint_ids);
builder builder
.push(" AND created_at >= ") .push(" AND created_at >= ")
.push_bind(since_ms) .push_bind(since_ms)
@@ -343,6 +339,20 @@ ON DUPLICATE KEY UPDATE
Ok(()) Ok(())
} }
fn push_endpoint_in_clause<'args>(
builder: &mut QueryBuilder<'args, MySql>,
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( fn merge_candidate(
candidate: UpsertRequestCandidateRecord, candidate: UpsertRequestCandidateRecord,
existing: Option<StoredRequestCandidate>, existing: Option<StoredRequestCandidate>,

View File

@@ -9,7 +9,6 @@ use super::types::{
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::{push_ci_contains_any, push_limit_offset, SqlDialect, WhereClause};
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct MysqlGeminiFileMappingRepository { pub struct MysqlGeminiFileMappingRepository {
@@ -243,9 +242,8 @@ LIMIT 1
fn build_list_count_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, MySql> { fn build_list_count_query(query: &GeminiFileMappingListQuery) -> QueryBuilder<'_, MySql> {
let mut builder = let mut builder =
QueryBuilder::<MySql>::new("SELECT COUNT(*) AS total FROM gemini_file_mappings"); QueryBuilder::<MySql>::new("SELECT COUNT(*) AS total FROM gemini_file_mappings WHERE 1=1");
let mut where_clause = WhereClause::new(); apply_list_filters(&mut builder, query);
apply_list_filters(&mut builder, &mut where_clause, query);
builder builder
} }
@@ -263,27 +261,20 @@ SELECT
created_at AS created_at_unix_ms, created_at AS created_at_unix_ms,
expires_at AS expires_at_unix_secs expires_at AS expires_at_unix_secs
FROM gemini_file_mappings FROM gemini_file_mappings
WHERE 1=1
"#, "#,
); );
let mut where_clause = WhereClause::new(); apply_list_filters(&mut builder, query);
apply_list_filters(&mut builder, &mut where_clause, query); builder.push(" ORDER BY created_at DESC, file_name ASC LIMIT ");
builder.push(" ORDER BY created_at DESC, file_name ASC"); builder.push_bind(i64::try_from(query.limit).unwrap_or(i64::MAX));
push_limit_offset( builder.push(" OFFSET ");
&mut builder, builder.push_bind(i64::try_from(query.offset).unwrap_or(i64::MAX));
i64::try_from(query.limit).unwrap_or(i64::MAX),
i64::try_from(query.offset).unwrap_or(i64::MAX),
);
builder builder
} }
fn apply_list_filters( fn apply_list_filters(builder: &mut QueryBuilder<'_, MySql>, query: &GeminiFileMappingListQuery) {
builder: &mut QueryBuilder<'_, MySql>,
where_clause: &mut WhereClause,
query: &GeminiFileMappingListQuery,
) {
if !query.include_expired { if !query.include_expired {
where_clause.push_next(builder); builder.push(" AND expires_at > ");
builder.push("expires_at > ");
builder.push_bind(query.now_unix_secs as i64); builder.push_bind(query.now_unix_secs as i64);
} }
if let Some(search) = query if let Some(search) = query
@@ -292,13 +283,12 @@ fn apply_list_filters(
.map(str::trim) .map(str::trim)
.filter(|value| !value.is_empty()) .filter(|value| !value.is_empty())
{ {
push_ci_contains_any( let pattern = format!("%{}%", search.to_ascii_lowercase());
builder, builder.push(" AND (LOWER(file_name) LIKE ");
where_clause, builder.push_bind(pattern.clone());
SqlDialect::Mysql, builder.push(" OR LOWER(COALESCE(display_name, '')) LIKE ");
&["file_name", "COALESCE(display_name, '')"], builder.push_bind(pattern);
search, builder.push(")");
);
} }
} }

View File

@@ -1,5 +1,5 @@
use async_trait::async_trait; use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use sqlx::{mysql::MySqlRow, Row};
use super::types::{ use super::types::{
CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository, CreateManagementTokenRecord, ManagementTokenListQuery, ManagementTokenReadRepository,
@@ -10,7 +10,6 @@ use super::types::{
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::{push_eq, push_limit, push_limit_offset, push_optional_eq, WhereClause};
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct MysqlManagementTokenRepository { pub struct MysqlManagementTokenRepository {
@@ -26,12 +25,8 @@ impl MysqlManagementTokenRepository {
&self, &self,
token_id: &str, token_id: &str,
) -> Result<Option<StoredManagementToken>, DataLayerError> { ) -> Result<Option<StoredManagementToken>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(TOKEN_COLUMNS); let row = sqlx::query(TOKEN_BY_ID_SQL)
let mut where_clause = WhereClause::new(); .bind(token_id)
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) .fetch_optional(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
@@ -39,7 +34,7 @@ impl MysqlManagementTokenRepository {
} }
} }
const TOKEN_COLUMNS: &str = r#" const TOKEN_BY_ID_SQL: &str = r#"
SELECT SELECT
id, id,
user_id, user_id,
@@ -56,9 +51,11 @@ SELECT
created_at AS created_at_unix_ms, created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs updated_at AS updated_at_unix_secs
FROM management_tokens FROM management_tokens
WHERE id = ?
LIMIT 1
"#; "#;
const TOKEN_WITH_USER_COLUMNS: &str = r#" const LIST_MANAGEMENT_TOKENS_SQL: &str = r#"
SELECT SELECT
mt.id, mt.id,
mt.user_id, mt.user_id,
@@ -80,6 +77,69 @@ SELECT
u.role AS user_role u.role AS user_role
FROM management_tokens mt FROM management_tokens mt
JOIN users u ON u.id = mt.user_id 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] #[async_trait]
@@ -88,27 +148,23 @@ impl ManagementTokenReadRepository for MysqlManagementTokenRepository {
&self, &self,
query: &ManagementTokenListQuery, query: &ManagementTokenListQuery,
) -> Result<StoredManagementTokenListPage, DataLayerError> { ) -> Result<StoredManagementTokenListPage, DataLayerError> {
let mut count_builder = let count_row = sqlx::query(COUNT_MANAGEMENT_TOKENS_SQL)
QueryBuilder::<MySql>::new("SELECT COUNT(mt.id) AS total FROM management_tokens mt"); .bind(query.user_id.as_deref())
let mut count_where = WhereClause::new(); .bind(query.user_id.as_deref())
apply_management_token_filters(&mut count_builder, &mut count_where, query); .bind(query.is_active)
let total = count_builder .bind(query.is_active)
.build_query_scalar::<i64>()
.fetch_one(&self.pool) .fetch_one(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
let total = count_row.try_get::<i64, _>("total").map_sql_err()?;
let mut list_builder = QueryBuilder::<MySql>::new(TOKEN_WITH_USER_COLUMNS); let rows = sqlx::query(LIST_MANAGEMENT_TOKENS_SQL)
let mut list_where = WhereClause::new(); .bind(query.user_id.as_deref())
apply_management_token_filters(&mut list_builder, &mut list_where, query); .bind(query.user_id.as_deref())
list_builder.push(" ORDER BY mt.created_at DESC, mt.id DESC"); .bind(query.is_active)
push_limit_offset( .bind(query.is_active)
&mut list_builder, .bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
i64::try_from(query.limit).unwrap_or(i64::MAX), .bind(i64::try_from(query.offset).unwrap_or(i64::MAX))
i64::try_from(query.offset).unwrap_or(i64::MAX),
);
let rows = list_builder
.build()
.fetch_all(&self.pool) .fetch_all(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
@@ -126,17 +182,8 @@ impl ManagementTokenReadRepository for MysqlManagementTokenRepository {
&self, &self,
token_id: &str, token_id: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> { ) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(TOKEN_WITH_USER_COLUMNS); let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_SQL)
let mut where_clause = WhereClause::new(); .bind(token_id)
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) .fetch_optional(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
@@ -147,17 +194,8 @@ impl ManagementTokenReadRepository for MysqlManagementTokenRepository {
&self, &self,
token_hash: &str, token_hash: &str,
) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> { ) -> Result<Option<StoredManagementTokenWithUser>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(TOKEN_WITH_USER_COLUMNS); let row = sqlx::query(GET_MANAGEMENT_TOKEN_WITH_USER_BY_HASH_SQL)
let mut where_clause = WhereClause::new(); .bind(token_hash)
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) .fetch_optional(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
@@ -165,15 +203,6 @@ impl ManagementTokenReadRepository for MysqlManagementTokenRepository {
} }
} }
fn apply_management_token_filters(
builder: &mut QueryBuilder<'_, MySql>,
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] #[async_trait]
impl ManagementTokenWriteRepository for MysqlManagementTokenRepository { impl ManagementTokenWriteRepository for MysqlManagementTokenRepository {
async fn create_management_token( async fn create_management_token(

View File

@@ -1,5 +1,5 @@
use async_trait::async_trait; use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use sqlx::{mysql::MySqlRow, Row};
use super::types::{ use super::types::{
OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig, OAuthProviderReadRepository, OAuthProviderWriteRepository, StoredOAuthProviderConfig,
@@ -8,7 +8,6 @@ use super::types::{
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::{push_eq, push_limit, WhereClause};
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct MysqlOAuthProviderRepository { pub struct MysqlOAuthProviderRepository {
@@ -24,17 +23,8 @@ impl MysqlOAuthProviderRepository {
&self, &self,
provider_type: &str, provider_type: &str,
) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> { ) -> Result<Option<StoredOAuthProviderConfig>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(OAUTH_PROVIDER_COLUMNS); let row = sqlx::query(GET_OAUTH_PROVIDER_CONFIG_SQL)
let mut where_clause = WhereClause::new(); .bind(provider_type)
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) .fetch_optional(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
@@ -42,7 +32,7 @@ impl MysqlOAuthProviderRepository {
} }
} }
const OAUTH_PROVIDER_COLUMNS: &str = r#" const LIST_OAUTH_PROVIDER_CONFIGS_SQL: &str = r#"
SELECT SELECT
provider_type, provider_type,
display_name, display_name,
@@ -60,6 +50,29 @@ SELECT
created_at AS created_at_unix_ms, created_at AS created_at_unix_ms,
updated_at AS updated_at_unix_secs updated_at AS updated_at_unix_secs
FROM oauth_providers 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#" const COUNT_LOCKED_USERS_IF_PROVIDER_DISABLED_SQL: &str = r#"
@@ -104,9 +117,10 @@ impl OAuthProviderReadRepository for MysqlOAuthProviderRepository {
async fn list_oauth_provider_configs( async fn list_oauth_provider_configs(
&self, &self,
) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> { ) -> Result<Vec<StoredOAuthProviderConfig>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(OAUTH_PROVIDER_COLUMNS); let rows = sqlx::query(LIST_OAUTH_PROVIDER_CONFIGS_SQL)
builder.push(" ORDER BY provider_type ASC"); .fetch_all(&self.pool)
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; .await
.map_sql_err()?;
rows.iter().map(map_oauth_provider_row).collect() rows.iter().map(map_oauth_provider_row).collect()
} }

View File

@@ -12,7 +12,6 @@ use super::{
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::repository::pool_scores::merge_score_reason_patch; use crate::repository::pool_scores::merge_score_reason_patch;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::{push_eq, push_in, push_limit, push_limit_offset, WhereClause};
const SCORE_COLUMNS: &str = r#" const SCORE_COLUMNS: &str = r#"
SELECT SELECT
@@ -58,54 +57,25 @@ impl MysqlPoolMemberScoreRepository {
scope: Option<&PoolScoreScope>, scope: Option<&PoolScoreScope>,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> { ) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS); let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
let mut where_clause = WhereClause::new(); builder
push_eq( .push(" WHERE pool_kind = ")
&mut builder, .push_bind(identity.pool_kind.clone())
&mut where_clause, .push(" AND pool_id = ")
"pool_kind", .push_bind(identity.pool_id.clone())
identity.pool_kind.clone(), .push(" AND member_kind = ")
); .push_bind(identity.member_kind.clone())
push_eq( .push(" AND member_id = ")
&mut builder, .push_bind(identity.member_id.clone());
&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 { if let Some(scope) = scope {
push_eq( builder
&mut builder, .push(" AND capability = ")
&mut where_clause, .push_bind(scope.capability.clone())
"capability", .push(" AND scope_kind = ")
scope.capability.clone(), .push_bind(scope.scope_kind.clone());
);
push_eq(
&mut builder,
&mut where_clause,
"scope_kind",
scope.scope_kind.clone(),
);
if let Some(scope_id) = &scope.scope_id { if let Some(scope_id) = &scope.scope_id {
push_eq( builder.push(" AND scope_id = ").push_bind(scope_id.clone());
&mut builder,
&mut where_clause,
"scope_id",
scope_id.clone(),
);
} else { } else {
where_clause.push_next(&mut builder); builder.push(" AND scope_id IS NULL");
builder.push("scope_id IS NULL");
} }
} }
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
@@ -120,65 +90,44 @@ impl PoolScoreReadRepository for MysqlPoolMemberScoreRepository {
query: &ListRankedPoolMembersQuery, query: &ListRankedPoolMembersQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> { ) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS); let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
let mut where_clause = WhereClause::new(); builder
push_eq( .push(" WHERE pool_kind = ")
&mut builder, .push_bind(query.pool_kind.clone())
&mut where_clause, .push(" AND pool_id = ")
"pool_kind", .push_bind(query.pool_id.clone())
query.pool_kind.clone(), .push(" AND capability = ")
); .push_bind(query.capability.clone())
push_eq( .push(" AND scope_kind = ")
&mut builder, .push_bind(query.scope_kind.clone());
&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 { if let Some(scope_id) = &query.scope_id {
push_eq( builder.push(" AND scope_id = ").push_bind(scope_id.clone());
&mut builder,
&mut where_clause,
"scope_id",
scope_id.clone(),
);
} else { } else {
where_clause.push_next(&mut builder); builder.push(" AND scope_id IS NULL");
builder.push("scope_id IS NULL");
} }
if !query.hard_states.is_empty() { if !query.hard_states.is_empty() {
let states = query builder.push(" AND hard_state IN (");
.hard_states let mut separated = builder.separated(", ");
.iter() for state in &query.hard_states {
.map(|state| state.as_database()) separated.push_bind(state.as_database());
.collect::<Vec<_>>(); }
push_in(&mut builder, &mut where_clause, "hard_state", &states); separated.push_unseparated(")");
} }
if let Some(statuses) = &query.probe_statuses { if let Some(statuses) = &query.probe_statuses {
if !statuses.is_empty() { if !statuses.is_empty() {
let statuses = statuses builder.push(" AND probe_status IN (");
.iter() let mut separated = builder.separated(", ");
.map(|status| status.as_database()) for status in statuses {
.collect::<Vec<_>>(); separated.push_bind(status.as_database());
push_in(&mut builder, &mut where_clause, "probe_status", &statuses); }
separated.push_unseparated(")");
} }
} }
builder.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC"); builder
push_limit_offset( .push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC")
&mut builder, .push(" LIMIT ")
i64_from_usize(query.limit.max(1), "pool score limit")?, .push_bind(i64_from_usize(query.limit.max(1), "pool score limit")?)
i64_from_usize(query.offset, "pool score offset")?, .push(" OFFSET ")
); .push_bind(i64_from_usize(query.offset, "pool score offset")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_score_row).collect() rows.iter().map(map_score_row).collect()
} }
@@ -188,66 +137,48 @@ impl PoolScoreReadRepository for MysqlPoolMemberScoreRepository {
query: &ListPoolMemberScoresQuery, query: &ListPoolMemberScoresQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> { ) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS); let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
let mut where_clause = WhereClause::new(); builder
push_eq( .push(" WHERE pool_kind = ")
&mut builder, .push_bind(query.pool_kind.clone())
&mut where_clause, .push(" AND pool_id = ")
"pool_kind", .push_bind(query.pool_id.clone());
query.pool_kind.clone(),
);
push_eq(
&mut builder,
&mut where_clause,
"pool_id",
query.pool_id.clone(),
);
if let Some(capability) = &query.capability { if let Some(capability) = &query.capability {
push_eq( builder
&mut builder, .push(" AND capability = ")
&mut where_clause, .push_bind(capability.clone());
"capability",
capability.clone(),
);
} }
if let Some(scope_kind) = &query.scope_kind { if let Some(scope_kind) = &query.scope_kind {
push_eq( builder
&mut builder, .push(" AND scope_kind = ")
&mut where_clause, .push_bind(scope_kind.clone());
"scope_kind",
scope_kind.clone(),
);
} }
if let Some(scope_id) = &query.scope_id { if let Some(scope_id) = &query.scope_id {
push_eq( builder.push(" AND scope_id = ").push_bind(scope_id.clone());
&mut builder,
&mut where_clause,
"scope_id",
scope_id.clone(),
);
} }
if !query.hard_states.is_empty() { if !query.hard_states.is_empty() {
let states = query builder.push(" AND hard_state IN (");
.hard_states let mut separated = builder.separated(", ");
.iter() for state in &query.hard_states {
.map(|state| state.as_database()) separated.push_bind(state.as_database());
.collect::<Vec<_>>(); }
push_in(&mut builder, &mut where_clause, "hard_state", &states); separated.push_unseparated(")");
} }
if let Some(statuses) = &query.probe_statuses { if let Some(statuses) = &query.probe_statuses {
if !statuses.is_empty() { if !statuses.is_empty() {
let statuses = statuses builder.push(" AND probe_status IN (");
.iter() let mut separated = builder.separated(", ");
.map(|status| status.as_database()) for status in statuses {
.collect::<Vec<_>>(); separated.push_bind(status.as_database());
push_in(&mut builder, &mut where_clause, "probe_status", &statuses); }
separated.push_unseparated(")");
} }
} }
builder.push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC"); builder
push_limit_offset( .push(" ORDER BY score DESC, last_ranked_at DESC, member_id ASC, id ASC")
&mut builder, .push(" LIMIT ")
i64_from_usize(query.limit.max(1), "pool score limit")?, .push_bind(i64_from_usize(query.limit.max(1), "pool score limit")?)
i64_from_usize(query.offset, "pool score offset")?, .push(" OFFSET ")
); .push_bind(i64_from_usize(query.offset, "pool score offset")?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_score_row).collect() rows.iter().map(map_score_row).collect()
} }
@@ -257,30 +188,18 @@ impl PoolScoreReadRepository for MysqlPoolMemberScoreRepository {
query: &ListPoolMemberProbeCandidatesQuery, query: &ListPoolMemberProbeCandidatesQuery,
) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> { ) -> Result<Vec<StoredPoolMemberScore>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS); let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
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 {
push_eq(
&mut builder,
&mut where_clause,
"capability",
capability.clone(),
);
}
where_clause.push_next(&mut builder);
builder builder
.push("hard_state IN ('available','unknown','cooldown','quota_exhausted')") .push(" WHERE pool_kind = ")
.push_bind(query.pool_kind.clone())
.push(" AND pool_id = ")
.push_bind(query.pool_id.clone());
if let Some(capability) = &query.capability {
builder
.push(" AND capability = ")
.push_bind(capability.clone());
}
builder
.push(" AND hard_state IN ('available','unknown','cooldown','quota_exhausted')")
.push(" AND (probe_status IN ('never','failed','stale')") .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(" OR (probe_status = 'ok' AND (last_probe_success_at IS NULL OR last_probe_success_at <= ")
.push_bind(i64_from_u64( .push_bind(i64_from_u64(
@@ -309,11 +228,12 @@ impl PoolScoreReadRepository for MysqlPoolMemberScoreRepository {
COALESCE(last_scheduled_at, 0) DESC, COALESCE(last_scheduled_at, 0) DESC,
member_id ASC member_id ASC
"#, "#,
); )
push_limit( .push(" LIMIT ")
&mut builder, .push_bind(i64_from_usize(
i64_from_usize(query.limit.max(1), "pool probe candidate limit")?, query.limit.max(1),
); "pool probe candidate limit",
)?);
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_score_row).collect() rows.iter().map(map_score_row).collect()
} }
@@ -326,8 +246,12 @@ impl PoolScoreReadRepository for MysqlPoolMemberScoreRepository {
return Ok(Vec::new()); return Ok(Vec::new());
} }
let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS); let mut builder = QueryBuilder::<MySql>::new(SCORE_COLUMNS);
let mut where_clause = WhereClause::new(); builder.push(" WHERE id IN (");
push_in(&mut builder, &mut where_clause, "id", &query.ids); let mut separated = builder.separated(", ");
for id in &query.ids {
separated.push_bind(id.clone());
}
separated.push_unseparated(")");
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_score_row).collect() rows.iter().map(map_score_row).collect()
} }

View File

@@ -1,5 +1,5 @@
use async_trait::async_trait; use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row}; use sqlx::{mysql::MySqlRow, Row};
use super::types::{ use super::types::{
bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample, bucket_start_unix_secs, build_tunnel_error_event_detail, build_tunnel_metrics_sample,
@@ -14,7 +14,6 @@ use super::types::{
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::{push_eq, push_limit, WhereClause};
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct MysqlProxyNodeReadRepository { pub struct MysqlProxyNodeReadRepository {
@@ -338,23 +337,13 @@ SELECT
FROM proxy_nodes 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] #[async_trait]
impl ProxyNodeReadRepository for MysqlProxyNodeReadRepository { impl ProxyNodeReadRepository for MysqlProxyNodeReadRepository {
async fn list_proxy_nodes(&self) -> Result<Vec<StoredProxyNode>, DataLayerError> { async fn list_proxy_nodes(&self) -> Result<Vec<StoredProxyNode>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(PROXY_NODE_COLUMNS); let rows = sqlx::query(&format!("{PROXY_NODE_COLUMNS} ORDER BY name ASC, id ASC"))
builder.push(" ORDER BY name ASC, id ASC"); .fetch_all(&self.pool)
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; .await
.map_sql_err()?;
rows.iter().map(map_proxy_node_row).collect() rows.iter().map(map_proxy_node_row).collect()
} }
@@ -362,12 +351,8 @@ impl ProxyNodeReadRepository for MysqlProxyNodeReadRepository {
&self, &self,
node_id: &str, node_id: &str,
) -> Result<Option<StoredProxyNode>, DataLayerError> { ) -> Result<Option<StoredProxyNode>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(PROXY_NODE_COLUMNS); let row = sqlx::query(&format!("{PROXY_NODE_COLUMNS} WHERE id = ? LIMIT 1"))
let mut where_clause = WhereClause::new(); .bind(node_id)
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) .fetch_optional(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
@@ -379,17 +364,26 @@ impl ProxyNodeReadRepository for MysqlProxyNodeReadRepository {
node_id: &str, node_id: &str,
limit: usize, limit: usize,
) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> { ) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(PROXY_NODE_EVENT_COLUMNS); let rows = sqlx::query(
let mut where_clause = WhereClause::new(); r#"
push_eq( SELECT
&mut builder, id,
&mut where_clause, node_id,
"node_id", event_type,
node_id.to_string(), detail,
); event_metadata,
builder.push(" ORDER BY created_at DESC, id DESC"); created_at AS created_at_unix_ms
push_limit(&mut builder, i64::try_from(limit).unwrap_or(i64::MAX)); FROM proxy_node_events
let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?; 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()?;
rows.iter().map(map_proxy_node_event_row).collect() rows.iter().map(map_proxy_node_event_row).collect()
} }
@@ -398,36 +392,51 @@ impl ProxyNodeReadRepository for MysqlProxyNodeReadRepository {
node_id: &str, node_id: &str,
query: &ProxyNodeEventQuery, query: &ProxyNodeEventQuery,
) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> { ) -> Result<Vec<StoredProxyNodeEvent>, DataLayerError> {
let mut builder = QueryBuilder::<MySql>::new(PROXY_NODE_EVENT_COLUMNS); let rows = sqlx::query(
let mut where_clause = WhereClause::new(); r#"
push_eq( SELECT
&mut builder, id,
&mut where_clause, node_id,
"node_id", event_type,
node_id.to_string(), detail,
); event_metadata,
if let Some(from_unix_secs) = query.from_unix_secs { created_at AS created_at_unix_ms
where_clause.push_next(&mut builder); FROM proxy_node_events
builder WHERE node_id = ?
.push("created_at >= ") AND (? IS NULL OR created_at >= ?)
.push_bind(i64::try_from(from_unix_secs).unwrap_or(i64::MAX)); AND (? IS NULL OR created_at <= ?)
} AND (? IS NULL OR LOWER(event_type) = LOWER(?))
if let Some(to_unix_secs) = query.to_unix_secs { ORDER BY created_at DESC, id DESC
where_clause.push_next(&mut builder); LIMIT ?
builder "#,
.push("created_at <= ") )
.push_bind(i64::try_from(to_unix_secs).unwrap_or(i64::MAX)); .bind(node_id)
} .bind(
if let Some(event_type) = query.event_type.as_deref() { query
where_clause.push_next(&mut builder); .from_unix_secs
builder .map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
.push("LOWER(event_type) = LOWER(") )
.push_bind(event_type.to_string()) .bind(
.push(")"); query
} .from_unix_secs
builder.push(" ORDER BY created_at DESC, id DESC"); .map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
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()?; .bind(
query
.to_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(
query
.to_unix_secs
.map(|v| i64::try_from(v).unwrap_or(i64::MAX)),
)
.bind(query.event_type.as_deref())
.bind(query.event_type.as_deref())
.bind(i64::try_from(query.limit).unwrap_or(i64::MAX))
.fetch_all(&self.pool)
.await
.map_sql_err()?;
rows.iter().map(map_proxy_node_event_row).collect() rows.iter().map(map_proxy_node_event_row).collect()
} }

View File

@@ -25,26 +25,22 @@ fn quota_snapshot_select() -> SelectQuery<'static> {
SelectColumn::expr(DialectSql::dialect( SelectColumn::expr(DialectSql::dialect(
"CAST(monthly_quota_usd AS DOUBLE PRECISION)", "CAST(monthly_quota_usd AS DOUBLE PRECISION)",
"CAST(monthly_quota_usd AS REAL)", "CAST(monthly_quota_usd AS REAL)",
"monthly_quota_usd",
)) ))
.alias("monthly_quota_usd"), .alias("monthly_quota_usd"),
SelectColumn::expr(DialectSql::dialect( SelectColumn::expr(DialectSql::dialect(
"CAST(COALESCE(monthly_used_usd, 0) AS DOUBLE PRECISION)", "CAST(COALESCE(monthly_used_usd, 0) AS DOUBLE PRECISION)",
"CAST(COALESCE(monthly_used_usd, 0) AS REAL)", "CAST(COALESCE(monthly_used_usd, 0) AS REAL)",
"COALESCE(monthly_used_usd, 0)",
)) ))
.alias("monthly_used_usd"), .alias("monthly_used_usd"),
SelectColumn::expr("quota_reset_day"), SelectColumn::expr("quota_reset_day"),
SelectColumn::expr(DialectSql::dialect( SelectColumn::expr(DialectSql::dialect(
"CAST(EXTRACT(EPOCH FROM quota_last_reset_at) AS BIGINT)", "CAST(EXTRACT(EPOCH FROM quota_last_reset_at) AS BIGINT)",
"quota_last_reset_at", "quota_last_reset_at",
"quota_last_reset_at",
)) ))
.alias("quota_last_reset_at_unix_secs"), .alias("quota_last_reset_at_unix_secs"),
SelectColumn::expr(DialectSql::dialect( SelectColumn::expr(DialectSql::dialect(
"CAST(EXTRACT(EPOCH FROM quota_expires_at) AS BIGINT)", "CAST(EXTRACT(EPOCH FROM quota_expires_at) AS BIGINT)",
"quota_expires_at", "quota_expires_at",
"quota_expires_at",
)) ))
.alias("quota_expires_at_unix_secs"), .alias("quota_expires_at_unix_secs"),
SelectColumn::expr("is_active"), SelectColumn::expr("is_active"),

View File

@@ -1,14 +1,25 @@
use async_trait::async_trait; use async_trait::async_trait;
use sqlx::{mysql::MySqlRow, MySql, Row}; use sqlx::{mysql::MySqlRow, MySql, QueryBuilder, Row};
use super::{ use super::{
quota_snapshot_select, ProviderQuotaReadRepository, ProviderQuotaWriteRepository, ProviderQuotaReadRepository, ProviderQuotaWriteRepository, StoredProviderQuotaSnapshot,
StoredProviderQuotaSnapshot,
}; };
use crate::driver::mysql::MysqlPool; use crate::driver::mysql::MysqlPool;
use crate::error::SqlResultExt; use crate::error::SqlResultExt;
use crate::DataLayerError; use crate::DataLayerError;
use aether_data_query::SqlDialect;
const QUOTA_COLUMNS: &str = r#"
SELECT
id AS provider_id,
billing_type,
monthly_quota_usd,
COALESCE(monthly_used_usd, 0) 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
"#;
#[derive(Debug, Clone)] #[derive(Debug, Clone)]
pub struct MysqlProviderQuotaRepository { pub struct MysqlProviderQuotaRepository {
@@ -27,11 +38,8 @@ impl ProviderQuotaReadRepository for MysqlProviderQuotaRepository {
&self, &self,
provider_id: &str, provider_id: &str,
) -> Result<Option<StoredProviderQuotaSnapshot>, DataLayerError> { ) -> Result<Option<StoredProviderQuotaSnapshot>, DataLayerError> {
let mut statement = quota_snapshot_select().statement::<MySql>(SqlDialect::Mysql); let row = sqlx::query(&format!("{QUOTA_COLUMNS} WHERE id = ? LIMIT 1"))
statement.where_eq("id", provider_id.to_string()).limit(1); .bind(provider_id)
let row = statement
.finish()
.build()
.fetch_optional(&self.pool) .fetch_optional(&self.pool)
.await .await
.map_sql_err()?; .map_sql_err()?;
@@ -46,16 +54,16 @@ impl ProviderQuotaReadRepository for MysqlProviderQuotaRepository {
return Ok(Vec::new()); return Ok(Vec::new());
} }
let mut statement = quota_snapshot_select().statement::<MySql>(SqlDialect::Mysql); let mut builder = QueryBuilder::<MySql>::new(QUOTA_COLUMNS);
statement builder.push(" WHERE id IN (");
.where_in("id", provider_ids) {
.order_by_sql("id ASC"); let mut separated = builder.separated(", ");
let rows = statement for provider_id in provider_ids {
.finish() separated.push_bind(provider_id);
.build() }
.fetch_all(&self.pool) }
.await builder.push(") ORDER BY id ASC");
.map_sql_err()?; let rows = builder.build().fetch_all(&self.pool).await.map_sql_err()?;
rows.iter().map(map_row).collect() rows.iter().map(map_row).collect()
} }
} }

View File

@@ -3,8 +3,7 @@
This inventory tracks repository read paths that are intended to use the This inventory tracks repository read paths that are intended to use the
internal `aether-data-query` helpers. The first layer centralizes SQL fragments; internal `aether-data-query` helpers. The first layer centralizes SQL fragments;
the newer `SelectQuery` layer lets repositories describe simple `SELECT` the newer `SelectQuery` layer lets repositories describe simple `SELECT`
queries once and render dialect-specific projections for Postgres, SQLite, and queries once and render dialect-specific projections for Postgres and SQLite.
MySQL.
## Included In This Pass ## Included In This Pass
@@ -27,7 +26,7 @@ MySQL.
- `find_by_provider_id` - `find_by_provider_id`
- `find_by_provider_ids` - `find_by_provider_ids`
- now uses one `SelectQuery` specification for the quota snapshot projection - now uses one `SelectQuery` specification for the quota snapshot projection
across Postgres, SQLite, and MySQL across Postgres and SQLite
- `provider_catalog` - `provider_catalog`
- provider by-id/provider list reads in PG/SQLite - provider by-id/provider list reads in PG/SQLite
- endpoint/key by-id and by-provider-id `IN` reads in PG/SQLite - endpoint/key by-id and by-provider-id `IN` reads in PG/SQLite
@@ -63,9 +62,6 @@ MySQL.
- `wallet` ledger, order, refund, callback, and redeem-code list logic. - `wallet` ledger, order, refund, callback, and redeem-code list logic.
- write/upsert/delete paths, transactions, `RETURNING`, CTEs, window functions, - write/upsert/delete paths, transactions, `RETURNING`, CTEs, window functions,
advisory locks, and schema compatibility probes. advisory locks, and schema compatibility probes.
- MySQL `provider_catalog` still delegates read paths through the existing
memory adapter; migrate it in a focused follow-up so MySQL parity can be tested
independently.
- `users/auth` and `global_models` still contain additional simple read paths. - `users/auth` and `global_models` still contain additional simple read paths.
`global_models/sqlite.rs` had pre-existing local edits and must be handled `global_models/sqlite.rs` had pre-existing local edits and must be handled
carefully in a dedicated slice. carefully in a dedicated slice.