Source code for codegen_database.functions

"""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