Source code for codegen_database.types.duration
from __future__ import annotations
import re
from typing import TYPE_CHECKING
from typing import cast as type_cast
from sqlalchemy import Integer, Text, case, cast, extract, func, type_coerce
from sqlalchemy import types as sa_types
from sqlalchemy.types import UserDefinedType
if TYPE_CHECKING:
from sqlalchemy.engine import Dialect
from sqlalchemy.sql.elements import ColumnElement
_DURATION = re.compile(r"(?:P(?=\d+[YMD])(?:\d+Y)?(?:\d+M)?(?:\d+D)?|P\d+W)\Z")
def validate_duration(value: object) -> str:
if not isinstance(value, str):
msg = f"Expected str, got {type(value).__name__}"
raise ValueError(msg)
if not _DURATION.fullmatch(value):
msg = f"Invalid duration: {value!r}"
raise ValueError(msg)
return value
class _DurationIntervalColumn(UserDefinedType[str]):
cache_ok = True
def get_col_spec(self, **_kw: object) -> str:
return "INTERVAL"
[docs]
class DurationInterval(sa_types.TypeDecorator[str]):
impl = _DurationIntervalColumn
cache_ok = True
[docs]
def bind_expression(
self, bindvalue: ColumnElement[str]
) -> ColumnElement[str]:
return type_cast(
"ColumnElement[str]",
cast(cast(bindvalue, Text()), _DurationIntervalColumn()),
)
[docs]
def column_expression(
self, column: ColumnElement[str]
) -> ColumnElement[str]:
value = func.concat(
"P",
cast(extract("year", column), Integer),
"Y",
cast(extract("month", column), Integer),
"M",
cast(extract("day", column), Integer),
"D",
)
return type_coerce(case((column.is_(None), None), else_=value), Text())
[docs]
def process_bind_param(
self,
value: str | None,
_dialect: Dialect,
) -> str | None:
if value is None:
return None
return validate_duration(value)