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)