import importlib from dataclasses import is_dataclass from typing import Any, TypeAlias, TypeVar, cast, is_typeddict import pydantic from factory.alchemy import SQLAlchemyModelFactory from polyfactory import BaseFactory from polyfactory.factories.dataclass_factory import DataclassFactory from polyfactory.factories.typed_dict_factory import TypedDictFactory from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import Session from tests.unit.factories import ModelFactory T = TypeVar("T") FactoryType: TypeAlias = type[BaseFactory[Any] | SQLAlchemyModelFactory] class FactoryService: def __init__(self) -> None: self._factories: dict[type, FactoryType] = {} def register(self, factory_type: FactoryType) -> None: if issubclass(factory_type, BaseFactory) and getattr( factory_type, "__model__", None ): self._factories[factory_type.__model__] = factory_type elif ( issubclass(factory_type, SQLAlchemyModelFactory) and hasattr(factory_type, "_meta") and factory_type._meta.model # noqa ): self._factories[factory_type._meta.model] = factory_type # noqa else: raise TypeError(f"Not supported factory type {factory_type}.") def scan(self, module_name: str) -> None: for name, value in importlib.import_module(module_name).__dict__.items(): if not name.endswith("Factory"): continue try: self.register(value) except TypeError: pass def build(self, model: type[T], **kwargs: Any) -> T: factory_cls = self._get_factory(model) if issubclass(factory_cls, BaseFactory): return cast(T, factory_cls.build(**kwargs)) elif issubclass(factory_cls, SQLAlchemyModelFactory): return cast(T, factory_cls.build(**kwargs)) raise ValueError(f"Cannot get factory for {model}.") def build_batch(self, model: type[T], *, size: int, **kwargs: Any) -> list[T]: factory_cls = self._get_factory(model) if issubclass(factory_cls, BaseFactory): return cast(list[T], factory_cls.batch(size, **kwargs)) elif issubclass(factory_cls, SQLAlchemyModelFactory): return cast(list[T], factory_cls.build_batch(size, **kwargs)) # type: ignore[union-attr] raise ValueError(f"Cannot get factory for {model}.") async def create( self, session: AsyncSession, model: type[T], *, persistence: str | None = None, **kwargs: Any, ) -> T: factory_cls = self._get_factory(model) if issubclass(factory_cls, SQLAlchemyModelFactory): return cast( T, await session.run_sync( self._create_sync, factory_cls, persistence=persistence, **kwargs ), ) raise TypeError( f"Not supported factory type {factory_cls} to perform `create` action." ) async def create_batch( self, session: AsyncSession, model: type[T], *, size: int, persistence: str | None = None, **kwargs: Any, ) -> list[T]: factory_cls = self._get_factory(model) if issubclass(factory_cls, SQLAlchemyModelFactory): return cast( list[T], await session.run_sync( self._create_batch_sync, factory_cls, size=size, persistence=persistence, **kwargs, ), ) raise TypeError( f"Not supported factory type {factory_cls} to perform " "`create_batch` action." ) def _get_factory(self, model: type[T]) -> FactoryType: if model not in self._factories: if is_typeddict(model): self._factories[model] = TypedDictFactory.create_factory(model) elif is_dataclass(model): self._factories[model] = DataclassFactory.create_factory(model) elif issubclass(model, pydantic.BaseModel): self._factories[model] = ModelFactory.create_factory(model) try: return self._factories[model] except KeyError as exc: raise ValueError(f"`{model}` has not registered factory.") from exc def _create_sync( self, session: Session, factory_cls: type[SQLAlchemyModelFactory], *, persistence: str | None = None, **kwargs: Any, ) -> Any: self._patch_session(session, persistence=persistence) return factory_cls.create(**kwargs) def _create_batch_sync( self, session: Session, factory_cls: type[SQLAlchemyModelFactory], *, size: int, persistence: str | None = None, **kwargs: Any, ) -> list[T]: self._patch_session(session, persistence=persistence) return cast(list[T], factory_cls.create_batch(size, **kwargs)) def _patch_session( self, session: Session, *, persistence: str | None = None ) -> None: for factory_cls in self._factories.values(): if issubclass(factory_cls, SQLAlchemyModelFactory): factory_cls._meta.sqlalchemy_session = session # type: ignore[union-attr] factory_cls._meta.sqlalchemy_session_persistence = persistence # type: ignore[union-attr]