Skip to content

Commit c12c02f

Browse files
committed
Add CAST expression support to parser
1 parent 82f6040 commit c12c02f

3 files changed

Lines changed: 198 additions & 3 deletions

File tree

‎pyiceberg/expressions/parser.py‎

Lines changed: 56 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
# under the License.
1717
import re
1818
from decimal import Decimal
19+
from typing import Any
1920

2021
from pyparsing import (
2122
CaselessKeyword,
@@ -66,6 +67,14 @@
6667
LongLiteral,
6768
StringLiteral,
6869
)
70+
from pyiceberg.transforms import (
71+
DayTransform,
72+
HourTransform,
73+
MonthTransform,
74+
Transform,
75+
UnboundTransform,
76+
YearTransform,
77+
)
6978
from pyiceberg.typedef import L
7079
from pyiceberg.types import strtobool
7180

@@ -80,6 +89,8 @@
8089
NAN = CaselessKeyword("nan")
8190
LIKE = CaselessKeyword("like")
8291
BETWEEN = CaselessKeyword("between")
92+
CAST = CaselessKeyword("cast")
93+
AS = CaselessKeyword("as")
8394

8495
unquoted_identifier = Word(alphas + "_", alphanums + "_$")
8596
quoted_identifier = QuotedString('"', esc_quote="\\", unquote_results=True)
@@ -265,7 +276,51 @@ def _evaluate_like_statement(result: ParseResults) -> BooleanExpression:
265276
return EqualTo(result.column, StringLiteral(literal_like.value.replace("\\%", "%")))
266277

267278

268-
predicate = (between | comparison | in_check | null_check | nan_check | starts_check | boolean).set_results_name("predicate")
279+
# CAST expression support: CAST(column AS type) maps to Iceberg transforms
280+
_CAST_TYPE_TO_TRANSFORM: dict[str, Transform[Any, Any]] = {
281+
"date": DayTransform(),
282+
"year": YearTransform(),
283+
"month": MonthTransform(),
284+
"hour": HourTransform(),
285+
}
286+
287+
cast_type = Word(alphas, alphanums + "_")
288+
cast_term = Suppress(CAST) + Suppress("(") + column + Suppress(AS) + cast_type + Suppress(")")
289+
290+
291+
@cast_term.set_parse_action
292+
def _(result: ParseResults) -> UnboundTransform:
293+
ref = result[0]
294+
target = str(result[1]).lower()
295+
if target not in _CAST_TYPE_TO_TRANSFORM:
296+
raise ValueError(f"Unsupported CAST target type: {target}")
297+
return UnboundTransform(ref, _CAST_TYPE_TO_TRANSFORM[target])
298+
299+
300+
cast_left_ref = cast_term + comparison_op + literal
301+
302+
303+
@cast_left_ref.set_parse_action
304+
def _(result: ParseResults) -> BooleanExpression:
305+
term = result[0]
306+
if result.op == "<":
307+
return LessThan(term, result.literal)
308+
elif result.op == "<=":
309+
return LessThanOrEqual(term, result.literal)
310+
elif result.op == ">":
311+
return GreaterThan(term, result.literal)
312+
elif result.op == ">=":
313+
return GreaterThanOrEqual(term, result.literal)
314+
if result.op in ("=", "=="):
315+
return EqualTo(term, result.literal)
316+
if result.op in ("!=", "<>"):
317+
return NotEqualTo(term, result.literal)
318+
raise ValueError(f"Unsupported operation type: {result.op}")
319+
320+
321+
predicate = (between | cast_left_ref | comparison | in_check | null_check | nan_check | starts_check | boolean).set_results_name(
322+
"predicate"
323+
)
269324

270325

271326
def handle_not(result: ParseResults) -> Not:

‎pyiceberg/transforms.py‎

Lines changed: 51 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@
4444
BoundNotIn,
4545
BoundNotStartsWith,
4646
BoundPredicate,
47+
BoundReference,
4748
BoundSetPredicate,
4849
BoundStartsWith,
4950
BoundTerm,
@@ -58,6 +59,7 @@
5859
Reference,
5960
StartsWith,
6061
UnboundPredicate,
62+
UnboundTerm,
6163
)
6264
from pyiceberg.expressions.literals import (
6365
DateLiteral,
@@ -67,7 +69,7 @@
6769
TimestampLiteral,
6870
literal,
6971
)
70-
from pyiceberg.typedef import IcebergRootModel, L
72+
from pyiceberg.typedef import IcebergRootModel, L, StructProtocol
7173
from pyiceberg.types import (
7274
BinaryType,
7375
DateType,
@@ -92,6 +94,8 @@
9294
if TYPE_CHECKING:
9395
import pyarrow as pa
9496

97+
from pyiceberg.schema import Schema
98+
9599
ArrayLike = TypeVar("ArrayLike", pa.Array, pa.ChunkedArray)
96100

97101
S = TypeVar("S")
@@ -1152,3 +1156,49 @@ class BoundTransform(BoundTerm):
11521156
def __init__(self, term: BoundTerm, transform: Transform[Any, Any]):
11531157
self.term: BoundTerm = term
11541158
self.transform = transform
1159+
1160+
def ref(self) -> BoundReference:
1161+
"""Return the bound reference of the underlying term."""
1162+
return self.term.ref()
1163+
1164+
def eval(self, struct: StructProtocol) -> Any:
1165+
"""Evaluate the transform on the struct value."""
1166+
return self.transform.transform(self.term.ref().field.field_type)(self.term.eval(struct))
1167+
1168+
1169+
class UnboundTransform(UnboundTerm):
1170+
"""An unbound transform expression that pairs a column reference with a transform.
1171+
1172+
When bound to a schema, validates that the transform is compatible with the
1173+
source column type and produces a BoundTransform.
1174+
"""
1175+
1176+
transform: Transform[Any, Any]
1177+
1178+
def __init__(self, term: UnboundTerm, transform: Transform[Any, Any]):
1179+
self.term: UnboundTerm = term
1180+
self.transform = transform
1181+
1182+
def bind(self, schema: "Schema", case_sensitive: bool = True) -> BoundTransform:
1183+
"""Bind this unbound transform to a schema, returning a BoundTransform."""
1184+
bound_term = self.term.bind(schema, case_sensitive)
1185+
source_type = bound_term.ref().field.field_type
1186+
if not self.transform.can_transform(source_type):
1187+
raise TypeError(f"Cannot apply {self.transform} to {source_type}")
1188+
return BoundTransform(bound_term, self.transform)
1189+
1190+
@property
1191+
def as_bound(self) -> type[BoundTransform]:
1192+
return BoundTransform
1193+
1194+
def __eq__(self, other: object) -> bool:
1195+
"""Return the equality of two instances of the UnboundTransform class."""
1196+
return isinstance(other, UnboundTransform) and self.term == other.term and self.transform == other.transform
1197+
1198+
def __repr__(self) -> str:
1199+
"""Return the string representation of the UnboundTransform class."""
1200+
return f"UnboundTransform(term={self.term}, transform={self.transform})"
1201+
1202+
def __hash__(self) -> int:
1203+
"""Return the hash of the UnboundTransform."""
1204+
return hash((self.term, self.transform))

‎tests/expressions/test_parser.py‎

Lines changed: 91 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -42,7 +42,14 @@
4242
Reference,
4343
StartsWith,
4444
)
45-
from pyiceberg.expressions.literals import DecimalLiteral, LongLiteral, literal
45+
from pyiceberg.expressions.literals import DecimalLiteral, LongLiteral, StringLiteral, literal
46+
from pyiceberg.transforms import (
47+
DayTransform,
48+
HourTransform,
49+
MonthTransform,
50+
UnboundTransform,
51+
YearTransform,
52+
)
4653

4754

4855
def test_always_true() -> None:
@@ -272,3 +279,86 @@ def test_valid_between_with_numerics() -> None:
272279
) == parser.parse("foo between '2025-01-01T00:00:00.000000' and '2025-01-10T12:00:00.000000'")
273280

274281
assert parser.parse("foo between 1 and 3") == parser.parse("1 <= foo and foo <= 3")
282+
283+
284+
def test_cast_date_comparison() -> None:
285+
expected = EqualTo(
286+
UnboundTransform(Reference("created_at"), DayTransform()),
287+
StringLiteral("2024-01-01"),
288+
)
289+
assert expected == parser.parse("CAST(created_at AS date) = '2024-01-01'")
290+
291+
292+
def test_cast_year_comparison() -> None:
293+
expected = GreaterThan(
294+
UnboundTransform(Reference("ts"), YearTransform()),
295+
LongLiteral(2020),
296+
)
297+
assert expected == parser.parse("CAST(ts AS year) > 2020")
298+
299+
300+
def test_cast_month_comparison() -> None:
301+
expected = EqualTo(
302+
UnboundTransform(Reference("ts"), MonthTransform()),
303+
LongLiteral(6),
304+
)
305+
assert expected == parser.parse("CAST(ts AS month) = 6")
306+
307+
308+
def test_cast_hour_comparison() -> None:
309+
expected = GreaterThanOrEqual(
310+
UnboundTransform(Reference("ts"), HourTransform()),
311+
LongLiteral(12),
312+
)
313+
assert expected == parser.parse("CAST(ts AS hour) >= 12")
314+
315+
316+
def test_cast_case_insensitive() -> None:
317+
"""CAST keyword and type name should be case-insensitive."""
318+
expected = EqualTo(
319+
UnboundTransform(Reference("created_at"), DayTransform()),
320+
StringLiteral("2024-01-01"),
321+
)
322+
assert expected == parser.parse("cast(created_at as DATE) = '2024-01-01'")
323+
324+
325+
def test_cast_nested_field() -> None:
326+
"""CAST should work with dotted column references."""
327+
expected = EqualTo(
328+
UnboundTransform(Reference("event.timestamp"), DayTransform()),
329+
StringLiteral("2024-01-01"),
330+
)
331+
assert expected == parser.parse("CAST(event.timestamp AS date) = '2024-01-01'")
332+
333+
334+
def test_cast_unsupported_type() -> None:
335+
with pytest.raises(ValueError, match="Unsupported CAST target type"):
336+
parser.parse("CAST(col AS foobar) = 5")
337+
338+
339+
def test_cast_not_equal() -> None:
340+
expected = NotEqualTo(
341+
UnboundTransform(Reference("ts"), YearTransform()),
342+
LongLiteral(2020),
343+
)
344+
assert expected == parser.parse("CAST(ts AS year) != 2020")
345+
346+
347+
def test_cast_less_than_or_equal() -> None:
348+
expected = LessThanOrEqual(
349+
UnboundTransform(Reference("ts"), MonthTransform()),
350+
LongLiteral(6),
351+
)
352+
assert expected == parser.parse("CAST(ts AS month) <= 6")
353+
354+
355+
def test_cast_with_and() -> None:
356+
"""CAST predicates compose with AND/OR via infix_notation."""
357+
result = parser.parse("CAST(ts AS date) = '2024-01-01' and status = 'active'")
358+
assert isinstance(result, And)
359+
360+
361+
def test_cast_with_not() -> None:
362+
"""NOT should negate a CAST predicate."""
363+
result = parser.parse("not CAST(ts AS date) = '2024-01-01'")
364+
assert isinstance(result, Not)

0 commit comments

Comments
 (0)