Init project

This commit is contained in:
2026-01-15 09:47:43 +03:00
commit 3e515c66ec
59 changed files with 5101 additions and 0 deletions

View 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",
)

View 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)

View 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

View 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

View File

@@ -0,0 +1,10 @@
from .cli import get_cli
from .migrator import Migrator
__all__ = (
# cli
"get_cli",
# migrator
"Migrator",
)

View 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

View 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()

View 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

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

View File

@@ -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 ###

View File

@@ -0,0 +1,13 @@
from .pools import PoolsRepository
from .subscriptions import SubscriptionsRepository
from .users import UsersRepository
__all__ = (
# pools
"PoolsRepository",
# subscriptions
"SubscriptionsRepository",
# users
"UsersRepository",
)

View 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)

View 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_)

View 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

View 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

View 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(),
)