start
This commit is contained in:
27
api/db/connection.py
Normal file
27
api/db/connection.py
Normal file
@@ -0,0 +1,27 @@
|
||||
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from sqlalchemy import text
|
||||
import asyncio
|
||||
from api.config import settings
|
||||
|
||||
DATABASE_URL = settings.DATABASE_URL
|
||||
engine = create_async_engine(DATABASE_URL, echo=False, future=True)
|
||||
SessionLocal = sessionmaker(engine, expire_on_commit=False, class_=AsyncSession)
|
||||
|
||||
|
||||
async def wait_for_db(retries: int = 12, delay: float = 1.0) -> None:
|
||||
"""Wait until the database is available. Retries with exponential backoff.
|
||||
|
||||
Raises RuntimeError if DB is still unavailable after retries.
|
||||
"""
|
||||
last_exc = None
|
||||
for attempt in range(1, retries + 1):
|
||||
try:
|
||||
async with engine.connect() as conn:
|
||||
await conn.execute(text("SELECT 1"))
|
||||
return
|
||||
except Exception as e:
|
||||
last_exc = e
|
||||
wait = min(delay * (2 ** (attempt - 1)), 5)
|
||||
await asyncio.sleep(wait)
|
||||
raise RuntimeError(f"Could not connect to DB after {retries} attempts") from last_exc
|
||||
121
api/db/subscription.py
Normal file
121
api/db/subscription.py
Normal file
@@ -0,0 +1,121 @@
|
||||
from api.db.connection import engine, SessionLocal
|
||||
from api.models.subscription import Subscription
|
||||
from api.models.base import Base
|
||||
from api.config import settings
|
||||
from datetime import datetime, timedelta
|
||||
from sqlalchemy import select
|
||||
import asyncpg
|
||||
|
||||
|
||||
async def init_db():
|
||||
try:
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
return
|
||||
except Exception as exc: # try to create database if it doesn't exist
|
||||
msg = str(exc).lower()
|
||||
if "does not exist" in msg or "invalidcatalogname" in msg or "database" in msg:
|
||||
try:
|
||||
admin_conn = await asyncpg.connect(
|
||||
user=settings.DB_USER,
|
||||
password=settings.DB_PASS,
|
||||
database="postgres",
|
||||
host=settings.DB_HOST,
|
||||
port=int(settings.DB_PORT),
|
||||
)
|
||||
try:
|
||||
await admin_conn.execute(f'CREATE DATABASE "{settings.DB_NAME}"')
|
||||
except Exception:
|
||||
pass
|
||||
await admin_conn.close()
|
||||
except Exception:
|
||||
raise
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
return
|
||||
raise
|
||||
|
||||
|
||||
async def create_subscription(user_id: str, uuid: str, until: datetime, subscription_link: str | None = None, connect_link: str | None = None, active: bool = True):
|
||||
async with SessionLocal() as session:
|
||||
# if user already has a subscription, remove it (replace)
|
||||
q = select(Subscription).where(Subscription.user_id == user_id)
|
||||
res = await session.execute(q)
|
||||
existing = res.scalars().first()
|
||||
if existing:
|
||||
await session.delete(existing)
|
||||
# avoid NULL in DB: replace None links with empty strings
|
||||
safe_subscription_link = subscription_link if subscription_link is not None else ""
|
||||
safe_connect_link = connect_link if connect_link is not None else ""
|
||||
sub = Subscription(user_id=user_id, uuid=uuid, until=until, active=active, subscription_link=safe_subscription_link, connect_link=safe_connect_link)
|
||||
session.add(sub)
|
||||
await session.commit()
|
||||
await session.refresh(sub)
|
||||
return sub
|
||||
|
||||
|
||||
async def get_subscription_by_user(user_id: str):
|
||||
async with SessionLocal() as session:
|
||||
q = select(Subscription).where(Subscription.user_id == user_id)
|
||||
res = await session.execute(q)
|
||||
sub = res.scalars().first()
|
||||
return sub
|
||||
|
||||
|
||||
async def get_subscription_by_uuid(sub_uuid: str):
|
||||
async with SessionLocal() as session:
|
||||
q = select(Subscription).where(Subscription.uuid == sub_uuid)
|
||||
res = await session.execute(q)
|
||||
return res.scalars().first()
|
||||
|
||||
|
||||
async def update_subscription(sub_uuid: str, until: datetime = None, user_id: str | None = None, active: bool | None = None, subscription_link: str | None = None, connect_link: str | None = None):
|
||||
async with SessionLocal() as session:
|
||||
q = select(Subscription).where(Subscription.uuid == sub_uuid)
|
||||
res = await session.execute(q)
|
||||
sub = res.scalars().first()
|
||||
if not sub:
|
||||
return None
|
||||
if until is not None:
|
||||
sub.until = until
|
||||
if user_id is not None:
|
||||
sub.user_id = user_id
|
||||
if active is not None:
|
||||
sub.active = active
|
||||
if subscription_link is not None:
|
||||
sub.subscription_link = subscription_link
|
||||
if connect_link is not None:
|
||||
sub.connect_link = connect_link
|
||||
await session.commit()
|
||||
return sub
|
||||
|
||||
|
||||
async def delete_subscription(sub_uuid: str):
|
||||
async with SessionLocal() as session:
|
||||
q = select(Subscription).where(Subscription.uuid == sub_uuid)
|
||||
res = await session.execute(q)
|
||||
sub = res.scalars().first()
|
||||
if not sub:
|
||||
return False
|
||||
await session.delete(sub)
|
||||
await session.commit()
|
||||
return True
|
||||
|
||||
|
||||
async def list_subscriptions():
|
||||
async with SessionLocal() as session:
|
||||
q = select(Subscription)
|
||||
res = await session.execute(q)
|
||||
return res.scalars().all()
|
||||
|
||||
|
||||
async def extend_subscription_by_user(user_id: str, months: int):
|
||||
sub = await get_subscription_by_user(user_id)
|
||||
if sub:
|
||||
new_until = max(sub.until, datetime.now()) + timedelta(days=30 * months)
|
||||
return await update_subscription(sub.uuid, until=new_until, active=True)
|
||||
return None
|
||||
|
||||
|
||||
async def deactivate_subscription(sub_uuid: str):
|
||||
return await update_subscription(sub_uuid, active=False)
|
||||
64
api/db/user.py
Normal file
64
api/db/user.py
Normal file
@@ -0,0 +1,64 @@
|
||||
from api.db.connection import SessionLocal
|
||||
from api.models.user import User
|
||||
from sqlalchemy import select
|
||||
|
||||
|
||||
import uuid
|
||||
|
||||
|
||||
async def create_user(telegram_id: str, name: str | None = None):
|
||||
async with SessionLocal() as session:
|
||||
# ensure no NULL in DB: replace None with empty string
|
||||
safe_name = name if name is not None else ""
|
||||
user = User(id=str(uuid.uuid4()), telegram_id=telegram_id, name=safe_name)
|
||||
session.add(user)
|
||||
await session.commit()
|
||||
await session.refresh(user)
|
||||
return user
|
||||
|
||||
|
||||
async def get_user(user_id: str):
|
||||
async with SessionLocal() as session:
|
||||
q = select(User).where(User.id == user_id)
|
||||
res = await session.execute(q)
|
||||
return res.scalars().first()
|
||||
|
||||
|
||||
async def get_user_by_telegram(telegram_id: str):
|
||||
async with SessionLocal() as session:
|
||||
q = select(User).where(User.telegram_id == telegram_id)
|
||||
res = await session.execute(q)
|
||||
return res.scalars().first()
|
||||
|
||||
|
||||
async def update_user(user_id: str, name: str | None = None):
|
||||
async with SessionLocal() as session:
|
||||
q = select(User).where(User.id == user_id)
|
||||
res = await session.execute(q)
|
||||
user = res.scalars().first()
|
||||
if not user:
|
||||
return None
|
||||
if name is not None:
|
||||
# replace None handled by caller; here update only when provided
|
||||
user.name = name
|
||||
await session.commit()
|
||||
return user
|
||||
|
||||
|
||||
async def delete_user(user_id: str):
|
||||
async with SessionLocal() as session:
|
||||
q = select(User).where(User.id == user_id)
|
||||
res = await session.execute(q)
|
||||
user = res.scalars().first()
|
||||
if not user:
|
||||
return False
|
||||
await session.delete(user)
|
||||
await session.commit()
|
||||
return True
|
||||
|
||||
|
||||
async def list_users():
|
||||
async with SessionLocal() as session:
|
||||
q = select(User)
|
||||
res = await session.execute(q)
|
||||
return res.scalars().all()
|
||||
Reference in New Issue
Block a user