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
10pub(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 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::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
200fn 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 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 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 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 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 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 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 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 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}