Skip to main content

alopex_sql/
nim_bridge.rs

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
46/// Exact Select wire adapter for the continuous-aggregate payload.
47///
48/// Existing top-level Select statements encode their variant through
49/// `StatementKind`. The nested continuous-aggregate query is a named Select
50/// payload in its own right, so it carries and validates an explicit
51/// `variant: Select` field.
52pub(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    /// Decode the staged 0.4 wire shape without adding it to `StatementKind`.
139    ///
140    /// This narrow seam lets the contract shape be proven while the checked-in
141    /// producer remains on 0.3. The linked-version check is the same one used
142    /// by the production parser path and always runs before payload preflight.
143    #[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
162/// Return the SQL/PromQL MessagePack wire contract version exported by Nim.
163pub 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    // Option (a): double-quoted tokens are identifiers under SQL standard and
176    // PostgreSQL rules. The currently deployed Nim lexer predates that contract
177    // and emits them as string literals, so normalize the FFI input until every
178    // parser binary has the corrected token kind.
179    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            // 正常時の payload は最低でも MessagePack の配列ヘッダ 1 バイトを
185            // 含む。空 payload はゼロ初期化された CParseResult、つまり Nim 側
186            // から例外が漏れた事故 (issue #40 の desync 経路) を意味するため、
187            // 汎用の decode エラーではなく原因が特定できるエラーにする。
188            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                // Replace each quote with a space rather than removing it, so
458                // every later token keeps its original offset and diagnostics
459                // point into the SQL the caller actually wrote.
460                normalized.push(' ');
461                while let Some(identifier_ch) = chars.next() {
462                    if identifier_ch == '"' {
463                        if chars.peek() == Some(&'"') {
464                            // An escaped quote is two characters in the input
465                            // and one in the identifier; pad to keep the width.
466                            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                // PostgreSQL folds bare identifiers to lowercase. Delimited
508                // identifiers take the `\"` branch above and keep their exact
509                // spelling for case-sensitive resolution.
510                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(&quote) {
571                chars.next();
572            } else {
573                break;
574            }
575        }
576    }
577}
578
579/// Apply the parser's NATURAL markers to the joins they belong to.
580///
581/// The markers arrive as a flat list alongside the AST, so they only line up
582/// while both sides walk the joins in the same order. A mismatch used to leave
583/// the remaining joins as plain joins, turning `NATURAL JOIN` into a cross
584/// product without any diagnostic. Treat it as the contract violation it is.
585fn 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
818// nim-sql-parser/src/alopex_sql_parser.nim の `internalDefectPrefix` と
819// 一致させる。Nim 側の `except Defect` 節が付与する接頭辞で、パーサー
820// 内部の不変条件違反 (通常の構文エラーではない) を機械的に区別するための
821// マーカー。ワイヤ契約 (MessagePack AST) には影響しない、エラー文言のみの
822// 合意。
823const 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}