Skip to main content

substrait_explain/parser/
types.rs

1use pest::iterators::Pair;
2use substrait::proto::r#type::{Kind, Nullability, Parameter};
3use substrait::proto::{self, Type};
4
5use super::{ParsePair, Rule, ScopedParsePair, iter_pairs, unwrap_single_pair};
6use crate::extensions::SimpleExtensions;
7use crate::extensions::simple::{ExtensionKind, MissingReference};
8use crate::parser::MessageParseError;
9
10/// Given a name (plain or compound) and an optional anchor, resolve and
11/// validate the extension anchor.
12///
13/// For `Function` kind, delegates to [`SimpleExtensions::resolve_function`]
14/// which encapsulates the full base-name fallback logic.
15///
16/// For other kinds, performs a direct anchor/name validation.
17pub(crate) fn get_and_validate_anchor(
18    extensions: &SimpleExtensions,
19    kind: ExtensionKind,
20    anchor: Option<u32>,
21    name: &str,
22    span: pest::Span,
23) -> Result<u32, MessageParseError> {
24    if kind == ExtensionKind::Function {
25        return extensions
26            .resolve_function(name, anchor)
27            .map(|r| r.anchor)
28            .map_err(|e| {
29                MessageParseError::lookup(kind.name(), e, span, "Error resolving function")
30            });
31    }
32    // For non-function kinds, validate the anchor/name pair directly.
33    match anchor {
34        Some(a) => match extensions.find_by_anchor(kind, a) {
35            Err(e) => Err(MessageParseError::lookup(
36                kind.name(),
37                e,
38                span,
39                "Error matching name to anchor",
40            )),
41            Ok((_, stored)) => {
42                if stored.full() == name || stored.base() == name {
43                    Ok(a)
44                } else {
45                    Err(MessageParseError::lookup(
46                        kind.name(),
47                        MissingReference::Mismatched(kind, name.to_string(), a),
48                        span,
49                        "Error matching name to anchor",
50                    ))
51                }
52            }
53        },
54        None => extensions.find_by_name(kind, name).map_err(|e| {
55            MessageParseError::lookup(kind.name(), e, span, "Error finding extension for name")
56        }),
57    }
58}
59
60impl ParsePair for Nullability {
61    fn rule() -> Rule {
62        Rule::nullability
63    }
64
65    fn message() -> &'static str {
66        "Nullability"
67    }
68
69    fn parse_pair(pair: Pair<Rule>) -> Self {
70        assert_eq!(pair.as_rule(), Rule::nullability);
71        match pair.as_str() {
72            "?" => Nullability::Nullable,
73            "" => Nullability::Required,
74            "⁉" => Nullability::Unspecified,
75            _ => panic!("Invalid nullability: {}", pair.as_str()),
76        }
77    }
78}
79
80impl ScopedParsePair for Parameter {
81    fn rule() -> Rule {
82        Rule::parameter
83    }
84
85    fn message() -> &'static str {
86        "Parameter"
87    }
88
89    fn parse_pair(
90        extensions: &SimpleExtensions,
91        pair: Pair<Rule>,
92    ) -> Result<Self, MessageParseError> {
93        assert_eq!(pair.as_rule(), Rule::parameter);
94        let inner = unwrap_single_pair(pair);
95        match inner.as_rule() {
96            Rule::r#type => Ok(Parameter {
97                parameter: Some(proto::r#type::parameter::Parameter::DataType(
98                    Type::parse_pair(extensions, inner)?,
99                )),
100            }),
101            _ => unimplemented!("{:?}", inner.as_rule()),
102        }
103    }
104}
105
106fn parse_simple_type(pair: Pair<Rule>) -> Type {
107    assert_eq!(pair.as_rule(), Rule::simple_type);
108    let mut iter = iter_pairs(pair.into_inner());
109    let name = iter.pop(Rule::simple_type_name).as_str();
110    let nullability = iter.parse_next::<Nullability>();
111    iter.done();
112
113    let kind = match name {
114        "boolean" => Kind::Bool(proto::r#type::Boolean {
115            nullability: nullability.into(),
116            type_variation_reference: 0,
117        }),
118        "i64" => Kind::I64(proto::r#type::I64 {
119            nullability: nullability.into(),
120            type_variation_reference: 0,
121        }),
122        "i32" => Kind::I32(proto::r#type::I32 {
123            nullability: nullability.into(),
124            type_variation_reference: 0,
125        }),
126        "i16" => Kind::I16(proto::r#type::I16 {
127            nullability: nullability.into(),
128            type_variation_reference: 0,
129        }),
130        "i8" => Kind::I8(proto::r#type::I8 {
131            nullability: nullability.into(),
132            type_variation_reference: 0,
133        }),
134        "fp32" => Kind::Fp32(proto::r#type::Fp32 {
135            nullability: nullability.into(),
136            type_variation_reference: 0,
137        }),
138        "fp64" => Kind::Fp64(proto::r#type::Fp64 {
139            nullability: nullability.into(),
140            type_variation_reference: 0,
141        }),
142        "string" => Kind::String(proto::r#type::String {
143            nullability: nullability.into(),
144            type_variation_reference: 0,
145        }),
146        "binary" => Kind::Binary(proto::r#type::Binary {
147            nullability: nullability.into(),
148            type_variation_reference: 0,
149        }),
150        #[allow(deprecated)]
151        "timestamp" => Kind::Timestamp(proto::r#type::Timestamp {
152            nullability: nullability.into(),
153            type_variation_reference: 0,
154        }),
155        #[allow(deprecated)]
156        "timestamp_tz" => Kind::TimestampTz(proto::r#type::TimestampTz {
157            nullability: nullability.into(),
158            type_variation_reference: 0,
159        }),
160        "date" => Kind::Date(proto::r#type::Date {
161            nullability: nullability.into(),
162            type_variation_reference: 0,
163        }),
164        #[allow(deprecated)]
165        "time" => Kind::Time(proto::r#type::Time {
166            nullability: nullability.into(),
167            type_variation_reference: 0,
168        }),
169        "interval_year" => Kind::IntervalYear(proto::r#type::IntervalYear {
170            nullability: nullability.into(),
171            type_variation_reference: 0,
172        }),
173        "uuid" => Kind::Uuid(proto::r#type::Uuid {
174            nullability: nullability.into(),
175            type_variation_reference: 0,
176        }),
177        _ => unreachable!("Type {} exists in parser but not implemented in code", name),
178    };
179    Type { kind: Some(kind) }
180}
181
182fn parse_compound_type(
183    extensions: &SimpleExtensions,
184    pair: Pair<Rule>,
185) -> Result<Type, MessageParseError> {
186    assert_eq!(pair.as_rule(), Rule::compound_type);
187    let inner = unwrap_single_pair(pair);
188    match inner.as_rule() {
189        Rule::list_type => parse_list_type(extensions, inner),
190        // Rule::map_type => parse_map_type(inner),
191        // Rule::struct_type => parse_struct_type(inner),
192        Rule::precision_timestamp_tz_type
193        | Rule::precision_timestamp_type
194        | Rule::precision_time_type => parse_precision_type(inner),
195        Rule::interval_day_type => parse_interval_day_type(inner),
196        _ => unimplemented!("{:?}", inner.as_rule()),
197    }
198}
199
200/// Parse a sub-second precision type parameter.
201///
202/// Substrait allows any precision from 0 to 12 on a type, so this stays a plain
203/// integer. Narrowing to a precision that a *value* can actually be written at
204/// is [`crate::precision::SupportedPrecision`]'s job, and only literals need it.
205fn parse_precision(
206    precision_pair: Pair<Rule>,
207    context: &'static str,
208) -> Result<i32, MessageParseError> {
209    let precision_span = precision_pair.as_span();
210    let precision = precision_pair.as_str().parse::<i32>().ok();
211    precision.filter(|p| (0..=12).contains(p)).ok_or_else(|| {
212        MessageParseError::invalid(
213            context,
214            precision_span,
215            format!(
216                "precision must be between 0 and 12, got {}",
217                precision_pair.as_str()
218            ),
219        )
220    })
221}
222
223fn parse_precision_type(pair: Pair<Rule>) -> Result<Type, MessageParseError> {
224    let rule = pair.as_rule();
225    let mut iter = iter_pairs(pair.into_inner());
226    let nullability = iter.parse_next::<Nullability>();
227    let precision_pair = iter.pop(Rule::integer);
228    let precision = parse_precision(precision_pair, "precision time type")?;
229    iter.done();
230    let kind = match rule {
231        Rule::precision_timestamp_type => {
232            Kind::PrecisionTimestamp(proto::r#type::PrecisionTimestamp {
233                precision,
234                nullability: nullability.into(),
235                type_variation_reference: 0,
236            })
237        }
238        Rule::precision_timestamp_tz_type => {
239            Kind::PrecisionTimestampTz(proto::r#type::PrecisionTimestampTz {
240                precision,
241                nullability: nullability.into(),
242                type_variation_reference: 0,
243            })
244        }
245        Rule::precision_time_type => Kind::PrecisionTime(proto::r#type::PrecisionTime {
246            precision,
247            nullability: nullability.into(),
248            type_variation_reference: 0,
249        }),
250        _ => unreachable!("parse_precision_type called with rule {:?}", rule),
251    };
252    Ok(Type { kind: Some(kind) })
253}
254
255fn parse_interval_day_type(pair: Pair<Rule>) -> Result<Type, MessageParseError> {
256    assert_eq!(pair.as_rule(), Rule::interval_day_type);
257    let mut iter = iter_pairs(pair.into_inner());
258    let nullability = iter.parse_next::<Nullability>();
259    let precision_pair = iter.pop(Rule::integer);
260    let precision = parse_precision(precision_pair, "interval day type")?;
261    iter.done();
262    Ok(Type {
263        kind: Some(Kind::IntervalDay(proto::r#type::IntervalDay {
264            nullability: nullability.into(),
265            type_variation_reference: 0,
266            precision: Some(precision),
267        })),
268    })
269}
270
271fn parse_list_type(
272    extensions: &SimpleExtensions,
273    pair: Pair<Rule>,
274) -> Result<Type, MessageParseError> {
275    assert_eq!(pair.as_rule(), Rule::list_type);
276    let mut iter = iter_pairs(pair.into_inner());
277    let nullability = iter.parse_next::<Nullability>();
278    let inner = iter.parse_next_scoped::<Type>(extensions)?;
279    iter.done();
280
281    Ok(Type {
282        kind: Some(Kind::List(Box::new(proto::r#type::List {
283            nullability: nullability.into(),
284            r#type: Some(Box::new(inner)),
285            type_variation_reference: 0,
286        }))),
287    })
288}
289
290fn parse_parameters(
291    extensions: &SimpleExtensions,
292    pair: Pair<Rule>,
293) -> Result<Vec<Parameter>, MessageParseError> {
294    assert_eq!(pair.as_rule(), Rule::parameters);
295    let mut iter = iter_pairs(pair.into_inner());
296    let mut params = Vec::new();
297    while let Some(param) = iter.parse_if_next_scoped::<Parameter>(extensions) {
298        params.push(param?);
299    }
300    iter.done();
301    Ok(params)
302}
303
304fn parse_user_defined_type(
305    extensions: &SimpleExtensions,
306    pair: Pair<Rule>,
307) -> Result<Type, MessageParseError> {
308    let span = pair.as_span();
309    assert_eq!(pair.as_rule(), Rule::user_defined_type);
310    let mut iter = iter_pairs(pair.into_inner());
311    // TODO: quoted names not yet supported — see grammar TODO at user_defined_type
312    let name = iter.pop(Rule::identifier).as_str().to_string();
313    let anchor = iter
314        .try_pop(Rule::anchor)
315        .map(|n| unwrap_single_pair(n).as_str().parse::<u32>().unwrap());
316
317    // TODO: Handle urn_anchor; validate that it matches the anchor
318    let _urn_anchor = iter
319        .try_pop(Rule::urn_anchor)
320        .map(|n| unwrap_single_pair(n).as_str().parse::<u32>().unwrap());
321
322    let nullability = iter.parse_next::<Nullability>();
323    let parameters = match iter.try_pop(Rule::parameters) {
324        Some(p) => parse_parameters(extensions, p)?,
325        None => Vec::new(),
326    };
327    iter.done();
328
329    let anchor = get_and_validate_anchor(extensions, ExtensionKind::Type, anchor, &name, span)?;
330
331    Ok(Type {
332        kind: Some(Kind::UserDefined(proto::r#type::UserDefined {
333            type_reference: anchor,
334            nullability: nullability.into(),
335            type_parameters: parameters,
336            type_variation_reference: 0,
337        })),
338    })
339}
340
341impl ScopedParsePair for Type {
342    fn rule() -> Rule {
343        Rule::r#type
344    }
345
346    fn message() -> &'static str {
347        "Type"
348    }
349
350    fn parse_pair(
351        extensions: &SimpleExtensions,
352        pair: Pair<Rule>,
353    ) -> Result<Self, MessageParseError> {
354        assert_eq!(pair.as_rule(), Rule::r#type);
355        let inner = unwrap_single_pair(pair);
356        match inner.as_rule() {
357            Rule::simple_type => Ok(parse_simple_type(inner)),
358            Rule::compound_type => parse_compound_type(extensions, inner),
359            Rule::user_defined_type => parse_user_defined_type(extensions, inner),
360            _ => unreachable!(
361                "Grammar guarantees type can only be simple_type, compound_type, or user_defined_type, got: {:?}",
362                inner.as_rule()
363            ),
364        }
365    }
366}
367
368#[cfg(test)]
369mod tests {
370    use pest::Parser;
371    use substrait::proto::r#type::{I64, Kind, Nullability};
372
373    use super::*;
374    use crate::parser::ExpressionParser;
375
376    #[test]
377    fn test_parse_simple_type() {
378        let mut pairs = ExpressionParser::parse(Rule::simple_type, "i64").unwrap();
379        let pair = pairs.next().unwrap();
380        assert_eq!(pairs.next(), None);
381        let t = parse_simple_type(pair);
382        assert_eq!(
383            t,
384            Type {
385                kind: Some(Kind::I64(I64 {
386                    nullability: Nullability::Required as i32,
387                    type_variation_reference: 0,
388                })),
389            }
390        );
391
392        let mut pairs = ExpressionParser::parse(Rule::simple_type, "string?").unwrap();
393        let pair = pairs.next().unwrap();
394        assert_eq!(pairs.next(), None);
395        let t = parse_simple_type(pair);
396        assert_eq!(
397            t,
398            Type {
399                kind: Some(Kind::String(proto::r#type::String {
400                    nullability: Nullability::Nullable as i32,
401                    type_variation_reference: 0,
402                })),
403            }
404        );
405    }
406
407    #[test]
408    fn test_parse_type() {
409        let extensions = SimpleExtensions::default();
410        let mut pairs = ExpressionParser::parse(Rule::r#type, "i64").unwrap();
411        let pair = pairs.next().unwrap();
412        assert_eq!(pairs.next(), None);
413        let t = Type::parse_pair(&extensions, pair).unwrap();
414        assert_eq!(
415            t,
416            Type {
417                kind: Some(Kind::I64(I64 {
418                    nullability: Nullability::Required as i32,
419                    type_variation_reference: 0,
420                }))
421            }
422        );
423    }
424
425    #[test]
426    fn test_parse_list_type() {
427        let extensions = SimpleExtensions::default();
428        let mut pairs = ExpressionParser::parse(Rule::list_type, "list<i64>").unwrap();
429        let pair = pairs.next().unwrap();
430        assert_eq!(pairs.next(), None);
431        let t = parse_list_type(&extensions, pair).unwrap();
432        assert_eq!(
433            t,
434            Type {
435                kind: Some(Kind::List(Box::new(proto::r#type::List {
436                    nullability: Nullability::Required as i32,
437                    r#type: Some(Box::new(Type {
438                        kind: Some(Kind::I64(I64 {
439                            nullability: Nullability::Required as i32,
440                            type_variation_reference: 0,
441                        }))
442                    })),
443                    type_variation_reference: 0,
444                })))
445            }
446        );
447    }
448
449    #[test]
450    fn test_parse_interval_day_type() {
451        for (input, nullability, precision) in [
452            ("interval_day<9>", Nullability::Required, 9),
453            ("interval_day?<6>", Nullability::Nullable, 6),
454        ] {
455            let mut pairs = ExpressionParser::parse(Rule::interval_day_type, input).unwrap();
456            let pair = pairs.next().unwrap();
457            assert_eq!(pairs.next(), None);
458            let t = parse_interval_day_type(pair).unwrap();
459            assert_eq!(
460                t,
461                Type {
462                    kind: Some(Kind::IntervalDay(proto::r#type::IntervalDay {
463                        nullability: nullability as i32,
464                        type_variation_reference: 0,
465                        precision: Some(precision),
466                    })),
467                },
468                "input: {input}"
469            );
470        }
471    }
472
473    #[test]
474    fn test_parse_interval_day_type_invalid_precision() {
475        for input in ["interval_day<13>", "interval_day<-1>"] {
476            let mut pairs = ExpressionParser::parse(Rule::interval_day_type, input).unwrap();
477            let pair = pairs.next().unwrap();
478            assert_eq!(pairs.next(), None);
479            assert!(parse_interval_day_type(pair).is_err(), "input: {input}");
480        }
481    }
482
483    #[test]
484    fn test_parse_parameters() {
485        let extensions = SimpleExtensions::default();
486        let mut pairs = ExpressionParser::parse(Rule::parameters, "<i64?,string>").unwrap();
487        let pair = pairs.next().unwrap();
488        assert_eq!(pairs.next(), None);
489        let t = parse_parameters(&extensions, pair).unwrap();
490        assert_eq!(
491            t,
492            vec![
493                Parameter {
494                    parameter: Some(proto::r#type::parameter::Parameter::DataType(Type {
495                        kind: Some(Kind::I64(proto::r#type::I64 {
496                            nullability: Nullability::Nullable as i32,
497                            type_variation_reference: 0,
498                        })),
499                    })),
500                },
501                Parameter {
502                    parameter: Some(proto::r#type::parameter::Parameter::DataType(Type {
503                        kind: Some(Kind::String(proto::r#type::String {
504                            nullability: Nullability::Required as i32,
505                            type_variation_reference: 0,
506                        })),
507                    })),
508                },
509            ]
510        );
511    }
512
513    #[test]
514    fn test_udts() {
515        let mut extensions = SimpleExtensions::default();
516        extensions
517            .add_extension_urn("some_source".to_string(), 4)
518            .unwrap();
519        extensions
520            .add_extension(ExtensionKind::Type, 4, 42, "udt".to_string())
521            .unwrap();
522        let mut pairs = ExpressionParser::parse(Rule::user_defined_type, "udt#42<i64?>").unwrap();
523        let pair = pairs.next().unwrap();
524        assert_eq!(pairs.next(), None);
525
526        let t = parse_user_defined_type(&extensions, pair).unwrap();
527        assert_eq!(
528            t,
529            Type {
530                kind: Some(Kind::UserDefined(proto::r#type::UserDefined {
531                    type_reference: 42,
532                    type_variation_reference: 0,
533                    nullability: Nullability::Required as i32,
534                    type_parameters: vec![Parameter {
535                        parameter: Some(proto::r#type::parameter::Parameter::DataType(Type {
536                            kind: Some(Kind::I64(proto::r#type::I64 {
537                                nullability: Nullability::Nullable as i32,
538                                type_variation_reference: 0,
539                            })),
540                        })),
541                    }],
542                }))
543            }
544        );
545    }
546
547    #[test]
548    fn test_udts_with_u_prefix() {
549        // Extensions declared with "u!" prefix (e.g. "# 7 @ 4: u!json") normalize to
550        // bare name at storage time. Both "u!json" and "json" in the plan resolve to the
551        // same anchor.
552        let mut extensions = SimpleExtensions::default();
553        extensions
554            .add_extension_urn("some_source".to_string(), 4)
555            .unwrap();
556        extensions
557            .add_extension(ExtensionKind::Type, 4, 7, "u!json".to_string())
558            .unwrap();
559
560        // Plan uses u! prefix with explicit anchor
561        let pair = ExpressionParser::parse(Rule::user_defined_type, "u!json#7")
562            .unwrap()
563            .next()
564            .unwrap();
565        let t = parse_user_defined_type(&extensions, pair).unwrap();
566        assert_eq!(
567            t,
568            Type {
569                kind: Some(Kind::UserDefined(proto::r#type::UserDefined {
570                    type_reference: 7,
571                    type_variation_reference: 0,
572                    nullability: Nullability::Required as i32,
573                    type_parameters: vec![],
574                }))
575            }
576        );
577
578        // Plan uses bare name — also resolves to the same anchor
579        let pair = ExpressionParser::parse(Rule::user_defined_type, "json?")
580            .unwrap()
581            .next()
582            .unwrap();
583        let t = parse_user_defined_type(&extensions, pair).unwrap();
584        assert_eq!(
585            t,
586            Type {
587                kind: Some(Kind::UserDefined(proto::r#type::UserDefined {
588                    type_reference: 7,
589                    type_variation_reference: 0,
590                    nullability: Nullability::Nullable as i32,
591                    type_parameters: vec![],
592                }))
593            }
594        );
595
596        // Plan uses u! prefix with no anchor — anchor-free lookup must strip u! and find the type
597        let pair = ExpressionParser::parse(Rule::user_defined_type, "u!json")
598            .unwrap()
599            .next()
600            .unwrap();
601        let t = parse_user_defined_type(&extensions, pair).unwrap();
602        assert_eq!(
603            t,
604            Type {
605                kind: Some(Kind::UserDefined(proto::r#type::UserDefined {
606                    type_reference: 7,
607                    type_variation_reference: 0,
608                    nullability: Nullability::Required as i32,
609                    type_parameters: vec![],
610                }))
611            }
612        );
613    }
614
615    #[test]
616    fn test_udt_u_prefix_matches_plain_registration() {
617        // Extension registered without "u!" and plan uses "u!json" — must succeed.
618        let mut extensions = SimpleExtensions::default();
619        extensions.add_extension_urn("src".to_string(), 1).unwrap();
620        extensions
621            .add_extension(ExtensionKind::Type, 1, 3, "json".to_string())
622            .unwrap();
623
624        let pair = ExpressionParser::parse(Rule::user_defined_type, "u!json")
625            .unwrap()
626            .next()
627            .unwrap();
628        let t = parse_user_defined_type(&extensions, pair).unwrap();
629        assert_eq!(
630            t.kind,
631            Some(Kind::UserDefined(proto::r#type::UserDefined {
632                type_reference: 3,
633                type_variation_reference: 0,
634                nullability: Nullability::Required as i32,
635                type_parameters: vec![],
636            }))
637        );
638    }
639
640    #[test]
641    fn test_udt_plain_matches_u_prefix_registration() {
642        // Extension registered with "u!" prefix and plan uses bare "json" — must succeed.
643        let mut extensions = SimpleExtensions::default();
644        extensions.add_extension_urn("src".to_string(), 1).unwrap();
645        extensions
646            .add_extension(ExtensionKind::Type, 1, 3, "u!json".to_string())
647            .unwrap();
648
649        let pair = ExpressionParser::parse(Rule::user_defined_type, "json")
650            .unwrap()
651            .next()
652            .unwrap();
653        let t = parse_user_defined_type(&extensions, pair).unwrap();
654        assert_eq!(
655            t.kind,
656            Some(Kind::UserDefined(proto::r#type::UserDefined {
657                type_reference: 3,
658                type_variation_reference: 0,
659                nullability: Nullability::Required as i32,
660                type_parameters: vec![],
661            }))
662        );
663    }
664}