import datetime import decimal import enum import importlib from collections.abc import Callable from dataclasses import is_dataclass from typing import Any, cast, is_typeddict import pydantic from polyfactory import BaseFactory from polyfactory.factories.dataclass_factory import DataclassFactory from polyfactory.factories.pydantic_factory import ModelFactory as BaseModelFactory from polyfactory.factories.sqlalchemy_factory import ( SQLAlchemyFactory as BaseSQLAlchemyFactory, ) from polyfactory.factories.typed_dict_factory import TypedDictFactory from sqlalchemy import Column from sqlalchemy.orm import DeclarativeBase, Session from fansifter_common.adapters.db.types import ( ChoiceType, NaiveUTCDateTime, PydanticType, SafeJSONType, ) from fansifter_common.utils import timezone type FactoryType = type[BaseFactory[Any]] # Base factories class ModelFactory[T: pydantic.BaseModel](BaseModelFactory[T]): __is_base_factory__ = True __allow_none_optionals__ = False __check_model__ = False @classmethod def get_provider_map(cls) -> dict[Any, Callable[..., Any]]: provider_map = super().get_provider_map() return { **provider_map, datetime.datetime: timezone.now, } class SQLAlchemyFactory[T](BaseSQLAlchemyFactory[T]): __is_base_factory__ = True __allow_none_optionals__ = False __check_model__ = False __set_foreign_keys__ = False __set_relationships__ = True @classmethod def get_sqlalchemy_types(cls) -> dict[Any, Callable[[], Any]]: """Get mapping of types where column type.""" sqlalchemy_types = super().get_sqlalchemy_types() return { NaiveUTCDateTime: lambda: cls.__faker__.date_time(tzinfo=timezone.UTC), **sqlalchemy_types, } @classmethod def get_type_from_column(cls, column: Column[Any]) -> type: if ( isinstance(column.type, ChoiceType) and isinstance(column.type.choices, type) and issubclass(column.type.choices, enum.Enum) ): return column.type.choices elif isinstance(column.type, PydanticType): return cast(type, column.type.type) elif isinstance(column.type, SafeJSONType): return dict return super().get_type_from_column(column) @classmethod def get_provider_map(cls) -> dict[Any, Callable[..., Any]]: provider_map = super().get_provider_map() return { **provider_map, decimal.Decimal: lambda: cls.__faker__.pydecimal( left_digits=10, right_digits=2, max_value=9999999, positive=True ), } class FactoryService: def __init__(self) -> None: self._factories: dict[type, FactoryType] = {} def register(self, factory_type: FactoryType) -> None: if (model := getattr(factory_type, "__model__", None)) is not None: self._factories[model] = factory_type 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[T](self, model: type[T], **kwargs: Any) -> T: factory_cls = self._get_factory(model) return cast(T, factory_cls.build(**kwargs)) def build_batch[T](self, model: type[T], *, size: int, **kwargs: Any) -> list[T]: factory_cls = self._get_factory(model) return cast(list[T], factory_cls.batch(size, **kwargs)) def create[T]( self, session: Session, model: type[T], **kwargs: Any, ) -> T: factory_cls = self._get_factory(model) if issubclass(factory_cls, SQLAlchemyFactory): instance = cast(T, factory_cls.build(**kwargs)) session.add(instance) session.commit() return instance raise TypeError( f"Not supported factory type {factory_cls} to perform `create` action." ) def create_batch[T]( self, session: Session, model: type[T], *, size: int, **kwargs: Any, ) -> list[T]: factory_cls = self._get_factory(model) if issubclass(factory_cls, SQLAlchemyFactory): instances = cast(list[T], factory_cls.batch(size, **kwargs)) session.add_all(instances) session.commit() return instances raise TypeError( f"Not supported factory type {factory_cls} to perform " "`create_batch` action." ) def _get_factory[T](self, model: type[T]) -> FactoryType: if model not in self._factories: if is_typeddict(model): self._factories[model] = TypedDictFactory.create_factory( model, __check_model__=False ) elif is_dataclass(model) and not issubclass(model, DeclarativeBase): self._factories[model] = DataclassFactory.create_factory( model, __check_model__=False ) elif issubclass(model, pydantic.BaseModel): self._factories[model] = ModelFactory.create_factory( model, __check_model__=False ) elif issubclass(model, DeclarativeBase): self._factories[model] = SQLAlchemyFactory.create_factory( model, __check_model__=False ) try: return self._factories[model] except KeyError as exc: raise ValueError(f"`{model}` has not registered factory.") from exc