"""Athena DATE and TIMESTAMP types and literal conversion."""
from __future__ import annotations
from collections.abc import Callable
from datetime import date, datetime
from functools import partial
from typing import TYPE_CHECKING, Any
from sqlalchemy import types
from sqlalchemy.sql.type_api import TypeEngine
from pyathena.formatter import _date_literal, _escape_trino, _timestamp_literal
from pyathena.util import override
if TYPE_CHECKING:
from sqlalchemy import Dialect
from sqlalchemy.sql.operators import OperatorType
from sqlalchemy.sql.type_api import _BindProcessorType, _LiteralProcessorType
[docs]
class AthenaTimestamp(TypeEngine[datetime]):
"""SQLAlchemy type for Athena TIMESTAMP values.
This type handles the conversion of Python datetime objects to Athena's
TIMESTAMP literal syntax. When used in queries, datetime values are
rendered as ``TIMESTAMP 'YYYY-MM-DD HH:MM:SS.mmm'``, or with six
fractional digits (``timestamp(6)``) when the value has a sub-millisecond
part. Iceberg tables store microseconds; Hive tables store milliseconds.
With a ``precision``, bound values and literals, including values compared
with a column of this type, are truncated to that many fractional digits,
and casts render ``TIMESTAMP(precision)``. Without one, casts render
``TIMESTAMP(6)``. ``CREATE TABLE`` always renders ``TIMESTAMP``.
Example:
>>> from sqlalchemy import Column, Table, MetaData, cast, select
>>> from pyathena.sqlalchemy.types import AthenaTimestamp
>>> metadata = MetaData()
>>> events = Table('events', metadata,
... Column('event_time', AthenaTimestamp)
... )
>>> millis = select(cast(events.c.event_time, AthenaTimestamp(precision=3)))
"""
__visit_name__ = "TIMESTAMP"
[docs]
def __init__(self, precision: int | None = None) -> None:
"""Initialize the type.
Args:
precision: The number of fractional-second digits, from 0 to 6,
or None for the default rendering.
Raises:
ValueError: If ``precision`` is not an integer from 0 to 6.
"""
if precision is not None and (
not isinstance(precision, int) or isinstance(precision, bool) or not 0 <= precision <= 6
):
raise ValueError(f"TIMESTAMP precision must be an integer from 0 to 6: {precision!r}")
self.precision = precision
@property
@override
def python_type(self) -> type[datetime]:
"""The Python type of TIMESTAMP values.
Returns:
``datetime.datetime``.
"""
return datetime
[docs]
@override
def bind_processor(self, dialect: Dialect) -> _BindProcessorType[datetime] | None:
"""Return a processor truncating bound datetimes to the precision.
Args:
dialect: The dialect binding the value.
Returns:
The processor, or None without a precision below 6.
"""
if self.precision is None or self.precision == 6:
return None
unit = 10 ** (6 - self.precision)
def process(value: datetime | Any | None) -> datetime | Any | None:
if isinstance(value, datetime):
return value.replace(microsecond=value.microsecond // unit * unit)
return value
return process
[docs]
@override
def coerce_compared_value(self, op: OperatorType | None, value: Any) -> TypeEngine[Any]:
"""Keep this type for a datetime compared with a column of it.
Args:
op: The comparison operator.
value: The compared value.
Returns:
This type for a datetime, so its precision applies.
"""
if isinstance(value, datetime):
return self
return super().coerce_compared_value(op, value)
[docs]
@staticmethod
def process(
value: datetime | Any | None,
quote: Callable[[str], str] = _escape_trino,
precision: int | None = None,
) -> str:
"""Render a value as an Athena TIMESTAMP literal.
Args:
value: A datetime, or any other value rendered with ``str()``.
quote: The function quoting a value that is not a datetime.
precision: The number of fractional-second digits for a datetime,
or None for the default rendering.
Returns:
The TIMESTAMP literal.
"""
if isinstance(value, datetime):
return _timestamp_literal(value, precision)
return f"TIMESTAMP {quote(str(value))}"
[docs]
@override
def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[datetime] | None:
"""Return the literal renderer for the dialect.
Args:
dialect: The dialect compiling the statement.
Returns:
A function rendering a value as a TIMESTAMP literal.
"""
return partial(self.process, quote=_string_quote(dialect), precision=self.precision)
[docs]
class AthenaDate(TypeEngine[date]):
"""SQLAlchemy type for Athena DATE values.
This type handles the conversion of Python date objects to Athena's
DATE literal syntax. When used in queries, date values are rendered
as ``DATE 'YYYY-MM-DD'``.
Example:
>>> from sqlalchemy import Column, Table, MetaData
>>> from pyathena.sqlalchemy.types import AthenaDate
>>> metadata = MetaData()
>>> orders = Table('orders', metadata,
... Column('order_date', AthenaDate)
... )
"""
__visit_name__ = "DATE"
@property
@override
def python_type(self) -> type[date]:
"""The Python type of DATE values.
Returns:
``datetime.date``.
"""
return date
[docs]
@staticmethod
def process(value: date | Any, quote: Callable[[str], str] = _escape_trino) -> str:
"""Render a value as an Athena DATE literal.
Args:
value: A date, or any other value rendered with ``str()``.
quote: The function quoting a value that is not a date.
Returns:
The DATE literal.
"""
# datetime is a subclass of date, so this branch also covers datetime,
# which is truncated to its date part.
if isinstance(value, date):
return _date_literal(value)
return f"DATE {quote(str(value))}"
[docs]
@override
def literal_processor(self, dialect: Dialect) -> _LiteralProcessorType[date] | None:
"""Return the literal renderer for the dialect.
Args:
dialect: The dialect compiling the statement.
Returns:
A function rendering a value as a DATE literal.
"""
return partial(self.process, quote=_string_quote(dialect))
def _string_quote(dialect: Dialect) -> Callable[[str], str]:
"""Return the dialect's string literal renderer.
It also doubles ``%`` for dialects whose paramstyle needs it.
Args:
dialect: The dialect compiling the statement.
Returns:
A function rendering a string as a quoted SQL literal.
"""
return types.String().literal_processor(dialect) or _escape_trino