1use crate::ast::ddl::CreateContinuousAggregate;
2use crate::ast::dml::{FromItem, Select, SelectItem};
3use crate::ast::expr::{Expr, ExprKind};
4use crate::ast::{Span, Statement, StatementKind};
5use crate::error::{ParserError, Result};
6use crate::nim_ffi::{self, OwnedBuffer, ParseResultKind};
7use serde::Deserialize;
8
9const MAX_SQL_INPUT_BYTES: usize = 1_048_576;
10const MAX_MESSAGEPACK_PAYLOAD_BYTES: usize = 1_048_576;
11const MAX_MESSAGEPACK_DEPTH: usize = 128;
12const MAX_MESSAGEPACK_VALUES: usize = 65_536;
13const SELECT_WRAPPER_PREFIX: &str = "SELECT ";
14const PARSER_CONTRACT_DESCRIPTOR: &str = include_str!("../nim-sql-parser/PARSER_CONTRACT_VERSION");
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
17enum InputPreflightError {
18 TooLarge,
19 LengthOverflow,
20 InteriorNul,
21}
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
24enum MessagePackPreflightError {
25 TooLarge,
26 TooDeep,
27 TooManyValues,
28 Truncated,
29 ReservedMarker,
30 TrailingBytes,
31}
32
33#[derive(Deserialize)]
34#[serde(deny_unknown_fields)]
35struct StagedContinuousAggregateStatement {
36 kind: StagedContinuousAggregateKind,
37 span: Span,
38}
39
40#[derive(Deserialize)]
41#[serde(tag = "variant")]
42enum StagedContinuousAggregateKind {
43 CreateContinuousAggregate(CreateContinuousAggregate),
44}
45
46pub(crate) mod continuous_aggregate_select_wire {
53 use crate::ast::{Expr, FromItem, OrderByExpr, Select, SelectItem, Span};
54 use serde::de::Error as _;
55 use serde::{Deserialize, Deserializer, Serialize, Serializer};
56
57 #[derive(Serialize)]
58 struct SelectWireRef<'a> {
59 variant: &'static str,
60 distinct: bool,
61 projection: &'a [SelectItem],
62 from: &'a [FromItem],
63 selection: &'a Option<Expr>,
64 group_by: &'a Option<Vec<Expr>>,
65 having: &'a Option<Expr>,
66 order_by: &'a [OrderByExpr],
67 limit: &'a Option<Expr>,
68 offset: &'a Option<Expr>,
69 span: Span,
70 }
71
72 #[derive(Deserialize)]
73 #[serde(deny_unknown_fields)]
74 struct SelectWire {
75 variant: String,
76 distinct: bool,
77 projection: Vec<SelectItem>,
78 from: Vec<FromItem>,
79 selection: Option<Expr>,
80 group_by: Option<Vec<Expr>>,
81 having: Option<Expr>,
82 order_by: Vec<OrderByExpr>,
83 limit: Option<Expr>,
84 offset: Option<Expr>,
85 span: Span,
86 }
87
88 pub(crate) fn serialize<S>(
89 select: &Select,
90 serializer: S,
91 ) -> std::result::Result<S::Ok, S::Error>
92 where
93 S: Serializer,
94 {
95 SelectWireRef {
96 variant: "Select",
97 distinct: select.distinct,
98 projection: &select.projection,
99 from: &select.from,
100 selection: &select.selection,
101 group_by: &select.group_by,
102 having: &select.having,
103 order_by: &select.order_by,
104 limit: &select.limit,
105 offset: &select.offset,
106 span: select.span,
107 }
108 .serialize(serializer)
109 }
110
111 pub(crate) fn deserialize<'de, D>(deserializer: D) -> std::result::Result<Select, D::Error>
112 where
113 D: Deserializer<'de>,
114 {
115 let wire = SelectWire::deserialize(deserializer)?;
116 if wire.variant != "Select" {
117 return Err(D::Error::custom(format!(
118 "expected nested query variant `Select`, found `{}`",
119 wire.variant
120 )));
121 }
122 Ok(Select {
123 distinct: wire.distinct,
124 projection: wire.projection,
125 from: wire.from,
126 selection: wire.selection,
127 group_by: wire.group_by,
128 having: wire.having,
129 order_by: wire.order_by,
130 limit: wire.limit,
131 offset: wire.offset,
132 span: wire.span,
133 })
134 }
135}
136
137impl CreateContinuousAggregate {
138 #[doc(hidden)]
144 pub fn decode_staged_messagepack(linked_parser_contract: &str, payload: &[u8]) -> Result<Self> {
145 ensure_linked_parser_contract(linked_parser_contract)?;
146 validate_bounded_messagepack(payload).map_err(messagepack_preflight_error)?;
147 let decoded = rmp_serde::from_slice::<StagedContinuousAggregateStatement>(payload)
148 .map_err(messagepack_decode_error)?;
149 let StagedContinuousAggregateKind::CreateContinuousAggregate(statement) = decoded.kind;
150 if decoded.span != statement.span {
151 return Err(ParserError::UnexpectedToken {
152 line: 0,
153 column: 0,
154 expected: "matching outer and kind spans in MessagePack AST".to_string(),
155 found: "continuous aggregate outer span differs from kind span".to_string(),
156 });
157 }
158 Ok(statement)
159 }
160}
161
162pub fn parser_contract_version() -> String {
164 nim_ffi::parser_contract_version()
165}
166
167pub fn parse_sql(sql: &str) -> Result<Vec<Statement>> {
168 preflight_input(sql, 0).map_err(parser_error_from_preflight)?;
169 parse_sql_preflighted(sql)
170}
171
172fn parse_sql_preflighted(sql: &str) -> Result<Vec<Statement>> {
173 ensure_linked_parser_contract(&nim_ffi::parser_contract_version())?;
174 let natural_join_markers = natural_join_markers(sql);
175 let normalized_sql = normalize_quoted_identifiers(sql);
180 let result = nim_ffi::parse_sql(&normalized_sql).map_err(parser_error_from_ffi_input)?;
181 match result.kind {
182 ParseResultKind::Ok => {
183 let buffer = OwnedBuffer::new(result.buffer_ptr, result.buffer_len);
184 if buffer.as_slice().is_empty() {
189 return Err(ParserError::UnexpectedToken {
190 line: 0,
191 column: 0,
192 expected: "MessagePack AST matching docs/ffi-ast-contract.md".to_string(),
193 found: "empty payload from Nim parser (leaked exception at FFI boundary; \
194 see issue #40)"
195 .to_string(),
196 });
197 }
198 validate_bounded_messagepack(buffer.as_slice()).map_err(messagepack_preflight_error)?;
199 let mut statements = rmp_serde::from_slice::<Vec<Statement>>(buffer.as_slice())
200 .map_err(messagepack_decode_error)?;
201 annotate_natural_joins(&mut statements, natural_join_markers)?;
202 Ok(statements)
203 }
204 ParseResultKind::Error => {
205 let buffer = OwnedBuffer::new(result.error_ptr.cast(), result.error_len);
206 Err(parser_error_from_nim(
207 String::from_utf8_lossy(buffer.as_slice()).as_ref(),
208 ))
209 }
210 }
211}
212
213fn expected_parser_contract() -> &'static str {
214 PARSER_CONTRACT_DESCRIPTOR.trim()
215}
216
217fn ensure_linked_parser_contract(linked_parser_contract: &str) -> Result<()> {
218 let expected = expected_parser_contract();
219 if linked_parser_contract == expected {
220 return Ok(());
221 }
222 Err(ParserError::UnexpectedToken {
223 line: 0,
224 column: 0,
225 expected: format!("linked Nim parser contract {expected}"),
226 found: format!("linked Nim parser contract {linked_parser_contract}"),
227 })
228}
229
230fn messagepack_decode_error(error: rmp_serde::decode::Error) -> ParserError {
231 ParserError::UnexpectedToken {
232 line: 0,
233 column: 0,
234 expected: "bounded MessagePack AST matching docs/ffi-ast-contract.md".to_string(),
235 found: error.to_string(),
236 }
237}
238
239fn messagepack_preflight_error(error: MessagePackPreflightError) -> ParserError {
240 let found = match error {
241 MessagePackPreflightError::TooLarge => {
242 format!("MessagePack payload exceeds {MAX_MESSAGEPACK_PAYLOAD_BYTES} bytes")
243 }
244 MessagePackPreflightError::TooDeep => {
245 format!("MessagePack nesting exceeds {MAX_MESSAGEPACK_DEPTH} levels")
246 }
247 MessagePackPreflightError::TooManyValues => {
248 format!("MessagePack collection limit of {MAX_MESSAGEPACK_VALUES} values exceeded")
249 }
250 MessagePackPreflightError::Truncated => "truncated MessagePack payload".to_string(),
251 MessagePackPreflightError::ReservedMarker => "reserved MessagePack marker 0xc1".to_string(),
252 MessagePackPreflightError::TrailingBytes => {
253 "trailing bytes after MessagePack payload".to_string()
254 }
255 };
256 ParserError::UnexpectedToken {
257 line: 0,
258 column: 0,
259 expected: "bounded MessagePack AST matching docs/ffi-ast-contract.md".to_string(),
260 found,
261 }
262}
263
264fn validate_bounded_messagepack(
265 payload: &[u8],
266) -> std::result::Result<(), MessagePackPreflightError> {
267 if payload.len() > MAX_MESSAGEPACK_PAYLOAD_BYTES {
268 return Err(MessagePackPreflightError::TooLarge);
269 }
270 let mut scanner = MessagePackScanner {
271 payload,
272 position: 0,
273 values: 0,
274 };
275 scanner.scan_value(1)?;
276 if scanner.position != payload.len() {
277 return Err(MessagePackPreflightError::TrailingBytes);
278 }
279 Ok(())
280}
281
282struct MessagePackScanner<'a> {
283 payload: &'a [u8],
284 position: usize,
285 values: usize,
286}
287
288impl MessagePackScanner<'_> {
289 fn scan_value(&mut self, depth: usize) -> std::result::Result<(), MessagePackPreflightError> {
290 if depth > MAX_MESSAGEPACK_DEPTH {
291 return Err(MessagePackPreflightError::TooDeep);
292 }
293 self.values = self
294 .values
295 .checked_add(1)
296 .ok_or(MessagePackPreflightError::TooManyValues)?;
297 if self.values > MAX_MESSAGEPACK_VALUES {
298 return Err(MessagePackPreflightError::TooManyValues);
299 }
300
301 let marker = self.read_byte()?;
302 match marker {
303 0x00..=0x7f | 0xc0 | 0xc2 | 0xc3 | 0xe0..=0xff => Ok(()),
304 0x80..=0x8f => self.scan_map(usize::from(marker & 0x0f), depth),
305 0x90..=0x9f => self.scan_children(usize::from(marker & 0x0f), depth),
306 0xa0..=0xbf => self.skip(usize::from(marker & 0x1f)),
307 0xc1 => Err(MessagePackPreflightError::ReservedMarker),
308 0xc4 | 0xd9 => {
309 let length = usize::from(self.read_byte()?);
310 self.skip(length)
311 }
312 0xc5 | 0xda => {
313 let length = usize::from(self.read_u16()?);
314 self.skip(length)
315 }
316 0xc6 | 0xdb => {
317 let length = usize::try_from(self.read_u32()?)
318 .map_err(|_| MessagePackPreflightError::TooLarge)?;
319 self.skip(length)
320 }
321 0xc7 => {
322 let length = usize::from(self.read_byte()?);
323 self.skip_ext(length)
324 }
325 0xc8 => {
326 let length = usize::from(self.read_u16()?);
327 self.skip_ext(length)
328 }
329 0xc9 => {
330 let length = usize::try_from(self.read_u32()?)
331 .map_err(|_| MessagePackPreflightError::TooLarge)?;
332 self.skip_ext(length)
333 }
334 0xca => self.skip(4),
335 0xcb => self.skip(8),
336 0xcc | 0xd0 => self.skip(1),
337 0xcd | 0xd1 => self.skip(2),
338 0xce | 0xd2 => self.skip(4),
339 0xcf | 0xd3 => self.skip(8),
340 0xd4 => self.skip_ext(1),
341 0xd5 => self.skip_ext(2),
342 0xd6 => self.skip_ext(4),
343 0xd7 => self.skip_ext(8),
344 0xd8 => self.skip_ext(16),
345 0xdc => {
346 let count = usize::from(self.read_u16()?);
347 self.scan_children(count, depth)
348 }
349 0xdd => {
350 let count = usize::try_from(self.read_u32()?)
351 .map_err(|_| MessagePackPreflightError::TooManyValues)?;
352 self.scan_children(count, depth)
353 }
354 0xde => {
355 let count = usize::from(self.read_u16()?);
356 self.scan_map(count, depth)
357 }
358 0xdf => {
359 let count = usize::try_from(self.read_u32()?)
360 .map_err(|_| MessagePackPreflightError::TooManyValues)?;
361 self.scan_map(count, depth)
362 }
363 }
364 }
365
366 fn scan_map(
367 &mut self,
368 entries: usize,
369 depth: usize,
370 ) -> std::result::Result<(), MessagePackPreflightError> {
371 let children = entries
372 .checked_mul(2)
373 .ok_or(MessagePackPreflightError::TooManyValues)?;
374 self.scan_children(children, depth)
375 }
376
377 fn scan_children(
378 &mut self,
379 children: usize,
380 depth: usize,
381 ) -> std::result::Result<(), MessagePackPreflightError> {
382 if children > MAX_MESSAGEPACK_VALUES {
383 return Err(MessagePackPreflightError::TooManyValues);
384 }
385 for _ in 0..children {
386 self.scan_value(depth + 1)?;
387 }
388 Ok(())
389 }
390
391 fn skip_ext(
392 &mut self,
393 payload_length: usize,
394 ) -> std::result::Result<(), MessagePackPreflightError> {
395 let total = payload_length
396 .checked_add(1)
397 .ok_or(MessagePackPreflightError::TooLarge)?;
398 self.skip(total)
399 }
400
401 fn read_byte(&mut self) -> std::result::Result<u8, MessagePackPreflightError> {
402 let byte = *self
403 .payload
404 .get(self.position)
405 .ok_or(MessagePackPreflightError::Truncated)?;
406 self.position += 1;
407 Ok(byte)
408 }
409
410 fn read_u16(&mut self) -> std::result::Result<u16, MessagePackPreflightError> {
411 let bytes = self.take(2)?;
412 Ok(u16::from_be_bytes([bytes[0], bytes[1]]))
413 }
414
415 fn read_u32(&mut self) -> std::result::Result<u32, MessagePackPreflightError> {
416 let bytes = self.take(4)?;
417 Ok(u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]]))
418 }
419
420 fn skip(&mut self, length: usize) -> std::result::Result<(), MessagePackPreflightError> {
421 self.take(length).map(|_| ())
422 }
423
424 fn take(&mut self, length: usize) -> std::result::Result<&[u8], MessagePackPreflightError> {
425 let end = self
426 .position
427 .checked_add(length)
428 .ok_or(MessagePackPreflightError::Truncated)?;
429 let bytes = self
430 .payload
431 .get(self.position..end)
432 .ok_or(MessagePackPreflightError::Truncated)?;
433 self.position = end;
434 Ok(bytes)
435 }
436}
437
438fn normalize_quoted_identifiers(sql: &str) -> String {
439 let mut normalized = String::with_capacity(sql.len());
440 let mut chars = sql.chars().peekable();
441 while let Some(ch) = chars.next() {
442 match ch {
443 '\'' => {
444 normalized.push(ch);
445 while let Some(string_ch) = chars.next() {
446 normalized.push(string_ch);
447 if string_ch == '\'' {
448 if chars.peek() == Some(&'\'') {
449 normalized.push(chars.next().expect("peeked quote"));
450 } else {
451 break;
452 }
453 }
454 }
455 }
456 '"' => {
457 normalized.push(' ');
461 while let Some(identifier_ch) = chars.next() {
462 if identifier_ch == '"' {
463 if chars.peek() == Some(&'"') {
464 normalized.push(chars.next().expect("peeked quote"));
467 normalized.push(' ');
468 } else {
469 normalized.push(' ');
470 break;
471 }
472 } else {
473 normalized.push(identifier_ch);
474 }
475 }
476 }
477 '-' if chars.peek() == Some(&'-') => {
478 normalized.push(ch);
479 normalized.push(chars.next().expect("peeked comment dash"));
480 for comment_ch in chars.by_ref() {
481 normalized.push(comment_ch);
482 if comment_ch == '\n' {
483 break;
484 }
485 }
486 }
487 '/' if chars.peek() == Some(&'*') => {
488 normalized.push(ch);
489 normalized.push(chars.next().expect("peeked comment star"));
490 let mut previous = '\0';
491 for comment_ch in chars.by_ref() {
492 normalized.push(comment_ch);
493 if previous == '*' && comment_ch == '/' {
494 break;
495 }
496 previous = comment_ch;
497 }
498 }
499 ch if ch.is_ascii_alphabetic() || ch == '_' => {
500 let mut identifier = String::from(ch);
501 while chars
502 .peek()
503 .is_some_and(|next| next.is_ascii_alphanumeric() || *next == '_')
504 {
505 identifier.push(chars.next().expect("peeked identifier character"));
506 }
507 normalized.push_str(&identifier.to_ascii_lowercase());
511 }
512 _ => normalized.push(ch),
513 }
514 }
515 normalized
516}
517
518fn natural_join_markers(sql: &str) -> Vec<bool> {
519 let mut markers = Vec::new();
520 let mut saw_natural = false;
521 let mut chars = sql.chars().peekable();
522 while let Some(ch) = chars.next() {
523 match ch {
524 '\'' | '"' => skip_quoted(&mut chars, ch),
525 '-' if chars.peek() == Some(&'-') => {
526 chars.next();
527 for comment_ch in chars.by_ref() {
528 if comment_ch == '\n' {
529 break;
530 }
531 }
532 }
533 '/' if chars.peek() == Some(&'*') => {
534 chars.next();
535 let mut previous = '\0';
536 for comment_ch in chars.by_ref() {
537 if previous == '*' && comment_ch == '/' {
538 break;
539 }
540 previous = comment_ch;
541 }
542 }
543 ';' => saw_natural = false,
544 c if c.is_ascii_alphabetic() || c == '_' => {
545 let mut word = String::from(c);
546 while chars
547 .peek()
548 .is_some_and(|next| next.is_ascii_alphanumeric() || *next == '_')
549 {
550 word.push(chars.next().expect("peeked identifier character"));
551 }
552 match word.to_ascii_lowercase().as_str() {
553 "natural" => saw_natural = true,
554 "join" => {
555 markers.push(saw_natural);
556 saw_natural = false;
557 }
558 _ => {}
559 }
560 }
561 _ => {}
562 }
563 }
564 markers
565}
566
567fn skip_quoted(chars: &mut std::iter::Peekable<std::str::Chars<'_>>, quote: char) {
568 while let Some(ch) = chars.next() {
569 if ch == quote {
570 if chars.peek() == Some("e) {
571 chars.next();
572 } else {
573 break;
574 }
575 }
576 }
577}
578
579fn annotate_natural_joins(statements: &mut [Statement], natural_markers: Vec<bool>) -> Result<()> {
586 let supplied = natural_markers.len();
587 let mut natural_markers = natural_markers.into_iter();
588 let mut consumed = 0usize;
589 for statement in statements {
590 if let StatementKind::Select(select) = &mut statement.kind {
591 annotate_select_natural_joins(select, &mut natural_markers, &mut consumed);
592 }
593 }
594
595 if consumed != supplied {
596 return Err(ParserError::UnexpectedToken {
597 line: 0,
598 column: 0,
599 expected: format!("{supplied} NATURAL join markers, one per join"),
600 found: format!("{consumed} joins in the AST"),
601 });
602 }
603 Ok(())
604}
605
606fn annotate_select_natural_joins(
607 select: &mut Select,
608 natural_markers: &mut impl Iterator<Item = bool>,
609 consumed: &mut usize,
610) {
611 for item in &mut select.projection {
612 if let SelectItem::Expr { expr, .. } = item {
613 annotate_expr_natural_joins(expr, natural_markers, consumed);
614 }
615 }
616 for from in &mut select.from {
617 annotate_from_natural_joins(from, natural_markers, consumed);
618 }
619 if let Some(selection) = &mut select.selection {
620 annotate_expr_natural_joins(selection, natural_markers, consumed);
621 }
622 if let Some(group_by) = &mut select.group_by {
623 for expression in group_by {
624 annotate_expr_natural_joins(expression, natural_markers, consumed);
625 }
626 }
627 if let Some(having) = &mut select.having {
628 annotate_expr_natural_joins(having, natural_markers, consumed);
629 }
630 for order_by in &mut select.order_by {
631 annotate_expr_natural_joins(&mut order_by.expr, natural_markers, consumed);
632 }
633 if let Some(limit) = &mut select.limit {
634 annotate_expr_natural_joins(limit, natural_markers, consumed);
635 }
636 if let Some(offset) = &mut select.offset {
637 annotate_expr_natural_joins(offset, natural_markers, consumed);
638 }
639}
640
641fn annotate_from_natural_joins(
642 from: &mut FromItem,
643 natural_markers: &mut impl Iterator<Item = bool>,
644 consumed: &mut usize,
645) {
646 match from {
647 FromItem::Join {
648 left,
649 right,
650 natural,
651 ..
652 } => {
653 annotate_from_natural_joins(left, natural_markers, consumed);
654 if let Some(marker) = natural_markers.next() {
655 *natural |= marker;
656 *consumed += 1;
657 }
658 annotate_from_natural_joins(right, natural_markers, consumed);
659 }
660 FromItem::Derived { subquery, .. } => {
661 if let StatementKind::Select(select) = &mut subquery.kind {
662 annotate_select_natural_joins(select, natural_markers, consumed);
663 }
664 }
665 FromItem::Table { .. } => {}
666 }
667}
668
669fn annotate_expr_natural_joins(
670 expr: &mut Expr,
671 natural_markers: &mut impl Iterator<Item = bool>,
672 consumed: &mut usize,
673) {
674 match &mut expr.kind {
675 ExprKind::ScalarSubquery { subquery } | ExprKind::Exists { subquery, .. } => {
676 if let StatementKind::Select(select) = &mut subquery.kind {
677 annotate_select_natural_joins(select, natural_markers, consumed);
678 }
679 }
680 ExprKind::InSubquery { expr, subquery, .. }
681 | ExprKind::Quantified { expr, subquery, .. } => {
682 annotate_expr_natural_joins(expr, natural_markers, consumed);
683 if let StatementKind::Select(select) = &mut subquery.kind {
684 annotate_select_natural_joins(select, natural_markers, consumed);
685 }
686 }
687 ExprKind::BinaryOp { left, right, .. } => {
688 annotate_expr_natural_joins(left, natural_markers, consumed);
689 annotate_expr_natural_joins(right, natural_markers, consumed);
690 }
691 ExprKind::UnaryOp { operand, .. } | ExprKind::IsNull { expr: operand, .. } => {
692 annotate_expr_natural_joins(operand, natural_markers, consumed);
693 }
694 ExprKind::FunctionCall { args, .. } => {
695 for argument in args {
696 annotate_expr_natural_joins(argument, natural_markers, consumed);
697 }
698 }
699 ExprKind::Between {
700 expr, low, high, ..
701 } => {
702 annotate_expr_natural_joins(expr, natural_markers, consumed);
703 annotate_expr_natural_joins(low, natural_markers, consumed);
704 annotate_expr_natural_joins(high, natural_markers, consumed);
705 }
706 ExprKind::Like {
707 expr,
708 pattern,
709 escape,
710 ..
711 } => {
712 annotate_expr_natural_joins(expr, natural_markers, consumed);
713 annotate_expr_natural_joins(pattern, natural_markers, consumed);
714 if let Some(escape) = escape {
715 annotate_expr_natural_joins(escape, natural_markers, consumed);
716 }
717 }
718 ExprKind::InList { expr, list, .. } => {
719 annotate_expr_natural_joins(expr, natural_markers, consumed);
720 for item in list {
721 annotate_expr_natural_joins(item, natural_markers, consumed);
722 }
723 }
724 ExprKind::Cast { expr, .. } => {
725 annotate_expr_natural_joins(expr, natural_markers, consumed);
726 }
727 ExprKind::Literal { .. } | ExprKind::ColumnRef { .. } | ExprKind::VectorLiteral { .. } => {}
728 }
729}
730
731pub fn parse_expression_sql(sql: &str) -> Result<crate::ast::Expr> {
732 let wrapped_len =
733 preflight_input(sql, SELECT_WRAPPER_PREFIX.len()).map_err(parser_error_from_preflight)?;
734 let mut wrapped = String::with_capacity(wrapped_len);
735 wrapped.push_str(SELECT_WRAPPER_PREFIX);
736 wrapped.push_str(sql);
737 let statements = parse_sql_preflighted(&wrapped)?;
738 let Some(statement) = statements.into_iter().next() else {
739 return Err(empty_expression_error());
740 };
741 let StatementKind::Select(select) = statement.kind else {
742 return Err(empty_expression_error());
743 };
744 let Some(crate::ast::SelectItem::Expr { expr, .. }) = select.projection.into_iter().next()
745 else {
746 return Err(empty_expression_error());
747 };
748 Ok(expr)
749}
750
751fn empty_expression_error() -> ParserError {
752 ParserError::UnexpectedToken {
753 line: 0,
754 column: 0,
755 expected: "expression".to_string(),
756 found: "empty parser result".to_string(),
757 }
758}
759
760fn checked_total_input_len(
761 input_len: usize,
762 wrapper_len: usize,
763) -> std::result::Result<usize, InputPreflightError> {
764 let total_len = input_len
765 .checked_add(wrapper_len)
766 .ok_or(InputPreflightError::LengthOverflow)?;
767 if total_len > MAX_SQL_INPUT_BYTES {
768 return Err(InputPreflightError::TooLarge);
769 }
770 Ok(total_len)
771}
772
773fn preflight_input(
774 sql: &str,
775 wrapper_len: usize,
776) -> std::result::Result<usize, InputPreflightError> {
777 let total_len = checked_total_input_len(sql.len(), wrapper_len)?;
778 if sql.as_bytes().contains(&0) {
779 return Err(InputPreflightError::InteriorNul);
780 }
781 Ok(total_len)
782}
783
784fn parser_error_from_preflight(error: InputPreflightError) -> ParserError {
785 match error {
786 InputPreflightError::TooLarge | InputPreflightError::LengthOverflow => {
787 input_too_large_error()
788 }
789 InputPreflightError::InteriorNul => interior_nul_error(),
790 }
791}
792
793fn parser_error_from_ffi_input(error: nim_ffi::ParseInputError) -> ParserError {
794 match error {
795 nim_ffi::ParseInputError::LengthOutOfRange => input_too_large_error(),
796 nim_ffi::ParseInputError::InteriorNul => interior_nul_error(),
797 }
798}
799
800fn input_too_large_error() -> ParserError {
801 ParserError::UnexpectedToken {
802 line: 0,
803 column: 0,
804 expected: "SQL input at most 1048576 UTF-8 bytes".to_string(),
805 found: "SQL input exceeds byte limit".to_string(),
806 }
807}
808
809fn interior_nul_error() -> ParserError {
810 ParserError::UnexpectedToken {
811 line: 0,
812 column: 0,
813 expected: "valid SQL without interior NUL bytes".to_string(),
814 found: "interior NUL byte".to_string(),
815 }
816}
817
818const INTERNAL_DEFECT_PREFIX: &str =
824 "internal parser defect (this is a parser bug, not invalid SQL): ";
825
826fn parser_error_from_nim(message: &str) -> ParserError {
827 if let Some(defect_message) = message.strip_prefix(INTERNAL_DEFECT_PREFIX) {
828 return ParserError::InternalParserDefect {
829 message: defect_message.to_string(),
830 };
831 }
832 let (line, column) = parse_nim_line_col(message).unwrap_or((0, 0));
833 ParserError::UnexpectedToken {
834 line,
835 column,
836 expected: "valid SQL".to_string(),
837 found: message.to_string(),
838 }
839}
840
841fn parse_nim_line_col(message: &str) -> Option<(u64, u64)> {
842 let after_line = message.strip_prefix("Parse error at line ")?;
843 let (line, rest) = after_line.split_once(", col ")?;
844 let (col, _) = rest.split_once(':')?;
845 Some((line.parse().ok()?, col.parse().ok()?))
846}
847
848#[cfg(test)]
849mod input_preflight_tests {
850 use super::*;
851
852 #[test]
853 fn raw_sql_guard_accepts_boundary_minus_and_exact_but_rejects_plus() {
854 assert_eq!(
855 checked_total_input_len(MAX_SQL_INPUT_BYTES - 1, 0),
856 Ok(MAX_SQL_INPUT_BYTES - 1)
857 );
858 assert_eq!(
859 checked_total_input_len(MAX_SQL_INPUT_BYTES, 0),
860 Ok(MAX_SQL_INPUT_BYTES)
861 );
862 assert_eq!(
863 checked_total_input_len(MAX_SQL_INPUT_BYTES + 1, 0),
864 Err(InputPreflightError::TooLarge)
865 );
866 }
867
868 #[test]
869 fn guard_counts_utf8_bytes_instead_of_characters() {
870 let exact = "é".repeat(MAX_SQL_INPUT_BYTES / "é".len());
871 assert_eq!(exact.chars().count(), MAX_SQL_INPUT_BYTES / 2);
872 assert_eq!(exact.len(), MAX_SQL_INPUT_BYTES);
873 assert_eq!(preflight_input(&exact, 0), Ok(MAX_SQL_INPUT_BYTES));
874
875 let plus = format!("{exact}é");
876 assert_eq!(
877 preflight_input(&plus, 0),
878 Err(InputPreflightError::TooLarge)
879 );
880 }
881
882 #[test]
883 fn expression_guard_includes_wrapper_and_detects_length_overflow() {
884 assert_eq!(
885 checked_total_input_len(
886 MAX_SQL_INPUT_BYTES - SELECT_WRAPPER_PREFIX.len(),
887 SELECT_WRAPPER_PREFIX.len(),
888 ),
889 Ok(MAX_SQL_INPUT_BYTES)
890 );
891 assert_eq!(
892 checked_total_input_len(
893 MAX_SQL_INPUT_BYTES - SELECT_WRAPPER_PREFIX.len() + 1,
894 SELECT_WRAPPER_PREFIX.len(),
895 ),
896 Err(InputPreflightError::TooLarge)
897 );
898 assert_eq!(
899 checked_total_input_len(usize::MAX, SELECT_WRAPPER_PREFIX.len()),
900 Err(InputPreflightError::LengthOverflow)
901 );
902 }
903
904 #[test]
905 fn preflight_rejects_nul_before_any_ffi_work() {
906 assert_eq!(
907 preflight_input("SELECT \0 1", 0),
908 Err(InputPreflightError::InteriorNul)
909 );
910 assert_eq!(
911 preflight_input("1 \0 2", SELECT_WRAPPER_PREFIX.len()),
912 Err(InputPreflightError::InteriorNul)
913 );
914 }
915}