Source code for codegen_database.types.enums

from __future__ import annotations

import enum
from typing import Any

from sqlalchemy import Integer, Text
from sqlalchemy import types as sa_types


[docs] class TextEnum(sa_types.TypeDecorator[enum.Enum]): impl = Text cache_ok = True def __init__(self, enum_class: type[enum.Enum]) -> None: if not ( isinstance(enum_class, type) and issubclass(enum_class, enum.Enum) ): msg = ( # type: ignore[unreachable] f"TextEnum expects an Enum subclass, got {enum_class!r}" ) raise TypeError(msg) self.enum_class = enum_class super().__init__()
[docs] def process_bind_param( self, value: Any, # noqa: ANN401 dialect: Any, # noqa: ANN401, ARG002 ) -> str | None: if value is None: return None if isinstance(value, self.enum_class): return value.value if isinstance(value, str): self.enum_class(value) # type: ignore[misc] return value msg = ( f"Expected {self.enum_class.__name__} or str, " f"got {type(value).__name__}" ) raise ValueError(msg)
[docs] def process_result_value( self, value: Any, # noqa: ANN401 dialect: Any, # noqa: ANN401, ARG002 ) -> enum.Enum | None: if value is None: return None return self.enum_class(value) # type: ignore[misc]
[docs] class IntEnum(sa_types.TypeDecorator[enum.Enum]): impl = Integer cache_ok = True def __init__(self, enum_class: type[enum.Enum]) -> None: if not ( isinstance(enum_class, type) and issubclass(enum_class, enum.Enum) ): msg = ( # type: ignore[unreachable] f"IntEnum expects an Enum subclass, got {enum_class!r}" ) raise TypeError(msg) self.enum_class = enum_class super().__init__()
[docs] def process_bind_param( self, value: Any, # noqa: ANN401 dialect: Any, # noqa: ANN401, ARG002 ) -> int | None: if value is None: return None if isinstance(value, self.enum_class): return value.value if isinstance(value, int): self.enum_class(value) # type: ignore[misc] return value msg = ( f"Expected {self.enum_class.__name__} or int, " f"got {type(value).__name__}" ) raise ValueError(msg)
[docs] def process_result_value( self, value: Any, # noqa: ANN401 dialect: Any, # noqa: ANN401, ARG002 ) -> enum.Enum | None: if value is None: return None return self.enum_class(value) # type: ignore[misc]