Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 11 additions & 0 deletions src/dialect/databricks.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
// under the License.

use crate::dialect::Dialect;
use crate::keywords::{self, Keyword};

/// A [`Dialect`] for [Databricks SQL](https://www.databricks.com/)
///
Expand Down Expand Up @@ -75,6 +76,16 @@ impl Dialect for DatabricksDialect {
true
}

/// See <https://docs.databricks.com/aws/en/sql/language-manual/data-types/interval-type>
fn supports_interval_string_without_qualifier(&self) -> bool {
true
}

/// `SELECT interval FROM t` names a column.
fn is_reserved_for_identifier(&self, kw: Keyword) -> bool {
kw != Keyword::INTERVAL && keywords::RESERVED_FOR_IDENTIFIER.contains(&kw)
}

// See https://docs.databricks.com/en/sql/language-manual/functions/struct.html
fn supports_struct_literal(&self) -> bool {
true
Expand Down
10 changes: 10 additions & 0 deletions src/dialect/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1228,6 +1228,16 @@ pub trait Dialect: Debug + Any {
false
}

/// Returns true if an interval string may omit the qualifier,
/// e.g. `INTERVAL '2 months'`. The value must then be a literal.
///
/// When `true`:
/// * `INTERVAL '2 months'` and `INTERVAL -'1' DAY` are VALID
/// * `INTERVAL 1` and `INTERVAL 1 + 1 DAY` are INVALID
fn supports_interval_string_without_qualifier(&self) -> bool {
false
}

/// Returns true if the dialect supports `EXPLAIN` statements with utility options
/// e.g. `EXPLAIN (ANALYZE TRUE, BUFFERS TRUE) SELECT * FROM tbl;`
fn supports_explain_with_utility_options(&self) -> bool {
Expand Down
12 changes: 11 additions & 1 deletion src/dialect/spark.rs
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ use alloc::boxed::Box;

use crate::ast::{BinaryOperator, Expr};
use crate::dialect::Dialect;
use crate::keywords::Keyword;
use crate::keywords::{self, Keyword};
use crate::parser::{Parser, ParserError};

/// A [`Dialect`] for [Apache Spark SQL](https://spark.apache.org/docs/latest/sql-ref.html).
Expand Down Expand Up @@ -104,6 +104,16 @@ impl Dialect for SparkSqlDialect {
true
}

/// See <https://spark.apache.org/docs/latest/sql-ref-literals.html#interval-literal>
fn supports_interval_string_without_qualifier(&self) -> bool {
true
}

/// `SELECT interval FROM t` names a column.
fn is_reserved_for_identifier(&self, kw: Keyword) -> bool {
kw != Keyword::INTERVAL && keywords::RESERVED_FOR_IDENTIFIER.contains(&kw)
}

fn supports_bang_not_operator(&self) -> bool {
true
}
Expand Down
60 changes: 57 additions & 3 deletions src/parser/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1902,7 +1902,16 @@ impl<'a> Parser<'a> {
Err(e) => {
self.failed_reserved_word_prefix_positions
.insert(next_token_index, (&e).into());
if !self.dialect.is_reserved_for_identifier(w.keyword) {
// `INTERVAL 1` is a malformed interval, not a name.
let interval_number = w.keyword == Keyword::INTERVAL
&& match &self.peek_token_ref().token {
Token::Number(..) => true,
Token::Plus | Token::Minus => {
matches!(self.peek_nth_token_ref(1).token, Token::Number(..))
}
_ => false,
};
if !interval_number && !self.dialect.is_reserved_for_identifier(w.keyword) {
if let Ok(Some(expr)) = self.maybe_parse(|parser| {
parser.parse_expr_prefix_by_unreserved_word(&w, span)
}) {
Expand Down Expand Up @@ -3431,14 +3440,30 @@ impl<'a> Parser<'a> {
// to match the different flavours of INTERVAL syntax, we only allow expressions
// if the dialect requires an interval qualifier,
// see https://github.com/sqlparser-rs/sqlparser-rs/pull/1398 for more details
let value = if self.dialect.require_interval_qualifier() {
let literal_value = self.dialect.supports_interval_string_without_qualifier();
let value = if literal_value {
// otherwise `interval` is read as a name
self.parse_interval_literal_value()?
} else if self.dialect.require_interval_qualifier() {
// parse a whole expression so `INTERVAL 1 + 1 DAY` is valid
self.parse_expr()?
} else {
// parse a prefix expression so `INTERVAL 1 DAY` is valid, but `INTERVAL 1 + 1 DAY` is not
// this also means that `INTERVAL '5 days' > INTERVAL '1 day'` treated properly
self.parse_prefix()?
};
// only an unsigned string may omit the qualifier
let qualifier_required = if literal_value {
!matches!(
&value,
Expr::Value(ValueWithSpan {
value: Value::SingleQuotedString(_),
..
})
)
} else {
self.dialect.require_interval_qualifier()
};

// Following the string literal is a qualifier which indicates the units
// of the duration specified in the string literal.
Expand All @@ -3447,7 +3472,7 @@ impl<'a> Parser<'a> {
// this more general implementation.
let leading_field = if self.next_token_is_temporal_unit() {
Some(self.parse_date_time_field()?)
} else if self.dialect.require_interval_qualifier() {
} else if qualifier_required {
return parser_err!(
"INTERVAL requires a unit after the literal value",
self.peek_token_ref().span.start
Expand Down Expand Up @@ -3490,6 +3515,35 @@ impl<'a> Parser<'a> {
}))
}

/// A signed number or string literal.
fn parse_interval_literal_value(&mut self) -> Result<Expr, ParserError> {
let sign = match self.peek_token_ref().token {
Token::Plus => Some(UnaryOperator::Plus),
Token::Minus => Some(UnaryOperator::Minus),
_ => None,
};
if sign.is_some() {
self.advance_token();
}
let next = self.next_token();
let value = match next.token {
Token::Number(n, l) => {
Expr::value(Value::Number(Self::parse(n, next.span.start)?, l).with_span(next.span))
}
Token::SingleQuotedString(s) => {
Expr::value(Value::SingleQuotedString(s).with_span(next.span))
}
_ => return self.expected("an expression", next),
};
Ok(match sign {
Some(op) => Expr::UnaryOp {
op,
expr: Box::new(value),
},
None => value,
})
}

/// Peek at the next token and determine if it is a temporal unit
/// like `second`.
pub fn next_token_is_temporal_unit(&mut self) -> bool {
Expand Down
63 changes: 61 additions & 2 deletions tests/sqlparser_common.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6755,7 +6755,9 @@ fn parse_interval_dont_require_unit() {

#[test]
fn parse_interval_require_unit() {
let dialects = all_dialects_where(|d| d.require_interval_qualifier());
let dialects = all_dialects_where(|d| {
d.require_interval_qualifier() && !d.supports_interval_string_without_qualifier()
});
let sql = "SELECT INTERVAL '1 DAY'";
let err = dialects.parse_sql_statements(sql).unwrap_err();
assert_eq!(
Expand All @@ -6764,9 +6766,66 @@ fn parse_interval_require_unit() {
)
}

#[test]
fn parse_interval_string_without_qualifier() {
let dialects = all_dialects_where(|d| d.supports_interval_string_without_qualifier());
for sql in [
"SELECT INTERVAL '2 months'",
"SELECT INTERVAL '1' DAY",
"SELECT INTERVAL '1-2' YEAR TO MONTH",
"SELECT INTERVAL -'1' DAY",
"SELECT INTERVAL 3 DAY",
"SELECT INTERVAL -1 DAY, INTERVAL +2 HOURS",
"SELECT INTERVAL '2 seconds' * 2",
] {
dialects.verified_stmt(sql);
}
for sql in [
"SELECT INTERVAL 1",
"SELECT INTERVAL 1 + 1 DAY",
"SELECT INTERVAL x DAY",
] {
assert!(
dialects.parse_sql_statements(sql).is_err(),
"{sql} should not parse"
);
}
dialects.verified_stmt("SELECT interval, x FROM t");
dialects.verified_stmt("SELECT max(interval) FROM t WHERE interval > 1");
dialects.one_statement_parses_to(
"SELECT INTERVAL -'2 months'",
"SELECT INTERVAL - '2 months'",
);
}

#[test]
fn parse_interval_number_without_unit_error() {
for sql in ["SELECT INTERVAL 1", "SELECT INTERVAL -1"] {
for dialect in all_dialects().dialects {
if let Err(e) = Parser::parse_sql(&*dialect, sql) {
assert_eq!(
e.to_string(),
"sql parser error: INTERVAL requires a unit after the literal value",
"{sql} with {dialect:?}"
);
}
}
}
}

#[test]
fn parse_interval_expression_value() {
let dialects = all_dialects_where(|d| {
d.require_interval_qualifier() && !d.supports_interval_string_without_qualifier()
});
dialects.verified_stmt("SELECT INTERVAL 1 + 1 DAY");
}

#[test]
fn parse_interval_require_qualifier() {
let dialects = all_dialects_where(|d| d.require_interval_qualifier());
let dialects = all_dialects_where(|d| {
d.require_interval_qualifier() && !d.supports_interval_string_without_qualifier()
});

let sql = "SELECT INTERVAL 1 + 1 DAY";
let select = dialects.verified_only_select(sql);
Expand Down
61 changes: 61 additions & 0 deletions tests/sqlparser_databricks.rs
Original file line number Diff line number Diff line change
Expand Up @@ -769,3 +769,64 @@ fn parse_databricks_collated_data_types() {
.parse_sql_statements("CREATE TABLE t (c ARRAY<STRING COLLATE UTF8_LCASE>)")
.is_err());
}

#[test]
fn test_interval_literals() {
for sql in [
"SELECT INTERVAL '2 months'",
"SELECT INTERVAL '-1 day 1 hour'",
"SELECT INTERVAL '1' DAY",
"SELECT INTERVAL 3 DAY",
"SELECT INTERVAL '1-2' YEAR TO MONTH",
"SELECT -INTERVAL '1 year'",
"SELECT INTERVAL '2 seconds' * 2",
"SELECT d + INTERVAL '1' DAY FROM t",
] {
databricks().verified_stmt(sql);
}
match databricks().verified_expr("INTERVAL '2 months'") {
Expr::Interval(i) => assert!(i.leading_field.is_none()),
other => panic!("Expected an interval, got {other:?}"),
}
}

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
}
}
#[test]
fn test_interval_numeric_requires_unit() {
assert_eq!(
databricks()
.parse_sql_statements("SELECT INTERVAL 1")
.unwrap_err()
.to_string(),
"sql parser error: INTERVAL requires a unit after the literal value"
);
}
#[test]
fn test_interval_rejects_value_arithmetic() {
assert!(databricks()
.parse_sql_statements("SELECT INTERVAL 1 + 1 DAY")
.is_err());
}
#[test]
fn test_interval_preserves_column_query() {
let query = databricks()
.run_parser_method("SELECT interval FROM t", |parser| parser.parse_query())
.unwrap();
let SetExpr::Select(select) = *query.body else {
panic!("Expected SELECT");
};
assert_eq!(
select.from,
vec![TableWithJoins {
relation: table("t"),
joins: vec![],
}]
);
assert_eq!(
select.projection,
vec![SelectItem::UnnamedExpr(Expr::Identifier(Ident::new(
"interval"
)))]
);
}


#[test]
fn test_interval_numeric_requires_unit() {
assert_eq!(
databricks()
.parse_sql_statements("SELECT INTERVAL 1")
.unwrap_err()
.to_string(),
"sql parser error: INTERVAL requires a unit after the literal value"
);
}

#[test]
fn test_interval_rejects_value_arithmetic() {
assert!(databricks()
.parse_sql_statements("SELECT INTERVAL 1 + 1 DAY")
.is_err());
}

#[test]
fn test_interval_preserves_column_query() {
let query = databricks()
.run_parser_method("SELECT interval FROM t", |parser| parser.parse_query())
.unwrap();
let SetExpr::Select(select) = *query.body else {
panic!("Expected SELECT");
};
assert_eq!(
select.from,
vec![TableWithJoins {
relation: table("t"),
joins: vec![],
}]
);
assert_eq!(
select.projection,
vec![SelectItem::UnnamedExpr(Expr::Identifier(Ident::new(
"interval"
)))]
);
}
61 changes: 61 additions & 0 deletions tests/sqlparser_spark.rs
Original file line number Diff line number Diff line change
Expand Up @@ -362,3 +362,64 @@ fn test_substring() {
fn test_pipe_operator() {
spark().verified_stmt("SELECT * FROM t |> WHERE x > 1 |> SELECT x AS y |> ORDER BY y");
}

#[test]
fn test_interval_literals() {
for sql in [
"SELECT INTERVAL '2 months'",
"SELECT INTERVAL '-1 day 1 hour'",
"SELECT INTERVAL '1' DAY",
"SELECT INTERVAL 3 DAY",
"SELECT INTERVAL '1-2' YEAR TO MONTH",
"SELECT -INTERVAL '1 year'",
"SELECT INTERVAL '2 seconds' * 2",
"SELECT d + INTERVAL '1' DAY FROM t",
] {
spark().verified_stmt(sql);
}
match spark().verified_expr("INTERVAL '2 months'") {
Expr::Interval(i) => assert!(i.leading_field.is_none()),
other => panic!("Expected an interval, got {other:?}"),
}
}

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
}
}
#[test]
fn test_interval_numeric_requires_unit() {
assert_eq!(
spark()
.parse_sql_statements("SELECT INTERVAL 1")
.unwrap_err()
.to_string(),
"sql parser error: INTERVAL requires a unit after the literal value"
);
}
#[test]
fn test_interval_rejects_value_arithmetic() {
assert!(spark()
.parse_sql_statements("SELECT INTERVAL 1 + 1 DAY")
.is_err());
}
#[test]
fn test_interval_preserves_column_query() {
let query = spark()
.run_parser_method("SELECT interval FROM t", |parser| parser.parse_query())
.unwrap();
let SetExpr::Select(select) = *query.body else {
panic!("Expected SELECT");
};
assert_eq!(
select.from,
vec![TableWithJoins {
relation: table("t"),
joins: vec![],
}]
);
assert_eq!(
select.projection,
vec![SelectItem::UnnamedExpr(Expr::Identifier(Ident::new(
"interval"
)))]
);
}


#[test]
fn test_interval_numeric_requires_unit() {
assert_eq!(
spark()
.parse_sql_statements("SELECT INTERVAL 1")
.unwrap_err()
.to_string(),
"sql parser error: INTERVAL requires a unit after the literal value"
);
}

#[test]
fn test_interval_rejects_value_arithmetic() {
assert!(spark()
.parse_sql_statements("SELECT INTERVAL 1 + 1 DAY")
.is_err());
}

#[test]
fn test_interval_preserves_column_query() {
let query = spark()
.run_parser_method("SELECT interval FROM t", |parser| parser.parse_query())
.unwrap();
let SetExpr::Select(select) = *query.body else {
panic!("Expected SELECT");
};
assert_eq!(
select.from,
vec![TableWithJoins {
relation: table("t"),
joins: vec![],
}]
);
assert_eq!(
select.projection,
vec![SelectItem::UnnamedExpr(Expr::Identifier(Ident::new(
"interval"
)))]
);
}
Loading