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, )