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)