Skip to main content

qail_core/
wire.rs

1//! QAIL wire codecs for command transport.
2//!
3//! - Text codecs (`QAIL-CMD/1`, `QAIL-CMDS/1`) round-trip through canonical text.
4//! - Binary codec (`QWB2`) transports framed AST bytes directly.
5
6use crate::ast::Qail;
7
8const CMD_TEXT_MAGIC: &str = "QAIL-CMD/1";
9const CMDS_TEXT_MAGIC: &str = "QAIL-CMDS/1";
10const CMD_BIN_MAGIC: [u8; 4] = *b"QWB2";
11const CMD_BIN_LEGACY_MAGIC: [u8; 4] = *b"QWB1";
12
13/// Maximum allowed QWB2 payload size (bytes).
14pub const MAX_CMD_BINARY_PAYLOAD_BYTES: usize = 64 * 1024;
15const MAX_AST_DEPTH: usize = 64;
16const MAX_AST_NODES: usize = 16_384;
17const MAX_AST_COLLECTION_LEN: usize = 2_048;
18const MAX_CMD_TEXT_BATCH_COMMANDS: usize = MAX_AST_COLLECTION_LEN;
19const MAX_AST_STRING_LEN: usize = 32 * 1024;
20const MAX_AST_VECTOR_LEN: usize = 8_192;
21const MAX_AST_BINARY_VALUE_LEN: usize = 32 * 1024;
22
23/// Encode one command into versioned text wire format.
24pub fn encode_cmd_text(cmd: &Qail) -> String {
25    let payload = cmd.to_string();
26    let mut out = String::with_capacity(CMD_TEXT_MAGIC.len() + payload.len() + 32);
27    out.push_str(CMD_TEXT_MAGIC);
28    out.push('\n');
29    out.push_str(&payload.len().to_string());
30    out.push('\n');
31    out.push_str(&payload);
32    out
33}
34
35/// Decode one command from text wire format.
36///
37/// Also accepts raw QAIL query text as fallback for convenience.
38pub fn decode_cmd_text(input: &str) -> Result<Qail, String> {
39    let bytes = input.as_bytes();
40    let mut idx = 0usize;
41
42    let Ok(magic) = read_line(bytes, &mut idx) else {
43        return crate::parse(input).map_err(|e| e.to_string());
44    };
45
46    if magic != CMD_TEXT_MAGIC {
47        return crate::parse(input).map_err(|e| e.to_string());
48    }
49
50    let len_line = read_line(bytes, &mut idx)?;
51    let payload_len = parse_usize("payload length", len_line)?;
52    let payload = read_exact_utf8(bytes, &mut idx, payload_len)?;
53    if idx != bytes.len() {
54        return Err("trailing bytes after command payload".to_string());
55    }
56
57    crate::parse(payload).map_err(|e| e.to_string())
58}
59
60/// Encode multiple commands into versioned text wire format.
61pub fn encode_cmds_text(cmds: &[Qail]) -> String {
62    let mut out = String::new();
63    out.push_str(CMDS_TEXT_MAGIC);
64    out.push('\n');
65    out.push_str(&cmds.len().to_string());
66    out.push('\n');
67
68    for cmd in cmds {
69        let payload = cmd.to_string();
70        out.push_str(&payload.len().to_string());
71        out.push('\n');
72        out.push_str(&payload);
73    }
74
75    out
76}
77
78/// Decode multiple commands from text wire format.
79pub fn decode_cmds_text(input: &str) -> Result<Vec<Qail>, String> {
80    let bytes = input.as_bytes();
81    let mut idx = 0usize;
82
83    let magic = read_line(bytes, &mut idx)?;
84    if magic != CMDS_TEXT_MAGIC {
85        return Err(format!(
86            "invalid wire magic: expected {CMDS_TEXT_MAGIC}, got {magic}"
87        ));
88    }
89
90    let count_line = read_line(bytes, &mut idx)?;
91    let count = parse_usize("command count", count_line)?;
92    if count > MAX_CMD_TEXT_BATCH_COMMANDS {
93        return Err(format!(
94            "command count exceeds limit: {count} > {MAX_CMD_TEXT_BATCH_COMMANDS}"
95        ));
96    }
97    let mut out = Vec::with_capacity(count);
98
99    for _ in 0..count {
100        let len_line = read_line(bytes, &mut idx)?;
101        let payload_len = parse_usize("payload length", len_line)?;
102        let payload = read_exact_utf8(bytes, &mut idx, payload_len)?;
103        let cmd = crate::parse(payload).map_err(|e| e.to_string())?;
104        out.push(cmd);
105    }
106
107    if idx != bytes.len() {
108        return Err("trailing bytes after batch payload".to_string());
109    }
110
111    Ok(out)
112}
113
114/// Encode one command into compact binary wire format (QWB2 AST binary).
115pub fn encode_cmd_binary(cmd: &Qail) -> Result<Vec<u8>, String> {
116    validate_binary_ast_limits(cmd)?;
117    crate::sanitize::validate_ast(cmd).map_err(|e| e.to_string())?;
118
119    let payload = serde_json::to_vec(cmd).map_err(|e| format!("binary AST encode failed: {e}"))?;
120    if payload.len() > MAX_CMD_BINARY_PAYLOAD_BYTES {
121        return Err(format!(
122            "binary AST payload too large: {} bytes (max {})",
123            payload.len(),
124            MAX_CMD_BINARY_PAYLOAD_BYTES
125        ));
126    }
127
128    let payload_len = u32::try_from(payload.len())
129        .map_err(|_| format!("binary AST payload exceeds u32 length: {}", payload.len()))?;
130    let mut out = Vec::with_capacity(8 + payload.len());
131    out.extend_from_slice(&CMD_BIN_MAGIC);
132    out.extend_from_slice(&payload_len.to_be_bytes());
133    out.extend_from_slice(&payload);
134    Ok(out)
135}
136
137/// Decode one command from strict QWB2 AST-binary wire format.
138///
139/// This path rejects legacy QWB1/raw-text payloads.
140pub fn decode_cmd_binary(input: &[u8]) -> Result<Qail, String> {
141    let payload = decode_cmd_binary_payload(input)?;
142    let mut deserializer = serde_json::Deserializer::from_slice(payload);
143    let cmd = serde::Deserialize::deserialize(&mut deserializer)
144        .map_err(|e| format!("binary AST decode failed: {e}"))?;
145    deserializer
146        .end()
147        .map_err(|_| "trailing bytes after AST payload".to_string())?;
148    validate_binary_ast_limits(&cmd)?;
149    crate::sanitize::validate_ast(&cmd).map_err(|e| e.to_string())?;
150    Ok(cmd)
151}
152
153/// Decode and validate strict QWB2-framed payload bytes.
154///
155/// This validates framing and payload-size limits only.
156pub fn decode_cmd_binary_payload(input: &[u8]) -> Result<&[u8], String> {
157    if input.len() < 8 {
158        return Err("invalid wire header".to_string());
159    }
160    if input[0..4] != CMD_BIN_MAGIC {
161        if input[0..4] == CMD_BIN_LEGACY_MAGIC {
162            return Err(
163                "legacy QWB1 text payload is not supported on parse-free binary path".to_string(),
164            );
165        }
166        return Err("invalid wire header".to_string());
167    }
168
169    let len = u32::from_be_bytes([input[4], input[5], input[6], input[7]]) as usize;
170    if len > MAX_CMD_BINARY_PAYLOAD_BYTES {
171        return Err(format!(
172            "binary AST payload too large: header={len}, max={MAX_CMD_BINARY_PAYLOAD_BYTES}"
173        ));
174    }
175    if input.len() != 8 + len {
176        return Err(format!(
177            "invalid payload length: header={len}, actual={}",
178            input.len().saturating_sub(8)
179        ));
180    }
181    Ok(&input[8..])
182}
183
184#[derive(Default)]
185struct AstLimitState {
186    nodes: usize,
187}
188
189impl AstLimitState {
190    fn bump(&mut self, kind: &str) -> Result<(), String> {
191        self.nodes = self
192            .nodes
193            .checked_add(1)
194            .ok_or_else(|| "AST node counter overflow".to_string())?;
195        if self.nodes > MAX_AST_NODES {
196            return Err(format!(
197                "AST node limit exceeded while walking {kind}: {} > {}",
198                self.nodes, MAX_AST_NODES
199            ));
200        }
201        Ok(())
202    }
203}
204
205fn ensure_depth(depth: usize, kind: &str) -> Result<(), String> {
206    if depth > MAX_AST_DEPTH {
207        return Err(format!(
208            "AST depth limit exceeded while walking {kind}: {depth} > {MAX_AST_DEPTH}"
209        ));
210    }
211    Ok(())
212}
213
214fn ensure_len(kind: &str, len: usize, max: usize) -> Result<(), String> {
215    if len > max {
216        return Err(format!("{kind} exceeds limit: {len} > {max}"));
217    }
218    Ok(())
219}
220
221fn ensure_str(kind: &str, value: &str) -> Result<(), String> {
222    ensure_len(kind, value.len(), MAX_AST_STRING_LEN)
223}
224
225fn validate_binary_ast_limits(cmd: &Qail) -> Result<(), String> {
226    let mut state = AstLimitState::default();
227    validate_qail_limits(cmd, 0, &mut state)
228}
229
230fn validate_qail_limits(cmd: &Qail, depth: usize, state: &mut AstLimitState) -> Result<(), String> {
231    use crate::ast::GroupByMode;
232
233    ensure_depth(depth, "Qail")?;
234    state.bump("Qail")?;
235
236    ensure_str("qail.table", &cmd.table)?;
237    ensure_len("qail.columns", cmd.columns.len(), MAX_AST_COLLECTION_LEN)?;
238    for expr in &cmd.columns {
239        validate_expr_limits(expr, depth + 1, state)?;
240    }
241
242    ensure_len("qail.joins", cmd.joins.len(), MAX_AST_COLLECTION_LEN)?;
243    for join in &cmd.joins {
244        validate_join_limits(join, depth + 1, state)?;
245    }
246
247    ensure_len("qail.cages", cmd.cages.len(), MAX_AST_COLLECTION_LEN)?;
248    for cage in &cmd.cages {
249        validate_cage_limits(cage, depth + 1, state)?;
250    }
251
252    if let Some(index_def) = &cmd.index_def {
253        validate_index_def_limits(index_def)?;
254    }
255
256    ensure_len(
257        "qail.table_constraints",
258        cmd.table_constraints.len(),
259        MAX_AST_COLLECTION_LEN,
260    )?;
261    for constraint in &cmd.table_constraints {
262        match constraint {
263            crate::ast::TableConstraint::Unique(cols)
264            | crate::ast::TableConstraint::PrimaryKey(cols) => {
265                ensure_len(
266                    "qail.table_constraint.columns",
267                    cols.len(),
268                    MAX_AST_COLLECTION_LEN,
269                )?;
270                for col in cols {
271                    ensure_str("qail.table_constraint.column", col)?;
272                }
273            }
274            crate::ast::TableConstraint::ForeignKey {
275                name,
276                columns,
277                ref_table,
278                ref_columns,
279                on_delete,
280                on_update,
281                deferrable,
282            } => {
283                if let Some(name) = name {
284                    ensure_str("qail.table_constraint.name", name)?;
285                }
286                ensure_len(
287                    "qail.table_constraint.columns",
288                    columns.len(),
289                    MAX_AST_COLLECTION_LEN,
290                )?;
291                for col in columns {
292                    ensure_str("qail.table_constraint.column", col)?;
293                }
294                ensure_str("qail.table_constraint.ref_table", ref_table)?;
295                ensure_len(
296                    "qail.table_constraint.ref_columns",
297                    ref_columns.len(),
298                    MAX_AST_COLLECTION_LEN,
299                )?;
300                for col in ref_columns {
301                    ensure_str("qail.table_constraint.ref_column", col)?;
302                }
303                if let Some(action) = on_delete {
304                    ensure_str("qail.table_constraint.on_delete", action)?;
305                }
306                if let Some(action) = on_update {
307                    ensure_str("qail.table_constraint.on_update", action)?;
308                }
309                if let Some(clause) = deferrable {
310                    ensure_str("qail.table_constraint.deferrable", clause)?;
311                }
312            }
313        }
314    }
315
316    ensure_len("qail.set_ops", cmd.set_ops.len(), MAX_AST_COLLECTION_LEN)?;
317    for (_, rhs) in &cmd.set_ops {
318        validate_qail_limits(rhs, depth + 1, state)?;
319    }
320
321    ensure_len("qail.having", cmd.having.len(), MAX_AST_COLLECTION_LEN)?;
322    for cond in &cmd.having {
323        validate_condition_limits(cond, depth + 1, state)?;
324    }
325
326    if let GroupByMode::GroupingSets(groups) = &cmd.group_by_mode {
327        ensure_len("qail.grouping_sets", groups.len(), MAX_AST_COLLECTION_LEN)?;
328        for group in groups {
329            ensure_len("qail.grouping_set", group.len(), MAX_AST_COLLECTION_LEN)?;
330            for col in group {
331                ensure_str("qail.grouping_set.column", col)?;
332            }
333        }
334    }
335
336    ensure_len("qail.ctes", cmd.ctes.len(), MAX_AST_COLLECTION_LEN)?;
337    for cte in &cmd.ctes {
338        ensure_str("qail.cte.name", &cte.name)?;
339        ensure_len(
340            "qail.cte.columns",
341            cte.columns.len(),
342            MAX_AST_COLLECTION_LEN,
343        )?;
344        for col in &cte.columns {
345            ensure_str("qail.cte.column", col)?;
346        }
347        validate_qail_limits(&cte.base_query, depth + 1, state)?;
348        if let Some(recursive) = &cte.recursive_query {
349            validate_qail_limits(recursive, depth + 1, state)?;
350        }
351        if let Some(source_table) = &cte.source_table {
352            ensure_str("qail.cte.source_table", source_table)?;
353        }
354    }
355
356    ensure_len(
357        "qail.distinct_on",
358        cmd.distinct_on.len(),
359        MAX_AST_COLLECTION_LEN,
360    )?;
361    for expr in &cmd.distinct_on {
362        validate_expr_limits(expr, depth + 1, state)?;
363    }
364
365    if let Some(returning) = &cmd.returning {
366        ensure_len("qail.returning", returning.len(), MAX_AST_COLLECTION_LEN)?;
367        for expr in returning {
368            validate_expr_limits(expr, depth + 1, state)?;
369        }
370    }
371
372    if let Some(on_conflict) = &cmd.on_conflict {
373        ensure_len(
374            "qail.on_conflict.columns",
375            on_conflict.columns.len(),
376            MAX_AST_COLLECTION_LEN,
377        )?;
378        for col in &on_conflict.columns {
379            ensure_str("qail.on_conflict.column", col)?;
380        }
381        if let Some(assignments) = on_conflict.action.update_assignments() {
382            ensure_len(
383                "qail.on_conflict.assignments",
384                assignments.len(),
385                MAX_AST_COLLECTION_LEN,
386            )?;
387            for (col, expr) in assignments {
388                ensure_str("qail.on_conflict.assignment.column", col)?;
389                validate_expr_limits(expr, depth + 1, state)?;
390            }
391        }
392        ensure_len(
393            "qail.on_conflict.where_conditions",
394            on_conflict.where_conditions.len(),
395            MAX_AST_COLLECTION_LEN,
396        )?;
397        for cond in &on_conflict.where_conditions {
398            validate_condition_limits(cond, depth + 1, state)?;
399        }
400    }
401
402    if let Some(merge) = &cmd.merge {
403        if let Some(alias) = &merge.target_alias {
404            ensure_str("qail.merge.target_alias", alias)?;
405        }
406        match &merge.source {
407            crate::ast::MergeSource::Table { name, alias } => {
408                ensure_str("qail.merge.source.table", name)?;
409                if let Some(alias) = alias {
410                    ensure_str("qail.merge.source.alias", alias)?;
411                }
412            }
413            crate::ast::MergeSource::Query { query, alias } => {
414                validate_qail_limits(query, depth + 1, state)?;
415                if let Some(alias) = alias {
416                    ensure_str("qail.merge.source.alias", alias)?;
417                }
418            }
419        }
420        ensure_len("qail.merge.on", merge.on.len(), MAX_AST_COLLECTION_LEN)?;
421        for condition in &merge.on {
422            validate_condition_limits(condition, depth + 1, state)?;
423        }
424        ensure_len(
425            "qail.merge.clauses",
426            merge.clauses.len(),
427            MAX_AST_COLLECTION_LEN,
428        )?;
429        for clause in &merge.clauses {
430            ensure_len(
431                "qail.merge.clause.condition",
432                clause.condition.len(),
433                MAX_AST_COLLECTION_LEN,
434            )?;
435            for condition in &clause.condition {
436                validate_condition_limits(condition, depth + 1, state)?;
437            }
438            match &clause.action {
439                crate::ast::MergeAction::Update { assignments } => {
440                    ensure_len(
441                        "qail.merge.update.assignments",
442                        assignments.len(),
443                        MAX_AST_COLLECTION_LEN,
444                    )?;
445                    for (col, expr) in assignments {
446                        ensure_str("qail.merge.update.column", col)?;
447                        validate_expr_limits(expr, depth + 1, state)?;
448                    }
449                }
450                crate::ast::MergeAction::Insert { columns, values } => {
451                    ensure_len(
452                        "qail.merge.insert.columns",
453                        columns.len(),
454                        MAX_AST_COLLECTION_LEN,
455                    )?;
456                    for col in columns {
457                        ensure_str("qail.merge.insert.column", col)?;
458                    }
459                    ensure_len(
460                        "qail.merge.insert.values",
461                        values.len(),
462                        MAX_AST_COLLECTION_LEN,
463                    )?;
464                    for expr in values {
465                        validate_expr_limits(expr, depth + 1, state)?;
466                    }
467                }
468                crate::ast::MergeAction::Delete | crate::ast::MergeAction::DoNothing => {}
469            }
470        }
471    }
472
473    if let Some(source_query) = &cmd.source_query {
474        validate_qail_limits(source_query, depth + 1, state)?;
475    }
476
477    if let Some(channel) = &cmd.channel {
478        ensure_str("qail.channel", channel)?;
479    }
480    if let Some(payload) = &cmd.payload {
481        ensure_str("qail.payload", payload)?;
482    }
483    if let Some(savepoint_name) = &cmd.savepoint_name {
484        ensure_str("qail.savepoint_name", savepoint_name)?;
485    }
486
487    ensure_len(
488        "qail.from_tables",
489        cmd.from_tables.len(),
490        MAX_AST_COLLECTION_LEN,
491    )?;
492    for table in &cmd.from_tables {
493        ensure_str("qail.from_table", table)?;
494    }
495
496    ensure_len(
497        "qail.using_tables",
498        cmd.using_tables.len(),
499        MAX_AST_COLLECTION_LEN,
500    )?;
501    for table in &cmd.using_tables {
502        ensure_str("qail.using_table", table)?;
503    }
504
505    if let Some((_, percent, _seed)) = cmd.sample
506        && !percent.is_finite()
507    {
508        return Err("qail.sample.percent must be finite".to_string());
509    }
510
511    if let Some(vector) = &cmd.vector {
512        ensure_len("qail.vector", vector.len(), MAX_AST_VECTOR_LEN)?;
513    }
514    if let Some(vector_name) = &cmd.vector_name {
515        ensure_str("qail.vector_name", vector_name)?;
516    }
517    if let Some(function_def) = &cmd.function_def {
518        validate_function_def_limits(function_def)?;
519    }
520    if let Some(trigger_def) = &cmd.trigger_def {
521        validate_trigger_def_limits(trigger_def)?;
522    }
523    if let Some(policy_def) = &cmd.policy_def {
524        validate_policy_def_limits(policy_def, depth + 1, state)?;
525    }
526
527    Ok(())
528}
529
530fn validate_join_limits(
531    join: &crate::ast::Join,
532    depth: usize,
533    state: &mut AstLimitState,
534) -> Result<(), String> {
535    ensure_depth(depth, "Join")?;
536    state.bump("Join")?;
537    ensure_str("join.table", &join.table)?;
538    if let Some(on) = &join.on {
539        ensure_len("join.on", on.len(), MAX_AST_COLLECTION_LEN)?;
540        for cond in on {
541            validate_condition_limits(cond, depth + 1, state)?;
542        }
543    }
544    Ok(())
545}
546
547fn validate_cage_limits(
548    cage: &crate::ast::Cage,
549    depth: usize,
550    state: &mut AstLimitState,
551) -> Result<(), String> {
552    use crate::ast::CageKind;
553
554    ensure_depth(depth, "Cage")?;
555    state.bump("Cage")?;
556    ensure_len(
557        "cage.conditions",
558        cage.conditions.len(),
559        MAX_AST_COLLECTION_LEN,
560    )?;
561    for cond in &cage.conditions {
562        validate_condition_limits(cond, depth + 1, state)?;
563    }
564    match cage.kind {
565        CageKind::Limit(v) | CageKind::Offset(v) | CageKind::Sample(v) => {
566            ensure_len("cage.numeric", v, usize::MAX)?;
567        }
568        _ => {}
569    }
570    Ok(())
571}
572
573fn validate_condition_limits(
574    cond: &crate::ast::Condition,
575    depth: usize,
576    state: &mut AstLimitState,
577) -> Result<(), String> {
578    ensure_depth(depth, "Condition")?;
579    state.bump("Condition")?;
580    validate_expr_limits(&cond.left, depth + 1, state)?;
581    validate_value_limits(&cond.value, depth + 1, state)
582}
583
584fn validate_expr_limits(
585    expr: &crate::ast::Expr,
586    depth: usize,
587    state: &mut AstLimitState,
588) -> Result<(), String> {
589    use crate::ast::{ColumnGeneration, Constraint, Expr, WindowFrame};
590
591    ensure_depth(depth, "Expr")?;
592    state.bump("Expr")?;
593
594    match expr {
595        Expr::Star => {}
596        Expr::Named(name) => ensure_str("expr.named", name)?,
597        Expr::Aliased { name, alias } => {
598            ensure_str("expr.aliased.name", name)?;
599            ensure_str("expr.aliased.alias", alias)?;
600        }
601        Expr::Aggregate {
602            col, filter, alias, ..
603        } => {
604            ensure_str("expr.aggregate.col", col)?;
605            if let Some(filters) = filter {
606                ensure_len(
607                    "expr.aggregate.filter",
608                    filters.len(),
609                    MAX_AST_COLLECTION_LEN,
610                )?;
611                for cond in filters {
612                    validate_condition_limits(cond, depth + 1, state)?;
613                }
614            }
615            if let Some(alias) = alias {
616                ensure_str("expr.aggregate.alias", alias)?;
617            }
618        }
619        Expr::Cast {
620            expr,
621            target_type,
622            alias,
623        } => {
624            validate_expr_limits(expr, depth + 1, state)?;
625            ensure_str("expr.cast.target_type", target_type)?;
626            if let Some(alias) = alias {
627                ensure_str("expr.cast.alias", alias)?;
628            }
629        }
630        Expr::Def {
631            name,
632            data_type,
633            constraints,
634        } => {
635            ensure_str("expr.def.name", name)?;
636            ensure_str("expr.def.data_type", data_type)?;
637            ensure_len(
638                "expr.def.constraints",
639                constraints.len(),
640                MAX_AST_COLLECTION_LEN,
641            )?;
642            for constraint in constraints {
643                match constraint {
644                    Constraint::PrimaryKey | Constraint::Unique | Constraint::Nullable => {}
645                    Constraint::Default(v) => ensure_str("expr.def.default", v)?,
646                    Constraint::Check(values) => {
647                        ensure_len("expr.def.check", values.len(), MAX_AST_COLLECTION_LEN)?;
648                        for value in values {
649                            ensure_str("expr.def.check.value", value)?;
650                        }
651                    }
652                    Constraint::Comment(v) | Constraint::References(v) => {
653                        ensure_str("expr.def.constraint", v)?;
654                    }
655                    Constraint::Generated(ColumnGeneration::Stored(v))
656                    | Constraint::Generated(ColumnGeneration::Virtual(v)) => {
657                        ensure_str("expr.def.generated", v)?;
658                    }
659                }
660            }
661        }
662        Expr::Mod { col, .. } => validate_expr_limits(col, depth + 1, state)?,
663        Expr::Window {
664            name,
665            func,
666            params,
667            partition,
668            order,
669            frame,
670        } => {
671            ensure_str("expr.window.name", name)?;
672            ensure_str("expr.window.func", func)?;
673            ensure_len("expr.window.params", params.len(), MAX_AST_COLLECTION_LEN)?;
674            for param in params {
675                validate_expr_limits(param, depth + 1, state)?;
676            }
677            ensure_len(
678                "expr.window.partition",
679                partition.len(),
680                MAX_AST_COLLECTION_LEN,
681            )?;
682            for col in partition {
683                ensure_str("expr.window.partition.column", col)?;
684            }
685            ensure_len("expr.window.order", order.len(), MAX_AST_COLLECTION_LEN)?;
686            for cage in order {
687                validate_cage_limits(cage, depth + 1, state)?;
688            }
689            if let Some(frame) = frame {
690                match frame {
691                    WindowFrame::Rows { .. } | WindowFrame::Range { .. } => {}
692                }
693            }
694        }
695        Expr::Case {
696            when_clauses,
697            else_value,
698            alias,
699        } => {
700            ensure_len("expr.case.when", when_clauses.len(), MAX_AST_COLLECTION_LEN)?;
701            for (cond, then_expr) in when_clauses {
702                validate_condition_limits(cond, depth + 1, state)?;
703                validate_expr_limits(then_expr, depth + 1, state)?;
704            }
705            if let Some(else_expr) = else_value {
706                validate_expr_limits(else_expr, depth + 1, state)?;
707            }
708            if let Some(alias) = alias {
709                ensure_str("expr.case.alias", alias)?;
710            }
711        }
712        Expr::JsonAccess {
713            column,
714            path_segments,
715            alias,
716        } => {
717            ensure_str("expr.json_access.column", column)?;
718            ensure_len(
719                "expr.json_access.path_segments",
720                path_segments.len(),
721                MAX_AST_COLLECTION_LEN,
722            )?;
723            for (segment, _) in path_segments {
724                ensure_str("expr.json_access.segment", segment)?;
725            }
726            if let Some(alias) = alias {
727                ensure_str("expr.json_access.alias", alias)?;
728            }
729        }
730        Expr::FunctionCall { name, args, alias } => {
731            ensure_str("expr.function_call.name", name)?;
732            ensure_len(
733                "expr.function_call.args",
734                args.len(),
735                MAX_AST_COLLECTION_LEN,
736            )?;
737            for arg in args {
738                validate_expr_limits(arg, depth + 1, state)?;
739            }
740            if let Some(alias) = alias {
741                ensure_str("expr.function_call.alias", alias)?;
742            }
743        }
744        Expr::SpecialFunction { name, args, alias } => {
745            ensure_str("expr.special_function.name", name)?;
746            ensure_len(
747                "expr.special_function.args",
748                args.len(),
749                MAX_AST_COLLECTION_LEN,
750            )?;
751            for (keyword, arg) in args {
752                if let Some(keyword) = keyword {
753                    ensure_str("expr.special_function.keyword", keyword)?;
754                }
755                validate_expr_limits(arg, depth + 1, state)?;
756            }
757            if let Some(alias) = alias {
758                ensure_str("expr.special_function.alias", alias)?;
759            }
760        }
761        Expr::Binary {
762            left, right, alias, ..
763        } => {
764            validate_expr_limits(left, depth + 1, state)?;
765            validate_expr_limits(right, depth + 1, state)?;
766            if let Some(alias) = alias {
767                ensure_str("expr.binary.alias", alias)?;
768            }
769        }
770        Expr::Literal(v) => validate_value_limits(v, depth + 1, state)?,
771        Expr::ArrayConstructor { elements, alias } | Expr::RowConstructor { elements, alias } => {
772            ensure_len("expr.elements", elements.len(), MAX_AST_COLLECTION_LEN)?;
773            for el in elements {
774                validate_expr_limits(el, depth + 1, state)?;
775            }
776            if let Some(alias) = alias {
777                ensure_str("expr.elements.alias", alias)?;
778            }
779        }
780        Expr::Subscript { expr, index, alias } => {
781            validate_expr_limits(expr, depth + 1, state)?;
782            validate_expr_limits(index, depth + 1, state)?;
783            if let Some(alias) = alias {
784                ensure_str("expr.subscript.alias", alias)?;
785            }
786        }
787        Expr::Collate {
788            expr,
789            collation,
790            alias,
791        } => {
792            validate_expr_limits(expr, depth + 1, state)?;
793            ensure_str("expr.collate.collation", collation)?;
794            if let Some(alias) = alias {
795                ensure_str("expr.collate.alias", alias)?;
796            }
797        }
798        Expr::FieldAccess { expr, field, alias } => {
799            validate_expr_limits(expr, depth + 1, state)?;
800            ensure_str("expr.field_access.field", field)?;
801            if let Some(alias) = alias {
802                ensure_str("expr.field_access.alias", alias)?;
803            }
804        }
805        Expr::Subquery { query, alias } => {
806            validate_qail_limits(query, depth + 1, state)?;
807            if let Some(alias) = alias {
808                ensure_str("expr.subquery.alias", alias)?;
809            }
810        }
811        Expr::Exists { query, alias, .. } => {
812            validate_qail_limits(query, depth + 1, state)?;
813            if let Some(alias) = alias {
814                ensure_str("expr.exists.alias", alias)?;
815            }
816        }
817    }
818
819    Ok(())
820}
821
822fn validate_value_limits(
823    value: &crate::ast::Value,
824    depth: usize,
825    state: &mut AstLimitState,
826) -> Result<(), String> {
827    use crate::ast::Value;
828
829    ensure_depth(depth, "Value")?;
830    state.bump("Value")?;
831
832    match value {
833        Value::Null | Value::Bool(_) | Value::Int(_) | Value::Float(_) | Value::Param(_) => {}
834        Value::String(v)
835        | Value::NamedParam(v)
836        | Value::Function(v)
837        | Value::Column(v)
838        | Value::Timestamp(v)
839        | Value::Json(v) => ensure_str("value.string", v)?,
840        Value::Array(values) => {
841            ensure_len("value.array", values.len(), MAX_AST_COLLECTION_LEN)?;
842            for v in values {
843                validate_value_limits(v, depth + 1, state)?;
844            }
845        }
846        Value::Subquery(q) => validate_qail_limits(q, depth + 1, state)?,
847        Value::Uuid(_) | Value::NullUuid | Value::Interval { .. } => {}
848        Value::Bytes(bytes) => ensure_len("value.bytes", bytes.len(), MAX_AST_BINARY_VALUE_LEN)?,
849        Value::Expr(expr) => validate_expr_limits(expr, depth + 1, state)?,
850        Value::Vector(values) => ensure_len("value.vector", values.len(), MAX_AST_VECTOR_LEN)?,
851    }
852
853    Ok(())
854}
855
856fn validate_index_def_limits(index_def: &crate::ast::IndexDef) -> Result<(), String> {
857    ensure_str("index_def.name", &index_def.name)?;
858    ensure_str("index_def.table", &index_def.table)?;
859    ensure_len(
860        "index_def.columns",
861        index_def.columns.len(),
862        MAX_AST_COLLECTION_LEN,
863    )?;
864    for col in &index_def.columns {
865        ensure_str("index_def.column", col)?;
866    }
867    if let Some(index_type) = &index_def.index_type {
868        ensure_str("index_def.index_type", index_type)?;
869    }
870    if let Some(where_clause) = &index_def.where_clause {
871        ensure_str("index_def.where_clause", where_clause)?;
872    }
873    Ok(())
874}
875
876fn validate_function_def_limits(function_def: &crate::ast::FunctionDef) -> Result<(), String> {
877    ensure_str("function_def.name", &function_def.name)?;
878    ensure_len(
879        "function_def.args",
880        function_def.args.len(),
881        MAX_AST_COLLECTION_LEN,
882    )?;
883    for arg in &function_def.args {
884        ensure_str("function_def.arg", arg)?;
885    }
886    ensure_str("function_def.returns", &function_def.returns)?;
887    ensure_str("function_def.body", &function_def.body)?;
888    if let Some(language) = &function_def.language {
889        ensure_str("function_def.language", language)?;
890    }
891    if let Some(volatility) = &function_def.volatility {
892        ensure_str("function_def.volatility", volatility)?;
893    }
894    Ok(())
895}
896
897fn validate_trigger_def_limits(trigger_def: &crate::ast::TriggerDef) -> Result<(), String> {
898    ensure_str("trigger_def.name", &trigger_def.name)?;
899    ensure_str("trigger_def.table", &trigger_def.table)?;
900    ensure_len(
901        "trigger_def.events",
902        trigger_def.events.len(),
903        MAX_AST_COLLECTION_LEN,
904    )?;
905    ensure_len(
906        "trigger_def.update_columns",
907        trigger_def.update_columns.len(),
908        MAX_AST_COLLECTION_LEN,
909    )?;
910    for col in &trigger_def.update_columns {
911        ensure_str("trigger_def.update_column", col)?;
912    }
913    ensure_str(
914        "trigger_def.execute_function",
915        &trigger_def.execute_function,
916    )?;
917    Ok(())
918}
919
920fn validate_policy_def_limits(
921    policy_def: &crate::migrate::policy::RlsPolicy,
922    depth: usize,
923    state: &mut AstLimitState,
924) -> Result<(), String> {
925    ensure_str("policy_def.name", &policy_def.name)?;
926    ensure_str("policy_def.table", &policy_def.table)?;
927    if let Some(using_expr) = &policy_def.using {
928        validate_expr_limits(using_expr, depth + 1, state)?;
929    }
930    if let Some(with_check_expr) = &policy_def.with_check {
931        validate_expr_limits(with_check_expr, depth + 1, state)?;
932    }
933    if let Some(role) = &policy_def.role {
934        ensure_str("policy_def.role", role)?;
935    }
936    Ok(())
937}
938
939fn read_line<'a>(bytes: &'a [u8], idx: &mut usize) -> Result<&'a str, String> {
940    if *idx >= bytes.len() {
941        return Err("unexpected EOF".to_string());
942    }
943
944    let start = *idx;
945    while *idx < bytes.len() && bytes[*idx] != b'\n' {
946        *idx += 1;
947    }
948
949    if *idx >= bytes.len() {
950        return Err("unterminated header line".to_string());
951    }
952
953    let line =
954        std::str::from_utf8(&bytes[start..*idx]).map_err(|_| "header is not UTF-8".to_string())?;
955    *idx += 1; // consume '\n'
956    Ok(line)
957}
958
959fn parse_usize(field: &str, line: &str) -> Result<usize, String> {
960    line.parse::<usize>()
961        .map_err(|_| format!("invalid {field}: {line}"))
962}
963
964fn read_exact_utf8<'a>(bytes: &'a [u8], idx: &mut usize, len: usize) -> Result<&'a str, String> {
965    let end = (*idx)
966        .checked_add(len)
967        .ok_or_else(|| "payload length overflow".to_string())?;
968    if end > bytes.len() {
969        return Err("payload truncated".to_string());
970    }
971    let start = *idx;
972    *idx = end;
973    std::str::from_utf8(&bytes[start..end]).map_err(|_| "payload is not UTF-8".to_string())
974}
975
976#[cfg(test)]
977mod tests {
978    use super::*;
979    use proptest::prelude::*;
980
981    fn encode_cmd_binary_unchecked_for_test(cmd: &Qail) -> Vec<u8> {
982        let payload = serde_json::to_vec(cmd).expect("test binary AST encode");
983        let mut out = Vec::with_capacity(8 + payload.len());
984        out.extend_from_slice(&CMD_BIN_MAGIC);
985        out.extend_from_slice(&(payload.len() as u32).to_be_bytes());
986        out.extend_from_slice(&payload);
987        out
988    }
989
990    #[test]
991    fn cmd_text_roundtrip() {
992        let cmd = crate::ast::Qail::get("users")
993            .columns(["id", "email"])
994            .where_eq("active", true)
995            .limit(10);
996
997        let encoded = encode_cmd_text(&cmd);
998        let decoded = decode_cmd_text(&encoded).unwrap();
999        assert_eq!(decoded.to_string(), cmd.to_string());
1000    }
1001
1002    #[test]
1003    fn cmd_binary_roundtrip() {
1004        let cmd = crate::ast::Qail::set("users")
1005            .set_value("active", true)
1006            .where_eq("id", 7);
1007
1008        let encoded = encode_cmd_binary(&cmd).expect("binary encode");
1009        let decoded = decode_cmd_binary(&encoded).unwrap();
1010        assert_eq!(decoded.to_string(), cmd.to_string());
1011    }
1012
1013    #[test]
1014    fn cmd_binary_roundtrip_preserves_merge_ast() {
1015        let cmd = crate::ast::Qail::merge_into("users")
1016            .target_alias("u")
1017            .using_table_as("staging_users", "s")
1018            .merge_on_column("u.id", crate::ast::Operator::Eq, "s.id")
1019            .when_matched_update(&[("name", crate::ast::Expr::Named("s.name".to_string()))])
1020            .when_not_matched_insert(
1021                &["id", "name"],
1022                &[
1023                    crate::ast::Expr::Named("s.id".to_string()),
1024                    crate::ast::Expr::Named("s.name".to_string()),
1025                ],
1026            );
1027
1028        let encoded = encode_cmd_binary(&cmd).expect("binary encode");
1029        let decoded = decode_cmd_binary(&encoded).unwrap();
1030        assert_eq!(decoded, cmd);
1031    }
1032
1033    #[test]
1034    fn cmd_binary_payload_rejects_oversized_header() {
1035        let mut payload = Vec::new();
1036        payload.extend_from_slice(&CMD_BIN_MAGIC);
1037        payload.extend_from_slice(&((MAX_CMD_BINARY_PAYLOAD_BYTES + 1) as u32).to_be_bytes());
1038        payload.extend_from_slice(&[]);
1039
1040        let err = decode_cmd_binary_payload(&payload).unwrap_err();
1041        assert!(err.contains("binary AST payload too large"));
1042    }
1043
1044    #[test]
1045    fn cmd_binary_payload_roundtrip() {
1046        let cmd = crate::ast::Qail::get("users").limit(3);
1047        let encoded = encode_cmd_binary(&cmd).expect("binary encode");
1048        let payload = decode_cmd_binary_payload(&encoded).unwrap();
1049        let mut deserializer = serde_json::Deserializer::from_slice(payload);
1050        let decoded: crate::ast::Qail = serde::Deserialize::deserialize(&mut deserializer).unwrap();
1051        deserializer.end().unwrap();
1052        assert_eq!(decoded.to_string(), cmd.to_string());
1053    }
1054
1055    #[test]
1056    fn cmd_binary_payload_rejects_legacy_qwb1() {
1057        let legacy_text = b"get users limit 1";
1058        let mut payload = Vec::new();
1059        payload.extend_from_slice(&CMD_BIN_LEGACY_MAGIC);
1060        payload.extend_from_slice(&(legacy_text.len() as u32).to_be_bytes());
1061        payload.extend_from_slice(legacy_text);
1062
1063        let err = decode_cmd_binary_payload(&payload).unwrap_err();
1064        assert!(err.contains("legacy QWB1"));
1065    }
1066
1067    #[test]
1068    fn cmd_binary_decode_rejects_raw_text_without_qwb2_header() {
1069        let err = decode_cmd_binary(b"get users limit 1").unwrap_err();
1070        assert!(err.contains("invalid wire header"));
1071    }
1072
1073    #[test]
1074    fn cmd_binary_decode_rejects_trailing_bytes() {
1075        let cmd = crate::ast::Qail::get("users").limit(1);
1076        let mut encoded = encode_cmd_binary(&cmd).expect("binary encode");
1077        encoded.extend_from_slice(&[0xAA, 0xBB]);
1078        let err = decode_cmd_binary(&encoded).unwrap_err();
1079        assert!(err.contains("invalid payload length"));
1080    }
1081
1082    #[test]
1083    fn cmd_binary_decode_rejects_unsafe_identifiers() {
1084        let cmd = crate::ast::Qail::get("users; DROP TABLE users; --").limit(1);
1085        let encoded = encode_cmd_binary_unchecked_for_test(&cmd);
1086
1087        let err = decode_cmd_binary(&encoded).unwrap_err();
1088
1089        assert!(err.contains("AST validation failed"));
1090        assert!(err.contains("table"));
1091    }
1092
1093    #[test]
1094    fn cmd_binary_encode_rejects_unsafe_identifiers() {
1095        let cmd = crate::ast::Qail::get("users; DROP TABLE users; --").limit(1);
1096
1097        let err = encode_cmd_binary(&cmd).unwrap_err();
1098
1099        assert!(err.contains("AST validation failed"));
1100        assert!(err.contains("table"));
1101    }
1102
1103    #[test]
1104    fn cmd_binary_decode_rejects_procedural_actions() {
1105        let cmd = crate::ast::Qail::call("refresh_materialized_views()");
1106        let encoded = encode_cmd_binary_unchecked_for_test(&cmd);
1107
1108        let err = decode_cmd_binary(&encoded).unwrap_err();
1109
1110        assert!(err.contains("AST validation failed"));
1111        assert!(err.contains("procedural/session actions"));
1112    }
1113
1114    #[test]
1115    fn cmd_binary_encode_rejects_procedural_actions() {
1116        let cmd = crate::ast::Qail::call("refresh_materialized_views()");
1117
1118        let err = encode_cmd_binary(&cmd).unwrap_err();
1119
1120        assert!(err.contains("AST validation failed"));
1121        assert!(err.contains("procedural/session actions"));
1122    }
1123
1124    #[test]
1125    fn cmd_binary_decode_rejects_unsafe_raw_function_values() {
1126        let cmd = crate::ast::Qail::get("users").filter(
1127            "updated_at",
1128            crate::ast::Operator::Lt,
1129            crate::ast::Value::Function("NOW(); DROP TABLE users; --".to_string()),
1130        );
1131        let encoded = encode_cmd_binary_unchecked_for_test(&cmd);
1132
1133        let err = decode_cmd_binary(&encoded).unwrap_err();
1134
1135        assert!(err.contains("AST validation failed"));
1136        assert!(err.contains("raw function values"));
1137    }
1138
1139    #[test]
1140    fn cmd_binary_encode_rejects_unsafe_raw_function_values() {
1141        let cmd = crate::ast::Qail::get("users").filter(
1142            "updated_at",
1143            crate::ast::Operator::Lt,
1144            crate::ast::Value::Function("NOW(); DROP TABLE users; --".to_string()),
1145        );
1146
1147        let err = encode_cmd_binary(&cmd).unwrap_err();
1148
1149        assert!(err.contains("AST validation failed"));
1150        assert!(err.contains("raw function values"));
1151    }
1152
1153    #[test]
1154    fn cmd_binary_encode_rejects_unsafe_check_constraint_payload() {
1155        let cmd = crate::ast::Qail {
1156            action: crate::ast::Action::AlterAddConstraint,
1157            table: "users".to_string(),
1158            channel: Some("users_active_check".to_string()),
1159            payload: Some("active); DROP TABLE users; --".to_string()),
1160            ..Default::default()
1161        };
1162
1163        let err = encode_cmd_binary(&cmd).unwrap_err();
1164
1165        assert!(err.contains("AST validation failed"));
1166        assert!(err.contains("SQL expression fragments"));
1167    }
1168
1169    #[test]
1170    fn cmd_binary_decode_enforces_depth_limits() {
1171        let mut nested = crate::ast::Qail::get("users").limit(1);
1172        for _ in 0..(MAX_AST_DEPTH + 2) {
1173            nested = crate::ast::Qail {
1174                action: crate::ast::Action::Get,
1175                table: "users".to_string(),
1176                columns: vec![crate::ast::Expr::Subquery {
1177                    query: Box::new(nested),
1178                    alias: None,
1179                }],
1180                ..crate::ast::Qail::default()
1181            };
1182        }
1183
1184        let encoded = encode_cmd_binary_unchecked_for_test(&nested);
1185        let err = decode_cmd_binary(&encoded).unwrap_err();
1186        assert!(
1187            err.contains("AST depth limit exceeded")
1188                || err.contains("binary AST decode failed")
1189                || err.contains("recursion limit exceeded")
1190        );
1191    }
1192
1193    #[test]
1194    fn cmd_binary_decode_bitflip_corpus_no_panic() {
1195        let seeds = vec![
1196            encode_cmd_binary(&crate::ast::Qail::get("users").limit(1)).expect("binary encode"),
1197            encode_cmd_binary(&crate::ast::Qail::set("users").set_value("active", true))
1198                .expect("binary encode"),
1199            vec![],
1200            b"QWB2garbage".to_vec(),
1201            vec![0u8; 32],
1202        ];
1203
1204        for seed in seeds {
1205            for i in 0..seed.len().min(128) {
1206                for bit in 0..8u8 {
1207                    let mut mutated = seed.clone();
1208                    mutated[i] ^= 1 << bit;
1209                    let _ = decode_cmd_binary(&mutated);
1210                }
1211            }
1212            let _ = decode_cmd_binary(&seed);
1213        }
1214    }
1215
1216    proptest! {
1217        #[test]
1218        fn cmd_binary_decode_fuzz_never_panics(data in proptest::collection::vec(any::<u8>(), 0..4096)) {
1219            let _ = decode_cmd_binary(&data);
1220        }
1221    }
1222
1223    #[test]
1224    fn cmds_text_roundtrip() {
1225        let cmds = vec![
1226            crate::ast::Qail::get("users").columns(["id", "email"]),
1227            crate::ast::Qail::get("users").limit(1),
1228            crate::ast::Qail::del("users").where_eq("id", 99),
1229        ];
1230
1231        let encoded = encode_cmds_text(&cmds);
1232        let decoded = decode_cmds_text(&encoded).unwrap();
1233        assert_eq!(decoded.len(), cmds.len());
1234        for (lhs, rhs) in decoded.iter().zip(cmds.iter()) {
1235            assert_eq!(lhs.to_string(), rhs.to_string());
1236        }
1237    }
1238
1239    #[test]
1240    fn decode_cmd_text_falls_back_to_raw_qail() {
1241        let decoded = decode_cmd_text("get users limit 1").unwrap();
1242        assert_eq!(decoded.action, crate::ast::Action::Get);
1243        assert_eq!(decoded.table, "users");
1244        assert!(
1245            decoded
1246                .cages
1247                .iter()
1248                .any(|c| matches!(c.kind, crate::ast::CageKind::Limit(1)))
1249        );
1250    }
1251
1252    #[test]
1253    fn cmd_text_rejects_overflowing_payload_length_without_panic() {
1254        let input = format!("{CMD_TEXT_MAGIC}\n{}\n", usize::MAX);
1255
1256        let err = decode_cmd_text(&input).expect_err("oversized length should fail closed");
1257        assert!(err.contains("payload length overflow") || err.contains("payload truncated"));
1258    }
1259
1260    #[test]
1261    fn cmds_text_rejects_oversized_command_count_without_allocation() {
1262        let input = format!("{CMDS_TEXT_MAGIC}\n{}\n", MAX_CMD_TEXT_BATCH_COMMANDS + 1);
1263
1264        let err = decode_cmds_text(&input).expect_err("oversized command count should fail closed");
1265        assert!(err.contains("command count exceeds limit"));
1266    }
1267}