Single part at the root, COMPOSE_FILE in .env
This commit is contained in:
+19
@@ -0,0 +1,19 @@
|
||||
{% if database == 'postgres' -%}
|
||||
async def init_db() -> None:
|
||||
from . import models # noqa: F401, PLC0415
|
||||
{%- else -%}
|
||||
from beanie import init_beanie
|
||||
from pymongo import AsyncMongoClient
|
||||
|
||||
from utils.env import env
|
||||
|
||||
client = AsyncMongoClient(env.db.connection_url)
|
||||
|
||||
|
||||
async def init_db() -> None:
|
||||
from .models import {% if dynamic_config %}DynamicConfig, {% endif %}User # noqa: PLC0415
|
||||
|
||||
{% if dynamic_config %} await init_beanie(
|
||||
database=client[env.db.db_name], document_models=[DynamicConfig, User]
|
||||
){% else %} await init_beanie(database=client[env.db.db_name], document_models=[User]){% endif %}
|
||||
{%- endif %}
|
||||
+4
@@ -0,0 +1,4 @@
|
||||
{% if dynamic_config %}from .config import BotConfig, DynamicConfig, DynamicConfigBase
|
||||
{% endif %}from .user import User
|
||||
|
||||
__all__ = [{% if dynamic_config %}"BotConfig", "DynamicConfig", "DynamicConfigBase", {% endif %}"User"]
|
||||
+54
@@ -0,0 +1,54 @@
|
||||
{% if database == 'postgres' -%}
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import BigInteger, Column
|
||||
from sqlmodel import Field as SQLField
|
||||
from sqlmodel import SQLModel
|
||||
|
||||
from .base import created_at_col
|
||||
|
||||
|
||||
class User(SQLModel, table=True):
|
||||
__tablename__ = "user_account"
|
||||
|
||||
id: int | None = SQLField(default=None, primary_key=True)
|
||||
tg_id: int = SQLField(sa_column=Column(BigInteger, unique=True, index=True))
|
||||
username: str | None = None
|
||||
first_name: str | None = None
|
||||
last_name: str | None = None
|
||||
|
||||
created_at: datetime | None = created_at_col()
|
||||
{%- else -%}
|
||||
from beanie import Document
|
||||
|
||||
|
||||
class User(Document):
|
||||
id: int
|
||||
balance: float = 0
|
||||
|
||||
class Settings:
|
||||
name = "users"
|
||||
|
||||
@classmethod
|
||||
async def get_by_id(cls, id_: int) -> "User | None":
|
||||
return await cls.find_one(cls.id == id_)
|
||||
|
||||
@classmethod
|
||||
async def get_or_create(cls, id_: int) -> "User":
|
||||
user = await cls.get_by_id(id_)
|
||||
if user is None:
|
||||
user = cls(id=id_)
|
||||
await user.insert()
|
||||
return user
|
||||
|
||||
async def add_balance(self, amount: float) -> None:
|
||||
self.balance += amount
|
||||
await self.save()
|
||||
|
||||
async def subtract_balance(self, amount: float) -> bool:
|
||||
if self.balance >= amount:
|
||||
self.balance -= amount
|
||||
await self.save()
|
||||
return True
|
||||
return False
|
||||
{%- endif %}
|
||||
+60
@@ -0,0 +1,60 @@
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import BigInteger, Column, DateTime, ForeignKey, func, text
|
||||
from sqlalchemy.dialects import postgresql
|
||||
from sqlmodel import Field as SQLField
|
||||
|
||||
|
||||
def uuid_pk() -> uuid.UUID:
|
||||
return SQLField(
|
||||
default_factory=uuid.uuid4,
|
||||
sa_column=Column(postgresql.UUID(as_uuid=True), primary_key=True),
|
||||
)
|
||||
|
||||
|
||||
def uuid_fk(target: str, *, nullable: bool = True) -> Column:
|
||||
return Column(
|
||||
postgresql.UUID(as_uuid=True),
|
||||
ForeignKey(target, ondelete="SET NULL"),
|
||||
nullable=nullable,
|
||||
)
|
||||
|
||||
|
||||
def bigint_col(*, nullable: bool = True) -> Column:
|
||||
return Column(BigInteger, nullable=nullable)
|
||||
|
||||
|
||||
def nullable_ts_col() -> datetime | None:
|
||||
return SQLField(
|
||||
default=None, sa_column=Column(DateTime(timezone=True), nullable=True)
|
||||
)
|
||||
|
||||
|
||||
def created_at_col() -> datetime:
|
||||
return SQLField(
|
||||
default=None,
|
||||
sa_column=Column(
|
||||
DateTime(timezone=True), nullable=False, server_default=func.now()
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def updated_at_col() -> datetime:
|
||||
return SQLField(
|
||||
default=None,
|
||||
sa_column=Column(
|
||||
DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=func.now(),
|
||||
onupdate=func.now(),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def jsonb_column(*, nullable: bool = False, default: str = "[]") -> Column:
|
||||
return Column(
|
||||
postgresql.JSONB,
|
||||
nullable=nullable,
|
||||
server_default=None if nullable else text(f"'{default}'::jsonb"),
|
||||
)
|
||||
+58
@@ -0,0 +1,58 @@
|
||||
{% if database == 'postgres' -%}
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from sqlalchemy import Column, DateTime, func
|
||||
from sqlalchemy.dialects import postgresql
|
||||
from sqlmodel import Field as SQLField
|
||||
from sqlmodel import SQLModel
|
||||
|
||||
|
||||
class BotConfig(BaseModel):
|
||||
admins: list[int] = Field(default_factory=list)
|
||||
|
||||
|
||||
class DynamicConfigBase(BaseModel):
|
||||
bot: BotConfig = Field(default_factory=BotConfig)
|
||||
|
||||
|
||||
class DynamicConfig(SQLModel, table=True):
|
||||
__tablename__ = "config"
|
||||
|
||||
id: int = SQLField(default=1, primary_key=True)
|
||||
data: dict[str, Any] = SQLField(sa_column=Column(postgresql.JSONB, nullable=False))
|
||||
updated_at: datetime | None = SQLField(
|
||||
default=None,
|
||||
sa_column=Column(
|
||||
DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=func.now(),
|
||||
onupdate=func.now(),
|
||||
),
|
||||
)
|
||||
{%- else -%}
|
||||
from beanie import Document
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class BotConfig(BaseModel):
|
||||
admins: list[int] = Field(default_factory=list)
|
||||
|
||||
|
||||
class DynamicConfigBase(BaseModel):
|
||||
bot: BotConfig = Field(default_factory=BotConfig)
|
||||
|
||||
|
||||
class DynamicConfig(DynamicConfigBase, Document):
|
||||
class Settings:
|
||||
name = "config"
|
||||
|
||||
@classmethod
|
||||
async def get_or_create(cls) -> "DynamicConfig":
|
||||
config = await cls.find_one()
|
||||
if config is None:
|
||||
config = cls()
|
||||
await config.save()
|
||||
return config
|
||||
{%- endif %}
|
||||
+5
@@ -0,0 +1,5 @@
|
||||
{% if use_redis %}from .cache import RedisCache
|
||||
{% endif %}{% if dynamic_config %}from .config import ConfigRepository
|
||||
{% endif %}from .user import UserRepository
|
||||
|
||||
__all__ = [{% if dynamic_config %}"ConfigRepository", {% endif %}{% if use_redis %}"RedisCache", {% endif %}"UserRepository"]
|
||||
+35
@@ -0,0 +1,35 @@
|
||||
from sqlmodel import select
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from utils.db.models import User
|
||||
|
||||
|
||||
class UserRepository:
|
||||
def __init__(self, session: AsyncSession) -> None:
|
||||
self._session = session
|
||||
|
||||
async def get(self, user_id: int) -> User | None:
|
||||
return await self._session.get(User, user_id)
|
||||
|
||||
async def by_tg_id(self, tg_id: int) -> User | None:
|
||||
stmt = select(User).where(User.tg_id == tg_id)
|
||||
return (await self._session.exec(stmt)).first()
|
||||
|
||||
async def upsert(
|
||||
self,
|
||||
tg_id: int,
|
||||
*,
|
||||
username: str | None = None,
|
||||
first_name: str | None = None,
|
||||
last_name: str | None = None,
|
||||
) -> User:
|
||||
user = await self.by_tg_id(tg_id)
|
||||
if user is None:
|
||||
user = User(tg_id=tg_id)
|
||||
user.username = username
|
||||
user.first_name = first_name
|
||||
user.last_name = last_name
|
||||
self._session.add(user)
|
||||
await self._session.commit()
|
||||
await self._session.refresh(user)
|
||||
return user
|
||||
+88
@@ -0,0 +1,88 @@
|
||||
{% if use_redis -%}
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from utils.db.models.config import DynamicConfig, DynamicConfigBase
|
||||
|
||||
from .cache import RedisCache
|
||||
|
||||
CONFIG_CACHE_KEY = "config:dynamic"
|
||||
CONFIG_ROW_ID = 1
|
||||
|
||||
|
||||
class ConfigRepository:
|
||||
def __init__(
|
||||
self, sessionmaker: async_sessionmaker[AsyncSession], cache: RedisCache
|
||||
) -> None:
|
||||
self._sessionmaker = sessionmaker
|
||||
self._cache = cache
|
||||
|
||||
async def get(self) -> DynamicConfigBase:
|
||||
return await self._cache.get_or_set_model(
|
||||
CONFIG_CACHE_KEY, DynamicConfigBase, self._load
|
||||
)
|
||||
|
||||
async def _load(self) -> DynamicConfigBase:
|
||||
async with self._sessionmaker() as session:
|
||||
row = await session.get(DynamicConfig, CONFIG_ROW_ID)
|
||||
if row is None:
|
||||
config = DynamicConfigBase()
|
||||
session.add(DynamicConfig(id=CONFIG_ROW_ID, data=config.model_dump()))
|
||||
await session.commit()
|
||||
return config
|
||||
return DynamicConfigBase.model_validate(row.data)
|
||||
|
||||
async def save(self, config: DynamicConfigBase) -> DynamicConfigBase:
|
||||
async with self._sessionmaker() as session:
|
||||
row = await session.get(DynamicConfig, CONFIG_ROW_ID)
|
||||
if row is None:
|
||||
session.add(DynamicConfig(id=CONFIG_ROW_ID, data=config.model_dump()))
|
||||
else:
|
||||
row.data = config.model_dump()
|
||||
session.add(row)
|
||||
await session.commit()
|
||||
await self._cache.set_model(CONFIG_CACHE_KEY, config)
|
||||
return config
|
||||
|
||||
async def reset(self) -> DynamicConfigBase:
|
||||
return await self.save(DynamicConfigBase())
|
||||
|
||||
async def invalidate(self) -> None:
|
||||
await self._cache.invalidate(CONFIG_CACHE_KEY)
|
||||
{%- else -%}
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker
|
||||
from sqlmodel.ext.asyncio.session import AsyncSession
|
||||
|
||||
from utils.db.models.config import DynamicConfig, DynamicConfigBase
|
||||
|
||||
CONFIG_ROW_ID = 1
|
||||
|
||||
|
||||
class ConfigRepository:
|
||||
def __init__(self, sessionmaker: async_sessionmaker[AsyncSession]) -> None:
|
||||
self._sessionmaker = sessionmaker
|
||||
|
||||
async def get(self) -> DynamicConfigBase:
|
||||
async with self._sessionmaker() as session:
|
||||
row = await session.get(DynamicConfig, CONFIG_ROW_ID)
|
||||
if row is None:
|
||||
config = DynamicConfigBase()
|
||||
session.add(DynamicConfig(id=CONFIG_ROW_ID, data=config.model_dump()))
|
||||
await session.commit()
|
||||
return config
|
||||
return DynamicConfigBase.model_validate(row.data)
|
||||
|
||||
async def save(self, config: DynamicConfigBase) -> DynamicConfigBase:
|
||||
async with self._sessionmaker() as session:
|
||||
row = await session.get(DynamicConfig, CONFIG_ROW_ID)
|
||||
if row is None:
|
||||
session.add(DynamicConfig(id=CONFIG_ROW_ID, data=config.model_dump()))
|
||||
else:
|
||||
row.data = config.model_dump()
|
||||
session.add(row)
|
||||
await session.commit()
|
||||
return config
|
||||
|
||||
async def reset(self) -> DynamicConfigBase:
|
||||
return await self.save(DynamicConfigBase())
|
||||
{%- endif %}
|
||||
+52
@@ -0,0 +1,52 @@
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import TypeVar
|
||||
|
||||
from pydantic import BaseModel
|
||||
from redis.asyncio import Redis
|
||||
|
||||
ModelT = TypeVar("ModelT", bound=BaseModel)
|
||||
NAMESPACE = "{{ project_slug }}"
|
||||
|
||||
|
||||
class RedisCache:
|
||||
def __init__(
|
||||
self, redis: Redis, *, namespace: str = NAMESPACE, ttl: int = 300
|
||||
) -> None:
|
||||
self._redis = redis
|
||||
self._namespace = namespace
|
||||
self._ttl = ttl
|
||||
|
||||
def _key(self, key: str) -> str:
|
||||
return f"{self._namespace}:{key}"
|
||||
|
||||
async def get_model(self, key: str, model: type[ModelT]) -> ModelT | None:
|
||||
raw = await self._redis.get(self._key(key))
|
||||
if raw is None:
|
||||
return None
|
||||
return model.model_validate_json(raw)
|
||||
|
||||
async def set_model(
|
||||
self, key: str, value: BaseModel, *, ttl: int | None = None
|
||||
) -> None:
|
||||
await self._redis.set(
|
||||
self._key(key), value.model_dump_json(), ex=ttl or self._ttl
|
||||
)
|
||||
|
||||
async def get_or_set_model(
|
||||
self,
|
||||
key: str,
|
||||
model: type[ModelT],
|
||||
loader: Callable[[], Awaitable[ModelT]],
|
||||
*,
|
||||
ttl: int | None = None,
|
||||
) -> ModelT:
|
||||
cached = await self.get_model(key, model)
|
||||
if cached is not None:
|
||||
return cached
|
||||
value = await loader()
|
||||
await self.set_model(key, value, ttl=ttl)
|
||||
return value
|
||||
|
||||
async def invalidate(self, *keys: str) -> None:
|
||||
if keys:
|
||||
await self._redis.delete(*(self._key(key) for key in keys))
|
||||
Reference in New Issue
Block a user