Skip to content

Commit fdd0315

Browse files
authored
fix:ensure that format for times can be accessed and supplied in duckdb casting (#135)
1 parent 948d42a commit fdd0315

4 files changed

Lines changed: 93 additions & 42 deletions

File tree

src/dve/core_engine/backends/implementations/duckdb/duckdb_helpers.py

Lines changed: 24 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@
2424
from dve.core_engine.backends.utilities import DEFAULT_ISO_FORMATS, datetime_format_to_regex
2525
from dve.core_engine.constants import RECORD_INDEX_COLUMN_NAME
2626
from dve.core_engine.type_hints import URI, EntityName
27+
from dve.metadata_parser.utilities import resilient_get
2728
from dve.parser.file_handling.service import LocalFilesystemImplementation, _get_implementation
2829

2930

@@ -451,23 +452,27 @@ def get_duckdb_cast_statement_from_annotation(
451452
raise ValueError(f"dict must be `typing.TypedDict` subclass, got {type_annotation!r}")
452453

453454
for type_ in type_annotation.mro():
454-
_date_format: str = getattr( # type: ignore
455-
type_, "DATE_FORMAT", DEFAULT_ISO_FORMATS.get(type_, DEFAULT_ISO_FORMATS.get(datetime))
456-
)
457-
dt_cast_statement = rf"CASE WHEN REGEXP_FULL_MATCH(TRIM({quoted_name}), '{datetime_format_to_regex(_date_format)}') THEN TRY_STRPTIME(TRIM({quoted_name}), '{_date_format}') ELSE NULL END" # pylint: disable=C0301
458-
459-
# datetime is subclass of date, so needs to be handled first
460-
if issubclass(type_, datetime):
461-
stmt = rf"TRY_CAST({dt_cast_statement} as TIMESTAMP)"
462-
return stmt
463-
if issubclass(type_, date):
464-
stmt = rf"TRY_CAST({dt_cast_statement} as DATE)"
465-
return stmt
466-
if issubclass(type_, time):
467-
stmt = rf"TRY_CAST({dt_cast_statement} as TIME)"
468-
return stmt
469-
duck_type = get_duckdb_type_from_annotation(type_)
470-
if duck_type:
471-
stmt = f"TRIM({quoted_name})"
472-
return _cast_as_ddb_type(stmt, type_) if parent_element else stmt
455+
if issubclass(type_, (date, time)):
456+
_date_format: str = resilient_get(
457+
type_, "DATE_FORMAT", "TIME_FORMAT"
458+
) or DEFAULT_ISO_FORMATS.get(
459+
type_, DEFAULT_ISO_FORMATS.get(datetime)
460+
) # type: ignore
461+
dt_cast_statement = rf"CASE WHEN REGEXP_FULL_MATCH(TRIM({quoted_name}), '{datetime_format_to_regex(_date_format)}') THEN TRY_STRPTIME(TRIM({quoted_name}), '{_date_format}') ELSE NULL END" # pylint: disable=C0301
462+
463+
# datetime is subclass of date, so needs to be handled first
464+
if issubclass(type_, datetime):
465+
stmt = rf"TRY_CAST({dt_cast_statement} as TIMESTAMP)"
466+
return stmt
467+
if issubclass(type_, date):
468+
stmt = rf"TRY_CAST({dt_cast_statement} as DATE)"
469+
return stmt
470+
if issubclass(type_, time):
471+
stmt = rf"TRY_CAST({dt_cast_statement} as TIME)"
472+
return stmt
473+
else:
474+
duck_type = get_duckdb_type_from_annotation(type_)
475+
if duck_type:
476+
stmt = f"TRIM({quoted_name})"
477+
return _cast_as_ddb_type(stmt, type_) if parent_element else stmt
473478
raise ValueError(f"No equivalent DuckDB type for {type_annotation!r}")

src/dve/core_engine/backends/implementations/spark/spark_helpers.py

Lines changed: 25 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -593,29 +593,31 @@ def get_spark_cast_statement_from_annotation(
593593
raise ValueError(f"dict must be `typing.TypedDict` subclass, got {type_annotation!r}")
594594

595595
for type_ in type_annotation.mro():
596-
_date_format: str = getattr( # type: ignore
597-
type_,
598-
"DATE_FORMAT",
599-
DEFAULT_ISO_FORMATS.get(type_, DEFAULT_ISO_FORMATS.get(dt.datetime)),
600-
)
601-
602-
# pylint: disable=C0301
603-
dt_cast_statement = f"CASE WHEN REGEXP(TRIM({quoted_name}), '{datetime_format_to_regex(_date_format)}') THEN TRY_TO_TIMESTAMP(TRIM({quoted_name}), \"{python_to_java_datetime_format(_date_format)}\") ELSE NULL END" # pylint: disable=C0301
604-
# datetime is subclass of date, so needs to be handled first
605-
if issubclass(type_, dt.datetime):
606-
return (
607-
_cast_as_spark_type(dt_cast_statement, type_)
608-
if parent_element
609-
else dt_cast_statement
610-
)
611596
if issubclass(type_, dt.date):
612-
return (
613-
_cast_as_spark_type(dt_cast_statement, type_)
614-
if parent_element
615-
else dt_cast_statement
597+
_date_format: str = getattr( # type: ignore
598+
type_,
599+
"DATE_FORMAT",
600+
DEFAULT_ISO_FORMATS.get(type_, DEFAULT_ISO_FORMATS.get(dt.datetime)),
616601
)
617-
spark_type = get_type_from_annotation(type_)
618-
if spark_type:
619-
stmt = f"TRIM({quoted_name})"
620-
return _cast_as_spark_type(stmt, type_) if parent_element else stmt
602+
603+
# pylint: disable=C0301
604+
dt_cast_statement = f"CASE WHEN REGEXP(TRIM({quoted_name}), '{datetime_format_to_regex(_date_format)}') THEN TRY_TO_TIMESTAMP(TRIM({quoted_name}), \"{python_to_java_datetime_format(_date_format)}\") ELSE NULL END" # pylint: disable=C0301
605+
# datetime is subclass of date, so needs to be handled first
606+
if issubclass(type_, dt.datetime):
607+
return (
608+
_cast_as_spark_type(dt_cast_statement, type_)
609+
if parent_element
610+
else dt_cast_statement
611+
)
612+
if issubclass(type_, dt.date):
613+
return (
614+
_cast_as_spark_type(dt_cast_statement, type_)
615+
if parent_element
616+
else dt_cast_statement
617+
)
618+
else:
619+
spark_type = get_type_from_annotation(type_)
620+
if spark_type:
621+
stmt = f"TRIM({quoted_name})"
622+
return _cast_as_spark_type(stmt, type_) if parent_element else stmt
621623
raise ValueError(f"No equivalent Spark type for {type_annotation!r}")

src/dve/metadata_parser/utilities.py

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -52,3 +52,23 @@ def chain_get(
5252
return result
5353

5454
raise exc.TypeNotFoundError(f"Callable or type ({item!r}) not found")
55+
56+
57+
def resilient_get(item: object, *attribute_names: str) -> Any:
58+
"""Given a number of attribute names, try to get attribute value
59+
sequentially. Returns the first value found, and if no attributes found
60+
returns None.
61+
62+
Args:
63+
item (object): The object to obtain attributes from (where possible)
64+
attribute_names (str): The attribute names to search for
65+
66+
Returns:
67+
Any: The first found attribute, otherwise None
68+
"""
69+
for attr in attribute_names:
70+
try:
71+
return getattr(item, attr)
72+
except AttributeError:
73+
continue
74+
return None

tests/test_parser/test_utils.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,24 @@
1+
import pytest
2+
from dve.metadata_parser.utilities import resilient_get
3+
4+
class MyParent:
5+
cls_attr = "hello"
6+
def __init__(self, my_attr:str, another_attr:str):
7+
self.my_attr = my_attr
8+
self.another_attr = another_attr
9+
10+
class MyObject(MyParent):
11+
sub_attr = "bye"
12+
def __init__(self, extra_attr:int):
13+
self.extra_attr = extra_attr
14+
super().__init__("from", "child")
15+
16+
17+
@pytest.mark.parametrize("obj,attrs,expected", [(MyParent, ("cls_attr",), "hello"),
18+
(MyObject, ("cls_attr", "sub_attr"), "hello"),
19+
(MyObject, ("sub_attr", "cls_attr"), "bye"),
20+
(MyParent, ("my_attr",), None),
21+
(MyParent("this", "test"), ("extra_attr", "my_attr"), "this"),
22+
(MyObject("this"), ("daft_attr", "another_daft_attr", "yet_another", "another_attr"), "child")])
23+
def test_resilient_get(obj, attrs, expected):
24+
assert resilient_get(obj, *attrs) == expected

0 commit comments

Comments
 (0)