import contextlib from collections.abc import Callable from functools import wraps from typing import ( Concatenate, Literal, ParamSpec, Protocol, TypeVar, overload, runtime_checkable, ) from .base import DefaultDB, ReportingDB @runtime_checkable class HasDefaultDB(Protocol): db: DefaultDB @runtime_checkable class HasReportingDB(Protocol): db: ReportingDB @runtime_checkable class HasDefaultDBAndReportingDB(Protocol): db: DefaultDB reporting_db: ReportingDB DBName = Literal["default", "reporting"] P = ParamSpec("P") R = TypeVar("R") @overload def transactional[R, HasDB: HasDefaultDB | HasReportingDB, **P]( obj: Callable[Concatenate[HasDB, P], R], ) -> Callable[Concatenate[HasDB, P], R]: ... @overload def transactional[R, HasDB: HasDefaultDB | HasReportingDB, **P]( *, ignore: tuple[type[Exception], ...] | None = None, ) -> Callable[ [Callable[Concatenate[HasDB, P], R]], Callable[Concatenate[HasDB, P], R], ]: ... def transactional[R, HasDB: HasDefaultDB | HasReportingDB, **P]( obj: Callable[Concatenate[HasDB, P], R] | None = None, ignore: tuple[type[Exception], ...] | None = None, ) -> ( Callable[ [Callable[Concatenate[HasDB, P], R]], Callable[Concatenate[HasDB, P], R], ] | Callable[Concatenate[HasDB, P], R] ): def decorator( func: Callable[Concatenate[HasDB, P], R], ) -> Callable[Concatenate[HasDB, P], R]: @wraps(func) def wrapper(self: HasDB, /, *args: P.args, **kwargs: P.kwargs) -> R: stack = contextlib.ExitStack() if isinstance(self, HasDefaultDBAndReportingDB): stack.enter_context(self.db.transaction(commit_on_error=ignore)) stack.enter_context( self.reporting_db.transaction(commit_on_error=ignore) ) elif isinstance(self, HasDefaultDB): stack.enter_context(self.db.transaction(commit_on_error=ignore)) elif isinstance(self, HasReportingDB): stack.enter_context(self.db.transaction(commit_on_error=ignore)) with stack: return func(self, *args, **kwargs) return wrapper if obj is None: return decorator return decorator(obj)