diff --git a/docs/crud-where-dsl.md b/docs/crud-where-dsl.md new file mode 100644 index 00000000..4018924a --- /dev/null +++ b/docs/crud-where-dsl.md @@ -0,0 +1,284 @@ +# CRUD Where DSL 设计 + +## 背景 + +当前 `crud_query_paged!` 要求用户手写原始 SQL: + +```rust +crud_query_paged!(pool, Order, + data_sql: "SELECT * FROM orders WHERE user_id = ?{tenant} ORDER BY created_at DESC", + count_sql: "SELECT COUNT(*) FROM orders WHERE user_id = ?{tenant}", + binds: [user_id], + tenant: tenant_id, + page: page, + page_size: page_size +) +``` + +### 问题 + +1. **占位符不兼容** — `?` 在 PostgreSQL 下必须写成 `$N`,用户手写 `?` 会导致 PG 运行时报错 +2. **`{tenant}` hack** — 字符串替换方式注入租户条件,脆弱且不优雅 +3. **无编译时校验** — 列名拼写错误只能在运行时发现 +4. **跨库语法陷阱** — 用户可能无意中写了 `ILIKE`、`::text`、`RETURNING` 等 PG 特有语法 + +## 设计目标 + +1. 用户**不再写原始 SQL 的 WHERE 子句**,改用结构化 DSL +2. 宏完全控制 SQL 生成,自动适配 `$N` / `?N` / `?` +3. 编译时校验列名 +4. `{tenant}` 自动注入,用户不感知 +5. 覆盖 80% 的常见查询场景,剩余 20% 用 `crud_query!` + `Driver::ph()` 手写 + +## DSL 语法 + +### 基本形式 + +``` +("column", value) → col = ? +("column", OP, value) → col OP ? +``` + +### 逻辑组合 + +``` +AND( cond1, cond2, ... ) → (cond1 AND cond2 AND ...) +OR( cond1, cond2, ... ) → (cond1 OR cond2 OR ...) +``` + +可任意嵌套: + +``` +AND( ("status", 1), OR(("role", "admin"), ("role", "editor")) ) +``` + +生成 SQL: + +```sql +(status = ?) AND ((role = ?) OR (role = ?)) +``` + +### 运算符 + +| 运算符 | SQL | 元素 | 示例 | +|---|---|---|---| +| `EQ`(默认,可省略) | `= ?` | 2-tuple 或 3-tuple | `("id", 1)` 或 `("id", EQ, 1)` | +| `NEQ` | `!= ?` | 3-tuple | `("status", NEQ, "deleted")` | +| `GT` / `GTE` | `> ?` / `>= ?` | 3-tuple | `("amount", GT, 100)` | +| `LT` / `LTE` | `< ?` / `<= ?` | 3-tuple | `("created_at", LTE, now)` | +| `LIKE` | `LIKE ?` | 3-tuple | `("title", LIKE, "%rust%")` | +| `NOT_LIKE` | `NOT LIKE ?` | 3-tuple | `("title", NOT_LIKE, "%spam%")` | +| `IN` | `IN (?,?,?)` | 3-tuple,value 为 Vec | `("status", IN, vec!["a","b"])` | +| `NOT_IN` | `NOT IN (?,?,?)` | 3-tuple,value 为 Vec | `("status", NOT_IN, vec!["x"])` | +| `IS_NULL` | `IS NULL` | 2-tuple,value 为 `()` | `("deleted_at", ())` | +| `NOT_NULL` | `IS NOT NULL` | 特殊标记 | 待定 | + +### 可选条件(动态 WHERE) + +用 `Option` 包装值,`None` 时跳过该条件: + +```rust +// status 为 None 时不加入 WHERE +where: AND(("user_id", uid), opt!("status", status)) +``` + +`opt!` 宏在编译时展开为条件绑定代码: + +```rust +// 生成的代码 +if let Some(ref __wv) = status { + __ph_idx += 1; + __where_sql.push_str(&format!(" AND status = {}", __ph(__ph_idx))); +} +``` + +### 完整调用示例 + +迁移前(手写 SQL): + +```rust +crud_query_paged!(pool, Comment, + data_sql: "SELECT * FROM comments WHERE post_id = ? AND status = ?{tenant} ORDER BY created_at ASC", + count_sql: "SELECT COUNT(*) FROM comments WHERE post_id = ? AND status = ?{tenant}", + binds: [post_id, CommentStatus::Approved], + tenant: tenant_id, + page: page, + page_size: page_size +) +``` + +迁移后(DSL): + +```rust +crud_query_paged!(pool, Comment, + table: "comments", + where: AND(("post_id", post_id), ("status", CommentStatus::Approved)), + order_by: "created_at ASC", + page: page, + page_size: page_size, + tenant: tenant_id +) +``` + +- 无 `data_sql` / `count_sql` / `binds` — 宏从 DSL 自动生成 +- 无 `{tenant}` — `tenant:` 参数自动注入 +- 无 `?` — 占位符由宏根据 dialect 自动生成 + +## 宏 API 对比 + +### 现有 API(保留) + +```rust +crud_query_paged!(pool, Type, + data_sql: "...", // 用户手写 SQL,含 ? 和 {tenant} + count_sql: "...", // 用户手写 SQL + binds: [val1, val2], // 绑定值 + where: [...], // 可选的动态条件 + tenant: expr, + page: expr, + page_size: expr +) +``` + +### 新 API(推荐) + +```rust +crud_query_paged!(pool, Type, + table: "orders", // 表名 + where: AND(("user_id", uid), opt!("status", status)), // DSL 条件 + order_by: "created_at DESC", // 排序 + tenant: expr, // 自动注入租户 + page: expr, + page_size: expr +) +``` + +**解析策略**:宏检测第一个命名参数。如果是 `data_sql:` → 走旧路径;如果是 `table:` → 走新路径。向后兼容。 + +## 现有调用迁移分析 + +| # | 文件 | 现有 WHERE | DSL 等价 | +|---|---|---|---| +| 1 | `wallet_transaction.rs` | `wallet_id = ?` | `("wallet_id", wallet_id)` | +| 2 | `wallet_transaction.rs` | `user_id = ?` | `("user_id", user_id)` | +| 3 | `wallet_transaction.rs` | `1=1{tenant}` | 无 where,仅 tenant | +| 4 | `page.rs` | `status = ?{tenant}` | `("status", status)` | +| 5 | `order.rs` | `user_id = ?{tenant}` | `("user_id", user_id)` | +| 6 | `media.rs` | `user_id = ?{tenant}` | `("user_id", user_id)` | +| 7 | `payment_order.rs` | `user_id = ?{tenant}` | `("user_id", user_id)` | +| 8 | `comment.rs` | `post_id = ? AND status = ?{tenant}` | `AND(("post_id", post_id), ("status", status))` | +| 9-28 | 其他 20 处 | `1=1{tenant}` + 可选 where | 无 where + `where: [...]` | + +**100% 现有用法可迁移。** + +## 不适合 DSL 的场景 + +以下场景继续使用 `crud_query!` + `Driver::ph()` 手写: + +| 场景 | 原因 | 示例 | +|---|---|---| +| 子查询 | `WHERE id IN (SELECT ...)` | `worker/job_queue.rs` dequeue | +| 聚合 | `GROUP BY` / `HAVING` / `SUM()` | `models/payment_refund.rs` sum_refunded | +| CASE 表达式 | `CASE WHEN ... THEN ... END` | `models/wallet_outbox.rs` mark_failed | +| 动态表名 | `FROM {dynamic_table}` | `services/stats.rs` | +| 动态列名 | `SET {col} = ?` | `models/order.rs` timestamp_col | +| 复杂 JOIN | 多表 JOIN + 动态 WHERE | `models/post.rs` | +| CAS 更新 | `version = version + 1` | `models/order.rs` tx_update_status_cas | + +## 实现计划 + +### Phase 1: 解析器 + +在 `raisfast-derive/src/crud.rs` 中新增 DSL 解析: + +- `WhereExpr` 枚举:`Condition` / `And` / `Or` / `Optional` +- `Operator` 枚举:`EQ` / `NEQ` / `GT` / `GTE` / `LT` / `LTE` / `LIKE` / `IN` / `IS_NULL` ... +- `Condition` 结构体:`(col, [op,] value)` + +### Phase 2: SQL 生成 + +`WhereExpr` → SQL 字符串 + 绑定参数列表: + +- 递归遍历 AST,生成 `col = ?N` / `col > ?N` 等 +- 自动追踪 `__ph_idx` +- tenant 条件自动追加 + +### Phase 3: 编译时校验 + +对 DSL 中的每个列名调用 `validate_column(table, col)`,编译时报错拼写错误。 + +### Phase 4: 旧 API 兼容 + +- `data_sql:` 路径保留,加编译警告建议迁移 +- `crud_join_paged!` 同步支持 DSL + +### Phase 5: 迁移现有调用 + +将 28 处 `crud_query_paged!` 从手写 SQL 迁移到 DSL。 + +## Rust 类型表达(待定) + +在 proc-macro 中解析 DSL 有几种方案: + +### 方案 A: Token 解析 + +```rust +// 用户写法 +where: AND(("post_id", post_id), ("status", status)) + +// proc-macro 解析 TokenStream +// 识别 AND(...) / OR(...) / (...) 三种模式 +``` + +优点:语法简洁,接近 SQL 思维。 +难点:Rust 宏里 `(expr, expr)` 是 tuple 字面量,需要自定义解析器。 + +### 方案 B: 数组语法 + +```rust +// 用户写法 +where: [AND, [("post_id", post_id), ("status", status)]] + +// proc-macro 更容易解析 +``` + +优点:解析简单。 +缺点:不直观。 + +### 方案 C: 关键字语法 + +```rust +// 用户写法 +where: "post_id" => post_id, and: ["status" => status] + +// 与现有 crud_find! 语法一致 +``` + +优点:已有先例,解析器已存在。 +缺点:不支持 OR / 嵌套 / 运算符。 + +**推荐方案 A**,因为它最接近 SQL 表达力,且 proc-macro 有完整的 TokenStream 解析能力。 + +## 占位符生成规则 + +| Dialect | `EQ` | `GT` | `IN(vec![a,b,c])` | `IS_NULL` | `LIKE` | +|---|---|---|---|---|---| +| SQLite | `?1` | `?2` | `IN (?3,?4,?5)` | `IS NULL` | `LIKE ?6` | +| PostgreSQL | `$1` | `$2` | `IN ($3,$4,$5)` | `IS NULL` | `LIKE $6` | +| MySQL | `?` | `?` | `IN (?,?,?)` | `IS NULL` | `LIKE ?` | + +所有占位符由 `Dialect::ph(idx)` 生成,idx 由宏在遍历 AST 时自动递增。 + +## 与 `crud_join_paged!` 的关系 + +`crud_join_paged!` 已自行生成 WHERE SQL(不依赖用户手写),但只有 `=` 运算符。 +DSL 可同时应用于两个宏,统一条件表达方式。 + +## 风险和缓解 + +| 风险 | 缓解 | +|---|---| +| DSL 解析复杂度 | 递归下降解析器,参考 `expand_find` 现有的 `and:` 解析 | +| proc-macro 编译变慢 | DSL 只在 `crud_query_paged!` 中使用,影响面小 | +| 用户学习成本 | DSL 元素与 SQL 一一对应,直觉易学 | +| 不支持的 SQL 回退 | 保留 `crud_query!` + `Driver::ph()` 手写路径 | diff --git a/raisfast-derive/src/crud.rs b/raisfast-derive/src/crud.rs index 6e7ef4b3..da1ea535 100644 --- a/raisfast-derive/src/crud.rs +++ b/raisfast-derive/src/crud.rs @@ -87,9 +87,15 @@ fn validate_table(table: &syn::LitStr) -> Option { /// Validate that a column exists in a table. Returns `None` if valid, or a compile error. fn validate_column(table: &syn::LitStr, col: &syn::LitStr) -> Option { - let table_str = table.value(); + validate_column_inner(&table.value(), col).map(Into::into) +} + +pub fn validate_column_inner( + table_str: &str, + col: &syn::LitStr, +) -> Option { let col_str = col.value(); - if let Some(ts) = get_schema().tables.get(&table_str) + if let Some(ts) = get_schema().tables.get(table_str) && ts.columns.iter().any(|c| c.name == col_str) { return None; @@ -102,8 +108,7 @@ fn validate_column(table: &syn::LitStr, col: &syn::LitStr) -> Option Option, - gt_cols: Vec, - gt_vals: Vec, - lt_cols: Vec, - lt_vals: Vec, - gte_cols: Vec, - gte_vals: Vec, - lte_cols: Vec, - lte_vals: Vec, - in_cols: Vec, - in_vals: Vec, -} - -impl ExtraConds { - fn is_empty(&self) -> bool { - self.null_cols.is_empty() - && self.gt_cols.is_empty() - && self.lt_cols.is_empty() - && self.gte_cols.is_empty() - && self.lte_cols.is_empty() - && self.in_cols.is_empty() - } -} - -/// Parse extra condition parameters from a `while` loop section dispatch. -/// Call this inside the section-matching branch for each supported keyword. -fn parse_extra_conds_section( - ecs: &mut ExtraConds, - section: &syn::Ident, - input: syn::parse::ParseStream, -) -> syn::Result { - let s = section.to_string(); - match s.as_str() { - "and_null" => { - let content; - syn::bracketed!(content in input); - while !content.is_empty() { - ecs.null_cols.push(content.parse()?); - let _ = content.parse::(); +fn emit_runtime_binds(binds: &[crate::where_dsl::BindKind]) -> (Vec, Vec) { + let local_stmts: Vec<_> = binds + .iter() + .enumerate() + .map(|(i, bk)| match bk { + crate::where_dsl::BindKind::Static(expr) => { + let ident = syn::Ident::new(&format!("__wb_{}", i), proc_macro2::Span::call_site()); + quote! { let #ident = #expr; } } - Ok(true) - } - "and_gt" => { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - ecs.gt_cols = c; - ecs.gt_vals = v; - Ok(true) - } - "and_lt" => { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - ecs.lt_cols = c; - ecs.lt_vals = v; - Ok(true) - } - "and_gte" => { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - ecs.gte_cols = c; - ecs.gte_vals = v; - Ok(true) - } - "and_lte" => { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - ecs.lte_cols = c; - ecs.lte_vals = v; - Ok(true) - } - "and_in" => { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - ecs.in_cols = c; - ecs.in_vals = v; - Ok(true) - } - _ => Ok(false), - } + crate::where_dsl::BindKind::InLoop(expr) => { + let ident = syn::Ident::new(&format!("__in_{}", i), proc_macro2::Span::call_site()); + quote! { let #ident = #expr; } + } + }) + .collect(); + + let bind_stmts: Vec<_> = binds + .iter() + .enumerate() + .map(|(i, bk)| match bk { + crate::where_dsl::BindKind::Static(_) => { + let ident = syn::Ident::new(&format!("__wb_{}", i), proc_macro2::Span::call_site()); + quote! { __q = __q.bind(#ident); } + } + crate::where_dsl::BindKind::InLoop(_) => { + let ident = syn::Ident::new(&format!("__in_{}", i), proc_macro2::Span::call_site()); + quote! { for __iv in #ident { __q = __q.bind(__iv.clone()); } } + } + }) + .collect(); + + (local_stmts, bind_stmts) } -/// Build SQL fragments for extra conditions, appending to `parts` and incrementing `ph_idx`. -/// Returns a list of value expressions to bind. -fn build_extra_conds_sql( - ecs: &ExtraConds, - d: Dialect, - ph_idx: &mut usize, -) -> (Vec, Vec) { - let mut parts = Vec::new(); - let mut vals = Vec::new(); +fn emit_tenant_code(tid: &Option, d: Dialect, tenant_alias: Option<&syn::LitStr>) -> (proc_macro2::TokenStream, proc_macro2::TokenStream) { + let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); + let numbered_lit = syn::LitBool::new(numbered, proc_macro2::Span::call_site()); + let ph_prefix_lit = syn::LitStr::new( + match d { + Dialect::Postgres => "$", + _ => "?", + }, + proc_macro2::Span::call_site(), + ); + let alias_prefix = match tenant_alias { + Some(a) => format!("{}.", a.value()), + None => String::new(), + }; + let alias_prefix_lit = syn::LitStr::new(&alias_prefix, proc_macro2::Span::call_site()); - for col in &ecs.null_cols { - parts.push(format!("AND {} IS NULL", col.value())); + if tid.is_some() { + let tid_expr = tid.as_ref().unwrap(); + let tenant_sql = quote! { + let (__tenant_sql, __tid_val) = match #tid_expr { + Some(_tid) => { + let __tph = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + __ph_idx += 1; + (format!(" AND {}tenant_id = {}", #alias_prefix_lit, __tph), Some(_tid)) + }, + None => (String::new(), None), + }; + }; + let tenant_bind = quote! { + if let Some(_tid) = __tid_val { + __q = __q.bind(_tid); + } + }; + (tenant_sql, tenant_bind) + } else { + (quote! { let __tenant_sql = String::new(); }, quote! {}) } - for (col, val) in ecs.gt_cols.iter().zip(&ecs.gt_vals) { - let ph = d.ph(*ph_idx); - *ph_idx += 1; - parts.push(format!("AND {} > {}", col.value(), ph)); - vals.push(val.clone()); - } - for (col, val) in ecs.lt_cols.iter().zip(&ecs.lt_vals) { - let ph = d.ph(*ph_idx); - *ph_idx += 1; - parts.push(format!("AND {} < {}", col.value(), ph)); - vals.push(val.clone()); - } - for (col, val) in ecs.gte_cols.iter().zip(&ecs.gte_vals) { - let ph = d.ph(*ph_idx); - *ph_idx += 1; - parts.push(format!("AND {} >= {}", col.value(), ph)); - vals.push(val.clone()); - } - for (col, val) in ecs.lte_cols.iter().zip(&ecs.lte_vals) { - let ph = d.ph(*ph_idx); - *ph_idx += 1; - parts.push(format!("AND {} <= {}", col.value(), ph)); - vals.push(val.clone()); - } - - (parts, vals) -} - -/// Build a list of extra column names for schema validation. -fn extra_conds_columns(ecs: &ExtraConds) -> Vec { - let mut cols = Vec::new(); - cols.extend_from_slice(&ecs.null_cols); - cols.extend_from_slice(&ecs.gt_cols); - cols.extend_from_slice(&ecs.lt_cols); - cols.extend_from_slice(&ecs.gte_cols); - cols.extend_from_slice(&ecs.lte_cols); - cols.extend_from_slice(&ecs.in_cols); - cols } // ── crud_delete! ──────────────────────────────────────────────────────── @@ -270,129 +211,41 @@ pub fn crud_delete(input: TokenStream) -> TokenStream { } fn expand_delete(input: TokenStream) -> TokenStream { - let parsed = parse_macro_input!(input as DeleteInput); + let parsed = parse_macro_input!(input as WhereOnlyInput); let table = &parsed.table; - let col = &parsed.col; if let Some(err) = validate_table(table) { return err; } - if let Some(err) = validate_column(table, col) { - return err; - } - if let Some(err) = validate_columns(table, &parsed.and_cols) { - return err; - } - if let Some(err) = validate_columns(table, &extra_conds_columns(&parsed.ecs)) { - return err; + if let Some(err) = crate::where_dsl::validate_where_columns(&parsed.dsl_where, &table.value()) { + return err.into(); } let pool = &parsed.pool; - let val = &parsed.val; - let tid = &parsed.tid; - let table_str = table.value(); - let col_str = col.value(); - let and_cols = &parsed.and_cols; - let and_vals = &parsed.and_vals; - let d = dialect(); - let mut ph_idx = 1usize; - let col_ph = d.ph(ph_idx); - ph_idx += 1; - let mut and_parts: Vec = and_cols - .iter() - .map(|ac| { - let ph = d.ph(ph_idx); - ph_idx += 1; - format!("AND {} = {}", ac.value(), ph) - }) - .collect(); - let (ecs_parts, ecs_vals) = build_extra_conds_sql(&parsed.ecs, d, &mut ph_idx); - and_parts.extend(ecs_parts); - let all_extra_vals: Vec = and_vals.iter().chain(ecs_vals.iter()).cloned().collect(); - let and_str = and_parts.join(" "); + let table_str = table.value(); + let table_lit = syn::LitStr::new(&table_str, table.span()); - let has_extra = !and_cols.is_empty() || !parsed.ecs.is_empty(); + let wr = crate::where_dsl::generate_where_runtime(&parsed.dsl_where, d); + let (local_stmts, bind_stmts) = emit_runtime_binds(&wr.binds); + let (tenant_sql, tenant_bind) = emit_tenant_code(&parsed.tid, d, None); + let sql_code = wr.sql_code; - if parsed.tid.is_some() { - let tid_ph = d.ph(ph_idx); - - let expanded = if !has_extra { - let sql_with_tenant = syn::LitStr::new( - &format!( - "DELETE FROM {} WHERE {} = {} AND tenant_id = {}", - table_str, col_str, col_ph, tid_ph - ), - table.span(), - ); - let sql_without_tenant = syn::LitStr::new( - &format!("DELETE FROM {} WHERE {} = {}", table_str, col_str, col_ph), - table.span(), - ); - quote! { - match #tid { - Some(_tid) => sqlx::query!(#sql_with_tenant, #val, _tid).execute(#pool).await, - None => sqlx::query!(#sql_without_tenant, #val).execute(#pool).await, - } - } - } else { - let extra_idents: Vec = (0..all_extra_vals.len()) - .map(|i| syn::Ident::new(&format!("__da_{}", i), proc_macro2::Span::call_site())) - .collect(); - let sql_with_tenant = syn::LitStr::new( - &format!( - "DELETE FROM {} WHERE {} = {} {} AND tenant_id = {}", - table_str, col_str, col_ph, and_str, tid_ph - ), - table.span(), - ); - let sql_without_tenant = syn::LitStr::new( - &format!( - "DELETE FROM {} WHERE {} = {} {}", - table_str, col_str, col_ph, and_str - ), - table.span(), - ); - quote! { - { - #(let #extra_idents = #all_extra_vals;)* - match #tid { - Some(_tid) => sqlx::query!(#sql_with_tenant, #val, #(#extra_idents),*, _tid).execute(#pool).await, - None => sqlx::query!(#sql_without_tenant, #val, #(#extra_idents),*).execute(#pool).await, - } - } - } - }; - TokenStream::from(expanded) - } else { - let expanded = if !has_extra { - let sql_lit = syn::LitStr::new( - &format!("DELETE FROM {} WHERE {} = {}", table_str, col_str, col_ph), - table.span(), - ); - quote! { - sqlx::query!(#sql_lit, #val).execute(#pool).await - } - } else { - let extra_idents: Vec = (0..all_extra_vals.len()) - .map(|i| syn::Ident::new(&format!("__da_{}", i), proc_macro2::Span::call_site())) - .collect(); - let sql_lit = syn::LitStr::new( - &format!( - "DELETE FROM {} WHERE {} = {} {}", - table_str, col_str, col_ph, and_str - ), - table.span(), - ); - quote! { - { - #(let #extra_idents = #all_extra_vals;)* - sqlx::query!(#sql_lit, #val, #(#extra_idents),*).execute(#pool).await - } - } - }; - TokenStream::from(expanded) - } + let expanded = quote! { + { + #(#local_stmts)* + let mut __ph_idx: usize = 1; + let mut __where_sql = String::new(); + #sql_code + #tenant_sql + let __sql = format!("DELETE FROM {} WHERE {}{}", #table_lit, __where_sql, __tenant_sql); + let mut __q = sqlx::query(&__sql); + #(#bind_stmts)* + #tenant_bind + __q.execute(#pool).await + } + }; + TokenStream::from(expanded) } // ── crud_insert! ──────────────────────────────────────────────────────── @@ -527,113 +380,44 @@ fn expand_select(input: TokenStream) -> TokenStream { if let Some(err) = validate_table(table) { return err; } - if let Some(err) = validate_column(table, &parsed.col) { - return err; - } if let Some(err) = validate_columns(table, &parsed.sel_cols) { return err; } - if let Some(err) = validate_columns(table, &parsed.and_cols) { - return err; + if let Some(err) = crate::where_dsl::validate_where_columns(&parsed.dsl_where, &table.value()) { + return err.into(); } let pool = &parsed.pool; - let val = &parsed.val; let table_str = table.value(); - let col_str = parsed.col.value(); + let table_lit = syn::LitStr::new(&table_str, table.span()); let sel_str: String = parsed .sel_cols .iter() .map(|l| l.value()) .collect::>() .join(", "); - let and_cols = &parsed.and_cols; - let and_vals = &parsed.and_vals; - - let table_lit = syn::LitStr::new(&table_str, table.span()); let sel_lit = syn::LitStr::new(&sel_str, table.span()); - let col_lit = syn::LitStr::new(&col_str, table.span()); - let and_col_lits: Vec = and_cols - .iter() - .enumerate() - .map(|(i, ac)| { - syn::LitStr::new( - &format!(" AND {} = {}", ac.value(), d.ph(2 + i)), - table.span(), - ) - }) - .collect(); + let wr = crate::where_dsl::generate_where_runtime(&parsed.dsl_where, d); + let (local_stmts, bind_stmts) = emit_runtime_binds(&wr.binds); + let (tenant_sql, tenant_bind) = emit_tenant_code(&parsed.tid, d, None); + let sql_code = wr.sql_code; - let and_val_idents: Vec = (0..and_vals.len()) - .map(|i| syn::Ident::new(&format!("__sav_{}", i), proc_macro2::Span::call_site())) - .collect(); - - if parsed.tid.is_some() { - let tid = &parsed.tid; - let tid_fmt = syn::LitStr::new( - &format!( - "SELECT {{}} FROM {{}} WHERE {{}} = {}{{}} AND tenant_id = {}", - d.ph(1), - d.ph(2 + and_cols.len()) - ), - table.span(), - ); - let no_tid_fmt = syn::LitStr::new( - &format!("SELECT {{}} FROM {{}} WHERE {{}} = {}{{}}", d.ph(1)), - table.span(), - ); - let expanded = quote! { - { - let __sv = #val; - #(let #and_val_idents = #and_vals;)* - let __and_sql: &str = concat!(#(#and_col_lits),*); - let __sql = match #tid { - Some(_tid) => format!(#tid_fmt, #sel_lit, #table_lit, #col_lit, __and_sql), - None => format!(#no_tid_fmt, #sel_lit, #table_lit, #col_lit, __and_sql), - }; - let mut _q = sqlx::query_as::<_, _>(&__sql).bind(__sv); - #(_q = _q.bind(#and_val_idents);)* - if let Some(_tid) = #tid { - _q = _q.bind(_tid); - } - _q.fetch_optional(#pool).await - } - }; - TokenStream::from(expanded) - } else { - let and_sql: String = and_cols - .iter() - .enumerate() - .map(|(i, ac)| format!(" AND {} = {}", ac.value(), d.ph(2 + i))) - .collect(); - let sql_str = format!( - "SELECT {} FROM {} WHERE {} = {}{}", - sel_str, - table_str, - col_str, - d.ph(1), - and_sql - ); - let sql = syn::LitStr::new(&sql_str, table.span()); - - if and_vals.is_empty() { - let expanded = quote! { - sqlx::query_as::<_, _>(#sql).bind(#val).fetch_optional(#pool).await - }; - TokenStream::from(expanded) - } else { - let expanded = quote! { - { - #(let #and_val_idents = #and_vals;)* - let mut _q = sqlx::query_as::<_, _>(#sql).bind(#val); - #(_q = _q.bind(#and_val_idents);)* - _q.fetch_optional(#pool).await - } - }; - TokenStream::from(expanded) + let expanded = quote! { + { + #(#local_stmts)* + let mut __ph_idx: usize = 1; + let mut __where_sql = String::new(); + #sql_code + #tenant_sql + let __sql = format!("SELECT {} FROM {} WHERE {}{}", #sel_lit, #table_lit, __where_sql, __tenant_sql); + let mut __q = sqlx::query_as::<_, _>(&__sql); + #(#bind_stmts)* + #tenant_bind + __q.fetch_optional(#pool).await } - } + }; + TokenStream::from(expanded) } // ── crud_query! ────────────────────────────────────────── @@ -705,41 +489,31 @@ fn expand_find(input: TokenStream, method: FindMethod) -> TokenStream { if let Some(err) = validate_table(table) { return err; } - if let Some(err) = validate_column(table, &parsed.col) { - return err; - } - if let Some(err) = validate_columns(table, &parsed.and_cols) { - return err; - } - if let Some(err) = validate_columns(table, &extra_conds_columns(&parsed.ecs)) { - return err; + if let Some(err) = crate::where_dsl::validate_where_columns(&parsed.dsl_where, &table.value()) { + return err.into(); } let pool = &parsed.pool; let ty = &parsed.ty; - let val = &parsed.val; - let tid = &parsed.tid; - let table_str = table.value(); - let col_str = parsed.col.value(); - let cols = get_select_columns(table); let d = dialect(); - let in_base: usize = 2 - + parsed.and_cols.len() - + parsed.ecs.gt_vals.len() - + parsed.ecs.lt_vals.len() - + parsed.ecs.gte_vals.len() - + parsed.ecs.lte_vals.len(); - let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); - let numbered_lit = syn::LitBool::new(numbered, proc_macro2::Span::call_site()); - let ph_prefix_lit = syn::LitStr::new( - match d { - Dialect::Postgres => "$", - _ => "?", - }, - proc_macro2::Span::call_site(), - ); - let and_cols = &parsed.and_cols; - let and_vals = &parsed.and_vals; + let table_str = table.value(); + let table_lit = syn::LitStr::new(&table_str, table.span()); + let cols = get_select_columns(table); + let cols_lit = syn::LitStr::new(&cols, table.span()); + + let wr = crate::where_dsl::generate_where_runtime(&parsed.dsl_where, d); + let (local_stmts, bind_stmts) = emit_runtime_binds(&wr.binds); + let (tenant_sql, tenant_bind) = emit_tenant_code(&parsed.tid, d, None); + let sql_code = wr.sql_code; + + let order_fragment = match &parsed.order_by { + Some(ob) => { + let ob_str = format!(" ORDER BY {}", ob.value()); + let ob_lit = syn::LitStr::new(&ob_str, table.span()); + quote! { #ob_lit } + } + None => quote! { "" }, + }; let method_call = match &method { FindMethod::FetchOptional => quote! { fetch_optional(#pool).await }, @@ -747,262 +521,21 @@ fn expand_find(input: TokenStream, method: FindMethod) -> TokenStream { FindMethod::FetchAll => quote! { fetch_all(#pool).await }, }; - if parsed.tid.is_some() { - let col_lit = syn::LitStr::new(&col_str, table.span()); - let table_lit = syn::LitStr::new(&table_str, table.span()); - let cols_lit = syn::LitStr::new(&cols, table.span()); - - let and_col_lits: Vec = and_cols - .iter() - .enumerate() - .map(|(i, ac)| { - syn::LitStr::new( - &format!(" AND {} = {}", ac.value(), d.ph(2 + i)), - table.span(), - ) - }) - .collect(); - - let ecs_null_lits: Vec = parsed - .ecs - .null_cols - .iter() - .map(|c| syn::LitStr::new(&format!(" AND {} IS NULL", c.value()), table.span())) - .collect(); - - let mut all_bind_vals: Vec = and_vals.to_vec(); - all_bind_vals.extend(parsed.ecs.gt_vals.iter().cloned()); - all_bind_vals.extend(parsed.ecs.lt_vals.iter().cloned()); - all_bind_vals.extend(parsed.ecs.gte_vals.iter().cloned()); - all_bind_vals.extend(parsed.ecs.lte_vals.iter().cloned()); - - let ecs_cmp_lits: Vec = { - let mut idx: usize = 2 + and_cols.len(); - let mut lits = Vec::new(); - for _ in &parsed.ecs.gt_vals { - let ph = d.ph(idx); - idx += 1; - let col = &parsed.ecs.gt_cols[lits.len()]; - lits.push(syn::LitStr::new( - &format!(" AND {} > {}", col.value(), ph), - table.span(), - )); - } - let mut lits2 = Vec::new(); - for _ in &parsed.ecs.lt_vals { - let ph = d.ph(idx); - idx += 1; - let col = &parsed.ecs.lt_cols[lits2.len()]; - lits2.push(syn::LitStr::new( - &format!(" AND {} < {}", col.value(), ph), - table.span(), - )); - } - let mut lits3 = Vec::new(); - for _ in &parsed.ecs.gte_vals { - let ph = d.ph(idx); - idx += 1; - let col = &parsed.ecs.gte_cols[lits3.len()]; - lits3.push(syn::LitStr::new( - &format!(" AND {} >= {}", col.value(), ph), - table.span(), - )); - } - let mut lits4 = Vec::new(); - for _ in &parsed.ecs.lte_vals { - let ph = d.ph(idx); - idx += 1; - let col = &parsed.ecs.lte_cols[lits4.len()]; - lits4.push(syn::LitStr::new( - &format!(" AND {} <= {}", col.value(), ph), - table.span(), - )); - } - lits.into_iter() - .chain(lits2) - .chain(lits3) - .chain(lits4) - .collect() - }; - - let in_col_lits: Vec = parsed - .ecs - .in_cols - .iter() - .map(|c| syn::LitStr::new(&c.value(), table.span())) - .collect(); - let in_vals = &parsed.ecs.in_vals; - - let and_idents: Vec = (0..all_bind_vals.len()) - .map(|i| syn::Ident::new(&format!("__fa_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let order_by_fragment = match &parsed.order_by { - Some(ob) => { - let ob_str = ob.value(); - let ob_lit = syn::LitStr::new(&format!(" ORDER BY {}", ob_str), table.span()); - quote! { #ob_lit } - } - None => quote! { "" }, - }; - - let ph1_lit = syn::LitStr::new(&d.ph(1), table.span()); - let in_base_lit = syn::LitInt::new(&in_base.to_string(), proc_macro2::Span::call_site()); - let expanded = quote! { - { - let __fv = #val; - #(let #and_idents = #all_bind_vals;)* - let __ob: &str = #order_by_fragment; - let mut __and_sql = String::new(); - #(__and_sql.push_str(#and_col_lits);)* - #(__and_sql.push_str(#ecs_null_lits);)* - #(__and_sql.push_str(#ecs_cmp_lits);)* - let mut __ph_idx: usize = #in_base_lit; - #( - if !#in_vals.is_empty() { - let __in_ph: String = if #numbered_lit { - (0..#in_vals.len()).map(|i| format!(concat!(#ph_prefix_lit, "{}"), __ph_idx + i)).collect::>().join(",") - } else { - (0..#in_vals.len()).map(|_| "?").collect::>().join(",") - }; - __ph_idx += #in_vals.len(); - __and_sql.push_str(&format!(" AND {} IN ({})", #in_col_lits, __in_ph)); - } - )* - let __sql = match #tid { - Some(_tid) => { - let __tid_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - format!("SELECT {} FROM {} WHERE {} = {}{} AND tenant_id = {}{}", #cols_lit, #table_lit, #col_lit, #ph1_lit, __and_sql, __tid_ph, __ob) - } - None => format!("SELECT {} FROM {} WHERE {} = {}{}{}", #cols_lit, #table_lit, #col_lit, #ph1_lit, __and_sql, __ob), - }; - let mut _q = sqlx::query_as::<_, #ty>(&__sql).bind(__fv); - #(_q = _q.bind(#and_idents);)* - #( - for __iv in #in_vals { - _q = _q.bind(__iv); - } - )* - if let Some(_tid) = #tid { - _q = _q.bind(_tid); - } - _q.#method_call - } - }; - TokenStream::from(expanded) - } else { - let mut sql_str = format!( - "SELECT {} FROM {} WHERE {} = {}", - cols, - table_str, - col_str, - d.ph(1) - ); - let mut all_extra_vals: Vec = and_vals.to_vec(); - for (i, ac) in and_cols.iter().enumerate() { - sql_str.push_str(&format!(" AND {} = {}", ac.value(), d.ph(2 + i))); - } - for c in &parsed.ecs.null_cols { - sql_str.push_str(&format!(" AND {} IS NULL", c.value())); - } + let expanded = quote! { { - let mut idx = 2 + and_cols.len(); - for (c, v) in parsed.ecs.gt_cols.iter().zip(&parsed.ecs.gt_vals) { - sql_str.push_str(&format!(" AND {} > {}", c.value(), d.ph(idx))); - idx += 1; - all_extra_vals.push(v.clone()); - } - for (c, v) in parsed.ecs.lt_cols.iter().zip(&parsed.ecs.lt_vals) { - sql_str.push_str(&format!(" AND {} < {}", c.value(), d.ph(idx))); - idx += 1; - all_extra_vals.push(v.clone()); - } - for (c, v) in parsed.ecs.gte_cols.iter().zip(&parsed.ecs.gte_vals) { - sql_str.push_str(&format!(" AND {} >= {}", c.value(), d.ph(idx))); - idx += 1; - all_extra_vals.push(v.clone()); - } - for (c, v) in parsed.ecs.lte_cols.iter().zip(&parsed.ecs.lte_vals) { - sql_str.push_str(&format!(" AND {} <= {}", c.value(), d.ph(idx))); - idx += 1; - all_extra_vals.push(v.clone()); - } + #(#local_stmts)* + let mut __ph_idx: usize = 1; + let mut __where_sql = String::new(); + #sql_code + #tenant_sql + let __sql = format!("SELECT {} FROM {} WHERE {}{}{}", #cols_lit, #table_lit, __where_sql, __tenant_sql, #order_fragment); + let mut __q = sqlx::query_as::<_, #ty>(&__sql); + #(#bind_stmts)* + #tenant_bind + __q.#method_call } - if let Some(ref ob) = parsed.order_by { - sql_str.push_str(&format!(" ORDER BY {}", ob.value())); - } - - let has_in = !parsed.ecs.in_cols.is_empty(); - - if has_in { - let in_col_lits: Vec = parsed - .ecs - .in_cols - .iter() - .map(|c| syn::LitStr::new(&c.value(), table.span())) - .collect(); - let in_vals = &parsed.ecs.in_vals; - let order_str = match &parsed.order_by { - Some(ob) => format!(" ORDER BY {}", ob.value()), - None => String::new(), - }; - let sql_prefix_lit = syn::LitStr::new(&sql_str, table.span()); - let order_lit = syn::LitStr::new(&order_str, table.span()); - - let and_idents: Vec = (0..all_extra_vals.len()) - .map(|i| syn::Ident::new(&format!("__fa_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let in_base_lit = - syn::LitInt::new(&in_base.to_string(), proc_macro2::Span::call_site()); - let expanded = quote! { - { - #(let #and_idents = #all_extra_vals;)* - let mut __sql = #sql_prefix_lit.to_string(); - let mut __ph_idx: usize = #in_base_lit; - #( - if !#in_vals.is_empty() { - let __in_ph: String = if #numbered_lit { - (0..#in_vals.len()).map(|i| format!(concat!(#ph_prefix_lit, "{}"), __ph_idx + i)).collect::>().join(",") - } else { - (0..#in_vals.len()).map(|_| "?").collect::>().join(",") - }; - __ph_idx += #in_vals.len(); - __sql.push_str(&format!(" AND {} IN ({})", #in_col_lits, __in_ph)); - } - )* - __sql.push_str(#order_lit); - let mut _q = sqlx::query_as::<_, #ty>(&__sql).bind(#val); - #(_q = _q.bind(#and_idents);)* - #( - for __iv in #in_vals { - _q = _q.bind(__iv); - } - )* - _q.#method_call - } - }; - TokenStream::from(expanded) - } else { - let sql = syn::LitStr::new(&sql_str, table.span()); - - let and_idents: Vec = (0..all_extra_vals.len()) - .map(|i| syn::Ident::new(&format!("__fa_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let expanded = quote! { - { - #(let #and_idents = #all_extra_vals;)* - sqlx::query_as::<_, #ty>(#sql).bind(#val)#(.bind(#and_idents))*.#method_call - } - }; - TokenStream::from(expanded) - } - } + }; + TokenStream::from(expanded) } // ── crud_count! ───────────────────────────────────────── @@ -1012,287 +545,41 @@ pub fn crud_count(input: TokenStream) -> TokenStream { } fn expand_count(input: TokenStream) -> TokenStream { - let parsed = parse_macro_input!(input as DeleteInput); + let parsed = parse_macro_input!(input as WhereOnlyInput); let table = &parsed.table; - let col = &parsed.col; if let Some(err) = validate_table(table) { return err; } - if let Some(err) = validate_column(table, col) { - return err; - } - if let Some(err) = validate_columns(table, &parsed.and_cols) { - return err; - } - if let Some(err) = validate_columns(table, &extra_conds_columns(&parsed.ecs)) { - return err; + if let Some(err) = crate::where_dsl::validate_where_columns(&parsed.dsl_where, &table.value()) { + return err.into(); } let pool = &parsed.pool; - let val = &parsed.val; - let tid = &parsed.tid; - let table_str = table.value(); - let col_str = col.value(); - let and_cols = &parsed.and_cols; - let and_vals = &parsed.and_vals; - - let has_in = !parsed.ecs.in_cols.is_empty(); - let d = dialect(); - let mut ph_idx = 1usize; - let col_ph = d.ph(ph_idx); - ph_idx += 1; - let mut and_parts: Vec = and_cols - .iter() - .map(|ac| { - let ph = d.ph(ph_idx); - ph_idx += 1; - format!("AND {} = {}", ac.value(), ph) - }) - .collect(); - let (ecs_parts, ecs_vals) = build_extra_conds_sql(&parsed.ecs, d, &mut ph_idx); - and_parts.extend(ecs_parts); - let all_extra_vals: Vec = and_vals.iter().chain(ecs_vals.iter()).cloned().collect(); - let and_str = and_parts.join(" "); + let table_str = table.value(); + let table_lit = syn::LitStr::new(&table_str, table.span()); - let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); - let numbered_lit = syn::LitBool::new(numbered, proc_macro2::Span::call_site()); - let ph_prefix_lit = syn::LitStr::new( - match d { - Dialect::Postgres => "$", - _ => "?", - }, - proc_macro2::Span::call_site(), - ); - let in_base: usize = ph_idx; + let wr = crate::where_dsl::generate_where_runtime(&parsed.dsl_where, d); + let (local_stmts, bind_stmts) = emit_runtime_binds(&wr.binds); + let (tenant_sql, tenant_bind) = emit_tenant_code(&parsed.tid, d, None); + let sql_code = wr.sql_code; - let has_extra = !and_cols.is_empty() || !parsed.ecs.is_empty(); - - let in_col_lits: Vec = parsed - .ecs - .in_cols - .iter() - .map(|c| syn::LitStr::new(&c.value(), table.span())) - .collect(); - let in_vals = &parsed.ecs.in_vals; - - if parsed.tid.is_some() { - if has_in { - let sql_prefix = format!( - "SELECT COUNT(*) FROM {} WHERE {} = {}{}", - table_str, - col_str, - col_ph, - if and_str.is_empty() { - String::new() - } else { - format!(" {}", and_str) - }, - ); - let sql_prefix_lit = syn::LitStr::new(&sql_prefix, table.span()); - let in_base_lit = - syn::LitInt::new(&in_base.to_string(), proc_macro2::Span::call_site()); - - let extra_idents: Vec = (0..all_extra_vals.len()) - .map(|i| syn::Ident::new(&format!("__cnt_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let and_bind = if all_extra_vals.is_empty() { - quote! {} - } else { - quote! { #(_q = _q.bind(#extra_idents);)* } - }; - - let expanded = quote! { - { - #(let #extra_idents = #all_extra_vals;)* - let mut __sql = #sql_prefix_lit.to_string(); - let mut __ph_idx: usize = #in_base_lit; - #( - if !#in_vals.is_empty() { - let __in_ph: String = if #numbered_lit { - (0..#in_vals.len()).map(|i| format!(concat!(#ph_prefix_lit, "{}"), __ph_idx + i)).collect::>().join(",") - } else { - (0..#in_vals.len()).map(|_| "?").collect::>().join(",") - }; - __ph_idx += #in_vals.len(); - __sql.push_str(&format!(" AND {} IN ({})", #in_col_lits, __in_ph)); - } - )* - let __tenant_sql: String = match #tid { - Some(_) => { - let __tid_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - format!(" AND tenant_id = {}", __tid_ph) - } - None => String::new(), - }; - __sql.push_str(&__tenant_sql); - let mut _q = sqlx::query_scalar::<_, i64>(&__sql).bind(#val); - #and_bind - #( - for __iv in #in_vals { - _q = _q.bind(__iv); - } - )* - if let Some(_tid) = #tid { - _q = _q.bind(_tid); - } - _q.fetch_one(#pool).await - } - }; - TokenStream::from(expanded) - } else { - let tid_ph = d.ph(ph_idx); - - let expanded = if !has_extra { - let sql_with = syn::LitStr::new( - &format!( - "SELECT COUNT(*) FROM {} WHERE {} = {} AND tenant_id = {}", - table_str, col_str, col_ph, tid_ph - ), - table.span(), - ); - let sql_without = syn::LitStr::new( - &format!( - "SELECT COUNT(*) FROM {} WHERE {} = {}", - table_str, col_str, col_ph - ), - table.span(), - ); - quote! { - { - match #tid { - Some(_tid) => sqlx::query_scalar::<_, i64>(#sql_with).bind(#val).bind(_tid).fetch_one(#pool).await, - None => sqlx::query_scalar::< _, i64>(#sql_without).bind(#val).fetch_one(#pool).await, - } - } - } - } else { - let extra_idents: Vec = (0..all_extra_vals.len()) - .map(|i| { - syn::Ident::new(&format!("__cnt_{}", i), proc_macro2::Span::call_site()) - }) - .collect(); - let sql_with = syn::LitStr::new( - &format!( - "SELECT COUNT(*) FROM {} WHERE {} = {} {} AND tenant_id = {}", - table_str, col_str, col_ph, and_str, tid_ph - ), - table.span(), - ); - let sql_without = syn::LitStr::new( - &format!( - "SELECT COUNT(*) FROM {} WHERE {} = {} {}", - table_str, col_str, col_ph, and_str - ), - table.span(), - ); - quote! { - { - #(let #extra_idents = #all_extra_vals;)* - match #tid { - Some(_tid) => sqlx::query_scalar::< _, i64>(#sql_with).bind(#val)#(.bind(#extra_idents))*.bind(_tid).fetch_one(#pool).await, - None => sqlx::query_scalar::< _, i64>(#sql_without).bind(#val)#(.bind(#extra_idents))*.fetch_one(#pool).await, - } - } - } - }; - TokenStream::from(expanded) + let expanded = quote! { + { + #(#local_stmts)* + let mut __ph_idx: usize = 1; + let mut __where_sql = String::new(); + #sql_code + #tenant_sql + let __sql = format!("SELECT COUNT(*) FROM {} WHERE {}{}", #table_lit, __where_sql, __tenant_sql); + let mut __q = sqlx::query_scalar::<_, i64>(&__sql); + #(#bind_stmts)* + #tenant_bind + __q.fetch_one(#pool).await } - } else { - if has_in { - let sql_prefix = format!( - "SELECT COUNT(*) FROM {} WHERE {} = {}{}", - table_str, - col_str, - col_ph, - if and_str.is_empty() { - String::new() - } else { - format!(" {}", and_str) - } - ); - let sql_prefix_lit = syn::LitStr::new(&sql_prefix, table.span()); - let in_base_lit = - syn::LitInt::new(&in_base.to_string(), proc_macro2::Span::call_site()); - - let extra_idents: Vec = (0..all_extra_vals.len()) - .map(|i| syn::Ident::new(&format!("__cnt_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let and_bind = if all_extra_vals.is_empty() { - quote! {} - } else { - quote! { #(_q = _q.bind(#extra_idents);)* } - }; - - let expanded = quote! { - { - #(let #extra_idents = #all_extra_vals;)* - let mut __sql = #sql_prefix_lit.to_string(); - let mut __ph_idx: usize = #in_base_lit; - #( - if !#in_vals.is_empty() { - let __in_ph: String = if #numbered_lit { - (0..#in_vals.len()).map(|i| format!(concat!(#ph_prefix_lit, "{}"), __ph_idx + i)).collect::>().join(",") - } else { - (0..#in_vals.len()).map(|_| "?").collect::>().join(",") - }; - __ph_idx += #in_vals.len(); - __sql.push_str(&format!(" AND {} IN ({})", #in_col_lits, __in_ph)); - } - )* - let mut _q = sqlx::query_scalar::<_, i64>(&__sql).bind(#val); - #and_bind - #( - for __iv in #in_vals { - _q = _q.bind(__iv); - } - )* - _q.fetch_one(#pool).await - } - }; - TokenStream::from(expanded) - } else { - let expanded = if !has_extra { - let sql_lit = syn::LitStr::new( - &format!( - "SELECT COUNT(*) FROM {} WHERE {} = {}", - table_str, col_str, col_ph - ), - table.span(), - ); - quote! { - sqlx::query_scalar::< _, i64>(#sql_lit).bind(#val).fetch_one(#pool).await - } - } else { - let extra_idents: Vec = (0..all_extra_vals.len()) - .map(|i| { - syn::Ident::new(&format!("__cnt_{}", i), proc_macro2::Span::call_site()) - }) - .collect(); - let sql_lit = syn::LitStr::new( - &format!( - "SELECT COUNT(*) FROM {} WHERE {} = {} {}", - table_str, col_str, col_ph, and_str - ), - table.span(), - ); - quote! { - { - #(let #extra_idents = #all_extra_vals;)* - sqlx::query_scalar::< _, i64>(#sql_lit).bind(#val)#(.bind(#extra_idents))*.fetch_one(#pool).await - } - } - }; - TokenStream::from(expanded) - } - } + }; + TokenStream::from(expanded) } // ── crud_list! ──────────────────────────────────────────────────────── @@ -1360,9 +647,6 @@ fn expand_update(input: TokenStream) -> TokenStream { if let Some(err) = validate_table(table) { return err; } - if let Some(err) = validate_column(table, &parsed.pk_col) { - return err; - } if let Some(err) = validate_columns(table, &parsed.bind_cols) { return err; } @@ -1374,21 +658,18 @@ fn expand_update(input: TokenStream) -> TokenStream { if let Some(err) = validate_columns(table, &parsed.opt_cols) { return err; } - if let Some(err) = validate_columns(table, &parsed.and_cols) { - return err; + if let Some(err) = crate::where_dsl::validate_where_columns(&parsed.dsl_where, &table.value()) { + return err.into(); } let pool = &parsed.pool; let table_str = table.value(); + let table_lit = syn::LitStr::new(&table_str, table.span()); let bind_cols = &parsed.bind_cols; let bind_vals = &parsed.bind_vals; let opt_cols = &parsed.opt_cols; let opt_vals = &parsed.opt_vals; let raw_pairs = &parsed.raw_pairs; - let pk_col = &parsed.pk_col; - let pk_val = &parsed.pk_val; - let and_cols = &parsed.and_cols; - let and_vals = &parsed.and_vals; let has_optional = !opt_cols.is_empty(); @@ -1400,409 +681,6 @@ fn expand_update(input: TokenStream) -> TokenStream { let opt_idents: Vec = (0..opt_vals.len()) .map(|i| syn::Ident::new(&format!("__ov_{}", i), proc_macro2::Span::call_site())) .collect(); - let and_idents: Vec = (0..and_vals.len()) - .map(|i| syn::Ident::new(&format!("__ua_{}", i), proc_macro2::Span::call_site())) - .collect(); - - if has_optional { - let d = dialect(); - let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); - let numbered_lit = syn::LitBool::new(numbered, proc_macro2::Span::call_site()); - let ph_prefix_lit = syn::LitStr::new( - match d { - Dialect::Postgres => "$", - _ => "?", - }, - proc_macro2::Span::call_site(), - ); - let dyn_start_lit = syn::LitInt::new( - &(bind_cols.len() + 1).to_string(), - proc_macro2::Span::call_site(), - ); - - let table_lit = syn::LitStr::new(&table_str, table.span()); - let pk_col_lit = syn::LitStr::new(&pk_col.value(), table.span()); - - let bind_set_lits: Vec = bind_col_strs - .iter() - .enumerate() - .map(|(i, c)| syn::LitStr::new(&format!("{} = {}", c, d.ph(1 + i)), table.span())) - .collect(); - - let raw_col_names: Vec = raw_pairs.iter().map(|(rc, _)| rc.clone()).collect(); - let raw_exprs: Vec<&syn::Expr> = raw_pairs.iter().map(|(_, rv)| rv).collect(); - - let opt_col_lits: Vec = opt_cols - .iter() - .map(|c| syn::LitStr::new(&c.value(), c.span())) - .collect(); - - let and_col_names: Vec = and_cols - .iter() - .map(|ac| syn::LitStr::new(&ac.value(), ac.span())) - .collect(); - - let tid = &parsed.tid; - - let opt_bind_idents: Vec = (0..opt_vals.len()) - .map(|i| syn::Ident::new(&format!("__obv_{}", i), proc_macro2::Span::call_site())) - .collect(); - let and_ph_idents: Vec = (0..and_vals.len()) - .map(|i| syn::Ident::new(&format!("__aph_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let expanded = if parsed.tid.is_some() { - quote! { - { - #(let #val_idents = #bind_vals;)* - #(let #opt_idents = &(#opt_vals);)* - #(let #and_idents = #and_vals;)* - let __pkv = #pk_val; - let mut __sets: Vec = Vec::new(); - #(__sets.push(#bind_set_lits.to_string());)* - #(__sets.push(format!("{} = {}", #raw_col_names, #raw_exprs));)* - let mut __ph_idx: usize = #dyn_start_lit; - #( - if #opt_idents.is_some() { - let __ph: String = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - __sets.push(format!("{} = {}", #opt_col_lits, __ph)); - } - )* - let __pk_ph: String = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - let mut __and_sql = String::new(); - #( - let #and_ph_idents: String = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - __and_sql.push_str(&format!(" AND {} = {}", #and_col_names, #and_ph_idents)); - )* - let _sql = match #tid { - Some(_tid) => { - let __tid_ph: String = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - format!( - "UPDATE {} SET {} WHERE {} = {}{} AND tenant_id = {}", - #table_lit, __sets.join(", "), #pk_col_lit, __pk_ph, __and_sql, __tid_ph - ) - } - None => format!( - "UPDATE {} SET {} WHERE {} = {}{}", - #table_lit, __sets.join(", "), #pk_col_lit, __pk_ph, __and_sql - ), - }; - let mut _q = sqlx::query(&_sql); - #(_q = _q.bind(#val_idents);)* - #( - if let Some(#opt_bind_idents) = #opt_idents { - _q = _q.bind(#opt_bind_idents); - } - )* - _q = _q.bind(__pkv); - #(_q = _q.bind(#and_idents);)* - if let Some(_tid) = #tid { - _q = _q.bind(_tid); - } - _q.execute(#pool).await - } - } - } else { - quote! { - { - #(let #val_idents = #bind_vals;)* - #(let #opt_idents = &(#opt_vals);)* - #(let #and_idents = #and_vals;)* - let __pkv = #pk_val; - let mut __sets: Vec = Vec::new(); - #(__sets.push(#bind_set_lits.to_string());)* - #(__sets.push(format!("{} = {}", #raw_col_names, #raw_exprs));)* - let mut __ph_idx: usize = #dyn_start_lit; - #( - if #opt_idents.is_some() { - let __ph: String = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - __sets.push(format!("{} = {}", #opt_col_lits, __ph)); - } - )* - let __pk_ph: String = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - let mut __and_sql = String::new(); - #( - let #and_ph_idents: String = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - __and_sql.push_str(&format!(" AND {} = {}", #and_col_names, #and_ph_idents)); - )* - let _sql = format!( - "UPDATE {} SET {} WHERE {} = {}{}", - #table_lit, __sets.join(", "), #pk_col_lit, __pk_ph, __and_sql - ); - let mut _q = sqlx::query(&_sql); - #(_q = _q.bind(#val_idents);)* - #( - if let Some(#opt_bind_idents) = #opt_idents { - _q = _q.bind(#opt_bind_idents); - } - )* - _q = _q.bind(__pkv); - #(_q = _q.bind(#and_idents);)* - _q.execute(#pool).await - } - } - }; - return TokenStream::from(expanded); - } - - // ── Static path (no optional columns — original logic) ── - let d = dialect(); - let mut ph_idx = 1usize; - let mut set_parts: Vec = Vec::new(); - for _ in &bind_col_strs { - set_parts.push(d.ph(ph_idx)); - ph_idx += 1; - } - let raw_col_names: Vec = raw_pairs.iter().map(|(rc, _)| rc.clone()).collect(); - let raw_exprs: Vec<&syn::Expr> = raw_pairs.iter().map(|(_, rv)| rv).collect(); - - let pk_ph = d.ph(ph_idx); - ph_idx += 1; - - let and_parts: Vec = and_cols - .iter() - .map(|ac| { - let ph = d.ph(ph_idx); - ph_idx += 1; - format!("AND {} = {}", ac.value(), ph) - }) - .collect(); - - let set_str_prefix = bind_col_strs - .iter() - .zip(set_parts.iter()) - .map(|(c, p)| format!("{} = {}", c, p)) - .collect::>() - .join(", "); - - let and_str = and_parts.join(""); - - if parsed.tid.is_some() { - let tid = &parsed.tid; - let tenant_ph = d.ph(ph_idx); - - // Store SQL fragments as literals for the generated format!() call - let table_lit = syn::LitStr::new(&table_str, table.span()); - let set_prefix_lit = syn::LitStr::new(&set_str_prefix, table.span()); - let pk_col_lit = syn::LitStr::new(&pk_col.value(), table.span()); - let pk_ph_lit = syn::LitStr::new(&pk_ph, table.span()); - let and_lit = syn::LitStr::new(&and_str, table.span()); - let tenant_ph_lit = syn::LitStr::new(&tenant_ph, table.span()); - - let val_idents: Vec = (0..bind_vals.len()) - .map(|i| syn::Ident::new(&format!("__uv_{}", i), proc_macro2::Span::call_site())) - .collect(); - let and_idents: Vec = (0..and_vals.len()) - .map(|i| syn::Ident::new(&format!("__ua_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let has_raw = !raw_pairs.is_empty(); - let expanded = if has_raw { - let sep = if set_str_prefix.is_empty() { "" } else { ", " }; - let raw_parts: Vec = raw_col_names - .iter() - .zip(raw_exprs.iter()) - .enumerate() - .map(|(i, (c, _))| { - if i == 0 { - format!("{} = {{}}", c.value()) - } else { - format!(", {} = {{}}", c.value()) - } - }) - .collect(); - let raw_fmt = format!("{}{}", sep, raw_parts.join("")); - let raw_fmt_lit = syn::LitStr::new(&raw_fmt, table.span()); - quote! { - { - #(let #val_idents = #bind_vals;)* - #(let #and_idents = #and_vals;)* - let __pkv = #pk_val; - let __raw_suffix = format!(#raw_fmt_lit #(, #raw_exprs)*); - let __set = format!("{}{}", #set_prefix_lit, __raw_suffix); - let _sql = match #tid { - Some(_tid) => format!( - "UPDATE {} SET {} WHERE {} = {}{} AND tenant_id = {}", - #table_lit, __set, #pk_col_lit, #pk_ph_lit, #and_lit, #tenant_ph_lit - ), - None => format!( - "UPDATE {} SET {} WHERE {} = {}{}", - #table_lit, __set, #pk_col_lit, #pk_ph_lit, #and_lit - ), - }; - let mut _q = sqlx::query(&_sql)#(.bind(#val_idents))*; - _q = _q.bind(__pkv); - #(_q = _q.bind(#and_idents);)* - if let Some(_tid) = #tid { - _q = _q.bind(_tid); - } - _q.execute(#pool).await - } - } - } else { - quote! { - { - #(let #val_idents = #bind_vals;)* - #(let #and_idents = #and_vals;)* - let __pkv = #pk_val; - let _sql = match #tid { - Some(_tid) => format!( - "UPDATE {} SET {} WHERE {} = {}{} AND tenant_id = {}", - #table_lit, #set_prefix_lit, #pk_col_lit, #pk_ph_lit, #and_lit, #tenant_ph_lit - ), - None => format!( - "UPDATE {} SET {} WHERE {} = {}{}", - #table_lit, #set_prefix_lit, #pk_col_lit, #pk_ph_lit, #and_lit - ), - }; - let mut _q = sqlx::query(&_sql)#(.bind(#val_idents))*; - _q = _q.bind(__pkv); - #(_q = _q.bind(#and_idents);)* - if let Some(_tid) = #tid { - _q = _q.bind(_tid); - } - _q.execute(#pool).await - } - } - }; - TokenStream::from(expanded) - } else { - let has_raw = !raw_pairs.is_empty(); - let sql = if has_raw { - syn::LitStr::new( - &format!( - "UPDATE {} SET {{}}{{}} WHERE {} = {}{}", - table_str, - pk_col.value(), - pk_ph, - and_str - ), - table.span(), - ) - } else { - syn::LitStr::new( - &format!( - "UPDATE {} SET {} WHERE {} = {}{}", - table_str, - set_str_prefix, - pk_col.value(), - pk_ph, - and_str - ), - table.span(), - ) - }; - // E0716 fix: pre-bind all values to named locals - let val_idents: Vec = (0..bind_vals.len()) - .map(|i| syn::Ident::new(&format!("__uv_{}", i), proc_macro2::Span::call_site())) - .collect(); - let and_idents: Vec = (0..and_vals.len()) - .map(|i| syn::Ident::new(&format!("__ua_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let expanded = if has_raw { - let sep = if set_str_prefix.is_empty() { "" } else { ", " }; - let raw_parts: Vec = raw_col_names - .iter() - .zip(raw_exprs.iter()) - .enumerate() - .map(|(i, (c, _))| { - if i == 0 { - format!("{} = {{}}", c.value()) - } else { - format!(", {} = {{}}", c.value()) - } - }) - .collect(); - let raw_fmt = format!("{}{}", sep, raw_parts.join("")); - let raw_fmt_lit = syn::LitStr::new(&raw_fmt, table.span()); - let set_prefix_lit = syn::LitStr::new(&set_str_prefix, table.span()); - quote! { - { - #(let #val_idents = #bind_vals;)* - #(let #and_idents = #and_vals;)* - let __pkv = #pk_val; - let __raw_suffix = format!(#raw_fmt_lit #(, #raw_exprs)*); - let __set = format!("{}{}", #set_prefix_lit, __raw_suffix); - let _sql = format!(#sql, __set, ""); - sqlx::query(&_sql)#(.bind(#val_idents))*.bind(__pkv)#(.bind(#and_idents))*.execute(#pool).await - } - } - } else { - quote! { - { - #(let #val_idents = #bind_vals;)* - #(let #and_idents = #and_vals;)* - let __pkv = #pk_val; - sqlx::query(#sql)#(.bind(#val_idents))*.bind(__pkv)#(.bind(#and_idents))*.execute(#pool).await - } - } - }; - TokenStream::from(expanded) - } -} - -// ── crud_query_paged! ────────────────────────────────────────────── - -pub fn crud_query_paged(input: TokenStream) -> TokenStream { - let parsed = parse_macro_input!(input as QueryPagedInput); - let pool = &parsed.pool; - let ty = &parsed.ty; - let data_sql = &parsed.data_sql; - let count_sql = &parsed.count_sql; - let binds = &parsed.binds; - let tid = &parsed.tid; - let page = &parsed.page; - let page_size = &parsed.page_size; - let where_cols = &parsed.where_cols; - let where_vals = &parsed.where_vals; - - let bind_idents: Vec = (0..binds.len()) - .map(|i| syn::Ident::new(&format!("__pb_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let where_bind_idents: Vec = (0..where_vals.len()) - .map(|i| syn::Ident::new(&format!("__wp_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let where_col_strs: Vec = where_cols - .iter() - .map(|c| syn::LitStr::new(&c.value(), c.span())) - .collect(); let d = dialect(); let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); @@ -1815,18 +693,249 @@ pub fn crud_query_paged(input: TokenStream) -> TokenStream { proc_macro2::Span::call_site(), ); - let base_idx = binds.len() + 1; - let base_idx_lit = syn::LitInt::new(&base_idx.to_string(), proc_macro2::Span::call_site()); + let bind_set_lits: Vec = bind_col_strs + .iter() + .enumerate() + .map(|(i, c)| { + let ph = d.ph(1 + i); + syn::LitStr::new(&format!("{} = {}", c, ph), table.span()) + }) + .collect(); + + let raw_col_names: Vec = raw_pairs.iter().map(|(rc, _)| rc.clone()).collect(); + let raw_exprs: Vec<&syn::Expr> = raw_pairs.iter().map(|(_, rv)| rv).collect(); + + let opt_col_lits: Vec = opt_cols + .iter() + .map(|c| syn::LitStr::new(&c.value(), c.span())) + .collect(); + + let opt_bind_idents: Vec = (0..opt_vals.len()) + .map(|i| syn::Ident::new(&format!("__obv_{}", i), proc_macro2::Span::call_site())) + .collect(); + + // SET clause building + let dyn_start = bind_cols.len() + 1; + let dyn_start_lit = syn::LitInt::new(&dyn_start.to_string(), proc_macro2::Span::call_site()); + + // WHERE clause via runtime codegen + let wr = crate::where_dsl::generate_where_runtime(&parsed.dsl_where, d); + let (where_local_stmts, where_bind_stmts) = emit_runtime_binds(&wr.binds); + let (tenant_sql, tenant_bind) = emit_tenant_code(&parsed.tid, d, None); + let where_sql_code = wr.sql_code; + + let set_code = if has_optional { + quote! { + let mut __sets: Vec = Vec::new(); + #(__sets.push(#bind_set_lits.to_string());)* + #(__sets.push(format!("{} = {}", #raw_col_names, #raw_exprs));)* + let mut __ph_idx: usize = #dyn_start_lit; + #( + if #opt_idents.is_some() { + let __ph: String = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + __ph_idx += 1; + __sets.push(format!("{} = {}", #opt_col_lits, __ph)); + } + )* + } + } else { + let static_set: String = { + let parts: Vec = bind_col_strs + .iter() + .enumerate() + .map(|(i, c)| format!("{} = {}", c, d.ph(1 + i))) + .collect(); + parts.join(", ") + }; + let has_raw = !raw_pairs.is_empty(); + if has_raw { + let sep = if static_set.is_empty() { "" } else { ", " }; + let raw_parts: Vec = raw_pairs + .iter() + .enumerate() + .map(|(i, (c, _))| { + if i == 0 { + format!("{} = {{}}", c.value()) + } else { + format!(", {} = {{}}", c.value()) + } + }) + .collect(); + let raw_fmt = format!("{}{}", sep, raw_parts.join("")); + let raw_fmt_lit = syn::LitStr::new(&raw_fmt, table.span()); + let static_set_lit = syn::LitStr::new(&static_set, table.span()); + quote! { + let __raw_suffix = format!(#raw_fmt_lit #(, #raw_exprs)*); + let __set_str = format!("{}{}", #static_set_lit, __raw_suffix); + } + } else { + let static_set_lit = syn::LitStr::new(&static_set, table.span()); + quote! { + let __set_str = #static_set_lit.to_string(); + } + } + }; + + let set_join = if has_optional { + quote! { __sets.join(", ") } + } else { + quote! { __set_str } + }; + + let bind_code = if has_optional { + quote! { + #(let #val_idents = #bind_vals;)* + #(let #opt_idents = &(#opt_vals);)* + } + } else { + quote! { + #(let #val_idents = #bind_vals;)* + } + }; + + let opt_bind_code = if has_optional { + quote! { + #( + if let Some(#opt_bind_idents) = #opt_idents { + __q = __q.bind(#opt_bind_idents); + } + )* + } + } else { + quote! {} + }; + + let where_ph_start_stmt = if has_optional { + quote! {} // __ph_idx is set by set_code and carries over + } else { + let start = 1 + bind_vals.len(); + let start_lit = syn::LitInt::new(&start.to_string(), proc_macro2::Span::call_site()); + quote! { let mut __ph_idx: usize = #start_lit; } + }; let expanded = quote! { { - let __page_size = #page_size; - let __offset = (#page - 1).max(0) * __page_size; - #(let #bind_idents = #binds;)* - #(let #where_bind_idents = #where_vals;)* - let mut __ph_idx: usize = #base_idx_lit; - let __tenant_sql = match #tid { - Some(_tid) => { + #bind_code + #(#where_local_stmts)* + #set_code + let mut __where_sql = String::new(); + #where_ph_start_stmt + #where_sql_code + #tenant_sql + let __sql = format!("UPDATE {} SET {} WHERE {}{}", #table_lit, #set_join, __where_sql, __tenant_sql); + let mut __q = sqlx::query(&__sql); + #(__q = __q.bind(#val_idents);)* + #opt_bind_code + #(#where_bind_stmts)* + #tenant_bind + __q.execute(#pool).await + } + }; + TokenStream::from(expanded) +} + +// ── crud_query_paged! ────────────────────────────────────────────── + +pub fn crud_query_paged(input: TokenStream) -> TokenStream { + let parsed = parse_macro_input!(input as QueryPagedInput); + expand_query_paged_dsl(&parsed) +} + +fn expand_query_paged_dsl(parsed: &QueryPagedInput) -> TokenStream { + use crate::where_dsl::WhereCodegen; + + let pool = &parsed.pool; + let ty = &parsed.ty; + let table = &parsed.table; + let tid = &parsed.tid; + let page = &parsed.page; + let page_size = &parsed.page_size; + + if let Some(err) = validate_table(table) { + return err; + } + + let d = dialect(); + let cols = get_select_columns(table); + let table_str = table.value(); + + let where_sql: String; + let mut dsl_bind_count: usize = 0; + let dsl_bind_exprs: Vec; + + if let Some(ref dsl_where) = parsed.dsl_where { + if let Some(err) = crate::where_dsl::validate_where_columns(dsl_where, &table_str) { + return err.into(); + } + let mut cg = WhereCodegen::new(d, 1); + cg.generate(dsl_where); + where_sql = format!(" WHERE {}", cg.sql); + dsl_bind_count = cg.next_idx - 1; + dsl_bind_exprs = cg.binds; + } else { + where_sql = String::new(); + dsl_bind_exprs = Vec::new(); + } + + let order_str = parsed + .order_by + .as_ref() + .map(|o| format!(" ORDER BY {}", o.value())) + .unwrap_or_default(); + let order_str_lit = syn::LitStr::new(&order_str, table.span()); + + let needs_where_base = where_sql.is_empty() && (tid.is_some() || !parsed.where_cols.is_empty()); + let where_base = if needs_where_base { " WHERE 1=1" } else { "" }; + + let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); + let numbered_lit = syn::LitBool::new(numbered, proc_macro2::Span::call_site()); + let ph_prefix_lit = syn::LitStr::new( + match d { + Dialect::Postgres => "$", + _ => "?", + }, + proc_macro2::Span::call_site(), + ); + + let dsl_bind_idents: Vec = (0..dsl_bind_exprs.len()) + .map(|i| syn::Ident::new(&format!("__db_{}", i), proc_macro2::Span::call_site())) + .collect(); + + let where_col_strs: Vec = parsed + .where_cols + .iter() + .map(|c| syn::LitStr::new(&c.value(), c.span())) + .collect(); + let where_bind_idents: Vec = (0..parsed.where_vals.len()) + .map(|i| syn::Ident::new(&format!("__wp_{}", i), proc_macro2::Span::call_site())) + .collect(); + let where_vals = &parsed.where_vals; + + let data_sql_str = format!( + "SELECT {} FROM {}{}{}", + cols, table_str, where_sql, where_base + ); + let count_sql_str = format!( + "SELECT COUNT(*) FROM {}{}{}", + table_str, where_sql, where_base + ); + + let data_sql_lit = syn::LitStr::new(&data_sql_str, table.span()); + let count_sql_lit = syn::LitStr::new(&count_sql_str, table.span()); + + let (tenant_sql_block, tenant_bind_data, _tenant_bind_count) = if let Some(tid_expr) = tid { + let base_after_dsl = dsl_bind_count + 1; + let base_lit = + syn::LitInt::new(&base_after_dsl.to_string(), proc_macro2::Span::call_site()); + let block: proc_macro2::TokenStream = quote! { + let mut __ph_idx: usize = #base_lit; + let __tenant_val = #tid_expr; + let __tenant_sql = match __tenant_val { + Some(_) => { let __tph = if #numbered_lit { format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) } else { @@ -1837,7 +946,22 @@ pub fn crud_query_paged(input: TokenStream) -> TokenStream { } None => String::new(), }; - let mut __where_sql = String::new(); + }; + (block, Some(tid_expr.clone()), 1) + } else { + let base_lit = syn::LitInt::new( + &(dsl_bind_count + 1).to_string(), + proc_macro2::Span::call_site(), + ); + let block: proc_macro2::TokenStream = quote! { + let mut __ph_idx: usize = #base_lit; + let __tenant_sql = String::new(); + }; + (block, None, 0) + }; + + let where_bind_block: proc_macro2::TokenStream = if !parsed.where_cols.is_empty() { + quote! { #( if let Some(ref __wv) = #where_bind_idents { let __wph = if #numbered_lit { @@ -1849,36 +973,50 @@ pub fn crud_query_paged(input: TokenStream) -> TokenStream { __where_sql.push_str(&format!(" AND {} = {}", #where_col_strs, __wph)); } )* - let __limit_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - let __offset_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - let __data_sql_raw = #data_sql.replace("{tenant}", &__tenant_sql); - let mut __data_sql = if let Some(__pos) = __data_sql_raw.find("ORDER BY") { - let mut __s = String::with_capacity(__data_sql_raw.len() + __where_sql.len() + 20); - __s.push_str(&__data_sql_raw[..__pos]); - __s.push_str(&__where_sql); - __s.push_str(&__data_sql_raw[__pos..]); - __s - } else { - let mut __s = __data_sql_raw; - __s.push_str(&__where_sql); - __s - }; + } + } else { + quote! {} + }; + + let limit_offset_block: proc_macro2::TokenStream = quote! { + let __limit_ph = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + __ph_idx += 1; + let __offset_ph = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + }; + + let (tenant_bind_dq, tenant_bind_cq): (proc_macro2::TokenStream, proc_macro2::TokenStream) = + if tenant_bind_data.is_some() { + ( + quote! { if let Some(ref _tid) = __tenant_val { __dq = __dq.bind(_tid); } }, + quote! { if let Some(ref _tid) = __tenant_val { __cq = __cq.bind(_tid); } }, + ) + } else { + (quote! {}, quote! {}) + }; + + let expanded = quote! { + { + let __page_size = #page_size; + let __offset = (#page - 1).max(0) * __page_size; + #(let #dsl_bind_idents = #dsl_bind_exprs;)* + #(let #where_bind_idents = #where_vals;)* + #tenant_sql_block + let mut __where_sql = String::new(); + #where_bind_block + #limit_offset_block + let mut __data_sql = format!("{}{}{}{}", #data_sql_lit, __tenant_sql, __where_sql, #order_str_lit); __data_sql.push_str(&format!(" LIMIT {} OFFSET {}", __limit_ph, __offset_ph)); - let mut __count_sql = #count_sql.replace("{tenant}", &__tenant_sql); - __count_sql.push_str(&__where_sql); - let mut __dq = sqlx::query_as::<_, #ty>(&__data_sql)#(.bind(#bind_idents))*; - if let Some(_tid) = #tid { - __dq = __dq.bind(_tid); - } + let mut __count_sql = format!("{}{}{}", #count_sql_lit, __tenant_sql, __where_sql); + let mut __dq = sqlx::query_as::<_, #ty>(&__data_sql)#(.bind(#dsl_bind_idents))*; + #tenant_bind_dq #( if let Some(ref __wv) = #where_bind_idents { __dq = __dq.bind(__wv); @@ -1886,10 +1024,8 @@ pub fn crud_query_paged(input: TokenStream) -> TokenStream { )* __dq = __dq.bind(__page_size).bind(__offset); let __data = __dq.fetch_all(#pool).await?; - let mut __cq = sqlx::query_scalar::<_, i64>(&__count_sql)#(.bind(#bind_idents))*; - if let Some(_tid) = #tid { - __cq = __cq.bind(_tid); - } + let mut __cq = sqlx::query_scalar::<_, i64>(&__count_sql)#(.bind(#dsl_bind_idents))*; + #tenant_bind_cq #( if let Some(ref __wv) = #where_bind_idents { __cq = __cq.bind(__wv); @@ -1909,135 +1045,41 @@ pub fn crud_exists(input: TokenStream) -> TokenStream { } fn expand_exists(input: TokenStream) -> TokenStream { - let parsed = parse_macro_input!(input as DeleteInput); + let parsed = parse_macro_input!(input as WhereOnlyInput); let table = &parsed.table; - let col = &parsed.col; if let Some(err) = validate_table(table) { return err; } - if let Some(err) = validate_column(table, col) { - return err; - } - if let Some(err) = validate_columns(table, &parsed.and_cols) { - return err; - } - if let Some(err) = validate_columns(table, &extra_conds_columns(&parsed.ecs)) { - return err; + if let Some(err) = crate::where_dsl::validate_where_columns(&parsed.dsl_where, &table.value()) { + return err.into(); } let pool = &parsed.pool; - let val = &parsed.val; - let tid = &parsed.tid; - let table_str = table.value(); - let col_str = col.value(); - let and_cols = &parsed.and_cols; - let and_vals = &parsed.and_vals; - let d = dialect(); - let mut ph_idx = 1usize; - let col_ph = d.ph(ph_idx); - ph_idx += 1; - let mut and_parts: Vec = and_cols - .iter() - .map(|ac| { - let ph = d.ph(ph_idx); - ph_idx += 1; - format!("AND {} = {}", ac.value(), ph) - }) - .collect(); - let (ecs_parts, ecs_vals) = build_extra_conds_sql(&parsed.ecs, d, &mut ph_idx); - and_parts.extend(ecs_parts); - let all_extra_vals: Vec = and_vals.iter().chain(ecs_vals.iter()).cloned().collect(); - let and_str = and_parts.join(" "); + let table_str = table.value(); + let table_lit = syn::LitStr::new(&table_str, table.span()); - let has_extra = !and_cols.is_empty() || !parsed.ecs.is_empty(); + let wr = crate::where_dsl::generate_where_runtime(&parsed.dsl_where, d); + let (local_stmts, bind_stmts) = emit_runtime_binds(&wr.binds); + let (tenant_sql, tenant_bind) = emit_tenant_code(&parsed.tid, d, None); + let sql_code = wr.sql_code; - if parsed.tid.is_some() { - let tid_ph = d.ph(ph_idx); - - let expanded = if !has_extra { - let sql_with = syn::LitStr::new( - &format!( - "SELECT EXISTS(SELECT 1 FROM {} WHERE {} = {} AND tenant_id = {}) as _e", - table_str, col_str, col_ph, tid_ph - ), - table.span(), - ); - let sql_without = syn::LitStr::new( - &format!( - "SELECT EXISTS(SELECT 1 FROM {} WHERE {} = {}) as _e", - table_str, col_str, col_ph - ), - table.span(), - ); - quote! { - match #tid { - Some(_tid) => sqlx::query_scalar::<_, bool>(#sql_with).bind(#val).bind(_tid).fetch_one(#pool).await, - None => sqlx::query_scalar::<_, bool>(#sql_without).bind(#val).fetch_one(#pool).await, - } - } - } else { - let extra_idents: Vec = (0..all_extra_vals.len()) - .map(|i| syn::Ident::new(&format!("__ex_{}", i), proc_macro2::Span::call_site())) - .collect(); - let sql_with = syn::LitStr::new( - &format!( - "SELECT EXISTS(SELECT 1 FROM {} WHERE {} = {} {} AND tenant_id = {}) as _e", - table_str, col_str, col_ph, and_str, tid_ph - ), - table.span(), - ); - let sql_without = syn::LitStr::new( - &format!( - "SELECT EXISTS(SELECT 1 FROM {} WHERE {} = {} {}) as _e", - table_str, col_str, col_ph, and_str - ), - table.span(), - ); - quote! { - { - #(let #extra_idents = #all_extra_vals;)* - match #tid { - Some(_tid) => sqlx::query_scalar::<_, bool>(#sql_with).bind(#val)#(.bind(#extra_idents))*.bind(_tid).fetch_one(#pool).await, - None => sqlx::query_scalar::<_, bool>(#sql_without).bind(#val)#(.bind(#extra_idents))*.fetch_one(#pool).await, - } - } - } - }; - TokenStream::from(expanded) - } else { - let expanded = if !has_extra { - let sql_lit = syn::LitStr::new( - &format!( - "SELECT EXISTS(SELECT 1 FROM {} WHERE {} = {}) as _e", - table_str, col_str, col_ph - ), - table.span(), - ); - quote! { - sqlx::query_scalar::<_, bool>(#sql_lit).bind(#val).fetch_one(#pool).await - } - } else { - let extra_idents: Vec = (0..all_extra_vals.len()) - .map(|i| syn::Ident::new(&format!("__ex_{}", i), proc_macro2::Span::call_site())) - .collect(); - let sql_lit = syn::LitStr::new( - &format!( - "SELECT EXISTS(SELECT 1 FROM {} WHERE {} = {} {}) as _e", - table_str, col_str, col_ph, and_str - ), - table.span(), - ); - quote! { - { - #(let #extra_idents = #all_extra_vals;)* - sqlx::query_scalar::<_, bool>(#sql_lit).bind(#val)#(.bind(#extra_idents))*.fetch_one(#pool).await - } - } - }; - TokenStream::from(expanded) - } + let expanded = quote! { + { + #(#local_stmts)* + let mut __ph_idx: usize = 1; + let mut __where_sql = String::new(); + #sql_code + #tenant_sql + let __sql = format!("SELECT EXISTS(SELECT 1 FROM {} WHERE {}{}{}) as _e", #table_lit, __where_sql, __tenant_sql, ""); + let mut __q = sqlx::query_scalar::<_, bool>(&__sql); + #(#bind_stmts)* + #tenant_bind + __q.fetch_one(#pool).await + } + }; + TokenStream::from(expanded) } // ── crud_upsert! ──────────────────────────────────────────────────── @@ -2128,189 +1170,6 @@ fn expand_upsert(input: TokenStream) -> TokenStream { } } -// ── crud_patch! ──────────────────────────────────────────────────── - -pub fn crud_patch(input: TokenStream) -> TokenStream { - expand_patch(input) -} - -fn expand_patch(input: TokenStream) -> TokenStream { - let parsed = parse_macro_input!(input as PatchInput); - let table = &parsed.table; - - if let Some(err) = validate_table(table) { - return err; - } - if let Some(err) = validate_columns(table, &parsed.bind_cols) { - return err; - } - if let Some(err) = validate_columns(table, &parsed.opt_cols) { - return err; - } - if let Some(err) = validate_column(table, &parsed.pk_col) { - return err; - } - if let Some(err) = validate_columns(table, &parsed.and_cols) { - return err; - } - - let pool = &parsed.pool; - let table_str = table.value(); - let table_lit = syn::LitStr::new(&table_str, table.span()); - let pk_col_lit = syn::LitStr::new(&parsed.pk_col.value(), table.span()); - - let bind_cols = &parsed.bind_cols; - let bind_vals = &parsed.bind_vals; - let opt_cols = &parsed.opt_cols; - let opt_vals = &parsed.opt_vals; - let raw_pairs = &parsed.raw_pairs; - let and_cols = &parsed.and_cols; - let and_vals = &parsed.and_vals; - - let d = dialect(); - let mut __ph_i: usize = 1; - let bind_col_lits: Vec = bind_cols - .iter() - .map(|c| { - let ph = d.ph(__ph_i); - __ph_i += 1; - syn::LitStr::new(&format!("{} = {}", c.value(), ph), table.span()) - }) - .collect(); - let opt_ph_start = __ph_i; - let opt_col_lits: Vec = opt_cols - .iter() - .map(|c| syn::LitStr::new(&c.value(), c.span())) - .collect(); - let raw_col_names: Vec = raw_pairs.iter().map(|(c, _)| c.clone()).collect(); - let raw_exprs: Vec<&syn::Expr> = raw_pairs.iter().map(|(_, v)| v).collect(); - let and_col_lits: Vec = and_cols - .iter() - .map(|ac| syn::LitStr::new(&ac.value(), table.span())) - .collect(); - - let bind_idents: Vec = (0..bind_vals.len()) - .map(|i| syn::Ident::new(&format!("__pb_{}", i), proc_macro2::Span::call_site())) - .collect(); - let opt_idents: Vec = (0..opt_vals.len()) - .map(|i| syn::Ident::new(&format!("__po_{}", i), proc_macro2::Span::call_site())) - .collect(); - let opt_bind_idents: Vec = (0..opt_vals.len()) - .map(|i| syn::Ident::new(&format!("__pov_{}", i), proc_macro2::Span::call_site())) - .collect(); - let and_idents: Vec = (0..and_vals.len()) - .map(|i| syn::Ident::new(&format!("__pa_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let opt_ph_start_lit = - syn::LitInt::new(&opt_ph_start.to_string(), proc_macro2::Span::call_site()); - let ph_fn: syn::Expr = if matches!(d, Dialect::Postgres) { - syn::parse_quote!(|__i: usize| format!("${}", __i)) - } else { - syn::parse_quote!(|_: usize| "?".to_string()) - }; - - let pk_val = &parsed.pk_val; - let tid = &parsed.tid; - - let expanded = if parsed.tid.is_some() { - quote! { - { - #(let #bind_idents = #bind_vals;)* - #(let #opt_idents = &(#opt_vals);)* - #(let #and_idents = #and_vals;)* - let __pkv = #pk_val; - let mut __sets: Vec = Vec::new(); - #(__sets.push(#bind_col_lits.to_string());)* - #(__sets.push(format!("{} = {}", #raw_col_names, #raw_exprs));)* - let mut __ph_idx: usize = #opt_ph_start_lit; - let __ph = #ph_fn; - #( - if #opt_idents.is_some() { - __sets.push(format!("{} = {}", #opt_col_lits, __ph(__ph_idx))); - __ph_idx += 1; - } - )* - let __pk_ph = __ph(__ph_idx); - __ph_idx += 1; - let mut __and_sql = String::new(); - #( - __and_sql.push_str(&format!("AND {} = {}", #and_col_lits, __ph(__ph_idx))); - __ph_idx += 1; - )* - let _sql = match #tid { - Some(_tid) => { - let __tid_ph = __ph(__ph_idx); - format!( - "UPDATE {} SET {} WHERE {} = {}{} AND tenant_id = {}", - #table_lit, __sets.join(", "), #pk_col_lit, __pk_ph, __and_sql, __tid_ph - ) - } - None => format!( - "UPDATE {} SET {} WHERE {} = {}{}", - #table_lit, __sets.join(", "), #pk_col_lit, __pk_ph, __and_sql - ), - }; - let mut _q = sqlx::query(&_sql); - #(_q = _q.bind(#bind_idents);)* - #( - if let Some(#opt_bind_idents) = #opt_idents { - _q = _q.bind(#opt_bind_idents); - } - )* - _q = _q.bind(__pkv); - #(_q = _q.bind(#and_idents);)* - if let Some(_tid) = #tid { - _q = _q.bind(_tid); - } - _q.execute(#pool).await - } - } - } else { - quote! { - { - #(let #bind_idents = #bind_vals;)* - #(let #opt_idents = &(#opt_vals);)* - #(let #and_idents = #and_vals;)* - let __pkv = #pk_val; - let mut __sets: Vec = Vec::new(); - #(__sets.push(#bind_col_lits.to_string());)* - #(__sets.push(format!("{} = {}", #raw_col_names, #raw_exprs));)* - let mut __ph_idx: usize = #opt_ph_start_lit; - let __ph = #ph_fn; - #( - if #opt_idents.is_some() { - __sets.push(format!("{} = {}", #opt_col_lits, __ph(__ph_idx))); - __ph_idx += 1; - } - )* - let __pk_ph = __ph(__ph_idx); - __ph_idx += 1; - let mut __and_sql = String::new(); - #( - __and_sql.push_str(&format!("AND {} = {}", #and_col_lits, __ph(__ph_idx))); - __ph_idx += 1; - )* - let _sql = format!( - "UPDATE {} SET {} WHERE {} = {}{}", - #table_lit, __sets.join(", "), #pk_col_lit, __pk_ph, __and_sql - ); - let mut _q = sqlx::query(&_sql); - #(_q = _q.bind(#bind_idents);)* - #( - if let Some(#opt_bind_idents) = #opt_idents { - _q = _q.bind(#opt_bind_idents); - } - )* - _q = _q.bind(__pkv); - #(_q = _q.bind(#and_idents);)* - _q.execute(#pool).await - } - } - }; - TokenStream::from(expanded) -} - // ── check_schema! ──────────────────────────────────────────────────── /// Expand `check_schema!("table", "col1", "col2", ...)`. @@ -2345,47 +1204,35 @@ pub fn check_schema(input: TokenStream) -> TokenStream { // Tenant*Input — has a `tid` (tenant_id) field // Crud*Input — no `tid` field -// ── Delete inputs ── +// ── Where-only inputs (used by crud_delete!, crud_count!, crud_exists!) ── -/// `tenant_delete!(pool, "table", "col" => val, tenant_id [, and: ["c" => v, ...]])` -/// `crud_delete!(pool, "table", "col" => val [, tenant: expr, and: ["c" => v, ...], and_null: [...], ...])` -struct DeleteInput { +/// `crud_delete!(pool, "table", where: WhereExpr [, tenant: expr])` +/// `crud_count!(pool, "table", where: WhereExpr [, tenant: expr])` +/// `crud_exists!(pool, "table", where: WhereExpr [, tenant: expr])` +struct WhereOnlyInput { pool: syn::Expr, table: syn::LitStr, - col: syn::LitStr, - val: syn::Expr, + dsl_where: crate::where_dsl::WhereExpr, tid: Option, - and_cols: Vec, - and_vals: Vec, - ecs: ExtraConds, } -impl syn::parse::Parse for DeleteInput { +impl syn::parse::Parse for WhereOnlyInput { fn parse(input: syn::parse::ParseStream) -> syn::Result { let pool: syn::Expr = input.parse()?; let _: syn::Token![,] = input.parse()?; let table: syn::LitStr = input.parse()?; - let _: syn::Token![,] = input.parse()?; - let col: syn::LitStr = input.parse()?; - let _: syn::Token![=>] = input.parse()?; - let val: syn::Expr = input.parse()?; + let mut dsl_where = None; let mut tid = None; - let mut and_cols = Vec::new(); - let mut and_vals = Vec::new(); - let mut ecs = ExtraConds::default(); + while input.parse::().is_ok() { let section: syn::Ident = input.call(syn::Ident::parse_any)?; let _: syn::Token![:] = input.parse()?; - if section == "tenant" { + if section == "where" { + dsl_where = Some(input.parse()?); + } else if section == "tenant" { tid = Some(input.parse()?); - } else if section == "and" { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - and_cols = c; - and_vals = v; - } else if !parse_extra_conds_section(&mut ecs, §ion, input)? { + } else { return Err(syn::Error::new( section.span(), format!("unknown section: {}", section), @@ -2393,15 +1240,14 @@ impl syn::parse::Parse for DeleteInput { } } + let dsl_where = + dsl_where.ok_or_else(|| syn::Error::new(table.span(), "missing `where:` section"))?; + Ok(Self { pool, table, - col, - val, + dsl_where, tid, - and_cols, - and_vals, - ecs, }) } } @@ -2431,19 +1277,13 @@ fn parse_insert_body( Ok((pool, table, cols, vals)) } -/// `tenant_select!(pool, "table", ["col1", "col2"], "where_col" => val, tenant_id [, and: ["col" => val]])` -/// `crud_select!(pool, "table", ["col1", "col2"], "where_col" => val [, and: ["col" => val]])` +/// `crud_select!(pool, "table", ["col1", "col2"], where: WhereExpr [, tenant: expr])` struct SelectInput { pool: syn::Expr, table: syn::LitStr, sel_cols: Vec, - col: syn::LitStr, - val: syn::Expr, + dsl_where: crate::where_dsl::WhereExpr, tid: Option, - and_cols: Vec, - and_vals: Vec, - #[allow(dead_code)] - ecs: ExtraConds, } impl syn::parse::Parse for SelectInput { @@ -2459,50 +1299,34 @@ impl syn::parse::Parse for SelectInput { sel_cols.push(content.parse()?); let _ = content.parse::(); } - let _: syn::Token![,] = input.parse()?; - let col: syn::LitStr = input.parse()?; - let _: syn::Token![=>] = input.parse()?; - let val: syn::Expr = input.parse()?; + let mut dsl_where = None; let mut tid = None; - let mut and_cols = Vec::new(); - let mut and_vals = Vec::new(); - let mut ecs = ExtraConds::default(); while input.parse::().is_ok() { - if input.peek(syn::Ident) - && input.peek2(syn::Token![:]) - && !input.peek2(syn::Token![::]) - { - let section: syn::Ident = input.parse()?; - let _: syn::Token![:] = input.parse()?; - if section == "and" { - let ac; - syn::bracketed!(ac in input); - let (c, v) = parse_kv_bracket(&ac)?; - and_cols = c; - and_vals = v; - } else if !parse_extra_conds_section(&mut ecs, §ion, input)? { - return Err(syn::Error::new( - section.span(), - format!("unknown section: {}", section), - )); - } - } else { + let section: syn::Ident = input.call(syn::Ident::parse_any)?; + let _: syn::Token![:] = input.parse()?; + if section == "where" { + dsl_where = Some(input.parse()?); + } else if section == "tenant" { tid = Some(input.parse()?); + } else { + return Err(syn::Error::new( + section.span(), + format!("unknown section: {}", section), + )); } } + let dsl_where = + dsl_where.ok_or_else(|| syn::Error::new(table.span(), "missing `where:` section"))?; + Ok(Self { pool, table, sel_cols, - col, - val, + dsl_where, tid, - and_cols, - and_vals, - ecs, }) } } @@ -2531,18 +1355,14 @@ fn parse_joins(content: syn::parse::ParseStream) -> syn::Result> Ok(joins) } -/// `tenant_join!(pool, Type, select: [...], from: "...", joins: [LEFT "table" ON "..."], where: "col" => val, tenant_alias: "...", tenant: tid, method: fetch_one)` +/// `crud_join!(pool, Type, select: [...], from: "...", joins: [LEFT "table" ON "..."], where: WhereExpr, tenant_alias: "...", tenant: tid, method: fetch_one)` struct JoinInput { pool: syn::Expr, ty: syn::Type, sel_cols: Vec, from: syn::LitStr, joins: Vec, - where_col: Option, - where_val: Option, - and_cols: Vec, - and_vals: Vec, - ecs: ExtraConds, + dsl_where: Option, tenant_alias: Option, tid: Option, method: syn::Ident, @@ -2560,11 +1380,7 @@ impl syn::parse::Parse for JoinInput { let mut sel_cols = Vec::new(); let mut from = None; let mut joins = Vec::new(); - let mut where_col = None; - let mut where_val = None; - let mut and_cols = Vec::new(); - let mut and_vals = Vec::new(); - let mut ecs = ExtraConds::default(); + let mut dsl_where = None; let mut tenant_alias = None; let mut tid = None; let mut method = None; @@ -2590,17 +1406,7 @@ impl syn::parse::Parse for JoinInput { syn::bracketed!(content in input); joins = parse_joins(&content)?; } else if section == "where" { - where_col = Some(input.parse()?); - let _: syn::Token![=>] = input.parse()?; - where_val = Some(input.parse()?); - } else if section == "and" { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - and_cols = c; - and_vals = v; - } else if parse_extra_conds_section(&mut ecs, §ion, input)? { - // handled + dsl_where = Some(input.parse()?); } else if section == "tenant_alias" { tenant_alias = Some(input.parse()?); } else if section == "tenant" { @@ -2622,11 +1428,7 @@ impl syn::parse::Parse for JoinInput { sel_cols, from: from.unwrap_or_else(|| syn::LitStr::new("", proc_macro2::Span::call_site())), joins, - where_col, - where_val, - and_cols, - and_vals, - ecs, + dsl_where, tenant_alias, tid, method: method.unwrap_or_else(|| { @@ -2667,87 +1469,142 @@ pub fn crud_join_paged(input: TokenStream) -> TokenStream { let join_str = join_parts.join(" "); let d = dialect(); - let mut where_parts: Vec = Vec::new(); - if let Some(ref wc) = parsed.where_col { - where_parts.push(format!("{} = {}", wc.value(), d.ph(1))); - } - for (i, ac) in parsed.and_cols.iter().enumerate() { - let idx = if parsed.where_col.is_some() { - 2 + i - } else { - 1 + i - }; - where_parts.push(format!("{} = {}", ac.value(), d.ph(idx))); - } - - let all_and_vals: Vec = parsed.and_vals.to_vec(); - - let has_primary_where = parsed.where_col.is_some(); - let has_where = !where_parts.is_empty(); - let where_str = if has_where { - where_parts.join(" AND ") - } else { - "1=1".to_string() - }; - - let order_str = match &parsed.order_by { - Some(ob) => format!(" ORDER BY {}", ob.value()), - None => String::new(), - }; let sel_lit = syn::LitStr::new(&sel_str, proc_macro2::Span::call_site()); let from_lit = syn::LitStr::new(&from_str, proc_macro2::Span::call_site()); let join_lit = syn::LitStr::new(&join_str, proc_macro2::Span::call_site()); - let where_lit = syn::LitStr::new(&where_str, proc_macro2::Span::call_site()); - let order_lit = syn::LitStr::new(&order_str, proc_macro2::Span::call_site()); let from_for_count = syn::LitStr::new( from_str.split_whitespace().next().unwrap_or(""), proc_macro2::Span::call_site(), ); - let where_val = &parsed.where_val; - let tid = &parsed.tid; + let order_str = match &parsed.order_by { + Some(ob) => format!(" ORDER BY {}", ob.value()), + None => String::new(), + }; + let order_lit = syn::LitStr::new(&order_str, proc_macro2::Span::call_site()); + let page = &parsed.page; let page_size = &parsed.page_size; - let tenant_alias = &parsed.tenant_alias; - let and_val_idents: Vec = (0..all_and_vals.len()) - .map(|i| syn::Ident::new(&format!("__jav_{}", i), proc_macro2::Span::call_site())) - .collect(); + let tenant_alias_ref = parsed.tenant_alias.as_ref(); - let n_where = where_parts.len(); - let tenant_sql_with = match tenant_alias { - Some(alias) => format!(" AND {}.tenant_id = {}", alias.value(), d.ph(n_where + 1)), - None => format!(" AND tenant_id = {}", d.ph(n_where + 1)), + let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); + let numbered_lit = syn::LitBool::new(numbered, proc_macro2::Span::call_site()); + let ph_prefix_lit = syn::LitStr::new( + match d { + Dialect::Postgres => "$", + _ => "?", + }, + proc_macro2::Span::call_site(), + ); + + let tenant_sql_code = if let Some(tid_expr) = &parsed.tid { + let alias_prefix = match tenant_alias_ref { + Some(a) => format!("{}.", a.value()), + None => String::new(), + }; + let alias_prefix_lit = syn::LitStr::new(&alias_prefix, proc_macro2::Span::call_site()); + quote! { + let (__tenant_sql, __tid_val) = match #tid_expr { + Some(_tid) => { + let __tph = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + __ph_idx += 1; + (format!(" AND {}tenant_id = {}", #alias_prefix_lit, __tph), Some(_tid)) + }, + None => (String::new(), None), + }; + } + } else { + quote! { let __tenant_sql = String::new(); let __tid_val: Option<_> = None; } }; - let tenant_sql_with_lit = syn::LitStr::new(&tenant_sql_with, proc_macro2::Span::call_site()); - let tenant_sql_empty_lit = syn::LitStr::new("", proc_macro2::Span::call_site()); - let where_bind = if has_primary_where { - quote! { __dq = __dq.bind(__wv); __cq = __cq.bind(__wv); } + let tenant_bind = if parsed.tid.is_some() { + quote! { + if let Some(_tid) = __tid_val { + __dq = __dq.bind(_tid); + __cq = __cq.bind(_tid); + } + } } else { quote! {} }; + let (where_sql_code, local_stmts, bind_stmts_dq, bind_stmts_cq) = if let Some(ref dsl_where) = parsed.dsl_where { + let wr = crate::where_dsl::generate_where_runtime(dsl_where, d); + let ls: Vec<_> = wr.binds.iter().enumerate().map(|(i, bk)| match bk { + crate::where_dsl::BindKind::Static(expr) => { + let ident = syn::Ident::new(&format!("__wb_{}", i), proc_macro2::Span::call_site()); + quote! { let #ident = #expr; } + } + crate::where_dsl::BindKind::InLoop(expr) => { + let ident = syn::Ident::new(&format!("__in_{}", i), proc_macro2::Span::call_site()); + quote! { let #ident = #expr; } + } + }).collect(); + let bs_dq: Vec<_> = wr.binds.iter().enumerate().map(|(i, bk)| match bk { + crate::where_dsl::BindKind::Static(_) => { + let ident = syn::Ident::new(&format!("__wb_{}", i), proc_macro2::Span::call_site()); + quote! { __dq = __dq.bind(#ident); } + } + crate::where_dsl::BindKind::InLoop(_) => { + let ident = syn::Ident::new(&format!("__in_{}", i), proc_macro2::Span::call_site()); + quote! { for __iv in #ident { __dq = __dq.bind(__iv.clone()); } } + } + }).collect(); + let bs_cq: Vec<_> = wr.binds.iter().enumerate().map(|(i, bk)| match bk { + crate::where_dsl::BindKind::Static(_) => { + let ident = syn::Ident::new(&format!("__wb_{}", i), proc_macro2::Span::call_site()); + quote! { __cq = __cq.bind(#ident); } + } + crate::where_dsl::BindKind::InLoop(_) => { + let ident = syn::Ident::new(&format!("__in_{}", i), proc_macro2::Span::call_site()); + quote! { for __iv in #ident { __cq = __cq.bind(__iv.clone()); } } + } + }).collect(); + (wr.sql_code, ls, bs_dq, bs_cq) + } else { + let fallback = syn::LitStr::new("1=1", proc_macro2::Span::call_site()); + let code = quote! { let __where_sql = #fallback.to_string(); }; + (code, Vec::new(), Vec::new(), Vec::new()) + }; + + let limit_sql_code = quote! { + let __limit_ph = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + __ph_idx += 1; + let __offset_ph = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + __ph_idx += 1; + }; + let expanded = quote! { { + #(#local_stmts)* let __page_size = #page_size; let __offset = (#page - 1).max(0) * __page_size; - let __tenant_sql: &str = match #tid { - Some(_) => #tenant_sql_with_lit, - None => #tenant_sql_empty_lit, - }; - #(let #and_val_idents = #all_and_vals;)* - let __data_sql = format!("SELECT {} FROM {} {} WHERE {}{}{}", #sel_lit, #from_lit, #join_lit, #where_lit, __tenant_sql, #order_lit); - let __count_sql = format!("SELECT COUNT(*) FROM {} WHERE {}{}", #from_for_count, #where_lit, __tenant_sql); + let mut __ph_idx: usize = 1; + let mut __where_sql = String::new(); + #where_sql_code + #tenant_sql_code + #limit_sql_code + let __data_sql = format!("SELECT {} FROM {} {} WHERE {}{}{}{} LIMIT {} OFFSET {}", #sel_lit, #from_lit, #join_lit, __where_sql, __tenant_sql, #order_lit, "", __limit_ph, __offset_ph); + let __count_sql = format!("SELECT COUNT(*) FROM {} WHERE {}{}", #from_for_count, __where_sql, __tenant_sql); let mut __dq = sqlx::query_as::<_, #ty>(&__data_sql); let mut __cq = sqlx::query_scalar::<_, i64>(&__count_sql); - #where_bind - #(__dq = __dq.bind(#and_val_idents); __cq = __cq.bind(#and_val_idents);)* - if let Some(_tid) = #tid { - __dq = __dq.bind(_tid); - __cq = __cq.bind(_tid); - } + #(#bind_stmts_dq)* + #(#bind_stmts_cq)* + #tenant_bind __dq = __dq.bind(__page_size).bind(__offset); let __data = __dq.fetch_all(#pool).await?; let __total = __cq.fetch_one(#pool).await?; @@ -2755,18 +1612,7 @@ pub fn crud_join_paged(input: TokenStream) -> TokenStream { } }; - let full_code = if has_primary_where { - quote! { - { - let __wv = #where_val; - #expanded - } - } - } else { - expanded - }; - - TokenStream::from(full_code) + TokenStream::from(expanded) } struct JoinPagedInput { @@ -2775,10 +1621,7 @@ struct JoinPagedInput { sel_cols: Vec, from: syn::LitStr, joins: Vec, - where_col: Option, - where_val: Option, - and_cols: Vec, - and_vals: Vec, + dsl_where: Option, tenant_alias: Option, tid: Option, order_by: Option, @@ -2795,10 +1638,7 @@ impl syn::parse::Parse for JoinPagedInput { let mut sel_cols = Vec::new(); let mut from = None; let mut joins = Vec::new(); - let mut where_col = None; - let mut where_val = None; - let mut and_cols = Vec::new(); - let mut and_vals = Vec::new(); + let mut dsl_where = None; let mut tenant_alias = None; let mut tid = None; let mut order_by = None; @@ -2823,15 +1663,7 @@ impl syn::parse::Parse for JoinPagedInput { syn::bracketed!(content in input); joins = parse_joins(&content)?; } else if section == "where" { - where_col = Some(input.parse()?); - let _: syn::Token![=>] = input.parse()?; - where_val = Some(input.parse()?); - } else if section == "and" { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - and_cols = c; - and_vals = v; + dsl_where = Some(input.parse()?); } else if section == "tenant_alias" { tenant_alias = Some(input.parse()?); } else if section == "tenant" { @@ -2856,10 +1688,7 @@ impl syn::parse::Parse for JoinPagedInput { sel_cols, from, joins, - where_col, - where_val, - and_cols, - and_vals, + dsl_where, tenant_alias, tid, order_by, @@ -2892,397 +1721,103 @@ fn expand_join(input: TokenStream) -> TokenStream { let join_str = join_parts.join(" "); let d = dialect(); - let mut where_parts: Vec = Vec::new(); - if let Some(ref wc) = parsed.where_col { - where_parts.push(format!("{} = {}", wc.value(), d.ph(1))); - } - for (i, ac) in parsed.and_cols.iter().enumerate() { - let idx = if parsed.where_col.is_some() { - 2 + i - } else { - 1 + i - }; - where_parts.push(format!("{} = {}", ac.value(), d.ph(idx))); - } - for c in &parsed.ecs.null_cols { - where_parts.push(format!("{} IS NULL", c.value())); - } - let mut all_and_vals: Vec = parsed.and_vals.to_vec(); - all_and_vals.extend(parsed.ecs.gt_vals.iter().cloned()); - all_and_vals.extend(parsed.ecs.lt_vals.iter().cloned()); - all_and_vals.extend(parsed.ecs.gte_vals.iter().cloned()); - all_and_vals.extend(parsed.ecs.lte_vals.iter().cloned()); - - { - let mut idx = parsed.where_col.is_some() as usize + parsed.and_cols.len() + 1; - for (c, _) in parsed.ecs.gt_cols.iter().zip(&parsed.ecs.gt_vals) { - where_parts.push(format!("{} > {}", c.value(), d.ph(idx))); - idx += 1; - } - for (c, _) in parsed.ecs.lt_cols.iter().zip(&parsed.ecs.lt_vals) { - where_parts.push(format!("{} < {}", c.value(), d.ph(idx))); - idx += 1; - } - for (c, _) in parsed.ecs.gte_cols.iter().zip(&parsed.ecs.gte_vals) { - where_parts.push(format!("{} >= {}", c.value(), d.ph(idx))); - idx += 1; - } - for (c, _) in parsed.ecs.lte_cols.iter().zip(&parsed.ecs.lte_vals) { - where_parts.push(format!("{} <= {}", c.value(), d.ph(idx))); - idx += 1; - } - } - - let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); - let numbered_lit = syn::LitBool::new(numbered, proc_macro2::Span::call_site()); - let ph_prefix_lit = syn::LitStr::new( - match d { - Dialect::Postgres => "$", - _ => "?", - }, - proc_macro2::Span::call_site(), - ); - let in_base: usize = parsed.where_col.is_some() as usize + all_and_vals.len() + 1; - let in_base_lit = syn::LitInt::new(&in_base.to_string(), proc_macro2::Span::call_site()); - - let has_in = !parsed.ecs.in_cols.is_empty(); - - let and_val_idents: Vec = (0..all_and_vals.len()) - .map(|i| syn::Ident::new(&format!("__jav_{}", i), proc_macro2::Span::call_site())) - .collect(); - - let has_primary_where = parsed.where_col.is_some(); - let has_where = !where_parts.is_empty(); - let where_str = if has_where { - where_parts.join(" AND ") - } else { - "1=1".to_string() - }; + let sel_lit = syn::LitStr::new(&sel_str, proc_macro2::Span::call_site()); + let from_lit = syn::LitStr::new(&from_str, proc_macro2::Span::call_site()); + let join_lit = syn::LitStr::new(&join_str, proc_macro2::Span::call_site()); + let method = &parsed.method; let order_str = match &parsed.order_by { Some(ob) => format!(" ORDER BY {}", ob.value()), None => String::new(), }; + let order_lit = syn::LitStr::new(&order_str, proc_macro2::Span::call_site()); - let sel_lit = syn::LitStr::new(&sel_str, proc_macro2::Span::call_site()); - let from_lit = syn::LitStr::new(&from_str, proc_macro2::Span::call_site()); - let join_lit = syn::LitStr::new(&join_str, proc_macro2::Span::call_site()); - let where_lit = syn::LitStr::new(&where_str, proc_macro2::Span::call_site()); - let _order_lit = syn::LitStr::new(&order_str, proc_macro2::Span::call_site()); + let tenant_alias_ref = parsed.tenant_alias.as_ref(); + let (tenant_sql, tenant_bind) = emit_tenant_code(&parsed.tid, d, tenant_alias_ref); - let where_val = &parsed.where_val; - let method = &parsed.method; - - let in_col_lits: Vec = parsed - .ecs - .in_cols - .iter() - .map(|c| syn::LitStr::new(&c.value(), proc_macro2::Span::call_site())) - .collect(); - let in_vals = &parsed.ecs.in_vals; - - let in_loop_code = if has_in { - quote! { - #( - if !#in_vals.is_empty() { - let __in_ph: String = if #numbered_lit { - (0..#in_vals.len()).map(|i| format!(concat!(#ph_prefix_lit, "{}"), __ph_idx + i)).collect::>().join(",") - } else { - (0..#in_vals.len()).map(|_| "?").collect::>().join(",") - }; - __ph_idx += #in_vals.len(); - __in_sql.push_str(&format!(" AND {} IN ({})", #in_col_lits, __in_ph)); - } - )* - } + let (where_sql_code, local_stmts, bind_stmts) = if let Some(ref dsl_where) = parsed.dsl_where { + let wr = crate::where_dsl::generate_where_runtime(dsl_where, d); + let (ls, bs) = emit_runtime_binds(&wr.binds); + (wr.sql_code, ls, bs) } else { - quote! {} + let fallback = syn::LitStr::new("1=1", proc_macro2::Span::call_site()); + let code = quote! { let __where_sql = #fallback.to_string(); }; + (code, Vec::new(), Vec::new()) }; - let in_bind_code = if has_in { - quote! { - #( - for __iv in #in_vals { - _q = _q.bind(__iv); - } - )* - } - } else { - quote! {} - }; - - if parsed.tid.is_some() { - let tid = &parsed.tid; - let tenant_alias = &parsed.tenant_alias; - - let tenant_sql_code = match tenant_alias { - Some(alias) => { - quote! { - let __tenant_sql = match #tid { - Some(_) => { - let __tid_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - format!(" AND {}.tenant_id = {}", #alias, __tid_ph) - } - None => String::new(), - }; - } - } - None => { - quote! { - let __tenant_sql = match #tid { - Some(_) => { - let __tid_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - format!(" AND tenant_id = {}", __tid_ph) - } - None => String::new(), - }; - } - } - }; - - let limit_sql_code = if parsed.limit.is_some() { - if parsed.offset.is_some() { - quote! { - { - let __limit_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - let __offset_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - __sql.push_str(&format!(" LIMIT {} OFFSET {}", __limit_ph, __offset_ph)); - } - } - } else { - quote! { - { - let __limit_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - __sql.push_str(&format!(" LIMIT {}", __limit_ph)); - } - } - } - } else { - quote! {} - }; - - let limit_bind = if let Some(lim) = &parsed.limit { - let off = parsed.offset.as_ref(); - match off { - Some(o) => quote! { _q = _q.bind(#lim).bind(#o); }, - None => quote! { _q = _q.bind(#lim); }, - } - } else { - quote! {} - }; - - let where_bind = if has_primary_where { - quote! { _q = _q.bind(__wv); } - } else { - quote! {} - }; - - let order_lit_for_sql = syn::LitStr::new(&order_str, proc_macro2::Span::call_site()); - - let expanded = quote! { - { - let mut __ph_idx: usize = #in_base_lit; - let mut __in_sql = String::new(); - #in_loop_code - #tenant_sql_code - let mut __sql = format!("SELECT {} FROM {} {} WHERE {}{}{}", #sel_lit, #from_lit, #join_lit, #where_lit, __in_sql, __tenant_sql); - __sql.push_str(#order_lit_for_sql); - #limit_sql_code - let mut _q = sqlx::query_as::<_, #ty>(&__sql); - #where_bind - #(_q = _q.bind(#and_val_idents);)* - #in_bind_code - if let Some(_tid) = #tid { - _q = _q.bind(_tid); - } - #limit_bind - _q.#method(#pool).await - } - }; - - let full_code = if has_primary_where { - quote! { - { - let __wv = #where_val; - #(let #and_val_idents = #all_and_vals;)* - #expanded - } - } - } else { - quote! { - { - #(let #and_val_idents = #all_and_vals;)* - #expanded - } - } - }; - - TokenStream::from(full_code) - } else { - let sql_prefix = format!( - "SELECT {} FROM {} {} WHERE {}", - sel_str, from_str, join_str, where_str + let limit_sql_code = if parsed.limit.is_some() { + let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); + let numbered_lit = syn::LitBool::new(numbered, proc_macro2::Span::call_site()); + let ph_prefix_lit = syn::LitStr::new( + match d { + Dialect::Postgres => "$", + _ => "?", + }, + proc_macro2::Span::call_site(), ); - - let lim = &parsed.limit; - let off = &parsed.offset; - - let limit_code = if let Some(l) = lim.as_ref() { - let o = off.as_ref().unwrap(); - quote! { _q = _q.bind(#l).bind(#o); } + if parsed.offset.is_some() { + quote! { + { + let __limit_ph = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + __ph_idx += 1; + let __offset_ph = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + __ph_idx += 1; + __sql.push_str(&format!(" LIMIT {} OFFSET {}", __limit_ph, __offset_ph)); + } + } } else { - quote! {} - }; - - if has_in { - let sql_prefix_lit = syn::LitStr::new(&sql_prefix, proc_macro2::Span::call_site()); - let order_lit_inner = syn::LitStr::new(&order_str, proc_macro2::Span::call_site()); - - let limit_sql_code = if parsed.limit.is_some() { - if parsed.offset.is_some() { - quote! { - { - let __limit_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - let __offset_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - __sql.push_str(&format!(" LIMIT {} OFFSET {}", __limit_ph, __offset_ph)); - } - } - } else { - quote! { - { - let __limit_ph = if #numbered_lit { - format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) - } else { - "?".to_string() - }; - __ph_idx += 1; - __sql.push_str(&format!(" LIMIT {}", __limit_ph)); - } - } + quote! { + { + let __limit_ph = if #numbered_lit { + format!(concat!(#ph_prefix_lit, "{}"), __ph_idx) + } else { + "?".to_string() + }; + __ph_idx += 1; + __sql.push_str(&format!(" LIMIT {}", __limit_ph)); } - } else { - quote! {} - }; - - let full_code = if has_primary_where { - quote! { - { - let __wv = #where_val; - #(let #and_val_idents = #all_and_vals;)* - let mut __ph_idx: usize = #in_base_lit; - let mut __in_sql = String::new(); - let mut __sql = #sql_prefix_lit.to_string(); - #in_loop_code - __sql.push_str(&__in_sql); - __sql.push_str(#order_lit_inner); - #limit_sql_code - let mut _q = sqlx::query_as::<_, #ty>(&__sql).bind(__wv); - #(_q = _q.bind(#and_val_idents);)* - #in_bind_code - #limit_code - _q.#method(#pool).await - } - } - } else { - quote! { - { - #(let #and_val_idents = #all_and_vals;)* - let mut __ph_idx: usize = #in_base_lit; - let mut __in_sql = String::new(); - let mut __sql = #sql_prefix_lit.to_string(); - #in_loop_code - __sql.push_str(&__in_sql); - __sql.push_str(#order_lit_inner); - #limit_sql_code - let mut _q = sqlx::query_as::<_, #ty>(&__sql); - #(_q = _q.bind(#and_val_idents);)* - #in_bind_code - #limit_code - _q.#method(#pool).await - } - } - }; - TokenStream::from(full_code) - } else { - let mut sql_str = format!("{}{}", sql_prefix, order_str); - - let limit_sql_code = if let Some(l) = lim.as_ref() { - let lo_idx = parsed.where_col.is_some() as usize + all_and_vals.len() + 1; - sql_str.push_str(&format!( - " LIMIT {} OFFSET {}", - d.ph(lo_idx), - d.ph(lo_idx + 1) - )); - let o = off.as_ref().unwrap(); - quote! { _q = _q.bind(#l).bind(#o); } - } else { - quote! {} - }; - - if has_primary_where { - let sql_lit = syn::LitStr::new(&sql_str, proc_macro2::Span::call_site()); - let expanded = quote! { - { - let __wv = #where_val; - #(let #and_val_idents = #all_and_vals;)* - let mut _q = sqlx::query_as::<_, #ty>(#sql_lit).bind(__wv); - #(_q = _q.bind(#and_val_idents);)* - #limit_sql_code - _q.#method(#pool).await - } - }; - TokenStream::from(expanded) - } else { - let sql_lit = syn::LitStr::new(&sql_str, proc_macro2::Span::call_site()); - let expanded = quote! { - { - #(let #and_val_idents = #all_and_vals;)* - let mut _q = sqlx::query_as::<_, #ty>(#sql_lit); - #(_q = _q.bind(#and_val_idents);)* - #limit_sql_code - _q.#method(#pool).await - } - }; - TokenStream::from(expanded) } } - } + } else { + quote! {} + }; + + let limit_bind = if let Some(lim) = &parsed.limit { + let off = parsed.offset.as_ref(); + match off { + Some(o) => quote! { __q = __q.bind(#lim).bind(#o); }, + None => quote! { __q = __q.bind(#lim); }, + } + } else { + quote! {} + }; + + let expanded = quote! { + { + #(#local_stmts)* + let mut __ph_idx: usize = 1; + let mut __where_sql = String::new(); + #where_sql_code + #tenant_sql + let mut __sql = format!("SELECT {} FROM {} {} WHERE {}{}{}", #sel_lit, #from_lit, #join_lit, __where_sql, __tenant_sql, #order_lit); + #limit_sql_code + let mut __q = sqlx::query_as::<_, #ty>(&__sql); + #(#bind_stmts)* + #tenant_bind + #limit_bind + __q.#method(#pool).await + } + }; + TokenStream::from(expanded) } /// `crud_insert!(pool, "table", ["col" => val, ...] [, tenant: expr])` @@ -3426,18 +1961,14 @@ impl syn::parse::Parse for QueryInput { // ── Find inputs ── -/// `crud_find!(pool, "table", Type, "col" => val [, tenant: expr, and: ["c" => v, ...], order_by: "expr"])` +/// `crud_find!(pool, "table", Type, where: WhereExpr [, tenant: expr, order_by: "expr"])` struct FindInput { pool: syn::Expr, table: syn::LitStr, ty: syn::Type, - col: syn::LitStr, - val: syn::Expr, + dsl_where: crate::where_dsl::WhereExpr, tid: Option, order_by: Option, - and_cols: Vec, - and_vals: Vec, - ecs: ExtraConds, } impl syn::parse::Parse for FindInput { @@ -3447,30 +1978,21 @@ impl syn::parse::Parse for FindInput { let table: syn::LitStr = input.parse()?; let _: syn::Token![,] = input.parse()?; let ty: syn::Type = input.parse()?; - let _: syn::Token![,] = input.parse()?; - let col: syn::LitStr = input.parse()?; - let _: syn::Token![=>] = input.parse()?; - let val: syn::Expr = input.parse()?; + let mut dsl_where = None; let mut tid = None; let mut order_by = None; - let mut and_cols = Vec::new(); - let mut and_vals = Vec::new(); - let mut ecs = ExtraConds::default(); + while input.parse::().is_ok() { let section: syn::Ident = input.call(syn::Ident::parse_any)?; let _: syn::Token![:] = input.parse()?; - if section == "tenant" { + if section == "where" { + dsl_where = Some(input.parse()?); + } else if section == "tenant" { tid = Some(input.parse()?); } else if section == "order_by" { order_by = Some(input.parse()?); - } else if section == "and" { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - and_cols = c; - and_vals = v; - } else if !parse_extra_conds_section(&mut ecs, §ion, input)? { + } else { return Err(syn::Error::new( section.span(), format!("unknown section: {}", section), @@ -3478,17 +2000,16 @@ impl syn::parse::Parse for FindInput { } } + let dsl_where = + dsl_where.ok_or_else(|| syn::Error::new(ty.span(), "missing `where:` section"))?; + Ok(Self { pool, table, ty, - col, - val, + dsl_where, tid, order_by, - and_cols, - and_vals, - ecs, }) } } @@ -3587,12 +2108,7 @@ struct UpdateInput { opt_cols: Vec, opt_vals: Vec, raw_pairs: Vec<(syn::LitStr, syn::Expr)>, - pk_col: syn::LitStr, - pk_val: syn::Expr, - and_cols: Vec, - and_vals: Vec, - #[allow(dead_code)] - ecs: ExtraConds, + dsl_where: crate::where_dsl::WhereExpr, tid: Option, } @@ -3641,11 +2157,7 @@ impl syn::parse::Parse for UpdateInput { let mut opt_cols = Vec::new(); let mut opt_vals = Vec::new(); let mut raw_pairs: Vec<(syn::LitStr, syn::Expr)> = Vec::new(); - let mut pk_col = None; - let mut pk_val = None; - let mut and_cols = Vec::new(); - let mut and_vals = Vec::new(); - let mut ecs = ExtraConds::default(); + let mut dsl_where = None; let mut tid = None; while !input.is_empty() { @@ -3673,38 +2185,23 @@ impl syn::parse::Parse for UpdateInput { raw_pairs = parse_raw_bracket(&content)?; } "where" => { - let col: syn::LitStr = input.parse()?; - let _: syn::Token![=>] = input.parse()?; - let val: syn::Expr = input.parse()?; - pk_col = Some(col); - pk_val = Some(val); - } - "and" => { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - and_cols = c; - and_vals = v; + dsl_where = Some(input.parse()?); } "tenant" => { tid = Some(input.parse()?); } other => { - if !parse_extra_conds_section(&mut ecs, §ion, input)? { - return Err(syn::Error::new( - section.span(), - format!("unknown section: {}", other), - )); - } + return Err(syn::Error::new( + section.span(), + format!("unknown section: {}", other), + )); } } let _ = input.parse::(); } - let pk_col = - pk_col.ok_or_else(|| syn::Error::new(table.span(), "missing `where:` section"))?; - let pk_val = - pk_val.ok_or_else(|| syn::Error::new(table.span(), "missing `where:` value"))?; + let dsl_where = + dsl_where.ok_or_else(|| syn::Error::new(table.span(), "missing `where:` section"))?; Ok(Self { pool, @@ -3714,11 +2211,7 @@ impl syn::parse::Parse for UpdateInput { opt_cols, opt_vals, raw_pairs, - pk_col, - pk_val, - and_cols, - and_vals, - ecs, + dsl_where, tid, }) } @@ -3726,13 +2219,13 @@ impl syn::parse::Parse for UpdateInput { // ── QueryPaged input ── -/// `crud_query_paged!(pool, Type, data_sql: "...", count_sql: "...", binds: [...], tenant: tid, page: page, page_size: page_size)` +/// `crud_query_paged!(pool, Type, table: "...", where: DSL, order_by: "...", tenant: tid, page: page, page_size: page_size)` struct QueryPagedInput { pool: syn::Expr, ty: syn::Type, - data_sql: syn::LitStr, - count_sql: syn::LitStr, - binds: Vec, + table: syn::LitStr, + dsl_where: Option, + order_by: Option, tid: Option, page: syn::Expr, page_size: syn::Expr, @@ -3747,9 +2240,9 @@ impl syn::parse::Parse for QueryPagedInput { let ty: syn::Type = input.parse()?; let _: syn::Token![,] = input.parse()?; - let mut data_sql = None; - let mut count_sql = None; - let mut binds = Vec::new(); + let mut table = None; + let mut dsl_where = None; + let mut order_by = None; let mut tid = None; let mut page = None; let mut page_size = None; @@ -3761,20 +2254,29 @@ impl syn::parse::Parse for QueryPagedInput { let _: syn::Token![:] = input.parse()?; match section.to_string().as_str() { - "data_sql" => { - data_sql = Some(input.parse()?); + "table" => { + table = Some(input.parse()?); } - "count_sql" => { - count_sql = Some(input.parse()?); - } - "binds" => { - let content; - syn::bracketed!(content in input); - while !content.is_empty() { - binds.push(content.parse()?); - let _ = content.parse::(); + "where" => { + let fork = input.fork(); + if fork.peek(syn::token::Bracket) { + let content; + syn::bracketed!(content in input); + while !content.is_empty() { + let col: syn::LitStr = content.parse()?; + let _: syn::Token![=>] = content.parse()?; + let val: syn::Expr = content.parse()?; + where_cols.push(col); + where_vals.push(val); + let _ = content.parse::(); + } + } else { + dsl_where = Some(input.parse()?); } } + "order_by" => { + order_by = Some(input.parse()?); + } "tenant" => { tid = Some(input.parse()?); } @@ -3784,18 +2286,6 @@ impl syn::parse::Parse for QueryPagedInput { "page_size" => { page_size = Some(input.parse()?); } - "where" => { - let content; - syn::bracketed!(content in input); - while !content.is_empty() { - let col: syn::LitStr = content.parse()?; - let _: syn::Token![=>] = content.parse()?; - let val: syn::Expr = content.parse()?; - where_cols.push(col); - where_vals.push(val); - let _ = content.parse::(); - } - } other => { return Err(syn::Error::new( section.span(), @@ -3806,10 +2296,8 @@ impl syn::parse::Parse for QueryPagedInput { let _ = input.parse::(); } - let data_sql = - data_sql.ok_or_else(|| syn::Error::new(ty.span(), "missing `data_sql:` section"))?; - let count_sql = - count_sql.ok_or_else(|| syn::Error::new(ty.span(), "missing `count_sql:` section"))?; + let table = + table.ok_or_else(|| syn::Error::new(ty.span(), "missing `table:` section"))?; let page = page.ok_or_else(|| syn::Error::new(ty.span(), "missing `page:` section"))?; let page_size = page_size.ok_or_else(|| syn::Error::new(ty.span(), "missing `page_size:` section"))?; @@ -3817,9 +2305,9 @@ impl syn::parse::Parse for QueryPagedInput { Ok(Self { pool, ty, - data_sql, - count_sql, - binds, + table, + dsl_where, + order_by, tid, page, page_size, @@ -3921,111 +2409,4 @@ impl syn::parse::Parse for UpsertInput { } } -// ── Patch input ── -/// `crud_patch!(pool, "table", bind: [...], optional: [...], raw: [...], where: "pk" => val, and: [...] [, tenant: tid])` -struct PatchInput { - pool: syn::Expr, - table: syn::LitStr, - bind_cols: Vec, - bind_vals: Vec, - opt_cols: Vec, - opt_vals: Vec, - raw_pairs: Vec<(syn::LitStr, syn::Expr)>, - pk_col: syn::LitStr, - pk_val: syn::Expr, - and_cols: Vec, - and_vals: Vec, - tid: Option, -} - -impl syn::parse::Parse for PatchInput { - fn parse(input: syn::parse::ParseStream) -> syn::Result { - let pool: syn::Expr = input.parse()?; - let _: syn::Token![,] = input.parse()?; - let table: syn::LitStr = input.parse()?; - let _: syn::Token![,] = input.parse()?; - - let mut bind_cols = Vec::new(); - let mut bind_vals = Vec::new(); - let mut opt_cols = Vec::new(); - let mut opt_vals = Vec::new(); - let mut raw_pairs: Vec<(syn::LitStr, syn::Expr)> = Vec::new(); - let mut pk_col = None; - let mut pk_val = None; - let mut and_cols = Vec::new(); - let mut and_vals = Vec::new(); - let mut tid = None; - - while !input.is_empty() { - let section: syn::Ident = input.call(syn::Ident::parse_any)?; - let _: syn::Token![:] = input.parse()?; - - match section.to_string().as_str() { - "bind" => { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - bind_cols = c; - bind_vals = v; - } - "optional" => { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - opt_cols = c; - opt_vals = v; - } - "raw" => { - let content; - syn::bracketed!(content in input); - raw_pairs = parse_raw_bracket(&content)?; - } - "where" => { - let col: syn::LitStr = input.parse()?; - let _: syn::Token![=>] = input.parse()?; - let val: syn::Expr = input.parse()?; - pk_col = Some(col); - pk_val = Some(val); - } - "and" => { - let content; - syn::bracketed!(content in input); - let (c, v) = parse_kv_bracket(&content)?; - and_cols = c; - and_vals = v; - } - "tenant" => { - tid = Some(input.parse()?); - } - other => { - return Err(syn::Error::new( - section.span(), - format!("unknown section: {}", other), - )); - } - } - let _ = input.parse::(); - } - - let pk_col = - pk_col.ok_or_else(|| syn::Error::new(table.span(), "missing `where:` section"))?; - let pk_val = - pk_val.ok_or_else(|| syn::Error::new(table.span(), "missing `where:` value"))?; - - Ok(Self { - pool, - table, - bind_cols, - bind_vals, - opt_cols, - opt_vals, - raw_pairs, - pk_col, - pk_val, - and_cols, - and_vals, - tid, - }) - } -} diff --git a/raisfast-derive/src/lib.rs b/raisfast-derive/src/lib.rs index 335c4562..87abdc77 100644 --- a/raisfast-derive/src/lib.rs +++ b/raisfast-derive/src/lib.rs @@ -21,16 +21,17 @@ //! //! | Macro | SQL operation | //! |-------|---------------| -//! | `crud_delete!` | `DELETE FROM ... WHERE col = ?` | +//! | `crud_delete!` | `DELETE FROM ... WHERE WhereExpr` | //! | `crud_insert!` | `INSERT INTO ... (...) VALUES (...)` | //! | `crud_scalar!` | `SELECT scalar ...` | //! | `crud_query!` | `SELECT ...` via `query_as` | -//! | `crud_find!` | `SELECT cols FROM ... WHERE col = ?` → `fetch_optional` | +//! | `crud_find!` | `SELECT cols FROM ... WHERE WhereExpr` → `fetch_optional` | //! | `crud_find_one!` | same → `fetch_one` | //! | `crud_find_all!` | same → `fetch_all` | //! | `crud_list!` | `SELECT cols FROM ...` → `fetch_all` (no WHERE) | -//! | `crud_update!` | `UPDATE ... SET ... WHERE pk = ?` | -//! | `crud_count!` | `SELECT COUNT(*) FROM ... WHERE col = ?` | +//! | `crud_update!` | `UPDATE ... SET ... WHERE WhereExpr` | +//! | `crud_count!` | `SELECT COUNT(*) FROM ... WHERE WhereExpr` | +//! | `crud_exists!` | `SELECT EXISTS(SELECT 1 ... WHERE WhereExpr)` | //! | `crud_query_paged!` | paginated data + COUNT | //! | `crud_join_paged!` | paginated JOIN + COUNT | //! @@ -52,6 +53,7 @@ mod aspect_service; mod crud; mod event_meta; mod schema; +mod where_dsl; use proc_macro::TokenStream; @@ -64,10 +66,10 @@ pub fn derive_event_meta(input: TokenStream) -> TokenStream { event_meta::derive_event_meta(input) } -/// `crud_delete!(pool, "table", "col" => val [, tenant: expr, and: ...])` +/// `crud_delete!(pool, "table", where: WhereExpr [, tenant: expr])` /// -/// Generates `DELETE FROM table WHERE col = ?1` via `sqlx::query!()`. -/// When `tenant:` is provided, adds `AND tenant_id = ?` filter. +/// Generates `DELETE FROM table WHERE ...` via `sqlx::query()`. +/// Uses Where DSL for conditions. When `tenant:` is provided, adds `AND tenant_id = ?` filter. #[proc_macro] pub fn crud_delete(input: TokenStream) -> TokenStream { crud::crud_delete(input) @@ -90,17 +92,17 @@ pub fn crud_scalar(input: TokenStream) -> TokenStream { crud::crud_scalar(input) } -/// `crud_select!(pool, "table", ["col1", "col2"], "where_col" => val [, tenant: expr, and: ...])` +/// `crud_select!(pool, "table", ["col1", "col2"], where: WhereExpr [, tenant: expr])` /// -/// Generates `SELECT col1, col2 FROM table WHERE where_col = ?` via `sqlx::query_as`. +/// Generates `SELECT col1, col2 FROM table WHERE ...` via `sqlx::query_as`. #[proc_macro] pub fn crud_select(input: TokenStream) -> TokenStream { crud::crud_select(input) } -/// `crud_join!(pool, Type, select: [...], from: "...", joins: [...], where: "col" => val [, tenant: expr, method: fetch_all])` +/// `crud_join!(pool, Type, select: [...], from: "...", joins: [...], where: WhereExpr [, tenant: expr, method: fetch_all, order_by: "...", limit: expr, offset: expr])` /// -/// Generates a JOIN query with optional tenant filtering. +/// Generates a JOIN query with optional Where DSL conditions and tenant filtering. #[proc_macro] pub fn crud_join(input: TokenStream) -> TokenStream { crud::crud_join(input) @@ -114,9 +116,9 @@ pub fn crud_join_paged(input: TokenStream) -> TokenStream { crud::crud_join_paged(input) } -/// `crud_count!(pool, "table", "col" => val [, tenant: expr, and: ["c" => v, ...]])` +/// `crud_count!(pool, "table", where: WhereExpr [, tenant: expr])` /// -/// `SELECT COUNT(*) FROM table WHERE col = ? [AND c = ? ...]` → `i64`. +/// `SELECT COUNT(*) FROM table WHERE ...` → `i64`. #[proc_macro] pub fn crud_count(input: TokenStream) -> TokenStream { crud::crud_count(input) @@ -130,9 +132,9 @@ pub fn crud_query(input: TokenStream) -> TokenStream { crud::crud_query(input) } -/// `crud_find!(pool, "table", Type, "col" => val [, tenant: expr, and: ...])` +/// `crud_find!(pool, "table", Type, where: WhereExpr [, tenant: expr, order_by: "expr"])` /// -/// `SELECT {all_columns} FROM table WHERE col = ?` → `fetch_optional`. +/// `SELECT {all_columns} FROM table WHERE ...` → `fetch_optional`. /// Column list is generated from schema (replaces `SELECT *`). #[proc_macro] pub fn crud_find(input: TokenStream) -> TokenStream { @@ -176,9 +178,9 @@ pub fn check_schema(input: TokenStream) -> TokenStream { crud::check_schema(input) } -/// `crud_exists!(pool, "table", "col" => val [, tenant: expr, and: ["c" => v, ...]])` +/// `crud_exists!(pool, "table", where: WhereExpr [, tenant: expr])` /// -/// `SELECT EXISTS(SELECT 1 FROM table WHERE col = ? [...])` → `bool`. +/// `SELECT EXISTS(SELECT 1 FROM table WHERE ...)` → `bool`. /// Uses `sqlx::query_scalar` with compile-time verified SQL. #[proc_macro] pub fn crud_exists(input: TokenStream) -> TokenStream { @@ -194,20 +196,10 @@ pub fn crud_upsert(input: TokenStream) -> TokenStream { crud::crud_upsert(input) } -/// `crud_patch!(pool, "table", bind: [...], optional: [...], raw: [...], where: "pk" => val, and: [...] [, tenant: tid])` +/// `crud_update!(pool, "table", bind: [...], optional: [...], raw: [...], where: WhereExpr [, tenant: tid])` /// -/// Dynamic partial UPDATE — only non-None `optional:` fields are included in SET. -/// `bind:` fields are always set. `raw:` fields use SQL expressions. -/// Generates runtime `sqlx::query()`. -#[proc_macro] -pub fn crud_patch(input: TokenStream) -> TokenStream { - crud::crud_patch(input) -} - -/// `crud_update!(pool, "table", bind: [...], raw: [...], where: "pk" => val, and: [...] [, tenant: tid])` -/// -/// Generates a runtime `sqlx::query()` UPDATE. -/// Values are pre-bound to `let` variables to avoid E0716 temporary lifetime issues. +/// Generates a runtime `sqlx::query()` UPDATE. Supports `bind:` (always-set), +/// `optional:` (set only when Some), `raw:` (SQL expressions). #[proc_macro] pub fn crud_update(input: TokenStream) -> TokenStream { crud::crud_update(input) diff --git a/raisfast-derive/src/where_dsl.rs b/raisfast-derive/src/where_dsl.rs new file mode 100644 index 00000000..0066d13f --- /dev/null +++ b/raisfast-derive/src/where_dsl.rs @@ -0,0 +1,471 @@ +use syn::parse::discouraged::Speculative; + +use crate::schema::Dialect; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum CmpOp { + Eq, + Neq, + Gt, + Gte, + Lt, + Lte, + Like, + NotLike, + In, + NotIn, + IsNull, + NotNull, +} + +const OP_KEYWORDS: &[&str] = &[ + "EQ", "NEQ", "GT", "GTE", "LT", "LTE", "LIKE", "NOT_LIKE", "IN", "NOT_IN", "IS_NULL", + "NOT_NULL", +]; + +fn is_operator_keyword(s: &str) -> bool { + OP_KEYWORDS.contains(&s) +} + +fn parse_op(s: &str, span: proc_macro2::Span) -> syn::Result { + match s { + "EQ" => Ok(CmpOp::Eq), + "NEQ" => Ok(CmpOp::Neq), + "GT" => Ok(CmpOp::Gt), + "GTE" => Ok(CmpOp::Gte), + "LT" => Ok(CmpOp::Lt), + "LTE" => Ok(CmpOp::Lte), + "LIKE" => Ok(CmpOp::Like), + "NOT_LIKE" => Ok(CmpOp::NotLike), + "IN" => Ok(CmpOp::In), + "NOT_IN" => Ok(CmpOp::NotIn), + "IS_NULL" => Ok(CmpOp::IsNull), + "NOT_NULL" => Ok(CmpOp::NotNull), + _ => Err(syn::Error::new(span, format!("unknown operator: {s}"))), + } +} + +#[derive(Debug)] +#[allow(clippy::large_enum_variant)] +pub enum WhereExpr { + Condition { + col: String, + col_span: proc_macro2::Span, + op: CmpOp, + value: Option, + }, + And(Vec), + Or(Vec), +} + +impl syn::parse::Parse for WhereExpr { + fn parse(input: syn::parse::ParseStream) -> syn::Result { + let lookahead = input.lookahead1(); + if lookahead.peek(syn::Ident) { + let ident: syn::Ident = input.parse()?; + match ident.to_string().as_str() { + "AND" => { + let content; + syn::parenthesized!(content in input); + let mut exprs = Vec::new(); + while !content.is_empty() { + exprs.push(content.parse()?); + let _ = content.parse::(); + } + if exprs.is_empty() { + return Err(syn::Error::new( + ident.span(), + "AND requires at least one condition", + )); + } + Ok(WhereExpr::And(exprs)) + } + "OR" => { + let content; + syn::parenthesized!(content in input); + let mut exprs = Vec::new(); + while !content.is_empty() { + exprs.push(content.parse()?); + let _ = content.parse::(); + } + if exprs.is_empty() { + return Err(syn::Error::new( + ident.span(), + "OR requires at least one condition", + )); + } + Ok(WhereExpr::Or(exprs)) + } + _ => Err(syn::Error::new( + ident.span(), + "expected AND or OR, or a condition tuple", + )), + } + } else if lookahead.peek(syn::token::Paren) { + let content; + syn::parenthesized!(content in input); + + let col: syn::LitStr = content.parse()?; + let _: syn::Token![,] = content.parse()?; + + let fork = content.fork(); + if let Ok(ident) = fork.parse::() + && is_operator_keyword(&ident.to_string()) + { + content.advance_to(&fork); + let op = parse_op(&ident.to_string(), ident.span())?; + + if matches!(op, CmpOp::IsNull | CmpOp::NotNull) { + return Ok(WhereExpr::Condition { + col: col.value(), + col_span: col.span(), + op, + value: None, + }); + } + + let _: syn::Token![,] = content.parse()?; + let value: syn::Expr = content.parse()?; + return Ok(WhereExpr::Condition { + col: col.value(), + col_span: col.span(), + op, + value: Some(value), + }); + } + + let value: syn::Expr = content.parse()?; + Ok(WhereExpr::Condition { + col: col.value(), + col_span: col.span(), + op: CmpOp::Eq, + value: Some(value), + }) + } else { + Err(lookahead.error()) + } + } +} + +pub struct WhereCodegen { + pub sql: String, + pub binds: Vec, + pub next_idx: usize, + pub d: Dialect, +} + +impl WhereCodegen { + pub fn new(d: Dialect, start_idx: usize) -> Self { + Self { + sql: String::new(), + binds: Vec::new(), + next_idx: start_idx, + d, + } + } + + pub fn generate(&mut self, expr: &WhereExpr) { + match expr { + WhereExpr::And(exprs) => { + self.sql.push('('); + for (i, e) in exprs.iter().enumerate() { + if i > 0 { + self.sql.push_str(" AND "); + } + self.generate(e); + } + self.sql.push(')'); + } + WhereExpr::Or(exprs) => { + self.sql.push('('); + for (i, e) in exprs.iter().enumerate() { + if i > 0 { + self.sql.push_str(" OR "); + } + self.generate(e); + } + self.sql.push(')'); + } + WhereExpr::Condition { col, op, value, .. } => { + self.sql.push_str(col); + match op { + CmpOp::Eq => { + self.sql.push('='); + self.push_ph(); + self.binds.push(value.as_ref().unwrap().clone()); + } + CmpOp::Neq => { + self.sql.push_str("!="); + self.push_ph(); + self.binds.push(value.as_ref().unwrap().clone()); + } + CmpOp::Gt => { + self.sql.push('>'); + self.push_ph(); + self.binds.push(value.as_ref().unwrap().clone()); + } + CmpOp::Gte => { + self.sql.push_str(">="); + self.push_ph(); + self.binds.push(value.as_ref().unwrap().clone()); + } + CmpOp::Lt => { + self.sql.push('<'); + self.push_ph(); + self.binds.push(value.as_ref().unwrap().clone()); + } + CmpOp::Lte => { + self.sql.push_str("<="); + self.push_ph(); + self.binds.push(value.as_ref().unwrap().clone()); + } + CmpOp::Like => { + self.sql.push_str(" LIKE "); + self.push_ph(); + self.binds.push(value.as_ref().unwrap().clone()); + } + CmpOp::NotLike => { + self.sql.push_str(" NOT LIKE "); + self.push_ph(); + self.binds.push(value.as_ref().unwrap().clone()); + } + CmpOp::In | CmpOp::NotIn => { + let kw = if *op == CmpOp::In { " IN " } else { " NOT IN " }; + self.sql.push_str(kw); + self.sql.push('('); + self.push_ph(); + self.sql.push(')'); + self.binds.push(value.as_ref().unwrap().clone()); + } + CmpOp::IsNull => { + self.sql.push_str(" IS NULL"); + } + CmpOp::NotNull => { + self.sql.push_str(" IS NOT NULL"); + } + } + } + } + } + + fn push_ph(&mut self) { + self.sql.push_str(&self.d.ph(self.next_idx)); + self.next_idx += 1; + } +} + +#[derive(Debug)] +pub enum BindKind { + Static(syn::Expr), + InLoop(syn::Expr), +} + +pub struct WhereRuntimeResult { + pub sql_code: proc_macro2::TokenStream, + pub binds: Vec, +} + +pub fn generate_where_runtime(expr: &WhereExpr, d: Dialect) -> WhereRuntimeResult { + let mut result = WhereRuntimeResult { + sql_code: proc_macro2::TokenStream::new(), + binds: Vec::new(), + }; + build_runtime(expr, d, &mut result); + result +} + +fn ph_code(d: Dialect) -> proc_macro2::TokenStream { + match d { + Dialect::Sqlite => quote::quote! { format!("?{}", __ph_idx) }, + Dialect::Postgres => quote::quote! { format!("${}", __ph_idx) }, + Dialect::Mysql => quote::quote! { "?".to_string() }, + } +} + +fn build_runtime(expr: &WhereExpr, d: Dialect, result: &mut WhereRuntimeResult) { + match expr { + WhereExpr::And(exprs) => { + result.sql_code.extend(quote::quote! { __where_sql.push('('); }); + for (i, e) in exprs.iter().enumerate() { + if i > 0 { + result.sql_code.extend(quote::quote! { __where_sql.push_str(" AND "); }); + } + build_runtime(e, d, result); + } + result.sql_code.extend(quote::quote! { __where_sql.push(')'); }); + } + WhereExpr::Or(exprs) => { + result.sql_code.extend(quote::quote! { __where_sql.push('('); }); + for (i, e) in exprs.iter().enumerate() { + if i > 0 { + result.sql_code.extend(quote::quote! { __where_sql.push_str(" OR "); }); + } + build_runtime(e, d, result); + } + result.sql_code.extend(quote::quote! { __where_sql.push(')'); }); + } + WhereExpr::Condition { col, op, value, .. } => { + let col_lit = syn::LitStr::new(col, proc_macro2::Span::call_site()); + match op { + CmpOp::Eq => { + let val = value.as_ref().unwrap(); + let ph = ph_code(d); + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push('='); + __where_sql.push_str(&{ #ph }); + __ph_idx += 1; + }); + result.binds.push(BindKind::Static(val.clone())); + } + CmpOp::Neq => { + let val = value.as_ref().unwrap(); + let ph = ph_code(d); + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push_str("!="); + __where_sql.push_str(&{ #ph }); + __ph_idx += 1; + }); + result.binds.push(BindKind::Static(val.clone())); + } + CmpOp::Gt => { + let val = value.as_ref().unwrap(); + let ph = ph_code(d); + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push('>'); + __where_sql.push_str(&{ #ph }); + __ph_idx += 1; + }); + result.binds.push(BindKind::Static(val.clone())); + } + CmpOp::Gte => { + let val = value.as_ref().unwrap(); + let ph = ph_code(d); + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push_str(">="); + __where_sql.push_str(&{ #ph }); + __ph_idx += 1; + }); + result.binds.push(BindKind::Static(val.clone())); + } + CmpOp::Lt => { + let val = value.as_ref().unwrap(); + let ph = ph_code(d); + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push('<'); + __where_sql.push_str(&{ #ph }); + __ph_idx += 1; + }); + result.binds.push(BindKind::Static(val.clone())); + } + CmpOp::Lte => { + let val = value.as_ref().unwrap(); + let ph = ph_code(d); + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push_str("<="); + __where_sql.push_str(&{ #ph }); + __ph_idx += 1; + }); + result.binds.push(BindKind::Static(val.clone())); + } + CmpOp::Like => { + let val = value.as_ref().unwrap(); + let ph = ph_code(d); + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push_str(" LIKE "); + __where_sql.push_str(&{ #ph }); + __ph_idx += 1; + }); + result.binds.push(BindKind::Static(val.clone())); + } + CmpOp::NotLike => { + let val = value.as_ref().unwrap(); + let ph = ph_code(d); + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push_str(" NOT LIKE "); + __where_sql.push_str(&{ #ph }); + __ph_idx += 1; + }); + result.binds.push(BindKind::Static(val.clone())); + } + CmpOp::In | CmpOp::NotIn => { + let val = value.as_ref().unwrap(); + let kw = if *op == CmpOp::In { " IN " } else { " NOT IN " }; + let kw_lit = syn::LitStr::new(kw, proc_macro2::Span::call_site()); + let numbered = matches!(d, Dialect::Sqlite | Dialect::Postgres); + let numbered_lit = syn::LitBool::new(numbered, proc_macro2::Span::call_site()); + let ph_prefix_lit = syn::LitStr::new( + match d { + Dialect::Postgres => "$", + _ => "?", + }, + proc_macro2::Span::call_site(), + ); + let bind_idx = result.binds.len(); + let in_ident = syn::Ident::new( + &format!("__in_{}", bind_idx), + proc_macro2::Span::call_site(), + ); + let fallback = if *op == CmpOp::In { "1=0" } else { "1=1" }; + let fallback_lit = syn::LitStr::new(fallback, proc_macro2::Span::call_site()); + result.sql_code.extend(quote::quote! { + if !#in_ident.is_empty() { + __where_sql.push_str(#col_lit); + __where_sql.push_str(#kw_lit); + __where_sql.push('('); + let __in_ph: String = if #numbered_lit { + (0..#in_ident.len()).map(|i| format!(concat!(#ph_prefix_lit, "{}"), __ph_idx + i)).collect::>().join(",") + } else { + (0..#in_ident.len()).map(|_| "?").collect::>().join(",") + }; + __ph_idx += #in_ident.len(); + __where_sql.push_str(&__in_ph); + __where_sql.push(')'); + } else { + __where_sql.push_str(#col_lit); + __where_sql.push_str(#kw_lit); + __where_sql.push_str(#fallback_lit); + } + }); + result.binds.push(BindKind::InLoop(val.clone())); + } + CmpOp::IsNull => { + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push_str(" IS NULL"); + }); + } + CmpOp::NotNull => { + result.sql_code.extend(quote::quote! { + __where_sql.push_str(#col_lit); + __where_sql.push_str(" IS NOT NULL"); + }); + } + } + } + } +} + +pub fn validate_where_columns(expr: &WhereExpr, table: &str) -> Option { + match expr { + WhereExpr::And(exprs) | WhereExpr::Or(exprs) => { + for e in exprs { + if let Some(err) = validate_where_columns(e, table) { + return Some(err); + } + } + None + } + WhereExpr::Condition { col, col_span, .. } => { + let col_lit = syn::LitStr::new(col, *col_span); + crate::crud::validate_column_inner(table, &col_lit) + } + } +} diff --git a/src/cli/db_cmd.rs b/src/cli/db_cmd.rs index f1eef495..736f46f0 100644 --- a/src/cli/db_cmd.rs +++ b/src/cli/db_cmd.rs @@ -40,7 +40,7 @@ pub async fn seed( ) -> anyhow::Result<()> { let pool = init_pool(&config.database_url, 1).await?; - if raisfast_derive::crud_exists!(&pool, "users", "username" => username)? { + if raisfast_derive::crud_exists!(&pool, "users", where: ("username", username))? { println!("seed: admin user already exists ({username}), skipping"); return Ok(()); } diff --git a/src/models/api_token.rs b/src/models/api_token.rs index 467b65a0..91eb364c 100644 --- a/src/models/api_token.rs +++ b/src/models/api_token.rs @@ -74,7 +74,7 @@ pub async fn create( /// Find API Token by token_hash pub async fn find_by_hash(pool: &crate::db::Pool, token_hash: &str) -> AppResult> { - raisfast_derive::crud_find!(pool, "api_tokens", ApiToken, "token_hash" => token_hash) + raisfast_derive::crud_find!(pool, "api_tokens", ApiToken, where: ("token_hash", token_hash)) .map_err(Into::into) } @@ -107,12 +107,12 @@ pub async fn list_by_user( /// Find API Token by id pub async fn find_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult> { - raisfast_derive::crud_find!(pool, "api_tokens", ApiToken, "id" => id).map_err(Into::into) + raisfast_derive::crud_find!(pool, "api_tokens", ApiToken, where: ("id", id)).map_err(Into::into) } /// Delete API Token by id pub async fn delete_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult<()> { - raisfast_derive::crud_delete!(pool, "api_tokens", "id" => id)?; + raisfast_derive::crud_delete!(pool, "api_tokens", where: ("id", id))?; Ok(()) } @@ -121,7 +121,7 @@ pub async fn touch_last_used(pool: &crate::db::Pool, id: SnowflakeId) -> AppResu let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(pool, "api_tokens", bind: ["last_used_at" => now], - where: "id" => id + where: ("id", id) )?; Ok(()) } diff --git a/src/models/audit_log.rs b/src/models/audit_log.rs index 56d34a8d..c7f0ac98 100644 --- a/src/models/audit_log.rs +++ b/src/models/audit_log.rs @@ -58,11 +58,9 @@ pub async fn find_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, AuditEntry, - data_sql: "SELECT * FROM audit_log WHERE 1=1 ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM audit_log WHERE 1=1", - binds: [], - where: ["tenant_id" => tenant_id, "action" => action, "actor_id" => actor_id], - tenant: None::<&str>, + table: "audit_log", + where: AND(("tenant_id", tenant_id), ("action", action), ("actor_id", actor_id)), + order_by: "created_at DESC", page: page, page_size: page_size ); @@ -71,5 +69,5 @@ pub async fn find_paginated( /// Find an audit log entry by ID pub async fn find_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult { - raisfast_derive::crud_find_one!(pool, "audit_log", AuditEntry, "id" => id).map_err(Into::into) + raisfast_derive::crud_find_one!(pool, "audit_log", AuditEntry, where: ("id", id)).map_err(Into::into) } diff --git a/src/models/cart_item.rs b/src/models/cart_item.rs index 13280cc9..0c5b9e05 100644 --- a/src/models/cart_item.rs +++ b/src/models/cart_item.rs @@ -29,7 +29,7 @@ pub async fn find_by_user_id( pool, "cart_items", CartItem, - "user_id" => user_id, + where: ("user_id", user_id), order_by: "created_at DESC", tenant: tenant_id )?) @@ -47,8 +47,7 @@ pub async fn find_by_user_and_product( pool, "cart_items", CartItem, - "user_id" => user_id, - and: ["product_id" => product_id, "variant_id" => vid], + where: AND(("user_id", user_id), ("product_id", product_id), ("variant_id", vid)), tenant: tenant_id ) .map_err(Into::into) @@ -57,9 +56,7 @@ pub async fn find_by_user_and_product( pool, "cart_items", CartItem, - "user_id" => user_id, - and: ["product_id" => product_id], - and_null: ["variant_id"], + where: AND(("user_id", user_id), ("product_id", product_id), ("variant_id", IS_NULL)), tenant: tenant_id ) .map_err(Into::into) @@ -94,7 +91,7 @@ pub async fn insert( ], tenant: tenant_id )?; - raisfast_derive::crud_find_one!(pool, "cart_items", CartItem, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find_one!(pool, "cart_items", CartItem, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -109,7 +106,7 @@ pub async fn update_quantity( pool, "cart_items", bind: ["quantity" => quantity, "updated_at" => &now], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; AppError::expect_affected(&result, "cart_item") @@ -120,7 +117,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - Ok(raisfast_derive::crud_find!(pool, "cart_items", CartItem, "id" => id, tenant: tenant_id)?) + Ok(raisfast_derive::crud_find!(pool, "cart_items", CartItem, where: ("id", id), tenant: tenant_id)?) } pub async fn delete_by_id( @@ -128,7 +125,7 @@ pub async fn delete_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult<()> { - let result = raisfast_derive::crud_delete!(pool, "cart_items", "id" => id, tenant: tenant_id)?; + let result = raisfast_derive::crud_delete!(pool, "cart_items", where: ("id", id), tenant: tenant_id)?; AppError::expect_affected(&result, "cart_item") } @@ -138,7 +135,7 @@ pub async fn delete_by_user_id( tenant_id: Option<&str>, ) -> AppResult<()> { let result = - raisfast_derive::crud_delete!(pool, "cart_items", "user_id" => user_id, tenant: tenant_id)?; + raisfast_derive::crud_delete!(pool, "cart_items", where: ("user_id", user_id), tenant: tenant_id)?; AppError::expect_affected(&result, "cart_item") } @@ -147,7 +144,7 @@ pub async fn tx_delete_by_user_id( user_id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult<()> { - let result = raisfast_derive::crud_delete!(&mut *tx, "cart_items", "user_id" => user_id, tenant: tenant_id)?; + let result = raisfast_derive::crud_delete!(&mut *tx, "cart_items", where: ("user_id", user_id), tenant: tenant_id)?; AppError::expect_affected(&result, "cart_item") } @@ -157,7 +154,7 @@ pub async fn count_by_user( tenant_id: Option<&str>, ) -> AppResult { let count = - raisfast_derive::crud_count!(pool, "cart_items", "user_id" => user_id, tenant: tenant_id)?; + raisfast_derive::crud_count!(pool, "cart_items", where: ("user_id", user_id), tenant: tenant_id)?; Ok(count) } diff --git a/src/models/category.rs b/src/models/category.rs index 8c896b95..e12139b6 100644 --- a/src/models/category.rs +++ b/src/models/category.rs @@ -46,9 +46,8 @@ pub async fn find_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Category, - data_sql: "SELECT * FROM categories WHERE 1=1{tenant} ORDER BY sort_order, name", - count_sql: "SELECT COUNT(*) FROM categories WHERE 1=1{tenant}", - binds: [], + table: "categories", + order_by: "sort_order, name", tenant: tenant_id, page: page, page_size: page_size @@ -61,7 +60,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult { - raisfast_derive::crud_find_one!(pool, "categories", Category, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find_one!(pool, "categories", Category, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -121,7 +120,7 @@ pub async fn update( let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(pool, "categories", bind: ["name" => name, "slug" => slug, "description" => desc, "parent_id" => parent, "sort_order" => sort, "updated_by" => updated_by, "updated_at" => &now], - where: "id" => cat_id, + where: ("id", cat_id), tenant: tenant_id )?; @@ -133,7 +132,7 @@ pub async fn delete( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult<()> { - let result = raisfast_derive::crud_delete!(pool, "categories", "id" => id, tenant: tenant_id)?; + let result = raisfast_derive::crud_delete!(pool, "categories", where: ("id", id), tenant: tenant_id)?; AppError::expect_affected(&result, "category") } diff --git a/src/models/comment.rs b/src/models/comment.rs index c1c8b554..0baafc84 100644 --- a/src/models/comment.rs +++ b/src/models/comment.rs @@ -64,7 +64,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - Ok(raisfast_derive::crud_find!(pool, "comments", Comment, "id" => id, tenant: tenant_id)?) + Ok(raisfast_derive::crud_find!(pool, "comments", Comment, where: ("id", id), tenant: tenant_id)?) } pub async fn create( @@ -107,7 +107,7 @@ pub async fn find_approved_by_post( tenant_id: Option<&str>, ) -> AppResult> { Ok( - raisfast_derive::crud_find_all!(pool, "comments", Comment, "post_id" => post_id, tenant: tenant_id, and: ["status" => CommentStatus::Approved], order_by: "created_at ASC")?, + raisfast_derive::crud_find_all!(pool, "comments", Comment, where: AND(("post_id", post_id), ("status", CommentStatus::Approved)), tenant: tenant_id, order_by: "created_at ASC")?, ) } @@ -120,9 +120,9 @@ pub async fn find_approved_by_post_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Comment, - data_sql: "SELECT * FROM comments WHERE post_id = ? AND status = ?{tenant} ORDER BY created_at ASC", - count_sql: "SELECT COUNT(*) FROM comments WHERE post_id = ? AND status = ?{tenant}", - binds: [post_id, CommentStatus::Approved], + table: "comments", + where: AND(("post_id", post_id), ("status", CommentStatus::Approved)), + order_by: "created_at ASC", tenant: tenant_id, page: page, page_size: page_size @@ -135,7 +135,7 @@ pub async fn find_all_by_post( post_id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find_all!(pool, "comments", Comment, "post_id" => post_id, tenant: tenant_id, order_by: "created_at ASC").map_err(Into::into) + raisfast_derive::crud_find_all!(pool, "comments", Comment, where: ("post_id", post_id), tenant: tenant_id, order_by: "created_at ASC").map_err(Into::into) } #[cfg_attr(feature = "export-types", derive(TS))] @@ -185,16 +185,17 @@ pub async fn find_all_paginated( page_size: i64, tenant_id: Option<&str>, ) -> AppResult<(Vec, i64)> { - let result = raisfast_derive::crud_query_paged!( + let (rows, total) = raisfast_derive::crud_join_paged!( pool, AdminCommentRowDb, - data_sql: "SELECT c.id, c.post_id, p.title AS post_title, c.created_by, c.nickname, c.email, c.content, c.parent_id, c.status, c.created_at FROM comments c JOIN posts p ON c.post_id = p.id WHERE 1=1{tenant} ORDER BY c.created_at DESC", - count_sql: "SELECT COUNT(*) FROM comments WHERE 1=1{tenant}", - binds: [], + select: ["c.id", "c.post_id", "p.title AS post_title", "c.created_by", "c.nickname", "c.email", "c.content", "c.parent_id", "c.status", "c.created_at"], + from: "comments c", + joins: [INNER "posts p" ON "c.post_id = p.id"], + tenant_alias: "c", tenant: tenant_id, + order_by: "c.created_at DESC", page: page, page_size: page_size ); - let (rows, total) = result; Ok((rows.into_iter().map(AdminCommentRow::from).collect(), total)) } @@ -207,7 +208,7 @@ pub async fn update_status( let now = crate::utils::tz::now_utc(); let result = raisfast_derive::crud_update!(pool, "comments", bind: ["status" => status, "updated_at" => &now], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; @@ -219,7 +220,7 @@ pub async fn delete( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult<()> { - let result = raisfast_derive::crud_delete!(pool, "comments", "id" => id, tenant: tenant_id)?; + let result = raisfast_derive::crud_delete!(pool, "comments", where: ("id", id), tenant: tenant_id)?; AppError::expect_affected(&result, "comment") } diff --git a/src/models/content_revision.rs b/src/models/content_revision.rs index ff7cfc8c..e9f16f8f 100644 --- a/src/models/content_revision.rs +++ b/src/models/content_revision.rs @@ -63,7 +63,7 @@ pub async fn create_revision( "created_at" => now, ])?; - Ok(raisfast_derive::crud_find_one!(pool, "content_revisions", ContentRevision, "id" => id)?) + Ok(raisfast_derive::crud_find_one!(pool, "content_revisions", ContentRevision, where: ("id", id))?) } async fn next_revision_number( @@ -97,7 +97,7 @@ pub async fn list_revisions( record_id: SnowflakeId, ) -> AppResult> { Ok( - raisfast_derive::crud_find_all!(pool, "content_revisions", ContentRevision, "content_type" => content_type, and: ["record_id" => record_id], order_by: "revision_number DESC")?, + raisfast_derive::crud_find_all!(pool, "content_revisions", ContentRevision, where: AND(("content_type", content_type), ("record_id", record_id)), order_by: "revision_number DESC")?, ) } @@ -108,7 +108,7 @@ pub async fn get_revision( revision_id: SnowflakeId, ) -> AppResult> { Ok( - raisfast_derive::crud_find!(pool, "content_revisions", ContentRevision, "id" => revision_id, and: ["content_type" => content_type, "record_id" => record_id])?, + raisfast_derive::crud_find!(pool, "content_revisions", ContentRevision, where: AND(("id", revision_id), ("content_type", content_type), ("record_id", record_id)))?, ) } @@ -158,7 +158,7 @@ pub async fn delete_revisions( content_type: &str, record_id: SnowflakeId, ) -> AppResult { - let result = raisfast_derive::crud_delete!(pool, "content_revisions", "content_type" => content_type, and: ["record_id" => record_id])?; + let result = raisfast_derive::crud_delete!(pool, "content_revisions", where: AND(("content_type", content_type), ("record_id", record_id)))?; Ok(result.rows_affected()) } diff --git a/src/models/currencies.rs b/src/models/currencies.rs index e38d11b2..5c6aef80 100644 --- a/src/models/currencies.rs +++ b/src/models/currencies.rs @@ -18,14 +18,14 @@ pub struct Currency { } pub async fn find_by_code(pool: &crate::db::Pool, code: &str) -> AppResult> { - raisfast_derive::crud_find!(pool, "currencies", Currency, "code" => code).map_err(Into::into) + raisfast_derive::crud_find!(pool, "currencies", Currency, where: ("code", code)).map_err(Into::into) } pub async fn find_active_by_code( pool: &crate::db::Pool, code: &str, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "currencies", Currency, "code" => code, and: ["is_active" => 1i64]) + raisfast_derive::crud_find!(pool, "currencies", Currency, where: AND(("code", code), ("is_active", 1i64))) .map_err(Into::into) } @@ -33,7 +33,7 @@ pub async fn find_by_code_tx( tx: &mut crate::db::pool::DbConnection, code: &str, ) -> AppResult> { - raisfast_derive::crud_find!(tx, "currencies", Currency, "code" => code, and: ["is_active" => 1i64]) + raisfast_derive::crud_find!(tx, "currencies", Currency, where: AND(("code", code), ("is_active", 1i64))) .map_err(Into::into) } @@ -75,7 +75,7 @@ pub async fn create( ] )?; - raisfast_derive::crud_find_one!(pool, "currencies", Currency, "id" => id).map_err(Into::into) + raisfast_derive::crud_find_one!(pool, "currencies", Currency, where: ("id", id)).map_err(Into::into) } pub async fn update( @@ -106,8 +106,7 @@ pub async fn update( let result = raisfast_derive::crud_update!(pool, "currencies", bind: ["name" => name, "is_active" => is_active, "updated_at" => now], raw: ["version" => "version + 1"], - where: "id" => existing.id, - and: ["version" => existing.version] + where: AND(("id", existing.id), ("version", existing.version)) )?; let affected = result.rows_affected(); @@ -127,7 +126,7 @@ pub async fn delete_by_code(pool: &crate::db::Pool, code: &str) -> AppResult return Ok(false), }; - let count: i64 = raisfast_derive::crud_count!(pool, "wallets", "currency" => code)?; + let count: i64 = raisfast_derive::crud_count!(pool, "wallets", where: ("currency", code))?; if count > 0 { return Err(crate::errors::app_error::AppError::BadRequest(format!( @@ -135,7 +134,7 @@ pub async fn delete_by_code(pool: &crate::db::Pool, code: &str) -> AppResult existing.id)?; + let result = raisfast_derive::crud_delete!(pool, "currencies", where: ("id", existing.id))?; Ok(result.rows_affected() > 0) } diff --git a/src/models/email_verification.rs b/src/models/email_verification.rs index 8bf7d64b..5b07c1f8 100644 --- a/src/models/email_verification.rs +++ b/src/models/email_verification.rs @@ -68,8 +68,7 @@ pub async fn find_by_token( pool, "email_verification_tokens", EmailVerificationToken, - "token" => token, - and_null: ["verified_at"] + where: AND(("token", token), ("verified_at", IS_NULL)) )?) } @@ -78,7 +77,7 @@ pub async fn mark_verified(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(pool, "email_verification_tokens", bind: ["verified_at" => now], - where: "id" => id + where: ("id", id) )?; Ok(()) } @@ -88,8 +87,7 @@ pub async fn delete_unused_by_user(pool: &crate::db::Pool, user_id: SnowflakeId) raisfast_derive::crud_delete!( pool, "email_verification_tokens", - "user_id" => user_id, - and_null: ["verified_at"] + where: AND(("user_id", user_id), ("verified_at", IS_NULL)) )?; Ok(()) } diff --git a/src/models/media.rs b/src/models/media.rs index a57c17d1..baea10c1 100644 --- a/src/models/media.rs +++ b/src/models/media.rs @@ -61,7 +61,7 @@ pub async fn create( )?; let media = - raisfast_derive::crud_find_one!(pool, "media", Media, "id" => id, tenant: tenant_id)?; + raisfast_derive::crud_find_one!(pool, "media", Media, where: ("id", id), tenant: tenant_id)?; Ok(media) } @@ -75,9 +75,9 @@ pub async fn find_all( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Media, - data_sql: "SELECT * FROM media WHERE user_id = ?{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM media WHERE user_id = ?{tenant}", - binds: [user_id], + table: "media", + where: ("user_id", user_id), + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -93,9 +93,8 @@ pub async fn find_all_admin( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Media, - data_sql: "SELECT * FROM media WHERE 1=1{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM media WHERE 1=1{tenant}", - binds: [], + table: "media", + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -108,7 +107,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - Ok(raisfast_derive::crud_find!(pool, "media", Media, "id" => id, tenant: tenant_id)?) + Ok(raisfast_derive::crud_find!(pool, "media", Media, where: ("id", id), tenant: tenant_id)?) } #[derive(Debug, Serialize, Clone)] @@ -179,7 +178,7 @@ pub async fn delete( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult<()> { - let result = raisfast_derive::crud_delete!(pool, "media", "id" => id, tenant: tenant_id)?; + let result = raisfast_derive::crud_delete!(pool, "media", where: ("id", id), tenant: tenant_id)?; AppError::expect_affected(&result, "media") } diff --git a/src/models/oauth.rs b/src/models/oauth.rs index b9edbf1d..fa47e05a 100644 --- a/src/models/oauth.rs +++ b/src/models/oauth.rs @@ -79,7 +79,7 @@ pub async fn consume_state( .await?; if state.is_some() { - raisfast_derive::crud_delete!(pool, "oauth_states", "id" => id)?; + raisfast_derive::crud_delete!(pool, "oauth_states", where: ("id", id))?; } Ok(state) @@ -101,7 +101,7 @@ pub async fn find_by_provider_user( provider: &str, provider_user_id: &str, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "oauth_accounts", OAuthAccount, "provider" => provider, and: ["provider_user_id" => provider_user_id]) + raisfast_derive::crud_find!(pool, "oauth_accounts", OAuthAccount, where: AND(("provider", provider), ("provider_user_id", provider_user_id))) .map_err(Into::into) } @@ -111,7 +111,7 @@ pub async fn find_by_user_id( user_id: SnowflakeId, ) -> AppResult> { raisfast_derive::check_schema!("oauth_accounts", "user_id", "created_at"); - let accounts = raisfast_derive::crud_find_all!(pool, "oauth_accounts", OAuthAccount, "user_id" => user_id, order_by: "created_at")?; + let accounts = raisfast_derive::crud_find_all!(pool, "oauth_accounts", OAuthAccount, where: ("user_id", user_id), order_by: "created_at")?; Ok(accounts) } @@ -155,7 +155,7 @@ pub async fn create_account( "updated_at" => now, ])?; - Ok(raisfast_derive::crud_find_one!(pool, "oauth_accounts", OAuthAccount, "id" => id)?) + Ok(raisfast_derive::crud_find_one!(pool, "oauth_accounts", OAuthAccount, where: ("id", id))?) } /// Parameters for updating an OAuth account binding @@ -187,7 +187,7 @@ pub async fn update_account( "token_expires_at" => params.token_expires_at, "profile" => params.profile, ], - where: "id" => params.id + where: ("id", params.id) )?; Ok(()) } @@ -198,13 +198,13 @@ pub async fn delete_account( user_id: SnowflakeId, provider: &str, ) -> AppResult { - let result = raisfast_derive::crud_delete!(pool, "oauth_accounts", "user_id" => user_id, and: ["provider" => provider])?; + let result = raisfast_derive::crud_delete!(pool, "oauth_accounts", where: AND(("user_id", user_id), ("provider", provider)))?; Ok(result.rows_affected() > 0) } /// Count the number of OAuth providers bound to a user pub async fn count_by_user(pool: &crate::db::Pool, user_id: SnowflakeId) -> AppResult { - Ok(raisfast_derive::crud_count!(pool, "oauth_accounts", "user_id" => user_id)?) + Ok(raisfast_derive::crud_count!(pool, "oauth_accounts", where: ("user_id", user_id))?) } #[cfg(test)] diff --git a/src/models/options.rs b/src/models/options.rs index 176deb8a..728394a2 100644 --- a/src/models/options.rs +++ b/src/models/options.rs @@ -42,7 +42,7 @@ pub struct OptionRow { /// Query all autoload options (preloaded at startup) pub async fn find_autoload(pool: &crate::db::Pool) -> AppResult> { - Ok(raisfast_derive::crud_find_all!(pool, "options", OptionRow, "autoload" => 1_i64)?) + Ok(raisfast_derive::crud_find_all!(pool, "options", OptionRow, where: ("autoload", 1_i64))?) } /// Query a single option by key @@ -52,7 +52,7 @@ pub async fn find_by_key( tenant_id: Option<&str>, ) -> AppResult> { Ok( - raisfast_derive::crud_find!(pool, "options", OptionRow, "option_key" => key, tenant: tenant_id)?, + raisfast_derive::crud_find!(pool, "options", OptionRow, where: ("option_key", key), tenant: tenant_id)?, ) } @@ -76,7 +76,7 @@ pub async fn upsert_value( let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(pool, "options", bind: ["value" => value, "updated_at" => now], - where: "option_key" => key, + where: ("option_key", key), tenant: tenant_id )?; Ok(()) @@ -88,7 +88,7 @@ pub async fn delete_by_key( key: &str, tenant_id: Option<&str>, ) -> AppResult<()> { - raisfast_derive::crud_delete!(pool, "options", "option_key" => key, tenant: tenant_id)?; + raisfast_derive::crud_delete!(pool, "options", where: ("option_key", key), tenant: tenant_id)?; Ok(()) } diff --git a/src/models/order.rs b/src/models/order.rs index 274234be..790f1a5e 100644 --- a/src/models/order.rs +++ b/src/models/order.rs @@ -62,7 +62,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "orders", Order, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find!(pool, "orders", Order, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -71,7 +71,7 @@ pub async fn find_by_order_no( order_no: &str, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "orders", Order, "order_no" => order_no, tenant: tenant_id) + raisfast_derive::crud_find!(pool, "orders", Order, where: ("order_no", order_no), tenant: tenant_id) .map_err(Into::into) } @@ -84,9 +84,9 @@ pub async fn find_by_user_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Order, - data_sql: "SELECT * FROM orders WHERE user_id = ?{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM orders WHERE user_id = ?{tenant}", - binds: [user_id], + table: "orders", + where: ("user_id", user_id), + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -103,10 +103,9 @@ pub async fn find_all_admin_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Order, - data_sql: "SELECT * FROM orders WHERE 1=1{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM orders WHERE 1=1{tenant}", - binds: [], + table: "orders", where: ["status" => status], + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -225,7 +224,7 @@ pub async fn update_shipped( "carrier" => carrier, ], raw: ["updated_at" => crate::db::Driver::now_fn()], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; Ok(()) @@ -241,7 +240,7 @@ pub async fn update_admin_remark( pool, "orders", bind: ["admin_remark" => admin_remark], raw: ["updated_at" => crate::db::Driver::now_fn()], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; Ok(()) @@ -257,7 +256,7 @@ pub async fn update_delivery_data( pool, "orders", bind: ["delivery_data" => delivery_data], raw: ["updated_at" => crate::db::Driver::now_fn()], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; Ok(()) @@ -361,8 +360,7 @@ pub async fn tx_update_shipped( "carrier" => carrier ], raw: ["updated_at" => crate::db::Driver::now_fn()], - where: "id" => id, - and: ["status" => OrderStatus::Paid.as_str()] + where: AND(("id", id), ("status", OrderStatus::Paid.as_str())) )?; Ok(result.rows_affected()) } diff --git a/src/models/order_item.rs b/src/models/order_item.rs index 81f54f5d..ad517928 100644 --- a/src/models/order_item.rs +++ b/src/models/order_item.rs @@ -31,7 +31,7 @@ pub async fn find_by_order_id( order_id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find_all!(pool, "order_items", OrderItem, "order_id" => order_id, tenant: tenant_id) + raisfast_derive::crud_find_all!(pool, "order_items", OrderItem, where: ("order_id", order_id), tenant: tenant_id) .map_err(Into::into) } @@ -65,7 +65,7 @@ pub async fn insert( ], tenant: tenant_id )?; - raisfast_derive::crud_find_one!(pool, "order_items", OrderItem, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find_one!(pool, "order_items", OrderItem, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -110,7 +110,7 @@ pub async fn tx_insert( ], tenant: tenant_id )?; - raisfast_derive::crud_find_one!(&mut *tx, "order_items", OrderItem, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find_one!(&mut *tx, "order_items", OrderItem, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } diff --git a/src/models/page.rs b/src/models/page.rs index 9d786232..dd618fe9 100644 --- a/src/models/page.rs +++ b/src/models/page.rs @@ -315,7 +315,7 @@ pub async fn find_by_slug( slug: &str, tenant_id: Option<&str>, ) -> AppResult> { - Ok(raisfast_derive::crud_find!(pool, "pages", Page, "slug" => slug, tenant: tenant_id)?) + Ok(raisfast_derive::crud_find!(pool, "pages", Page, where: ("slug", slug), tenant: tenant_id)?) } pub async fn find_by_id( @@ -323,7 +323,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - Ok(raisfast_derive::crud_find!(pool, "pages", Page, "id" => id, tenant: tenant_id)?) + Ok(raisfast_derive::crud_find!(pool, "pages", Page, where: ("id", id), tenant: tenant_id)?) } pub async fn list_published( @@ -334,9 +334,9 @@ pub async fn list_published( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Page, - data_sql: "SELECT * FROM pages WHERE status = ?{tenant} ORDER BY sort_order ASC, created_at DESC", - count_sql: "SELECT COUNT(*) FROM pages WHERE status = ?{tenant}", - binds: [PageStatus::Published], + table: "pages", + where: ("status", PageStatus::Published), + order_by: "sort_order ASC, created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -353,10 +353,9 @@ pub async fn list_all( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Page, - data_sql: "SELECT * FROM pages WHERE 1=1{tenant} ORDER BY sort_order ASC, created_at DESC", - count_sql: "SELECT COUNT(*) FROM pages WHERE 1=1{tenant}", - binds: [], + table: "pages", where: ["status" => status], + order_by: "sort_order ASC, created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -570,7 +569,7 @@ pub async fn delete( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult<()> { - let result = raisfast_derive::crud_delete!(pool, "pages", "id" => id, tenant: tenant_id)?; + let result = raisfast_derive::crud_delete!(pool, "pages", where: ("id", id), tenant: tenant_id)?; AppError::expect_affected(&result, "page") } @@ -651,7 +650,7 @@ pub async fn reorder( raisfast_derive::crud_update!( pool, "pages", bind: ["sort_order" => sort_order, "updated_at" => now], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; } diff --git a/src/models/password_reset.rs b/src/models/password_reset.rs index 7d8743d4..72fdef12 100644 --- a/src/models/password_reset.rs +++ b/src/models/password_reset.rs @@ -69,8 +69,7 @@ pub async fn find_by_token( pool, "password_reset_tokens", PasswordResetToken, - "token" => token, - and_null: ["used_at"] + where: AND(("token", token), ("used_at", IS_NULL)) )?) } @@ -79,7 +78,7 @@ pub async fn mark_used(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult<()> let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(pool, "password_reset_tokens", bind: ["used_at" => now], - where: "id" => id + where: ("id", id) )?; Ok(()) } @@ -89,8 +88,7 @@ pub async fn delete_unused_by_user(pool: &crate::db::Pool, user_id: SnowflakeId) raisfast_derive::crud_delete!( pool, "password_reset_tokens", - "user_id" => user_id, - and_null: ["used_at"] + where: AND(("user_id", user_id), ("used_at", IS_NULL)) )?; Ok(()) } diff --git a/src/models/payment_channel.rs b/src/models/payment_channel.rs index a6b83f5c..0d6c2d00 100644 --- a/src/models/payment_channel.rs +++ b/src/models/payment_channel.rs @@ -27,7 +27,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "payment_channels", PaymentChannel, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find!(pool, "payment_channels", PaymentChannel, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -35,7 +35,7 @@ pub async fn find_all_active( pool: &crate::db::Pool, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find_all!(pool, "payment_channels", PaymentChannel, "is_active" => 1_i64, tenant: tenant_id, order_by: "sort_order, created_at DESC") + raisfast_derive::crud_find_all!(pool, "payment_channels", PaymentChannel, where: ("is_active", 1_i64), tenant: tenant_id, order_by: "sort_order, created_at DESC") .map_err(Into::into) } @@ -49,10 +49,9 @@ pub async fn find_all_admin_paginated( let active_val = is_active.map(|a| if a { 1_i64 } else { 0_i64 }); let result = raisfast_derive::crud_query_paged!( pool, PaymentChannel, - data_sql: "SELECT * FROM payment_channels WHERE 1=1{tenant} ORDER BY sort_order, created_at DESC", - count_sql: "SELECT COUNT(*) FROM payment_channels WHERE 1=1{tenant}", - binds: [], + table: "payment_channels", where: ["is_active" => active_val], + order_by: "sort_order, created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -116,8 +115,7 @@ pub async fn update( "sort_order" => cmd.sort_order, ], raw: ["updated_at" => crate::db::Driver::now_fn(), "version" => "version + 1"], - where: "id" => cmd.id, - and: ["version" => cmd.version], + where: AND(("id", cmd.id), ("version", cmd.version)), tenant: tenant_id )? .rows_affected(); @@ -130,7 +128,7 @@ pub async fn delete_by_id( tenant_id: Option<&str>, ) -> AppResult { let affected = - raisfast_derive::crud_delete!(pool, "payment_channels", "id" => id, tenant: tenant_id)? + raisfast_derive::crud_delete!(pool, "payment_channels", where: ("id", id), tenant: tenant_id)? .rows_affected(); Ok(affected > 0) } diff --git a/src/models/payment_order.rs b/src/models/payment_order.rs index c2d497e3..5d67608c 100644 --- a/src/models/payment_order.rs +++ b/src/models/payment_order.rs @@ -72,7 +72,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "payment_orders", PaymentOrder, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find!(pool, "payment_orders", PaymentOrder, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -81,7 +81,7 @@ pub async fn find_by_idempotency_key( key: &str, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "payment_orders", PaymentOrder, "idempotency_key" => key, tenant: tenant_id) + raisfast_derive::crud_find!(pool, "payment_orders", PaymentOrder, where: ("idempotency_key", key), tenant: tenant_id) .map_err(Into::into) } @@ -90,7 +90,7 @@ pub async fn find_by_provider_order_id( provider_order_id: &str, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "payment_orders", PaymentOrder, "provider_order_id" => provider_order_id, tenant: tenant_id).map_err(Into::into) + raisfast_derive::crud_find!(pool, "payment_orders", PaymentOrder, where: ("provider_order_id", provider_order_id), tenant: tenant_id).map_err(Into::into) } pub async fn find_by_user_paginated( @@ -102,9 +102,9 @@ pub async fn find_by_user_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, PaymentOrder, - data_sql: "SELECT * FROM payment_orders WHERE user_id = ?{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM payment_orders WHERE user_id = ?{tenant}", - binds: [user_id], + table: "payment_orders", + where: ("user_id", user_id), + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -121,10 +121,9 @@ pub async fn find_all_admin_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, PaymentOrder, - data_sql: "SELECT * FROM payment_orders WHERE 1=1{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM payment_orders WHERE 1=1{tenant}", - binds: [], + table: "payment_orders", where: ["status" => status], + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -186,7 +185,7 @@ pub async fn update_provider_order_id( pool, "payment_orders", bind: ["provider_order_id" => provider_order_id, "provider_data" => provider_data], raw: ["updated_at" => crate::db::Driver::now_fn()], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; Ok(()) diff --git a/src/models/payment_refund.rs b/src/models/payment_refund.rs index 2dc1dab2..e30502e8 100644 --- a/src/models/payment_refund.rs +++ b/src/models/payment_refund.rs @@ -29,7 +29,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "payment_refunds", PaymentRefund, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find!(pool, "payment_refunds", PaymentRefund, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -38,7 +38,7 @@ pub async fn find_by_payment_order_id( payment_order_id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find_all!(pool, "payment_refunds", PaymentRefund, "payment_order_id" => payment_order_id, tenant: tenant_id, order_by: "created_at DESC") + raisfast_derive::crud_find_all!(pool, "payment_refunds", PaymentRefund, where: ("payment_order_id", payment_order_id), tenant: tenant_id, order_by: "created_at DESC") .map_err(Into::into) } @@ -47,7 +47,7 @@ pub async fn find_by_order_id( order_id: &str, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find_all!(pool, "payment_refunds", PaymentRefund, "order_id" => order_id, tenant: tenant_id, order_by: "created_at DESC") + raisfast_derive::crud_find_all!(pool, "payment_refunds", PaymentRefund, where: ("order_id", order_id), tenant: tenant_id, order_by: "created_at DESC") .map_err(Into::into) } @@ -78,7 +78,7 @@ pub async fn insert( ], tenant: tenant_id )?; - raisfast_derive::crud_find_one!(pool, "payment_refunds", PaymentRefund, "id" => id, tenant: tenant_id).map_err(Into::into) + raisfast_derive::crud_find_one!(pool, "payment_refunds", PaymentRefund, where: ("id", id), tenant: tenant_id).map_err(Into::into) } pub async fn update_status( @@ -91,7 +91,7 @@ pub async fn update_status( pool, "payment_refunds", bind: ["status" => status], raw: ["updated_at" => crate::db::Driver::now_fn()], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; Ok(()) @@ -126,9 +126,8 @@ pub async fn find_all_admin_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, PaymentRefund, - data_sql: "SELECT * FROM payment_refunds WHERE 1=1{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM payment_refunds WHERE 1=1{tenant}", - binds: [], + table: "payment_refunds", + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size diff --git a/src/models/payment_transaction.rs b/src/models/payment_transaction.rs index 3ed35293..cf6a5619 100644 --- a/src/models/payment_transaction.rs +++ b/src/models/payment_transaction.rs @@ -25,7 +25,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "payment_transactions", PaymentTransaction, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find!(pool, "payment_transactions", PaymentTransaction, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -34,7 +34,7 @@ pub async fn find_by_payment_order_id( payment_order_id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find_all!(pool, "payment_transactions", PaymentTransaction, "payment_order_id" => payment_order_id, tenant: tenant_id, order_by: "created_at DESC") + raisfast_derive::crud_find_all!(pool, "payment_transactions", PaymentTransaction, where: ("payment_order_id", payment_order_id), tenant: tenant_id, order_by: "created_at DESC") .map_err(Into::into) } @@ -43,7 +43,7 @@ pub async fn find_by_order_id( order_id: &str, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find_all!(pool, "payment_transactions", PaymentTransaction, "order_id" => order_id, tenant: tenant_id, order_by: "created_at DESC") + raisfast_derive::crud_find_all!(pool, "payment_transactions", PaymentTransaction, where: ("order_id", order_id), tenant: tenant_id, order_by: "created_at DESC") .map_err(Into::into) } @@ -52,7 +52,7 @@ pub async fn find_by_provider_tx_id( provider_tx_id: &str, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "payment_transactions", PaymentTransaction, "provider_tx_id" => provider_tx_id, tenant: tenant_id) + raisfast_derive::crud_find!(pool, "payment_transactions", PaymentTransaction, where: ("provider_tx_id", provider_tx_id), tenant: tenant_id) .map_err(Into::into) } @@ -64,9 +64,8 @@ pub async fn find_all_admin_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, PaymentTransaction, - data_sql: "SELECT * FROM payment_transactions WHERE 1=1{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM payment_transactions WHERE 1=1{tenant}", - binds: [], + table: "payment_transactions", + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size diff --git a/src/models/plugin_storage.rs b/src/models/plugin_storage.rs index 1745f9ae..c14d4fa5 100644 --- a/src/models/plugin_storage.rs +++ b/src/models/plugin_storage.rs @@ -38,8 +38,7 @@ pub async fn get(pool: &Pool, plugin_id: &str, key: &str) -> AppResult plugin_id, - and: ["storage_key" => key] + pool, "plugin_storage", where: AND(("plugin_id", plugin_id), ("storage_key", key)) ); return Ok(None); } @@ -80,15 +79,14 @@ pub async fn delete(pool: &Pool, plugin_id: &str, key: &str) -> AppResult<()> { raisfast_derive::crud_delete!( pool, "plugin_storage", - "plugin_id" => plugin_id, - and: ["storage_key" => key] + where: AND(("plugin_id", plugin_id), ("storage_key", key)) )?; Ok(()) } /// Delete all data for a plugin pub async fn delete_all(pool: &Pool, plugin_id: &str) -> AppResult<()> { - raisfast_derive::crud_delete!(pool, "plugin_storage", "plugin_id" => plugin_id)?; + raisfast_derive::crud_delete!(pool, "plugin_storage", where: ("plugin_id", plugin_id))?; Ok(()) } diff --git a/src/models/post.rs b/src/models/post.rs index 17645973..d48f7f50 100644 --- a/src/models/post.rs +++ b/src/models/post.rs @@ -78,7 +78,7 @@ pub async fn find_by_slug( slug: &str, tenant_id: Option<&str>, ) -> AppResult> { - let post = raisfast_derive::crud_find!(pool, "posts", Post, "slug" => slug, tenant: tenant_id)?; + let post = raisfast_derive::crud_find!(pool, "posts", Post, where: ("slug", slug), tenant: tenant_id)?; Ok(post) } @@ -87,7 +87,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - let post = raisfast_derive::crud_find!(pool, "posts", Post, "id" => id, tenant: tenant_id)?; + let post = raisfast_derive::crud_find!(pool, "posts", Post, where: ("id", id), tenant: tenant_id)?; Ok(post) } @@ -152,7 +152,7 @@ async fn find_by_id_tx( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(&mut **tx, "posts", Post, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find!(&mut **tx, "posts", Post, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -174,7 +174,7 @@ pub async fn update_tx( ) -> AppResult { let post_id = cmd.id; let existing = - raisfast_derive::crud_find_one!(&mut **tx, "posts", Post, "id" => post_id, tenant: tenant_id) + raisfast_derive::crud_find_one!(&mut **tx, "posts", Post, where: ("id", post_id), tenant: tenant_id) .map_err(|_| AppError::not_found("post"))?; let now = crate::utils::tz::now_utc(); @@ -213,7 +213,7 @@ pub async fn update_tx( "published_at" => published_at, "updated_by" => updated_by, "updated_at" => now ], - where: "id" => post_id, + where: ("id", post_id), tenant: tenant_id )?; @@ -253,7 +253,7 @@ pub async fn delete( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult<()> { - let result = raisfast_derive::crud_delete!(pool, "posts", "id" => id, tenant: tenant_id)?; + let result = raisfast_derive::crud_delete!(pool, "posts", where: ("id", id), tenant: tenant_id)?; AppError::expect_affected(&result, "post") } @@ -265,8 +265,7 @@ pub async fn increment_view_count_joined( raisfast_derive::crud_update!( pool, "posts", raw: ["view_count" => "view_count + 1"], - where: "slug" => slug, - and: ["status" => PostStatus::Published], + where: AND(("slug", slug), ("status", PostStatus::Published)), tenant: tenant_id )?; @@ -289,7 +288,7 @@ pub async fn sync_tags_tx( post_id: SnowflakeId, tag_ids: &[i64], ) -> AppResult<()> { - raisfast_derive::crud_delete!(&mut **tx, "posts_tags", "post_id" => post_id)?; + raisfast_derive::crud_delete!(&mut **tx, "posts_tags", where: ("post_id", post_id))?; for tag_id in tag_ids { raisfast_derive::crud_insert!(&mut **tx, "posts_tags", ["post_id" => post_id, "tag_id" => *tag_id])?; @@ -315,7 +314,7 @@ pub async fn get_post_tags( select: ["t.id", "t.name", "t.slug"], from: "tags t", joins: [INNER "posts_tags pt" ON "t.id = pt.tag_id"], - where: "pt.post_id" => post_id, + where: ("pt.post_id", post_id), tenant_alias: "t", tenant: tenant_id, method: fetch_all @@ -337,7 +336,7 @@ pub async fn get_author_name( tenant_id: Option<&str>, ) -> AppResult> { let row: Option<(String,)> = raisfast_derive::crud_select!( - pool, "users", ["username"], "id" => created_by, tenant_id + pool, "users", ["username"], where: ("id", created_by), tenant: tenant_id )?; Ok(row.map(|(s,)| s)) } @@ -348,7 +347,7 @@ pub async fn get_category_name( tenant_id: Option<&str>, ) -> AppResult> { let row: Option<(String,)> = raisfast_derive::crud_select!( - pool, "categories", ["name"], "id" => category_id, tenant_id + pool, "categories", ["name"], where: ("id", category_id), tenant: tenant_id )?; Ok(row.map(|(s,)| s)) } @@ -529,7 +528,7 @@ pub async fn find_all_joined( LEFT "users u" ON "p.created_by = u.id", LEFT "categories c" ON "p.category_id = c.id" ], - and: ["p.status" => s], + where: ("p.status", s), tenant_alias: "p", tenant: tenant_id, order_by: "p.is_pinned DESC, p.created_at DESC", @@ -636,7 +635,7 @@ pub async fn find_joined_by_id( LEFT "users u" ON "p.created_by = u.id", LEFT "categories c" ON "p.category_id = c.id" ], - where: "p.id" => id, + where: ("p.id", id), tenant_alias: "p", tenant: tenant_id, method: fetch_one @@ -656,8 +655,7 @@ pub async fn find_published_joined_by_slug( LEFT "users u" ON "p.created_by = u.id", LEFT "categories c" ON "p.category_id = c.id" ], - where: "p.slug" => slug, - and: ["p.status" => PostStatus::Published], + where: AND(("p.slug", slug), ("p.status", PostStatus::Published)), tenant_alias: "p", tenant: tenant_id, method: fetch_one @@ -687,7 +685,7 @@ pub async fn get_tags_for_posts( select: ["pt.post_id", "t.id", "t.name", "t.slug"], from: "posts_tags pt", joins: [JOIN "tags" ON "pt.tag_id = t.id"], - and_in: ["pt.post_id" => post_ids], + where: ("pt.post_id", IN, post_ids), tenant_alias: "t", tenant: tenant_id, method: fetch_all @@ -732,8 +730,7 @@ pub async fn find_joined_by_ids( LEFT "users u" ON "p.created_by = u.id", LEFT "categories c" ON "p.category_id = c.id" ], - and: ["p.status" => PostStatus::Published], - and_in: ["p.id" => ids], + where: AND(("p.status", PostStatus::Published), ("p.id", IN, ids)), tenant_alias: "p", tenant: tenant_id, order_by: "p.is_pinned DESC, p.created_at DESC", @@ -754,9 +751,8 @@ pub async fn count_published_by_ids( raisfast_derive::crud_count!( pool, "posts", - "status" => PostStatus::Published, - tenant: tenant_id, - and_in: ["id" => ids] + where: AND(("status", PostStatus::Published), ("id", IN, ids)), + tenant: tenant_id ) .map_err(Into::into) } diff --git a/src/models/product.rs b/src/models/product.rs index 627c48ec..3a0f296d 100644 --- a/src/models/product.rs +++ b/src/models/product.rs @@ -83,7 +83,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "products", Product, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find!(pool, "products", Product, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -95,9 +95,9 @@ pub async fn find_active_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Product, - data_sql: "SELECT * FROM products WHERE status = 'active'{tenant} ORDER BY sort_order, created_at DESC", - count_sql: "SELECT COUNT(*) FROM products WHERE status = 'active'{tenant}", - binds: [], + table: "products", + where: ("status", "active"), + order_by: "sort_order, created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -114,10 +114,9 @@ pub async fn find_all_admin( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Product, - data_sql: "SELECT * FROM products WHERE 1=1{tenant} ORDER BY sort_order, created_at DESC", - count_sql: "SELECT COUNT(*) FROM products WHERE 1=1{tenant}", - binds: [], + table: "products", where: ["status" => status], + order_by: "sort_order, created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -216,8 +215,7 @@ pub async fn update( "has_variants" => cmd.has_variants, ], raw: ["updated_at" => crate::db::Driver::now_fn(), "version" => "version + 1"], - where: "id" => cmd.id, - and: ["version" => cmd.version], + where: AND(("id", cmd.id), ("version", cmd.version)), tenant: tenant_id )? .rows_affected(); @@ -229,7 +227,7 @@ pub async fn delete_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult { - let result = raisfast_derive::crud_delete!(pool, "products", "id" => id, tenant: tenant_id)?; + let result = raisfast_derive::crud_delete!(pool, "products", where: ("id", id), tenant: tenant_id)?; Ok(result.rows_affected() > 0) } diff --git a/src/models/product_variant.rs b/src/models/product_variant.rs index ff38576c..01fc1828 100644 --- a/src/models/product_variant.rs +++ b/src/models/product_variant.rs @@ -34,7 +34,7 @@ pub async fn find_by_id( pool, "product_variants", ProductVariant, - "id" => id, + where: ("id", id), tenant: tenant_id ) .map_err(Into::into) @@ -49,7 +49,7 @@ pub async fn find_by_sku( pool, "product_variants", ProductVariant, - "sku" => sku, + where: ("sku", sku), tenant: tenant_id ) .map_err(Into::into) @@ -64,7 +64,7 @@ pub async fn find_by_product_id( pool, "product_variants", ProductVariant, - "product_id" => product_id, + where: ("product_id", product_id), order_by: "sort_order, created_at", tenant: tenant_id )?) @@ -79,8 +79,7 @@ pub async fn find_active_by_product_id( pool, "product_variants", ProductVariant, - "product_id" => product_id, - and: ["is_active" => true], + where: AND(("product_id", product_id), ("is_active", true)), order_by: "sort_order, created_at", tenant: tenant_id )?) @@ -138,7 +137,7 @@ pub async fn update( "is_active" => cmd.is_active, ], raw: ["updated_at" => crate::db::Driver::now_fn()], - where: "id" => cmd.id, + where: ("id", cmd.id), tenant: tenant_id )?; Ok(result.rows_affected() > 0) @@ -150,7 +149,7 @@ pub async fn delete_by_id( tenant_id: Option<&str>, ) -> AppResult { let result: crate::db::DbQueryResult = - raisfast_derive::crud_delete!(pool, "product_variants", "id" => id, tenant: tenant_id)?; + raisfast_derive::crud_delete!(pool, "product_variants", where: ("id", id), tenant: tenant_id)?; Ok(result.rows_affected() > 0) } @@ -162,7 +161,7 @@ pub async fn delete_by_product_id( let result = raisfast_derive::crud_delete!( pool, "product_variants", - "product_id" => product_id, + where: ("product_id", product_id), tenant: tenant_id )?; AppError::expect_affected(&result, "product_variant") @@ -176,7 +175,7 @@ pub async fn count_by_product( let count = raisfast_derive::crud_count!( pool, "product_variants", - "product_id" => product_id, + where: ("product_id", product_id), tenant: tenant_id )?; Ok(count) diff --git a/src/models/rbac.rs b/src/models/rbac.rs index f939c6ed..c6a6cc89 100644 --- a/src/models/rbac.rs +++ b/src/models/rbac.rs @@ -55,7 +55,7 @@ pub async fn list_roles(pool: &crate::db::Pool) -> AppResult> { /// Find role by id pub async fn find_role_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult> { - let role = raisfast_derive::crud_find!(pool, "roles", Role, "id" => id)?; + let role = raisfast_derive::crud_find!(pool, "roles", Role, where: ("id", id))?; Ok(role) } @@ -112,7 +112,7 @@ pub async fn update_role( pool, "roles", bind: ["updated_at" => now], optional: ["name" => name, "description" => description], - where: "id" => id + where: ("id", id) )?; find_role_by_id(pool, id) @@ -122,7 +122,7 @@ pub async fn update_role( /// Delete role pub async fn delete_role(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult<()> { - raisfast_derive::crud_delete!(pool, "roles", "id" => id)?; + raisfast_derive::crud_delete!(pool, "roles", where: ("id", id))?; Ok(()) } @@ -141,7 +141,7 @@ pub async fn find_permissions_by_role_id( "conditions", "created_at" ); - let perms = raisfast_derive::crud_find_all!(pool, "permissions", Permission, "role_id" => role_id, order_by: "action")?; + let perms = raisfast_derive::crud_find_all!(pool, "permissions", Permission, where: ("role_id", role_id), order_by: "action")?; Ok(perms) } @@ -150,7 +150,7 @@ pub async fn delete_permissions_by_role_id( pool: &crate::db::Pool, role_id: SnowflakeId, ) -> AppResult<()> { - raisfast_derive::crud_delete!(pool, "permissions", "role_id" => role_id)?; + raisfast_derive::crud_delete!(pool, "permissions", where: ("role_id", role_id))?; Ok(()) } diff --git a/src/models/refresh_token.rs b/src/models/refresh_token.rs index 4ba544df..9ecf79a9 100644 --- a/src/models/refresh_token.rs +++ b/src/models/refresh_token.rs @@ -51,7 +51,7 @@ pub async fn create_token( /// /// Returns `Ok(Some(token))` or `Ok(None)` when not found. pub async fn find_by_token(pool: &crate::db::Pool, token: &str) -> AppResult> { - raisfast_derive::crud_find!(pool, "refresh_tokens", RefreshToken, "token" => token) + raisfast_derive::crud_find!(pool, "refresh_tokens", RefreshToken, where: ("token", token)) .map_err(Into::into) } @@ -59,7 +59,7 @@ pub async fn find_by_token(pool: &crate::db::Pool, token: &str) -> AppResult AppResult<()> { - raisfast_derive::crud_delete!(pool, "refresh_tokens", "token" => token)?; + raisfast_derive::crud_delete!(pool, "refresh_tokens", where: ("token", token))?; Ok(()) } @@ -67,7 +67,7 @@ pub async fn delete_by_token(pool: &crate::db::Pool, token: &str) -> AppResult<( /// /// Used for logging out all devices or forcing re-login after a password change. pub async fn delete_by_user(pool: &crate::db::Pool, user_id: SnowflakeId) -> AppResult<()> { - raisfast_derive::crud_delete!(pool, "refresh_tokens", "user_id" => user_id)?; + raisfast_derive::crud_delete!(pool, "refresh_tokens", where: ("user_id", user_id))?; Ok(()) } diff --git a/src/models/reusable_block.rs b/src/models/reusable_block.rs index 9605636d..dfc90a3c 100644 --- a/src/models/reusable_block.rs +++ b/src/models/reusable_block.rs @@ -31,7 +31,7 @@ pub async fn find_reusable_by_id( tenant_id: Option<&str>, ) -> AppResult> { Ok( - raisfast_derive::crud_find!(pool, "reusable_blocks", ReusableBlock, "id" => id, tenant: tenant_id)?, + raisfast_derive::crud_find!(pool, "reusable_blocks", ReusableBlock, where: ("id", id), tenant: tenant_id)?, ) } @@ -83,7 +83,7 @@ pub async fn update_reusable( pool, "reusable_blocks", bind: ["updated_at" => now], optional: ["updated_by" => cmd.updated_by, "name" => cmd.name, "block_type" => cmd.block_type, "content" => cmd.content, "description" => cmd.description], - where: "id" => cmd.id, + where: ("id", cmd.id), tenant: tenant_id )?; @@ -98,7 +98,7 @@ pub async fn delete_reusable( tenant_id: Option<&str>, ) -> AppResult<()> { let result = - raisfast_derive::crud_delete!(pool, "reusable_blocks", "id" => id, tenant: tenant_id)?; + raisfast_derive::crud_delete!(pool, "reusable_blocks", where: ("id", id), tenant: tenant_id)?; AppError::expect_affected(&result, "reusable_block") } diff --git a/src/models/sms_code.rs b/src/models/sms_code.rs index 14625213..64a4e880 100644 --- a/src/models/sms_code.rs +++ b/src/models/sms_code.rs @@ -65,14 +65,14 @@ pub async fn create( "created_at" => now ])?; - raisfast_derive::crud_find!(pool, "sms_codes", SmsCode, "id" => id)?.ok_or_else(|| { + raisfast_derive::crud_find!(pool, "sms_codes", SmsCode, where: ("id", id))?.ok_or_else(|| { crate::errors::app_error::AppError::Internal(anyhow::anyhow!("failed to fetch sms code")) }) } /// Find a verification code by ID pub async fn find_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult> { - Ok(raisfast_derive::crud_find!(pool, "sms_codes", SmsCode, "id" => id)?) + Ok(raisfast_derive::crud_find!(pool, "sms_codes", SmsCode, where: ("id", id))?) } /// Find the most recent unverified code for a phone number @@ -85,9 +85,7 @@ pub async fn find_latest_unverified( pool, "sms_codes", SmsCode, - "phone" => phone, - and: ["purpose" => purpose], - and_null: ["verified_at"], + where: AND(("phone", phone), ("purpose", purpose), ("verified_at", IS_NULL)), order_by: "created_at DESC LIMIT 1" )?) } @@ -103,9 +101,7 @@ pub async fn is_rate_limited( let cnt = raisfast_derive::crud_count!( pool, "sms_codes", - "phone" => phone, - and: ["purpose" => purpose], - and_gt: ["created_at" => cutoff] + where: AND(("phone", phone), ("purpose", purpose), ("created_at", GT, cutoff)) )?; Ok(cnt > 0) } @@ -136,7 +132,7 @@ pub async fn verify_code( raisfast_derive::crud_update!(pool, "sms_codes", bind: [], raw: ["attempts" => "attempts + 1"], - where: "id" => id + where: ("id", id) )?; return Ok(VerifyResult::WrongCode); } @@ -144,7 +140,7 @@ pub async fn verify_code( let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(pool, "sms_codes", bind: ["verified_at" => now], - where: "id" => id + where: ("id", id) )?; Ok(VerifyResult::Verified) diff --git a/src/models/tag.rs b/src/models/tag.rs index 6c6ce34c..81c3c2b0 100644 --- a/src/models/tag.rs +++ b/src/models/tag.rs @@ -44,9 +44,8 @@ pub async fn find_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, Tag, - data_sql: "SELECT * FROM tags WHERE 1=1{tenant} ORDER BY name", - count_sql: "SELECT COUNT(*) FROM tags WHERE 1=1{tenant}", - binds: [], + table: "tags", + order_by: "name", tenant: tenant_id, page: page, page_size: page_size @@ -59,7 +58,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult { - raisfast_derive::crud_find_one!(pool, "tags", Tag, "id" => id, tenant: tenant_id) + raisfast_derive::crud_find_one!(pool, "tags", Tag, where: ("id", id), tenant: tenant_id) .map_err(Into::into) } @@ -103,7 +102,7 @@ pub async fn update( let now = crate::utils::tz::now_utc(); let result = raisfast_derive::crud_update!(pool, "tags", bind: ["name" => name, "slug" => slug, "updated_at" => &now], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; AppError::expect_affected(&result, "tag")?; @@ -115,7 +114,7 @@ pub async fn delete( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult<()> { - let result = raisfast_derive::crud_delete!(pool, "tags", "id" => id, tenant: tenant_id)?; + let result = raisfast_derive::crud_delete!(pool, "tags", where: ("id", id), tenant: tenant_id)?; AppError::expect_affected(&result, "tag") } diff --git a/src/models/tenant.rs b/src/models/tenant.rs index 7dacad21..3f0d92a0 100644 --- a/src/models/tenant.rs +++ b/src/models/tenant.rs @@ -49,13 +49,13 @@ pub async fn find_all(pool: &crate::db::Pool) -> AppResult> { /// Find a tenant by integer primary key pub async fn find_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult> { - let tenant = raisfast_derive::crud_find!(pool, "tenants", Tenant, "id" => id)?; + let tenant = raisfast_derive::crud_find!(pool, "tenants", Tenant, where: ("id", id))?; Ok(tenant) } /// Find a tenant by domain pub async fn find_by_domain(pool: &crate::db::Pool, domain: &str) -> AppResult> { - let tenant = raisfast_derive::crud_find!(pool, "tenants", Tenant, "domain" => domain)?; + let tenant = raisfast_derive::crud_find!(pool, "tenants", Tenant, where: ("domain", domain))?; Ok(tenant) } @@ -104,7 +104,7 @@ pub async fn update( pool, "tenants", bind: ["updated_at" => now], optional: ["name" => name, "domain" => domain, "config" => config, "status" => status], - where: "id" => id + where: ("id", id) )?; find_by_id(pool, id) @@ -114,7 +114,7 @@ pub async fn update( /// Delete a tenant pub async fn delete(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult<()> { - raisfast_derive::crud_delete!(pool, "tenants", "id" => id)?; + raisfast_derive::crud_delete!(pool, "tenants", where: ("id", id))?; Ok(()) } diff --git a/src/models/user.rs b/src/models/user.rs index f0ae0b29..63a44978 100644 --- a/src/models/user.rs +++ b/src/models/user.rs @@ -81,7 +81,7 @@ pub fn encode_metadata(meta: &Option) -> Option { /// Find user by username pub async fn find_by_username(pool: &crate::db::Pool, username: &str) -> AppResult> { - Ok(raisfast_derive::crud_find!(pool, "users", User, "username" => username)?) + Ok(raisfast_derive::crud_find!(pool, "users", User, where: ("username", username))?) } /// Find user by integer primary key (internal FK lookup) @@ -90,7 +90,7 @@ pub async fn find_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult> { - Ok(raisfast_derive::crud_find!(pool, "users", User, "id" => id, tenant: tenant_id)?) + Ok(raisfast_derive::crud_find!(pool, "users", User, where: ("id", id), tenant: tenant_id)?) } /// Create a new user @@ -119,7 +119,7 @@ pub async fn create( tenant: tenant_id )?; - let user = raisfast_derive::crud_find!(pool, "users", User, "id" => id, tenant: tenant_id)? + let user = raisfast_derive::crud_find!(pool, "users", User, where: ("id", id), tenant: tenant_id)? .ok_or_else(|| AppError::Internal(anyhow::anyhow!("failed to fetch newly created user")))?; Ok(user) } @@ -162,7 +162,7 @@ pub async fn update_profile( let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(pool, "users", bind: ["username" => username, "bio" => bio, "website" => website, "avatar" => avatar, "social_links" => social_links, "metadata" => metadata, "updated_at" => &now], - where: "id" => user.id, + where: ("id", user.id), tenant: tenant_id )?; find_by_id(pool, cmd.id, tenant_id) @@ -179,9 +179,8 @@ pub async fn find_all( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, User, - data_sql: "SELECT * FROM users WHERE 1=1{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM users WHERE 1=1{tenant}", - binds: [], + table: "users", + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -199,7 +198,7 @@ pub async fn update_role( let now = crate::utils::tz::now_utc(); let result = raisfast_derive::crud_update!(pool, "users", bind: ["role" => role, "updated_at" => &now], - where: "id" => id, + where: ("id", id), tenant: tenant_id )?; AppError::expect_affected(&result, "user")?; @@ -213,7 +212,7 @@ pub async fn delete_by_id( id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult<()> { - raisfast_derive::crud_delete!(pool, "users", "id" => id, tenant: tenant_id)?; + raisfast_derive::crud_delete!(pool, "users", where: ("id", id), tenant: tenant_id)?; Ok(()) } diff --git a/src/models/user_address.rs b/src/models/user_address.rs index 229a8aa6..4e4c0685 100644 --- a/src/models/user_address.rs +++ b/src/models/user_address.rs @@ -38,7 +38,7 @@ pub async fn find_by_id( pool, "user_addresses", UserAddress, - "id" => id, + where: ("id", id), tenant: tenant_id ) .map_err(Into::into) @@ -53,7 +53,7 @@ pub async fn find_by_user_id( pool, "user_addresses", UserAddress, - "user_id" => user_id, + where: ("user_id", user_id), order_by: "is_default DESC, created_at DESC", tenant: tenant_id )?) @@ -68,8 +68,7 @@ pub async fn find_default_by_user( pool, "user_addresses", UserAddress, - "user_id" => user_id, - and: ["is_default" => true], + where: AND(("user_id", user_id), ("is_default", true)), tenant: tenant_id ) .map_err(Into::into) @@ -135,8 +134,7 @@ pub async fn update( "address_type" => &cmd.address_type, ], raw: ["updated_at" => crate::db::Driver::now_fn()], - where: "id" => cmd.id, - and: ["user_id" => cmd.user_id], + where: AND(("id", cmd.id), ("user_id", cmd.user_id)), tenant: tenant_id )?; Ok(result.rows_affected() > 0) @@ -151,8 +149,7 @@ pub async fn delete_by_id( let result: crate::db::DbQueryResult = raisfast_derive::crud_delete!( pool, "user_addresses", - "id" => id, - and: ["user_id" => user_id], + where: AND(("id", id), ("user_id", user_id)), tenant: tenant_id )?; Ok(result.rows_affected() > 0) @@ -163,7 +160,7 @@ pub async fn count_by_user( user_id: SnowflakeId, tenant_id: Option<&str>, ) -> AppResult { - let count = raisfast_derive::crud_count!(pool, "user_addresses", "user_id" => user_id, tenant: tenant_id)?; + let count = raisfast_derive::crud_count!(pool, "user_addresses", where: ("user_id", user_id), tenant: tenant_id)?; Ok(count) } diff --git a/src/models/user_credential.rs b/src/models/user_credential.rs index 01f17d3d..ad3f2540 100644 --- a/src/models/user_credential.rs +++ b/src/models/user_credential.rs @@ -55,7 +55,7 @@ pub async fn find_by_auth_type_and_identifier( auth_type: AuthType, identifier: &str, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "user_credentials", UserCredential, "auth_type" => auth_type, and: ["identifier" => identifier]) + raisfast_derive::crud_find!(pool, "user_credentials", UserCredential, where: AND(("auth_type", auth_type), ("identifier", identifier))) .map_err(Into::into) } @@ -63,12 +63,12 @@ pub async fn find_by_user_id( pool: &crate::db::Pool, user_id: SnowflakeId, ) -> AppResult> { - raisfast_derive::crud_find_all!(pool, "user_credentials", UserCredential, "user_id" => user_id) + raisfast_derive::crud_find_all!(pool, "user_credentials", UserCredential, where: ("user_id", user_id)) .map_err(Into::into) } pub async fn count_by_user(pool: &crate::db::Pool, user_id: SnowflakeId) -> AppResult { - Ok(raisfast_derive::crud_count!(pool, "user_credentials", "user_id" => user_id)?) + Ok(raisfast_derive::crud_count!(pool, "user_credentials", where: ("user_id", user_id))?) } pub async fn create( @@ -94,7 +94,7 @@ pub async fn create( "updated_at" => now ])?; let cred = - raisfast_derive::crud_find_one!(pool, "user_credentials", UserCredential, "id" => id)?; + raisfast_derive::crud_find_one!(pool, "user_credentials", UserCredential, where: ("id", id))?; Ok(cred) } @@ -106,7 +106,7 @@ pub async fn update_credential_data( let now = crate::utils::tz::now_str(); raisfast_derive::crud_update!(pool, "user_credentials", bind: ["credential_data" => credential_data, "updated_at" => &now], - where: "id" => id + where: ("id", id) )?; Ok(()) } @@ -119,13 +119,13 @@ pub async fn update_verified( let now = crate::utils::tz::now_str(); raisfast_derive::crud_update!(pool, "user_credentials", bind: ["verified" => if verified { 1 } else { 0 }, "updated_at" => &now], - where: "id" => id + where: ("id", id) )?; Ok(()) } pub async fn delete_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult { - let result = raisfast_derive::crud_delete!(pool, "user_credentials", "id" => id)?; + let result = raisfast_derive::crud_delete!(pool, "user_credentials", where: ("id", id))?; Ok(result.rows_affected() > 0) } @@ -133,6 +133,6 @@ pub async fn find_by_id( pool: &crate::db::Pool, id: SnowflakeId, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "user_credentials", UserCredential, "id" => id) + raisfast_derive::crud_find!(pool, "user_credentials", UserCredential, where: ("id", id)) .map_err(Into::into) } diff --git a/src/models/wallet.rs b/src/models/wallet.rs index 0ad98163..0a89c157 100644 --- a/src/models/wallet.rs +++ b/src/models/wallet.rs @@ -29,16 +29,16 @@ pub async fn find_by_user_and_currency( user_id: SnowflakeId, currency: &str, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "wallets", Wallet, "user_id" => user_id, and: ["currency" => currency]) + raisfast_derive::crud_find!(pool, "wallets", Wallet, where: AND(("user_id", user_id), ("currency", currency))) .map_err(Into::into) } pub async fn find_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult> { - raisfast_derive::crud_find!(pool, "wallets", Wallet, "id" => id).map_err(Into::into) + raisfast_derive::crud_find!(pool, "wallets", Wallet, where: ("id", id)).map_err(Into::into) } pub async fn find_by_user(pool: &crate::db::Pool, user_id: SnowflakeId) -> AppResult> { - raisfast_derive::crud_find_all!(pool, "wallets", Wallet, "user_id" => user_id) + raisfast_derive::crud_find_all!(pool, "wallets", Wallet, where: ("user_id", user_id)) .map_err(Into::into) } @@ -58,7 +58,7 @@ pub async fn create( "created_at" => now, "updated_at" => now ])?; - raisfast_derive::crud_find_one!(pool, "wallets", Wallet, "id" => id).map_err(Into::into) + raisfast_derive::crud_find_one!(pool, "wallets", Wallet, where: ("id", id)).map_err(Into::into) } pub async fn find_or_create( @@ -78,12 +78,10 @@ pub async fn find_all_wallets( page_size: i64, tenant_id: Option<&str>, ) -> AppResult<(Vec, i64)> { - raisfast_derive::check_schema!("wallets", "tenant_id", "created_at"); let result = raisfast_derive::crud_query_paged!( pool, Wallet, - data_sql: "SELECT * FROM wallets WHERE 1=1{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM wallets WHERE 1=1{tenant}", - binds: [], + table: "wallets", + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size diff --git a/src/models/wallet_outbox.rs b/src/models/wallet_outbox.rs index 17f4565d..66bbcdca 100644 --- a/src/models/wallet_outbox.rs +++ b/src/models/wallet_outbox.rs @@ -100,7 +100,7 @@ pub async fn mark_processing(pool: &crate::db::Pool, id: SnowflakeId) -> AppResu pub async fn mark_completed(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult<()> { raisfast_derive::crud_update!(pool, "wallet_outbox", raw: ["status" => "'completed'", "updated_at" => crate::db::Driver::now_fn()], - where: "id" => id + where: ("id", id) )?; Ok(()) } diff --git a/src/models/wallet_transaction.rs b/src/models/wallet_transaction.rs index 5a0c6b2f..fb591ceb 100644 --- a/src/models/wallet_transaction.rs +++ b/src/models/wallet_transaction.rs @@ -63,10 +63,9 @@ pub async fn find_transactions_by_wallet( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, WalletTransaction, - data_sql: "SELECT * FROM wallet_transactions WHERE wallet_id = ? ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM wallet_transactions WHERE wallet_id = ?", - binds: [wallet_id], - tenant: None::<&str>, + table: "wallet_transactions", + where: ("wallet_id", wallet_id), + order_by: "created_at DESC", page: page, page_size: page_size ); @@ -81,10 +80,9 @@ pub async fn find_transactions_by_user( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, WalletTransaction, - data_sql: "SELECT * FROM wallet_transactions WHERE user_id = ? ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM wallet_transactions WHERE user_id = ?", - binds: [user_id], - tenant: None::<&str>, + table: "wallet_transactions", + where: ("user_id", user_id), + order_by: "created_at DESC", page: page, page_size: page_size ); @@ -97,12 +95,10 @@ pub async fn find_all_transactions( page_size: i64, tenant_id: Option<&str>, ) -> AppResult<(Vec, i64)> { - raisfast_derive::check_schema!("wallet_transactions", "tenant_id", "created_at"); let result = raisfast_derive::crud_query_paged!( pool, WalletTransaction, - data_sql: "SELECT * FROM wallet_transactions WHERE 1=1{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM wallet_transactions WHERE 1=1{tenant}", - binds: [], + table: "wallet_transactions", + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -114,7 +110,7 @@ pub async fn find_tx_by_transaction_no( pool: &crate::db::Pool, transaction_no: &str, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "wallet_transactions", WalletTransaction, "transaction_no" => transaction_no) + raisfast_derive::crud_find!(pool, "wallet_transactions", WalletTransaction, where: ("transaction_no", transaction_no)) .map_err(Into::into) } @@ -122,7 +118,7 @@ pub async fn find_tx_by_id( pool: &crate::db::Pool, id: SnowflakeId, ) -> AppResult> { - raisfast_derive::crud_find!(pool, "wallet_transactions", WalletTransaction, "id" => id) + raisfast_derive::crud_find!(pool, "wallet_transactions", WalletTransaction, where: ("id", id)) .map_err(Into::into) } @@ -131,8 +127,7 @@ pub async fn has_reversal_for( related_tx_id: SnowflakeId, ) -> AppResult { Ok(raisfast_derive::crud_exists!( - pool, "wallet_transactions", "related_tx_id" => related_tx_id, - and: ["tx_type" => WalletTxType::Refund] + pool, "wallet_transactions", where: AND(("related_tx_id", related_tx_id), ("tx_type", WalletTxType::Refund)) )?) } diff --git a/src/services/auth.rs b/src/services/auth.rs index b943290a..de2ed25c 100644 --- a/src/services/auth.rs +++ b/src/services/auth.rs @@ -391,7 +391,7 @@ pub async fn refresh( let now = crate::utils::tz::now_str(); in_transaction!(pool, tx, { - raisfast_derive::crud_delete!(&mut *tx, "refresh_tokens", "token" => refresh_token_str)?; + raisfast_derive::crud_delete!(&mut *tx, "refresh_tokens", where: ("token", refresh_token_str))?; raisfast_derive::crud_insert!(&mut *tx, "refresh_tokens", [ "id" => new_id, diff --git a/src/services/email_verification.rs b/src/services/email_verification.rs index 8b83bba7..70b61a79 100644 --- a/src/services/email_verification.rs +++ b/src/services/email_verification.rs @@ -43,14 +43,13 @@ pub async fn verify_email(pool: &crate::db::Pool, token: &str) -> AppResult<()> let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(&mut *tx, "email_verification_tokens", bind: ["verified_at" => now], - where: "id" => verification.id + where: ("id", verification.id) )?; let now = crate::utils::tz::now_str(); raisfast_derive::crud_update!(&mut *tx, "user_credentials", bind: ["verified" => 1i64, "updated_at" => &now], - where: "user_id" => verification.user_id, - and: ["auth_type" => crate::models::user_credential::AuthType::Email] + where: AND(("user_id", verification.user_id), ("auth_type", crate::models::user_credential::AuthType::Email)) )?; Ok::<_, crate::errors::app_error::AppError>(()) })?; diff --git a/src/services/oauth.rs b/src/services/oauth.rs index 5dd588a6..81f2cbb6 100644 --- a/src/services/oauth.rs +++ b/src/services/oauth.rs @@ -327,7 +327,7 @@ async fn auto_register_user( let now = crate::utils::tz::now_str(); raisfast_derive::crud_update!(pool, "users", bind: ["avatar" => avatar, "updated_at" => &now], - where: "id" => user.id + where: ("id", user.id) )?; } diff --git a/src/services/password_reset.rs b/src/services/password_reset.rs index ad4c3bcc..573d107f 100644 --- a/src/services/password_reset.rs +++ b/src/services/password_reset.rs @@ -68,24 +68,23 @@ pub async fn reset_password( &mut *tx, "user_credentials", (SnowflakeId, crate::models::user_credential::AuthType), - "user_id" => reset_token.user_id, - and: ["auth_type" => crate::models::user_credential::AuthType::Email] + where: AND(("user_id", reset_token.user_id), ("auth_type", crate::models::user_credential::AuthType::Email)) )?; if let Some((cred_id, _)) = row.into_iter().next() { let now = crate::utils::tz::now_str(); raisfast_derive::crud_update!(&mut *tx, "user_credentials", bind: ["credential_data" => crate::models::user_credential::wrap_password_hash(&new_hash), "updated_at" => &now], - where: "id" => cred_id + where: ("id", cred_id) )?; } let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(&mut *tx, "password_reset_tokens", bind: ["used_at" => now], - where: "id" => reset_token.id + where: ("id", reset_token.id) )?; - raisfast_derive::crud_delete!(&mut *tx, "refresh_tokens", "user_id" => reset_token.user_id)?; + raisfast_derive::crud_delete!(&mut *tx, "refresh_tokens", where: ("user_id", reset_token.user_id))?; Ok::<_, crate::errors::app_error::AppError>(()) })?; diff --git a/src/services/wallet.rs b/src/services/wallet.rs index 04d2bd2b..ce1b11f3 100644 --- a/src/services/wallet.rs +++ b/src/services/wallet.rs @@ -165,7 +165,7 @@ async fn tx_find_wallet_by_id( tx: &mut DbConnection, id: SnowflakeId, ) -> AppResult> { - Ok(raisfast_derive::crud_find!(tx, "wallets", wallet::Wallet, "id" => id)?) + Ok(raisfast_derive::crud_find!(tx, "wallets", wallet::Wallet, where: ("id", id))?) } async fn tx_find_or_create( @@ -174,8 +174,7 @@ async fn tx_find_or_create( currency: &str, ) -> AppResult { if let Some(w) = raisfast_derive::crud_find!( - &mut *tx, "wallets", wallet::Wallet, "user_id" => user_id, - and: ["currency" => currency] + &mut *tx, "wallets", wallet::Wallet, where: AND(("user_id", user_id), ("currency", currency)) )? { return Ok(w); } @@ -190,11 +189,10 @@ async fn tx_find_or_create( match insert_result { Ok(_) => { - Ok(raisfast_derive::crud_find_one!(&mut *tx, "wallets", wallet::Wallet, "id" => id)?) + Ok(raisfast_derive::crud_find_one!(&mut *tx, "wallets", wallet::Wallet, where: ("id", id))?) } Err(_) => Ok(raisfast_derive::crud_find_one!( - &mut *tx, "wallets", wallet::Wallet, "user_id" => user_id, - and: ["currency" => currency] + &mut *tx, "wallets", wallet::Wallet, where: AND(("user_id", user_id), ("currency", currency)) )?), } } @@ -203,7 +201,7 @@ async fn tx_find_tx_by_id( tx: &mut DbConnection, id: SnowflakeId, ) -> AppResult> { - Ok(raisfast_derive::crud_find!(tx, "wallet_transactions", WalletTransaction, "id" => id)?) + Ok(raisfast_derive::crud_find!(tx, "wallet_transactions", WalletTransaction, where: ("id", id))?) } async fn tx_find_tx_by_transaction_no( @@ -211,14 +209,13 @@ async fn tx_find_tx_by_transaction_no( transaction_no: &str, ) -> AppResult> { Ok( - raisfast_derive::crud_find!(tx, "wallet_transactions", WalletTransaction, "transaction_no" => transaction_no)?, + raisfast_derive::crud_find!(tx, "wallet_transactions", WalletTransaction, where: ("transaction_no", transaction_no))?, ) } async fn tx_has_reversal_for(tx: &mut DbConnection, related_tx_id: SnowflakeId) -> AppResult { Ok(raisfast_derive::crud_exists!( - tx, "wallet_transactions", "related_tx_id" => related_tx_id, - and: ["tx_type" => WalletTxType::Refund] + tx, "wallet_transactions", where: AND(("related_tx_id", related_tx_id), ("tx_type", WalletTxType::Refund)) )?) } @@ -684,7 +681,7 @@ async fn insert_tx( ["id" => id, "wallet_id" => wallet_id, "user_id" => user_id, "entry_type" => entry_type, "amount" => amount, "balance_after" => balance_after, "tx_type" => tx_type, "currency" => currency, "transaction_no" => transaction_no, "related_tx_id" => related_tx_id, "reference_type" => reference_type, "reference_id" => reference_id, "counterparty_wallet_id" => counterparty_wallet_id, "metadata" => metadata, "created_at" => now] )?; - let row = raisfast_derive::crud_find_one!(&mut *tx, "wallet_transactions", WalletTransaction, "id" => id)?; + let row = raisfast_derive::crud_find_one!(&mut *tx, "wallet_transactions", WalletTransaction, where: ("id", id))?; Ok(row) } diff --git a/src/webhook/model.rs b/src/webhook/model.rs index 25b6e42c..bce60ea0 100644 --- a/src/webhook/model.rs +++ b/src/webhook/model.rs @@ -74,9 +74,8 @@ pub async fn find_paginated( ) -> AppResult<(Vec, i64)> { let result = raisfast_derive::crud_query_paged!( pool, WebhookSubscription, - data_sql: "SELECT * FROM webhook_subscriptions WHERE 1=1{tenant} ORDER BY created_at DESC", - count_sql: "SELECT COUNT(*) FROM webhook_subscriptions WHERE 1=1{tenant}", - binds: [], + table: "webhook_subscriptions", + order_by: "created_at DESC", tenant: tenant_id, page: page, page_size: page_size @@ -85,7 +84,7 @@ pub async fn find_paginated( } pub async fn find_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult { - raisfast_derive::crud_find_one!(pool, "webhook_subscriptions", WebhookSubscription, "id" => id) + raisfast_derive::crud_find_one!(pool, "webhook_subscriptions", WebhookSubscription, where: ("id", id)) .map_err(Into::into) } @@ -94,14 +93,14 @@ pub async fn update(pool: &crate::db::Pool, sub: &WebhookSubscription) -> AppRes let result = raisfast_derive::crud_update!( pool, "webhook_subscriptions", bind: ["url" => &sub.url, "secret" => &sub.secret, "events" => &sub.events, "enabled" => sub.enabled, "description" => &sub.description, "updated_at" => now], - where: "id" => sub.id + where: ("id", sub.id) )?; AppError::expect_affected(&result, "webhook_subscription")?; Ok(()) } pub async fn delete_by_id(pool: &crate::db::Pool, id: SnowflakeId) -> AppResult<()> { - let result = raisfast_derive::crud_delete!(pool, "webhook_subscriptions", "id" => id)?; + let result = raisfast_derive::crud_delete!(pool, "webhook_subscriptions", where: ("id", id))?; AppError::expect_affected(&result, "webhook_subscription")?; Ok(()) } @@ -111,6 +110,6 @@ pub async fn find_enabled_by_tenant( tenant_id: Option<&str>, ) -> AppResult> { Ok( - raisfast_derive::crud_find_all!(pool, "webhook_subscriptions", WebhookSubscription, "enabled" => true, tenant: tenant_id)?, + raisfast_derive::crud_find_all!(pool, "webhook_subscriptions", WebhookSubscription, where: ("enabled", true), tenant: tenant_id)?, ) } diff --git a/src/worker/job_queue.rs b/src/worker/job_queue.rs index 0a2d156a..78d944ed 100644 --- a/src/worker/job_queue.rs +++ b/src/worker/job_queue.rs @@ -115,7 +115,7 @@ impl JobQueue for DefaultJobQueue { .map_err(|e| AppError::Internal(anyhow::anyhow!("invalid id: {e}")))?; raisfast_derive::crud_update!(&self.pool, "jobs", bind: ["status" => JobStatus::Completed, "updated_at" => now], - where: "id" => id + where: ("id", id) )?; tracing::debug!("job {id} completed"); Ok(()) @@ -144,7 +144,7 @@ impl JobQueue for DefaultJobQueue { if attempts >= max_attempts { raisfast_derive::crud_update!(&mut *tx, "jobs", bind: ["status" => JobStatus::Dead, "error" => error, "updated_at" => now], - where: "id" => id + where: ("id", id) )?; tracing::error!("job {id} dead: {error}"); return Ok::<_, AppError>(()); @@ -156,7 +156,7 @@ impl JobQueue for DefaultJobQueue { raisfast_derive::crud_update!(&mut *tx, "jobs", bind: ["status" => JobStatus::Pending, "error" => error, "run_after" => run_after, "updated_at" => now], - where: "id" => id + where: ("id", id) )?; tracing::warn!( @@ -173,7 +173,7 @@ impl JobQueue for DefaultJobQueue { .map_err(|e| AppError::Internal(anyhow::anyhow!("invalid id: {e}")))?; raisfast_derive::crud_update!(&self.pool, "jobs", bind: ["status" => JobStatus::Dead, "error" => error, "updated_at" => now], - where: "id" => id + where: ("id", id) )?; tracing::error!("job {id} dead: {error}"); Ok(()) @@ -314,8 +314,7 @@ impl JobQueue for DefaultJobQueue { "run_after" => None::, "updated_at" => now ], - where: "id" => id, - and: ["status" => JobStatus::Dead] + where: AND(("id", id), ("status", JobStatus::Dead)) )?; if result.rows_affected() == 0 { @@ -331,7 +330,7 @@ impl JobQueue for DefaultJobQueue { .parse() .map_err(|e| AppError::Internal(anyhow::anyhow!("invalid id: {e}")))?; let result: crate::db::DbQueryResult = - raisfast_derive::crud_delete!(&self.pool, "jobs", "id" => id)?; + raisfast_derive::crud_delete!(&self.pool, "jobs", where: ("id", id))?; if result.rows_affected() == 0 { return Err(AppError::not_found("job")); diff --git a/src/worker/scheduler.rs b/src/worker/scheduler.rs index d3d84aa3..c6f94d24 100644 --- a/src/worker/scheduler.rs +++ b/src/worker/scheduler.rs @@ -235,7 +235,7 @@ pub async fn toggle_schedule(pool: &Pool, id: SnowflakeId, enabled: bool) -> App let now = crate::utils::tz::now_utc(); let result = raisfast_derive::crud_update!(pool, "cron_schedules", bind: ["enabled" => enabled, "updated_at" => now], - where: "id" => id + where: ("id", id) )?; if result.rows_affected() == 0 { @@ -289,7 +289,7 @@ pub async fn update_schedule( "next_run_at" => next, "updated_at" => now ], - where: "id" => id + where: ("id", id) )?; Ok(find_by_id(pool, id).await?.unwrap_or(schedule)) @@ -298,7 +298,7 @@ pub async fn update_schedule( /// Delete a schedule pub async fn delete_schedule(pool: &Pool, id: SnowflakeId) -> AppResult<()> { let result: crate::db::DbQueryResult = - raisfast_derive::crud_delete!(pool, "cron_schedules", "id" => id)?; + raisfast_derive::crud_delete!(pool, "cron_schedules", where: ("id", id))?; if result.rows_affected() == 0 { return Err(AppError::not_found("cron_schedule")); @@ -404,7 +404,7 @@ impl CronScheduler { if let Some(ref lid) = log_id { raisfast_derive::crud_update!(&mut *tx, "cron_execution_log", bind: ["status" => CronExecStatus::Completed, "duration_ms" => elapsed, "finished_at" => now], - where: "id" => lid + where: ("id", lid) )?; } } @@ -413,7 +413,7 @@ impl CronScheduler { let err_str = e.to_string(); raisfast_derive::crud_update!(&mut *tx, "cron_execution_log", bind: ["status" => CronExecStatus::Failed, "duration_ms" => elapsed, "error" => &err_str, "finished_at" => now], - where: "id" => lid + where: ("id", lid) )?; } tracing::error!("cron dispatch failed for '{}': {e}", schedule.label); @@ -423,7 +423,7 @@ impl CronScheduler { if let Some(next) = &next { raisfast_derive::crud_update!(&mut *tx, "cron_schedules", bind: ["last_run_at" => now, "next_run_at" => next, "updated_at" => now], - where: "id" => schedule.id + where: ("id", schedule.id) )?; } @@ -542,7 +542,7 @@ pub async fn complete_execution_log( let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(pool, "cron_execution_log", bind: ["status" => CronExecStatus::Completed, "duration_ms" => duration_ms, "finished_at" => now], - where: "id" => log_id + where: ("id", log_id) )?; Ok(()) } @@ -557,7 +557,7 @@ pub async fn fail_execution_log( let now = crate::utils::tz::now_utc(); raisfast_derive::crud_update!(pool, "cron_execution_log", bind: ["status" => CronExecStatus::Failed, "duration_ms" => duration_ms, "error" => error, "finished_at" => now], - where: "id" => log_id + where: ("id", log_id) )?; Ok(()) } @@ -647,7 +647,7 @@ pub async fn sync_plugin_crons( for row in &old { if !new_types.contains(&row.job_type.as_str()) { - raisfast_derive::crud_delete!(&mut *tx, "cron_schedules", "id" => row.id)?; + raisfast_derive::crud_delete!(&mut *tx, "cron_schedules", where: ("id", row.id))?; tracing::info!( "removed stale cron '{}' for plugin {plugin_id}", row.job_type @@ -671,7 +671,7 @@ pub async fn sync_plugin_crons( let next = next_run(&entry.cron_expr, crate::utils::tz::now_utc())?; raisfast_derive::crud_update!(&mut *tx, "cron_schedules", bind: ["label" => &entry.label, "payload" => &entry.payload, "cron_expr" => &entry.cron_expr, "enabled" => entry.enabled, "next_run_at" => next, "updated_at" => now], - where: "id" => existing_row.0 + where: ("id", existing_row.0) )?; tracing::debug!("updated cron '{}' for plugin {plugin_id}", entry.job_type); @@ -705,7 +705,7 @@ pub async fn sync_plugin_crons( /// Called when a plugin is unloaded. pub async fn remove_plugin_crons(pool: &Pool, plugin_id: &str) -> AppResult<()> { let result: crate::db::DbQueryResult = - raisfast_derive::crud_delete!(pool, "cron_schedules", "plugin_id" => plugin_id)?; + raisfast_derive::crud_delete!(pool, "cron_schedules", where: ("plugin_id", plugin_id))?; let count = result.rows_affected(); if count > 0 { diff --git a/src/workflow/model.rs b/src/workflow/model.rs index d475879a..271279f7 100644 --- a/src/workflow/model.rs +++ b/src/workflow/model.rs @@ -148,7 +148,7 @@ pub async fn get_definition( pool: &Pool, id: SnowflakeId, ) -> anyhow::Result> { - Ok(raisfast_derive::crud_find!(pool, "workflow_definitions", WorkflowDefinition, "id" => id)?) + Ok(raisfast_derive::crud_find!(pool, "workflow_definitions", WorkflowDefinition, where: ("id", id))?) } /// List all workflow definitions @@ -160,7 +160,7 @@ pub async fn list_definitions(pool: &Pool) -> anyhow::Result anyhow::Result<()> { - raisfast_derive::crud_delete!(pool, "workflow_definitions", "id" => id)?; + raisfast_derive::crud_delete!(pool, "workflow_definitions", where: ("id", id))?; Ok(()) } @@ -188,7 +188,7 @@ pub async fn get_instance( pool: &Pool, id: SnowflakeId, ) -> anyhow::Result> { - Ok(raisfast_derive::crud_find!(pool, "workflow_instances", WorkflowInstance, "id" => id)?) + Ok(raisfast_derive::crud_find!(pool, "workflow_instances", WorkflowInstance, where: ("id", id))?) } /// List workflow instances @@ -291,7 +291,7 @@ pub async fn create_step_log( pool, "workflow_step_logs", ["id" => id, "instance_id" => instance_id, "step_id" => step_id, "step_name" => step_name, "status" => WorkflowStepStatus::Running, "input" => input_str, "started_at" => now] )?; - raisfast_derive::crud_find_one!(pool, "workflow_step_logs", StepLog, "id" => id) + raisfast_derive::crud_find_one!(pool, "workflow_step_logs", StepLog, where: ("id", id)) .map_err(|e| anyhow::anyhow!("failed to fetch created step log: {e}")) } @@ -306,7 +306,7 @@ pub async fn complete_step_log( raisfast_derive::crud_update!( pool, "workflow_step_logs", bind: ["status" => WorkflowStepStatus::Completed, "output" => output_str, "completed_at" => now], - where: "id" => id + where: ("id", id) )?; Ok(()) } @@ -317,7 +317,7 @@ pub async fn fail_step_log(pool: &Pool, id: SnowflakeId, error: &str) -> anyhow: raisfast_derive::crud_update!( pool, "workflow_step_logs", bind: ["status" => WorkflowStepStatus::Failed, "error" => error, "completed_at" => now], - where: "id" => id + where: ("id", id) )?; Ok(()) } @@ -325,7 +325,7 @@ pub async fn fail_step_log(pool: &Pool, id: SnowflakeId, error: &str) -> anyhow: /// List step logs for an instance pub async fn list_step_logs(pool: &Pool, instance_id: SnowflakeId) -> anyhow::Result> { Ok( - raisfast_derive::crud_find_all!(pool, "workflow_step_logs", StepLog, "instance_id" => instance_id, order_by: "started_at ASC")?, + raisfast_derive::crud_find_all!(pool, "workflow_step_logs", StepLog, where: ("instance_id", instance_id), order_by: "started_at ASC")?, ) }