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]