diff --git a/src/parser/mod.rs b/src/parser/mod.rs index 294a0bed9..0860179e3 100644 --- a/src/parser/mod.rs +++ b/src/parser/mod.rs @@ -12699,9 +12699,12 @@ impl<'a> Parser<'a> { Ok(ty) } + #[cfg_attr(feature = "recursive-protection", recursive::recursive)] fn parse_data_type_helper( &mut self, ) -> Result<(DataType, MatchedTrailingBracket), ParserError> { + let _guard = self.recursion_counter.try_decrease()?; + let dialect = self.dialect; self.advance_token(); let next_token = self.get_current_token(); diff --git a/tests/sqlparser_common.rs b/tests/sqlparser_common.rs index 5ff5e5c47..6489d4306 100644 --- a/tests/sqlparser_common.rs +++ b/tests/sqlparser_common.rs @@ -11359,6 +11359,24 @@ fn parse_deeply_nested_subquery_expr_hits_recursion_limits() { assert_eq!(res, Err(ParserError::RecursionLimitExceeded)); } +#[test] +fn parse_deeply_nested_data_type_hits_recursion_limits() { + let dialect = GenericDialect {}; + + let sql = format!( + "SELECT CAST(x AS {}INT64{})", + "ARRAY<".repeat(1000), + ">".repeat(1000) + ); + + let res = Parser::new(&dialect) + .try_with_sql(&sql) + .expect("tokenize to work") + .parse_statements(); + + assert_eq!(res, Err(ParserError::RecursionLimitExceeded)); +} + #[test] fn parse_with_recursion_limit() { let dialect = GenericDialect {};