Source code for codegen_database.types.ranges
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Literal, cast
import infinity
import intervals
from sqlalchemy.dialects.postgresql import TSTZRANGE
from sqlalchemy.dialects.postgresql.ranges import Range as PostgresRange
from sqlalchemy_utils.types.range import RangeType
if TYPE_CHECKING:
from sqlalchemy.engine import Dialect
_BoundsType = Literal["()", "[)", "(]", "[]"]
def _is_infinite(bound: object) -> bool:
return bound is None or infinity.is_infinite(bound)
[docs]
class TZDateTimeRangeType(RangeType):
impl = TSTZRANGE
cache_ok = True
def __init__(self) -> None:
super().__init__()
self.interval_class = intervals.DateTimeInterval
[docs]
def process_bind_param(
self,
value: intervals.DateTimeInterval | None,
dialect: Dialect, # noqa: ARG002
) -> PostgresRange | None:
if value is None:
return None
lower = None if _is_infinite(value.lower) else value.lower
upper = None if _is_infinite(value.upper) else value.upper
open_bracket = "[" if value.lower_inc and lower is not None else "("
close_bracket = "]" if value.upper_inc and upper is not None else ")"
bounds = cast("_BoundsType", f"{open_bracket}{close_bracket}")
return PostgresRange(
lower,
upper,
bounds=bounds,
)
[docs]
def process_result_value(
self,
value: Any, # noqa: ANN401
dialect: Dialect, # noqa: ARG002
) -> intervals.DateTimeInterval | None:
if value is None:
return None
if isinstance(value, str):
return self.interval_class.from_string(value)
return self.interval_class(
[value.lower, value.upper],
lower_inc=value.lower_inc,
upper_inc=value.upper_inc,
)