Source code for codegen_database.ext.text_range.types

from __future__ import annotations

import re
from typing import Any, cast

from asyncpg import Range
from sqlalchemy.types import TypeDecorator, UserDefinedType

__all__ = ["TextMultiRangeType"]


class _TextMultiRangeColumn(UserDefinedType):
    cache_ok = True

    def __init__(self, schema: str | None = None) -> None:
        self.pg_schema = schema
        super().__init__()

    def get_col_spec(self, **_kw: Any) -> str:  # noqa: ANN401
        if self.pg_schema:
            return f"{self.pg_schema}.textmultirange"

        return "textmultirange"


_RANGE_RE = re.compile(r"([\[(])([^,\]\)]*),([^,\]\)]*)([\])])")
_MIN_LITERAL_LEN = 2


def parse_multirange_literal(
    value: str,
) -> list[tuple[str | None, str | None, bool, bool]]:
    literal = value.strip()
    well_formed = (
        len(literal) >= _MIN_LITERAL_LEN
        and literal[0] == "{"
        and literal[-1] == "}"
    )

    if not well_formed:
        msg = f"malformed textmultirange literal: {value!r}"
        raise ValueError(msg)

    body = literal[1:-1].strip()

    if not body:
        return []

    matches = list(_RANGE_RE.finditer(body))
    residue = _RANGE_RE.sub("", body)

    if (
        not matches
        or residue.count(",") != len(matches) - 1
        or residue.strip(" ,")
    ):
        msg = f"malformed textmultirange literal: {value!r}"
        raise ValueError(msg)

    ranges: list[tuple[str | None, str | None, bool, bool]] = []

    for match in matches:
        lower = match.group(2) or None
        upper = match.group(3) or None
        ranges.append(
            (
                lower,
                upper,
                match.group(1) == "[",
                match.group(4) == "]",
            )
        )

    return ranges


def format_multirange(ranges: Any) -> str:  # noqa: ANN401
    parts: list[str] = []

    for rng in ranges:
        if getattr(rng, "empty", False) or getattr(rng, "isempty", False):
            continue

        lower_inc = bool(rng.lower_inc)
        upper_inc = bool(rng.upper_inc)
        left = "[" if lower_inc else "("
        right = "]" if upper_inc else ")"
        parts.append(f"{left}{rng.lower},{rng.upper}{right}")

    return "{" + ",".join(parts) + "}"


def literal_to_ranges(value: str) -> list[Any]:
    return [
        Range(
            lower=lower, upper=upper, lower_inc=lower_inc, upper_inc=upper_inc
        )
        for lower, upper, lower_inc, upper_inc in parse_multirange_literal(
            value
        )
    ]


[docs] class TextMultiRangeType(TypeDecorator[str | None]): impl = _TextMultiRangeColumn cache_ok = True def __init__(self, schema: str | None = None) -> None: super().__init__(schema=schema) @property def schema(self) -> str | None: return cast("_TextMultiRangeColumn", self.impl_instance).pg_schema
[docs] def process_bind_param(self, value: Any, dialect: Any) -> Any: # noqa: ANN401 if value is None: return None if isinstance(value, str): # asyncpg decodes/encodes multiranges as Range objects only; # psycopg binds the canonical literal as text natively. if dialect.driver == "asyncpg": return literal_to_ranges(value) return value return value
[docs] def process_result_value(self, value: Any, dialect: Any) -> Any: # noqa: ANN401, ARG002 if value is None or isinstance(value, str): return value return format_multirange(value)