Init project
This commit is contained in:
14
birthday_pool_bot/repositories/__init__.py
Normal file
14
birthday_pool_bot/repositories/__init__.py
Normal file
@@ -0,0 +1,14 @@
|
||||
from .container import RepositoriesContainer
|
||||
from .exceptions import AlreadyExistsError, HaveNoSessionError
|
||||
from .settings import RepositorySettings
|
||||
|
||||
|
||||
__all__ = (
|
||||
# container
|
||||
"RepositoriesContainer",
|
||||
# exceptions
|
||||
"AlreadyExistsError",
|
||||
"HaveNoSessionError",
|
||||
# settings
|
||||
"RepositorySettings",
|
||||
)
|
||||
150
birthday_pool_bot/repositories/base.py
Normal file
150
birthday_pool_bot/repositories/base.py
Normal file
@@ -0,0 +1,150 @@
|
||||
import contextvars
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from typing import Generic, TypeVar
|
||||
|
||||
from pydantic_filters.drivers.sqlalchemy import append_to_statement
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from sqlmodel import delete, insert, select, update
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from birthday_pool_bot.interfaces import RepositoryInterface
|
||||
|
||||
|
||||
DTOType = TypeVar("DTOType")
|
||||
FilterType = TypeVar("FilterType")
|
||||
|
||||
|
||||
class BaseRepository(Generic[DTOType, FilterType]):
|
||||
def __init__(self, settings: RepositorySettings):
|
||||
self._settings = settings
|
||||
|
||||
engine_parameters = {
|
||||
"url": str(self._settings.dsn),
|
||||
"pool_recycle": self._settings.pool_recycle,
|
||||
}
|
||||
if not str(self._settings.dsn).startswith("sqlite"):
|
||||
engine_parameters["pool_timeout"] = self._settings.pool_timeout
|
||||
engine_parameters["pool_size"] = self._settings.pool_size
|
||||
|
||||
self._engine = create_async_engine(**engine_parameters)
|
||||
self._session_context = contextvars.ContextVar("session_context")
|
||||
|
||||
@property
|
||||
def session(self) -> AsyncSession:
|
||||
session = self._session_context.get(None)
|
||||
if session is None:
|
||||
raise HaveNoSessionError()
|
||||
|
||||
return session
|
||||
|
||||
@asynccontextmanager
|
||||
async def transaction(self) -> AsyncGenerator[None, None]:
|
||||
session_args = {"bind": self._engine, "expire_on_commit": False}
|
||||
async with AsyncSession(**session_args) as session:
|
||||
with self._set_session_to_session_context(session=session):
|
||||
async with session.begin():
|
||||
yield
|
||||
|
||||
@contextmanager
|
||||
def _set_session_to_session_context(self, session: AsyncSession) -> Generator[None, None, None]:
|
||||
token = self._session_context.set(session)
|
||||
yield
|
||||
self._session_context.reset(token)
|
||||
|
||||
def get_db_table(self) -> type[BaseSQLModel]:
|
||||
raise NotImplementedError
|
||||
|
||||
async def get_items_count(
|
||||
self,
|
||||
filter_: FilterType | None = None,
|
||||
) -> int:
|
||||
table = self.get_db_table()
|
||||
statement = select(func.count())
|
||||
statement = append_to_statement(
|
||||
statement=statement,
|
||||
model=table,
|
||||
filter_=filter_,
|
||||
)
|
||||
|
||||
result = await self.session.exec(statement)
|
||||
items_count = result.first()
|
||||
|
||||
return items_count
|
||||
|
||||
async def get_items(
|
||||
self,
|
||||
filter_: FilterType | None = None,
|
||||
pagination: BasePagination | None = None,
|
||||
sort: BaseSort | None = None,
|
||||
options: list | None = None,
|
||||
) -> AsyncGenerator[DTOType]:
|
||||
table = self.get_db_table()
|
||||
statement = select(table)
|
||||
statement = append_to_statement(
|
||||
statement=statement,
|
||||
model=table,
|
||||
filter_=filter_,
|
||||
pagination=pagination,
|
||||
sort=sort,
|
||||
)
|
||||
if options:
|
||||
statement = statement.options(*options)
|
||||
|
||||
result = await self.session.exec(statement)
|
||||
db_items = result.all()
|
||||
|
||||
for db_item in db_items:
|
||||
yield db_item.to_item()
|
||||
|
||||
async def create_item(self, item: DTOType) -> DTOType:
|
||||
table = self.get_db_table()
|
||||
values = table.from_item(item=item).to_values()
|
||||
statement = insert(table).values(values).returning(table)
|
||||
|
||||
try:
|
||||
result = await self.session.exec(statement)
|
||||
except IntegrityError:
|
||||
raise AlreadyExistsError(model=table, values=values)
|
||||
db_item = result.first()[0]
|
||||
|
||||
return db_item.to_item()
|
||||
|
||||
async def update_items(
|
||||
self,
|
||||
filter_: FilterType | None = None,
|
||||
**values,
|
||||
) -> AsyncGenerator[DTOType]:
|
||||
table = self.get_db_table()
|
||||
statement = update(table)
|
||||
statement = append_to_statement(
|
||||
statement=statement,
|
||||
model=table,
|
||||
filter_=filter_,
|
||||
)
|
||||
statement = statement.values(**values).returning(table)
|
||||
|
||||
result = await self.session.exec(statement)
|
||||
db_items = result.all()
|
||||
|
||||
for db_item, *_ in db_items:
|
||||
yield db_item.to_item()
|
||||
|
||||
async def delete_items(
|
||||
self,
|
||||
filter_: FilterType | None = None,
|
||||
pagination: BasePagination | None = None,
|
||||
sort: BaseSort | None = None,
|
||||
):
|
||||
table = self.get_db_table()
|
||||
statement = delete(table)
|
||||
statement = append_to_statement(
|
||||
statement=statement,
|
||||
model=table,
|
||||
filter_=filter_,
|
||||
pagination=pagination,
|
||||
sort=sort,
|
||||
)
|
||||
|
||||
await self.session.exec(statement)
|
||||
25
birthday_pool_bot/repositories/container.py
Normal file
25
birthday_pool_bot/repositories/container.py
Normal file
@@ -0,0 +1,25 @@
|
||||
from .repositories import (
|
||||
PoolsRepository,
|
||||
SubscriptionsRepository,
|
||||
UsersRepository,
|
||||
)
|
||||
from .settings import RepositorySettings
|
||||
|
||||
|
||||
class RepositoriesContainer:
|
||||
def __init__(self, settings: RepositorySettings):
|
||||
self._pools_repository = PoolsRepository(settings=settings)
|
||||
self._subscriptions_repository = SubscriptionsRepository(settings=settings)
|
||||
self._users_repository = UsersRepository(settings=settings)
|
||||
|
||||
@property
|
||||
def pools(self) -> PoolsRepository:
|
||||
return self._pools_repository
|
||||
|
||||
@property
|
||||
def subscriptions(self) -> SubscriptionsRepository:
|
||||
return self._subscriptions_repository
|
||||
|
||||
@property
|
||||
def users(self) -> UsersRepository:
|
||||
return self._users_repository
|
||||
20
birthday_pool_bot/repositories/exceptions.py
Normal file
20
birthday_pool_bot/repositories/exceptions.py
Normal file
@@ -0,0 +1,20 @@
|
||||
from typing import Any
|
||||
|
||||
from sqlmodel import SQLModel
|
||||
|
||||
|
||||
class RepositoryException(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class HaveNoSessionError(RepositoryException):
|
||||
def __init__(self):
|
||||
super().__init__("Have no actual session")
|
||||
|
||||
|
||||
class AlreadyExistsError(RepositoryException):
|
||||
def __init__(self, model: SQLModel, values: dict[str, Any]):
|
||||
super().__init__(f"Record for model '{model.__name__}' already exists")
|
||||
|
||||
self.model = model
|
||||
self.values = values
|
||||
10
birthday_pool_bot/repositories/migrations/__init__.py
Normal file
10
birthday_pool_bot/repositories/migrations/__init__.py
Normal file
@@ -0,0 +1,10 @@
|
||||
from .cli import get_cli
|
||||
from .migrator import Migrator
|
||||
|
||||
|
||||
__all__ = (
|
||||
# cli
|
||||
"get_cli",
|
||||
# migrator
|
||||
"Migrator",
|
||||
)
|
||||
75
birthday_pool_bot/repositories/migrations/cli.py
Normal file
75
birthday_pool_bot/repositories/migrations/cli.py
Normal file
@@ -0,0 +1,75 @@
|
||||
from typing import Optional
|
||||
|
||||
import typer
|
||||
from rich.console import Console
|
||||
from rich.table import Table
|
||||
|
||||
from birthday_pool_bot.settings import Settings
|
||||
|
||||
from .migrator import Migrator
|
||||
|
||||
|
||||
def callback(ctx: typer.Context):
|
||||
ctx.obj = ctx.obj or {}
|
||||
|
||||
settings = ctx.obj["settings"]
|
||||
ctx.obj["migrator"] = Migrator(settings=settings.repository)
|
||||
|
||||
|
||||
def show_migrations_list(ctx: typer.Context):
|
||||
migrator: Migrator = ctx.obj["migrator"]
|
||||
|
||||
table = Table("ID")
|
||||
for migration_id in migrator.get_migrations():
|
||||
table.add_row(migration_id)
|
||||
Console().print(table)
|
||||
|
||||
|
||||
def apply_migrations(ctx: typer.Context):
|
||||
migrator: Migrator = ctx.obj["migrator"]
|
||||
|
||||
migrator.migrate()
|
||||
|
||||
|
||||
def rollback_migrations(
|
||||
ctx: typer.Context,
|
||||
revision: Optional[str] = typer.Argument(
|
||||
None,
|
||||
help="Revision id or relative revision (`-1`, `-2`)",
|
||||
),
|
||||
):
|
||||
migrator: Migrator = ctx.obj["migrator"]
|
||||
|
||||
migrator.rollback(migration_id=revision)
|
||||
|
||||
|
||||
def squash_migrations(ctx: typer.Context):
|
||||
migrator: Migrator = ctx.obj["migrator"]
|
||||
|
||||
migrator.squash_migrations(new_message="init")
|
||||
|
||||
|
||||
def create_migration(
|
||||
ctx: typer.Context,
|
||||
message: Optional[str] = typer.Option(
|
||||
None,
|
||||
"-m", "--message",
|
||||
help="Migration short message",
|
||||
),
|
||||
):
|
||||
migrator: Migrator = ctx.obj["migrator"]
|
||||
|
||||
migrator.create_migration(message=message)
|
||||
|
||||
|
||||
def get_cli() -> typer.Typer:
|
||||
cli = typer.Typer(name="Migration")
|
||||
|
||||
cli.callback()(callback)
|
||||
cli.command(name="apply")(apply_migrations)
|
||||
cli.command(name="rollback")(rollback_migrations)
|
||||
cli.command(name="create")(create_migration)
|
||||
cli.command(name="squash")(squash_migrations)
|
||||
cli.command(name="list")(show_migrations_list)
|
||||
|
||||
return cli
|
||||
93
birthday_pool_bot/repositories/migrations/env.py
Normal file
93
birthday_pool_bot/repositories/migrations/env.py
Normal file
@@ -0,0 +1,93 @@
|
||||
import asyncio
|
||||
from logging.config import fileConfig
|
||||
|
||||
from alembic import context
|
||||
from sqlalchemy import pool
|
||||
from sqlalchemy.engine import Connection
|
||||
from sqlalchemy.ext.asyncio import async_engine_from_config
|
||||
|
||||
from birthday_pool_bot.repositories.tables import User
|
||||
|
||||
# this is the Alembic Config object, which provides
|
||||
# access to the values within the .ini file in use.
|
||||
config = context.config
|
||||
|
||||
# Interpret the config file for Python logging.
|
||||
# This line sets up loggers basically.
|
||||
if config.config_file_name is not None:
|
||||
fileConfig(config.config_file_name)
|
||||
|
||||
# add your model's MetaData object here
|
||||
# for 'autogenerate' support
|
||||
# from myapp import mymodel
|
||||
# target_metadata = mymodel.Base.metadata
|
||||
target_metadata = User.metadata
|
||||
|
||||
# other values from the config, defined by the needs of env.py,
|
||||
# can be acquired:
|
||||
# my_important_option = config.get_main_option("my_important_option")
|
||||
# ... etc.
|
||||
|
||||
|
||||
def run_migrations_offline() -> None:
|
||||
"""Run migrations in 'offline' mode.
|
||||
|
||||
This configures the context with just a URL
|
||||
and not an Engine, though an Engine is acceptable
|
||||
here as well. By skipping the Engine creation
|
||||
we don't even need a DBAPI to be available.
|
||||
|
||||
Calls to context.execute() here emit the given string to the
|
||||
script output.
|
||||
|
||||
"""
|
||||
url = config.get_main_option("sqlalchemy.url")
|
||||
context.configure(
|
||||
url=url,
|
||||
target_metadata=target_metadata,
|
||||
literal_binds=True,
|
||||
dialect_opts={"paramstyle": "named"},
|
||||
)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
def do_run_migrations(connection: Connection) -> None:
|
||||
context.configure(
|
||||
connection=connection,
|
||||
target_metadata=target_metadata,
|
||||
)
|
||||
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
async def run_async_migrations() -> None:
|
||||
"""In this scenario we need to create an Engine
|
||||
and associate a connection with the context.
|
||||
|
||||
"""
|
||||
|
||||
connectable = async_engine_from_config(
|
||||
config.get_section(config.config_ini_section, {}),
|
||||
prefix="sqlalchemy.",
|
||||
poolclass=pool.NullPool,
|
||||
)
|
||||
|
||||
async with connectable.connect() as connection:
|
||||
await connection.run_sync(do_run_migrations)
|
||||
|
||||
await connectable.dispose()
|
||||
|
||||
|
||||
def run_migrations_online() -> None:
|
||||
"""Run migrations in 'online' mode."""
|
||||
|
||||
asyncio.run(run_async_migrations())
|
||||
|
||||
|
||||
if context.is_offline_mode():
|
||||
run_migrations_offline()
|
||||
else:
|
||||
run_migrations_online()
|
||||
146
birthday_pool_bot/repositories/migrations/migrator.py
Normal file
146
birthday_pool_bot/repositories/migrations/migrator.py
Normal file
@@ -0,0 +1,146 @@
|
||||
import importlib.util
|
||||
import pathlib
|
||||
from types import ModuleType
|
||||
from typing import Generator, Iterable
|
||||
|
||||
from alembic import command as alembic_command
|
||||
from alembic.config import Config as AlembicConfig
|
||||
|
||||
from birthday_pool_bot.repositories import RepositorySettings
|
||||
|
||||
|
||||
class Migrator:
|
||||
MIGRATION_FILENAME_TEMPLATE = (
|
||||
"%%(year)d_%%(month)02d_%%(day)02d_%%(hour)02dh%%(minute)02dm%%(second)02ds_%%(slug)s"
|
||||
)
|
||||
|
||||
def __init__(self, settings: RepositorySettings):
|
||||
self._settings = settings
|
||||
|
||||
self._alembic_config = self._generate_alembic_config()
|
||||
|
||||
def _generate_alembic_config(self) -> AlembicConfig:
|
||||
migrations_path = self.get_migrations_directory_path()
|
||||
|
||||
config = AlembicConfig()
|
||||
config.set_main_option("script_location", str(migrations_path))
|
||||
config.set_main_option("file_template", self.MIGRATION_FILENAME_TEMPLATE)
|
||||
config.set_main_option("sqlalchemy.url", str(self._settings.dsn).replace("%", "%%"))
|
||||
|
||||
return config
|
||||
|
||||
def squash_migrations(
|
||||
self,
|
||||
new_migration_id: str | None = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
new_message = kwargs.get("new_message", None)
|
||||
|
||||
if new_migration_id is None:
|
||||
last_migration_id = self.get_latest_migration()
|
||||
new_migration_id = last_migration_id
|
||||
|
||||
self.rollback(migration_id="base")
|
||||
self._delete_all_migrations()
|
||||
new_migration_id = self.create_migration(
|
||||
migration_id=new_migration_id,
|
||||
message=new_message,
|
||||
)
|
||||
return new_migration_id
|
||||
|
||||
def migrate(self, migration_id: str | None = None, **kwargs) -> str | None:
|
||||
alembic_revision = migration_id or "head"
|
||||
alembic_command.upgrade(config=self._alembic_config, revision=alembic_revision)
|
||||
|
||||
return self.get_current_migration()
|
||||
|
||||
def rollback(self, migration_id: str | None = None, **kwargs) -> str | None:
|
||||
alembic_revision = migration_id or "-1"
|
||||
alembic_command.downgrade(config=self._alembic_config, revision=alembic_revision)
|
||||
|
||||
return self.get_current_migration()
|
||||
|
||||
def get_current_migration(self) -> str | None:
|
||||
return None
|
||||
|
||||
def create_migration(
|
||||
self,
|
||||
migration_id: str | None = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
message = kwargs.get("message", None)
|
||||
|
||||
alembic_command.revision(
|
||||
self._alembic_config,
|
||||
message=message,
|
||||
rev_id=migration_id,
|
||||
autogenerate=True,
|
||||
)
|
||||
|
||||
last_migration_id = self.get_latest_migration()
|
||||
if last_migration_id is None:
|
||||
raise RuntimeError("Migration not created")
|
||||
|
||||
return last_migration_id
|
||||
|
||||
def get_latest_migration(self) -> str | None:
|
||||
last_migration_id = None
|
||||
for migration_id in self.get_migrations():
|
||||
last_migration_id = migration_id
|
||||
return last_migration_id
|
||||
|
||||
def get_migrations(self) -> Generator[str, None, None]:
|
||||
migrations_modules = self._get_migrations_modules()
|
||||
sorted_migrations_modules = self._sort_migrations_modules(modules=migrations_modules)
|
||||
for migration_module in sorted_migrations_modules:
|
||||
yield migration_module.revision
|
||||
|
||||
def _get_migrations_modules(self) -> Iterable[ModuleType]:
|
||||
migrations_path = self.get_migrations_directory_path()
|
||||
versions_path = migrations_path / "versions"
|
||||
|
||||
for index, migration_module_path in enumerate(versions_path.glob("*.py")):
|
||||
migration_module_name = f"migration_{index:08}"
|
||||
migration_module_spec = importlib.util.spec_from_file_location(
|
||||
name=migration_module_name,
|
||||
location=migration_module_path,
|
||||
)
|
||||
if migration_module_spec is None or migration_module_spec.loader is None:
|
||||
raise ValueError(f"Cannot load migration module '{migration_module_path}'")
|
||||
migration_module = importlib.util.module_from_spec(spec=migration_module_spec)
|
||||
migration_module_spec.loader.exec_module(migration_module)
|
||||
yield migration_module
|
||||
|
||||
def _sort_migrations_modules(self, modules: Iterable[ModuleType]) -> list[ModuleType]:
|
||||
next_module_map = {}
|
||||
first_module = None
|
||||
for module in modules:
|
||||
if module.down_revision is None:
|
||||
if first_module is not None:
|
||||
raise ValueError((
|
||||
"Found multiple first migrations: "
|
||||
f"'{module.revision}' and '{first_module.revision}'"
|
||||
))
|
||||
first_module = module
|
||||
else:
|
||||
next_module_map[module.down_revision] = module
|
||||
if first_module is None:
|
||||
raise ValueError("Doesn't found first migration")
|
||||
|
||||
current_module = first_module
|
||||
sorted_modules = [current_module]
|
||||
while next_module_map:
|
||||
current_module = next_module_map.pop(current_module.revision)
|
||||
sorted_modules.append(current_module)
|
||||
|
||||
return sorted_modules
|
||||
|
||||
def _delete_all_migrations(self):
|
||||
migrations_path = self.get_migrations_directory_path()
|
||||
versions_path = migrations_path / "versions"
|
||||
for migration_file_path in versions_path.glob("*.py"):
|
||||
migration_file_path.unlink()
|
||||
|
||||
def get_migrations_directory_path(self) -> pathlib.Path:
|
||||
directory_path = pathlib.Path(__file__).parent
|
||||
return directory_path
|
||||
28
birthday_pool_bot/repositories/migrations/script.py.mako
Normal file
28
birthday_pool_bot/repositories/migrations/script.py.mako
Normal file
@@ -0,0 +1,28 @@
|
||||
"""${message}
|
||||
|
||||
Revision ID: ${up_revision}
|
||||
Revises: ${down_revision | comma,n}
|
||||
Create Date: ${create_date}
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
${imports if imports else ""}
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = ${repr(up_revision)}
|
||||
down_revision: Union[str, Sequence[str], None] = ${repr(down_revision)}
|
||||
branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)}
|
||||
depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)}
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Upgrade schema."""
|
||||
${upgrades if upgrades else "pass"}
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Downgrade schema."""
|
||||
${downgrades if downgrades else "pass"}
|
||||
@@ -0,0 +1,99 @@
|
||||
"""init
|
||||
|
||||
Revision ID: c1060d90df61
|
||||
Revises:
|
||||
Create Date: 2026-01-13 10:06:43.435312
|
||||
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
import sqlalchemy as sa
|
||||
import sqlmodel
|
||||
from alembic import op
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "c1060d90df61"
|
||||
down_revision: Union[str, Sequence[str], None] = None
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Upgrade schema."""
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.create_table("users",
|
||||
sa.Column("id", sa.Uuid(), nullable=False),
|
||||
sa.Column("name", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||
sa.Column("birthday", sa.Date(), nullable=True),
|
||||
sa.Column("phone", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||
sa.Column("telegram_id", sa.Integer(), nullable=True),
|
||||
sa.Column("gift_payment_data", sa.JSON(), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("id"),
|
||||
sa.UniqueConstraint("phone", name="uq__users__phone"),
|
||||
sa.UniqueConstraint("telegram_id", name="uq__users__telegram_id",)
|
||||
)
|
||||
op.create_index("ix__users__phone", "users", ["phone"], unique=False)
|
||||
op.create_index("ix__users__telegram_id", "users", ["telegram_id"], unique=False)
|
||||
op.create_table("pools",
|
||||
sa.Column("id", sa.Uuid(), nullable=False),
|
||||
sa.Column("owner_id", sa.Uuid(), nullable=False),
|
||||
sa.Column("birthday_user_id", sa.Uuid(), nullable=False),
|
||||
sa.Column("description", sqlmodel.sql.sqltypes.AutoString(), nullable=True),
|
||||
sa.Column("payment_data", sa.JSON(), nullable=False),
|
||||
sa.CheckConstraint(
|
||||
"owner_id <> birthday_user_id",
|
||||
name="ck__pools__owner_not_birthday_user",
|
||||
),
|
||||
sa.ForeignKeyConstraint(["birthday_user_id"], ["users.id"]),
|
||||
sa.ForeignKeyConstraint(["owner_id"], ["users.id"]),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("id"),
|
||||
)
|
||||
op.create_index("ix__pools__birthday_user_id", "pools", ["birthday_user_id"], unique=False)
|
||||
op.create_index("ix__pools__owner_id", "pools", ["owner_id"], unique=False)
|
||||
op.create_table("subscriptions",
|
||||
sa.Column("from_user_id", sa.Uuid(), nullable=False),
|
||||
sa.Column("to_user_id", sa.Uuid(), nullable=False),
|
||||
sa.Column("name", sqlmodel.sql.sqltypes.AutoString(), nullable=False),
|
||||
sa.Column("pool_id", sa.Uuid(), nullable=True),
|
||||
sa.CheckConstraint(
|
||||
"from_user_id <> to_user_id",
|
||||
name="ck__subscriptions__from_user_not_to_user",
|
||||
),
|
||||
sa.ForeignKeyConstraint(["from_user_id"], ["users.id"]),
|
||||
sa.ForeignKeyConstraint(["pool_id"], ["pools.id"], ondelete="SET NULL"),
|
||||
sa.ForeignKeyConstraint(["to_user_id"], ["users.id"]),
|
||||
sa.PrimaryKeyConstraint("from_user_id", "to_user_id"),
|
||||
sa.UniqueConstraint("from_user_id", "to_user_id", name="uq__subscriptions__from_to_user"),
|
||||
)
|
||||
op.create_index("ix__subscriptions__from_user_id", "subscriptions", ["from_user_id"], unique=False)
|
||||
op.create_index(
|
||||
"ix__subscriptions__from_user_id__to_user_id",
|
||||
"subscriptions",
|
||||
["from_user_id", "to_user_id"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index("ix__subscriptions__name", "subscriptions", ["name"], unique=False)
|
||||
op.create_index("ix__subscriptions__pool_id", "subscriptions", ["pool_id"], unique=False)
|
||||
op.create_index("ix__subscriptions__to_user_id", "subscriptions", ["to_user_id"], unique=False)
|
||||
# ### end Alembic commands ###
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Downgrade schema."""
|
||||
# ### commands auto generated by Alembic - please adjust! ###
|
||||
op.drop_index("ix__subscriptions__to_user_id", table_name="subscriptions")
|
||||
op.drop_index("ix__subscriptions__pool_id", table_name="subscriptions")
|
||||
op.drop_index("ix__subscriptions__name", table_name="subscriptions")
|
||||
op.drop_index("ix__subscriptions__from_user_id__to_user_id", table_name="subscriptions")
|
||||
op.drop_index("ix__subscriptions__from_user_id", table_name="subscriptions")
|
||||
op.drop_table("subscriptions")
|
||||
op.drop_index("ix__pools__owner_id", table_name="pools")
|
||||
op.drop_index("ix__pools__birthday_user_id", table_name="pools")
|
||||
op.drop_table("pools")
|
||||
op.drop_index("ix__users__telegram_id", table_name="users")
|
||||
op.drop_index("ix__users__phone", table_name="users")
|
||||
op.drop_table("users")
|
||||
# ### end Alembic commands ###
|
||||
13
birthday_pool_bot/repositories/repositories/__init__.py
Normal file
13
birthday_pool_bot/repositories/repositories/__init__.py
Normal file
@@ -0,0 +1,13 @@
|
||||
from .pools import PoolsRepository
|
||||
from .subscriptions import SubscriptionsRepository
|
||||
from .users import UsersRepository
|
||||
|
||||
|
||||
__all__ = (
|
||||
# pools
|
||||
"PoolsRepository",
|
||||
# subscriptions
|
||||
"SubscriptionsRepository",
|
||||
# users
|
||||
"UsersRepository",
|
||||
)
|
||||
41
birthday_pool_bot/repositories/repositories/pools.py
Normal file
41
birthday_pool_bot/repositories/repositories/pools.py
Normal file
@@ -0,0 +1,41 @@
|
||||
from pydantic_filters import OffsetPagination
|
||||
from sqlalchemy.orm import joinedload
|
||||
|
||||
from birthday_pool_bot.dto import Pool, PoolFilter
|
||||
from birthday_pool_bot.repositories.base import BaseRepository
|
||||
from birthday_pool_bot.repositories.tables import BaseSQLModel, Pool as DBPool
|
||||
|
||||
|
||||
class PoolsRepository(BaseRepository[Pool, PoolFilter]):
|
||||
def get_db_table(self) -> type[BaseSQLModel]:
|
||||
return DBPool
|
||||
|
||||
async def get_pools_by_birthday_user_id_count(self, birthday_user_id: uuid.UUID) -> int:
|
||||
filter_ = PoolFilter(birthday_user_id={birthday_user_id})
|
||||
return await self.get_items_count(filter_=filter_)
|
||||
|
||||
async def get_pools_by_birthday_user_id(self, birthday_user_id: uuid.UUID) -> list[Pool]:
|
||||
filter_ = PoolFilter(birthday_user_id={birthday_user_id})
|
||||
return [pool async for pool in self.get_items(filter_=filter_)]
|
||||
|
||||
async def get_pool_by_id(
|
||||
self,
|
||||
pool_id: uuid.UUID,
|
||||
with_owner: bool = False,
|
||||
) -> Pool | None:
|
||||
filter_ = PoolFilter(id={pool_id})
|
||||
pagination = OffsetPagination(limit=1)
|
||||
options = []
|
||||
if with_owner:
|
||||
options.append(joinedload(DBPool.owner))
|
||||
|
||||
pools_generator = self.get_items(
|
||||
filter_=filter_,
|
||||
pagination=pagination,
|
||||
options=options,
|
||||
)
|
||||
async for pool in pools_generator:
|
||||
return pool
|
||||
|
||||
async def create_pool(self, pool: Pool) -> Pool:
|
||||
return await self.create_item(item=pool)
|
||||
135
birthday_pool_bot/repositories/repositories/subscriptions.py
Normal file
135
birthday_pool_bot/repositories/repositories/subscriptions.py
Normal file
@@ -0,0 +1,135 @@
|
||||
import uuid
|
||||
|
||||
from pydantic_filters import BasePagination, OffsetPagination
|
||||
from sqlalchemy.orm import joinedload
|
||||
from sqlalchemy import delete, update, select
|
||||
|
||||
from birthday_pool_bot.dto import Subscription, SubscriptionFilter
|
||||
from birthday_pool_bot.repositories.base import BaseRepository
|
||||
from birthday_pool_bot.repositories.tables import (
|
||||
BaseSQLModel,
|
||||
Pool as DBPool,
|
||||
Subscription as DBSubscription,
|
||||
)
|
||||
|
||||
|
||||
class SubscriptionsRepository(BaseRepository[Subscription, SubscriptionFilter]):
|
||||
def get_db_table(self) -> type[BaseSQLModel]:
|
||||
return DBSubscription
|
||||
|
||||
async def get_user_subscriptions_count(
|
||||
self,
|
||||
user_id: uuid.UUID,
|
||||
) -> int:
|
||||
filter_ = SubscriptionFilter(from_user_id={user_id})
|
||||
return await self.get_items_count(filter_=filter_)
|
||||
|
||||
def get_user_subscriptions(
|
||||
self,
|
||||
user_id: uuid.UUID,
|
||||
pagination: BasePagination | None = None,
|
||||
with_from_user: bool = False,
|
||||
with_to_user: bool = False,
|
||||
with_pool: bool = False,
|
||||
with_pool_owner: bool = False,
|
||||
) -> AsyncGenerator[Subscription, None]:
|
||||
filter_ = SubscriptionFilter(from_user_id={user_id})
|
||||
options = []
|
||||
if with_from_user:
|
||||
options.append(joinedload(DBSubscription.from_user))
|
||||
if with_to_user:
|
||||
options.append(joinedload(DBSubscription.to_user))
|
||||
if with_pool:
|
||||
options.append(joinedload(DBSubscription.pool))
|
||||
if with_pool_owner:
|
||||
options.append(joinedload(DBSubscription.pool).joinedload(DBPool.owner))
|
||||
|
||||
return self.get_items(
|
||||
filter_=filter_,
|
||||
pagination=pagination,
|
||||
options=options,
|
||||
)
|
||||
|
||||
async def get_to_users_subscriptions(
|
||||
self,
|
||||
to_users_ids: Iterable[uuid.UUID],
|
||||
with_from_user: bool = False,
|
||||
with_to_user: bool = False,
|
||||
with_pool: bool = False,
|
||||
with_pool_owner: bool = False,
|
||||
) -> list[Subscription]:
|
||||
filter_ = SubscriptionFilter(to_user_id=set(to_users_ids))
|
||||
options = []
|
||||
if with_from_user:
|
||||
options.append(joinedload(DBSubscription.from_user))
|
||||
if with_to_user:
|
||||
options.append(joinedload(DBSubscription.to_user))
|
||||
if with_pool:
|
||||
options.append(joinedload(DBSubscription.pool))
|
||||
if with_pool_owner:
|
||||
options.append(joinedload(DBSubscription.pool).joinedload(DBPool.owner))
|
||||
|
||||
subscriptions_generator = self.get_items(
|
||||
filter_=filter_,
|
||||
options=options,
|
||||
)
|
||||
|
||||
return [subscription async for subscription in subscriptions_generator]
|
||||
|
||||
async def get_subscription(
|
||||
self,
|
||||
from_user_id: uuid.UUID,
|
||||
to_user_id: uuid.UUID,
|
||||
with_from_user: bool = False,
|
||||
with_to_user: bool = False,
|
||||
with_pool: bool = False,
|
||||
with_pool_owner: bool = False,
|
||||
) -> Subscription | None:
|
||||
filter_ = SubscriptionFilter(from_user_id={from_user_id}, to_user_id={to_user_id})
|
||||
pagination = OffsetPagination(limit=1)
|
||||
options = []
|
||||
if with_from_user:
|
||||
options.append(joinedload(DBSubscription.from_user))
|
||||
if with_to_user:
|
||||
options.append(joinedload(DBSubscription.to_user))
|
||||
if with_pool:
|
||||
options.append(joinedload(DBSubscription.pool))
|
||||
if with_pool_owner:
|
||||
options.append(joinedload(DBSubscription.pool).joinedload(DBPool.owner))
|
||||
|
||||
subscriptions_generator = self.get_items(
|
||||
filter_=filter_,
|
||||
pagination=pagination,
|
||||
options=options,
|
||||
)
|
||||
async for subscription in subscriptions_generator:
|
||||
return subscription
|
||||
|
||||
async def create_subscription(self, subscription: Subscription) -> Subscription:
|
||||
return await self.create_item(subscription)
|
||||
|
||||
async def delete_subscription(self, from_user_id: uuid.UUID, to_user_id: uuid.UUID):
|
||||
statement = select(DBPool.id).where(
|
||||
DBPool.owner_id == from_user_id,
|
||||
DBPool.birthday_user_id == to_user_id,
|
||||
)
|
||||
result = await self.session.exec(statement)
|
||||
pool_id = result.one_or_none()
|
||||
if pool_id is not None:
|
||||
pool_id = pool_id[0]
|
||||
statement = delete(DBPool).where(DBPool.id == pool_id)
|
||||
await self.session.exec(statement)
|
||||
|
||||
if pool_id is not None:
|
||||
statement = (
|
||||
update(DBSubscription)
|
||||
.where(DBSubscription.pool_id == pool_id)
|
||||
.values(pool_id=None)
|
||||
)
|
||||
await self.session.exec(statement)
|
||||
|
||||
filter_ = SubscriptionFilter(
|
||||
from_user_id={from_user_id},
|
||||
to_user_id={to_user_id},
|
||||
)
|
||||
await self.delete_items(filter_=filter_)
|
||||
154
birthday_pool_bot/repositories/repositories/users.py
Normal file
154
birthday_pool_bot/repositories/repositories/users.py
Normal file
@@ -0,0 +1,154 @@
|
||||
from datetime import date
|
||||
from typing import Container, Sequence, TypeVar
|
||||
|
||||
from pydantic_filters import OffsetPagination
|
||||
from sqlalchemy import Column, func, inspect, select
|
||||
from sqlmodel import SQLModel
|
||||
|
||||
from birthday_pool_bot.dto import User, UserFilter
|
||||
from birthday_pool_bot.repositories.base import BaseRepository
|
||||
from birthday_pool_bot.repositories.tables import (
|
||||
BaseSQLModel,
|
||||
Pool as DBPool,
|
||||
Subscription as DBSubscription,
|
||||
User as DBUser,
|
||||
)
|
||||
|
||||
|
||||
class UsersRepository(BaseRepository[User, UserFilter]):
|
||||
def get_db_table(self) -> type[BaseSQLModel]:
|
||||
return DBUser
|
||||
|
||||
async def get_user_by_id(self, user_id: uuid.UUID) -> User | None:
|
||||
filter_ = UserFilter(id={user_id})
|
||||
pagination = OffsetPagination(limit=1)
|
||||
async for user in self.get_items(filter_=filter_, pagination=pagination):
|
||||
return user
|
||||
|
||||
async def get_user_by_telegram_id(self, telegram_id: int) -> User | None:
|
||||
filter_ = UserFilter(telegram_id={telegram_id})
|
||||
pagination = OffsetPagination(limit=1)
|
||||
async for user in self.get_items(filter_=filter_, pagination=pagination):
|
||||
return user
|
||||
|
||||
async def get_user_by_phone(self, phone: str) -> User | None:
|
||||
filter_ = UserFilter(phone={phone})
|
||||
pagination = OffsetPagination(limit=1)
|
||||
async for user in self.get_items(filter_=filter_, pagination=pagination):
|
||||
return user
|
||||
|
||||
async def get_users_by_ids(
|
||||
self,
|
||||
user_ids: Container[uuid.UUID],
|
||||
pagination: BasePagination | None = None,
|
||||
) -> list[User]:
|
||||
filter_ = UserFilter(id=set(user_ids))
|
||||
users_generator = self.get_items(filter_=filter_, pagination=pagination)
|
||||
users = [user async for user in users_generator]
|
||||
return users
|
||||
|
||||
async def get_users_by_primary_keys(
|
||||
self,
|
||||
user_id: uuid.UUID | None = None,
|
||||
telegram_id: int | None = None,
|
||||
phone: str | None = None,
|
||||
) -> list[User]:
|
||||
filters = []
|
||||
if user_id is not None:
|
||||
filters.append(DBUser.id == user_id)
|
||||
if telegram_id is not None:
|
||||
filters.append(DBUser.telegram_id == telegram_id)
|
||||
if phone is not None:
|
||||
filters.append(DBUser.phone == phone)
|
||||
if not filters:
|
||||
return []
|
||||
|
||||
statement = select(DBUser).where(or_(filters))
|
||||
result = await self.session.exec(statement)
|
||||
|
||||
return [db_user.to_item() for db_user in result.all()]
|
||||
|
||||
async def get_users_by_birthdays(
|
||||
self,
|
||||
birthday: date | None = None,
|
||||
) -> list[User]:
|
||||
statement = select(DBUser).where(
|
||||
func.extract("month", DBUser.birthday) == birthday.month,
|
||||
func.extract("day", DBUser.birthday) == birthday.day,
|
||||
)
|
||||
result = await self.session.execute(statement)
|
||||
|
||||
return [db_user.to_item() for (db_user,) in result.all()]
|
||||
|
||||
async def create_user(self, user: User) -> User:
|
||||
return await self.create_item(item=user)
|
||||
|
||||
async def update_user(self, user: User) -> User:
|
||||
db_user = self.get_db_table().from_item(item=user)
|
||||
filter_ = UserFilter(id={user.id})
|
||||
users_generator = self.update_items(filter_=filter_, **db_user.to_values())
|
||||
async for user in users_generator:
|
||||
return user
|
||||
|
||||
async def merge_users(self, *users: User) -> User | None:
|
||||
if not users:
|
||||
return None
|
||||
|
||||
merged_user, *users = users
|
||||
merged_user = merge(model=User, first_item=merged_user, items=users)
|
||||
|
||||
if users:
|
||||
users_ids = {user.id for user in users}
|
||||
for model in (DBPool, DBSubscription):
|
||||
foreign_key_columns = get_model_foreign_keys_columns(
|
||||
model=model,
|
||||
table_name="users",
|
||||
column_name="id",
|
||||
)
|
||||
for foreign_key_column in foreign_key_columns:
|
||||
statement = (
|
||||
update(model).values(**{foreign_key_column.key: merged_user.id})
|
||||
.where(foreign_key_column.in_(users_ids))
|
||||
)
|
||||
await self.session.exec(statement)
|
||||
|
||||
filter_ = UserFilter(id={user.id for user in users})
|
||||
await self.delete_items(filter_=filter_)
|
||||
|
||||
merged_user = await self.update_user(user=merged_user)
|
||||
return merged_user
|
||||
|
||||
|
||||
ModelType = TypeVar("ModelType")
|
||||
|
||||
|
||||
def merge(
|
||||
model: type[ModelType],
|
||||
first_item: ModelType,
|
||||
items: Sequence[ModelType],
|
||||
) -> ModelType:
|
||||
data = first_item.model_dump()
|
||||
for item in items:
|
||||
item_data = item.model_dump()
|
||||
for field_name, value in data.items():
|
||||
item_value = item_data[field_name]
|
||||
data[field_name] = value or item_value
|
||||
|
||||
return model(**data)
|
||||
|
||||
|
||||
def get_model_foreign_keys_columns(
|
||||
model: type[SQLModel],
|
||||
table_name: str | None = None,
|
||||
column_name: str | None = None,
|
||||
) -> list[Column]:
|
||||
mapper = inspect(model)
|
||||
columns = []
|
||||
|
||||
for column in mapper.columns:
|
||||
for foreign_key in column.foreign_keys:
|
||||
table_name = table_name or foreign_key.column.table.name
|
||||
column_name = column_name or foreign_key.column.name
|
||||
if (foreign_key.column.table.name, foreign_key.column.name) == (table_name, column_name):
|
||||
columns.append(column)
|
||||
return columns
|
||||
8
birthday_pool_bot/repositories/settings.py
Normal file
8
birthday_pool_bot/repositories/settings.py
Normal file
@@ -0,0 +1,8 @@
|
||||
from pydantic import AnyUrl, BaseModel, PositiveInt
|
||||
|
||||
|
||||
class RepositorySettings(BaseModel):
|
||||
dsn: AnyUrl = "sqlite+aiosqlite:///db.sqlite3"
|
||||
pool_size: PositiveInt = 5
|
||||
pool_recycle: PositiveInt = 60 # in seconds: 1 minute
|
||||
pool_timeout: PositiveInt = 60 # in seconds: 1 minute
|
||||
221
birthday_pool_bot/repositories/tables.py
Normal file
221
birthday_pool_bot/repositories/tables.py
Normal file
@@ -0,0 +1,221 @@
|
||||
import uuid
|
||||
import enum
|
||||
from datetime import date
|
||||
from typing import Self, List
|
||||
|
||||
import sqlalchemy as sa
|
||||
from sqlmodel import SQLModel, Field, Relationship
|
||||
|
||||
from birthday_pool_bot.dto import (
|
||||
BankEnum,
|
||||
PaymentData as DTOPaymentData,
|
||||
Pool as DTOPool,
|
||||
Subscription as DTOSubscription,
|
||||
User as DTOUser,
|
||||
)
|
||||
|
||||
|
||||
class BaseSQLModel(SQLModel):
|
||||
@classmethod
|
||||
def from_item(cls, item: BaseModel) -> Self:
|
||||
raise NotImplementedError
|
||||
|
||||
def to_item(self) -> BaseModel:
|
||||
raise NotImplementedError
|
||||
|
||||
def to_values(self) -> dict[str, Any]:
|
||||
return {column.name: getattr(self, column.name) for column in self.__table__.columns}
|
||||
|
||||
|
||||
class User(BaseSQLModel, table=True):
|
||||
__tablename__ = "users"
|
||||
|
||||
__table_args__ = (
|
||||
sa.Index("ix__users__phone", "phone"),
|
||||
sa.Index("ix__users__telegram_id", "telegram_id"),
|
||||
sa.UniqueConstraint("phone", name="uq__users__phone"),
|
||||
sa.UniqueConstraint("telegram_id", name="uq__users__telegram_id"),
|
||||
)
|
||||
|
||||
id: uuid.UUID = Field(
|
||||
default_factory=uuid.uuid4,
|
||||
primary_key=True,
|
||||
nullable=False,
|
||||
sa_column_kwargs={"unique": True},
|
||||
)
|
||||
name: str | None = Field(nullable=True)
|
||||
birthday: date | None = Field(nullable=True)
|
||||
phone: str | None = Field(default=None, nullable=True)
|
||||
telegram_id: int | None = Field(default=None, nullable=True)
|
||||
gift_payment_data: dict | None = Field(
|
||||
sa_column=sa.Column(sa.JSON, nullable=True),
|
||||
default_factory=dict,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_item(cls, item: DTOUser) -> Self:
|
||||
return cls(
|
||||
id=item.id,
|
||||
name=item.name,
|
||||
birthday=item.birthday,
|
||||
phone=item.phone,
|
||||
telegram_id=item.telegram_id,
|
||||
gift_payment_data=(
|
||||
item.gift_payment_data.model_dump_json()
|
||||
if item.gift_payment_data is not None else
|
||||
None
|
||||
),
|
||||
)
|
||||
|
||||
def to_item(self) -> DTOUser:
|
||||
return DTOUser(
|
||||
id=self.id,
|
||||
name=self.name,
|
||||
birthday=self.birthday,
|
||||
phone=self.phone,
|
||||
telegram_id=self.telegram_id,
|
||||
gift_payment_data=(
|
||||
DTOPaymentData.model_validate_json(self.gift_payment_data)
|
||||
if self.gift_payment_data is not None else
|
||||
None
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class Pool(BaseSQLModel, table=True):
|
||||
__tablename__ = "pools"
|
||||
|
||||
__table_args__ = (
|
||||
sa.Index("ix__pools__owner_id", "owner_id"),
|
||||
sa.Index("ix__pools__birthday_user_id", "birthday_user_id"),
|
||||
sa.CheckConstraint(
|
||||
"owner_id <> birthday_user_id",
|
||||
name="ck__pools__owner_not_birthday_user",
|
||||
),
|
||||
)
|
||||
|
||||
id: uuid.UUID = Field(
|
||||
default_factory=uuid.uuid4,
|
||||
primary_key=True,
|
||||
nullable=False,
|
||||
sa_column_kwargs={"unique": True},
|
||||
)
|
||||
owner_id: uuid.UUID = Field(
|
||||
foreign_key="users.id", nullable=False,
|
||||
)
|
||||
birthday_user_id: uuid.UUID = Field(
|
||||
foreign_key="users.id", nullable=False,
|
||||
)
|
||||
description: str | None = Field(nullable=True)
|
||||
payment_data: dict = Field(
|
||||
sa_column=sa.Column(sa.JSON, nullable=False),
|
||||
default_factory=dict,
|
||||
)
|
||||
|
||||
owner: User = Relationship(
|
||||
sa_relationship_kwargs={
|
||||
"primaryjoin": "User.id == Pool.owner_id",
|
||||
"lazy": None,
|
||||
},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_item(cls, item: DTOPool) -> Self:
|
||||
return cls(
|
||||
id=item.id,
|
||||
owner_id=item.owner_id,
|
||||
birthday_user_id=item.birthday_user_id,
|
||||
description=item.description,
|
||||
payment_data=item.payment_data.model_dump_json(),
|
||||
owner=None if item.owner is None else DTOUser.from_item(item.owner),
|
||||
)
|
||||
|
||||
def to_item(self) -> DTOPool:
|
||||
return DTOPool(
|
||||
id=self.id,
|
||||
owner_id=self.owner_id,
|
||||
birthday_user_id=self.birthday_user_id,
|
||||
description=self.description,
|
||||
payment_data=DTOPaymentData.model_validate_json(self.payment_data),
|
||||
owner=None if self.owner is None else self.owner.to_item()
|
||||
)
|
||||
|
||||
|
||||
class Subscription(BaseSQLModel, table=True):
|
||||
__tablename__ = "subscriptions"
|
||||
|
||||
__table_args__ = (
|
||||
sa.Index("ix__subscriptions__from_user_id", "from_user_id"),
|
||||
sa.Index("ix__subscriptions__to_user_id", "to_user_id"),
|
||||
sa.Index("ix__subscriptions__name", "name"),
|
||||
sa.Index("ix__subscriptions__pool_id", "pool_id"),
|
||||
sa.Index(
|
||||
"ix__subscriptions__from_user_id__to_user_id",
|
||||
"from_user_id",
|
||||
"to_user_id",
|
||||
),
|
||||
sa.CheckConstraint(
|
||||
"from_user_id <> to_user_id",
|
||||
name="ck__subscriptions__from_user_not_to_user",
|
||||
),
|
||||
sa.UniqueConstraint(
|
||||
"from_user_id",
|
||||
"to_user_id",
|
||||
name="uq__subscriptions__from_to_user",
|
||||
),
|
||||
)
|
||||
|
||||
from_user_id: uuid.UUID = Field(
|
||||
foreign_key="users.id", primary_key=True, nullable=False,
|
||||
)
|
||||
to_user_id: uuid.UUID = Field(
|
||||
foreign_key="users.id", primary_key=True, nullable=False,
|
||||
)
|
||||
name: str = Field(nullable=False)
|
||||
pool_id: uuid.UUID = Field(
|
||||
foreign_key="pools.id",
|
||||
nullable=True,
|
||||
ondelete="SET NULL",
|
||||
)
|
||||
|
||||
from_user: "User" = Relationship(
|
||||
sa_relationship_kwargs={
|
||||
"primaryjoin": "User.id == Subscription.from_user_id",
|
||||
"lazy": None,
|
||||
},
|
||||
)
|
||||
to_user: "User" = Relationship(
|
||||
sa_relationship_kwargs={
|
||||
"primaryjoin": "User.id == Subscription.to_user_id",
|
||||
"lazy": None,
|
||||
},
|
||||
)
|
||||
pool: "Pool" = Relationship(
|
||||
sa_relationship_kwargs={
|
||||
"primaryjoin": "Pool.id == Subscription.pool_id",
|
||||
"lazy": None,
|
||||
},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_item(cls, item: DTOSubscription) -> Self:
|
||||
return cls(
|
||||
from_user_id=item.from_user_id,
|
||||
to_user_id=item.to_user_id,
|
||||
name=item.name,
|
||||
pool_id=item.pool_id,
|
||||
from_user=None if item.from_user is None else User.from_item(item.from_user),
|
||||
to_user=None if item.to_user is None else User.from_item(item.to_user),
|
||||
pool=None if item.pool is None else Pool.from_item(item.pool),
|
||||
)
|
||||
|
||||
def to_item(self) -> DTOSubscription:
|
||||
return DTOSubscription(
|
||||
from_user_id=self.from_user_id,
|
||||
to_user_id=self.to_user_id,
|
||||
name=self.name,
|
||||
pool_id=self.pool_id,
|
||||
from_user=None if self.from_user is None else self.from_user.to_item(),
|
||||
to_user=None if self.to_user is None else self.to_user.to_item(),
|
||||
pool=None if self.pool is None else self.pool.to_item(),
|
||||
)
|
||||
Reference in New Issue
Block a user