"""Explicit PostgreSQL function registration."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import TYPE_CHECKING, ClassVar
from sqlalchemy_declarative_extensions import register_function
from sqlalchemy_declarative_extensions.dialects.postgresql import (
Function,
FunctionParam,
FunctionSecurity,
FunctionVolatility,
)
from codegen_database.errors import CodegenDatabaseValidationError
if TYPE_CHECKING:
from collections.abc import Callable
from sqlalchemy import MetaData
[docs]
@dataclass
class CodegenDatabaseFunctionSpec:
name: str
definition: str | Callable[[], str]
language: str
parameters: list[FunctionParam]
returns: str = "void"
security: FunctionSecurity = FunctionSecurity.invoker
volatility: FunctionVolatility = FunctionVolatility.VOLATILE
def construct_function(self, schema: str) -> Function:
return Function(
self.name,
self.definition() if callable(self.definition) else self.definition,
returns=self.returns,
language=self.language,
schema=schema,
parameters=self.parameters,
security=self.security,
volatility=self.volatility,
)
[docs]
@dataclass(frozen=True)
class FunctionOptions:
"""Configure a declarative PostgreSQL function."""
returns: str = "void"
language: str = "sql"
parameters: list[FunctionParam] = field(default_factory=list)
security: FunctionSecurity = FunctionSecurity.invoker
volatility: FunctionVolatility = FunctionVolatility.VOLATILE
def _schema_from_table_args(cls: type) -> str:
raw = getattr(cls, "__table_args__", None)
if isinstance(raw, dict):
schema = raw.get("schema")
elif isinstance(raw, tuple) and raw and isinstance(raw[-1], dict):
schema = raw[-1].get("schema")
else:
schema = None
if schema is None:
msg = f"{cls.__name__}: must specify a schema in __table_args__."
raise CodegenDatabaseValidationError(msg)
return schema
def _register_function_class(class_: type) -> None:
name = class_.__dict__["__funcname__"]
if "__definition__" not in class_.__dict__:
msg = f"{class_.__name__}: must define __definition__."
raise CodegenDatabaseValidationError(msg)
options = class_.__dict__.get("__options__", FunctionOptions())
definition = class_.__dict__["__definition__"]
resolved_definition = definition() if callable(definition) else definition
function = Function(
name,
resolved_definition,
returns=options.returns,
language=options.language,
schema=_schema_from_table_args(class_),
parameters=options.parameters,
security=options.security,
volatility=options.volatility,
)
register_function(class_.metadata, function)
class_.function = function
class CodegenDatabaseFunction:
metadata: ClassVar[MetaData]
function: Function
def __init__(
self,
name: str,
schema: str,
definition: str | Callable[[], str],
*,
metadata: MetaData | None = None,
returns: str = "void",
language: str = "sql",
parameters: list[FunctionParam] | None = None,
security: FunctionSecurity = FunctionSecurity.invoker,
volatility: FunctionVolatility = FunctionVolatility.VOLATILE,
) -> None:
metadata = metadata or self._find_metadata(type(self))
if metadata is None:
msg = f"{type(self).__name__}: metadata must be provided."
raise CodegenDatabaseValidationError(msg)
resolved_definition = (
definition() if callable(definition) else definition
)
function = Function(
name,
resolved_definition,
returns=returns,
language=language,
schema=schema,
parameters=parameters or [],
security=security,
volatility=volatility,
)
register_function(metadata, function)
self.function = function
self.name = name
self.schema = schema
@staticmethod
def _find_metadata(class_: type) -> MetaData | None:
for klass in class_.__mro__:
if "metadata" in klass.__dict__:
return klass.__dict__["metadata"]
return None