1use 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
13pub 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
23pub 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
35pub 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
60pub 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
78pub 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
114pub 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
137pub 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
153pub 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; 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}