diff --git a/.gitignore b/.gitignore index 3b25f6e9..33d6982c 100644 --- a/.gitignore +++ b/.gitignore @@ -27,15 +27,19 @@ env/ *.db # Sensitive configuration files -config.py -config.ini -alembic.ini -.env +/config.py +/config.ini +/alembic.ini +/.env # Backup files *.bak *.swp *~ +backups/ +backup_bot/ +Solo_backup/ +.cursor/ # Specific project files vpn_users.db @@ -51,15 +55,23 @@ handlers/texts.py Thumbs.db nginx.conf -scripts +/scripts/load_balancer.py +/scripts/__pycache__ .csv /logs -setup.py .ruff_cache -.github/workflows/ modules/ storage/ static/web_uploads/ -/web-app/ -.license_state \ No newline at end of file +.license_state +Solo_backup/ + +.cursor/hooks.json +.cursor/hooks/after-agent-response.cjs +.cursor/hooks/before-submit-prompt.cjs +.cursor/hooks/after-agent-response.js +.cursor/hooks/before-submit-prompt.js +.cursor/cursor-notifier.json +.cursor/cursor-notifier-start.json +nuitka-crash-report* diff --git a/Makefile b/Makefile index 16aac4e6..ffb0eaaf 100644 --- a/Makefile +++ b/Makefile @@ -8,3 +8,12 @@ lint: format-payments: @echo "Running Ruff format ONLY on handlers/payments..." && ruff format handlers/payments --config pyproject.toml @echo "Running Ruff check ONLY on handlers/payments..." && ruff check handlers/payments --config pyproject.toml --fix + +test: + @echo "Running unit tests..." && cd /tmp && PYTHONPATH="$(CURDIR)" "$(CURDIR)/venv/bin/python" -m unittest discover -s "$(CURDIR)/tests" -q + +test-sudo: + @echo "Running unit tests with sudo..." && cd /tmp && sudo env PYTHONPATH="$(CURDIR)" "$(CURDIR)/venv/bin/python" -m unittest discover -s "$(CURDIR)/tests" -q + +smoke: + @echo "Running smoke checks..." && bash "$(CURDIR)/tests/smoke_runner.sh" diff --git a/api/depends.py b/api/depends.py index 9764683e..fec6baa6 100644 --- a/api/depends.py +++ b/api/depends.py @@ -1,14 +1,16 @@ import hashlib +from urllib.parse import urlparse from collections.abc import AsyncGenerator -from fastapi import Depends, HTTPException, Header, Query, Request +from fastapi import Depends, HTTPException, Header, Query, Request, Response from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from audit import set_api_actor from database import async_session_maker, identities as idb -from database.models import Admin +from database.access.resolution import ResolvedActor, resolve_actor_from_identity +from database.models import Admin, Identity async def get_session() -> AsyncGenerator[AsyncSession, None]: @@ -25,6 +27,24 @@ def hash_token(token: str) -> str: return hashlib.sha256(token.encode()).hexdigest() +async def bind_identity_actor( + request: Request | None, + session: AsyncSession, + identity: Identity, +) -> ResolvedActor: + actor = await resolve_actor_from_identity(session, identity) + set_api_actor(request, identity_id=actor.identity_id, tg_id=actor.telegram_chat_id) + if request is not None: + request.state.actor = actor + return actor + + +def get_request_actor(request: Request | None) -> ResolvedActor | None: + if request is None: + return None + return getattr(request.state, "actor", None) + + async def verify_admin_token( admin_id: int = Query(..., alias="tg_id"), token: str = Header(..., alias="X-Token"), @@ -40,50 +60,146 @@ async def verify_admin_token( return admin +AUTH_COOKIE_NAME = "auth_token" + + +IS_ADMIN_COOKIE_NAME = "is_admin" + + +AUTH_COOKIE_MAX_AGE_SECONDS = 30 * 24 * 60 * 60 + + +def _is_secure_request(request: Request | None) -> bool: + if request is None: + return False + if request.url.scheme == "https": + return True + forwarded_proto = request.headers.get("x-forwarded-proto", "").lower() + return forwarded_proto == "https" + + +def set_auth_cookie(response: Response, token: str, request: Request | None = None) -> None: + """Устанавливает HttpOnly cookie с auth-токеном на ответ. Используется во всех login-ручках.""" + response.set_cookie( + key=AUTH_COOKIE_NAME, + value=token, + max_age=AUTH_COOKIE_MAX_AGE_SECONDS, + path="/", + httponly=True, + secure=_is_secure_request(request), + samesite="lax", + ) + + +def clear_auth_cookie(response: Response, request: Request | None = None) -> None: + """Удаляет auth cookie на стороне браузера.""" + response.delete_cookie( + key=AUTH_COOKIE_NAME, + path="/", + httponly=True, + secure=_is_secure_request(request), + samesite="lax", + ) + + clear_is_admin_cookie(response, request) + + +def set_is_admin_cookie(response: Response, identity: Identity, request: Request | None = None) -> None: + """Ставит/гасит `is_admin` cookie в зависимости от текущей identity.""" + if getattr(identity, "is_admin", False): + response.set_cookie( + key=IS_ADMIN_COOKIE_NAME, + value="1", + max_age=AUTH_COOKIE_MAX_AGE_SECONDS, + path="/", + httponly=True, + secure=_is_secure_request(request), + samesite="lax", + ) + else: + clear_is_admin_cookie(response, request) + + +def clear_is_admin_cookie(response: Response, request: Request | None = None) -> None: + response.delete_cookie( + key=IS_ADMIN_COOKIE_NAME, + path="/", + httponly=True, + secure=_is_secure_request(request), + samesite="lax", + ) + + +def _read_auth_cookie(request: Request | None) -> str | None: + if request is None: + return None + raw = request.cookies.get(AUTH_COOKIE_NAME) + if not raw: + return None + raw = raw.strip() + return raw or None + + +async def _identity_from_cookie(session: AsyncSession, request: Request | None) -> Identity | None: + token = _read_auth_cookie(request) + if not token: + return None + token_hash = hash_token(token) + identity = await idb.get_identity_by_token_hash(session, token_hash) + if identity is None: + return None + if idb._is_token_expired(identity): + return None + return identity + + async def verify_identity_token( - x_identity_id: str = Header(..., alias="X-Identity-Id"), - token: str = Header(..., alias="X-Token"), - request: Request = None, + request: Request, session: AsyncSession = Depends(get_session), ): - """Проверяет пару identity_id + token; возвращает Identity. Для использования в API v2.""" - identity = await idb.verify_identity_token(session, x_identity_id, token) - if not identity: + """Проверяет токен из HttpOnly cookie `auth_token`; возвращает Identity.""" + identity = await _identity_from_cookie(session, request) + if identity is None: raise HTTPException(status_code=401, detail="Unauthorized") - set_api_actor(request, identity_id=identity.id, tg_id=identity.tg_id) + await bind_identity_actor(request, session, identity) return identity async def verify_identity_admin( - x_identity_id: str = Header(..., alias="X-Identity-Id"), - token: str = Header(..., alias="X-Token"), - request: Request = None, + request: Request, session: AsyncSession = Depends(get_session), ): - """Проверяет identity + token и что identity.is_admin; для админских ручек v2.""" - identity = await idb.verify_identity_token(session, x_identity_id, token) - if not identity: + """Проверяет токен из cookie и что identity.is_admin; для админских ручек v2.""" + identity = await _identity_from_cookie(session, request) + if identity is None: raise HTTPException(status_code=401, detail="Unauthorized") if not identity.is_admin: raise HTTPException(status_code=403, detail="Forbidden") - set_api_actor(request, identity_id=identity.id, tg_id=identity.tg_id) + await bind_identity_actor(request, session, identity) return identity async def verify_identity_admin_short( - x_identity_id: str = Header(..., alias="X-Identity-Id"), - token: str = Header(..., alias="X-Token"), - request: Request = None, + request: Request, ): """Проверка админа с короткой сессией (для broadcast и др.), чтобы не держать соединение с БД.""" + identity = None + actor = None async with async_session_maker() as session: - identity = await idb.verify_identity_token(session, x_identity_id, token) + identity = await _identity_from_cookie(session, request) + if identity: + actor = await resolve_actor_from_identity(session, identity) await session.commit() if not identity: raise HTTPException(status_code=401, detail="Unauthorized") if not identity.is_admin: raise HTTPException(status_code=403, detail="Forbidden") - set_api_actor(request, identity_id=identity.id, tg_id=identity.tg_id) + if actor is not None: + set_api_actor(request, identity_id=identity.id, tg_id=actor.telegram_chat_id) + if request is not None: + request.state.actor = actor + else: + set_api_actor(request, identity_id=identity.id, tg_id=identity.tg_id) return identity @@ -102,3 +218,20 @@ async def verify_admin_token_short( raise HTTPException(status_code=401, detail="Unauthorized") set_api_actor(request, tg_id=admin.tg_id) return admin + + +def validate_redirect_url(url: str, base_url: str) -> str: + """Validate redirect URL is same-origin or relative. Returns safe URL or base_url fallback.""" + url = url.strip() + if not url: + return base_url + if url.startswith("/"): + return url + try: + parsed = urlparse(url) + base_parsed = urlparse(base_url) + if parsed.scheme in ("http", "https") and parsed.netloc == base_parsed.netloc: + return url + except Exception: + pass + return base_url diff --git a/api/main.py b/api/main.py index f6de29a0..502b1dc3 100644 --- a/api/main.py +++ b/api/main.py @@ -14,7 +14,8 @@ from logger import logger if API_VERSION == 1: from api.v1 import router as api_router, VERSION as API_DOC_VERSION else: - from api.v2 import router as api_router, VERSION as API_DOC_VERSION + from api.v2 import VERSION as API_DOC_VERSION + from api.v2.router import router as api_router app = FastAPI( title=f"SoloBot API (Alpha) — API v{API_DOC_VERSION}", @@ -25,15 +26,28 @@ app = FastAPI( openapi_url="/api/openapi.json", ) +_cors_origins = API_CORS_ORIGINS if API_CORS_ORIGINS != ["*"] else API_CORS_ORIGINS +_cors_credentials = API_CORS_ORIGINS != ["*"] + app.add_middleware( CORSMiddleware, - allow_origins=API_CORS_ORIGINS, - allow_credentials=True, - allow_methods=["*"], - allow_headers=["*"], + allow_origins=_cors_origins, + allow_credentials=_cors_credentials, + allow_methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS"], + allow_headers=["X-Identity-Id", "X-Token", "Content-Type", "Authorization"], ) +@app.middleware("http") +async def security_headers_middleware(request: Request, call_next): + response = await call_next(request) + response.headers.setdefault("X-Content-Type-Options", "nosniff") + response.headers.setdefault("X-Frame-Options", "DENY") + response.headers.setdefault("Referrer-Policy", "strict-origin-when-cross-origin") + response.headers.setdefault("X-XSS-Protection", "1; mode=block") + return response + + @app.middleware("http") async def api_access_log_middleware(request: Request, call_next): context = ensure_api_context(request) @@ -86,6 +100,11 @@ async def api_access_log_middleware(request: Request, call_next): return response +@app.get("/api/health", include_in_schema=False) +async def health(): + return {"status": "ok"} + + app.include_router(api_router) _web_uploads_dir = "static/web_uploads" diff --git a/api/v1/routes/base_crud.py b/api/v1/routes/base_crud.py index 0d9a6d73..57534871 100644 --- a/api/v1/routes/base_crud.py +++ b/api/v1/routes/base_crud.py @@ -1,12 +1,22 @@ from typing import Any from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query +from sqlalchemy import inspect as sa_inspect from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession +from sqlalchemy.orm import selectinload from sqlalchemy.orm.attributes import InstrumentedAttribute from api.depends import get_session, verify_admin_token from database.models import Admin +from database.access.resolution import resolve_user_optional +from handlers.texts import get_site_gift_link, get_telegram_gift_link + + +def _apply_user_relationship_loader(model: type, stmt): + if model.__name__ in ("ManualBan", "BlockedUser", "TemporaryData"): + return stmt.options(selectinload(model.user)) + return stmt def cast_identifier_type(field: InstrumentedAttribute, value: int | str): @@ -19,6 +29,21 @@ def cast_identifier_type(field: InstrumentedAttribute, value: int | str): def normalize_outgoing_object(obj: object) -> None: if hasattr(obj, "vless") and getattr(obj, "vless") is None: setattr(obj, "vless", False) + cls_name = type(obj).__name__ + if cls_name == "Gift": + gift_id = getattr(obj, "gift_id", None) + if gift_id: + setattr(obj, "telegram_gift_link", get_telegram_gift_link(gift_id)) + setattr(obj, "site_gift_link", get_site_gift_link(gift_id)) + if cls_name in ("ManualBan", "BlockedUser", "TemporaryData"): + insp = sa_inspect(obj) + stored = getattr(obj, "tg_id", None) + if "user" in insp.unloaded: + setattr(obj, "tg_id", stored) + return + rel = getattr(obj, "user", None) + rel_tg = getattr(rel, "tg_id", None) if rel is not None else None + setattr(obj, "tg_id", stored if stored is not None else rel_tg) def to_schema(schema_response: type, obj: object): @@ -35,10 +60,20 @@ def generate_crud_router( identifier_field: str = "tg_id", parameter_name: str = "tg_id", extra_get_by_email: bool = False, + telegram_path_to_user_id: bool = False, enabled_methods: list[str] = ("get_all", "get_one", "get_by_email", "create", "update", "delete"), ) -> APIRouter: router = APIRouter() + async def _path_filter(session: AsyncSession, value: int | str): + if telegram_path_to_user_id: + u = await resolve_user_optional(session, int(value)) + if u is None: + return None + return getattr(model, "user_id"), u.id + field = getattr(model, identifier_field) + return field, cast_identifier_type(field, value) + if "get_all" in enabled_methods: @router.get("/", response_model=list[schema_response]) @@ -46,7 +81,7 @@ def generate_crud_router( admin: Admin = Depends(verify_admin_token), session: AsyncSession = Depends(get_session), ): - result = await session.execute(select(model)) + result = await session.execute(_apply_user_relationship_loader(model, select(model))) items = result.scalars().all() for item in items: normalize_outgoing_object(item) @@ -74,9 +109,13 @@ def generate_crud_router( admin: Admin = Depends(verify_admin_token), session: AsyncSession = Depends(get_session), ): - field = getattr(model, identifier_field) - casted = cast_identifier_type(field, value) - result = await session.execute(select(model).where(field == casted)) + resolved = await _path_filter(session, value) + if resolved is None: + raise HTTPException(status_code=404, detail=f"{model.__name__} not found") + field, casted = resolved + result = await session.execute( + _apply_user_relationship_loader(model, select(model).where(field == casted)) + ) obj = result.scalar_one_or_none() if not obj: raise HTTPException(status_code=404, detail=f"{model.__name__} not found") @@ -90,9 +129,13 @@ def generate_crud_router( admin: Admin = Depends(verify_admin_token), session: AsyncSession = Depends(get_session), ): - field = getattr(model, identifier_field) - casted = cast_identifier_type(field, value) - result = await session.execute(select(model).where(field == casted)) + resolved = await _path_filter(session, value) + if resolved is None: + raise HTTPException(status_code=404, detail=f"{model.__name__} not found") + field, casted = resolved + result = await session.execute( + _apply_user_relationship_loader(model, select(model).where(field == casted)) + ) objs = result.scalars().all() if not objs: raise HTTPException(status_code=404, detail=f"{model.__name__} not found") @@ -127,8 +170,10 @@ def generate_crud_router( admin: Admin = Depends(verify_admin_token), session: AsyncSession = Depends(get_session), ): - field = getattr(model, identifier_field) - casted = cast_identifier_type(field, value) + resolved = await _path_filter(session, value) + if resolved is None: + raise HTTPException(status_code=404, detail=f"{model.__name__} not found") + field, casted = resolved result = await session.execute(select(model).where(field == casted)) obj = result.scalar_one_or_none() if not obj: @@ -150,8 +195,10 @@ def generate_crud_router( admin: Admin = Depends(verify_admin_token), session: AsyncSession = Depends(get_session), ): - field = getattr(model, identifier_field) - casted = cast_identifier_type(field, value) + resolved = await _path_filter(session, value) + if resolved is None: + raise HTTPException(status_code=404, detail=f"{model.__name__} not found") + field, casted = resolved result = await session.execute(select(model).where(field == casted)) obj = result.scalar_one_or_none() if not obj: diff --git a/api/v1/routes/gifts.py b/api/v1/routes/gifts.py index 5a6aee73..50f17a72 100644 --- a/api/v1/routes/gifts.py +++ b/api/v1/routes/gifts.py @@ -6,6 +6,7 @@ from api.depends import get_session, verify_admin_token from api.v1.routes.base_crud import generate_crud_router from api.v1.schemas import GiftBase, GiftResponse, GiftUpdate, GiftUsageResponse from database.models import Admin, Gift, GiftUsage +from database.access.resolution import resolve_user_optional router = APIRouter() @@ -29,7 +30,10 @@ async def get_gifts_by_tg_id( admin: Admin = Depends(verify_admin_token), session: AsyncSession = Depends(get_session), ): - result = await session.execute(select(Gift).where(Gift.sender_tg_id == tg_id)) + u = await resolve_user_optional(session, tg_id) + if u is None: + raise HTTPException(status_code=404, detail="Gifts not found") + result = await session.execute(select(Gift).where(Gift.sender_user_id == u.id)) gifts = result.scalars().all() if not gifts: raise HTTPException(status_code=404, detail="Gifts not found") diff --git a/api/v1/routes/keys.py b/api/v1/routes/keys.py index befe0160..509486ea 100644 --- a/api/v1/routes/keys.py +++ b/api/v1/routes/keys.py @@ -8,7 +8,8 @@ from api.depends import get_session, verify_admin_token from api.v1.routes.base_crud import generate_crud_router from api.v1.schemas.keys import KeyBase, KeyCreateRequest, KeyResponse, KeyUpdate from database.models import Admin, Key, Tariff -from handlers.keys.operations import create_key_on_cluster, delete_key_from_cluster, renew_key_in_cluster +from database.access.resolution import resolve_user_optional +from services.operations import create_key_on_cluster, delete_key_from_cluster, renew_key_in_cluster from logger import logger @@ -63,7 +64,10 @@ async def get_router_keys_by_tg_id( if not tariff_ids: return [] - keys_result = await session.execute(select(Key).where(Key.tg_id == tg_id, Key.tariff_id.in_(tariff_ids))) + u = await resolve_user_optional(session, tg_id) + if u is None: + return [] + keys_result = await session.execute(select(Key).where(Key.user_id == u.id, Key.tariff_id.in_(tariff_ids))) keys = keys_result.scalars().all() return keys diff --git a/api/v1/routes/management.py b/api/v1/routes/management.py index 8b7db907..424389e0 100644 --- a/api/v1/routes/management.py +++ b/api/v1/routes/management.py @@ -208,7 +208,7 @@ async def restore_trials( update(User) .where( User.trial == 1, - ~exists(select(Key.tg_id).where(Key.tg_id == User.tg_id)), + ~exists(select(Key.user_id).where(Key.user_id == User.id)), ) .values(trial=0) ) diff --git a/api/v1/routes/misc.py b/api/v1/routes/misc.py index 7ee133b7..924746f5 100644 --- a/api/v1/routes/misc.py +++ b/api/v1/routes/misc.py @@ -13,6 +13,7 @@ from api.v1.schemas import ( TrackingSourceResponse, ) from database import get_tracking_source_stats +from database.access.resolution import resolve_user_optional from database.models import ( Admin, BlockedUser, @@ -47,7 +48,10 @@ async def get_payments_by_tg_id( admin: Admin = Depends(verify_admin_token), session: AsyncSession = Depends(get_session), ): - result = await session.execute(select(Payment).where(Payment.tg_id == tg_id)) + u = await resolve_user_optional(session, tg_id) + if u is None: + raise HTTPException(status_code=404, detail="Payments not found") + result = await session.execute(select(Payment).where(Payment.user_id == u.id)) payments = result.scalars().all() if not payments: raise HTTPException(status_code=404, detail="Payments not found") @@ -60,7 +64,8 @@ router.include_router( schema_response=NotificationResponse, schema_create=None, schema_update=None, - identifier_field="tg_id", + identifier_field="user_id", + telegram_path_to_user_id=True, enabled_methods=["get_all", "get_one", "delete"], ), prefix="/notifications", @@ -75,7 +80,9 @@ router.include_router( schema_response=ManualBanResponse, schema_create=None, schema_update=None, - identifier_field="tg_id", + identifier_field="user_id", + parameter_name="tg_id", + telegram_path_to_user_id=True, enabled_methods=["get_all", "get_one", "delete"], ), prefix="/manual-bans", @@ -89,7 +96,9 @@ router.include_router( schema_response=BlockedUserResponse, schema_create=None, schema_update=None, - identifier_field="tg_id", + identifier_field="user_id", + parameter_name="tg_id", + telegram_path_to_user_id=True, enabled_methods=["get_all", "get_one", "delete"], ), prefix="/blocked-users", @@ -103,7 +112,9 @@ router.include_router( schema_response=TemporaryDataResponse, schema_create=None, schema_update=None, - identifier_field="tg_id", + identifier_field="user_id", + parameter_name="tg_id", + telegram_path_to_user_id=True, enabled_methods=["get_all", "get_one", "delete"], ), prefix="/temporary-data", diff --git a/api/v1/routes/referrals.py b/api/v1/routes/referrals.py index 4dac1115..f420ca9a 100644 --- a/api/v1/routes/referrals.py +++ b/api/v1/routes/referrals.py @@ -6,6 +6,7 @@ from api.depends import get_session, verify_admin_token from api.v1.routes.base_crud import generate_crud_router from api.v1.schemas import ReferralResponse from database.models import Admin, Referral +from database.access.resolution import resolve_user_optional router = generate_crud_router( @@ -13,7 +14,9 @@ router = generate_crud_router( schema_response=ReferralResponse, schema_create=None, schema_update=None, - identifier_field="referrer_tg_id", + identifier_field="referrer_user_id", + parameter_name="referrer_tg_id", + telegram_path_to_user_id=True, enabled_methods=["get_all", "get_one", "get_all_by_field"], ) @@ -25,8 +28,15 @@ async def delete_one_referral( admin: Admin = Depends(verify_admin_token), session: AsyncSession = Depends(get_session), ): + ru_ref = await resolve_user_optional(session, referrer_tg_id) + rd_ref = await resolve_user_optional(session, referred_tg_id) + if ru_ref is None or rd_ref is None: + raise HTTPException(status_code=404, detail="Referral not found") result = await session.execute( - select(Referral).where(Referral.referrer_tg_id == referrer_tg_id, Referral.referred_tg_id == referred_tg_id) + select(Referral).where( + Referral.referrer_user_id == ru_ref.id, + Referral.referred_user_id == rd_ref.id, + ) ) obj = result.scalar_one_or_none() if not obj: diff --git a/api/v1/routes/users.py b/api/v1/routes/users.py index 9bdb23b7..63969db5 100644 --- a/api/v1/routes/users.py +++ b/api/v1/routes/users.py @@ -9,7 +9,8 @@ from api.v1.routes.base_crud import generate_crud_router from api.v1.schemas.users import UserBase, UserResponse, UserUpdate from database import async_session_maker, delete_user_data, get_servers from database.models import Key, User -from handlers.keys.operations import delete_key_from_cluster +from database.access.resolution import resolve_user_optional +from services.operations import delete_key_from_cluster from logger import logger @@ -30,7 +31,10 @@ async def delete_user( session: AsyncSession = Depends(get_session), ): try: - result = await session.execute(select(Key.email, Key.client_id).where(Key.tg_id == tg_id)) + u = await resolve_user_optional(session, tg_id) + if u is None: + raise HTTPException(status_code=404, detail="Пользователь не найден") + result = await session.execute(select(Key.email, Key.client_id).where(Key.user_id == u.id)) key_records = result.all() async with async_session_maker() as s: diff --git a/api/v1/schemas/gifts.py b/api/v1/schemas/gifts.py index b6260c2b..b503d759 100644 --- a/api/v1/schemas/gifts.py +++ b/api/v1/schemas/gifts.py @@ -4,11 +4,13 @@ from pydantic import BaseModel class GiftBase(BaseModel): - sender_tg_id: int - recipient_tg_id: int | None = None + sender_user_id: int + recipient_user_id: int | None = None selected_months: int | None = None expiry_time: datetime gift_link: str + telegram_gift_link: str | None = None + site_gift_link: str | None = None is_used: bool = False is_unlimited: bool | None = False max_usages: int | None = None @@ -33,10 +35,12 @@ class GiftUsageResponse(BaseModel): class GiftUpdate(BaseModel): - recipient_tg_id: int | None = None + recipient_user_id: int | None = None selected_months: int | None = None expiry_time: datetime | None = None gift_link: str | None = None + telegram_gift_link: str | None = None + site_gift_link: str | None = None is_used: bool | None = None is_unlimited: bool | None = None max_usages: int | None = None diff --git a/api/v1/schemas/keys.py b/api/v1/schemas/keys.py index caae97bb..5c69d0d7 100644 --- a/api/v1/schemas/keys.py +++ b/api/v1/schemas/keys.py @@ -2,8 +2,9 @@ from pydantic import BaseModel, Field class KeyBase(BaseModel): - tg_id: int + user_id: int client_id: str + tg_id: int | None = None email: str | None = None created_at: int | None = None expiry_time: int diff --git a/api/v1/schemas/misc.py b/api/v1/schemas/misc.py index 467fd364..caef8334 100644 --- a/api/v1/schemas/misc.py +++ b/api/v1/schemas/misc.py @@ -4,7 +4,8 @@ from pydantic import BaseModel class PaymentBase(BaseModel): - tg_id: int + user_id: int + tg_id: int | None = None amount: float payment_system: str status: str @@ -19,8 +20,8 @@ class PaymentResponse(PaymentBase): class ReferralResponse(BaseModel): - referred_tg_id: int - referrer_tg_id: int + referred_user_id: int + referrer_user_id: int reward_issued: bool = False class Config: @@ -37,11 +38,13 @@ class NotificationResponse(BaseModel): class GiftBase(BaseModel): - sender_tg_id: int - recipient_tg_id: int | None = None + sender_user_id: int + recipient_user_id: int | None = None selected_months: int expiry_time: datetime gift_link: str + telegram_gift_link: str | None = None + site_gift_link: str | None = None is_used: bool = False is_unlimited: bool = False max_usages: int | None = None @@ -66,7 +69,8 @@ class GiftUsageResponse(BaseModel): class ManualBanResponse(BaseModel): - tg_id: int + user_id: int + tg_id: int | None = None banned_at: datetime reason: str banned_by: int @@ -77,7 +81,8 @@ class ManualBanResponse(BaseModel): class TemporaryDataResponse(BaseModel): - tg_id: int + user_id: int + tg_id: int | None = None state: str data: dict updated_at: datetime @@ -87,7 +92,8 @@ class TemporaryDataResponse(BaseModel): class BlockedUserResponse(BaseModel): - tg_id: int + user_id: int + tg_id: int | None = None class Config: from_attributes = True diff --git a/api/v1/schemas/referrals.py b/api/v1/schemas/referrals.py index 39c60d42..dab27b20 100644 --- a/api/v1/schemas/referrals.py +++ b/api/v1/schemas/referrals.py @@ -2,8 +2,8 @@ from pydantic import BaseModel class ReferralResponse(BaseModel): - referred_tg_id: int - referrer_tg_id: int + referred_user_id: int + referrer_user_id: int reward_issued: bool = False class Config: diff --git a/api/v2/__init__.py b/api/v2/__init__.py index 0e439d72..2cfb418c 100644 --- a/api/v2/__init__.py +++ b/api/v2/__init__.py @@ -1,5 +1,9 @@ -from api.v2.router import router - VERSION = "2.0.0" - __all__ = ("router", "VERSION") + + +def __getattr__(name: str): + if name == "router": + from api.v2.router import router + return router + raise AttributeError(name) diff --git a/api/v2/base_crud.py b/api/v2/base_crud.py index c98cda8d..e12d3ac5 100644 --- a/api/v2/base_crud.py +++ b/api/v2/base_crud.py @@ -6,7 +6,13 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm.attributes import InstrumentedAttribute from api.depends import get_session, verify_identity_admin -from api.v1.routes.base_crud import cast_identifier_type, normalize_outgoing_object, to_schema +from api.v1.routes.base_crud import ( + _apply_user_relationship_loader, + cast_identifier_type, + normalize_outgoing_object, + to_schema, +) +from database.access.resolution import resolve_user_optional def generate_crud_router( @@ -18,10 +24,20 @@ def generate_crud_router( identifier_field: str = "tg_id", parameter_name: str = "tg_id", extra_get_by_email: bool = False, + telegram_path_to_user_id: bool = False, enabled_methods: list[str] = ("get_all", "get_one", "get_by_email", "create", "update", "delete"), ) -> APIRouter: router = APIRouter() + async def _path_filter(session: AsyncSession, value: int | str): + if telegram_path_to_user_id: + u = await resolve_user_optional(session, int(value)) + if u is None: + return None + return getattr(model, "user_id"), u.id + field = getattr(model, identifier_field) + return field, cast_identifier_type(field, value) + if "get_all" in enabled_methods: @router.get("/", response_model=list[schema_response]) @@ -29,7 +45,7 @@ def generate_crud_router( identity=Depends(verify_identity_admin), session: AsyncSession = Depends(get_session), ): - result = await session.execute(select(model)) + result = await session.execute(_apply_user_relationship_loader(model, select(model))) items = result.scalars().all() for item in items: normalize_outgoing_object(item) @@ -57,9 +73,13 @@ def generate_crud_router( identity=Depends(verify_identity_admin), session: AsyncSession = Depends(get_session), ): - field = getattr(model, identifier_field) - casted = cast_identifier_type(field, value) - result = await session.execute(select(model).where(field == casted)) + resolved = await _path_filter(session, value) + if resolved is None: + raise HTTPException(status_code=404, detail=f"{model.__name__} not found") + field, casted = resolved + result = await session.execute( + _apply_user_relationship_loader(model, select(model).where(field == casted)) + ) obj = result.scalar_one_or_none() if not obj: raise HTTPException(status_code=404, detail=f"{model.__name__} not found") @@ -73,9 +93,13 @@ def generate_crud_router( identity=Depends(verify_identity_admin), session: AsyncSession = Depends(get_session), ): - field = getattr(model, identifier_field) - casted = cast_identifier_type(field, value) - result = await session.execute(select(model).where(field == casted)) + resolved = await _path_filter(session, value) + if resolved is None: + raise HTTPException(status_code=404, detail=f"{model.__name__} not found") + field, casted = resolved + result = await session.execute( + _apply_user_relationship_loader(model, select(model).where(field == casted)) + ) objs = result.scalars().all() if not objs: raise HTTPException(status_code=404, detail=f"{model.__name__} not found") @@ -110,8 +134,10 @@ def generate_crud_router( identity=Depends(verify_identity_admin), session: AsyncSession = Depends(get_session), ): - field = getattr(model, identifier_field) - casted = cast_identifier_type(field, value) + resolved = await _path_filter(session, value) + if resolved is None: + raise HTTPException(status_code=404, detail=f"{model.__name__} not found") + field, casted = resolved result = await session.execute(select(model).where(field == casted)) obj = result.scalar_one_or_none() if not obj: @@ -131,8 +157,10 @@ def generate_crud_router( identity=Depends(verify_identity_admin), session: AsyncSession = Depends(get_session), ): - field = getattr(model, identifier_field) - casted = cast_identifier_type(field, value) + resolved = await _path_filter(session, value) + if resolved is None: + raise HTTPException(status_code=404, detail=f"{model.__name__} not found") + field, casted = resolved result = await session.execute(select(model).where(field == casted)) obj = result.scalar_one_or_none() if not obj: diff --git a/api/v2/router.py b/api/v2/router.py index 2134506e..ad4932c0 100644 --- a/api/v2/router.py +++ b/api/v2/router.py @@ -18,6 +18,8 @@ from api.v2.routes import ( payment_links, identities, web, + flows, + notifications, ) router = APIRouter() @@ -25,10 +27,12 @@ router = APIRouter() router.include_router(root_router) router.include_router(auth.router, prefix="/api") router.include_router(users.router, prefix="/api/users", tags=["Users"]) -router.include_router(keys.router, prefix="/api/keys", tags=["Keys"]) +router.include_router(keys.user_router, prefix="/api/keys", tags=["Keys"]) +router.include_router(keys.router, prefix="/api/admin/keys", tags=["AdminKeys"]) router.include_router(coupons.router, prefix="/api/coupons", tags=["Coupons"]) router.include_router(servers.router, prefix="/api/servers", tags=["Servers"]) router.include_router(tariffs.public_router, prefix="/api/tariffs", tags=["Tariffs"]) +router.include_router(tariffs.user_tariff_router, prefix="/api/tariffs", tags=["Tariffs"]) router.include_router(tariffs.router, prefix="/api/tariffs", tags=["Tariffs"]) router.include_router(gifts.router, prefix="/api/gifts", tags=["Gifts"]) router.include_router(referrals.router, prefix="/api/referrals", tags=["Referrals"]) @@ -39,4 +43,6 @@ router.include_router(misc.router, prefix="/api") router.include_router(modules.router, prefix="/api") router.include_router(management.router, prefix="/api/management", tags=["Management"]) router.include_router(settings.router, prefix="/api/settings", tags=["Settings"]) -router.include_router(web.router, prefix="", tags=["Web"]) \ No newline at end of file +router.include_router(web.router, prefix="", tags=["Web"]) +router.include_router(flows.router, prefix="/api", tags=["Flows"]) +router.include_router(notifications.router, prefix="/api", tags=["Notifications"]) diff --git a/api/v2/routes/__init__.py b/api/v2/routes/__init__.py index 31b027c0..e7bad74b 100644 --- a/api/v2/routes/__init__.py +++ b/api/v2/routes/__init__.py @@ -1,3 +1,9 @@ -from api.v2.routes.root import router as root_router - __all__ = ("root_router",) + + +def __getattr__(name: str): + if name == "root_router": + from api.v2.routes.root import router + + return router + raise AttributeError(name) diff --git a/api/v2/routes/auth.py b/api/v2/routes/auth.py deleted file mode 100644 index c9105120..00000000 --- a/api/v2/routes/auth.py +++ /dev/null @@ -1,173 +0,0 @@ -from fastapi import APIRouter, Depends, HTTPException, Request -from sqlalchemy.ext.asyncio import AsyncSession - -from audit import set_api_actor -from api.depends import get_session, verify_identity_token -from api.v2.schemas.identities import ( - IdentityResponse, - LinkTelegramRequest, - LoginByCodeRequest, - LoginRequest, - LoginResponse, - LoginTelegramRequest, - RegisterByEmailRequest, - RegisterResponse, - SendLoginCodeRequest, -) -from config import API_TOKEN_TTL_DAYS, API_TOKEN -from database import identities as idb -from utils.telegram_login import verify_telegram_login - -router = APIRouter(prefix="/auth", tags=["Auth"]) -TOKEN_TTL_HINT = "бессрочно" if API_TOKEN_TTL_DAYS is None else f"{API_TOKEN_TTL_DAYS} дн." -TELEGRAM_LOGIN_MAX_AGE = 86400 - - -@router.post("/register", response_model=RegisterResponse) -async def register_by_email( - body: RegisterByEmailRequest, - request: Request, - session: AsyncSession = Depends(get_session), -): - ( - """Регистрация по почте и паролю: создаётся идентичность, выдаётся токен. Срок действия токена: """ - + TOKEN_TTL_HINT - + "." - ) - email = body.email.strip().lower() - if not email: - raise HTTPException(status_code=400, detail="Email обязателен") - if not body.password or len(body.password) < 8: - raise HTTPException(status_code=400, detail="Пароль минимум 8 символов") - existing = await idb.get_identity_by_email(session, email) - if existing: - raise HTTPException(status_code=409, detail="Идентичность с таким email уже существует") - identity, token = await idb.create_identity_with_token(session, email=email, password=body.password) - set_api_actor(request, identity_id=identity.id, tg_id=identity.tg_id) - return RegisterResponse(identity_id=identity.id, token=token) - - -@router.post("/login", response_model=LoginResponse) -async def login( - body: LoginRequest, - request: Request, - session: AsyncSession = Depends(get_session), -): - """Вход по email и паролю. Возвращает identity_id и новый токен. Срок действия токена: """ + TOKEN_TTL_HINT + "." - email = body.email.strip().lower() - if not email: - raise HTTPException(status_code=400, detail="Email обязателен") - result = await idb.login_by_email(session, email, body.password) - if not result: - raise HTTPException(status_code=401, detail="Неверный email или пароль") - identity, token = result - set_api_actor(request, identity_id=identity.id, tg_id=identity.tg_id) - return LoginResponse(identity_id=identity.id, token=token) - - -_LOGIN_CODES: dict[str, tuple[str, float]] = {} -_LOGIN_CODE_TTL = 600.0 - - -def _clean_login_codes() -> None: - import time - now = time.time() - for k in list(_LOGIN_CODES): - if now - _LOGIN_CODES[k][1] > _LOGIN_CODE_TTL: - del _LOGIN_CODES[k] - - -@router.post("/send-login-code") -async def send_login_code( - body: SendLoginCodeRequest, - request: Request, - session: AsyncSession = Depends(get_session), -): - """Отправить код входа на email. Код хранится на сервере 10 мин (для демо — без реальной отправки письма).""" - _clean_login_codes() - email = body.email.strip().lower() - if not email: - raise HTTPException(status_code=400, detail="Email обязателен") - identity = await idb.get_identity_by_email(session, email) - if not identity: - raise HTTPException(status_code=404, detail="Аккаунт с таким email не найден") - import secrets - import time - code = "".join(secrets.choice("0123456789") for _ in range(6)) - _LOGIN_CODES[email] = (code, time.time()) - return {"ok": True, "message": "Код отправлен на почту"} - - -@router.post("/login-by-code", response_model=LoginResponse) -async def login_by_code( - body: LoginByCodeRequest, - request: Request, - session: AsyncSession = Depends(get_session), -): - """Вход по email и коду из письма.""" - _clean_login_codes() - email = body.email.strip().lower() - if not email or not body.code or not body.code.strip(): - raise HTTPException(status_code=400, detail="Email и код обязательны") - stored = _LOGIN_CODES.get(email) - if not stored: - raise HTTPException(status_code=400, detail="Код не найден или истёк. Запросите новый.") - code_value, _ = stored - if body.code.strip() != code_value: - raise HTTPException(status_code=401, detail="Неверный код") - del _LOGIN_CODES[email] - identity = await idb.get_identity_by_email(session, email) - if not identity: - raise HTTPException(status_code=401, detail="Аккаунт не найден") - token = await idb.issue_token_for_identity(session, identity) - set_api_actor(request, identity_id=identity.id, tg_id=identity.tg_id) - return LoginResponse(identity_id=identity.id, token=token) - - -@router.post("/login-telegram", response_model=LoginResponse) -async def login_telegram( - body: LoginTelegramRequest, - request: Request, - session: AsyncSession = Depends(get_session), -): - ( - """Вход через Telegram Login Widget (кнопка на сайте). По tg_id находим или создаём Identity, выдаём токен. Срок действия токена: """ - + TOKEN_TTL_HINT - + "." - ) - payload = body.model_dump(mode="json") - if not verify_telegram_login(payload, API_TOKEN, max_age_seconds=TELEGRAM_LOGIN_MAX_AGE): - raise HTTPException(status_code=401, detail="Неверная подпись или устаревшие данные от Telegram") - identity = await idb.get_or_create_identity_for_tg(session, body.id) - token = await idb.issue_token_for_identity(session, identity) - set_api_actor(request, identity_id=identity.id, tg_id=identity.tg_id) - return LoginResponse(identity_id=identity.id, token=token) - - -@router.post("/link-telegram", response_model=IdentityResponse) -async def link_telegram( - body: LinkTelegramRequest, - request: Request, - session: AsyncSession = Depends(get_session), - identity=Depends(verify_identity_token), -): - """Привязывает Telegram к текущей идентичности. Требуется подпись от Telegram Login Widget (доказательство владения аккаунтом).""" - payload = body.model_dump(mode="json") - if not verify_telegram_login(payload, API_TOKEN, max_age_seconds=TELEGRAM_LOGIN_MAX_AGE): - raise HTTPException(status_code=401, detail="Неверная подпись или устаревшие данные от Telegram") - result = await idb.attach_telegram(session, identity.id, body.id) - if not result: - raise HTTPException( - status_code=409, - detail="Этот Telegram уже привязан к другой идентичности", - ) - set_api_actor(request, identity_id=result.id, tg_id=result.tg_id) - return IdentityResponse.model_validate(result) - - -@router.get("/me", response_model=IdentityResponse) -async def me( - identity=Depends(verify_identity_token), -): - """Текущая идентичность по заголовкам X-Identity-Id и X-Token.""" - return IdentityResponse.model_validate(identity) diff --git a/api/v2/routes/auth/__init__.py b/api/v2/routes/auth/__init__.py new file mode 100644 index 00000000..4886dcde --- /dev/null +++ b/api/v2/routes/auth/__init__.py @@ -0,0 +1,13 @@ +from fastapi import APIRouter + +from api.v2.routes.auth import email_verify, link, password, session, telegram + + +router = APIRouter(prefix="/auth", tags=["Auth"]) +router.include_router(password.router) +router.include_router(telegram.router) +router.include_router(link.router) +router.include_router(email_verify.router) +router.include_router(session.router) + +__all__ = ["router"] diff --git a/api/v2/routes/auth/_common.py b/api/v2/routes/auth/_common.py new file mode 100644 index 00000000..9ba06c8e --- /dev/null +++ b/api/v2/routes/auth/_common.py @@ -0,0 +1,112 @@ +from fastapi import Request +from sqlalchemy import text +from sqlalchemy.ext.asyncio import AsyncSession + +from config import API_TOKEN_TTL_DAYS +from logger import logger +from utils.referral_codes import encode_partner_code + + +TOKEN_TTL_HINT = "бессрочно" if API_TOKEN_TTL_DAYS is None else f"{API_TOKEN_TTL_DAYS} дн." +TELEGRAM_LOGIN_MAX_AGE = 86400 + +_TRUSTED_PROXY_CIDRS: list[str] = [] + + +def _client_ip(request: Request) -> str: + client_host = (request.client.host if request.client else "") or "" + forwarded = request.headers.get("x-forwarded-for") or request.headers.get("X-Forwarded-For") + if not forwarded: + return client_host + if not _TRUSTED_PROXY_CIDRS and client_host not in ("127.0.0.1", "::1"): + return client_host + return forwarded.split(",")[0].strip() or client_host + + +async def _resolve_partner_snapshot(session: AsyncSession, billing_user_id: int) -> dict[str, object]: + partner_feature_enabled = False + default_percent = 0.0 + try: + from modules.partner_program import settings as partner_settings + + partner_feature_enabled = True + raw_percent = getattr(partner_settings, "PARTNER_BONUS_PERCENTAGES", {}).get(1, 0.0) + default_percent = float(raw_percent) * 100.0 + except Exception: + partner_feature_enabled = False + default_percent = 0.0 + payload: dict[str, object] = { + "partner_enabled": partner_feature_enabled, + "partner_code": "", + "partner_balance": 0.0, + "partner_percent": default_percent, + "partner_percent_custom": False, + "partner_referred_total": 0, + "partner_payout_method": None, + } + try: + partner_row = ( + await session.execute( + text( + """ + SELECT + tg_id, + COALESCE(partner_balance, 0), + partner_percent, + COALESCE(partner_percent_custom, false), + partner_code, + payout_method + FROM users + WHERE id = :user_id + LIMIT 1 + """ + ), + {"user_id": int(billing_user_id)}, + ) + ).first() + except Exception: + partner_row = None + if partner_row is None: + return payload + tg_id = int(partner_row[0]) if partner_row[0] is not None else None + balance = float(partner_row[1] or 0.0) + percent_raw = partner_row[2] + percent_custom = bool(partner_row[3]) + percent_value = float(percent_raw) if (percent_custom and percent_raw is not None) else float(default_percent) + code = str(partner_row[4] or "").strip() + if (not code or code.isdigit() or code.startswith("r1_")) and int(billing_user_id) > 0: + generated_code = encode_partner_code(int(billing_user_id)) + code = generated_code + try: + await session.execute( + text("UPDATE users SET partner_code = :code WHERE id = :id"), + {"code": generated_code, "id": int(billing_user_id)}, + ) + await session.flush() + except Exception as e: + logger.warning("[Auth] Ошибка сохранения partner_code для billing_user_id={}: {}", billing_user_id, e) + payout_method = str(partner_row[5] or "").strip() or None + referred_total = 0 + if tg_id is not None: + try: + referred_total = int( + ( + await session.execute( + text("SELECT COUNT(*) FROM partners WHERE partner_tg_id = :tg_id"), + {"tg_id": int(tg_id)}, + ) + ).scalar() + or 0 + ) + except Exception: + referred_total = 0 + payload.update({ + "partner_enabled": bool(partner_feature_enabled or code or referred_total > 0 or balance > 0), + "partner_code": code, + "partner_balance": balance, + "partner_percent": percent_value, + "partner_percent_custom": percent_custom, + "partner_referred_total": referred_total, + "partner_payout_method": payout_method, + }) + return payload diff --git a/api/v2/routes/auth/_fallback_limiter.py b/api/v2/routes/auth/_fallback_limiter.py new file mode 100644 index 00000000..f19e8c53 --- /dev/null +++ b/api/v2/routes/auth/_fallback_limiter.py @@ -0,0 +1,43 @@ +"""In-memory rate limit fallback для случаев когда Redis недоступен. + +Используется только когда Redis не ответил — чтобы критичные auth-эндпоинты +не теряли защиту при кратковременных Redis-сбоях. Per-process, не шарится +между репликами — поэтому на нескольких инстансах лимит будет N*limit. +""" + +import time +from collections import deque +from threading import Lock + + +_BUCKETS: dict[str, deque[float]] = {} +_LOCK = Lock() +_MAX_KEYS = 10000 + + +def _prune(bucket: deque[float], window_sec: int) -> None: + threshold = time.monotonic() - window_sec + while bucket and bucket[0] < threshold: + bucket.popleft() + + +def _evict_if_full() -> None: + if len(_BUCKETS) < _MAX_KEYS: + return + now = time.monotonic() + dead = [k for k, b in _BUCKETS.items() if not b or b[-1] < now - 3600] + for k in dead: + del _BUCKETS[k] + if len(_BUCKETS) >= _MAX_KEYS: + oldest = min(_BUCKETS.keys(), key=lambda k: _BUCKETS[k][0] if _BUCKETS[k] else 0) + del _BUCKETS[oldest] + + +def check_and_increment(key: str, limit: int, window_sec: int) -> int: + """Возвращает текущее значение счётчика после инкремента. Если >= limit — превышение.""" + with _LOCK: + _evict_if_full() + bucket = _BUCKETS.setdefault(key, deque()) + _prune(bucket, window_sec) + bucket.append(time.monotonic()) + return len(bucket) diff --git a/api/v2/routes/auth/email_verify.py b/api/v2/routes/auth/email_verify.py new file mode 100644 index 00000000..26c22c8c --- /dev/null +++ b/api/v2/routes/auth/email_verify.py @@ -0,0 +1,79 @@ +import secrets + +from fastapi import APIRouter, Depends, HTTPException, Request +from pydantic import ( + BaseModel, + Field as PydanticField, +) +from sqlalchemy.ext.asyncio import AsyncSession + +from api.depends import get_session, verify_identity_token +from api.v2.routes.auth._common import _client_ip +from mail import send_email_verify_code_email, smtp_configured +from utils import web_email_verify_code as verify_util + + +router = APIRouter() + + +class VerifyEmailRequest(BaseModel): + code: str = PydanticField(..., min_length=1, max_length=10) + + +@router.post("/send-verify-code") +async def send_email_verify_code( + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + """Отправить код подтверждения email. Требует авторизации.""" + if not smtp_configured(): + raise HTTPException(status_code=503, detail="Почтовый сервер не настроен") + email = (identity.email or "").strip().lower() + if not email: + raise HTTPException(status_code=400, detail="Email не привязан к аккаунту") + if getattr(identity, "email_verified", False): + return {"ok": True, "detail": "Email уже подтверждён"} + if not await verify_util.redis_ready(): + raise HTTPException(status_code=503, detail="Сервис временно недоступен") + ip = _client_ip(request) + if not await verify_util.try_consume_ip_send_budget(ip): + raise HTTPException(status_code=429, detail="Слишком много запросов, попробуйте позже") + if not await verify_util.try_consume_email_send_budget(email): + raise HTTPException(status_code=429, detail="Слишком много запросов на этот email") + if not await verify_util.try_acquire_resend_cooldown(email): + raise HTTPException(status_code=429, detail="Подождите минуту перед повторной отправкой") + code = f"{secrets.randbelow(900000) + 100000}" + await verify_util.store_code(email, code) + try: + await send_email_verify_code_email(email, code) + except Exception: + await verify_util.delete_code(email) + raise HTTPException(status_code=503, detail="Не удалось отправить письмо") + return {"ok": True} + + +@router.post("/verify-email") +async def verify_email( + body: VerifyEmailRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + """Подтвердить email по коду.""" + email = (identity.email or "").strip().lower() + if not email: + raise HTTPException(status_code=400, detail="Email не привязан к аккаунту") + if getattr(identity, "email_verified", False): + return {"ok": True, "detail": "Email уже подтверждён"} + if not await verify_util.try_consume_verify_budget(email): + raise HTTPException(status_code=429, detail="Слишком много попыток, попробуйте позже") + if not await verify_util.verify_and_consume_code(email, body.code.strip()): + raise HTTPException(status_code=400, detail="Неверный или просроченный код") + from sqlalchemy import update + + from database.models import Identity as IdentityModel + await session.execute( + update(IdentityModel).where(IdentityModel.id == identity.id).values(email_verified=True) + ) + return {"ok": True} diff --git a/api/v2/routes/auth/link.py b/api/v2/routes/auth/link.py new file mode 100644 index 00000000..00a87e72 --- /dev/null +++ b/api/v2/routes/auth/link.py @@ -0,0 +1,114 @@ +import secrets + +from fastapi import APIRouter, Depends, HTTPException, Request +from sqlalchemy.ext.asyncio import AsyncSession + +from api.depends import ( + bind_identity_actor, + get_session, + verify_identity_token, +) +from api.v2.routes.auth._common import _client_ip +from api.v2.schemas.identities import ( + IdentityResponse, + LinkEmailConfirmRequest, + LinkEmailSendCodeRequest, +) +from database import identities as idb +from mail import send_email_link_code_email, smtp_configured +from utils import web_email_link_code as email_link_code + + +router = APIRouter() + + +@router.post("/link-email/send-code") +async def link_email_send_code( + body: LinkEmailSendCodeRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + email_norm = email_link_code.normalize_email(body.email) + if not email_norm: + raise HTTPException(status_code=400, detail="Укажите корректный email") + if identity.email and str(identity.email).strip().lower() == email_norm: + raise HTTPException(status_code=409, detail="Этот email уже привязан к аккаунту") + if not smtp_configured(): + raise HTTPException( + status_code=503, + detail="Отправка кода недоступна: почта не настроена на сервере", + ) + if not await email_link_code.redis_ready(): + raise HTTPException( + status_code=503, + detail="Сервис временно недоступен. Попробуйте позже.", + ) + existing = await idb.get_identity_by_email(session, email_norm) + if existing and existing.id != identity.id: + raise HTTPException(status_code=400, detail="Не удалось привязать email") + ip = _client_ip(request) + if not await email_link_code.try_consume_ip_budget(ip): + raise HTTPException( + status_code=429, + detail="Слишком много запросов с вашего адреса. Попробуйте позже.", + ) + if not await email_link_code.try_consume_email_send_budget(email_norm): + raise HTTPException( + status_code=429, + detail="Слишком много запросов для этого адреса. Попробуйте позже.", + ) + if not await email_link_code.try_acquire_cooldown(email_norm): + raise HTTPException( + status_code=429, + detail="Код уже отправлен. Подождите перед повторной отправкой.", + ) + code = "".join(secrets.choice("0123456789") for _ in range(6)) + if not await email_link_code.store_code(email_norm, code): + await email_link_code.release_cooldown(email_norm) + raise HTTPException( + status_code=503, + detail="Не удалось сохранить код. Попробуйте позже.", + ) + try: + await send_email_link_code_email(email_norm, code) + except Exception: + await email_link_code.release_cooldown(email_norm) + await email_link_code.delete_code(email_norm) + raise HTTPException( + status_code=503, + detail="Не удалось отправить письмо. Попробуйте позже.", + ) from None + return {"ok": True, "message": "Код подтверждения отправлен на почту"} + + +@router.post("/link-email/confirm", response_model=IdentityResponse) +async def link_email_confirm( + body: LinkEmailConfirmRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + email_norm = email_link_code.normalize_email(body.email) + if not email_norm or not body.code or not str(body.code).strip(): + raise HTTPException(status_code=400, detail="Email и код обязательны") + if not await email_link_code.redis_ready(): + raise HTTPException( + status_code=503, + detail="Сервис временно недоступен. Попробуйте позже.", + ) + if not await email_link_code.try_consume_email_verify_budget(email_norm): + raise HTTPException( + status_code=429, + detail="Слишком много попыток. Запросите новый код.", + ) + if not await email_link_code.verify_and_consume_code(email_norm, str(body.code).strip()): + raise HTTPException(status_code=401, detail="Неверный код или срок действия истёк") + result = await idb.attach_email(session, identity.id, email_norm) + if not result: + raise HTTPException( + status_code=409, + detail="Этот email уже привязан к другой идентичности", + ) + await bind_identity_actor(request, session, result) + return IdentityResponse.model_validate(result) diff --git a/api/v2/routes/auth/password.py b/api/v2/routes/auth/password.py new file mode 100644 index 00000000..da096e36 --- /dev/null +++ b/api/v2/routes/auth/password.py @@ -0,0 +1,389 @@ +import secrets + +from fastapi import APIRouter, Depends, HTTPException, Request, Response +from sqlalchemy.ext.asyncio import AsyncSession + +from api.depends import ( + bind_identity_actor, + get_session, + set_auth_cookie, + set_is_admin_cookie, +) +from api.v2.routes.auth._common import TOKEN_TTL_HINT, _client_ip +from api.v2.schemas.identities import ( + ConfirmPasswordResetRequest, + LoginByCodeRequest, + LoginRequest, + LoginResponse, + RegisterByEmailRequest, + RegisterResponse, + SendLoginCodeRequest, +) +from database import ( + add_referral, + get_referral_by_referred_id, + identities as idb, +) +from database.access.resolution import resolve_user_optional +from logger import logger +from mail import ( + send_email_verify_code_email, + send_login_code_email, + send_password_reset_code_email, + smtp_configured, +) +from utils import ( + web_email_verify_code as email_verify, + web_password_reset_code as pwd_reset, +) +from utils.disposable_emails import is_disposable_email +from utils.referral_codes import decode_referral_code +from utils.turnstile import turnstile_enabled, verify_turnstile_token +from utils.web_login_code import ( + delete_code, + normalize_login_email, + redis_ready_for_login_codes, + release_resend_cooldown, + store_code, + try_acquire_resend_cooldown, + try_consume_email_send_budget, + try_consume_email_verify_budget, + try_consume_ip_send_budget, + verify_and_consume_code, +) + + +router = APIRouter() + +_RESET_OK_MESSAGE = { + "ok": True, + "message": "Если для этого адреса есть аккаунт с паролем, мы отправили код. Проверьте почту.", +} + + +@router.post("/register", response_model=RegisterResponse) +async def register_by_email( + body: RegisterByEmailRequest, + request: Request, + response: Response, + session: AsyncSession = Depends(get_session), +): + ( + """Регистрация по почте и паролю: создаётся идентичность, выдаётся токен. Срок действия токена: """ + + TOKEN_TTL_HINT + + "." + ) + ip = _client_ip(request) + try: + from core.redis_cache import cache_incr_checked + from api.v2.routes.auth._fallback_limiter import check_and_increment + count, redis_ok = await cache_incr_checked(f"register_rate:{ip}", 3600) + if not redis_ok: + count = check_and_increment(f"register_rate:{ip}", 5, 3600) + if count > 5: + raise HTTPException(status_code=429, detail="Слишком много регистраций с этого IP. Попробуйте позже.") + except HTTPException: + raise + except Exception: + pass + if turnstile_enabled(): + if not await verify_turnstile_token(body.turnstile_token, ip): + raise HTTPException(status_code=400, detail="Проверка CAPTCHA не пройдена") + email = body.email.strip().lower() + if not email: + raise HTTPException(status_code=400, detail="Email обязателен") + if is_disposable_email(email): + raise HTTPException(status_code=400, detail="Одноразовые email-адреса не поддерживаются") + if not body.password or len(body.password) < 8: + raise HTTPException(status_code=400, detail="Пароль минимум 8 символов") + existing = await idb.get_identity_by_email(session, email) + if existing: + raise HTTPException(status_code=409, detail="Идентичность с таким email уже существует") + raw_referral = str(body.referral_code or "").strip() + if "/referral/" in raw_referral: + raw_referral = raw_referral.split("/referral/", 1)[-1] + if "start=referral_" in raw_referral: + raw_referral = raw_referral.split("start=referral_", 1)[-1] + raw_referral = raw_referral.split("?", 1)[0].split("#", 1)[0].strip() + referrer_legacy = decode_referral_code(raw_referral) + referrer_user = None + if body.referral_code and referrer_legacy is None: + raise HTTPException(status_code=400, detail="Код приглашения недействителен") + if referrer_legacy is not None: + referrer_user = await resolve_user_optional(session, referrer_legacy) + if referrer_user is None: + raise HTTPException(status_code=400, detail="Код приглашения недействителен") + identity, token = await idb.create_identity_with_token(session, email=email, password=body.password) + await bind_identity_actor(request, session, identity) + billing_user_id = await idb.ensure_billing_user_for_identity(session, identity) + if referrer_user is not None and not await get_referral_by_referred_id(session, billing_user_id): + await add_referral(session, billing_user_id, referrer_user.id) + if smtp_configured(): + try: + code = f"{secrets.randbelow(900000) + 100000}" + await email_verify.store_code(email, code) + await send_email_verify_code_email(email, code) + except Exception as e: + logger.warning("[Auth] Не удалось отправить код подтверждения email при регистрации: {}", e) + logger.info("[Auth] Register success: identity={}, email={}, ip={}", identity.id, email, _client_ip(request)) + set_auth_cookie(response, token, request) + set_is_admin_cookie(response, identity, request) + return RegisterResponse(identity_id=identity.id) + + +@router.post("/login", response_model=LoginResponse) +async def login( + body: LoginRequest, + request: Request, + response: Response, + session: AsyncSession = Depends(get_session), +): + """Вход по email и паролю. Возвращает identity_id и новый токен. Срок действия токена: """ + TOKEN_TTL_HINT + "." + email = body.email.strip().lower() + if not email: + raise HTTPException(status_code=400, detail="Email обязателен") + ip = _client_ip(request) + try: + from core.redis_cache import cache_get, cache_incr_checked + from api.v2.routes.auth._fallback_limiter import check_and_increment + lockout_key = f"login_lockout:{email}" + locked = await cache_get(lockout_key) + if locked: + raise HTTPException(status_code=429, detail="Аккаунт временно заблокирован. Попробуйте через 15 минут.") + rkey = f"login_pwd_rate:{ip}" + count, redis_ok = await cache_incr_checked(rkey, 900) + if not redis_ok: + count = check_and_increment(rkey, 10, 900) + if count > 10: + raise HTTPException(status_code=429, detail="Слишком много попыток. Попробуйте позже.") + except HTTPException: + raise + except Exception as e: + logger.warning("[Auth] Ошибка rate-limit проверки для email-логина: {}", e) + result = await idb.login_by_email(session, email, body.password) + if not result: + try: + from core.redis_cache import cache_incr, cache_set + fail_key = f"login_fail:{email}" + fails = await cache_incr(fail_key, 900) + if fails >= 10: + await cache_set(f"login_lockout:{email}", "1", 900) + except Exception: + pass + raise HTTPException(status_code=401, detail="Неверный email или пароль") + try: + from core.redis_cache import cache_delete + await cache_delete(f"login_fail:{email}") + except Exception: + pass + identity, token = result + await bind_identity_actor(request, session, identity) + logger.info("[Auth] Login success: identity={}, email={}, ip={}, method=password", identity.id, email, ip) + set_auth_cookie(response, token, request) + set_is_admin_cookie(response, identity, request) + return LoginResponse(identity_id=identity.id) + + +@router.post("/send-login-code") +async def send_login_code( + body: SendLoginCodeRequest, + request: Request, + session: AsyncSession = Depends(get_session), +): + """Отправить код входа на email (SMTP + Redis).""" + ip = _client_ip(request) + try: + from core.redis_cache import cache_incr_checked + from api.v2.routes.auth._fallback_limiter import check_and_increment + count, redis_ok = await cache_incr_checked(f"send_code_rate:{ip}", 3600) + if not redis_ok: + count = check_and_increment(f"send_code_rate:{ip}", 10, 3600) + if count > 10: + raise HTTPException(status_code=429, detail="Слишком много запросов кодов. Попробуйте позже.") + except HTTPException: + raise + except Exception: + pass + if turnstile_enabled(): + if not await verify_turnstile_token(body.turnstile_token, ip): + raise HTTPException(status_code=400, detail="Проверка CAPTCHA не пройдена") + email_norm = normalize_login_email(body.email) + if not email_norm: + raise HTTPException(status_code=400, detail="Email обязателен") + if is_disposable_email(email_norm): + raise HTTPException(status_code=400, detail="Одноразовые email-адреса не поддерживаются") + if not smtp_configured(): + raise HTTPException( + status_code=503, + detail="Отправка кода недоступна: почта не настроена на сервере", + ) + if not await redis_ready_for_login_codes(): + raise HTTPException( + status_code=503, + detail="Сервис временно недоступен. Попробуйте позже.", + ) + identity = await idb.get_identity_by_email(session, email_norm) + if not identity: + if not body.allow_register: + return {"ok": True, "message": "Код отправлен на почту"} + identity = await idb.create_identity(session, email=email_norm) + ip = _client_ip(request) + if not await try_consume_ip_send_budget(ip): + raise HTTPException( + status_code=429, + detail="Слишком много запросов с вашего адреса. Попробуйте позже.", + ) + if not await try_consume_email_send_budget(email_norm): + raise HTTPException( + status_code=429, + detail="Слишком много запросов для этого адреса. Попробуйте позже.", + ) + if not await try_acquire_resend_cooldown(email_norm): + raise HTTPException( + status_code=429, + detail="Код уже отправлен. Подождите перед повторной отправкой.", + ) + code = "".join(secrets.choice("0123456789") for _ in range(6)) + if not await store_code(email_norm, code): + await release_resend_cooldown(email_norm) + raise HTTPException( + status_code=503, + detail="Не удалось сохранить код. Попробуйте позже.", + ) + try: + await send_login_code_email(email_norm, code) + except Exception: + await release_resend_cooldown(email_norm) + await delete_code(email_norm) + raise HTTPException( + status_code=503, + detail="Не удалось отправить письмо. Попробуйте позже.", + ) from None + return {"ok": True, "message": "Код отправлен на почту"} + + +@router.post("/login-by-code", response_model=LoginResponse) +async def login_by_code( + body: LoginByCodeRequest, + request: Request, + response: Response, + session: AsyncSession = Depends(get_session), +): + """Вход по email и коду из письма.""" + email_norm = normalize_login_email(body.email) + if not email_norm or not body.code or not body.code.strip(): + raise HTTPException(status_code=400, detail="Email и код обязательны") + if not await redis_ready_for_login_codes(): + raise HTTPException( + status_code=503, + detail="Сервис временно недоступен. Попробуйте позже.", + ) + if not await try_consume_email_verify_budget(email_norm): + raise HTTPException( + status_code=429, + detail="Слишком много попыток. Запросите новый код.", + ) + if not await verify_and_consume_code(email_norm, body.code.strip()): + raise HTTPException(status_code=401, detail="Неверный код или срок действия истёк") + identity = await idb.get_identity_by_email(session, email_norm) + if not identity: + raise HTTPException(status_code=401, detail="Аккаунт не найден") + if not getattr(identity, "email_verified", False): + from sqlalchemy import update as sa_update + + from database.models import Identity as IdentityModel + await session.execute(sa_update(IdentityModel).where(IdentityModel.id == identity.id).values(email_verified=True)) + await bind_identity_actor(request, session, identity) + token = await idb.issue_token_for_identity(session, identity) + logger.info("[Auth] Login success: identity={}, email={}, method=code", identity.id, email_norm) + set_auth_cookie(response, token, request) + set_is_admin_cookie(response, identity, request) + return LoginResponse(identity_id=identity.id) + + +@router.post("/request-password-reset") +async def request_password_reset( + body: SendLoginCodeRequest, + request: Request, + session: AsyncSession = Depends(get_session), +): + email_norm = normalize_login_email(body.email) + if not email_norm: + raise HTTPException(status_code=400, detail="Email обязателен") + if not smtp_configured() or not await pwd_reset.redis_ready(): + return _RESET_OK_MESSAGE + identity = await idb.get_identity_by_email(session, email_norm) + if not identity or not identity.password_hash: + return _RESET_OK_MESSAGE + ip = _client_ip(request) + if not await pwd_reset.try_consume_ip_budget(ip): + raise HTTPException( + status_code=429, + detail="Слишком много запросов с вашего адреса. Попробуйте позже.", + ) + if not await pwd_reset.try_consume_email_send_budget(email_norm): + raise HTTPException( + status_code=429, + detail="Слишком много запросов для этого адреса. Попробуйте позже.", + ) + if not await pwd_reset.try_acquire_cooldown(email_norm): + raise HTTPException( + status_code=429, + detail="Код уже отправлен. Подождите перед повторной отправкой.", + ) + code = "".join(secrets.choice("0123456789") for _ in range(6)) + if not await pwd_reset.store_code(email_norm, code): + await pwd_reset.release_cooldown(email_norm) + raise HTTPException( + status_code=503, + detail="Не удалось сохранить код. Попробуйте позже.", + ) + try: + await send_password_reset_code_email(email_norm, code) + except Exception: + await pwd_reset.release_cooldown(email_norm) + await pwd_reset.delete_code(email_norm) + raise HTTPException( + status_code=503, + detail="Не удалось отправить письмо. Попробуйте позже.", + ) from None + return _RESET_OK_MESSAGE + + +@router.post("/confirm-password-reset", response_model=LoginResponse) +async def confirm_password_reset( + body: ConfirmPasswordResetRequest, + request: Request, + response: Response, + session: AsyncSession = Depends(get_session), +): + email_norm = normalize_login_email(body.email) + if not email_norm or not body.code or not body.code.strip(): + raise HTTPException(status_code=400, detail="Email и код обязательны") + if body.password != body.password_confirm: + raise HTTPException(status_code=400, detail="Пароли не совпадают") + if len(body.password) < 8: + raise HTTPException(status_code=400, detail="Пароль минимум 8 символов") + if not await pwd_reset.redis_ready(): + raise HTTPException( + status_code=503, + detail="Сервис временно недоступен. Попробуйте позже.", + ) + if not await pwd_reset.try_consume_email_verify_budget(email_norm): + raise HTTPException( + status_code=429, + detail="Слишком много попыток. Запросите новый код.", + ) + if not await pwd_reset.verify_and_consume_code(email_norm, body.code.strip()): + raise HTTPException(status_code=401, detail="Неверный код или срок действия истёк") + identity = await idb.get_identity_by_email(session, email_norm) + if not identity: + raise HTTPException(status_code=400, detail="Аккаунт не найден") + updated = await idb.set_password_for_identity(session, identity.id, body.password) + if not updated: + raise HTTPException(status_code=400, detail="Не удалось обновить пароль") + await bind_identity_actor(request, session, updated) + token = await idb.issue_token_for_identity(session, updated) + set_auth_cookie(response, token, request) + set_is_admin_cookie(response, updated, request) + return LoginResponse(identity_id=updated.id) diff --git a/api/v2/routes/auth/session.py b/api/v2/routes/auth/session.py new file mode 100644 index 00000000..34ecc478 --- /dev/null +++ b/api/v2/routes/auth/session.py @@ -0,0 +1,151 @@ +from fastapi import APIRouter, Depends, HTTPException, Request, Response +from sqlalchemy import func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from api.depends import ( + bind_identity_actor, + clear_auth_cookie, + get_request_actor, + get_session, + verify_identity_token, +) +from api.v2.routes.auth._common import _resolve_partner_snapshot +from api.v2.schemas.identities import ( + ChangePasswordRequest, + IdentityResponse, + SetPasswordRequest, +) +from api.v2.schemas.web_public import AccountSummaryResponse +from database import ( + get_balance, + get_keys, + get_trial, + identities as idb, +) +from database.models import CouponUsage, Gift, GiftUsage +from database.referrals import get_referral_stats +from database.web_notifications import count_unread_for_identity +from utils.referral_codes import encode_referral_code + + +router = APIRouter() + + +@router.get("/me", response_model=IdentityResponse) +async def me( + identity=Depends(verify_identity_token), +): + """Текущая идентичность по HttpOnly cookie `auth_token`.""" + return IdentityResponse.model_validate(identity) + + +@router.post("/logout") +async def logout( + request: Request, + response: Response, +): + """Очищает auth cookie. Не требует валидной сессии — всегда возвращает ok.""" + clear_auth_cookie(response, request) + return {"ok": True} + + +@router.get("/summary", response_model=AccountSummaryResponse) +async def auth_summary( + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + actor = get_request_actor(request) + billing_user_id = actor.billing_user_id if actor and actor.billing_user_id is not None else None + if billing_user_id is None: + billing_user_id = await idb.ensure_billing_user_for_identity(session, identity) + balance = float(await get_balance(session, billing_user_id)) + trial_status = await get_trial(session, billing_user_id) + keys = await get_keys(session, billing_user_id) + keys_total = len(keys) if keys else 0 + gifts_sent_r = await session.execute( + select(func.count()).select_from(Gift).where(Gift.sender_user_id == billing_user_id) + ) + gifts_sent = gifts_sent_r.scalar_one() or 0 + gifts_claimed_r = await session.execute( + select(func.count()).select_from(GiftUsage).where(GiftUsage.user_id == billing_user_id) + ) + gifts_claimed = gifts_claimed_r.scalar_one() or 0 + coupons_r = await session.execute( + select(func.count()).select_from(CouponUsage).where(CouponUsage.user_id == billing_user_id) + ) + coupons_used = coupons_r.scalar_one() or 0 + ref = await get_referral_stats(session, billing_user_id) + partner = await _resolve_partner_snapshot(session, int(billing_user_id)) + unread_notifications = await count_unread_for_identity(session, identity.id) + return AccountSummaryResponse( + identity_id=identity.id, + email=identity.email, + tg_id=identity.tg_id, + linked_telegram=identity.tg_id is not None, + referral_code=encode_referral_code(int(billing_user_id)), + balance=balance, + trial_status=int(trial_status), + keys_total=keys_total, + referrals_total=int(ref.get("total_referrals") or 0), + referrals_active=int(ref.get("active_referrals") or 0), + referral_bonus_total=float(ref.get("total_referral_bonus") or 0), + gifts_sent=int(gifts_sent), + gifts_claimed=int(gifts_claimed), + coupons_used=int(coupons_used), + partner_enabled=bool(partner.get("partner_enabled", False)), + partner_code=str(partner.get("partner_code") or ""), + partner_balance=float(partner.get("partner_balance") or 0.0), + partner_percent=float(partner.get("partner_percent") or 0.0), + partner_percent_custom=bool(partner.get("partner_percent_custom", False)), + partner_referred_total=int(partner.get("partner_referred_total") or 0), + partner_payout_method=partner.get("partner_payout_method"), + unread_notifications=int(unread_notifications), + ) + + +@router.post("/set-password") +async def set_password( + body: SetPasswordRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + if body.password != body.password_confirm: + raise HTTPException(status_code=400, detail="Пароли не совпадают") + updated = await idb.set_initial_password(session, identity.id, body.password) + if not updated: + raise HTTPException( + status_code=409, + detail="Пароль уже установлен или аккаунт недоступен", + ) + await bind_identity_actor(request, session, updated) + return {"ok": True} + + +@router.post("/change-password") +async def change_password( + body: ChangePasswordRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + if body.password != body.password_confirm: + raise HTTPException(status_code=400, detail="Новые пароли не совпадают") + err = await idb.change_identity_password( + session, + identity.id, + body.current_password, + body.password, + ) + if err == "no_password": + raise HTTPException( + status_code=409, + detail="Пароль ещё не установлен. Сначала задайте пароль в кабинете.", + ) + if err == "wrong_password": + raise HTTPException(status_code=401, detail="Неверный текущий пароль") + refreshed = await idb.get_identity_by_id(session, identity.id) + if refreshed: + await bind_identity_actor(request, session, refreshed) + return {"ok": True} diff --git a/api/v2/routes/auth/telegram.py b/api/v2/routes/auth/telegram.py new file mode 100644 index 00000000..0e7a32a6 --- /dev/null +++ b/api/v2/routes/auth/telegram.py @@ -0,0 +1,101 @@ +from fastapi import APIRouter, Depends, HTTPException, Request, Response +from pydantic import ( + BaseModel, + Field as PydanticField, +) +from sqlalchemy.ext.asyncio import AsyncSession + +from api.depends import ( + bind_identity_actor, + get_session, + set_auth_cookie, + set_is_admin_cookie, + verify_identity_token, +) +from api.v2.routes.auth._common import TELEGRAM_LOGIN_MAX_AGE, TOKEN_TTL_HINT, _client_ip +from api.v2.schemas.identities import ( + IdentityResponse, + LinkTelegramRequest, + LoginResponse, + LoginTelegramRequest, +) +from config import API_TOKEN +from database import identities as idb +from logger import logger +from utils.telegram_login import verify_telegram_login + + +router = APIRouter() + + +class LoginTelegramWebAppRequest(BaseModel): + init_data: str = PydanticField(..., min_length=1) + + +@router.post("/login-telegram", response_model=LoginResponse) +async def login_telegram( + body: LoginTelegramRequest, + request: Request, + response: Response, + session: AsyncSession = Depends(get_session), +): + ( + """Вход через Telegram Login Widget (кнопка на сайте). По tg_id находим или создаём Identity, выдаём токен. Срок действия токена: """ + + TOKEN_TTL_HINT + + "." + ) + payload = body.model_dump(mode="json") + if not verify_telegram_login(payload, API_TOKEN, max_age_seconds=TELEGRAM_LOGIN_MAX_AGE): + raise HTTPException(status_code=401, detail="Неверная подпись или устаревшие данные от Telegram") + identity = await idb.get_or_create_identity_for_tg(session, body.id) + await bind_identity_actor(request, session, identity) + token = await idb.issue_token_for_identity(session, identity) + logger.info("[Auth] Login success: identity={}, tg_id={}, ip={}, method=telegram_widget", identity.id, body.id, _client_ip(request)) + set_auth_cookie(response, token, request) + set_is_admin_cookie(response, identity, request) + return LoginResponse(identity_id=identity.id) + + +@router.post("/login-telegram-webapp", response_model=LoginResponse) +async def login_telegram_webapp( + body: LoginTelegramWebAppRequest, + request: Request, + response: Response, + session: AsyncSession = Depends(get_session), +): + """Вход через Telegram WebApp initData. Валидирует HMAC, находит/создаёт Identity по tg_id.""" + from utils.telegram_login import verify_webapp_init_data + result = verify_webapp_init_data(body.init_data, API_TOKEN) + if not result: + raise HTTPException(status_code=401, detail="Неверная подпись initData") + tg_id = result.get("user_id") + if not tg_id: + raise HTTPException(status_code=401, detail="Не удалось определить пользователя из initData") + identity = await idb.get_or_create_identity_for_tg(session, int(tg_id)) + await bind_identity_actor(request, session, identity) + token = await idb.issue_token_for_identity(session, identity) + logger.info("[Auth] Login success: identity={}, tg_id={}, ip={}, method=telegram_webapp", identity.id, tg_id, _client_ip(request)) + set_auth_cookie(response, token, request) + set_is_admin_cookie(response, identity, request) + return LoginResponse(identity_id=identity.id) + + +@router.post("/link-telegram", response_model=IdentityResponse) +async def link_telegram( + body: LinkTelegramRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + """Привязывает Telegram к текущей идентичности. Требуется подпись от Telegram Login Widget (доказательство владения аккаунтом).""" + payload = body.model_dump(mode="json") + if not verify_telegram_login(payload, API_TOKEN, max_age_seconds=TELEGRAM_LOGIN_MAX_AGE): + raise HTTPException(status_code=401, detail="Неверная подпись или устаревшие данные от Telegram") + result = await idb.attach_telegram(session, identity.id, body.id) + if not result: + raise HTTPException( + status_code=409, + detail="Этот Telegram уже привязан к другой идентичности", + ) + await bind_identity_actor(request, session, result) + return IdentityResponse.model_validate(result) diff --git a/api/v2/routes/coupon_pricing.py b/api/v2/routes/coupon_pricing.py new file mode 100644 index 00000000..901d4798 --- /dev/null +++ b/api/v2/routes/coupon_pricing.py @@ -0,0 +1,31 @@ +from fastapi import HTTPException +from sqlalchemy.ext.asyncio import AsyncSession + +from services.coupons import resolve_percent_coupon +from services.errors import ServiceError + + +async def resolve_percent_coupon_pricing( + session: AsyncSession, + billing_user_id: int, + base_price_rub: int, + coupon_code: str | None, +) -> tuple[int, int, int | None, str | None]: + """Применяет процентный купон. Бросает HTTPException при ошибке.""" + try: + return await resolve_percent_coupon( + session=session, + billing_user_id=billing_user_id, + base_price_rub=base_price_rub, + coupon_code=coupon_code, + ) + except ServiceError as e: + status_map = { + "not_found": 404, + "limit_exceeded": 409, + "validation_error": 400, + } + raise HTTPException( + status_code=status_map.get(e.code, 400), + detail=e.message, + ) diff --git a/api/v2/routes/coupons.py b/api/v2/routes/coupons.py index 2cf50d78..cf8cde4c 100644 --- a/api/v2/routes/coupons.py +++ b/api/v2/routes/coupons.py @@ -1,8 +1,14 @@ -from fastapi import APIRouter +from fastapi import Depends, HTTPException, Request from api.v2.base_crud import generate_crud_router from api.v2.schemas import CouponBase, CouponResponse, CouponUpdate +from api.v2.schemas.web_public import CouponApplyRequest, CouponApplyResponse +from api.depends import get_request_actor, get_session, verify_identity_token +from database import identities as idb from database.models import Coupon +from services.coupons import apply_fixed_coupon +from services.errors import LimitExceededError, NotFoundError, ServiceError, ValidationError +from sqlalchemy.ext.asyncio import AsyncSession router = generate_crud_router( model=Coupon, @@ -13,3 +19,53 @@ router = generate_crud_router( parameter_name="code", enabled_methods=["get_all", "get_one", "create", "update", "delete"], ) + + +async def _resolve_coupon_user_id(session: AsyncSession, request: Request, identity) -> tuple[int, int | None]: + actor = get_request_actor(request) + billing_user_id = actor.billing_user_id if actor and actor.billing_user_id is not None else None + if billing_user_id is None: + billing_user_id = await idb.ensure_billing_user_for_identity(session, identity) + tg_id = actor.telegram_chat_id if actor else None + return int(billing_user_id), tg_id + + +def _service_error_to_http(e: ServiceError) -> HTTPException: + status_map = { + "not_found": 404, + "limit_exceeded": 409, + "validation_error": 400, + "forbidden": 403, + } + return HTTPException(status_code=status_map.get(e.code, 400), detail=e.message) + + +@router.post("/apply", response_model=CouponApplyResponse) +async def apply_coupon( + body: CouponApplyRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + user_id, tg_id = await _resolve_coupon_user_id(session, request, identity) + try: + result = await apply_fixed_coupon( + session=session, + user_id=user_id, + tg_id=tg_id, + code=str(body.code or ""), + ) + await session.commit() + return CouponApplyResponse( + ok=True, + message="Купон успешно активирован", + coupon_code=result.coupon_code, + amount=result.amount, + balance=result.balance, + ) + except ServiceError as e: + await session.rollback() + raise _service_error_to_http(e) + except Exception: + await session.rollback() + raise HTTPException(status_code=500, detail="Ошибка активации купона") diff --git a/api/v2/routes/flows.py b/api/v2/routes/flows.py new file mode 100644 index 00000000..a9671f16 --- /dev/null +++ b/api/v2/routes/flows.py @@ -0,0 +1,115 @@ +from __future__ import annotations + +from datetime import datetime, UTC + +from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from api.depends import get_session, verify_identity_admin +from api.v2.schemas.flows import FlowCreate, FlowResponse, FlowUpdate +from database.models import WebFlow + +router = APIRouter() + + +def _flow_to_response(flow: WebFlow) -> FlowResponse: + return FlowResponse( + id=flow.id, + name=flow.name, + nodes=flow.nodes or [], + edges=flow.edges or [], + entry_node_id=flow.entry_node_id, + version=flow.version, + ) + + +@router.get("/flows/{flow_id}", response_model=FlowResponse) +async def get_flow_public(flow_id: str, session: AsyncSession = Depends(get_session)): + flow = await session.get(WebFlow, flow_id) + if not flow: + raise HTTPException(404, "Flow not found") + return _flow_to_response(flow) + + +@router.get("/admin/flows", response_model=list[FlowResponse]) +async def list_flows( + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + result = await session.execute(select(WebFlow)) + return [_flow_to_response(f) for f in result.scalars().all()] + + +@router.get("/admin/flows/{flow_id}", response_model=FlowResponse) +async def get_flow_admin( + flow_id: str, + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + flow = await session.get(WebFlow, flow_id) + if not flow: + raise HTTPException(404, "Flow not found") + return _flow_to_response(flow) + + +@router.post("/admin/flows", response_model=FlowResponse, status_code=201) +async def create_flow( + body: FlowCreate, + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + existing = await session.get(WebFlow, body.id) + if existing: + raise HTTPException(409, "Flow with this ID already exists") + + flow = WebFlow( + id=body.id, + name=body.name, + nodes=[n.model_dump() for n in body.nodes], + edges=[e.model_dump() for e in body.edges], + entry_node_id=body.entry_node_id, + version=1, + updated_at=datetime.now(UTC), + ) + session.add(flow) + await session.commit() + await session.refresh(flow) + return _flow_to_response(flow) + + +@router.put("/admin/flows/{flow_id}", response_model=FlowResponse) +async def update_flow( + flow_id: str, + body: FlowUpdate, + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + flow = await session.get(WebFlow, flow_id) + if not flow: + raise HTTPException(404, "Flow not found") + + if body.name is not None: + flow.name = body.name + flow.nodes = [n.model_dump() for n in body.nodes] + flow.edges = [e.model_dump() for e in body.edges] + flow.entry_node_id = body.entry_node_id + flow.version = flow.version + 1 + flow.updated_at = datetime.now(UTC) + + await session.commit() + await session.refresh(flow) + return _flow_to_response(flow) + + +@router.delete("/admin/flows/{flow_id}", status_code=204) +async def delete_flow( + flow_id: str, + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + flow = await session.get(WebFlow, flow_id) + if not flow: + raise HTTPException(404, "Flow not found") + await session.delete(flow) + await session.commit() diff --git a/api/v2/routes/gifts.py b/api/v2/routes/gifts.py index 9ef431d3..718ec457 100644 --- a/api/v2/routes/gifts.py +++ b/api/v2/routes/gifts.py @@ -1,14 +1,251 @@ -from fastapi import APIRouter, Depends, HTTPException, Path -from sqlalchemy import delete, select +from math import ceil + +from fastapi import APIRouter, Depends, HTTPException, Path, Query, Request +from sqlalchemy import delete, func, select from sqlalchemy.ext.asyncio import AsyncSession -from api.depends import get_session, verify_identity_admin +from api.depends import ( + get_request_actor, + get_session, + validate_redirect_url, + verify_identity_admin, + verify_identity_token, +) from api.v2.base_crud import generate_crud_router +from api.v2.routes.tariffs import _resolve_default_web_payment_provider, _resolve_public_base_url from api.v2.schemas import GiftBase, GiftResponse, GiftUpdate, GiftUsageResponse -from database.models import Gift, GiftUsage +from api.v2.schemas.web_public import ( + GiftCreatePreviewResponse, + GiftCreateRequest, + GiftCreateResponse, + GiftRedeemRequest, + GiftRedeemResponse, + GiftUsageEntry, + MyGiftItem, + MyGiftsResponse, +) +from config import GIFT_BUTTON +from core.bootstrap import BUTTONS_CONFIG +from database import ( + get_balance, + identities as idb, +) +from database.access.resolution import resolve_user_optional +from database.models import Gift, GiftUsage, Tariff +from database.tariffs import get_tariff_by_id +from database.temporary_data import create_temporary_data +from services.errors import NotFoundError, ValidationError +from services.formatting import get_site_gift_link +from services.gifts import create_gift as service_create_gift +from services.gifts import redeem_gift as service_redeem_gift +from services.payments.payment_links import PaymentLinkRequest, create_payment_link +from services.tariffs import calculate_config_price + router = APIRouter() + +def _check_gifts_enabled(): + if not bool(BUTTONS_CONFIG.get("GIFT_BUTTON_ENABLE", GIFT_BUTTON)): + raise HTTPException(status_code=403, detail="Подарки отключены") + + +@router.post("/create", tags=["Gifts"]) +async def create_gift_for_user( + body: GiftCreateRequest, + request: Request, + preview: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + _check_gifts_enabled() + actor = get_request_actor(request) + billing_user_id = actor.billing_user_id if actor and actor.billing_user_id is not None else None + if billing_user_id is None: + billing_user_id = await idb.ensure_billing_user_for_identity(session, identity) + + tariff = await get_tariff_by_id(session, body.tariff_id) + if not tariff or tariff.get("group_code") != "gifts" or not tariff.get("is_active", True): + raise HTTPException(status_code=404, detail="Тариф не найден") + + price = int(calculate_config_price(tariff, body.selected_device_limit, body.selected_traffic_gb)) + balance = float(await get_balance(session, billing_user_id)) + + required_amount = int(max(0, ceil(float(price) - balance))) + + if preview: + return GiftCreatePreviewResponse( + ok=True, + price_rub=price, + balance_rub=balance, + sufficient_funds=balance >= price, + tariff_name=str(tariff.get("name", "")), + duration_days=int(tariff.get("duration_days") or 0), + ) + + if required_amount > 0: + provider_id = str(body.provider_id or _resolve_default_web_payment_provider() or "").strip().upper() + if not provider_id: + raise HTTPException(status_code=503, detail="Нет доступных провайдеров оплаты") + base_url = _resolve_public_base_url(request) + success_url = validate_redirect_url(str(body.success_url or ""), f"{base_url}/payment-success") + failure_url = validate_redirect_url(str(body.failure_url or ""), f"{base_url}/payment-failure") + payment_request = PaymentLinkRequest( + legacy_user_ref=int(billing_user_id), + amount=required_amount, + currency="RUB", + provider_id=provider_id, + success_url=success_url, + failure_url=failure_url, + metadata={ + "payment_flow": "gift_create", + "tariff_id": int(body.tariff_id), + "selected_device_limit": body.selected_device_limit, + "selected_traffic_gb": body.selected_traffic_gb, + "selected_price_rub": int(price), + }, + ) + payment_result = await create_payment_link(session, payment_request) + if not payment_result.success or not payment_result.payment_url or not payment_result.payment_id: + raise HTTPException(status_code=400, detail=payment_result.error or "Не удалось создать ссылку оплаты") + await create_temporary_data( + session, + int(billing_user_id), + "waiting_for_payment", + { + "payment_flow": "gift_create", + "tariff_id": int(body.tariff_id), + "required_amount": int(required_amount), + "selected_price_rub": int(price), + "selected_device_limit": body.selected_device_limit, + "selected_traffic_gb": body.selected_traffic_gb, + }, + ) + return GiftCreateResponse( + ok=True, + message="Требуется оплата для создания подарка", + payment_required=True, + required_amount_rub=required_amount, + payment_id=payment_result.payment_id, + payment_url=payment_result.payment_url, + ) + + from services.errors import InsufficientFundsError, NotFoundError + + try: + result = await service_create_gift( + session=session, + sender_user_ref=billing_user_id, + tariff_id=body.tariff_id, + selected_device_limit=body.selected_device_limit, + selected_traffic_gb=body.selected_traffic_gb, + selected_price_rub=price, + ) + except InsufficientFundsError as e: + raise HTTPException(status_code=400, detail=str(e)) from None + except NotFoundError as e: + raise HTTPException(status_code=404, detail=str(e)) from None + + new_balance = float(await get_balance(session, billing_user_id)) + return GiftCreateResponse( + ok=True, + message=f"Подарок создан — {result.tariff_name} на {result.duration_text}", + gift_id=result.gift_id, + site_gift_link=result.site_gift_link, + tariff_name=result.tariff_name, + duration_days=result.duration_days, + price_charged=result.price_charged, + balance_rub=new_balance, + ) + + +@router.get("/my", response_model=MyGiftsResponse, tags=["Gifts"]) +async def get_my_gifts( + request: Request, + limit: int = Query(20, ge=1, le=100), + offset: int = Query(0, ge=0), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + _check_gifts_enabled() + actor = get_request_actor(request) + billing_user_id = actor.billing_user_id if actor and actor.billing_user_id is not None else None + if billing_user_id is None: + billing_user_id = await idb.ensure_billing_user_for_identity(session, identity) + + base_filter = Gift.sender_user_id == billing_user_id + total = (await session.execute(select(func.count()).select_from(Gift).where(base_filter))).scalar_one() + + result = await session.execute( + select(Gift).where(base_filter).order_by(Gift.created_at.desc()).limit(limit).offset(offset) + ) + gifts = result.scalars().all() + + tariff_ids = {g.tariff_id for g in gifts if g.tariff_id} + tariff_map: dict[int, str] = {} + duration_map: dict[int, int] = {} + if tariff_ids: + tariff_rows = await session.execute(select(Tariff).where(Tariff.id.in_(tariff_ids))) + for t in tariff_rows.scalars().all(): + tariff_map[t.id] = t.name or "" + duration_map[t.id] = int(t.duration_days or 0) + + gift_ids = [g.gift_id for g in gifts] + usages_map: dict[str, list[GiftUsageEntry]] = {gid: [] for gid in gift_ids} + if gift_ids: + usage_rows = await session.execute(select(GiftUsage).where(GiftUsage.gift_id.in_(gift_ids))) + for u in usage_rows.scalars().all(): + usages_map.setdefault(u.gift_id, []).append( + GiftUsageEntry( + user_id=int(u.user_id), + used_at=u.used_at.isoformat() if u.used_at else None, + ) + ) + + items = [] + for g in gifts: + items.append( + MyGiftItem( + gift_id=g.gift_id, + tariff_name=tariff_map.get(g.tariff_id, ""), + duration_days=duration_map.get(g.tariff_id, 0), + price_rub=int(g.selected_price_rub or 0), + created_at=g.created_at.isoformat() if g.created_at else None, + expiry_time=g.expiry_time.isoformat() if g.expiry_time else None, + is_used=bool(g.is_used), + is_unlimited=bool(g.is_unlimited), + max_usages=g.max_usages, + site_gift_link=get_site_gift_link(g.gift_id), + usages=usages_map.get(g.gift_id, []), + ) + ) + + return MyGiftsResponse(ok=True, gifts=items, total=total, limit=limit, offset=offset) + + +@router.delete("/my/{gift_id}", response_model=dict, tags=["Gifts"]) +async def delete_my_gift( + request: Request, + gift_id: str = Path(...), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + """Удаляет свой подарок.""" + actor = get_request_actor(request) + billing_user_id = actor.billing_user_id if actor and actor.billing_user_id is not None else None + if billing_user_id is None: + billing_user_id = await idb.ensure_billing_user_for_identity(session, identity) + + result = await session.execute(select(Gift).where(Gift.gift_id == gift_id)) + gift = result.scalar_one_or_none() + if not gift or gift.sender_user_id != billing_user_id: + raise HTTPException(status_code=404, detail="Подарок не найден") + await session.execute(delete(GiftUsage).where(GiftUsage.gift_id == gift_id)) + await session.delete(gift) + await session.commit() + return {"ok": True, "message": "Подарок удалён"} + + gift_router = generate_crud_router( model=Gift, schema_response=GiftResponse, @@ -21,6 +258,35 @@ gift_router = generate_crud_router( router.include_router(gift_router, prefix="", tags=["Gifts"]) +@router.post("/redeem", response_model=GiftRedeemResponse, tags=["Gifts"]) +async def redeem_gift( + body: GiftRedeemRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + _check_gifts_enabled() + actor = get_request_actor(request) + billing_user_id = actor.billing_user_id if actor and actor.billing_user_id is not None else None + if billing_user_id is None: + billing_user_id = await idb.ensure_billing_user_for_identity(session, identity) + try: + result = await service_redeem_gift(session, body.gift_code, billing_user_id) + except ValidationError as e: + raise HTTPException(status_code=400, detail=str(e)) from None + except NotFoundError as e: + raise HTTPException(status_code=404, detail=str(e)) from None + except Exception: + raise HTTPException(status_code=500, detail="Не удалось активировать подарок") from None + return GiftRedeemResponse( + ok=True, + message=result.message, + gift_id=result.gift_id, + tariff_id=result.tariff_id, + duration_days=result.duration_days, + ) + + @router.get("/by_tg_id/{tg_id}", response_model=list[GiftResponse], tags=["Gifts"]) async def get_gifts_by_tg_id( tg_id: int = Path(...), @@ -28,7 +294,10 @@ async def get_gifts_by_tg_id( session: AsyncSession = Depends(get_session), ): """Список подарков по tg_id отправителя.""" - result = await session.execute(select(Gift).where(Gift.sender_tg_id == tg_id)) + u = await resolve_user_optional(session, tg_id) + if u is None: + raise HTTPException(status_code=404, detail="Gifts not found") + result = await session.execute(select(Gift).where(Gift.sender_user_id == u.id)) gifts = result.scalars().all() if not gifts: raise HTTPException(status_code=404, detail="Gifts not found") diff --git a/api/v2/routes/keys/__init__.py b/api/v2/routes/keys/__init__.py new file mode 100644 index 00000000..a68ebb7e --- /dev/null +++ b/api/v2/routes/keys/__init__.py @@ -0,0 +1,4 @@ +from ._common import router, user_router +from . import admin, user # noqa: F401 — import triggers endpoint registration + +__all__ = ["router", "user_router"] diff --git a/api/v2/routes/keys/_common.py b/api/v2/routes/keys/_common.py new file mode 100644 index 00000000..dcc04ba4 --- /dev/null +++ b/api/v2/routes/keys/_common.py @@ -0,0 +1,279 @@ +import asyncio +import re + +from base64 import b64encode +from datetime import datetime +from io import BytesIO +from math import ceil +from typing import Any +from urllib.parse import urlsplit + +import qrcode + +from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query, Request, status +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from api.depends import ( + get_request_actor, + get_session, + validate_redirect_url, + verify_identity_admin, + verify_identity_token, +) +from api.v2.base_crud import generate_crud_router +from api.v2.routes.coupon_pricing import resolve_percent_coupon_pricing +from api.v2.schemas import KeyBase, KeyCreateRequest, KeyResponse, KeyUpdate +from api.v2.schemas.web_public import ( + AccountKeyActionResponse, + AccountKeyActionsAvailability, + AccountKeyActionsConfigResponse, + AccountKeyAddonOptionResponse, + AccountKeyAddonsPreviewRequest, + AccountKeyAddonsPreviewResponse, + AccountKeyAliasUpdateRequest, + AccountKeyApplyAddonsResponse, + AccountKeyChangeLocationRequest, + AccountKeyChangeLocationResponse, + AccountKeyDetailsResponse, + AccountKeyLocationOptionResponse, + AccountKeyLocationsResponse, + AccountKeyQrResponse, + AccountKeyRenewRequest, + AccountKeyRenewResponse, + AccountKeyResetHwidResponse, + AccountKeyResponse, +) +from config import ( + ENABLE_DELETE_KEY_BUTTON, + HWID_RESET_BUTTON, + INSTRUCTIONS_BUTTON, + QRCODE, + REMNAWAVE_LOGIN, + REMNAWAVE_PASSWORD, + USE_COUNTRY_SELECTION, +) +from core.bootstrap import BUTTONS_CONFIG, MODES_CONFIG, PAYMENTS_CONFIG, TARIFFS_CONFIG +from core.settings.tariffs_config import normalize_tariff_config +from database import ( + check_server_name_by_cluster, + filter_cluster_by_subgroup, + get_balance, + get_key_details, + get_keys, + get_tariff_by_id, + identities as idb, + save_key_config_with_mode, + update_balance, +) +from database.access.resolution import resolve_user_optional +from database.coupons import mark_coupon_used +from database.models import Key, Server, ServerSpecialgroup, Tariff +from database.temporary_data import create_temporary_data +from handlers.buttons import CONNECT_DEVICE, ROUTER_BUTTON, TV_BUTTON +from handlers.keys.key_view import build_key_view_payload +from handlers.tariffs.addons.key_addons_pack import calc_pack_full_price_rub, get_pack_flags +from handlers.tariffs.addons.utils import calc_remaining_ratio_seconds, is_not_downgrade +from handlers.utils import ALLOWED_GROUP_CODES, is_full_remnawave_cluster +from logger import logger +from panels._3xui import delete_client, get_xui_instance +from panels.remnawave import RemnawaveAPI, get_vless_link_for_remnawave_by_username +from panels.remnawave_runtime import get_remnawave_profile, invalidate_remnawave_profile, with_remnawave_api +from services.operations import ( + create_client_on_server, + create_key_on_cluster, + delete_key_from_cluster, + renew_key_in_cluster, +) +from services.operations.aggregated_links import make_aggregated_link +from services.payments.payment_links import PaymentLinkRequest, create_payment_link +from services.payments.providers import WEB_LINK_PROVIDER_IDS +from services.tariffs import calculate_config_price +from services.tariffs.tariff_display import GB, get_effective_limits_for_key, get_key_tariff_addons_state + + + +router = generate_crud_router( + model=Key, + schema_response=KeyResponse, + schema_create=KeyBase, + schema_update=KeyUpdate, + identifier_field="tg_id", + extra_get_by_email=True, + enabled_methods=["get_all", "get_one", "get_by_email", "get_all_by_field"], +) +user_router = APIRouter() + + + + +def _key_actions_config() -> AccountKeyActionsConfigResponse: + addons_mode = str(TARIFFS_CONFIG.get("KEY_ADDONS_PACK_MODE", "") or "").strip().lower() + if addons_mode not in {"", "traffic", "devices", "all"}: + addons_mode = "" + addons_enabled_default = addons_mode in {"", "traffic", "devices", "all"} + return AccountKeyActionsConfigResponse( + renew_enabled=True, + delete_enabled=bool(BUTTONS_CONFIG.get("DELETE_KEY_BUTTON_ENABLE", ENABLE_DELETE_KEY_BUTTON)), + qr_enabled=bool(BUTTONS_CONFIG.get("QRCODE_BUTTON_ENABLE", QRCODE)), + hwid_reset_enabled=bool(BUTTONS_CONFIG.get("HWID_RESET_BUTTON_ENABLE", HWID_RESET_BUTTON)), + country_change_enabled=bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)), + instructions_enabled=bool(BUTTONS_CONFIG.get("INSTRUCTIONS_BUTTON_ENABLE", INSTRUCTIONS_BUTTON)), + addons_enabled=addons_enabled_default, + addons_mode=addons_mode, + tv_connect_enabled=bool(BUTTONS_CONFIG.get("ANDROID_TV_BUTTON_ENABLE")), + ) + + +def _extract_key_actions_from_markup(markup) -> AccountKeyActionsAvailability: + actions = AccountKeyActionsAvailability() + rows = getattr(markup, "inline_keyboard", None) or [] + for row in rows: + for button in row: + callback_data = str(getattr(button, "callback_data", "") or "") + text = str(getattr(button, "text", "") or "") + has_url = bool(getattr(button, "url", None)) + has_web_app = bool(getattr(button, "web_app", None)) + if callback_data.startswith("connect_router|") or text == ROUTER_BUTTON: + actions.can_connect_router = True + if callback_data.startswith("connect_tv|") or text == TV_BUTTON: + actions.can_connect_tv = True + if callback_data.startswith("connect_device|") or (text == CONNECT_DEVICE and (has_url or has_web_app)): + actions.can_connect_device = True + if callback_data.startswith("renew_key|"): + actions.can_renew = True + if callback_data.startswith("key_addons|"): + actions.can_addons = True + if callback_data.startswith("reset_hwid|"): + actions.can_reset_hwid = True + if callback_data.startswith("show_qr|"): + actions.can_qr = True + if callback_data.startswith("delete_key|"): + actions.can_delete = True + if callback_data.startswith("change_location|"): + actions.can_change_location = True + return actions + + +async def _resolve_available_location_servers(session: AsyncSession, db_key: Key) -> list[str]: + current_server = str(getattr(db_key, "server_id", "") or "") + if not current_server: + return [] + cluster_info = await check_server_name_by_cluster(session, current_server) + if not cluster_info: + return [] + cluster_name = str(cluster_info.get("cluster_name") or "") + if not cluster_name: + return [] + q = ( + select( + Server.id, + Server.server_name, + Server.api_url, + Server.panel_type, + Server.enabled, + Server.max_keys, + ) + .where(Server.cluster_name == cluster_name) + .where(Server.server_name != current_server) + ) + servers = [dict(m) for m in (await session.execute(q)).mappings().all()] + if not servers: + return [] + server_ids = [s["id"] for s in servers if s.get("id") is not None] + groups_map: dict[int, list[str]] = {} + if server_ids: + r = await session.execute( + select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where( + ServerSpecialgroup.server_id.in_(server_ids) + ) + ) + for sid, gc in r.all(): + groups_map.setdefault(int(sid), []).append(gc) + for server in servers: + sid_raw = server.get("id") + sid = int(sid_raw) if sid_raw is not None else -1 + server["special_groups"] = [g for g in groups_map.get(sid, []) if g in ALLOWED_GROUP_CODES] + key_tariff_id = getattr(db_key, "tariff_id", None) + subgroup_title = None + tariff_dict = None + if key_tariff_id: + tariff_dict = await get_tariff_by_id(session, int(key_tariff_id)) + if tariff_dict: + subgroup_title = tariff_dict.get("subgroup_title") + available_servers = [s for s in servers if bool(s.get("enabled", True))] + if subgroup_title and available_servers: + filtered_servers = await filter_cluster_by_subgroup( + session=session, + cluster=available_servers, + target_subgroup=str(subgroup_title).strip(), + cluster_id=cluster_name, + tariff_id=int(key_tariff_id) if key_tariff_id else None, + ) + if filtered_servers: + available_servers = filtered_servers + else: + available_servers = [] + if available_servers and tariff_dict: + special = None + gc = str(tariff_dict.get("group_code") or "").lower() + if gc and gc in ALLOWED_GROUP_CODES: + special = gc + if special: + bound_servers = [s for s in available_servers if special in (s.get("special_groups") or [])] + if bound_servers: + available_servers = bound_servers + names = sorted( + { + str(s.get("server_name") or "").strip() + for s in available_servers + if str(s.get("server_name") or "").strip() + } + ) + return names + + +async def _resolve_billing_user_id(request: Request, identity, session: AsyncSession) -> int: + actor = get_request_actor(request) + billing_user_id = actor.billing_user_id if actor and actor.billing_user_id is not None else None + if billing_user_id is None: + billing_user_id = await idb.ensure_billing_user_for_identity(session, identity) + return int(billing_user_id) + + +def _resolve_public_base_url(request: Request) -> str: + origin = str(request.headers.get("origin") or "").strip() + if origin.startswith(("http://", "https://")): + return origin.rstrip("/") + referer = str(request.headers.get("referer") or request.headers.get("referrer") or "").strip() + if referer.startswith(("http://", "https://")): + parsed = urlsplit(referer) + if parsed.scheme and parsed.netloc: + return f"{parsed.scheme}://{parsed.netloc}".rstrip("/") + forwarded_host = str(request.headers.get("x-forwarded-host") or "").strip() + host = forwarded_host or str(request.headers.get("host") or "").strip() + forwarded_proto = str(request.headers.get("x-forwarded-proto") or "").split(",", 1)[0].strip().lower() + scheme = forwarded_proto if forwarded_proto in {"http", "https"} else request.url.scheme + if host: + return f"{scheme}://{host}".rstrip("/") + return str(request.base_url).rstrip("/") + + +def _resolve_default_web_payment_provider() -> str | None: + for provider_id in WEB_LINK_PROVIDER_IDS: + if bool(PAYMENTS_CONFIG.get(provider_id)): + return provider_id + return WEB_LINK_PROVIDER_IDS[0] if WEB_LINK_PROVIDER_IDS else None + + +def _normalize_expiry_ms(raw_value: int | float | None) -> int: + if not raw_value: + return 0 + value = int(raw_value) + if value > 10**13: + value //= 1000 + elif value < 10**10: + value *= 1000 + return value + + diff --git a/api/v2/routes/keys.py b/api/v2/routes/keys/admin.py similarity index 84% rename from api/v2/routes/keys.py rename to api/v2/routes/keys/admin.py index 9ab70733..06df1847 100644 --- a/api/v2/routes/keys.py +++ b/api/v2/routes/keys/admin.py @@ -1,26 +1,5 @@ -from datetime import datetime - -from fastapi import Body, Depends, HTTPException, Path, status -from sqlalchemy import select -from sqlalchemy.ext.asyncio import AsyncSession - -from api.depends import get_session, verify_identity_admin -from api.v2.schemas import KeyBase, KeyCreateRequest, KeyResponse, KeyUpdate -from api.v2.base_crud import generate_crud_router -from database.models import Key, Tariff -from handlers.keys.operations import create_key_on_cluster, delete_key_from_cluster, renew_key_in_cluster -from logger import logger - -router = generate_crud_router( - model=Key, - schema_response=KeyResponse, - schema_create=KeyBase, - schema_update=KeyUpdate, - identifier_field="tg_id", - extra_get_by_email=True, - enabled_methods=["get_all", "get_one", "get_by_email", "get_all_by_field"], -) - +from ._common import * # noqa: F401,F403 +from ._common import router, user_router # noqa: F401 @router.delete("/by_email/{email}", response_model=dict) async def delete_key_by_email( @@ -60,7 +39,10 @@ async def get_router_keys_by_tg_id( tariff_ids = [row[0] for row in tariffs_result.all()] if not tariff_ids: return [] - keys_result = await session.execute(select(Key).where(Key.tg_id == tg_id, Key.tariff_id.in_(tariff_ids))) + u = await resolve_user_optional(session, tg_id) + if u is None: + return [] + keys_result = await session.execute(select(Key).where(Key.user_id == u.id, Key.tariff_id.in_(tariff_ids))) return keys_result.scalars().all() @@ -132,3 +114,4 @@ async def create_key_api( except Exception as e: logger.error(f"[API] Ошибка при создании ключа: {e}") raise HTTPException(status_code=500, detail="Ошибка при создании ключа") + diff --git a/api/v2/routes/keys/user/__init__.py b/api/v2/routes/keys/user/__init__.py new file mode 100644 index 00000000..e834a77a --- /dev/null +++ b/api/v2/routes/keys/user/__init__.py @@ -0,0 +1 @@ +from . import addons, core, hwid, location, renew # noqa: F401 — trigger registration diff --git a/api/v2/routes/keys/user/addons.py b/api/v2/routes/keys/user/addons.py new file mode 100644 index 00000000..288b254c --- /dev/null +++ b/api/v2/routes/keys/user/addons.py @@ -0,0 +1,619 @@ +"""User-facing key endpoints (/api/keys/*). + +Регистрирует эндпоинты на ``user_router`` из ``_common``. Импорт этого модуля +из ``__init__.py`` запускает регистрацию декораторов. +""" + +from .._common import * # noqa: F401,F403 — подтягиваем все имена для endpoints +from .._common import ( + _key_actions_config, + _resolve_available_location_servers, + _resolve_billing_user_id, + _resolve_default_web_payment_provider, + _resolve_public_base_url, + _normalize_expiry_ms, + router, + user_router, +) + + +@user_router.get("/{client_id}/addons-preview", response_model=AccountKeyAddonsPreviewResponse) +async def user_key_addons_preview( + client_id: str, + request: Request, + selected_device_limit: int | None = Query(None), + selected_traffic_gb: int | None = Query(None), + include_device: bool | None = Query(None), + include_traffic: bool | None = Query(None), + coupon_code: str | None = Query(None), + force_web: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + actions = _key_actions_config() + if not force_web and not actions.addons_enabled: + raise HTTPException(status_code=403, detail="Доп. опции отключены в настройках") + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + tariff_id = getattr(db_key, "tariff_id", None) + if not tariff_id: + raise HTTPException(status_code=400, detail="Для подписки не назначен тариф") + tariff = await get_tariff_by_id(session, int(tariff_id)) + if not tariff: + raise HTTPException(status_code=404, detail="Тариф не найден") + key_details = await get_key_details(session, str(getattr(db_key, "email", "") or "")) + if not key_details: + raise HTTPException(status_code=404, detail="Подписка не найдена") + ( + _tariff_name, + _subgroup_title, + _traffic_limit_gb, + _device_limit, + _panel, + is_tariff_configurable, + addons_devices_enabled, + addons_traffic_enabled, + ) = await get_key_tariff_addons_state(session=session, key_record=key_details, db_key=db_key) + if not is_tariff_configurable: + raise HTTPException(status_code=400, detail="Тариф не поддерживает доп. опции") + cfg = normalize_tariff_config(tariff) + raw_device_options = cfg.get("device_options") or tariff.get("device_options") or [] + raw_traffic_options = cfg.get("traffic_options_gb") or tariff.get("traffic_options_gb") or [] + device_options: list[int] = [] + for value in raw_device_options: + try: + device_options.append(int(value)) + except (TypeError, ValueError): + continue + traffic_options: list[int] = [] + for value in raw_traffic_options: + try: + traffic_options.append(int(value)) + except (TypeError, ValueError): + continue + device_options = sorted(set(device_options), key=lambda val: (int(val == 0), val)) + traffic_options = sorted(set(traffic_options), key=lambda val: (int(val == 0), val)) + has_device_option = bool(device_options) and bool(addons_devices_enabled) + has_traffic_option = bool(traffic_options) and bool(addons_traffic_enabled) + pack_devices, pack_traffic, pack_mode = get_pack_flags() + if pack_mode: + has_device_option = has_device_option and bool(pack_devices) + has_traffic_option = has_traffic_option and bool(pack_traffic) + if not has_device_option: + device_options = [] + if not has_traffic_option: + traffic_options = [] + if not has_device_option and not has_traffic_option: + raise HTTPException(status_code=400, detail="Доп. опции для этой подписки недоступны") + selected_device_limit_db = key_details.get("selected_device_limit") + selected_traffic_limit_db = key_details.get("selected_traffic_limit") + current_device_limit_db = key_details.get("current_device_limit") + current_traffic_limit_db = key_details.get("current_traffic_limit") + base_devices = tariff.get("device_limit") + base_devices = int(base_devices) if base_devices is not None else None + base_traffic_bytes = tariff.get("traffic_limit") + base_traffic_gb_from_tariff = int(base_traffic_bytes / GB) if base_traffic_bytes else None + current_device_limit = ( + int(current_device_limit_db) + if current_device_limit_db is not None + else (int(selected_device_limit_db) if selected_device_limit_db is not None else base_devices) + ) + current_traffic_gb = ( + int(current_traffic_limit_db) + if current_traffic_limit_db is not None + else (int(selected_traffic_limit_db) if selected_traffic_limit_db is not None else base_traffic_gb_from_tariff) + ) + if pack_mode and current_device_limit is not None and int(current_device_limit) == 0: + has_device_option = False + device_options = [] + if pack_mode and current_traffic_gb is not None and int(current_traffic_gb) == 0: + has_traffic_option = False + traffic_options = [] + if not has_device_option and not has_traffic_option: + raise HTTPException(status_code=400, detail="Доп. опции для этой подписки недоступны") + if pack_mode: + include_device_effective = bool(include_device) if include_device is not None else selected_device_limit is not None + include_traffic_effective = bool(include_traffic) if include_traffic is not None else selected_traffic_gb is not None + selected_device = ( + selected_device_limit + if selected_device_limit is not None + else None + ) + selected_traffic = ( + selected_traffic_gb + if selected_traffic_gb is not None + else None + ) + else: + include_device_effective = has_device_option + include_traffic_effective = has_traffic_option + selected_device = selected_device_limit if selected_device_limit is not None else current_device_limit + selected_traffic = selected_traffic_gb if selected_traffic_gb is not None else current_traffic_gb + if has_device_option and include_device_effective and selected_device is not None and int(selected_device) not in device_options: + raise HTTPException(status_code=400, detail="Выбранный пакет устройств недоступен") + if has_traffic_option and include_traffic_effective and selected_traffic is not None and int(selected_traffic) not in traffic_options: + raise HTTPException(status_code=400, detail="Выбранный пакет трафика недоступен") + current_devices_for_price = int(current_device_limit) if current_device_limit is not None else None + current_traffic_for_price = int(current_traffic_gb) if current_traffic_gb is not None else None + base_price_for_current = int( + calculate_config_price( + tariff=tariff, + selected_device_limit=current_devices_for_price, + selected_traffic_gb=current_traffic_for_price, + ) + ) + if pack_mode: + diff_full = int( + calc_pack_full_price_rub( + tariff=tariff, + has_device_option=bool(has_device_option and include_device_effective), + has_traffic_option=bool(has_traffic_option and include_traffic_effective), + selected_devices=int(selected_device) if has_device_option and include_device_effective and selected_device is not None else None, + selected_traffic_gb=int(selected_traffic) if has_traffic_option and include_traffic_effective and selected_traffic is not None else None, + ) + ) + recalc_enabled = bool( + MODES_CONFIG.get( + "KEY_ADDONS_RECALC_PRICE", + TARIFFS_CONFIG.get("KEY_ADDONS_RECALC_PRICE", False), + ) + ) + if recalc_enabled: + remaining_seconds, total_seconds = calc_remaining_ratio_seconds( + key_details.get("expiry_time"), + tariff, + ) + extra_price_rub = int((diff_full * remaining_seconds + total_seconds - 1) // total_seconds) + else: + extra_price_rub = diff_full + total_price_rub = int(base_price_for_current + diff_full) + else: + total_price_rub = int( + calculate_config_price( + tariff=tariff, + selected_device_limit=int(selected_device) if has_device_option and include_device_effective and selected_device is not None else None, + selected_traffic_gb=int(selected_traffic) if has_traffic_option and include_traffic_effective and selected_traffic is not None else None, + ) + ) + extra_price_rub = int(max(0, total_price_rub - base_price_for_current)) + final_extra_price_rub, discount_rub, _coupon_id, applied_coupon_code = await resolve_percent_coupon_pricing( + session=session, + billing_user_id=int(billing_user_id), + base_price_rub=int(max(0, extra_price_rub)), + coupon_code=coupon_code, + ) + return AccountKeyAddonsPreviewResponse( + client_id=str(getattr(db_key, "client_id", "") or ""), + tariff_id=int(tariff_id), + addons_mode=str(pack_mode or ""), + has_device_option=bool(has_device_option), + has_traffic_option=bool(has_traffic_option), + current_device_limit=int(current_device_limit) if current_device_limit is not None else None, + current_traffic_gb=int(current_traffic_gb) if current_traffic_gb is not None else None, + selected_device_limit=int(selected_device) if has_device_option and include_device_effective and selected_device is not None else None, + selected_traffic_gb=int(selected_traffic) if has_traffic_option and include_traffic_effective and selected_traffic is not None else None, + device_options=[ + AccountKeyAddonOptionResponse( + value=int(val), + label=( + "Безлимит устройств" + if int(val) <= 0 + else (f"+{int(val)} устройств" if pack_mode else f"{int(val)} устройств") + ), + ) + for val in device_options + ], + traffic_options=[ + AccountKeyAddonOptionResponse( + value=int(val), + label=( + "Безлимит трафика" + if int(val) <= 0 + else (f"+{int(val)} ГБ" if pack_mode else f"{int(val)} ГБ") + ), + ) + for val in traffic_options + ], + total_price_rub=int(total_price_rub), + extra_price_rub=int(max(0, extra_price_rub)), + discount_rub=int(discount_rub), + final_price_rub=int(max(0, final_extra_price_rub)), + applied_coupon_code=applied_coupon_code, + balance_rub=float(await get_balance(session, int(billing_user_id))), + ) + + +@user_router.post("/{client_id}/apply-addons", response_model=AccountKeyApplyAddonsResponse) +async def user_key_apply_addons( + client_id: str, + body: AccountKeyAddonsPreviewRequest, + request: Request, + force_web: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + actions = _key_actions_config() + if not force_web and not actions.addons_enabled: + raise HTTPException(status_code=403, detail="Доп. опции отключены в настройках") + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + tariff_id = getattr(db_key, "tariff_id", None) + if not tariff_id: + raise HTTPException(status_code=400, detail="Для подписки не назначен тариф") + tariff = await get_tariff_by_id(session, int(tariff_id)) + if not tariff: + raise HTTPException(status_code=404, detail="Тариф не найден") + key_details = await get_key_details(session, str(getattr(db_key, "email", "") or "")) + if not key_details: + raise HTTPException(status_code=404, detail="Подписка не найдена") + ( + _tariff_name, + _subgroup_title, + _traffic_limit_gb, + _device_limit, + _panel, + is_tariff_configurable, + addons_devices_enabled, + addons_traffic_enabled, + ) = await get_key_tariff_addons_state(session=session, key_record=key_details, db_key=db_key) + if not is_tariff_configurable: + raise HTTPException(status_code=400, detail="Тариф не поддерживает доп. опции") + cfg = normalize_tariff_config(tariff) + raw_device_options = cfg.get("device_options") or tariff.get("device_options") or [] + raw_traffic_options = cfg.get("traffic_options_gb") or tariff.get("traffic_options_gb") or [] + device_options: list[int] = [] + for value in raw_device_options: + try: + device_options.append(int(value)) + except (TypeError, ValueError): + continue + traffic_options: list[int] = [] + for value in raw_traffic_options: + try: + traffic_options.append(int(value)) + except (TypeError, ValueError): + continue + device_options = sorted(set(device_options), key=lambda val: (int(val == 0), val)) + traffic_options = sorted(set(traffic_options), key=lambda val: (int(val == 0), val)) + has_device_option = bool(device_options) and bool(addons_devices_enabled) + has_traffic_option = bool(traffic_options) and bool(addons_traffic_enabled) + pack_devices, pack_traffic, pack_mode = get_pack_flags() + if pack_mode: + has_device_option = has_device_option and bool(pack_devices) + has_traffic_option = has_traffic_option and bool(pack_traffic) + if not has_device_option: + device_options = [] + if not has_traffic_option: + traffic_options = [] + if not has_device_option and not has_traffic_option: + raise HTTPException(status_code=400, detail="Доп. опции для этой подписки недоступны") + selected_device_limit_db = key_details.get("selected_device_limit") + selected_traffic_limit_db = key_details.get("selected_traffic_limit") + current_device_limit_db = key_details.get("current_device_limit") + current_traffic_limit_db = key_details.get("current_traffic_limit") + base_devices = tariff.get("device_limit") + base_devices = int(base_devices) if base_devices is not None else None + base_traffic_bytes = tariff.get("traffic_limit") + base_traffic_gb_from_tariff = int(base_traffic_bytes / GB) if base_traffic_bytes else None + current_device_limit = ( + int(current_device_limit_db) + if current_device_limit_db is not None + else (int(selected_device_limit_db) if selected_device_limit_db is not None else base_devices) + ) + current_traffic_gb = ( + int(current_traffic_limit_db) + if current_traffic_limit_db is not None + else (int(selected_traffic_limit_db) if selected_traffic_limit_db is not None else base_traffic_gb_from_tariff) + ) + if pack_mode and current_device_limit is not None and int(current_device_limit) == 0: + has_device_option = False + device_options = [] + if pack_mode and current_traffic_gb is not None and int(current_traffic_gb) == 0: + has_traffic_option = False + traffic_options = [] + if not has_device_option and not has_traffic_option: + raise HTTPException(status_code=400, detail="Доп. опции для этой подписки недоступны") + if pack_mode: + include_device_effective = ( + bool(body.include_device) if body.include_device is not None else body.selected_device_limit is not None + ) + include_traffic_effective = ( + bool(body.include_traffic) if body.include_traffic is not None else body.selected_traffic_gb is not None + ) + selected_device = ( + body.selected_device_limit + if body.selected_device_limit is not None + else None + ) + selected_traffic = ( + body.selected_traffic_gb + if body.selected_traffic_gb is not None + else None + ) + else: + include_device_effective = has_device_option + include_traffic_effective = has_traffic_option + selected_device = body.selected_device_limit if body.selected_device_limit is not None else current_device_limit + selected_traffic = body.selected_traffic_gb if body.selected_traffic_gb is not None else current_traffic_gb + if has_device_option and include_device_effective and selected_device is not None and int(selected_device) not in device_options: + raise HTTPException(status_code=400, detail="Выбранный пакет устройств недоступен") + if has_traffic_option and include_traffic_effective and selected_traffic is not None and int(selected_traffic) not in traffic_options: + raise HTTPException(status_code=400, detail="Выбранный пакет трафика недоступен") + current_devices_for_price = int(current_device_limit) if current_device_limit is not None else None + current_traffic_for_price = int(current_traffic_gb) if current_traffic_gb is not None else None + base_price_for_current = int( + calculate_config_price( + tariff=tariff, + selected_device_limit=current_devices_for_price, + selected_traffic_gb=current_traffic_for_price, + ) + ) + total_price_after_purchase = base_price_for_current + if pack_mode: + diff_full = int( + calc_pack_full_price_rub( + tariff=tariff, + has_device_option=bool(has_device_option and include_device_effective), + has_traffic_option=bool(has_traffic_option and include_traffic_effective), + selected_devices=int(selected_device) if has_device_option and include_device_effective and selected_device is not None else None, + selected_traffic_gb=int(selected_traffic) if has_traffic_option and include_traffic_effective and selected_traffic is not None else None, + ) + ) + recalc_enabled = bool( + MODES_CONFIG.get( + "KEY_ADDONS_RECALC_PRICE", + TARIFFS_CONFIG.get("KEY_ADDONS_RECALC_PRICE", False), + ) + ) + if recalc_enabled: + remaining_seconds, total_seconds = calc_remaining_ratio_seconds( + key_details.get("expiry_time"), + tariff, + ) + extra_price_rub = int((diff_full * remaining_seconds + total_seconds - 1) // total_seconds) + else: + extra_price_rub = diff_full + total_price_after_purchase = int(base_price_for_current + diff_full) + else: + selected_total_price = int( + calculate_config_price( + tariff=tariff, + selected_device_limit=int(selected_device) if has_device_option and selected_device is not None else None, + selected_traffic_gb=int(selected_traffic) if has_traffic_option and selected_traffic is not None else None, + ) + ) + extra_price_rub = int(max(0, selected_total_price - base_price_for_current)) + total_price_after_purchase = selected_total_price + allow_downgrade = bool(TARIFFS_CONFIG.get("ALLOW_DOWNGRADE", True)) + device_downgrade = ( + allow_downgrade + and has_device_option + and current_device_limit is not None + and selected_device is not None + and not is_not_downgrade(current_device_limit, selected_device) + ) + traffic_downgrade = ( + allow_downgrade + and has_traffic_option + and current_traffic_gb is not None + and selected_traffic is not None + and not is_not_downgrade(current_traffic_gb, selected_traffic) + ) + if device_downgrade or traffic_downgrade: + raise HTTPException(status_code=400, detail="Снижение параметров через сайт пока не поддерживается") + final_extra_price_rub, discount_rub, coupon_id, applied_coupon_code = await resolve_percent_coupon_pricing( + session=session, + billing_user_id=int(billing_user_id), + base_price_rub=int(max(0, extra_price_rub)), + coupon_code=body.coupon_code, + ) + balance = float(await get_balance(session, int(billing_user_id))) + required_amount = int(max(0, ceil(float(final_extra_price_rub) - balance))) + if extra_price_rub <= 0: + return AccountKeyApplyAddonsResponse( + ok=True, + message="Доплата не требуется", + client_id=str(getattr(db_key, "client_id", "") or ""), + tariff_id=int(tariff_id), + total_price_rub=int(total_price_after_purchase), + extra_price_rub=0, + discount_rub=0, + final_price_rub=0, + applied_coupon_code=None, + charged_rub=0, + balance_rub=balance, + ) + expiry_time = int(getattr(db_key, "expiry_time", 0) or 0) + email = str(getattr(db_key, "email", "") or "") + server_id = str(getattr(db_key, "server_id", "") or "") + if not email or not server_id: + raise HTTPException(status_code=400, detail="Некорректные данные подписки") + if required_amount > 0: + provider_id = str(body.provider_id or _resolve_default_web_payment_provider() or "").strip().upper() + if not provider_id: + raise HTTPException(status_code=503, detail="Нет доступных провайдеров оплаты") + base_url = _resolve_public_base_url(request) + success_url = validate_redirect_url(str(body.success_url or ""), f"{base_url}/payment-success") + failure_url = validate_redirect_url(str(body.failure_url or ""), f"{base_url}/payment-failure") + payment_request = PaymentLinkRequest( + legacy_user_ref=int(billing_user_id), + amount=required_amount, + currency="RUB", + provider_id=provider_id, + success_url=success_url, + failure_url=failure_url, + metadata={ + "payment_flow": "key_addons", + "tariff_id": int(tariff_id), + "email": email, + "selected_device_limit": int(selected_device) if has_device_option and include_device_effective and selected_device is not None else None, + "selected_traffic_gb": int(selected_traffic) if has_traffic_option and include_traffic_effective and selected_traffic is not None else None, + "current_device_limit": int(current_device_limit) if current_device_limit is not None else None, + "current_traffic_gb": int(current_traffic_gb) if current_traffic_gb is not None else None, + "original_price": int(base_price_for_current), + "base_price_rub": int(max(0, extra_price_rub)), + "discount_rub": int(discount_rub), + "applied_coupon_code": applied_coupon_code, + "coupon_id": int(coupon_id) if coupon_id is not None else None, + }, + ) + payment_result = await create_payment_link(session, payment_request) + if not payment_result.success or not payment_result.payment_url or not payment_result.payment_id: + raise HTTPException(status_code=400, detail=payment_result.error or "Не удалось создать ссылку оплаты") + await create_temporary_data( + session, + int(billing_user_id), + "waiting_for_addons_payment", + { + "tariff_id": int(tariff_id), + "email": email, + "required_amount": int(required_amount), + "selected_device_limit": int(selected_device) if has_device_option and include_device_effective and selected_device is not None else None, + "selected_traffic_gb": int(selected_traffic) if has_traffic_option and include_traffic_effective and selected_traffic is not None else None, + "current_device_limit": int(current_device_limit) if current_device_limit is not None else None, + "current_traffic_gb": int(current_traffic_gb) if current_traffic_gb is not None else None, + "original_price": int(base_price_for_current), + "base_price_rub": int(max(0, extra_price_rub)), + "discount_rub": int(discount_rub), + "applied_coupon_code": applied_coupon_code, + "coupon_id": int(coupon_id) if coupon_id is not None else None, + }, + ) + return AccountKeyApplyAddonsResponse( + ok=True, + message="Требуется оплата для применения доп. опций", + client_id=str(getattr(db_key, "client_id", "") or ""), + tariff_id=int(tariff_id), + total_price_rub=int(total_price_after_purchase), + extra_price_rub=int(extra_price_rub), + discount_rub=int(discount_rub), + final_price_rub=int(final_extra_price_rub), + applied_coupon_code=applied_coupon_code, + charged_rub=0, + balance_rub=balance, + payment_required=True, + required_amount_rub=required_amount, + payment_id=payment_result.payment_id, + payment_url=payment_result.payment_url, + ) + target_subgroup = tariff.get("subgroup_title") + current_subgroup = None + current_tariff_id = key_details.get("tariff_id") + if current_tariff_id: + current_tariff = await get_tariff_by_id(session, int(current_tariff_id)) + if current_tariff: + current_subgroup = current_tariff.get("subgroup_title") + if pack_mode: + device_limit_effective_current, traffic_limit_bytes_effective_current = await get_effective_limits_for_key( + session=session, + tariff_id=int(tariff_id), + selected_device_limit=int(current_device_limit) if current_device_limit is not None else None, + selected_traffic_gb=int(current_traffic_gb) if current_traffic_gb is not None else None, + ) + traffic_limit_gb_effective_current = ( + int(traffic_limit_bytes_effective_current / GB) if traffic_limit_bytes_effective_current else 0 + ) + new_device_limit_effective = device_limit_effective_current + new_traffic_limit_gb_effective = traffic_limit_gb_effective_current + if has_device_option and include_device_effective and selected_device is not None: + pack_devices_val = int(selected_device) + if pack_devices_val <= 0 or ( + new_device_limit_effective is not None and int(new_device_limit_effective) <= 0 + ): + new_device_limit_effective = 0 + else: + if new_device_limit_effective is None: + new_device_limit_effective = pack_devices_val + else: + new_device_limit_effective = int(new_device_limit_effective) + pack_devices_val + if has_traffic_option and include_traffic_effective and selected_traffic is not None: + pack_traffic_val = int(selected_traffic) + if pack_traffic_val <= 0 or int(new_traffic_limit_gb_effective) <= 0: + new_traffic_limit_gb_effective = 0 + else: + new_traffic_limit_gb_effective = int(new_traffic_limit_gb_effective) + pack_traffic_val + await renew_key_in_cluster( + cluster_id=server_id, + email=email, + client_id=str(getattr(db_key, "client_id", "") or ""), + new_expiry_time=expiry_time, + total_gb=int(new_traffic_limit_gb_effective), + session=session, + hwid_device_limit=int(new_device_limit_effective) if new_device_limit_effective is not None else 0, + reset_traffic=False, + target_subgroup=target_subgroup, + old_subgroup=current_subgroup, + plan=int(tariff_id), + ) + await save_key_config_with_mode( + session=session, + email=email, + selected_devices=int(new_device_limit_effective) if new_device_limit_effective is not None else None, + selected_traffic_gb=int(new_traffic_limit_gb_effective) if new_traffic_limit_gb_effective is not None else None, + total_price=int(total_price_after_purchase), + has_device_choice=bool(has_device_option and include_device_effective), + has_traffic_choice=bool(has_traffic_option and include_traffic_effective), + config_mode="pack", + ) + else: + selected_device_for_effective = int(selected_device) if has_device_option and include_device_effective and selected_device is not None else None + selected_traffic_for_effective = int(selected_traffic) if has_traffic_option and include_traffic_effective and selected_traffic is not None else 0 + device_limit_effective_new, traffic_limit_bytes_effective_new = await get_effective_limits_for_key( + session=session, + tariff_id=int(tariff_id), + selected_device_limit=selected_device_for_effective, + selected_traffic_gb=selected_traffic_for_effective, + ) + traffic_limit_gb_effective = int(traffic_limit_bytes_effective_new / GB) if traffic_limit_bytes_effective_new else 0 + await renew_key_in_cluster( + cluster_id=server_id, + email=email, + client_id=str(getattr(db_key, "client_id", "") or ""), + new_expiry_time=expiry_time, + total_gb=int(traffic_limit_gb_effective), + session=session, + hwid_device_limit=int(device_limit_effective_new) if device_limit_effective_new is not None else 0, + reset_traffic=False, + target_subgroup=target_subgroup, + old_subgroup=current_subgroup, + plan=int(tariff_id), + ) + await save_key_config_with_mode( + session=session, + email=email, + selected_devices=int(selected_device) if has_device_option and include_device_effective and selected_device is not None else None, + selected_traffic_gb=int(selected_traffic) if has_traffic_option and include_traffic_effective and selected_traffic is not None else None, + total_price=int(total_price_after_purchase), + has_device_choice=bool(has_device_option and include_device_effective), + has_traffic_choice=bool(has_traffic_option and include_traffic_effective), + config_mode="addon", + ) + await update_balance(session, int(billing_user_id), -int(final_extra_price_rub)) + if coupon_id is not None: + await mark_coupon_used(session, int(coupon_id), int(billing_user_id)) + await session.commit() + return AccountKeyApplyAddonsResponse( + ok=True, + message="Доп. опции применены", + client_id=str(getattr(db_key, "client_id", "") or ""), + tariff_id=int(tariff_id), + total_price_rub=int(total_price_after_purchase), + extra_price_rub=int(extra_price_rub), + discount_rub=int(discount_rub), + final_price_rub=int(final_extra_price_rub), + applied_coupon_code=applied_coupon_code, + charged_rub=int(final_extra_price_rub), + balance_rub=float(await get_balance(session, int(billing_user_id))), + ) diff --git a/api/v2/routes/keys/user/core.py b/api/v2/routes/keys/user/core.py new file mode 100644 index 00000000..3aaff74e --- /dev/null +++ b/api/v2/routes/keys/user/core.py @@ -0,0 +1,249 @@ +"""User-facing key endpoints (/api/keys/*). + +Регистрирует эндпоинты на ``user_router`` из ``_common``. Импорт этого модуля +из ``__init__.py`` запускает регистрацию декораторов. +""" + +from .._common import * # noqa: F401,F403 — подтягиваем все имена для endpoints +from .._common import ( + _key_actions_config, + _resolve_available_location_servers, + _resolve_billing_user_id, + _resolve_default_web_payment_provider, + _resolve_public_base_url, + _normalize_expiry_ms, + router, + user_router, +) + + +@user_router.get("", response_model=list[AccountKeyResponse]) +async def user_keys( + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + billing_user_id = await _resolve_billing_user_id(request, identity, session) + keys = await get_keys(session, billing_user_id) + result: list[AccountKeyResponse] = [] + for key in keys: + key_actions = AccountKeyActionsAvailability() + try: + key_ref = str(getattr(key, "client_id", "") or getattr(key, "email", "") or "") + _, markup, _ = await build_key_view_payload(session, int(billing_user_id), key_ref) + key_actions = _extract_key_actions_from_markup(markup) + except Exception: + key_actions = AccountKeyActionsAvailability() + result.append( + AccountKeyResponse( + email=str(getattr(key, "email", "") or ""), + alias=getattr(key, "alias", None), + client_id=str(getattr(key, "client_id", "") or ""), + tariff_id=getattr(key, "tariff_id", None), + server_id=str(getattr(key, "server_id", "") or ""), + created_at=int(getattr(key, "created_at", 0) or 0), + expiry_time=int(getattr(key, "expiry_time", 0) or 0), + key=getattr(key, "key", None), + remnawave_link=getattr(key, "remnawave_link", None), + is_frozen=bool(getattr(key, "is_frozen", False)), + actions=key_actions, + ) + ) + return result + + +@user_router.get("/actions-config", response_model=AccountKeyActionsConfigResponse) +async def user_keys_actions_config( + identity=Depends(verify_identity_token), +): + _ = identity + return _key_actions_config() + + +@user_router.get("/{client_id}/details", response_model=AccountKeyDetailsResponse) +async def user_key_details( + client_id: str, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + key_details = await get_key_details(session, str(getattr(db_key, "email", "") or "")) + if not key_details: + raise HTTPException(status_code=404, detail="Подписка не найдена") + tariff_name = "" + subgroup_title = "" + traffic_limit_gb = 0 + device_limit = 0 + is_tariff_configurable = False + addons_devices_enabled = False + addons_traffic_enabled = False + ( + tariff_name, + subgroup_title, + traffic_limit_gb, + device_limit, + _, + is_tariff_configurable, + addons_devices_enabled, + addons_traffic_enabled, + ) = await get_key_tariff_addons_state( + session=session, + key_record=key_details, + db_key=db_key, + ) + connected_devices = 0 + used_traffic_gb = None + try: + profile = await get_remnawave_profile( + session, + str(getattr(db_key, "server_id", "") or ""), + client_id, + fallback_any=True, + ) + if profile: + connected_devices = int(profile.get("hwid_count") or 0) + used_raw = profile.get("used_gb") + used_traffic_gb = float(used_raw) if used_raw is not None else None + traffic_limit_bytes_actual = profile.get("traffic_limit_bytes") + if traffic_limit_bytes_actual is not None: + try: + traffic_limit_bytes_actual = int(traffic_limit_bytes_actual) + traffic_limit_gb = int(traffic_limit_bytes_actual / GB) if traffic_limit_bytes_actual > 0 else 0 + except (TypeError, ValueError): + pass + except Exception: + connected_devices = 0 + used_traffic_gb = None + return AccountKeyDetailsResponse( + client_id=str(getattr(db_key, "client_id", "") or ""), + email=str(getattr(db_key, "email", "") or ""), + alias=getattr(db_key, "alias", None), + expiry_time=int(getattr(db_key, "expiry_time", 0) or 0), + is_frozen=bool(getattr(db_key, "is_frozen", False)), + tariff_name=str(tariff_name or ""), + subgroup_title=str(subgroup_title or ""), + traffic_limit_gb=int(traffic_limit_gb or 0), + used_traffic_gb=used_traffic_gb, + device_limit=int(device_limit or 0), + connected_devices=int(connected_devices or 0), + is_tariff_configurable=bool(is_tariff_configurable), + addons_devices_enabled=bool(addons_devices_enabled), + addons_traffic_enabled=bool(addons_traffic_enabled), + ) + + +@user_router.get("/{client_id}/qr", response_model=AccountKeyQrResponse) +async def user_key_qr( + client_id: str, + request: Request, + force_web: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + actions = _key_actions_config() + if not force_web and not actions.qr_enabled: + raise HTTPException(status_code=403, detail="QR для подписок отключен в настройках") + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + qr_data = str(getattr(db_key, "key", "") or "").strip() or str(getattr(db_key, "remnawave_link", "") or "").strip() + if not qr_data: + raise HTTPException(status_code=400, detail="Ссылка для подключения отсутствует") + qr = qrcode.QRCode(version=1, box_size=10, border=4) + qr.add_data(qr_data) + qr.make(fit=True) + img = qr.make_image(fill_color="black", back_color="white") + buffer = BytesIO() + img.save(buffer, format="PNG") + image_data = b64encode(buffer.getvalue()).decode("ascii") + return AccountKeyQrResponse( + ok=True, + message="QR-код готов", + link=qr_data, + image_data_url=f"data:image/png;base64,{image_data}", + ) + + +@user_router.patch("/{client_id}/alias", response_model=AccountKeyResponse) +async def user_key_update_alias( + client_id: str, + body: AccountKeyAliasUpdateRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + alias = str(body.alias or "").strip() + if not alias: + raise HTTPException(status_code=400, detail="Укажите alias") + if len(alias) > 10: + raise HTTPException(status_code=400, detail="Alias должен быть не длиннее 10 символов") + if not re.match(r"^[a-zA-Zа-яА-ЯёЁ0-9@._-]+$", alias): + raise HTTPException(status_code=400, detail="Alias содержит недопустимые символы") + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + db_key.alias = alias + await session.commit() + return AccountKeyResponse( + email=str(getattr(db_key, "email", "") or ""), + alias=getattr(db_key, "alias", None), + client_id=str(getattr(db_key, "client_id", "") or ""), + tariff_id=getattr(db_key, "tariff_id", None), + server_id=str(getattr(db_key, "server_id", "") or ""), + created_at=int(getattr(db_key, "created_at", 0) or 0), + expiry_time=int(getattr(db_key, "expiry_time", 0) or 0), + key=getattr(db_key, "key", None), + remnawave_link=getattr(db_key, "remnawave_link", None), + is_frozen=bool(getattr(db_key, "is_frozen", False)), + ) + + +@user_router.delete("/{client_id}", response_model=AccountKeyActionResponse) +async def user_key_delete( + client_id: str, + request: Request, + force_web: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + actions = _key_actions_config() + if not force_web and not actions.delete_enabled: + raise HTTPException(status_code=403, detail="Удаление подписки отключено в настройках") + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + cluster_id = str(getattr(db_key, "server_id", "") or "") + email = str(getattr(db_key, "email", "") or "") + if cluster_id and email: + await delete_key_from_cluster( + cluster_id=cluster_id, + email=email, + client_id=client_id, + session=session, + ) + await session.delete(db_key) + await session.commit() + return AccountKeyActionResponse(ok=True, message="Подписка удалена") diff --git a/api/v2/routes/keys/user/hwid.py b/api/v2/routes/keys/user/hwid.py new file mode 100644 index 00000000..574414c3 --- /dev/null +++ b/api/v2/routes/keys/user/hwid.py @@ -0,0 +1,75 @@ +"""User-facing key endpoints (/api/keys/*). + +Регистрирует эндпоинты на ``user_router`` из ``_common``. Импорт этого модуля +из ``__init__.py`` запускает регистрацию декораторов. +""" + +from .._common import * # noqa: F401,F403 — подтягиваем все имена для endpoints +from .._common import ( + _key_actions_config, + _resolve_available_location_servers, + _resolve_billing_user_id, + _resolve_default_web_payment_provider, + _resolve_public_base_url, + _normalize_expiry_ms, + router, + user_router, +) + + +@user_router.post("/{client_id}/reset-hwid", response_model=AccountKeyResetHwidResponse) +async def user_key_reset_hwid( + client_id: str, + request: Request, + force_web: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + actions = _key_actions_config() + if not force_web and not actions.hwid_reset_enabled: + raise HTTPException(status_code=403, detail="Сброс устройств отключен в настройках") + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + server_id = str(getattr(db_key, "server_id", "") or "") + if not server_id: + raise HTTPException(status_code=400, detail="У подписки не указан сервер") + + async def _reset_devices(api): + devices = await api.get_user_hwid_devices(client_id) + if not devices: + return 0, 0 + reset_local = 0 + for device in devices: + hwid = device.get("hwid") + if hwid and await api.delete_user_hwid_device(client_id, hwid): + reset_local += 1 + return len(devices), reset_local + + reset_result = await with_remnawave_api( + session, + server_id, + _reset_devices, + fallback_any=True, + timeout_sec=12.0, + ) + if reset_result is None: + raise HTTPException(status_code=502, detail="Не удалось выполнить сброс устройств") + total_devices, reset_devices = reset_result + await invalidate_remnawave_profile( + session, + server_id, + str(client_id), + fallback_any=True, + ) + return AccountKeyResetHwidResponse( + ok=True, + message="Устройства сброшены" if total_devices > 0 else "Устройства не были привязаны", + total_devices=int(total_devices), + reset_devices=int(reset_devices), + ) diff --git a/api/v2/routes/keys/user/location.py b/api/v2/routes/keys/user/location.py new file mode 100644 index 00000000..5af17b50 --- /dev/null +++ b/api/v2/routes/keys/user/location.py @@ -0,0 +1,230 @@ +"""User-facing key endpoints (/api/keys/*). + +Регистрирует эндпоинты на ``user_router`` из ``_common``. Импорт этого модуля +из ``__init__.py`` запускает регистрацию декораторов. +""" + +from .._common import * # noqa: F401,F403 — подтягиваем все имена для endpoints +from .._common import ( + _key_actions_config, + _resolve_available_location_servers, + _resolve_billing_user_id, + _resolve_default_web_payment_provider, + _resolve_public_base_url, + _normalize_expiry_ms, + router, + user_router, +) + + +@user_router.get("/{client_id}/locations", response_model=AccountKeyLocationsResponse) +async def user_key_locations( + client_id: str, + request: Request, + force_web: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + actions = _key_actions_config() + if not force_web and not actions.country_change_enabled: + raise HTTPException(status_code=403, detail="Смена локации отключена в настройках") + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + names = await _resolve_available_location_servers(session, db_key) + return AccountKeyLocationsResponse( + client_id=str(getattr(db_key, "client_id", "") or ""), + current_server=str(getattr(db_key, "server_id", "") or ""), + locations=[AccountKeyLocationOptionResponse(server_name=name) for name in names], + ) + + +@user_router.post("/{client_id}/change-location", response_model=AccountKeyChangeLocationResponse) +async def user_key_change_location( + client_id: str, + body: AccountKeyChangeLocationRequest, + request: Request, + force_web: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + actions = _key_actions_config() + if not force_web and not actions.country_change_enabled: + raise HTTPException(status_code=403, detail="Смена локации отключена в настройках") + target_server = str(body.server_name or "").strip() + if not target_server: + raise HTTPException(status_code=400, detail="Укажите целевую локацию") + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + current_server = str(getattr(db_key, "server_id", "") or "") + if not current_server: + raise HTTPException(status_code=400, detail="У подписки не указан текущий сервер") + if current_server == target_server: + raise HTTPException(status_code=400, detail="Подписка уже в этой локации") + available_names = await _resolve_available_location_servers(session, db_key) + if target_server not in available_names: + raise HTTPException(status_code=400, detail="Выбранная локация недоступна") + email = str(getattr(db_key, "email", "") or "") + if not email: + raise HTTPException(status_code=400, detail="У подписки отсутствует email") + key_details = await get_key_details(session, email) + if not key_details: + raise HTTPException(status_code=404, detail="Подписка не найдена") + old_server_info = ( + await session.execute(select(Server).where(Server.server_name == current_server).limit(1)) + ).scalar_one_or_none() + if old_server_info: + old_panel_type = str(getattr(old_server_info, "panel_type", "") or "").lower() + try: + if old_panel_type == "3x-ui": + xui = await get_xui_instance(str(getattr(old_server_info, "api_url", "") or "")) + await delete_client( + xui, + int(getattr(old_server_info, "inbound_id", 0) or 0), + email, + str(getattr(db_key, "client_id", "") or ""), + ) + elif old_panel_type == "remnawave": + remna_del = RemnawaveAPI(str(getattr(old_server_info, "api_url", "") or "")) + if await remna_del.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD): + await remna_del.delete_user(str(getattr(db_key, "client_id", "") or "")) + except Exception: + pass + target_server_info = ( + await session.execute(select(Server).where(Server.server_name == target_server).limit(1)) + ).scalar_one_or_none() + if target_server_info is None: + raise HTTPException(status_code=404, detail="Целевая локация не найдена") + tariff_id = getattr(db_key, "tariff_id", None) + tariff = await get_tariff_by_id(session, int(tariff_id)) if tariff_id else None + need_vless_key = bool(tariff.get("vless")) if tariff else False + external_squad_uuid = (tariff.get("external_squad") if tariff else None) or None + selected_traffic_gb = getattr(db_key, "selected_traffic_limit", None) + selected_device_limit = getattr(db_key, "selected_device_limit", None) + if selected_traffic_gb is not None: + traffic_limit_bytes = int(selected_traffic_gb) * GB + else: + raw_traffic_limit = int(tariff.get("traffic_limit") or 0) if tariff else 0 + traffic_limit_bytes = raw_traffic_limit * GB if raw_traffic_limit > 0 else 0 + if selected_device_limit is not None: + device_limit = int(selected_device_limit) + else: + device_limit = int(tariff.get("device_limit") or 0) if tariff else 0 + key_client_id = str(getattr(db_key, "client_id", "") or "") + expiry_timestamp = int(getattr(db_key, "expiry_time", 0) or 0) + target_cluster_info = await check_server_name_by_cluster(session, target_server) + target_cluster_name = str((target_cluster_info or {}).get("cluster_name") or "") + full_remnawave_cluster = ( + await is_full_remnawave_cluster(target_cluster_name, session) if target_cluster_name else False + ) + panel_type = str(getattr(target_server_info, "panel_type", "") or "").lower() + remnawave_link = None + if panel_type == "remnawave" or full_remnawave_cluster: + remna = RemnawaveAPI(str(getattr(target_server_info, "api_url", "") or "")) + if not await remna.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD): + raise HTTPException(status_code=502, detail="Не удалось авторизоваться в Remnawave") + expire_at = datetime.utcfromtimestamp(expiry_timestamp / 1000).isoformat() + "Z" + user_data: dict[str, Any] = { + "username": email, + "trafficLimitStrategy": "NO_RESET", + "expireAt": expire_at, + "telegramId": int(key_details.get("tg_id") or 0), + "activeInternalSquads": [getattr(target_server_info, "inbound_id", None)], + "uuid": key_client_id, + } + if traffic_limit_bytes > 0: + user_data["trafficLimitBytes"] = traffic_limit_bytes + if device_limit > 0: + user_data["hwidDeviceLimit"] = device_limit + if external_squad_uuid: + user_data["externalSquadUuid"] = external_squad_uuid + result = await remna.create_user(user_data) + if not result: + raise HTTPException(status_code=502, detail="Не удалось создать подписку в новой локации") + key_client_id = str(result.get("uuid") or result.get("id") or key_client_id) + if need_vless_key: + try: + remnawave_link = await get_vless_link_for_remnawave_by_username(remna, email, email) + except Exception: + remnawave_link = None + if not remnawave_link: + try: + sub = await remna.get_subscription_by_username(email) + except Exception: + sub = None + if sub: + links = sub.get("links") or [] + remnawave_link = ( + next( + (link for link in links if isinstance(link, str) and link.lower().startswith("vless://")), + None, + ) + if need_vless_key + else None + ) + if not remnawave_link: + remnawave_link = sub.get("subscriptionUrl") + if panel_type == "3x-ui": + await create_client_on_server( + { + "api_url": str(getattr(target_server_info, "api_url", "") or ""), + "inbound_id": getattr(target_server_info, "inbound_id", None), + "server_name": str(getattr(target_server_info, "server_name", "") or ""), + "panel_type": str(getattr(target_server_info, "panel_type", "") or ""), + }, + int(key_details.get("tg_id") or 0), + key_client_id, + email, + expiry_timestamp, + asyncio.Semaphore(1), + plan=int(tariff_id) if tariff_id else None, + session=session, + is_trial=False, + total_traffic_limit_bytes=traffic_limit_bytes, + device_limit_value=device_limit, + ) + subgroup_code = tariff.get("subgroup_title") if tariff and tariff.get("subgroup_title") else None + public_link = await make_aggregated_link( + session=session, + cluster_all=[ + { + "server_name": str(getattr(target_server_info, "server_name", "") or ""), + "api_url": str(getattr(target_server_info, "api_url", "") or ""), + "panel_type": str(getattr(target_server_info, "panel_type", "") or ""), + "inbound_id": getattr(target_server_info, "inbound_id", None), + "enabled": True, + "max_keys": getattr(target_server_info, "max_keys", None), + } + ], + cluster_id=target_cluster_name or target_server, + email=email, + client_id=key_client_id, + tg_id=int(key_details.get("tg_id") or 0), + subgroup_code=subgroup_code, + remna_link_override=remnawave_link, + plan=int(tariff_id) if tariff_id else None, + ) + db_key.server_id = target_server + db_key.client_id = key_client_id + db_key.key = public_link if isinstance(public_link, str) and public_link.strip() else None + db_key.remnawave_link = remnawave_link + await session.commit() + return AccountKeyChangeLocationResponse( + ok=True, + message="Локация успешно изменена", + client_id=str(getattr(db_key, "client_id", "") or ""), + server_id=str(getattr(db_key, "server_id", "") or ""), + link=str(getattr(db_key, "key", "") or ""), + remnawave_link=getattr(db_key, "remnawave_link", None), + ) diff --git a/api/v2/routes/keys/user/renew.py b/api/v2/routes/keys/user/renew.py new file mode 100644 index 00000000..87d39a6f --- /dev/null +++ b/api/v2/routes/keys/user/renew.py @@ -0,0 +1,191 @@ +"""User-facing key endpoints (/api/keys/*). + +Регистрирует эндпоинты на ``user_router`` из ``_common``. Импорт этого модуля +из ``__init__.py`` запускает регистрацию декораторов. +""" + +from .._common import * # noqa: F401,F403 — подтягиваем все имена для endpoints +from .._common import ( + _key_actions_config, + _resolve_available_location_servers, + _resolve_billing_user_id, + _resolve_default_web_payment_provider, + _resolve_public_base_url, + _normalize_expiry_ms, + router, + user_router, +) + + +@user_router.post("/{client_id}/renew", response_model=AccountKeyRenewResponse) +async def user_key_renew( + client_id: str, + body: AccountKeyRenewRequest, + request: Request, + force_web: bool = Query(False), + preview: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + from services.errors import ServiceError + from services.keys import ( + calculate_renewal_pricing, + execute_renewal, + normalize_expiry_ms as _svc_normalize_expiry, + ) + + actions = _key_actions_config() + if not force_web and not actions.renew_enabled: + raise HTTPException(status_code=403, detail="Продление подписки отключено в настройках") + billing_user_id = await _resolve_billing_user_id(request, identity, session) + db_key = ( + await session.execute( + select(Key).where(Key.user_id == billing_user_id, Key.client_id == client_id).limit(1) + ) + ).scalar_one_or_none() + if db_key is None: + raise HTTPException(status_code=404, detail="Подписка не найдена") + if bool(getattr(db_key, "is_frozen", False)): + raise HTTPException(status_code=400, detail="Продление для замороженной подписки недоступно") + tariff_id = getattr(db_key, "tariff_id", None) + if not tariff_id: + raise HTTPException(status_code=400, detail="Для подписки не назначен тариф") + key_email = str(getattr(db_key, "email", "") or "") + key_server_id = str(getattr(db_key, "server_id", "") or "") + + try: + pricing = await calculate_renewal_pricing( + session=session, + billing_user_id=int(billing_user_id), + key_email=key_email, + tariff_id=int(tariff_id), + coupon_code=body.coupon_code, + ) + except ServiceError as e: + raise HTTPException(status_code=400, detail=e.message) + + if preview: + return AccountKeyRenewResponse( + ok=True, + message="Расчет обновлен", + client_id=str(client_id), + tariff_id=int(tariff_id), + charged_rub=0, + balance_rub=pricing.balance, + base_price_rub=pricing.base_price_rub, + discount_rub=pricing.discount_rub, + final_price_rub=pricing.final_price_rub, + applied_coupon_code=pricing.applied_coupon_code, + payment_required=pricing.payment_required, + required_amount_rub=pricing.required_amount, + payment_id=None, + payment_url=None, + ) + if pricing.payment_required: + provider_id = str(body.provider_id or _resolve_default_web_payment_provider() or "").strip().upper() + if not provider_id: + raise HTTPException(status_code=503, detail="Нет доступных провайдеров оплаты") + base_url = _resolve_public_base_url(request) + success_url = validate_redirect_url(str(body.success_url or ""), f"{base_url}/payment-success") + failure_url = validate_redirect_url(str(body.failure_url or ""), f"{base_url}/payment-failure") + payment_request = PaymentLinkRequest( + legacy_user_ref=int(billing_user_id), + amount=pricing.required_amount, + currency="RUB", + provider_id=provider_id, + success_url=success_url, + failure_url=failure_url, + metadata={ + "payment_flow": "key_renewal", + "tariff_id": int(tariff_id), + "client_id": str(client_id), + "email": key_email, + "cost": pricing.final_price_rub, + "selected_duration_days": pricing.duration_days, + "selected_device_limit": pricing.selected_device_limit, + "selected_traffic_limit": pricing.selected_traffic_limit, + "selected_price_rub": pricing.final_price_rub, + "total_gb": pricing.total_gb, + "base_price_rub": pricing.base_price_rub, + "discount_rub": pricing.discount_rub, + "applied_coupon_code": pricing.applied_coupon_code, + "coupon_id": pricing.coupon_id, + }, + ) + payment_result = await create_payment_link(session, payment_request) + if not payment_result.success or not payment_result.payment_url or not payment_result.payment_id: + raise HTTPException(status_code=400, detail=payment_result.error or "Не удалось создать ссылку оплаты") + await create_temporary_data( + session, + int(billing_user_id), + "waiting_for_renewal_payment", + { + "tariff_id": int(tariff_id), + "client_id": str(client_id), + "email": key_email, + "cost": pricing.final_price_rub, + "required_amount": pricing.required_amount, + "selected_duration_days": pricing.duration_days, + "selected_device_limit": pricing.selected_device_limit, + "selected_traffic_limit": pricing.selected_traffic_limit, + "selected_price_rub": pricing.final_price_rub, + "total_gb": pricing.total_gb, + "base_price_rub": pricing.base_price_rub, + "discount_rub": pricing.discount_rub, + "applied_coupon_code": pricing.applied_coupon_code, + "coupon_id": pricing.coupon_id, + }, + ) + return AccountKeyRenewResponse( + ok=True, + message="Требуется оплата для продления подписки", + client_id=str(client_id), + tariff_id=int(tariff_id), + charged_rub=0, + balance_rub=pricing.balance, + base_price_rub=pricing.base_price_rub, + discount_rub=pricing.discount_rub, + final_price_rub=pricing.final_price_rub, + applied_coupon_code=pricing.applied_coupon_code, + payment_required=True, + required_amount_rub=pricing.required_amount, + payment_id=payment_result.payment_id, + payment_url=payment_result.payment_url, + ) + expiry_raw = _normalize_expiry_ms(getattr(db_key, "expiry_time", None)) + now_ms = int(datetime.utcnow().timestamp() * 1000) + base_expiry = now_ms if expiry_raw <= now_ms else expiry_raw + new_expiry_time = int(base_expiry + pricing.duration_days * 24 * 60 * 60 * 1000) + if not key_email or not key_server_id: + raise HTTPException(status_code=400, detail="Некорректные данные подписки") + try: + result = await execute_renewal( + session=session, + billing_user_id=int(billing_user_id), + client_id=str(client_id), + key_email=key_email, + key_server_id=key_server_id, + tariff_id=int(tariff_id), + new_expiry_time=new_expiry_time, + total_gb=pricing.total_gb, + cost=float(pricing.final_price_rub), + selected_device_limit=pricing.selected_device_limit, + selected_traffic_limit=pricing.selected_traffic_limit, + selected_price_rub=pricing.final_price_rub, + coupon_id=pricing.coupon_id, + ) + except ServiceError as e: + raise HTTPException(status_code=400, detail=e.message) + await session.commit() + return AccountKeyRenewResponse( + ok=True, + message="Подписка продлена", + client_id=result.client_id, + tariff_id=result.tariff_id, + charged_rub=result.charged_rub, + balance_rub=result.balance_rub, + base_price_rub=pricing.base_price_rub, + discount_rub=pricing.discount_rub, + final_price_rub=pricing.final_price_rub, + applied_coupon_code=pricing.applied_coupon_code, + ) diff --git a/api/v2/routes/management.py b/api/v2/routes/management.py index 91da90d5..44289062 100644 --- a/api/v2/routes/management.py +++ b/api/v2/routes/management.py @@ -3,10 +3,12 @@ import os import re import subprocess import sys -from datetime import datetime, timezone, timedelta + +from datetime import datetime, timedelta, timezone from typing import Literal import psutil + from aiogram import Bot from aiogram.client.default import DefaultBotProperties from aiogram.enums import ParseMode @@ -23,10 +25,11 @@ from api.v2.schemas.audit import ( ) from audit import drain_audit_redis_to_db, get_audit_funnel, get_audit_stats, list_audit_events from config import API_TOKEN, BOT_SERVICE -from database import async_session_maker from core.bootstrap import MANAGEMENT_CONFIG from core.executor import run_io +from core.redis_cache import cache_incr from core.settings.management_config import update_management_config +from database import async_session_maker from database.models import Key, ScheduledBroadcast, Server, User from database.scheduled_broadcasts import ( cancel_scheduled_broadcast, @@ -48,9 +51,18 @@ from handlers.admin.sender.scheduled_service import ( from logger import logger from utils.backup import backup_database + router = APIRouter() +async def _admin_rate_limit(request_or_identity, action: str, max_calls: int, window_sec: int) -> None: + identity_id = getattr(request_or_identity, "id", "unknown") + key = f"admin_rl:{action}:{identity_id}" + count = await cache_incr(key, window_sec) + if count > max_calls: + raise HTTPException(status_code=429, detail="Слишком много запросов. Попробуйте позже.") + + class MaintenanceUpdate(BaseModel): enabled: bool @@ -174,6 +186,7 @@ async def restart_bot( identity=Depends(verify_identity_admin), ): """Запуск перезапуска бота в фоне.""" + await _admin_rate_limit(identity, "restart", max_calls=3, window_sec=60) background.add_task(_restart_bot) return {"status": "restarting"} @@ -185,6 +198,7 @@ async def change_domain( session: AsyncSession = Depends(get_session), ): """Массовая замена домена в ключах и remnawave_link.""" + await _admin_rate_limit(identity, "change_domain", max_calls=3, window_sec=300) domain = payload.domain.strip() if not domain or " " in domain or not re.fullmatch(r"[a-zA-Z0-9.-]+", domain): raise HTTPException(status_code=400, detail="Invalid domain") @@ -211,11 +225,12 @@ async def restore_trials( session: AsyncSession = Depends(get_session), ): """Сбрасывает trial=0 у пользователей без ключей.""" + await _admin_rate_limit(identity, "restore_trials", max_calls=3, window_sec=300) stmt = ( update(User) .where( User.trial == 1, - ~exists(select(Key.tg_id).where(Key.tg_id == User.tg_id)), + ~exists(select(Key.user_id).where(Key.user_id == User.id)), ) .values(trial=0) ) @@ -227,6 +242,7 @@ async def restore_trials( @router.post("/backup") async def trigger_backup(identity=Depends(verify_identity_admin)): """Запуск бэкапа БД в фоне.""" + await _admin_rate_limit(identity, "backup", max_calls=2, window_sec=300) async def _run_backup() -> None: exception = await backup_database() @@ -357,7 +373,7 @@ async def post_audit_drain(identity=Depends(verify_identity_admin_short)): return {"success": True, "drained": count} except Exception as exc: logger.warning("audit-drain failed: {}", exc) - raise HTTPException(status_code=500, detail=str(exc)) from exc + raise HTTPException(status_code=500, detail="Внутренняя ошибка при дренаже аудита") from exc @router.post("/broadcast") @@ -366,6 +382,7 @@ async def launch_broadcast( identity=Depends(verify_identity_admin_short), ): """Запуск рассылки по выбранной аудитории. Сессия БД не держится на время рассылки.""" + await _admin_rate_limit(identity, "broadcast", max_calls=5, window_sec=300) try: prepared = prepare_broadcast_payload( send_to=payload.send_to, @@ -475,14 +492,16 @@ async def send_broadcast_schedule_now( identity=Depends(verify_identity_admin_short), session: AsyncSession = Depends(get_session), ): + await _admin_rate_limit(identity, "broadcast_now", max_calls=5, window_sec=300) item = await start_scheduled_broadcast(session, broadcast_id) if item is None: raise HTTPException(status_code=409, detail="Scheduled broadcast can no longer be sent now") try: result = await execute_scheduled_broadcast(item, bot=_get_broadcast_bot()) except Exception as exc: + logger.error("[Broadcast] send-now failed for {}: {}", broadcast_id, exc) await mark_scheduled_broadcast_failed(session, broadcast_id, str(exc)) - raise HTTPException(status_code=500, detail=str(exc)) from exc + raise HTTPException(status_code=500, detail="Ошибка при выполнении рассылки") from exc if result.get("success"): item = await mark_scheduled_broadcast_sent(session, broadcast_id, result) else: diff --git a/api/v2/routes/misc.py b/api/v2/routes/misc.py index b97d2d50..f6a9b29b 100644 --- a/api/v2/routes/misc.py +++ b/api/v2/routes/misc.py @@ -13,6 +13,7 @@ from api.v2.schemas import ( ) from api.v2.base_crud import generate_crud_router from database import get_tracking_source_stats +from database.access.resolution import resolve_user_optional from database.models import ( BlockedUser, ManualBan, @@ -46,7 +47,10 @@ async def get_payments_by_tg_id( session: AsyncSession = Depends(get_session), ): """Список платежей по tg_id пользователя.""" - result = await session.execute(select(Payment).where(Payment.tg_id == tg_id)) + u = await resolve_user_optional(session, tg_id) + if u is None: + raise HTTPException(status_code=404, detail="Payments not found") + result = await session.execute(select(Payment).where(Payment.user_id == u.id)) payments = result.scalars().all() if not payments: raise HTTPException(status_code=404, detail="Payments not found") @@ -59,7 +63,9 @@ router.include_router( schema_response=NotificationResponse, schema_create=None, schema_update=None, - identifier_field="tg_id", + identifier_field="user_id", + parameter_name="tg_id", + telegram_path_to_user_id=True, enabled_methods=["get_all", "get_one", "delete"], ), prefix="/notifications", @@ -73,7 +79,9 @@ router.include_router( schema_response=ManualBanResponse, schema_create=None, schema_update=None, - identifier_field="tg_id", + identifier_field="user_id", + parameter_name="tg_id", + telegram_path_to_user_id=True, enabled_methods=["get_all", "get_one", "delete"], ), prefix="/manual-bans", @@ -87,7 +95,9 @@ router.include_router( schema_response=BlockedUserResponse, schema_create=None, schema_update=None, - identifier_field="tg_id", + identifier_field="user_id", + parameter_name="tg_id", + telegram_path_to_user_id=True, enabled_methods=["get_all", "get_one", "delete"], ), prefix="/blocked-users", @@ -101,7 +111,9 @@ router.include_router( schema_response=TemporaryDataResponse, schema_create=None, schema_update=None, - identifier_field="tg_id", + identifier_field="user_id", + parameter_name="tg_id", + telegram_path_to_user_id=True, enabled_methods=["get_all", "get_one", "delete"], ), prefix="/temporary-data", diff --git a/api/v2/routes/modules.py b/api/v2/routes/modules.py index 5f24035f..0248380e 100644 --- a/api/v2/routes/modules.py +++ b/api/v2/routes/modules.py @@ -1,4 +1,5 @@ import pkgutil + from pathlib import Path from typing import Literal @@ -7,9 +8,11 @@ from pydantic import BaseModel from api.depends import verify_identity_admin from core.executor import run_io +from logger import logger from utils.modules_loader import _is_safe_module_name from utils.modules_manager import manager + router = APIRouter(prefix="/modules", tags=["Modules"]) MODULES_DIR = Path(__file__).resolve().parents[3] / "modules" @@ -117,5 +120,6 @@ async def control_module(module_name: str, payload: ModuleAction, identity=Depen except ValueError as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc except RuntimeError as exc: - raise HTTPException(status_code=500, detail=str(exc)) from exc + logger.error("[Modules] action failed for {}: {}", name, exc) + raise HTTPException(status_code=500, detail="Ошибка при выполнении операции модуля") from exc return {"item": _module_state(name)} diff --git a/api/v2/routes/notifications.py b/api/v2/routes/notifications.py new file mode 100644 index 00000000..d23ff1cc --- /dev/null +++ b/api/v2/routes/notifications.py @@ -0,0 +1,83 @@ +from fastapi import APIRouter, Depends, Query +from pydantic import BaseModel +from sqlalchemy.ext.asyncio import AsyncSession + +from api.depends import get_session, verify_identity_token +from database.models import Identity +from database import web_notifications as wn_db + +router = APIRouter() + + +class PushSubscribeRequest(BaseModel): + endpoint: str + keys: dict + + +class NotificationItem(BaseModel): + id: str + type: str + title: str + message: str + read: bool + created_at: str + data: dict | None = None + + +class NotificationsResponse(BaseModel): + ok: bool = True + notifications: list[NotificationItem] + unread_count: int + + +@router.post("/push/subscribe", tags=["Notifications"]) +async def push_subscribe( + body: PushSubscribeRequest, + session: AsyncSession = Depends(get_session), + identity: Identity = Depends(verify_identity_token), +): + user_id = identity.tg_id or 0 + + await wn_db.upsert_push_subscription( + session, + user_id=user_id, + identity_id=identity.id, + endpoint=body.endpoint, + keys_json=body.keys, + ) + return {"ok": True} + + +@router.get("/notifications", response_model=NotificationsResponse, tags=["Notifications"]) +async def get_notifications( + limit: int = Query(20, ge=1, le=100), + session: AsyncSession = Depends(get_session), + identity: Identity = Depends(verify_identity_token), +): + notifications = await wn_db.get_notifications_for_identity( + session, identity.id, limit=limit, + ) + unread_count = await wn_db.count_unread_for_identity(session, identity.id) + + items = [ + NotificationItem( + id=n.id, + type=n.type, + title=n.title, + message=n.message, + read=n.read, + created_at=n.created_at.isoformat() if n.created_at else "", + data=n.data, + ) + for n in notifications + ] + return NotificationsResponse(notifications=items, unread_count=unread_count) + + +@router.post("/notifications/read-all", tags=["Notifications"]) +async def read_all_notifications( + session: AsyncSession = Depends(get_session), + identity: Identity = Depends(verify_identity_token), +): + count = await wn_db.mark_all_read_for_identity(session, identity.id) + return {"ok": True, "updated": count} diff --git a/api/v2/routes/partners.py b/api/v2/routes/partners.py index 17dedd47..85ccf04f 100644 --- a/api/v2/routes/partners.py +++ b/api/v2/routes/partners.py @@ -1,14 +1,31 @@ import csv +from base64 import b64encode from datetime import datetime -from io import StringIO +from io import BytesIO, StringIO import re +from urllib.parse import urlsplit -from fastapi import APIRouter, Depends, Path, Query +from fastapi import APIRouter, Depends, HTTPException, Path, Query, Request from fastapi.responses import JSONResponse, StreamingResponse +import qrcode from sqlalchemy import text from sqlalchemy.ext.asyncio import AsyncSession -from api.depends import get_session, verify_identity_admin +from api.depends import get_request_actor, get_session, verify_identity_admin, verify_identity_token +from api.v2.schemas.web_public import ( + PartnerApplyRequest, + PartnerApplyResponse, + PartnerConditionsResponse, + PartnerQrResponse, + PartnerTopEntryResponse, + PartnerTopResponse, + PartnerPayoutEntryResponse, + PartnerPayoutHistoryResponse, + PartnerPayoutRequestCreate, + PartnerPayoutRequestResponse, +) +from database import identities as idb +from utils.referral_codes import decode_partner_code, encode_partner_code try: from modules.partner_program.settings import PARTNER_BONUS_PERCENTAGES @@ -44,6 +61,451 @@ def _row_dt_iso(value) -> str | None: return None +def _resolve_public_base_url(request: Request) -> str: + origin = str(request.headers.get("origin") or "").strip() + if origin.startswith("http://") or origin.startswith("https://"): + return origin.rstrip("/") + referer = str(request.headers.get("referer") or request.headers.get("referrer") or "").strip() + if referer.startswith("http://") or referer.startswith("https://"): + parsed = urlsplit(referer) + if parsed.scheme and parsed.netloc: + return f"{parsed.scheme}://{parsed.netloc}".rstrip("/") + forwarded_host = str(request.headers.get("x-forwarded-host") or "").strip() + host = forwarded_host or str(request.headers.get("host") or "").strip() + forwarded_proto = str(request.headers.get("x-forwarded-proto") or "").split(",", 1)[0].strip().lower() + scheme = forwarded_proto if forwarded_proto in {"http", "https"} else request.url.scheme + if host: + return f"{scheme}://{host}".rstrip("/") + return str(request.base_url).rstrip("/") + + +async def _ensure_partner_code(session: AsyncSession, user_id: int, raw_code: str | None) -> str: + code = str(raw_code or "").strip() + if code and not code.isdigit() and not code.startswith("r1_"): + return code + generated = encode_partner_code(int(user_id)) + try: + await session.execute( + text("UPDATE users SET partner_code = :code WHERE id = :id"), + {"code": generated, "id": int(user_id)}, + ) + await session.flush() + except Exception: + pass + return generated + + +async def _resolve_partner_user(session: AsyncSession, request: Request, identity) -> tuple[int, int]: + actor = get_request_actor(request) + billing_user_id = actor.billing_user_id if actor and actor.billing_user_id is not None else None + if billing_user_id is None: + billing_user_id = await idb.ensure_billing_user_for_identity(session, identity) + row = ( + await session.execute( + text("SELECT id, tg_id FROM users WHERE id = :user_id LIMIT 1"), + {"user_id": int(billing_user_id)}, + ) + ).first() + if row is None or row[1] is None: + raise HTTPException(status_code=400, detail="Партнерский профиль недоступен") + return int(row[0]), int(row[1]) + + +async def _resolve_referrer_by_partner_code(session: AsyncSession, partner_code: str) -> tuple[int, int] | None: + code = str(partner_code or "").strip() + if not code: + return None + by_code_row = ( + await session.execute( + text( + """ + SELECT id, tg_id + FROM users + WHERE lower(COALESCE(partner_code, '')) = lower(:code) + LIMIT 1 + """ + ), + {"code": code}, + ) + ).first() + if by_code_row is not None: + user_id = int(by_code_row[0]) + tg_id = int(by_code_row[1] if by_code_row[1] is not None else by_code_row[0]) + return user_id, tg_id + decoded = decode_partner_code(code) + if decoded is None: + return None + by_id_row = ( + await session.execute( + text("SELECT id, tg_id FROM users WHERE id = :id LIMIT 1"), + {"id": int(decoded)}, + ) + ).first() + if by_id_row is not None: + user_id = int(by_id_row[0]) + tg_id = int(by_id_row[1] if by_id_row[1] is not None else by_id_row[0]) + return user_id, tg_id + by_tg_row = ( + await session.execute( + text("SELECT id, tg_id FROM users WHERE tg_id = :tg_id LIMIT 1"), + {"tg_id": int(decoded)}, + ) + ).first() + if by_tg_row is not None: + user_id = int(by_tg_row[0]) + tg_id = int(by_tg_row[1] if by_tg_row[1] is not None else by_tg_row[0]) + return user_id, tg_id + return None + + +@router.post("/apply", response_model=PartnerApplyResponse) +async def partner_apply( + body: PartnerApplyRequest, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + joined_user_id, joined_tg_id = await _resolve_partner_user(session, request, identity) + code_value = str(body.partner_code or "").strip() + referrer_user_id: int | None = None + referrer_tg_id: int | None = None + if code_value: + resolved = await _resolve_referrer_by_partner_code(session, code_value) + if resolved is not None: + referrer_user_id, referrer_tg_id = resolved + if referrer_tg_id is None and body.partner_tg_id is not None: + referrer_tg_id = int(body.partner_tg_id) + referrer_user_id_row = ( + await session.execute( + text("SELECT id FROM users WHERE tg_id = :tg_id LIMIT 1"), + {"tg_id": int(referrer_tg_id)}, + ) + ).first() + if referrer_user_id_row is not None: + referrer_user_id = int(referrer_user_id_row[0]) + if referrer_tg_id is None: + raise HTTPException(status_code=400, detail="Партнерский код не найден") + if int(referrer_tg_id) == int(joined_tg_id): + raise HTTPException(status_code=400, detail="Нельзя применить свой партнерский код") + already_row = ( + await session.execute( + text("SELECT partner_tg_id FROM partners WHERE joined_tg_id = :joined_tg_id LIMIT 1"), + {"joined_tg_id": int(joined_tg_id)}, + ) + ).first() + if already_row is not None and already_row[0] is not None: + raise HTTPException(status_code=409, detail="Партнер уже привязан") + await session.execute( + text( + """ + INSERT INTO partners (partner_tg_id, joined_tg_id) + VALUES (:partner_tg_id, :joined_tg_id) + """ + ), + {"partner_tg_id": int(referrer_tg_id), "joined_tg_id": int(joined_tg_id)}, + ) + await session.commit() + return PartnerApplyResponse( + ok=True, + message="Партнерский код применен", + partner_code=code_value, + partner_user_id=int(referrer_user_id or 0), + partner_tg_id=int(referrer_tg_id), + joined_user_id=int(joined_user_id), + joined_tg_id=int(joined_tg_id), + ) + + +@router.get("/qr", response_model=PartnerQrResponse) +async def partner_qr( + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + user_id, _ = await _resolve_partner_user(session, request, identity) + code_row = ( + await session.execute( + text("SELECT partner_code FROM users WHERE id = :id LIMIT 1"), + {"id": int(user_id)}, + ) + ).first() + partner_code = await _ensure_partner_code(session, int(user_id), code_row[0] if code_row else None) + base_url = _resolve_public_base_url(request) + partner_link = f"{base_url}/partner/{partner_code}" + qr = qrcode.QRCode(version=1, box_size=10, border=4) + qr.add_data(partner_link) + qr.make(fit=True) + image = qr.make_image(fill_color="black", back_color="white") + png_buffer = BytesIO() + image.save(png_buffer, format="PNG") + image_data = b64encode(png_buffer.getvalue()).decode("ascii") + return PartnerQrResponse( + ok=True, + link=partner_link, + image_data_url=f"data:image/png;base64,{image_data}", + ) + + +@router.get("/top", response_model=PartnerTopResponse) +async def partner_top( + request: Request, + limit: int = Query(5, ge=1, le=20), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + _, joined_tg_id = await _resolve_partner_user(session, request, identity) + user_referred_count_row = ( + await session.execute( + text( + """ + SELECT COUNT(DISTINCT joined_tg_id) + FROM partners + WHERE partner_tg_id = :partner_tg_id + """ + ), + {"partner_tg_id": int(joined_tg_id)}, + ) + ).first() + user_referred_count = int(user_referred_count_row[0] or 0) if user_referred_count_row else 0 + user_position: int | None = None + if user_referred_count > 0: + user_position_row = ( + await session.execute( + text( + """ + SELECT COUNT(*) + 1 + FROM ( + SELECT partner_tg_id, COUNT(DISTINCT joined_tg_id) AS referred_count + FROM partners + WHERE partner_tg_id IS NOT NULL + GROUP BY partner_tg_id + ) ranked + WHERE ranked.referred_count > :referred_count + """ + ), + {"referred_count": int(user_referred_count)}, + ) + ).first() + user_position = int(user_position_row[0] or 1) if user_position_row else 1 + top_rows = ( + await session.execute( + text( + """ + SELECT + COALESCE(u.id, 0) AS partner_user_id, + p.partner_tg_id AS partner_tg_id, + COUNT(DISTINCT p.joined_tg_id) AS referred_count + FROM partners p + LEFT JOIN users u ON u.tg_id = p.partner_tg_id + WHERE p.partner_tg_id IS NOT NULL + GROUP BY p.partner_tg_id, u.id + ORDER BY referred_count DESC, p.partner_tg_id ASC + LIMIT :limit + """ + ), + {"limit": int(limit)}, + ) + ).all() + top: list[PartnerTopEntryResponse] = [] + for index, row in enumerate(top_rows, 1): + partner_user_id = int(row[0] or 0) + partner_tg_id = int(row[1] or 0) + referred_count = int(row[2] or 0) + if partner_user_id > 0: + display_id = encode_partner_code(partner_user_id) + else: + tg_tail = str(partner_tg_id) + display_id = f"p_{tg_tail[:2]}***{tg_tail[-2:]}" if tg_tail else "p_***" + top.append( + PartnerTopEntryResponse( + position=index, + partner_user_id=partner_user_id, + referred_count=referred_count, + display_id=display_id, + ) + ) + return PartnerTopResponse( + user_referred_count=user_referred_count, + user_position=user_position, + top=top, + ) + + +@router.get("/conditions", response_model=PartnerConditionsResponse) +async def partner_conditions( + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + try: + from modules.partner_program import settings as partner_settings + except Exception: + partner_settings = None + mode = str(getattr(partner_settings, "REFERRAL_REWARD_MODE", "percent_only") or "percent_only") + percent_levels_raw = getattr(partner_settings, "PARTNER_BONUS_PERCENTAGES", {}) or {} + flat_levels_raw = getattr(partner_settings, "PARTNER_FLAT_BONUSES", {}) or {} + min_payout = float(getattr(partner_settings, "MIN_PARTNER_PAYOUT", 0) or 0) + custom_amount_enabled = bool(getattr(partner_settings, "ENABLE_CUSTOM_WITHDRAW_AMOUNT", False)) + method_map = [ + ("ENABLE_PAYOUT_CARD", "Карта"), + ("ENABLE_PAYOUT_SBP", "СБП"), + ("ENABLE_PAYOUT_USDT", "USDT"), + ("ENABLE_PAYOUT_TON", "TON"), + ] + payout_methods = [ + title for key, title in method_map if bool(getattr(partner_settings, key, False)) + ] if partner_settings else [] + level_lines: list[str] = [] + all_levels = sorted({int(k) for k in [*percent_levels_raw.keys(), *flat_levels_raw.keys()] if str(k).isdigit()}) + for level in all_levels: + parts: list[str] = [] + if level in percent_levels_raw: + try: + parts.append(f"{float(percent_levels_raw[level]) * 100:.0f}%") + except Exception: + pass + if level in flat_levels_raw: + try: + parts.append(f"{float(flat_levels_raw[level]):.0f} RUB") + except Exception: + pass + if parts: + level_lines.append(f"{level} уровень: {' + '.join(parts)}") + if not level_lines: + level_lines = ["1 уровень: бонус определяется настройками проекта"] + mode_labels = { + "percent_only": "Процент с каждого пополнения приглашенного", + "flat_only": "Фиксированный бонус за первую оплату приглашенного", + "flat_plus_percent": "Фиксированный бонус за первую оплату и процент с пополнений", + } + rules = [ + "Вознаграждение начисляется только после успешной оплаты приглашенного пользователя.", + "Самореферал и самопартнерство недоступны.", + f"Минимальная сумма вывода: {min_payout:.0f} RUB." if min_payout > 0 else "Вывод доступен по правилам проекта.", + ] + if payout_methods: + rules.append(f"Доступные способы вывода: {', '.join(payout_methods)}.") + examples = [ + "Пример: приглашенный пополнил на 1000 RUB, а ставка 15% — вы получаете 150 RUB.", + "Пример: приглашенный сделал несколько пополнений, бонус считается по каждой успешной операции.", + ] + return PartnerConditionsResponse( + title="Условия партнерской программы", + summary="Актуальные условия и режим начислений для партнеров.", + bonus_mode=mode, + bonus_mode_label=mode_labels.get(mode, mode_labels["percent_only"]), + level_lines=level_lines, + rules=rules, + examples=examples, + min_payout_rub=min_payout, + payout_methods=payout_methods, + custom_amount_enabled=custom_amount_enabled, + ) + + +@router.get("/payouts/me", response_model=PartnerPayoutHistoryResponse) +async def partner_payouts_me( + request: Request, + limit: int = Query(20, ge=1, le=100), + offset: int = Query(0, ge=0), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + _, tg_id = await _resolve_partner_user(session, request, identity) + count_sql = text("SELECT COUNT(*) FROM payout_requests WHERE tg_id = :tg_id") + rows_sql = text( + """ + SELECT id, amount, status, created_at, method, destination + FROM payout_requests + WHERE tg_id = :tg_id + ORDER BY created_at DESC, id DESC + LIMIT :limit OFFSET :offset + """ + ) + total = int((await session.scalar(count_sql, {"tg_id": tg_id})) or 0) + rows = ( + await session.execute(rows_sql, {"tg_id": tg_id, "limit": int(limit), "offset": int(offset)}) + ).fetchall() + items = [ + PartnerPayoutEntryResponse( + id=int(row[0]), + amount_rub=float(row[1] or 0.0), + status=str(row[2] or ""), + created_at=_row_dt_iso(row[3]), + method=row[4] or None, + destination=row[5] or None, + ) + for row in rows + ] + return PartnerPayoutHistoryResponse(total=total, items=items) + + +@router.post("/payouts/me", response_model=PartnerPayoutRequestResponse) +async def partner_create_payout_request( + body: PartnerPayoutRequestCreate, + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + user_id, tg_id = await _resolve_partner_user(session, request, identity) + row = ( + await session.execute( + text("SELECT COALESCE(partner_balance, 0), payout_method, card_number FROM users WHERE id = :id"), + {"id": user_id}, + ) + ).first() + balance = float(row[0] or 0.0) if row else 0.0 + requested = float(body.amount_rub) + if requested <= 0: + raise HTTPException(status_code=400, detail="Сумма должна быть больше нуля") + try: + from modules.partner_program.settings import ENABLE_CUSTOM_WITHDRAW_AMOUNT, MIN_PARTNER_PAYOUT + except Exception: + ENABLE_CUSTOM_WITHDRAW_AMOUNT = True + MIN_PARTNER_PAYOUT = 0 + min_payout = float(MIN_PARTNER_PAYOUT or 0) + if requested < min_payout: + raise HTTPException(status_code=400, detail=f"Минимальная сумма вывода — {min_payout:.0f} RUB") + if not bool(ENABLE_CUSTOM_WITHDRAW_AMOUNT): + requested = balance + if requested > balance: + raise HTTPException(status_code=400, detail="Недостаточно партнерского баланса") + if requested <= 0: + raise HTTPException(status_code=400, detail="Недостаточно средств для заявки") + payout_method = (row[1] if row else None) or "card" + destination = (row[2] if row else None) or None + inserted = ( + await session.execute( + text( + """ + INSERT INTO payout_requests (tg_id, amount, status, created_at, method, destination) + VALUES (:tg_id, :amount, 'pending', NOW(), :method, :destination) + RETURNING id + """ + ), + { + "tg_id": int(tg_id), + "amount": float(requested), + "method": payout_method, + "destination": destination, + }, + ) + ).scalar() + new_balance = balance - requested + await session.execute( + text("UPDATE users SET partner_balance = :balance WHERE id = :id"), + {"balance": new_balance, "id": int(user_id)}, + ) + await session.commit() + return PartnerPayoutRequestResponse( + ok=True, + message="Заявка на вывод создана", + request_id=int(inserted) if inserted is not None else None, + amount_rub=float(requested), + status="pending", + balance_rub=float(new_balance), + ) + + @router.get("/all") async def get_all_partners( limit: int = Query(1000, ge=1, le=10000, description="Лимит результатов"), diff --git a/api/v2/routes/payment_links.py b/api/v2/routes/payment_links.py index be630383..8c10d32b 100644 --- a/api/v2/routes/payment_links.py +++ b/api/v2/routes/payment_links.py @@ -1,28 +1,143 @@ +import asyncio +import json + from fastapi import APIRouter, Depends, HTTPException, Request +from fastapi.responses import StreamingResponse from sqlalchemy.ext.asyncio import AsyncSession from api.depends import get_session, verify_identity_token -from api.v2.schemas.payment_links import PaymentLinkCreateRequest, PaymentLinkCreateResponse -from database import identities as idb -from handlers.payments import create_payment_link -from handlers.payments.payment_links import PaymentLinkRequest +from api.v2.schemas.payment_links import PaymentLinkCreateRequest, PaymentLinkCreateResponse, PaymentLinkStatusResponse +from config import REDIS_URL +from database import ( + async_session_maker, + get_payment_by_payment_id, + get_payment_from_db_by_payment_id, + identities as idb, +) +from database.temporary_data import create_temporary_data +from logger import logger +from services.payments.payment_events import payment_events_channel +from services.payments.payment_links import PaymentLinkRequest, create_payment_link + router = APIRouter(tags=["PaymentLinks"]) -async def _resolve_tg_id(body: PaymentLinkCreateRequest, session: AsyncSession) -> int: - """Возвращает tg_id из body.tg_id или из identity_id; иначе исключение.""" - if body.tg_id is not None: - return body.tg_id - if body.identity_id: - tg_id = await idb.resolve_tg_id(session, body.identity_id) - if tg_id is not None: - return tg_id - raise HTTPException( - status_code=400, - detail="У идентичности не привязан Telegram. Привяжите tg_id для создания платёжной ссылки.", - ) - raise HTTPException(status_code=400, detail="Укажите tg_id или identity_id") +async def _store_payment_intent( + session: AsyncSession, + billing_user_ref: int, + metadata: dict | None, + amount: int | float, +) -> None: + if not isinstance(metadata, dict): + return + payment_flow = str(metadata.get("payment_flow") or "").strip().lower() + required_amount = int(round(float(amount))) + if payment_flow == "tariff_purchase": + tariff_id = metadata.get("tariff_id") + if tariff_id in (None, ""): + return + payload: dict[str, int | str] = { + "tariff_id": int(tariff_id), + "required_amount": required_amount, + "selected_price_rub": int(metadata.get("selected_price_rub") or required_amount), + } + selected_device_limit = metadata.get("selected_device_limit") + if selected_device_limit not in (None, ""): + payload["selected_device_limit"] = int(selected_device_limit) + selected_traffic_gb = metadata.get("selected_traffic_gb") + if selected_traffic_gb not in (None, ""): + payload["selected_traffic_limit_gb"] = int(selected_traffic_gb) + selected_duration_days = metadata.get("selected_duration_days") + if selected_duration_days not in (None, ""): + payload["selected_duration_days"] = int(selected_duration_days) + coupon_id = metadata.get("coupon_id") + if coupon_id not in (None, ""): + payload["coupon_id"] = int(coupon_id) + discount_rub = metadata.get("discount_rub") + if discount_rub not in (None, ""): + payload["discount_rub"] = int(discount_rub) + base_price_rub = metadata.get("base_price_rub") + if base_price_rub not in (None, ""): + payload["base_price_rub"] = int(base_price_rub) + applied_coupon_code = metadata.get("applied_coupon_code") + if applied_coupon_code not in (None, ""): + payload["applied_coupon_code"] = str(applied_coupon_code) + await create_temporary_data(session, billing_user_ref, "waiting_for_payment", payload) + return + if payment_flow == "key_renewal": + required_fields = ("tariff_id", "client_id", "email", "cost") + if any(metadata.get(field) in (None, "") for field in required_fields): + return + payload: dict[str, int | str] = { + "tariff_id": int(metadata["tariff_id"]), + "client_id": str(metadata["client_id"]), + "email": str(metadata["email"]), + "cost": int(metadata["cost"]), + "required_amount": required_amount, + "selected_price_rub": int(metadata.get("selected_price_rub") or metadata["cost"]), + } + selected_duration_days = metadata.get("selected_duration_days") + if selected_duration_days not in (None, ""): + payload["selected_duration_days"] = int(selected_duration_days) + selected_device_limit = metadata.get("selected_device_limit") + if selected_device_limit not in (None, ""): + payload["selected_device_limit"] = int(selected_device_limit) + selected_traffic_limit = metadata.get("selected_traffic_limit") + if selected_traffic_limit not in (None, ""): + payload["selected_traffic_limit"] = int(selected_traffic_limit) + total_gb = metadata.get("total_gb") + if total_gb not in (None, ""): + payload["total_gb"] = int(total_gb) + coupon_id = metadata.get("coupon_id") + if coupon_id not in (None, ""): + payload["coupon_id"] = int(coupon_id) + discount_rub = metadata.get("discount_rub") + if discount_rub not in (None, ""): + payload["discount_rub"] = int(discount_rub) + base_price_rub = metadata.get("base_price_rub") + if base_price_rub not in (None, ""): + payload["base_price_rub"] = int(base_price_rub) + applied_coupon_code = metadata.get("applied_coupon_code") + if applied_coupon_code not in (None, ""): + payload["applied_coupon_code"] = str(applied_coupon_code) + await create_temporary_data(session, billing_user_ref, "waiting_for_renewal_payment", payload) + return + if payment_flow == "key_addons": + required_fields = ("tariff_id", "email", "original_price") + if any(metadata.get(field) in (None, "") for field in required_fields): + return + payload: dict[str, int | str] = { + "tariff_id": int(metadata["tariff_id"]), + "email": str(metadata["email"]), + "original_price": int(metadata["original_price"]), + "required_amount": required_amount, + } + selected_device_limit = metadata.get("selected_device_limit") + if selected_device_limit not in (None, ""): + payload["selected_device_limit"] = int(selected_device_limit) + selected_traffic_gb = metadata.get("selected_traffic_gb") + if selected_traffic_gb not in (None, ""): + payload["selected_traffic_gb"] = int(selected_traffic_gb) + current_device_limit = metadata.get("current_device_limit") + if current_device_limit not in (None, ""): + payload["current_device_limit"] = int(current_device_limit) + current_traffic_gb = metadata.get("current_traffic_gb") + if current_traffic_gb not in (None, ""): + payload["current_traffic_gb"] = int(current_traffic_gb) + coupon_id = metadata.get("coupon_id") + if coupon_id not in (None, ""): + payload["coupon_id"] = int(coupon_id) + discount_rub = metadata.get("discount_rub") + if discount_rub not in (None, ""): + payload["discount_rub"] = int(discount_rub) + base_price_rub = metadata.get("base_price_rub") + if base_price_rub not in (None, ""): + payload["base_price_rub"] = int(base_price_rub) + applied_coupon_code = metadata.get("applied_coupon_code") + if applied_coupon_code not in (None, ""): + payload["applied_coupon_code"] = str(applied_coupon_code) + await create_temporary_data(session, billing_user_ref, "waiting_for_addons_payment", payload) @router.post("/", response_model=PaymentLinkCreateResponse) @@ -32,13 +147,10 @@ async def create_link( session: AsyncSession = Depends(get_session), identity=Depends(verify_identity_token), ): - """Создаёт платёжную ссылку через выбранную кассу (единая точка входа). Принимает identity_id или tg_id.""" - try: - tg_id = await _resolve_tg_id(body, session) - except HTTPException: - raise + """Создаёт платёжную ссылку для текущего авторизованного пользователя.""" + billing_user_ref = await idb.ensure_billing_user_for_identity(session, identity) payment_request = PaymentLinkRequest( - tg_id=tg_id, + legacy_user_ref=billing_user_ref, amount=body.amount, currency=body.currency or "RUB", provider_id=body.provider_id, @@ -47,9 +159,114 @@ async def create_link( metadata=body.metadata, ) result = await create_payment_link(session, payment_request) + if result.success: + await _store_payment_intent( + session=session, + billing_user_ref=billing_user_ref, + metadata=body.metadata, + amount=body.amount, + ) return PaymentLinkCreateResponse( success=result.success, payment_id=result.payment_id, payment_url=result.payment_url, error=result.error, ) + + +@router.get("/stream") +async def payment_events_stream( + request: Request, + x_identity_id: str = "", + token: str = "", +): + identity_id = str(request.headers.get("X-Identity-Id") or x_identity_id or "").strip() + token = str(request.headers.get("X-Token") or token or "").strip() + if not identity_id or not token: + raise HTTPException(status_code=401, detail="Unauthorized") + + async with async_session_maker() as session: + identity = await idb.verify_identity_token(session, identity_id, token) + if not identity: + raise HTTPException(status_code=401, detail="Unauthorized") + billing_user_ref = await idb.ensure_billing_user_for_identity(session, identity) + await session.commit() + + async def event_generator(): + redis_client = None + pubsub = None + channel = payment_events_channel(int(billing_user_ref)) + try: + from redis.asyncio import from_url + + redis_client = from_url(REDIS_URL, encoding="utf-8", decode_responses=True, max_connections=8) + pubsub = redis_client.pubsub(ignore_subscribe_messages=True) + await pubsub.subscribe(channel) + logger.info(f"[Payments] SSE subscribed: user_ref={billing_user_ref}, channel={channel}") + yield "retry: 1500\n\n" + while True: + if await request.is_disconnected(): + logger.info(f"[Payments] SSE disconnected by client: user_ref={billing_user_ref}") + break + message = await pubsub.get_message(ignore_subscribe_messages=True, timeout=15.0) + if message and message.get("type") == "message": + raw_data = message.get("data") + payload = json.loads(raw_data) if isinstance(raw_data, str) else raw_data + if isinstance(payload, dict): + logger.info( + f"[Payments] SSE emit: user_ref={billing_user_ref}, " + f"status={payload.get('status')}, flow={payload.get('flow')}" + ) + yield f"data: {json.dumps(payload, ensure_ascii=False)}\n\n" + continue + yield ": keepalive\n\n" + await asyncio.sleep(0.1) + finally: + if pubsub is not None: + try: + await pubsub.unsubscribe(channel) + await pubsub.close() + except Exception: + pass + if redis_client is not None: + try: + await redis_client.aclose() + except Exception: + pass + + return StreamingResponse( + event_generator(), + media_type="text/event-stream", + headers={ + "Cache-Control": "no-cache, no-transform", + "Connection": "keep-alive", + "X-Accel-Buffering": "no", + }, + ) + + +@router.get("/{payment_id}", response_model=PaymentLinkStatusResponse) +async def get_link_status( + payment_id: str, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + billing_user_ref = await idb.ensure_billing_user_for_identity(session, identity) + payment = await get_payment_from_db_by_payment_id(session, payment_id) + if payment is None: + payment = await get_payment_by_payment_id(session, payment_id) + if not payment: + raise HTTPException(status_code=404, detail="Payment not found") + owner_ref = payment.get("user_id") + if owner_ref is None: + owner_ref = payment.get("tg_id") + if owner_ref is None or int(owner_ref) != int(billing_user_ref): + raise HTTPException(status_code=404, detail="Payment not found") + status = str(payment.get("status") or "").lower() or None + return PaymentLinkStatusResponse( + success=True, + payment_id=payment_id, + status=status, + completed=status in {"success", "failed", "cancelled"}, + paid=status == "success", + ) diff --git a/api/v2/routes/referrals.py b/api/v2/routes/referrals.py index db3764a8..647d82bd 100644 --- a/api/v2/routes/referrals.py +++ b/api/v2/routes/referrals.py @@ -1,37 +1,187 @@ -from fastapi import Depends, HTTPException, Query -from sqlalchemy import select +from base64 import b64encode +from io import BytesIO +from urllib.parse import urlsplit + +import qrcode +from fastapi import APIRouter, Depends, HTTPException, Query, Request from sqlalchemy.ext.asyncio import AsyncSession -from api.depends import get_session, verify_identity_admin -from api.v2.base_crud import generate_crud_router -from api.v2.schemas import ReferralResponse -from database.models import Referral - -router = generate_crud_router( - model=Referral, - schema_response=ReferralResponse, - schema_create=None, - schema_update=None, - identifier_field="referrer_tg_id", - parameter_name="referrer_tg_id", - enabled_methods=["get_all", "get_one", "get_all_by_field"], +from api.depends import get_session, verify_identity_token +from api.v2.schemas.web_public import ( + ReferralApplyRequest, + ReferralApplyResponse, + ReferralConditionsResponse, + ReferralQrResponse, + ReferralTopEntryResponse, + ReferralTopResponse, ) +from config import CHECK_REFERRAL_REWARD_ISSUED, REFERRAL_BONUS_PERCENTAGES, REFERRAL_BUTTON, REFERRAL_QR, TOP_REFERRAL_BUTTON +from core.bootstrap import BUTTONS_CONFIG +from database import add_referral, get_referral_by_referred_id, get_user_referral_count +from database.referrals import get_referral_position, get_top_referrals +from database import identities as idb +from database.access.resolution import resolve_user_optional +from utils.referral_codes import decode_referral_code, encode_referral_code + +router = APIRouter() -@router.delete("/one") -async def delete_one_referral( - referrer_tg_id: int = Query(..., description="ID пригласившего"), - referred_tg_id: int = Query(..., description="ID приглашённого"), - identity=Depends(verify_identity_admin), +def _normalize_referrer_code(value: str | None, fallback_tg_id: int | None) -> int | None: + raw = str(value or "").strip() + if raw: + if "/referral/" in raw: + raw = raw.split("/referral/", 1)[-1] + if "start=referral_" in raw: + raw = raw.split("start=referral_", 1)[-1] + raw = raw.split("?", 1)[0].split("#", 1)[0].strip() + parsed = decode_referral_code(raw) + if parsed is not None: + return parsed + if fallback_tg_id is not None and int(fallback_tg_id) > 0: + return int(fallback_tg_id) + return None + + +def _resolve_public_base_url(request: Request) -> str: + origin = str(request.headers.get("origin") or "").strip() + if origin.startswith("http://") or origin.startswith("https://"): + return origin.rstrip("/") + referer = str(request.headers.get("referer") or request.headers.get("referrer") or "").strip() + if referer.startswith("http://") or referer.startswith("https://"): + parsed = urlsplit(referer) + if parsed.scheme and parsed.netloc: + return f"{parsed.scheme}://{parsed.netloc}".rstrip("/") + forwarded_host = str(request.headers.get("x-forwarded-host") or "").strip() + host = forwarded_host or str(request.headers.get("host") or "").strip() + forwarded_proto = str(request.headers.get("x-forwarded-proto") or "").split(",", 1)[0].strip().lower() + scheme = forwarded_proto if forwarded_proto in {"http", "https"} else request.url.scheme + if host: + return f"{scheme}://{host}".rstrip("/") + return str(request.base_url).rstrip("/") + + +@router.post("/apply", response_model=ReferralApplyResponse, tags=["Referrals"]) +async def apply_referral( + body: ReferralApplyRequest, session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), ): - """Удаляет одну связь реферала по паре referrer/referred.""" - result = await session.execute( - select(Referral).where(Referral.referrer_tg_id == referrer_tg_id, Referral.referred_tg_id == referred_tg_id) + if not bool(BUTTONS_CONFIG.get("REFERRAL_BUTTON_ENABLED", REFERRAL_BUTTON)): + raise HTTPException(status_code=403, detail="Реферальная программа отключена") + billing_uid = await idb.ensure_billing_user_for_identity(session, identity) + referrer_legacy = _normalize_referrer_code(body.referrer_code, body.referrer_tg_id) + if referrer_legacy is None: + raise HTTPException(status_code=400, detail="Приглашение недействительно") + referrer_u = await resolve_user_optional(session, referrer_legacy) + if referrer_u is None: + raise HTTPException(status_code=400, detail="Приглашение недействительно") + if billing_uid == referrer_u.id: + raise HTTPException(status_code=400, detail="Нельзя использовать собственную ссылку") + if await get_referral_by_referred_id(session, billing_uid): + raise HTTPException(status_code=409, detail="Реферальная связь уже сохранена") + await add_referral(session, billing_uid, referrer_u.id) + referred_u = await resolve_user_optional(session, billing_uid) + return ReferralApplyResponse( + ok=True, + message="Приглашение применено", + referrer_code=str(referrer_u.id), + referrer_user_id=int(referrer_u.id), + referrer_tg_id=referrer_u.tg_id, + referred_user_id=int(billing_uid), + referred_tg_id=referred_u.tg_id if referred_u is not None else None, + ) + + +@router.get("/top", response_model=ReferralTopResponse, tags=["Referrals"]) +async def referral_top( + limit: int = Query(5, ge=1, le=20), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + if not bool(BUTTONS_CONFIG.get("REFERRAL_BUTTON_ENABLED", REFERRAL_BUTTON)): + raise HTTPException(status_code=403, detail="Реферальная программа отключена") + if not bool(BUTTONS_CONFIG.get("TOP_REFERRAL_BUTTON_ENABLE", TOP_REFERRAL_BUTTON)): + raise HTTPException(status_code=403, detail="Топ рефералов отключен в настройках") + billing_uid = await idb.ensure_billing_user_for_identity(session, identity) + user_referral_count = int(await get_user_referral_count(session, billing_uid)) + user_position = int(await get_referral_position(session, user_referral_count)) if user_referral_count > 0 else None + top_rows = await get_top_referrals(session, limit=limit) + top: list[ReferralTopEntryResponse] = [] + for index, row in enumerate(top_rows, 1): + referrer_user_id = int(row.get("referrer_user_id") or 0) + referrals_count = int(row.get("referral_count") or 0) + display_id = encode_referral_code(referrer_user_id) + top.append( + ReferralTopEntryResponse( + position=index, + referrer_user_id=referrer_user_id, + referrals_count=referrals_count, + display_id=display_id, + ) + ) + return ReferralTopResponse( + user_referrals_count=user_referral_count, + user_position=user_position, + top=top, + ) + + +@router.get("/qr", response_model=ReferralQrResponse, tags=["Referrals"]) +async def referral_qr( + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + if not bool(BUTTONS_CONFIG.get("REFERRAL_BUTTON_ENABLED", REFERRAL_BUTTON)): + raise HTTPException(status_code=403, detail="Реферальная программа отключена") + if not bool(BUTTONS_CONFIG.get("REFERRAL_QR_BUTTON_ENABLE", REFERRAL_QR)): + raise HTTPException(status_code=403, detail="QR реферальной ссылки отключен в настройках") + billing_uid = await idb.ensure_billing_user_for_identity(session, identity) + base_url = _resolve_public_base_url(request) + referral_link = f"{base_url}/referral/{encode_referral_code(int(billing_uid))}" + qr = qrcode.QRCode(version=1, box_size=10, border=4) + qr.add_data(referral_link) + qr.make(fit=True) + img = qr.make_image(fill_color="black", back_color="white") + buffer = BytesIO() + img.save(buffer, format="PNG") + image_data = b64encode(buffer.getvalue()).decode("ascii") + return ReferralQrResponse( + ok=True, + link=referral_link, + image_data_url=f"data:image/png;base64,{image_data}", + ) + + +@router.get("/conditions", response_model=ReferralConditionsResponse, tags=["Referrals"]) +async def referral_conditions( + identity=Depends(verify_identity_token), +): + if not bool(BUTTONS_CONFIG.get("REFERRAL_BUTTON_ENABLED", REFERRAL_BUTTON)): + raise HTTPException(status_code=403, detail="Реферальная программа отключена") + del identity + level_lines: list[str] = [] + for level in sorted(REFERRAL_BONUS_PERCENTAGES.keys()): + value = REFERRAL_BONUS_PERCENTAGES[level] + if isinstance(value, float): + label = f"{int(value * 100)}% от суммы оплаты" + else: + label = f"{float(value):g} RUB" + level_lines.append(f"{level} уровень: {label}") + one_time_mode = bool(CHECK_REFERRAL_REWARD_ISSUED) + bonus_mode = "one_time" if one_time_mode else "each_payment" + bonus_mode_label = "Бонус за первую успешную оплату реферала" if one_time_mode else "Бонус за каждую успешную оплату реферала" + rules = [ + "Бонус начисляется только за реальных приглашённых пользователей.", + "Нельзя использовать собственную реферальную ссылку.", + "Реферальную связь можно применить только один раз.", + "Размер бонуса зависит от уровня реферальной программы.", + ] + return ReferralConditionsResponse( + title="Условия реферальной программы", + summary=f"Режим начисления: {bonus_mode_label}.", + bonus_mode=bonus_mode, + bonus_mode_label=bonus_mode_label, + level_lines=level_lines, + rules=rules, ) - obj = result.scalar_one_or_none() - if not obj: - raise HTTPException(status_code=404, detail="Referral not found") - await session.delete(obj) - await session.commit() - return {"status": "deleted_one"} diff --git a/api/v2/routes/root.py b/api/v2/routes/root.py index 43feae24..64ee8d0d 100644 --- a/api/v2/routes/root.py +++ b/api/v2/routes/root.py @@ -1,10 +1,63 @@ from fastapi import APIRouter -from config import PROJECT_NAME, USERNAME_BOT +from config import ( + BALANCE_BUTTON, + CAPTCHA_ENABLE, + CHANNEL_EXISTS, + CHANNEL_REQUIRED, + DONATIONS_ENABLE, + GIFT_BUTTON, + HAPP_CRYPTOLINK, + HWID_RESET_BUTTON, + INSTRUCTIONS_BUTTON, + PROJECT_NAME, + REFERRAL_BUTTON, + REFERRAL_QR, + REMNAWAVE_WEBAPP, + REMNAWAVE_WEBAPP_OPEN_IN_BROWSER, + TOP_REFERRAL_BUTTON, + TRIAL_TIME_DISABLE, + USE_COUNTRY_SELECTION, + USERNAME_BOT, + TELEGRAM_WEBAPP_DIRECT_LINK, + TELEGRAM_WEBAPP_SHORT_NAME, +) +from core.bootstrap import BUTTONS_CONFIG, MODES_CONFIG, MONEY_CONFIG, PAYMENTS_CONFIG +from core.settings.web_config import WEB_CONFIG +from core.settings.money_config import get_currency_mode +from services.payments.providers import PROVIDERS_BASE, TELEGRAM_ONLY_PROVIDER_IDS, WEB_LINK_PROVIDER_IDS router = APIRouter(tags=["Root"]) +def _telegram_web_app_return_base() -> str | None: + direct = str(TELEGRAM_WEBAPP_DIRECT_LINK or "").strip().rstrip("/") + if direct: + if direct.lower().startswith("http://"): + direct = "https://" + direct[7:] + if direct.lower().startswith("https://t.me/"): + return direct + bot = USERNAME_BOT.replace("@", "").strip() + sn = str(TELEGRAM_WEBAPP_SHORT_NAME or "").strip() + if bot and sn: + return f"https://t.me/{bot}/{sn}" + if bot: + return f"https://t.me/{bot}" + return None + + +def _partner_feature_enabled() -> bool: + try: + from modules.partner_program import settings as partner_settings + except Exception: + return False + for key in ("PARTNER_PROGRAM_ENABLED", "PARTNER_BUTTON_ENABLED", "PARTNER_ENABLED"): + value = getattr(partner_settings, key, None) + if isinstance(value, bool): + return value + return True + + @router.get("/api", include_in_schema=False) async def root(): return {"message": "SoloBot API v2", "docs": "/api/docs"} @@ -22,3 +75,87 @@ async def telegram_widget_bot(): "bot_username": USERNAME_BOT.replace("@", ""), "project_name": (PROJECT_NAME or "Solo").strip() if isinstance(PROJECT_NAME, str) else "Solo", } + + +@router.get("/api/site-config", include_in_schema=True) +async def site_config(): + """Настройки витрины и кабинета для веб-клиента (флаги из runtime-конфигов бота).""" + bot_username = USERNAME_BOT.replace("@", "").strip() + pay_flags = {name: bool(PAYMENTS_CONFIG.get(name)) for name in PROVIDERS_BASE} + any_pay = any(pay_flags.values()) + web_link_provider_ids = [provider_id for provider_id in WEB_LINK_PROVIDER_IDS if pay_flags.get(provider_id, False)] + telegram_only_provider_ids = [ + provider_id for provider_id in TELEGRAM_ONLY_PROVIDER_IDS if pay_flags.get(provider_id, False) + ] + currency_mode, currency_one_screen = get_currency_mode() + try: + cb_raw = MONEY_CONFIG.get("CASHBACK", 0) + cashback_percent = float(cb_raw) if cb_raw not in (None, False) else 0.0 + except (TypeError, ValueError): + cashback_percent = 0.0 + + webapp_short = str(TELEGRAM_WEBAPP_SHORT_NAME or "").strip() or None + webapp_return_base = _telegram_web_app_return_base() + return { + "bot_username": bot_username or None, + "telegram_web_app_short_name": webapp_short, + "telegram_web_app_return_base": webapp_return_base, + "project_name": (PROJECT_NAME or "Solo").strip() if isinstance(PROJECT_NAME, str) else "Solo", + "site_mode": str(WEB_CONFIG.get("SITE_MODE", "full")).strip() or "full", + "auth": { + "telegram_login_enabled": bool(bot_username), + "email_code_login_enabled": bool(MODES_CONFIG.get("WEB_EMAIL_CODE_LOGIN_ENABLED", True)), + }, + "mobile": { + "prefer_mini_app_on_telegram_mobile": bool( + MODES_CONFIG.get("PREFER_MINI_APP_ON_TELEGRAM_MOBILE", False) + ), + }, + "features": { + "channel_enabled": bool(BUTTONS_CONFIG.get("CHANNEL_BUTTON_ENABLE", CHANNEL_EXISTS)), + "donations_enabled": bool(BUTTONS_CONFIG.get("DONATIONS_BUTTON_ENABLE", DONATIONS_ENABLE)), + "balance_enabled": bool(BUTTONS_CONFIG.get("BALANCE_BUTTON_ENABLE", BALANCE_BUTTON)), + "referral_qr_enabled": bool(BUTTONS_CONFIG.get("REFERRAL_QR_BUTTON_ENABLE", REFERRAL_QR)), + "instructions_enabled": bool(BUTTONS_CONFIG.get("INSTRUCTIONS_BUTTON_ENABLE", INSTRUCTIONS_BUTTON)), + "gift_enabled": bool(BUTTONS_CONFIG.get("GIFT_BUTTON_ENABLE", GIFT_BUTTON)), + "referral_enabled": bool(BUTTONS_CONFIG.get("REFERRAL_BUTTON_ENABLED", REFERRAL_BUTTON)), + "top_referral_enabled": bool(BUTTONS_CONFIG.get("TOP_REFERRAL_BUTTON_ENABLE", TOP_REFERRAL_BUTTON)), + "coupon_enabled": bool(BUTTONS_CONFIG.get("COUPON_BUTTON_ENABLE", True)), + "qr_subscription_enabled": bool(MODES_CONFIG.get("HAPP_CRYPTOLINK_ENABLED", HAPP_CRYPTOLINK)), + "hwid_reset_enabled": bool(BUTTONS_CONFIG.get("HWID_RESET_BUTTON_ENABLE", HWID_RESET_BUTTON)), + "country_selection_enabled": bool( + MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION) + ), + "captcha_enabled": bool(MODES_CONFIG.get("CAPTCHA_ENABLED", CAPTCHA_ENABLE)), + "channel_check_enabled": bool(MODES_CONFIG.get("CHANNEL_CHECK_ENABLED", CHANNEL_REQUIRED)), + "trial_enabled": not bool(MODES_CONFIG.get("TRIAL_TIME_DISABLED", TRIAL_TIME_DISABLE)), + "mini_app_enabled": bool(MODES_CONFIG.get("REMNAWAVE_WEBAPP_ENABLED", REMNAWAVE_WEBAPP)), + "mini_app_open_in_browser": bool( + MODES_CONFIG.get("REMNAWAVE_WEBAPP_OPEN_IN_BROWSER", REMNAWAVE_WEBAPP_OPEN_IN_BROWSER) + ), + "partner_enabled": bool(_partner_feature_enabled()), + }, + "payments": { + "any_enabled": any_pay, + "any_web_link_enabled": bool(web_link_provider_ids), + "any_telegram_only_enabled": bool(telegram_only_provider_ids), + "web_link_provider_ids": web_link_provider_ids, + "telegram_only_provider_ids": telegram_only_provider_ids, + "yookassa_enabled": pay_flags.get("YOOKASSA", False), + "yoomoney_enabled": pay_flags.get("YOOMONEY", False), + "robokassa_enabled": pay_flags.get("ROBOKASSA", False), + "kassai_cards_enabled": pay_flags.get("KASSAI_CARDS", False), + "kassai_sbp_enabled": pay_flags.get("KASSAI_SBP", False), + "tribute_enabled": pay_flags.get("TRIBUTE", False), + "heleket_enabled": pay_flags.get("HELEKET", False), + "cryptobot_enabled": pay_flags.get("CRYPTOBOT", False), + "freekassa_enabled": pay_flags.get("FREEKASSA", False), + "stars_enabled": pay_flags.get("STARS", False), + }, + "money": { + "currency_mode": currency_mode, + "currency_one_screen": currency_one_screen, + "cashback_enabled": cashback_percent > 0, + "cashback_percent": cashback_percent, + }, + } diff --git a/api/v2/routes/tariffs.py b/api/v2/routes/tariffs.py index b1b69783..36fad4c6 100644 --- a/api/v2/routes/tariffs.py +++ b/api/v2/routes/tariffs.py @@ -1,15 +1,62 @@ -from fastapi import APIRouter, Depends, Query +from datetime import datetime, timedelta +from math import ceil +from urllib.parse import urlsplit + +from fastapi import APIRouter, Depends, HTTPException, Query, Request +from pytz import timezone as tz_moscow from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession -from api.depends import get_session +from api.depends import get_session, validate_redirect_url, verify_identity_token +from api.v2.base_crud import generate_crud_router +from api.v2.routes.coupon_pricing import resolve_percent_coupon_pricing from api.v2.schemas import TariffBase, TariffResponse, TariffUpdate from api.v2.schemas.tariffs import TariffGroup, TariffPublic -from api.v2.base_crud import generate_crud_router +from api.v2.schemas.web_public import ( + TariffConfigPriceResponse, + TariffPurchaseRequest, + TariffPurchaseResponse, +) +from core.bootstrap import PAYMENTS_CONFIG +from core.redis_cache import cache_get, cache_key, cache_set +from database import ( + get_balance, + identities as idb, +) +from database.coupons import mark_coupon_used from database.models import Tariff +from database.tariffs import get_tariff_by_id +from database.temporary_data import create_temporary_data +from services.keys import create_vpn_key_headless +from logger import logger +from services.payments.payment_links import PaymentLinkRequest, create_payment_link +from services.payments.providers import WEB_LINK_PROVIDER_IDS +from services.tariffs import calculate_config_price def _tariff_to_public(t: Tariff) -> TariffPublic: + dev_opts = getattr(t, "device_options", None) + tr_opts = getattr(t, "traffic_options_gb", None) + device_options: list[int] | None = None + traffic_options_gb: list[int] | None = None + if isinstance(dev_opts, list): + device_options = [] + for x in dev_opts: + try: + device_options.append(int(x)) + except (TypeError, ValueError): + continue + if not device_options: + device_options = None + if isinstance(tr_opts, list): + traffic_options_gb = [] + for x in tr_opts: + try: + traffic_options_gb.append(int(x)) + except (TypeError, ValueError): + continue + if not traffic_options_gb: + traffic_options_gb = None return TariffPublic( id=t.id, name=t.name or "", @@ -21,18 +68,57 @@ def _tariff_to_public(t: Tariff) -> TariffPublic: subgroup_title=t.subgroup_title, sort_order=t.sort_order, vless=bool(getattr(t, "vless", False)), + configurable=bool(getattr(t, "configurable", False)), + device_options=device_options, + traffic_options_gb=traffic_options_gb, ) public_router = APIRouter() +def _resolve_public_base_url(request: Request) -> str: + origin = str(request.headers.get("origin") or "").strip() + if origin.startswith(("http://", "https://")): + return origin.rstrip("/") + referer = str(request.headers.get("referer") or request.headers.get("referrer") or "").strip() + if referer.startswith(("http://", "https://")): + parsed = urlsplit(referer) + if parsed.scheme and parsed.netloc: + return f"{parsed.scheme}://{parsed.netloc}".rstrip("/") + forwarded_host = str(request.headers.get("x-forwarded-host") or "").strip() + host = forwarded_host or str(request.headers.get("host") or "").strip() + forwarded_proto = str(request.headers.get("x-forwarded-proto") or "").split(",", 1)[0].strip().lower() + scheme = forwarded_proto if forwarded_proto in {"http", "https"} else request.url.scheme + if host: + return f"{scheme}://{host}".rstrip("/") + return str(request.base_url).rstrip("/") + + +def _resolve_default_web_payment_provider() -> str | None: + for provider_id in WEB_LINK_PROVIDER_IDS: + if bool(PAYMENTS_CONFIG.get(provider_id)): + return provider_id + return WEB_LINK_PROVIDER_IDS[0] if WEB_LINK_PROVIDER_IDS else None + + +def _public_tariffs_cache_key( + group_code: str | None, + tariff_ids: str | None, + filter_vless: str | None, +) -> str: + normalized_group = (group_code or "").strip().lower() + normalized_ids = ",".join(part.strip() for part in (tariff_ids or "").split(",") if part.strip()) + normalized_vless = (filter_vless or "").strip().lower() + return cache_key("tariffs_public", normalized_group or "-", normalized_ids or "-", normalized_vless or "-") + + @public_router.get("/groups", response_model=list[TariffGroup]) async def get_tariff_groups(session: AsyncSession = Depends(get_session)): """Публичный список групп тарифов — уникальные значения колонки group_code.""" q = ( select(Tariff.group_code) - .where(Tariff.is_active == True, Tariff.group_code.isnot(None), Tariff.group_code != "") + .where(Tariff.is_active is True, Tariff.group_code.isnot(None), Tariff.group_code != "") .distinct() .order_by(Tariff.group_code) ) @@ -52,23 +138,304 @@ async def get_tariffs_public( session: AsyncSession = Depends(get_session), ): """Публичный список активных тарифов (без авторизации).""" - q = select(Tariff).where(Tariff.is_active == True).order_by(Tariff.sort_order.asc().nulls_last(), Tariff.price_rub.asc()) + cache_token = _public_tariffs_cache_key(group_code, tariff_ids, filter_vless) + cached = await cache_get(cache_token) + if isinstance(cached, list): + return cached + + q = select(Tariff).where(Tariff.is_active is True).order_by(Tariff.sort_order.asc().nulls_last(), Tariff.price_rub.asc()) if tariff_ids: try: ids = [int(x.strip()) for x in tariff_ids.split(",") if x.strip()] - if ids: - q = q.where(Tariff.id.in_(ids)) + if not ids: + return [] + q = q.where(Tariff.id.in_(ids)) except ValueError: - pass + raise HTTPException(status_code=422, detail="Некорректный параметр tariff_ids") elif group_code: q = q.where(Tariff.group_code == group_code) if filter_vless == "router": - q = q.where(Tariff.vless == True) + q = q.where(Tariff.vless is True) elif filter_vless == "app": - q = q.where(Tariff.vless == False) + q = q.where(Tariff.vless is False) result = await session.execute(q) rows = result.scalars().all() - return [_tariff_to_public(t) for t in rows] + payload = [_tariff_to_public(t).model_dump() for t in rows] + await cache_set(cache_token, payload, 30) + return payload + + +@public_router.get("/config-price", response_model=TariffConfigPriceResponse) +async def get_tariff_config_price( + tariff_id: int = Query(..., ge=1), + selected_device_limit: int | None = Query(None), + selected_traffic_gb: int | None = Query(None), + session: AsyncSession = Depends(get_session), +): + tariff = await get_tariff_by_id(session, tariff_id) + if not tariff or not tariff.get("is_active", True): + raise HTTPException(status_code=404, detail="Тариф не найден") + price = int(calculate_config_price(tariff, selected_device_limit, selected_traffic_gb)) + return TariffConfigPriceResponse(price_rub=price) + + +user_tariff_router = APIRouter() + + +@user_tariff_router.post("/purchase", response_model=TariffPurchaseResponse) +async def purchase_tariff_with_balance( + body: TariffPurchaseRequest, + request: Request, + preview: bool = Query(False), + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + tg_id = await idb.ensure_billing_user_for_identity(session, identity) + tariff = await get_tariff_by_id(session, body.tariff_id) + if not tariff or not tariff.get("is_active", True): + raise HTTPException(status_code=404, detail="Тариф не найден") + price = int(calculate_config_price(tariff, body.selected_device_limit, body.selected_traffic_gb)) + if price <= 0: + raise HTTPException(status_code=400, detail="Некорректная цена тарифа") + final_price, discount_rub, coupon_id, applied_coupon_code = await resolve_percent_coupon_pricing( + session=session, + billing_user_id=int(tg_id), + base_price_rub=int(price), + coupon_code=body.coupon_code, + ) + balance = float(await get_balance(session, tg_id)) + duration = int(tariff.get("duration_days") or 0) + if duration <= 0: + raise HTTPException(status_code=400, detail="Некорректная длительность тарифа") + required_amount = int(max(0, ceil(float(final_price) - balance))) + if preview: + return TariffPurchaseResponse( + ok=True, + message="Расчет обновлен", + key_email=None, + charged_rub=0, + base_price_rub=int(price), + discount_rub=int(discount_rub), + final_price_rub=int(final_price), + applied_coupon_code=applied_coupon_code, + payment_required=required_amount > 0, + required_amount_rub=int(required_amount), + payment_id=None, + payment_url=None, + ) + if required_amount > 0: + provider_id = str(body.provider_id or _resolve_default_web_payment_provider() or "").strip().upper() + if not provider_id: + raise HTTPException(status_code=503, detail="Нет доступных провайдеров оплаты") + base_url = _resolve_public_base_url(request) + success_url = validate_redirect_url(str(body.success_url or ""), f"{base_url}/payment-success") + failure_url = validate_redirect_url(str(body.failure_url or ""), f"{base_url}/payment-failure") + payment_request = PaymentLinkRequest( + legacy_user_ref=int(tg_id), + amount=required_amount, + currency="RUB", + provider_id=provider_id, + success_url=success_url, + failure_url=failure_url, + metadata={ + "payment_flow": "tariff_purchase", + "tariff_id": int(body.tariff_id), + "selected_device_limit": body.selected_device_limit, + "selected_traffic_gb": body.selected_traffic_gb, + "selected_duration_days": int(duration), + "selected_price_rub": int(final_price), + "base_price_rub": int(price), + "discount_rub": int(discount_rub), + "applied_coupon_code": applied_coupon_code, + "coupon_id": int(coupon_id) if coupon_id is not None else None, + }, + ) + payment_result = await create_payment_link(session, payment_request) + if not payment_result.success or not payment_result.payment_url or not payment_result.payment_id: + raise HTTPException(status_code=400, detail=payment_result.error or "Не удалось создать ссылку оплаты") + await create_temporary_data( + session, + int(tg_id), + "waiting_for_payment", + { + "tariff_id": int(body.tariff_id), + "required_amount": int(required_amount), + "selected_price_rub": int(final_price), + "selected_device_limit": body.selected_device_limit, + "selected_traffic_limit_gb": body.selected_traffic_gb, + "selected_duration_days": int(duration), + "base_price_rub": int(price), + "discount_rub": int(discount_rub), + "applied_coupon_code": applied_coupon_code, + "coupon_id": int(coupon_id) if coupon_id is not None else None, + }, + ) + return TariffPurchaseResponse( + ok=True, + message="Требуется оплата для оформления подписки", + key_email=None, + charged_rub=0, + base_price_rub=int(price), + discount_rub=int(discount_rub), + final_price_rub=int(final_price), + applied_coupon_code=applied_coupon_code, + payment_required=True, + required_amount_rub=required_amount, + payment_id=payment_result.payment_id, + payment_url=payment_result.payment_url, + ) + moscow_tz = tz_moscow("Europe/Moscow") + expiry = datetime.now(moscow_tz) + timedelta(days=duration) + try: + await create_vpn_key_headless( + session=session, + tg_id=tg_id, + expiry_time=expiry, + plan=body.tariff_id, + selected_device_limit=body.selected_device_limit, + selected_traffic_gb=body.selected_traffic_gb, + selected_price_rub=final_price, + ) + if coupon_id is not None: + await mark_coupon_used(session, int(coupon_id), int(tg_id)) + await session.commit() + except Exception: + logger.exception("web tariff purchase failed") + raise HTTPException(status_code=500, detail="Не удалось оформить подписку") from None + return TariffPurchaseResponse( + ok=True, + message="Подписка оформлена. Ключ в разделе «Мои ключи».", + key_email=None, + charged_rub=final_price, + base_price_rub=int(price), + discount_rub=int(discount_rub), + final_price_rub=int(final_price), + applied_coupon_code=applied_coupon_code, + ) + + +@user_tariff_router.post("/trial", response_model=TariffPurchaseResponse) +async def activate_trial( + request: Request, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_token), +): + """Активация триала (бесплатного или платного). Доступно 1 раз.""" + from database import get_trial, update_trial + from database.tariffs import get_tariffs + + tg_id = await idb.ensure_billing_user_for_identity(session, identity) + + trial_status = await get_trial(session, tg_id) + if trial_status not in (0, -1): + raise HTTPException(status_code=409, detail="Пробная подписка уже использована") + + trial_tariffs = await get_tariffs(session, group_code="trial") + if not trial_tariffs: + raise HTTPException(status_code=404, detail="Пробный тариф не найден") + + tariff = trial_tariffs[0] + price = int(tariff.get("price_rub", 0) or 0) + duration = int(tariff.get("duration_days") or 0) + if duration <= 0: + raise HTTPException(status_code=400, detail="Некорректная длительность триала") + + if price <= 0: + moscow_tz = tz_moscow("Europe/Moscow") + expiry = datetime.now(moscow_tz) + timedelta(days=duration) + try: + await create_vpn_key_headless( + session=session, + tg_id=tg_id, + expiry_time=expiry, + plan=int(tariff["id"]), + selected_price_rub=0, + skip_balance_charge=True, + is_trial=True, + ) + await update_trial(session, tg_id, 1) + except Exception: + logger.exception("web trial activation failed") + raise HTTPException(status_code=500, detail="Ошибка активации триала") from None + return TariffPurchaseResponse( + ok=True, + message="Пробная подписка активирована!", + charged_rub=0, + base_price_rub=0, + final_price_rub=0, + ) + + balance = float(await get_balance(session, tg_id)) + required_amount = int(max(0, ceil(float(price) - balance))) + + if required_amount <= 0: + moscow_tz = tz_moscow("Europe/Moscow") + expiry = datetime.now(moscow_tz) + timedelta(days=duration) + try: + await create_vpn_key_headless( + session=session, + tg_id=tg_id, + expiry_time=expiry, + plan=int(tariff["id"]), + selected_price_rub=price, + is_trial=True, + ) + await update_trial(session, tg_id, 1) + except Exception: + logger.exception("web paid trial activation failed") + raise HTTPException(status_code=500, detail="Ошибка активации триала") from None + return TariffPurchaseResponse( + ok=True, + message="Пробная подписка активирована!", + charged_rub=price, + base_price_rub=price, + final_price_rub=price, + ) + + provider_id = str(_resolve_default_web_payment_provider() or "").strip().upper() + if not provider_id: + raise HTTPException(status_code=503, detail="Нет доступных провайдеров оплаты") + base_url = _resolve_public_base_url(request) + payment_request = PaymentLinkRequest( + legacy_user_ref=int(tg_id), + amount=required_amount, + currency="RUB", + provider_id=provider_id, + success_url=f"{base_url}/payment-success", + failure_url=f"{base_url}/payment-failure", + metadata={ + "payment_flow": "trial_purchase", + "tariff_id": int(tariff["id"]), + "selected_price_rub": price, + "selected_duration_days": duration, + }, + ) + payment_result = await create_payment_link(session, payment_request) + if not payment_result.success or not payment_result.payment_url or not payment_result.payment_id: + raise HTTPException(status_code=400, detail=payment_result.error or "Не удалось создать ссылку оплаты") + await create_temporary_data( + session, + int(tg_id), + "waiting_for_payment", + { + "payment_flow": "trial_purchase", + "tariff_id": int(tariff["id"]), + "required_amount": required_amount, + "selected_price_rub": price, + "selected_duration_days": duration, + }, + ) + return TariffPurchaseResponse( + ok=True, + message="Требуется оплата для активации пробной подписки", + charged_rub=0, + base_price_rub=price, + final_price_rub=price, + payment_required=True, + required_amount_rub=required_amount, + payment_id=payment_result.payment_id, + payment_url=payment_result.payment_url, + ) router = generate_crud_router( diff --git a/api/v2/routes/users.py b/api/v2/routes/users.py index 0f334090..9fd37334 100644 --- a/api/v2/routes/users.py +++ b/api/v2/routes/users.py @@ -9,7 +9,8 @@ from api.v2.schemas import UserBase, UserResponse, UserUpdate from api.v2.base_crud import generate_crud_router from database import async_session_maker, delete_user_data, get_servers from database.models import Key, User -from handlers.keys.operations import delete_key_from_cluster +from database.access.resolution import resolve_user_optional +from services.operations import delete_key_from_cluster from logger import logger router = generate_crud_router( @@ -30,7 +31,10 @@ async def delete_user( ): """Удаляет пользователя и его ключи на серверах.""" try: - result = await session.execute(select(Key.email, Key.client_id).where(Key.tg_id == tg_id)) + u = await resolve_user_optional(session, tg_id) + if u is None: + raise HTTPException(status_code=404, detail="Пользователь не найден") + result = await session.execute(select(Key.email, Key.client_id).where(Key.user_id == u.id)) key_records = result.all() async with async_session_maker() as s: diff --git a/api/v2/routes/web.py b/api/v2/routes/web.py index 75b2f0f3..1a8b1b58 100644 --- a/api/v2/routes/web.py +++ b/api/v2/routes/web.py @@ -1,20 +1,72 @@ +import re import uuid + from pathlib import Path -from fastapi import APIRouter, Depends, File, HTTPException, UploadFile +from fastapi import APIRouter, Depends, File, HTTPException, Query, Request, UploadFile from pydantic import BaseModel -from sqlalchemy import select, delete +from datetime import datetime, timedelta, timezone + +from sqlalchemy import delete, func, select from sqlalchemy.ext.asyncio import AsyncSession from api.depends import get_session, verify_identity_admin -from api.v2.schemas import WebPageResponse, WebPageUpdate, WebBlockResponse, WebTheme -from api.v2.schemas.web import WebUploadResponse -from database.models import WebPage, WebBlock, WebTheme as WebThemeModel +from api.v2.schemas import WebBlockResponse, WebPageResponse, WebPageUpdate, WebTheme +from api.v2.schemas.web import ( + WebPageVariantCreate, + WebPageVariantSummary, + WebPageVariantUpdate, + WebPageVariantsResponse, + WebUploadResponse, +) +from database.models import ( + WebBlock, + WebCustomElementBuild, + WebFlow, + WebFlowEvent, + WebPage, + WebPageVariant, + WebPageVariantBlock, + WebTheme as WebThemeModel, +) +from logger import logger + UPLOAD_DIR = Path("static/web_uploads") ALLOWED_EXTENSIONS = frozenset({".png", ".jpg", ".jpeg", ".gif", ".webp", ".svg", ".mp4", ".webm"}) MAX_FILE_SIZE = 100 * 1024 * 1024 +_SLUG_RE = re.compile(r"^[a-z0-9][a-z0-9\-]*$") + +EXTENSION_CONTENT_TYPES: dict[str, frozenset[str]] = { + ".png": frozenset({"image/png"}), + ".jpg": frozenset({"image/jpeg"}), + ".jpeg": frozenset({"image/jpeg"}), + ".gif": frozenset({"image/gif"}), + ".webp": frozenset({"image/webp"}), + ".svg": frozenset({"image/svg+xml", "text/xml", "application/xml", "text/plain"}), + ".mp4": frozenset({"video/mp4"}), + ".webm": frozenset({"video/webm"}), +} + + +def _sanitize_svg(data: bytes) -> bytes: + import re as _re + text = data.decode("utf-8", errors="replace") + text = _re.sub(r"]*>.*?", "", text, flags=_re.DOTALL | _re.IGNORECASE) + text = _re.sub(r"]*>.*?", "", text, flags=_re.DOTALL | _re.IGNORECASE) + text = _re.sub(r"\bon\w+\s*=\s*[\"'][^\"']*[\"']", "", text, flags=_re.IGNORECASE) + text = _re.sub(r"\bon\w+\s*=\s*\S+", "", text, flags=_re.IGNORECASE) + text = _re.sub(r"(?:href|xlink:href)\s*=\s*[\"']\s*javascript:[^\"']*[\"']", "", text, flags=_re.IGNORECASE) + text = _re.sub(r"(?:href|xlink:href)\s*=\s*[\"']\s*data:\s*text/html[^\"']*[\"']", "", text, flags=_re.IGNORECASE) + text = _re.sub(r"(?:href|xlink:href)\s*=\s*[\"']\s*vbscript:[^\"']*[\"']", "", text, flags=_re.IGNORECASE) + text = _re.sub(r"]*>.*?", "", text, flags=_re.DOTALL | _re.IGNORECASE) + text = _re.sub(r"]*>.*?", "", text, flags=_re.DOTALL | _re.IGNORECASE) + text = _re.sub(r"]*>", "", text, flags=_re.IGNORECASE) + text = _re.sub(r"]*>.*?", "", text, flags=_re.DOTALL | _re.IGNORECASE) + return text.encode("utf-8") + + router = APIRouter(tags=["Web"]) @@ -22,7 +74,47 @@ class WebPagesListResponse(BaseModel): slugs: list[str] -KNOWN_PAGE_SLUGS = ["landing", "tariffs", "faq", "login", "dashboard"] +KNOWN_PAGE_SLUGS = [ + "landing", + "tariffs", + "faq", + "login", + "dashboard", + "checkout", + "gift-entry", + "referral-entry", + "partner-entry", + "payment-success", + "payment-failure", + "dashboard-keys", + "dashboard-profile", + "dashboard-instructions", + "dashboard-referrals", +] + +DEFAULT_VARIANT_KEY = "default" +DEFAULT_VARIANT_NAME = "Основной" + + +def _normalize_variant_key(value: str | None) -> str: + raw = (value or "").strip().lower() + normalized = re.sub(r"[^a-z0-9]+", "-", raw).strip("-") + if not normalized: + return DEFAULT_VARIANT_KEY + return normalized[:64].strip("-") or DEFAULT_VARIANT_KEY + + +def _normalize_variant_name(value: str | None, fallback: str) -> str: + name = (value or "").strip() + return name[:255] if name else fallback + + +def _variant_summary(row: WebPageVariant) -> WebPageVariantSummary: + return WebPageVariantSummary( + key=row.variant_key, + name=row.name or row.variant_key, + is_active=bool(row.is_active), + ) @router.get("/api/web/pages", response_model=WebPagesListResponse) @@ -46,69 +138,296 @@ async def get_or_create_page(session: AsyncSession, slug: str) -> WebPage: return page -@router.get("/api/web/pages/{slug}", response_model=WebPageResponse) -async def get_web_page( - slug: str, - session: AsyncSession = Depends(get_session), -): - await get_or_create_page(session, slug) +async def _list_variants(session: AsyncSession, slug: str) -> list[WebPageVariant]: + result = await session.execute( + select(WebPageVariant) + .where(WebPageVariant.page_slug == slug) + .order_by(WebPageVariant.is_active.desc(), WebPageVariant.created_at, WebPageVariant.variant_key) + ) + return list(result.scalars().all()) + +async def _get_theme_tokens_for_legacy_page(session: AsyncSession, slug: str) -> dict: + theme_result = await session.execute(select(WebThemeModel).where(WebThemeModel.page_slug == slug)) + theme_row = theme_result.scalar_one_or_none() + return dict(theme_row.tokens or {}) if theme_row else {} + + +async def _get_legacy_blocks(session: AsyncSession, slug: str) -> list[WebBlock]: blocks_result = await session.execute( select(WebBlock).where(WebBlock.page_slug == slug).order_by(WebBlock.order, WebBlock.id) ) - blocks = [WebBlockResponse.model_validate(b) for b in blocks_result.scalars().all()] + return list(blocks_result.scalars().all()) - theme_result = await session.execute(select(WebThemeModel).where(WebThemeModel.page_slug == slug)) - theme_row = theme_result.scalar_one_or_none() - theme = WebTheme(tokens=theme_row.tokens) if theme_row else None - return WebPageResponse(slug=slug, blocks=blocks, theme=theme) +async def _ensure_page_variants(session: AsyncSession, slug: str) -> list[WebPageVariant]: + await get_or_create_page(session, slug) + variants = await _list_variants(session, slug) + if variants: + if not any(variant.is_active for variant in variants): + variants[0].is_active = True + await session.flush() + variants = await _list_variants(session, slug) + return variants + + legacy_blocks = await _get_legacy_blocks(session, slug) + theme_tokens = await _get_theme_tokens_for_legacy_page(session, slug) + variant = WebPageVariant( + page_slug=slug, + variant_key=DEFAULT_VARIANT_KEY, + name=DEFAULT_VARIANT_NAME, + is_active=True, + theme_tokens=theme_tokens, + ) + session.add(variant) + await session.flush() + for legacy_block in legacy_blocks: + session.add( + WebPageVariantBlock( + variant_id=variant.id, + order=legacy_block.order, + type=legacy_block.type, + data=legacy_block.data, + ) + ) + await session.flush() + return await _list_variants(session, slug) + + +async def _resolve_variant( + session: AsyncSession, + slug: str, + variant_key: str | None, +) -> tuple[WebPageVariant, list[WebPageVariant]]: + variants = await _ensure_page_variants(session, slug) + desired_key = _normalize_variant_key(variant_key) if variant_key else "" + current = None + if desired_key: + current = next((variant for variant in variants if variant.variant_key == desired_key), None) + if current is None: + raise HTTPException(404, "Вариант страницы не найден") + else: + current = next((variant for variant in variants if variant.is_active), variants[0]) + return current, variants + + +async def _get_variant_blocks(session: AsyncSession, variant_id: str) -> list[WebBlockResponse]: + blocks_result = await session.execute( + select(WebPageVariantBlock) + .where(WebPageVariantBlock.variant_id == variant_id) + .order_by(WebPageVariantBlock.order, WebPageVariantBlock.id) + ) + return [WebBlockResponse.model_validate(block) for block in blocks_result.scalars().all()] + + +async def _build_page_response( + session: AsyncSession, + slug: str, + current: WebPageVariant, + variants: list[WebPageVariant] | None = None, +) -> WebPageResponse: + current_variants = variants or await _list_variants(session, slug) + active = next((variant for variant in current_variants if variant.is_active), current) + blocks = await _get_variant_blocks(session, current.id) + theme = WebTheme(tokens=dict(current.theme_tokens or {})) + return WebPageResponse( + slug=slug, + blocks=blocks, + theme=theme, + variant_key=current.variant_key, + active_variant_key=active.variant_key, + variants=[_variant_summary(variant) for variant in current_variants], + ) + + +async def _set_active_variant(session: AsyncSession, slug: str, variant_key: str) -> list[WebPageVariant]: + variants = await _list_variants(session, slug) + matched = False + for variant in variants: + is_target = variant.variant_key == variant_key + variant.is_active = is_target + matched = matched or is_target + if not matched: + raise HTTPException(404, "Вариант страницы не найден") + await session.flush() + return await _list_variants(session, slug) + + +def _generate_variant_key(existing_keys: set[str], requested_key: str | None, requested_name: str | None) -> str: + base = _normalize_variant_key(requested_key or requested_name) + if not base: + base = DEFAULT_VARIANT_KEY + if base not in existing_keys: + return base + suffix = 2 + while True: + candidate = f"{base}-{suffix}" + if candidate not in existing_keys: + return candidate[:64] + suffix += 1 + + +@router.get("/api/web/pages/{slug}", response_model=WebPageResponse) +async def get_web_page( + slug: str, + variant: str | None = Query(default=None), + session: AsyncSession = Depends(get_session), +): + if not slug or len(slug) > 64 or not _SLUG_RE.match(slug): + raise HTTPException(400, "Некорректный slug страницы") + current, variants = await _resolve_variant(session, slug, variant) + return await _build_page_response(session, slug, current, variants) @router.put("/api/web/pages/{slug}", response_model=WebPageResponse) async def update_web_page( slug: str, body: WebPageUpdate, + variant: str | None = Query(default=None), session: AsyncSession = Depends(get_session), identity=Depends(verify_identity_admin), ): - await get_or_create_page(session, slug) - - await session.execute(delete(WebBlock).where(WebBlock.page_slug == slug)) + current, _ = await _resolve_variant(session, slug, variant) + await session.execute(delete(WebPageVariantBlock).where(WebPageVariantBlock.variant_id == current.id)) for block in body.blocks: session.add( - WebBlock( - page_slug=slug, + WebPageVariantBlock( + variant_id=current.id, order=block.order, type=block.type, data=block.data, ) ) - theme_row = None if body.theme is not None: - result = await session.execute(select(WebThemeModel).where(WebThemeModel.page_slug == slug)) - theme_row = result.scalar_one_or_none() - if theme_row is None: - theme_row = WebThemeModel(page_slug=slug, tokens=body.theme.tokens) - session.add(theme_row) - else: - theme_row.tokens = body.theme.tokens + current.theme_tokens = body.theme.tokens await session.flush() + refreshed_variants = await _list_variants(session, slug) + refreshed_current = next((item for item in refreshed_variants if item.id == current.id), current) + return await _build_page_response(session, slug, refreshed_current, refreshed_variants) - blocks_result = await session.execute( - select(WebBlock).where(WebBlock.page_slug == slug).order_by(WebBlock.order, WebBlock.id) + +@router.get("/api/web/pages/{slug}/variants", response_model=WebPageVariantsResponse) +async def get_web_page_variants( + slug: str, + variant: str | None = Query(default=None), + session: AsyncSession = Depends(get_session), +): + current, variants = await _resolve_variant(session, slug, variant) + active = next((item for item in variants if item.is_active), current) + return WebPageVariantsResponse( + slug=slug, + active_variant_key=active.variant_key, + current_variant_key=current.variant_key, + variants=[_variant_summary(item) for item in variants], ) - blocks = [WebBlockResponse.model_validate(b) for b in blocks_result.scalars().all()] - if theme_row is None: - theme_result = await session.execute(select(WebThemeModel).where(WebThemeModel.page_slug == slug)) - theme_row = theme_result.scalar_one_or_none() - theme = WebTheme(tokens=theme_row.tokens) if theme_row else None - return WebPageResponse(slug=slug, blocks=blocks, theme=theme) +@router.post("/api/web/pages/{slug}/variants", response_model=WebPageVariantsResponse) +async def create_web_page_variant( + slug: str, + body: WebPageVariantCreate, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_admin), +): + source_variant, variants = await _resolve_variant(session, slug, body.from_variant_key) + existing_keys = {variant.variant_key for variant in variants} + variant_key = _generate_variant_key(existing_keys, body.key, body.name) + if variant_key in existing_keys: + raise HTTPException(400, "Вариант с таким ключом уже существует") + + variant_name = _normalize_variant_name(body.name, f"Вариант {len(variants) + 1}") + new_variant = WebPageVariant( + page_slug=slug, + variant_key=variant_key, + name=variant_name, + is_active=False, + theme_tokens=dict(source_variant.theme_tokens or {}), + ) + session.add(new_variant) + await session.flush() + + source_blocks = await _get_variant_blocks(session, source_variant.id) + for block in source_blocks: + session.add( + WebPageVariantBlock( + variant_id=new_variant.id, + order=block.order, + type=block.type, + data=block.data, + ) + ) + await session.flush() + + refreshed = await _list_variants(session, slug) + return WebPageVariantsResponse( + slug=slug, + active_variant_key=next((item.variant_key for item in refreshed if item.is_active), DEFAULT_VARIANT_KEY), + current_variant_key=new_variant.variant_key, + variants=[_variant_summary(item) for item in refreshed], + ) + + +@router.patch("/api/web/pages/{slug}/variants/{variant_key}", response_model=WebPageVariantsResponse) +async def update_web_page_variant( + slug: str, + variant_key: str, + body: WebPageVariantUpdate, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_admin), +): + current, variants = await _resolve_variant(session, slug, variant_key) + if body.name is not None: + current.name = _normalize_variant_name(body.name, current.name or current.variant_key) + if body.make_active is True: + variants = await _set_active_variant(session, slug, current.variant_key) + current = next((item for item in variants if item.variant_key == current.variant_key), current) + else: + await session.flush() + variants = await _list_variants(session, slug) + + active = next((item for item in variants if item.is_active), current) + return WebPageVariantsResponse( + slug=slug, + active_variant_key=active.variant_key, + current_variant_key=current.variant_key, + variants=[_variant_summary(item) for item in variants], + ) + + +@router.delete("/api/web/pages/{slug}/variants/{variant_key}", response_model=WebPageVariantsResponse) +async def delete_web_page_variant( + slug: str, + variant_key: str, + session: AsyncSession = Depends(get_session), + identity=Depends(verify_identity_admin), +): + current, variants = await _resolve_variant(session, slug, variant_key) + if len(variants) <= 1: + raise HTTPException(400, "Нельзя удалить единственный вариант страницы") + + replacement = next((item for item in variants if item.variant_key != current.variant_key), None) + await session.execute(delete(WebPageVariant).where(WebPageVariant.id == current.id)) + await session.flush() + + if current.is_active and replacement is not None: + replacement_variants = await _set_active_variant(session, slug, replacement.variant_key) + else: + replacement_variants = await _list_variants(session, slug) + + current_variant_key = replacement.variant_key if replacement is not None else DEFAULT_VARIANT_KEY + active_variant_key = next( + (item.variant_key for item in replacement_variants if item.is_active), + current_variant_key, + ) + return WebPageVariantsResponse( + slug=slug, + active_variant_key=active_variant_key, + current_variant_key=current_variant_key, + variants=[_variant_summary(item) for item in replacement_variants], + ) @router.post("/api/web/upload", response_model=WebUploadResponse) @@ -125,18 +444,271 @@ async def upload_media( 400, f"Разрешены только: {', '.join(sorted(ALLOWED_EXTENSIONS))}", ) + if file.content_type: + allowed_types = EXTENSION_CONTENT_TYPES.get(ext) + if allowed_types and file.content_type.lower() not in allowed_types: + raise HTTPException( + 400, + f"Тип файла ({file.content_type}) не соответствует расширению ({ext})", + ) UPLOAD_DIR.mkdir(parents=True, exist_ok=True) + chunks: list[bytes] = [] size = 0 for chunk in file.file: size += len(chunk) if size > MAX_FILE_SIZE: - raise HTTPException(400, f"Размер файла не более {MAX_FILE_SIZE // (1024*1024)} МБ") - await file.seek(0) + raise HTTPException(400, f"Размер файла не более {MAX_FILE_SIZE // (1024 * 1024)} МБ") + chunks.append(chunk) name = f"{uuid.uuid4().hex}{ext}" path = UPLOAD_DIR / name + file_data = b"".join(chunks) + if ext == ".svg": + file_data = _sanitize_svg(file_data) with open(path, "wb") as f: - while chunk := await file.read(64 * 1024): - f.write(chunk) + f.write(file_data) url = f"/api/web/uploads/{name}" + logger.info( + "[WebUpload] admin={} file={} -> {} ({} bytes)", + identity.id, + file.filename, + name, + len(file_data), + ) return WebUploadResponse(url=url) + +# ── Custom Element Builds ── + + +class CustomElementBuildCreate(BaseModel): + label: str = "" + slug: str = "" + runtime: str = "react-component" + source_kind: str = "inline-code" + source_value: str = "" + export_name: str = "default" + props_schema_text: str = "" + sample_props_text: str = "" + events_text: str = "" + notes: str = "" + + +class CustomElementBuildUpdate(BaseModel): + status: str | None = None + summary: str | None = None + next_steps: list[str] | None = None + artifact: dict | None = None + upload_meta: dict | None = None + worker_id: str | None = None + + +def _build_to_dict(b: WebCustomElementBuild) -> dict: + return { + "id": b.id, + "label": b.label, + "slug": b.slug, + "runtime": b.runtime, + "sourceKind": b.source_kind, + "sourceValue": b.source_value, + "exportName": b.export_name, + "propsSchemaText": b.props_schema_text, + "samplePropsText": b.sample_props_text, + "eventsText": b.events_text, + "notes": b.notes, + "status": b.status, + "summary": b.summary, + "nextSteps": b.next_steps or [], + "artifact": b.artifact, + "upload": b.upload_meta, + "workerId": b.worker_id, + "workerClaimedAt": b.worker_claimed_at.isoformat() if b.worker_claimed_at else None, + "completedAt": b.completed_at.isoformat() if b.completed_at else None, + "createdAt": b.created_at.isoformat() if b.created_at else None, + "updatedAt": b.updated_at.isoformat() if b.updated_at else None, + } + + +@router.get("/custom-element-builds") +async def list_custom_element_builds( + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + result = await session.execute( + select(WebCustomElementBuild).order_by(WebCustomElementBuild.created_at.desc()) + ) + builds = result.scalars().all() + return [_build_to_dict(b) for b in builds] + + +@router.post("/custom-element-builds") +async def create_custom_element_build( + body: CustomElementBuildCreate, + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + build = WebCustomElementBuild( + id=str(uuid.uuid4()), + label=body.label, + slug=body.slug, + runtime=body.runtime, + source_kind=body.source_kind, + source_value=body.source_value, + export_name=body.export_name, + props_schema_text=body.props_schema_text, + sample_props_text=body.sample_props_text, + events_text=body.events_text, + notes=body.notes, + status="queued", + ) + session.add(build) + return _build_to_dict(build) + + +@router.get("/custom-element-builds/{build_id}") +async def get_custom_element_build( + build_id: str, + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + build = await session.get(WebCustomElementBuild, build_id) + if not build: + raise HTTPException(404, "Build not found") + return _build_to_dict(build) + + +@router.patch("/custom-element-builds/{build_id}") +async def update_custom_element_build( + build_id: str, + body: CustomElementBuildUpdate, + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + build = await session.get(WebCustomElementBuild, build_id) + if not build: + raise HTTPException(404, "Build not found") + if body.status is not None: + build.status = body.status + if body.summary is not None: + build.summary = body.summary + if body.next_steps is not None: + build.next_steps = body.next_steps + if body.artifact is not None: + build.artifact = body.artifact + if body.upload_meta is not None: + build.upload_meta = body.upload_meta + if body.worker_id is not None: + build.worker_id = body.worker_id + return _build_to_dict(build) + + +@router.delete("/custom-element-builds/{build_id}") +async def delete_custom_element_build( + build_id: str, + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + build = await session.get(WebCustomElementBuild, build_id) + if not build: + raise HTTPException(404, "Build not found") + await session.delete(build) + return {"ok": True} + + +# ── Flow Analytics ── + + +class FlowEventBatch(BaseModel): + events: list[dict] + + +@router.post("/analytics/flow-events") +async def ingest_flow_events( + body: FlowEventBatch, + request: Request, + session: AsyncSession = Depends(get_session), +): + try: + from core.redis_cache import cache_incr_checked + from api.v2.routes.auth._fallback_limiter import check_and_increment + ip = (request.client.host if request.client else "") or "unknown" + count, redis_ok = await cache_incr_checked(f"analytics_rate:{ip}", 60) + if not redis_ok: + count = check_and_increment(f"analytics_rate:{ip}", 60, 60) + if count > 60: + raise HTTPException(status_code=429, detail="Too many events") + except HTTPException: + raise + except Exception: + pass + created = 0 + for raw in body.events[:100]: + flow_id = str(raw.get("flowId", "")) + node_id = str(raw.get("nodeId", "")) + event_type = str(raw.get("eventType", "")) + if not flow_id or not node_id or not event_type: + continue + ev = WebFlowEvent( + id=str(uuid.uuid4()), + flow_id=flow_id, + node_id=node_id, + node_type=str(raw.get("nodeType", "")), + event_type=event_type, + ab_variant=raw.get("abVariant") or None, + device=raw.get("device") or None, + locale=raw.get("locale") or None, + authenticated=raw.get("authenticated"), + event_metadata=raw.get("collectedDataSnapshot") or None, + ) + session.add(ev) + created += 1 + return {"ingested": created} + + +@router.get("/analytics/flow-funnel/{flow_id}") +async def get_flow_funnel( + flow_id: str, + days: int = Query(default=30, ge=1, le=365), + session: AsyncSession = Depends(get_session), + _identity=Depends(verify_identity_admin), +): + since = datetime.now(timezone.utc) - timedelta(days=days) + rows = ( + await session.execute( + select( + WebFlowEvent.node_id, + WebFlowEvent.node_type, + WebFlowEvent.event_type, + func.count().label("cnt"), + ) + .where(WebFlowEvent.flow_id == flow_id) + .where(WebFlowEvent.created_at >= since) + .group_by(WebFlowEvent.node_id, WebFlowEvent.node_type, WebFlowEvent.event_type) + ) + ).all() + + nodes: dict[str, dict] = {} + for node_id, node_type, event_type, cnt in rows: + if node_id not in nodes: + nodes[node_id] = {"nodeId": node_id, "nodeType": node_type, "entered": 0, "exited": 0, "completed": 0} + if event_type == "flow_step_entered": + nodes[node_id]["entered"] = cnt + elif event_type == "flow_step_exited": + nodes[node_id]["exited"] = cnt + elif event_type == "flow_completed": + nodes[node_id]["completed"] = cnt + + flow = await session.get(WebFlow, flow_id) + if flow and flow.nodes: + node_order = {n["id"]: i for i, n in enumerate(flow.nodes) if isinstance(n, dict)} + else: + node_order = {} + + funnel = sorted(nodes.values(), key=lambda n: node_order.get(n["nodeId"], 999)) + + for i, node in enumerate(funnel): + prev_entered = funnel[i - 1]["entered"] if i > 0 else node["entered"] + node["dropOff"] = round( + (1 - node["entered"] / prev_entered) * 100, 1 + ) if prev_entered > 0 else 0 + + return {"flowId": flow_id, "days": days, "funnel": funnel} diff --git a/api/v2/schemas/__init__.py b/api/v2/schemas/__init__.py index eb7c015b..f0bb5f43 100644 --- a/api/v2/schemas/__init__.py +++ b/api/v2/schemas/__init__.py @@ -28,4 +28,13 @@ from api.v1.schemas import ( ) from api.v1.schemas.keys import KeyBase, KeyCreateRequest, KeyUpdate from api.v1.schemas.settings import SettingResponse, SettingUpsert -from api.v2.schemas.web import WebBlockResponse, WebTheme, WebPageResponse, WebPageUpdate \ No newline at end of file +from api.v2.schemas.web import ( + WebBlockResponse, + WebTheme, + WebPageResponse, + WebPageUpdate, + WebPageVariantCreate, + WebPageVariantSummary, + WebPageVariantUpdate, + WebPageVariantsResponse, +) diff --git a/api/v2/schemas/flows.py b/api/v2/schemas/flows.py new file mode 100644 index 00000000..fadbc528 --- /dev/null +++ b/api/v2/schemas/flows.py @@ -0,0 +1,55 @@ +from __future__ import annotations + +from typing import Any + +from pydantic import BaseModel + + +class EdgeConditionSchema(BaseModel): + field: str + operator: str + value: Any = None + + +class FlowEdgeSchema(BaseModel): + id: str + source: str + target: str + condition: EdgeConditionSchema | None = None + label: str | None = None + priority: int | None = None + + +class FlowNodeSchema(BaseModel): + id: str + type: str + label: str + label_en: str | None = None + enabled: bool = True + page_slug: str | None = None + config: dict = {} + position: dict + + +class FlowResponse(BaseModel): + id: str + name: str + nodes: list[FlowNodeSchema] + edges: list[FlowEdgeSchema] + entry_node_id: str | None + version: int + + +class FlowUpdate(BaseModel): + name: str | None = None + nodes: list[FlowNodeSchema] + edges: list[FlowEdgeSchema] + entry_node_id: str | None = None + + +class FlowCreate(BaseModel): + id: str + name: str + nodes: list[FlowNodeSchema] = [] + edges: list[FlowEdgeSchema] = [] + entry_node_id: str | None = None diff --git a/api/v2/schemas/identities.py b/api/v2/schemas/identities.py index 36a1bdaa..4e95bab7 100644 --- a/api/v2/schemas/identities.py +++ b/api/v2/schemas/identities.py @@ -13,6 +13,8 @@ class IdentityResponse(BaseModel): email: str | None tg_id: int | None is_admin: bool = False + email_verified: bool = False + password_set: bool = False created_at: datetime | None updated_at: datetime | None @@ -23,11 +25,12 @@ class IdentityResponse(BaseModel): class RegisterByEmailRequest(BaseModel): email: str = Field(..., min_length=1) password: str = Field(..., min_length=8, description="Пароль (минимум 8 символов)") + referral_code: str | None = Field(None, min_length=1) + turnstile_token: str | None = Field(default=None, description="Cloudflare Turnstile CAPTCHA token") class RegisterResponse(BaseModel): identity_id: str - token: str class LoginRequest(BaseModel): @@ -35,13 +38,28 @@ class LoginRequest(BaseModel): password: str = Field(...) +class SetPasswordRequest(BaseModel): + password: str = Field(..., min_length=8, description="Новый пароль (минимум 8 символов)") + password_confirm: str = Field(..., min_length=8) + + +class ChangePasswordRequest(BaseModel): + current_password: str = Field(...) + password: str = Field(..., min_length=8, description="Новый пароль (минимум 8 символов)") + password_confirm: str = Field(..., min_length=8) + + class LoginResponse(BaseModel): identity_id: str - token: str class SendLoginCodeRequest(BaseModel): email: str = Field(..., min_length=1) + allow_register: bool = Field( + default=False, + description="Если true и email новый — создать идентичность и отправить код (гостевой вход с сайта)", + ) + turnstile_token: str | None = Field(default=None, description="Cloudflare Turnstile CAPTCHA token") class LoginByCodeRequest(BaseModel): @@ -49,6 +67,13 @@ class LoginByCodeRequest(BaseModel): code: str = Field(..., min_length=1) +class ConfirmPasswordResetRequest(BaseModel): + email: str = Field(..., min_length=1) + code: str = Field(..., min_length=1) + password: str = Field(..., min_length=8) + password_confirm: str = Field(..., min_length=8) + + class LoginTelegramRequest(BaseModel): """Данные от Telegram Login Widget (кнопка «Войти через Telegram»).""" @@ -77,5 +102,14 @@ class IdentityAttachEmail(BaseModel): email: str = Field(..., min_length=1) +class LinkEmailSendCodeRequest(BaseModel): + email: str = Field(..., min_length=1) + + +class LinkEmailConfirmRequest(BaseModel): + email: str = Field(..., min_length=1) + code: str = Field(..., min_length=1, max_length=16) + + class IdentityAttachTelegram(BaseModel): tg_id: int = Field(...) diff --git a/api/v2/schemas/payment_links.py b/api/v2/schemas/payment_links.py index cdfc46a6..1b9efc4c 100644 --- a/api/v2/schemas/payment_links.py +++ b/api/v2/schemas/payment_links.py @@ -22,3 +22,11 @@ class PaymentLinkCreateResponse(BaseModel): payment_id: str | None = None payment_url: str | None = None error: str | None = None + + +class PaymentLinkStatusResponse(BaseModel): + success: bool + payment_id: str + status: str | None = None + completed: bool = False + paid: bool = False diff --git a/api/v2/schemas/tariffs.py b/api/v2/schemas/tariffs.py index c9f8442f..0552187b 100644 --- a/api/v2/schemas/tariffs.py +++ b/api/v2/schemas/tariffs.py @@ -17,4 +17,7 @@ class TariffPublic(BaseModel): device_limit: int | None subgroup_title: str | None sort_order: int | None - vless: bool = False \ No newline at end of file + vless: bool = False + configurable: bool = False + device_options: list[int] | None = None + traffic_options_gb: list[int] | None = None diff --git a/api/v2/schemas/web.py b/api/v2/schemas/web.py index 7721ebbe..c023d6d8 100644 --- a/api/v2/schemas/web.py +++ b/api/v2/schemas/web.py @@ -1,6 +1,9 @@ +import json from typing import Any -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator + +_MAX_BLOCK_DATA_SIZE = 256 * 1024 class WebBlockBase(BaseModel): @@ -8,6 +11,12 @@ class WebBlockBase(BaseModel): order: int data: dict[str, Any] + @model_validator(mode="after") + def _check_data_size(self) -> "WebBlockBase": + if len(json.dumps(self.data, ensure_ascii=False)) > _MAX_BLOCK_DATA_SIZE: + raise ValueError(f"Размер data блока не должен превышать {_MAX_BLOCK_DATA_SIZE // 1024} КБ") + return self + class WebBlockResponse(WebBlockBase): id: str @@ -20,10 +29,19 @@ class WebTheme(BaseModel): tokens: dict[str, Any] +class WebPageVariantSummary(BaseModel): + key: str = Field(..., max_length=64) + name: str = Field(..., max_length=255) + is_active: bool = False + + class WebPageResponse(BaseModel): slug: str blocks: list[WebBlockResponse] theme: WebTheme | None = None + variant_key: str = "default" + active_variant_key: str = "default" + variants: list[WebPageVariantSummary] = Field(default_factory=list) class WebPageUpdate(BaseModel): @@ -31,6 +49,58 @@ class WebPageUpdate(BaseModel): theme: WebTheme | None = None +class WebPageVariantCreate(BaseModel): + key: str | None = Field(default=None, max_length=64) + name: str | None = Field(default=None, max_length=255) + from_variant_key: str | None = Field(default=None, max_length=64) + + +class WebPageVariantUpdate(BaseModel): + name: str | None = Field(default=None, max_length=255) + make_active: bool | None = None + + +class WebPageVariantsResponse(BaseModel): + slug: str + active_variant_key: str = "default" + current_variant_key: str = "default" + variants: list[WebPageVariantSummary] = Field(default_factory=list) + + class WebUploadResponse(BaseModel): url: str + +class FlowStepConfig(BaseModel): + provider_ids: list[str] | None = None + tariff_group_code: str | None = None + tariff_ids: list[int] | None = None + display_mode: str | None = None + skippable: bool = False + auto_advance_if_single: bool = False + + +class FlowStepSchema(BaseModel): + id: str = Field(..., max_length=64) + type: str = Field(..., max_length=32) + label: str = Field(..., max_length=255) + label_en: str | None = Field(default=None, max_length=255) + enabled: bool = True + page_slug: str | None = Field(default=None, max_length=64) + config: FlowStepConfig = Field(default_factory=FlowStepConfig) + + +class FlowDefinitionSchema(BaseModel): + id: str = Field(..., max_length=64) + name: str = Field(..., max_length=255) + steps: list[FlowStepSchema] = Field(default_factory=list) + version: int = 1 + + +class FlowDefinitionResponse(FlowDefinitionSchema): + pass + + +class FlowDefinitionUpdate(BaseModel): + name: str | None = Field(default=None, max_length=255) + steps: list[FlowStepSchema] = Field(default_factory=list) diff --git a/api/v2/schemas/web_public.py b/api/v2/schemas/web_public.py new file mode 100644 index 00000000..1f2ef4a1 --- /dev/null +++ b/api/v2/schemas/web_public.py @@ -0,0 +1,425 @@ +from pydantic import BaseModel, Field + + +class AccountSummaryResponse(BaseModel): + identity_id: str + email: str | None = None + tg_id: int | None = None + linked_telegram: bool = False + referral_code: str = "" + balance: float = 0.0 + trial_status: int = 0 + keys_total: int = 0 + referrals_total: int = 0 + referrals_active: int = 0 + referral_bonus_total: float = 0.0 + gifts_sent: int = 0 + gifts_claimed: int = 0 + coupons_used: int = 0 + partner_enabled: bool = False + partner_code: str = "" + partner_balance: float = 0.0 + partner_percent: float = 0.0 + partner_percent_custom: bool = False + partner_referred_total: int = 0 + partner_payout_method: str | None = None + unread_notifications: int = 0 + + +class AccountKeyActionsAvailability(BaseModel): + can_connect_device: bool = False + can_connect_router: bool = False + can_connect_tv: bool = False + can_renew: bool = False + can_addons: bool = False + can_reset_hwid: bool = False + can_qr: bool = False + can_delete: bool = False + can_change_location: bool = False + + +class AccountKeyDetailsResponse(BaseModel): + client_id: str + email: str + alias: str | None = None + expiry_time: int = 0 + is_frozen: bool = False + tariff_name: str = "" + subgroup_title: str = "" + traffic_limit_gb: int = 0 + used_traffic_gb: float | None = None + device_limit: int = 0 + connected_devices: int = 0 + is_tariff_configurable: bool = False + addons_devices_enabled: bool = False + addons_traffic_enabled: bool = False + + +class AccountKeyResponse(BaseModel): + email: str + alias: str | None = None + client_id: str + tariff_id: int | None = None + server_id: str + created_at: int = 0 + expiry_time: int = 0 + key: str | None = None + remnawave_link: str | None = None + is_frozen: bool = False + actions: AccountKeyActionsAvailability | None = None + + +class AccountKeyAliasUpdateRequest(BaseModel): + alias: str = Field(..., min_length=1, max_length=10) + + +class AccountKeyActionResponse(BaseModel): + ok: bool = True + message: str = "" + + +class AccountKeyRenewRequest(BaseModel): + provider_id: str | None = None + success_url: str | None = None + failure_url: str | None = None + coupon_code: str | None = None + + +class AccountKeyRenewResponse(AccountKeyActionResponse): + client_id: str + tariff_id: int + charged_rub: int = 0 + balance_rub: float = 0.0 + base_price_rub: int = 0 + discount_rub: int = 0 + final_price_rub: int = 0 + applied_coupon_code: str | None = None + payment_required: bool = False + required_amount_rub: int = 0 + payment_id: str | None = None + payment_url: str | None = None + + +class AccountKeyResetHwidResponse(AccountKeyActionResponse): + total_devices: int = 0 + reset_devices: int = 0 + + +class AccountKeyQrResponse(AccountKeyActionResponse): + link: str = "" + image_data_url: str = "" + + +class AccountKeyLocationOptionResponse(BaseModel): + server_name: str + + +class AccountKeyLocationsResponse(BaseModel): + client_id: str + current_server: str = "" + locations: list[AccountKeyLocationOptionResponse] = [] + + +class AccountKeyChangeLocationRequest(BaseModel): + server_name: str = Field(..., min_length=1) + + +class AccountKeyChangeLocationResponse(AccountKeyActionResponse): + client_id: str + server_id: str = "" + link: str = "" + remnawave_link: str | None = None + + +class AccountKeyAddonOptionResponse(BaseModel): + value: int + label: str = "" + + +class AccountKeyAddonsPreviewRequest(BaseModel): + selected_device_limit: int | None = None + selected_traffic_gb: int | None = None + include_device: bool | None = None + include_traffic: bool | None = None + provider_id: str | None = None + success_url: str | None = None + failure_url: str | None = None + coupon_code: str | None = None + + +class AccountKeyAddonsPreviewResponse(BaseModel): + client_id: str + tariff_id: int + addons_mode: str = "" + has_device_option: bool = False + has_traffic_option: bool = False + current_device_limit: int | None = None + current_traffic_gb: int | None = None + selected_device_limit: int | None = None + selected_traffic_gb: int | None = None + device_options: list[AccountKeyAddonOptionResponse] = [] + traffic_options: list[AccountKeyAddonOptionResponse] = [] + total_price_rub: int = 0 + extra_price_rub: int = 0 + discount_rub: int = 0 + final_price_rub: int = 0 + applied_coupon_code: str | None = None + balance_rub: float = 0.0 + + +class AccountKeyApplyAddonsResponse(AccountKeyActionResponse): + client_id: str + tariff_id: int + total_price_rub: int = 0 + extra_price_rub: int = 0 + discount_rub: int = 0 + final_price_rub: int = 0 + applied_coupon_code: str | None = None + charged_rub: int = 0 + balance_rub: float = 0.0 + payment_required: bool = False + required_amount_rub: int = 0 + payment_id: str | None = None + payment_url: str | None = None + + +class AccountKeyActionsConfigResponse(BaseModel): + renew_enabled: bool = True + delete_enabled: bool = False + qr_enabled: bool = False + hwid_reset_enabled: bool = False + country_change_enabled: bool = False + instructions_enabled: bool = False + addons_enabled: bool = False + addons_mode: str = "" + tv_connect_enabled: bool = False + + +class TariffConfigPriceResponse(BaseModel): + price_rub: int + + +class TariffPurchaseRequest(BaseModel): + tariff_id: int = Field(..., ge=1) + selected_device_limit: int | None = None + selected_traffic_gb: int | None = None + provider_id: str | None = None + success_url: str | None = None + failure_url: str | None = None + coupon_code: str | None = None + + +class TariffPurchaseResponse(BaseModel): + ok: bool = True + message: str = "" + key_email: str | None = None + charged_rub: int | None = None + base_price_rub: int = 0 + discount_rub: int = 0 + final_price_rub: int = 0 + applied_coupon_code: str | None = None + payment_required: bool = False + required_amount_rub: int = 0 + payment_id: str | None = None + payment_url: str | None = None + + +class GiftCreateRequest(BaseModel): + tariff_id: int = Field(..., ge=1) + selected_device_limit: int | None = None + selected_traffic_gb: int | None = None + provider_id: str | None = None + success_url: str | None = None + failure_url: str | None = None + + +class GiftCreatePreviewResponse(BaseModel): + ok: bool = True + price_rub: int = 0 + balance_rub: float = 0.0 + sufficient_funds: bool = True + tariff_name: str = "" + duration_days: int = 0 + + +class GiftCreateResponse(BaseModel): + ok: bool = True + message: str = "" + gift_id: str = "" + site_gift_link: str = "" + tariff_name: str = "" + duration_days: int = 0 + price_charged: int = 0 + balance_rub: float = 0.0 + payment_required: bool = False + required_amount_rub: int = 0 + payment_id: str | None = None + payment_url: str | None = None + + +class GiftUsageEntry(BaseModel): + user_id: int + used_at: str | None = None + + +class MyGiftItem(BaseModel): + gift_id: str + tariff_name: str = "" + duration_days: int = 0 + price_rub: int = 0 + created_at: str | None = None + expiry_time: str | None = None + is_used: bool = False + is_unlimited: bool = False + max_usages: int | None = None + site_gift_link: str = "" + usages: list[GiftUsageEntry] = [] + + +class MyGiftsResponse(BaseModel): + ok: bool = True + gifts: list[MyGiftItem] = [] + total: int = 0 + limit: int = 20 + offset: int = 0 + + +class GiftRedeemRequest(BaseModel): + gift_code: str = Field(..., min_length=1) + + +class GiftRedeemResponse(BaseModel): + ok: bool = True + message: str = "" + gift_id: str = "" + tariff_id: int = 0 + duration_days: int = 0 + + +class ReferralApplyRequest(BaseModel): + referrer_code: str | None = Field(None, min_length=1) + referrer_tg_id: int | None = Field(None, ge=1) + + +class ReferralApplyResponse(BaseModel): + ok: bool = True + message: str = "" + referrer_code: str = "" + referrer_user_id: int = 0 + referrer_tg_id: int | None = None + referred_user_id: int = 0 + referred_tg_id: int | None = None + + +class ReferralTopEntryResponse(BaseModel): + position: int + referrer_user_id: int + referrals_count: int + display_id: str + + +class ReferralTopResponse(BaseModel): + user_referrals_count: int = 0 + user_position: int | None = None + top: list[ReferralTopEntryResponse] = [] + + +class ReferralQrResponse(BaseModel): + ok: bool = True + link: str = "" + image_data_url: str = "" + + +class ReferralConditionsResponse(BaseModel): + title: str = "" + summary: str = "" + bonus_mode: str = "" + bonus_mode_label: str = "" + level_lines: list[str] = [] + rules: list[str] = [] + + +class PartnerConditionsResponse(BaseModel): + title: str = "" + summary: str = "" + bonus_mode: str = "" + bonus_mode_label: str = "" + level_lines: list[str] = [] + rules: list[str] = [] + examples: list[str] = [] + min_payout_rub: float = 0.0 + payout_methods: list[str] = [] + custom_amount_enabled: bool = False + + +class PartnerQrResponse(BaseModel): + ok: bool = True + link: str = "" + image_data_url: str = "" + + +class PartnerApplyRequest(BaseModel): + partner_code: str | None = Field(None, min_length=1) + partner_tg_id: int | None = Field(None, ge=1) + + +class PartnerApplyResponse(BaseModel): + ok: bool = True + message: str = "" + partner_code: str = "" + partner_user_id: int = 0 + partner_tg_id: int | None = None + joined_user_id: int = 0 + joined_tg_id: int | None = None + + +class PartnerTopEntryResponse(BaseModel): + position: int + partner_user_id: int + referred_count: int + display_id: str + + +class PartnerTopResponse(BaseModel): + user_referred_count: int = 0 + user_position: int | None = None + top: list[PartnerTopEntryResponse] = [] + + +class CouponApplyRequest(BaseModel): + code: str = Field(..., min_length=1, max_length=128) + + +class CouponApplyResponse(BaseModel): + ok: bool = True + message: str = "" + coupon_code: str = "" + amount: int = 0 + balance: float = 0.0 + + +class PartnerPayoutRequestCreate(BaseModel): + amount_rub: float = Field(..., gt=0) + + +class PartnerPayoutRequestResponse(BaseModel): + ok: bool = True + message: str = "" + request_id: int | None = None + amount_rub: float = 0.0 + status: str = "pending" + balance_rub: float = 0.0 + + +class PartnerPayoutEntryResponse(BaseModel): + id: int + amount_rub: float = 0.0 + status: str = "" + created_at: str | None = None + method: str | None = None + destination: str | None = None + + +class PartnerPayoutHistoryResponse(BaseModel): + total: int = 0 + items: list[PartnerPayoutEntryResponse] = [] diff --git a/audit/__init__.py b/audit/__init__.py index 5a88743f..4302d9b5 100644 --- a/audit/__init__.py +++ b/audit/__init__.py @@ -873,6 +873,9 @@ async def list_audit_events( limit=min(5000, need), ) merged = _dedupe_event_like(redis_events + db_events) + for ev in merged: + if ev.created_at is not None and ev.created_at.tzinfo is None: + ev.created_at = ev.created_at.replace(tzinfo=timezone.utc) merged.sort(key=lambda e: (e.created_at, getattr(e, "id", 0)), reverse=True) return merged[offset : offset + limit] diff --git a/audit/rules.py b/audit/rules.py index ed8152c0..5a030e0c 100644 --- a/audit/rules.py +++ b/audit/rules.py @@ -256,6 +256,16 @@ _HANDLER_CONTAINS: list[tuple[str, str] | tuple[str, str, str]] = [ ("auth/send-login", "login"), ("auth/login-by-code", "login"), ("auth/login-telegram", "login"), + ("auth/set-password", "login"), + ("auth/change-password", "login"), + ("auth/request-password-reset", "login"), + ("auth/confirm-password-reset", "login"), + ("auth/summary", "login"), + ("site-config", "api_other"), + ("tariffs/purchase", "pay_start"), + ("tariffs/config-price", "tariff_config"), + ("gifts/redeem", "key_create"), + ("referrals/apply", "referral"), ] @@ -265,6 +275,16 @@ _API_CONTAINS: list[tuple[str, str]] = [ ("auth/send-login", "login"), ("auth/login-by-code", "login"), ("auth/login-telegram", "login"), + ("auth/set-password", "login"), + ("auth/change-password", "login"), + ("auth/request-password-reset", "login"), + ("auth/confirm-password-reset", "login"), + ("auth/summary", "login"), + ("site-config", "api_other"), + ("tariffs/purchase", "pay_start"), + ("tariffs/config-price", "tariff_config"), + ("gifts/redeem", "key_create"), + ("referrals/apply", "referral"), ("/keys/create", "key_create"), ] diff --git a/bot.py b/bot.py index df8c61d5..8341f3bf 100644 --- a/bot.py +++ b/bot.py @@ -1,5 +1,7 @@ from importlib import import_module +version = "0.5.3" + from aiogram import Bot, Dispatcher from aiogram.client.default import DefaultBotProperties from aiogram.enums import ParseMode diff --git a/cli_launcher.py b/cli_launcher.py index 82946044..2794370e 100755 --- a/cli_launcher.py +++ b/cli_launcher.py @@ -456,7 +456,7 @@ def initialize_database() -> bool: [ VENV_PYTHON, "-c", - "import asyncio; from database.init_db import init_db; asyncio.run(init_db())", + "import asyncio; from database.setup.init_db import init_db; asyncio.run(init_db())", ], cwd=PROJECT_DIR, check=True, @@ -1100,6 +1100,558 @@ def update_from_release(): console.print(f"[red]❌ Ошибка при обновлении: {e}[/red]") +WEB_IMAGE = "ghcr.io/vladless/solo-brick:latest" +WEB_CONTAINER_NAME = "solo-brick" +WEB_DIR = os.path.join(os.path.expanduser("~"), "solo-brick") +WEB_REMOTE_ARCHIVE = "https://github.com/Vladless/Solo_bot/archive/refs/heads/dev.tar.gz" +WEB_REMOTE_SUBDIR = "web-app" + + +def _find_local_web_source() -> str | None: + candidates = [ + os.path.join(PROJECT_DIR, "web-app"), + os.path.join(os.path.dirname(PROJECT_DIR), "web-app"), + os.path.join(os.path.expanduser("~"), "Solo_bot", "web-app"), + ] + for path in candidates: + if ( + os.path.isdir(path) + and os.path.isfile(os.path.join(path, "package.json")) + and os.path.isfile(os.path.join(path, "Dockerfile")) + ): + return path + return None + + +def _copy_local_web_source(src: str, dst: str) -> bool: + subprocess.run(["rm", "-rf", dst], check=False) + if shutil.which("rsync"): + result = subprocess.run( + [ + "rsync", "-a", + "--exclude=node_modules", + "--exclude=.next", + "--exclude=.git", + "--exclude=.env", + "--exclude=.env.local", + "--exclude=.env.production", + "--exclude=logs", + "--exclude=.deploy", + "--exclude=.data", + "--exclude=.claude", + f"{src}/", + f"{dst}/", + ], + check=False, + ) + if result.returncode != 0: + return False + else: + try: + shutil.copytree( + src, + dst, + ignore=shutil.ignore_patterns( + "node_modules", ".next", ".git", ".env", ".env.local", + ".env.production", "logs", ".deploy", ".data", ".claude", + ), + ) + except Exception: + return False + return os.path.isfile(os.path.join(dst, "package.json")) + + +def _download_web_from_github(dst: str) -> bool: + import urllib.request + import tarfile + import tempfile + + subprocess.run(["rm", "-rf", dst], check=False) + os.makedirs(dst, exist_ok=True) + + tmp_path = "" + try: + with tempfile.NamedTemporaryFile(suffix=".tar.gz", delete=False) as tmp: + tmp_path = tmp.name + urllib.request.urlretrieve(WEB_REMOTE_ARCHIVE, tmp_path) + + with tarfile.open(tmp_path, "r:gz") as tar: + prefix_marker = f"/{WEB_REMOTE_SUBDIR}/" + extracted = 0 + for member in tar.getmembers(): + idx = member.name.find(prefix_marker) + if idx == -1: + continue + relative = member.name[idx + len(prefix_marker):] + if not relative: + continue + member.name = relative + tar.extract(member, dst) + extracted += 1 + if extracted == 0: + return False + except Exception as e: + console.print(f"[red]❌ Ошибка загрузки архива: {e}[/red]") + return False + finally: + if tmp_path and os.path.exists(tmp_path): + try: + os.unlink(tmp_path) + except Exception: + pass + + return os.path.isfile(os.path.join(dst, "package.json")) + + +def _prepare_web_sources(dst: str) -> bool: + local = _find_local_web_source() + if local: + console.print(f"[cyan]Найден локальный web-app: {local}[/cyan]") + if _copy_local_web_source(local, dst): + console.print("[green]✓ Локальные исходники скопированы[/green]") + return True + console.print("[yellow]Не удалось скопировать локальные исходники. Пробую загрузку из GitHub.[/yellow]") + + console.print("[cyan]Загрузка web-app из публичного репозитория Vladless/Solo_bot (dev)...[/cyan]") + if _download_web_from_github(dst): + console.print("[green]✓ Исходники загружены из публичного репозитория[/green]") + return True + + console.print("[red]❌ Не удалось получить исходники web-app[/red]") + return False + + +def _pull_web_image() -> bool: + console.print(f"[cyan]Загрузка готового образа: {WEB_IMAGE}[/cyan]") + result = subprocess.run( + ["docker", "pull", WEB_IMAGE], + check=False, + ) + return result.returncode == 0 + + +def _build_web_image(src_dir: str) -> bool: + if not os.path.isfile(os.path.join(src_dir, "package.json")): + if not _prepare_web_sources(src_dir): + return False + if not os.path.isfile(os.path.join(src_dir, "Dockerfile")): + console.print("[red]❌ В исходниках нет Dockerfile[/red]") + return False + console.print("[cyan]Сборка Docker-образа (несколько минут)...[/cyan]") + result = subprocess.run( + ["docker", "build", "-t", WEB_IMAGE, "."], + cwd=src_dir, check=False, + ) + if result.returncode != 0: + console.print("[red]❌ Ошибка сборки. Проверьте логи выше.[/red]") + return False + return True + + +def _ensure_web_image(src_dir: str, force_pull: bool = False) -> bool: + if _pull_web_image(): + console.print(f"[green]✓ Образ {WEB_IMAGE} получен из GHCR[/green]") + return True + + console.print("[yellow]Не удалось скачать образ из GHCR. Пробую локальную сборку.[/yellow]") + return _build_web_image(src_dir) + + +def _check_feature(name: str) -> bool: + try: + from core.rpc import check_feature + return check_feature(name) + except Exception: + return False + + +def _verify_license_for_web(code: str, password: str) -> tuple[bool, str]: + try: + from core.rpc import verify_web_license + return verify_web_license(code, password) + except Exception: + return False, "" + + +def _ensure_docker(): + """Проверяет/устанавливает Docker.""" + if shutil.which("docker"): + try: + subprocess.run(["docker", "info"], stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, check=True) + return True + except subprocess.CalledProcessError: + console.print("[yellow]Docker установлен, но не запущен.[/yellow]") + subprocess.run(["sudo", "systemctl", "start", "docker"], check=False) + return True + console.print("[cyan]Установка Docker...[/cyan]") + try: + subprocess.run("curl -fsSL https://get.docker.com | sh", shell=True, check=True) + subprocess.run(["sudo", "systemctl", "enable", "docker"], check=False) + subprocess.run(["sudo", "systemctl", "start", "docker"], check=False) + return True + except subprocess.CalledProcessError: + console.print("[red]❌ Не удалось установить Docker.[/red]") + return False + + +def _ensure_nginx(): + """Проверяет/устанавливает nginx.""" + if shutil.which("nginx"): + return True + console.print("[cyan]Установка nginx...[/cyan]") + try: + subprocess.run(["sudo", "apt-get", "update", "-qq"], check=True, stdout=subprocess.DEVNULL) + subprocess.run(["sudo", "apt-get", "install", "-y", "-qq", "nginx"], check=True, stdout=subprocess.DEVNULL) + subprocess.run(["sudo", "systemctl", "enable", "nginx"], check=False) + subprocess.run(["sudo", "systemctl", "start", "nginx"], check=False) + return True + except subprocess.CalledProcessError: + console.print("[yellow]Не удалось установить nginx автоматически.[/yellow]") + return False + + +def _setup_nginx(domain, web_port=3000): + """Настраивает nginx reverse proxy.""" + conf = f"""server {{ + listen 80; + server_name {domain}; + client_max_body_size 100m; + + location /_next/static/ {{ + proxy_pass http://127.0.0.1:{web_port}; + proxy_cache_valid 200 365d; + add_header Cache-Control "public, immutable, max-age=31536000"; + }} + + location = /sw.js {{ + proxy_pass http://127.0.0.1:{web_port}; + add_header Cache-Control "no-cache"; + }} + + location / {{ + proxy_pass http://127.0.0.1:{web_port}; + proxy_http_version 1.1; + proxy_set_header Upgrade $http_upgrade; + proxy_set_header Connection "upgrade"; + proxy_set_header Host $host; + proxy_set_header X-Real-IP $remote_addr; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + proxy_set_header X-Forwarded-Proto $scheme; + proxy_read_timeout 90s; + }} +}}""" + conf_path = f"/etc/nginx/sites-available/solo-{domain}" + enabled_path = f"/etc/nginx/sites-enabled/solo-{domain}" + try: + with open("/tmp/_solo_nginx.conf", "w") as f: + f.write(conf) + subprocess.run(["sudo", "cp", "/tmp/_solo_nginx.conf", conf_path], check=True) + subprocess.run(["sudo", "ln", "-sf", conf_path, enabled_path], check=True) + subprocess.run(["sudo", "rm", "-f", "/etc/nginx/sites-enabled/default"], check=False) + subprocess.run(["sudo", "nginx", "-t"], check=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) + subprocess.run(["sudo", "systemctl", "reload", "nginx"], check=True) + return True + except subprocess.CalledProcessError: + console.print("[yellow]Не удалось настроить nginx.[/yellow]") + return False + + +def _setup_ssl(domain): + """Получает SSL сертификат через certbot.""" + if not shutil.which("certbot"): + try: + subprocess.run(["sudo", "apt-get", "install", "-y", "-qq", "certbot", "python3-certbot-nginx"], + check=True, stdout=subprocess.DEVNULL) + except subprocess.CalledProcessError: + console.print("[yellow]Не удалось установить certbot.[/yellow]") + return False + try: + subprocess.run([ + "sudo", "certbot", "--nginx", "-d", domain, + "--non-interactive", "--agree-tos", "--register-unsafely-without-email", "--redirect", + ], check=True) + return True + except subprocess.CalledProcessError: + console.print(f"[yellow]Не удалось получить SSL. Убедитесь что {domain} указывает на этот сервер.[/yellow]") + console.print(f"[dim]Повторите вручную: sudo certbot --nginx -d {domain}[/dim]") + return False + + +def install_website(): + """Устанавливает веб-приложение (сайт) через Docker.""" + if not _check_feature("web"): + console.print("[yellow]Эта функция недоступна в текущей версии. Обновите бота.[/yellow]") + return + + console.print( + Panel( + "[white]CLI установит Docker, скачает готовый образ сайта, настроит nginx и SSL.\n" + "Бэкенд (бот) может быть на этом же сервере или на другом.[/white]", + border_style="green", + title="[bold green]Установка веб-приложения[/bold green]", + padding=(1, 2), + ) + ) + + console.print( + Panel( + "[bold cyan]Вариант A:[/bold cyan] Бот и сайт на одном сервере\n" + " → Адрес API: http://localhost:8000 (по умолчанию)\n\n" + "[bold cyan]Вариант B:[/bold cyan] Сайт на отдельном сервере\n" + " → Адрес API: http://IP-бота:8000 (укажите IP сервера с ботом)\n" + " → На сервере бота должен быть открыт порт 8000", + border_style="dim", + title="[dim]Варианты размещения[/dim]", + padding=(1, 2), + ) + ) + + if not safe_confirm("[bold green]Начать установку сайта?[/bold green]", default=True): + return + + console.print("\n[bold][0/5] Авторизация[/bold]") + console.print("[dim]Введите логин и пароль от вашего кабинета на сайте Solo.[/dim]") + console.print("[dim]Данные используются только для проверки лицензии и нигде не сохраняются.[/dim]\n") + + lc_code = safe_prompt("[cyan]Логин (Client Code)[/cyan]") + if not lc_code or not lc_code.strip(): + console.print("[red]Логин обязателен.[/red]") + return + + try: + import getpass + lc_pass = getpass.getpass(" Пароль: ") + except Exception: + lc_pass = safe_prompt("[cyan]Пароль[/cyan]") + + if not lc_pass or not lc_pass.strip(): + console.print("[red]Пароль обязателен.[/red]") + return + + console.print("[dim]Проверка лицензии...[/dim]") + lc_ok, lc_msg = _verify_license_for_web(lc_code.strip(), lc_pass.strip()) + lc_code = None + lc_pass = None + + if not lc_ok: + console.print(f"[red]❌ {lc_msg or 'Авторизация не пройдена'}[/red]") + return + console.print("[green]✓ Авторизация пройдена[/green]") + + console.print("\n[bold][1/5] Docker[/bold]") + if not _ensure_docker(): + return + + console.print("\n[bold][2/5] Настройки[/bold]\n") + + console.print("[dim]Домен, по которому будет открываться сайт.") + console.print("DNS (A-запись) должна уже указывать на IP этого сервера.[/dim]") + domain = safe_prompt("[cyan]Домен сайта[/cyan] (например vpn.example.com)") + if not domain or not domain.strip(): + console.print("[red]Домен обязателен.[/red]") + return + domain = domain.strip() + + console.print("\n[dim]Адрес API вашего бота (FastAPI).") + console.print("Если бот на этом же сервере — оставьте по умолчанию.") + console.print("Если на другом — укажите полный адрес, например http://123.45.67.89:8000[/dim]") + api_url = safe_prompt("[cyan]Адрес backend API[/cyan]", default="http://localhost:8000") + + console.print("\n[dim]Внутренний порт, на котором запустится сайт.") + console.print("Nginx проксирует на него запросы. Менять нужно только если порт занят.[/dim]") + web_port = safe_prompt("[cyan]Порт сайта[/cyan]", default="3000") + + console.print("\n[dim]Для push-уведомлений на сайте (колокольчик).") + console.print("Генерируется командой: npx web-push generate-vapid-keys") + console.print("Если не нужны — пропустите.[/dim]") + vapid_key = safe_prompt("[cyan]VAPID Public Key[/cyan] (Enter — пропустить)", default="") + + console.print("\n[dim]Cloudflare Turnstile защищает формы логина от ботов.") + console.print("Получите ключ на dash.cloudflare.com → Turnstile.") + console.print("Если не нужно — пропустите, формы будут работать без CAPTCHA.[/dim]") + turnstile_key = safe_prompt("[cyan]Turnstile Site Key[/cyan] (Enter — пропустить)", default="") + + console.print("\n[dim]Username Telegram-бота (без @) для кнопки «Войти через Telegram» на сайте.") + console.print("Если не нужно — пропустите.[/dim]") + tg_bot_username = safe_prompt("[cyan]Telegram Bot Username[/cyan] (Enter — пропустить)", default="") + + console.print("\n[dim]Для отправки email-кодов (логин, подтверждение, сброс пароля).") + console.print("Если не нужно — пропустите, регистрация по email+паролю будет работать без этого.[/dim]") + smtp_host = safe_prompt("[cyan]SMTP Host[/cyan] (Enter — пропустить)", default="") + smtp_user = "" + smtp_password = "" + smtp_from = "" + if smtp_host: + smtp_user = safe_prompt("[cyan]SMTP User[/cyan]", default="") + try: + import getpass + smtp_password = getpass.getpass(" SMTP Password: ") + except Exception: + smtp_password = safe_prompt("[cyan]SMTP Password[/cyan]", default="") + smtp_from = safe_prompt("[cyan]Email From[/cyan]", default=smtp_user) + + setup_ssl = safe_confirm("[cyan]Установить SSL (Let's Encrypt)?[/cyan]", default=True) + + site_url = f"https://{domain}" if setup_ssl else f"http://{domain}" + + console.print(f"\n Домен: [green]{domain}[/green]") + console.print(f" Backend: [green]{api_url}[/green]") + console.print(f" SSL: [green]{'Да' if setup_ssl else 'Нет'}[/green]") + + if not safe_confirm("\n[yellow]Всё верно?[/yellow]", default=True): + return + + console.print("\n[bold][3/5] Запуск сайта[/bold]") + os.makedirs(WEB_DIR, exist_ok=True) + + from urllib.parse import urlparse + parsed_api = urlparse(api_url) + api_port_from_url = "" + if parsed_api.port is not None: + api_port_from_url = str(parsed_api.port) + elif parsed_api.scheme == "https": + api_port_from_url = "443" + elif parsed_api.scheme == "http": + api_port_from_url = "80" + + env_path = os.path.join(WEB_DIR, ".env") + with open(env_path, "w") as f: + f.write(f"API_URL={api_url}\n") + f.write(f"API_BASE_URL={api_url}\n") + f.write(f"NEXT_PUBLIC_API_URL={api_url}\n") + f.write(f"NEXT_PUBLIC_API_BASE_URL={api_url}\n") + f.write(f"NEXT_PUBLIC_API_PORT={api_port_from_url}\n") + f.write(f"NEXT_PUBLIC_SITE_URL={site_url}\n") + f.write(f"NEXT_PUBLIC_VAPID_PUBLIC_KEY={vapid_key}\n") + f.write(f"NEXT_PUBLIC_TURNSTILE_SITE_KEY={turnstile_key}\n") + f.write(f"NEXT_PUBLIC_LOG_LEVEL=info\n") + f.write(f"WEB_PORT={web_port}\n") + if tg_bot_username: + f.write(f"NEXT_PUBLIC_TELEGRAM_BOT_USERNAME={tg_bot_username}\n") + if smtp_host: + f.write(f"EMAIL_SMTP_HOST={smtp_host}\n") + f.write(f"EMAIL_SMTP_PORT=465\n") + f.write(f"EMAIL_SMTP_USER={smtp_user}\n") + f.write(f"EMAIL_SMTP_PASSWORD={smtp_password}\n") + f.write(f"EMAIL_FROM={smtp_from}\n") + + src_dir = os.path.join(WEB_DIR, "src") + if not _ensure_web_image(src_dir): + return + + compose_path = os.path.join(WEB_DIR, "docker-compose.yml") + with open(compose_path, "w") as f: + f.write(f"""name: {WEB_CONTAINER_NAME} + +services: + web: + image: {WEB_IMAGE} + container_name: {WEB_CONTAINER_NAME} + ports: + - "127.0.0.1:{web_port}:3000" + env_file: + - .env + restart: unless-stopped + healthcheck: + test: ["CMD", "node", "-e", "fetch('http://127.0.0.1:3000/api/health').then(r=>process.exit(r.ok?0:1)).catch(()=>process.exit(1))"] + interval: 30s + timeout: 5s + retries: 3 + start_period: 10s + volumes: + - ./logs:/app/logs +""") + + console.print("[cyan]Запуск контейнера...[/cyan]") + subprocess.run(["docker", "compose", "up", "-d"], cwd=WEB_DIR, check=True) + console.print(f"[green]✅ Контейнер запущен на порту {web_port}[/green]") + + console.print("\n[bold][4/5] Nginx[/bold]") + if _ensure_nginx(): + _setup_nginx(domain, int(web_port)) + console.print(f"[green]✅ nginx настроен для {domain}[/green]") + + console.print("\n[bold][5/5] SSL[/bold]") + if setup_ssl: + if _setup_ssl(domain): + console.print("[green]✅ SSL сертификат установлен[/green]") + else: + console.print("[dim]SSL пропущен[/dim]") + + smtp_hint = "" + if not smtp_host: + smtp_hint = "\n\n[yellow]⚠ SMTP не настроен — вход по email-коду и сброс пароля не будут работать.\n Настройте позже через: меню → Управление сайтом → Изменить настройки[/yellow]" + + console.print( + Panel( + f"[bold green]Сайт доступен: {site_url}[/bold green]{smtp_hint}\n\n" + f"[white]Управление:[/white]\n" + f" cd {WEB_DIR}\n" + f" docker compose logs -f [dim]— логи[/dim]\n" + f" docker compose restart [dim]— перезапуск[/dim]\n" + f" docker compose down [dim]— остановка[/dim]\n" + f" nano .env [dim]— настройки[/dim]", + border_style="green", + title="[bold green]✅ Установка завершена[/bold green]", + padding=(1, 2), + ) + ) + + +def manage_website(): + """Меню управления сайтом.""" + if not _check_feature("web"): + console.print("[yellow]Эта функция недоступна в текущей версии. Обновите бота.[/yellow]") + return + if not os.path.exists(os.path.join(WEB_DIR, "docker-compose.yml")): + console.print("[yellow]Сайт не установлен.[/yellow]") + if safe_confirm("[green]Установить сейчас?[/green]", default=True): + install_website() + return + + table = Table(title="Управление сайтом", title_style="bold cyan", header_style="bold blue") + table.add_column("№", justify="center", style="cyan", no_wrap=True) + table.add_column("Действие", style="white") + table.add_row("1", "Показать статус") + table.add_row("2", "Показать логи") + table.add_row("3", "Перезапустить") + table.add_row("4", "Остановить") + table.add_row("5", "Обновить (пересборка + restart)") + table.add_row("6", "Изменить настройки (.env)") + table.add_row("7", "Переустановить") + table.add_row("8", "Назад") + console.print(table) + + choice = safe_prompt("[bold blue]👉 Выберите действие[/bold blue]", + choices=[str(i) for i in range(1, 9)], show_choices=False) + + if choice == "1": + subprocess.run(["docker", "compose", "ps"], cwd=WEB_DIR) + elif choice == "2": + subprocess.run(["docker", "compose", "logs", "--tail", "80", "-f"], cwd=WEB_DIR) + elif choice == "3": + subprocess.run(["docker", "compose", "restart"], cwd=WEB_DIR) + console.print("[green]✅ Перезапущено[/green]") + elif choice == "4": + subprocess.run(["docker", "compose", "down"], cwd=WEB_DIR) + console.print("[yellow]Сайт остановлен[/yellow]") + elif choice == "5": + src_dir = os.path.join(WEB_DIR, "src") + console.print("[cyan]Обновление образа...[/cyan]") + if not _ensure_web_image(src_dir, force_pull=True): + return + subprocess.run(["docker", "compose", "up", "-d", "--force-recreate"], cwd=WEB_DIR) + console.print("[green]✅ Обновлено[/green]") + elif choice == "6": + env_path = os.path.join(WEB_DIR, ".env") + editor = os.environ.get("EDITOR", "nano") + subprocess.run([editor, env_path]) + if safe_confirm("[cyan]Перезапустить сайт с новыми настройками?[/cyan]", default=True): + subprocess.run(["docker", "compose", "restart"], cwd=WEB_DIR) + elif choice == "7": + install_website() + + def show_update_menu(): if IS_ROOT_DIR: console.print("[red]Обновление невозможно: бот находится в /root[/red]") @@ -1123,7 +1675,7 @@ def show_update_menu(): def show_menu(): - table = Table(title="Solobot CLI v0.5.0", title_style="bold magenta", header_style="bold blue") + table = Table(title="Solobot CLI v0.5.3", title_style="bold magenta", header_style="bold blue") table.add_column("№", justify="center", style="cyan", no_wrap=True) table.add_column("Операция", style="white") table.add_row("1", "Запустить бота (systemd)") @@ -1135,7 +1687,8 @@ def show_menu(): table.add_row("7", "Обновить Solobot") table.add_row("8", "Восстановить из бэкапа") table.add_row("9", "Установить / переустановить бота") - table.add_row("10", "Выход") + table.add_row("10", "🌐 Веб-сайт (установка / управление)") + table.add_row("11", "Выход") console.print(table) @@ -1150,7 +1703,7 @@ def main(): show_menu() choice = safe_prompt( "[bold blue]👉 Введите номер действия[/bold blue]", - choices=[str(i) for i in range(1, 11)], + choices=[str(i) for i in range(1, 12)], show_choices=False, ) if choice == "1": @@ -1205,6 +1758,8 @@ def main(): elif choice == "9": install_bot() elif choice == "10": + manage_website() + elif choice == "11": console.print("[bold cyan]Выход из CLI. Удачного дня![/bold cyan]") break except KeyboardInterrupt: diff --git a/core/__init__.py b/core/__init__.py index e69de29b..8b137891 100644 --- a/core/__init__.py +++ b/core/__init__.py @@ -0,0 +1 @@ + diff --git a/core/bootstrap.py b/core/bootstrap.py index 014318a4..f4374ce9 100644 --- a/core/bootstrap.py +++ b/core/bootstrap.py @@ -12,6 +12,7 @@ from .settings.payments_config import PAYMENTS_CONFIG, load_payments_config, upd from .settings.providers_order_config import PROVIDERS_ORDER, load_providers_order, update_providers_order from .settings.runtime_sync import publish_runtime_snapshot from .settings.tariffs_config import TARIFFS_CONFIG, load_tariffs_config, update_tariffs_config +from .settings.web_config import WEB_CONFIG, load_web_config, update_web_config async def bootstrap() -> None: @@ -26,6 +27,7 @@ async def bootstrap() -> None: await load_money_config(session) await load_management_config(session) await load_tariffs_config(session) + await load_web_config(session) await session.commit() await settings_cache.load(session) await publish_runtime_snapshot() diff --git a/core/cache_config.py b/core/cache_config.py index ba76e512..2817dc3c 100644 --- a/core/cache_config.py +++ b/core/cache_config.py @@ -1,44 +1,3 @@ -""" -Сводка по Redis: префиксы ключей и окна жизни (TTL). - -Ключ/префикс │ TTL (сек) │ Назначение -──────────────────────────┼───────────┼──────────────────────────────────────── -throttle_counter │ 1 │ Счётчик троттлинга (окно сброса) -throttle_notice │ 1 │ Показ уведомления «подождите» -concurrency_notice │ 5 │ Сообщение «слишком много запросов» -utm_exists │ 300 │ UTM-код уже обработан (старт) -user_middleware │ 60 │ Снапшот пользователя (debounce*2) -user_snapshot │ 30 │ Снапшот пользователя -user_exists │ 60 │ Пользователь есть в БД -balance │ 25 │ Баланс в профиле -profile_data │ 25 │ Данные профиля -key_count │ 25 │ Количество ключей -ban_status │ 60 │ Статус бана -direct_start_user_exists │ 20 │ Пользователь есть (direct start blocker) -admin_access │ 60 │ Доступ в админку (да/нет) -remna_server │ 300 │ URL сервера Remnawave -remna_profile │ 20/45 │ Профиль Remnawave (45 при ошибке) -runtime_configs │ 86400 │ Рантайм-конфиг (1 сутки) -sub_response │ 20 │ Ответ подписки (subscription) -servers │ 60 │ Список серверов -tariff │ 120 │ Тариф по ID -tariffs_cluster │ 120 │ Тарифы по кластеру -keys_list │ 25 │ Список ключей -key_details │ 45 │ Детали ключа -key_email │ 45 │ email по client_id -payment_pending │ 3600 │ Ожидающий платёж (1 ч) -audit_history │ 300 │ История действий клиента (админка) -audit:flush │ — │ Буфер аудита (список для выгрузки в БД в 00:00) -audit:flush:processing │ — │ Батч аудита, перенесённый в processing до commit в БД -audit:flush:drain_lock │ 900 │ Лок nightly/manual drain аудита -audit:user:tg:* │ 25 ч │ События по tg_id для чтения до выгрузки -audit:user:identity:* │ 25 ч │ События по identity для чтения до выгрузки -webhook_abuse_fail │ 60 │ Счётчик неудачных вебхуков по IP -webhook_abuse_block │ 300 │ Блокировка IP по злоупотреблению - -При недоступности Redis: повторная попытка подключения через 5 сек -(REDIS_BACKOFF_SEC в redis_cache). -""" UPDATE_STALE_AGE_SEC = 60 CONCURRENCY_MAX_WAIT_SEC = 300 diff --git a/core/redis_cache.py b/core/redis_cache.py index 0028f175..def0c2a4 100644 --- a/core/redis_cache.py +++ b/core/redis_cache.py @@ -2,12 +2,14 @@ import asyncio import json import os import time + from importlib import import_module from typing import Any from config import REDIS_URL from logger import logger + _REDIS_CLIENTS: dict[tuple[int, int], Any] = {} _REDIS_UNAVAILABLE_UNTIL = 0.0 _REDIS_BACKOFF_SEC = 5.0 @@ -148,6 +150,22 @@ async def cache_incr(key: str, ttl_sec: float) -> int: return 1 +async def cache_incr_checked(key: str, ttl_sec: float) -> tuple[int, bool]: + """Возвращает (value, redis_available). redis_available=False значит клиент + должен применить fallback-логику (например, in-memory limiter). + """ + client = await _get_redis() + if client is None: + return 1, False + try: + value = await client.incr(key) + if value == 1: + await client.expire(key, max(1, int(ttl_sec))) + return int(value), True + except Exception: + return 1, False + + async def cache_delete_pattern(pattern: str) -> int: client = await _get_redis() if client is None: @@ -266,3 +284,7 @@ async def cache_lmove_batch(source: str, destination: str, count: int) -> list[A except Exception as exc: logger.warning(f"[Redis] lmove_batch({source}->{destination}) не удался: {exc}") return [] + + +async def redis_connection_ok() -> bool: + return await _get_redis() is not None diff --git a/core/settings/web_config.py b/core/settings/web_config.py new file mode 100644 index 00000000..9b287aae --- /dev/null +++ b/core/settings/web_config.py @@ -0,0 +1,78 @@ +from typing import Any + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from database.models import Setting + +from database.settings_cache import settings_cache +from ..defaults import DEFAULT_WEB_CONFIG +from .runtime_sync import publish_runtime_config, register_runtime_config + + +WEB_CONFIG: dict[str, Any] = DEFAULT_WEB_CONFIG.copy() +WEB_SETTING_KEY = "WEB_CONFIG" +register_runtime_config(WEB_SETTING_KEY, WEB_CONFIG) + + +async def load_web_config(session: AsyncSession) -> None: + stmt = select(Setting).where(Setting.key == WEB_SETTING_KEY) + result = await session.execute(stmt) + setting = result.scalar_one_or_none() + + if setting is None: + web_config = DEFAULT_WEB_CONFIG.copy() + setting = Setting( + key=WEB_SETTING_KEY, + value=web_config, + description="Конфигурация веб-сайта", + ) + session.add(setting) + else: + stored = setting.value or {} + web_config = DEFAULT_WEB_CONFIG.copy() + web_config.update(stored) + setting.value = web_config + + WEB_CONFIG.clear() + WEB_CONFIG.update(web_config) + await session.flush() + + +async def update_web_config(session: AsyncSession, new_values: dict[str, Any]) -> None: + stmt = select(Setting).where(Setting.key == WEB_SETTING_KEY) + result = await session.execute(stmt) + setting = result.scalar_one_or_none() + + if setting is None: + setting = Setting( + key=WEB_SETTING_KEY, + value=new_values, + description="Конфигурация веб-сайта", + ) + session.add(setting) + else: + setting.value = new_values + + await session.commit() + + web_config = DEFAULT_WEB_CONFIG.copy() + web_config.update(new_values) + + WEB_CONFIG.clear() + WEB_CONFIG.update(web_config) + settings_cache.update(WEB_SETTING_KEY, web_config) + await publish_runtime_config(WEB_SETTING_KEY, web_config) + + +def get_site_url() -> str: + """Возвращает SITE_URL из WEB_CONFIG, или из config.py как fallback.""" + url = str(WEB_CONFIG.get("SITE_URL") or "").strip() + if url: + return url.rstrip("/") + from config import SITE_URL + return SITE_URL.rstrip("/") if SITE_URL else "" + + +def is_web_enabled() -> bool: + return bool(WEB_CONFIG.get("WEB_ENABLED", False)) diff --git a/core/tasks/cron_tasks.py b/core/tasks/cron_tasks.py index cf9da1b5..2eab6850 100644 --- a/core/tasks/cron_tasks.py +++ b/core/tasks/cron_tasks.py @@ -26,6 +26,7 @@ async def scheduled_stats_report() -> None: async def sweep_stale_payments_job() -> None: async with async_session_maker() as session: await cancel_expired_pending_payments(session) + await session.commit() def scheduled_audit_drain_process_runner() -> None: @@ -40,6 +41,33 @@ def sweep_stale_payments_process_runner() -> None: asyncio.run(sweep_stale_payments_job()) +async def cleanup_expired_gifts_job() -> None: + from datetime import datetime + + from sqlalchemy import update as sa_update + + from database.models import Gift + + async with async_session_maker() as session: + try: + result = await session.execute( + sa_update(Gift) + .where(Gift.expiry_time < datetime.utcnow(), Gift.is_used == False) + .values(is_used=True) + ) + count = result.rowcount + await session.commit() + if count: + logger.info("[GiftCleanup] Просроченных подарков помечено использованными: {}", count) + except Exception as error: + logger.error("[GiftCleanup] Ошибка очистки подарков: {}", error) + + +def cleanup_expired_gifts_process_runner() -> None: + asyncio.run(cleanup_expired_gifts_job()) + + AUDIT_DRAIN_TRIGGER = CronTrigger(hour=0, minute=0, timezone="Europe/Moscow") DAILY_STATS_REPORT_TRIGGER = CronTrigger(hour=0, minute=1, timezone="Europe/Moscow") STALE_PAYMENTS_SWEEP_TRIGGER = CronTrigger(minute=0, timezone="Europe/Moscow") +EXPIRED_GIFTS_CLEANUP_TRIGGER = CronTrigger(hour=3, minute=0, timezone="Europe/Moscow") diff --git a/core/tasks/periodic_manager.py b/core/tasks/periodic_manager.py index 6e2d19ab..57a2b470 100644 --- a/core/tasks/periodic_manager.py +++ b/core/tasks/periodic_manager.py @@ -3,6 +3,7 @@ import fcntl import inspect import multiprocessing import os +import tempfile import threading from collections.abc import Awaitable, Callable @@ -109,6 +110,19 @@ class PeriodicTaskManager: self._process_lock_file = None self._process_lock_path = "/tmp/solo_bot_periodic_manager.lock" + def _process_lock_candidates(self) -> list[str]: + candidates = [self._process_lock_path] + uid_suffix = f"solo_bot_periodic_manager_{os.getuid()}.lock" + runtime_dir = os.environ.get("XDG_RUNTIME_DIR", "").strip() + if runtime_dir: + candidates.append(os.path.join(runtime_dir, uid_suffix)) + candidates.append(os.path.join(tempfile.gettempdir(), uid_suffix)) + unique_candidates: list[str] = [] + for candidate in candidates: + if candidate not in unique_candidates: + unique_candidates.append(candidate) + return unique_candidates + def register_loop_task(self, task_id: str, runner: LoopRunner) -> None: self._loop_tasks[task_id] = ManagedLoopTask(task_id=task_id, runner=runner) @@ -144,18 +158,26 @@ class PeriodicTaskManager: def _acquire_process_lock(self) -> bool: if self._process_lock_file is not None: return True - lock_file = open(self._process_lock_path, "a+", encoding="utf-8") - try: - fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB) - lock_file.seek(0) - lock_file.truncate() - lock_file.write(str(os.getpid())) - lock_file.flush() - self._process_lock_file = lock_file - return True - except OSError: - lock_file.close() - return False + for candidate_path in self._process_lock_candidates(): + try: + lock_file = open(candidate_path, "a+", encoding="utf-8") + except OSError as error: + logger.warning("[PeriodicManager] Не удалось открыть lock-файл {}: {}", candidate_path, error) + continue + try: + fcntl.flock(lock_file.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB) + lock_file.seek(0) + lock_file.truncate() + lock_file.write(str(os.getpid())) + lock_file.flush() + self._process_lock_file = lock_file + self._process_lock_path = candidate_path + return True + except OSError: + lock_file.close() + return False + logger.warning("[PeriodicManager] Не удалось создать lock-файл, запуск менеджера пропущен") + return False def _release_process_lock(self) -> None: if self._process_lock_file is None: diff --git a/core/tasks/registry.py b/core/tasks/registry.py index f2f3b710..5dc92615 100644 --- a/core/tasks/registry.py +++ b/core/tasks/registry.py @@ -1,7 +1,10 @@ from core.tasks.cron_tasks import ( AUDIT_DRAIN_TRIGGER, DAILY_STATS_REPORT_TRIGGER, + EXPIRED_GIFTS_CLEANUP_TRIGGER, STALE_PAYMENTS_SWEEP_TRIGGER, + cleanup_expired_gifts_job, + cleanup_expired_gifts_process_runner, scheduled_audit_drain, scheduled_audit_drain_process_runner, scheduled_stats_report, @@ -98,4 +101,18 @@ def register_periodic_tasks() -> None: STALE_PAYMENTS_SWEEP_TRIGGER, ) + if process_budget > 0: + periodic_task_manager.register_cron_task( + "cleanup_expired_gifts", + cleanup_expired_gifts_process_runner, + EXPIRED_GIFTS_CLEANUP_TRIGGER, + execution_mode="process", + ) + else: + periodic_task_manager.register_cron_task( + "cleanup_expired_gifts", + cleanup_expired_gifts_job, + EXPIRED_GIFTS_CLEANUP_TRIGGER, + ) + _TASKS_REGISTERED = True diff --git a/core/webhook_abuse.py b/core/webhook_abuse.py index dd3d862f..8b76b232 100644 --- a/core/webhook_abuse.py +++ b/core/webhook_abuse.py @@ -1,5 +1,7 @@ from aiohttp import web +from logger import logger + from core.cache_config import ( WEBHOOK_ABUSE_BLOCK_TTL_SEC, WEBHOOK_ABUSE_FAIL_THRESHOLD, @@ -48,5 +50,5 @@ async def record_webhook_signature_failure(ip: str) -> None: block_key = cache_key("webhook_abuse_block", ip) await cache_set(block_key, 1, WEBHOOK_ABUSE_BLOCK_TTL_SEC) await cache_delete(fail_key) - except Exception: - pass + except Exception as e: + logger.warning("[WebhookAbuse] Ошибка записи fail-счётчика для IP={}: {}", ip, e) diff --git a/database/__init__.py b/database/__init__.py index d9ec33ca..7edf4400 100644 --- a/database/__init__.py +++ b/database/__init__.py @@ -5,7 +5,7 @@ from .db import Base, async_session_maker, engine, reset_async_db_engine from .gifts import * from . import identities from .hot_leads import * -from .init_db import * +from .setup.init_db import * from .keys import * from .notifications import * from .payments import * diff --git a/database/access/__init__.py b/database/access/__init__.py new file mode 100644 index 00000000..48ac878b --- /dev/null +++ b/database/access/__init__.py @@ -0,0 +1,2 @@ +from .resolution import * +from .tg_mirror import * diff --git a/database/access/resolution.py b/database/access/resolution.py new file mode 100644 index 00000000..e8eadaf8 --- /dev/null +++ b/database/access/resolution.py @@ -0,0 +1,87 @@ +from dataclasses import dataclass +from enum import Enum + +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from database.models import Identity, User + + +class ActorSurface(str, Enum): + TELEGRAM = "telegram" + WEB = "web" + UNKNOWN = "unknown" + + +@dataclass(frozen=True) +class ResolvedActor: + surface: ActorSurface + billing_user_id: int | None + telegram_chat_id: int | None + identity_id: str | None + + +def telegram_chat_id(user: User | None) -> int | None: + if user is None: + return None + return user.tg_id + + +async def resolve_user_optional(session: AsyncSession, legacy_id: int) -> User | None: + r = await session.execute(select(User).where(User.tg_id == legacy_id)) + u = r.scalar_one_or_none() + if u is not None: + return u + r2 = await session.execute(select(User).where(User.id == legacy_id)) + return r2.scalar_one_or_none() + + +async def notify_telegram_chat_id(session: AsyncSession, legacy_ref: int) -> int | None: + payer = await resolve_user_optional(session, legacy_ref) + tg = telegram_chat_id(payer) + if tg is not None: + return tg + if payer is None: + return legacy_ref + return None + + +async def resolve_actor_from_legacy_ref(session: AsyncSession, legacy_ref: int) -> ResolvedActor: + user = await resolve_user_optional(session, legacy_ref) + if user is None: + return ResolvedActor( + surface=ActorSurface.UNKNOWN, + billing_user_id=None, + telegram_chat_id=legacy_ref, + identity_id=None, + ) + + user_tg = telegram_chat_id(user) + if user_tg is not None and int(user_tg) == int(legacy_ref): + surface = ActorSurface.TELEGRAM + elif int(user.id) == int(legacy_ref): + surface = ActorSurface.WEB + elif user_tg is None: + surface = ActorSurface.WEB + else: + surface = ActorSurface.UNKNOWN + + return ResolvedActor( + surface=surface, + billing_user_id=int(user.id), + telegram_chat_id=user_tg, + identity_id=user.identity_id, + ) + + +async def resolve_actor_from_identity(session: AsyncSession, identity: Identity) -> ResolvedActor: + from database.identities import ensure_billing_user_for_identity + + billing_uid = await ensure_billing_user_for_identity(session, identity) + user = await resolve_user_optional(session, billing_uid) + return ResolvedActor( + surface=ActorSurface.WEB, + billing_user_id=billing_uid, + telegram_chat_id=telegram_chat_id(user), + identity_id=identity.id, + ) diff --git a/database/access/tg_mirror.py b/database/access/tg_mirror.py new file mode 100644 index 00000000..6433842f --- /dev/null +++ b/database/access/tg_mirror.py @@ -0,0 +1,50 @@ +from __future__ import annotations + +from sqlalchemy import select, update +from sqlalchemy.ext.asyncio import AsyncSession + +from database.models import ( + BlockedUser, + CouponUsage, + Gift, + GiftUsage, + Key, + ManualBan, + Notification, + Payment, + Referral, + TemporaryData, + User, +) + + +def mirror_telegram_id(user: User | None) -> int | None: + if user is None: + return None + return user.tg_id + + +async def refresh_tg_mirrors_for_user(session: AsyncSession, user_id: int) -> None: + r = await session.execute(select(User.tg_id).where(User.id == user_id)) + tg = r.scalar_one_or_none() + + await session.execute(update(Key).where(Key.user_id == user_id).values(tg_id=tg)) + await session.execute(update(Payment).where(Payment.user_id == user_id).values(tg_id=tg)) + await session.execute(update(Notification).where(Notification.user_id == user_id).values(tg_id=tg)) + await session.execute(update(GiftUsage).where(GiftUsage.user_id == user_id).values(tg_id=tg)) + await session.execute(update(CouponUsage).where(CouponUsage.user_id == user_id).values(tg_id=tg)) + await session.execute(update(TemporaryData).where(TemporaryData.user_id == user_id).values(tg_id=tg)) + await session.execute(update(BlockedUser).where(BlockedUser.user_id == user_id).values(tg_id=tg)) + await session.execute(update(ManualBan).where(ManualBan.user_id == user_id).values(tg_id=tg)) + + await session.execute( + update(Referral).where(Referral.referred_user_id == user_id).values(referred_tg_id=tg) + ) + await session.execute( + update(Referral).where(Referral.referrer_user_id == user_id).values(referrer_tg_id=tg) + ) + + await session.execute(update(Gift).where(Gift.sender_user_id == user_id).values(sender_tg_id=tg)) + await session.execute( + update(Gift).where(Gift.recipient_user_id == user_id).values(recipient_tg_id=tg) + ) diff --git a/database/audit.py b/database/audit.py index ef30be54..016fbf13 100644 --- a/database/audit.py +++ b/database/audit.py @@ -77,7 +77,7 @@ async def fetch_successful_payment_rows_db( Payment.created_at, ) stmt = ( - select(Payment.payment_system, Payment.payment_id, Payment.tg_id) + select(Payment.payment_system, Payment.payment_id, Payment.user_id) .where( Payment.status == "success", Payment.payment_system.notin_(PAYMENT_SYSTEMS_EXCLUDED), diff --git a/database/bans.py b/database/bans.py index d16e96f3..b4a9b7f6 100644 --- a/database/bans.py +++ b/database/bans.py @@ -1,27 +1,43 @@ +from sqlalchemy import select from sqlalchemy.dialects.postgresql import insert from sqlalchemy.ext.asyncio import AsyncSession -from database.models import BlockedUser +from database.models import BlockedUser, User +from database.access.resolution import resolve_user_optional from logger import logger -async def create_blocked_user(session: AsyncSession, tg_id: int): - stmt = insert(BlockedUser).values(tg_id=tg_id).on_conflict_do_nothing(index_elements=[BlockedUser.tg_id]) +async def create_blocked_user(session: AsyncSession, legacy_user_ref: int): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return + stmt = ( + insert(BlockedUser) + .values(user_id=u.id, tg_id=u.tg_id) + .on_conflict_do_nothing(index_elements=[BlockedUser.user_id]) + ) await session.execute(stmt) - await session.commit() async def save_blocked_user_ids(session: AsyncSession, tg_ids: list[int]) -> None: - """Вставка списка tg_id в таблицу BlockedUser батчами по 500. Вызывать только из основного event loop.""" + """Вставка списка telegram id в таблицу blocked_users батчами по 500.""" if not tg_ids: return batch_size = 500 total = 0 for i in range(0, len(tg_ids), batch_size): batch = tg_ids[i : i + batch_size] - values = [{"tg_id": tg_id} for tg_id in batch] - stmt = insert(BlockedUser).values(values).on_conflict_do_nothing(index_elements=[BlockedUser.tg_id]) + res = await session.execute(select(User.id, User.tg_id).where(User.tg_id.in_(batch))) + rows = res.all() + uid_by_tg = {int(tgid): int(uid) for uid, tgid in rows if tgid is not None} + values = [ + {"user_id": uid_by_tg[int(tg)], "tg_id": int(tg)} + for tg in batch + if int(tg) in uid_by_tg + ] + if not values: + continue + stmt = insert(BlockedUser).values(values).on_conflict_do_nothing(index_elements=[BlockedUser.user_id]) await session.execute(stmt) - await session.commit() - total += len(batch) - logger.info(f"📝 Добавлено {total} пользователей в blocked_users") + total += len(values) + logger.info(f"📝 Добавлено до {total} пользователей в blocked_users") diff --git a/database/coupons.py b/database/coupons.py index 322de6e5..68061f60 100644 --- a/database/coupons.py +++ b/database/coupons.py @@ -1,9 +1,9 @@ from datetime import datetime -from sqlalchemy import case, delete, func, insert, select, update -from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy import case, delete, func, insert, or_, select, update from sqlalchemy.ext.asyncio import AsyncSession +from database.access.resolution import resolve_user_optional from database.models import Coupon, CouponUsage from logger import logger @@ -19,48 +19,42 @@ async def create_coupon( max_discount_amount: int | None = None, min_order_amount: int | None = None, ) -> bool: - try: - exists = await session.scalar(select(Coupon.id).where(Coupon.code == code)) - if exists: - logger.warning(f"[Coupon] ⚠️ Купон с кодом {code} уже существует.") + exists = await session.scalar(select(Coupon.id).where(Coupon.code == code)) + if exists: + logger.warning(f"[Coupon] ⚠️ Купон с кодом {code} уже существует.") + return False + + if percent is not None: + try: + percent_value = int(percent) + except (TypeError, ValueError): + logger.warning(f"[Coupon] ⚠️ Некорректный процент для купона {code}.") return False - if percent is not None: - try: - percent_value = int(percent) - except (TypeError, ValueError): - logger.warning(f"[Coupon] ⚠️ Некорректный процент для купона {code}.") - return False + if percent_value <= 0 or percent_value > 100: + logger.warning(f"[Coupon] ⚠️ процент должен быть в диапазоне 1..100 для купона {code}.") + return False - if percent_value <= 0 or percent_value > 100: - logger.warning(f"[Coupon] ⚠️ процент должен быть в диапазоне 1..100 для купона {code}.") - return False + if (amount or 0) > 0 or (days or 0) > 0: + logger.warning(f"[Coupon] ⚠️ Купон {code} не может одновременно иметь percent и amount/days.") + return False - if (amount or 0) > 0 or (days or 0) > 0: - logger.warning(f"[Coupon] ⚠️ Купон {code} не может одновременно иметь percent и amount/days.") - return False - - await session.execute( - insert(Coupon).values( - code=code, - amount=int(amount) if amount is not None else 0, - usage_limit=usage_limit, - usage_count=0, - is_used=False, - days=days, - new_users_only=new_users_only, - percent=percent, - max_discount_amount=max_discount_amount, - min_order_amount=min_order_amount, - ) + await session.execute( + insert(Coupon).values( + code=code, + amount=int(amount) if amount is not None else 0, + usage_limit=usage_limit, + usage_count=0, + is_used=False, + days=days, + new_users_only=new_users_only, + percent=percent, + max_discount_amount=max_discount_amount, + min_order_amount=min_order_amount, ) - await session.commit() - logger.info(f"[Coupon] ✅ Купон {code} успешно создан.") - return True - except SQLAlchemyError as e: - await session.rollback() - logger.error(f"[Coupon] ❌ Ошибка при создании купона {code}: {e}") - return False + ) + logger.info(f"[Coupon] ✅ Купон {code} успешно создан.") + return True async def get_coupon_by_code(session: AsyncSession, code: str) -> Coupon | None: @@ -69,6 +63,15 @@ async def get_coupon_by_code(session: AsyncSession, code: str) -> Coupon | None: return result.scalar_one_or_none() +async def get_coupon_by_code_ci(session: AsyncSession, code: str) -> Coupon | None: + normalized = str(code or "").strip() + if not normalized: + return None + stmt = select(Coupon).where(func.lower(Coupon.code) == normalized.lower()) + result = await session.execute(stmt) + return result.scalar_one_or_none() + + async def get_all_coupons(session: AsyncSession, page: int = 1, per_page: int = 10) -> dict: offset = (page - 1) * per_page @@ -99,45 +102,89 @@ async def delete_coupon(session: AsyncSession, code: str) -> bool: await session.execute(delete(CouponUsage).where(CouponUsage.coupon_id == coupon.id)) await session.delete(coupon) - await session.commit() logger.info(f"🗑 Купон {code} удалён вместе с его использованиями") return True +async def _coupon_usage_billing_match(session: AsyncSession, legacy_user_ref: int): + u = await resolve_user_optional(session, legacy_user_ref) + if u is not None: + opts = [CouponUsage.user_id == u.id] + if u.tg_id is not None: + opts.append(CouponUsage.tg_id == u.tg_id) + return or_(*opts) + return or_(CouponUsage.user_id == legacy_user_ref, CouponUsage.tg_id == legacy_user_ref) + + async def create_coupon_usage(session: AsyncSession, coupon_id: int, user_id: int): - try: - stmt = insert(CouponUsage).values(coupon_id=coupon_id, user_id=user_id, used_at=datetime.utcnow()) - await session.execute(stmt) - await session.commit() - logger.info(f"✅ Купон {coupon_id} использован пользователем {user_id}") - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при сохранении использования купона: {e}") - await session.rollback() - raise + u = await resolve_user_optional(session, user_id) + uid = u.id if u is not None else user_id + stmt = insert(CouponUsage).values( + coupon_id=coupon_id, + user_id=uid, + tg_id=u.tg_id if u is not None else None, + used_at=datetime.utcnow(), + ) + await session.execute(stmt) + logger.info(f"✅ Купон {coupon_id} использован пользователем {user_id}") -async def check_coupon_usage(session: AsyncSession, coupon_id: int, user_id: int) -> bool: - stmt = select(CouponUsage).where(CouponUsage.coupon_id == coupon_id, CouponUsage.user_id == user_id) +async def check_coupon_usage(session: AsyncSession, coupon_id: int, legacy_user_ref: int) -> bool: + m = await _coupon_usage_billing_match(session, legacy_user_ref) + stmt = select(CouponUsage).where(CouponUsage.coupon_id == coupon_id).where(m) result = await session.execute(stmt) return result.scalar_one_or_none() is not None +async def has_any_coupon_usage(session: AsyncSession, legacy_user_ref: int) -> bool: + m = await _coupon_usage_billing_match(session, legacy_user_ref) + stmt = select(CouponUsage.coupon_id).where(m).limit(1) + result = await session.execute(stmt) + return result.first() is not None + + async def update_coupon_usage_count(session: AsyncSession, coupon_id: int): - try: - await session.execute( - update(Coupon) - .where(Coupon.id == coupon_id) - .values( - usage_count=Coupon.usage_count + 1, - is_used=case((Coupon.usage_count + 1 >= Coupon.usage_limit, True), else_=False), - ) + await session.execute( + update(Coupon) + .where(Coupon.id == coupon_id) + .values( + usage_count=Coupon.usage_count + 1, + is_used=case((Coupon.usage_count + 1 >= Coupon.usage_limit, True), else_=False), ) - await session.commit() - logger.info(f"🔁 Обновлён счётчик купона {coupon_id}") - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при обновлении купона {coupon_id}: {e}") - await session.rollback() - raise + ) + logger.info(f"🔁 Обновлён счётчик купона {coupon_id}") + + +async def mark_coupon_used(session: AsyncSession, coupon_id: int, legacy_user_ref: int): + u = await resolve_user_optional(session, legacy_user_ref) + uid = u.id if u is not None else legacy_user_ref + match = [CouponUsage.user_id == int(uid)] + if u is not None and u.tg_id is not None: + match.append(CouponUsage.tg_id == int(u.tg_id)) + existing = await session.execute( + select(CouponUsage).where( + CouponUsage.coupon_id == int(coupon_id), + or_(*match), + ) + ) + if existing.scalar_one_or_none() is not None: + return + await session.execute( + insert(CouponUsage).values( + coupon_id=coupon_id, + user_id=uid, + tg_id=u.tg_id if u is not None else None, + used_at=datetime.utcnow(), + ) + ) + await session.execute( + update(Coupon) + .where(Coupon.id == coupon_id) + .values( + usage_count=Coupon.usage_count + 1, + is_used=case((Coupon.usage_count + 1 >= Coupon.usage_limit, True), else_=False), + ) + ) def apply_percent_coupon(price_rub: int, coupon: Coupon) -> tuple[int, int]: diff --git a/database/db.py b/database/db.py index d1264343..23d08292 100644 --- a/database/db.py +++ b/database/db.py @@ -20,6 +20,12 @@ if USE_PGBOUNCER and "+asyncpg" in DATABASE_URL: _pool_recycle = 60 if USE_PGBOUNCER else 300 +_QUERY_TIMEOUT_SEC = 30 + +if "+asyncpg" in _db_url: + _connect_args.setdefault("command_timeout", _QUERY_TIMEOUT_SEC) + _connect_args.setdefault("timeout", _QUERY_TIMEOUT_SEC) + def _create_engine(): return create_async_engine( diff --git a/database/gifts.py b/database/gifts.py index e7754b86..d0aa1fa3 100644 --- a/database/gifts.py +++ b/database/gifts.py @@ -1,17 +1,17 @@ from datetime import datetime -from sqlalchemy import insert -from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy import func, insert, select, update from sqlalchemy.ext.asyncio import AsyncSession -from database.models import Gift +from database.access.resolution import resolve_user_optional +from database.models import Gift, GiftUsage from logger import logger async def store_gift_link( session: AsyncSession, gift_id: str, - sender_tg_id: int, + sender_legacy_ref: int, selected_months: int, expiry_time: datetime, gift_link: str, @@ -21,33 +21,97 @@ async def store_gift_link( selected_device_limit: int | None = None, selected_traffic_gb: int | None = None, selected_price_rub: int | None = None, -): - try: - stmt = insert(Gift).values( +) -> bool: + u = await resolve_user_optional(session, sender_legacy_ref) + if u is None: + raise ValueError(f"sender not found for gift: {sender_legacy_ref}") + stmt = insert(Gift).values( + gift_id=gift_id, + sender_user_id=u.id, + sender_tg_id=u.tg_id, + recipient_user_id=None, + selected_months=selected_months, + expiry_time=expiry_time, + gift_link=gift_link, + created_at=datetime.utcnow(), + is_used=False, + tariff_id=tariff_id, + is_unlimited=is_unlimited, + max_usages=max_usages, + selected_device_limit=selected_device_limit, + selected_traffic_gb=selected_traffic_gb, + selected_price_rub=selected_price_rub, + ) + await session.execute(stmt) + logger.info( + f"🎁 Подарок {gift_id} сохранён " + f"(tariff_id={tariff_id}, max_usages={max_usages}, " + f"device={selected_device_limit}, traffic={selected_traffic_gb}, price={selected_price_rub})" + ) + return True + + +async def get_gift_locked(session: AsyncSession, gift_id: str) -> Gift | None: + """SELECT FOR UPDATE по gift_id — берёт row-lock для atomic redemption. + + Используется в `services.gifts.redeem_gift` чтобы два параллельных запроса + на активацию одного и того же подарка не смогли обойти проверку `is_used`. + """ + result = await session.execute(select(Gift).where(Gift.gift_id == gift_id).with_for_update()) + return result.scalar_one_or_none() + + +async def get_gift_usage(session: AsyncSession, gift_id: str, user_id: int) -> GiftUsage | None: + """Возвращает запись об использовании подарка конкретным пользователем, если есть.""" + result = await session.execute( + select(GiftUsage).where( + GiftUsage.gift_id == gift_id, + GiftUsage.user_id == user_id, + ) + ) + return result.scalar_one_or_none() + + +async def count_gift_usages(session: AsyncSession, gift_id: str) -> int: + """Сколько раз подарок был активирован (для `is_unlimited=False` с лимитом).""" + result = await session.execute( + select(func.count()).select_from(GiftUsage).where(GiftUsage.gift_id == gift_id) + ) + return int(result.scalar_one() or 0) + + +async def record_gift_usage( + session: AsyncSession, + gift_id: str, + user_id: int, + tg_id: int | None, +) -> None: + """Вставляет запись о применении подарка. Композитный ключ (gift_id, user_id).""" + await session.execute( + insert(GiftUsage).values( gift_id=gift_id, - sender_tg_id=sender_tg_id, - recipient_tg_id=None, - selected_months=selected_months, - expiry_time=expiry_time, - gift_link=gift_link, - created_at=datetime.utcnow(), - is_used=False, - tariff_id=tariff_id, - is_unlimited=is_unlimited, - max_usages=max_usages, - selected_device_limit=selected_device_limit, - selected_traffic_gb=selected_traffic_gb, - selected_price_rub=selected_price_rub, + user_id=user_id, + tg_id=tg_id, ) - await session.execute(stmt) - await session.commit() - logger.info( - f"🎁 Подарок {gift_id} сохранён " - f"(tariff_id={tariff_id}, max_usages={max_usages}, " - f"device={selected_device_limit}, traffic={selected_traffic_gb}, price={selected_price_rub})" + ) + + +async def mark_gift_fully_redeemed( + session: AsyncSession, + gift_id: str, + recipient_user_id: int, + recipient_tg_id: int | None, +) -> None: + """Помечает подарок как полностью использованный (is_used=True) и фиксирует получателя. + + Вызывается для non-unlimited подарков, когда набрали max_usages. + """ + await session.execute( + update(Gift) + .where(Gift.gift_id == gift_id) + .values( + is_used=True, + recipient_user_id=recipient_user_id, + recipient_tg_id=recipient_tg_id, ) - return True - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при сохранении подарка {gift_id}: {e}") - await session.rollback() - raise + ) diff --git a/database/hot_leads.py b/database/hot_leads.py index 32b19ff2..a728d850 100644 --- a/database/hot_leads.py +++ b/database/hot_leads.py @@ -8,17 +8,17 @@ from database.models import Key, Payment, User async def get_hot_leads(session: AsyncSession): now_ms = func.extract("epoch", func.now()) * 1000 - sub_active = select(Key.tg_id).where(Key.expiry_time > now_ms).distinct() + sub_active = select(Key.user_id).where(Key.expiry_time > now_ms).distinct() stmt = ( - select(Payment.tg_id) - .join(User, User.tg_id == Payment.tg_id) + select(Payment.user_id) + .join(User, User.id == Payment.user_id) .distinct() .where(User.trial == 1) .where(Payment.amount > 0) .where(Payment.status == "success") .where(Payment.payment_system.notin_(PAYMENT_SYSTEMS_EXCLUDED)) - .where(~Payment.tg_id.in_(sub_active)) + .where(~Payment.user_id.in_(sub_active)) ) result = await session.execute(stmt) diff --git a/database/identities.py b/database/identities.py index dbec7a97..8e0aea4b 100644 --- a/database/identities.py +++ b/database/identities.py @@ -3,11 +3,12 @@ import secrets from datetime import datetime, timedelta import bcrypt -from sqlalchemy import select +from sqlalchemy import delete, func, select, text, update from sqlalchemy.ext.asyncio import AsyncSession from config import API_TOKEN_TTL_DAYS from core.executor import run_cpu, run_io +from database.access.tg_mirror import refresh_tg_mirrors_for_user from database.models import Admin, Identity, User @@ -57,7 +58,6 @@ async def create_identity( await session.flush() if tg_id: await session.execute(User.__table__.update().where(User.tg_id == tg_id).values(identity_id=identity.id)) - await session.commit() await session.refresh(identity) return identity @@ -93,7 +93,6 @@ async def issue_token_for_identity(session: AsyncSession, identity: Identity) -> token = generate_token() identity.api_token_hash = await run_io(hash_token, token) identity.token_issued_at = datetime.utcnow() - await session.commit() await session.refresh(identity) return token @@ -116,7 +115,6 @@ async def create_identity_with_token( identity = await create_identity(session, email=email, tg_id=tg_id) if password: identity.password_hash = await run_cpu(hash_password, password) - await session.commit() await session.refresh(identity) token = await issue_token_for_identity(session, identity) return identity, token @@ -146,10 +144,209 @@ async def login_by_email(session: AsyncSession, email: str, password: str) -> tu return identity, token -async def resolve_tg_id(session: AsyncSession, identity_id: str) -> int | None: - """По identity_id возвращает tg_id, если привязан.""" +async def set_initial_password( + session: AsyncSession, + identity_id: str, + password: str, +) -> Identity | None: identity = await get_identity_by_id(session, identity_id) - return identity.tg_id if identity else None + if not identity or identity.password_hash: + return None + identity.password_hash = await run_cpu(hash_password, password) + await session.refresh(identity) + return identity + + +async def set_password_for_identity( + session: AsyncSession, + identity_id: str, + new_password: str, +) -> Identity | None: + identity = await get_identity_by_id(session, identity_id) + if not identity: + return None + identity.password_hash = await run_cpu(hash_password, new_password) + await session.refresh(identity) + return identity + + +async def change_identity_password( + session: AsyncSession, + identity_id: str, + current_password: str, + new_password: str, +) -> str | None: + """Возвращает None при успехе, иначе код: no_password | wrong_password.""" + identity = await get_identity_by_id(session, identity_id) + if not identity: + return "wrong_password" + if not identity.password_hash: + return "no_password" + if not await run_cpu(check_password, current_password, identity.password_hash): + return "wrong_password" + identity.password_hash = await run_cpu(hash_password, new_password) + await session.refresh(identity) + return None + + +async def ensure_billing_user_for_identity(session: AsyncSession, identity: Identity) -> int: + from database.users import add_user, check_user_exists + + if identity.tg_id is not None: + tid = int(identity.tg_id) + if not await check_user_exists(session, tid): + await add_user(session, tid) + ur = await session.execute(select(User).where(User.tg_id == tid).limit(1)) + u = ur.scalar_one() + await session.execute(update(User).where(User.id == u.id).values(identity_id=identity.id)) + return int(u.id) + res = await session.execute(select(User).where(User.identity_id == identity.id)) + row = res.scalars().first() + if row is not None: + return int(row.id) + new_u = User(identity_id=identity.id, tg_id=None) + session.add(new_u) + await session.flush() + return int(new_u.id) + + +async def merge_billing_user_into_telegram(session: AsyncSession, identity_id: str, telegram_tg_id: int) -> None: + from database.models import ( + CouponUsage, + Gift, + GiftUsage, + Key, + Notification, + Payment, + Referral, + ScheduledBroadcast, + TemporaryData, + ) + from database.access.resolution import resolve_user_optional + from database.users import invalidate_balance_cache, invalidate_profile_cache, update_balance + + res = await session.execute(select(User).where(User.identity_id == identity_id)) + rows = res.scalars().all() + if not rows: + return + billing = rows[0] + src_uid = int(billing.id) + dst_tg = int(telegram_tg_id) + if billing.tg_id is not None and int(billing.tg_id) > 0: + return + + dst_u = await resolve_user_optional(session, dst_tg) + if dst_u is None: + new_u = User( + tg_id=dst_tg, + identity_id=identity_id, + username=billing.username, + first_name=billing.first_name, + last_name=billing.last_name, + language_code=billing.language_code, + is_bot=billing.is_bot or False, + balance=float(billing.balance or 0.0), + trial=int(billing.trial or 0), + preferred_currency=billing.preferred_currency or "RUB", + source_code=billing.source_code, + ) + session.add(new_u) + await session.flush() + dst_uid = int(new_u.id) + else: + dst_uid = int(dst_u.id) + bal = float(billing.balance or 0.0) + if bal: + await update_balance(session, dst_uid, bal) + st = int(billing.trial or 0) + dt_r = await session.execute(select(User.trial).where(User.id == dst_uid)) + dt_val = dt_r.scalar_one_or_none() + if dt_val is not None and st > int(dt_val or 0): + await session.execute(update(User).where(User.id == dst_uid).values(trial=st)) + + await session.execute(update(Key).where(Key.user_id == src_uid).values(user_id=dst_uid)) + await session.execute(update(Payment).where(Payment.user_id == src_uid).values(user_id=dst_uid)) + + await session.execute( + text( + "DELETE FROM notifications AS n1 USING notifications AS n2 " + "WHERE n1.user_id = :src AND n2.user_id = :dst AND n1.notification_type = n2.notification_type" + ), + {"src": src_uid, "dst": dst_uid}, + ) + await session.execute(update(Notification).where(Notification.user_id == src_uid).values(user_id=dst_uid)) + + await session.execute(update(Gift).where(Gift.sender_user_id == src_uid).values(sender_user_id=dst_uid)) + await session.execute( + update(Gift).where(Gift.recipient_user_id == src_uid).values(recipient_user_id=dst_uid) + ) + + await session.execute( + text( + "DELETE FROM gift_usages AS g1 USING gift_usages AS g2 " + "WHERE g1.user_id = :src AND g2.user_id = :dst AND g1.gift_id = g2.gift_id" + ), + {"src": src_uid, "dst": dst_uid}, + ) + await session.execute(update(GiftUsage).where(GiftUsage.user_id == src_uid).values(user_id=dst_uid)) + + await session.execute( + text( + "DELETE FROM coupon_usages AS c1 USING coupon_usages AS c2 " + "WHERE c1.user_id = :src AND c2.user_id = :dst AND c1.coupon_id = c2.coupon_id" + ), + {"src": src_uid, "dst": dst_uid}, + ) + await session.execute(update(CouponUsage).where(CouponUsage.user_id == src_uid).values(user_id=dst_uid)) + + await session.execute(update(TemporaryData).where(TemporaryData.user_id == src_uid).values(user_id=dst_uid)) + + await session.execute( + update(ScheduledBroadcast) + .where(ScheduledBroadcast.created_by_user_id == src_uid) + .values(created_by_user_id=dst_uid) + ) + + await session.execute( + text( + "DELETE FROM referrals AS r1 USING referrals AS r2 " + "WHERE r1.referred_user_id = :src AND r2.referred_user_id = :dst " + "AND r1.referrer_user_id = r2.referrer_user_id" + ), + {"src": src_uid, "dst": dst_uid}, + ) + await session.execute( + text( + "DELETE FROM referrals AS r1 USING referrals AS r2 " + "WHERE r1.referrer_user_id = :src AND r2.referrer_user_id = :dst " + "AND r1.referred_user_id = r2.referred_user_id" + ), + {"src": src_uid, "dst": dst_uid}, + ) + await session.execute( + update(Referral).where(Referral.referred_user_id == src_uid).values(referred_user_id=dst_uid) + ) + await session.execute( + update(Referral).where(Referral.referrer_user_id == src_uid).values(referrer_user_id=dst_uid) + ) + + await refresh_tg_mirrors_for_user(session, dst_uid) + + await session.execute(delete(User).where(User.id == src_uid)) + await session.execute(update(User).where(User.id == dst_uid).values(identity_id=identity_id)) + + await invalidate_balance_cache(src_uid) + await invalidate_profile_cache(src_uid) + await invalidate_balance_cache(dst_uid) + await invalidate_profile_cache(dst_uid) + + +async def resolve_tg_id(session: AsyncSession, identity_id: str) -> int | None: + """По identity_id возвращает внутренний user id (users.id) для биллинга и ключей.""" + identity = await get_identity_by_id(session, identity_id) + if not identity: + return None + return await ensure_billing_user_for_identity(session, identity) async def attach_email(session: AsyncSession, identity_id: str, email: str) -> Identity | None: @@ -164,7 +361,6 @@ async def attach_email(session: AsyncSession, identity_id: str, email: str) -> I if existing and existing.id != identity_id: return None identity.email = email_clean - await session.commit() await session.refresh(identity) return identity @@ -177,12 +373,15 @@ async def attach_telegram(session: AsyncSession, identity_id: str, tg_id: int) - existing = await get_identity_by_tg_id(session, tg_id) if existing and existing.id != identity_id: return None + await merge_billing_user_into_telegram(session, identity_id, tg_id) + identity = await get_identity_by_id(session, identity_id) + if not identity: + return None identity.tg_id = tg_id admin_row = await session.execute(select(Admin).where(Admin.tg_id == tg_id)) if admin_row.scalar_one_or_none(): identity.is_admin = True await session.execute(User.__table__.update().where(User.tg_id == tg_id).values(identity_id=identity_id)) - await session.commit() await session.refresh(identity) return identity @@ -196,6 +395,5 @@ async def get_or_create_identity_for_tg(session: AsyncSession, tg_id: int) -> Id session.add(identity) await session.flush() await session.execute(User.__table__.update().where(User.tg_id == tg_id).values(identity_id=identity.id)) - await session.commit() await session.refresh(identity) return identity diff --git a/database/importer.py b/database/importer.py index 674ff7ed..c7e8c236 100644 --- a/database/importer.py +++ b/database/importer.py @@ -6,7 +6,6 @@ from datetime import datetime from itertools import cycle from sqlalchemy import select -from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from config import USE_COUNTRY_SELECTION @@ -76,53 +75,49 @@ async def import_keys_from_3xui_db(db_path: str, session: AsyncSession) -> tuple user_exists = await session.execute(select(User).where(User.tg_id == tg_id)) if not user_exists.scalar(): - try: - session.add( - User( - tg_id=tg_id, - username=None, - first_name=None, - last_name=None, - language_code=None, - is_bot=False, - balance=0.0, - trial=1, - source_code=None, - created_at=datetime.utcnow(), - updated_at=datetime.utcnow(), - ) + session.add( + User( + tg_id=tg_id, + username=None, + first_name=None, + last_name=None, + language_code=None, + is_bot=False, + balance=0.0, + trial=1, + source_code=None, + created_at=datetime.utcnow(), + updated_at=datetime.utcnow(), ) - except SQLAlchemyError as e: - await session.rollback() - raise RuntimeError(f"Ошибка при импорте пользователя tg_id={tg_id}") from e + ) + + await session.flush() + user_row = await session.execute(select(User.id).where(User.tg_id == tg_id)) + bill_uid = user_row.scalar_one() key_exists = await session.execute(select(Key).where(Key.client_id == client_id)) if key_exists.scalar(): skipped += 1 continue - try: - session.add( - Key( - tg_id=tg_id, - client_id=client_id, - email=email, - created_at=created_at, - expiry_time=expiry_time, - key="", - server_id=server_id, - remnawave_link=None, - tariff_id=None, - is_frozen=False, - alias=None, - notified=False, - notified_24h=False, - ) + session.add( + Key( + user_id=bill_uid, + tg_id=tg_id, + client_id=client_id, + email=email, + created_at=created_at, + expiry_time=expiry_time, + key="", + server_id=server_id, + remnawave_link=None, + tariff_id=None, + is_frozen=False, + alias=None, + notified=False, + notified_24h=False, ) - imported += 1 - except SQLAlchemyError as e: - await session.rollback() - raise RuntimeError(f"Ошибка при импорте ключа client_id={client_id}") from e + ) + imported += 1 - await session.commit() return imported, skipped diff --git a/database/keys.py b/database/keys.py index ac32f31a..7029f4e9 100644 --- a/database/keys.py +++ b/database/keys.py @@ -1,9 +1,8 @@ import asyncio -from datetime import datetime +from datetime import UTC, datetime from types import SimpleNamespace from sqlalchemy import delete, func, select, text, update -from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from core.cache_config import ( @@ -12,6 +11,7 @@ from core.cache_config import ( KEYS_LIST_CACHE_TTL_SEC, ) from core.redis_cache import cache_delete, cache_get, cache_key, cache_set +from database.access.resolution import resolve_user_optional from database.models import Key, Tariff, User from database.users import invalidate_profile_cache, invalidate_user_snapshot from logger import logger @@ -25,10 +25,22 @@ async def invalidate_key_email(client_id: str) -> None: await cache_delete(cache_key("key_email", client_id)) -async def invalidate_keys_list(tg_id: int) -> None: - await cache_delete(cache_key("keys_list", tg_id)) - await cache_delete(cache_key("key_count", tg_id)) - await invalidate_profile_cache(tg_id) +async def _purge_keys_cache_ids(*ids: int) -> None: + for i in ids: + await cache_delete(cache_key("keys_list", i)) + await cache_delete(cache_key("key_count", i)) + await invalidate_profile_cache(i) + + +async def invalidate_keys_list(session: AsyncSession, legacy_user_ref: int) -> None: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + await _purge_keys_cache_ids(legacy_user_ref) + return + if u.tg_id is not None: + await _purge_keys_cache_ids(u.id, u.tg_id) + else: + await _purge_keys_cache_ids(u.id) async def invalidate_key_details_by_client_id(session: AsyncSession, client_id: str) -> None: @@ -45,7 +57,7 @@ async def invalidate_key_details_by_client_id(session: AsyncSession, client_id: async def store_key( session: AsyncSession, - tg_id: int, + legacy_user_ref: int, client_id: str, email: str, expiry_time: int, @@ -61,71 +73,72 @@ async def store_key( current_traffic_limit: int | None = None, ): """Сохраняет или обновляет ключ подписки.""" - try: - exists = await session.execute(select(Key).where(Key.tg_id == tg_id, Key.client_id == client_id)) - existing_key = exists.scalar_one_or_none() + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + raise ValueError(f"Пользователь не найден для ключа: {legacy_user_ref}") + uid = u.id + exists = await session.execute(select(Key).where(Key.user_id == uid, Key.client_id == client_id)) + existing_key = exists.scalar_one_or_none() - if existing_key: - values: dict = { - "email": email, - "expiry_time": expiry_time, - "key": key, - "server_id": server_id, - "remnawave_link": remnawave_link, - "tariff_id": tariff_id, - "alias": alias, - } + if existing_key: + values: dict = { + "email": email, + "expiry_time": expiry_time, + "key": key, + "server_id": server_id, + "remnawave_link": remnawave_link, + "tariff_id": tariff_id, + "alias": alias, + "tg_id": u.tg_id, + } - if selected_device_limit is not None: - values["selected_device_limit"] = selected_device_limit - if selected_traffic_limit is not None: - values["selected_traffic_limit"] = selected_traffic_limit - if selected_price_rub is not None: - values["selected_price_rub"] = selected_price_rub - if current_device_limit is not None: - values["current_device_limit"] = current_device_limit - if current_traffic_limit is not None: - values["current_traffic_limit"] = current_traffic_limit + if selected_device_limit is not None: + values["selected_device_limit"] = selected_device_limit + if selected_traffic_limit is not None: + values["selected_traffic_limit"] = selected_traffic_limit + if selected_price_rub is not None: + values["selected_price_rub"] = selected_price_rub + if current_device_limit is not None: + values["current_device_limit"] = current_device_limit + if current_traffic_limit is not None: + values["current_traffic_limit"] = current_traffic_limit - await session.execute(update(Key).where(Key.tg_id == tg_id, Key.client_id == client_id).values(**values)) - logger.info(f"[Store Key] Ключ обновлён: tg_id={tg_id}, client_id={client_id}, server_id={server_id}") - else: - if current_device_limit is None: - current_device_limit = selected_device_limit - if current_traffic_limit is None: - current_traffic_limit = selected_traffic_limit + await session.execute(update(Key).where(Key.user_id == uid, Key.client_id == client_id).values(**values)) + logger.info(f"[Store Key] Ключ обновлён: user_id={uid}, client_id={client_id}, server_id={server_id}") + else: + if current_device_limit is None: + current_device_limit = selected_device_limit + if current_traffic_limit is None: + current_traffic_limit = selected_traffic_limit - new_key = Key( - tg_id=tg_id, - client_id=client_id, - email=email, - created_at=int(datetime.utcnow().timestamp() * 1000), - expiry_time=expiry_time, - key=key, - server_id=server_id, - remnawave_link=remnawave_link, - tariff_id=tariff_id, - alias=alias, - selected_device_limit=selected_device_limit, - selected_traffic_limit=selected_traffic_limit, - selected_price_rub=selected_price_rub, - current_device_limit=current_device_limit, - current_traffic_limit=current_traffic_limit, - ) - add_result = session.add(new_key) - if asyncio.iscoroutine(add_result): - await add_result - logger.info(f"[Store Key] Ключ создан: tg_id={tg_id}, client_id={client_id}, server_id={server_id}") + new_key = Key( + user_id=uid, + tg_id=u.tg_id, + client_id=client_id, + email=email, + created_at=int(datetime.now(UTC).timestamp() * 1000), + expiry_time=expiry_time, + key=key, + server_id=server_id, + remnawave_link=remnawave_link, + tariff_id=tariff_id, + alias=alias, + selected_device_limit=selected_device_limit, + selected_traffic_limit=selected_traffic_limit, + selected_price_rub=selected_price_rub, + current_device_limit=current_device_limit, + current_traffic_limit=current_traffic_limit, + ) + add_result = session.add(new_key) + if asyncio.iscoroutine(add_result): + await add_result + logger.info(f"[Store Key] Ключ создан: user_id={uid}, client_id={client_id}, server_id={server_id}") - await session.commit() - invalidate_user_snapshot(tg_id) - await invalidate_keys_list(tg_id) - await invalidate_key_details(email) - - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при сохранении ключа: {e}") - await session.rollback() - raise + invalidate_user_snapshot(uid) + if u.tg_id is not None: + invalidate_user_snapshot(u.tg_id) + await invalidate_keys_list(session, uid) + await invalidate_key_details(email) def _key_to_cache_dict(k: Key) -> dict: @@ -143,12 +156,16 @@ def _key_to_cache_dict(k: Key) -> dict: } -async def get_keys(session: AsyncSession, tg_id: int): - ckey = cache_key("keys_list", tg_id) +async def get_keys(session: AsyncSession, legacy_user_ref: int): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return [] + uid = u.id + ckey = cache_key("keys_list", uid) cached = await cache_get(ckey) if isinstance(cached, list): return [SimpleNamespace(**d) for d in cached] - result = await session.execute(select(Key).where(Key.tg_id == tg_id)) + result = await session.execute(select(Key).where(Key.user_id == uid)) rows = result.scalars().all() serialized = [_key_to_cache_dict(k) for k in rows] await cache_set(ckey, serialized, KEYS_LIST_CACHE_TTL_SEC) @@ -160,24 +177,33 @@ async def get_all_keys(session: AsyncSession): return result.scalars().all() -async def get_key_by_server(session: AsyncSession, tg_id: int, client_id: str): - stmt = select(Key).where(Key.tg_id == tg_id, Key.client_id == client_id) +async def get_key_by_server(session: AsyncSession, legacy_user_ref: int, client_id: str): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return None + stmt = select(Key).where(Key.user_id == u.id, Key.client_id == client_id) result = await session.execute(stmt) return result.scalar_one_or_none() -async def get_key_by_email(session: AsyncSession, email: str, tg_id: int | None = None) -> Key | None: +async def get_key_by_email(session: AsyncSession, email: str, legacy_user_ref: int | None = None) -> Key | None: stmt = select(Key).where(Key.email == email) - if tg_id is not None: - stmt = stmt.where(Key.tg_id == tg_id) + if legacy_user_ref is not None: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return None + stmt = stmt.where(Key.user_id == u.id) result = await session.execute(stmt.limit(1)) return result.scalar_one_or_none() -async def get_key_by_client_id(session: AsyncSession, client_id: str, tg_id: int | None = None) -> Key | None: +async def get_key_by_client_id(session: AsyncSession, client_id: str, legacy_user_ref: int | None = None) -> Key | None: stmt = select(Key).where(Key.client_id == client_id) - if tg_id is not None: - stmt = stmt.where(Key.tg_id == tg_id) + if legacy_user_ref is not None: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return None + stmt = stmt.where(Key.user_id == u.id) result = await session.execute(stmt.limit(1)) return result.scalar_one_or_none() @@ -218,15 +244,15 @@ async def get_key_details(session: AsyncSession, email: str) -> dict | None: if isinstance(cached, dict): return cached - stmt = select(Key, User).join(User, Key.tg_id == User.tg_id).where(Key.email == email) + stmt = select(Key, User).join(User, Key.user_id == User.id).where(Key.email == email) result = await session.execute(stmt) row = result.first() if not row: return None key, user = row - expiry_date = datetime.utcfromtimestamp(key.expiry_time / 1000) - current_date = datetime.utcnow() + expiry_date = datetime.fromtimestamp(key.expiry_time / 1000, UTC) + current_date = datetime.now(UTC) time_left = expiry_date - current_date if time_left.total_seconds() <= 0: @@ -267,39 +293,155 @@ async def get_key_details(session: AsyncSession, email: str) -> dict | None: return out -async def get_key_count(session: AsyncSession, tg_id: int) -> int: - cached = await cache_get(cache_key("key_count", tg_id)) +async def get_key_count(session: AsyncSession, legacy_user_ref: int) -> int: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return 0 + uid = u.id + cached = await cache_get(cache_key("key_count", uid)) if cached is not None: try: return int(cached) except (TypeError, ValueError): pass - result = await session.execute(select(func.count()).select_from(Key).where(Key.tg_id == tg_id)) + result = await session.execute(select(func.count()).select_from(Key).where(Key.user_id == uid)) count = result.scalar() or 0 - await cache_set(cache_key("key_count", tg_id), count, KEY_COUNT_CACHE_TTL_SEC) + await cache_set(cache_key("key_count", uid), count, KEY_COUNT_CACHE_TTL_SEC) return count -async def delete_key(session: AsyncSession, identifier: int | str, commit: bool = True): - tg_id_for_cache = None +async def get_key_by_user_and_email(session: AsyncSession, user_id: int, email: str) -> Key | None: + """Возвращает ORM-объект Key по паре (users.id, email) или None.""" + result = await session.execute( + select(Key).where(Key.user_id == int(user_id), Key.email == email) + ) + return result.scalar_one_or_none() + + +async def delete_key_by_user_and_email(session: AsyncSession, user_id: int, email: str) -> None: + """Удаляет ключ по паре (users.id, email). Commit — ответственность caller'а.""" + await session.execute( + delete(Key).where(Key.user_id == int(user_id), Key.email == email) + ) + + +async def get_user_keys_with_servers_by_email( + session: AsyncSession, user_id: int, email: str +) -> list[tuple[str, str, dict]]: + """Возвращает ключи пользователя + инфо о серверах (join Key × Server). + + Каждый элемент — ``(client_id, server_id, server_info_dict)``. Join + делается по (Key.server_id == Server.server_name OR Server.cluster_name), + чтобы поддержать и country-mode (server_id = cluster), и cluster-mode + (server_id = server_name). + + Используется в ``services.operations.traffic.get_user_traffic``. + """ + from sqlalchemy import or_ + + from database.models import Server + + join_cond = or_( + Key.server_id == Server.server_name, + Key.server_id == Server.cluster_name, + ) + result = await session.execute( + select(Key.client_id, Key.server_id, Server) + .select_from(Key) + .join(Server, join_cond) + .where(Server.enabled.is_(True), Key.user_id == int(user_id), Key.email == email) + ) + rows = [] + for client_id, server_id, server in result.all(): + rows.append(( + client_id, + server_id, + { + "server_name": server.server_name, + "cluster_name": server.cluster_name, + "api_url": server.api_url, + "panel_type": server.panel_type, + }, + )) + return rows + + +async def get_key_client_id_by_email_and_server( + session: AsyncSession, email: str, server_id: str +) -> str | None: + """Возвращает ``client_id`` первого ключа для пары (email, server_id). + + Используется для remnawave traffic reset, где нам нужен только client_id, + без остальных полей ключа. + """ + result = await session.execute( + select(Key.client_id) + .where(Key.email == email, Key.server_id == server_id) + .limit(1) + ) + return result.scalar() + + +async def count_keys_by_server_id(session: AsyncSession, server_id: str) -> int: + """Сколько всего ключей привязано к указанному server_id (кластеру или серверу). + + Используется для проверки max_keys лимита. ``server_id`` — строка + (у ``keys.server_id`` колонка типа String, содержит либо cluster_name, + либо server_name в зависимости от страны/кластера). + """ + result = await session.execute( + select(func.count()).select_from(Key).where(Key.server_id == server_id) + ) + return int(result.scalar() or 0) + + +async def get_all_key_server_ids(session: AsyncSession) -> list[str]: + """Список всех ``server_id`` из таблицы keys (с повторениями). + + Используется в ``services.clusters.select_cluster`` для подсчёта загрузки + кластеров. Возвращаем только server_id строки без подгрузки остальных + полей, чтобы не тянуть сотни мегабайт для огромных deployments. + """ + result = await session.execute(select(Key.server_id)) + return [row[0] for row in result.all() if row[0] is not None] + + +async def count_active_keys_for_user(session: AsyncSession, user_id: int) -> int: + """Количество незамороженных ключей у пользователя (по internal users.id). + + Отличается от `get_key_count`: не кэшируется и явно исключает замороженные. + Используется в проверке "новый пользователь" для купонных правил. + """ + result = await session.execute( + select(func.count()) + .select_from(Key) + .where(Key.user_id == int(user_id), Key.is_frozen.is_(False)) + ) + return int(result.scalar() or 0) + + +async def delete_key(session: AsyncSession, identifier: int | str): + legacy_for_cache = None email_for_cache = None if isinstance(identifier, str): res = await session.execute( - select(Key.tg_id, Key.email).where(Key.client_id == identifier).limit(1) + select(Key.user_id, Key.email).where(Key.client_id == identifier).limit(1) ) row = res.first() if row: - tg_id_for_cache, email_for_cache = row[0], row[1] + legacy_for_cache, email_for_cache = row[0], row[1] await cache_delete(cache_key("key_email", identifier)) + await session.execute(delete(Key).where(Key.client_id == identifier)) else: - tg_id_for_cache = identifier - stmt = delete(Key).where(Key.tg_id == identifier if isinstance(identifier, int) else Key.client_id == identifier) - await session.execute(stmt) - if commit: - await session.commit() - if tg_id_for_cache is not None: - invalidate_user_snapshot(tg_id_for_cache) - await invalidate_keys_list(tg_id_for_cache) + u = await resolve_user_optional(session, identifier) + if u is None: + logger.info(f"Ключ не удалён: пользователь {identifier} не найден") + return + legacy_for_cache = u.id + await session.execute(delete(Key).where(Key.user_id == u.id)) + if legacy_for_cache is not None: + invalidate_user_snapshot(legacy_for_cache) + await invalidate_keys_list(session, legacy_for_cache) if email_for_cache is not None: await invalidate_key_details(str(email_for_cache)) logger.info(f"Ключ с идентификатором {identifier} удалён") @@ -307,7 +449,6 @@ async def delete_key(session: AsyncSession, identifier: int | str, commit: bool async def update_key_expiry(session: AsyncSession, client_id: str, new_expiry_time: int): await session.execute(update(Key).where(Key.client_id == client_id).values(expiry_time=new_expiry_time)) - await session.commit() await invalidate_key_details_by_client_id(session, client_id) logger.info(f"Срок действия ключа {client_id} обновлён до {new_expiry_time}") @@ -317,59 +458,120 @@ async def get_client_id_by_email(session: AsyncSession, email: str): return result.scalar_one_or_none() -async def update_key_notified(session: AsyncSession, tg_id: int, client_id: str): - await session.execute(update(Key).where(Key.tg_id == tg_id, Key.client_id == client_id).values(notified=True)) - await session.commit() - await invalidate_keys_list(tg_id) +async def update_key_notified(session: AsyncSession, legacy_user_ref: int, client_id: str): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return + await session.execute(update(Key).where(Key.user_id == u.id, Key.client_id == client_id).values(notified=True)) + await invalidate_keys_list(session, u.id) await invalidate_key_details_by_client_id(session, client_id) -async def mark_key_as_frozen(session: AsyncSession, tg_id: int, client_id: str, time_left: int): +async def mark_key_as_frozen(session: AsyncSession, legacy_user_ref: int, client_id: str, time_left: int): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return await session.execute( text( """ UPDATE keys SET expiry_time = :expiry, is_frozen = TRUE - WHERE tg_id = :tg_id + WHERE user_id = :user_id AND client_id = :client_id """ ), - {"expiry": time_left, "tg_id": tg_id, "client_id": client_id}, + {"expiry": time_left, "user_id": u.id, "client_id": client_id}, ) - await invalidate_keys_list(tg_id) + await invalidate_keys_list(session, u.id) await invalidate_key_details_by_client_id(session, client_id) async def mark_key_as_unfrozen( session: AsyncSession, - tg_id: int, + legacy_user_ref: int, client_id: str, new_expiry_time: int, ): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return await session.execute( text( """ UPDATE keys SET expiry_time = :expiry, is_frozen = FALSE - WHERE tg_id = :tg_id + WHERE user_id = :user_id AND client_id = :client_id """ ), - {"expiry": new_expiry_time, "tg_id": tg_id, "client_id": client_id}, + {"expiry": new_expiry_time, "user_id": u.id, "client_id": client_id}, ) - await invalidate_keys_list(tg_id) + await invalidate_keys_list(session, u.id) await invalidate_key_details_by_client_id(session, client_id) async def update_key_tariff(session: AsyncSession, client_id: str, tariff_id: int): await session.execute(update(Key).where(Key.client_id == client_id).values(tariff_id=tariff_id)) - await session.commit() await invalidate_key_details_by_client_id(session, client_id) logger.info(f"Тариф ключа {client_id} обновлён на {tariff_id}") +async def update_key_renewal_snapshot( + session: AsyncSession, + email: str, + *, + tariff_id: int, + selected_device_limit: int | None = None, + current_device_limit: int | None = None, + selected_traffic_limit: int | None = None, + current_traffic_limit: int | None = None, + apply_limits: bool = True, +) -> None: + """Обновляет tariff_id и (опционально) лимиты ключа после продления. + + ``apply_limits=True`` — выставить все четыре лимита (для non-configurable + тарифов). ``apply_limits=False`` — обновить только ``tariff_id``, лимиты + не трогать (configurable-тарифы обновляют их через `save_key_config_with_mode`). + """ + values: dict = {"tariff_id": tariff_id} + if apply_limits: + values["selected_device_limit"] = selected_device_limit + values["current_device_limit"] = current_device_limit + values["selected_traffic_limit"] = selected_traffic_limit + values["current_traffic_limit"] = current_traffic_limit + await session.execute(update(Key).where(Key.email == email).values(**values)) + await invalidate_key_details(email) + + +async def update_key_post_creation_snapshot( + session: AsyncSession, + *, + user_id: int, + email: str, + selected_device_limit: int | None, + selected_traffic_limit: int | None, + selected_price_rub: int | None, +) -> None: + """Дозаписывает выбранные пользователем параметры ключа сразу после создания. + + Используется из `services.keys.create_vpn_key_headless` — тариф/лимиты не + всегда известны на момент `create_key_on_cluster`, поэтому после него + идёт snapshot-апдейт для полей, которые нужны для отображения в UI. + """ + await session.execute( + update(Key) + .where(Key.user_id == int(user_id), Key.email == email) + .values( + selected_device_limit=selected_device_limit, + selected_traffic_limit=selected_traffic_limit, + selected_price_rub=selected_price_rub, + ) + ) + await invalidate_key_details(email) + + async def get_subscription_link(session: AsyncSession, email: str) -> str | None: result = await session.execute(select(func.coalesce(Key.key, Key.remnawave_link)).where(Key.email == email)) return result.scalar_one_or_none() @@ -377,7 +579,6 @@ async def get_subscription_link(session: AsyncSession, email: str) -> str | None async def update_key_client_id(session: AsyncSession, email: str, new_client_id: str): await session.execute(update(Key).where(Key.email == email).values(client_id=new_client_id)) - await session.commit() await invalidate_key_details(email) logger.info(f"client_id обновлён для {email} -> {new_client_id}") @@ -385,7 +586,6 @@ async def update_key_client_id(session: AsyncSession, email: str, new_client_id: async def update_key_link(session: AsyncSession, email: str, link: str) -> bool: q = update(Key).where(Key.email == email).values(key=link).returning(Key.client_id) res = await session.execute(q) - await session.commit() ok = res.scalar_one_or_none() is not None if ok: await invalidate_key_details(email) @@ -403,7 +603,6 @@ async def update_key_subscription_links(session: AsyncSession, email: str, link: .returning(Key.client_id) ) res = await session.execute(stmt) - await session.commit() ok = res.scalar_one_or_none() is not None if ok: await invalidate_key_details(email) @@ -444,10 +643,13 @@ async def save_key_config_with_mode( await invalidate_key_details(email) -async def reset_key_tariff_state(session: AsyncSession, tg_id: int, email: str, tariff_id: int) -> None: +async def reset_key_tariff_state(session: AsyncSession, legacy_user_ref: int, email: str, tariff_id: int) -> None: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return await session.execute( update(Key) - .where(Key.tg_id == tg_id, Key.email == email) + .where(Key.user_id == u.id, Key.email == email) .values( tariff_id=tariff_id, selected_device_limit=None, @@ -457,25 +659,27 @@ async def reset_key_tariff_state(session: AsyncSession, tg_id: int, email: str, selected_price_rub=None, ) ) - await session.commit() - await invalidate_keys_list(tg_id) + await invalidate_keys_list(session, u.id) await invalidate_key_details(email) async def save_key_tariff_selection( session: AsyncSession, - tg_id: int, + legacy_user_ref: int, email: str, tariff_id: int, selected_devices: int | None, selected_traffic_gb: int | None, ) -> None: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return selected_devices_val = int(selected_devices) if selected_devices is not None else None selected_traffic_val = int(selected_traffic_gb) if selected_traffic_gb is not None and int(selected_traffic_gb) > 0 else None await session.execute( update(Key) - .where(Key.tg_id == tg_id, Key.email == email) + .where(Key.user_id == u.id, Key.email == email) .values( tariff_id=tariff_id, selected_device_limit=selected_devices_val, @@ -485,8 +689,7 @@ async def save_key_tariff_selection( selected_price_rub=None, ) ) - await session.commit() - await invalidate_keys_list(tg_id) + await invalidate_keys_list(session, u.id) await invalidate_key_details(email) @@ -510,7 +713,6 @@ async def save_admin_key_config( selected_price_rub=selected_price, ) ) - await session.commit() await invalidate_key_details(email) @@ -527,6 +729,5 @@ async def reset_key_current_limits_to_selected(session: AsyncSession, client_id: ), {"client_id": client_id}, ) - await session.commit() await invalidate_key_details_by_client_id(session, client_id) logger.info(f"Текущие лимиты ключа {client_id} сброшены к выбранным") diff --git a/database/migrations/__init__.py b/database/migrations/__init__.py new file mode 100644 index 00000000..4fc2cf2f --- /dev/null +++ b/database/migrations/__init__.py @@ -0,0 +1 @@ +from .schema_upgrade import * diff --git a/database/migrations/schema_upgrade.py b/database/migrations/schema_upgrade.py new file mode 100644 index 00000000..072bb77f --- /dev/null +++ b/database/migrations/schema_upgrade.py @@ -0,0 +1,1206 @@ +from __future__ import annotations + +import re + +from sqlalchemy import text +from sqlalchemy.ext.asyncio import AsyncConnection + +from config import DATABASE_URL +from logger import logger + + +def _is_postgresql() -> bool: + u = (DATABASE_URL or "").lower() + return "+asyncpg" in u or u.startswith("postgresql") + + +async def _table_exists(conn: AsyncConnection, table: str) -> bool: + r = await conn.execute( + text( + """ + SELECT 1 + FROM information_schema.tables + WHERE table_schema = 'public' AND table_name = :t + """ + ), + {"t": table}, + ) + return r.first() is not None + + +async def _ensure_migrations_table(conn: AsyncConnection) -> None: + if not await _table_exists(conn, "schema_migrations"): + await conn.execute( + text( + """ + CREATE TABLE schema_migrations ( + version INTEGER PRIMARY KEY, + applied_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + description TEXT + ) + """ + ) + ) + + +async def _get_current_version(conn: AsyncConnection) -> int: + await _ensure_migrations_table(conn) + r = await conn.execute( + text("SELECT COALESCE(MAX(version), 0) FROM schema_migrations") + ) + row = r.first() + return int(row[0]) if row else 0 + + +async def _mark_migration_applied(conn: AsyncConnection, version: int, description: str) -> None: + await conn.execute( + text( + """ + INSERT INTO schema_migrations (version, description) + VALUES (:v, :d) + ON CONFLICT (version) DO NOTHING + """ + ), + {"v": version, "d": description}, + ) + + +async def _users_pk_columns(conn: AsyncConnection) -> list[str]: + r = await conn.execute( + text( + """ + SELECT a.attname + FROM pg_constraint c + JOIN pg_class t ON t.oid = c.conrelid + JOIN pg_namespace n ON n.oid = t.relnamespace + JOIN unnest(c.conkey) WITH ORDINALITY AS u(attnum, ord) ON true + JOIN pg_attribute a ON a.attrelid = t.oid AND a.attnum = u.attnum + WHERE n.nspname = 'public' + AND t.relname = 'users' + AND c.contype = 'p' + ORDER BY u.ord + """ + ) + ) + return [row[0] for row in r.all()] + + +async def _column_exists(conn: AsyncConnection, table: str, column: str) -> bool: + r = await conn.execute( + text( + """ + SELECT 1 + FROM information_schema.columns + WHERE table_schema = 'public' AND table_name = :t AND column_name = :c + """ + ), + {"t": table, "c": column}, + ) + return r.first() is not None + + +async def _column_is_identity(conn: AsyncConnection, table: str, column: str) -> bool: + r = await conn.execute( + text( + """ + SELECT is_identity + FROM information_schema.columns + WHERE table_schema = 'public' AND table_name = :t AND column_name = :c + """ + ), + {"t": table, "c": column}, + ) + row = r.first() + return bool(row and str(row[0]).upper() == "YES") + + +async def _constraint_exists(conn: AsyncConnection, table: str, constraint: str) -> bool: + r = await conn.execute( + text( + """ + SELECT 1 + FROM pg_constraint c + JOIN pg_class t ON t.oid = c.conrelid + JOIN pg_namespace n ON n.oid = t.relnamespace + WHERE n.nspname = 'public' + AND t.relname = :t + AND c.conname = :c + """ + ), + {"t": table, "c": constraint}, + ) + return r.first() is not None + + +async def _exec_ignore(conn: AsyncConnection, sql: str) -> None: + try: + async with conn.begin_nested(): + await conn.execute(text(sql)) + except Exception as e: + logger.debug(f"[schema_upgrade] skip: {e}") + + +async def _add_constraint_if_missing(conn: AsyncConnection, table: str, name: str, sql: str) -> None: + if await _constraint_exists(conn, table, name): + return + await _exec_ignore(conn, sql) + + +async def _drop_fkeys_to_users(conn: AsyncConnection) -> None: + r = await conn.execute( + text( + """ + SELECT con.conname, rel.relname AS src_table + FROM pg_constraint con + JOIN pg_class rel ON rel.oid = con.conrelid + JOIN pg_namespace nsp ON nsp.oid = rel.relnamespace + WHERE con.confrelid = 'users'::regclass + AND con.contype = 'f' + AND nsp.nspname = 'public' + """ + ) + ) + for row in r.all(): + cname, src = row[0], row[1] + await conn.execute(text(f'ALTER TABLE "{src}" DROP CONSTRAINT IF EXISTS "{cname}"')) + + +async def _drop_pk(conn: AsyncConnection, table: str) -> None: + if not await _table_exists(conn, table): + return + r = await conn.execute( + text( + """ + SELECT tc.constraint_name + FROM information_schema.table_constraints tc + WHERE tc.table_schema = 'public' + AND tc.table_name = :t + AND tc.constraint_type = 'PRIMARY KEY' + """ + ), + {"t": table}, + ) + row = r.first() + if row: + await conn.execute(text(f'ALTER TABLE "{table}" DROP CONSTRAINT IF EXISTS "{row[0]}"')) + + +async def _column_has_nulls(conn: AsyncConnection, table: str, column: str) -> bool: + r = await conn.execute( + text(f'SELECT 1 FROM "{table}" WHERE "{column}" IS NULL LIMIT 1') + ) + return r.first() is not None + + +async def _safe_set_not_null(conn: AsyncConnection, table: str, column: str) -> bool: + if await _column_has_nulls(conn, table, column): + logger.warning(f"[schema_upgrade] {table}.{column} содержит NULL, пропуск SET NOT NULL") + return False + await conn.execute(text(f'ALTER TABLE "{table}" ALTER COLUMN "{column}" SET NOT NULL')) + return True + + +async def _index_exists(conn: AsyncConnection, table: str, index: str) -> bool: + r = await conn.execute( + text( + """ + SELECT 1 + FROM pg_indexes + WHERE schemaname = 'public' + AND tablename = :t + AND indexname = :i + """ + ), + {"t": table, "i": index}, + ) + return r.first() is not None + + +async def _ensure_users_id_referenceable(conn: AsyncConnection) -> None: + if not await _column_exists(conn, "users", "id"): + return + if await _column_is_identity(conn, "users", "id"): + await _exec_ignore(conn, "UPDATE users SET id = DEFAULT WHERE id IS NULL") + await _exec_ignore( + conn, + """ + WITH d AS ( + SELECT ctid, row_number() OVER (PARTITION BY id ORDER BY ctid) AS rn + FROM users + WHERE id IS NOT NULL + ) + UPDATE users u + SET id = DEFAULT + FROM d + WHERE u.ctid = d.ctid AND d.rn > 1 + """, + ) + else: + await _exec_ignore(conn, "CREATE SEQUENCE IF NOT EXISTS users_id_seq") + await _exec_ignore(conn, "ALTER TABLE users ALTER COLUMN id SET DEFAULT nextval('users_id_seq')") + await _exec_ignore(conn, "ALTER SEQUENCE users_id_seq OWNED BY users.id") + await _exec_ignore(conn, "UPDATE users SET id = nextval('users_id_seq') WHERE id IS NULL") + await _exec_ignore( + conn, + """ + WITH d AS ( + SELECT ctid, row_number() OVER (PARTITION BY id ORDER BY ctid) AS rn + FROM users + WHERE id IS NOT NULL + ) + UPDATE users u + SET id = nextval('users_id_seq') + FROM d + WHERE u.ctid = d.ctid AND d.rn > 1 + """, + ) + await _exec_ignore(conn, "CREATE UNIQUE INDEX IF NOT EXISTS ix_users_id ON users (id)") + + +async def _migration_v1_add_users_id(conn: AsyncConnection) -> None: + if not await _table_exists(conn, "users"): + return + + pk = await _users_pk_columns(conn) + if not pk: + return + if pk == ["id"]: + return + + if pk != ["tg_id"]: + logger.warning(f"[schema_upgrade] users PK неожиданен {pk}, пропуск v1") + return + + logger.info("[schema_upgrade] v1: Добавление users.id") + + if not await _column_exists(conn, "users", "id"): + await conn.execute( + text( + """ + ALTER TABLE users + ADD COLUMN id BIGINT GENERATED BY DEFAULT AS IDENTITY NOT NULL + """ + ) + ) + await _ensure_users_id_referenceable(conn) + + +async def _migration_v2_add_user_id_columns(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v2: Добавление user_id колонок в связанные таблицы") + + tables_columns = [ + ("keys", "user_id"), + ("payments", "user_id"), + ("referrals", "referred_user_id"), + ("referrals", "referrer_user_id"), + ("notifications", "user_id"), + ("scheduled_broadcasts", "created_by_user_id"), + ("gifts", "sender_user_id"), + ("gifts", "recipient_user_id"), + ("gift_usages", "user_id"), + ("coupon_usages", "account_user_id"), + ("temporary_data", "user_id"), + ("manual_bans", "user_id"), + ("blocked_users", "user_id"), + ] + + for table, column in tables_columns: + if await _table_exists(conn, table) and not await _column_exists(conn, table, column): + await conn.execute(text(f'ALTER TABLE "{table}" ADD COLUMN {column} BIGINT')) + + +async def _backfill_users_from_table( + conn: AsyncConnection, table: str, tg_col: str = "tg_id" +) -> int: + """Auto-создание users для orphan tg_id'ов из указанной таблицы. + + Legacy клиенты обновляются с TG-only схемы (где только tg_id), и в связанных + таблицах могут быть строки, ссылающиеся на tg_id, которого нет в users. Вместо + удаления таких строк — создаём минимальную users-запись, чтобы FK/NOT NULL + проходили и данные сохранялись. + """ + if not await _table_exists(conn, table): + return 0 + if not await _column_exists(conn, table, tg_col): + return 0 + if not await _table_exists(conn, "users"): + return 0 + if not await _column_exists(conn, "users", "tg_id"): + return 0 + + has_created_at = await _column_exists(conn, "users", "created_at") + has_updated_at = await _column_exists(conn, "users", "updated_at") + + cols = ["tg_id"] + vals = [f't."{tg_col}"'] + if has_created_at: + cols.append("created_at") + vals.append("NOW()") + if has_updated_at: + cols.append("updated_at") + vals.append("NOW()") + + cols_sql = ", ".join(cols) + vals_sql = ", ".join(vals) + + result = await conn.execute( + text( + f""" + INSERT INTO users ({cols_sql}) + SELECT DISTINCT {vals_sql} + FROM "{table}" t + WHERE t."{tg_col}" IS NOT NULL + AND NOT EXISTS ( + SELECT 1 FROM users u WHERE u.tg_id = t."{tg_col}" + ) + """ + ) + ) + created = result.rowcount or 0 + if created > 0: + logger.info( + f"[schema_upgrade] users backfill: создано {created} юзеров из orphan {table}.{tg_col}" + ) + return created + + +async def _migration_v3_populate_user_ids(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v3: Заполнение user_id из tg_id") + + if not await _table_exists(conn, "users") or not await _column_exists(conn, "users", "id"): + return + + updates = [ + ("keys", "user_id", "tg_id"), + ("payments", "user_id", "tg_id"), + ("notifications", "user_id", "tg_id"), + ("scheduled_broadcasts", "created_by_user_id", "created_by_tg_id"), + ("gift_usages", "user_id", "tg_id"), + ("temporary_data", "user_id", "tg_id"), + ("manual_bans", "user_id", "tg_id"), + ("blocked_users", "user_id", "tg_id"), + ] + + for table, _user_col, tg_col in updates: + await _backfill_users_from_table(conn, table, tg_col) + + for table, user_col, tg_col in updates: + if await _table_exists(conn, table) and await _column_exists(conn, table, user_col): + result = await conn.execute( + text( + f""" + UPDATE "{table}" t + SET {user_col} = u.id + FROM users u + WHERE t.{user_col} IS NULL AND t.{tg_col} = u.tg_id + """ + ) + ) + updated = result.rowcount + if updated > 0: + logger.debug(f"[schema_upgrade] v3: заполнено {updated} записей {user_col} в {table}") + + null_count = await conn.execute( + text(f'SELECT COUNT(*) FROM "{table}" WHERE {user_col} IS NULL') + ) + nulls = null_count.scalar() + if nulls > 0: + logger.warning(f"[schema_upgrade] v3: в {table} осталось {nulls} записей с NULL {user_col}") + + if await _table_exists(conn, "referrals"): + await _backfill_users_from_table(conn, "referrals", "referred_tg_id") + await _backfill_users_from_table(conn, "referrals", "referrer_tg_id") + if await _column_exists(conn, "referrals", "referred_user_id"): + await conn.execute( + text( + """ + UPDATE referrals r + SET referred_user_id = u.id + FROM users u + WHERE r.referred_user_id IS NULL AND r.referred_tg_id = u.tg_id + """ + ) + ) + if await _column_exists(conn, "referrals", "referrer_user_id"): + await conn.execute( + text( + """ + UPDATE referrals r + SET referrer_user_id = u.id + FROM users u + WHERE r.referrer_user_id IS NULL AND r.referrer_tg_id = u.tg_id + """ + ) + ) + + if await _table_exists(conn, "gifts"): + await _backfill_users_from_table(conn, "gifts", "sender_tg_id") + await _backfill_users_from_table(conn, "gifts", "recipient_tg_id") + if await _column_exists(conn, "gifts", "sender_user_id"): + await conn.execute( + text( + """ + UPDATE gifts g + SET sender_user_id = u.id + FROM users u + WHERE g.sender_user_id IS NULL AND g.sender_tg_id = u.tg_id + """ + ) + ) + if await _column_exists(conn, "gifts", "recipient_user_id"): + await conn.execute( + text( + """ + UPDATE gifts g + SET recipient_user_id = u.id + FROM users u + WHERE g.recipient_user_id IS NULL AND g.recipient_tg_id IS NOT NULL + AND g.recipient_tg_id = u.tg_id + """ + ) + ) + + if await _table_exists(conn, "coupon_usages") and await _column_exists(conn, "coupon_usages", "account_user_id"): + await _backfill_users_from_table(conn, "coupon_usages", "user_id") + await conn.execute( + text( + """ + UPDATE coupon_usages c + SET account_user_id = u.id + FROM users u + WHERE c.account_user_id IS NULL AND c.user_id = u.tg_id + """ + ) + ) + + +async def _migration_v4_add_tg_id_mirrors(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v4: Добавление tg_id mirror колонок") + + mirrors = [ + ("referrals", "referred_tg_id"), + ("referrals", "referrer_tg_id"), + ("notifications", "tg_id"), + ("gift_usages", "tg_id"), + ("keys", "tg_id"), + ("payments", "tg_id"), + ("gifts", "sender_tg_id"), + ("gifts", "recipient_tg_id"), + ("scheduled_broadcasts", "created_by_tg_id"), + ("coupon_usages", "tg_id"), + ("temporary_data", "tg_id"), + ("manual_bans", "tg_id"), + ("blocked_users", "tg_id"), + ] + + for table, column in mirrors: + if await _table_exists(conn, table) and not await _column_exists(conn, table, column): + await conn.execute(text(f'ALTER TABLE "{table}" ADD COLUMN {column} BIGINT')) + + +async def _migration_v5_switch_pks_to_user_id(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v5: Переключение PK на user_id где возможно") + + await _drop_fkeys_to_users(conn) + await _ensure_users_id_referenceable(conn) + + if await _table_exists(conn, "referrals"): + can_harden = not await _column_has_nulls(conn, "referrals", "referred_user_id") + can_harden = can_harden and not await _column_has_nulls(conn, "referrals", "referrer_user_id") + if can_harden: + await _drop_pk(conn, "referrals") + await conn.execute(text("ALTER TABLE referrals ALTER COLUMN referred_user_id SET NOT NULL")) + await conn.execute(text("ALTER TABLE referrals ALTER COLUMN referrer_user_id SET NOT NULL")) + await conn.execute( + text("ALTER TABLE referrals ADD PRIMARY KEY (referred_user_id, referrer_user_id)") + ) + else: + logger.warning("[schema_upgrade] referrals содержит NULL user_id, пропуск перевода PK") + + if await _table_exists(conn, "notifications") and await _safe_set_not_null(conn, "notifications", "user_id"): + await _drop_pk(conn, "notifications") + await conn.execute( + text("ALTER TABLE notifications ADD PRIMARY KEY (user_id, notification_type)") + ) + + if await _table_exists(conn, "gift_usages") and await _safe_set_not_null(conn, "gift_usages", "user_id"): + await _drop_pk(conn, "gift_usages") + await conn.execute(text("ALTER TABLE gift_usages ADD PRIMARY KEY (gift_id, user_id)")) + + if await _table_exists(conn, "coupon_usages"): + await _drop_pk(conn, "coupon_usages") + has_user_id = await _column_exists(conn, "coupon_usages", "user_id") + has_account_user_id = await _column_exists(conn, "coupon_usages", "account_user_id") + if has_account_user_id and not has_user_id: + await conn.execute(text("ALTER TABLE coupon_usages RENAME COLUMN account_user_id TO user_id")) + elif has_account_user_id and has_user_id: + await conn.execute( + text( + """ + UPDATE coupon_usages + SET user_id = account_user_id + WHERE account_user_id IS NOT NULL + """ + ) + ) + await conn.execute(text("ALTER TABLE coupon_usages DROP COLUMN account_user_id")) + elif not has_user_id: + await conn.execute(text("ALTER TABLE coupon_usages ADD COLUMN user_id BIGINT")) + if await _safe_set_not_null(conn, "coupon_usages", "user_id"): + await conn.execute(text("ALTER TABLE coupon_usages ADD PRIMARY KEY (coupon_id, user_id)")) + + for tbl in ("temporary_data", "manual_bans", "blocked_users"): + if not await _table_exists(conn, tbl): + continue + if not await _column_exists(conn, tbl, "user_id"): + continue + logger.info(f"[schema_upgrade] {tbl} оставлен на legacy PK по tg_id") + + if await _table_exists(conn, "users"): + await _drop_pk(conn, "users") + await conn.execute(text("ALTER TABLE users ADD PRIMARY KEY (id)")) + await conn.execute(text("ALTER TABLE users ALTER COLUMN tg_id DROP NOT NULL")) + await conn.execute(text("CREATE UNIQUE INDEX IF NOT EXISTS uq_users_tg_id ON users (tg_id)")) + + +async def _migration_v6_add_foreign_keys(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v6: Добавление foreign key constraints") + + if await _table_exists(conn, "referrals"): + await _add_constraint_if_missing( + conn, + "referrals", + "fk_referrals_referred_user", + """ + ALTER TABLE referrals + ADD CONSTRAINT fk_referrals_referred_user + FOREIGN KEY (referred_user_id) REFERENCES users (id) ON DELETE CASCADE + """, + ) + await _add_constraint_if_missing( + conn, + "referrals", + "fk_referrals_referrer_user", + """ + ALTER TABLE referrals + ADD CONSTRAINT fk_referrals_referrer_user + FOREIGN KEY (referrer_user_id) REFERENCES users (id) ON DELETE CASCADE + """, + ) + + if await _table_exists(conn, "notifications"): + await _add_constraint_if_missing( + conn, + "notifications", + "fk_notifications_user", + """ + ALTER TABLE notifications + ADD CONSTRAINT fk_notifications_user + FOREIGN KEY (user_id) REFERENCES users (id) ON DELETE CASCADE + """, + ) + + if await _table_exists(conn, "gift_usages"): + await _add_constraint_if_missing( + conn, + "gift_usages", + "fk_gift_usages_user", + """ + ALTER TABLE gift_usages + ADD CONSTRAINT fk_gift_usages_user + FOREIGN KEY (user_id) REFERENCES users (id) ON DELETE CASCADE + """, + ) + + if await _table_exists(conn, "coupon_usages"): + await _add_constraint_if_missing( + conn, + "coupon_usages", + "fk_coupon_usages_user", + """ + ALTER TABLE coupon_usages + ADD CONSTRAINT fk_coupon_usages_user + FOREIGN KEY (user_id) REFERENCES users (id) ON DELETE CASCADE + """, + ) + + for tbl in ("temporary_data", "manual_bans", "blocked_users"): + if await _table_exists(conn, tbl): + safe = re.sub(r"[^a-z_]", "_", tbl) + await _add_constraint_if_missing( + conn, + tbl, + f"fk_{safe}_user", + f""" + ALTER TABLE "{tbl}" + ADD CONSTRAINT fk_{safe}_user + FOREIGN KEY (user_id) REFERENCES users (id) ON DELETE CASCADE + """, + ) + + if await _table_exists(conn, "keys") and await _safe_set_not_null(conn, "keys", "user_id"): + await _add_constraint_if_missing( + conn, + "keys", + "fk_keys_user", + """ + ALTER TABLE keys + ADD CONSTRAINT fk_keys_user FOREIGN KEY (user_id) REFERENCES users (id) ON DELETE CASCADE + """, + ) + + if await _table_exists(conn, "payments") and await _safe_set_not_null(conn, "payments", "user_id"): + await _add_constraint_if_missing( + conn, + "payments", + "fk_payments_user", + """ + ALTER TABLE payments + ADD CONSTRAINT fk_payments_user FOREIGN KEY (user_id) REFERENCES users (id) ON DELETE CASCADE + """, + ) + + if await _table_exists(conn, "gifts"): + if await _safe_set_not_null(conn, "gifts", "sender_user_id"): + await _add_constraint_if_missing( + conn, + "gifts", + "fk_gifts_sender_user", + """ + ALTER TABLE gifts + ADD CONSTRAINT fk_gifts_sender_user FOREIGN KEY (sender_user_id) REFERENCES users (id) ON DELETE CASCADE + """, + ) + await _add_constraint_if_missing( + conn, + "gifts", + "fk_gifts_recipient_user", + """ + ALTER TABLE gifts + ADD CONSTRAINT fk_gifts_recipient_user FOREIGN KEY (recipient_user_id) REFERENCES users (id) ON DELETE SET NULL + """, + ) + + if await _table_exists(conn, "scheduled_broadcasts"): + await _add_constraint_if_missing( + conn, + "scheduled_broadcasts", + "fk_scheduled_broadcasts_creator_user", + """ + ALTER TABLE scheduled_broadcasts + ADD CONSTRAINT fk_scheduled_broadcasts_creator_user + FOREIGN KEY (created_by_user_id) REFERENCES users (id) ON DELETE SET NULL + """, + ) + + +async def _migration_v7_backfill_tg_mirrors(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v7: Backfill tg_id mirrors") + + backfills = [ + "UPDATE keys k SET tg_id = u.tg_id FROM users u WHERE k.user_id = u.id", + "UPDATE payments p SET tg_id = u.tg_id FROM users u WHERE p.user_id = u.id", + "UPDATE referrals r SET referred_tg_id = ur.tg_id, referrer_tg_id = ux.tg_id FROM users ur, users ux WHERE r.referred_user_id = ur.id AND r.referrer_user_id = ux.id", + "UPDATE notifications n SET tg_id = u.tg_id FROM users u WHERE n.user_id = u.id", + "UPDATE gift_usages gu SET tg_id = u.tg_id FROM users u WHERE gu.user_id = u.id", + "UPDATE manual_bans m SET tg_id = u.tg_id FROM users u WHERE m.user_id = u.id", + "UPDATE temporary_data t SET tg_id = u.tg_id FROM users u WHERE t.user_id = u.id", + "UPDATE blocked_users b SET tg_id = u.tg_id FROM users u WHERE b.user_id = u.id", + "UPDATE scheduled_broadcasts s SET created_by_tg_id = u.tg_id FROM users u WHERE s.created_by_user_id = u.id", + "UPDATE coupon_usages c SET tg_id = u.tg_id FROM users u WHERE c.user_id = u.id", + ] + + for sql in backfills: + try: + async with conn.begin_nested(): + await conn.execute(text(sql)) + except Exception as e: + logger.debug(f"[schema_upgrade] backfill skip: {e}") + + if await _table_exists(conn, "gifts"): + try: + async with conn.begin_nested(): + await conn.execute( + text( + """ + UPDATE gifts g + SET sender_tg_id = u.tg_id + FROM users u + WHERE g.sender_user_id = u.id + """ + ) + ) + except Exception as e: + logger.debug(f"[schema_upgrade] backfill gifts sender skip: {e}") + + try: + async with conn.begin_nested(): + await conn.execute( + text( + """ + UPDATE gifts g + SET recipient_tg_id = u.tg_id + FROM users u + WHERE g.recipient_user_id = u.id + """ + ) + ) + except Exception as e: + logger.debug(f"[schema_upgrade] backfill gifts recipient skip: {e}") + + +async def _migration_v8_fix_notification_timezone(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v8: Исправление timezone для last_notification_time") + + if not await _table_exists(conn, "notifications"): + return + + r = await conn.execute( + text( + """ + SELECT data_type + FROM information_schema.columns + WHERE table_schema = 'public' + AND table_name = 'notifications' + AND column_name = 'last_notification_time' + """ + ) + ) + row = r.first() + if not row: + return + + current_type = row[0] + if current_type == "timestamp with time zone": + return + + try: + await conn.execute( + text( + """ + ALTER TABLE notifications + ALTER COLUMN last_notification_time + TYPE TIMESTAMP WITH TIME ZONE + USING last_notification_time AT TIME ZONE 'UTC' + """ + ) + ) + except Exception as e: + logger.warning(f"[schema_upgrade] v8: не удалось изменить тип колонки: {e}") + + +async def _migration_v9_cleanup_orphaned_records(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v5: Мягкий backfill user_id для legacy таблиц") + + tables_to_clean = [ + ("blocked_users", "user_id"), + ("manual_bans", "user_id"), + ("temporary_data", "user_id"), + ] + + for table, user_col in tables_to_clean: + if not await _table_exists(conn, table): + continue + + if not await _column_exists(conn, table, user_col): + continue + + result = await conn.execute( + text( + f""" + UPDATE "{table}" AS t + SET "{user_col}" = u.id + FROM users AS u + WHERE t."{user_col}" IS NULL + AND t.tg_id IS NOT NULL + AND t.tg_id = u.tg_id + """ + ) + ) + updated = result.rowcount + if updated > 0: + logger.info(f"[schema_upgrade] v5: заполнено {updated} записей в {table}") + + +async def _migration_v10_finalize_user_id_not_null(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v10: legacy таблицы сохраняют nullable user_id и PK по tg_id") + + +async def _migration_v11_finalize_legacy_tables(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v11: финализация legacy таблиц на user_id") + + for table in ("blocked_users", "manual_bans", "temporary_data"): + if not await _table_exists(conn, table): + continue + if not await _column_exists(conn, table, "user_id"): + continue + + await _backfill_users_from_table(conn, table, "tg_id") + + await conn.execute( + text( + f""" + UPDATE "{table}" AS t + SET "user_id" = u.id + FROM users AS u + WHERE t."user_id" IS NULL + AND t."tg_id" IS NOT NULL + AND t."tg_id" = u.tg_id + """ + ) + ) + + deleted = await conn.execute(text(f'DELETE FROM "{table}" WHERE "user_id" IS NULL')) + if deleted.rowcount > 0: + logger.warning( + f"[schema_upgrade] v11: удалено {deleted.rowcount} записей из {table} " + "без tg_id и user_id (невосстановимы)" + ) + + await _drop_pk(conn, table) + if await _column_exists(conn, table, "tg_id"): + await _exec_ignore(conn, f'ALTER TABLE "{table}" ALTER COLUMN "tg_id" DROP NOT NULL') + if not await _safe_set_not_null(conn, table, "user_id"): + continue + await _exec_ignore(conn, f'ALTER TABLE "{table}" ADD PRIMARY KEY ("user_id")') + await _exec_ignore(conn, f'DROP INDEX IF EXISTS "ix_{table}_user_id"') + + tg_index_name = f"ix_{table}_tg_id" + if await _column_exists(conn, table, "tg_id") and not await _index_exists(conn, table, tg_index_name): + await conn.execute(text(f'CREATE INDEX "{tg_index_name}" ON "{table}" ("tg_id")')) + + +async def _migration_v12_relax_legacy_tg_id_nullability(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v12: приведение tg_id к nullable в legacy таблицах") + + for table in ("blocked_users", "manual_bans", "temporary_data"): + if not await _table_exists(conn, table): + continue + if not await _column_exists(conn, table, "tg_id"): + continue + await _exec_ignore(conn, f'ALTER TABLE "{table}" ALTER COLUMN "tg_id" DROP NOT NULL') + + +async def _migration_v13_add_web_page_variants(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v13: добавление таблиц вариантов web-страниц") + + await _exec_ignore( + conn, + """ + CREATE TABLE IF NOT EXISTS web_page_variants ( + id VARCHAR(36) PRIMARY KEY, + page_slug VARCHAR(64) NOT NULL REFERENCES web_pages(slug) ON DELETE CASCADE, + variant_key VARCHAR(64) NOT NULL, + name VARCHAR(255) NOT NULL DEFAULT 'Default', + is_active BOOLEAN NOT NULL DEFAULT FALSE, + theme_tokens JSONB NOT NULL DEFAULT '{}'::jsonb, + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP + ) + """, + ) + await _exec_ignore( + conn, + """ + CREATE UNIQUE INDEX IF NOT EXISTS uq_web_page_variants_page_slug_variant_key + ON web_page_variants (page_slug, variant_key) + """, + ) + await _exec_ignore( + conn, + """ + CREATE INDEX IF NOT EXISTS ix_web_page_variants_page_slug_is_active + ON web_page_variants (page_slug, is_active) + """, + ) + await _exec_ignore( + conn, + """ + CREATE TABLE IF NOT EXISTS web_page_variant_blocks ( + id VARCHAR(36) PRIMARY KEY, + variant_id VARCHAR(36) NOT NULL REFERENCES web_page_variants(id) ON DELETE CASCADE, + "order" INTEGER NOT NULL DEFAULT 0, + type VARCHAR(64) NOT NULL, + data JSONB NOT NULL DEFAULT '{}'::jsonb + ) + """, + ) + await _exec_ignore( + conn, + """ + CREATE INDEX IF NOT EXISTS ix_web_page_variant_blocks_variant_id_order + ON web_page_variant_blocks (variant_id, "order") + """, + ) + + +async def _migration_v14_web_flow_graph_model(conn: AsyncConnection) -> None: + """Переход web_flows со steps[] на nodes[] + edges[] (граф-модель).""" + if not await _table_exists(conn, "web_flows"): + return + + if await _column_exists(conn, "web_flows", "steps"): + if not await _column_exists(conn, "web_flows", "nodes"): + await _exec_ignore(conn, "ALTER TABLE web_flows RENAME COLUMN steps TO nodes") + else: + await _exec_ignore(conn, "ALTER TABLE web_flows DROP COLUMN steps") + + if not await _column_exists(conn, "web_flows", "edges"): + await _exec_ignore(conn, "ALTER TABLE web_flows ADD COLUMN edges JSONB NOT NULL DEFAULT '[]'::jsonb") + + if not await _column_exists(conn, "web_flows", "entry_node_id"): + await _exec_ignore(conn, "ALTER TABLE web_flows ADD COLUMN entry_node_id VARCHAR(64)") + + rows = (await conn.execute(text("SELECT id, nodes FROM web_flows"))).all() + for row in rows: + flow_id = row[0] + raw_nodes = row[1] + if not isinstance(raw_nodes, list) or len(raw_nodes) == 0: + continue + first = raw_nodes[0] if raw_nodes else {} + if isinstance(first, dict) and "position" not in first: + new_nodes = [] + new_edges = [] + entry_id = None + for i, step in enumerate(raw_nodes): + if not isinstance(step, dict): + continue + node_id = step.get("id", f"node-{i}") + if i == 0: + entry_id = node_id + new_nodes.append({ + **step, + "position": {"x": 300, "y": i * 180}, + }) + if i > 0: + prev_id = raw_nodes[i - 1].get("id", f"node-{i - 1}") + new_edges.append({ + "id": f"edge-migrated-{i}", + "source": prev_id, + "target": node_id, + }) + import json + await conn.execute( + text("UPDATE web_flows SET nodes = :nodes, edges = :edges, entry_node_id = :entry WHERE id = :fid"), + {"nodes": json.dumps(new_nodes), "edges": json.dumps(new_edges), "entry": entry_id, "fid": flow_id}, + ) + + +async def _migration_v15_recover_orphan_users(conn: AsyncConnection) -> None: + """Safety net для клиентов, прошедших v3/v11 со старой логикой. + + Обходит все таблицы, где может быть orphan tg_id, создаёт недостающих юзеров + и повторно заполняет user_id. Идемпотентно: если orphan'ов нет — no-op. + """ + logger.info("[schema_upgrade] v15: Восстановление orphan tg_ids в users") + + if not await _table_exists(conn, "users") or not await _column_exists(conn, "users", "id"): + return + + orphan_sources = [ + ("keys", "tg_id"), + ("payments", "tg_id"), + ("notifications", "tg_id"), + ("scheduled_broadcasts", "created_by_tg_id"), + ("gift_usages", "tg_id"), + ("temporary_data", "tg_id"), + ("manual_bans", "tg_id"), + ("blocked_users", "tg_id"), + ("referrals", "referred_tg_id"), + ("referrals", "referrer_tg_id"), + ("gifts", "sender_tg_id"), + ("gifts", "recipient_tg_id"), + ] + + total_created = 0 + for table, tg_col in orphan_sources: + total_created += await _backfill_users_from_table(conn, table, tg_col) + + if total_created > 0: + logger.info(f"[schema_upgrade] v15: всего создано {total_created} orphan-юзеров") + + repopulate = [ + ("keys", "user_id", "tg_id"), + ("payments", "user_id", "tg_id"), + ("notifications", "user_id", "tg_id"), + ("scheduled_broadcasts", "created_by_user_id", "created_by_tg_id"), + ("gift_usages", "user_id", "tg_id"), + ("temporary_data", "user_id", "tg_id"), + ("manual_bans", "user_id", "tg_id"), + ("blocked_users", "user_id", "tg_id"), + ("referrals", "referred_user_id", "referred_tg_id"), + ("referrals", "referrer_user_id", "referrer_tg_id"), + ("gifts", "sender_user_id", "sender_tg_id"), + ("gifts", "recipient_user_id", "recipient_tg_id"), + ] + + for table, user_col, tg_col in repopulate: + if not await _table_exists(conn, table): + continue + if not await _column_exists(conn, table, user_col): + continue + if not await _column_exists(conn, table, tg_col): + continue + result = await conn.execute( + text( + f""" + UPDATE "{table}" t + SET "{user_col}" = u.id + FROM users u + WHERE t."{user_col}" IS NULL + AND t."{tg_col}" IS NOT NULL + AND t."{tg_col}" = u.tg_id + """ + ) + ) + if result.rowcount and result.rowcount > 0: + logger.info( + f"[schema_upgrade] v15: повторно заполнено {result.rowcount} записей " + f"{table}.{user_col}" + ) + + +async def _migration_v16b_web_flow_events(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v17: таблица web_flow_events") + + await _exec_ignore( + conn, + """ + CREATE TABLE IF NOT EXISTS web_flow_events ( + id VARCHAR(36) PRIMARY KEY, + flow_id VARCHAR(64) NOT NULL, + node_id VARCHAR(64) NOT NULL, + node_type VARCHAR(32) NOT NULL DEFAULT '', + event_type VARCHAR(32) NOT NULL, + ab_variant VARCHAR(16), + device VARCHAR(16), + locale VARCHAR(8), + authenticated BOOLEAN, + metadata JSONB, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP + ) + """, + ) + await _exec_ignore( + conn, + "CREATE INDEX IF NOT EXISTS ix_web_flow_events_flow_node ON web_flow_events (flow_id, node_id)", + ) + await _exec_ignore( + conn, + "CREATE INDEX IF NOT EXISTS ix_web_flow_events_created ON web_flow_events (created_at)", + ) + + +async def _migration_v16_custom_element_builds(conn: AsyncConnection) -> None: + logger.info("[schema_upgrade] v16: таблица web_custom_element_builds") + + await _exec_ignore( + conn, + """ + CREATE TABLE IF NOT EXISTS web_custom_element_builds ( + id VARCHAR(36) PRIMARY KEY, + label VARCHAR(255) NOT NULL DEFAULT '', + slug VARCHAR(128) NOT NULL DEFAULT '', + runtime VARCHAR(32) NOT NULL DEFAULT 'react-component', + source_kind VARCHAR(32) NOT NULL DEFAULT 'inline-code', + source_value TEXT NOT NULL DEFAULT '', + export_name VARCHAR(128) NOT NULL DEFAULT 'default', + props_schema_text TEXT NOT NULL DEFAULT '', + sample_props_text TEXT NOT NULL DEFAULT '', + events_text TEXT NOT NULL DEFAULT '', + notes TEXT NOT NULL DEFAULT '', + status VARCHAR(32) NOT NULL DEFAULT 'queued', + summary TEXT NOT NULL DEFAULT '', + next_steps JSONB NOT NULL DEFAULT '[]'::jsonb, + artifact JSONB, + upload_meta JSONB, + worker_id VARCHAR(64), + worker_claimed_at TIMESTAMPTZ, + completed_at TIMESTAMPTZ, + created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP + ) + """, + ) + await _exec_ignore( + conn, + "CREATE INDEX IF NOT EXISTS ix_web_custom_element_builds_status ON web_custom_element_builds (status)", + ) + await _exec_ignore( + conn, + "CREATE INDEX IF NOT EXISTS ix_web_custom_element_builds_created ON web_custom_element_builds (created_at)", + ) + + +_MIGRATIONS = [ + (1, "Добавление users.id", _migration_v1_add_users_id), + (2, "Добавление user_id колонок", _migration_v2_add_user_id_columns), + (3, "Заполнение user_id из tg_id", _migration_v3_populate_user_ids), + (4, "Добавление tg_id mirrors", _migration_v4_add_tg_id_mirrors), + (5, "Очистка записей с NULL user_id", _migration_v9_cleanup_orphaned_records), + (6, "Переключение PK на user_id", _migration_v5_switch_pks_to_user_id), + (7, "Добавление foreign keys", _migration_v6_add_foreign_keys), + (8, "Backfill tg_id mirrors", _migration_v7_backfill_tg_mirrors), + (9, "Исправление timezone для notifications", _migration_v8_fix_notification_timezone), + (10, "Финальная установка NOT NULL на user_id", _migration_v10_finalize_user_id_not_null), + (11, "Финализация legacy таблиц на user_id", _migration_v11_finalize_legacy_tables), + (12, "Снятие NOT NULL с tg_id в legacy таблицах", _migration_v12_relax_legacy_tg_id_nullability), + (13, "Таблицы вариантов web-страниц", _migration_v13_add_web_page_variants), + (14, "WebFlow граф-модель (nodes + edges)", _migration_v14_web_flow_graph_model), + (15, "Восстановление orphan tg_ids в users", _migration_v15_recover_orphan_users), + (16, "Таблица custom element builds", _migration_v16_custom_element_builds), + (17, "Таблица flow analytics events", _migration_v16b_web_flow_events), +] + + +async def apply_all_migrations(conn: AsyncConnection) -> None: + if not _is_postgresql(): + return + + await _ensure_migrations_table(conn) + current_version = await _get_current_version(conn) + + for version, description, migration_func in _MIGRATIONS: + if version <= current_version: + continue + + logger.info(f"[schema_upgrade] Применение миграции v{version}: {description}") + try: + await migration_func(conn) + await _mark_migration_applied(conn, version, description) + logger.info(f"[schema_upgrade] Миграция v{version} применена успешно") + except Exception as e: + logger.error(f"[schema_upgrade] Ошибка при применении миграции v{version}: {e}") + raise + + logger.info(f"[schema_upgrade] Все миграции применены, текущая версия: {await _get_current_version(conn)}") + + +async def apply_account_schema_if_needed(conn: AsyncConnection) -> None: + await apply_all_migrations(conn) + + +_TG_MIRROR_TABLE_COLUMNS = ( + ("keys", "tg_id"), + ("payments", "tg_id"), + ("referrals", "referred_tg_id"), + ("referrals", "referrer_tg_id"), + ("notifications", "tg_id"), + ("gift_usages", "tg_id"), + ("gifts", "sender_tg_id"), + ("gifts", "recipient_tg_id"), + ("manual_bans", "tg_id"), + ("temporary_data", "tg_id"), + ("blocked_users", "tg_id"), + ("scheduled_broadcasts", "created_by_tg_id"), + ("coupon_usages", "tg_id"), +) + + +async def ensure_tg_mirror_columns_and_backfill(conn: AsyncConnection) -> None: + if not _is_postgresql(): + return + for table, col in _TG_MIRROR_TABLE_COLUMNS: + if await _table_exists(conn, table) and not await _column_exists(conn, table, col): + await conn.execute(text(f'ALTER TABLE "{table}" ADD COLUMN {col} BIGINT')) + await _migration_v7_backfill_tg_mirrors(conn) diff --git a/database/models.py b/database/models.py deleted file mode 100644 index d1f039df..00000000 --- a/database/models.py +++ /dev/null @@ -1,405 +0,0 @@ -import secrets -import uuid - -from datetime import datetime - -from sqlalchemy import ( - JSON, - BigInteger, - Boolean, - Column, - DateTime, - Float, - ForeignKey, - Index, - Integer, - Numeric, - String, - Text, - UniqueConstraint, - text as sql_text, -) -from sqlalchemy.dialects.postgresql import JSONB -from sqlalchemy.orm import Mapped, declarative_base, mapped_column, relationship - - -Base = declarative_base() - - -class DictLikeMixin: - def __getitem__(self, key): - return getattr(self, key) - - def get(self, key, default=None): - return getattr(self, key, default) - - def to_dict(self): - return {column.name: getattr(self, column.name) for column in self.__table__.columns} - - -class Identity(DictLikeMixin, Base): - """Слой идентификации: к одному identity можно привязать email и/или Telegram (tg_id).""" - - __tablename__ = "identities" - - id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) - email = Column(String(255), unique=True, nullable=True, index=True) - tg_id = Column(BigInteger, unique=True, nullable=True, index=True) - api_token_hash = Column(String(64), nullable=True, index=True) - token_issued_at = Column(DateTime, nullable=True) - password_hash = Column(String(64), nullable=True) - is_admin = Column(Boolean, nullable=False, server_default=sql_text("false")) - created_at = Column(DateTime, default=datetime.utcnow) - updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow) - - -class User(DictLikeMixin, Base): - __tablename__ = "users" - - tg_id = Column(BigInteger, primary_key=True) - identity_id = Column( - String(36), - ForeignKey("identities.id", ondelete="SET NULL", onupdate="CASCADE"), - nullable=True, - index=True, - ) - username = Column(String) - first_name = Column(String) - last_name = Column(String) - language_code = Column(String) - is_bot = Column(Boolean, default=False) - balance = Column(Float, default=0.0) - trial = Column(Integer, default=0) - preferred_currency = Column(String(10), nullable=False, server_default="RUB", index=True) - source_code = Column( - String, - ForeignKey( - "tracking_sources.code", - ondelete="SET NULL", - onupdate="CASCADE", - ), - nullable=True, - ) - created_at = Column(DateTime, default=datetime.utcnow) - updated_at = Column(DateTime, default=datetime.utcnow) - - -class Key(DictLikeMixin, Base): - __tablename__ = "keys" - - tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=False, index=True) - client_id = Column(String, primary_key=True) - email = Column(String, unique=True) - created_at = Column(BigInteger) - expiry_time = Column(BigInteger) - key = Column(String) - server_id = Column(String) - remnawave_link = Column(String) - tariff_id = Column(Integer, ForeignKey("tariffs.id", ondelete="SET NULL")) - is_frozen = Column(Boolean, default=False) - alias = Column(String) - notified = Column(Boolean, default=False) - notified_24h = Column(Boolean, default=False) - - selected_device_limit = Column(Integer, nullable=True) - selected_traffic_limit = Column(BigInteger, nullable=True) - selected_price_rub = Column(Integer, nullable=True) - - current_device_limit = Column(Integer, nullable=True) - current_traffic_limit = Column(BigInteger, nullable=True) - - -class Tariff(DictLikeMixin, Base): - __tablename__ = "tariffs" - - id = Column(Integer, primary_key=True) - name = Column(String) - group_code = Column(String) - duration_days = Column(Integer) - price_rub = Column(Integer) - traffic_limit = Column(BigInteger, nullable=True) - device_limit = Column(Integer, nullable=True) - is_active = Column(Boolean, default=True) - created_at = Column(DateTime, default=datetime.utcnow) - updated_at = Column(DateTime, default=datetime.utcnow) - subgroup_title = Column(String, nullable=True) - sort_order = Column(Integer, nullable=True) - vless = Column(Boolean, default=False) - external_squad: Mapped[str | None] = mapped_column(String(64), nullable=True) - - configurable = Column(Boolean, nullable=False, server_default="false") - - device_options = Column(JSONB, nullable=True) - traffic_options_gb = Column(JSONB, nullable=True) - - device_step_rub = Column(Integer, nullable=True) - device_overrides = Column(JSONB, nullable=True) - - traffic_step_rub = Column(Integer, nullable=True) - traffic_overrides = Column(JSONB, nullable=True) - - -class Server(DictLikeMixin, Base): - __tablename__ = "servers" - - id = Column(Integer, primary_key=True, autoincrement=True) - cluster_name = Column(String) - server_name = Column(String, unique=True) - api_url = Column(String) - subscription_url = Column(String) - inbound_id = Column(String) - panel_type = Column(String) - max_keys = Column(Integer) - tariff_group = Column(String) - enabled = Column(Boolean, default=True) - - subgroups = relationship("ServerSubgroup", back_populates="server", cascade="all, delete-orphan") - groups = relationship("ServerSpecialgroup", back_populates="server", cascade="all, delete-orphan") - - -class ServerSubgroup(DictLikeMixin, Base): - __tablename__ = "server_subgroups" - - id = Column(Integer, primary_key=True, autoincrement=True) - server_id = Column(Integer, ForeignKey("servers.id", ondelete="CASCADE"), index=True, nullable=False) - group_code = Column(String, nullable=False) - subgroup_title = Column(String, nullable=False) - - server = relationship("Server", back_populates="subgroups") - - __table_args__ = (UniqueConstraint("server_id", "subgroup_title", name="uq_server_subgroup"),) - - -class ServerSpecialgroup(DictLikeMixin, Base): - __tablename__ = "server_specialgroups" - - id = Column(Integer, primary_key=True, autoincrement=True) - server_id = Column(Integer, ForeignKey("servers.id", ondelete="CASCADE"), index=True, nullable=False) - group_code = Column(String, nullable=False) - - server = relationship("Server") - - __table_args__ = (UniqueConstraint("server_id", "group_code", name="uq_server_group"),) - - -class Payment(DictLikeMixin, Base): - __tablename__ = "payments" - - id = Column(Integer, primary_key=True, autoincrement=True) - tg_id = Column(BigInteger, ForeignKey("users.tg_id")) - amount = Column(Float) - payment_system = Column(String) - status = Column(String) - created_at = Column(DateTime, default=datetime.utcnow) - original_amount = Column(Numeric(18, 8), nullable=True) - currency = Column(String(10), nullable=False, server_default="RUB") - payment_id = Column(String(128), nullable=True, index=True) - metadata_ = Column("metadata", JSONB, nullable=True) - - -class Coupon(DictLikeMixin, Base): - __tablename__ = "coupons" - - id = Column(Integer, primary_key=True) - code = Column(String, unique=True) - amount = Column(Integer) - usage_limit = Column(Integer) - usage_count = Column(Integer, default=0) - is_used = Column(Boolean, default=False) - days = Column(Integer, nullable=True) - new_users_only = Column(Boolean, nullable=False, server_default=sql_text("false")) - - percent = Column(Integer, nullable=True) - max_discount_amount = Column(Integer, nullable=True) - min_order_amount = Column(Integer, nullable=True) - - -class CouponUsage(DictLikeMixin, Base): - __tablename__ = "coupon_usages" - - coupon_id = Column(Integer, ForeignKey("coupons.id", ondelete="CASCADE"), primary_key=True) - user_id = Column(BigInteger, primary_key=True) - used_at = Column(DateTime, default=datetime.utcnow) - - -class Referral(DictLikeMixin, Base): - __tablename__ = "referrals" - - referred_tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="CASCADE"), primary_key=True) - referrer_tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="CASCADE"), primary_key=True) - reward_issued = Column(Boolean, default=False) - - -class Notification(DictLikeMixin, Base): - __tablename__ = "notifications" - - tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="CASCADE"), primary_key=True) - notification_type = Column(String, primary_key=True) - last_notification_time = Column(DateTime, default=datetime.utcnow) - - -class ScheduledBroadcast(DictLikeMixin, Base): - __tablename__ = "scheduled_broadcasts" - __table_args__ = ( - Index("ix_scheduled_broadcasts_status_time", "status", "scheduled_for"), - Index("ix_scheduled_broadcasts_creator_time", "created_by_tg_id", "created_at"), - ) - - id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) - created_by_tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="SET NULL"), nullable=True, index=True) - status = Column(String(32), nullable=False, server_default=sql_text("'scheduled'"), index=True) - send_to = Column(String(32), nullable=False, index=True) - cluster_name = Column(String, nullable=True) - text = Column(Text, nullable=False) - photo = Column(String, nullable=True) - keyboard_json = Column(JSONB, nullable=True) - scheduled_for = Column(DateTime(timezone=True), nullable=False, index=True) - workers = Column(Integer, nullable=False, server_default=sql_text("5")) - messages_per_second = Column(Integer, nullable=False, server_default=sql_text("35")) - stats_json = Column(JSONB, nullable=True) - error_text = Column(Text, nullable=True) - started_at = Column(DateTime(timezone=True), nullable=True) - sent_at = Column(DateTime(timezone=True), nullable=True) - cancelled_at = Column(DateTime(timezone=True), nullable=True) - created_at = Column(DateTime, default=datetime.utcnow) - updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow) - - -class Gift(DictLikeMixin, Base): - __tablename__ = "gifts" - - gift_id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex) - sender_tg_id = Column(BigInteger, ForeignKey("users.tg_id")) - recipient_tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=True) - selected_months = Column(Integer) - expiry_time = Column(DateTime) - gift_link = Column(String) - created_at = Column(DateTime, default=datetime.utcnow) - is_used = Column(Boolean, default=False) - is_unlimited = Column(Boolean, default=False) - max_usages = Column(Integer, nullable=True) - tariff_id: Mapped[int | None] = mapped_column(ForeignKey("tariffs.id")) - - selected_device_limit = Column(Integer, nullable=True) - selected_traffic_gb = Column(Integer, nullable=True) - selected_price_rub = Column(Integer, nullable=True) - - -class GiftUsage(DictLikeMixin, Base): - __tablename__ = "gift_usages" - - gift_id = Column(String, ForeignKey("gifts.gift_id"), primary_key=True) - tg_id = Column(BigInteger, primary_key=True) - used_at = Column(DateTime, default=datetime.utcnow) - - -class ManualBan(DictLikeMixin, Base): - __tablename__ = "manual_bans" - - tg_id = Column(BigInteger, primary_key=True) - banned_at = Column(DateTime(timezone=True), default=datetime.utcnow) - reason = Column(Text) - banned_by = Column(BigInteger) - until = Column(DateTime(timezone=True), nullable=True) - - -class TemporaryData(DictLikeMixin, Base): - __tablename__ = "temporary_data" - - tg_id = Column(BigInteger, primary_key=True) - state = Column(String) - data = Column(JSON) - updated_at = Column(DateTime, default=datetime.utcnow) - - -class BlockedUser(DictLikeMixin, Base): - __tablename__ = "blocked_users" - - tg_id = Column(BigInteger, primary_key=True) - - -class TrackingSource(DictLikeMixin, Base): - __tablename__ = "tracking_sources" - - id = Column(Integer, primary_key=True) - name = Column(String) - code = Column(String, unique=True) - type = Column(String) - created_by = Column(BigInteger) - created_at = Column(DateTime, default=datetime.utcnow) - - -class AuditEvent(DictLikeMixin, Base): - """События аудита (флоу пользователя).""" - __tablename__ = "audit_events" - __table_args__ = ( - Index("ix_audit_events_tg_created", "actor_tg_id", "created_at"), - Index("ix_audit_events_identity_created", "actor_identity_id", "created_at"), - ) - - id = Column(Integer, primary_key=True, autoincrement=True) - event_type = Column(String(64), nullable=False, index=True) - channel = Column(String(32), nullable=False, index=True) - actor_identity_id = Column( - String(36), - ForeignKey("identities.id", ondelete="SET NULL", onupdate="CASCADE"), - nullable=True, - index=True, - ) - actor_tg_id = Column(BigInteger, nullable=True, index=True) - path_or_handler = Column(String(255), nullable=False) - entity_type = Column(String(64), nullable=True, index=True) - entity_id = Column(String(255), nullable=True, index=True) - result = Column(String(32), nullable=False, server_default=sql_text("'success'")) - reason = Column(Text, nullable=True) - metadata_ = Column("metadata", JSONB, nullable=True) - request_id = Column(String(64), nullable=True, index=True) - created_at = Column(DateTime, default=datetime.utcnow, index=True) - - -class Admin(Base): - __tablename__ = "admins" - - tg_id = Column(BigInteger, primary_key=True) - token = Column(String, unique=True, nullable=True) - description = Column(String, nullable=True) - role = Column(String, nullable=False, default="admin") - added_at = Column(DateTime, default=datetime.utcnow) - - @staticmethod - def generate_token() -> str: - return secrets.token_urlsafe(32) - - -class Setting(DictLikeMixin, Base): - __tablename__ = "settings" - - key = Column(String, primary_key=True) - value = Column(JSONB, nullable=True) - description = Column(Text, nullable=True) - created_at = Column(DateTime, default=datetime.utcnow) - updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow) - - -class WebPage(DictLikeMixin, Base): - __tablename__ = "web_pages" - - slug = Column(String(64), primary_key=True) - title = Column(String(255), nullable=True) - - -class WebTheme(DictLikeMixin, Base): - __tablename__ = "web_themes" - - page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), primary_key=True) - tokens = Column(JSONB, nullable=False, default=dict) - - -class WebBlock(DictLikeMixin, Base): - __tablename__ = "web_blocks" - - id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) - page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), index=True, nullable=False) - order = Column(Integer, nullable=False, default=0) - type = Column(String(64), nullable=False) - data = Column(JSONB, nullable=False, default=dict) diff --git a/database/models/__init__.py b/database/models/__init__.py new file mode 100644 index 00000000..aa351772 --- /dev/null +++ b/database/models/__init__.py @@ -0,0 +1,61 @@ +from ._base import Base, DictLikeMixin +from .admin import Admin, Setting +from .audit import AuditEvent +from .coupons import Coupon, CouponUsage +from .gifts import Gift, GiftUsage +from .identity import Identity +from .keys import Key +from .notifications import Notification, ScheduledBroadcast +from .payments import Payment +from .referrals import Referral +from .servers import Server, ServerSpecialgroup, ServerSubgroup +from .tariffs import Tariff +from .users import BlockedUser, ManualBan, TemporaryData, TrackingSource, User +from .web import ( + WebBlock, + WebCustomElementBuild, + WebFlow, + WebFlowEvent, + WebNotification, + WebPage, + WebPageVariant, + WebPageVariantBlock, + WebPushSubscription, + WebTheme, +) + + +__all__ = [ + "Base", + "DictLikeMixin", + "Identity", + "User", + "ManualBan", + "TemporaryData", + "BlockedUser", + "TrackingSource", + "Key", + "Tariff", + "Server", + "ServerSubgroup", + "ServerSpecialgroup", + "Payment", + "Coupon", + "CouponUsage", + "Referral", + "Notification", + "ScheduledBroadcast", + "Gift", + "GiftUsage", + "AuditEvent", + "Admin", + "Setting", + "WebPage", + "WebTheme", + "WebBlock", + "WebPageVariant", + "WebPageVariantBlock", + "WebPushSubscription", + "WebNotification", + "WebFlow", +] diff --git a/database/models/_base.py b/database/models/_base.py new file mode 100644 index 00000000..0e5e2aca --- /dev/null +++ b/database/models/_base.py @@ -0,0 +1,21 @@ +from sqlalchemy.orm import declarative_base + + +Base = declarative_base() + + +class DictLikeMixin: + """Позволяет обращаться к ORM-объектам как к словарю. + + Используется legacy-кодом, который мигрировал с dict-результатов asyncpg + на ORM и не хочет переписывать все `row["field"]` / `row.get("field")`. + """ + + def __getitem__(self, key): + return getattr(self, key) + + def get(self, key, default=None): + return getattr(self, key, default) + + def to_dict(self): + return {column.name: getattr(self, column.name) for column in self.__table__.columns} diff --git a/database/models/admin.py b/database/models/admin.py new file mode 100644 index 00000000..05b12345 --- /dev/null +++ b/database/models/admin.py @@ -0,0 +1,32 @@ +import secrets + +from datetime import datetime + +from sqlalchemy import BigInteger, Column, DateTime, String, Text +from sqlalchemy.dialects.postgresql import JSONB + +from ._base import Base, DictLikeMixin + + +class Admin(Base): + __tablename__ = "admins" + + tg_id = Column(BigInteger, primary_key=True) + token = Column(String, unique=True, nullable=True) + description = Column(String, nullable=True) + role = Column(String, nullable=False, default="admin") + added_at = Column(DateTime, default=datetime.utcnow) + + @staticmethod + def generate_token() -> str: + return secrets.token_urlsafe(32) + + +class Setting(DictLikeMixin, Base): + __tablename__ = "settings" + + key = Column(String, primary_key=True) + value = Column(JSONB, nullable=True) + description = Column(Text, nullable=True) + created_at = Column(DateTime, default=datetime.utcnow) + updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow) diff --git a/database/models/audit.py b/database/models/audit.py new file mode 100644 index 00000000..82f7cfe2 --- /dev/null +++ b/database/models/audit.py @@ -0,0 +1,44 @@ +from datetime import datetime + +from sqlalchemy import ( + BigInteger, + Column, + DateTime, + ForeignKey, + Index, + Integer, + String, + Text, + text as sql_text, +) +from sqlalchemy.dialects.postgresql import JSONB + +from ._base import Base, DictLikeMixin + + +class AuditEvent(DictLikeMixin, Base): + """События аудита (флоу пользователя).""" + __tablename__ = "audit_events" + __table_args__ = ( + Index("ix_audit_events_tg_created", "actor_tg_id", "created_at"), + Index("ix_audit_events_identity_created", "actor_identity_id", "created_at"), + ) + + id = Column(Integer, primary_key=True, autoincrement=True) + event_type = Column(String(64), nullable=False, index=True) + channel = Column(String(32), nullable=False, index=True) + actor_identity_id = Column( + String(36), + ForeignKey("identities.id", ondelete="SET NULL", onupdate="CASCADE"), + nullable=True, + index=True, + ) + actor_tg_id = Column(BigInteger, nullable=True, index=True) + path_or_handler = Column(String(255), nullable=False) + entity_type = Column(String(64), nullable=True, index=True) + entity_id = Column(String(255), nullable=True, index=True) + result = Column(String(32), nullable=False, server_default=sql_text("'success'")) + reason = Column(Text, nullable=True) + metadata_ = Column("metadata", JSONB, nullable=True) + request_id = Column(String(64), nullable=True, index=True) + created_at = Column(DateTime, default=datetime.utcnow, index=True) diff --git a/database/models/coupons.py b/database/models/coupons.py new file mode 100644 index 00000000..2d6cbc01 --- /dev/null +++ b/database/models/coupons.py @@ -0,0 +1,40 @@ +from datetime import datetime + +from sqlalchemy import ( + BigInteger, + Boolean, + Column, + DateTime, + ForeignKey, + Integer, + String, + text as sql_text, +) + +from ._base import Base, DictLikeMixin + + +class Coupon(DictLikeMixin, Base): + __tablename__ = "coupons" + + id = Column(Integer, primary_key=True) + code = Column(String, unique=True) + amount = Column(Integer) + usage_limit = Column(Integer) + usage_count = Column(Integer, default=0) + is_used = Column(Boolean, default=False) + days = Column(Integer, nullable=True) + new_users_only = Column(Boolean, nullable=False, server_default=sql_text("false")) + + percent = Column(Integer, nullable=True) + max_discount_amount = Column(Integer, nullable=True) + min_order_amount = Column(Integer, nullable=True) + + +class CouponUsage(DictLikeMixin, Base): + __tablename__ = "coupon_usages" + + coupon_id = Column(Integer, ForeignKey("coupons.id", ondelete="CASCADE"), primary_key=True) + user_id = Column(BigInteger, primary_key=True) + tg_id = Column(BigInteger, nullable=True, index=True) + used_at = Column(DateTime, default=datetime.utcnow) diff --git a/database/models/gifts.py b/database/models/gifts.py new file mode 100644 index 00000000..99b1228e --- /dev/null +++ b/database/models/gifts.py @@ -0,0 +1,39 @@ +import uuid + +from datetime import datetime + +from sqlalchemy import BigInteger, Boolean, Column, DateTime, ForeignKey, Integer, String +from sqlalchemy.orm import Mapped, mapped_column + +from ._base import Base, DictLikeMixin + + +class Gift(DictLikeMixin, Base): + __tablename__ = "gifts" + + gift_id = Column(String, primary_key=True, default=lambda: uuid.uuid4().hex) + sender_user_id = Column(BigInteger, nullable=True) + recipient_user_id = Column(BigInteger, nullable=True) + sender_tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=True, index=True) + recipient_tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=True, index=True) + selected_months = Column(Integer) + expiry_time = Column(DateTime) + gift_link = Column(String) + created_at = Column(DateTime, default=datetime.utcnow) + is_used = Column(Boolean, default=False) + is_unlimited = Column(Boolean, default=False) + max_usages = Column(Integer, nullable=True) + tariff_id: Mapped[int | None] = mapped_column(ForeignKey("tariffs.id")) + + selected_device_limit = Column(Integer, nullable=True) + selected_traffic_gb = Column(Integer, nullable=True) + selected_price_rub = Column(Integer, nullable=True) + + +class GiftUsage(DictLikeMixin, Base): + __tablename__ = "gift_usages" + + gift_id = Column(String, ForeignKey("gifts.gift_id"), primary_key=True) + user_id = Column(BigInteger, nullable=False, primary_key=True) + tg_id = Column(BigInteger, nullable=True, index=True) + used_at = Column(DateTime, default=datetime.utcnow) diff --git a/database/models/identity.py b/database/models/identity.py new file mode 100644 index 00000000..c3eb770d --- /dev/null +++ b/database/models/identity.py @@ -0,0 +1,31 @@ +import uuid + +from datetime import datetime + +from sqlalchemy import ( + BigInteger, + Boolean, + Column, + DateTime, + String, + text as sql_text, +) + +from ._base import Base, DictLikeMixin + + +class Identity(DictLikeMixin, Base): + """Слой идентификации: к одному identity можно привязать email и/или Telegram (tg_id).""" + + __tablename__ = "identities" + + id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) + email = Column(String(255), unique=True, nullable=True, index=True) + tg_id = Column(BigInteger, unique=True, nullable=True, index=True) + api_token_hash = Column(String(64), nullable=True, index=True) + token_issued_at = Column(DateTime, nullable=True) + password_hash = Column(String(64), nullable=True) + email_verified = Column(Boolean, nullable=False, server_default=sql_text("false")) + is_admin = Column(Boolean, nullable=False, server_default=sql_text("false")) + created_at = Column(DateTime, default=datetime.utcnow) + updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow) diff --git a/database/models/keys.py b/database/models/keys.py new file mode 100644 index 00000000..0b84f355 --- /dev/null +++ b/database/models/keys.py @@ -0,0 +1,29 @@ +from sqlalchemy import BigInteger, Boolean, Column, ForeignKey, Integer, String + +from ._base import Base, DictLikeMixin + + +class Key(DictLikeMixin, Base): + __tablename__ = "keys" + + tg_id = Column(BigInteger, ForeignKey("users.tg_id"), primary_key=True, nullable=False, index=True) + user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=True, index=True) + client_id = Column(String, primary_key=True) + email = Column(String, unique=True) + created_at = Column(BigInteger) + expiry_time = Column(BigInteger) + key = Column(String) + server_id = Column(String) + remnawave_link = Column(String) + tariff_id = Column(Integer, ForeignKey("tariffs.id", ondelete="SET NULL")) + is_frozen = Column(Boolean, default=False) + alias = Column(String) + notified = Column(Boolean, default=False) + notified_24h = Column(Boolean, default=False) + + selected_device_limit = Column(Integer, nullable=True) + selected_traffic_limit = Column(BigInteger, nullable=True) + selected_price_rub = Column(Integer, nullable=True) + + current_device_limit = Column(Integer, nullable=True) + current_traffic_limit = Column(BigInteger, nullable=True) diff --git a/database/models/notifications.py b/database/models/notifications.py new file mode 100644 index 00000000..db73363d --- /dev/null +++ b/database/models/notifications.py @@ -0,0 +1,55 @@ +import uuid + +from datetime import UTC, datetime + +from sqlalchemy import ( + BigInteger, + Column, + DateTime, + ForeignKey, + Index, + Integer, + String, + Text, + text as sql_text, +) +from sqlalchemy.dialects.postgresql import JSONB + +from ._base import Base, DictLikeMixin + + +class Notification(DictLikeMixin, Base): + __tablename__ = "notifications" + + tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="CASCADE"), nullable=True, index=True) + user_id = Column(BigInteger, nullable=False, primary_key=True) + notification_type = Column(String, primary_key=True) + last_notification_time = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) + + +class ScheduledBroadcast(DictLikeMixin, Base): + __tablename__ = "scheduled_broadcasts" + __table_args__ = ( + Index("ix_scheduled_broadcasts_status_time", "status", "scheduled_for"), + Index("ix_scheduled_broadcasts_creator_time", "created_by_tg_id", "created_at"), + ) + + id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) + created_by_user_id = Column(BigInteger, nullable=True, index=True) + created_by_tg_id = Column(BigInteger, ForeignKey("users.tg_id", ondelete="SET NULL"), nullable=True, index=True) + status = Column(String(32), nullable=False, server_default=sql_text("'scheduled'"), index=True) + send_to = Column(String(32), nullable=False, index=True) + cluster_name = Column(String, nullable=True) + text = Column(Text, nullable=False) + photo = Column(String, nullable=True) + keyboard_json = Column(JSONB, nullable=True) + scheduled_for = Column(DateTime(timezone=True), nullable=False, index=True) + workers = Column(Integer, nullable=False, server_default=sql_text("5")) + messages_per_second = Column(Integer, nullable=False, server_default=sql_text("35")) + stats_json = Column(JSONB, nullable=True) + error_text = Column(Text, nullable=True) + started_at = Column(DateTime(timezone=True), nullable=True) + sent_at = Column(DateTime(timezone=True), nullable=True) + cancelled_at = Column(DateTime(timezone=True), nullable=True) + created_at = Column(DateTime, default=datetime.utcnow) + updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow) diff --git a/database/models/payments.py b/database/models/payments.py new file mode 100644 index 00000000..24cd305e --- /dev/null +++ b/database/models/payments.py @@ -0,0 +1,22 @@ +from datetime import datetime + +from sqlalchemy import BigInteger, Column, DateTime, Float, ForeignKey, Integer, Numeric, String +from sqlalchemy.dialects.postgresql import JSONB + +from ._base import Base, DictLikeMixin + + +class Payment(DictLikeMixin, Base): + __tablename__ = "payments" + + id = Column(Integer, primary_key=True, autoincrement=True) + user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=True, index=True) + tg_id = Column(BigInteger, ForeignKey("users.tg_id"), nullable=True, index=True) + amount = Column(Float) + payment_system = Column(String) + status = Column(String) + created_at = Column(DateTime, default=datetime.utcnow) + original_amount = Column(Numeric(18, 8), nullable=True) + currency = Column(String(10), nullable=False, server_default="RUB") + payment_id = Column(String(128), nullable=True, index=True) + metadata_ = Column("metadata", JSONB, nullable=True) diff --git a/database/models/referrals.py b/database/models/referrals.py new file mode 100644 index 00000000..34aabad7 --- /dev/null +++ b/database/models/referrals.py @@ -0,0 +1,13 @@ +from sqlalchemy import BigInteger, Boolean, Column, ForeignKey + +from ._base import Base, DictLikeMixin + + +class Referral(DictLikeMixin, Base): + __tablename__ = "referrals" + + referred_user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), primary_key=True) + referrer_user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), primary_key=True) + referred_tg_id = Column(BigInteger, nullable=True, index=True) + referrer_tg_id = Column(BigInteger, nullable=True, index=True) + reward_issued = Column(Boolean, default=False) diff --git a/database/models/servers.py b/database/models/servers.py new file mode 100644 index 00000000..3af444ac --- /dev/null +++ b/database/models/servers.py @@ -0,0 +1,47 @@ +from sqlalchemy import Boolean, Column, ForeignKey, Integer, String, UniqueConstraint +from sqlalchemy.orm import relationship + +from ._base import Base, DictLikeMixin + + +class Server(DictLikeMixin, Base): + __tablename__ = "servers" + + id = Column(Integer, primary_key=True, autoincrement=True) + cluster_name = Column(String) + server_name = Column(String, unique=True) + api_url = Column(String) + subscription_url = Column(String) + inbound_id = Column(String) + panel_type = Column(String) + max_keys = Column(Integer) + tariff_group = Column(String) + enabled = Column(Boolean, default=True) + + subgroups = relationship("ServerSubgroup", back_populates="server", cascade="all, delete-orphan") + groups = relationship("ServerSpecialgroup", back_populates="server", cascade="all, delete-orphan") + + +class ServerSubgroup(DictLikeMixin, Base): + __tablename__ = "server_subgroups" + + id = Column(Integer, primary_key=True, autoincrement=True) + server_id = Column(Integer, ForeignKey("servers.id", ondelete="CASCADE"), index=True, nullable=False) + group_code = Column(String, nullable=False) + subgroup_title = Column(String, nullable=False) + + server = relationship("Server", back_populates="subgroups") + + __table_args__ = (UniqueConstraint("server_id", "subgroup_title", name="uq_server_subgroup"),) + + +class ServerSpecialgroup(DictLikeMixin, Base): + __tablename__ = "server_specialgroups" + + id = Column(Integer, primary_key=True, autoincrement=True) + server_id = Column(Integer, ForeignKey("servers.id", ondelete="CASCADE"), index=True, nullable=False) + group_code = Column(String, nullable=False) + + server = relationship("Server") + + __table_args__ = (UniqueConstraint("server_id", "group_code", name="uq_server_group"),) diff --git a/database/models/tariffs.py b/database/models/tariffs.py new file mode 100644 index 00000000..9f8093ce --- /dev/null +++ b/database/models/tariffs.py @@ -0,0 +1,37 @@ +from datetime import datetime + +from sqlalchemy import BigInteger, Boolean, Column, DateTime, Integer, String +from sqlalchemy.dialects.postgresql import JSONB +from sqlalchemy.orm import Mapped, mapped_column + +from ._base import Base, DictLikeMixin + + +class Tariff(DictLikeMixin, Base): + __tablename__ = "tariffs" + + id = Column(Integer, primary_key=True) + name = Column(String) + group_code = Column(String) + duration_days = Column(Integer) + price_rub = Column(Integer) + traffic_limit = Column(BigInteger, nullable=True) + device_limit = Column(Integer, nullable=True) + is_active = Column(Boolean, default=True) + created_at = Column(DateTime, default=datetime.utcnow) + updated_at = Column(DateTime, default=datetime.utcnow) + subgroup_title = Column(String, nullable=True) + sort_order = Column(Integer, nullable=True) + vless = Column(Boolean, default=False) + external_squad: Mapped[str | None] = mapped_column(String(64), nullable=True) + + configurable = Column(Boolean, nullable=False, server_default="false") + + device_options = Column(JSONB, nullable=True) + traffic_options_gb = Column(JSONB, nullable=True) + + device_step_rub = Column(Integer, nullable=True) + device_overrides = Column(JSONB, nullable=True) + + traffic_step_rub = Column(Integer, nullable=True) + traffic_overrides = Column(JSONB, nullable=True) diff --git a/database/models/users.py b/database/models/users.py new file mode 100644 index 00000000..a6b400cd --- /dev/null +++ b/database/models/users.py @@ -0,0 +1,88 @@ +from datetime import datetime + +from sqlalchemy import ( + JSON, + BigInteger, + Boolean, + Column, + DateTime, + Float, + ForeignKey, + Identity as SAIdentity, + Integer, + String, + Text, +) + +from ._base import Base, DictLikeMixin + + +class User(DictLikeMixin, Base): + __tablename__ = "users" + + id = Column(BigInteger, SAIdentity(always=False), primary_key=True) + tg_id = Column(BigInteger, nullable=True, unique=True, index=True) + identity_id = Column( + String(36), + ForeignKey("identities.id", ondelete="SET NULL", onupdate="CASCADE"), + nullable=True, + index=True, + ) + username = Column(String) + first_name = Column(String) + last_name = Column(String) + language_code = Column(String) + is_bot = Column(Boolean, default=False) + balance = Column(Float, default=0.0) + trial = Column(Integer, default=0) + preferred_currency = Column(String(10), nullable=False, server_default="RUB", index=True) + source_code = Column( + String, + ForeignKey( + "tracking_sources.code", + ondelete="SET NULL", + onupdate="CASCADE", + ), + nullable=True, + ) + created_at = Column(DateTime, default=datetime.utcnow) + updated_at = Column(DateTime, default=datetime.utcnow) + + +class ManualBan(DictLikeMixin, Base): + __tablename__ = "manual_bans" + + user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=False, primary_key=True) + tg_id = Column(BigInteger, nullable=True, index=True) + banned_at = Column(DateTime(timezone=True), default=datetime.utcnow) + reason = Column(Text) + banned_by = Column(BigInteger) + until = Column(DateTime(timezone=True), nullable=True) + + +class TemporaryData(DictLikeMixin, Base): + __tablename__ = "temporary_data" + + user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=False, primary_key=True) + tg_id = Column(BigInteger, nullable=True, index=True) + state = Column(String) + data = Column(JSON) + updated_at = Column(DateTime, default=datetime.utcnow) + + +class BlockedUser(DictLikeMixin, Base): + __tablename__ = "blocked_users" + + user_id = Column(BigInteger, ForeignKey("users.id", ondelete="CASCADE"), nullable=False, primary_key=True) + tg_id = Column(BigInteger, nullable=True, index=True) + + +class TrackingSource(DictLikeMixin, Base): + __tablename__ = "tracking_sources" + + id = Column(Integer, primary_key=True) + name = Column(String) + code = Column(String, unique=True) + type = Column(String) + created_by = Column(BigInteger) + created_at = Column(DateTime, default=datetime.utcnow) diff --git a/database/models/web.py b/database/models/web.py new file mode 100644 index 00000000..3a3fa32b --- /dev/null +++ b/database/models/web.py @@ -0,0 +1,166 @@ +import uuid + +from datetime import UTC, datetime + +from sqlalchemy import ( + BigInteger, + Boolean, + Column, + DateTime, + ForeignKey, + Index, + Integer, + String, + Text, + UniqueConstraint, + text as sql_text, +) +from sqlalchemy.dialects.postgresql import JSONB + +from ._base import Base, DictLikeMixin + + +class WebPage(DictLikeMixin, Base): + __tablename__ = "web_pages" + + slug = Column(String(64), primary_key=True) + title = Column(String(255), nullable=True) + + +class WebTheme(DictLikeMixin, Base): + __tablename__ = "web_themes" + + page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), primary_key=True) + tokens = Column(JSONB, nullable=False, default=dict) + + +class WebBlock(DictLikeMixin, Base): + __tablename__ = "web_blocks" + + id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) + page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), index=True, nullable=False) + order = Column(Integer, nullable=False, default=0) + type = Column(String(64), nullable=False) + data = Column(JSONB, nullable=False, default=dict) + + +class WebPageVariant(DictLikeMixin, Base): + __tablename__ = "web_page_variants" + __table_args__ = ( + UniqueConstraint("page_slug", "variant_key", name="uq_web_page_variants_page_slug_variant_key"), + Index("ix_web_page_variants_page_slug_is_active", "page_slug", "is_active"), + ) + + id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) + page_slug = Column(String(64), ForeignKey("web_pages.slug", ondelete="CASCADE"), index=True, nullable=False) + variant_key = Column(String(64), nullable=False) + name = Column(String(255), nullable=False, default="Default") + is_active = Column(Boolean, nullable=False, server_default=sql_text("false")) + theme_tokens = Column(JSONB, nullable=False, default=dict) + created_at = Column(DateTime, default=datetime.utcnow) + updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow) + + +class WebPageVariantBlock(DictLikeMixin, Base): + __tablename__ = "web_page_variant_blocks" + __table_args__ = (Index("ix_web_page_variant_blocks_variant_id_order", "variant_id", "order"),) + + id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) + variant_id = Column(String(36), ForeignKey("web_page_variants.id", ondelete="CASCADE"), index=True, nullable=False) + order = Column(Integer, nullable=False, default=0) + type = Column(String(64), nullable=False) + data = Column(JSONB, nullable=False, default=dict) + + +class WebPushSubscription(DictLikeMixin, Base): + __tablename__ = "web_push_subscriptions" + __table_args__ = ( + Index("ix_web_push_subscriptions_user_id", "user_id"), + ) + + id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) + user_id = Column(BigInteger, nullable=False) + identity_id = Column(String(36), nullable=True, index=True) + endpoint = Column(Text, nullable=False, unique=True) + keys_json = Column(JSONB, nullable=False) + created_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) + + +class WebNotification(DictLikeMixin, Base): + __tablename__ = "web_notifications" + __table_args__ = ( + Index("ix_web_notifications_user_read", "user_id", "read"), + Index("ix_web_notifications_created", "created_at"), + ) + + id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) + user_id = Column(BigInteger, nullable=False, index=True) + identity_id = Column(String(36), nullable=True, index=True) + type = Column(String(32), nullable=False, default="system") + title = Column(String(255), nullable=False) + message = Column(Text, nullable=False, default="") + read = Column(Boolean, nullable=False, default=False) + data = Column(JSONB, nullable=True) + created_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) + + +class WebFlowEvent(DictLikeMixin, Base): + __tablename__ = "web_flow_events" + __table_args__ = ( + Index("ix_web_flow_events_flow_node", "flow_id", "node_id"), + Index("ix_web_flow_events_created", "created_at"), + ) + + id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) + flow_id = Column(String(64), nullable=False) + node_id = Column(String(64), nullable=False) + node_type = Column(String(32), nullable=False, default="") + event_type = Column(String(32), nullable=False) + ab_variant = Column(String(16), nullable=True) + device = Column(String(16), nullable=True) + locale = Column(String(8), nullable=True) + authenticated = Column(Boolean, nullable=True) + event_metadata = Column("metadata", JSONB, nullable=True) + created_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) + + +class WebCustomElementBuild(DictLikeMixin, Base): + __tablename__ = "web_custom_element_builds" + __table_args__ = ( + Index("ix_web_custom_element_builds_status", "status"), + Index("ix_web_custom_element_builds_created", "created_at"), + ) + + id = Column(String(36), primary_key=True, default=lambda: str(uuid.uuid4())) + label = Column(String(255), nullable=False, default="") + slug = Column(String(128), nullable=False, default="") + runtime = Column(String(32), nullable=False, default="react-component") + source_kind = Column(String(32), nullable=False, default="inline-code") + source_value = Column(Text, nullable=False, default="") + export_name = Column(String(128), nullable=False, default="default") + props_schema_text = Column(Text, nullable=False, default="") + sample_props_text = Column(Text, nullable=False, default="") + events_text = Column(Text, nullable=False, default="") + notes = Column(Text, nullable=False, default="") + status = Column(String(32), nullable=False, default="queued") + summary = Column(Text, nullable=False, default="") + next_steps = Column(JSONB, nullable=False, default=list) + artifact = Column(JSONB, nullable=True) + upload_meta = Column(JSONB, nullable=True) + worker_id = Column(String(64), nullable=True) + worker_claimed_at = Column(DateTime(timezone=True), nullable=True) + completed_at = Column(DateTime(timezone=True), nullable=True) + created_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC)) + updated_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC), onupdate=lambda: datetime.now(UTC)) + + +class WebFlow(DictLikeMixin, Base): + __tablename__ = "web_flows" + + id = Column(String(64), primary_key=True, default="default") + name = Column(String(255), nullable=False, default="Основной flow") + nodes = Column(JSONB, nullable=False, default=list) + edges = Column(JSONB, nullable=False, default=list) + entry_node_id = Column(String(64), nullable=True) + version = Column(Integer, nullable=False, default=1) + updated_at = Column(DateTime(timezone=True), default=lambda: datetime.now(UTC), onupdate=lambda: datetime.now(UTC)) diff --git a/database/notifications.py b/database/notifications.py index f12389b9..031648b0 100644 --- a/database/notifications.py +++ b/database/notifications.py @@ -1,111 +1,155 @@ from collections import defaultdict -from datetime import datetime, timedelta +from datetime import UTC, datetime, timedelta from sqlalchemy import and_, delete, func, select, tuple_ from sqlalchemy.dialects.postgresql import insert -from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from config import DISCOUNT_ACTIVE_HOURS from core.bootstrap import NOTIFICATIONS_CONFIG -from database.models import Key, Notification, User +from database.models import BlockedUser, Key, Notification, User +from database.access.resolution import resolve_user_optional from logger import logger _NOTIFICATION_TIME_BATCH_SIZE = 300 _BULK_ADD_NOTIFICATIONS_BATCH_SIZE = 1000 -async def add_notification(session: AsyncSession, tg_id: int, notification_type: str): - try: - stmt = ( - insert(Notification) - .values( - tg_id=tg_id, - notification_type=notification_type, - last_notification_time=datetime.utcnow(), - ) - .on_conflict_do_update( - index_elements=[Notification.tg_id, Notification.notification_type], - set_={"last_notification_time": datetime.utcnow()}, - ) - ) - await session.execute(stmt) - await session.commit() - logger.info(f"✅ Добавлено уведомление {notification_type} для пользователя {tg_id}") - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при добавлении уведомления: {e}") - await session.rollback() - raise + +def _utc_now() -> datetime: + return datetime.now(UTC) -async def delete_notification(session: AsyncSession, tg_id: int, notification_type: str): +def _as_utc(value: datetime | None) -> datetime | None: + if value is None: + return None + if value.tzinfo is None: + return value.replace(tzinfo=UTC) + return value.astimezone(UTC) + + +async def _map_legacy_refs_to_user_ids(session: AsyncSession, refs: list[int]) -> dict[int, int]: + if not refs: + return {} + from sqlalchemy import or_ + + uniq = list(dict.fromkeys(refs)) + r = await session.execute(select(User.id, User.tg_id).where(or_(User.tg_id.in_(uniq), User.id.in_(uniq)))) + m: dict[int, int] = {} + for uid, tgid in r.all(): + m[int(uid)] = int(uid) + if tgid is not None: + m[int(tgid)] = int(uid) + return {ref: m[ref] for ref in uniq if ref in m} + + +async def add_notification(session: AsyncSession, legacy_user_ref: int, notification_type: str): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return + ins = insert(Notification).values( + user_id=u.id, + tg_id=u.tg_id, + notification_type=notification_type, + last_notification_time=_utc_now(), + ) + stmt = ins.on_conflict_do_update( + index_elements=[Notification.user_id, Notification.notification_type], + set_={ + "last_notification_time": ins.excluded.last_notification_time, + "tg_id": ins.excluded.tg_id, + }, + ) + await session.execute(stmt) + logger.info(f"✅ Добавлено уведомление {notification_type} для пользователя {u.id}") + + +async def delete_notification(session: AsyncSession, legacy_user_ref: int, notification_type: str): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return + uid = u.id await session.execute( delete(Notification).where( - Notification.tg_id == tg_id, + Notification.user_id == uid, Notification.notification_type == notification_type, ) ) - await session.commit() - logger.debug(f"🗑 Уведомление {notification_type} для пользователя {tg_id} удалено") + logger.debug(f"🗑 Уведомление {notification_type} для пользователя {uid} удалено") -async def bulk_add_notifications( - session: AsyncSession, items: list[tuple[int, str]], *, commit: bool = False -) -> None: - """Вставка/обновление многих (tg_id, notification_type) батчами (лимит параметров PostgreSQL). Без commit, если commit=False.""" +async def bulk_add_notifications(session: AsyncSession, items: list[tuple[int, str]]) -> None: + """Вставка/обновление многих (legacy_user_ref, notification_type) батчами (лимит параметров PostgreSQL).""" if not items: return - now = datetime.utcnow() + id_map = await _map_legacy_refs_to_user_ids(session, [p[0] for p in items]) + mapped = [(id_map[r], n) for r, n in items if r in id_map] + if not mapped: + return + uids = list({uid for uid, _ in mapped}) + tg_map_r = await session.execute(select(User.id, User.tg_id).where(User.id.in_(uids))) + tg_by_uid = {int(r.id): r.tg_id for r in tg_map_r.all()} + now = _utc_now() total = 0 - for i in range(0, len(items), _BULK_ADD_NOTIFICATIONS_BATCH_SIZE): - batch = items[i : i + _BULK_ADD_NOTIFICATIONS_BATCH_SIZE] - stmt = insert(Notification).values( + for i in range(0, len(mapped), _BULK_ADD_NOTIFICATIONS_BATCH_SIZE): + batch = mapped[i : i + _BULK_ADD_NOTIFICATIONS_BATCH_SIZE] + ins = insert(Notification).values( [ - {"tg_id": tg_id, "notification_type": ntype, "last_notification_time": now} - for tg_id, ntype in batch + { + "user_id": uid, + "tg_id": tg_by_uid.get(uid), + "notification_type": ntype, + "last_notification_time": now, + } + for uid, ntype in batch ] - ).on_conflict_do_update( - index_elements=[Notification.tg_id, Notification.notification_type], - set_={"last_notification_time": now}, + ) + stmt = ins.on_conflict_do_update( + index_elements=[Notification.user_id, Notification.notification_type], + set_={ + "last_notification_time": ins.excluded.last_notification_time, + "tg_id": ins.excluded.tg_id, + }, ) await session.execute(stmt) total += len(batch) - if commit: - await session.commit() logger.info(f"✅ Bulk: добавлено/обновлено {total} уведомлений") INACTIVE_TRIAL_REGISTERED_TYPE = "inactive_trial_registered" -async def bulk_delete_notifications( - session: AsyncSession, items: list[tuple[int, str]], *, commit: bool = False -) -> None: - """Удаление многих (tg_id, notification_type) батчами (лимит параметров PostgreSQL). Без commit, если commit=False.""" +async def bulk_delete_notifications(session: AsyncSession, items: list[tuple[int, str]]) -> None: + """Удаление многих (legacy_user_ref, notification_type) батчами (лимит параметров PostgreSQL).""" if not items: return + id_map = await _map_legacy_refs_to_user_ids(session, [p[0] for p in items]) + mapped = [(id_map[r], n) for r, n in items if r in id_map] + if not mapped: + return total = 0 - for i in range(0, len(items), _BULK_ADD_NOTIFICATIONS_BATCH_SIZE): - batch = items[i : i + _BULK_ADD_NOTIFICATIONS_BATCH_SIZE] + for i in range(0, len(mapped), _BULK_ADD_NOTIFICATIONS_BATCH_SIZE): + batch = mapped[i : i + _BULK_ADD_NOTIFICATIONS_BATCH_SIZE] stmt = delete(Notification).where( - tuple_(Notification.tg_id, Notification.notification_type).in_(batch) + tuple_(Notification.user_id, Notification.notification_type).in_(batch) ) await session.execute(stmt) total += len(batch) - if commit: - await session.commit() logger.debug(f"🗑 Bulk: удалено {total} уведомлений") -async def check_notification_time(session: AsyncSession, tg_id: int, notification_type: str, hours: int = 12) -> bool: +async def check_notification_time(session: AsyncSession, legacy_user_ref: int, notification_type: str, hours: int = 12) -> bool: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return True stmt = select(Notification.last_notification_time).where( - Notification.tg_id == tg_id, Notification.notification_type == notification_type + Notification.user_id == u.id, Notification.notification_type == notification_type ) result = await session.execute(stmt) last_time = result.scalar_one_or_none() if not last_time: return True - return datetime.utcnow() - last_time > timedelta(hours=hours) + return _utc_now() - _as_utc(last_time) > timedelta(hours=hours) async def check_notification_time_bulk( @@ -121,37 +165,46 @@ async def check_notification_time_bulk( """ if not items: return set() - now = datetime.utcnow() + now = _utc_now() threshold = now - timedelta(hours=hours) can_notify = set() found = set() - try: - for batch in ( - items[i : i + _NOTIFICATION_TIME_BATCH_SIZE] - for i in range(0, len(items), _NOTIFICATION_TIME_BATCH_SIZE) - ): - stmt = select( - Notification.tg_id, - Notification.notification_type, - Notification.last_notification_time, - ).where(tuple_(Notification.tg_id, Notification.notification_type).in_(batch)) - result = await session.execute(stmt) - for row in result: - found.add((row.tg_id, row.notification_type)) - if row.last_notification_time is None or row.last_notification_time < threshold: - can_notify.add((row.tg_id, row.notification_type)) - for pair in items: - if pair not in found: - can_notify.add(pair) - except SQLAlchemyError: - await session.rollback() - raise + for batch in ( + items[i : i + _NOTIFICATION_TIME_BATCH_SIZE] + for i in range(0, len(items), _NOTIFICATION_TIME_BATCH_SIZE) + ): + id_map = await _map_legacy_refs_to_user_ids(session, [p[0] for p in batch]) + mapped_batch = [(id_map[r], n) for r, n in batch if r in id_map] + if not mapped_batch: + continue + stmt = select( + Notification.user_id, + Notification.notification_type, + Notification.last_notification_time, + ).where(tuple_(Notification.user_id, Notification.notification_type).in_(mapped_batch)) + result = await session.execute(stmt) + uid_to_ref: dict[int, int] = {} + for r, _n in batch: + if r in id_map: + uid_to_ref[id_map[r]] = r + for row in result: + ref = uid_to_ref.get(row.user_id, row.user_id) + found.add((ref, row.notification_type)) + row_time = _as_utc(row.last_notification_time) + if row_time is None or row_time < threshold: + can_notify.add((ref, row.notification_type)) + for pair in items: + if pair not in found: + can_notify.add(pair) return can_notify -async def get_last_notification_time(session: AsyncSession, tg_id: int, notification_type: str) -> int | None: +async def get_last_notification_time(session: AsyncSession, legacy_user_ref: int, notification_type: str) -> int | None: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return None stmt = select(Notification.last_notification_time).where( - Notification.tg_id == tg_id, Notification.notification_type == notification_type + Notification.user_id == u.id, Notification.notification_type == notification_type ) result = await session.execute(stmt) ts = result.scalar_one_or_none() @@ -173,15 +226,24 @@ async def get_last_notification_times_bulk( out = {} for chunk in _batched_list(pairs, _BULK_NOTIFICATION_BATCH_SIZE): + id_map = await _map_legacy_refs_to_user_ids(session, [p[0] for p in chunk]) + mapped = [(id_map[r], n) for r, n in chunk if r in id_map] + if not mapped: + continue stmt = select( - Notification.tg_id, + Notification.user_id, Notification.notification_type, Notification.last_notification_time, - ).where(tuple_(Notification.tg_id, Notification.notification_type).in_(chunk)) + ).where(tuple_(Notification.user_id, Notification.notification_type).in_(mapped)) result = await session.execute(stmt) - for tg_id, ntype, last_time in result.all(): + uid_to_ref: dict[int, int] = {} + for r, _n in chunk: + if r in id_map: + uid_to_ref[id_map[r]] = r + for uid, ntype, last_time in result.all(): if last_time: - out[(tg_id, ntype)] = int(last_time.timestamp() * 1000) + ref = uid_to_ref.get(uid, uid) + out[(ref, ntype)] = int(last_time.timestamp() * 1000) return out @@ -202,54 +264,51 @@ async def get_hot_lead_notification_flags( """ if not tg_ids: return {} - stmt = select(Notification.tg_id, Notification.notification_type).where( - Notification.tg_id.in_(tg_ids), + stmt = select(Notification.user_id, Notification.notification_type).where( + Notification.user_id.in_(tg_ids), Notification.notification_type.in_(_HOT_LEAD_NOTIFICATION_TYPES), ) result = await session.execute(stmt) out = defaultdict(set) - for tg_id, ntype in result.all(): - out[tg_id].add(ntype) + for uid, ntype in result.all(): + out[uid].add(ntype) return dict(out) -async def check_hot_lead_discount(session: AsyncSession, tg_id: int) -> dict: - try: - result = await session.execute( - select(Notification.notification_type, Notification.last_notification_time) - .where(Notification.tg_id == tg_id) - .where(Notification.notification_type.in_(["hot_lead_step_2", "hot_lead_step_3"])) - .order_by(Notification.last_notification_time.desc()) - .limit(1) - ) - - row = result.first() - if not row: - return {"available": False} - - notification_type, last_time = row - - hours = int(NOTIFICATIONS_CONFIG.get("DISCOUNT_ACTIVE_HOURS", DISCOUNT_ACTIVE_HOURS)) - - expires_at = last_time + timedelta(hours=hours) - current_time = datetime.utcnow() - - if current_time > expires_at: - return {"available": False} - - tariff_group = "discounts" if notification_type == "hot_lead_step_2" else "discounts_max" - - return { - "available": True, - "type": notification_type, - "tariff_group": tariff_group, - "expires_at": expires_at, - } - - except Exception as e: - logger.error(f"❌ Ошибка при проверке скидки горячего лида для {tg_id}: {e}") - await session.rollback() +async def check_hot_lead_discount(session: AsyncSession, legacy_user_ref: int) -> dict: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: return {"available": False} + result = await session.execute( + select(Notification.notification_type, Notification.last_notification_time) + .where(Notification.user_id == u.id) + .where(Notification.notification_type.in_(["hot_lead_step_2", "hot_lead_step_3"])) + .order_by(Notification.last_notification_time.desc()) + .limit(1) + ) + + row = result.first() + if not row: + return {"available": False} + + notification_type, last_time = row + + hours = int(NOTIFICATIONS_CONFIG.get("DISCOUNT_ACTIVE_HOURS", DISCOUNT_ACTIVE_HOURS)) + + expires_at = last_time + timedelta(hours=hours) + current_time = _utc_now() + + if current_time > expires_at: + return {"available": False} + + tariff_group = "discounts" if notification_type == "hot_lead_step_2" else "discounts_max" + + return { + "available": True, + "type": notification_type, + "tariff_group": tariff_group, + "expires_at": expires_at, + } _BULK_NOTIFICATION_BATCH_SIZE = 250 @@ -274,144 +333,160 @@ async def check_notifications_bulk( tg_ids: list[int] = None, emails: list[str] = None, ) -> list[dict]: - from sqlalchemy import select + now = _utc_now() - from database.models import BlockedUser, Notification - - try: - now = datetime.utcnow() - - if notification_type == "inactive_trial": - stmt_inactive = ( - select(User.tg_id) - .where( - and_( - User.trial.in_([0, -1]), - ~User.tg_id.in_(select(BlockedUser.tg_id)), - ~User.tg_id.in_(select(Key.tg_id.distinct())), + if notification_type == "inactive_trial": + stmt_inactive = ( + select(User.id) + .where( + and_( + User.trial.in_([0, -1]), + User.tg_id.isnot(None), + ~User.id.in_(select(BlockedUser.user_id)), + ~User.id.in_(select(Key.user_id.distinct())), + ) + ) + ) + result_inactive = await session.execute(stmt_inactive) + inactive_user_ids = [r[0] for r in result_inactive.all()] + if inactive_user_ids: + already = set() + for chunk in _batched_list(inactive_user_ids, _NOTIFICATION_TIME_BATCH_SIZE): + result_existing = await session.execute( + select(Notification.user_id).where( + Notification.notification_type == INACTIVE_TRIAL_REGISTERED_TYPE, + Notification.user_id.in_(chunk), ) ) - ) - result_inactive = await session.execute(stmt_inactive) - inactive_tg_ids = [r[0] for r in result_inactive.all()] - if inactive_tg_ids: - already = set() - for chunk in _batched_list(inactive_tg_ids, _NOTIFICATION_TIME_BATCH_SIZE): - result_existing = await session.execute( - select(Notification.tg_id).where( - Notification.notification_type == INACTIVE_TRIAL_REGISTERED_TYPE, - Notification.tg_id.in_(chunk), - ) + already.update(r[0] for r in result_existing.all()) + to_register = [uid for uid in inactive_user_ids if uid not in already] + if to_register: + for batch in _batched_list(to_register, _BULK_ADD_NOTIFICATIONS_BATCH_SIZE): + await bulk_add_notifications( + session, + [(uid, INACTIVE_TRIAL_REGISTERED_TYPE) for uid in batch], ) - already.update(r[0] for r in result_existing.all()) - to_register = [tid for tid in inactive_tg_ids if tid not in already] - if to_register: - for batch in _batched_list(to_register, _BULK_ADD_NOTIFICATIONS_BATCH_SIZE): - await bulk_add_notifications( - session, - [(tid, INACTIVE_TRIAL_REGISTERED_TYPE) for tid in batch], - commit=True, - ) - logger.info(f"Зарегистрировано как неактивные (шаг 1): {len(to_register)} пользователей.") + logger.info(f"Зарегистрировано как неактивные (шаг 1): {len(to_register)} пользователей.") - subq_registered = ( - select( - Notification.tg_id, - func.max(Notification.last_notification_time).label("registered_time"), - ) - .where(Notification.notification_type == INACTIVE_TRIAL_REGISTERED_TYPE) - .group_by(Notification.tg_id) - .subquery() + subq_registered = ( + select( + Notification.user_id, + func.max(Notification.last_notification_time).label("registered_time"), ) - subq_sent = ( - select( - Notification.tg_id, - func.max(Notification.last_notification_time).label("last_notification_time"), - ) - .where(Notification.notification_type == notification_type) - .group_by(Notification.tg_id) - .subquery() + .where(Notification.notification_type == INACTIVE_TRIAL_REGISTERED_TYPE) + .group_by(Notification.user_id) + .subquery() + ) + subq_sent = ( + select( + Notification.user_id, + func.max(Notification.last_notification_time).label("last_notification_time"), ) - stmt = ( - select( - User.tg_id, - Key.email, - User.username, - User.first_name, - User.last_name, - subq_registered.c.registered_time, - subq_sent.c.last_notification_time, - ) - .select_from(User) - .outerjoin(Key, Key.tg_id == User.tg_id) - .outerjoin(subq_registered, subq_registered.c.tg_id == User.tg_id) - .outerjoin(subq_sent, subq_sent.c.tg_id == User.tg_id) - .where( - and_( - User.trial.in_([0, -1]), - ~User.tg_id.in_(select(BlockedUser.tg_id)), - ~User.tg_id.in_(select(Key.tg_id.distinct())), - ) + .where(Notification.notification_type == notification_type) + .group_by(Notification.user_id) + .subquery() + ) + stmt = ( + select( + User.tg_id, + Key.email, + User.username, + User.first_name, + User.last_name, + subq_registered.c.registered_time, + subq_sent.c.last_notification_time, + ) + .select_from(User) + .outerjoin(Key, Key.user_id == User.id) + .outerjoin(subq_registered, subq_registered.c.user_id == User.id) + .outerjoin(subq_sent, subq_sent.c.user_id == User.id) + .where( + and_( + User.trial.in_([0, -1]), + User.tg_id.isnot(None), + ~User.id.in_(select(BlockedUser.user_id)), + ~User.id.in_(select(Key.user_id.distinct())), ) ) + ) + result = await session.execute(stmt) + users = [] + for row in result: + registered_time = row.registered_time + last_sent_time = row.last_notification_time + first_ok = ( + registered_time is not None + and (now - _as_utc(registered_time)) >= timedelta(hours=hours) + and last_sent_time is None + ) + second_ok = last_sent_time is not None and (now - _as_utc(last_sent_time)) > timedelta(hours=hours) + if first_ok or second_ok: + users.append({ + "tg_id": row.tg_id, + "email": row.email, + "username": row.username, + "first_name": row.first_name, + "last_name": row.last_name, + "last_notification_time": int(last_sent_time.timestamp() * 1000) if last_sent_time else None, + }) + logger.info(f"Найдено {len(users)} пользователей, готовых к уведомлению типа {notification_type}") + return users + + subq_last_notification = ( + select(Notification.user_id, func.max(Notification.last_notification_time).label("last_notification_time")) + .where(Notification.notification_type == notification_type) + .group_by(Notification.user_id) + .subquery() + ) + + def make_stmt(tg_ids_batch: list[int] | None, emails_batch: list[str] | None): + stmt = ( + select( + User.tg_id, + Key.email, + User.username, + User.first_name, + User.last_name, + subq_last_notification.c.last_notification_time, + ) + .select_from(User) + .outerjoin(Key, Key.user_id == User.id) + .outerjoin(subq_last_notification, subq_last_notification.c.user_id == User.id) + ) + if tg_ids_batch: + stmt = stmt.where(User.tg_id.in_(tg_ids_batch)) + if emails_batch: + stmt = stmt.where(Key.email.in_(emails_batch)) + return stmt + + def _can_notify(last_time): + return last_time is None or (now - _as_utc(last_time)) > timedelta(hours=hours) + + users: list[dict] = [] + seen: set[tuple[int, str | None]] = set() + + if tg_ids and emails and len(tg_ids) == len(emails): + for tg_ids_chunk, emails_chunk in _batched_pairs(tg_ids, emails, _BULK_NOTIFICATION_BATCH_SIZE): + stmt = make_stmt(tg_ids_chunk, emails_chunk) result = await session.execute(stmt) - users = [] for row in result: - registered_time = row.registered_time - last_sent_time = row.last_notification_time - first_ok = ( - registered_time is not None - and (now - registered_time) >= timedelta(hours=hours) - and last_sent_time is None - ) - second_ok = last_sent_time is not None and (now - last_sent_time) > timedelta(hours=hours) - if first_ok or second_ok: + key = (row.tg_id, row.email) + if key in seen: + continue + seen.add(key) + last_time = row.last_notification_time + if _can_notify(last_time): users.append({ "tg_id": row.tg_id, "email": row.email, "username": row.username, "first_name": row.first_name, "last_name": row.last_name, - "last_notification_time": int(last_sent_time.timestamp() * 1000) if last_sent_time else None, + "last_notification_time": int(last_time.timestamp() * 1000) if last_time else None, }) - logger.info(f"Найдено {len(users)} пользователей, готовых к уведомлению типа {notification_type}") - return users - - subq_last_notification = ( - select(Notification.tg_id, func.max(Notification.last_notification_time).label("last_notification_time")) - .where(Notification.notification_type == notification_type) - .group_by(Notification.tg_id) - .subquery() - ) - - def make_stmt(tg_ids_batch: list[int] | None, emails_batch: list[str] | None): - stmt = ( - select( - User.tg_id, - Key.email, - User.username, - User.first_name, - User.last_name, - subq_last_notification.c.last_notification_time, - ) - .select_from(User) - .outerjoin(Key, Key.tg_id == User.tg_id) - .outerjoin(subq_last_notification, subq_last_notification.c.tg_id == User.tg_id) - ) - if tg_ids_batch: - stmt = stmt.where(User.tg_id.in_(tg_ids_batch)) - if emails_batch: - stmt = stmt.where(Key.email.in_(emails_batch)) - return stmt - - def _can_notify(last_time): - return last_time is None or (now - last_time) > timedelta(hours=hours) - - users: list[dict] = [] - seen: set[tuple[int, str | None]] = set() - - if tg_ids and emails and len(tg_ids) == len(emails): - for tg_ids_chunk, emails_chunk in _batched_pairs(tg_ids, emails, _BULK_NOTIFICATION_BATCH_SIZE): + elif tg_ids and emails: + for tg_ids_chunk in _batched_list(tg_ids, _BULK_NOTIFICATION_BATCH_SIZE): + for emails_chunk in _batched_list(emails, _BULK_NOTIFICATION_BATCH_SIZE): stmt = make_stmt(tg_ids_chunk, emails_chunk) result = await session.execute(stmt) for row in result: @@ -429,68 +504,15 @@ async def check_notifications_bulk( "last_name": row.last_name, "last_notification_time": int(last_time.timestamp() * 1000) if last_time else None, }) - elif tg_ids and emails: - for tg_ids_chunk in _batched_list(tg_ids, _BULK_NOTIFICATION_BATCH_SIZE): - for emails_chunk in _batched_list(emails, _BULK_NOTIFICATION_BATCH_SIZE): - stmt = make_stmt(tg_ids_chunk, emails_chunk) - result = await session.execute(stmt) - for row in result: - key = (row.tg_id, row.email) - if key in seen: - continue - seen.add(key) - last_time = row.last_notification_time - if _can_notify(last_time): - users.append({ - "tg_id": row.tg_id, - "email": row.email, - "username": row.username, - "first_name": row.first_name, - "last_name": row.last_name, - "last_notification_time": int(last_time.timestamp() * 1000) if last_time else None, - }) - elif tg_ids: - for tg_ids_chunk in _batched_list(tg_ids, _BULK_NOTIFICATION_BATCH_SIZE): - stmt = make_stmt(tg_ids_chunk, None) - result = await session.execute(stmt) - for row in result: - key = (row.tg_id, row.email) - if key in seen: - continue - seen.add(key) - last_time = row.last_notification_time - if _can_notify(last_time): - users.append({ - "tg_id": row.tg_id, - "email": row.email, - "username": row.username, - "first_name": row.first_name, - "last_name": row.last_name, - "last_notification_time": int(last_time.timestamp() * 1000) if last_time else None, - }) - elif emails: - for emails_chunk in _batched_list(emails, _BULK_NOTIFICATION_BATCH_SIZE): - stmt = make_stmt(None, emails_chunk) - result = await session.execute(stmt) - for row in result: - key = (row.tg_id, row.email) - if key in seen: - continue - seen.add(key) - last_time = row.last_notification_time - if _can_notify(last_time): - users.append({ - "tg_id": row.tg_id, - "email": row.email, - "username": row.username, - "first_name": row.first_name, - "last_name": row.last_name, - "last_notification_time": int(last_time.timestamp() * 1000) if last_time else None, - }) - else: - stmt = make_stmt(None, None) + elif tg_ids: + for tg_ids_chunk in _batched_list(tg_ids, _BULK_NOTIFICATION_BATCH_SIZE): + stmt = make_stmt(tg_ids_chunk, None) result = await session.execute(stmt) for row in result: + key = (row.tg_id, row.email) + if key in seen: + continue + seen.add(key) last_time = row.last_notification_time if _can_notify(last_time): users.append({ @@ -501,11 +523,40 @@ async def check_notifications_bulk( "last_name": row.last_name, "last_notification_time": int(last_time.timestamp() * 1000) if last_time else None, }) + elif emails: + for emails_chunk in _batched_list(emails, _BULK_NOTIFICATION_BATCH_SIZE): + stmt = make_stmt(None, emails_chunk) + result = await session.execute(stmt) + for row in result: + key = (row.tg_id, row.email) + if key in seen: + continue + seen.add(key) + last_time = row.last_notification_time + if _can_notify(last_time): + users.append({ + "tg_id": row.tg_id, + "email": row.email, + "username": row.username, + "first_name": row.first_name, + "last_name": row.last_name, + "last_notification_time": int(last_time.timestamp() * 1000) if last_time else None, + }) + else: + stmt = make_stmt(None, None) + result = await session.execute(stmt) + for row in result: + last_time = row.last_notification_time + if _can_notify(last_time): + users.append({ + "tg_id": row.tg_id, + "email": row.email, + "username": row.username, + "first_name": row.first_name, + "last_name": row.last_name, + "last_notification_time": int(last_time.timestamp() * 1000) if last_time else None, + }) - logger.info(f"Найдено {len(users)} пользователей, готовых к уведомлению типа {notification_type}") - return users + logger.info(f"Найдено {len(users)} пользователей, готовых к уведомлению типа {notification_type}") + return users - except Exception as e: - logger.error(f"Ошибка при массовой проверке уведомлений типа {notification_type}: {e}") - await session.rollback() - return [] diff --git a/database/payments.py b/database/payments.py index 8590039f..3279c4c9 100644 --- a/database/payments.py +++ b/database/payments.py @@ -1,12 +1,12 @@ from datetime import datetime, timedelta from pytz import timezone -from sqlalchemy import and_, insert, select, update -from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy import and_, func, insert, select, update from sqlalchemy.ext.asyncio import AsyncSession from core.cache_config import PAYMENT_PENDING_CACHE_TTL_SEC from core.redis_cache import cache_delete, cache_get, cache_key, cache_set +from database.access.resolution import resolve_user_optional from database.models import Payment from logger import logger @@ -52,52 +52,59 @@ async def invalidate_payment_cache(payment_id: str) -> None: async def add_payment( session: AsyncSession, - tg_id: int, - amount: float, - payment_system: str, + legacy_user_ref: int | None = None, + amount: float = 0, + payment_system: str = "", *, + tg_id: int | None = None, status: str = "success", currency: str = "RUB", payment_id: str | None = None, metadata: dict | None = None, original_amount: float | None = None, ) -> int: - try: - now_moscow = datetime.now(MOSCOW_TZ).replace(tzinfo=None) - stmt = ( - insert(Payment) - .values( - tg_id=tg_id, - amount=amount, - payment_system=payment_system, - status=status, - created_at=now_moscow, - currency=currency, - payment_id=payment_id, - metadata_=metadata, - original_amount=original_amount, - ) - .returning(Payment.id) + if legacy_user_ref is None: + legacy_user_ref = tg_id + if legacy_user_ref is None: + raise ValueError("legacy_user_ref is required for payment") + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + raise ValueError(f"user not found for payment: {legacy_user_ref}") + now_moscow = datetime.now(MOSCOW_TZ).replace(tzinfo=None) + stmt = ( + insert(Payment) + .values( + user_id=u.id, + tg_id=u.tg_id, + amount=amount, + payment_system=payment_system, + status=status, + created_at=now_moscow, + currency=currency, + payment_id=payment_id, + metadata_=metadata, + original_amount=original_amount, ) - result = await session.execute(stmt) - internal_id = result.scalar_one() - logger.info( - f"Добавлен платёж id={internal_id}: tg_id={tg_id}, amount={amount}, system={payment_system}, status={status}" - ) - return internal_id - except SQLAlchemyError as e: - await session.rollback() - logger.error(f"Ошибка при добавлении платежа: {e}") - raise + .returning(Payment.id) + ) + result = await session.execute(stmt) + internal_id = result.scalar_one() + logger.info( + f"Добавлен платёж id={internal_id}: user_id={u.id}, amount={amount}, system={payment_system}, status={status}" + ) + return internal_id async def get_last_payments( session: AsyncSession, - tg_id: int, + legacy_user_ref: int, limit: int = 3, statuses: list[str] | None = None, ): - query = select(Payment).where(Payment.tg_id == tg_id) + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return [] + query = select(Payment).where(Payment.user_id == u.id) if statuses: query = query.where(Payment.status.in_(statuses)) @@ -109,7 +116,8 @@ async def get_last_payments( return [ { "id": p.id, - "tg_id": p.tg_id, + "tg_id": p.user_id, + "user_id": p.user_id, "amount": p.amount, "currency": p.currency, "status": p.status, @@ -124,27 +132,23 @@ async def get_last_payments( async def get_payment_by_id(session: AsyncSession, internal_id: int) -> dict | None: - try: - result = await session.execute(select(Payment).where(Payment.id == internal_id).limit(1)) - payment = result.scalar_one_or_none() - if not payment: - return None - return { - "id": payment.id, - "tg_id": payment.tg_id, - "amount": payment.amount, - "currency": payment.currency, - "status": payment.status, - "payment_system": payment.payment_system, - "payment_id": payment.payment_id, - "created_at": payment.created_at, - "metadata": payment.metadata_, - "original_amount": payment.original_amount, - } - except SQLAlchemyError as e: - logger.error(f"Ошибка при поиске платежа id={internal_id}: {e}") - await session.rollback() + result = await session.execute(select(Payment).where(Payment.id == internal_id).limit(1)) + payment = result.scalar_one_or_none() + if not payment: return None + return { + "id": payment.id, + "tg_id": payment.user_id, + "user_id": payment.user_id, + "amount": payment.amount, + "currency": payment.currency, + "status": payment.status, + "payment_system": payment.payment_system, + "payment_id": payment.payment_id, + "created_at": payment.created_at, + "metadata": payment.metadata_, + "original_amount": payment.original_amount, + } async def update_payment_status( @@ -155,32 +159,49 @@ async def update_payment_status( payment_id: str | None = None, metadata_patch: dict | None = None, ) -> bool: - try: - result = await session.execute(select(Payment).where(Payment.id == internal_id).limit(1)) - payment = result.scalar_one_or_none() - if not payment: - logger.info(f"Не удалось сменить статус: платёж id={internal_id} не найден") - return False - - payment.status = new_status - if payment_id is not None: - payment.payment_id = payment_id - base = payment.metadata_ or {} - if new_status == "success" and "status_changed_at" not in base: - base["status_changed_at"] = datetime.utcnow().replace(tzinfo=None).isoformat() - if metadata_patch: - base.update(metadata_patch) - if base: - payment.metadata_ = base - - await session.commit() - logger.info(f"Статус платежа id={internal_id} изменён на {new_status}") - return True - except SQLAlchemyError as e: - await session.rollback() - logger.error(f"Ошибка при смене статуса платежа id={internal_id}: {e}") + result = await session.execute(select(Payment).where(Payment.id == internal_id).limit(1)) + payment = result.scalar_one_or_none() + if not payment: + logger.info(f"Не удалось сменить статус: платёж id={internal_id} не найден") return False + payment.status = new_status + if payment_id is not None: + payment.payment_id = payment_id + base = payment.metadata_ or {} + if new_status == "success" and "status_changed_at" not in base: + base["status_changed_at"] = datetime.utcnow().replace(tzinfo=None).isoformat() + if metadata_patch: + base.update(metadata_patch) + if base: + payment.metadata_ = base + + await session.flush() + logger.info(f"Статус платежа id={internal_id} изменён на {new_status}") + return True + + +async def get_payment_from_db_by_payment_id(session: AsyncSession, pid: str) -> dict | None: + if not str(pid or "").strip(): + return None + result = await session.execute(select(Payment).where(Payment.payment_id == pid).limit(1)) + payment = result.scalar_one_or_none() + if not payment: + return None + return { + "id": payment.id, + "tg_id": payment.user_id, + "user_id": payment.user_id, + "amount": payment.amount, + "currency": payment.currency, + "status": payment.status, + "payment_system": payment.payment_system, + "payment_id": payment.payment_id, + "created_at": payment.created_at, + "metadata": payment.metadata_, + "original_amount": payment.original_amount, + } + async def get_payment_by_payment_id(session: AsyncSession, pid: str) -> dict | None: """Сначала Redis (pending), затем БД. Из кэша возвращается запись без id — вебхук делает add_payment.""" @@ -198,27 +219,36 @@ async def get_payment_by_payment_id(session: AsyncSession, pid: str) -> dict | N "metadata": cached.get("metadata"), "original_amount": cached.get("original_amount"), } - try: - result = await session.execute(select(Payment).where(Payment.payment_id == pid).limit(1)) - payment = result.scalar_one_or_none() - if not payment: - return None - return { - "id": payment.id, - "tg_id": payment.tg_id, - "amount": payment.amount, - "currency": payment.currency, - "status": payment.status, - "payment_system": payment.payment_system, - "payment_id": payment.payment_id, - "created_at": payment.created_at, - "metadata": payment.metadata_, - "original_amount": payment.original_amount, - } - except SQLAlchemyError as e: - logger.error(f"Ошибка при поиске платежа payment_id={pid}: {e}") - await session.rollback() + result = await session.execute(select(Payment).where(Payment.payment_id == pid).limit(1)) + payment = result.scalar_one_or_none() + if not payment: return None + return { + "id": payment.id, + "tg_id": payment.user_id, + "user_id": payment.user_id, + "amount": payment.amount, + "currency": payment.currency, + "status": payment.status, + "payment_system": payment.payment_system, + "payment_id": payment.payment_id, + "created_at": payment.created_at, + "metadata": payment.metadata_, + "original_amount": payment.original_amount, + } + + +async def count_successful_payments(session: AsyncSession, user_id: int) -> int: + """Сколько успешных платежей у пользователя (по internal user id). + + Используется для проверки "новый пользователь" в купонных правилах. + """ + result = await session.execute( + select(func.count()) + .select_from(Payment) + .where(Payment.user_id == int(user_id), func.lower(Payment.status) == "success") + ) + return int(result.scalar() or 0) async def cancel_expired_pending_payments(session: AsyncSession) -> int: @@ -234,17 +264,19 @@ async def cancel_expired_pending_payments(session: AsyncSession) -> int: .values(status="cancelled") ) res = await session.execute(stmt) - await session.commit() affected = res.rowcount or 0 return affected async def get_all_payments( session: AsyncSession, - tg_id: int, + legacy_user_ref: int, statuses: list[str] | None = None, ) -> list[dict]: - query = select(Payment).where(Payment.tg_id == tg_id) + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return [] + query = select(Payment).where(Payment.user_id == u.id) if statuses: query = query.where(Payment.status.in_(statuses)) @@ -256,7 +288,8 @@ async def get_all_payments( return [ { "id": p.id, - "tg_id": p.tg_id, + "tg_id": p.user_id, + "user_id": p.user_id, "amount": p.amount, "currency": p.currency, "status": p.status, diff --git a/database/referrals.py b/database/referrals.py index 10eb5466..c8288fc8 100644 --- a/database/referrals.py +++ b/database/referrals.py @@ -1,50 +1,62 @@ from sqlalchemy import and_, desc, func, insert, select, text, update -from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from config import CHECK_REFERRAL_REWARD_ISSUED, REFERRAL_BONUS_PERCENTAGES from core.bootstrap import BUTTONS_CONFIG +from database.access.resolution import resolve_user_optional from database.models import Referral from logger import logger -async def add_referral(session: AsyncSession, referred_tg_id: int, referrer_tg_id: int): - try: - if referred_tg_id == referrer_tg_id: - logger.warning(f"⚠️ Попытка самореферала: {referred_tg_id}") - return +async def add_referral(session: AsyncSession, referred_legacy: int, referrer_legacy: int): + ru = await resolve_user_optional(session, referred_legacy) + rf = await resolve_user_optional(session, referrer_legacy) + if ru is None or rf is None: + return + if ru.id == rf.id: + logger.warning(f"⚠️ Попытка самореферала: {referred_legacy}") + return - stmt = insert(Referral).values(referred_tg_id=referred_tg_id, referrer_tg_id=referrer_tg_id) - await session.execute(stmt) - await session.commit() - logger.info(f"✅ Добавлена реферальная связь: {referred_tg_id} → {referrer_tg_id}") - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при добавлении реферала: {e}") - await session.rollback() - raise + stmt = insert(Referral).values( + referred_user_id=ru.id, + referrer_user_id=rf.id, + referred_tg_id=ru.tg_id, + referrer_tg_id=rf.tg_id, + ) + await session.execute(stmt) + logger.info(f"✅ Добавлена реферальная связь: {ru.id} → {rf.id}") -async def get_referral_by_referred_id(session: AsyncSession, referred_tg_id: int) -> dict | None: - stmt = select(Referral).where(Referral.referred_tg_id == referred_tg_id) +async def get_referral_by_referred_id(session: AsyncSession, referred_legacy: int) -> dict | None: + ru = await resolve_user_optional(session, referred_legacy) + if ru is None: + return None + stmt = select(Referral).where(Referral.referred_user_id == ru.id) result = await session.execute(stmt) row = result.scalar_one_or_none() return dict(row.__dict__) if row else None -async def get_total_referrals(session: AsyncSession, referrer_tg_id: int) -> int: - stmt = select(func.count()).select_from(Referral).where(Referral.referrer_tg_id == referrer_tg_id) +async def get_total_referrals(session: AsyncSession, referrer_legacy: int) -> int: + ru = await resolve_user_optional(session, referrer_legacy) + if ru is None: + return 0 + stmt = select(func.count()).select_from(Referral).where(Referral.referrer_user_id == ru.id) result = await session.execute(stmt) return result.scalar() -async def get_active_referrals(session: AsyncSession, referrer_tg_id: int) -> int: +async def get_active_referrals(session: AsyncSession, referrer_legacy: int) -> int: + ru = await resolve_user_optional(session, referrer_legacy) + if ru is None: + return 0 stmt = ( select(func.count()) .select_from(Referral) .where( and_( - Referral.referrer_tg_id == referrer_tg_id, - Referral.reward_issued is True, + Referral.referrer_user_id == ru.id, + Referral.reward_issued.is_(True), ) ) ) @@ -52,44 +64,51 @@ async def get_active_referrals(session: AsyncSession, referrer_tg_id: int) -> in return result.scalar() -async def mark_referral_reward_issued(session: AsyncSession, referred_tg_id: int): - await session.execute(update(Referral).where(Referral.referred_tg_id == referred_tg_id).values(reward_issued=True)) - await session.commit() +async def mark_referral_reward_issued(session: AsyncSession, referred_legacy: int): + ru = await resolve_user_optional(session, referred_legacy) + if ru is None: + return + await session.execute(update(Referral).where(Referral.referred_user_id == ru.id).values(reward_issued=True)) -async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, max_levels: int) -> float: +async def get_total_referral_bonus(session: AsyncSession, referrer_legacy: int, max_levels: int) -> float: referral_enabled = bool(BUTTONS_CONFIG.get("REFERRAL_BUTTON_ENABLED", True)) if not referral_enabled: logger.debug("Реферальная программа отключена, бонусы не начисляются") return 0.0 + ru = await resolve_user_optional(session, referrer_legacy) + if ru is None: + return 0.0 + uid = ru.id + if CHECK_REFERRAL_REWARD_ISSUED: bonus_cte = """ WITH RECURSIVE referral_levels AS ( SELECT - referred_tg_id, - referrer_tg_id, + referred_user_id, + referrer_user_id, 1 AS level FROM referrals - WHERE referrer_tg_id = :tg_id AND reward_issued = TRUE + WHERE referrer_user_id = :user_id AND reward_issued = TRUE UNION SELECT - r.referred_tg_id, - r.referrer_tg_id, + r.referred_user_id, + r.referrer_user_id, rl.level + 1 FROM referrals r - JOIN referral_levels rl ON r.referrer_tg_id = rl.referred_tg_id + JOIN referral_levels rl ON r.referrer_user_id = rl.referred_user_id WHERE rl.level < :max_levels AND r.reward_issued = TRUE ), earliest_payments AS ( - SELECT DISTINCT ON (tg_id) tg_id, amount, created_at + SELECT DISTINCT ON (user_id) user_id, amount, created_at FROM payments WHERE status = 'success' AND payment_system NOT IN ('coupon', 'admin', 'referral') - ORDER BY tg_id, created_at + ORDER BY user_id, created_at ) """ bonus_query = ( @@ -110,7 +129,7 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m END ), 0) AS total_bonus FROM referral_levels rl - JOIN earliest_payments ep ON rl.referred_tg_id = ep.tg_id + JOIN earliest_payments ep ON rl.referred_user_id = ep.user_id WHERE rl.level <= :max_levels """ ) @@ -119,20 +138,20 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m WITH RECURSIVE referral_levels AS ( SELECT - referred_tg_id, - referrer_tg_id, + referred_user_id, + referrer_user_id, 1 AS level FROM referrals - WHERE referrer_tg_id = :tg_id + WHERE referrer_user_id = :user_id UNION SELECT - r.referred_tg_id, - r.referrer_tg_id, + r.referred_user_id, + r.referrer_user_id, rl.level + 1 FROM referrals r - JOIN referral_levels rl ON r.referrer_tg_id = rl.referred_tg_id + JOIN referral_levels rl ON r.referrer_user_id = rl.referred_user_id WHERE rl.level < :max_levels ) """ @@ -154,7 +173,7 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m END ), 0) AS total_bonus FROM referral_levels rl - JOIN payments p ON rl.referred_tg_id = p.tg_id + JOIN payments p ON rl.referred_user_id = p.user_id WHERE p.status = 'success' AND p.payment_system NOT IN ('coupon', 'admin', 'referral') AND rl.level <= :max_levels @@ -163,7 +182,7 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m result = await session.execute( text(bonus_query), - {"tg_id": referrer_tg_id, "max_levels": max_levels}, + {"user_id": uid, "max_levels": max_levels}, ) total_bonus_raw = result.scalar() total_bonus = round(float(total_bonus_raw or 0), 2) @@ -172,29 +191,32 @@ async def get_total_referral_bonus(session: AsyncSession, referrer_tg_id: int, m return total_bonus -async def get_referrals_by_level(session: AsyncSession, referrer_tg_id: int, max_levels: int) -> dict: +async def get_referrals_by_level(session: AsyncSession, referrer_legacy: int, max_levels: int) -> dict: + ru = await resolve_user_optional(session, referrer_legacy) + if ru is None: + return {} query = """ WITH RECURSIVE referral_levels AS ( - SELECT referred_tg_id, referrer_tg_id, 1 AS level + SELECT referred_user_id, referrer_user_id, 1 AS level FROM referrals - WHERE referrer_tg_id = :referrer_tg_id + WHERE referrer_user_id = :referrer_user_id UNION - SELECT r.referred_tg_id, r.referrer_tg_id, rl.level + 1 + SELECT r.referred_user_id, r.referrer_user_id, rl.level + 1 FROM referrals r - JOIN referral_levels rl ON r.referrer_tg_id = rl.referred_tg_id + JOIN referral_levels rl ON r.referrer_user_id = rl.referred_user_id WHERE rl.level < :max_levels ) SELECT level, COUNT(*) AS level_count, COUNT(CASE WHEN reward_issued THEN 1 END) AS active_level_count FROM referral_levels rl - JOIN referrals r ON rl.referred_tg_id = r.referred_tg_id + JOIN referrals r ON rl.referred_user_id = r.referred_user_id GROUP BY level ORDER BY level """ result = await session.execute( text(query), - {"referrer_tg_id": referrer_tg_id, "max_levels": max_levels}, + {"referrer_user_id": ru.id, "max_levels": max_levels}, ) return { row["level"]: { @@ -205,38 +227,35 @@ async def get_referrals_by_level(session: AsyncSession, referrer_tg_id: int, max } -async def get_referral_stats(session: AsyncSession, referrer_tg_id: int): - try: - logger.info(f"[ReferralStats] Получение статистики для пользователя {referrer_tg_id}") +async def get_referral_stats(session: AsyncSession, referrer_legacy: int): + logger.info(f"[ReferralStats] Получение статистики для пользователя {referrer_legacy}") - total_referrals = await get_total_referrals(session, referrer_tg_id) - active_referrals = await get_active_referrals(session, referrer_tg_id) - max_levels = len(REFERRAL_BONUS_PERCENTAGES) - referrals_by_level = await get_referrals_by_level(session, referrer_tg_id, max_levels) - total_referral_bonus = await get_total_referral_bonus(session, referrer_tg_id, max_levels) + total_referrals = await get_total_referrals(session, referrer_legacy) + active_referrals = await get_active_referrals(session, referrer_legacy) + max_levels = len(REFERRAL_BONUS_PERCENTAGES) + referrals_by_level = await get_referrals_by_level(session, referrer_legacy, max_levels) + total_referral_bonus = await get_total_referral_bonus(session, referrer_legacy, max_levels) - return { - "total_referrals": total_referrals, - "active_referrals": active_referrals, - "referrals_by_level": referrals_by_level, - "total_referral_bonus": total_referral_bonus, - } - - except Exception as e: - logger.error(f"[ReferralStats] Ошибка при получении статистики для пользователя {referrer_tg_id}: {e}") - await session.rollback() - raise + return { + "total_referrals": total_referrals, + "active_referrals": active_referrals, + "referrals_by_level": referrals_by_level, + "total_referral_bonus": total_referral_bonus, + } -async def get_user_referral_count(session: AsyncSession, tg_id: int) -> int: - result = await session.execute(select(func.count()).select_from(Referral).where(Referral.referrer_tg_id == tg_id)) +async def get_user_referral_count(session: AsyncSession, legacy: int) -> int: + ru = await resolve_user_optional(session, legacy) + if ru is None: + return 0 + result = await session.execute(select(func.count()).select_from(Referral).where(Referral.referrer_user_id == ru.id)) return result.scalar_one() or 0 async def get_referral_position(session: AsyncSession, referral_count: int) -> int: subq = ( - select(Referral.referrer_tg_id) - .group_by(Referral.referrer_tg_id) + select(Referral.referrer_user_id) + .group_by(Referral.referrer_user_id) .having(func.count() > referral_count) .subquery() ) @@ -248,10 +267,10 @@ async def get_referral_position(session: AsyncSession, referral_count: int) -> i async def get_top_referrals(session: AsyncSession, limit: int = 5): query = ( - select(Referral.referrer_tg_id, func.count().label("referral_count")) - .group_by(Referral.referrer_tg_id) + select(Referral.referrer_user_id, func.count().label("referral_count")) + .group_by(Referral.referrer_user_id) .order_by(desc("referral_count")) .limit(limit) ) result = await session.execute(query) - return [{"referrer_tg_id": row.referrer_tg_id, "referral_count": row.referral_count} for row in result.all()] + return [{"referrer_user_id": row.referrer_user_id, "referral_count": row.referral_count} for row in result.all()] diff --git a/database/scheduled_broadcasts.py b/database/scheduled_broadcasts.py index 9bfddf2e..4b57b7b3 100644 --- a/database/scheduled_broadcasts.py +++ b/database/scheduled_broadcasts.py @@ -3,6 +3,7 @@ from datetime import datetime, timezone from sqlalchemy import select, update from sqlalchemy.ext.asyncio import AsyncSession +from database.access.resolution import resolve_user_optional from database.models import ScheduledBroadcast @@ -34,8 +35,16 @@ async def create_scheduled_broadcast( messages_per_second: int, status: str = SCHEDULED_BROADCAST_STATUS_SCHEDULED, ) -> ScheduledBroadcast: + created_by_uid = None + mirror_tg = created_by_tg_id + if created_by_tg_id is not None: + cu = await resolve_user_optional(session, created_by_tg_id) + if cu is not None: + created_by_uid = cu.id + mirror_tg = cu.tg_id broadcast = ScheduledBroadcast( - created_by_tg_id=created_by_tg_id, + created_by_user_id=created_by_uid, + created_by_tg_id=mirror_tg, send_to=send_to, cluster_name=cluster_name, text=text, @@ -47,7 +56,7 @@ async def create_scheduled_broadcast( status=status, ) session.add(broadcast) - await session.commit() + await session.flush() await session.refresh(broadcast) return broadcast @@ -91,9 +100,7 @@ async def update_scheduled_broadcast( .values(**values) ) if not result.rowcount: - await session.rollback() return None - await session.commit() return await get_scheduled_broadcast(session, broadcast_id) @@ -112,9 +119,7 @@ async def cancel_scheduled_broadcast(session: AsyncSession, broadcast_id: str) - ) ) if not result.rowcount: - await session.rollback() return None - await session.commit() return await get_scheduled_broadcast(session, broadcast_id) @@ -148,9 +153,7 @@ async def claim_due_scheduled_broadcasts(session: AsyncSession, limit: int = 10) if claim_result.rowcount: claimed_ids.append(broadcast_id) if not claimed_ids: - await session.rollback() return [] - await session.commit() result = await session.execute( select(ScheduledBroadcast) .where(ScheduledBroadcast.id.in_(claimed_ids)) @@ -177,9 +180,7 @@ async def start_scheduled_broadcast(session: AsyncSession, broadcast_id: str) -> ) ) if not result.rowcount: - await session.rollback() return None - await session.commit() return await get_scheduled_broadcast(session, broadcast_id) @@ -200,7 +201,6 @@ async def mark_scheduled_broadcast_sent( updated_at=datetime.utcnow(), ) ) - await session.commit() return await get_scheduled_broadcast(session, broadcast_id) @@ -218,5 +218,4 @@ async def mark_scheduled_broadcast_failed( updated_at=datetime.utcnow(), ) ) - await session.commit() return await get_scheduled_broadcast(session, broadcast_id) diff --git a/database/servers.py b/database/servers.py index 99eddad1..9d06af34 100644 --- a/database/servers.py +++ b/database/servers.py @@ -1,5 +1,4 @@ from sqlalchemy import delete, func, insert, select, update -from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from core.cache_config import SERVERS_CACHE_TTL_SEC @@ -20,35 +19,23 @@ async def create_server( subscription_url: str, inbound_id: str, ): - try: - stmt = insert(Server).values( - cluster_name=cluster_name, - server_name=server_name, - api_url=api_url, - subscription_url=subscription_url, - inbound_id=inbound_id, - ) - await session.execute(stmt) - await session.commit() - await _invalidate_servers_cache() - logger.info(f"✅ Сервер {server_name} добавлен в кластер {cluster_name}") - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при добавлении сервера {server_name}: {e}") - await session.rollback() - raise + stmt = insert(Server).values( + cluster_name=cluster_name, + server_name=server_name, + api_url=api_url, + subscription_url=subscription_url, + inbound_id=inbound_id, + ) + await session.execute(stmt) + await _invalidate_servers_cache() + logger.info(f"✅ Сервер {server_name} добавлен в кластер {cluster_name}") async def delete_server(session: AsyncSession, server_name: str): - try: - stmt = delete(Server).where(Server.server_name == server_name) - await session.execute(stmt) - await session.commit() - await _invalidate_servers_cache() - logger.info(f"🗑 Сервер {server_name} удалён") - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при удалении сервера {server_name}: {e}") - await session.rollback() - raise + stmt = delete(Server).where(Server.server_name == server_name) + await session.execute(stmt) + await _invalidate_servers_cache() + logger.info(f"🗑 Сервер {server_name} удалён") async def get_servers(session: AsyncSession, include_enabled: bool = False) -> dict: @@ -59,63 +46,58 @@ async def get_servers(session: AsyncSession, include_enabled: bool = False) -> d if isinstance(cached, dict): return cached - try: - stmt = select(Server) - result = await session.execute(stmt) - servers = result.scalars().all() + stmt = select(Server) + result = await session.execute(stmt) + servers = result.scalars().all() - ids = [s.id for s in servers] - subs_map = {} - tariffs_map = {} - if ids: - r = await session.execute( - select(ServerSubgroup.server_id, ServerSubgroup.subgroup_title).where(ServerSubgroup.server_id.in_(ids)) + ids = [s.id for s in servers] + subs_map = {} + tariffs_map = {} + if ids: + r = await session.execute( + select(ServerSubgroup.server_id, ServerSubgroup.subgroup_title).where(ServerSubgroup.server_id.in_(ids)) + ) + for sid, sg in r.all(): + if sg and sg.isdigit(): + tariffs_map.setdefault(sid, []).append(int(sg)) + else: + subs_map.setdefault(sid, []).append(sg) + + groups_map = {} + if ids: + r2 = await session.execute( + select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where( + ServerSpecialgroup.server_id.in_(ids) ) - for sid, sg in r.all(): - if sg and sg.isdigit(): - tariffs_map.setdefault(sid, []).append(int(sg)) - else: - subs_map.setdefault(sid, []).append(sg) + ) + for sid, gc in r2.all(): + groups_map.setdefault(sid, []).append(gc) - groups_map = {} - if ids: - r2 = await session.execute( - select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where( - ServerSpecialgroup.server_id.in_(ids) - ) - ) - for sid, gc in r2.all(): - groups_map.setdefault(sid, []).append(gc) + allowed = set(ALLOWED_GROUP_CODES) - allowed = set(ALLOWED_GROUP_CODES) - - grouped = {} - for s in servers: - if not include_enabled and not s.enabled: - continue - cluster = s.cluster_name - special = sorted({g for g in groups_map.get(s.id, []) if g in allowed}) - grouped.setdefault(cluster, []).append({ - "server_name": s.server_name, - "api_url": s.api_url, - "subscription_url": s.subscription_url, - "inbound_id": s.inbound_id, - "panel_type": s.panel_type, - "enabled": s.enabled, - "max_keys": s.max_keys, - "tariff_group": s.tariff_group, - "tariff_subgroups": subs_map.get(s.id, []), - "tariff_ids": tariffs_map.get(s.id, []), - "special_groups": special, - "cluster_name": cluster, - "server_id": s.id, - }) - await cache_set(cache_key_servers, grouped, SERVERS_CACHE_TTL_SEC) - return grouped - except SQLAlchemyError as e: - logger.error(f"Ошибка при получении серверов: {e}") - await session.rollback() - return {} + grouped = {} + for s in servers: + if not include_enabled and not s.enabled: + continue + cluster = s.cluster_name + special = sorted({g for g in groups_map.get(s.id, []) if g in allowed}) + grouped.setdefault(cluster, []).append({ + "server_name": s.server_name, + "api_url": s.api_url, + "subscription_url": s.subscription_url, + "inbound_id": s.inbound_id, + "panel_type": s.panel_type, + "enabled": s.enabled, + "max_keys": s.max_keys, + "tariff_group": s.tariff_group, + "tariff_subgroups": subs_map.get(s.id, []), + "tariff_ids": tariffs_map.get(s.id, []), + "special_groups": special, + "cluster_name": cluster, + "server_id": s.id, + }) + await cache_set(cache_key_servers, grouped, SERVERS_CACHE_TTL_SEC) + return grouped async def get_clusters(session: AsyncSession) -> list[str]: @@ -139,14 +121,49 @@ async def check_unique_server_name(session: AsyncSession, server_name: str, clus async def check_server_name_by_cluster(session: AsyncSession, server_name: str) -> dict | None: - try: - result = await session.execute(select(Server.cluster_name).where(Server.server_name == server_name)) - row = result.first() - return {"cluster_name": row[0]} if row else None - except SQLAlchemyError as e: - logger.error(f"Ошибка при поиске кластера для сервера {server_name}: {e}") - await session.rollback() - return None + result = await session.execute(select(Server.cluster_name).where(Server.server_name == server_name)) + row = result.first() + return {"cluster_name": row[0]} if row else None + + +async def get_panel_types_for_cluster(session: AsyncSession, cluster_name: str) -> list[str]: + """Список panel_type всех серверов кластера (для проверки "весь remnawave").""" + result = await session.execute( + select(Server.panel_type).where(Server.cluster_name == cluster_name) + ) + return list(result.scalars().all()) + + +async def get_panel_type_for_server(session: AsyncSession, server_name: str) -> str | None: + """Возвращает panel_type конкретного сервера по его имени.""" + result = await session.execute( + select(Server.panel_type).where(Server.server_name == server_name) + ) + return result.scalar_one_or_none() + + +async def get_enabled_server_subscription_url(session: AsyncSession, server_name: str) -> str | None: + """Возвращает ``subscription_url`` для включённого сервера по его имени.""" + result = await session.execute( + select(Server.subscription_url).where(Server.server_name == server_name, Server.enabled.is_(True)) + ) + return result.scalar() + + +async def cluster_name_exists(session: AsyncSession, cluster_name: str) -> bool: + """Есть ли хоть один сервер с таким cluster_name.""" + result = await session.execute( + select(Server).where(Server.cluster_name == cluster_name).limit(1) + ) + return result.scalars().first() is not None + + +async def get_cluster_name_for_server_name(session: AsyncSession, server_name: str) -> str | None: + """Возвращает cluster_name для указанного server_name (строго по server_name).""" + result = await session.execute( + select(Server.cluster_name).where(Server.server_name == server_name).limit(1) + ) + return result.scalar() async def get_cluster_name_by_server(session: AsyncSession, server_id_or_name: str) -> str | None: @@ -162,132 +179,100 @@ async def get_cluster_name_by_server(session: AsyncSession, server_id_or_name: s async def get_server_by_name(session: AsyncSession, server_name: str) -> dict | None: - try: - stmt = select(Server).where(Server.server_name == server_name) - result = await session.execute(stmt) - server = result.scalar_one_or_none() + stmt = select(Server).where(Server.server_name == server_name) + result = await session.execute(stmt) + server = result.scalar_one_or_none() - if server: - return { - "id": server.id, - "cluster_name": server.cluster_name, - "server_name": server.server_name, - "api_url": server.api_url, - "subscription_url": server.subscription_url, - "inbound_id": server.inbound_id, - "panel_type": server.panel_type, - "enabled": server.enabled, - "max_keys": server.max_keys, - "tariff_group": server.tariff_group, - } - return None - except SQLAlchemyError as e: - logger.error(f"Ошибка при получении сервера {server_name}: {e}") - await session.rollback() - return None + if server: + return { + "id": server.id, + "cluster_name": server.cluster_name, + "server_name": server.server_name, + "api_url": server.api_url, + "subscription_url": server.subscription_url, + "inbound_id": server.inbound_id, + "panel_type": server.panel_type, + "enabled": server.enabled, + "max_keys": server.max_keys, + "tariff_group": server.tariff_group, + } + return None async def update_server_field(session: AsyncSession, server_name: str, field: str, value: any) -> bool: - try: - stmt = update(Server).where(Server.server_name == server_name).values(**{field: value}) - await session.execute(stmt) - await session.commit() - await _invalidate_servers_cache() - logger.info(f"✅ Поле {field} сервера {server_name} обновлено на {value}") - return True - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при обновлении поля {field} сервера {server_name}: {e}") - await session.rollback() - return False + stmt = update(Server).where(Server.server_name == server_name).values(**{field: value}) + await session.execute(stmt) + await _invalidate_servers_cache() + logger.info(f"✅ Поле {field} сервера {server_name} обновлено на {value}") + return True async def update_server_name_with_keys(session: AsyncSession, old_name: str, new_name: str) -> bool: - try: - from sqlalchemy import update - - from database.models import Key - - if not await check_unique_server_name(session, new_name): - logger.error(f"❌ Сервер с именем {new_name} уже существует") - return False - - stmt_server = update(Server).where(Server.server_name == old_name).values(server_name=new_name) - await session.execute(stmt_server) - - stmt_keys = update(Key).where(Key.server_id == old_name).values(server_id=new_name) - await session.execute(stmt_keys) - - await session.commit() - await _invalidate_servers_cache() - logger.info(f"✅ Сервер переименован с {old_name} на {new_name}") - return True - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при переименовании сервера {old_name}: {e}") - await session.rollback() + if not await check_unique_server_name(session, new_name): + logger.error(f"❌ Сервер с именем {new_name} уже существует") return False + stmt_server = update(Server).where(Server.server_name == old_name).values(server_name=new_name) + await session.execute(stmt_server) + + stmt_keys = update(Key).where(Key.server_id == old_name).values(server_id=new_name) + await session.execute(stmt_keys) + + await _invalidate_servers_cache() + logger.info(f"✅ Сервер переименован с {old_name} на {new_name}") + return True + async def get_available_clusters(session: AsyncSession) -> list[str]: - try: - stmt = select(Server.cluster_name).distinct().order_by(Server.cluster_name) - result = await session.execute(stmt) - return [row[0] for row in result.all()] - except SQLAlchemyError as e: - logger.error(f"Ошибка при получении списка кластеров: {e}") - await session.rollback() - return [] + stmt = select(Server.cluster_name).distinct().order_by(Server.cluster_name) + result = await session.execute(stmt) + return [row[0] for row in result.all()] async def update_server_cluster(session: AsyncSession, server_name: str, new_cluster: str) -> bool: - try: - server_data = await get_server_by_name(session, server_name) - if not server_data: - return False - - old_cluster = server_data["cluster_name"] - - stmt_remaining = select(func.count()).where( - (Server.cluster_name == old_cluster) & (Server.server_name != server_name) - ) - result = await session.execute(stmt_remaining) - remaining_servers = result.scalar_one() - - if remaining_servers == 0: - stmt_update_keys = update(Key).where(Key.server_id == old_cluster).values(server_id=new_cluster) - await session.execute(stmt_update_keys) - - stmt_new_cluster = select(Server.tariff_group).where(Server.cluster_name == new_cluster).limit(1) - result = await session.execute(stmt_new_cluster) - new_tariff_group = result.scalar_one_or_none() - - await session.execute( - update(Server) - .where(Server.server_name == server_name) - .values(cluster_name=new_cluster, tariff_group=new_tariff_group) - ) - - if server_data.get("id") is None: - rid = await session.execute(select(Server.id).where(Server.server_name == server_name).limit(1)) - server_id = rid.scalar_one_or_none() - else: - server_id = server_data["id"] - - if server_id is not None and new_tariff_group is not None: - await session.execute( - update(ServerSubgroup).where(ServerSubgroup.server_id == server_id).values(group_code=new_tariff_group) - ) - - await session.commit() - await _invalidate_servers_cache() - logger.info( - f"✅ Сервер {server_name} перемещен в кластер {new_cluster} с обновлением тарифной группы и привязок подгрупп" - ) - return True - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при обновлении кластера сервера {server_name}: {e}") - await session.rollback() + server_data = await get_server_by_name(session, server_name) + if not server_data: return False + old_cluster = server_data["cluster_name"] + + stmt_remaining = select(func.count()).where( + (Server.cluster_name == old_cluster) & (Server.server_name != server_name) + ) + result = await session.execute(stmt_remaining) + remaining_servers = result.scalar_one() + + if remaining_servers == 0: + stmt_update_keys = update(Key).where(Key.server_id == old_cluster).values(server_id=new_cluster) + await session.execute(stmt_update_keys) + + stmt_new_cluster = select(Server.tariff_group).where(Server.cluster_name == new_cluster).limit(1) + result = await session.execute(stmt_new_cluster) + new_tariff_group = result.scalar_one_or_none() + + await session.execute( + update(Server) + .where(Server.server_name == server_name) + .values(cluster_name=new_cluster, tariff_group=new_tariff_group) + ) + + if server_data.get("id") is None: + rid = await session.execute(select(Server.id).where(Server.server_name == server_name).limit(1)) + server_id = rid.scalar_one_or_none() + else: + server_id = server_data["id"] + + if server_id is not None and new_tariff_group is not None: + await session.execute( + update(ServerSubgroup).where(ServerSubgroup.server_id == server_id).values(group_code=new_tariff_group) + ) + + await _invalidate_servers_cache() + logger.info( + f"✅ Сервер {server_name} перемещен в кластер {new_cluster} с обновлением тарифной группы и привязок подгрупп" + ) + return True + async def resolve_device_limit_from_group(session: AsyncSession, server_id: str) -> int | None: r = await session.execute(select(Server.tariff_group).where(Server.server_name == server_id)) diff --git a/database/setup/__init__.py b/database/setup/__init__.py new file mode 100644 index 00000000..3f80c2fe --- /dev/null +++ b/database/setup/__init__.py @@ -0,0 +1 @@ +from .init_db import * diff --git a/database/init_db.py b/database/setup/init_db.py similarity index 81% rename from database/init_db.py rename to database/setup/init_db.py index e64284a1..933d8a69 100644 --- a/database/init_db.py +++ b/database/setup/init_db.py @@ -4,12 +4,19 @@ from sqlalchemy import select from config import ADMIN_ID from database import db +from database.migrations.schema_upgrade import ( + apply_all_migrations, + ensure_tg_mirror_columns_and_backfill, +) from database.models import Admin, Base, User async def init_db(): async with db.engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) + await apply_all_migrations(conn) + async with db.engine.begin() as conn: + await ensure_tg_mirror_columns_and_backfill(conn) async with db.async_session_maker() as session: result = await session.execute(select(User).where(User.tg_id == 0)) diff --git a/database/statistics.py b/database/statistics.py index cfc75ab9..726ca3b5 100644 --- a/database/statistics.py +++ b/database/statistics.py @@ -152,15 +152,15 @@ async def sum_total_payments(session: AsyncSession) -> float: async def count_hot_leads(session: AsyncSession) -> int: subquery_active_keys = ( - select(Key.tg_id).where(Key.expiry_time > int(datetime.utcnow().timestamp() * 1000)).distinct() + select(Key.user_id).where(Key.expiry_time > int(datetime.utcnow().timestamp() * 1000)).distinct() ) stmt = ( - select(Payment.tg_id) + select(Payment.user_id) .where(Payment.amount > 0) .where(Payment.status == "success") .where(Payment.payment_system.notin_(PAYMENT_SYSTEMS_EXCLUDED)) - .where(not_(exists(subquery_active_keys.where(Key.tg_id == Payment.tg_id)))) + .where(not_(exists(subquery_active_keys.where(Key.user_id == Payment.user_id)))) .distinct() ) diff --git a/database/tariffs.py b/database/tariffs.py index a17f722a..9c5dee63 100644 --- a/database/tariffs.py +++ b/database/tariffs.py @@ -4,7 +4,6 @@ from collections import defaultdict from datetime import datetime from sqlalchemy import delete, func, insert, select, update -from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from core.cache_config import TARIFF_BY_ID_CACHE_TTL_SEC, TARIFFS_FOR_CLUSTER_CACHE_TTL_SEC @@ -31,6 +30,7 @@ async def _invalidate_tariff_cache(tariff_id: int | None = None) -> None: if tariff_id is not None: await cache_delete(cache_key("tariff", tariff_id)) await cache_delete_pattern("tariffs_cluster:*") + await cache_delete_pattern("tariffs_public:*") def create_subgroup_hash(subgroup_title: str, group_code: str) -> str: @@ -60,43 +60,37 @@ async def find_subgroup_by_hash(session: AsyncSession, subgroup_hash: str, group async def get_tariffs( session: AsyncSession, tariff_id: int = None, group_code: str = None, with_subgroup_weights: bool = False ): - try: - if tariff_id: - result = await session.execute(select(Tariff).where(Tariff.id == tariff_id)) - elif group_code: - result = await session.execute( - select(Tariff).where(Tariff.group_code == group_code).order_by(Tariff.sort_order, Tariff.id) - ) - else: - result = await session.execute(select(Tariff).order_by(Tariff.sort_order, Tariff.id)) + if tariff_id: + result = await session.execute(select(Tariff).where(Tariff.id == tariff_id)) + elif group_code: + result = await session.execute( + select(Tariff).where(Tariff.group_code == group_code).order_by(Tariff.sort_order, Tariff.id) + ) + else: + result = await session.execute(select(Tariff).order_by(Tariff.sort_order, Tariff.id)) - tariffs = [dict(r.__dict__) for r in result.scalars().all()] + tariffs = [dict(r.__dict__) for r in result.scalars().all()] - if with_subgroup_weights and group_code: - tariffs_without_order = [t for t in tariffs if t.get("sort_order") is None] - if tariffs_without_order: - for tariff in tariffs_without_order: - tariff["sort_order"] = 1 - await session.execute(update(Tariff).where(Tariff.id == tariff["id"]).values(sort_order=1)) - await session.commit() + if with_subgroup_weights and group_code: + tariffs_without_order = [t for t in tariffs if t.get("sort_order") is None] + if tariffs_without_order: + for tariff in tariffs_without_order: + tariff["sort_order"] = 1 + await session.execute(update(Tariff).where(Tariff.id == tariff["id"]).values(sort_order=1)) - grouped = defaultdict(list) - for t in tariffs: - grouped[t.get("subgroup_title")].append(t) + grouped = defaultdict(list) + for t in tariffs: + grouped[t.get("subgroup_title")].append(t) - subgroup_weights = {} - for subgroup, tariffs_list in grouped.items(): - if subgroup: - total_weight = sum(t.get("sort_order", 1) for t in tariffs_list) - subgroup_weights[subgroup] = total_weight + subgroup_weights = {} + for subgroup, tariffs_list in grouped.items(): + if subgroup: + total_weight = sum(t.get("sort_order", 1) for t in tariffs_list) + subgroup_weights[subgroup] = total_weight - return {"tariffs": tariffs, "subgroup_weights": subgroup_weights} + return {"tariffs": tariffs, "subgroup_weights": subgroup_weights} - return tariffs - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при получении тарифов: {e}") - await session.rollback() - return [] + return tariffs async def get_tariff_names_groups_subgroups_durations( @@ -133,18 +127,13 @@ async def get_tariff_by_id(session: AsyncSession, tariff_id: int): cached = await cache_get(key) if isinstance(cached, dict): return cached - try: - result = await session.execute(select(Tariff).where(Tariff.id == tariff_id)) - tariff = result.scalar_one_or_none() - if not tariff: - return None - row = _row_to_cache_dict(dict(tariff.__dict__)) - await cache_set(key, row, TARIFF_BY_ID_CACHE_TTL_SEC) - return row - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при получении тарифа по ID {tariff_id}: {e}") - await session.rollback() + result = await session.execute(select(Tariff).where(Tariff.id == tariff_id)) + tariff = result.scalar_one_or_none() + if not tariff: return None + row = _row_to_cache_dict(dict(tariff.__dict__)) + await cache_set(key, row, TARIFF_BY_ID_CACHE_TTL_SEC) + return row async def get_tariff_group_codes(session: AsyncSession) -> list[str]: @@ -152,6 +141,14 @@ async def get_tariff_group_codes(session: AsyncSession) -> list[str]: return [row[0] for row in result.fetchall() if row[0]] +async def get_active_tariff_by_id(session: AsyncSession, tariff_id: int) -> Tariff | None: + """Возвращает ORM-объект Tariff по id, если тариф активен (is_active=True).""" + result = await session.execute( + select(Tariff).where(Tariff.id == int(tariff_id), Tariff.is_active.is_(True)) + ) + return result.scalar_one_or_none() + + async def get_active_tariffs_by_group_code(session: AsyncSession, group_code: str) -> list[Tariff]: result = await session.execute( select(Tariff).where(Tariff.group_code == group_code, Tariff.is_active.is_(True)).order_by(Tariff.id) @@ -164,105 +161,78 @@ async def get_tariffs_for_cluster(session: AsyncSession, cluster_name: str): cached = await cache_get(key) if isinstance(cached, list): return cached - try: + server_row = await session.execute( + select(Server.tariff_group).where(Server.cluster_name == cluster_name).limit(1) + ) + row = server_row.first() + + if not row: server_row = await session.execute( - select(Server.tariff_group).where(Server.cluster_name == cluster_name).limit(1) + select(Server.tariff_group).where(Server.server_name == cluster_name).limit(1) ) row = server_row.first() - if not row: - server_row = await session.execute( - select(Server.tariff_group).where(Server.server_name == cluster_name).limit(1) - ) - row = server_row.first() - - if not row or not row[0]: - return [] - - group_code = row[0] - result = await session.execute( - select(Tariff) - .where(Tariff.group_code == group_code, Tariff.is_active.is_(True)) - .order_by(Tariff.sort_order, Tariff.id) - ) - rows = [_row_to_cache_dict(dict(r.__dict__)) for r in result.scalars().all()] - await cache_set(key, rows, TARIFFS_FOR_CLUSTER_CACHE_TTL_SEC) - return rows - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при получении тарифов для кластера {cluster_name}: {e}") + if not row or not row[0]: return [] + group_code = row[0] + result = await session.execute( + select(Tariff) + .where(Tariff.group_code == group_code, Tariff.is_active.is_(True)) + .order_by(Tariff.sort_order, Tariff.id) + ) + rows = [_row_to_cache_dict(dict(r.__dict__)) for r in result.scalars().all()] + await cache_set(key, rows, TARIFFS_FOR_CLUSTER_CACHE_TTL_SEC) + return rows + async def create_tariff(session: AsyncSession, data: dict): - try: - data["created_at"] = datetime.utcnow() - data["updated_at"] = datetime.utcnow() + data["created_at"] = datetime.utcnow() + data["updated_at"] = datetime.utcnow() - if "sort_order" not in data: - group_code = data.get("group_code") - if group_code: - result = await session.execute( - select(func.max(Tariff.sort_order)).where( - Tariff.group_code == group_code, Tariff.sort_order.isnot(None) - ) + if "sort_order" not in data: + group_code = data.get("group_code") + if group_code: + result = await session.execute( + select(func.max(Tariff.sort_order)).where( + Tariff.group_code == group_code, Tariff.sort_order.isnot(None) ) - max_order = result.scalar() or 0 - else: - result = await session.execute(select(func.max(Tariff.sort_order)).where(Tariff.sort_order.isnot(None))) - max_order = result.scalar() or 0 + ) + max_order = result.scalar() or 0 + else: + result = await session.execute(select(func.max(Tariff.sort_order)).where(Tariff.sort_order.isnot(None))) + max_order = result.scalar() or 0 - data["sort_order"] = max_order + 1 + data["sort_order"] = max_order + 1 - stmt = insert(Tariff).values(**data).returning(Tariff) - result = await session.execute(stmt) - await session.commit() - await _invalidate_tariff_cache() - return result.scalar_one() - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при создании тарифа: {e}") - await session.rollback() - return None + stmt = insert(Tariff).values(**data).returning(Tariff) + result = await session.execute(stmt) + await _invalidate_tariff_cache() + return result.scalar_one() async def update_tariff(session: AsyncSession, tariff_id: int, updates: dict): if not updates: return False - try: - updates["updated_at"] = datetime.utcnow() - await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(**updates)) - await session.commit() - await _invalidate_tariff_cache(tariff_id) - return True - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при обновлении тарифа ID={tariff_id}: {e}") - await session.rollback() - return False + updates["updated_at"] = datetime.utcnow() + await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(**updates)) + await _invalidate_tariff_cache(tariff_id) + return True async def delete_tariff(session: AsyncSession, tariff_id: int): - try: - await session.execute(delete(Tariff).where(Tariff.id == tariff_id)) - await session.commit() - await _invalidate_tariff_cache(tariff_id) - return True - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при удалении тарифа ID={tariff_id}: {e}") - await session.rollback() - return False + await session.execute(delete(Tariff).where(Tariff.id == tariff_id)) + await _invalidate_tariff_cache(tariff_id) + return True async def check_tariff_exists(session: AsyncSession, tariff_id: int): - try: - result = await session.execute(select(Tariff).where(Tariff.id == tariff_id, Tariff.is_active.is_(True))) - tariff = result.scalar_one_or_none() - if tariff: - return True - logger.warning(f"[TARIFF] Тариф {tariff_id} не найден в БД") - return False - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при проверке тарифа {tariff_id}: {e}") - await session.rollback() - return False + result = await session.execute(select(Tariff).where(Tariff.id == tariff_id, Tariff.is_active.is_(True))) + tariff = result.scalar_one_or_none() + if tariff: + return True + logger.warning(f"[TARIFF] Тариф {tariff_id} не найден в БД") + return False async def get_vless_enabled(session: AsyncSession, tariff_id: int | None) -> bool: @@ -292,90 +262,59 @@ async def get_vless_enabled_batch( async def get_tariff_sort_order(session: AsyncSession, tariff_id: int) -> int: - try: - result = await session.execute(select(Tariff.sort_order).where(Tariff.id == tariff_id)) - sort_order = result.scalar_one_or_none() + result = await session.execute(select(Tariff.sort_order).where(Tariff.id == tariff_id)) + sort_order = result.scalar_one_or_none() - if sort_order is None: - await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=1)) - await session.commit() - await _invalidate_tariff_cache(tariff_id) - return 1 + if sort_order is None: + await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=1)) + await _invalidate_tariff_cache(tariff_id) + return 1 - return sort_order - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при получении sort_order для тарифа {tariff_id}: {e}") - await session.rollback() - return None + return sort_order async def move_tariff_up(session: AsyncSession, tariff_id: int) -> bool: - try: - current_order = await get_tariff_sort_order(session, tariff_id) - new_order = max(1, current_order - 1) + current_order = await get_tariff_sort_order(session, tariff_id) + new_order = max(1, current_order - 1) - await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=new_order)) - await session.commit() - await _invalidate_tariff_cache(tariff_id) - return True - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при перемещении тарифа {tariff_id} вверх: {e}") - await session.rollback() - return False + await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=new_order)) + await _invalidate_tariff_cache(tariff_id) + return True async def move_tariff_down(session: AsyncSession, tariff_id: int) -> bool: - try: - current_order = await get_tariff_sort_order(session, tariff_id) - new_order = current_order + 1 + current_order = await get_tariff_sort_order(session, tariff_id) + new_order = current_order + 1 - await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=new_order)) - await session.commit() - await _invalidate_tariff_cache(tariff_id) - return True - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при перемещении тарифа {tariff_id} вниз: {e}") - await session.rollback() - return False + await session.execute(update(Tariff).where(Tariff.id == tariff_id).values(sort_order=new_order)) + await _invalidate_tariff_cache(tariff_id) + return True async def initialize_tariff_sort_orders(session: AsyncSession, group_code: str) -> bool: - try: - result = await session.execute(select(Tariff).where(Tariff.group_code == group_code).order_by(Tariff.id)) - tariffs = result.scalars().all() + result = await session.execute(select(Tariff).where(Tariff.group_code == group_code).order_by(Tariff.id)) + tariffs = result.scalars().all() - if not tariffs: - return True - - for i, tariff in enumerate(tariffs): - new_sort_order = 1 + i - await session.execute(update(Tariff).where(Tariff.id == tariff.id).values(sort_order=new_sort_order)) - - await session.commit() - await _invalidate_tariff_cache() + if not tariffs: return True - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при инициализации sort_order для группы {group_code}: {e}") - await session.rollback() - return False + + for i, tariff in enumerate(tariffs): + new_sort_order = 1 + i + await session.execute(update(Tariff).where(Tariff.id == tariff.id).values(sort_order=new_sort_order)) + + await _invalidate_tariff_cache() + return True async def initialize_all_tariff_weights(session: AsyncSession) -> bool: - try: - result = await session.execute(select(Tariff).where(Tariff.sort_order.is_(None))) - tariffs_without_weight = result.scalars().all() + result = await session.execute(select(Tariff).where(Tariff.sort_order.is_(None))) + tariffs_without_weight = result.scalars().all() - if not tariffs_without_weight: - return True - - for tariff in tariffs_without_weight: - await session.execute(update(Tariff).where(Tariff.id == tariff.id).values(sort_order=1)) - - await session.commit() - await _invalidate_tariff_cache() + if not tariffs_without_weight: return True - except SQLAlchemyError as e: - logger.error(f"[TARIFF] Ошибка при инициализации весов тарифов: {e}") - await session.rollback() - return False + for tariff in tariffs_without_weight: + await session.execute(update(Tariff).where(Tariff.id == tariff.id).values(sort_order=1)) + + await _invalidate_tariff_cache() + return True diff --git a/database/temporary_data.py b/database/temporary_data.py index db40d8ea..323707bc 100644 --- a/database/temporary_data.py +++ b/database/temporary_data.py @@ -1,35 +1,49 @@ from datetime import datetime -from sqlalchemy import delete, select +from sqlalchemy import delete, or_, select from sqlalchemy.dialects.postgresql import insert -from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession +from database.access.resolution import resolve_user_optional from database.models import TemporaryData from logger import logger -async def create_temporary_data(session: AsyncSession, tg_id: int, state: str, data: dict): - try: - stmt = ( - insert(TemporaryData) - .values(tg_id=tg_id, state=state, data=data, updated_at=datetime.utcnow()) - .on_conflict_do_update( - index_elements=[TemporaryData.tg_id], - set_={"state": state, "data": data, "updated_at": datetime.utcnow()}, +async def create_temporary_data(session: AsyncSession, legacy_user_ref: int, state: str, data: dict): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + raise ValueError(f"user not found for temporary data: {legacy_user_ref}") + ins = insert(TemporaryData).values( + user_id=u.id, + tg_id=u.tg_id, + state=state, + data=data, + updated_at=datetime.utcnow(), + ) + stmt = ins.on_conflict_do_update( + index_elements=[TemporaryData.user_id], + set_={ + "state": state, + "data": data, + "updated_at": datetime.utcnow(), + "tg_id": ins.excluded.tg_id, + }, + ) + await session.execute(stmt) + logger.info(f"📝 Временные данные сохранены для user_id={u.id}") + + +async def get_temporary_data(session: AsyncSession, legacy_user_ref: int) -> dict | None: + u = await resolve_user_optional(session, legacy_user_ref) + if u is not None: + stmt = select(TemporaryData).where( + or_( + TemporaryData.user_id == u.id, + TemporaryData.tg_id == u.tg_id, ) ) - await session.execute(stmt) - await session.commit() - logger.info(f"📝 Временные данные сохранены для {tg_id}") - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при сохранении временных данных для {tg_id}: {e}") - await session.rollback() - raise - - -async def get_temporary_data(session: AsyncSession, tg_id: int) -> dict | None: - stmt = select(TemporaryData).where(TemporaryData.tg_id == tg_id) + else: + stmt = select(TemporaryData).where(TemporaryData.tg_id == legacy_user_ref) result = await session.execute(stmt) row = result.scalar_one_or_none() if row: @@ -37,7 +51,17 @@ async def get_temporary_data(session: AsyncSession, tg_id: int) -> dict | None: return None -async def clear_temporary_data(session: AsyncSession, tg_id: int): - await session.execute(delete(TemporaryData).where(TemporaryData.tg_id == tg_id)) - await session.commit() - logger.info(f"🗑 Временные данные очищены для {tg_id}") +async def clear_temporary_data(session: AsyncSession, legacy_user_ref: int): + u = await resolve_user_optional(session, legacy_user_ref) + if u is not None: + await session.execute( + delete(TemporaryData).where( + or_( + TemporaryData.user_id == u.id, + TemporaryData.tg_id == u.tg_id, + ) + ) + ) + else: + await session.execute(delete(TemporaryData).where(TemporaryData.tg_id == legacy_user_ref)) + logger.info(f"🗑 Временные данные очищены для {legacy_user_ref}") diff --git a/database/tracking_sources.py b/database/tracking_sources.py index 17f9da9d..d0bf2042 100644 --- a/database/tracking_sources.py +++ b/database/tracking_sources.py @@ -1,5 +1,4 @@ -from sqlalchemy import and_, func, insert, not_, select -from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy import and_, func, insert, select from sqlalchemy.ext.asyncio import AsyncSession from core.constants import PAYMENT_SYSTEMS_EXCLUDED @@ -8,40 +7,34 @@ from logger import logger async def create_tracking_source(session: AsyncSession, name: str, code: str, type_: str, created_by: int): - try: - stmt = insert(TrackingSource).values( - name=name, - code=code, - type=type_, - created_by=created_by, - ) - await session.execute(stmt) - await session.commit() - logger.info(f"🆕 Источник трафика {code} создан") - except SQLAlchemyError as e: - logger.error(f"❌ Ошибка при создании источника {code}: {e}") - await session.rollback() - raise + stmt = insert(TrackingSource).values( + name=name, + code=code, + type=type_, + created_by=created_by, + ) + await session.execute(stmt) + logger.info(f"🆕 Источник трафика {code} создан") async def get_all_tracking_sources(session: AsyncSession) -> list[dict]: registrations_subq = ( - select(func.count(func.distinct(User.tg_id))) + select(func.count(func.distinct(User.id))) .where(User.source_code == TrackingSource.code) .correlate(TrackingSource) .scalar_subquery() ) trials_subq = ( - select(func.count(func.distinct(User.tg_id))) + select(func.count(func.distinct(User.id))) .where((User.source_code == TrackingSource.code) & (User.trial == 1)) .correlate(TrackingSource) .scalar_subquery() ) payments_subq = ( - select(func.count(func.distinct(Payment.tg_id))) - .join(User, Payment.tg_id == User.tg_id) + select(func.count(func.distinct(Payment.user_id))) + .join(User, Payment.user_id == User.id) .where( (User.source_code == TrackingSource.code) & (Payment.status == "success") @@ -89,20 +82,20 @@ async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict | _src_name, _src_code, created_at = src reg_subq = ( - select(func.count(func.distinct(User.tg_id))) + select(func.count(func.distinct(User.id))) .where((User.source_code == code) & (User.created_at >= created_at)) .scalar_subquery() ) trial_subq = ( - select(func.count(func.distinct(User.tg_id))) + select(func.count(func.distinct(User.id))) .where((User.source_code == code) & (User.trial == 1) & (User.created_at >= created_at)) .scalar_subquery() ) payments_subq = ( - select(func.count(func.distinct(Payment.tg_id))) - .join(User, Payment.tg_id == User.tg_id) + select(func.count(func.distinct(Payment.user_id))) + .join(User, Payment.user_id == User.id) .where( (User.source_code == code) & (Payment.status == "success") @@ -114,7 +107,7 @@ async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict | amount_subq = ( select(func.coalesce(func.sum(Payment.amount), 0.0)) - .join(User, Payment.tg_id == User.tg_id) + .join(User, Payment.user_id == User.id) .where( (User.source_code == code) & (Payment.status == "success") @@ -141,11 +134,11 @@ async def get_tracking_source_stats(session: AsyncSession, code: str) -> dict | payments_base = ( select( - Payment.tg_id.label("tg_id"), + Payment.user_id.label("tg_id"), Payment.amount.label("amount"), Payment.created_at.label("dt"), ) - .join(User, Payment.tg_id == User.tg_id) + .join(User, Payment.user_id == User.id) .where( (User.source_code == code) & (Payment.status == "success") diff --git a/database/users.py b/database/users.py index 0f2d2e80..8b06fc32 100644 --- a/database/users.py +++ b/database/users.py @@ -1,8 +1,7 @@ from datetime import datetime -from sqlalchemy import delete, exists, func, or_, select, update +from sqlalchemy import delete, func, or_, select, update from sqlalchemy.dialects.postgresql import insert -from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession from core.cache_config import ( @@ -11,6 +10,7 @@ from core.cache_config import ( USER_SNAPSHOT_CACHE_TTL_SEC, ) from core.redis_cache import cache_delete, cache_get, cache_key, cache_set +from database.access.resolution import resolve_user_optional from database.models import ( BlockedUser, CouponUsage, @@ -45,36 +45,28 @@ async def add_user( language_code: str = None, is_bot: bool = False, source_code: str = None, - commit: bool = True, -) -> bool: - try: - stmt = ( - insert(User) - .values( - tg_id=tg_id, - username=username, - first_name=first_name, - last_name=last_name, - language_code=language_code, - is_bot=is_bot, - source_code=source_code, - ) - .on_conflict_do_nothing(index_elements=["tg_id"]) - .returning(User.tg_id) +) -> int | None: + stmt = ( + insert(User) + .values( + tg_id=tg_id, + username=username, + first_name=first_name, + last_name=last_name, + language_code=language_code, + is_bot=is_bot, + source_code=source_code, ) - res = await session.execute(stmt) - inserted_tg_id = res.scalar_one_or_none() - if inserted_tg_id is None: - return False - if commit: - await session.commit() - await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC) - logger.info(f"[DB] Новый пользователь добавлен: {tg_id} (source: {source_code})") - return True - except SQLAlchemyError as e: - logger.error(f"[DB] Ошибка при добавлении пользователя {tg_id}: {e}") - await session.rollback() - raise + .on_conflict_do_nothing(index_elements=["tg_id"]) + .returning(User.id) + ) + res = await session.execute(stmt) + inserted_id = res.scalar_one_or_none() + if inserted_id is None: + return None + await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC) + logger.info(f"[DB] Новый пользователь добавлен: tg_id={tg_id} id={inserted_id} (source: {source_code})") + return int(inserted_id) async def invalidate_balance_cache(tg_id: int) -> None: @@ -87,109 +79,130 @@ async def invalidate_profile_cache(tg_id: int) -> None: async def update_balance( session: AsyncSession, - tg_id: int, + legacy_user_ref: int, amount: float, ) -> None: - try: - amount = float(amount) - res = await session.execute( - update(User) - .where(User.tg_id == tg_id) - .values(balance=func.coalesce(User.balance, 0) + amount) - .returning(User.balance) - ) - new_balance = res.scalar_one_or_none() - if new_balance is not None: - old_balance = new_balance - amount - await session.commit() - if new_balance is not None: - logger.info(f"[DB] Баланс пользователя {tg_id} обновлён: {old_balance} → {new_balance}") - else: - logger.info(f"[DB] Баланс пользователя {tg_id} не изменён: пользователь не найден") - await invalidate_balance_cache(tg_id) - await invalidate_profile_cache(tg_id) - except SQLAlchemyError as e: - logger.error(f"[DB] Ошибка при обновлении баланса пользователя {tg_id}: {e}") - await session.rollback() + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + logger.info(f"[DB] Баланс не изменён: пользователь {legacy_user_ref} не найден") + return + uid = u.id + amount = float(amount) + res = await session.execute( + update(User) + .where(User.id == uid) + .values(balance=func.coalesce(User.balance, 0) + amount) + .returning(User.balance) + ) + new_balance = res.scalar_one_or_none() + if new_balance is not None: + old_balance = new_balance - amount + logger.info(f"[DB] Баланс пользователя id={uid} обновлён: {old_balance} → {new_balance}") + else: + logger.info(f"[DB] Баланс пользователя id={uid} не изменён: пользователь не найден") + await invalidate_balance_cache(uid) + await invalidate_profile_cache(uid) + if u.tg_id is not None: + await invalidate_balance_cache(u.tg_id) + await invalidate_profile_cache(u.tg_id) -async def check_user_exists(session: AsyncSession, tg_id: int) -> bool: - cached = await cache_get(cache_key("user_exists", tg_id)) +async def check_user_exists(session: AsyncSession, legacy_user_ref: int) -> bool: + cached = await cache_get(cache_key("user_exists", legacy_user_ref)) if isinstance(cached, bool): return cached - stmt = select(exists().where(User.tg_id == tg_id)) - result = await session.execute(stmt) - value = result.scalar() - await cache_set(cache_key("user_exists", tg_id), bool(value), USER_EXISTS_CACHE_TTL_SEC) + u = await resolve_user_optional(session, legacy_user_ref) + value = u is not None + await cache_set(cache_key("user_exists", legacy_user_ref), value, USER_EXISTS_CACHE_TTL_SEC) return value -async def get_balance(session: AsyncSession, tg_id: int) -> float: - cached = await cache_get(cache_key("balance", tg_id)) +async def get_balance(session: AsyncSession, legacy_user_ref: int) -> float: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return 0.0 + uid = u.id + cached = await cache_get(cache_key("balance", uid)) if cached is not None: try: return round(float(cached), 1) except (TypeError, ValueError): pass - result = await session.execute(select(func.coalesce(User.balance, 0.0)).where(User.tg_id == tg_id)) + result = await session.execute(select(func.coalesce(User.balance, 0.0)).where(User.id == uid)) balance = result.scalar_one_or_none() value = round(float(balance or 0.0), 1) - await cache_set(cache_key("balance", tg_id), value, BALANCE_CACHE_TTL_SEC) + await cache_set(cache_key("balance", uid), value, BALANCE_CACHE_TTL_SEC) return value async def set_user_balance( session: AsyncSession, - tg_id: int, + legacy_user_ref: int, balance: float, ) -> None: - try: - old_balance_result = await session.execute(select(func.coalesce(User.balance, 0.0)).where(User.tg_id == tg_id)) - old_balance = old_balance_result.scalar_one_or_none() - if old_balance is None: - await session.execute(update(User).where(User.tg_id == tg_id).values(balance=balance)) - await session.commit() - await invalidate_balance_cache(tg_id) - await invalidate_profile_cache(tg_id) - return - - balance = float(balance) - await session.execute(update(User).where(User.tg_id == tg_id).values(balance=balance)) - await session.commit() - await invalidate_balance_cache(tg_id) - await invalidate_profile_cache(tg_id) - except SQLAlchemyError as e: - logger.error(f"Ошибка при установке баланса для пользователя {tg_id}: {e}") - await session.rollback() - raise + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return + uid = u.id + balance = float(balance) + await session.execute(update(User).where(User.id == uid).values(balance=balance)) + await invalidate_balance_cache(uid) + await invalidate_profile_cache(uid) -async def update_trial(session: AsyncSession, tg_id: int, status: int): - try: - await session.execute(update(User).where(User.tg_id == tg_id).values(trial=status)) - await session.commit() - await invalidate_profile_cache(tg_id) - invalidate_user_snapshot(tg_id) - logger.info(f"[DB] Триал статус обновлён для пользователя {tg_id}: {status}") - except SQLAlchemyError as e: - logger.error(f"[DB] Ошибка при обновлении триала пользователя {tg_id}: {e}") - await session.rollback() - raise +async def get_user_preferred_currency(session: AsyncSession, tg_id: int) -> str | None: + """Предпочитаемая валюта пользователя по ``tg_id``, если установлена.""" + result = await session.execute( + select(User.preferred_currency).where(User.tg_id == int(tg_id)) + ) + return result.scalar() -async def get_trial(session: AsyncSession, tg_id: int) -> int: - result = await session.execute(select(func.coalesce(User.trial, 0)).where(User.tg_id == tg_id)) +async def mark_trial_started_if_eligible(session: AsyncSession, tg_id: int) -> None: + """Переводит `trial` в 1, только если текущее значение in [0, -1] (пользователь + ещё не использовал триал). Условный update без пред-чтения — атомарно на уровне БД. + + Используется в `services.operations.creation.create_key_on_cluster` после + успешного создания ключа. + """ + await session.execute( + update(User).where(User.tg_id == tg_id, User.trial.in_([0, -1])).values(trial=1) + ) + + +async def update_trial(session: AsyncSession, legacy_user_ref: int, status: int): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return + uid = u.id + await session.execute(update(User).where(User.id == uid).values(trial=status)) + await invalidate_profile_cache(uid) + invalidate_user_snapshot(uid) + if u.tg_id is not None: + await invalidate_profile_cache(u.tg_id) + invalidate_user_snapshot(u.tg_id) + logger.info(f"[DB] Триал статус обновлён для пользователя id={uid}: {status}") + + +async def get_trial(session: AsyncSession, legacy_user_ref: int) -> int: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return 0 + result = await session.execute(select(func.coalesce(User.trial, 0)).where(User.id == u.id)) trial = result.scalar_one_or_none() return int(trial or 0) -async def get_balance_and_trial(session: AsyncSession, tg_id: int) -> tuple[float, int]: +async def get_balance_and_trial(session: AsyncSession, legacy_user_ref: int) -> tuple[float, int]: """Один запрос к БД для баланса и триала (профиль при промахе кэша).""" + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return 0.0, 0 result = await session.execute( select( func.coalesce(User.balance, 0.0), func.coalesce(User.trial, 0), - ).where(User.tg_id == tg_id) + ).where(User.id == u.id) ) row = result.one_or_none() if row is None: @@ -198,18 +211,21 @@ async def get_balance_and_trial(session: AsyncSession, tg_id: int) -> tuple[floa return round(float(balance or 0.0), 1), int(trial or 0) -async def get_balance_trial_key_count(session: AsyncSession, tg_id: int) -> tuple[float, int, int]: +async def get_balance_trial_key_count(session: AsyncSession, legacy_user_ref: int) -> tuple[float, int, int]: """ Один запрос: баланс, триал и число ключей пользователя (для профиля при промахе кэша). Возвращает (balance_rub, trial_status, key_count). """ - key_count_subq = select(func.count()).select_from(Key).where(Key.tg_id == User.tg_id).scalar_subquery() + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return 0.0, 0, 0 + key_count_subq = select(func.count()).select_from(Key).where(Key.user_id == User.id).scalar_subquery() result = await session.execute( select( func.coalesce(User.balance, 0.0), func.coalesce(User.trial, 0), key_count_subq, - ).where(User.tg_id == tg_id) + ).where(User.id == u.id) ) row = result.one_or_none() if row is None: @@ -233,117 +249,131 @@ async def upsert_user( only_if_exists: bool = False, ) -> dict | None: """Создаёт пользователя или обновляет поля профиля.""" - try: - now = datetime.utcnow() - returning_cols = list(User.__table__.c) + now = datetime.utcnow() + returning_cols = list(User.__table__.c) - if only_if_exists: - username_value = username if username else User.username - first_name_value = first_name if first_name else User.first_name - last_name_value = last_name if last_name else User.last_name - language_code_value = language_code if language_code else User.language_code - - res = await session.execute( - update(User) - .where(User.tg_id == tg_id) - .values( - username=username_value, - first_name=first_name_value, - last_name=last_name_value, - language_code=language_code_value, - is_bot=is_bot, - updated_at=now, - ) - .returning(*returning_cols) - ) - row = res.mappings().one_or_none() - if row is None: - return None - await session.commit() - await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC) - return dict(row) + if only_if_exists: + username_value = username if username else User.username + first_name_value = first_name if first_name else User.first_name + last_name_value = last_name if last_name else User.last_name + language_code_value = language_code if language_code else User.language_code res = await session.execute( - insert(User) + update(User) + .where(User.tg_id == tg_id) .values( - tg_id=tg_id, - username=username, - first_name=first_name, - last_name=last_name, - language_code=language_code, + username=username_value, + first_name=first_name_value, + last_name=last_name_value, + language_code=language_code_value, is_bot=is_bot, - created_at=now, updated_at=now, ) - .on_conflict_do_update( - index_elements=[User.tg_id], - set_={ - "username": username, - "first_name": first_name, - "last_name": last_name, - "language_code": language_code, - "is_bot": is_bot, - "updated_at": now, - }, - ) .returning(*returning_cols) ) - row = res.mappings().one() - await session.commit() + row = res.mappings().one_or_none() + if row is None: + return None await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC) return dict(row) - except SQLAlchemyError as e: - logger.error(f"[DB] Ошибка при UPSERT пользователя {tg_id}: {e}") - await session.rollback() - raise - -async def delete_user_data(session: AsyncSession, tg_id: int): - try: - from database.keys import delete_key - - await session.execute(delete(Notification).where(Notification.tg_id == tg_id)) - await session.execute( - delete(GiftUsage).where(GiftUsage.gift_id.in_(select(Gift.gift_id).where(Gift.sender_tg_id == tg_id))) + res = await session.execute( + insert(User) + .values( + tg_id=tg_id, + username=username, + first_name=first_name, + last_name=last_name, + language_code=language_code, + is_bot=is_bot, + created_at=now, + updated_at=now, ) - await session.execute(delete(Gift).where(Gift.sender_tg_id == tg_id)) - await session.execute(update(Gift).where(Gift.recipient_tg_id == tg_id).values(recipient_tg_id=None)) - await session.execute(delete(Payment).where(Payment.tg_id == tg_id)) - await session.execute( - delete(Referral).where(or_(Referral.referrer_tg_id == tg_id, Referral.referred_tg_id == tg_id)) + .on_conflict_do_update( + index_elements=[User.tg_id], + set_={ + "username": username, + "first_name": first_name, + "last_name": last_name, + "language_code": language_code, + "is_bot": is_bot, + "updated_at": now, + }, ) - await session.execute(delete(CouponUsage).where(CouponUsage.user_id == tg_id)) - await delete_key(session, tg_id, commit=False) - await session.execute(delete(TemporaryData).where(TemporaryData.tg_id == tg_id)) - await session.execute(delete(BlockedUser).where(BlockedUser.tg_id == tg_id)) - await session.execute(delete(User).where(User.tg_id == tg_id)) - await session.commit() - logger.info(f"[DB] Данные пользователя {tg_id} полностью удалены") - except SQLAlchemyError as e: - await session.rollback() - logger.error(f"[DB] Ошибка при удалении данных пользователя {tg_id}: {e}") - raise + .returning(*returning_cols) + ) + row = res.mappings().one() + await cache_set(cache_key("user_exists", tg_id), True, USER_EXISTS_CACHE_TTL_SEC) + return dict(row) -async def mark_trial_extended(tg_id: int, session: AsyncSession): - await session.execute(update(User).where(User.tg_id == tg_id).values(trial=-1)) - await session.commit() - invalidate_user_snapshot(tg_id) +async def delete_user_data(session: AsyncSession, legacy_user_ref: int): + from database.keys import delete_key + + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return + uid = u.id + + await session.execute(delete(Notification).where(Notification.user_id == uid)) + await session.execute( + delete(GiftUsage).where(GiftUsage.gift_id.in_(select(Gift.gift_id).where(Gift.sender_user_id == uid))) + ) + await session.execute(delete(Gift).where(Gift.sender_user_id == uid)) + await session.execute(update(Gift).where(Gift.recipient_user_id == uid).values(recipient_user_id=None)) + await session.execute(delete(Payment).where(Payment.user_id == uid)) + await session.execute( + delete(Referral).where(or_(Referral.referrer_user_id == uid, Referral.referred_user_id == uid)) + ) + await session.execute(delete(CouponUsage).where(CouponUsage.user_id == uid)) + await delete_key(session, uid) + await session.execute( + delete(TemporaryData).where( + or_( + TemporaryData.user_id == uid, + TemporaryData.tg_id == u.tg_id, + ) + ) + ) + await session.execute( + delete(BlockedUser).where( + or_( + BlockedUser.user_id == uid, + BlockedUser.tg_id == u.tg_id, + ) + ) + ) + await session.execute(delete(User).where(User.id == uid)) + logger.info(f"[DB] Данные пользователя id={uid} полностью удалены") -async def get_user_snapshot(session: AsyncSession, tg_id: int) -> tuple[int, int] | None: - cached = await cache_get(cache_key("user_snapshot", tg_id)) +async def mark_trial_extended(legacy_user_ref: int, session: AsyncSession): + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return + await session.execute(update(User).where(User.id == u.id).values(trial=-1)) + invalidate_user_snapshot(u.id) + if u.tg_id is not None: + invalidate_user_snapshot(u.tg_id) + + +async def get_user_snapshot(session: AsyncSession, legacy_user_ref: int) -> tuple[int, int] | None: + u = await resolve_user_optional(session, legacy_user_ref) + if u is None: + return None + uid = u.id + cached = await cache_get(cache_key("user_snapshot", uid)) if isinstance(cached, list) and len(cached) == 2: return (int(cached[0]), int(cached[1])) if isinstance(cached, tuple) and len(cached) == 2: return (int(cached[0]), int(cached[1])) - keys_count_sq = select(func.count(Key.client_id)).where(Key.tg_id == tg_id).scalar_subquery() - res = await session.execute(select(func.coalesce(User.trial, 0), keys_count_sq).where(User.tg_id == tg_id)) + keys_count_sq = select(func.count(Key.client_id)).where(Key.user_id == uid).scalar_subquery() + res = await session.execute(select(func.coalesce(User.trial, 0), keys_count_sq).where(User.id == uid)) row = res.first() if row is None: return None value = (int(row[0]), int(row[1])) - await cache_set(cache_key("user_snapshot", tg_id), [value[0], value[1]], USER_SNAPSHOT_CACHE_TTL_SEC) + await cache_set(cache_key("user_snapshot", uid), [value[0], value[1]], USER_SNAPSHOT_CACHE_TTL_SEC) return value @@ -351,7 +381,6 @@ async def upsert_source_if_empty( session: AsyncSession, tg_id: int, source_code: str, - commit: bool = True, ) -> bool: if not source_code: return False @@ -367,8 +396,4 @@ async def upsert_source_if_empty( ) res = await session.execute(stmt) changed_tg_id = res.scalar_one_or_none() - if changed_tg_id is None: - return False - if commit: - await session.commit() - return True + return changed_tg_id is not None diff --git a/database/web_notifications.py b/database/web_notifications.py new file mode 100644 index 00000000..1e03b01e --- /dev/null +++ b/database/web_notifications.py @@ -0,0 +1,234 @@ +from datetime import UTC, datetime + +from sqlalchemy import delete, func, select, update +from sqlalchemy.dialects.postgresql import insert as pg_insert +from sqlalchemy.ext.asyncio import AsyncSession + +from database.models import User, WebNotification, WebPushSubscription +from logger import logger + + +async def upsert_push_subscription( + session: AsyncSession, + *, + user_id: int, + identity_id: str | None, + endpoint: str, + keys_json: dict, +) -> WebPushSubscription: + """Upsert push subscription by endpoint (unique).""" + stmt = pg_insert(WebPushSubscription).values( + user_id=user_id, + identity_id=identity_id, + endpoint=endpoint, + keys_json=keys_json, + created_at=datetime.now(UTC), + ).on_conflict_do_update( + index_elements=["endpoint"], + set_={ + "user_id": user_id, + "identity_id": identity_id, + "keys_json": keys_json, + "created_at": datetime.now(UTC), + }, + ).returning(WebPushSubscription) + result = await session.execute(stmt) + return result.scalar_one() + + +async def get_push_subscriptions_by_user( + session: AsyncSession, user_id: int, +) -> list[WebPushSubscription]: + result = await session.execute( + select(WebPushSubscription).where(WebPushSubscription.user_id == user_id) + ) + return list(result.scalars().all()) + + +async def get_push_subscriptions_by_identity( + session: AsyncSession, identity_id: str, +) -> list[WebPushSubscription]: + result = await session.execute( + select(WebPushSubscription).where(WebPushSubscription.identity_id == identity_id) + ) + return list(result.scalars().all()) + + +async def delete_push_subscription_by_endpoint( + session: AsyncSession, endpoint: str, +) -> None: + await session.execute( + delete(WebPushSubscription).where(WebPushSubscription.endpoint == endpoint) + ) + + +async def get_notifications_for_identity( + session: AsyncSession, + identity_id: str, + limit: int = 20, + offset: int = 0, +) -> list[WebNotification]: + result = await session.execute( + select(WebNotification) + .where(WebNotification.identity_id == identity_id) + .order_by(WebNotification.created_at.desc()) + .limit(limit) + .offset(offset) + ) + return list(result.scalars().all()) + + +async def count_unread_for_identity( + session: AsyncSession, identity_id: str, +) -> int: + result = await session.execute( + select(func.count()) + .select_from(WebNotification) + .where( + WebNotification.identity_id == identity_id, + WebNotification.read is False, + ) + ) + return result.scalar() or 0 + + +async def mark_all_read_for_identity( + session: AsyncSession, identity_id: str, +) -> int: + result = await session.execute( + update(WebNotification) + .where( + WebNotification.identity_id == identity_id, + WebNotification.read is False, + ) + .values(read=True) + ) + return result.rowcount + + +async def resolve_identity_id_by_tg_id( + session: AsyncSession, tg_id: int, +) -> str | None: + """Resolve identity_id from user's tg_id.""" + result = await session.execute( + select(User.identity_id).where(User.tg_id == tg_id) + ) + return result.scalar_one_or_none() + + +async def create_notification( + session: AsyncSession, + *, + user_id: int, + identity_id: str | None, + type: str = "system", + title: str, + message: str = "", + data: dict | None = None, +) -> WebNotification: + notif = WebNotification( + user_id=user_id, + identity_id=identity_id, + type=type, + title=title, + message=message, + data=data, + ) + session.add(notif) + await session.flush() + return notif + + +def _render_template(template: str, **kwargs: object) -> str: + """Safe format — unknown placeholders stay as-is.""" + try: + return template.format_map( + {k: str(v) for k, v in kwargs.items() if v is not None} + | type("_Defaults", (), {"__missing__": lambda self, k: f"{{{k}}}"})() + ) + except Exception: + return template + + +def _get_web_config_str(key: str, default: str) -> str: + try: + from core.settings.web_config import WEB_CONFIG + val = WEB_CONFIG.get(key) + return str(val).strip() if val else default + except Exception: + return default + + +async def notify_web( + session: AsyncSession, + *, + tg_id: int, + type: str = "system", + title: str | None = None, + message: str | None = None, + data: dict | None = None, + template_vars: dict | None = None, +) -> WebNotification | None: + """Создаёт web-уведомление по tg_id. + + title/message — если None, берутся из WEB_CONFIG шаблонов по type. + template_vars — подстановки в шаблон ({email}, {amount}, {name}, {duration}). + """ + try: + identity_id = await resolve_identity_id_by_tg_id(session, tg_id) + if not identity_id: + return None + + vars_ = template_vars or {} + + type_key_map = { + "payment": ("WEB_NOTIFY_PAYMENT_TITLE", "WEB_NOTIFY_PAYMENT_MESSAGE"), + "key_created": ("WEB_NOTIFY_KEY_CREATED_TITLE", "WEB_NOTIFY_KEY_CREATED_MESSAGE"), + "key_expiry": ("WEB_NOTIFY_KEY_EXPIRY_TITLE", "WEB_NOTIFY_KEY_EXPIRY_MESSAGE"), + "gift_received": ("WEB_NOTIFY_GIFT_TITLE", "WEB_NOTIFY_GIFT_MESSAGE"), + } + title_key, msg_key = type_key_map.get(type, (None, None)) + + resolved_title = title + if resolved_title is None and title_key: + resolved_title = _render_template(_get_web_config_str(title_key, ""), **vars_) + resolved_title = resolved_title or type + + resolved_message = message + if resolved_message is None and msg_key: + resolved_message = _render_template(_get_web_config_str(msg_key, ""), **vars_) + resolved_message = resolved_message or "" + + notif = await create_notification( + session, + user_id=tg_id, + identity_id=identity_id, + type=type, + title=resolved_title, + message=resolved_message, + data=data, + ) + + + try: + from services.web_push import push_enabled, send_push_to_many + if push_enabled(): + subs = await get_push_subscriptions_by_identity(session, identity_id) + if subs: + sub_infos = [ + {"endpoint": s.endpoint, "keys": s.keys_json} + for s in subs + ] + sent = await send_push_to_many( + sub_infos, + title=resolved_title, + body=resolved_message, + url="/dashboard/notifications", + ) + logger.debug("[notify_web] push sent to {}/{} subscriptions", sent, len(sub_infos)) + except Exception as push_err: + logger.warning("[notify_web] push delivery failed: {}", push_err) + + return notif + except Exception: + return None diff --git a/docker-compose.local.yml b/docker-compose.local.yml new file mode 100644 index 00000000..707e0f06 --- /dev/null +++ b/docker-compose.local.yml @@ -0,0 +1,37 @@ +services: + postgres: + image: postgres:16-alpine + container_name: solobot-postgres + restart: unless-stopped + environment: + POSTGRES_DB: solobot + POSTGRES_USER: myuser + POSTGRES_PASSWORD: "6901" + ports: + - "5432:5432" + volumes: + - postgres_data:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U myuser -d solobot"] + interval: 5s + timeout: 5s + retries: 20 + + redis: + image: redis:7-alpine + container_name: solobot-redis + restart: unless-stopped + command: ["redis-server", "--appendonly", "yes", "--maxmemory", "256mb", "--maxmemory-policy", "allkeys-lru"] + ports: + - "6379:6379" + volumes: + - redis_data:/data + healthcheck: + test: ["CMD", "redis-cli", "ping"] + interval: 5s + timeout: 3s + retries: 20 + +volumes: + postgres_data: + redis_data: diff --git a/docker-compose.yml b/docker-compose.yml index 0c238056..338654b9 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -13,6 +13,12 @@ services: volumes: - /:/host:ro - backups_data:/app/backups + healthcheck: + test: ["CMD", "python", "-c", "import urllib.request; urllib.request.urlopen('http://127.0.0.1:3004/api/health')"] + interval: 30s + timeout: 5s + retries: 3 + start_period: 15s redis: image: redis:7-alpine diff --git a/handlers/admin/bans/bans_handler.py b/handlers/admin/bans/bans_handler.py index 1307b0ec..7f41cee3 100644 --- a/handlers/admin/bans/bans_handler.py +++ b/handlers/admin/bans/bans_handler.py @@ -13,7 +13,9 @@ from sqlalchemy.dialects.postgresql import insert as pg_insert from sqlalchemy.ext.asyncio import AsyncSession from database import delete_user_data -from database.models import BlockedUser, Key, ManualBan +from database.models import BlockedUser, Key, ManualBan, User +from database.access.resolution import resolve_user_optional +from database.users import add_user from filters.admin import IsAdminFilter from logger import logger from middlewares.ban_checker import invalidate_ban_cache @@ -83,8 +85,10 @@ async def handle_manual_bans_menu(callback_query: CallbackQuery): async def handle_bans_export(callback_query: CallbackQuery, session: AsyncSession): kb = build_blocked_users_kb() try: - result = await session.execute(select(BlockedUser.tg_id)) - banned_users = result.scalars().all() + result = await session.execute( + select(User.tg_id).join(BlockedUser, BlockedUser.user_id == User.id).where(User.tg_id.isnot(None)) + ) + banned_users = [row[0] for row in result.all()] csv_output = io.StringIO() writer = csv.writer(csv_output) @@ -111,7 +115,7 @@ async def handle_bans_export(callback_query: CallbackQuery, session: AsyncSessio async def handle_bans_delete_banned(callback_query: CallbackQuery, session: AsyncSession): kb = build_blocked_users_kb() try: - stmt = select(BlockedUser.tg_id).outerjoin(Key, BlockedUser.tg_id == Key.tg_id).where(Key.tg_id.is_(None)) + stmt = select(BlockedUser.user_id).outerjoin(Key, BlockedUser.user_id == Key.user_id).where(Key.user_id.is_(None)) result = await session.execute(stmt) blocked_ids = [row[0] for row in result.all()] @@ -141,9 +145,10 @@ async def handle_shadow_bans_export(callback_query: CallbackQuery, session: Asyn kb = build_shadow_bans_kb() try: result = await session.execute( - select(ManualBan.tg_id, ManualBan.banned_at, ManualBan.banned_by, ManualBan.until).where( - ManualBan.reason == "shadow" - ) + select(User.tg_id, ManualBan.user_id, ManualBan.banned_at, ManualBan.banned_by, ManualBan.until) + .select_from(ManualBan) + .join(User, ManualBan.user_id == User.id) + .where(ManualBan.reason == "shadow") ) rows = result.all() @@ -151,8 +156,9 @@ async def handle_shadow_bans_export(callback_query: CallbackQuery, session: Asyn writer = csv.writer(csv_output) writer.writerow(["tg_id", "banned_at", "banned_by", "until"]) - for user in rows: - writer.writerow([user.tg_id, user.banned_at, user.banned_by, user.until]) + for row in rows: + display_id = row.tg_id if row.tg_id is not None else row.user_id + writer.writerow([display_id, row.banned_at, row.banned_by, row.until]) csv_output.seek(0) document = BufferedInputFile(file=csv_output.getvalue().encode("utf-8"), filename="shadow_bans.csv") @@ -173,9 +179,17 @@ async def handle_manual_bans_export(callback_query: CallbackQuery, session: Asyn kb = build_manual_bans_kb() try: result = await session.execute( - select(ManualBan.tg_id, ManualBan.banned_at, ManualBan.reason, ManualBan.until, ManualBan.banned_by).where( - or_(ManualBan.reason != "shadow", ManualBan.reason.is_(None)) + select( + User.tg_id, + ManualBan.user_id, + ManualBan.banned_at, + ManualBan.reason, + ManualBan.until, + ManualBan.banned_by, ) + .select_from(ManualBan) + .join(User, ManualBan.user_id == User.id) + .where(or_(ManualBan.reason != "shadow", ManualBan.reason.is_(None))) ) rows = result.all() @@ -183,8 +197,9 @@ async def handle_manual_bans_export(callback_query: CallbackQuery, session: Asyn writer = csv.writer(csv_output) writer.writerow(["tg_id", "banned_at", "reason", "until", "banned_by"]) - for user in rows: - writer.writerow([user.tg_id, user.banned_at, user.reason, user.until, user.banned_by]) + for row in rows: + display_id = row.tg_id if row.tg_id is not None else row.user_id + writer.writerow([display_id, row.banned_at, row.reason, row.until, row.banned_by]) csv_output.seek(0) document = BufferedInputFile(file=csv_output.getvalue().encode("utf-8"), filename="manual_bans.csv") @@ -246,12 +261,17 @@ async def handle_clear_shadow_bans(callback_query: CallbackQuery, session: Async ) return - tg_ids_result = await session.execute(select(ManualBan.tg_id).where(ManualBan.reason == "shadow")) - tg_ids_to_invalidate = [r[0] for r in tg_ids_result.all()] + tg_ids_result = await session.execute( + select(User.tg_id) + .select_from(ManualBan) + .join(User, ManualBan.user_id == User.id) + .where(ManualBan.reason == "shadow") + ) + tg_to_invalidate = [r[0] for r in tg_ids_result.all() if r[0] is not None] await session.execute(delete(ManualBan).where(ManualBan.reason == "shadow")) await session.commit() - for uid in tg_ids_to_invalidate: - await invalidate_ban_cache(uid) + for tid in tg_to_invalidate: + await invalidate_ban_cache(tid) await callback_query.message.answer( text=f"🗑️ Очищено {total_count} записей теневых банов из базы данных.", @@ -285,13 +305,16 @@ async def handle_clear_manual_bans(callback_query: CallbackQuery, session: Async return tg_ids_result = await session.execute( - select(ManualBan.tg_id).where(or_(ManualBan.reason != "shadow", ManualBan.reason.is_(None))) + select(User.tg_id) + .select_from(ManualBan) + .join(User, ManualBan.user_id == User.id) + .where(or_(ManualBan.reason != "shadow", ManualBan.reason.is_(None))) ) - tg_ids_to_invalidate = [r[0] for r in tg_ids_result.all()] + tg_to_invalidate = [r[0] for r in tg_ids_result.all() if r[0] is not None] await session.execute(delete(ManualBan).where(or_(ManualBan.reason != "shadow", ManualBan.reason.is_(None)))) await session.commit() - for uid in tg_ids_to_invalidate: - await invalidate_ban_cache(uid) + for tid in tg_to_invalidate: + await invalidate_ban_cache(tid) await callback_query.message.answer( text=f"🗑️ Очищено {total_count} записей ручных банов из базы данных.", @@ -344,36 +367,53 @@ async def handle_preemptive_ids_input(message: Message, state: FSMContext, sessi now = datetime.now(timezone.utc) - stmt = ( - pg_insert(ManualBan) - .values([ + rows = [] + cache_tg_ids = [] + for raw_tg in tg_ids: + u = await resolve_user_optional(session, raw_tg) + if u is None: + await add_user(session, raw_tg) + await session.flush() + u = await resolve_user_optional(session, raw_tg) + if u is None: + continue + rows.append( { - "tg_id": tg_id, + "user_id": u.id, + "tg_id": u.tg_id, "reason": "shadow", "banned_by": message.from_user.id, "until": None, "banned_at": now, } - for tg_id in tg_ids - ]) - .on_conflict_do_update( - index_elements=[ManualBan.tg_id], - set_={ - "reason": "shadow", - "until": None, - "banned_by": message.from_user.id, - "banned_at": now, - }, ) + if u.tg_id is not None: + cache_tg_ids.append(u.tg_id) + + if not rows: + await message.answer("❌ Не удалось сопоставить ни одного пользователя.") + await state.clear() + return + + ins = pg_insert(ManualBan).values(rows) + stmt = ins.on_conflict_do_update( + index_elements=[ManualBan.user_id], + set_={ + "tg_id": ins.excluded.tg_id, + "reason": "shadow", + "until": None, + "banned_by": message.from_user.id, + "banned_at": now, + }, ) await session.execute(stmt) await session.commit() - for uid in tg_ids: - await invalidate_ban_cache(uid) + for tid in cache_tg_ids: + await invalidate_ban_cache(tid) await message.answer( - f"✅ Успешно добавлено в теневой бан: {len(tg_ids)} пользователей.", + f"✅ Успешно добавлено в теневой бан: {len(rows)} пользователей.", reply_markup=build_shadow_bans_kb(), ) await state.clear() diff --git a/handlers/admin/clusters/cluster_manage.py b/handlers/admin/clusters/cluster_manage.py index f7c8c96b..8fd47715 100644 --- a/handlers/admin/clusters/cluster_manage.py +++ b/handlers/admin/clusters/cluster_manage.py @@ -10,7 +10,7 @@ from database import get_servers, update_key_expiry from database.models import Key, Server, Tariff from filters.admin import IsAdminFilter from middlewares.session import release_session_early -from handlers.keys.operations import renew_key_in_cluster +from services.operations import renew_key_in_cluster from logger import logger from ..panel.keyboard import build_admin_back_kb @@ -45,7 +45,7 @@ async def handle_clusters_manage( result = await session.execute(select(Server.server_name).where(Server.cluster_name == cluster_name)) server_names = [row[0] for row in result.all()] result = await session.execute( - select(func.count(func.distinct(Key.tg_id))).where( + select(func.count(func.distinct(Key.user_id))).where( (Key.server_id == cluster_name) | (Key.server_id.in_(server_names)) ) ) diff --git a/handlers/admin/clusters/cluster_sync.py b/handlers/admin/clusters/cluster_sync.py index 53e18cdf..364441e7 100644 --- a/handlers/admin/clusters/cluster_sync.py +++ b/handlers/admin/clusters/cluster_sync.py @@ -18,14 +18,14 @@ from config import ( ) from core.bootstrap import MODES_CONFIG from database import get_servers -from database.models import Key, Server, Tariff +from database.models import Key, Server, Tariff, User from filters.admin import IsAdminFilter -from handlers.keys.operations import ( +from services.operations import ( create_client_on_server, create_key_on_cluster, delete_key_from_cluster, ) -from handlers.keys.operations.aggregated_links import make_aggregated_link +from services.operations.aggregated_links import make_aggregated_link from handlers.utils import ALLOWED_GROUP_CODES from logger import logger from panels.remnawave import RemnawaveAPI @@ -215,7 +215,8 @@ async def handle_sync_server( Server.inbound_id, Server.server_name, Server.panel_type, - Key.tg_id, + Key.user_id, + User.tg_id.label("owner_tg_id"), Key.client_id, Key.email, Key.expiry_time, @@ -227,6 +228,7 @@ async def handle_sync_server( Key.current_traffic_limit, ) .join(Key, Server.server_name == Key.server_id) + .join(User, Key.user_id == User.id) .where(Server.server_name == server_name) ) else: @@ -236,7 +238,8 @@ async def handle_sync_server( Server.inbound_id, Server.server_name, Server.panel_type, - Key.tg_id, + Key.user_id, + User.tg_id.label("owner_tg_id"), Key.client_id, Key.email, Key.expiry_time, @@ -248,6 +251,7 @@ async def handle_sync_server( Key.current_traffic_limit, ) .join(Key, Server.cluster_name == Key.server_id) + .join(User, Key.user_id == User.id) .where(Server.server_name == server_name) ) @@ -336,7 +340,7 @@ async def handle_sync_server( success = await remna.update_user( uuid=key["client_id"], expire_at=expire_iso, - telegram_id=key["tg_id"], + telegram_id=int(key.get("owner_tg_id") or 0), email=f"{key['email']}@fake.local", active_user_inbounds=[key["inbound_id"]], traffic_limit_bytes=traffic_limit_bytes, @@ -356,14 +360,14 @@ async def handle_sync_server( cluster_id=cluster_name, email=key["email"], client_id=key["client_id"], - tg_id=key["tg_id"], + tg_id=key["user_id"], remna_link_override=None, plan=tariff, ) await session.execute( update(Key) - .where(Key.tg_id == key["tg_id"], Key.client_id == key["client_id"]) + .where(Key.user_id == key["user_id"], Key.client_id == key["client_id"]) .values(remnawave_link=new_remnawave_link, key=key_value) ) await session.commit() @@ -378,7 +382,7 @@ async def handle_sync_server( await create_key_on_cluster( cluster_id=server_name, - tg_id=key["tg_id"], + tg_id=key["user_id"], client_id=key["client_id"], email=key["email"], expiry_timestamp=key["expiry_time"], @@ -400,7 +404,7 @@ async def handle_sync_server( "inbound_id": key["inbound_id"], "server_name": key["server_name"], }, - key["tg_id"], + int(key.get("owner_tg_id") or 0) or key["user_id"], key["client_id"], key["email"], key["expiry_time"], @@ -448,7 +452,8 @@ async def handle_sync_cluster( return result = await session.execute( select( - Key.tg_id, + Key.user_id, + User.tg_id.label("owner_tg_id"), Key.client_id, Key.email, Key.expiry_time, @@ -459,12 +464,15 @@ async def handle_sync_cluster( Key.selected_traffic_limit, Key.current_device_limit, Key.current_traffic_limit, - ).where(Key.server_id.in_(server_names), Key.is_frozen.is_(False)) + ) + .join(User, Key.user_id == User.id) + .where(Key.server_id.in_(server_names), Key.is_frozen.is_(False)) ) else: result = await session.execute( select( - Key.tg_id, + Key.user_id, + User.tg_id.label("owner_tg_id"), Key.client_id, Key.email, Key.expiry_time, @@ -475,7 +483,9 @@ async def handle_sync_cluster( Key.selected_traffic_limit, Key.current_device_limit, Key.current_traffic_limit, - ).where(Key.server_id == cluster_name, Key.is_frozen.is_(False)) + ) + .join(User, Key.user_id == User.id) + .where(Key.server_id == cluster_name, Key.is_frozen.is_(False)) ) keys_to_sync = result.mappings().all() @@ -590,7 +600,7 @@ async def handle_sync_cluster( success = await remna.update_user( uuid=key["client_id"], expire_at=expire_iso, - telegram_id=key["tg_id"], + telegram_id=int(key.get("owner_tg_id") or 0), email=f"{key['email']}@fake.local", active_user_inbounds=inbound_ids, traffic_limit_bytes=traffic_limit_bytes, @@ -651,7 +661,7 @@ async def handle_sync_cluster( cluster_id=cluster_name, email=key["email"], client_id=key["client_id"], - tg_id=key["tg_id"], + tg_id=key["user_id"], remna_link_override=None, plan=tariff, ) @@ -696,14 +706,14 @@ async def handle_sync_cluster( logger.warning(f"[Sync] Пересоздание {key['email']}") await delete_key_from_cluster(cluster_name, key["email"], key["client_id"], session) await session.execute( - delete(Key).where(Key.tg_id == key["tg_id"], Key.client_id == key["client_id"]) + delete(Key).where(Key.user_id == key["user_id"], Key.client_id == key["client_id"]) ) await session.commit() cluster_id_for_recreate = key["server_id"] if use_country_selection else cluster_name await create_key_on_cluster( cluster_id_for_recreate, - key["tg_id"], + key["user_id"], key["client_id"], key["email"], key["expiry_time"], @@ -776,13 +786,13 @@ async def handle_sync_cluster( await delete_key_from_cluster(cluster_name, key["email"], key["client_id"], session) await session.execute( - delete(Key).where(Key.tg_id == key["tg_id"], Key.client_id == key["client_id"]) + delete(Key).where(Key.user_id == key["user_id"], Key.client_id == key["client_id"]) ) cluster_id_for_recreate = key["server_id"] if use_country_selection else cluster_name await create_key_on_cluster( cluster_id_for_recreate, - key["tg_id"], + key["user_id"], key["client_id"], key["email"], key["expiry_time"], diff --git a/handlers/admin/gifts/gifts_handler.py b/handlers/admin/gifts/gifts_handler.py index 62a8154a..236b24bd 100644 --- a/handlers/admin/gifts/gifts_handler.py +++ b/handlers/admin/gifts/gifts_handler.py @@ -17,6 +17,7 @@ from logger import logger from ..panel.keyboard import AdminPanelCallback from .keyboard import build_admin_gifts_kb, build_gifts_list_kb from handlers.buttons import BACK +from handlers.texts import get_site_gift_link router = Router() @@ -225,7 +226,8 @@ async def view_gift(callback: CallbackQuery, session: AsyncSession): f"ID: {gift.gift_id}\n" f"Срок: {duration_text}\n" f"Активаций: {usage_text}\n" - f"Ссылка для активации:\n
{gift.gift_link}
" + f"Ссылки для активации:\n" + f"
🌐 {get_site_gift_link(gift.gift_id)}\n🤖 {gift.gift_link}
" ) builder = InlineKeyboardBuilder() diff --git a/handlers/admin/management/database.py b/handlers/admin/management/database.py index 122d39cd..65d6f8d3 100644 --- a/handlers/admin/management/database.py +++ b/handlers/admin/management/database.py @@ -1,4 +1,6 @@ import os +import re +import shutil import subprocess import sys import traceback @@ -10,15 +12,27 @@ from aiogram.fsm.context import FSMContext from aiogram.fsm.state import State, StatesGroup from aiogram.types import CallbackQuery, Message -from config import DB_NAME, DB_PASSWORD, DB_USER, PG_HOST, PG_PORT +from config import DB_NAME, DB_PASSWORD, DB_USER, PG_HOST, PG_IN_DOCKER, PG_PORT from core.executor import run_io from filters.admin import IsAdminFilter from logger import logger +from utils.backup import _find_docker_postgres_container + +_PG_IDENT_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") + + +def _safe_pg_identifier(value: str, label: str) -> str: + if not _PG_IDENT_RE.match(value): + raise ValueError(f"Недопустимый PostgreSQL-идентификатор ({label}): {value!r}") + return value from . import router from .keyboard import AdminPanelCallback, build_back_to_db_menu, build_database_kb, build_export_db_sources_kb +DOCKER_POSTGRES_CONTAINER = "solobot-postgres" + + def sync_restore_database( tmp_path: str, db_name: str, @@ -33,46 +47,153 @@ def sync_restore_database( if f.read(5) == b"PGDMP": is_custom_dump = True - try: + use_docker = PG_IN_DOCKER + docker_container = _find_docker_postgres_container() if use_docker else None + + if use_docker and not docker_container: + return False, f"Контейнер PostgreSQL '{DOCKER_POSTGRES_CONTAINER}' не найден или не запущен" + + def _run_admin_psql(sql: str) -> None: + if use_docker: + subprocess.run( + [ + "docker", + "exec", + "-e", + f"PGPASSWORD={db_password}", + docker_container, + "psql", + "-U", + db_user, + "-h", + "127.0.0.1", + "-p", + "5432", + "-d", + "postgres", + "-c", + sql, + ], + check=True, + capture_output=True, + text=True, + ) + return + + if shutil.which("psql") is None: + raise FileNotFoundError("psql не найден на хосте и контейнер PostgreSQL не обнаружен") + + env = os.environ.copy() + env["PGPASSWORD"] = db_password subprocess.run( [ - "sudo", "-u", "postgres", "psql", "-d", "postgres", + "psql", + "-U", + db_user, + "-h", + pg_host, + "-p", + pg_port, + "-d", + "postgres", "-c", - f"SELECT pg_terminate_backend(pid) FROM pg_stat_activity WHERE datname = '{db_name}' AND pid <> pg_backend_pid();", + sql, ], check=True, + capture_output=True, + text=True, + env=env, ) - subprocess.run( - ["sudo", "-u", "postgres", "psql", "-d", "postgres", "-c", f"DROP DATABASE IF EXISTS {db_name};"], - check=True, - ) - subprocess.run( - ["sudo", "-u", "postgres", "psql", "-d", "postgres", "-c", f"CREATE DATABASE {db_name} OWNER {db_user};"], - check=True, + + try: + safe_name = _safe_pg_identifier(db_name, "db_name") + safe_user = _safe_pg_identifier(db_user, "db_user") + _run_admin_psql( + f"SELECT pg_terminate_backend(pid) FROM pg_stat_activity WHERE datname = '{safe_name}' AND pid <> pg_backend_pid();" ) + _run_admin_psql(f"DROP DATABASE IF EXISTS {safe_name};") + _run_admin_psql(f"CREATE DATABASE {safe_name} OWNER {safe_user};") + except ValueError as e: + return False, str(e) except subprocess.CalledProcessError as e: return False, (e.stderr or e.stdout or str(e)) - os.environ["PGPASSWORD"] = db_password try: - if is_custom_dump: - result = subprocess.run( - [ - "pg_restore", f"--dbname={db_name}", "-U", db_user, - "-h", pg_host, "-p", pg_port, "--no-owner", "--exit-on-error", tmp_path, - ], - capture_output=True, - text=True, - ) + if use_docker: + with open(tmp_path, "rb") as dump_file: + if is_custom_dump: + result = subprocess.run( + [ + "docker", + "exec", + "-i", + "-e", + f"PGPASSWORD={db_password}", + docker_container, + "pg_restore", + f"--dbname={db_name}", + "-U", + db_user, + "-h", + "127.0.0.1", + "-p", + "5432", + "--no-owner", + "--exit-on-error", + ], + stdin=dump_file, + capture_output=True, + ) + else: + result = subprocess.run( + [ + "docker", + "exec", + "-i", + "-e", + f"PGPASSWORD={db_password}", + docker_container, + "psql", + "-U", + db_user, + "-h", + "127.0.0.1", + "-p", + "5432", + "-d", + db_name, + ], + stdin=dump_file, + capture_output=True, + ) else: - result = subprocess.run( - ["psql", "-U", db_user, "-h", pg_host, "-p", pg_port, "-d", db_name, "-f", tmp_path], - capture_output=True, - text=True, - ) - return result.returncode == 0, result.stderr or "" - finally: - del os.environ["PGPASSWORD"] + env = os.environ.copy() + env["PGPASSWORD"] = db_password + if is_custom_dump: + if shutil.which("pg_restore") is None: + return False, "pg_restore не найден на хосте и контейнер PostgreSQL не обнаружен" + result = subprocess.run( + [ + "pg_restore", f"--dbname={db_name}", "-U", db_user, + "-h", pg_host, "-p", pg_port, "--no-owner", "--exit-on-error", tmp_path, + ], + capture_output=True, + text=True, + env=env, + ) + else: + if shutil.which("psql") is None: + return False, "psql не найден на хосте и контейнер PostgreSQL не обнаружен" + result = subprocess.run( + ["psql", "-U", db_user, "-h", pg_host, "-p", pg_port, "-d", db_name, "-f", tmp_path], + capture_output=True, + text=True, + env=env, + ) + stderr = result.stderr.decode("utf-8", errors="replace") if isinstance(result.stderr, bytes) else result.stderr + return result.returncode == 0, stderr or "" + except Exception as e: + return False, str(e) class DatabaseState(StatesGroup): diff --git a/handlers/admin/management/file_upload.py b/handlers/admin/management/file_upload.py index 1d97d3c8..f8c81583 100644 --- a/handlers/admin/management/file_upload.py +++ b/handlers/admin/management/file_upload.py @@ -81,7 +81,11 @@ async def handle_admin_file_upload(message: Message, state: FSMContext): base_dir = os.path.abspath(".") await run_io(lambda: os.makedirs(base_dir, exist_ok=True)) - dest_path = os.path.join(base_dir, file_name) + safe_name = os.path.basename(file_name) + dest_path = os.path.join(base_dir, safe_name) + if not os.path.abspath(dest_path).startswith(base_dir + os.sep): + await message.answer("❌ Недопустимое имя файла.") + return try: await message.bot.download(document, destination=dest_path) diff --git a/handlers/admin/management/import_3xui.py b/handlers/admin/management/import_3xui.py index b57673ce..583d0427 100644 --- a/handlers/admin/management/import_3xui.py +++ b/handlers/admin/management/import_3xui.py @@ -6,8 +6,8 @@ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from filters.admin import IsAdminFilter -from database.models import Key -from handlers.keys.operations import update_subscription +from database.models import Key, User +from services.operations import update_subscription from logger import logger from . import router @@ -69,7 +69,12 @@ async def handle_3xui_db_upload(message: Message, state: FSMContext, session: As async def handle_resync_after_import(callback: CallbackQuery, session: AsyncSession): await callback.answer("🔁 Начинаю перевыпуск подписок...") - result = await session.execute(select(Key.tg_id, Key.email)) + result = await session.execute( + select(User.tg_id, Key.email) + .select_from(Key) + .join(User, Key.user_id == User.id) + .where(User.tg_id.isnot(None)) + ) keys = result.all() success = 0 diff --git a/handlers/admin/sender/scheduled_service.py b/handlers/admin/sender/scheduled_service.py index 96d88021..4bbb0c64 100644 --- a/handlers/admin/sender/scheduled_service.py +++ b/handlers/admin/sender/scheduled_service.py @@ -157,8 +157,9 @@ async def execute_broadcast_payload(payload: dict, bot: Bot | None = None) -> di async with async_session_maker() as session: try: await save_blocked_user_ids(session, blocked_ids) - except Exception: - pass + await session.commit() + except Exception as e: + logger.warning("[Broadcast] Ошибка сохранения blocked_ids: {}", e) return { "success": True, "message": "Broadcast completed", @@ -186,6 +187,7 @@ async def execute_scheduled_broadcast(broadcast: ScheduledBroadcast, bot: Bot | async def process_due_scheduled_broadcasts_once(bot: Bot, limit: int = 3) -> int: async with async_session_maker() as session: broadcasts = await claim_due_scheduled_broadcasts(session, limit=limit) + await session.commit() processed = 0 for broadcast in broadcasts: processed += 1 @@ -196,9 +198,11 @@ async def process_due_scheduled_broadcasts_once(bot: Bot, limit: int = 3) -> int await mark_scheduled_broadcast_sent(session, broadcast.id, result) else: await mark_scheduled_broadcast_failed(session, broadcast.id, result.get("message", "Broadcast failed")) + await session.commit() except Exception as exc: async with async_session_maker() as session: await mark_scheduled_broadcast_failed(session, broadcast.id, str(exc)) + await session.commit() return processed diff --git a/handlers/admin/sender/sender_service.py b/handlers/admin/sender/sender_service.py index 51207bf7..1cb01a78 100644 --- a/handlers/admin/sender/sender_service.py +++ b/handlers/admin/sender/sender_service.py @@ -214,6 +214,7 @@ class BroadcastService: else: async with async_session_maker() as session: await save_blocked_user_ids(session, list(self.blocked_users)) + await session.commit() except Exception as e: logger.error(f"❌ Ошибка при сохранении заблокированных пользователей: {e}") if self._session is not None: @@ -334,4 +335,36 @@ class BroadcastService: f"скорость: {avg_speed:.1f} сообщений/сек, время: {total_duration:.1f} сек" ) + await self._create_web_notifications(messages) + return stats + + async def _create_web_notifications(self, messages: list[dict]) -> None: + """Создаёт web-уведомления для всех получателей рассылки.""" + if not messages: + return + try: + from database.web_notifications import notify_web + text = messages[0].get("text", "") + import re + clean = re.sub(r"<[^>]+>", "", text).strip() + lines = clean.split("\n", 1) + title = (lines[0][:120] + "…") if len(lines[0]) > 120 else lines[0] + body = lines[1].strip()[:300] if len(lines) > 1 else "" + + session = self._session + if session is None: + from database import async_session_maker + async with async_session_maker() as session: + for msg in messages: + tg_id = msg.get("tg_id") + if tg_id and tg_id not in self.blocked_users: + await notify_web(session, tg_id=tg_id, type="broadcast", title=title, message=body) + await session.commit() + else: + for msg in messages: + tg_id = msg.get("tg_id") + if tg_id and tg_id not in self.blocked_users: + await notify_web(session, tg_id=tg_id, type="broadcast", title=title, message=body) + except Exception as e: + logger.warning(f"[Broadcast] Ошибка создания web-уведомлений: {e}") diff --git a/handlers/admin/sender/sender_utils.py b/handlers/admin/sender/sender_utils.py index c4b0767b..47c7e9d9 100644 --- a/handlers/admin/sender/sender_utils.py +++ b/handlers/admin/sender/sender_utils.py @@ -4,7 +4,7 @@ import re from datetime import datetime from aiogram.types import InlineKeyboardButton, InlineKeyboardMarkup -from sqlalchemy import distinct, exists, func, not_, select +from sqlalchemy import and_, distinct, exists, func, not_, select from sqlalchemy.ext.asyncio import AsyncSession from core.constants import PAYMENT_SYSTEMS_EXCLUDED @@ -12,11 +12,11 @@ from database.models import BlockedUser, Key, ManualBan, Payment, Server, Tariff from logger import logger -def _not_banned(tg_id_col): +def _not_banned(user_id_col): return ( - ~exists().where(BlockedUser.tg_id == tg_id_col) + ~exists().where(BlockedUser.user_id == user_id_col) & ~exists().where( - ManualBan.tg_id == tg_id_col, + ManualBan.user_id == user_id_col, (ManualBan.until.is_(None)) | (ManualBan.until > datetime.utcnow()), ) ) @@ -30,74 +30,89 @@ async def get_recipients(session: AsyncSession, send_to: str, cluster_name: str if send_to == "subscribed": query = ( select(distinct(User.tg_id)) - .join(Key) + .join(Key, Key.user_id == User.id) .where(Key.expiry_time > now_ms) - .where(_not_banned(User.tg_id)) + .where(User.tg_id.isnot(None)) + .where(_not_banned(User.id)) ) elif send_to == "unsubscribed": - subquery = ( - select(User.tg_id) - .outerjoin(Key, User.tg_id == Key.tg_id) - .group_by(User.tg_id) - .having(func.count(Key.tg_id) == 0) + unsub_base = ( + select(User.id.label("uid"), User.tg_id) + .outerjoin(Key, User.id == Key.user_id) + .group_by(User.id, User.tg_id) + .having(func.count(Key.client_id) == 0) .union_all( - select(User.tg_id) - .join(Key, User.tg_id == Key.tg_id) - .group_by(User.tg_id) + select(User.id.label("uid"), User.tg_id) + .join(Key, User.id == Key.user_id) + .group_by(User.id, User.tg_id) .having(func.max(Key.expiry_time) <= now_ms) ) - ) + ).subquery() query = ( - select(distinct(subquery.c.tg_id)) + select(distinct(unsub_base.c.tg_id)) + .select_from(unsub_base) + .where(unsub_base.c.tg_id.isnot(None)) .where( - ~exists().where(BlockedUser.tg_id == subquery.c.tg_id), + ~exists().where(BlockedUser.user_id == unsub_base.c.uid), ~exists().where( - ManualBan.tg_id == subquery.c.tg_id, + ManualBan.user_id == unsub_base.c.uid, (ManualBan.until.is_(None)) | (ManualBan.until > datetime.utcnow()), ), ) ) elif send_to == "untrial": - subquery = select(Key.tg_id) + key_user_ids = select(Key.user_id).distinct() query = ( select(distinct(User.tg_id)) - .where(~User.tg_id.in_(subquery) & User.trial.in_([0, -1])) - .where(_not_banned(User.tg_id)) + .where(~User.id.in_(key_user_ids) & User.trial.in_([0, -1])) + .where(User.tg_id.isnot(None)) + .where(_not_banned(User.id)) ) elif send_to == "cluster": query = ( select(distinct(User.tg_id)) - .join(Key, User.tg_id == Key.tg_id) + .join(Key, Key.user_id == User.id) .join(Server, Key.server_id == Server.cluster_name) .where(Server.cluster_name == cluster_name) - .where(_not_banned(User.tg_id)) + .where(User.tg_id.isnot(None)) + .where(_not_banned(User.id)) ) elif send_to == "hotleads": - subquery_active_keys = select(Key.tg_id).where(Key.expiry_time > now_ms).distinct() query = ( select(distinct(User.tg_id)) - .join(Payment, User.tg_id == Payment.tg_id) + .join(Payment, User.id == Payment.user_id) .where(Payment.status == "success") .where(Payment.amount > 0) .where(Payment.payment_system.notin_(PAYMENT_SYSTEMS_EXCLUDED)) - .where(not_(exists(subquery_active_keys.where(Key.tg_id == User.tg_id)))) - .where(_not_banned(User.tg_id)) + .where( + not_( + exists( + select(1).select_from(Key).where( + and_(Key.user_id == User.id, Key.expiry_time > now_ms) + ) + ) + ) + ) + .where(User.tg_id.isnot(None)) + .where(_not_banned(User.id)) ) elif send_to == "trial": trial_tariff_subquery = select(Tariff.id).where(Tariff.group_code == "trial") query = ( - select(distinct(Key.tg_id)) + select(distinct(User.tg_id)) + .join(Key, Key.user_id == User.id) .where(Key.tariff_id.in_(trial_tariff_subquery)) - .where(_not_banned(Key.tg_id)) + .where(User.tg_id.isnot(None)) + .where(_not_banned(User.id)) ) else: - query = select(distinct(User.tg_id)).where(_not_banned(User.tg_id)) + query = select(distinct(User.tg_id)).where(User.tg_id.isnot(None)).where(_not_banned(User.id)) result = await session.execute(query) tg_ids = [row[0] for row in result.all()] diff --git a/handlers/admin/settings/__init__.py b/handlers/admin/settings/__init__.py index c7817806..e12f5f46 100644 --- a/handlers/admin/settings/__init__.py +++ b/handlers/admin/settings/__init__.py @@ -9,6 +9,7 @@ from .settings_modes import router as settings_modes_router from .settings_money import router as settings_panels_router from .settings_notifications import router as settings_notifications_router from .settings_tariffs import router as settings_tariffs_router +from .settings_web import router as settings_web_router router = Router(name="admin_settings") @@ -21,3 +22,4 @@ router.include_router(settings_panels_router) router.include_router(settings_notifications_router) router.include_router(settings_modes_router) router.include_router(settings_tariffs_router) +router.include_router(settings_web_router) diff --git a/handlers/admin/settings/keyboard.py b/handlers/admin/settings/keyboard.py index 36a80ac5..ae81fe9c 100644 --- a/handlers/admin/settings/keyboard.py +++ b/handlers/admin/settings/keyboard.py @@ -80,8 +80,12 @@ def build_settings_kb() -> InlineKeyboardMarkup: text="Тарификация", callback_data=AdminPanelCallback(action="settings_tariffs").pack(), ) + builder.button( + text="🌐 Сайт", + callback_data=AdminPanelCallback(action="settings_web").pack(), + ) - builder.adjust(2, 2, 2) + builder.adjust(2, 2, 2, 1) builder.row(build_admin_back_btn()) return builder.as_markup() diff --git a/handlers/admin/settings/settings_cashboxes.py b/handlers/admin/settings/settings_cashboxes.py index 437d9f34..1d6ce6ef 100644 --- a/handlers/admin/settings/settings_cashboxes.py +++ b/handlers/admin/settings/settings_cashboxes.py @@ -5,7 +5,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from core.bootstrap import PAYMENTS_CONFIG, update_payments_config from core.settings.providers_order_config import PROVIDERS_ORDER, update_providers_order from filters.admin import IsAdminFilter -from handlers.payments.providers import PROVIDERS_BASE, _get_effective_order +from services.payments.providers import PROVIDERS_BASE, _get_effective_order from ..panel.keyboard import AdminPanelCallback from .keyboard import PAYMENT_PROVIDER_TITLES, build_providers_order_kb, build_settings_cashboxes_kb diff --git a/handlers/admin/settings/settings_web.py b/handlers/admin/settings/settings_web.py new file mode 100644 index 00000000..21cc61e6 --- /dev/null +++ b/handlers/admin/settings/settings_web.py @@ -0,0 +1,137 @@ +from aiogram import F, Router +from aiogram.fsm.context import FSMContext +from aiogram.fsm.state import State, StatesGroup +from aiogram.types import CallbackQuery, InlineKeyboardButton, Message +from aiogram.utils.keyboard import InlineKeyboardBuilder + +from database import async_session_maker +from core.settings.web_config import WEB_CONFIG, update_web_config +from handlers.buttons import BACK + +from ..panel.keyboard import AdminPanelCallback, build_admin_back_btn + + +router = Router(name="admin_settings_web") + + +class WebSettingsState(StatesGroup): + waiting_for_url = State() + + +def build_settings_web_kb() -> InlineKeyboardBuilder: + builder = InlineKeyboardBuilder() + + enabled = bool(WEB_CONFIG.get("WEB_ENABLED", False)) + url = str(WEB_CONFIG.get("SITE_URL") or "не указан") + + builder.row( + InlineKeyboardButton( + text=f"{'✅' if enabled else '❌'} Сайт {'включён' if enabled else 'выключен'}", + callback_data=AdminPanelCallback(action="settings_web_toggle").pack(), + ) + ) + builder.row( + InlineKeyboardButton( + text=f"🌐 URL: {url}", + callback_data=AdminPanelCallback(action="settings_web_url").pack(), + ) + ) + builder.row(build_admin_back_btn("settings")) + + return builder + + +@router.callback_query(AdminPanelCallback.filter(F.action == "settings_web")) +async def open_web_settings(callback: CallbackQuery) -> None: + enabled = bool(WEB_CONFIG.get("WEB_ENABLED", False)) + url = str(WEB_CONFIG.get("SITE_URL") or "не указан") + + text = ( + "🌐 Настройки веб-сайта\n\n" + f"Статус: {'✅ Включён' if enabled else '❌ Выключен'}\n" + f"URL: {url}\n\n" + "Сайт может работать на отдельном домене и сервере.\n" + "При выключении кнопка «Личный кабинет» скрывается из бота." + ) + await callback.message.edit_text( + text=text, + reply_markup=build_settings_web_kb().as_markup(), + ) + await callback.answer() + + +@router.callback_query(AdminPanelCallback.filter(F.action == "settings_web_toggle")) +async def toggle_web_enabled(callback: CallbackQuery) -> None: + current = bool(WEB_CONFIG.get("WEB_ENABLED", False)) + new_config = dict(WEB_CONFIG) + new_config["WEB_ENABLED"] = not current + + async with async_session_maker() as session: + await update_web_config(session, new_config) + + status = "✅ Сайт включён" if new_config["WEB_ENABLED"] else "❌ Сайт выключен" + await callback.answer(status, show_alert=True) + + enabled = new_config["WEB_ENABLED"] + url = str(new_config.get("SITE_URL") or "не указан") + text = ( + "🌐 Настройки веб-сайта\n\n" + f"Статус: {'✅ Включён' if enabled else '❌ Выключен'}\n" + f"URL: {url}\n\n" + "Сайт может работать на отдельном домене и сервере.\n" + "При выключении кнопка «Личный кабинет» скрывается из бота." + ) + await callback.message.edit_text( + text=text, + reply_markup=build_settings_web_kb().as_markup(), + ) + + +@router.callback_query(AdminPanelCallback.filter(F.action == "settings_web_url")) +async def prompt_web_url(callback: CallbackQuery, state: FSMContext) -> None: + current = str(WEB_CONFIG.get("SITE_URL") or "") + text = ( + "🌐 Введите URL сайта\n\n" + f"Текущий: {current or 'не указан'}\n\n" + "Отправьте полный URL (с https://).\n" + "Пример: https://my-vpn.com\n\n" + "Отправьте - чтобы очистить." + ) + await callback.message.edit_text(text=text) + await state.set_state(WebSettingsState.waiting_for_url) + await callback.answer() + + +@router.message(WebSettingsState.waiting_for_url) +async def set_web_url(message: Message, state: FSMContext) -> None: + url = message.text.strip() if message.text else "" + + if url == "-": + url = "" + elif url and not url.startswith("http"): + await message.answer("❌ URL должен начинаться с http:// или https://") + return + + url = url.rstrip("/") + + new_config = dict(WEB_CONFIG) + new_config["SITE_URL"] = url + + async with async_session_maker() as session: + await update_web_config(session, new_config) + + await state.clear() + + enabled = new_config.get("WEB_ENABLED", False) + display_url = url or "не указан" + text = ( + "🌐 Настройки веб-сайта\n\n" + f"Статус: {'✅ Включён' if enabled else '❌ Выключен'}\n" + f"URL: {display_url}\n\n" + "Сайт может работать на отдельном домене и сервере.\n" + "При выключении кнопка «Личный кабинет» скрывается из бота." + ) + await message.answer( + text=text, + reply_markup=build_settings_web_kb().as_markup(), + ) diff --git a/handlers/admin/users/keyboard.py b/handlers/admin/users/keyboard.py index 5e3ed96d..e1774cc7 100644 --- a/handlers/admin/users/keyboard.py +++ b/handlers/admin/users/keyboard.py @@ -14,7 +14,7 @@ from hooks.hook_buttons import insert_hook_buttons from hooks.hooks import run_hooks from ..panel.keyboard import build_admin_back_btn -from .utils import build_admin_key_ref +from services.users_utils import build_admin_key_ref class AdminUserEditorCallback(CallbackData, prefix="admin_users"): diff --git a/handlers/admin/users/users_balance.py b/handlers/admin/users/users_balance.py index 9953b498..7f26c81f 100644 --- a/handlers/admin/users/users_balance.py +++ b/handlers/admin/users/users_balance.py @@ -8,6 +8,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from database import get_balance, set_user_balance, update_balance from database.models import Payment +from database.access.resolution import resolve_user_optional from database.payments import add_payment from filters.admin import IsAdminFilter from utils.csv_export import export_user_all_payments_csv @@ -54,8 +55,11 @@ async def _render_balance_page( balance = await get_balance(session, tg_id) balance = int(balance or 0) + u = await resolve_user_optional(session, tg_id) + uid = u.id if u is not None else tg_id + total_count_result = await session.execute( - select(func.count()).where(Payment.tg_id == tg_id) + select(func.count()).where(Payment.user_id == uid) ) total = total_count_result.scalar() or 0 @@ -70,7 +74,7 @@ async def _render_balance_page( Payment.status, Payment.payment_id, ) - .where(Payment.tg_id == tg_id) + .where(Payment.user_id == uid) .order_by(Payment.created_at.desc()) .offset(page * 5) .limit(5) diff --git a/handlers/admin/users/users_bans.py b/handlers/admin/users/users_bans.py index da9d506d..d5f25446 100644 --- a/handlers/admin/users/users_bans.py +++ b/handlers/admin/users/users_bans.py @@ -10,6 +10,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from handlers.buttons import BACK from database.models import ManualBan +from database.access.resolution import resolve_user_optional from filters.admin import IsAdminFilter from middlewares.ban_checker import invalidate_ban_cache @@ -60,18 +61,26 @@ async def handle_ban_forever_reason_input(message: Message, state: FSMContext, s user_data = await state.get_data() tg_id = user_data.get("tg_id") + u = await resolve_user_optional(session, tg_id) + if u is None: + await message.answer("❌ Пользователь не найден.") + await state.clear() + return + stmt = ( pg_insert(ManualBan) .values( - tg_id=tg_id, + user_id=u.id, + tg_id=u.tg_id, reason=reason, banned_by=message.from_user.id, until=None, banned_at=datetime.now(timezone.utc), ) .on_conflict_do_update( - index_elements=[ManualBan.tg_id], + index_elements=[ManualBan.user_id], set_={ + "tg_id": u.tg_id, "reason": reason, "until": None, "banned_by": message.from_user.id, @@ -82,7 +91,8 @@ async def handle_ban_forever_reason_input(message: Message, state: FSMContext, s await session.execute(stmt) await session.commit() - await invalidate_ban_cache(tg_id) + if u.tg_id is not None: + await invalidate_ban_cache(u.tg_id) await state.clear() await message.answer( @@ -141,18 +151,25 @@ async def handle_ban_duration_input(message: Message, state: FSMContext, session until = datetime.now(timezone.utc) + timedelta(days=days) + u = await resolve_user_optional(session, tg_id) + if u is None: + await message.answer("❌ Пользователь не найден.") + return + stmt = ( pg_insert(ManualBan) .values( - tg_id=tg_id, + user_id=u.id, + tg_id=u.tg_id, reason=reason, banned_by=message.from_user.id, until=until, banned_at=datetime.now(timezone.utc), ) .on_conflict_do_update( - index_elements=[ManualBan.tg_id], + index_elements=[ManualBan.user_id], set_={ + "tg_id": u.tg_id, "reason": reason, "until": until, "banned_at": datetime.now(timezone.utc), @@ -163,7 +180,8 @@ async def handle_ban_duration_input(message: Message, state: FSMContext, session await session.execute(stmt) await session.commit() - await invalidate_ban_cache(tg_id) + if u.tg_id is not None: + await invalidate_ban_cache(u.tg_id) text = ( f"✅ Пользователь {tg_id} временно забанен до {until:%Y-%m-%d %H:%M} по UTC." @@ -182,18 +200,25 @@ async def handle_ban_duration_input(message: Message, state: FSMContext, session IsAdminFilter(), ) async def handle_ban_shadow(callback: CallbackQuery, callback_data: AdminUserEditorCallback, session: AsyncSession): + u = await resolve_user_optional(session, callback_data.tg_id) + if u is None: + await callback.answer("Пользователь не найден", show_alert=True) + return + stmt = ( pg_insert(ManualBan) .values( - tg_id=callback_data.tg_id, + user_id=u.id, + tg_id=u.tg_id, reason="shadow", banned_by=callback.from_user.id, until=None, banned_at=datetime.now(timezone.utc), ) .on_conflict_do_update( - index_elements=[ManualBan.tg_id], + index_elements=[ManualBan.user_id], set_={ + "tg_id": u.tg_id, "reason": "shadow", "until": None, "banned_by": callback.from_user.id, @@ -203,7 +228,8 @@ async def handle_ban_shadow(callback: CallbackQuery, callback_data: AdminUserEdi ) await session.execute(stmt) await session.commit() - await invalidate_ban_cache(callback_data.tg_id) + if u.tg_id is not None: + await invalidate_ban_cache(u.tg_id) await callback.message.edit_text( text=f"👻 Пользователь {callback_data.tg_id} получил теневой бан.", @@ -220,9 +246,15 @@ async def handle_user_unban( callback_data: AdminUserEditorCallback, session: AsyncSession, ): - await session.execute(delete(ManualBan).where(ManualBan.tg_id == callback_data.tg_id)) + u = await resolve_user_optional(session, callback_data.tg_id) + if u is None: + await callback.answer("Пользователь не найден", show_alert=True) + return + + await session.execute(delete(ManualBan).where(ManualBan.user_id == u.id)) await session.commit() - await invalidate_ban_cache(callback_data.tg_id) + if u.tg_id is not None: + await invalidate_ban_cache(u.tg_id) text = ( f"✅ Пользователь {callback_data.tg_id} разблокирован. Нажмите кнопку ниже для возврата в профиль." diff --git a/handlers/admin/users/users_gifts.py b/handlers/admin/users/users_gifts.py index 24e65c64..2f4832cc 100644 --- a/handlers/admin/users/users_gifts.py +++ b/handlers/admin/users/users_gifts.py @@ -20,7 +20,12 @@ router = Router() async def get_user_gifts(session: AsyncSession, tg_id: int) -> list: - stmt = select(Gift).where(Gift.sender_tg_id == tg_id).order_by(Gift.created_at.desc()) + from database.access.resolution import resolve_user_optional + + u = await resolve_user_optional(session, tg_id) + if u is None: + return [] + stmt = select(Gift).where(Gift.sender_user_id == u.id).order_by(Gift.created_at.desc()) result = await session.execute(stmt) return result.scalars().all() diff --git a/handlers/admin/users/users_hwid.py b/handlers/admin/users/users_hwid.py index 13e20a53..d6b44950 100644 --- a/handlers/admin/users/users_hwid.py +++ b/handlers/admin/users/users_hwid.py @@ -10,7 +10,7 @@ from panels.remnawave_runtime import ( from filters.admin import IsAdminFilter from .keyboard import AdminUserEditorCallback, build_editor_kb, build_hwid_menu_kb -from .utils import resolve_admin_key +from services.users_utils import resolve_admin_key router = Router() diff --git a/handlers/admin/users/users_keys.py b/handlers/admin/users/users_keys.py deleted file mode 100644 index 1ffc115e..00000000 --- a/handlers/admin/users/users_keys.py +++ /dev/null @@ -1,1931 +0,0 @@ -import asyncio -import time -import uuid - -from datetime import datetime, timedelta, timezone -from handlers.buttons import BACK - -import pytz - -from aiogram import F, Router, types -from aiogram.exceptions import TelegramBadRequest -from aiogram.fsm.context import FSMContext -from aiogram.types import CallbackQuery, InlineKeyboardButton, Message -from aiogram.utils.keyboard import InlineKeyboardBuilder -from sqlalchemy.ext.asyncio import AsyncSession - -from config import REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD, REMNAWAVE_TOKEN_LOGIN_ENABLED, USE_COUNTRY_SELECTION -from core.bootstrap import MODES_CONFIG -from database import ( - check_server_name_by_cluster, - delete_key, - delete_user_data, - get_active_tariffs_by_group_code, - get_key_by_email, - get_key_details, - get_keys, - get_server_names, - get_servers, - get_tariff_by_id, - get_tariffs_for_cluster, - mark_key_as_frozen, - mark_key_as_unfrozen, - save_admin_key_config, - update_key_subscription_links, - update_key_expiry, -) -from database.models import Key -from filters.admin import IsAdminFilter -from middlewares.session import release_session_early -from handlers.keys.operations import ( - create_key_on_cluster, - delete_key_from_cluster, - get_user_traffic, - renew_key_in_cluster, - reset_traffic_in_cluster, - toggle_client_on_cluster, - update_subscription, -) -from handlers.utils import generate_random_email, handle_error -from hooks.hook_buttons import insert_hook_buttons -from hooks.processors import process_admin_key_edit_menu -from logger import logger -from panels.remnawave import RemnawaveAPI - -from ..panel.keyboard import AdminPanelCallback, build_admin_back_btn, build_admin_back_kb -from .keyboard import ( - AdminUserEditorCallback, - AdminUserKeyEditorCallback, - build_cluster_selection_kb, - build_editor_kb, - build_key_delete_kb, - build_key_edit_kb, - build_reissue_menu_kb, - build_user_delete_kb, - build_users_key_expiry_kb, - build_users_key_show_kb, -) -from .utils import resolve_admin_key -from .users_states import RenewTariffState, UserEditorState - - -MOSCOW_TZ = pytz.timezone("Europe/Moscow") - -router = Router() - - -async def _resolve_callback_key( - session: AsyncSession, - tg_id: int, - key_ref: str | int | None, -) -> Key | None: - return await resolve_admin_key(session, tg_id, key_ref) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_key_edit"), - IsAdminFilter(), -) -async def handle_key_edit( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback | AdminUserKeyEditorCallback, - session: AsyncSession, - update: bool = False, -): - key_ref = callback_data.data - key_obj = await _resolve_callback_key(session, callback_data.tg_id, key_ref) - - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Информация о подписке не найдена.", - reply_markup=build_editor_kb(callback_data.tg_id), - ) - return - - email = key_obj.email - key_details = await get_key_details(session, email) - is_frozen = bool(key_details.get("is_frozen")) if key_details else bool(getattr(key_obj, "is_frozen", False)) - - key_value = key_obj.key or key_obj.remnawave_link or "—" - alias_part = f" ({key_obj.alias})" if key_obj.alias else "" - - if key_obj.created_at: - created_at_dt = datetime.fromtimestamp(int(key_obj.created_at) / 1000, tz=MOSCOW_TZ) - created_at = created_at_dt.strftime("%d %B %Y года %H:%M") - else: - created_at = "—" - - if is_frozen: - frozen_left_ms = int((key_details or {}).get("expiry_time") or 0) - total_minutes = max(frozen_left_ms // 60000, 0) - days, rem_minutes = divmod(total_minutes, 24 * 60) - hours, minutes = divmod(rem_minutes, 60) - frozen_parts: list[str] = [] - if days: - frozen_parts.append(f"{days} дн.") - if hours: - frozen_parts.append(f"{hours} ч.") - if minutes or not frozen_parts: - frozen_parts.append(f"{minutes} мин.") - expiry_label = "⏳ Остаток:" - expiry_date = " ".join(frozen_parts) - elif key_obj.expiry_time: - expiry_dt = datetime.fromtimestamp(int(key_obj.expiry_time) / 1000, tz=MOSCOW_TZ) - expiry_label = "⏰ Истекает:" - expiry_date = expiry_dt.strftime("%d %B %Y года %H:%M") - else: - expiry_label = "⏰ Истекает:" - expiry_date = "—" - - tariff_name = "—" - subgroup_title = "—" - group_code = "—" - base_devices = None - base_traffic = None - is_configurable = False - if key_obj.tariff_id: - tariff = await get_tariff_by_id(session, key_obj.tariff_id) - if tariff: - tariff_name = tariff.get("name", "—") - subgroup_title = tariff.get("subgroup_title") or "—" - group_code = tariff.get("group_code") or "—" - base_devices = tariff.get("device_limit") - base_traffic = tariff.get("traffic_limit") - is_configurable = bool(tariff.get("configurable")) - - devices_line = "" - traffic_line = "" - if is_configurable: - sel_dev, cur_dev = key_obj.selected_device_limit, key_obj.current_device_limit - if sel_dev is not None or cur_dev is not None: - base_dev = sel_dev if sel_dev is not None else (base_devices if base_devices is not None else cur_dev) - extra = ( - f" + {cur_dev - base_dev} (докуплено)" - if (base_dev is not None and cur_dev is not None and cur_dev > base_dev) - else "" - ) - devices_line = f"📱 Устройства: {base_dev}{extra}\n" - - sel_traf, cur_traf = key_obj.selected_traffic_limit, key_obj.current_traffic_limit - if sel_traf is not None or cur_traf is not None: - base_traf = sel_traf if sel_traf is not None else (base_traffic if base_traffic is not None else cur_traf) - extra = ( - f" + {cur_traf - base_traf} ГБ (докуплено)" - if (base_traf is not None and cur_traf is not None and cur_traf > base_traf) - else "" - ) - traffic_line = f"📊 Трафик: {base_traf} ГБ{extra}\n" - - text = ( - "🔑 Информация о подписке\n\n" - "
" - f"🔗 Ключ{alias_part}: {key_value}\n" - f"📆 Создан: {created_at} (МСК)\n" - f"{'⛔ Статус: отключена\n' if is_frozen else ''}" - f"{expiry_label} {expiry_date}{' (МСК)' if not is_frozen and expiry_date != '—' else ''}\n" - f"🌐 Кластер: {key_obj.server_id or '—'}\n" - f"🆔 ID клиента: {key_obj.tg_id or '—'}\n" - f"🏷️ Тарифная группа: {group_code}\n" - f"📁 Подгруппа: {subgroup_title}\n" - f"📦 Тариф: {tariff_name}\n" - f"{devices_line}" - f"{traffic_line}" - "
" - ) - - if not update or not getattr(callback_data, "edit", False): - kb_key_details = dict(key_obj.__dict__) - kb_key_details["is_frozen"] = is_frozen - kb_markup = build_key_edit_kb(kb_key_details, email, is_configurable=is_configurable, key_ref=str(key_ref)) - kb_builder = InlineKeyboardBuilder.from_markup(kb_markup) - hook_buttons = await process_admin_key_edit_menu( - email=email, - session=session, - client_id=key_obj.client_id, - tg_id=key_obj.tg_id, - ) - kb_builder = insert_hook_buttons(kb_builder, hook_buttons) - try: - await callback_query.message.edit_text( - text=text, - reply_markup=kb_builder.as_markup(), - ) - except TelegramBadRequest as e: - if "message is not modified" not in str(e): - raise - else: - try: - await callback_query.message.edit_text( - text=text, - reply_markup=await build_users_key_expiry_kb( - session, - callback_data.tg_id, - email, - key_ref=str(key_ref), - ), - ) - except TelegramBadRequest as e: - if "message is not modified" not in str(e): - raise - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_expiry_edit"), - IsAdminFilter(), -) -async def handle_change_expiry( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_ref = str(callback_data.data) - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - email = key_obj.email - - await callback_query.message.edit_reply_markup( - reply_markup=await build_users_key_expiry_kb(session, tg_id, email, key_ref=key_ref) - ) - - -@router.callback_query( - AdminUserKeyEditorCallback.filter(F.action == "add"), - IsAdminFilter(), -) -async def handle_expiry_add( - callback_query: CallbackQuery, - callback_data: AdminUserKeyEditorCallback, - state: FSMContext, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_ref = str(callback_data.data) - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - email = key_obj.email - days = callback_data.month - - key_details = await get_key_details(session, email) - - if not key_details: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - - if days: - await change_expiry_time(key_details["expiry_time"] + days * 24 * 3600 * 1000, email, session) - await handle_key_edit(callback_query, callback_data, session, True) - return - - await state.update_data(tg_id=tg_id, email=email, key_ref=key_ref, op_type="add") - await state.set_state(UserEditorState.waiting_for_expiry_time) - - await callback_query.message.edit_text( - text="✍️ Введите количество дней, которое хотите добавить к времени действия ключа:", - reply_markup=build_users_key_show_kb(tg_id, key_ref), - ) - - -@router.callback_query( - AdminUserKeyEditorCallback.filter(F.action == "take"), - IsAdminFilter(), -) -async def handle_expiry_take( - callback_query: CallbackQuery, - callback_data: AdminUserKeyEditorCallback, - state: FSMContext, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_ref = str(callback_data.data) - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - email = key_obj.email - - await state.update_data(tg_id=tg_id, email=email, key_ref=key_ref, op_type="take") - await state.set_state(UserEditorState.waiting_for_expiry_time) - - await callback_query.message.edit_text( - text="✍️ Введите количество дней, которое хотите вычесть из времени действия ключа:", - reply_markup=build_users_key_show_kb(tg_id, key_ref), - ) - - -@router.callback_query( - AdminUserKeyEditorCallback.filter(F.action == "set"), - IsAdminFilter(), -) -async def handle_expiry_set( - callback_query: CallbackQuery, - callback_data: AdminUserKeyEditorCallback, - state: FSMContext, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_ref = str(callback_data.data) - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - email = key_obj.email - - key_details = await get_key_details(session, email) - - if not key_details: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - - await state.update_data(tg_id=tg_id, email=email, key_ref=key_ref, op_type="set") - await state.set_state(UserEditorState.waiting_for_expiry_time) - - text = ( - "✍️ Введите новое время действия ключа:" - "\n\n📌 Формат: год-месяц-день час:минута" - f"\n\n📄 Текущая дата: {datetime.fromtimestamp(key_details['expiry_time'] / 1000, tz=MOSCOW_TZ).strftime('%Y-%m-%d %H:%M')} (МСК)" - ) - - await callback_query.message.edit_text( - text=text, - reply_markup=build_users_key_show_kb(tg_id, key_ref), - ) - - -@router.message(UserEditorState.waiting_for_expiry_time, IsAdminFilter()) -async def handle_expiry_time_input(message: Message, state: FSMContext, session: AsyncSession): - data = await state.get_data() - tg_id = data.get("tg_id") - email = data.get("email") - key_ref = data.get("key_ref") - op_type = data.get("op_type") - - if op_type != "set" and (not message.text.isdigit() or int(message.text) < 0): - await message.answer( - text="🚫 Пожалуйста, введите корректное количество дней!", - reply_markup=build_users_key_show_kb(tg_id, key_ref) if key_ref else build_editor_kb(tg_id), - ) - return - - key_details = await get_key_details(session, email) - - if not key_details: - await message.answer( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - - try: - current_expiry_time = datetime.fromtimestamp( - key_details["expiry_time"] / 1000, - tz=MOSCOW_TZ, - ) - - if op_type == "add": - days = int(message.text) - new_expiry_time = current_expiry_time + timedelta(days=days) - text = f"✅ Ко времени действия ключа добавлено {days} дн." - elif op_type == "take": - days = int(message.text) - new_expiry_time = current_expiry_time - timedelta(days=days) - text = f"✅ Из времени действия ключа вычтено {days} дн." - else: - new_expiry_time = datetime.strptime(message.text, "%Y-%m-%d %H:%M") - new_expiry_time = MOSCOW_TZ.localize(new_expiry_time) - text = f"✅ Время действия ключа изменено на {message.text} (МСК)" - - new_expiry_timestamp = int(new_expiry_time.timestamp() * 1000) - await change_expiry_time(new_expiry_timestamp, email, session) - except ValueError: - text = "🚫 Пожалуйста, используйте корректный формат даты (ГГГГ-ММ-ДД ЧЧ:ММ)!" - except Exception as e: - text = f"❗ Произошла ошибка во время изменения времени действия ключа: {e}" - - await message.answer( - text=text, - reply_markup=build_users_key_show_kb(tg_id, key_ref) if key_ref else build_editor_kb(tg_id), - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_reissue_menu"), - IsAdminFilter(), -) -async def handle_reissue_menu( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_ref = str(callback_data.data) - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - if not key_obj: - await callback_query.message.edit_text("🚫 Ключ не найден.", reply_markup=build_editor_kb(tg_id)) - return - - text = ( - "🔄 Перевыпуск подписки\n\n" - "📦 Полный перевыпуск\n" - "Пересоздаёт подписку на сервере с возможностью выбора кластера. " - "Используйте для переноса на другой сервер или обновления данных.\n\n" - "🔗 Сменить ссылку\n" - "Генерирует новую ссылку подписки. Старая ссылка перестанет работать. " - "Все данные подписки сохранятся." - ) - - await callback_query.message.edit_text( - text=text, - reply_markup=build_reissue_menu_kb(key_ref, tg_id), - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_update_key"), - IsAdminFilter(), -) -async def handle_update_key( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_ref = str(callback_data.data) - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - if not key_obj: - await callback_query.message.edit_text("🚫 Ключ не найден.", reply_markup=build_editor_kb(tg_id)) - return - email = key_obj.email - - await callback_query.message.edit_text( - text=f"📡 Выберите кластер, на котором пересоздать ключ {email}:", - reply_markup=await build_cluster_selection_kb( - session, - tg_id, - key_ref, - action="confirm_admin_key_reissue", - ), - ) - - -@router.callback_query(F.data.startswith("confirm_admin_key_reissue|"), IsAdminFilter()) -async def confirm_admin_key_reissue(callback_query: CallbackQuery, session: AsyncSession, state: FSMContext): - _, tg_id, key_ref, cluster_id = callback_query.data.split("|") - tg_id = int(tg_id) - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - if not key_obj: - await callback_query.message.edit_text("🚫 Ключ не найден.", reply_markup=build_editor_kb(tg_id)) - return - email = key_obj.email - - try: - servers = await get_servers(session) - cluster_servers = servers.get(cluster_id, []) - - tariffs = await get_tariffs_for_cluster(session, cluster_id) - if not tariffs: - builder = InlineKeyboardBuilder() - builder.row( - InlineKeyboardButton( - text="🔗 Привязать тариф", - callback_data=AdminPanelCallback(action="clusters").pack(), - ) - ) - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=AdminUserEditorCallback( - action="users_key_edit", - tg_id=tg_id, - data=key_ref, - ).pack(), - ) - ) - await callback_query.message.edit_text( - f"🚫 Невозможно пересоздать подписку\n\n" - f"📊 Информация о кластере:\n
" - f"🌐 Кластер: {cluster_id}\n" - f"⚠️ Статус: Нет привязанного тарифа\n
" - f"💡 Привяжите тариф к кластеру", - reply_markup=builder.as_markup(), - ) - return - - use_country_selection = bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)) - - if use_country_selection: - unique_countries = {srv["server_name"] for srv in cluster_servers} - await state.update_data(tg_id=tg_id, email=email, key_ref=key_ref, cluster_id=cluster_id) - builder = InlineKeyboardBuilder() - for country in sorted(unique_countries): - builder.button( - text=country, - callback_data=f"admin_reissue_country|{tg_id}|{key_ref}|{country}", - ) - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=AdminUserEditorCallback( - action="users_key_edit", - tg_id=tg_id, - data=key_ref, - ).pack(), - ) - ) - await callback_query.message.edit_text( - "🌍 Выберите сервер (страну) для пересоздания подписки:", - reply_markup=builder.as_markup(), - ) - return - - key_link = await get_key_by_email(session, email) - remnawave_link = key_link.remnawave_link if key_link else None - - await update_subscription( - tg_id, - email, - session, - cluster_override=cluster_id, - remnawave_link=remnawave_link, - ) - - await handle_key_edit( - callback_query, - AdminUserEditorCallback(tg_id=tg_id, data=key_ref, action="view_key"), - session, - True, - ) - except Exception as e: - logger.error(f"Ошибка при перевыпуске ключа {email}: {e}") - await callback_query.message.answer(f"❗ Ошибка: {e}") - - -@router.callback_query(F.data.startswith("admin_reissue_country|"), IsAdminFilter()) -async def admin_reissue_country(callback_query: CallbackQuery, session: AsyncSession, state: FSMContext): - _, tg_id, key_ref, country = callback_query.data.split("|") - tg_id = int(tg_id) - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - if not key_obj: - await callback_query.message.edit_text("🚫 Ключ не найден.", reply_markup=build_editor_kb(tg_id)) - return - email = key_obj.email - - try: - data = await state.get_data() - cluster_id = data.get("cluster_id") - - if cluster_id: - tariffs = await get_tariffs_for_cluster(session, cluster_id) - if not tariffs: - builder = InlineKeyboardBuilder() - builder.row( - InlineKeyboardButton( - text="🔗 Привязать тариф", - callback_data=AdminPanelCallback(action="clusters").pack(), - ) - ) - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=AdminUserEditorCallback( - action="users_key_edit", - tg_id=tg_id, - data=key_ref, - ).pack(), - ) - ) - await callback_query.message.edit_text( - f"🚫 Невозможно пересоздать подписку\n\n" - f"📊 Информация о кластере:\n
" - f"🌐 Кластер: {cluster_id}\n" - f"⚠️ Статус: Нет привязанного тарифа\n
" - f"💡 Привяжите тариф к кластеру", - reply_markup=builder.as_markup(), - ) - return - - key_link = await get_key_by_email(session, email) - remnawave_link = key_link.remnawave_link if key_link else None - - await update_subscription( - tg_id=tg_id, - email=email, - session=session, - country_override=country, - remnawave_link=remnawave_link, - ) - - await handle_key_edit( - callback_query, - AdminUserEditorCallback(tg_id=tg_id, data=key_ref, action="view_key"), - session, - True, - ) - except Exception as e: - logger.error(f"Ошибка при перевыпуске ключа для страны {country}: {e}") - await callback_query.message.answer(f"❗ Ошибка: {e}") - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_recreate_key"), - IsAdminFilter(), -) -async def handle_recreate_key_start( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_ref = str(callback_data.data) - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Ключ не найден.", - reply_markup=build_editor_kb(tg_id), - ) - return - - email = key_obj.email - - tariff_name = "—" - if key_obj.tariff_id: - tariff = await get_tariff_by_id(session, key_obj.tariff_id) - if tariff: - tariff_name = tariff.get("name", "—") - - text = ( - "🔁 Пересоздание ссылки подписки\n\n" - f"📦 Тариф: {tariff_name}\n\n" - "⚠️ Будет сгенерирована новая ссылка подписки.\n" - "Старая ссылка перестанет работать.\n\n" - "✅ Все данные подписки сохранятся." - ) - - builder = InlineKeyboardBuilder() - builder.row( - InlineKeyboardButton( - text="✅ Пересоздать", - callback_data=f"confirm_recreate|{tg_id}|{key_ref}", - ) - ) - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=AdminUserEditorCallback(action="users_key_edit", tg_id=tg_id, data=key_ref).pack(), - ) - ) - - await callback_query.message.edit_text(text=text, reply_markup=builder.as_markup()) - - -@router.callback_query(F.data.startswith("confirm_recreate|"), IsAdminFilter()) -async def handle_recreate_key_confirm( - callback_query: CallbackQuery, - session: AsyncSession, -): - _, tg_id, key_ref = callback_query.data.split("|") - tg_id = int(tg_id) - - try: - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Ключ не найден.", - reply_markup=build_editor_kb(tg_id), - ) - return - - old_email = key_obj.email - - await callback_query.message.edit_text("⏳ Пересоздание ссылки подписки...") - - client_id = key_obj.client_id - cluster_id = key_obj.server_id - old_link = key_obj.remnawave_link or key_obj.key - - servers = await get_servers(session) - cluster = servers.get(cluster_id) - - if not cluster: - for _, server_list in servers.items(): - for server_info in server_list: - if server_info.get("server_name", "").lower() == cluster_id.lower(): - cluster = [server_info] - break - if cluster: - break - - if not cluster: - await callback_query.message.edit_text( - text=f"❗ Кластер {cluster_id} не найден.", - reply_markup=build_editor_kb(tg_id), - ) - return - - remnawave_servers = [s for s in cluster if s.get("panel_type", "3x-ui").lower() == "remnawave"] - - if not remnawave_servers: - await callback_query.message.edit_text( - text="❗ Revoke доступен только для Remnawave. Для 3x-ui используйте перевыпуск.", - reply_markup=build_editor_kb(tg_id), - ) - return - - api_url = remnawave_servers[0].get("api_url") - if not api_url: - await callback_query.message.edit_text( - text="❗ У Remnawave сервера не задан api_url.", - reply_markup=build_editor_kb(tg_id), - ) - return - - api = RemnawaveAPI(api_url) - try: - if not REMNAWAVE_TOKEN_LOGIN_ENABLED: - await api.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD) - - user_data = await api.revoke_user_subscription(client_id) - finally: - await api.aclose() - - if not user_data: - await callback_query.message.edit_text( - text="❗ Не удалось выполнить revoke. Проверьте логи.", - reply_markup=build_editor_kb(tg_id), - ) - return - - new_link = user_data.get("subscriptionUrl") - - if not new_link: - await callback_query.message.edit_text( - text="❗ Revoke выполнен, но новая ссылка не получена.", - reply_markup=build_editor_kb(tg_id), - ) - return - - await update_key_subscription_links(session, old_email, new_link) - - try: - user_text = ( - "🔄 Ваша подписка была перевыпущена\n\n" - f"🔗 Новая ссылка подписки:\n{new_link}\n\n" - "Старая ссылка больше не работает." - ) - user_kb = InlineKeyboardBuilder() - user_kb.row( - InlineKeyboardButton( - text="📱 Мои подписки", - callback_data="view_keys", - ) - ) - user_kb.row( - InlineKeyboardButton( - text="👤 Личный кабинет", - callback_data="profile", - ) - ) - - await callback_query.bot.send_message( - chat_id=tg_id, - text=user_text, - reply_markup=user_kb.as_markup(), - ) - notification_sent = True - except Exception as e: - logger.warning(f"Не удалось отправить уведомление клиенту {tg_id}: {e}") - notification_sent = False - - text = ( - "✅ Ссылка подписки пересоздана\n\n" - f"🔗 Старая ссылка:\n{old_link}\n\n" - f"🔗 Новая ссылка:\n{new_link}\n\n" - ) - if notification_sent: - text += "📨 Клиент уведомлён о новой ссылке." - else: - text += "⚠️ Не удалось уведомить клиента." - - builder = InlineKeyboardBuilder() - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=AdminUserEditorCallback( - action="users_key_edit", - tg_id=tg_id, - data=key_ref, - ).pack(), - ) - ) - - await callback_query.message.edit_text( - text=text, - reply_markup=builder.as_markup(), - ) - - except Exception as e: - logger.error(f"Ошибка при revoke ключа {old_email}: {e}") - await callback_query.message.edit_text( - text=f"❗ Ошибка при пересоздании: {e}", - reply_markup=build_editor_kb(tg_id), - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_delete_key"), - IsAdminFilter(), -) -async def handle_delete_key( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - state: FSMContext, - session: AsyncSession, -): - key_obj = await _resolve_callback_key(session, callback_data.tg_id, callback_data.data) - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Ключ не найден!", - reply_markup=build_editor_kb(callback_data.tg_id), - ) - return - - email = key_obj.email - client_id = key_obj.client_id - - if client_id is None: - await callback_query.message.edit_text( - text="🚫 Ключ не найден!", - reply_markup=build_editor_kb(callback_data.tg_id), - ) - return - - await state.set_state(UserEditorState.confirm_delete_key) - await state.update_data( - delete_key_email=email, - delete_key_tg_id=int(callback_data.tg_id), - delete_key_client_id=client_id, - ) - - await callback_query.message.edit_text( - text="❓ Вы уверены, что хотите удалить ключ?", - reply_markup=build_key_delete_kb(callback_data.tg_id), - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_delete_key_confirm"), - UserEditorState.confirm_delete_key, - IsAdminFilter(), -) -async def handle_delete_key_confirm( - callback_query: types.CallbackQuery, - callback_data: AdminUserEditorCallback, - state: FSMContext, - session: AsyncSession, -): - data = await state.get_data() - email = data.get("delete_key_email") - expected_tg_id = data.get("delete_key_tg_id") - client_id = data.get("delete_key_client_id") - await state.clear() - - if not email or int(expected_tg_id or 0) != int(callback_data.tg_id): - await callback_query.answer("Данные устарели", show_alert=True) - return - - if not client_id: - key_obj = await get_key_by_email(session, email, int(callback_data.tg_id)) - client_id = key_obj.client_id if key_obj else None - - kb = build_editor_kb(callback_data.tg_id) - - if client_id: - clusters = await get_servers(session=session) - await release_session_early(session) - - async def delete_key_from_servers(): - tasks = [] - for cluster_name, cluster_servers in clusters.items(): - for _ in cluster_servers: - tasks.append(delete_key_from_cluster(cluster_name, email, client_id, session)) - await asyncio.gather(*tasks, return_exceptions=True) - - await delete_key_from_servers() - await delete_key(session, client_id) - - await callback_query.message.edit_text(text="✅ Ключ успешно удален.", reply_markup=kb) - else: - await callback_query.message.edit_text( - text="🚫 Ключ не найден или уже удален.", - reply_markup=kb, - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_delete_user"), - IsAdminFilter(), -) -async def handle_delete_user( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, -): - tg_id = callback_data.tg_id - await callback_query.message.edit_text( - text=f"❗️ Вы уверены, что хотите удалить пользователя с ID {tg_id}?", - reply_markup=build_user_delete_kb(tg_id), - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_delete_user_confirm"), - IsAdminFilter(), -) -async def handle_delete_user_confirm( - callback_query: types.CallbackQuery, - callback_data: AdminUserEditorCallback, - session: AsyncSession, -): - tg_id = callback_data.tg_id - - key_records = [(row.email, row.client_id) for row in await get_keys(session, tg_id)] - await release_session_early(session) - - async def delete_keys_from_servers(): - try: - tasks = [] - servers = await get_servers(session=session) - for email, client_id in key_records: - for cluster_id, _cluster in servers.items(): - tasks.append(delete_key_from_cluster(cluster_id, email, client_id, session)) - await asyncio.gather(*tasks, return_exceptions=True) - except Exception as e: - logger.error(f"Ошибка при удалении ключей с серверов для пользователя {tg_id}: {e}") - - await delete_keys_from_servers() - - try: - await delete_user_data(session, tg_id) - await callback_query.message.edit_text( - text=f"🗑️ Пользователь с ID {tg_id} был удален.", - reply_markup=build_admin_back_kb(), - ) - except Exception as e: - logger.error(f"Ошибка при удалении данных из базы данных для пользователя {tg_id}: {e}") - await callback_query.message.edit_text( - text=f"❌ Произошла ошибка при удалении пользователя с ID {tg_id}. Попробуйте снова.", - reply_markup=build_admin_back_kb(), - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_traffic"), - IsAdminFilter(), -) -async def handle_user_traffic( - callback_query: types.CallbackQuery, - callback_data: AdminUserEditorCallback, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_obj = await _resolve_callback_key(session, tg_id, callback_data.data) - if not key_obj: - await callback_query.message.edit_text("❌ Ключ не найден.", reply_markup=build_editor_kb(tg_id)) - return - email = key_obj.email - - await callback_query.message.edit_text("⏳ Получаем данные о трафике, пожалуйста, подождите...") - - traffic_data = await get_user_traffic(session, tg_id, email) - - if traffic_data["status"] == "error": - await callback_query.message.edit_text( - traffic_data["message"], - reply_markup=build_editor_kb(tg_id, True), - ) - return - - total_traffic = 0 - result_text = f"📊 Трафик подписки {email}:\n\n" - - for server, traffic in traffic_data["traffic"].items(): - if isinstance(traffic, str): - result_text += f"❌ {server}: {traffic}\n" - else: - result_text += f"🌍 {server}: {traffic} ГБ\n" - total_traffic += traffic - - result_text += f"\n🔢 Общий трафик: {total_traffic:.2f} ГБ" - - await callback_query.message.edit_text( - result_text, - reply_markup=build_editor_kb(tg_id, True), - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_create_key"), - IsAdminFilter(), -) -async def handle_create_key_start( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - state: FSMContext, - session: AsyncSession, -): - tg_id = callback_data.tg_id - await state.update_data(tg_id=tg_id) - - use_country_selection = bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)) - - if use_country_selection: - await state.set_state(UserEditorState.selecting_country) - - countries = await get_server_names(session) - - if not countries: - await callback_query.message.edit_text( - "❌ Нет доступных стран для создания ключа.", - reply_markup=build_editor_kb(tg_id), - ) - return - - builder = InlineKeyboardBuilder() - for country in countries: - builder.button(text=country, callback_data=country) - builder.adjust(1) - builder.row(build_admin_back_btn()) - - await callback_query.message.edit_text( - "🌍 Выберите страну для создания ключа:", - reply_markup=builder.as_markup(), - ) - return - - await state.set_state(UserEditorState.selecting_cluster) - - servers = await get_servers(session=session) - cluster_names = list(servers.keys()) - - if not cluster_names: - await callback_query.message.edit_text( - "❌ Нет доступных кластеров для создания ключа.", - reply_markup=build_editor_kb(tg_id), - ) - return - - builder = InlineKeyboardBuilder() - for cluster in cluster_names: - builder.button(text=f"🌐 {cluster}", callback_data=cluster) - builder.adjust(2) - builder.row(build_admin_back_btn()) - - await callback_query.message.edit_text( - "🌐 Выберите кластер для создания ключа:", - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(UserEditorState.selecting_country, IsAdminFilter()) -async def handle_create_key_country(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - country = callback_query.data - await state.update_data(country=country) - await state.set_state(UserEditorState.selecting_duration) - - builder = InlineKeyboardBuilder() - - cluster_info = await check_server_name_by_cluster(session, country) - - if not cluster_info: - await callback_query.message.edit_text("❌ Сервер не найден.") - return - - cluster_name = cluster_info["cluster_name"] - await state.update_data(cluster_name=cluster_name) - - tariffs = await get_tariffs_for_cluster(session, cluster_name) - - for tariff in tariffs: - if tariff["duration_days"] < 1: - continue - builder.button( - text=f"{tariff['name']} — {tariff['price_rub']}₽", - callback_data=f"tariff_{tariff['id']}", - ) - - builder.adjust(1) - builder.row(build_admin_back_btn()) - - await callback_query.message.edit_text( - text=f"🕒 Выберите срок действия ключа для страны {country}:", - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(UserEditorState.selecting_cluster, IsAdminFilter()) -async def handle_create_key_cluster(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - cluster_name = callback_query.data - - data = await state.get_data() - tg_id = data.get("tg_id") - - if not tg_id: - await callback_query.message.edit_text("❌ Ошибка: tg_id клиента не найден.") - return - - await state.update_data(cluster_name=cluster_name) - await state.set_state(UserEditorState.selecting_duration) - - tariffs = await get_tariffs_for_cluster(session, cluster_name) - - builder = InlineKeyboardBuilder() - for tariff in tariffs: - if tariff["duration_days"] < 1: - continue - builder.button( - text=f"{tariff['name']} — {tariff['price_rub']}₽", - callback_data=f"tariff_{tariff['id']}", - ) - - builder.adjust(1) - builder.row(build_admin_back_btn()) - - await callback_query.message.edit_text( - text=f"🕒 Выберите срок действия ключа для кластера {cluster_name}:", - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(UserEditorState.selecting_duration, IsAdminFilter()) -async def handle_create_key_duration(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - data = await state.get_data() - tg_id = data.get("tg_id", callback_query.from_user.id) - - use_country_selection = bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)) - - try: - if not callback_query.data.startswith("tariff_"): - raise ValueError("Некорректный callback_data") - tariff_id = int(callback_query.data.replace("tariff_", "")) - - tariff = await get_tariff_by_id(session, tariff_id) - if not tariff: - raise ValueError("Тариф не найден.") - - duration_days = tariff["duration_days"] - client_id = str(uuid.uuid4()) - email = await generate_random_email(session=session) - expiry = datetime.now(tz=timezone.utc) + timedelta(days=duration_days) - expiry_ms = int(expiry.timestamp() * 1000) - - if use_country_selection and "country" in data: - country = data["country"] - await create_key_on_cluster( - country, - tg_id, - client_id, - email, - expiry_ms, - plan=tariff_id, - session=session, - ) - - await state.clear() - await callback_query.message.edit_text( - f"✅ Ключ успешно создан для страны {country} на {duration_days} дней.", - reply_markup=build_editor_kb(tg_id), - ) - elif "cluster_name" in data: - cluster_name = data["cluster_name"] - await create_key_on_cluster( - cluster_name, - tg_id, - client_id, - email, - expiry_ms, - plan=tariff_id, - session=session, - ) - - await state.clear() - await callback_query.message.edit_text( - f"✅ Ключ успешно создан в кластере {cluster_name} на {duration_days} дней.", - reply_markup=build_editor_kb(tg_id), - ) - else: - await callback_query.message.edit_text("❌ Не удалось определить источник — страна или кластер.") - except Exception as e: - logger.error(f"[CreateKey] Ошибка при создании ключа: {e}") - await callback_query.message.edit_text( - "❌ Не удалось создать ключ. Попробуйте позже.", - reply_markup=build_editor_kb(tg_id), - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_reset_traffic"), - IsAdminFilter(), -) -async def handle_reset_traffic( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_obj = await _resolve_callback_key(session, tg_id, callback_data.data) - if not key_obj: - await callback_query.message.edit_text( - "❌ Ключ не найден в базе данных.", - reply_markup=build_editor_kb(tg_id), - ) - return - - email = key_obj.email - cluster_id = key_obj.server_id - - try: - await reset_traffic_in_cluster(cluster_id, email, session) - await callback_query.message.edit_text( - f"✅ Трафик для ключа {email} успешно сброшен.", - reply_markup=build_editor_kb(tg_id), - ) - except Exception as e: - logger.error(f"Ошибка при сбросе трафика: {e}") - await callback_query.message.edit_text( - "❌ Произошла ошибка при сбросе трафика. Попробуйте позже.", - reply_markup=build_editor_kb(tg_id), - ) - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_freeze"), - IsAdminFilter(), -) -async def handle_admin_freeze_subscription( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_obj = await _resolve_callback_key(session, tg_id, callback_data.data) - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - email = key_obj.email - - try: - record = await get_key_details(session, email) - if not record: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - - client_id = record["client_id"] - cluster_id = record["server_id"] - - result = await toggle_client_on_cluster(cluster_id, email, client_id, enable=False, session=session) - if result["status"] != "success": - text_error = ( - f"Произошла ошибка при отключении подписки.\nДетали: {result.get('error') or result.get('results')}" - ) - await callback_query.message.edit_text( - text_error, - reply_markup=build_editor_kb(tg_id, True), - ) - return - - now_ms = int(time.time() * 1000) - time_left = record["expiry_time"] - now_ms - if time_left < 0: - time_left = 0 - - await mark_key_as_frozen(session, record["tg_id"], client_id, time_left) - await session.commit() - session.expire_all() - - await callback_query.answer("✅ Подписка отключена") - - await handle_key_edit( - callback_query=callback_query, - callback_data=callback_data, - session=session, - update=False, - ) - except Exception as e: - await handle_error(tg_id, callback_query, f"Ошибка при отключении подписки: {e}") - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_unfreeze"), - IsAdminFilter(), -) -async def handle_admin_unfreeze_subscription( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - session: AsyncSession, -): - tg_id = callback_data.tg_id - key_obj = await _resolve_callback_key(session, tg_id, callback_data.data) - if not key_obj: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - email = key_obj.email - - try: - record = await get_key_details(session, email) - if not record: - await callback_query.message.edit_text( - text="🚫 Информация о ключе не найдена.", - reply_markup=build_editor_kb(tg_id), - ) - return - - client_id = record["client_id"] - cluster_id = record["server_id"] - - result = await toggle_client_on_cluster(cluster_id, email, client_id, enable=True, session=session) - if result["status"] != "success": - text_error = ( - f"Произошла ошибка при включении подписки.\nДетали: {result.get('error') or result.get('results')}" - ) - await callback_query.message.edit_text( - text_error, - reply_markup=build_editor_kb(tg_id, True), - ) - return - - tariff = await get_tariff_by_id(session, record["tariff_id"]) if record.get("tariff_id") else None - if not tariff: - total_gb = 0 - hwid_limit = 0 - else: - total_gb = int(tariff.get("traffic_limit") or 0) - hwid_limit = int(tariff.get("device_limit") or 0) - - if record.get("current_traffic_limit") is not None: - total_gb = record["current_traffic_limit"] - if record.get("current_device_limit") is not None: - hwid_limit = record["current_device_limit"] - - now_ms = int(time.time() * 1000) - leftover = record["expiry_time"] - if leftover < 0: - leftover = 0 - new_expiry_time = now_ms + leftover - - await mark_key_as_unfrozen(session, record["tg_id"], client_id, new_expiry_time) - await session.commit() - session.expire_all() - await release_session_early(session) - - await renew_key_in_cluster( - cluster_id=cluster_id, - email=email, - client_id=client_id, - new_expiry_time=new_expiry_time, - total_gb=total_gb, - session=session, - hwid_device_limit=hwid_limit, - reset_traffic=False, - plan=record.get("tariff_id"), - ) - - await callback_query.answer("✅ Подписка включена") - - await handle_key_edit( - callback_query=callback_query, - callback_data=callback_data, - session=session, - update=False, - ) - except Exception as e: - await handle_error(tg_id, callback_query, f"Ошибка при включении подписки: {e}") - - -async def change_expiry_time(expiry_time: int, email: str, session: AsyncSession) -> Exception | None: - key_obj = await get_key_by_email(session, email) - if not key_obj: - return ValueError(f"User with email {email} was not found") - - client_id = key_obj.client_id - tariff_id = key_obj.tariff_id - server_id = key_obj.server_id - key_device_limit = key_obj.current_device_limit - key_traffic_limit = key_obj.current_traffic_limit - if server_id is None: - return ValueError(f"Key with client_id {client_id} was not found") - - traffic_limit = 0 - device_limit = None - key_subgroup = None - if tariff_id: - tariff = await get_tariff_by_id(session, tariff_id) - if tariff: - traffic_limit = int(tariff.get("traffic_limit") or 0) - raw_device_limit = tariff.get("device_limit") - device_limit = int(raw_device_limit) if raw_device_limit is not None else 0 - key_subgroup = tariff.get("subgroup_title") - - if key_device_limit is not None: - device_limit = key_device_limit - if key_traffic_limit is not None: - traffic_limit = key_traffic_limit - - servers = await get_servers(session=session) - - if server_id in servers: - target_cluster = server_id - else: - target_cluster = None - for cluster_name, cluster_servers in servers.items(): - if any(s.get("server_name") == server_id for s in cluster_servers): - target_cluster = cluster_name - break - - if not target_cluster: - return ValueError(f"No suitable cluster found for server {server_id}") - - await release_session_early(session) - - await renew_key_in_cluster( - cluster_id=target_cluster, - email=email, - client_id=client_id, - new_expiry_time=expiry_time, - total_gb=traffic_limit, - session=session, - hwid_device_limit=device_limit, - reset_traffic=False, - target_subgroup=key_subgroup, - old_subgroup=key_subgroup, - plan=tariff_id, - ) - - await update_key_expiry(session, client_id, expiry_time) - return None - - -@router.callback_query( - AdminUserEditorCallback.filter(F.action == "users_edit_config"), - IsAdminFilter(), -) -async def handle_edit_config_start( - callback_query: CallbackQuery, - callback_data: AdminUserEditorCallback, - state: FSMContext, - session: AsyncSession, -): - key_ref = str(callback_data.data) - tg_id = callback_data.tg_id - - key_obj = await _resolve_callback_key(session, tg_id, key_ref) - - if not key_obj: - await callback_query.message.edit_text("❌ Ключ не найден.", reply_markup=build_editor_kb(tg_id)) - return - - email = key_obj.email - - if not key_obj.tariff_id: - await callback_query.message.edit_text( - "❌ У ключа не назначен тариф.", - reply_markup=build_key_edit_kb(key_obj.__dict__, email), - ) - return - - tariff = await get_tariff_by_id(session, key_obj.tariff_id) - if not tariff or not tariff.get("configurable"): - await callback_query.message.edit_text( - "❌ Тариф не поддерживает конфигурацию.", - reply_markup=build_key_edit_kb(key_obj.__dict__, email), - ) - return - - base_devices = key_obj.selected_device_limit or tariff.get("device_limit") or 1 - current_devices = key_obj.current_device_limit or base_devices - extra_devices = max(0, current_devices - base_devices) - - base_traffic = key_obj.selected_traffic_limit - current_traffic = key_obj.current_traffic_limit - extra_traffic = max(0, (current_traffic or 0) - (base_traffic or 0)) if current_traffic and base_traffic else 0 - - await state.set_state(UserEditorState.config_menu) - await state.update_data( - email=email, - key_ref=key_ref, - tg_id=tg_id, - tariff_id=key_obj.tariff_id, - cfg_base_devices=base_devices, - cfg_extra_devices=extra_devices, - cfg_base_traffic=base_traffic, - cfg_extra_traffic=extra_traffic, - ) - - await render_config_menu(callback_query, state, session) - - -async def render_config_menu(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - data = await state.get_data() - email = data.get("email") - key_ref = data.get("key_ref") - tg_id = data.get("tg_id") - tariff_id = data.get("tariff_id") - - tariff = await get_tariff_by_id(session, tariff_id) - if not tariff: - await callback_query.message.edit_text("❌ Тариф не найден.") - await state.clear() - return - - base_devices = data.get("cfg_base_devices") or 1 - extra_devices = data.get("cfg_extra_devices") or 0 - base_traffic = data.get("cfg_base_traffic") - extra_traffic = data.get("cfg_extra_traffic") or 0 - - traffic_to_show = base_traffic - if traffic_to_show is None and email: - key_obj = await get_key_by_email(session, email) - if key_obj: - traffic_to_show = key_obj.selected_traffic_limit or key_obj.current_traffic_limit - if traffic_to_show is None and tariff: - raw = tariff.get("traffic_limit") - if raw is not None: - try: - val = int(raw) - if val > 0: - traffic_to_show = val - except (TypeError, ValueError): - pass - - text = ( - f"⚙️ Конфигурация ключа\n\n" - f"🔑 Ключ: {email}\n" - f"📦 Тариф: {tariff.get('name')}\n\n" - ) - - extra_dev_str = f" + {extra_devices} (докуплено)" if extra_devices > 0 else "" - text += f"📱 Устройства: {base_devices}{extra_dev_str}\n" - - if traffic_to_show: - extra_traf_str = f" + {extra_traffic} ГБ (докуплено)" if extra_traffic > 0 else "" - text += f"📊 Трафик: {traffic_to_show} ГБ{extra_traf_str}\n" - else: - text += "📊 Трафик: безлимит\n" - - text += "\nВыберите что редактировать:" - - builder = InlineKeyboardBuilder() - builder.row( - InlineKeyboardButton(text="📦 Тариф (база)", callback_data="cfg_edit_base"), - InlineKeyboardButton(text="➕ Докупка", callback_data="cfg_edit_addon"), - ) - builder.row(InlineKeyboardButton(text="💾 Сохранить", callback_data="cfg_save")) - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=AdminUserEditorCallback(action="users_key_edit", data=key_ref, tg_id=tg_id).pack(), - ) - ) - - await state.set_state(UserEditorState.config_menu) - await callback_query.message.edit_text(text=text, reply_markup=builder.as_markup()) - - -@router.callback_query(F.data == "cfg_edit_base", UserEditorState.config_menu, IsAdminFilter()) -async def handle_cfg_edit_base(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - data = await state.get_data() - tariff = await get_tariff_by_id(session, data.get("tariff_id")) - device_options = tariff.get("device_options") or [] if tariff else [] - traffic_options = tariff.get("traffic_options_gb") or [] if tariff else [] - - builder = InlineKeyboardBuilder() - if device_options: - builder.row(InlineKeyboardButton(text="📱 Устройства", callback_data="cfg_base_devices")) - if traffic_options: - builder.row(InlineKeyboardButton(text="📊 Трафик", callback_data="cfg_base_traffic")) - builder.row(InlineKeyboardButton(text=BACK, callback_data="cfg_back_menu")) - - await callback_query.message.edit_text( - "📦 Редактирование базы тарифа\n\nВыберите параметр:", - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(F.data == "cfg_edit_addon", UserEditorState.config_menu, IsAdminFilter()) -async def handle_cfg_edit_addon(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - data = await state.get_data() - tariff = await get_tariff_by_id(session, data.get("tariff_id")) - device_options = tariff.get("device_options") or [] if tariff else [] - traffic_options = tariff.get("traffic_options_gb") or [] if tariff else [] - - builder = InlineKeyboardBuilder() - if device_options: - builder.row(InlineKeyboardButton(text="📱 Устройства", callback_data="cfg_addon_devices")) - if traffic_options: - builder.row(InlineKeyboardButton(text="📊 Трафик", callback_data="cfg_addon_traffic")) - builder.row(InlineKeyboardButton(text=BACK, callback_data="cfg_back_menu")) - - await callback_query.message.edit_text( - "➕ Редактирование докупки\n\nВыберите параметр:", - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(F.data == "cfg_back_menu", UserEditorState.config_menu, IsAdminFilter()) -async def handle_cfg_back_menu(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - await render_config_menu(callback_query, state, session) - - -@router.callback_query(F.data == "cfg_base_devices", UserEditorState.config_menu, IsAdminFilter()) -async def handle_cfg_base_devices(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - data = await state.get_data() - tariff = await get_tariff_by_id(session, data.get("tariff_id")) - device_options = tariff.get("device_options") or [] if tariff else [] - base_devices = data.get("cfg_base_devices") or 1 - - builder = InlineKeyboardBuilder() - for opt in sorted(device_options): - mark = " ✅" if int(opt) == int(base_devices) else "" - builder.button(text=f"{opt} устр.{mark}", callback_data=f"cfg_set_base_dev:{opt}") - builder.adjust(3) - builder.row(InlineKeyboardButton(text=BACK, callback_data="cfg_back_menu")) - - await state.set_state(UserEditorState.config_select_base) - await state.update_data(cfg_param="devices") - await callback_query.message.edit_text( - "📱 Выберите базу устройств:", - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(F.data == "cfg_base_traffic", UserEditorState.config_menu, IsAdminFilter()) -async def handle_cfg_base_traffic(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - data = await state.get_data() - tariff = await get_tariff_by_id(session, data.get("tariff_id")) - traffic_options = tariff.get("traffic_options_gb") or [] if tariff else [] - base_traffic = data.get("cfg_base_traffic") - - builder = InlineKeyboardBuilder() - for opt in sorted(traffic_options): - is_sel = (base_traffic is None and opt == 0) or (base_traffic is not None and int(opt) == int(base_traffic)) - mark = " ✅" if is_sel else "" - label = "безлимит" if opt == 0 else f"{opt} ГБ" - builder.button(text=f"{label}{mark}", callback_data=f"cfg_set_base_traf:{opt}") - builder.adjust(2) - builder.row(InlineKeyboardButton(text=BACK, callback_data="cfg_back_menu")) - - await state.set_state(UserEditorState.config_select_base) - await state.update_data(cfg_param="traffic") - await callback_query.message.edit_text( - "📊 Выберите базу трафика:", - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(F.data.startswith("cfg_set_base_dev:"), UserEditorState.config_select_base, IsAdminFilter()) -async def handle_cfg_set_base_dev(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - base_devices = int(callback_query.data.split(":")[1]) - await state.update_data(cfg_base_devices=base_devices) - await callback_query.answer(f"✅ База устройств: {base_devices}") - await state.set_state(UserEditorState.config_menu) - await render_config_menu(callback_query, state, session) - - -@router.callback_query(F.data.startswith("cfg_set_base_traf:"), UserEditorState.config_select_base, IsAdminFilter()) -async def handle_cfg_set_base_traf(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - traffic_gb = int(callback_query.data.split(":")[1]) - await state.update_data(cfg_base_traffic=traffic_gb if traffic_gb > 0 else None) - label = "безлимит" if traffic_gb == 0 else f"{traffic_gb} ГБ" - await callback_query.answer(f"✅ База трафика: {label}") - await state.set_state(UserEditorState.config_menu) - await render_config_menu(callback_query, state, session) - - -@router.callback_query(F.data == "cfg_addon_devices", UserEditorState.config_menu, IsAdminFilter()) -async def handle_cfg_addon_devices(callback_query: CallbackQuery, state: FSMContext): - data = await state.get_data() - extra_devices = data.get("cfg_extra_devices") or 0 - - await state.set_state(UserEditorState.config_input_addon) - await state.update_data(cfg_param="devices") - - builder = InlineKeyboardBuilder() - builder.row(InlineKeyboardButton(text="🔙 Отмена", callback_data="cfg_cancel_input")) - - await callback_query.message.edit_text( - f"📱 Докупка устройств\n\n" - f"Текущее значение: {extra_devices}\n\n" - f"Введите новое количество докупленных устройств (число):", - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(F.data == "cfg_addon_traffic", UserEditorState.config_menu, IsAdminFilter()) -async def handle_cfg_addon_traffic(callback_query: CallbackQuery, state: FSMContext): - data = await state.get_data() - extra_traffic = data.get("cfg_extra_traffic") or 0 - - await state.set_state(UserEditorState.config_input_addon) - await state.update_data(cfg_param="traffic") - - builder = InlineKeyboardBuilder() - builder.row(InlineKeyboardButton(text="🔙 Отмена", callback_data="cfg_cancel_input")) - - await callback_query.message.edit_text( - f"📊 Докупка трафика\n\n" - f"Текущее значение: {extra_traffic} ГБ\n\n" - f"Введите новое количество докупленного трафика в ГБ (число):", - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(F.data == "cfg_cancel_input", UserEditorState.config_input_addon, IsAdminFilter()) -async def handle_cfg_cancel_input(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - await state.set_state(UserEditorState.config_menu) - await render_config_menu(callback_query, state, session) - - -@router.message(UserEditorState.config_input_addon, IsAdminFilter()) -async def handle_cfg_input_addon(message: Message, state: FSMContext, session: AsyncSession): - data = await state.get_data() - param = data.get("cfg_param") - email = data.get("email") - key_ref = data.get("key_ref") - tg_id = data.get("tg_id") - tariff_id = data.get("tariff_id") - - if not message.text or not message.text.isdigit(): - await message.answer("❌ Введите корректное число.") - return - - value = int(message.text) - if value < 0: - await message.answer("❌ Значение не может быть отрицательным.") - return - - if param == "devices": - await state.update_data(cfg_extra_devices=value) - else: - await state.update_data(cfg_extra_traffic=value) - - await state.set_state(UserEditorState.config_menu) - - data = await state.get_data() - tariff = await get_tariff_by_id(session, tariff_id) - - base_devices = data.get("cfg_base_devices") or 1 - extra_devices = data.get("cfg_extra_devices") or 0 - base_traffic = data.get("cfg_base_traffic") - extra_traffic = data.get("cfg_extra_traffic") or 0 - - text = ( - f"⚙️ Конфигурация ключа\n\n" - f"🔑 Ключ: {email}\n" - f"📦 Тариф: {tariff.get('name') if tariff else '—'}\n\n" - ) - - extra_dev_str = f" + {extra_devices} (докуплено)" if extra_devices > 0 else "" - text += f"📱 Устройства: {base_devices}{extra_dev_str}\n" - - if base_traffic: - extra_traf_str = f" + {extra_traffic} ГБ (докуплено)" if extra_traffic > 0 else "" - text += f"📊 Трафик: {base_traffic} ГБ{extra_traf_str}\n" - else: - text += "📊 Трафик: безлимит\n" - - text += "\nВыберите что редактировать:" - - builder = InlineKeyboardBuilder() - builder.row( - InlineKeyboardButton(text="📦 Тариф (база)", callback_data="cfg_edit_base"), - InlineKeyboardButton(text="➕ Докупка", callback_data="cfg_edit_addon"), - ) - builder.row(InlineKeyboardButton(text="💾 Сохранить", callback_data="cfg_save")) - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=AdminUserEditorCallback(action="users_key_edit", data=key_ref, tg_id=tg_id).pack(), - ) - ) - - await message.answer(text=text, reply_markup=builder.as_markup()) - - -@router.callback_query(F.data == "cfg_save", UserEditorState.config_menu, IsAdminFilter()) -async def handle_cfg_save(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - data = await state.get_data() - email = data.get("email") - tg_id = data.get("tg_id") - tariff_id = data.get("tariff_id") - - base_devices = data.get("cfg_base_devices") or 1 - extra_devices = data.get("cfg_extra_devices") or 0 - total_devices = base_devices + extra_devices - - base_traffic = data.get("cfg_base_traffic") - extra_traffic = data.get("cfg_extra_traffic") or 0 - total_traffic = (base_traffic + extra_traffic) if base_traffic else None - - tariff = await get_tariff_by_id(session, tariff_id) - selected_price = None - if tariff: - base_price = tariff.get("price_rub") or 0 - - device_step = tariff.get("device_step_rub") or 0 - tariff_base_devices = tariff.get("device_limit") or 1 - extra_base_devices = max(0, base_devices - tariff_base_devices) - devices_extra_price = extra_base_devices * device_step - - traffic_step = tariff.get("traffic_step_rub") or 0 - tariff_base_traffic = tariff.get("traffic_limit") or 0 - extra_base_traffic = max(0, (base_traffic or 0) - tariff_base_traffic) if base_traffic else 0 - traffic_extra_price = extra_base_traffic * traffic_step - - selected_price = base_price + devices_extra_price + traffic_extra_price - - key_obj = await get_key_by_email(session, email) - - if not key_obj: - await callback_query.message.edit_text("❌ Ключ не найден.", reply_markup=build_editor_kb(tg_id)) - await state.clear() - return - - try: - await release_session_early(session) - await renew_key_in_cluster( - cluster_id=key_obj.server_id, - email=email, - client_id=key_obj.client_id, - new_expiry_time=key_obj.expiry_time, - total_gb=total_traffic or 0, - session=session, - hwid_device_limit=total_devices, - reset_traffic=False, - plan=tariff_id, - ) - - await save_admin_key_config( - session, - email=email, - base_devices=base_devices, - total_devices=total_devices, - base_traffic=base_traffic, - total_traffic=total_traffic, - selected_price=selected_price, - ) - - await state.clear() - await callback_query.answer("✅ Конфигурация сохранена", show_alert=True) - - callback_data_back = AdminUserEditorCallback(action="users_key_edit", data=email, tg_id=tg_id) - await handle_key_edit( - callback_query=callback_query, - callback_data=callback_data_back, - session=session, - update=False, - ) - - except Exception as e: - logger.error(f"[EditConfig] Ошибка при сохранении конфигурации: {e}") - await callback_query.message.edit_text( - "❌ Не удалось сохранить конфигурацию. Попробуйте позже.", - reply_markup=build_editor_kb(tg_id), - ) - await state.clear() - - -@router.callback_query(F.data == "cfg_back_menu", IsAdminFilter()) -async def handle_cfg_back_menu_any(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): - await state.set_state(UserEditorState.config_menu) - await render_config_menu(callback_query, state, session) diff --git a/handlers/admin/users/users_keys/__init__.py b/handlers/admin/users/users_keys/__init__.py new file mode 100644 index 00000000..1ef6e39b --- /dev/null +++ b/handlers/admin/users/users_keys/__init__.py @@ -0,0 +1,5 @@ +from ._common import router +from .edit import handle_key_edit # noqa: F401 — re-exported for users_tariffs.py +from . import config, edit, lifecycle, operations # noqa: F401 — trigger endpoint registration + +__all__ = ["router", "handle_key_edit"] diff --git a/handlers/admin/users/users_keys/_common.py b/handlers/admin/users/users_keys/_common.py new file mode 100644 index 00000000..3f452090 --- /dev/null +++ b/handlers/admin/users/users_keys/_common.py @@ -0,0 +1,81 @@ +import asyncio +import time +import uuid + +from datetime import datetime, timedelta, timezone +from handlers.buttons import BACK + +import pytz + +from aiogram import F, Router, types +from aiogram.exceptions import TelegramBadRequest +from aiogram.fsm.context import FSMContext +from aiogram.types import CallbackQuery, InlineKeyboardButton, Message +from aiogram.utils.keyboard import InlineKeyboardBuilder +from sqlalchemy.ext.asyncio import AsyncSession + +from config import REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD, REMNAWAVE_TOKEN_LOGIN_ENABLED, USE_COUNTRY_SELECTION +from core.bootstrap import MODES_CONFIG +from database import ( + check_server_name_by_cluster, + delete_key, + delete_user_data, + get_active_tariffs_by_group_code, + get_key_by_email, + get_key_details, + get_keys, + get_server_names, + get_servers, + get_tariff_by_id, + get_tariffs_for_cluster, + mark_key_as_frozen, + mark_key_as_unfrozen, + save_admin_key_config, + update_key_subscription_links, + update_key_expiry, +) +from database.models import Key +from filters.admin import IsAdminFilter +from middlewares.session import release_session_early +from services.operations import ( + create_key_on_cluster, + delete_key_from_cluster, + get_user_traffic, + renew_key_in_cluster, + reset_traffic_in_cluster, + toggle_client_on_cluster, + update_subscription, +) +from handlers.utils import generate_random_email, handle_error +from hooks.hook_buttons import insert_hook_buttons +from hooks.processors import process_admin_key_edit_menu +from logger import logger +from panels.remnawave import RemnawaveAPI + +from ...panel.keyboard import AdminPanelCallback, build_admin_back_btn, build_admin_back_kb +from ..keyboard import ( + AdminUserEditorCallback, + AdminUserKeyEditorCallback, + build_cluster_selection_kb, + build_editor_kb, + build_key_delete_kb, + build_key_edit_kb, + build_reissue_menu_kb, + build_user_delete_kb, + build_users_key_expiry_kb, + build_users_key_show_kb, +) +from services.users_utils import resolve_admin_key +from ..users_states import RenewTariffState, UserEditorState + + +MOSCOW_TZ = pytz.timezone("Europe/Moscow") + +router = Router() + +async def resolve_callback_key( + session: AsyncSession, + tg_id: int, + key_ref: str | int | None, +) -> Key | None: + return await resolve_admin_key(session, tg_id, key_ref) diff --git a/handlers/admin/users/users_keys/config.py b/handlers/admin/users/users_keys/config.py new file mode 100644 index 00000000..f0476ca0 --- /dev/null +++ b/handlers/admin/users/users_keys/config.py @@ -0,0 +1,438 @@ +"""Key config editor (base/addon devices + traffic limits).""" + +from ._common import * # noqa: F401,F403 +from .edit import handle_key_edit + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_edit_config"), + IsAdminFilter(), +) +async def handle_edit_config_start( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + state: FSMContext, + session: AsyncSession, +): + key_ref = str(callback_data.data) + tg_id = callback_data.tg_id + + key_obj = await resolve_callback_key(session, tg_id, key_ref) + + if not key_obj: + await callback_query.message.edit_text("❌ Ключ не найден.", reply_markup=build_editor_kb(tg_id)) + return + + email = key_obj.email + + if not key_obj.tariff_id: + await callback_query.message.edit_text( + "❌ У ключа не назначен тариф.", + reply_markup=build_key_edit_kb(key_obj.__dict__, email), + ) + return + + tariff = await get_tariff_by_id(session, key_obj.tariff_id) + if not tariff or not tariff.get("configurable"): + await callback_query.message.edit_text( + "❌ Тариф не поддерживает конфигурацию.", + reply_markup=build_key_edit_kb(key_obj.__dict__, email), + ) + return + + base_devices = key_obj.selected_device_limit or tariff.get("device_limit") or 1 + current_devices = key_obj.current_device_limit or base_devices + extra_devices = max(0, current_devices - base_devices) + + base_traffic = key_obj.selected_traffic_limit + current_traffic = key_obj.current_traffic_limit + extra_traffic = max(0, (current_traffic or 0) - (base_traffic or 0)) if current_traffic and base_traffic else 0 + + await state.set_state(UserEditorState.config_menu) + await state.update_data( + email=email, + key_ref=key_ref, + tg_id=tg_id, + tariff_id=key_obj.tariff_id, + cfg_base_devices=base_devices, + cfg_extra_devices=extra_devices, + cfg_base_traffic=base_traffic, + cfg_extra_traffic=extra_traffic, + ) + + await render_config_menu(callback_query, state, session) + + +async def render_config_menu(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + data = await state.get_data() + email = data.get("email") + key_ref = data.get("key_ref") + tg_id = data.get("tg_id") + tariff_id = data.get("tariff_id") + + tariff = await get_tariff_by_id(session, tariff_id) + if not tariff: + await callback_query.message.edit_text("❌ Тариф не найден.") + await state.clear() + return + + base_devices = data.get("cfg_base_devices") or 1 + extra_devices = data.get("cfg_extra_devices") or 0 + base_traffic = data.get("cfg_base_traffic") + extra_traffic = data.get("cfg_extra_traffic") or 0 + + traffic_to_show = base_traffic + if traffic_to_show is None and email: + key_obj = await get_key_by_email(session, email) + if key_obj: + traffic_to_show = key_obj.selected_traffic_limit or key_obj.current_traffic_limit + if traffic_to_show is None and tariff: + raw = tariff.get("traffic_limit") + if raw is not None: + try: + val = int(raw) + if val > 0: + traffic_to_show = val + except (TypeError, ValueError): + pass + + text = ( + f"⚙️ Конфигурация ключа\n\n" + f"🔑 Ключ: {email}\n" + f"📦 Тариф: {tariff.get('name')}\n\n" + ) + + extra_dev_str = f" + {extra_devices} (докуплено)" if extra_devices > 0 else "" + text += f"📱 Устройства: {base_devices}{extra_dev_str}\n" + + if traffic_to_show: + extra_traf_str = f" + {extra_traffic} ГБ (докуплено)" if extra_traffic > 0 else "" + text += f"📊 Трафик: {traffic_to_show} ГБ{extra_traf_str}\n" + else: + text += "📊 Трафик: безлимит\n" + + text += "\nВыберите что редактировать:" + + builder = InlineKeyboardBuilder() + builder.row( + InlineKeyboardButton(text="📦 Тариф (база)", callback_data="cfg_edit_base"), + InlineKeyboardButton(text="➕ Докупка", callback_data="cfg_edit_addon"), + ) + builder.row(InlineKeyboardButton(text="💾 Сохранить", callback_data="cfg_save")) + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=AdminUserEditorCallback(action="users_key_edit", data=key_ref, tg_id=tg_id).pack(), + ) + ) + + await state.set_state(UserEditorState.config_menu) + await callback_query.message.edit_text(text=text, reply_markup=builder.as_markup()) + + +@router.callback_query(F.data == "cfg_edit_base", UserEditorState.config_menu, IsAdminFilter()) +async def handle_cfg_edit_base(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + data = await state.get_data() + tariff = await get_tariff_by_id(session, data.get("tariff_id")) + device_options = tariff.get("device_options") or [] if tariff else [] + traffic_options = tariff.get("traffic_options_gb") or [] if tariff else [] + + builder = InlineKeyboardBuilder() + if device_options: + builder.row(InlineKeyboardButton(text="📱 Устройства", callback_data="cfg_base_devices")) + if traffic_options: + builder.row(InlineKeyboardButton(text="📊 Трафик", callback_data="cfg_base_traffic")) + builder.row(InlineKeyboardButton(text=BACK, callback_data="cfg_back_menu")) + + await callback_query.message.edit_text( + "📦 Редактирование базы тарифа\n\nВыберите параметр:", + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(F.data == "cfg_edit_addon", UserEditorState.config_menu, IsAdminFilter()) +async def handle_cfg_edit_addon(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + data = await state.get_data() + tariff = await get_tariff_by_id(session, data.get("tariff_id")) + device_options = tariff.get("device_options") or [] if tariff else [] + traffic_options = tariff.get("traffic_options_gb") or [] if tariff else [] + + builder = InlineKeyboardBuilder() + if device_options: + builder.row(InlineKeyboardButton(text="📱 Устройства", callback_data="cfg_addon_devices")) + if traffic_options: + builder.row(InlineKeyboardButton(text="📊 Трафик", callback_data="cfg_addon_traffic")) + builder.row(InlineKeyboardButton(text=BACK, callback_data="cfg_back_menu")) + + await callback_query.message.edit_text( + "➕ Редактирование докупки\n\nВыберите параметр:", + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(F.data == "cfg_back_menu", UserEditorState.config_menu, IsAdminFilter()) +async def handle_cfg_back_menu(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + await render_config_menu(callback_query, state, session) + + +@router.callback_query(F.data == "cfg_base_devices", UserEditorState.config_menu, IsAdminFilter()) +async def handle_cfg_base_devices(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + data = await state.get_data() + tariff = await get_tariff_by_id(session, data.get("tariff_id")) + device_options = tariff.get("device_options") or [] if tariff else [] + base_devices = data.get("cfg_base_devices") or 1 + + builder = InlineKeyboardBuilder() + for opt in sorted(device_options): + mark = " ✅" if int(opt) == int(base_devices) else "" + builder.button(text=f"{opt} устр.{mark}", callback_data=f"cfg_set_base_dev:{opt}") + builder.adjust(3) + builder.row(InlineKeyboardButton(text=BACK, callback_data="cfg_back_menu")) + + await state.set_state(UserEditorState.config_select_base) + await state.update_data(cfg_param="devices") + await callback_query.message.edit_text( + "📱 Выберите базу устройств:", + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(F.data == "cfg_base_traffic", UserEditorState.config_menu, IsAdminFilter()) +async def handle_cfg_base_traffic(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + data = await state.get_data() + tariff = await get_tariff_by_id(session, data.get("tariff_id")) + traffic_options = tariff.get("traffic_options_gb") or [] if tariff else [] + base_traffic = data.get("cfg_base_traffic") + + builder = InlineKeyboardBuilder() + for opt in sorted(traffic_options): + is_sel = (base_traffic is None and opt == 0) or (base_traffic is not None and int(opt) == int(base_traffic)) + mark = " ✅" if is_sel else "" + label = "безлимит" if opt == 0 else f"{opt} ГБ" + builder.button(text=f"{label}{mark}", callback_data=f"cfg_set_base_traf:{opt}") + builder.adjust(2) + builder.row(InlineKeyboardButton(text=BACK, callback_data="cfg_back_menu")) + + await state.set_state(UserEditorState.config_select_base) + await state.update_data(cfg_param="traffic") + await callback_query.message.edit_text( + "📊 Выберите базу трафика:", + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(F.data.startswith("cfg_set_base_dev:"), UserEditorState.config_select_base, IsAdminFilter()) +async def handle_cfg_set_base_dev(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + base_devices = int(callback_query.data.split(":")[1]) + await state.update_data(cfg_base_devices=base_devices) + await callback_query.answer(f"✅ База устройств: {base_devices}") + await state.set_state(UserEditorState.config_menu) + await render_config_menu(callback_query, state, session) + + +@router.callback_query(F.data.startswith("cfg_set_base_traf:"), UserEditorState.config_select_base, IsAdminFilter()) +async def handle_cfg_set_base_traf(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + traffic_gb = int(callback_query.data.split(":")[1]) + await state.update_data(cfg_base_traffic=traffic_gb if traffic_gb > 0 else None) + label = "безлимит" if traffic_gb == 0 else f"{traffic_gb} ГБ" + await callback_query.answer(f"✅ База трафика: {label}") + await state.set_state(UserEditorState.config_menu) + await render_config_menu(callback_query, state, session) + + +@router.callback_query(F.data == "cfg_addon_devices", UserEditorState.config_menu, IsAdminFilter()) +async def handle_cfg_addon_devices(callback_query: CallbackQuery, state: FSMContext): + data = await state.get_data() + extra_devices = data.get("cfg_extra_devices") or 0 + + await state.set_state(UserEditorState.config_input_addon) + await state.update_data(cfg_param="devices") + + builder = InlineKeyboardBuilder() + builder.row(InlineKeyboardButton(text="🔙 Отмена", callback_data="cfg_cancel_input")) + + await callback_query.message.edit_text( + f"📱 Докупка устройств\n\n" + f"Текущее значение: {extra_devices}\n\n" + f"Введите новое количество докупленных устройств (число):", + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(F.data == "cfg_addon_traffic", UserEditorState.config_menu, IsAdminFilter()) +async def handle_cfg_addon_traffic(callback_query: CallbackQuery, state: FSMContext): + data = await state.get_data() + extra_traffic = data.get("cfg_extra_traffic") or 0 + + await state.set_state(UserEditorState.config_input_addon) + await state.update_data(cfg_param="traffic") + + builder = InlineKeyboardBuilder() + builder.row(InlineKeyboardButton(text="🔙 Отмена", callback_data="cfg_cancel_input")) + + await callback_query.message.edit_text( + f"📊 Докупка трафика\n\n" + f"Текущее значение: {extra_traffic} ГБ\n\n" + f"Введите новое количество докупленного трафика в ГБ (число):", + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(F.data == "cfg_cancel_input", UserEditorState.config_input_addon, IsAdminFilter()) +async def handle_cfg_cancel_input(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + await state.set_state(UserEditorState.config_menu) + await render_config_menu(callback_query, state, session) + + +@router.message(UserEditorState.config_input_addon, IsAdminFilter()) +async def handle_cfg_input_addon(message: Message, state: FSMContext, session: AsyncSession): + data = await state.get_data() + param = data.get("cfg_param") + email = data.get("email") + key_ref = data.get("key_ref") + tg_id = data.get("tg_id") + tariff_id = data.get("tariff_id") + + if not message.text or not message.text.isdigit(): + await message.answer("❌ Введите корректное число.") + return + + value = int(message.text) + if value < 0: + await message.answer("❌ Значение не может быть отрицательным.") + return + + if param == "devices": + await state.update_data(cfg_extra_devices=value) + else: + await state.update_data(cfg_extra_traffic=value) + + await state.set_state(UserEditorState.config_menu) + + data = await state.get_data() + tariff = await get_tariff_by_id(session, tariff_id) + + base_devices = data.get("cfg_base_devices") or 1 + extra_devices = data.get("cfg_extra_devices") or 0 + base_traffic = data.get("cfg_base_traffic") + extra_traffic = data.get("cfg_extra_traffic") or 0 + + text = ( + f"⚙️ Конфигурация ключа\n\n" + f"🔑 Ключ: {email}\n" + f"📦 Тариф: {tariff.get('name') if tariff else '—'}\n\n" + ) + + extra_dev_str = f" + {extra_devices} (докуплено)" if extra_devices > 0 else "" + text += f"📱 Устройства: {base_devices}{extra_dev_str}\n" + + if base_traffic: + extra_traf_str = f" + {extra_traffic} ГБ (докуплено)" if extra_traffic > 0 else "" + text += f"📊 Трафик: {base_traffic} ГБ{extra_traf_str}\n" + else: + text += "📊 Трафик: безлимит\n" + + text += "\nВыберите что редактировать:" + + builder = InlineKeyboardBuilder() + builder.row( + InlineKeyboardButton(text="📦 Тариф (база)", callback_data="cfg_edit_base"), + InlineKeyboardButton(text="➕ Докупка", callback_data="cfg_edit_addon"), + ) + builder.row(InlineKeyboardButton(text="💾 Сохранить", callback_data="cfg_save")) + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=AdminUserEditorCallback(action="users_key_edit", data=key_ref, tg_id=tg_id).pack(), + ) + ) + + await message.answer(text=text, reply_markup=builder.as_markup()) + + +@router.callback_query(F.data == "cfg_save", UserEditorState.config_menu, IsAdminFilter()) +async def handle_cfg_save(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + data = await state.get_data() + email = data.get("email") + tg_id = data.get("tg_id") + tariff_id = data.get("tariff_id") + + base_devices = data.get("cfg_base_devices") or 1 + extra_devices = data.get("cfg_extra_devices") or 0 + total_devices = base_devices + extra_devices + + base_traffic = data.get("cfg_base_traffic") + extra_traffic = data.get("cfg_extra_traffic") or 0 + total_traffic = (base_traffic + extra_traffic) if base_traffic else None + + tariff = await get_tariff_by_id(session, tariff_id) + selected_price = None + if tariff: + base_price = tariff.get("price_rub") or 0 + + device_step = tariff.get("device_step_rub") or 0 + tariff_base_devices = tariff.get("device_limit") or 1 + extra_base_devices = max(0, base_devices - tariff_base_devices) + devices_extra_price = extra_base_devices * device_step + + traffic_step = tariff.get("traffic_step_rub") or 0 + tariff_base_traffic = tariff.get("traffic_limit") or 0 + extra_base_traffic = max(0, (base_traffic or 0) - tariff_base_traffic) if base_traffic else 0 + traffic_extra_price = extra_base_traffic * traffic_step + + selected_price = base_price + devices_extra_price + traffic_extra_price + + key_obj = await get_key_by_email(session, email) + + if not key_obj: + await callback_query.message.edit_text("❌ Ключ не найден.", reply_markup=build_editor_kb(tg_id)) + await state.clear() + return + + try: + await release_session_early(session) + await renew_key_in_cluster( + cluster_id=key_obj.server_id, + email=email, + client_id=key_obj.client_id, + new_expiry_time=key_obj.expiry_time, + total_gb=total_traffic or 0, + session=session, + hwid_device_limit=total_devices, + reset_traffic=False, + plan=tariff_id, + ) + + await save_admin_key_config( + session, + email=email, + base_devices=base_devices, + total_devices=total_devices, + base_traffic=base_traffic, + total_traffic=total_traffic, + selected_price=selected_price, + ) + + await state.clear() + await callback_query.answer("✅ Конфигурация сохранена", show_alert=True) + + callback_data_back = AdminUserEditorCallback(action="users_key_edit", data=email, tg_id=tg_id) + await handle_key_edit( + callback_query=callback_query, + callback_data=callback_data_back, + session=session, + update=False, + ) + + except Exception as e: + logger.error(f"[EditConfig] Ошибка при сохранении конфигурации: {e}") + await callback_query.message.edit_text( + "❌ Не удалось сохранить конфигурацию. Попробуйте позже.", + reply_markup=build_editor_kb(tg_id), + ) + await state.clear() + + +@router.callback_query(F.data == "cfg_back_menu", IsAdminFilter()) +async def handle_cfg_back_menu_any(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + await state.set_state(UserEditorState.config_menu) + await render_config_menu(callback_query, state, session) diff --git a/handlers/admin/users/users_keys/edit.py b/handlers/admin/users/users_keys/edit.py new file mode 100644 index 00000000..c81479e1 --- /dev/null +++ b/handlers/admin/users/users_keys/edit.py @@ -0,0 +1,411 @@ +from ._common import * # noqa: F401,F403 + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_key_edit"), + IsAdminFilter(), +) +async def handle_key_edit( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback | AdminUserKeyEditorCallback, + session: AsyncSession, + update: bool = False, +): + key_ref = callback_data.data + key_obj = await resolve_callback_key(session, callback_data.tg_id, key_ref) + + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Информация о подписке не найдена.", + reply_markup=build_editor_kb(callback_data.tg_id), + ) + return + + email = key_obj.email + key_details = await get_key_details(session, email) + is_frozen = bool(key_details.get("is_frozen")) if key_details else bool(getattr(key_obj, "is_frozen", False)) + + key_value = key_obj.key or key_obj.remnawave_link or "—" + alias_part = f" ({key_obj.alias})" if key_obj.alias else "" + + if key_obj.created_at: + created_at_dt = datetime.fromtimestamp(int(key_obj.created_at) / 1000, tz=MOSCOW_TZ) + created_at = created_at_dt.strftime("%d %B %Y года %H:%M") + else: + created_at = "—" + + if is_frozen: + frozen_left_ms = int((key_details or {}).get("expiry_time") or 0) + total_minutes = max(frozen_left_ms // 60000, 0) + days, rem_minutes = divmod(total_minutes, 24 * 60) + hours, minutes = divmod(rem_minutes, 60) + frozen_parts: list[str] = [] + if days: + frozen_parts.append(f"{days} дн.") + if hours: + frozen_parts.append(f"{hours} ч.") + if minutes or not frozen_parts: + frozen_parts.append(f"{minutes} мин.") + expiry_label = "⏳ Остаток:" + expiry_date = " ".join(frozen_parts) + elif key_obj.expiry_time: + expiry_dt = datetime.fromtimestamp(int(key_obj.expiry_time) / 1000, tz=MOSCOW_TZ) + expiry_label = "⏰ Истекает:" + expiry_date = expiry_dt.strftime("%d %B %Y года %H:%M") + else: + expiry_label = "⏰ Истекает:" + expiry_date = "—" + + tariff_name = "—" + subgroup_title = "—" + group_code = "—" + base_devices = None + base_traffic = None + is_configurable = False + if key_obj.tariff_id: + tariff = await get_tariff_by_id(session, key_obj.tariff_id) + if tariff: + tariff_name = tariff.get("name", "—") + subgroup_title = tariff.get("subgroup_title") or "—" + group_code = tariff.get("group_code") or "—" + base_devices = tariff.get("device_limit") + base_traffic = tariff.get("traffic_limit") + is_configurable = bool(tariff.get("configurable")) + + devices_line = "" + traffic_line = "" + if is_configurable: + sel_dev, cur_dev = key_obj.selected_device_limit, key_obj.current_device_limit + if sel_dev is not None or cur_dev is not None: + base_dev = sel_dev if sel_dev is not None else (base_devices if base_devices is not None else cur_dev) + extra = ( + f" + {cur_dev - base_dev} (докуплено)" + if (base_dev is not None and cur_dev is not None and cur_dev > base_dev) + else "" + ) + devices_line = f"📱 Устройства: {base_dev}{extra}\n" + + sel_traf, cur_traf = key_obj.selected_traffic_limit, key_obj.current_traffic_limit + if sel_traf is not None or cur_traf is not None: + base_traf = sel_traf if sel_traf is not None else (base_traffic if base_traffic is not None else cur_traf) + extra = ( + f" + {cur_traf - base_traf} ГБ (докуплено)" + if (base_traf is not None and cur_traf is not None and cur_traf > base_traf) + else "" + ) + traffic_line = f"📊 Трафик: {base_traf} ГБ{extra}\n" + + text = ( + "🔑 Информация о подписке\n\n" + "
" + f"🔗 Ключ{alias_part}: {key_value}\n" + f"📆 Создан: {created_at} (МСК)\n" + f"{'⛔ Статус: отключена\n' if is_frozen else ''}" + f"{expiry_label} {expiry_date}{' (МСК)' if not is_frozen and expiry_date != '—' else ''}\n" + f"🌐 Кластер: {key_obj.server_id or '—'}\n" + f"🆔 ID клиента: {key_obj.tg_id or '—'}\n" + f"🏷️ Тарифная группа: {group_code}\n" + f"📁 Подгруппа: {subgroup_title}\n" + f"📦 Тариф: {tariff_name}\n" + f"{devices_line}" + f"{traffic_line}" + "
" + ) + + if not update or not getattr(callback_data, "edit", False): + kb_key_details = dict(key_obj.__dict__) + kb_key_details["is_frozen"] = is_frozen + kb_markup = build_key_edit_kb(kb_key_details, email, is_configurable=is_configurable, key_ref=str(key_ref)) + kb_builder = InlineKeyboardBuilder.from_markup(kb_markup) + hook_buttons = await process_admin_key_edit_menu( + email=email, + session=session, + client_id=key_obj.client_id, + tg_id=key_obj.tg_id, + ) + kb_builder = insert_hook_buttons(kb_builder, hook_buttons) + try: + await callback_query.message.edit_text( + text=text, + reply_markup=kb_builder.as_markup(), + ) + except TelegramBadRequest as e: + if "message is not modified" not in str(e): + raise + else: + try: + await callback_query.message.edit_text( + text=text, + reply_markup=await build_users_key_expiry_kb( + session, + callback_data.tg_id, + email, + key_ref=str(key_ref), + ), + ) + except TelegramBadRequest as e: + if "message is not modified" not in str(e): + raise + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_expiry_edit"), + IsAdminFilter(), +) +async def handle_change_expiry( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_ref = str(callback_data.data) + key_obj = await resolve_callback_key(session, tg_id, key_ref) + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + email = key_obj.email + + await callback_query.message.edit_reply_markup( + reply_markup=await build_users_key_expiry_kb(session, tg_id, email, key_ref=key_ref) + ) + + +@router.callback_query( + AdminUserKeyEditorCallback.filter(F.action == "add"), + IsAdminFilter(), +) +async def handle_expiry_add( + callback_query: CallbackQuery, + callback_data: AdminUserKeyEditorCallback, + state: FSMContext, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_ref = str(callback_data.data) + key_obj = await resolve_callback_key(session, tg_id, key_ref) + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + email = key_obj.email + days = callback_data.month + + key_details = await get_key_details(session, email) + + if not key_details: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + + if days: + await change_expiry_time(key_details["expiry_time"] + days * 24 * 3600 * 1000, email, session) + await handle_key_edit(callback_query, callback_data, session, True) + return + + await state.update_data(tg_id=tg_id, email=email, key_ref=key_ref, op_type="add") + await state.set_state(UserEditorState.waiting_for_expiry_time) + + await callback_query.message.edit_text( + text="✍️ Введите количество дней, которое хотите добавить к времени действия ключа:", + reply_markup=build_users_key_show_kb(tg_id, key_ref), + ) + + +@router.callback_query( + AdminUserKeyEditorCallback.filter(F.action == "take"), + IsAdminFilter(), +) +async def handle_expiry_take( + callback_query: CallbackQuery, + callback_data: AdminUserKeyEditorCallback, + state: FSMContext, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_ref = str(callback_data.data) + key_obj = await resolve_callback_key(session, tg_id, key_ref) + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + email = key_obj.email + + await state.update_data(tg_id=tg_id, email=email, key_ref=key_ref, op_type="take") + await state.set_state(UserEditorState.waiting_for_expiry_time) + + await callback_query.message.edit_text( + text="✍️ Введите количество дней, которое хотите вычесть из времени действия ключа:", + reply_markup=build_users_key_show_kb(tg_id, key_ref), + ) + + +@router.callback_query( + AdminUserKeyEditorCallback.filter(F.action == "set"), + IsAdminFilter(), +) +async def handle_expiry_set( + callback_query: CallbackQuery, + callback_data: AdminUserKeyEditorCallback, + state: FSMContext, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_ref = str(callback_data.data) + key_obj = await resolve_callback_key(session, tg_id, key_ref) + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + email = key_obj.email + + key_details = await get_key_details(session, email) + + if not key_details: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + + await state.update_data(tg_id=tg_id, email=email, key_ref=key_ref, op_type="set") + await state.set_state(UserEditorState.waiting_for_expiry_time) + + text = ( + "✍️ Введите новое время действия ключа:" + "\n\n📌 Формат: год-месяц-день час:минута" + f"\n\n📄 Текущая дата: {datetime.fromtimestamp(key_details['expiry_time'] / 1000, tz=MOSCOW_TZ).strftime('%Y-%m-%d %H:%M')} (МСК)" + ) + + await callback_query.message.edit_text( + text=text, + reply_markup=build_users_key_show_kb(tg_id, key_ref), + ) + + +@router.message(UserEditorState.waiting_for_expiry_time, IsAdminFilter()) +async def handle_expiry_time_input(message: Message, state: FSMContext, session: AsyncSession): + data = await state.get_data() + tg_id = data.get("tg_id") + email = data.get("email") + key_ref = data.get("key_ref") + op_type = data.get("op_type") + + if op_type != "set" and (not message.text.isdigit() or int(message.text) < 0): + await message.answer( + text="🚫 Пожалуйста, введите корректное количество дней!", + reply_markup=build_users_key_show_kb(tg_id, key_ref) if key_ref else build_editor_kb(tg_id), + ) + return + + key_details = await get_key_details(session, email) + + if not key_details: + await message.answer( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + + try: + current_expiry_time = datetime.fromtimestamp( + key_details["expiry_time"] / 1000, + tz=MOSCOW_TZ, + ) + + if op_type == "add": + days = int(message.text) + new_expiry_time = current_expiry_time + timedelta(days=days) + text = f"✅ Ко времени действия ключа добавлено {days} дн." + elif op_type == "take": + days = int(message.text) + new_expiry_time = current_expiry_time - timedelta(days=days) + text = f"✅ Из времени действия ключа вычтено {days} дн." + else: + new_expiry_time = datetime.strptime(message.text, "%Y-%m-%d %H:%M") + new_expiry_time = MOSCOW_TZ.localize(new_expiry_time) + text = f"✅ Время действия ключа изменено на {message.text} (МСК)" + + new_expiry_timestamp = int(new_expiry_time.timestamp() * 1000) + await change_expiry_time(new_expiry_timestamp, email, session) + except ValueError: + text = "🚫 Пожалуйста, используйте корректный формат даты (ГГГГ-ММ-ДД ЧЧ:ММ)!" + except Exception as e: + text = f"❗ Произошла ошибка во время изменения времени действия ключа: {e}" + + await message.answer( + text=text, + reply_markup=build_users_key_show_kb(tg_id, key_ref) if key_ref else build_editor_kb(tg_id), + ) + + +async def change_expiry_time(expiry_time: int, email: str, session: AsyncSession) -> Exception | None: + key_obj = await get_key_by_email(session, email) + if not key_obj: + return ValueError(f"User with email {email} was not found") + + client_id = key_obj.client_id + tariff_id = key_obj.tariff_id + server_id = key_obj.server_id + key_device_limit = key_obj.current_device_limit + key_traffic_limit = key_obj.current_traffic_limit + if server_id is None: + return ValueError(f"Key with client_id {client_id} was not found") + + traffic_limit = 0 + device_limit = None + key_subgroup = None + if tariff_id: + tariff = await get_tariff_by_id(session, tariff_id) + if tariff: + traffic_limit = int(tariff.get("traffic_limit") or 0) + raw_device_limit = tariff.get("device_limit") + device_limit = int(raw_device_limit) if raw_device_limit is not None else 0 + key_subgroup = tariff.get("subgroup_title") + + if key_device_limit is not None: + device_limit = key_device_limit + if key_traffic_limit is not None: + traffic_limit = key_traffic_limit + + servers = await get_servers(session=session) + + if server_id in servers: + target_cluster = server_id + else: + target_cluster = None + for cluster_name, cluster_servers in servers.items(): + if any(s.get("server_name") == server_id for s in cluster_servers): + target_cluster = cluster_name + break + + if not target_cluster: + return ValueError(f"No suitable cluster found for server {server_id}") + + await release_session_early(session) + + await renew_key_in_cluster( + cluster_id=target_cluster, + email=email, + client_id=client_id, + new_expiry_time=expiry_time, + total_gb=traffic_limit, + session=session, + hwid_device_limit=device_limit, + reset_traffic=False, + target_subgroup=key_subgroup, + old_subgroup=key_subgroup, + plan=tariff_id, + ) + + await update_key_expiry(session, client_id, expiry_time) + return None diff --git a/handlers/admin/users/users_keys/lifecycle.py b/handlers/admin/users/users_keys/lifecycle.py new file mode 100644 index 00000000..244512cc --- /dev/null +++ b/handlers/admin/users/users_keys/lifecycle.py @@ -0,0 +1,771 @@ +from ._common import * # noqa: F401,F403 +from .edit import handle_key_edit + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_reissue_menu"), + IsAdminFilter(), +) +async def handle_reissue_menu( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_ref = str(callback_data.data) + key_obj = await resolve_callback_key(session, tg_id, key_ref) + if not key_obj: + await callback_query.message.edit_text("🚫 Ключ не найден.", reply_markup=build_editor_kb(tg_id)) + return + + text = ( + "🔄 Перевыпуск подписки\n\n" + "📦 Полный перевыпуск\n" + "Пересоздаёт подписку на сервере с возможностью выбора кластера. " + "Используйте для переноса на другой сервер или обновления данных.\n\n" + "🔗 Сменить ссылку\n" + "Генерирует новую ссылку подписки. Старая ссылка перестанет работать. " + "Все данные подписки сохранятся." + ) + + await callback_query.message.edit_text( + text=text, + reply_markup=build_reissue_menu_kb(key_ref, tg_id), + ) + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_update_key"), + IsAdminFilter(), +) +async def handle_update_key( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_ref = str(callback_data.data) + key_obj = await resolve_callback_key(session, tg_id, key_ref) + if not key_obj: + await callback_query.message.edit_text("🚫 Ключ не найден.", reply_markup=build_editor_kb(tg_id)) + return + email = key_obj.email + + await callback_query.message.edit_text( + text=f"📡 Выберите кластер, на котором пересоздать ключ {email}:", + reply_markup=await build_cluster_selection_kb( + session, + tg_id, + key_ref, + action="confirm_admin_key_reissue", + ), + ) + + +@router.callback_query(F.data.startswith("confirm_admin_key_reissue|"), IsAdminFilter()) +async def confirm_admin_key_reissue(callback_query: CallbackQuery, session: AsyncSession, state: FSMContext): + _, tg_id, key_ref, cluster_id = callback_query.data.split("|") + tg_id = int(tg_id) + key_obj = await resolve_callback_key(session, tg_id, key_ref) + if not key_obj: + await callback_query.message.edit_text("🚫 Ключ не найден.", reply_markup=build_editor_kb(tg_id)) + return + email = key_obj.email + + try: + servers = await get_servers(session) + cluster_servers = servers.get(cluster_id, []) + + tariffs = await get_tariffs_for_cluster(session, cluster_id) + if not tariffs: + builder = InlineKeyboardBuilder() + builder.row( + InlineKeyboardButton( + text="🔗 Привязать тариф", + callback_data=AdminPanelCallback(action="clusters").pack(), + ) + ) + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=AdminUserEditorCallback( + action="users_key_edit", + tg_id=tg_id, + data=key_ref, + ).pack(), + ) + ) + await callback_query.message.edit_text( + f"🚫 Невозможно пересоздать подписку\n\n" + f"📊 Информация о кластере:\n
" + f"🌐 Кластер: {cluster_id}\n" + f"⚠️ Статус: Нет привязанного тарифа\n
" + f"💡 Привяжите тариф к кластеру", + reply_markup=builder.as_markup(), + ) + return + + use_country_selection = bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)) + + if use_country_selection: + unique_countries = {srv["server_name"] for srv in cluster_servers} + await state.update_data(tg_id=tg_id, email=email, key_ref=key_ref, cluster_id=cluster_id) + builder = InlineKeyboardBuilder() + for country in sorted(unique_countries): + builder.button( + text=country, + callback_data=f"admin_reissue_country|{tg_id}|{key_ref}|{country}", + ) + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=AdminUserEditorCallback( + action="users_key_edit", + tg_id=tg_id, + data=key_ref, + ).pack(), + ) + ) + await callback_query.message.edit_text( + "🌍 Выберите сервер (страну) для пересоздания подписки:", + reply_markup=builder.as_markup(), + ) + return + + key_link = await get_key_by_email(session, email) + remnawave_link = key_link.remnawave_link if key_link else None + + await update_subscription( + tg_id, + email, + session, + cluster_override=cluster_id, + remnawave_link=remnawave_link, + ) + + await handle_key_edit( + callback_query, + AdminUserEditorCallback(tg_id=tg_id, data=key_ref, action="view_key"), + session, + True, + ) + except Exception as e: + logger.error(f"Ошибка при перевыпуске ключа {email}: {e}") + await callback_query.message.answer(f"❗ Ошибка: {e}") + + +@router.callback_query(F.data.startswith("admin_reissue_country|"), IsAdminFilter()) +async def admin_reissue_country(callback_query: CallbackQuery, session: AsyncSession, state: FSMContext): + _, tg_id, key_ref, country = callback_query.data.split("|") + tg_id = int(tg_id) + key_obj = await resolve_callback_key(session, tg_id, key_ref) + if not key_obj: + await callback_query.message.edit_text("🚫 Ключ не найден.", reply_markup=build_editor_kb(tg_id)) + return + email = key_obj.email + + try: + data = await state.get_data() + cluster_id = data.get("cluster_id") + + if cluster_id: + tariffs = await get_tariffs_for_cluster(session, cluster_id) + if not tariffs: + builder = InlineKeyboardBuilder() + builder.row( + InlineKeyboardButton( + text="🔗 Привязать тариф", + callback_data=AdminPanelCallback(action="clusters").pack(), + ) + ) + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=AdminUserEditorCallback( + action="users_key_edit", + tg_id=tg_id, + data=key_ref, + ).pack(), + ) + ) + await callback_query.message.edit_text( + f"🚫 Невозможно пересоздать подписку\n\n" + f"📊 Информация о кластере:\n
" + f"🌐 Кластер: {cluster_id}\n" + f"⚠️ Статус: Нет привязанного тарифа\n
" + f"💡 Привяжите тариф к кластеру", + reply_markup=builder.as_markup(), + ) + return + + key_link = await get_key_by_email(session, email) + remnawave_link = key_link.remnawave_link if key_link else None + + await update_subscription( + tg_id=tg_id, + email=email, + session=session, + country_override=country, + remnawave_link=remnawave_link, + ) + + await handle_key_edit( + callback_query, + AdminUserEditorCallback(tg_id=tg_id, data=key_ref, action="view_key"), + session, + True, + ) + except Exception as e: + logger.error(f"Ошибка при перевыпуске ключа для страны {country}: {e}") + await callback_query.message.answer(f"❗ Ошибка: {e}") + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_recreate_key"), + IsAdminFilter(), +) +async def handle_recreate_key_start( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_ref = str(callback_data.data) + key_obj = await resolve_callback_key(session, tg_id, key_ref) + + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Ключ не найден.", + reply_markup=build_editor_kb(tg_id), + ) + return + + email = key_obj.email + + tariff_name = "—" + if key_obj.tariff_id: + tariff = await get_tariff_by_id(session, key_obj.tariff_id) + if tariff: + tariff_name = tariff.get("name", "—") + + text = ( + "🔁 Пересоздание ссылки подписки\n\n" + f"📦 Тариф: {tariff_name}\n\n" + "⚠️ Будет сгенерирована новая ссылка подписки.\n" + "Старая ссылка перестанет работать.\n\n" + "✅ Все данные подписки сохранятся." + ) + + builder = InlineKeyboardBuilder() + builder.row( + InlineKeyboardButton( + text="✅ Пересоздать", + callback_data=f"confirm_recreate|{tg_id}|{key_ref}", + ) + ) + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=AdminUserEditorCallback(action="users_key_edit", tg_id=tg_id, data=key_ref).pack(), + ) + ) + + await callback_query.message.edit_text(text=text, reply_markup=builder.as_markup()) + + +@router.callback_query(F.data.startswith("confirm_recreate|"), IsAdminFilter()) +async def handle_recreate_key_confirm( + callback_query: CallbackQuery, + session: AsyncSession, +): + _, tg_id, key_ref = callback_query.data.split("|") + tg_id = int(tg_id) + + try: + key_obj = await resolve_callback_key(session, tg_id, key_ref) + + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Ключ не найден.", + reply_markup=build_editor_kb(tg_id), + ) + return + + old_email = key_obj.email + + await callback_query.message.edit_text("⏳ Пересоздание ссылки подписки...") + + client_id = key_obj.client_id + cluster_id = key_obj.server_id + old_link = key_obj.remnawave_link or key_obj.key + + servers = await get_servers(session) + cluster = servers.get(cluster_id) + + if not cluster: + for _, server_list in servers.items(): + for server_info in server_list: + if server_info.get("server_name", "").lower() == cluster_id.lower(): + cluster = [server_info] + break + if cluster: + break + + if not cluster: + await callback_query.message.edit_text( + text=f"❗ Кластер {cluster_id} не найден.", + reply_markup=build_editor_kb(tg_id), + ) + return + + remnawave_servers = [s for s in cluster if s.get("panel_type", "3x-ui").lower() == "remnawave"] + + if not remnawave_servers: + await callback_query.message.edit_text( + text="❗ Revoke доступен только для Remnawave. Для 3x-ui используйте перевыпуск.", + reply_markup=build_editor_kb(tg_id), + ) + return + + api_url = remnawave_servers[0].get("api_url") + if not api_url: + await callback_query.message.edit_text( + text="❗ У Remnawave сервера не задан api_url.", + reply_markup=build_editor_kb(tg_id), + ) + return + + api = RemnawaveAPI(api_url) + try: + if not REMNAWAVE_TOKEN_LOGIN_ENABLED: + await api.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD) + + user_data = await api.revoke_user_subscription(client_id) + finally: + await api.aclose() + + if not user_data: + await callback_query.message.edit_text( + text="❗ Не удалось выполнить revoke. Проверьте логи.", + reply_markup=build_editor_kb(tg_id), + ) + return + + new_link = user_data.get("subscriptionUrl") + + if not new_link: + await callback_query.message.edit_text( + text="❗ Revoke выполнен, но новая ссылка не получена.", + reply_markup=build_editor_kb(tg_id), + ) + return + + await update_key_subscription_links(session, old_email, new_link) + + try: + user_text = ( + "🔄 Ваша подписка была перевыпущена\n\n" + f"🔗 Новая ссылка подписки:\n{new_link}\n\n" + "Старая ссылка больше не работает." + ) + user_kb = InlineKeyboardBuilder() + user_kb.row( + InlineKeyboardButton( + text="📱 Мои подписки", + callback_data="view_keys", + ) + ) + user_kb.row( + InlineKeyboardButton( + text="👤 Личный кабинет", + callback_data="profile", + ) + ) + + await callback_query.bot.send_message( + chat_id=tg_id, + text=user_text, + reply_markup=user_kb.as_markup(), + ) + notification_sent = True + except Exception as e: + logger.warning(f"Не удалось отправить уведомление клиенту {tg_id}: {e}") + notification_sent = False + + text = ( + "✅ Ссылка подписки пересоздана\n\n" + f"🔗 Старая ссылка:\n{old_link}\n\n" + f"🔗 Новая ссылка:\n{new_link}\n\n" + ) + if notification_sent: + text += "📨 Клиент уведомлён о новой ссылке." + else: + text += "⚠️ Не удалось уведомить клиента." + + builder = InlineKeyboardBuilder() + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=AdminUserEditorCallback( + action="users_key_edit", + tg_id=tg_id, + data=key_ref, + ).pack(), + ) + ) + + await callback_query.message.edit_text( + text=text, + reply_markup=builder.as_markup(), + ) + + except Exception as e: + logger.error(f"Ошибка при revoke ключа {old_email}: {e}") + await callback_query.message.edit_text( + text=f"❗ Ошибка при пересоздании: {e}", + reply_markup=build_editor_kb(tg_id), + ) + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_delete_key"), + IsAdminFilter(), +) +async def handle_delete_key( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + state: FSMContext, + session: AsyncSession, +): + key_obj = await resolve_callback_key(session, callback_data.tg_id, callback_data.data) + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Ключ не найден!", + reply_markup=build_editor_kb(callback_data.tg_id), + ) + return + + email = key_obj.email + client_id = key_obj.client_id + + if client_id is None: + await callback_query.message.edit_text( + text="🚫 Ключ не найден!", + reply_markup=build_editor_kb(callback_data.tg_id), + ) + return + + await state.set_state(UserEditorState.confirm_delete_key) + await state.update_data( + delete_key_email=email, + delete_key_tg_id=int(callback_data.tg_id), + delete_key_client_id=client_id, + ) + + await callback_query.message.edit_text( + text="❓ Вы уверены, что хотите удалить ключ?", + reply_markup=build_key_delete_kb(callback_data.tg_id), + ) + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_delete_key_confirm"), + UserEditorState.confirm_delete_key, + IsAdminFilter(), +) +async def handle_delete_key_confirm( + callback_query: types.CallbackQuery, + callback_data: AdminUserEditorCallback, + state: FSMContext, + session: AsyncSession, +): + data = await state.get_data() + email = data.get("delete_key_email") + expected_tg_id = data.get("delete_key_tg_id") + client_id = data.get("delete_key_client_id") + await state.clear() + + if not email or int(expected_tg_id or 0) != int(callback_data.tg_id): + await callback_query.answer("Данные устарели", show_alert=True) + return + + if not client_id: + key_obj = await get_key_by_email(session, email, int(callback_data.tg_id)) + client_id = key_obj.client_id if key_obj else None + + kb = build_editor_kb(callback_data.tg_id) + + if client_id: + clusters = await get_servers(session=session) + await release_session_early(session) + + async def delete_key_from_servers(): + tasks = [] + for cluster_name, cluster_servers in clusters.items(): + for _ in cluster_servers: + tasks.append(delete_key_from_cluster(cluster_name, email, client_id, session)) + await asyncio.gather(*tasks, return_exceptions=True) + + await delete_key_from_servers() + await delete_key(session, client_id) + + await callback_query.message.edit_text(text="✅ Ключ успешно удален.", reply_markup=kb) + else: + await callback_query.message.edit_text( + text="🚫 Ключ не найден или уже удален.", + reply_markup=kb, + ) + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_delete_user"), + IsAdminFilter(), +) +async def handle_delete_user( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, +): + tg_id = callback_data.tg_id + await callback_query.message.edit_text( + text=f"❗️ Вы уверены, что хотите удалить пользователя с ID {tg_id}?", + reply_markup=build_user_delete_kb(tg_id), + ) + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_delete_user_confirm"), + IsAdminFilter(), +) +async def handle_delete_user_confirm( + callback_query: types.CallbackQuery, + callback_data: AdminUserEditorCallback, + session: AsyncSession, +): + tg_id = callback_data.tg_id + + key_records = [(row.email, row.client_id) for row in await get_keys(session, tg_id)] + await release_session_early(session) + + async def delete_keys_from_servers(): + try: + tasks = [] + servers = await get_servers(session=session) + for email, client_id in key_records: + for cluster_id, _cluster in servers.items(): + tasks.append(delete_key_from_cluster(cluster_id, email, client_id, session)) + await asyncio.gather(*tasks, return_exceptions=True) + except Exception as e: + logger.error(f"Ошибка при удалении ключей с серверов для пользователя {tg_id}: {e}") + + await delete_keys_from_servers() + + try: + await delete_user_data(session, tg_id) + await callback_query.message.edit_text( + text=f"🗑️ Пользователь с ID {tg_id} был удален.", + reply_markup=build_admin_back_kb(), + ) + except Exception as e: + logger.error(f"Ошибка при удалении данных из базы данных для пользователя {tg_id}: {e}") + await callback_query.message.edit_text( + text=f"❌ Произошла ошибка при удалении пользователя с ID {tg_id}. Попробуйте снова.", + reply_markup=build_admin_back_kb(), + ) + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_create_key"), + IsAdminFilter(), +) +async def handle_create_key_start( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + state: FSMContext, + session: AsyncSession, +): + tg_id = callback_data.tg_id + await state.update_data(tg_id=tg_id) + + use_country_selection = bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)) + + if use_country_selection: + await state.set_state(UserEditorState.selecting_country) + + countries = await get_server_names(session) + + if not countries: + await callback_query.message.edit_text( + "❌ Нет доступных стран для создания ключа.", + reply_markup=build_editor_kb(tg_id), + ) + return + + builder = InlineKeyboardBuilder() + for country in countries: + builder.button(text=country, callback_data=country) + builder.adjust(1) + builder.row(build_admin_back_btn()) + + await callback_query.message.edit_text( + "🌍 Выберите страну для создания ключа:", + reply_markup=builder.as_markup(), + ) + return + + await state.set_state(UserEditorState.selecting_cluster) + + servers = await get_servers(session=session) + cluster_names = list(servers.keys()) + + if not cluster_names: + await callback_query.message.edit_text( + "❌ Нет доступных кластеров для создания ключа.", + reply_markup=build_editor_kb(tg_id), + ) + return + + builder = InlineKeyboardBuilder() + for cluster in cluster_names: + builder.button(text=f"🌐 {cluster}", callback_data=cluster) + builder.adjust(2) + builder.row(build_admin_back_btn()) + + await callback_query.message.edit_text( + "🌐 Выберите кластер для создания ключа:", + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(UserEditorState.selecting_country, IsAdminFilter()) +async def handle_create_key_country(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + country = callback_query.data + await state.update_data(country=country) + await state.set_state(UserEditorState.selecting_duration) + + builder = InlineKeyboardBuilder() + + cluster_info = await check_server_name_by_cluster(session, country) + + if not cluster_info: + await callback_query.message.edit_text("❌ Сервер не найден.") + return + + cluster_name = cluster_info["cluster_name"] + await state.update_data(cluster_name=cluster_name) + + tariffs = await get_tariffs_for_cluster(session, cluster_name) + + for tariff in tariffs: + if tariff["duration_days"] < 1: + continue + builder.button( + text=f"{tariff['name']} — {tariff['price_rub']}₽", + callback_data=f"tariff_{tariff['id']}", + ) + + builder.adjust(1) + builder.row(build_admin_back_btn()) + + await callback_query.message.edit_text( + text=f"🕒 Выберите срок действия ключа для страны {country}:", + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(UserEditorState.selecting_cluster, IsAdminFilter()) +async def handle_create_key_cluster(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + cluster_name = callback_query.data + + data = await state.get_data() + tg_id = data.get("tg_id") + + if not tg_id: + await callback_query.message.edit_text("❌ Ошибка: tg_id клиента не найден.") + return + + await state.update_data(cluster_name=cluster_name) + await state.set_state(UserEditorState.selecting_duration) + + tariffs = await get_tariffs_for_cluster(session, cluster_name) + + builder = InlineKeyboardBuilder() + for tariff in tariffs: + if tariff["duration_days"] < 1: + continue + builder.button( + text=f"{tariff['name']} — {tariff['price_rub']}₽", + callback_data=f"tariff_{tariff['id']}", + ) + + builder.adjust(1) + builder.row(build_admin_back_btn()) + + await callback_query.message.edit_text( + text=f"🕒 Выберите срок действия ключа для кластера {cluster_name}:", + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(UserEditorState.selecting_duration, IsAdminFilter()) +async def handle_create_key_duration(callback_query: CallbackQuery, state: FSMContext, session: AsyncSession): + data = await state.get_data() + tg_id = data.get("tg_id", callback_query.from_user.id) + + use_country_selection = bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)) + + try: + if not callback_query.data.startswith("tariff_"): + raise ValueError("Некорректный callback_data") + tariff_id = int(callback_query.data.replace("tariff_", "")) + + tariff = await get_tariff_by_id(session, tariff_id) + if not tariff: + raise ValueError("Тариф не найден.") + + duration_days = tariff["duration_days"] + client_id = str(uuid.uuid4()) + email = await generate_random_email(session=session) + expiry = datetime.now(tz=timezone.utc) + timedelta(days=duration_days) + expiry_ms = int(expiry.timestamp() * 1000) + + if use_country_selection and "country" in data: + country = data["country"] + await create_key_on_cluster( + country, + tg_id, + client_id, + email, + expiry_ms, + plan=tariff_id, + session=session, + ) + + await state.clear() + await callback_query.message.edit_text( + f"✅ Ключ успешно создан для страны {country} на {duration_days} дней.", + reply_markup=build_editor_kb(tg_id), + ) + elif "cluster_name" in data: + cluster_name = data["cluster_name"] + await create_key_on_cluster( + cluster_name, + tg_id, + client_id, + email, + expiry_ms, + plan=tariff_id, + session=session, + ) + + await state.clear() + await callback_query.message.edit_text( + f"✅ Ключ успешно создан в кластере {cluster_name} на {duration_days} дней.", + reply_markup=build_editor_kb(tg_id), + ) + else: + await callback_query.message.edit_text("❌ Не удалось определить источник — страна или кластер.") + except Exception as e: + logger.error(f"[CreateKey] Ошибка при создании ключа: {e}") + await callback_query.message.edit_text( + "❌ Не удалось создать ключ. Попробуйте позже.", + reply_markup=build_editor_kb(tg_id), + ) diff --git a/handlers/admin/users/users_keys/operations.py b/handlers/admin/users/users_keys/operations.py new file mode 100644 index 00000000..eae14e33 --- /dev/null +++ b/handlers/admin/users/users_keys/operations.py @@ -0,0 +1,234 @@ +from ._common import * # noqa: F401,F403 +from .edit import handle_key_edit + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_traffic"), + IsAdminFilter(), +) +async def handle_user_traffic( + callback_query: types.CallbackQuery, + callback_data: AdminUserEditorCallback, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_obj = await resolve_callback_key(session, tg_id, callback_data.data) + if not key_obj: + await callback_query.message.edit_text("❌ Ключ не найден.", reply_markup=build_editor_kb(tg_id)) + return + email = key_obj.email + + await callback_query.message.edit_text("⏳ Получаем данные о трафике, пожалуйста, подождите...") + + traffic_data = await get_user_traffic(session, tg_id, email) + + if traffic_data["status"] == "error": + await callback_query.message.edit_text( + traffic_data["message"], + reply_markup=build_editor_kb(tg_id, True), + ) + return + + total_traffic = 0 + result_text = f"📊 Трафик подписки {email}:\n\n" + + for server, traffic in traffic_data["traffic"].items(): + if isinstance(traffic, str): + result_text += f"❌ {server}: {traffic}\n" + else: + result_text += f"🌍 {server}: {traffic} ГБ\n" + total_traffic += traffic + + result_text += f"\n🔢 Общий трафик: {total_traffic:.2f} ГБ" + + await callback_query.message.edit_text( + result_text, + reply_markup=build_editor_kb(tg_id, True), + ) + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_reset_traffic"), + IsAdminFilter(), +) +async def handle_reset_traffic( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_obj = await resolve_callback_key(session, tg_id, callback_data.data) + if not key_obj: + await callback_query.message.edit_text( + "❌ Ключ не найден в базе данных.", + reply_markup=build_editor_kb(tg_id), + ) + return + + email = key_obj.email + cluster_id = key_obj.server_id + + try: + await reset_traffic_in_cluster(cluster_id, email, session) + await callback_query.message.edit_text( + f"✅ Трафик для ключа {email} успешно сброшен.", + reply_markup=build_editor_kb(tg_id), + ) + except Exception as e: + logger.error(f"Ошибка при сбросе трафика: {e}") + await callback_query.message.edit_text( + "❌ Произошла ошибка при сбросе трафика. Попробуйте позже.", + reply_markup=build_editor_kb(tg_id), + ) + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_freeze"), + IsAdminFilter(), +) +async def handle_admin_freeze_subscription( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_obj = await resolve_callback_key(session, tg_id, callback_data.data) + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + email = key_obj.email + + try: + record = await get_key_details(session, email) + if not record: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + + client_id = record["client_id"] + cluster_id = record["server_id"] + + result = await toggle_client_on_cluster(cluster_id, email, client_id, enable=False, session=session) + if result["status"] != "success": + text_error = ( + f"Произошла ошибка при отключении подписки.\nДетали: {result.get('error') or result.get('results')}" + ) + await callback_query.message.edit_text( + text_error, + reply_markup=build_editor_kb(tg_id, True), + ) + return + + now_ms = int(time.time() * 1000) + time_left = record["expiry_time"] - now_ms + if time_left < 0: + time_left = 0 + + await mark_key_as_frozen(session, record["tg_id"], client_id, time_left) + await session.commit() + session.expire_all() + + await callback_query.answer("✅ Подписка отключена") + + await handle_key_edit( + callback_query=callback_query, + callback_data=callback_data, + session=session, + update=False, + ) + except Exception as e: + await handle_error(tg_id, callback_query, f"Ошибка при отключении подписки: {e}") + + +@router.callback_query( + AdminUserEditorCallback.filter(F.action == "users_unfreeze"), + IsAdminFilter(), +) +async def handle_admin_unfreeze_subscription( + callback_query: CallbackQuery, + callback_data: AdminUserEditorCallback, + session: AsyncSession, +): + tg_id = callback_data.tg_id + key_obj = await resolve_callback_key(session, tg_id, callback_data.data) + if not key_obj: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + email = key_obj.email + + try: + record = await get_key_details(session, email) + if not record: + await callback_query.message.edit_text( + text="🚫 Информация о ключе не найдена.", + reply_markup=build_editor_kb(tg_id), + ) + return + + client_id = record["client_id"] + cluster_id = record["server_id"] + + result = await toggle_client_on_cluster(cluster_id, email, client_id, enable=True, session=session) + if result["status"] != "success": + text_error = ( + f"Произошла ошибка при включении подписки.\nДетали: {result.get('error') or result.get('results')}" + ) + await callback_query.message.edit_text( + text_error, + reply_markup=build_editor_kb(tg_id, True), + ) + return + + tariff = await get_tariff_by_id(session, record["tariff_id"]) if record.get("tariff_id") else None + if not tariff: + total_gb = 0 + hwid_limit = 0 + else: + total_gb = int(tariff.get("traffic_limit") or 0) + hwid_limit = int(tariff.get("device_limit") or 0) + + if record.get("current_traffic_limit") is not None: + total_gb = record["current_traffic_limit"] + if record.get("current_device_limit") is not None: + hwid_limit = record["current_device_limit"] + + now_ms = int(time.time() * 1000) + leftover = record["expiry_time"] + if leftover < 0: + leftover = 0 + new_expiry_time = now_ms + leftover + + await mark_key_as_unfrozen(session, record["tg_id"], client_id, new_expiry_time) + await session.commit() + session.expire_all() + await release_session_early(session) + + await renew_key_in_cluster( + cluster_id=cluster_id, + email=email, + client_id=client_id, + new_expiry_time=new_expiry_time, + total_gb=total_gb, + session=session, + hwid_device_limit=hwid_limit, + reset_traffic=False, + plan=record.get("tariff_id"), + ) + + await callback_query.answer("✅ Подписка включена") + + await handle_key_edit( + callback_query=callback_query, + callback_data=callback_data, + session=session, + update=False, + ) + except Exception as e: + await handle_error(tg_id, callback_query, f"Ошибка при включении подписки: {e}") diff --git a/handlers/admin/users/users_manage.py b/handlers/admin/users/users_manage.py index 7149aa8d..cc151fca 100644 --- a/handlers/admin/users/users_manage.py +++ b/handlers/admin/users/users_manage.py @@ -1,6 +1,7 @@ import pytz from aiogram import F, Router, types +from logger import logger from aiogram.exceptions import TelegramBadRequest from aiogram.fsm.context import FSMContext from aiogram.types import ( @@ -19,6 +20,7 @@ from database import ( update_trial, ) from database.models import Admin, Key, ManualBan, Payment, Referral, User +from database.access.resolution import resolve_user_optional from filters.admin import IsAdminFilter from handlers.utils import sanitize_key_name from utils.csv_export import export_referrals_csv @@ -216,6 +218,20 @@ async def handle_send_user_message(callback_query: CallbackQuery, state: FSMCont text=text_message, parse_mode="HTML", ) + try: + import re + from database import async_session_maker + from database.web_notifications import notify_web + clean = re.sub(r"<[^>]+>", "", text_message or "").strip() + lines = clean.split("\n", 1) + title = lines[0][:120] + body = lines[1].strip()[:300] if len(lines) > 1 else "" + async with async_session_maker() as session: + await notify_web(session, tg_id=tg_id, type="message", title=title, message=body) + await session.commit() + except Exception as e: + logger.warning("[UserManage] Ошибка web-уведомления для tg_id={}: {}", tg_id, e) + await callback_query.message.edit_text( text="✅ Сообщение успешно отправлено.", reply_markup=build_editor_kb(tg_id), @@ -291,7 +307,7 @@ async def restore_trials(callback_query: types.CallbackQuery, session: AsyncSess update(User) .where( User.trial == 1, - ~exists(select(Key.tg_id).where(Key.tg_id == User.tg_id)), + ~exists(select(Key.user_id).where(Key.user_id == User.id)), ) .values(trial=0) ) @@ -340,9 +356,16 @@ async def process_user_search( ) -> None: await state.clear() - stmt_user = select(User.username, User.balance, User.created_at, User.updated_at, User.trial).where( - User.tg_id == tg_id - ) + u = await resolve_user_optional(session, tg_id) + if u is None: + await message.answer( + text="🚫 Пользователь с указанным ID не найден!", + reply_markup=build_admin_back_kb(), + ) + return + uid = u.id + + stmt_user = select(User.username, User.balance, User.created_at, User.updated_at, User.trial).where(User.id == uid) result_user = await session.execute(stmt_user) user_data = result_user.first() @@ -360,40 +383,43 @@ async def process_user_search( trial_status = "использован" if trial == 1 else "доступен" - stmt_ref_count = select(func.count()).select_from(Referral).where(Referral.referrer_tg_id == tg_id) + stmt_ref_count = select(func.count()).select_from(Referral).where(Referral.referrer_user_id == uid) result_ref = await session.execute(stmt_ref_count) referral_count = result_ref.scalar_one() - stmt_ref_by = select(Referral.referrer_tg_id).where(Referral.referred_tg_id == tg_id).limit(1) + stmt_ref_by = select(Referral.referrer_user_id).where(Referral.referred_user_id == uid).limit(1) result_ref_by = await session.execute(stmt_ref_by) - referrer_tg_id = result_ref_by.scalar_one_or_none() + referrer_uid = result_ref_by.scalar_one_or_none() referrer_text = None - if referrer_tg_id: - stmt_referrer = select(User.username).where(User.tg_id == referrer_tg_id) + if referrer_uid: + stmt_referrer = select(User.username, User.tg_id).where(User.id == referrer_uid) result_referrer = await session.execute(stmt_referrer) - ref_username = result_referrer.scalar_one_or_none() + ref_row = result_referrer.first() + ref_username = ref_row[0] if ref_row else None + ref_tg = ref_row[1] if ref_row else None + ref_label = int(ref_tg) if ref_tg is not None else int(referrer_uid) if ref_username: - referrer_text = f"🤝 Пригласил: @{ref_username} ({referrer_tg_id})" + referrer_text = f"🤝 Пригласил: @{ref_username} ({ref_label})" else: - referrer_text = f"🤝 Пригласил: {referrer_tg_id}" + referrer_text = f"🤝 Пригласил: {ref_label}" stmt = select( func.count(Payment.id), func.coalesce(func.sum(Payment.amount), 0), ).where( Payment.status == "success", - Payment.tg_id == tg_id, + Payment.user_id == uid, Payment.payment_system != "admin", ) result = await session.execute(stmt) topups_amount, topups_sum = result.one_or_none() or (0, 0) - stmt_keys = select(Key).where(Key.tg_id == tg_id) + stmt_keys = select(Key).where(Key.user_id == uid) result_keys = await session.execute(stmt_keys) key_records = result_keys.scalars().all() - stmt_ban = select(ManualBan).where(ManualBan.tg_id == tg_id).limit(1) + stmt_ban = select(ManualBan).where(ManualBan.user_id == uid).limit(1) result_ban = await session.execute(stmt_ban) ban_record = result_ban.scalar_one_or_none() diff --git a/handlers/admin/users/users_tariffs.py b/handlers/admin/users/users_tariffs.py index 2ef72c45..8ebb36e6 100644 --- a/handlers/admin/users/users_tariffs.py +++ b/handlers/admin/users/users_tariffs.py @@ -20,11 +20,11 @@ from database import ( from database.models import Tariff from filters.admin import IsAdminFilter from middlewares.session import release_session_early -from handlers.keys.operations import renew_key_in_cluster +from services.operations import renew_key_in_cluster from logger import logger from .keyboard import AdminUserEditorCallback, build_editor_kb -from .utils import resolve_admin_key +from services.users_utils import resolve_admin_key from .users_states import RenewTariffState from .users_keys import handle_key_edit diff --git a/handlers/buttons.py b/handlers/buttons.py index 0e1e98fe..352adb68 100644 --- a/handlers/buttons.py +++ b/handlers/buttons.py @@ -1,16 +1,13 @@ BUTTON_ICON_CONFIG: dict[str, dict[str, str]] = { } -# Общие кнопки BACK = "⬅️ Назад" APPLY = "✅ Подтвердить" CANCEL = "❌ Отмена" -# Кнопки подписки на канал SUB_CHANELL = "📢 Подписаться" SUB_CHANELL_DONE = "✅ Я подписался" -# Стартовое меню MAIN_MENU = "👤 Личный кабинет" TRIAL_SUB = "🎁 Пробная подписка" SUPPORT = "💬 Поддержка" @@ -18,7 +15,6 @@ CHANNEL = "📢 Канал" ABOUT_VPN = "💬 О сервисе" ADMIN_BTN = "📊 Администратор" -# Профиль ADD_SUB = "Купить новую подписку" MY_SUB = "🔐 Моя подписка" MY_SUBS = "📱 Мои подписки" @@ -27,7 +23,6 @@ INVITE = "👥 Пригласить" GIFTS = "🎁 Подарить" INSTRUCTIONS = "📘 Инструкции" -# Подписки CONNECT_DEVICE = "📲 Подключить устройство" ROUTER_BUTTON = "Подключить роутер" TV_BUTTON = "📺 Подключить Андроид TV" @@ -43,23 +38,19 @@ FREEZE = "Отключить подписку" UNFREEZE = "Включить подписку" ALIAS = "✏️" -# Реферальная система TOP_FIVE = "🏆 Топ-5" -# Меню баланса PAYMENT = "💳 Пополнить баланс" BALANCE_HISTORY = "📊 История пополнения" COUPON = "🎟️ Активировать купон" COUPON_RESTART = "🎟️ Попробовать другой купон" -# Подарки GIFT = "🎁 Подарить подписку" MY_GIFTS = "🎁 Мои подарки" GIFTS_MENU = "В меню подарков" SHARE_GIFT = "🎁 Поделиться подарком" GET_GIFT = "🎁 Получить подарок" -# Инструкции PC_PC = "💻 Windows" PC_MACOS = "🍏 macOS" DOWNLOAD_PC_BUTTON = "💻 Скачать Windows" @@ -69,7 +60,6 @@ CONNECT_MACOS_BUTTON = "🍏 Подключить" TV_CONTINUE = "▶ Продолжить" TV_INSTRUCTIONS = "📖 Полная инструкция" -# Подключение IPHONE = "🍏 Айфон" ANDROID = "🤖 Андроид" PC = "💻 Компьютер" @@ -80,13 +70,11 @@ IMPORT_IOS = "🍏 Подключить" IMPORT_ANDROID = "🤖 Подключить" MANUAL_INSTRUCTIONS = "📖 Ручная установка" -# Конфигуратор тарифов CONFIG_PAY_BUTTON_TEXT = "Оплатить {amount}" CONFIRM_ADDON_BUTTON_TEXT = "Подтвердить доплату {amount}" DOWNGRADE_ADDON_BUTTON_TEXT = "Понизить условия" DOWNGRADE_CONFIRM_BUTTON_TEXT = "Подтвердить понижение" -# Провайдеры оплаты YOOKASSA = "💳 ЮКасса: быстрая оплата" YOOMONEY = "💳 ЮМани: перевод по карте" FREEKASSA = "💰 FreeKassa: межд. платежи" @@ -103,18 +91,15 @@ KASSAI_SBP = "🏦 KassaAI: СБП" TRIBUTE = "💳 Tribute" HELEKET = "Heleket Crypto" -# Кнопки оплаты PAY = "Пополнить" PAY_2 = "Оплатить" CUSTOM_AMOUNT = "💰 Ввести свою сумму" STARS_BOT = "🤖 Бот для покупки звезд" DONAT_BUTTON = "💰 Поддержать проект" -# Выбор валюты RUB_CURRENCY = "₽ Рубли (RUB)" USD_CURRENCY = "$ USD / Cryptowallet" -# Уведомления TRIAL_BONUS = "🚀 Активировать пробный период" RENEW_KEY_NOTIFICATION = "🔄 Продлить подписку" CHANGE_TARIFF = "🔄 Сменить тариф" diff --git a/handlers/coupons.py b/handlers/coupons.py index 90960210..23bf6e3a 100644 --- a/handlers/coupons.py +++ b/handlers/coupons.py @@ -28,8 +28,8 @@ from database import ( ) from handlers.buttons import MAIN_MENU from middlewares.session import release_session_early -from handlers.keys.operations import renew_key_in_cluster -from handlers.payments.currency_rates import format_for_user +from services.operations import renew_key_in_cluster +from services.payments.currency_rates import format_for_user from handlers.profile import process_callback_view_profile from handlers.texts import ( COUPONS_DAYS_MESSAGE, @@ -132,13 +132,21 @@ async def activate_coupon( if coupon.amount > 0: try: - await update_balance(session, user_id, coupon.amount) - await update_coupon_usage_count(session, coupon.id) - await create_coupon_usage(session, coupon.id, user_id) - await add_payment(session, tg_id=user_id, amount=coupon.amount, payment_system="coupon") - amount_txt = await format_for_user(session, user_id, coupon.amount, language_code) + from services.coupons import apply_fixed_coupon + from services.errors import ServiceError + + result = await apply_fixed_coupon( + session=session, + user_id=user_id, + tg_id=user_id, + code=coupon_code, + ) + amount_txt = await format_for_user(session, user_id, result.amount, language_code) await message.answer(f"✅ Купон активирован, на баланс начислено {amount_txt}.") await state.clear() + except ServiceError as e: + await message.answer(f"❌ {e.message}") + await state.clear() except Exception as e: logger.error(f"Ошибка при активации купона на баланс: {e}") await message.answer("❌ Ошибка при активации купона.") @@ -195,6 +203,7 @@ async def handle_key_extension( admin: bool = False, ): from database.models import Coupon, Key, User + from database.access.resolution import resolve_user_optional parts = callback_query.data.split("|") client_id = parts[1] @@ -222,7 +231,12 @@ async def handle_key_extension( await state.clear() return - result = await session.execute(select(Key).where(Key.tg_id == tg_id, Key.client_id == client_id)) + owner = await resolve_user_optional(session, tg_id) + if owner is None: + await callback_query.message.edit_text("❌ Выбранная подписка не найдена или заморожена.") + await state.clear() + return + result = await session.execute(select(Key).where(Key.user_id == owner.id, Key.client_id == client_id)) key = result.scalar_one_or_none() if not key or key.is_frozen: await callback_query.message.edit_text("❌ Выбранная подписка не найдена или заморожена.") diff --git a/handlers/keys/key_create.py b/handlers/keys/key_create.py index e145ee78..a6361a22 100644 --- a/handlers/keys/key_create.py +++ b/handlers/keys/key_create.py @@ -24,8 +24,11 @@ from database import ( get_tariffs_for_cluster, get_trial, ) +from database.users import get_balance +from handlers.payments.fast_payment_flow import try_fast_payment_flow from database.models import Admin from database.notifications import check_hot_lead_discount +from database.access.resolution import notify_telegram_chat_id from database.tariffs import create_subgroup_hash, find_subgroup_by_hash, get_tariffs from handlers.admin.panel.keyboard import AdminPanelCallback from handlers.buttons import MAIN_MENU, BACK @@ -118,13 +121,46 @@ async def handle_key_creation( return trial_tariff = trial_tariffs[0] + trial_price = float(trial_tariff.get("price_rub", 0) or 0) base_days = trial_tariff["duration_days"] extra_days_value = int(NOTIFICATIONS_CONFIG.get("EXTRA_DAYS_AFTER_EXPIRY", NOTIFY_EXTRA_DAYS)) extra_days = extra_days_value if trial_status == -1 else 0 total_days = base_days + extra_days expiry_time = current_time + timedelta(days=total_days) - logger.info(f"[Trial] Доступен {total_days}-дневный триал для пользователя {tg_id}") + if trial_price > 0: + balance = await get_balance(session, tg_id) + if balance < trial_price: + shortfall = int(trial_price - balance) + handled = await try_fast_payment_flow( + message_or_query if isinstance(message_or_query, CallbackQuery) else None, + session, + state, + tg_id=tg_id, + temp_key="waiting_for_payment", + temp_payload={ + "payment_flow": "trial_purchase", + "tariff_id": trial_tariff["id"], + "selected_price_rub": int(trial_price), + "selected_duration_days": total_days, + }, + required_amount=shortfall, + ) + if handled: + return + builder = InlineKeyboardBuilder() + builder.row(InlineKeyboardButton(text="💳 Пополнить баланс", callback_data="pay")) + builder.row(InlineKeyboardButton(text=MAIN_MENU, callback_data="profile")) + await edit_or_send_message( + target_message=target_message, + text=f"💰 Пробная подписка стоит {int(trial_price)} ₽.\n" + f"Ваш баланс: {int(balance)} ₽.\n\n" + f"Пополните баланс для активации.", + reply_markup=builder.as_markup(), + ) + return + + logger.info(f"[Trial] Доступен {total_days}-дневный триал для пользователя {tg_id} (цена: {trial_price})") await edit_or_send_message( target_message=target_message, @@ -458,6 +494,7 @@ async def create_key( selected_traffic_gb: int | None = None, selected_price_rub: int | None = None, skip_balance_charge: bool | None = None, + is_trial: bool = False, ): from_user = message_or_query.from_user if isinstance(message_or_query, CallbackQuery | Message) else None if from_user: @@ -471,16 +508,25 @@ async def create_key( session=session, ) - use_country_selection = bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)) + country_cfg = bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)) + if country_cfg: + tg_notify = await notify_telegram_chat_id(session, tg_id) + can_deliver_country_ui = message_or_query is not None or tg_notify is not None + use_country_selection = can_deliver_country_ui + else: + use_country_selection = False if state and ( - skip_balance_charge is not None + is_trial + or skip_balance_charge is not None or any( value is not None for value in (selected_duration_days, selected_device_limit, selected_traffic_gb, selected_price_rub) ) ): state_data = await state.get_data() + if is_trial: + state_data["is_trial"] = True if selected_duration_days is not None: state_data["config_selected_duration_days"] = selected_duration_days if selected_device_limit is not None: @@ -519,4 +565,5 @@ async def create_key( selected_traffic_gb=selected_traffic_gb, selected_price_rub=selected_price_rub, skip_balance_charge=skip_balance_charge, + is_trial=is_trial or None, ) diff --git a/handlers/keys/key_mode/key_cluster_mode.py b/handlers/keys/key_mode/key_cluster_mode.py index 9c57ee27..e218be65 100644 --- a/handlers/keys/key_mode/key_cluster_mode.py +++ b/handlers/keys/key_mode/key_cluster_mode.py @@ -24,6 +24,7 @@ from database import ( update_trial, ) from database.models import Key +from database.access.resolution import notify_telegram_chat_id, resolve_user_optional from handlers.buttons import ( CONNECT_DEVICE, MAIN_MENU, @@ -32,10 +33,10 @@ from handlers.buttons import ( SUPPORT, TV_BUTTON, ) -from handlers.keys.operations import create_key_on_cluster +from services.operations import create_key_on_cluster from handlers.keys.utils import build_key_callback from database import get_vless_enabled -from handlers.tariffs.tariff_display import ( +from services.tariffs.tariff_display import ( build_key_created_message, get_effective_limits_for_key, resolve_price_to_charge, @@ -71,6 +72,7 @@ async def key_cluster_mode( selected_traffic_gb: int | None = None, selected_price_rub: int | None = None, skip_balance_charge: bool | None = None, + is_trial: bool | None = None, ): target_message = None safe_to_edit = False @@ -82,6 +84,8 @@ async def key_cluster_mode( target_message = message_or_query safe_to_edit = True + tg_notify = await notify_telegram_chat_id(session, tg_id) + while True: key_name = await generate_random_email(session=session) existing_key = await get_key_details(session, key_name) @@ -93,8 +97,19 @@ async def key_cluster_mode( expiry_timestamp = int(expiry_time.timestamp() * 1000) try: + owner = await resolve_user_optional(session, tg_id) + if owner is None: + error_message = "Пользователь не найден." + if safe_to_edit: + await edit_or_send_message(target_message=target_message, text=error_message, reply_markup=None) + elif tg_notify is not None: + await bot.send_message(chat_id=tg_notify, text=error_message) + return + uid = owner.id + data = await state.get_data() if state else {} - is_trial = data.get("is_trial", False) + if is_trial is None: + is_trial = data.get("is_trial", False) skip_balance_charge = bool(skip_balance_charge) if selected_device_limit is None: @@ -134,8 +149,8 @@ async def key_cluster_mode( text=error_message, reply_markup=None, ) - else: - await bot.send_message(chat_id=tg_id, text=error_message) + elif tg_notify is not None: + await bot.send_message(chat_id=tg_notify, text=error_message) return if device_limit is None: @@ -165,7 +180,7 @@ async def key_cluster_mode( await session.execute( update(Key) - .where(Key.tg_id == tg_id, Key.email == email) + .where(Key.user_id == uid, Key.email == email) .values( selected_device_limit=selected_device_limit, selected_traffic_limit=selected_traffic_gb, @@ -198,8 +213,8 @@ async def key_cluster_mode( text=error_message, reply_markup=None, ) - else: - await bot.send_message(chat_id=tg_id, text=error_message) + elif tg_notify is not None: + await bot.send_message(chat_id=tg_notify, text=error_message) return vless_enabled = False @@ -250,19 +265,23 @@ async def key_cluster_mode( builder.row(InlineKeyboardButton(text=SUPPORT, url=SUPPORT_CHAT_URL)) builder.row(InlineKeyboardButton(text=MAIN_MENU, callback_data="profile")) - if await process_intercept_key_creation_message( - chat_id=tg_id, + if tg_notify is not None and await process_intercept_key_creation_message( + chat_id=tg_notify, session=session, target_message=message_or_query, ): return - hook_commands = await process_key_creation_complete( - chat_id=tg_id, - admin=False, - session=session, - email=email, - key_name=key_name, + hook_commands = ( + await process_key_creation_complete( + chat_id=tg_notify, + admin=False, + session=session, + email=email, + key_name=key_name, + ) + if tg_notify is not None + else [] ) if hook_commands: builder = insert_hook_buttons(builder, hook_commands) @@ -283,9 +302,9 @@ async def key_cluster_mode( reply_markup=builder.as_markup(), media_path=default_media_path, ) - else: + elif tg_notify is not None: await bot.send_message( - chat_id=tg_id, + chat_id=tg_notify, text=key_message_text, reply_markup=builder.as_markup(), ) diff --git a/handlers/keys/key_mode/key_country_mode.py b/handlers/keys/key_mode/key_country_mode.py deleted file mode 100644 index 971c1362..00000000 --- a/handlers/keys/key_mode/key_country_mode.py +++ /dev/null @@ -1,903 +0,0 @@ -import asyncio -import uuid - -from datetime import datetime -from typing import Any - -import pytz - -from aiogram import F, Router -from aiogram.fsm.context import FSMContext -from aiogram.types import CallbackQuery, InlineKeyboardButton, Message, WebAppInfo -from aiogram.utils.keyboard import InlineKeyboardBuilder -from py3xui import AsyncApi -from sqlalchemy import func, select, update -from sqlalchemy.exc import SQLAlchemyError -from sqlalchemy.ext.asyncio import AsyncSession - -from bot import bot -from config import ( - ADMIN_PASSWORD, - ADMIN_USERNAME, - REMNAWAVE_LOGIN, - REMNAWAVE_PASSWORD, - REMNAWAVE_WEBAPP, - REMNAWAVE_WEBAPP_OPEN_IN_BROWSER, - SUPPORT_CHAT_URL, -) -from core.bootstrap import BUTTONS_CONFIG, MODES_CONFIG -from database import ( - add_user, - check_server_name_by_cluster, - check_user_exists, - filter_cluster_by_subgroup, - get_key_details, - get_tariff_by_id, - get_trial, - update_balance, - update_trial, -) -from database.models import Key, Server, ServerSpecialgroup -from handlers.buttons import ( - BACK, - CONNECT_DEVICE, - MAIN_MENU, - MY_SUB, - ROUTER_BUTTON, - SUPPORT, - TV_BUTTON, -) -from handlers.keys.utils import build_key_callback, resolve_key -from handlers.keys.operations import create_client_on_server -from handlers.keys.operations.aggregated_links import make_aggregated_link -from handlers.tariffs.tariff_display import ( - build_key_created_message, - get_effective_limits_for_key, -) -from handlers.texts import SELECT_COUNTRY_MSG -from handlers.utils import ( - ALLOWED_GROUP_CODES, - edit_or_send_message, - generate_random_email, - get_least_loaded_cluster, - is_full_remnawave_cluster, -) -from hooks.hook_buttons import insert_hook_buttons -from hooks.processors import ( - process_cluster_override, - process_intercept_key_creation_message, - process_key_creation_complete, - process_remnawave_webapp_override, -) -from logger import logger -from panels._3xui import delete_client, get_xui_instance -from panels.remnawave import RemnawaveAPI, get_vless_link_for_remnawave_by_username - - -router = Router() -moscow_tz = pytz.timezone("Europe/Moscow") -GB = 1024 * 1024 * 1024 - - -async def key_country_mode( - tg_id: int, - expiry_time: datetime, - state: FSMContext, - session: AsyncSession, - message_or_query: Message | CallbackQuery | None = None, - old_key_name: str | None = None, - plan: int | None = None, - selected_device_limit: int | None = None, - selected_traffic_gb: int | None = None, - selected_price_rub: int | None = None, - skip_balance_charge: bool | None = None, -): - target_message = None - safe_to_edit = False - - if state and plan: - await state.update_data(tariff_id=plan) - - if state and ( - skip_balance_charge is not None - or any(value is not None for value in (selected_device_limit, selected_traffic_gb, selected_price_rub)) - ): - data = await state.get_data() - if selected_device_limit is not None: - data["config_selected_device_limit"] = selected_device_limit - if selected_traffic_gb is not None: - data["config_selected_traffic_gb"] = selected_traffic_gb - if selected_price_rub is not None: - data["config_selected_price_rub"] = selected_price_rub - if skip_balance_charge is not None: - data["skip_balance_charge"] = skip_balance_charge - await state.set_data(data) - - if isinstance(message_or_query, CallbackQuery) and message_or_query.message: - target_message = message_or_query.message - safe_to_edit = True - elif isinstance(message_or_query, Message): - target_message = message_or_query - safe_to_edit = True - - data = await state.get_data() if state else {} - - forced_cluster = await process_cluster_override( - tg_id=tg_id, - state_data=data, - session=session, - plan=plan, - ) - if forced_cluster: - least_loaded_cluster = forced_cluster - else: - try: - least_loaded_cluster = await get_least_loaded_cluster(session) - except ValueError as e: - text = str(e) - if safe_to_edit: - await edit_or_send_message(target_message=target_message, text=text, reply_markup=None) - else: - await bot.send_message(chat_id=tg_id, text=text) - return - - subgroup_title = None - tariff: dict[str, Any] | None = None - if plan: - tariff = await get_tariff_by_id(session, int(plan)) - if tariff: - subgroup_title = tariff.get("subgroup_title") - - q = select( - Server.id, - Server.server_name, - Server.api_url, - Server.panel_type, - Server.enabled, - Server.max_keys, - ).where(Server.cluster_name == least_loaded_cluster) - servers = [dict(m) for m in (await session.execute(q)).mappings().all()] - - if not servers: - text = "❌ Нет доступных серверов в выбранном кластере." - if safe_to_edit: - await edit_or_send_message(target_message=target_message, text=text, reply_markup=None) - else: - await bot.send_message(chat_id=tg_id, text=text) - return - - server_ids = [s["id"] for s in servers] - groups_map: dict[int, list[str]] = {} - if server_ids: - r = await session.execute( - select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where( - ServerSpecialgroup.server_id.in_(server_ids) - ) - ) - for sid, gc in r.all(): - groups_map.setdefault(sid, []).append(gc) - - for server in servers: - server["special_groups"] = [g for g in groups_map.get(server["id"], []) if g in ALLOWED_GROUP_CODES] - - if subgroup_title: - servers = await filter_cluster_by_subgroup( - session, servers, subgroup_title, least_loaded_cluster, tariff_id=plan - ) - if not servers: - text = "❌ Нет доступных серверов в выбранном кластере." - if safe_to_edit: - await edit_or_send_message(target_message=target_message, text=text, reply_markup=None) - else: - await bot.send_message(chat_id=tg_id, text=text) - return - - special = None - if tariff: - gc = (tariff.get("group_code") or "").lower() - if gc in ALLOWED_GROUP_CODES: - special = gc - - if special: - bound_servers = [s for s in servers if special in (s.get("special_groups") or [])] - if bound_servers: - servers = bound_servers - - available_servers: list[str] = [] - tasks = [asyncio.create_task(check_server_availability(dict(server), session)) for server in servers] - results = await asyncio.gather(*tasks, return_exceptions=True) - - for server, result_ok in zip(servers, results, strict=False): - if result_ok is True: - available_servers.append(server["server_name"]) - - if not available_servers: - text = "❌ Нет доступных серверов в выбранном кластере." - if safe_to_edit: - await edit_or_send_message(target_message=target_message, text=text, reply_markup=None) - else: - await bot.send_message(chat_id=tg_id, text=text) - return - - builder = InlineKeyboardBuilder() - ts = int(expiry_time.timestamp()) - - for i in range(0, len(available_servers), 2): - row_buttons = [] - for server_name in available_servers[i : i + 2]: - if old_key_name: - callback_data = f"select_country|{server_name}|{ts}|{old_key_name}" - else: - if plan: - callback_data = f"select_country|{server_name}|{ts}||{plan}" - else: - callback_data = f"select_country|{server_name}|{ts}" - row_buttons.append(InlineKeyboardButton(text=server_name, callback_data=callback_data)) - builder.row(*row_buttons) - - builder.row(InlineKeyboardButton(text=MAIN_MENU, callback_data="profile")) - - if safe_to_edit: - await edit_or_send_message( - target_message=target_message, - text=SELECT_COUNTRY_MSG, - reply_markup=builder.as_markup(), - ) - else: - await bot.send_message( - chat_id=tg_id, - text=SELECT_COUNTRY_MSG, - reply_markup=builder.as_markup(), - ) - - -@router.callback_query(F.data.startswith("change_location|")) -async def change_location_callback(callback_query: CallbackQuery, session: Any): - try: - data = callback_query.data.split("|") - if len(data) < 2: - await callback_query.answer("❌ Некорректные данные", show_alert=True) - return - - old_key_ref = data[1] - key_obj = await resolve_key(session, callback_query.from_user.id, old_key_ref) - old_key_name = key_obj.email if key_obj else old_key_ref - record = await get_key_details(session, old_key_name) - if not record: - await callback_query.answer("❌ Ключ не найден", show_alert=True) - return - if record.get("tg_id") != callback_query.from_user.id: - await callback_query.answer("Доступ запрещён.", show_alert=True) - return - - expiry_timestamp = record["expiry_time"] - ts = int(expiry_timestamp / 1000) - current_server = record["server_id"] - - cluster_info = await check_server_name_by_cluster(session, current_server) - if not cluster_info: - await callback_query.answer("❌ Кластер для текущего сервера не найден", show_alert=True) - return - - cluster_name = cluster_info["cluster_name"] - - key_tariff_id = record.get("tariff_id") - tariff_dict: dict[str, Any] | None = None - subgroup_title = None - if key_tariff_id: - tariff_dict = await get_tariff_by_id(session, int(key_tariff_id)) - if tariff_dict: - subgroup_title = tariff_dict.get("subgroup_title") - - q = ( - select( - Server.id, - Server.server_name, - Server.api_url, - Server.panel_type, - Server.enabled, - Server.max_keys, - ) - .where(Server.cluster_name == cluster_name) - .where(Server.server_name != current_server) - ) - servers = [dict(m) for m in (await session.execute(q)).mappings().all()] - if not servers: - await callback_query.answer("❌ Доступных серверов в кластере не найдено", show_alert=True) - return - - server_ids = [s["id"] for s in servers] - groups_map: dict[int, list[str]] = {} - if server_ids: - r = await session.execute( - select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where( - ServerSpecialgroup.server_id.in_(server_ids) - ) - ) - for sid, gc in r.all(): - groups_map.setdefault(sid, []).append(gc) - - for server in servers: - server["special_groups"] = [g for g in groups_map.get(server["id"], []) if g in ALLOWED_GROUP_CODES] - - available_servers: list[str] = [] - tasks = [ - asyncio.create_task( - check_server_availability( - { - "server_name": s["server_name"], - "api_url": s["api_url"], - "panel_type": s["panel_type"], - "enabled": s.get("enabled", True), - "max_keys": s.get("max_keys"), - }, - session, - ) - ) - for s in servers - ] - results = await asyncio.gather(*tasks, return_exceptions=True) - for server, result_ok in zip(servers, results, strict=False): - if result_ok is True: - available_servers.append(server["server_name"]) - - if subgroup_title and available_servers: - available_servers_dict = [s for s in servers if s["server_name"] in available_servers] - filtered_servers = await filter_cluster_by_subgroup( - session, - available_servers_dict, - subgroup_title.strip(), - cluster_name, - tariff_id=key_tariff_id, - ) - if filtered_servers: - available_servers = [s["server_name"] for s in filtered_servers] - else: - builder = InlineKeyboardBuilder() - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=build_key_callback("view_key", record.get("client_id"), old_key_name), - ) - ) - await edit_or_send_message( - target_message=callback_query.message, - text="❌ Нет доступных стран для смены локации.", - reply_markup=builder.as_markup(), - ) - return - - if available_servers and tariff_dict: - special = None - gc = (tariff_dict.get("group_code") or "").lower() - if gc and gc in ALLOWED_GROUP_CODES: - special = gc - - if special: - available_servers_dict = [s for s in servers if s["server_name"] in available_servers] - bound_servers = [s for s in available_servers_dict if special in (s.get("special_groups") or [])] - if bound_servers: - available_servers = [s["server_name"] for s in bound_servers] - - if not available_servers: - builder = InlineKeyboardBuilder() - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=build_key_callback("view_key", record.get("client_id"), old_key_name), - ) - ) - await edit_or_send_message( - target_message=callback_query.message, - text="❌ Нет доступных стран для смены локации.", - reply_markup=builder.as_markup(), - ) - return - - builder = InlineKeyboardBuilder() - - for i in range(0, len(available_servers), 2): - row_buttons = [] - for country in available_servers[i : i + 2]: - callback_data = f"select_country|{country}|{ts}|{old_key_ref}" - row_buttons.append(InlineKeyboardButton(text=country, callback_data=callback_data)) - builder.row(*row_buttons) - - builder.row( - InlineKeyboardButton( - text=BACK, - callback_data=build_key_callback("view_key", record.get("client_id"), old_key_name), - ) - ) - - await edit_or_send_message( - target_message=callback_query.message, - text="🌍 Пожалуйста, выберите новую локацию для вашей подписки:", - reply_markup=builder.as_markup(), - media_path=None, - ) - except Exception as e: - logger.error(f"Ошибка при смене локации для пользователя {callback_query.from_user.id}: {e}") - await callback_query.answer("❌ Ошибка смены локации. Попробуйте снова.", show_alert=True) - - -@router.callback_query(F.data.startswith("select_country|")) -async def handle_country_selection(callback_query: CallbackQuery, session: Any, state: FSMContext): - data = callback_query.data.split("|") - if len(data) < 3: - await callback_query.message.answer("❌ Некорректные данные. Попробуйте снова.") - return - - selected_country = data[1] - try: - ts = int(data[2]) - except ValueError: - await callback_query.message.answer("❌ Некорректное время истечения. Попробуйте снова.") - return - - old_key_name = data[3] if len(data) > 3 and data[3] else None - try: - tariff_id = int(data[4]) if len(data) > 4 and data[4] else None - except (ValueError, IndexError): - tariff_id = None - - tg_id = callback_query.from_user.id - - fsm_data = await state.get_data() - if fsm_data.get("creating_key"): - try: - await callback_query.answer("⏳ Уже обрабатываю…") - except Exception: - pass - return - - await state.update_data(creating_key=True) - - try: - await callback_query.answer("Обрабатываю…") - if callback_query.message: - await callback_query.message.edit_reply_markup(reply_markup=None) - except Exception: - pass - - try: - expiry_time = datetime.fromtimestamp(ts, tz=moscow_tz) - await finalize_key_creation( - tg_id=tg_id, - expiry_time=expiry_time, - selected_country=selected_country, - state=state, - session=session, - callback_query=callback_query, - old_key_name=old_key_name, - tariff_id=tariff_id, - ) - finally: - fsm_data = await state.get_data() - if fsm_data.get("creating_key"): - await state.update_data(creating_key=False) - - -async def finalize_key_creation( - tg_id: int, - expiry_time: datetime, - selected_country: str, - state: FSMContext | None, - session: AsyncSession, - callback_query: CallbackQuery, - old_key_name: str | None = None, - tariff_id: int | None = None, -): - from_user = callback_query.from_user - - if not await check_user_exists(session, tg_id): - await add_user( - session=session, - tg_id=from_user.id, - username=from_user.username, - first_name=from_user.first_name, - last_name=from_user.last_name, - language_code=from_user.language_code, - is_bot=from_user.is_bot, - ) - - expiry_time = expiry_time.astimezone(moscow_tz) - - old_key_details: dict[str, Any] | None = None - if old_key_name: - key_obj = await resolve_key(session, tg_id, old_key_name) - old_key_name = key_obj.email if key_obj else old_key_name - old_key_details = await get_key_details(session, old_key_name) - if not old_key_details: - await callback_query.message.answer("❌ Ключ не найден. Попробуйте снова.") - return - key_name = old_key_name - client_id = old_key_details["client_id"] - email = old_key_details["email"] - expiry_timestamp = old_key_details["expiry_time"] - tariff_id = old_key_details.get("tariff_id") or tariff_id - else: - while True: - key_name = await generate_random_email(session=session) - existing_key = await get_key_details(session, key_name) - if not existing_key: - break - client_id = str(uuid.uuid4()) - email = key_name.lower() - expiry_timestamp = int(expiry_time.timestamp() * 1000) - - data = await state.get_data() if state else {} - is_trial = data.get("is_trial", False) - skip_balance_charge = bool(data.get("skip_balance_charge", False)) - - selected_traffic_gb = data.get("config_selected_traffic_gb") - if selected_traffic_gb is None: - selected_traffic_gb = data.get("selected_traffic_limit_gb") - - selected_device_limit = data.get("config_selected_device_limit") - if selected_device_limit is None: - selected_device_limit = data.get("selected_device_limit") - - if old_key_details: - if selected_traffic_gb is None: - stored_traffic = old_key_details.get("selected_traffic_limit") - if stored_traffic is not None: - selected_traffic_gb = int(stored_traffic) - if selected_device_limit is None: - stored_devices = old_key_details.get("selected_device_limit") - if stored_devices is not None: - selected_device_limit = int(stored_devices) - - price_to_charge = data.get("selected_price_rub") - - effective_tariff_id = data.get("tariff_id") or tariff_id - tariff: dict[str, Any] | None = None - if effective_tariff_id: - tariff_id = int(effective_tariff_id) - tariff = await get_tariff_by_id(session, tariff_id) - - device_limit, traffic_limit_bytes = await get_effective_limits_for_key( - session=session, - tariff_id=tariff_id, - selected_device_limit=selected_device_limit, - selected_traffic_gb=selected_traffic_gb, - ) - - if selected_traffic_gb is not None: - traffic_limit_gb = int(selected_traffic_gb) - else: - traffic_limit_gb = int(traffic_limit_bytes / GB) if traffic_limit_bytes else 0 - - if price_to_charge is None and tariff and not old_key_name: - price_to_charge = tariff.get("price_rub") - - need_vless_key = bool(tariff.get("vless")) if tariff else False - - public_link = None - remnawave_link = None - created_at = int(datetime.now(moscow_tz).timestamp() * 1000) - - try: - result = await session.execute(select(Server).where(Server.server_name == selected_country)) - server_info = result.scalar_one_or_none() - if not server_info: - raise ValueError(f"Сервер {selected_country} не найден") - - cluster_info = await check_server_name_by_cluster(session, server_info.server_name) - if not cluster_info: - raise ValueError(f"Кластер для сервера {server_info.server_name} не найден") - - cluster_name = cluster_info["cluster_name"] - is_full_remnawave = await is_full_remnawave_cluster(cluster_name, session) - - if old_key_name and old_key_details: - old_server_id = old_key_details["server_id"] - if old_server_id: - result = await session.execute(select(Server).where(Server.server_name == old_server_id)) - old_server_info = result.scalar_one_or_none() - if old_server_info: - try: - if old_server_info.panel_type.lower() == "3x-ui": - xui = await get_xui_instance(old_server_info.api_url) - await delete_client(xui, old_server_info.inbound_id, email, client_id) - await session.execute( - update(Key).where(Key.tg_id == tg_id, Key.email == email).values(key=None) - ) - elif old_server_info.panel_type.lower() == "remnawave": - remna_del = RemnawaveAPI(old_server_info.api_url) - if await remna_del.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD): - await remna_del.delete_user(client_id) - await session.execute( - update(Key) - .where(Key.tg_id == tg_id, Key.email == email) - .values(remnawave_link=None) - ) - except Exception as e: - logger.warning(f"[Delete] Ошибка при удалении клиента: {e}") - - panel_type = server_info.panel_type.lower() - - if panel_type == "remnawave" or is_full_remnawave: - remna = RemnawaveAPI(server_info.api_url) - if not await remna.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD): - raise ValueError(f"❌ Не удалось авторизоваться в Remnawave ({server_info.server_name})") - - expire_at = datetime.utcfromtimestamp(expiry_timestamp / 1000).isoformat() + "Z" - user_data: dict[str, Any] = { - "username": email, - "trafficLimitStrategy": "NO_RESET", - "expireAt": expire_at, - "telegramId": tg_id, - "activeInternalSquads": [server_info.inbound_id], - "uuid": client_id, - } - if traffic_limit_bytes: - user_data["trafficLimitBytes"] = traffic_limit_bytes - if device_limit: - user_data["hwidDeviceLimit"] = device_limit - - result = await remna.create_user(user_data) - if not result: - raise ValueError("❌ Ошибка при создании пользователя в Remnawave") - - client_id = result.get("uuid") or result.get("id") or client_id - - remnawave_link = None - if need_vless_key: - try: - vless_link = await get_vless_link_for_remnawave_by_username(remna, email, email) - except Exception: - vless_link = None - if vless_link: - remnawave_link = vless_link - - if not remnawave_link: - try: - sub = await remna.get_subscription_by_username(email) - except Exception: - sub = None - - if sub: - if need_vless_key and not remnawave_link: - links = sub.get("links") or [] - remnawave_link = next( - (l for l in links if isinstance(l, str) and l.lower().startswith("vless://")), - None, - ) - - if not remnawave_link: - remnawave_link = sub.get("subscriptionUrl") - - if old_key_name: - await session.execute( - update(Key).where(Key.tg_id == tg_id, Key.email == email).values(client_id=client_id) - ) - - if panel_type == "3x-ui": - semaphore = asyncio.Semaphore(2) - await create_client_on_server( - server_info={ - "api_url": server_info.api_url, - "inbound_id": server_info.inbound_id, - "server_name": server_info.server_name, - "panel_type": server_info.panel_type, - }, - tg_id=tg_id, - client_id=client_id, - email=email, - expiry_timestamp=expiry_timestamp, - semaphore=semaphore, - session=session, - plan=tariff_id, - is_trial=is_trial, - total_traffic_limit_bytes=traffic_limit_bytes, - device_limit_value=device_limit, - ) - - subgroup_code = tariff.get("subgroup_title") if tariff and tariff.get("subgroup_title") else None - cluster_all = [ - { - "server_name": server_info.server_name, - "api_url": server_info.api_url, - "panel_type": server_info.panel_type, - "inbound_id": getattr(server_info, "inbound_id", None), - "enabled": True, - "max_keys": getattr(server_info, "max_keys", None), - } - ] - - link_to_show = await make_aggregated_link( - session=session, - cluster_all=cluster_all, - cluster_id=cluster_name, - email=email, - client_id=client_id, - tg_id=tg_id, - subgroup_code=subgroup_code, - remna_link_override=remnawave_link, - plan=tariff_id, - ) - - public_link = link_to_show - - if old_key_name: - update_data: dict[str, Any] = { - "server_id": selected_country, - "key": None, - "remnawave_link": None, - } - if public_link and public_link.startswith("vless://"): - update_data["key"] = public_link - elif public_link and public_link.startswith("http"): - update_data["key"] = public_link - if remnawave_link: - update_data["remnawave_link"] = remnawave_link - await session.execute(update(Key).where(Key.tg_id == tg_id, Key.email == email).values(**update_data)) - else: - new_key = Key( - tg_id=tg_id, - client_id=client_id, - email=email, - created_at=created_at, - expiry_time=expiry_timestamp, - key=public_link if public_link else None, - remnawave_link=remnawave_link, - server_id=selected_country, - tariff_id=tariff_id, - selected_device_limit=int(selected_device_limit) if selected_device_limit is not None else None, - selected_traffic_limit=int(selected_traffic_gb) if selected_traffic_gb is not None else None, - selected_price_rub=int(price_to_charge) if price_to_charge is not None else None, - ) - session.add(new_key) - if is_trial: - trial_status = await get_trial(session, tg_id) - if trial_status in [0, -1]: - await update_trial(session, tg_id, 1) - if not is_trial and price_to_charge and not skip_balance_charge: - await update_balance(session, tg_id, -int(price_to_charge)) - - if state: - await state.update_data(skip_balance_charge=False) - - await session.commit() - - except Exception as e: - logger.error(f"[Key Finalize] Ошибка при создании ключа для пользователя {tg_id}: {e}") - await callback_query.message.answer("❌ Произошла ошибка при создании подписки. Попробуйте снова.") - return - - builder = InlineKeyboardBuilder() - is_full_remnawave = await is_full_remnawave_cluster(cluster_name, session) - is_vless = bool(public_link and public_link.lower().startswith("vless://")) or bool(need_vless_key) - final_link = public_link or remnawave_link - webapp_url = ( - final_link - if isinstance(final_link, str) and final_link.strip().lower().startswith(("http://", "https://")) - else None - ) - - use_webapp = bool(MODES_CONFIG.get("REMNAWAVE_WEBAPP_ENABLED", REMNAWAVE_WEBAPP)) - open_in_browser = bool(MODES_CONFIG.get("REMNAWAVE_WEBAPP_OPEN_IN_BROWSER", REMNAWAVE_WEBAPP_OPEN_IN_BROWSER)) - if use_webapp and webapp_url: - use_webapp = await process_remnawave_webapp_override( - remnawave_webapp=use_webapp, - final_link=final_link, - session=session, - ) - - tv_button_enabled = bool(BUTTONS_CONFIG.get("ANDROID_TV_BUTTON_ENABLE")) - - if panel_type == "remnawave" or is_full_remnawave: - if is_vless: - builder.row(InlineKeyboardButton(text=ROUTER_BUTTON, callback_data=build_key_callback("connect_router", client_id, key_name))) - else: - if use_webapp and webapp_url: - if open_in_browser: - builder.row(InlineKeyboardButton(text=CONNECT_DEVICE, url=webapp_url)) - else: - builder.row(InlineKeyboardButton(text=CONNECT_DEVICE, web_app=WebAppInfo(url=webapp_url))) - if tv_button_enabled: - builder.row(InlineKeyboardButton(text=TV_BUTTON, callback_data=build_key_callback("connect_tv", client_id, key_name))) - else: - builder.row( - InlineKeyboardButton( - text=CONNECT_DEVICE, - callback_data=build_key_callback("connect_device", client_id, key_name), - ) - ) - else: - builder.row( - InlineKeyboardButton( - text=CONNECT_DEVICE, - callback_data=build_key_callback("connect_device", client_id, key_name), - ) - ) - - builder.row(InlineKeyboardButton(text=MY_SUB, callback_data=build_key_callback("view_key", client_id, key_name))) - builder.row(InlineKeyboardButton(text=SUPPORT, url=SUPPORT_CHAT_URL)) - builder.row(InlineKeyboardButton(text=MAIN_MENU, callback_data="profile")) - - if await process_intercept_key_creation_message( - chat_id=tg_id, - session=session, - target_message=callback_query, - ): - return - - hook_commands = await process_key_creation_complete( - chat_id=tg_id, - admin=False, - session=session, - email=email, - key_name=key_name, - ) - if hook_commands: - builder = insert_hook_buttons(builder, hook_commands) - - key_record = await get_key_details(session, key_name) - final_link_for_message = final_link or (key_record.get("link") if key_record else None) or "Ссылка не найдена" - message_text = await build_key_created_message( - session=session, - key_record=key_record, - final_link=final_link_for_message, - selected_device_limit=selected_device_limit, - selected_traffic_gb=selected_traffic_gb, - ) - - await edit_or_send_message( - target_message=callback_query.message, - text=message_text, - reply_markup=builder.as_markup(), - media_path="img/pic.jpg", - ) - - if state: - await state.clear() - - -async def check_server_availability(server_info: dict, session: AsyncSession) -> bool: - server_name = server_info.get("server_name", "unknown") - panel_type = server_info.get("panel_type", "3x-ui").lower() - enabled = server_info.get("enabled", True) - max_keys = server_info.get("max_keys") - - if not enabled: - logger.info(f"[Ping] Сервер {server_name} выключен (enabled = FALSE).") - return False - - try: - if max_keys is not None: - result = await session.execute(select(func.count()).select_from(Key).where(Key.server_id == server_name)) - key_count = result.scalar() - - if key_count >= max_keys: - logger.info(f"[Ping] Сервер {server_name} достиг лимита ключей: {key_count}/{max_keys}.") - return False - - except SQLAlchemyError as e: - logger.warning(f"[Ping] Ошибка при проверке лимита ключей на сервере {server_name}: {e}") - return False - - try: - if panel_type == "remnawave": - remna = RemnawaveAPI(server_info["api_url"]) - await asyncio.wait_for(remna.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD), timeout=5.0) - logger.info(f"[Ping] Remnawave сервер {server_name} доступен.") - return True - - xui = AsyncApi( - server_info["api_url"], - username=ADMIN_USERNAME, - password=ADMIN_PASSWORD, - logger=logger, - ) - await asyncio.wait_for(xui.login(), timeout=5.0) - logger.info(f"[Ping] 3x-ui сервер {server_name} доступен.") - return True - - except TimeoutError: - logger.warning(f"[Ping] Сервер {server_name} не ответил вовремя.") - return False - except Exception as e: - logger.warning(f"[Ping] Ошибка при проверке сервера {server_name}: {e}") - return False diff --git a/handlers/keys/key_mode/key_country_mode/__init__.py b/handlers/keys/key_mode/key_country_mode/__init__.py new file mode 100644 index 00000000..a1fcb8b8 --- /dev/null +++ b/handlers/keys/key_mode/key_country_mode/__init__.py @@ -0,0 +1,16 @@ +from ._common import router +from .entry import handle_country_selection, key_country_mode +from .finalize import ( + _legacy_check_server_availability, + check_server_availability, + finalize_key_creation, +) +from . import change_location # noqa: F401 — trigger endpoint registration + +__all__ = [ + "router", + "key_country_mode", + "handle_country_selection", + "finalize_key_creation", + "check_server_availability", +] diff --git a/handlers/keys/key_mode/key_country_mode/_common.py b/handlers/keys/key_mode/key_country_mode/_common.py new file mode 100644 index 00000000..f2d3feee --- /dev/null +++ b/handlers/keys/key_mode/key_country_mode/_common.py @@ -0,0 +1,81 @@ +import asyncio +import uuid + +from datetime import datetime +from typing import Any + +import pytz + +from aiogram import F, Router +from aiogram.fsm.context import FSMContext +from aiogram.types import CallbackQuery, InlineKeyboardButton, Message, WebAppInfo +from aiogram.utils.keyboard import InlineKeyboardBuilder +from py3xui import AsyncApi +from sqlalchemy import func, select, update +from sqlalchemy.exc import SQLAlchemyError +from sqlalchemy.ext.asyncio import AsyncSession + +from bot import bot +from config import ( + ADMIN_PASSWORD, + ADMIN_USERNAME, + REMNAWAVE_LOGIN, + REMNAWAVE_PASSWORD, + REMNAWAVE_WEBAPP, + REMNAWAVE_WEBAPP_OPEN_IN_BROWSER, + SUPPORT_CHAT_URL, +) +from core.bootstrap import BUTTONS_CONFIG, MODES_CONFIG +from database import ( + add_user, + check_server_name_by_cluster, + check_user_exists, + filter_cluster_by_subgroup, + get_key_details, + get_tariff_by_id, + get_trial, + update_balance, + update_trial, +) +from database.models import Key, Server, ServerSpecialgroup +from database.access.resolution import notify_telegram_chat_id, resolve_user_optional +from handlers.buttons import ( + BACK, + CONNECT_DEVICE, + MAIN_MENU, + MY_SUB, + ROUTER_BUTTON, + SUPPORT, + TV_BUTTON, +) +from handlers.keys.utils import build_key_callback, resolve_key +from services.operations import create_client_on_server +from services.operations.aggregated_links import make_aggregated_link +from services.tariffs.tariff_display import ( + build_key_created_message, + get_effective_limits_for_key, +) +from handlers.texts import SELECT_COUNTRY_MSG +from handlers.utils import ( + ALLOWED_GROUP_CODES, + edit_or_send_message, + generate_random_email, + get_least_loaded_cluster, + is_full_remnawave_cluster, +) +from hooks.hook_buttons import insert_hook_buttons +from hooks.processors import ( + process_cluster_override, + process_intercept_key_creation_message, + process_key_creation_complete, + process_remnawave_webapp_override, +) +from logger import logger +from panels._3xui import delete_client, get_xui_instance +from panels.remnawave import RemnawaveAPI, get_vless_link_for_remnawave_by_username + + +router = Router() +moscow_tz = pytz.timezone("Europe/Moscow") +GB = 1024 * 1024 * 1024 + diff --git a/handlers/keys/key_mode/key_country_mode/change_location.py b/handlers/keys/key_mode/key_country_mode/change_location.py new file mode 100644 index 00000000..52c57e22 --- /dev/null +++ b/handlers/keys/key_mode/key_country_mode/change_location.py @@ -0,0 +1,171 @@ +from ._common import * # noqa: F401,F403 +from ._common import router # noqa: F401 + +@router.callback_query(F.data.startswith("change_location|")) +async def change_location_callback(callback_query: CallbackQuery, session: Any): + try: + data = callback_query.data.split("|") + if len(data) < 2: + await callback_query.answer("❌ Некорректные данные", show_alert=True) + return + + old_key_ref = data[1] + key_obj = await resolve_key(session, callback_query.from_user.id, old_key_ref) + old_key_name = key_obj.email if key_obj else old_key_ref + record = await get_key_details(session, old_key_name) + if not record: + await callback_query.answer("❌ Ключ не найден", show_alert=True) + return + if record.get("tg_id") != callback_query.from_user.id: + await callback_query.answer("Доступ запрещён.", show_alert=True) + return + + expiry_timestamp = record["expiry_time"] + ts = int(expiry_timestamp / 1000) + current_server = record["server_id"] + + cluster_info = await check_server_name_by_cluster(session, current_server) + if not cluster_info: + await callback_query.answer("❌ Кластер для текущего сервера не найден", show_alert=True) + return + + cluster_name = cluster_info["cluster_name"] + + key_tariff_id = record.get("tariff_id") + tariff_dict: dict[str, Any] | None = None + subgroup_title = None + if key_tariff_id: + tariff_dict = await get_tariff_by_id(session, int(key_tariff_id)) + if tariff_dict: + subgroup_title = tariff_dict.get("subgroup_title") + + q = ( + select( + Server.id, + Server.server_name, + Server.api_url, + Server.panel_type, + Server.enabled, + Server.max_keys, + ) + .where(Server.cluster_name == cluster_name) + .where(Server.server_name != current_server) + ) + servers = [dict(m) for m in (await session.execute(q)).mappings().all()] + if not servers: + await callback_query.answer("❌ Доступных серверов в кластере не найдено", show_alert=True) + return + + server_ids = [s["id"] for s in servers] + groups_map: dict[int, list[str]] = {} + if server_ids: + r = await session.execute( + select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where( + ServerSpecialgroup.server_id.in_(server_ids) + ) + ) + for sid, gc in r.all(): + groups_map.setdefault(sid, []).append(gc) + + for server in servers: + server["special_groups"] = [g for g in groups_map.get(server["id"], []) if g in ALLOWED_GROUP_CODES] + + available_servers: list[str] = [] + tasks = [ + asyncio.create_task( + check_server_availability( + { + "server_name": s["server_name"], + "api_url": s["api_url"], + "panel_type": s["panel_type"], + "enabled": s.get("enabled", True), + "max_keys": s.get("max_keys"), + }, + session, + ) + ) + for s in servers + ] + results = await asyncio.gather(*tasks, return_exceptions=True) + for server, result_ok in zip(servers, results, strict=False): + if result_ok is True: + available_servers.append(server["server_name"]) + + if subgroup_title and available_servers: + available_servers_dict = [s for s in servers if s["server_name"] in available_servers] + filtered_servers = await filter_cluster_by_subgroup( + session, + available_servers_dict, + subgroup_title.strip(), + cluster_name, + tariff_id=key_tariff_id, + ) + if filtered_servers: + available_servers = [s["server_name"] for s in filtered_servers] + else: + builder = InlineKeyboardBuilder() + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=build_key_callback("view_key", record.get("client_id"), old_key_name), + ) + ) + await edit_or_send_message( + target_message=callback_query.message, + text="❌ Нет доступных стран для смены локации.", + reply_markup=builder.as_markup(), + ) + return + + if available_servers and tariff_dict: + special = None + gc = (tariff_dict.get("group_code") or "").lower() + if gc and gc in ALLOWED_GROUP_CODES: + special = gc + + if special: + available_servers_dict = [s for s in servers if s["server_name"] in available_servers] + bound_servers = [s for s in available_servers_dict if special in (s.get("special_groups") or [])] + if bound_servers: + available_servers = [s["server_name"] for s in bound_servers] + + if not available_servers: + builder = InlineKeyboardBuilder() + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=build_key_callback("view_key", record.get("client_id"), old_key_name), + ) + ) + await edit_or_send_message( + target_message=callback_query.message, + text="❌ Нет доступных стран для смены локации.", + reply_markup=builder.as_markup(), + ) + return + + builder = InlineKeyboardBuilder() + + for i in range(0, len(available_servers), 2): + row_buttons = [] + for country in available_servers[i : i + 2]: + callback_data = f"select_country|{country}|{ts}|{old_key_ref}" + row_buttons.append(InlineKeyboardButton(text=country, callback_data=callback_data)) + builder.row(*row_buttons) + + builder.row( + InlineKeyboardButton( + text=BACK, + callback_data=build_key_callback("view_key", record.get("client_id"), old_key_name), + ) + ) + + await edit_or_send_message( + target_message=callback_query.message, + text="🌍 Пожалуйста, выберите новую локацию для вашей подписки:", + reply_markup=builder.as_markup(), + media_path=None, + ) + except Exception as e: + logger.error(f"Ошибка при смене локации для пользователя {callback_query.from_user.id}: {e}") + await callback_query.answer("❌ Ошибка смены локации. Попробуйте снова.", show_alert=True) diff --git a/handlers/keys/key_mode/key_country_mode/entry.py b/handlers/keys/key_mode/key_country_mode/entry.py new file mode 100644 index 00000000..e98079d7 --- /dev/null +++ b/handlers/keys/key_mode/key_country_mode/entry.py @@ -0,0 +1,232 @@ +from ._common import * # noqa: F401,F403 +from ._common import router # noqa: F401 + +async def key_country_mode( + tg_id: int, + expiry_time: datetime, + state: FSMContext, + session: AsyncSession, + message_or_query: Message | CallbackQuery | None = None, + old_key_name: str | None = None, + plan: int | None = None, + selected_device_limit: int | None = None, + selected_traffic_gb: int | None = None, + selected_price_rub: int | None = None, + skip_balance_charge: bool | None = None, +): + target_message = None + safe_to_edit = False + + if state and plan: + await state.update_data(tariff_id=plan) + + if state and ( + skip_balance_charge is not None + or any(value is not None for value in (selected_device_limit, selected_traffic_gb, selected_price_rub)) + ): + data = await state.get_data() + if selected_device_limit is not None: + data["config_selected_device_limit"] = selected_device_limit + if selected_traffic_gb is not None: + data["config_selected_traffic_gb"] = selected_traffic_gb + if selected_price_rub is not None: + data["config_selected_price_rub"] = selected_price_rub + if skip_balance_charge is not None: + data["skip_balance_charge"] = skip_balance_charge + await state.set_data(data) + + if isinstance(message_or_query, CallbackQuery) and message_or_query.message: + target_message = message_or_query.message + safe_to_edit = True + elif isinstance(message_or_query, Message): + target_message = message_or_query + safe_to_edit = True + + tg_notify = await notify_telegram_chat_id(session, tg_id) + + data = await state.get_data() if state else {} + + forced_cluster = await process_cluster_override( + tg_id=tg_id, + state_data=data, + session=session, + plan=plan, + ) + if forced_cluster: + least_loaded_cluster = forced_cluster + else: + try: + least_loaded_cluster = await get_least_loaded_cluster(session) + except ValueError as e: + text = str(e) + if safe_to_edit: + await edit_or_send_message(target_message=target_message, text=text, reply_markup=None) + elif tg_notify is not None: + await bot.send_message(chat_id=tg_notify, text=text) + return + + subgroup_title = None + tariff: dict[str, Any] | None = None + if plan: + tariff = await get_tariff_by_id(session, int(plan)) + if tariff: + subgroup_title = tariff.get("subgroup_title") + + q = select( + Server.id, + Server.server_name, + Server.api_url, + Server.panel_type, + Server.enabled, + Server.max_keys, + ).where(Server.cluster_name == least_loaded_cluster) + servers = [dict(m) for m in (await session.execute(q)).mappings().all()] + + if not servers: + text = "❌ Нет доступных серверов в выбранном кластере." + if safe_to_edit: + await edit_or_send_message(target_message=target_message, text=text, reply_markup=None) + elif tg_notify is not None: + await bot.send_message(chat_id=tg_notify, text=text) + return + + server_ids = [s["id"] for s in servers] + groups_map: dict[int, list[str]] = {} + if server_ids: + r = await session.execute( + select(ServerSpecialgroup.server_id, ServerSpecialgroup.group_code).where( + ServerSpecialgroup.server_id.in_(server_ids) + ) + ) + for sid, gc in r.all(): + groups_map.setdefault(sid, []).append(gc) + + for server in servers: + server["special_groups"] = [g for g in groups_map.get(server["id"], []) if g in ALLOWED_GROUP_CODES] + + if subgroup_title: + servers = await filter_cluster_by_subgroup( + session, servers, subgroup_title, least_loaded_cluster, tariff_id=plan + ) + if not servers: + text = "❌ Нет доступных серверов в выбранном кластере." + if safe_to_edit: + await edit_or_send_message(target_message=target_message, text=text, reply_markup=None) + elif tg_notify is not None: + await bot.send_message(chat_id=tg_notify, text=text) + return + + special = None + if tariff: + gc = (tariff.get("group_code") or "").lower() + if gc in ALLOWED_GROUP_CODES: + special = gc + + if special: + bound_servers = [s for s in servers if special in (s.get("special_groups") or [])] + if bound_servers: + servers = bound_servers + + available_servers: list[str] = [] + tasks = [asyncio.create_task(check_server_availability(dict(server), session)) for server in servers] + results = await asyncio.gather(*tasks, return_exceptions=True) + + for server, result_ok in zip(servers, results, strict=False): + if result_ok is True: + available_servers.append(server["server_name"]) + + if not available_servers: + text = "❌ Нет доступных серверов в выбранном кластере." + if safe_to_edit: + await edit_or_send_message(target_message=target_message, text=text, reply_markup=None) + elif tg_notify is not None: + await bot.send_message(chat_id=tg_notify, text=text) + return + + builder = InlineKeyboardBuilder() + ts = int(expiry_time.timestamp()) + + for i in range(0, len(available_servers), 2): + row_buttons = [] + for server_name in available_servers[i : i + 2]: + if old_key_name: + callback_data = f"select_country|{server_name}|{ts}|{old_key_name}" + else: + if plan: + callback_data = f"select_country|{server_name}|{ts}||{plan}" + else: + callback_data = f"select_country|{server_name}|{ts}" + row_buttons.append(InlineKeyboardButton(text=server_name, callback_data=callback_data)) + builder.row(*row_buttons) + + builder.row(InlineKeyboardButton(text=MAIN_MENU, callback_data="profile")) + + if safe_to_edit: + await edit_or_send_message( + target_message=target_message, + text=SELECT_COUNTRY_MSG, + reply_markup=builder.as_markup(), + ) + elif tg_notify is not None: + await bot.send_message( + chat_id=tg_notify, + text=SELECT_COUNTRY_MSG, + reply_markup=builder.as_markup(), + ) + + +@router.callback_query(F.data.startswith("select_country|")) +async def handle_country_selection(callback_query: CallbackQuery, session: Any, state: FSMContext): + data = callback_query.data.split("|") + if len(data) < 3: + await callback_query.message.answer("❌ Некорректные данные. Попробуйте снова.") + return + + selected_country = data[1] + try: + ts = int(data[2]) + except ValueError: + await callback_query.message.answer("❌ Некорректное время истечения. Попробуйте снова.") + return + + old_key_name = data[3] if len(data) > 3 and data[3] else None + try: + tariff_id = int(data[4]) if len(data) > 4 and data[4] else None + except (ValueError, IndexError): + tariff_id = None + + tg_id = callback_query.from_user.id + + fsm_data = await state.get_data() + if fsm_data.get("creating_key"): + try: + await callback_query.answer("⏳ Уже обрабатываю…") + except Exception: + pass + return + + await state.update_data(creating_key=True) + + try: + await callback_query.answer("Обрабатываю…") + if callback_query.message: + await callback_query.message.edit_reply_markup(reply_markup=None) + except Exception: + pass + + try: + expiry_time = datetime.fromtimestamp(ts, tz=moscow_tz) + await finalize_key_creation( + tg_id=tg_id, + expiry_time=expiry_time, + selected_country=selected_country, + state=state, + session=session, + callback_query=callback_query, + old_key_name=old_key_name, + tariff_id=tariff_id, + ) + finally: + fsm_data = await state.get_data() + if fsm_data.get("creating_key"): + await state.update_data(creating_key=False) diff --git a/handlers/keys/key_mode/key_country_mode/finalize.py b/handlers/keys/key_mode/key_country_mode/finalize.py new file mode 100644 index 00000000..53f45b18 --- /dev/null +++ b/handlers/keys/key_mode/key_country_mode/finalize.py @@ -0,0 +1,439 @@ +from ._common import * # noqa: F401,F403 + +async def finalize_key_creation( + tg_id: int, + expiry_time: datetime, + selected_country: str, + state: FSMContext | None, + session: AsyncSession, + callback_query: CallbackQuery, + old_key_name: str | None = None, + tariff_id: int | None = None, +): + from_user = callback_query.from_user + + if not await check_user_exists(session, tg_id): + await add_user( + session=session, + tg_id=from_user.id, + username=from_user.username, + first_name=from_user.first_name, + last_name=from_user.last_name, + language_code=from_user.language_code, + is_bot=from_user.is_bot, + ) + + owner = await resolve_user_optional(session, tg_id) + if owner is None: + await callback_query.message.answer("❌ Пользователь не найден.") + return + uid = owner.id + + expiry_time = expiry_time.astimezone(moscow_tz) + + old_key_details: dict[str, Any] | None = None + if old_key_name: + key_obj = await resolve_key(session, tg_id, old_key_name) + old_key_name = key_obj.email if key_obj else old_key_name + old_key_details = await get_key_details(session, old_key_name) + if not old_key_details: + await callback_query.message.answer("❌ Ключ не найден. Попробуйте снова.") + return + key_name = old_key_name + client_id = old_key_details["client_id"] + email = old_key_details["email"] + expiry_timestamp = old_key_details["expiry_time"] + tariff_id = old_key_details.get("tariff_id") or tariff_id + else: + while True: + key_name = await generate_random_email(session=session) + existing_key = await get_key_details(session, key_name) + if not existing_key: + break + client_id = str(uuid.uuid4()) + email = key_name.lower() + expiry_timestamp = int(expiry_time.timestamp() * 1000) + + data = await state.get_data() if state else {} + is_trial = data.get("is_trial", False) + skip_balance_charge = bool(data.get("skip_balance_charge", False)) + + selected_traffic_gb = data.get("config_selected_traffic_gb") + if selected_traffic_gb is None: + selected_traffic_gb = data.get("selected_traffic_limit_gb") + + selected_device_limit = data.get("config_selected_device_limit") + if selected_device_limit is None: + selected_device_limit = data.get("selected_device_limit") + + if old_key_details: + if selected_traffic_gb is None: + stored_traffic = old_key_details.get("selected_traffic_limit") + if stored_traffic is not None: + selected_traffic_gb = int(stored_traffic) + if selected_device_limit is None: + stored_devices = old_key_details.get("selected_device_limit") + if stored_devices is not None: + selected_device_limit = int(stored_devices) + + price_to_charge = data.get("selected_price_rub") + + effective_tariff_id = data.get("tariff_id") or tariff_id + tariff: dict[str, Any] | None = None + if effective_tariff_id: + tariff_id = int(effective_tariff_id) + tariff = await get_tariff_by_id(session, tariff_id) + + device_limit, traffic_limit_bytes = await get_effective_limits_for_key( + session=session, + tariff_id=tariff_id, + selected_device_limit=selected_device_limit, + selected_traffic_gb=selected_traffic_gb, + ) + + if selected_traffic_gb is not None: + traffic_limit_gb = int(selected_traffic_gb) + else: + traffic_limit_gb = int(traffic_limit_bytes / GB) if traffic_limit_bytes else 0 + + if price_to_charge is None and tariff and not old_key_name: + price_to_charge = tariff.get("price_rub") + + need_vless_key = bool(tariff.get("vless")) if tariff else False + + public_link = None + remnawave_link = None + created_at = int(datetime.now(moscow_tz).timestamp() * 1000) + + try: + result = await session.execute(select(Server).where(Server.server_name == selected_country)) + server_info = result.scalar_one_or_none() + if not server_info: + raise ValueError(f"Сервер {selected_country} не найден") + + cluster_info = await check_server_name_by_cluster(session, server_info.server_name) + if not cluster_info: + raise ValueError(f"Кластер для сервера {server_info.server_name} не найден") + + cluster_name = cluster_info["cluster_name"] + is_full_remnawave = await is_full_remnawave_cluster(cluster_name, session) + + if old_key_name and old_key_details: + old_server_id = old_key_details["server_id"] + if old_server_id: + result = await session.execute(select(Server).where(Server.server_name == old_server_id)) + old_server_info = result.scalar_one_or_none() + if old_server_info: + try: + if old_server_info.panel_type.lower() == "3x-ui": + xui = await get_xui_instance(old_server_info.api_url) + await delete_client(xui, old_server_info.inbound_id, email, client_id) + await session.execute( + update(Key).where(Key.user_id == uid, Key.email == email).values(key=None) + ) + elif old_server_info.panel_type.lower() == "remnawave": + remna_del = RemnawaveAPI(old_server_info.api_url) + if await remna_del.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD): + await remna_del.delete_user(client_id) + await session.execute( + update(Key) + .where(Key.user_id == uid, Key.email == email) + .values(remnawave_link=None) + ) + except Exception as e: + logger.warning(f"[Delete] Ошибка при удалении клиента: {e}") + + panel_type = server_info.panel_type.lower() + + if panel_type == "remnawave" or is_full_remnawave: + remna = RemnawaveAPI(server_info.api_url) + if not await remna.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD): + raise ValueError(f"❌ Не удалось авторизоваться в Remnawave ({server_info.server_name})") + + expire_at = datetime.utcfromtimestamp(expiry_timestamp / 1000).isoformat() + "Z" + user_data: dict[str, Any] = { + "username": email, + "trafficLimitStrategy": "NO_RESET", + "expireAt": expire_at, + "telegramId": tg_id, + "activeInternalSquads": [server_info.inbound_id], + "uuid": client_id, + } + if traffic_limit_bytes: + user_data["trafficLimitBytes"] = traffic_limit_bytes + if device_limit: + user_data["hwidDeviceLimit"] = device_limit + + result = await remna.create_user(user_data) + if not result: + raise ValueError("❌ Ошибка при создании пользователя в Remnawave") + + client_id = result.get("uuid") or result.get("id") or client_id + + remnawave_link = None + if need_vless_key: + try: + vless_link = await get_vless_link_for_remnawave_by_username(remna, email, email) + except Exception: + vless_link = None + if vless_link: + remnawave_link = vless_link + + if not remnawave_link: + try: + sub = await remna.get_subscription_by_username(email) + except Exception: + sub = None + + if sub: + if need_vless_key and not remnawave_link: + links = sub.get("links") or [] + remnawave_link = next( + (l for l in links if isinstance(l, str) and l.lower().startswith("vless://")), + None, + ) + + if not remnawave_link: + remnawave_link = sub.get("subscriptionUrl") + + if old_key_name: + await session.execute( + update(Key).where(Key.user_id == uid, Key.email == email).values(client_id=client_id) + ) + + if panel_type == "3x-ui": + semaphore = asyncio.Semaphore(2) + await create_client_on_server( + server_info={ + "api_url": server_info.api_url, + "inbound_id": server_info.inbound_id, + "server_name": server_info.server_name, + "panel_type": server_info.panel_type, + }, + tg_id=tg_id, + client_id=client_id, + email=email, + expiry_timestamp=expiry_timestamp, + semaphore=semaphore, + session=session, + plan=tariff_id, + is_trial=is_trial, + total_traffic_limit_bytes=traffic_limit_bytes, + device_limit_value=device_limit, + ) + + subgroup_code = tariff.get("subgroup_title") if tariff and tariff.get("subgroup_title") else None + cluster_all = [ + { + "server_name": server_info.server_name, + "api_url": server_info.api_url, + "panel_type": server_info.panel_type, + "inbound_id": getattr(server_info, "inbound_id", None), + "enabled": True, + "max_keys": getattr(server_info, "max_keys", None), + } + ] + + link_to_show = await make_aggregated_link( + session=session, + cluster_all=cluster_all, + cluster_id=cluster_name, + email=email, + client_id=client_id, + tg_id=tg_id, + subgroup_code=subgroup_code, + remna_link_override=remnawave_link, + plan=tariff_id, + ) + + public_link = link_to_show + + if old_key_name: + update_data: dict[str, Any] = { + "server_id": selected_country, + "key": None, + "remnawave_link": None, + } + if public_link and public_link.startswith("vless://"): + update_data["key"] = public_link + elif public_link and public_link.startswith("http"): + update_data["key"] = public_link + if remnawave_link: + update_data["remnawave_link"] = remnawave_link + await session.execute(update(Key).where(Key.user_id == uid, Key.email == email).values(**update_data)) + else: + new_key = Key( + user_id=uid, + client_id=client_id, + email=email, + created_at=created_at, + expiry_time=expiry_timestamp, + key=public_link if public_link else None, + remnawave_link=remnawave_link, + server_id=selected_country, + tariff_id=tariff_id, + selected_device_limit=int(selected_device_limit) if selected_device_limit is not None else None, + selected_traffic_limit=int(selected_traffic_gb) if selected_traffic_gb is not None else None, + selected_price_rub=int(price_to_charge) if price_to_charge is not None else None, + ) + session.add(new_key) + if is_trial: + trial_status = await get_trial(session, tg_id) + if trial_status in [0, -1]: + await update_trial(session, tg_id, 1) + if not is_trial and price_to_charge and not skip_balance_charge: + await update_balance(session, tg_id, -int(price_to_charge)) + + if state: + await state.update_data(skip_balance_charge=False) + + await session.commit() + + except Exception as e: + logger.error(f"[Key Finalize] Ошибка при создании ключа для пользователя {tg_id}: {e}") + await callback_query.message.answer("❌ Произошла ошибка при создании подписки. Попробуйте снова.") + return + + builder = InlineKeyboardBuilder() + is_full_remnawave = await is_full_remnawave_cluster(cluster_name, session) + is_vless = bool(public_link and public_link.lower().startswith("vless://")) or bool(need_vless_key) + final_link = public_link or remnawave_link + webapp_url = ( + final_link + if isinstance(final_link, str) and final_link.strip().lower().startswith(("http://", "https://")) + else None + ) + + use_webapp = bool(MODES_CONFIG.get("REMNAWAVE_WEBAPP_ENABLED", REMNAWAVE_WEBAPP)) + open_in_browser = bool(MODES_CONFIG.get("REMNAWAVE_WEBAPP_OPEN_IN_BROWSER", REMNAWAVE_WEBAPP_OPEN_IN_BROWSER)) + if use_webapp and webapp_url: + use_webapp = await process_remnawave_webapp_override( + remnawave_webapp=use_webapp, + final_link=final_link, + session=session, + ) + + tv_button_enabled = bool(BUTTONS_CONFIG.get("ANDROID_TV_BUTTON_ENABLE")) + + if panel_type == "remnawave" or is_full_remnawave: + if is_vless: + builder.row(InlineKeyboardButton(text=ROUTER_BUTTON, callback_data=build_key_callback("connect_router", client_id, key_name))) + else: + if use_webapp and webapp_url: + if open_in_browser: + builder.row(InlineKeyboardButton(text=CONNECT_DEVICE, url=webapp_url)) + else: + builder.row(InlineKeyboardButton(text=CONNECT_DEVICE, web_app=WebAppInfo(url=webapp_url))) + if tv_button_enabled: + builder.row(InlineKeyboardButton(text=TV_BUTTON, callback_data=build_key_callback("connect_tv", client_id, key_name))) + else: + builder.row( + InlineKeyboardButton( + text=CONNECT_DEVICE, + callback_data=build_key_callback("connect_device", client_id, key_name), + ) + ) + else: + builder.row( + InlineKeyboardButton( + text=CONNECT_DEVICE, + callback_data=build_key_callback("connect_device", client_id, key_name), + ) + ) + + builder.row(InlineKeyboardButton(text=MY_SUB, callback_data=build_key_callback("view_key", client_id, key_name))) + builder.row(InlineKeyboardButton(text=SUPPORT, url=SUPPORT_CHAT_URL)) + builder.row(InlineKeyboardButton(text=MAIN_MENU, callback_data="profile")) + + if await process_intercept_key_creation_message( + chat_id=tg_id, + session=session, + target_message=callback_query, + ): + return + + hook_commands = await process_key_creation_complete( + chat_id=tg_id, + admin=False, + session=session, + email=email, + key_name=key_name, + ) + if hook_commands: + builder = insert_hook_buttons(builder, hook_commands) + + key_record = await get_key_details(session, key_name) + final_link_for_message = final_link or (key_record.get("link") if key_record else None) or "Ссылка не найдена" + message_text = await build_key_created_message( + session=session, + key_record=key_record, + final_link=final_link_for_message, + selected_device_limit=selected_device_limit, + selected_traffic_gb=selected_traffic_gb, + ) + + await edit_or_send_message( + target_message=callback_query.message, + text=message_text, + reply_markup=builder.as_markup(), + media_path="img/pic.jpg", + ) + + if state: + await state.clear() + + +async def check_server_availability(server_info: dict, session: AsyncSession) -> bool: + """Делегирует в services.clusters.check_server_availability().""" + from services.clusters import check_server_availability as _svc_check + result = await _svc_check(server_info, session) + return result.available + + +async def _legacy_check_server_availability(server_info: dict, session: AsyncSession) -> bool: + """Legacy — оставлено для reference.""" + server_name = server_info.get("server_name", "unknown") + panel_type = server_info.get("panel_type", "3x-ui").lower() + enabled = server_info.get("enabled", True) + max_keys = server_info.get("max_keys") + + if not enabled: + logger.info(f"[Ping] Сервер {server_name} выключен (enabled = FALSE).") + return False + + try: + if max_keys is not None: + result = await session.execute(select(func.count()).select_from(Key).where(Key.server_id == server_name)) + key_count = result.scalar() + + if key_count >= max_keys: + logger.info(f"[Ping] Сервер {server_name} достиг лимита ключей: {key_count}/{max_keys}.") + return False + + except SQLAlchemyError as e: + logger.warning(f"[Ping] Ошибка при проверке лимита ключей на сервере {server_name}: {e}") + return False + + try: + if panel_type == "remnawave": + remna = RemnawaveAPI(server_info["api_url"]) + await asyncio.wait_for(remna.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD), timeout=5.0) + logger.info(f"[Ping] Remnawave сервер {server_name} доступен.") + return True + + xui = AsyncApi( + server_info["api_url"], + username=ADMIN_USERNAME, + password=ADMIN_PASSWORD, + logger=logger, + ) + await asyncio.wait_for(xui.login(), timeout=5.0) + logger.info(f"[Ping] 3x-ui сервер {server_name} доступен.") + return True + + except TimeoutError: + logger.warning(f"[Ping] Сервер {server_name} не ответил вовремя.") + return False + except Exception as e: + logger.warning(f"[Ping] Ошибка при проверке сервера {server_name}: {e}") + return False diff --git a/handlers/keys/key_mode/key_discount_mode.py b/handlers/keys/key_mode/key_discount_mode.py index d732ab01..8d3956f6 100644 --- a/handlers/keys/key_mode/key_discount_mode.py +++ b/handlers/keys/key_mode/key_discount_mode.py @@ -11,6 +11,7 @@ from config import DISCOUNT_ACTIVE_HOURS from core.bootstrap import NOTIFICATIONS_CONFIG from database import get_keys, get_tariffs, get_tariffs_for_cluster from database.models import Notification +from database.access.resolution import resolve_user_optional from handlers.buttons import MAIN_MENU, RENEW_KEY_NOTIFICATION from handlers.keys.utils import build_key_callback from handlers.notifications.notify_kb import build_tariffs_keyboard @@ -26,10 +27,14 @@ router = Router() @router.callback_query(F.data == "hot_lead_discount") async def handle_discount_entry(callback: CallbackQuery, session: AsyncSession): tg_id = callback.from_user.id + u = await resolve_user_optional(session, tg_id) + if u is None: + await callback.message.edit_text("❌ Скидка недоступна.") + return result = await session.execute( select(Notification.last_notification_time).where( - Notification.tg_id == tg_id, + Notification.user_id == u.id, Notification.notification_type == "hot_lead_step_2", ) ) @@ -109,10 +114,14 @@ async def handle_discount_tariff_selection(callback: CallbackQuery, session: Asy @router.callback_query(F.data == "hot_lead_final_discount") async def handle_ultra_discount(callback: CallbackQuery, session: AsyncSession): tg_id = callback.from_user.id + u = await resolve_user_optional(session, tg_id) + if u is None: + await callback.message.edit_text("❌ Скидка недоступна.") + return result = await session.execute( select(Notification.last_notification_time).where( - Notification.tg_id == tg_id, + Notification.user_id == u.id, Notification.notification_type == "hot_lead_step_3", ) ) diff --git a/handlers/keys/key_renew.py b/handlers/keys/key_renew.py index 6e4a3b9e..bb868734 100644 --- a/handlers/keys/key_renew.py +++ b/handlers/keys/key_renew.py @@ -28,13 +28,14 @@ from database import ( update_key_expiry, ) from database.models import Key, Server +from database.access.resolution import notify_telegram_chat_id from database.notifications import check_hot_lead_discount from database.tariffs import create_subgroup_hash, find_subgroup_by_hash, get_tariffs from handlers.buttons import BACK, MAIN_MENU, MY_SUB, PAYMENT -from handlers.keys.operations import renew_key_in_cluster -from handlers.payments.currency_rates import format_for_user +from services.operations import renew_key_in_cluster +from services.payments.currency_rates import format_for_user from handlers.payments.fast_payment_flow import try_fast_payment_flow -from handlers.tariffs.tariff_display import GB, get_effective_limits_for_key +from services.tariffs.tariff_display import GB, get_effective_limits_for_key from handlers.texts import ( DISCOUNT_OFFER_MESSAGE, DISCOUNT_OFFER_STEP2, @@ -69,8 +70,6 @@ def normalize_expiry_ms(raw_value: int | float | None) -> int: return 0 value = int(raw_value) if value > 10**13: - value //= 1000 - elif value < 10**10: value *= 1000 return value @@ -548,6 +547,7 @@ async def process_callback_renew_plan(callback_query: CallbackQuery, state: FSMC "cost": cost, "required_amount": required_amount, "new_expiry_time": new_expiry_time, + "selected_duration_days": duration_days, "total_gb": total_gb, "email": email, }, @@ -605,7 +605,7 @@ async def handle_renew_config_confirm(callback_query: CallbackQuery, state: FSMC await callback_query.message.answer("❌ Данные для продления не найдены.") return - from handlers.tariffs.buy.key_tariffs import calculate_config_price + from services.tariffs import calculate_config_price tariff = await get_tariff_by_id(session, int(tariff_id)) if not tariff or not tariff.get("configurable"): @@ -637,6 +637,7 @@ async def handle_renew_config_confirm(callback_query: CallbackQuery, state: FSMC "cost": cost, "required_amount": required_amount, "new_expiry_time": int(new_expiry_time), + "selected_duration_days": int(tariff["duration_days"]), "total_gb": int(selected_traffic_gb or 0), "email": email, "selected_device_limit": selected_devices, @@ -709,13 +710,18 @@ async def complete_key_renewal( selected_traffic_limit: int | None = None, selected_price_rub: int | None = None, ): - """Продлевает подписку, обновляет лимиты и данные в БД.""" + """Продлевает подписку через сервис и отправляет Telegram-уведомление.""" + from services.keys import execute_renewal + from services.errors import ServiceError + try: logger.info(f"[Info] Продление ключа {client_id} по тарифу ID={tariff_id} (Start)") + tg_notify = await notify_telegram_chat_id(session, tg_id) + renewal_hook_chat = tg_notify if tg_notify is not None else tg_id + waiting_message = None wait_text = "⏳ Подождите. Идет продление подписки…" - try: if callback_query: await edit_or_send_message( @@ -723,85 +729,56 @@ async def complete_key_renewal( text=wait_text, reply_markup=None, ) + elif tg_notify is not None: + waiting_message = await bot.send_message(tg_notify, wait_text) else: - waiting_message = await bot.send_message(tg_id, wait_text) + logger.info(f"[Renew] Нет Telegram-чата для экрана ожидания (ref={tg_id}), пропуск") except Exception as e: logger.warning(f"[Renew] Не удалось показать экран ожидания: {e}") - tariff = await get_tariff_by_id(session, tariff_id) - if not tariff: - logger.error(f"[Error] Тариф с id={tariff_id} не найден.") - return - key_info = await get_key_details(session, email) if not key_info: logger.error(f"[Error] Ключ с client_id={client_id} не найден в БД.") return - new_tariff_device_limit = tariff.get("device_limit") - new_tariff_traffic_limit_gb = tariff.get("traffic_limit") + server_or_cluster = key_info["server_id"] - if new_tariff_device_limit is None: - final_device_limit = None - elif selected_device_limit is not None: - final_device_limit = int(selected_device_limit) - else: - final_device_limit = new_tariff_device_limit - - if new_tariff_traffic_limit_gb is None: - final_traffic_limit = None - elif selected_traffic_limit is not None: - final_traffic_limit = int(selected_traffic_limit) - else: - final_traffic_limit = int(new_tariff_traffic_limit_gb) - - if tariff.get("configurable"): - selected_traffic_gb_effective = int(final_traffic_limit) if final_traffic_limit is not None else None - selected_device_limit_effective = int(final_device_limit) if final_device_limit is not None else None - - device_limit_effective, traffic_limit_bytes_effective = await get_effective_limits_for_key( + try: + await release_session_early(session) + result = await execute_renewal( session=session, - tariff_id=int(tariff_id), - selected_device_limit=selected_device_limit_effective, - selected_traffic_gb=selected_traffic_gb_effective, + billing_user_id=tg_id, + client_id=client_id, + key_email=email, + key_server_id=server_or_cluster, + tariff_id=tariff_id, + new_expiry_time=new_expiry_time, + total_gb=total_gb, + cost=cost, + selected_device_limit=selected_device_limit, + selected_traffic_limit=selected_traffic_limit, + selected_price_rub=selected_price_rub, ) + except ServiceError as e: + logger.error(f"[Error] Сервис продления: {e.message}") + return - traffic_limit_gb_effective = int(traffic_limit_bytes_effective / GB) if traffic_limit_bytes_effective else 0 - total_gb = int(traffic_limit_gb_effective) - else: - device_limit_effective = final_device_limit - traffic_limit_gb_effective = int(final_traffic_limit) if final_traffic_limit is not None else 0 - total_gb = int(traffic_limit_gb_effective) + tariff = await get_tariff_by_id(session, tariff_id) + tariff_name = tariff["name"] if tariff else "" + subgroup_title = tariff.get("subgroup_title", "") if tariff else "" - cfg = normalize_tariff_config(tariff) - raw_device_options = cfg.get("device_options") or tariff.get("device_options") or [] - raw_traffic_options = cfg.get("traffic_options_gb") or tariff.get("traffic_options_gb") or [] + device_limit_effective = selected_device_limit + traffic_limit_gb_effective = selected_traffic_limit or 0 - device_int_options: list[int] = [] - for value in raw_device_options: - try: - device_int_options.append(int(value)) - except (TypeError, ValueError): - continue - - traffic_int_options: list[int] = [] - for value in raw_traffic_options: - try: - traffic_int_options.append(int(value)) - except (TypeError, ValueError): - continue - - has_device_choice = len(device_int_options) > 1 - has_traffic_choice = len(traffic_int_options) > 1 - - if selected_price_rub is None: - stored_price = key_info.get("selected_price_rub") - if stored_price is not None: - final_price_rub = int(stored_price) - else: - final_price_rub = int(cost) - else: - final_price_rub = int(selected_price_rub) + if tariff and tariff.get("configurable"): + sel_dev = int(selected_device_limit) if selected_device_limit is not None else None + sel_trf = int(selected_traffic_limit) if selected_traffic_limit is not None else None + dev_eff, trf_bytes = await get_effective_limits_for_key( + session=session, tariff_id=tariff_id, + selected_device_limit=sel_dev, selected_traffic_gb=sel_trf, + ) + device_limit_effective = dev_eff + traffic_limit_gb_effective = int(trf_bytes / GB) if trf_bytes else 0 formatted_expiry_date = datetime.fromtimestamp(new_expiry_time / 1000, tz=moscow_tz).strftime("%d %B %Y, %H:%M") formatted_expiry_date = formatted_expiry_date.replace( @@ -810,90 +787,17 @@ async def complete_key_renewal( ) response_message = get_renewal_message( - tariff_name=tariff["name"], + tariff_name=tariff_name, traffic_limit=traffic_limit_gb_effective, device_limit=device_limit_effective, expiry_date=formatted_expiry_date, - subgroup_title=tariff.get("subgroup_title", ""), + subgroup_title=subgroup_title, ) - current_subgroup = None - try: - current_tariff_id = key_info.get("tariff_id") - if current_tariff_id: - current_tariff = await get_tariff_by_id(session, int(current_tariff_id)) - if current_tariff: - current_subgroup = current_tariff.get("subgroup_title") - except Exception as e: - logger.warning(f"[Renew] Не удалось определить текущую подгруппу: {e}") - - target_subgroup = tariff.get("subgroup_title") - old_subgroup = current_subgroup - - server_or_cluster = key_info["server_id"] - cluster_id = await resolve_cluster_name(session, server_or_cluster) - if not cluster_id: - logger.error(f"[Error] Кластер для {server_or_cluster} не найден.") - return - - await release_session_early(session) - await renew_key_in_cluster( - cluster_id=cluster_id, - email=email, - client_id=client_id, - new_expiry_time=new_expiry_time, - total_gb=total_gb, - session=session, - hwid_device_limit=device_limit_effective, - reset_traffic=True, - target_subgroup=target_subgroup, - old_subgroup=old_subgroup, - plan=tariff_id, - ) - - key_row = await get_key_details(session, email) - effective_client_id = key_row["client_id"] if key_row else client_id - - await update_key_expiry(session, effective_client_id, new_expiry_time) - - update_values = {"tariff_id": tariff_id} - - if not tariff.get("configurable"): - if new_tariff_device_limit is None: - update_values["selected_device_limit"] = None - update_values["current_device_limit"] = None - else: - update_values["selected_device_limit"] = new_tariff_device_limit - update_values["current_device_limit"] = final_device_limit - - if new_tariff_traffic_limit_gb is None: - update_values["selected_traffic_limit"] = None - update_values["current_traffic_limit"] = None - else: - update_values["selected_traffic_limit"] = new_tariff_traffic_limit_gb - update_values["current_traffic_limit"] = final_traffic_limit - - await session.execute(update(Key).where(Key.email == email).values(**update_values)) - await update_balance(session, tg_id, -cost) - - if tariff.get("configurable"): - await save_key_config_with_mode( - session=session, - email=email, - selected_devices=final_device_limit, - selected_traffic_gb=final_traffic_limit, - total_price=int(final_price_rub), - has_device_choice=has_device_choice, - has_traffic_choice=has_traffic_choice, - config_mode="renewal", - ) - if has_device_choice or has_traffic_choice: - await reset_key_current_limits_to_selected(session, effective_client_id) - builder = InlineKeyboardBuilder() builder.row(InlineKeyboardButton(text=MY_SUB, callback_data=build_key_callback("view_key", client_id, email))) hook_commands = await process_renewal_complete( - chat_id=tg_id, admin=False, session=session, email=email, client_id=client_id + chat_id=renewal_hook_chat, admin=False, session=session, email=email, client_id=client_id ) if hook_commands: builder = insert_hook_buttons(builder, hook_commands) @@ -915,11 +819,14 @@ async def complete_key_renewal( reply_markup=builder.as_markup(), media_path=renewal_media_path, ) + elif tg_notify is not None: + await bot.send_message(tg_notify, response_message, reply_markup=builder.as_markup()) else: - await bot.send_message(tg_id, response_message, reply_markup=builder.as_markup()) + logger.info(f"[Renew] Нет Telegram-чата для итогового сообщения (ref={tg_id}), пропуск") except Exception as e: logger.error(f"[Error] Ошибка при выводе финального сообщения: {e}") - await bot.send_message(tg_id, response_message, reply_markup=builder.as_markup()) + if tg_notify is not None: + await bot.send_message(tg_notify, response_message, reply_markup=builder.as_markup()) logger.info(f"[Info] Продление ключа {client_id} завершено успешно (User: {tg_id})") diff --git a/handlers/keys/key_view.py b/handlers/keys/key_view.py index a58d4b7f..59507102 100644 --- a/handlers/keys/key_view.py +++ b/handlers/keys/key_view.py @@ -33,6 +33,7 @@ from panels.remnawave_runtime import ( ) from database import get_key_details, get_keys from database.models import Key +from database.access.resolution import resolve_user_optional from handlers.buttons import ( ADDONS_BUTTON_DEVICES, ADDONS_BUTTON_DEVICES_TRAFFIC, @@ -50,7 +51,7 @@ from handlers.buttons import ( TV_BUTTON, ) from database import get_vless_enabled_batch -from handlers.tariffs.tariff_display import GB, get_key_tariff_addons_state +from services.tariffs.tariff_display import GB, get_key_tariff_addons_state from handlers.texts import ( DAYS_LEFT_MESSAGE, FROZEN_SUBSCRIPTION_MSG, @@ -274,8 +275,13 @@ async def handle_new_alias_input(message: Message, state: FSMContext, session: A client_id = data.get("client_id") try: + u = await resolve_user_optional(session, message.chat.id) + if u is None: + await message.answer("❌ Не удалось переименовать подписку.") + await state.clear() + return await session.execute( - update(Key).where(Key.tg_id == message.chat.id, Key.client_id == client_id).values(alias=alias) + update(Key).where(Key.user_id == u.id, Key.client_id == client_id).values(alias=alias) ) await session.commit() except Exception as error: diff --git a/handlers/keys/keys.py b/handlers/keys/keys.py index 9a278ee2..06fe4c5a 100644 --- a/handlers/keys/keys.py +++ b/handlers/keys/keys.py @@ -6,7 +6,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from database import delete_key, get_key_details from handlers.buttons import APPLY, BACK, CANCEL from handlers.keys.key_view import process_callback_view_key -from handlers.keys.operations import delete_key_from_cluster, update_subscription +from services.operations import delete_key_from_cluster, update_subscription from handlers.keys.utils import build_key_callback, key_owned_by_user, resolve_key from handlers.texts import DELETE_KEY_CONFIRM_MSG, KEY_DELETED_MSG_SIMPLE from handlers.utils import edit_or_send_message, handle_error diff --git a/handlers/keys/utils.py b/handlers/keys/utils.py index f97d6e18..d1a4a1f0 100644 --- a/handlers/keys/utils.py +++ b/handlers/keys/utils.py @@ -8,7 +8,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from database import get_key_by_client_id, get_key_by_email, get_keys from database.models import Key -from handlers.payments.currency_rates import format_for_user +from services.payments.currency_rates import format_for_user def key_owned_by_user(record: dict | None, user_id: int) -> bool: diff --git a/handlers/notifications/general_notifications.py b/handlers/notifications/general_notifications.py index 2dffe615..b598e17a 100644 --- a/handlers/notifications/general_notifications.py +++ b/handlers/notifications/general_notifications.py @@ -47,13 +47,13 @@ from database.tariffs import ( get_tariff_by_id, get_tariffs_for_cluster, ) -from handlers.keys.operations import delete_key_from_cluster, renew_key_in_cluster +from services.operations import delete_key_from_cluster, renew_key_in_cluster from handlers.notifications.notify_kb import ( build_change_tariff_kb, build_notification_expired_kb, build_notification_kb, ) -from handlers.tariffs.tariff_display import GB, get_effective_limits_for_key, resolve_price_to_charge +from services.tariffs.tariff_display import GB, get_effective_limits_for_key, resolve_price_to_charge from handlers.texts import ( KEY_CANNOT_RENEW_CURRENT, KEY_DELETED_MSG, @@ -109,7 +109,7 @@ async def preload_notification_data(session: AsyncSession) -> dict[str, Any]: User.balance.label("user_balance"), ) .outerjoin(Tariff, Key.tariff_id == Tariff.id) - .outerjoin(User, Key.tg_id == User.tg_id) + .outerjoin(User, Key.user_id == User.id) .where(Key.is_frozen.is_(False)) ) @@ -182,11 +182,11 @@ async def execute_bulk_updates(session: AsyncSession, bulk_updates: dict[str, An to_add = bulk_updates.get("notifications_to_add") or [] if to_add: - await bulk_add_notifications(session, to_add, commit=False) + await bulk_add_notifications(session, to_add) to_delete = bulk_updates.get("notifications_to_delete") or [] if to_delete: - await bulk_delete_notifications(session, to_delete, commit=False) + await bulk_delete_notifications(session, to_delete) if to_add or to_delete: logger.info( @@ -507,9 +507,14 @@ async def notify_expiring_keys( "notification_id": notification_id, "email": email, }) + try: + from database.web_notifications import notify_web + await notify_web(ctx.session, tg_id=tg_id, type="key_expiry", template_vars={"email": email}, data={"email": email}) + except Exception as e: + logger.warning("[Notifications] Ошибка web-уведомления key_expiry tg_id={}: {}", tg_id, e) - - renew_results: list[tuple[Any, str, bool, Optional[dict], Optional[int]]] = [] + + renew_results: list[tuple[Any, str, bool, Optional[dict], Optional[int]]] = [] use_parallel = ( notify_renew_enabled and renew_candidates diff --git a/handlers/notifications/notify_utils.py b/handlers/notifications/notify_utils.py index dc6e2323..d42ac438 100644 --- a/handlers/notifications/notify_utils.py +++ b/handlers/notifications/notify_utils.py @@ -18,7 +18,7 @@ from aiogram.types import BufferedInputFile, InlineKeyboardMarkup from sqlalchemy.ext.asyncio import AsyncSession from database import async_session_maker, create_blocked_user -from handlers.tariffs.tariff_display import get_key_tariff_display +from services.tariffs.tariff_display import get_key_tariff_display from handlers.utils import format_hours, format_minutes, get_russian_month from logger import logger @@ -150,15 +150,12 @@ class FastNotificationSender: if not self.blocked_users: return try: - from sqlalchemy.dialects.postgresql import insert - from database.models import BlockedUser + from database.bans import save_blocked_user_ids - values = [{"tg_id": tg_id} for tg_id in self.blocked_users] - stmt = insert(BlockedUser).values(values).on_conflict_do_nothing(index_elements=[BlockedUser.tg_id]) async with async_session_maker() as session: - await session.execute(stmt) + await save_blocked_user_ids(session, list(self.blocked_users)) await session.commit() - logger.info(f"📝 Добавлено {len(self.blocked_users)} пользователей в blocked_users") + logger.info(f"📝 Добавлено до {len(self.blocked_users)} пользователей в blocked_users") except Exception as e: logger.error(f"❌ Ошибка при сохранении заблокированных пользователей: {e}") diff --git a/handlers/notifications/special_notifications.py b/handlers/notifications/special_notifications.py index 79fcd7af..b9a31d65 100644 --- a/handlers/notifications/special_notifications.py +++ b/handlers/notifications/special_notifications.py @@ -22,7 +22,7 @@ from database.models import Key, User from database.tariffs import get_tariffs from handlers.buttons import CONNECT_DEVICE, MAIN_MENU, SUPPORT, TRIAL_BONUS from handlers.keys.utils import build_key_callback -from handlers.keys.operations import get_user_traffic +from services.operations import get_user_traffic from handlers.notifications.notify_utils import send_messages_with_limit from handlers.texts import ( TRIAL_INACTIVE_BONUS_MSG, @@ -126,6 +126,7 @@ async def notify_inactive_trial_users( async with sessionmaker() as fresh_session: for tg_id in sent_tg_ids: await add_notification(fresh_session, tg_id, "inactive_trial") + await fresh_session.commit() else: for tg_id in sent_tg_ids: await add_notification(session, tg_id, "inactive_trial") diff --git a/handlers/payments/__init__.py b/handlers/payments/__init__.py index f4b666fa..5f2cd6c3 100644 --- a/handlers/payments/__init__.py +++ b/handlers/payments/__init__.py @@ -9,14 +9,14 @@ __all__ = ( from aiogram import Router from config import PROVIDERS_ENABLED -from handlers.payments.providers import get_providers - -from .payment_links import ( +from services.payments.payment_links import ( PaymentLinkRequest, PaymentLinkResult, create_payment_link, register_payment_creator, ) +from services.payments.providers import get_providers + from .cryptobot import router as cryptobot_router from .fast_payment_flow import router as fast_payment_flow_router from .freekassa.freekassa_pay import router as freekassa_router diff --git a/handlers/payments/currency_flow.py b/handlers/payments/currency_flow.py index 9f3f9ff3..a87a2fc5 100644 --- a/handlers/payments/currency_flow.py +++ b/handlers/payments/currency_flow.py @@ -1,11 +1,13 @@ -from typing import Iterable, List, Any +from collections.abc import Iterable +from typing import Any, List + from aiogram.types import InlineKeyboardButton from aiogram.utils.keyboard import InlineKeyboardBuilder -from handlers.texts import FAST_PAY_NOT_ENOUGH -from handlers.buttons import RUB_CURRENCY, USD_CURRENCY, STARS, MAIN_MENU from config import TRIBUTE_LINK -from .currency_rates import format_for_user +from handlers.buttons import MAIN_MENU, RUB_CURRENCY, STARS, USD_CURRENCY +from handlers.texts import FAST_PAY_NOT_ENOUGH +from services.payments.currency_rates import format_for_user def build_currency_choice_kb( @@ -41,7 +43,7 @@ async def shortfall_lead_text( *, force_currency: str | None = None, ) -> str: - if not isinstance(required_amount, (int, float)) or required_amount <= 0: + if not isinstance(required_amount, int | float) or required_amount <= 0: return "💳" amount_txt = await format_for_user( session, tg_id, float(required_amount), language_code, force_currency=force_currency @@ -53,9 +55,9 @@ def filter_providers_by_currency( currency: str, providers: Iterable[str], rub_providers: Iterable[str], -) -> List[str]: +) -> list[str]: rub_set = {p.upper() for p in rub_providers} - out: List[str] = [] + out: list[str] = [] for p in providers: up = p.upper() if currency == "RUB": diff --git a/handlers/payments/fast_payment_flow.py b/handlers/payments/fast_payment_flow.py index f7b3b869..6e256a4d 100644 --- a/handlers/payments/fast_payment_flow.py +++ b/handlers/payments/fast_payment_flow.py @@ -1,5 +1,5 @@ -from typing import Any from math import ceil +from typing import Any from aiogram import F, Router from aiogram.fsm.context import FSMContext @@ -8,7 +8,7 @@ from aiogram.types import CallbackQuery, InlineKeyboardButton, Message from aiogram.utils.keyboard import InlineKeyboardBuilder from sqlalchemy import select -from config import USE_NEW_PAYMENT_FLOW, TRIBUTE_LINK +from config import TRIBUTE_LINK, USE_NEW_PAYMENT_FLOW from core.bootstrap import PAYMENTS_CONFIG from core.settings.buttons_config import BUTTONS_CONFIG from core.settings.money_config import get_currency_mode @@ -17,21 +17,21 @@ from database import ( create_coupon_usage, get_balance, get_coupon_by_code, + has_any_coupon_usage, update_coupon_usage_count, ) from database.coupons import apply_percent_coupon -from database.models import CouponUsage from database.temporary_data import create_temporary_data from handlers import buttons as btn from handlers.payments.currency_flow import ( build_currency_choice_kb, - shortfall_lead_text, currency_label, + shortfall_lead_text, ) -from handlers.payments.providers import get_providers_with_hooks, sort_provider_names -from handlers.texts import FAST_PAY_CHOOSE_CURRENCY, FAST_PAY_CHOOSE_PROVIDER, FASTFLOW_COUPON_APPLIED_TEMPLATE +from handlers.texts import FASTFLOW_COUPON_APPLIED_TEMPLATE, FAST_PAY_CHOOSE_CURRENCY, FAST_PAY_CHOOSE_PROVIDER from handlers.utils import edit_or_send_message from logger import logger +from services.payments.providers import get_providers_with_hooks, sort_provider_names router = Router() @@ -392,10 +392,7 @@ async def fastflow_apply_coupon(message: Message, state: FSMContext, session: An return if bool(getattr(coupon, "new_users_only", False)): - used_any = await session.scalar( - select(CouponUsage.id).where(CouponUsage.user_id == message.from_user.id).limit(1) - ) - if used_any: + if await has_any_coupon_usage(session, message.from_user.id): await message.answer(new_users_only_text, reply_markup=back_markup) return diff --git a/handlers/payments/freekassa/__init__.py b/handlers/payments/freekassa/__init__.py index e69de29b..8b137891 100644 --- a/handlers/payments/freekassa/__init__.py +++ b/handlers/payments/freekassa/__init__.py @@ -0,0 +1 @@ + diff --git a/handlers/payments/freekassa/freekassa_pay.py b/handlers/payments/freekassa/freekassa_pay.py index 11ce96ce..0b15ce6c 100644 --- a/handlers/payments/freekassa/freekassa_pay.py +++ b/handlers/payments/freekassa/freekassa_pay.py @@ -31,14 +31,15 @@ from database import ( get_payment_by_payment_id, get_temporary_data, invalidate_payment_cache, + register_pending_payment, update_balance, ) from handlers.buttons import BACK, PAY_2 -from handlers.payments.payment_links import register_payment_creator from handlers.payments.utils import send_payment_success_notification from handlers.texts import DEFAULT_PAYMENT_MESSAGE, ENTER_SUM, PAYMENT_OPTIONS from handlers.utils import edit_or_send_message from logger import logger +from services.payments.payment_links import register_payment_creator router = Router() @@ -203,6 +204,16 @@ def verify_signature(params: dict) -> bool: async def freekassa_webhook(request: web.Request): + """Freekassa webhook через общий pipeline. + + Провайдер-специфичная часть: MD5 подпись + проверка merchant_id + + извлечение tg_id из ``us_tg_id`` (custom param) либо из формата + ``order__`` fallback. После pipeline.process_success_payment + отдельно очищаем ``temporary_data`` (FSM state), т.к. freekassa используется + из Telegram-bot flow в отличие от остальных web-провайдеров. + """ + from services.payments.pipeline import ParsedPayment, process_success_payment + try: ip = get_webhook_client_ip(request) if await is_webhook_ip_blocked(ip): @@ -244,17 +255,23 @@ async def freekassa_webhook(request: web.Request): logger.error(f"Error parsing parameters: {e}") return web.Response(status=400, text="Invalid parameter format") - async with async_session_maker() as session: - existing = await get_payment_by_payment_id(session, merchant_order_id) - if existing and existing.get("status") == "success": - logger.warning(f"[Freekassa] Повторный webhook. Платёж уже обработан: order_id={merchant_order_id}") - return web.Response(text="YES") + parsed = ParsedPayment( + payment_id=merchant_order_id, + tg_id=tg_id_int, + amount=amount_float, + currency="RUB", + ) + result = await process_success_payment("freekassa", parsed) + if not result.ok: + return web.Response(status=500, text="Internal server error") - await update_balance(session, tg_id_int, amount_float) - await send_payment_success_notification(tg_id_int, amount_float, session) - await add_payment(session, tg_id_int, amount_float, "freekassa", payment_id=merchant_order_id) - await clear_temporary_data(session, tg_id_int) - await invalidate_payment_cache(merchant_order_id) + + try: + async with async_session_maker() as session: + await clear_temporary_data(session, tg_id_int) + await session.commit() + except Exception as e: + logger.warning(f"[Freekassa] Не удалось очистить temporary_data: {e}") logger.info(f"Payment processed successfully. User: {tg_id_int}, Amount: {amount_float}") return web.Response(text="YES") @@ -367,11 +384,20 @@ async def create_link( currency: str, success_url: str | None, failure_url: str | None, + metadata: dict | None, ) -> tuple[str, str]: if currency not in ("RUB", "USD"): raise ValueError("Freekassa поддерживает только RUB или USD") order_id = f"order_{tg_id}_{int(amount)}_{hash(str(tg_id) + str(amount))}" url = generate_payment_link(amount, order_id, tg_id, currency) + await register_pending_payment( + payment_id=order_id, + tg_id=tg_id, + amount=float(amount), + payment_system="freekassa", + currency=currency, + metadata=metadata, + ) return (url, order_id) diff --git a/handlers/payments/heleket/handlers.py b/handlers/payments/heleket/handlers.py index 8676f29e..51597a08 100644 --- a/handlers/payments/heleket/handlers.py +++ b/handlers/payments/heleket/handlers.py @@ -7,12 +7,12 @@ from sqlalchemy.ext.asyncio import AsyncSession from database import get_temporary_data from database.models import User from handlers.buttons import MAIN_MENU, PAY_2 -from handlers.payments.currency_rates import format_for_user from handlers.texts import DEFAULT_PAYMENT_MESSAGE from handlers.utils import edit_or_send_message from logger import logger -from ..constants import ALLOWED_TEMP_PAYMENT_STATES +from services.payments.currency_rates import format_for_user +from ..constants import ALLOWED_TEMP_PAYMENT_STATES from .service import ( HELEKET_METHODS, generate_heleket_payment_link, @@ -20,6 +20,7 @@ from .service import ( router as service_router, ) + router = Router(name="heleket_router") router.include_router(service_router) diff --git a/handlers/payments/heleket/service.py b/handlers/payments/heleket/service.py index 4da2c8c5..61c58e05 100644 --- a/handlers/payments/heleket/service.py +++ b/handlers/payments/heleket/service.py @@ -2,9 +2,11 @@ import base64 import hashlib import json import time + from decimal import ROUND_HALF_UP, Decimal import aiohttp + from aiogram import F, Router, types from aiogram.fsm.context import FSMContext from aiogram.fsm.state import State, StatesGroup @@ -24,20 +26,12 @@ from config import ( from database import async_session_maker, register_pending_payment from database.models import User from handlers.buttons import BACK, HELEKET, PAY_2 -from handlers.payments.currency_rates import ( - format_for_user, - get_rub_rate, - pick_currency, - to_rub, -) -from handlers.payments.payment_links import register_payment_creator from handlers.payments.keyboards import ( build_amounts_keyboard, parse_amount_from_callback, pay_keyboard, payment_options_for_user, ) -from handlers.payments.providers import get_providers from handlers.texts import ( ENTER_SUM, HELEKET_CRYPTO_DESCRIPTION, @@ -45,6 +39,14 @@ from handlers.texts import ( ) from handlers.utils import edit_or_send_message from logger import logger +from services.payments.currency_rates import ( + format_for_user, + get_rub_rate, + pick_currency, + to_rub, +) +from services.payments.payment_links import register_payment_creator +from services.payments.providers import get_providers router = Router() @@ -326,7 +328,15 @@ async def process_amount_selection(callback_query: types.CallbackQuery, state: F async def generate_heleket_payment_link( - amount: int, tg_id: int, method: dict, session: AsyncSession | None = None + amount: int, + tg_id: int, + method: dict, + session: AsyncSession | None = None, + *, + order_id: str | None = None, + success_url: str | None = None, + failure_url: str | None = None, + metadata: dict | None = None, ) -> str: """ Создание платежа в Heleket и получение ссылки на оплату. @@ -334,8 +344,7 @@ async def generate_heleket_payment_link( session — сессия из хендлера; если не передана, создаётся своя (лишняя нагрузка на пул). """ url = "https://api.heleket.com/v1/payment" - unique_order_id = f"{int(time.time())}_{tg_id}" - db_session = session + unique_order_id = order_id or f"{int(time.time())}_{tg_id}" timeout = aiohttp.ClientTimeout(total=30, connect=10) try: @@ -352,8 +361,8 @@ async def generate_heleket_payment_link( "amount": str(payment_amount), "currency": method["currency"], "order_id": unique_order_id, - "url_success": HELEKET_SUCCESS_URL, - "url_return": HELEKET_RETURN_URL, + "url_success": success_url or HELEKET_SUCCESS_URL, + "url_return": failure_url or HELEKET_RETURN_URL, "url_callback": HELEKET_CALLBACK_URL, "additional_data": f"tg_id:{tg_id},rub_amount:{amount}", } @@ -384,6 +393,7 @@ async def generate_heleket_payment_link( amount=float(amount), payment_system="heleket", currency="RUB", + metadata=metadata, ) logger.info(f"Heleket payment URL created for user {tg_id}") return payment_url @@ -418,17 +428,28 @@ async def create_link( currency: str, success_url: str | None, failure_url: str | None, + metadata: dict | None, ) -> tuple[str, str | None]: method = HELEKET_METHODS.get("crypto") if not method or not method.get("enable"): raise ValueError("Heleket недоступен") amount_int = int(amount) + order_id = f"{int(time.time())}_{tg_id}" if amount_int < 10: raise ValueError("Минимальная сумма для Heleket — 10₽") - url = await generate_heleket_payment_link(amount_int, tg_id, method, session) + url = await generate_heleket_payment_link( + amount_int, + tg_id, + method, + session, + order_id=order_id, + success_url=success_url, + failure_url=failure_url, + metadata=metadata, + ) if not url or url == "https://heleket.com/": raise ValueError("Не удалось создать платёж Heleket") - return (url, None) + return (url, order_id) register_payment_creator("HELEKET", create_link) diff --git a/handlers/payments/heleket/webhook.py b/handlers/payments/heleket/webhook.py deleted file mode 100644 index dc2657b4..00000000 --- a/handlers/payments/heleket/webhook.py +++ /dev/null @@ -1,196 +0,0 @@ -import base64 -import hashlib -import json - -from aiohttp import web - -from config import HELEKET_API_KEY -from core.webhook_abuse import ( - get_webhook_client_ip, - is_webhook_ip_blocked, - record_webhook_signature_failure, -) -from database import ( - add_payment, - async_session_maker, - get_payment_by_payment_id, - invalidate_payment_cache, - update_balance, - update_payment_status, -) -from handlers.payments.utils import send_payment_success_notification -from logger import logger - - -def verify_heleket_signature(data: dict) -> bool: - """Проверяет подпись webhook от Heleket. - - Args: - data: Данные от webhook - - Returns: - True если подпись валидна, иначе False - """ - try: - received_signature = data.get("sign") - if not received_signature: - logger.error("Heleket webhook: отсутствует подпись") - return False - - data_without_sign = data.copy() - del data_without_sign["sign"] - - json_data = json.dumps(data_without_sign, ensure_ascii=False, separators=(",", ":")) - json_data = json_data.replace("/", "\\/") - base64_data = base64.b64encode(json_data.encode("utf-8")).decode("utf-8") - sign_string = base64_data + HELEKET_API_KEY - calculated_signature = hashlib.md5(sign_string.encode("utf-8")).hexdigest() - is_valid = calculated_signature.lower() == received_signature.lower() - - if not is_valid: - logger.error( - f"Heleket webhook: неверная подпись. Ожидалось: {calculated_signature}, получено: {received_signature}" - ) - logger.error(f"Heleket webhook: строка для подписи: {sign_string}") - else: - logger.info("Heleket webhook: подпись успешно проверена") - return is_valid - except Exception as e: - logger.error(f"Ошибка проверки подписи Heleket webhook: {e}") - return False - - -async def process_heleket_webhook(data: dict) -> bool: - """Обрабатывает webhook от Heleket. - - Args: - data: Данные от webhook - - Returns: - True если обработка успешна, иначе False - """ - try: - logger.info(f"Processing Heleket webhook: {data}") - - webhook_type = data.get("type") - order_id = data.get("order_id") - status = data.get("status") - payment_amount = data.get("payment_amount") - merchant_amount = data.get("merchant_amount") - payer_currency = data.get("payer_currency") - additional_data = data.get("additional_data") - - logger.info(f"Heleket webhook - Type: {webhook_type}, Order: {order_id}, Status: {status}") - if webhook_type != "payment": - logger.warning(f"Heleket webhook: неизвестный тип {webhook_type}") - return False - if status in ["paid", "paid_over"]: - logger.info(f"Heleket: успешный платёж {order_id} на сумму {payment_amount} {payer_currency}") - tg_id = None - rub_amount = None - if additional_data: - try: - for part in additional_data.split(","): - if part.startswith("tg_id:"): - tg_id = int(part.split(":")[1]) - elif part.startswith("rub_amount:"): - rub_amount = float(part.split(":")[1]) - except Exception as e: - logger.error(f"Ошибка парсинга additional_data: {e}") - if not tg_id and "_" in order_id: - try: - tg_id = int(order_id.split("_")[1]) - except Exception as e: - logger.error(f"Ошибка извлечения tg_id из order_id: {e}") - if not tg_id: - logger.error(f"Не удалось извлечь tg_id из Heleket webhook: {data}") - return False - balance_amount = rub_amount if rub_amount else float(merchant_amount) - async with async_session_maker() as session: - payment = await get_payment_by_payment_id(session, order_id) - if payment: - if payment.get("status") == "success": - logger.info(f"Heleket: платёж {order_id} уже обработан") - return True - if payment.get("id") is not None: - ok = await update_payment_status( - session=session, internal_id=int(payment["id"]), new_status="success" - ) - if not ok: - logger.error(f"Heleket: не удалось обновить статус платежа {order_id}") - return False - else: - await add_payment( - session=session, - tg_id=tg_id, - amount=balance_amount, - payment_system="HELEKET", - status="success", - currency="USD", - payment_id=order_id, - metadata=None, - ) - else: - await add_payment( - session=session, - tg_id=tg_id, - amount=balance_amount, - payment_system="heleket", - status="success", - currency="USD", - payment_id=order_id, - metadata=None, - ) - - await update_balance(session, tg_id, balance_amount) - await send_payment_success_notification(tg_id, balance_amount, session) - await invalidate_payment_cache(order_id) - logger.info( - f"Heleket: платёж {order_id} для пользователя {tg_id} " - f"успешно обработан, баланс пополнен на {balance_amount} RUB" - ) - return True - elif status in ["fail", "wrong_amount", "cancel", "system_fail"]: - logger.warning(f"Heleket: неудачный платёж {order_id}, статус: {status}") - - async with async_session_maker() as session: - payment = await get_payment_by_payment_id(session, order_id) - if payment and payment.get("id") is not None: - await update_payment_status( - session=session, - internal_id=int(payment["id"]), - new_status="failed", - ) - await session.commit() - await invalidate_payment_cache(order_id) - return True - else: - logger.info(f"Heleket: промежуточный статус {status} для платежа {order_id}") - return True - except Exception as e: - logger.error(f"Ошибка обработки Heleket webhook: {e}") - return False - - -async def heleket_webhook(request: web.Request): - """Обработчик webhook от Heleket для aiohttp.""" - try: - ip = get_webhook_client_ip(request) - if await is_webhook_ip_blocked(ip): - return web.Response(status=429) - data = await request.json() - logger.info(f"Heleket webhook received from {request.remote}") - - if not verify_heleket_signature(data): - logger.error("Heleket webhook: неверная подпись") - await record_webhook_signature_failure(ip) - return web.Response(status=400, text="Invalid signature") - - success = await process_heleket_webhook(data) - if success: - return web.Response(status=200, text="OK") - else: - return web.Response(status=400, text="Processing failed") - except Exception as e: - logger.error(f"Ошибка обработки Heleket webhook: {e}") - return web.Response(status=500, text="Internal server error") diff --git a/handlers/payments/kassai/handlers.py b/handlers/payments/kassai/handlers.py index 704f440f..72567e90 100644 --- a/handlers/payments/kassai/handlers.py +++ b/handlers/payments/kassai/handlers.py @@ -7,12 +7,12 @@ from sqlalchemy.ext.asyncio import AsyncSession from database import get_temporary_data from database.models import User from handlers.buttons import MAIN_MENU, PAY_2 -from handlers.payments.currency_rates import format_for_user from handlers.texts import DEFAULT_PAYMENT_MESSAGE from handlers.utils import edit_or_send_message from logger import logger -from ..constants import ALLOWED_TEMP_PAYMENT_STATES +from services.payments.currency_rates import format_for_user +from ..constants import ALLOWED_TEMP_PAYMENT_STATES from .service import ( KASSAI_METHODS, generate_kassai_payment_link, @@ -20,6 +20,7 @@ from .service import ( router as service_router, ) + router = Router(name="kassai_router") router.include_router(service_router) diff --git a/handlers/payments/kassai/service.py b/handlers/payments/kassai/service.py index 6d2cb0cb..cc5c9967 100644 --- a/handlers/payments/kassai/service.py +++ b/handlers/payments/kassai/service.py @@ -3,6 +3,7 @@ import hmac import time import aiohttp + from aiogram import F, Router, types from aiogram.fsm.context import FSMContext from aiogram.fsm.state import State, StatesGroup @@ -23,19 +24,12 @@ from config import ( from database import async_session_maker, register_pending_payment from database.models import User from handlers.buttons import BACK, KASSAI_CARDS, KASSAI_SBP, PAY_2 -from handlers.payments.currency_rates import ( - format_for_user, - pick_currency, - to_rub, -) from handlers.payments.keyboards import ( build_amounts_keyboard, parse_amount_from_callback, pay_keyboard, payment_options_for_user, ) -from handlers.payments.payment_links import register_payment_creator -from handlers.payments.providers import get_providers from handlers.texts import ( ENTER_SUM, KASSAI_CARDS_DESCRIPTION, @@ -44,6 +38,14 @@ from handlers.texts import ( ) from handlers.utils import edit_or_send_message from logger import logger +from services.payments.currency_rates import ( + format_for_user, + pick_currency, + to_rub, +) +from services.payments.payment_links import register_payment_creator +from services.payments.providers import get_providers + router = Router() @@ -356,14 +358,22 @@ async def process_amount_selection(callback_query: types.CallbackQuery, state: F async def generate_kassai_payment_link( - amount: int, tg_id: int, method: dict, session: AsyncSession | None = None + amount: int, + tg_id: int, + method: dict, + session: AsyncSession | None = None, + *, + payment_id: str | None = None, + success_url: str | None = None, + failure_url: str | None = None, + metadata: dict | None = None, ) -> str: """ Создание заказа в KassaAI и получение ссылки на оплату. session — сессия из хендлера; если не передана, создаётся своя (лишняя нагрузка на пул). """ nonce = int(time.time()) - unique_payment_id = f"{nonce}_{tg_id}" + unique_payment_id = payment_id or f"{nonce}_{tg_id}" url = "https://api.fk.life/v1/orders/create" headers = {"Content-Type": "application/json"} @@ -379,8 +389,8 @@ async def generate_kassai_payment_link( "ip": client_ip, "amount": int(amount), "currency": "RUB", - "success_url": KASSAI_SUCCESS_URL, - "failure_url": KASSAI_FAILURE_URL, + "success_url": success_url or KASSAI_SUCCESS_URL, + "failure_url": failure_url or KASSAI_FAILURE_URL, "paymentId": unique_payment_id, } @@ -389,8 +399,6 @@ async def generate_kassai_payment_link( data = {**data_for_signature, "signature": signature} - db_session = session - timeout = aiohttp.ClientTimeout(total=60, connect=10) try: async with aiohttp.ClientSession(timeout=timeout) as http_session: @@ -407,6 +415,7 @@ async def generate_kassai_payment_link( amount=float(amount), payment_system="kassai", currency="RUB", + metadata=metadata, ) logger.info(f"KassaAI payment URL created for user {tg_id}") return payment_url @@ -440,6 +449,7 @@ def create_link_factory(method_name: str): currency: str, success_url: str | None, failure_url: str | None, + metadata: dict | None, ) -> tuple[str, str | None]: if currency != "RUB": raise ValueError("KassaI поддерживает только RUB") @@ -447,14 +457,24 @@ def create_link_factory(method_name: str): if not method or not method.get("enable"): raise ValueError("Способ оплаты KassaI недоступен") amount_int = int(amount) + payment_id = f"{int(time.time())}_{tg_id}" if method_name == "cards" and amount_int < 50: raise ValueError("Минимальная сумма для карт — 50₽") if method_name == "sbp" and amount_int < 10: raise ValueError("Минимальная сумма для СБП — 10₽") - url = await generate_kassai_payment_link(amount_int, tg_id, method, session) + url = await generate_kassai_payment_link( + amount_int, + tg_id, + method, + session, + payment_id=payment_id, + success_url=success_url, + failure_url=failure_url, + metadata=metadata, + ) if not url or url == "https://fk.life/": raise ValueError("Не удалось создать платёж KassaI") - return (url, None) + return (url, payment_id) return create_link diff --git a/handlers/payments/kassai/webhook.py b/handlers/payments/kassai/webhook.py deleted file mode 100644 index a330e7f1..00000000 --- a/handlers/payments/kassai/webhook.py +++ /dev/null @@ -1,130 +0,0 @@ -import hashlib - -from aiohttp import web - -from config import KASSAI_SECRET_KEY, KASSAI_SHOP_ID, KASSAI_WEBHOOK_RESPONSE -from core.webhook_abuse import ( - get_webhook_client_ip, - is_webhook_ip_blocked, - record_webhook_signature_failure, -) -from database import ( - add_payment, - async_session_maker, - get_payment_by_payment_id, - invalidate_payment_cache, - update_balance, - update_payment_status, -) -from handlers.payments.utils import send_payment_success_notification -from logger import logger - - -def verify_kassai_signature(data: dict, signature: str) -> bool: - """Проверяет подпись webhook от KassaAI. - - Args: - data: Данные от webhook - signature: Подпись для проверки - - Returns: - True если подпись валидна, иначе False - """ - try: - sign_string = ( - f"{KASSAI_SHOP_ID}:{data.get('AMOUNT', '')}:{KASSAI_SECRET_KEY}:{data.get('MERCHANT_ORDER_ID', '')}" - ) - expected_signature = hashlib.md5(sign_string.encode("utf-8")).hexdigest() - result = signature.upper() == expected_signature.upper() - if not result: - logger.error(f"KassaAI signature mismatch. Expected: {expected_signature}, Got: {signature}") - logger.error(f"Sign string: {sign_string}") - else: - logger.info("KassaAI webhook: подпись успешно проверена") - return result - except Exception as e: - logger.error(f"Ошибка проверки подписи KassaAI: {e}") - return False - - -async def kassai_webhook(request: web.Request): - """Обработчик webhook от KassaAI для aiohttp.""" - try: - ip = get_webhook_client_ip(request) - if await is_webhook_ip_blocked(ip): - return web.Response(status=429) - data = await request.post() - logger.info(f"KassaAI webhook received: {dict(data)}") - signature = data.get("SIGN", "") - if not signature: - logger.error("KassaAI webhook: отсутствует подпись") - await record_webhook_signature_failure(ip) - return web.Response(status=400) - if not verify_kassai_signature(data, signature): - logger.error("KassaAI webhook: неверная подпись") - await record_webhook_signature_failure(ip) - return web.Response(status=400) - - amount_raw = data.get("AMOUNT") - order_id = data.get("MERCHANT_ORDER_ID") - - if not amount_raw or not order_id: - logger.error("KassaAI webhook: отсутствуют обязательные параметры") - return web.Response(status=400) - - amount = float(amount_raw) - - try: - tg_id = int(order_id.split("_")[1]) - except (IndexError, ValueError) as e: - logger.error(f"KassaAI webhook: не удалось извлечь tg_id из order_id {order_id}: {e}") - return web.Response(status=400) - - logger.info(f"KassaAI: успешный платёж {order_id} на сумму {amount} RUB для пользователя {tg_id}") - - async with async_session_maker() as session: - payment = await get_payment_by_payment_id(session, order_id) - if payment: - if payment.get("status") == "success": - logger.info(f"KassaAI: платёж {order_id} уже обработан") - return web.Response(text=KASSAI_WEBHOOK_RESPONSE) - if payment.get("id") is not None: - ok = await update_payment_status( - session=session, internal_id=int(payment["id"]), new_status="success" - ) - if not ok: - logger.error(f"KassaAI: не удалось обновить статус платежа {order_id}") - return web.Response(status=500) - else: - await add_payment( - session=session, - tg_id=tg_id, - amount=amount, - payment_system="kassai", - status="success", - currency="RUB", - payment_id=order_id, - metadata=None, - ) - else: - await add_payment( - session=session, - tg_id=tg_id, - amount=amount, - payment_system="KASSAI", - status="success", - currency="RUB", - payment_id=order_id, - metadata=None, - ) - - await update_balance(session, tg_id, amount) - await send_payment_success_notification(tg_id, amount, session) - await invalidate_payment_cache(order_id) - logger.info( - f"KassaAI: платёж {order_id} успешно обработан, баланс пользователя {tg_id} пополнен на {amount} RUB" - ) - return web.Response(text=KASSAI_WEBHOOK_RESPONSE) - except Exception as e: - logger.error(f"Ошибка обработки KassaAI webhook: {e}") - return web.Response(status=500) diff --git a/handlers/payments/keyboards.py b/handlers/payments/keyboards.py index 6d690f84..18cddbde 100644 --- a/handlers/payments/keyboards.py +++ b/handlers/payments/keyboards.py @@ -5,7 +5,7 @@ from aiogram.utils.keyboard import InlineKeyboardBuilder from config import RENEWAL_PRICES from handlers.buttons import BACK, CUSTOM_AMOUNT -from handlers.payments.currency_rates import format_for_user +from services.payments.currency_rates import format_for_user async def payment_options_for_user( diff --git a/handlers/payments/pay.py b/handlers/payments/pay.py index 33c93196..d35ee42a 100644 --- a/handlers/payments/pay.py +++ b/handlers/payments/pay.py @@ -1,4 +1,5 @@ import os + from typing import Any from aiogram import F, Router @@ -9,24 +10,24 @@ from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from config import DONATIONS_ENABLE, TRIBUTE_LINK -from core.bootstrap import PAYMENTS_CONFIG, BUTTONS_CONFIG +from core.bootstrap import BUTTONS_CONFIG, PAYMENTS_CONFIG from core.settings.money_config import get_currency_mode from database import get_last_payments from database.models import User from handlers import buttons as btn from handlers.payments.currency_flow import build_currency_choice_kb -from handlers.payments.currency_rates import format_for_user -from handlers.payments.providers import get_providers_with_hooks from handlers.payments.stars.handlers import process_callback_pay_stars from handlers.payments.tribute.handlers import process_callback_pay_tribute from handlers.texts import ( - FAST_PAY_CHOOSE_CURRENCY, - BALANCE_MANAGEMENT_TEXT, - PAYMENT_METHODS_MSG, BALANCE_HISTORY_HEADER, + BALANCE_MANAGEMENT_TEXT, + FAST_PAY_CHOOSE_CURRENCY, + PAYMENT_METHODS_MSG, ) from hooks.hook_buttons import insert_hook_buttons from hooks.hooks import run_hooks +from services.payments.currency_rates import format_for_user +from services.payments.providers import get_providers_with_hooks from ..utils import edit_or_send_message diff --git a/handlers/payments/robokassa/handlers.py b/handlers/payments/robokassa/handlers.py index ecdb7c55..3eed4543 100644 --- a/handlers/payments/robokassa/handlers.py +++ b/handlers/payments/robokassa/handlers.py @@ -3,11 +3,11 @@ from typing import Any from aiogram import F, Router, types from aiogram.fsm.context import FSMContext from aiogram.fsm.state import State, StatesGroup -from aiogram.types import InlineKeyboardMarkup, InlineKeyboardButton +from aiogram.types import InlineKeyboardButton, InlineKeyboardMarkup from sqlalchemy.ext.asyncio import AsyncSession from database import add_user, check_user_exists, get_key_count, get_temporary_data -from handlers.buttons import CUSTOM_AMOUNT, PAY_2, MAIN_MENU +from handlers.buttons import CUSTOM_AMOUNT, MAIN_MENU, PAY_2 from handlers.payments.keyboards import ( back_keyboard, build_amounts_keyboard, @@ -16,12 +16,12 @@ from handlers.payments.keyboards import ( payment_options_for_user, ) from handlers.texts import DEFAULT_PAYMENT_MESSAGE, ENTER_SUM -from handlers.payments.currency_rates import format_for_user from handlers.utils import edit_or_send_message from logger import logger -from ..constants import ALLOWED_TEMP_PAYMENT_STATES +from services.payments.currency_rates import format_for_user +from services.payments.robokassa.service import create_and_store_robokassa_payment -from .service import create_and_store_robokassa_payment +from ..constants import ALLOWED_TEMP_PAYMENT_STATES router = Router() diff --git a/handlers/payments/robokassa/service.py b/handlers/payments/robokassa/service.py deleted file mode 100644 index 821f7bfc..00000000 --- a/handlers/payments/robokassa/service.py +++ /dev/null @@ -1,103 +0,0 @@ -import hashlib -import json -import uuid - -from decimal import ROUND_DOWN, Decimal -from urllib.parse import quote_plus, urlencode - -from sqlalchemy.ext.asyncio import AsyncSession - -from config import ROBOKASSA_LOGIN, ROBOKASSA_PASSWORD1, ROBOKASSA_PASSWORD2, ROBOKASSA_TEST_MODE -from database import register_pending_payment -from handlers.payments.payment_links import register_payment_creator - - -def _build_receipt(amount: float, sno: str = "usn_income") -> dict: - return { - "items": [ - { - "name": "Пополнение баланса", - "quantity": 1, - "sum": float(amount), - "payment_method": "full_payment", - "payment_object": "payment", - "tax": "none", - } - ], - "sno": sno, - } - - -def _format_amount(amount: float | int) -> str: - s = str(Decimal(str(amount)).quantize(Decimal("0.01"), rounding=ROUND_DOWN)) - return s.rstrip("0").rstrip(".") if "." in s else s - - -def generate_payment_link(amount: int | float, inv_id: int, description: str, tg_id: int) -> tuple[str, str]: - out_sum = _format_amount(amount) - receipt_json = json.dumps(_build_receipt(amount), ensure_ascii=False, separators=(",", ":")) - receipt_enc = quote_plus(receipt_json, safe="") - pid = str(uuid.uuid4()) - shp = {"Shp_id": str(tg_id), "Shp_pid": pid} - base = f"{ROBOKASSA_LOGIN}:{out_sum}:{inv_id}:{receipt_enc}:{ROBOKASSA_PASSWORD1}" - for k in sorted(shp.keys(), key=str.lower): - base += f":{k}={shp[k]}" - signature = hashlib.md5(base.encode("utf-8")).hexdigest().upper() - query = { - "MrchLogin": ROBOKASSA_LOGIN, - "OutSum": out_sum, - "InvId": inv_id, - "Description": description, - "Receipt": receipt_enc, - "SignatureValue": signature, - **shp, - } - if ROBOKASSA_TEST_MODE: - query["IsTest"] = 1 - return "https://auth.robokassa.ru/Merchant/Index.aspx?" + urlencode(query), pid - - -async def create_and_store_robokassa_payment( - session: AsyncSession, tg_id: int, amount: int | float, description: str, inv_id: int = 0 -) -> tuple[str, str]: - url, pid = generate_payment_link(amount, inv_id, description, tg_id) - await register_pending_payment( - payment_id=pid, - tg_id=tg_id, - amount=float(amount), - payment_system="robokassa", - currency="RUB", - ) - return url, pid - - -def check_payment_signature(params) -> bool: - out_sum = params.get("OutSum") or params.get("out_summ") or params.get("outsumm") - inv_id = params.get("InvId") or params.get("inv_id") or params.get("invid") - received_sig = (params.get("SignatureValue") or params.get("signaturevalue") or "").upper() - if not out_sum or not inv_id or not received_sig: - return False - shp_items = [(k, params[k]) for k in params.keys() if k.lower().startswith("shp_")] - shp_items.sort(key=lambda kv: kv[0].lower()) - shp_suffix = "".join(f":{k}={v}" for k, v in shp_items) - base = f"{out_sum}:{inv_id}:{ROBOKASSA_PASSWORD2}{shp_suffix}" - expected_sig = hashlib.md5(base.encode("utf-8")).hexdigest().upper() - return received_sig == expected_sig - - -async def create_link( - session: AsyncSession, - tg_id: int, - amount: float, - currency: str, - success_url: str | None, - failure_url: str | None, -) -> tuple[str, str]: - if currency != "RUB": - raise ValueError("Robokassa поддерживает только RUB") - amount_val = int(amount) if amount == int(amount) else amount - url, pid = await create_and_store_robokassa_payment(session, tg_id, amount_val, "Пополнение баланса", inv_id=0) - return (url, pid) - - -register_payment_creator("ROBOKASSA", create_link) diff --git a/handlers/payments/robokassa/webhook.py b/handlers/payments/robokassa/webhook.py deleted file mode 100644 index 89c69cee..00000000 --- a/handlers/payments/robokassa/webhook.py +++ /dev/null @@ -1,84 +0,0 @@ -from aiohttp import web - -from core.webhook_abuse import ( - get_webhook_client_ip, - is_webhook_ip_blocked, - record_webhook_signature_failure, -) -from database import ( - add_payment, - async_session_maker, - get_payment_by_payment_id, - invalidate_payment_cache, - update_balance, - update_payment_status, -) -from handlers.payments.utils import send_payment_success_notification -from logger import logger - -from .service import check_payment_signature - - -async def robokassa_webhook(request: web.Request): - try: - ip = get_webhook_client_ip(request) - if await is_webhook_ip_blocked(ip): - return web.Response(status=429) - params = await request.post() - if not check_payment_signature(params): - await record_webhook_signature_failure(ip) - return web.Response(status=400) - - amount_raw = params.get("OutSum") - inv_id = params.get("InvId") - shp_id = params.get("Shp_id") or params.get("shp_id") or params.get("id") - shp_pid = params.get("Shp_pid") or params.get("shp_pid") or params.get("pid") - - if not amount_raw or not inv_id or not shp_id or not shp_pid: - return web.Response(status=400) - - tg_id = int(shp_id) - amount = float(amount_raw) - - async with async_session_maker() as session: - payment = await get_payment_by_payment_id(session, shp_pid) - if payment: - if payment.get("status") == "success": - return web.Response(text=f"OK{inv_id}") - if payment.get("id") is not None: - ok = await update_payment_status( - session=session, internal_id=int(payment["id"]), new_status="success" - ) - if not ok: - return web.Response(status=500) - else: - await add_payment( - session=session, - tg_id=tg_id, - amount=amount, - payment_system="robokassa", - status="success", - currency="RUB", - payment_id=shp_pid, - metadata=None, - ) - else: - await add_payment( - session=session, - tg_id=tg_id, - amount=amount, - payment_system="robokassa", - status="success", - currency="RUB", - payment_id=shp_pid, - metadata=None, - ) - - await update_balance(session, tg_id, amount) - await send_payment_success_notification(tg_id, amount, session) - await invalidate_payment_cache(shp_pid) - - return web.Response(text=f"OK{inv_id}") - except Exception as e: - logger.error(f"Error processing ROBOKASSA webhook: {e}") - return web.Response(status=500) diff --git a/handlers/profile.py b/handlers/profile.py index 528f3d70..4b83bb23 100644 --- a/handlers/profile.py +++ b/handlers/profile.py @@ -2,7 +2,7 @@ import os from aiogram import F, Router from aiogram.fsm.context import FSMContext -from aiogram.types import CallbackQuery, InlineKeyboardButton, Message +from aiogram.types import CallbackQuery, InlineKeyboardButton, Message, WebAppInfo from aiogram.utils.keyboard import InlineKeyboardBuilder from config import ( @@ -31,7 +31,7 @@ from handlers.buttons import ( TRIAL_SUB, ADMIN_BTN, ) -from handlers.payments.currency_rates import format_for_user +from services.payments.currency_rates import format_for_user from handlers.texts import ADD_SUBSCRIPTION_HINT from hooks.hook_buttons import insert_hook_buttons from hooks.hooks import run_hooks @@ -108,6 +108,13 @@ async def process_callback_view_profile( builder = InlineKeyboardBuilder() + from core.settings.web_config import get_site_url, is_web_enabled + if is_web_enabled(): + site_url = get_site_url() + if site_url: + webapp_url = f"{site_url}/dashboard" + builder.row(InlineKeyboardButton(text="🌐 Личный кабинет", web_app=WebAppInfo(url=webapp_url))) + trial_time_disabled = bool(MODES_CONFIG.get("TRIAL_TIME_DISABLED", TRIAL_TIME_DISABLE)) if key_count > 0: diff --git a/handlers/refferal.py b/handlers/refferal.py index 9e4bffb8..e8313622 100644 --- a/handlers/refferal.py +++ b/handlers/refferal.py @@ -29,9 +29,10 @@ from database import ( get_referral_stats, ) from database.models import Referral +from database.access.resolution import resolve_user_optional from database.tariffs import get_tariffs from handlers.buttons import BACK, INVITE, MAIN_MENU, QR, TOP_FIVE -from handlers.payments.currency_rates import format_for_user +from services.payments.currency_rates import format_for_user from handlers.texts import ( INVITE_MESSAGE_TEMPLATE, INVITE_TEXT_NON_INLINE, @@ -200,16 +201,22 @@ async def show_referral_qr(callback_query: CallbackQuery): @router.callback_query(F.data == "top_referrals") async def top_referrals_handler(callback_query: CallbackQuery, session: AsyncSession): user_id = callback_query.from_user.id + u = await resolve_user_optional(session, user_id) + uid = u.id if u is not None else None - result = await session.execute(select(func.count()).select_from(Referral).where(Referral.referrer_tg_id == user_id)) - user_referral_count = result.scalar_one() or 0 + user_referral_count = 0 + if uid is not None: + result = await session.execute( + select(func.count()).select_from(Referral).where(Referral.referrer_user_id == uid) + ) + user_referral_count = result.scalar_one() or 0 personal_block = "Твоё место в рейтинге:\n" if user_referral_count > 0: subquery = ( select(func.count().label("cnt")) .select_from(Referral) - .group_by(Referral.referrer_tg_id) + .group_by(Referral.referrer_user_id) .having(func.count() > user_referral_count) .subquery() ) @@ -221,10 +228,10 @@ async def top_referrals_handler(callback_query: CallbackQuery, session: AsyncSes result = await session.execute( select( - Referral.referrer_tg_id, - func.count(Referral.referred_tg_id).label("referral_count"), + Referral.referrer_user_id, + func.count(Referral.referred_user_id).label("referral_count"), ) - .group_by(Referral.referrer_tg_id) + .group_by(Referral.referrer_user_id) .order_by(desc("referral_count")) .limit(5) ) @@ -233,7 +240,7 @@ async def top_referrals_handler(callback_query: CallbackQuery, session: AsyncSes is_admin = user_id in ADMIN_ID rows = "" for index, row in enumerate(top_referrals, 1): - referrer_id = str(row.referrer_tg_id) + referrer_id = str(row.referrer_user_id) count = row.referral_count display_id = referrer_id if is_admin else f"{referrer_id[:5]}*****" rows += f"{index}. {display_id} - {count} чел.\n" @@ -287,17 +294,19 @@ async def handle_referral_link( is_bot=getattr(user, "is_bot", False), ) - if not inserted: + if inserted is None: await message.answer("❌ Вы уже зарегистрированы и не можете стать рефералом.") return await add_referral(session, user_id, referrer_tg_id) try: - await bot.send_message( - referrer_tg_id, - NEW_REFERRAL_NOTIFICATION.format(referred_id=user_id), - ) + ref_notifier = await resolve_user_optional(session, referrer_tg_id) + if ref_notifier is not None and ref_notifier.tg_id is not None: + await bot.send_message( + int(ref_notifier.tg_id), + NEW_REFERRAL_NOTIFICATION.format(referred_id=user_id), + ) except Exception as error: logger.error(f"Не удалось отправить уведомление пригласившему ({referrer_tg_id}): {error}") diff --git a/handlers/start.py b/handlers/start.py index 657223ab..a3255b06 100644 --- a/handlers/start.py +++ b/handlers/start.py @@ -257,8 +257,8 @@ async def handle_referral_link_safe(part, message, state, session, user_data): try: referrer_id = int(part.split("referral")[1].strip("_")) await handle_referral_link(referrer_id, message, state, session, user_data) - except Exception: - pass + except Exception as e: + logger.warning("[Referral] Ошибка обработки реферальной ссылки '{}': {}", part, e) async def prompt_subscription(callback: CallbackQuery): diff --git a/handlers/tariffs/addons/key_addons_main.py b/handlers/tariffs/addons/key_addons_main.py index ecef8ab4..fb0b10ea 100644 --- a/handlers/tariffs/addons/key_addons_main.py +++ b/handlers/tariffs/addons/key_addons_main.py @@ -18,9 +18,9 @@ from handlers.buttons import ( ) from middlewares.session import release_session_early from handlers.keys.key_view import render_key_info -from handlers.payments.currency_rates import format_for_user +from services.payments.currency_rates import format_for_user from handlers.payments.fast_payment_flow import try_fast_payment_flow -from handlers.tariffs.tariff_display import GB, get_effective_limits_for_key +from services.tariffs.tariff_display import GB, get_effective_limits_for_key from handlers.texts import ( ADDONS_APPLIED_TEXT, DOWNGRADE_INLINE_WARNING_TEXT, @@ -678,7 +678,7 @@ async def handle_addons_downgrade_apply(callback: CallbackQuery, state: FSMConte @router.callback_query(F.data == "key_addons_confirm", KeyAddonConfigState.configuring) async def handle_addons_confirm(callback: CallbackQuery, state: FSMContext, session: AsyncSession): - from handlers.keys.operations import renew_key_in_cluster + from services.operations import renew_key_in_cluster tg_id = callback.from_user.id data = await state.get_data() diff --git a/handlers/tariffs/addons/key_addons_pack.py b/handlers/tariffs/addons/key_addons_pack.py index 9b86842b..ed6e9667 100644 --- a/handlers/tariffs/addons/key_addons_pack.py +++ b/handlers/tariffs/addons/key_addons_pack.py @@ -21,9 +21,9 @@ from database.models import User from handlers.buttons import BACK, CONFIRM_ADDON_BUTTON_TEXT, PAYMENT from middlewares.session import release_session_early from handlers.keys.key_view import render_key_info -from handlers.payments.currency_rates import format_for_user +from services.payments.currency_rates import format_for_user from handlers.payments.fast_payment_flow import try_fast_payment_flow -from handlers.tariffs.tariff_display import GB, get_effective_limits_for_key +from services.tariffs.tariff_display import GB, get_effective_limits_for_key from handlers.texts import ( ADDONS_NO_EXTRA_PAYMENT_TEXT, ADDONS_PACK_SUCCESS_TEXT, @@ -622,7 +622,7 @@ async def handle_addons_traffic_choice(callback: CallbackQuery, state: FSMContex @router.callback_query(F.data == "key_addons_confirm", KeyAddonConfigState.configuring) async def handle_addons_confirm(callback: CallbackQuery, state: FSMContext, session: AsyncSession): - from handlers.keys.operations import renew_key_in_cluster + from services.operations import renew_key_in_cluster tg_id = callback.from_user.id data = await state.get_data() diff --git a/handlers/tariffs/buy/key_tariffs.py b/handlers/tariffs/buy/key_tariffs.py index 88349ca1..aa74cd86 100644 --- a/handlers/tariffs/buy/key_tariffs.py +++ b/handlers/tariffs/buy/key_tariffs.py @@ -13,9 +13,9 @@ from core.settings.tariffs_config import normalize_tariff_config from database import get_balance, get_tariff_by_id from database.notifications import check_hot_lead_discount from handlers.buttons import BACK, CONFIG_PAY_BUTTON_TEXT, MAIN_MENU, PAYMENT -from handlers.payments.currency_rates import format_for_user +from services.payments.currency_rates import format_for_user from handlers.payments.fast_payment_flow import try_fast_payment_flow -from handlers.tariffs.tariff_display import GB +from services.tariffs.tariff_display import GB from handlers.texts import ( CONFIG_SCREEN_TEMPLATE, CREATING_CONNECTION_MSG, @@ -890,7 +890,7 @@ async def handle_user_devices_choice(callback: CallbackQuery, state: FSMContext, ) async def handle_user_traffic_choice(callback: CallbackQuery, state: FSMContext, session: Any): """Обрабатывает выбор лимита трафика в конфигураторе.""" - await safe_answer_callback(callback) + await safe_answer_callback(callback) _, _tariff_id_str, traffic_str = callback.data.split("|", 2) traffic = int(traffic_str) diff --git a/handlers/utils.py b/handlers/utils.py index bc73fe9a..5002b3b3 100644 --- a/handlers/utils.py +++ b/handlers/utils.py @@ -26,6 +26,7 @@ from bot import bot from config import ADMIN_ID from database import get_servers from database.models import Key, Notification, Server +from database.access.resolution import resolve_user_optional from hooks.processors import process_cluster_balancer from logger import logger @@ -89,98 +90,47 @@ async def generate_random_email( async def get_least_loaded_cluster(session: AsyncSession) -> str: - servers = await get_servers(session) - server_to_cluster = {} - cluster_loads = {} - - for cluster_name, cluster_servers in servers.items(): - cluster_loads[cluster_name] = 0 - for server in cluster_servers: - server_to_cluster[server["server_name"]] = cluster_name - - result = await session.execute(select(Key)) - keys = result.scalars().all() - - for key in keys: - server_id = key.server_id - cluster_id = server_to_cluster.get(server_id, server_id) - if cluster_id in cluster_loads: - cluster_loads[cluster_id] += 1 - - available_clusters = {} - for cluster_name, cluster_servers in servers.items(): - enabled_servers = [server for server in cluster_servers if server.get("enabled", True)] - - if not enabled_servers: - continue - - available_servers = [] - for server in enabled_servers: - if await check_server_key_limit(server, session): - available_servers.append(server) - - if available_servers: - available_clusters[cluster_name] = cluster_loads[cluster_name] - else: - continue - - filtered_clusters = await process_cluster_balancer(available_clusters=available_clusters, session=session) - if filtered_clusters: - available_clusters = filtered_clusters - - if not available_clusters: - logger.warning("❌ Нет доступных кластеров с лимитом ключей!") - raise ValueError("⚠️ Сервисы временно недоступны. Попробуйте позже.") - - least_loaded_cluster = min(available_clusters, key=lambda k: (available_clusters[k], k)) - logger.info( - f"Выбран наименее загруженный кластер: {least_loaded_cluster} (загрузка: {available_clusters[least_loaded_cluster]})" - ) - return least_loaded_cluster + """Делегирует в services.clusters.select_cluster().""" + from services.clusters import select_cluster + result = await select_cluster(session) + return result.cluster_name async def check_server_key_limit(server_info: dict, session: AsyncSession) -> bool: - server_name = server_info.get("server_name") - cluster_name = server_info.get("cluster_name") - max_keys = server_info.get("max_keys") + """Делегирует в services.clusters.check_server_key_limit() с Telegram-callback для уведомлений.""" + from services.clusters import check_server_key_limit as _svc_check - if not max_keys: - return True - - identifier = cluster_name if cluster_name else server_name - - result = await session.execute(select(func.count()).select_from(Key).where(Key.server_id == identifier)) - total_keys = result.scalar() or 0 - - if total_keys >= max_keys: - logger.warning(f"[Key Limit] Сервер {server_name} достиг лимита: {total_keys}/{max_keys}") - return False - - usage_percent = total_keys / max_keys - - if usage_percent >= 0.9: + async def _notify_admin_capacity(server_name: str, total_keys: int, max_keys: int) -> None: notif_key = f"server_warn_{server_name}" - - result = await session.execute( - select(Notification).where(Notification.tg_id == 0, Notification.notification_type == notif_key) - ) - already_sent = result.scalar_one_or_none() - + anchor_uid = None + if ADMIN_ID: + au = await resolve_user_optional(session, int(ADMIN_ID[0])) + if au is not None: + anchor_uid = au.id + already_sent = None + if anchor_uid is not None: + result = await session.execute( + select(Notification).where( + Notification.user_id == anchor_uid, + Notification.notification_type == notif_key, + ) + ) + already_sent = result.scalar_one_or_none() if not already_sent: for admin_id in ADMIN_ID: try: await bot.send_message( admin_id, - f"⚠️ Сервер {server_name} почти заполнен ({int(usage_percent * 100)}%)." + f"⚠️ Сервер {server_name} почти заполнен ({int(total_keys / max_keys * 100)}%)." f"\nРекомендуется создать новый для балансировки.", ) except Exception: pass + if anchor_uid is not None: + session.add(Notification(user_id=anchor_uid, notification_type=notif_key)) + await session.commit() - session.add(Notification(tg_id=0, notification_type=notif_key)) - await session.commit() - - return True + return await _svc_check(server_info, session, on_capacity_warning=_notify_admin_capacity) async def handle_error(tg_id: int, callback_query: object | None = None, message: str = "") -> None: @@ -444,15 +394,8 @@ def convert_to_bytes(value: float, unit: str) -> int: async def is_full_remnawave_cluster(cluster_id: str, session: AsyncSession) -> bool: - result = await session.execute(select(Server.panel_type).where(Server.cluster_name == cluster_id)) - panel_types = result.scalars().all() - - if panel_types: - return all(pt.lower() == "remnawave" for pt in panel_types) - - result = await session.execute(select(Server.panel_type).where(Server.server_name == cluster_id)) - panel_type = result.scalar_one_or_none() - return panel_type and panel_type.lower() == "remnawave" + from services.clusters import is_full_remnawave_cluster as _svc + return await _svc(cluster_id, session) def sanitize_key_name(key_name: str) -> str: diff --git a/hooks/processors.py b/hooks/processors.py index 474cee5c..84a6a93e 100644 --- a/hooks/processors.py +++ b/hooks/processors.py @@ -6,17 +6,21 @@ from .hooks import run_hooks async def process_cluster_override( - tg_id: int, - state_data: dict, session: Any, + tg_id: int | None = None, + state_data: dict | None = None, plan: int | None = None, **kwargs, ) -> str | None: - """Обрабатывает хук cluster_override и возвращает название кластера.""" + """Обрабатывает хук cluster_override и возвращает название кластера. + + tg_id и state_data опциональны — хук может вызываться из контекстов + балансировщика (services.clusters.select_cluster), где этих данных нет. + """ results = await run_hooks( "cluster_override", tg_id=tg_id, - state_data=state_data, + state_data=state_data if state_data is not None else {}, session=session, plan=plan, **kwargs, @@ -158,7 +162,7 @@ async def process_get_cryptolink_after_renewal( try: from config import REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD from database import get_tariff_by_id - from handlers.keys.operations.utils import is_plan_vless + from services.operations.utils import is_plan_vless from panels.remnawave import RemnawaveAPI remna = RemnawaveAPI(remnawave_nodes[0]["api_url"]) diff --git a/install.sh b/install.sh new file mode 100644 index 00000000..564a3e96 --- /dev/null +++ b/install.sh @@ -0,0 +1,36 @@ +#!/bin/bash +set -e + +REPO="Vladless/Solo_bot" +BRANCH="main" +INSTALL_DIR="/root/solobot" +CLI="cli_launcher.py" +CMD_NAME="solobot" +CMD_PATH="/usr/local/bin/$CMD_NAME" + +command -v python3 >/dev/null 2>&1 || { + apt-get update -qq >/dev/null 2>&1 + apt-get install -y -qq python3 git curl >/dev/null 2>&1 +} + +command -v git >/dev/null 2>&1 || { + apt-get update -qq >/dev/null 2>&1 + apt-get install -y -qq git >/dev/null 2>&1 +} + +if [ ! -f "$INSTALL_DIR/$CLI" ]; then + echo "Загрузка Solo Bot..." + git clone --depth 1 --branch "$BRANCH" "https://github.com/$REPO.git" "$INSTALL_DIR" 2>/dev/null || { + echo "Ошибка загрузки. Проверьте подключение к интернету." + exit 1 + } +fi + +cat > "$CMD_PATH" << EOF +#!/bin/bash +cd $INSTALL_DIR && python3 $CLI "\$@" +EOF +chmod +x "$CMD_PATH" + +cd "$INSTALL_DIR" +exec python3 "$CLI" diff --git a/mail/__init__.py b/mail/__init__.py new file mode 100644 index 00000000..a9aab49d --- /dev/null +++ b/mail/__init__.py @@ -0,0 +1,3 @@ +from .smtp import send_email_link_code_email, send_email_verify_code_email, send_login_code_email, send_password_reset_code_email, smtp_configured + +__all__ = ["send_login_code_email", "send_password_reset_code_email", "send_email_link_code_email", "send_email_verify_code_email", "smtp_configured"] diff --git a/mail/smtp.py b/mail/smtp.py new file mode 100644 index 00000000..41433feb --- /dev/null +++ b/mail/smtp.py @@ -0,0 +1,132 @@ +from email.message import EmailMessage + +import aiosmtplib + +from config import ( + EMAIL_FROM, + EMAIL_SMTP_HOST, + EMAIL_SMTP_PASSWORD, + EMAIL_SMTP_PORT, + EMAIL_SMTP_USER, + PROJECT_NAME, +) +from logger import logger + + +_SMTP_TIMEOUT_SEC = 30.0 +_SMTP_VALIDATE_CERTS = True + + +def smtp_configured() -> bool: + return bool(EMAIL_SMTP_HOST and (EMAIL_FROM or EMAIL_SMTP_USER)) + + +def _get_email_template(key: str, default: str) -> str: + """Читает шаблон из WEB_CONFIG (настраиваемый админом), fallback на default.""" + try: + from core.settings.web_config import WEB_CONFIG + val = WEB_CONFIG.get(key) + return str(val).strip() if val else default + except Exception: + return default + + +def _render(template: str, **kwargs: str) -> str: + try: + return template.format(**kwargs) + except (KeyError, ValueError): + return template + + +def _smtp_kwargs() -> dict: + kwargs: dict = { + "hostname": EMAIL_SMTP_HOST, + "port": EMAIL_SMTP_PORT, + "username": EMAIL_SMTP_USER or None, + "password": EMAIL_SMTP_PASSWORD or None, + "timeout": _SMTP_TIMEOUT_SEC, + "validate_certs": _SMTP_VALIDATE_CERTS, + } + if EMAIL_SMTP_PORT == 465: + kwargs["use_tls"] = True + kwargs["start_tls"] = False + else: + kwargs["use_tls"] = False + kwargs["start_tls"] = True + return kwargs + + +async def send_login_code_email(to_addr: str, code: str) -> None: + if not smtp_configured(): + raise RuntimeError("smtp_not_configured") + from_addr = EMAIL_FROM or EMAIL_SMTP_USER + project = PROJECT_NAME + subject = _render(_get_email_template("EMAIL_LOGIN_SUBJECT", "{project}: код для входа"), project=project, code=code) + body = _render(_get_email_template("EMAIL_LOGIN_BODY", "Код для входа: {code}"), project=project, code=code) + msg = EmailMessage() + msg["Subject"] = subject + msg["From"] = f"{project} <{from_addr}>" + msg["To"] = to_addr + msg.set_content(body) + try: + await aiosmtplib.send(msg, **_smtp_kwargs()) + except Exception as exc: + logger.warning(f"[SMTP] Отправка кода входа на {to_addr} не удалась: {exc}") + raise + + +async def send_password_reset_code_email(to_addr: str, code: str) -> None: + if not smtp_configured(): + raise RuntimeError("smtp_not_configured") + from_addr = EMAIL_FROM or EMAIL_SMTP_USER + project = PROJECT_NAME + subject = _render(_get_email_template("EMAIL_RESET_SUBJECT", "{project}: сброс пароля"), project=project, code=code) + body = _render(_get_email_template("EMAIL_RESET_BODY", "Код для сброса пароля: {code}"), project=project, code=code) + msg = EmailMessage() + msg["Subject"] = subject + msg["From"] = f"{project} <{from_addr}>" + msg["To"] = to_addr + msg.set_content(body) + try: + await aiosmtplib.send(msg, **_smtp_kwargs()) + except Exception as exc: + logger.warning(f"[SMTP] Отправка кода сброса пароля на {to_addr} не удалась: {exc}") + raise + + +async def send_email_verify_code_email(to_addr: str, code: str) -> None: + if not smtp_configured(): + raise RuntimeError("smtp_not_configured") + from_addr = EMAIL_FROM or EMAIL_SMTP_USER + project = PROJECT_NAME + subject = _render(_get_email_template("EMAIL_VERIFY_SUBJECT", "{project}: подтверждение email"), project=project, code=code) + body = _render(_get_email_template("EMAIL_VERIFY_BODY", "Код подтверждения email: {code}"), project=project, code=code) + msg = EmailMessage() + msg["Subject"] = subject + msg["From"] = f"{project} <{from_addr}>" + msg["To"] = to_addr + msg.set_content(body) + try: + await aiosmtplib.send(msg, **_smtp_kwargs()) + except Exception as exc: + logger.warning(f"[SMTP] Отправка кода подтверждения email на {to_addr} не удалась: {exc}") + raise + + +async def send_email_link_code_email(to_addr: str, code: str) -> None: + if not smtp_configured(): + raise RuntimeError("smtp_not_configured") + from_addr = EMAIL_FROM or EMAIL_SMTP_USER + project = PROJECT_NAME + subject = _render(_get_email_template("EMAIL_LINK_SUBJECT", "{project}: подтверждение привязки email"), project=project, code=code) + body = _render(_get_email_template("EMAIL_LINK_BODY", "Код для привязки email: {code}"), project=project, code=code) + msg = EmailMessage() + msg["Subject"] = subject + msg["From"] = f"{project} <{from_addr}>" + msg["To"] = to_addr + msg.set_content(body) + try: + await aiosmtplib.send(msg, **_smtp_kwargs()) + except Exception as exc: + logger.warning(f"[SMTP] Отправка кода привязки email на {to_addr} не удалась: {exc}") + raise diff --git a/main.py b/main.py old mode 100644 new mode 100755 diff --git a/middlewares/__init__.py b/middlewares/__init__.py index c2a1ee25..502b06e7 100644 --- a/middlewares/__init__.py +++ b/middlewares/__init__.py @@ -6,6 +6,7 @@ from core.bootstrap import MODES_CONFIG from middlewares.ban_checker import BanCheckerMiddleware from middlewares.subscription import SubscriptionMiddleware +from .actor import ActorMiddleware from .admin import AdminMiddleware from .answer import CallbackAnswerMiddleware, EarlyCallbackAnswerMiddleware from .concurrency import ConcurrencyLimiterMiddleware @@ -46,6 +47,7 @@ def register_middleware( "logging": "LOGGING_MIDDLEWARE_ENABLED", "throttling": "THROTTLING_MIDDLEWARE_ENABLED", "user": "USER_MIDDLEWARE_ENABLED", + "actor": "ACTOR_MIDDLEWARE_ENABLED", "answer": "ANSWER_MIDDLEWARE_ENABLED", } @@ -83,6 +85,7 @@ def register_middleware( "logging": LoggingMiddleware(sessionmaker) if sessionmaker else LoggingMiddleware(), "throttling": ThrottlingMiddleware(), "user": UserMiddleware(), + "actor": ActorMiddleware(), "answer": CallbackAnswerMiddleware(), } middlewares = [ diff --git a/middlewares/actor.py b/middlewares/actor.py new file mode 100644 index 00000000..a8f5f3bf --- /dev/null +++ b/middlewares/actor.py @@ -0,0 +1,30 @@ +from collections.abc import Awaitable, Callable +from typing import Any + +from aiogram import BaseMiddleware +from aiogram.types import TelegramObject, User + +from database.access.resolution import resolve_actor_from_legacy_ref +from logger import logger + + +class ActorMiddleware(BaseMiddleware): + async def __call__( + self, + handler: Callable[[TelegramObject, dict[str, Any]], Awaitable[Any]], + event: TelegramObject, + data: dict[str, Any], + ) -> Any: + try: + from_user: User | None = data.get("event_from_user") + session = data.get("session") + if ( + from_user is not None + and not from_user.is_bot + and session is not None + and getattr(session, "execute", None) is not None + ): + data["actor"] = await resolve_actor_from_legacy_ref(session, int(from_user.id)) + except Exception as error: + logger.error(f"[ActorMiddleware] Ошибка резолва actor: {error}") + return await handler(event, data) diff --git a/middlewares/ban_checker.py b/middlewares/ban_checker.py index 32aabaec..a536fa7e 100644 --- a/middlewares/ban_checker.py +++ b/middlewares/ban_checker.py @@ -12,7 +12,7 @@ from config import ADMIN_ID, SUPPORT_CHAT_URL from core.cache_config import BAN_CACHE_TTL_SEC from core.redis_cache import cache_delete, cache_get, cache_key, cache_set from database import async_session_maker -from database.models import ManualBan +from database.models import ManualBan, User from logger import logger @@ -31,8 +31,9 @@ class BanCheckerMiddleware(BaseMiddleware): async def _load_ban_info(self, session: AsyncSession, tg_id: int) -> dict[str, Any] | None: query = ( select(ManualBan.reason, ManualBan.until) + .join(User, ManualBan.user_id == User.id) .where( - ManualBan.tg_id == tg_id, + User.tg_id == tg_id, (ManualBan.until.is_(None)) | (ManualBan.until > datetime.utcnow()), ) .limit(1) diff --git a/middlewares/loggings.py b/middlewares/loggings.py index 00988fae..084e2b70 100644 --- a/middlewares/loggings.py +++ b/middlewares/loggings.py @@ -62,6 +62,13 @@ class LoggingMiddleware(BaseMiddleware): try: result = await handler(event, data) db_user = data.get("user") + actor = data.get("actor") + if actor is not None: + set_telegram_actor( + audit_context, + identity_id=getattr(actor, "identity_id", None), + tg_id=getattr(actor, "telegram_chat_id", None), + ) if isinstance(db_user, dict): set_telegram_actor( audit_context, diff --git a/platform_loader.py b/platform_loader.py new file mode 100644 index 00000000..c1849aa1 --- /dev/null +++ b/platform_loader.py @@ -0,0 +1,40 @@ +import importlib +import os +import platform +import shutil +import sys + +COMPILED_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "compiled") + + +def _detect_platform(): + system = platform.system().lower() + if system == "darwin": + return "macos" + return "linux" + + +def install_platform_binaries(): + plat = _detect_platform() + plat_dir = os.path.join(COMPILED_DIR, plat) + + if not os.path.isdir(plat_dir): + return + + root = os.path.dirname(os.path.abspath(__file__)) + + for dirpath, _, filenames in os.walk(plat_dir): + for fname in filenames: + if not fname.endswith(".so"): + continue + src = os.path.join(dirpath, fname) + rel = os.path.relpath(dirpath, plat_dir) + dst_dir = os.path.join(root, rel) + dst = os.path.join(dst_dir, fname) + if os.path.exists(dst): + continue + os.makedirs(dst_dir, exist_ok=True) + shutil.copy2(src, dst) + + +install_platform_binaries() diff --git a/pyproject.toml b/pyproject.toml index 8a2b77b5..3fe84834 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "Solo_bot" -version = "0.0.1" +version = "0.5.3" dependencies = [ "aiofiles==24.1.0", "aiogram>=3.24.0", @@ -11,7 +11,7 @@ dependencies = [ "async-timeout==4.0.3", "asyncpg==0.30.0", "attrs==24.2.0", - "certifi>=2023.5.7,<2024.0.0", # ограничение из aiocryptopay; CVE-2024-39689 исправлен в 2024.7.4 + "certifi>=2023.5.7,<2024.0.0", "charset-normalizer==3.4.0", "deprecated==1.2.14", "distro==1.9.0", @@ -22,8 +22,6 @@ dependencies = [ "netaddr==1.3.0", "propcache==0.2.0", "pydantic", - "pydantic-core", - "requests==2.32.4", "typing-extensions==4.12.2", "urllib3==2.6.3", "wrapt==1.16.0", @@ -55,10 +53,9 @@ package = false select = ["E", "F", "W", "I", "N", "UP", "ANN", "ASYNC", "S", "BLE", "FBT", "B", "A", "C4", "DTZ", "T10", "ISC", "ICN", "G", "PIE"] ignore = ["ANN101", "ANN102", "S101", "ANN201", "ANN001", "BLE001", "W291", "ANN401", "DTZ003", "DTZ005", "F401", "FBT002", "FBT001", "FBT003", "A005", "E501", "UP017", "DTZ004", "W293", "ANN202", "DTZ007"] exclude = [ - ".git", + ".git", "venv", "main.py", - "handlers/payments", ] [tool.ruff.format] diff --git a/requirements.txt b/requirements.txt index 88839962..187b0e96 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,6 +1,7 @@ aiocryptopay==0.4.7 aiofiles==24.1.0 aiogram==3.24.0 +aiosmtplib==3.0.2 aiohappyeyeballs==2.5.0 aiohttp==3.13.3 aiosignal>=1.4.0 @@ -44,6 +45,7 @@ psycopg2-binary==2.9.10 py3xui==0.3.4 pycparser==2.22 pydantic==2.9.2 +pywebpush>=2.0.0 pydantic_core==2.23.4 Pygments==2.19.2 python-dateutil==2.9.0.post0 diff --git a/scripts/install-hooks.sh b/scripts/install-hooks.sh new file mode 100755 index 00000000..1d0c3e62 --- /dev/null +++ b/scripts/install-hooks.sh @@ -0,0 +1,25 @@ +#!/usr/bin/env bash +set -e + +REPO_ROOT="$(git rev-parse --show-toplevel)" +HOOKS_DIR="$REPO_ROOT/.git/hooks" +SCRIPTS_DIR="$REPO_ROOT/scripts" + +if [ ! -d "$HOOKS_DIR" ]; then + echo "Error: $HOOKS_DIR не найден. Это git-репозиторий?" + exit 1 +fi + +cat > "$HOOKS_DIR/pre-commit" <<'HOOK' +#!/usr/bin/env bash +set -e +REPO_ROOT="$(git rev-parse --show-toplevel)" +exec "$REPO_ROOT/scripts/pre-commit-secrets.sh" +HOOK + +chmod +x "$HOOKS_DIR/pre-commit" +echo "✓ pre-commit hook установлен" +echo " → $HOOKS_DIR/pre-commit" +echo "" +echo "Проверка секретов при каждом git commit." +echo "Если нужно пропустить (осознанно): git commit --no-verify" diff --git a/scripts/load-test.sh b/scripts/load-test.sh new file mode 100755 index 00000000..05a14923 --- /dev/null +++ b/scripts/load-test.sh @@ -0,0 +1,89 @@ +#!/usr/bin/env bash +set -e + +BASE_URL="${BASE_URL:-http://127.0.0.1:3004}" +REQUESTS="${REQUESTS:-1000}" +CONCURRENCY="${CONCURRENCY:-50}" + +RED='\033[0;31m' +GREEN='\033[0;32m' +YELLOW='\033[1;33m' +BLUE='\033[0;34m' +NC='\033[0m' + +echo -e "${BLUE}═══════════════════════════════════════════${NC}" +echo -e "${BLUE} Load Test — $BASE_URL${NC}" +echo -e "${BLUE} $REQUESTS requests, $CONCURRENCY concurrent${NC}" +echo -e "${BLUE}═══════════════════════════════════════════${NC}" +echo "" + +run_endpoint() { + local name="$1" + local path="$2" + local method="${3:-GET}" + local body="${4:-}" + + echo -e "${YELLOW}▶ $name${NC} ($method $path)" + local out + if [ "$method" = "POST" ] && [ -n "$body" ]; then + local tmpfile=$(mktemp) + echo "$body" > "$tmpfile" + out=$(ab -q -n "$REQUESTS" -c "$CONCURRENCY" -T "application/json" -p "$tmpfile" "$BASE_URL$path" 2>&1) + rm -f "$tmpfile" + else + out=$(ab -q -n "$REQUESTS" -c "$CONCURRENCY" "$BASE_URL$path" 2>&1) + fi + + local rps=$(echo "$out" | awk '/Requests per second/ {print $4}') + local p50=$(echo "$out" | awk '/^ 50%/ {print $2}') + local p95=$(echo "$out" | awk '/^ 95%/ {print $2}') + local p99=$(echo "$out" | awk '/^ 99%/ {print $2}') + local failed=$(echo "$out" | awk '/^Failed requests/ {print $3}') + local non200=$(echo "$out" | awk '/^Non-2xx responses/ {print $3}') + + local status_text="OK" + local status_color="$GREEN" + if [ -n "$non200" ] && [ "$non200" -gt 0 ]; then + status_text="FAIL (non-2xx: $non200)" + status_color="$RED" + elif [ -n "$failed" ] && [ "$failed" -gt "$((REQUESTS / 20))" ]; then + status_text="FAIL (failed: $failed)" + status_color="$RED" + fi + + echo -e " ${status_color}${status_text}${NC} ${rps:-?} rps | p50=${p50:-?}ms p95=${p95:-?}ms p99=${p99:-?}ms" + echo "" +} + +if ! curl -s -o /dev/null -m 3 "$BASE_URL/api/health"; then + echo -e "${RED}✗ Backend $BASE_URL не отвечает на /api/health${NC}" + exit 1 +fi +echo -e "${GREEN}✓ Backend жив${NC}" +echo "" + +run_endpoint "health" "/api/health" +run_endpoint "landing page data" "/api/web/pages/landing" +run_endpoint "tariffs public" "/api/tariffs/public" +run_endpoint "site config" "/api/site-config" +run_endpoint "flow default" "/api/flows/default" + +echo -e "${YELLOW}▶ Rate limit stress test (register endpoint, same IP)${NC}" +hit_counts=$(for i in $(seq 1 20); do + curl -s -o /dev/null -w "%{http_code}\n" -X POST "$BASE_URL/api/auth/register" \ + -H "Content-Type: application/json" \ + -d '{"email":"loadtest@example.com","password":"testpass123"}' 2>/dev/null +done | sort | uniq -c | sort -rn) +echo "$hit_counts" | sed 's/^/ /' +echo "" + +if echo "$hit_counts" | grep -q "429"; then + echo -e " ${GREEN}✓ Rate limit сработал (получено 429)${NC}" +else + echo -e " ${YELLOW}⚠ Rate limit не сработал (нет 429) — проверь Redis/конфиг${NC}" +fi + +echo "" +echo -e "${BLUE}═══════════════════════════════════════════${NC}" +echo -e "${BLUE} Load test завершён${NC}" +echo -e "${BLUE}═══════════════════════════════════════════${NC}" diff --git a/scripts/pre-commit-secrets.sh b/scripts/pre-commit-secrets.sh new file mode 100755 index 00000000..eab97d62 --- /dev/null +++ b/scripts/pre-commit-secrets.sh @@ -0,0 +1,82 @@ +#!/usr/bin/env bash +set -e + +RED='\033[0;31m' +YELLOW='\033[1;33m' +NC='\033[0m' + +FORBIDDEN_FILES=( + 'config\.py$' + '\.env$' + '\.env\.local$' + '\.env\.production$' + '\.env\.prod$' + 'handlers/texts\.py$' + 'alembic\.ini$' + '\.license_state$' +) + +SECRET_PATTERNS=( + '[0-9]{9,}:[A-Za-z0-9_-]{35}' + '-----BEGIN [A-Z ]+PRIVATE KEY-----' + 'AKIA[0-9A-Z]{16}' + 'sk-[A-Za-z0-9]{20,}' + 'sk-ant-[A-Za-z0-9_-]{20,}' + 'ghp_[A-Za-z0-9]{36}' + 'gho_[A-Za-z0-9]{36}' + 'postgres(ql)?(\+asyncpg)?://[^:]+:[^@]{3,}@[^/]+' + 'mysql://[^:]+:[^@]{3,}@[^/]+' + 'redis://[^:]+:[^@]{3,}@[^/]+' +) + +staged_files=$(git diff --cached --name-only --diff-filter=ACM) + +if [ -z "$staged_files" ]; then + exit 0 +fi + +violations=() + +while IFS= read -r file; do + [ -z "$file" ] && continue + + for forbidden in "${FORBIDDEN_FILES[@]}"; do + if echo "$file" | grep -qE "$forbidden"; then + violations+=("FORBIDDEN FILE: $file") + break + fi + done + + if [ -f "$file" ]; then + case "$file" in + *.png|*.jpg|*.jpeg|*.gif|*.webp|*.svg|*.ico|*.woff|*.woff2|*.ttf|*.eot|*.zip|*.tar|*.gz|*.pdf|*.mp4|*.webm) + continue + ;; + esac + + for pattern in "${SECRET_PATTERNS[@]}"; do + if git show ":$file" 2>/dev/null | grep -E "$pattern" > /dev/null 2>&1; then + matched=$(git show ":$file" 2>/dev/null | grep -nE "$pattern" | head -1) + violations+=("SECRET in $file: $(echo "$matched" | cut -c1-100)") + fi + done + fi +done <<< "$staged_files" + +if [ ${#violations[@]} -gt 0 ]; then + echo -e "${RED}════════════════════════════════════════════════${NC}" + echo -e "${RED}║ COMMIT REJECTED: обнаружены секреты/forbidden files ║${NC}" + echo -e "${RED}════════════════════════════════════════════════${NC}" + for v in "${violations[@]}"; do + echo -e " ${RED}✗${NC} $v" + done + echo "" + echo -e "${YELLOW}Что делать:${NC}" + echo -e " 1. Проверь staged файлы: ${YELLOW}git diff --cached${NC}" + echo -e " 2. Убери чувствительные данные или используй .env (gitignored)" + echo -e " 3. Если false positive и нужно форсировать коммит:" + echo -e " ${YELLOW}git commit --no-verify${NC}" + exit 1 +fi + +exit 0 diff --git a/services/__init__.py b/services/__init__.py new file mode 100644 index 00000000..8b137891 --- /dev/null +++ b/services/__init__.py @@ -0,0 +1 @@ + diff --git a/services/addons.py b/services/addons.py new file mode 100644 index 00000000..47e61eb3 --- /dev/null +++ b/services/addons.py @@ -0,0 +1,265 @@ +from __future__ import annotations + +from dataclasses import dataclass +from math import ceil +from typing import TYPE_CHECKING, Any + +from core.bootstrap import TARIFFS_CONFIG +from core.settings.tariffs_config import normalize_tariff_config +from database import get_balance, get_key_details, get_tariff_by_id, save_key_config_with_mode +from database.coupons import mark_coupon_used +from database.users import update_balance +from logger import logger + +from .errors import InsufficientFundsError, NotFoundError, ValidationError +from .keys import resolve_cluster_name + +if TYPE_CHECKING: + from sqlalchemy.ext.asyncio import AsyncSession + + +def get_pack_flags() -> tuple[bool, bool, str]: + """Определяет доступные опции pack mode из конфига.""" + mode = TARIFFS_CONFIG.get("KEY_ADDONS_PACK_MODE") or "" + if not mode: + return False, False, "" + if mode == "traffic": + return False, True, mode + if mode == "devices": + return True, False, mode + if mode == "all": + return True, True, mode + return False, False, mode + + +def get_override_value(overrides: Any, key: int) -> Any: + if not isinstance(overrides, dict): + return None + if key in overrides: + return overrides.get(key) + return overrides.get(str(key)) + + +def calc_pack_devices_price_rub(tariff: dict[str, Any], pack_devices: int | None) -> int: + if pack_devices is None: + return 0 + pack_devices = int(pack_devices) + overrides = tariff.get("device_overrides") or {} + if pack_devices == 0: + override = get_override_value(overrides, 0) + return int(ceil(float(override))) if override is not None else 0 + if pack_devices < 0: + return 0 + override = get_override_value(overrides, pack_devices) + if override is not None: + return int(ceil(float(override))) + step_price = int(tariff.get("device_step_rub") or 0) + return int(ceil(pack_devices * step_price)) + + +def calc_pack_traffic_price_rub(tariff: dict[str, Any], pack_traffic_gb: int | None) -> int: + if pack_traffic_gb is None: + return 0 + pack_traffic_gb = int(pack_traffic_gb) + overrides = tariff.get("traffic_overrides") or {} + if pack_traffic_gb == 0: + override = get_override_value(overrides, 0) + return int(ceil(float(override))) if override is not None else 0 + if pack_traffic_gb < 0: + return 0 + override = get_override_value(overrides, pack_traffic_gb) + if override is not None: + return int(ceil(float(override))) + step_price = int(tariff.get("traffic_step_rub") or 0) + return int(ceil(pack_traffic_gb * step_price)) + + +def calc_pack_full_price_rub( + tariff: dict[str, Any], + has_device_option: bool, + has_traffic_option: bool, + selected_devices: int | None, + selected_traffic_gb: int | None, +) -> int: + total = 0 + if has_device_option: + total += calc_pack_devices_price_rub(tariff, selected_devices) + if has_traffic_option: + total += calc_pack_traffic_price_rub(tariff, selected_traffic_gb) + return int(total) + + +@dataclass +class AddonsPreviewResult: + """Результат предпросмотра аддонов.""" + total_price_rub: int + extra_price_rub: int + balance: float + required_amount: int + payment_required: bool + current_device_limit: int | None + current_traffic_gb: int | None + selected_device_limit: int | None + selected_traffic_gb: int | None + has_device_option: bool + has_traffic_option: bool + + +@dataclass +class AddonsApplyResult: + """Результат применения аддонов.""" + ok: bool + client_id: str + tariff_id: int + total_price_rub: int + extra_price_rub: int + charged_rub: int + balance_rub: float + + +async def preview_addons( + session: AsyncSession, + billing_user_id: int, + client_id: str, + key_email: str, + tariff_id: int, + selected_device_limit: int | None = None, + selected_traffic_gb: int | None = None, +) -> AddonsPreviewResult: + """Считает стоимость аддонов без применения. + + Raises: NotFoundError, ValidationError + """ + tariff = await get_tariff_by_id(session, tariff_id) + if not tariff: + raise NotFoundError("Тариф не найден") + + cfg = normalize_tariff_config(tariff) + has_device, has_traffic, _ = get_pack_flags() + + key_details = await get_key_details(session, key_email) + if not key_details: + raise NotFoundError("Подписка не найдена") + + current_device = key_details.get("current_device_limit") + current_traffic = key_details.get("current_traffic_limit") + + total_price = calc_pack_full_price_rub( + tariff, has_device, has_traffic, + selected_device_limit, selected_traffic_gb, + ) + + current_price = calc_pack_full_price_rub( + tariff, has_device, has_traffic, + current_device, current_traffic, + ) + extra_price = max(0, total_price - current_price) + + balance = float(await get_balance(session, billing_user_id)) + required = max(0, int(ceil(float(extra_price) - balance))) + + return AddonsPreviewResult( + total_price_rub=total_price, + extra_price_rub=extra_price, + balance=balance, + required_amount=required, + payment_required=required > 0, + current_device_limit=current_device, + current_traffic_gb=current_traffic, + selected_device_limit=selected_device_limit, + selected_traffic_gb=selected_traffic_gb, + has_device_option=has_device, + has_traffic_option=has_traffic, + ) + + +async def apply_addons( + session: AsyncSession, + billing_user_id: int, + client_id: str, + key_email: str, + key_server_id: str, + tariff_id: int, + extra_price_rub: int, + selected_device_limit: int | None = None, + selected_traffic_gb: int | None = None, + coupon_id: int | None = None, +) -> AddonsApplyResult: + """Применяет аддоны к ключу: обновляет лимиты на кластере и в БД. + + Raises: NotFoundError, InsufficientFundsError + """ + tariff = await get_tariff_by_id(session, tariff_id) + if not tariff: + raise NotFoundError("Тариф не найден") + + key_details = await get_key_details(session, key_email) + if not key_details: + raise NotFoundError("Подписка не найдена") + + from services.tariffs.tariff_display import GB, get_effective_limits_for_key + from services.operations import renew_key_in_cluster + from .keys import resolve_cluster_name + dev_eff, trf_bytes = await get_effective_limits_for_key( + session=session, + tariff_id=tariff_id, + selected_device_limit=selected_device_limit, + selected_traffic_gb=selected_traffic_gb if selected_traffic_gb is not None else 0, + ) + total_gb = int(trf_bytes / GB) if trf_bytes else 0 + + cluster_id = await resolve_cluster_name(session, key_server_id) + if not cluster_id: + raise NotFoundError(f"Кластер для {key_server_id} не найден") + + expiry_ms = int(key_details.get("expiry_time") or 0) + + await renew_key_in_cluster( + cluster_id=cluster_id, + email=key_email, + client_id=client_id, + new_expiry_time=expiry_ms, + total_gb=total_gb, + session=session, + hwid_device_limit=dev_eff, + reset_traffic=False, + target_subgroup=tariff.get("subgroup_title"), + old_subgroup=tariff.get("subgroup_title"), + plan=tariff_id, + ) + + if extra_price_rub > 0: + await update_balance(session, billing_user_id, -extra_price_rub) + + cfg = normalize_tariff_config(tariff) + has_device, has_traffic, _ = get_pack_flags() + total_price = calc_pack_full_price_rub( + tariff, has_device, has_traffic, + selected_device_limit, selected_traffic_gb, + ) + + await save_key_config_with_mode( + session=session, + email=key_email, + selected_devices=selected_device_limit, + selected_traffic_gb=selected_traffic_gb, + total_price=total_price, + has_device_choice=has_device, + has_traffic_choice=has_traffic, + config_mode="addons", + ) + + if coupon_id is not None: + await mark_coupon_used(session, coupon_id, billing_user_id) + + new_balance = float(await get_balance(session, billing_user_id)) + + return AddonsApplyResult( + ok=True, + client_id=client_id, + tariff_id=tariff_id, + total_price_rub=total_price, + extra_price_rub=extra_price_rub, + charged_rub=extra_price_rub, + balance_rub=new_balance, + ) diff --git a/services/clusters.py b/services/clusters.py new file mode 100644 index 00000000..9ea6b4f2 --- /dev/null +++ b/services/clusters.py @@ -0,0 +1,229 @@ +from __future__ import annotations + +import asyncio +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, Callable, Coroutine + +from config import ADMIN_USERNAME, ADMIN_PASSWORD, REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD +from database.keys import count_keys_by_server_id, get_all_key_server_ids +from database.servers import ( + filter_cluster_by_subgroup, + filter_cluster_by_tariff, + get_panel_type_for_server, + get_panel_types_for_cluster, + get_servers, +) +from hooks.processors import process_cluster_balancer, process_cluster_override +from logger import logger + +from .errors import NotFoundError, ValidationError + +if TYPE_CHECKING: + from sqlalchemy.ext.asyncio import AsyncSession + +ALLOWED_GROUP_CODES = ["trial", "discounts", "discounts_max", "gifts"] + + +@dataclass +class ClusterSelection: + """Результат выбора кластера.""" + cluster_name: str + load: int + available_servers: list[dict[str, Any]] + + +@dataclass +class ServerAvailability: + """Результат проверки доступности сервера.""" + server_name: str + available: bool + panel_type: str + + +async def check_server_key_limit( + server_info: dict[str, Any], + session: AsyncSession, + on_capacity_warning: Callable[..., Coroutine] | None = None, +) -> bool: + """Проверяет, не превышен ли лимит ключей на сервере. + + on_capacity_warning — опциональный callback при >=90% заполненности + (бот передаёт функцию уведомления админа, API может логировать). + """ + server_name = server_info.get("server_name") + cluster_name = server_info.get("cluster_name") + max_keys = server_info.get("max_keys") + + if not max_keys: + return True + + identifier = cluster_name if cluster_name else server_name + total_keys = await count_keys_by_server_id(session, identifier) + + if total_keys >= max_keys: + logger.warning(f"[Key Limit] Сервер {server_name} достиг лимита: {total_keys}/{max_keys}") + return False + + usage_percent = total_keys / max_keys + if usage_percent >= 0.9 and on_capacity_warning: + try: + await on_capacity_warning(server_name, total_keys, max_keys) + except Exception: + pass + + return True + + +async def check_server_availability(server_info: dict[str, Any], session: AsyncSession) -> ServerAvailability: + """Проверяет доступность сервера (enabled + лимит + API ping).""" + server_name = server_info.get("server_name", "unknown") + panel_type = (server_info.get("panel_type") or "3x-ui").lower() + enabled = server_info.get("enabled", True) + + if not enabled: + return ServerAvailability(server_name=server_name, available=False, panel_type=panel_type) + + max_keys = server_info.get("max_keys") + if max_keys is not None: + try: + total = await count_keys_by_server_id(session, server_name) + if total >= max_keys: + return ServerAvailability(server_name=server_name, available=False, panel_type=panel_type) + except Exception: + return ServerAvailability(server_name=server_name, available=False, panel_type=panel_type) + + try: + if panel_type == "remnawave": + from panels.remnawave import RemnawaveAPI + remna = RemnawaveAPI(server_info["api_url"]) + await asyncio.wait_for(remna.login(REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD), timeout=5.0) + else: + from panels._3xui import AsyncApi + xui = AsyncApi( + server_info["api_url"], + username=ADMIN_USERNAME, + password=ADMIN_PASSWORD, + logger=logger, + ) + await asyncio.wait_for(xui.login(), timeout=5.0) + return ServerAvailability(server_name=server_name, available=True, panel_type=panel_type) + except Exception: + logger.warning(f"[Ping] Сервер {server_name} недоступен") + return ServerAvailability(server_name=server_name, available=False, panel_type=panel_type) + + +async def select_cluster( + session: AsyncSession, + on_capacity_warning: Callable[..., Coroutine] | None = None, +) -> ClusterSelection: + """Выбирает наименее нагруженный кластер. + + Raises: ValidationError если нет доступных кластеров. + """ + forced = await process_cluster_override(session=session) + if isinstance(forced, str) and forced.strip(): + servers = await get_servers(session) + cluster_servers = servers.get(forced.strip(), []) + enabled = [s for s in cluster_servers if s.get("enabled", True)] + if enabled: + return ClusterSelection(cluster_name=forced.strip(), load=0, available_servers=enabled) + + servers = await get_servers(session) + server_to_cluster: dict[str, str] = {} + cluster_loads: dict[str, int] = {} + + for cluster_name, cluster_servers in servers.items(): + cluster_loads[cluster_name] = 0 + for server in cluster_servers: + server_to_cluster[server["server_name"]] = cluster_name + + key_server_ids = await get_all_key_server_ids(session) + for sid in key_server_ids: + cid = server_to_cluster.get(sid, sid) + if cid in cluster_loads: + cluster_loads[cid] += 1 + + available: dict[str, int] = {} + cluster_available_servers: dict[str, list] = {} + + for cluster_name, cluster_servers in servers.items(): + enabled = [s for s in cluster_servers if s.get("enabled", True)] + if not enabled: + continue + ok_servers = [] + for s in enabled: + if await check_server_key_limit(s, session, on_capacity_warning): + ok_servers.append(s) + if ok_servers: + available[cluster_name] = cluster_loads[cluster_name] + cluster_available_servers[cluster_name] = ok_servers + + filtered = await process_cluster_balancer(available_clusters=available, session=session) + if filtered: + available = {k: v for k, v in available.items() if k in filtered} + + if not available: + raise ValidationError("Сервисы временно недоступны. Попробуйте позже.") + + best = min(available, key=lambda k: (available[k], k)) + logger.info(f"Выбран кластер: {best} (загрузка: {available[best]})") + + return ClusterSelection( + cluster_name=best, + load=available[best], + available_servers=cluster_available_servers.get(best, []), + ) + + +async def filter_servers_for_key( + session: AsyncSession, + cluster_servers: list[dict[str, Any]], + cluster_id: str, + tariff_id: int | None = None, + subgroup_title: str | None = None, + special_group: str | None = None, +) -> list[dict[str, Any]]: + """Фильтрует серверы кластера по тарифу, подгруппе и special group. + + Возвращает отфильтрованный список серверов. + """ + enabled = [s for s in cluster_servers if s.get("enabled", True)] + + if tariff_id: + filtered = await filter_cluster_by_tariff(session, enabled, tariff_id, cluster_id) + if filtered: + enabled = filtered + + if subgroup_title: + filtered = await filter_cluster_by_subgroup( + session, enabled, subgroup_title, cluster_id, tariff_id=tariff_id, + ) + if filtered: + enabled = filtered + + if special_group and special_group in ALLOWED_GROUP_CODES: + bound = [s for s in enabled if special_group in (s.get("special_groups") or [])] + if bound: + enabled = bound + + return enabled + + +async def is_full_remnawave_cluster(cluster_id: str, session: AsyncSession) -> bool: + """Проверяет, состоит ли кластер полностью из Remnawave-серверов.""" + panel_types = await get_panel_types_for_cluster(session, cluster_id) + if panel_types: + return all(pt.lower() == "remnawave" for pt in panel_types) + pt = await get_panel_type_for_server(session, cluster_id) + return bool(pt and pt.lower() == "remnawave") + + +def resolve_special_group(tariff: dict[str, Any] | None, is_trial: bool = False) -> str | None: + """Определяет special group для фильтрации серверов.""" + if is_trial: + return "trial" + if tariff: + gc = (tariff.get("group_code") or "").lower() + if gc in ALLOWED_GROUP_CODES: + return gc + return None diff --git a/services/coupons.py b/services/coupons.py new file mode 100644 index 00000000..5887d2a5 --- /dev/null +++ b/services/coupons.py @@ -0,0 +1,141 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +from database.coupons import ( + apply_percent_coupon, + check_coupon_usage, + create_coupon_usage, + get_coupon_by_code_ci, + update_coupon_usage_count, +) +from database.keys import count_active_keys_for_user +from database.models import Coupon +from database.payments import add_payment, count_successful_payments +from database.users import get_balance, update_balance + +from .errors import LimitExceededError, NotFoundError, ValidationError + +if TYPE_CHECKING: + from sqlalchemy.ext.asyncio import AsyncSession + + +async def resolve_percent_coupon( + session: AsyncSession, + billing_user_id: int, + base_price_rub: int, + coupon_code: str | None, +) -> tuple[int, int, int | None, str | None]: + """Применяет процентный купон к цене. + + Returns: (final_price, discount, coupon_id, coupon_code) + Raises: NotFoundError, LimitExceededError, ValidationError + """ + normalized = (coupon_code or "").strip() + if not normalized: + return int(base_price_rub), 0, None, None + + coupon = await get_coupon_by_code_ci(session, normalized) + if coupon is None: + raise NotFoundError("Купон не найден") + + _check_coupon_limits(coupon) + + if await check_coupon_usage(session, int(coupon.id), int(billing_user_id)): + raise LimitExceededError("Вы уже использовали этот купон") + + percent = getattr(coupon, "percent", None) + if percent is None: + raise ValidationError("Поддерживаются только процентные купоны") + + if bool(getattr(coupon, "new_users_only", False)): + await _check_new_user(session, billing_user_id) + + discounted_price, discount_rub = apply_percent_coupon(int(base_price_rub), coupon) + if int(discount_rub) <= 0: + raise ValidationError("Купон не применим к текущей сумме") + + return int(discounted_price), int(discount_rub), int(coupon.id), str(coupon.code or normalized) + + +class CouponApplyResult: + __slots__ = ("coupon_code", "amount", "balance") + + def __init__(self, coupon_code: str, amount: int, balance: float): + self.coupon_code = coupon_code + self.amount = amount + self.balance = balance + + +async def apply_fixed_coupon( + session: AsyncSession, + user_id: int, + tg_id: int | None, + code: str, +) -> CouponApplyResult: + """Активирует купон с фиксированной суммой — зачисляет на баланс. + + Raises: NotFoundError, LimitExceededError, ValidationError + """ + normalized = code.strip() + if not normalized: + raise ValidationError("Введите код купона") + + coupon = await get_coupon_by_code_ci(session, normalized) + if coupon is None: + raise NotFoundError("Купон не найден") + + _check_coupon_limits(coupon) + + if await check_coupon_usage(session, int(coupon.id), int(user_id)): + raise LimitExceededError("Вы уже использовали этот купон") + + percent = getattr(coupon, "percent", None) + if percent is not None: + raise ValidationError("Этот купон применяется при оплате тарифа") + + days = int(getattr(coupon, "days", 0) or 0) + if days > 0: + raise ValidationError("Купон на продление применяйте через Telegram-бота") + + amount = int(getattr(coupon, "amount", 0) or 0) + if amount <= 0: + raise ValidationError("Купон недействителен") + + if bool(getattr(coupon, "new_users_only", False)): + await _check_new_user(session, user_id) + + await update_balance(session, int(user_id), float(amount)) + await create_coupon_usage(session, int(coupon.id), int(user_id), tg_id) + await update_coupon_usage_count(session, int(coupon.id)) + await add_payment( + session=session, + legacy_user_ref=int(user_id), + amount=float(amount), + payment_system="coupon", + status="success", + currency="RUB", + ) + + balance = float(await get_balance(session, int(user_id))) + + return CouponApplyResult( + coupon_code=str(coupon.code or normalized), + amount=amount, + balance=balance, + ) + + +def _check_coupon_limits(coupon: Coupon) -> None: + usage_limit = int(getattr(coupon, "usage_limit", 0) or 0) + usage_count = int(getattr(coupon, "usage_count", 0) or 0) + is_used = bool(getattr(coupon, "is_used", False)) + if usage_limit > 0 and (is_used or usage_count >= usage_limit): + raise LimitExceededError("Лимит активаций купона исчерпан") + + +async def _check_new_user(session: AsyncSession, user_id: int) -> None: + payments_count = await count_successful_payments(session, int(user_id)) + keys_count = await count_active_keys_for_user(session, int(user_id)) + if payments_count > 0 or keys_count > 0: + raise ValidationError("Этот купон доступен только для новых пользователей") diff --git a/services/errors.py b/services/errors.py new file mode 100644 index 00000000..90746d69 --- /dev/null +++ b/services/errors.py @@ -0,0 +1,44 @@ +class ServiceError(Exception): + """Базовая ошибка бизнес-логики.""" + + def __init__(self, message: str, code: str = "error"): + self.message = message + self.code = code + super().__init__(message) + + +class NotFoundError(ServiceError): + """Ресурс не найден.""" + + def __init__(self, message: str = "Не найдено"): + super().__init__(message, code="not_found") + + +class ValidationError(ServiceError): + """Невалидные входные данные или бизнес-правило нарушено.""" + + def __init__(self, message: str = "Ошибка валидации"): + super().__init__(message, code="validation_error") + + +class InsufficientFundsError(ServiceError): + """Недостаточно средств на балансе.""" + + def __init__(self, message: str = "Недостаточно средств", required: float = 0, balance: float = 0): + self.required = required + self.balance = balance + super().__init__(message, code="insufficient_funds") + + +class LimitExceededError(ServiceError): + """Превышен лимит (купоны, использования, и т.д.).""" + + def __init__(self, message: str = "Лимит исчерпан"): + super().__init__(message, code="limit_exceeded") + + +class ForbiddenError(ServiceError): + """Операция запрещена настройками или правами.""" + + def __init__(self, message: str = "Операция запрещена"): + super().__init__(message, code="forbidden") diff --git a/services/formatting.py b/services/formatting.py new file mode 100644 index 00000000..92dae461 --- /dev/null +++ b/services/formatting.py @@ -0,0 +1,40 @@ +from __future__ import annotations + +from config import USERNAME_BOT + + +def get_plural_form(num: int, form1: str, form2: str, form3: str) -> str: + n = abs(num) % 100 + if 10 < n < 20: + return form3 + return {1: form1, 2: form2, 3: form2, 4: form2}.get(n % 10, form3) + + +def format_months(months: int) -> str: + if months <= 0: + return "0 месяцев" + return f"{months} {get_plural_form(months, 'месяц', 'месяца', 'месяцев')}" + + +def format_days(days: int) -> str: + if days <= 0: + return "0 дней" + return f"{days} {get_plural_form(days, 'день', 'дня', 'дней')}" + + +def format_duration_days(days: int) -> str: + return format_months(days // 30) if days % 30 == 0 else format_days(days) + + +def get_telegram_gift_link(gift_id: str) -> str: + return f"https://t.me/{USERNAME_BOT}?start=gift_{gift_id}" + + +def get_gift_link(user_id: int, gift_id: str) -> str: + return get_telegram_gift_link(gift_id) + + +def get_site_gift_link(gift_id: str) -> str: + from core.settings.web_config import get_site_url + site_url = get_site_url() + return f"{site_url}/gift/{gift_id}" diff --git a/services/gifts.py b/services/gifts.py new file mode 100644 index 00000000..e7e40416 --- /dev/null +++ b/services/gifts.py @@ -0,0 +1,239 @@ +from __future__ import annotations + +import uuid +from dataclasses import dataclass +from datetime import datetime, timedelta +from typing import TYPE_CHECKING + +from pytz import timezone + +from database import ( + add_referral, + get_balance, + get_referral_by_referred_id, + store_gift_link, + update_balance, + update_trial, +) +from database.access.resolution import resolve_user_optional +from database.gifts import ( + count_gift_usages, + get_gift_locked, + get_gift_usage, + mark_gift_fully_redeemed, + record_gift_usage, +) +from database.tariffs import get_tariff_by_id +from logger import logger +from services.formatting import format_days, format_months, get_gift_link, get_site_gift_link + +from .errors import InsufficientFundsError, NotFoundError, ValidationError + +if TYPE_CHECKING: + from sqlalchemy.ext.asyncio import AsyncSession + + +@dataclass +class GiftRedeemResult: + message: str + gift_id: str + tariff_id: int + duration_days: int + + +@dataclass +class GiftCreateResult: + gift_id: str + gift_link: str + site_gift_link: str + tariff_name: str + duration_days: int + duration_text: str + expiry_time: datetime + price_charged: int + + +def normalize_gift_code(raw: str) -> str: + s = raw.strip() + if not s: + return "" + if "/gift/" in s: + from_url = s.split("/gift/")[-1] + s = from_url.split("?")[0].split("#")[0].strip() + if s.startswith("gift_"): + s = s[5:].strip() + for prefix in ("start=gift_", "start="): + idx = s.find(prefix) + if idx >= 0: + token = s[idx + len(prefix):] + return token.split("&")[0].split("?")[0].split("#")[0].strip() + return s.split("?")[0].split("#")[0].strip() + + +def _format_duration(days: int) -> str: + return format_months(days // 30) if days % 30 == 0 else format_days(days) + + +async def redeem_gift( + session: AsyncSession, + gift_code: str, + billing_user_ref: int, +) -> GiftRedeemResult: + """Активирует подарок для пользователя. + + Создаёт ключ, записывает usage, привязывает реферала. + Raises: ValidationError, NotFoundError + """ + code = normalize_gift_code(gift_code) + if not code: + raise ValidationError("Укажите ссылку или код подарка") + + gift_info = await get_gift_locked(session, code) + if not gift_info: + raise NotFoundError("Подарок не найден или срок ссылки истёк") + + wu = await resolve_user_optional(session, billing_user_ref) + if wu is None: + raise NotFoundError("Пользователь не найден") + + if gift_info.expiry_time and gift_info.expiry_time < datetime.utcnow(): + raise ValidationError("Срок действия подарка истёк") + + if gift_info.sender_user_id == wu.id: + raise ValidationError("Нельзя активировать подарок, который вы создали сами") + + if await get_gift_usage(session, gift_info.gift_id, wu.id) is not None: + raise ValidationError("Вы уже активировали этот подарок") + + if gift_info.recipient_user_id and not gift_info.is_unlimited: + raise ValidationError("Этот подарок уже был активирован другим пользователем") + + if not gift_info.is_unlimited: + usage_count = await count_gift_usages(session, gift_info.gift_id) + if ( + gift_info.is_used + or (gift_info.max_usages is not None and usage_count >= gift_info.max_usages) + or (gift_info.max_usages is None and gift_info.recipient_user_id is not None) + ): + raise ValidationError("Этот подарок уже был использован") + + existing_referral = await get_referral_by_referred_id(session, wu.id) + if not existing_referral and gift_info.sender_user_id: + await add_referral(session, wu.id, gift_info.sender_user_id) + + await update_trial(session, wu.id, 1) + + tariff = await get_tariff_by_id(session, int(gift_info.tariff_id)) if gift_info.tariff_id else None + if not tariff: + raise NotFoundError("Тариф, связанный с подарком, не найден") + + from services.keys import create_vpn_key_headless + + now_moscow = datetime.now(timezone("Europe/Moscow")).replace(tzinfo=None) + duration_days = int(tariff["duration_days"] or 0) + expiry_time = now_moscow + timedelta(days=duration_days) + + selected_device_limit = getattr(gift_info, "selected_device_limit", None) + selected_traffic_gb = getattr(gift_info, "selected_traffic_gb", None) + selected_price_rub = getattr(gift_info, "selected_price_rub", None) + + await create_vpn_key_headless( + session=session, + tg_id=wu.id, + expiry_time=expiry_time, + plan=int(tariff["id"]), + selected_device_limit=int(selected_device_limit) if selected_device_limit is not None else None, + selected_traffic_gb=int(selected_traffic_gb) if selected_traffic_gb is not None else None, + selected_price_rub=int(selected_price_rub) if selected_price_rub is not None else None, + skip_balance_charge=True, + ) + + await record_gift_usage(session, gift_info.gift_id, wu.id, wu.tg_id) + + if not gift_info.is_unlimited: + usage_count = await count_gift_usages(session, gift_info.gift_id) + if gift_info.max_usages and usage_count >= gift_info.max_usages: + await mark_gift_fully_redeemed(session, code, wu.id, wu.tg_id) + + duration_text = _format_duration(duration_days) + + try: + from database.web_notifications import notify_web + await notify_web( + session, + tg_id=wu.tg_id, + type="gift_received", + template_vars={"name": tariff["name"], "duration": duration_text}, + data={"gift_id": gift_info.gift_id, "tariff_id": int(tariff["id"])}, + ) + except Exception as e: + logger.warning("[Gifts] Ошибка отправки уведомления о подарке: {}", e) + + return GiftRedeemResult( + message=f"Подарок активирован — подписка на {duration_text}", + gift_id=gift_info.gift_id, + tariff_id=int(tariff["id"]), + duration_days=duration_days, + ) + + +async def create_gift( + session: AsyncSession, + sender_user_ref: int, + tariff_id: int, + selected_device_limit: int | None = None, + selected_traffic_gb: int | None = None, + selected_price_rub: int | None = None, +) -> GiftCreateResult: + """Создаёт подарок: списывает баланс, сохраняет в БД. + + Raises: NotFoundError, InsufficientFundsError + """ + tariff = await get_tariff_by_id(session, int(tariff_id)) + if not tariff or tariff.get("group_code") != "gifts": + raise NotFoundError("Тариф не найден") + + price_to_charge = ( + int(selected_price_rub) if selected_price_rub is not None else int(tariff["price_rub"]) + ) + + balance = await get_balance(session, sender_user_ref) + if balance < price_to_charge: + raise InsufficientFundsError( + "Недостаточно средств для создания подарка", + required=price_to_charge, + balance=balance, + ) + + await update_balance(session, sender_user_ref, -price_to_charge) + + duration_days = int(tariff["duration_days"] or 0) + expiry_time = datetime.utcnow() + timedelta(days=duration_days) + gift_id = uuid.uuid4().hex + gift_link = get_gift_link(sender_user_ref, gift_id) + site_gift_link = get_site_gift_link(gift_id) + + await store_gift_link( + session=session, + gift_id=gift_id, + sender_legacy_ref=sender_user_ref, + selected_months=duration_days // 30, + expiry_time=expiry_time, + gift_link=gift_link, + tariff_id=int(tariff["id"]), + max_usages=1, + selected_device_limit=selected_device_limit, + selected_traffic_gb=selected_traffic_gb, + selected_price_rub=price_to_charge, + ) + + return GiftCreateResult( + gift_id=gift_id, + gift_link=gift_link, + site_gift_link=site_gift_link, + tariff_name=tariff["name"], + duration_days=duration_days, + duration_text=_format_duration(duration_days), + expiry_time=expiry_time, + price_charged=price_to_charge, + ) diff --git a/services/keys.py b/services/keys.py new file mode 100644 index 00000000..e7dfd297 --- /dev/null +++ b/services/keys.py @@ -0,0 +1,464 @@ +from __future__ import annotations + +import uuid +from dataclasses import dataclass, field +from datetime import datetime +from math import ceil +from typing import TYPE_CHECKING, Any + +from core.settings.tariffs_config import normalize_tariff_config +from database import ( + get_balance, + get_key_details, + get_tariff_by_id, + reset_key_current_limits_to_selected, + save_key_config_with_mode, + update_balance, + update_key_expiry, + update_trial, +) +from database.access.resolution import resolve_user_optional +from database.coupons import mark_coupon_used +from database.keys import ( + update_key_post_creation_snapshot, + update_key_renewal_snapshot, +) +from database.servers import cluster_name_exists, get_cluster_name_for_server_name +from database.users import get_trial +from logger import logger + +from .errors import NotFoundError, ValidationError + +if TYPE_CHECKING: + from sqlalchemy.ext.asyncio import AsyncSession + + +@dataclass +class RenewalPricing: + """Расчёт цены продления (без фактического продления).""" + + base_price_rub: int + discount_rub: int + final_price_rub: int + coupon_id: int | None + applied_coupon_code: str | None + total_gb: int + balance: float + required_amount: int + payment_required: bool + duration_days: int + selected_device_limit: int | None + selected_traffic_limit: int | None + + +@dataclass +class RenewalResult: + """Результат фактического продления.""" + + ok: bool + client_id: str + tariff_id: int + charged_rub: int + balance_rub: float + new_expiry_time: int + base_price_rub: int = 0 + discount_rub: int = 0 + final_price_rub: int = 0 + applied_coupon_code: str | None = None + + +async def resolve_cluster_name(session: AsyncSession, server_or_cluster: str) -> str | None: + """Определяет имя кластера: если имя — уже кластер, возвращаем его; + иначе ищем сервер с таким именем и отдаём его ``cluster_name``.""" + if await cluster_name_exists(session, server_or_cluster): + return server_or_cluster + return await get_cluster_name_for_server_name(session, server_or_cluster) + + +def normalize_expiry_ms(raw_value: int | float | None) -> int: + """Нормализует таймстамп истечения в миллисекунды. + + Единая реализация — обрабатывает секунды, миллисекунды и микросекунды. + """ + if not raw_value: + return 0 + value = int(raw_value) + if value > 10**13: + value //= 1000 + elif value < 10**10: + value *= 1000 + return value + + +def _resolve_effective_limits( + tariff: dict[str, Any], + selected_device_limit: int | None, + selected_traffic_limit: int | None, +) -> tuple[int | None, int | None]: + """Определяет финальные device/traffic лимиты по тарифу и выбору пользователя.""" + new_tariff_device = tariff.get("device_limit") + new_tariff_traffic = tariff.get("traffic_limit") + + if new_tariff_device is None: + final_device = None + elif selected_device_limit is not None: + final_device = int(selected_device_limit) + else: + final_device = new_tariff_device + + if new_tariff_traffic is None: + final_traffic = None + elif selected_traffic_limit is not None: + final_traffic = int(selected_traffic_limit) + else: + final_traffic = int(new_tariff_traffic) + + return final_device, final_traffic + + +async def calculate_renewal_pricing( + session: AsyncSession, + billing_user_id: int, + key_email: str, + tariff_id: int, + coupon_code: str | None = None, +) -> RenewalPricing: + """Считает цену продления без фактического выполнения. + + Raises: NotFoundError, ValidationError + """ + from services.coupons import resolve_percent_coupon + + tariff = await get_tariff_by_id(session, int(tariff_id)) + if not tariff or not tariff.get("is_active", True): + raise NotFoundError("Тариф не найден") + + duration_days = int(tariff.get("duration_days") or 0) + if duration_days <= 0: + raise ValidationError("Некорректная длительность тарифа") + + key_details = await get_key_details(session, key_email) + if not key_details: + raise NotFoundError("Подписка не найдена") + + selected_device = key_details.get("selected_device_limit") + selected_traffic = key_details.get("selected_traffic_limit") + sel_dev_int = int(selected_device) if selected_device is not None else None + sel_trf_int = int(selected_traffic) if selected_traffic is not None else None + + from services.tariffs import calculate_config_price as calc_price_buy + total_price_rub = int(calc_price_buy( + tariff=tariff, + selected_device_limit=sel_dev_int, + selected_traffic_gb=sel_trf_int, + )) + if total_price_rub <= 0: + raise ValidationError("Некорректная стоимость продления") + + final_price_rub, discount_rub, coupon_id, applied_code = await resolve_percent_coupon( + session=session, + billing_user_id=billing_user_id, + base_price_rub=total_price_rub, + coupon_code=coupon_code, + ) + + from services.tariffs.tariff_display import GB, get_effective_limits_for_key + _, traffic_bytes = await get_effective_limits_for_key( + session=session, + tariff_id=int(tariff_id), + selected_device_limit=sel_dev_int, + selected_traffic_gb=sel_trf_int if sel_trf_int is not None else 0, + ) + total_gb = int(traffic_bytes / GB) if traffic_bytes else 0 + + balance = float(await get_balance(session, billing_user_id)) + required = int(max(0, ceil(float(final_price_rub) - balance))) + + return RenewalPricing( + base_price_rub=total_price_rub, + discount_rub=discount_rub, + final_price_rub=final_price_rub, + coupon_id=coupon_id, + applied_coupon_code=applied_code, + total_gb=total_gb, + balance=balance, + required_amount=required, + payment_required=required > 0, + duration_days=duration_days, + selected_device_limit=sel_dev_int, + selected_traffic_limit=sel_trf_int, + ) + + +async def execute_renewal( + session: AsyncSession, + billing_user_id: int, + client_id: str, + key_email: str, + key_server_id: str, + tariff_id: int, + new_expiry_time: int, + total_gb: int, + cost: float, + selected_device_limit: int | None = None, + selected_traffic_limit: int | None = None, + selected_price_rub: int | None = None, + coupon_id: int | None = None, +) -> RenewalResult: + """Выполняет продление ключа на кластере и обновляет БД. + + Не отправляет сообщений в Telegram — это делает вызывающий код. + Raises: NotFoundError, ValidationError + """ + tariff = await get_tariff_by_id(session, tariff_id) + if not tariff: + raise NotFoundError(f"Тариф с id={tariff_id} не найден") + + key_info = await get_key_details(session, key_email) + if not key_info: + raise NotFoundError(f"Ключ {client_id} не найден в БД") + + final_device, final_traffic = _resolve_effective_limits( + tariff, selected_device_limit, selected_traffic_limit, + ) + + from services.tariffs.tariff_display import GB, get_effective_limits_for_key + from services.operations import renew_key_in_cluster + if tariff.get("configurable"): + sel_trf = int(final_traffic) if final_traffic is not None else None + sel_dev = int(final_device) if final_device is not None else None + device_eff, traffic_bytes_eff = await get_effective_limits_for_key( + session=session, + tariff_id=tariff_id, + selected_device_limit=sel_dev, + selected_traffic_gb=sel_trf, + ) + traffic_gb_eff = int(traffic_bytes_eff / GB) if traffic_bytes_eff else 0 + total_gb = traffic_gb_eff + else: + device_eff = final_device + traffic_gb_eff = int(final_traffic) if final_traffic is not None else 0 + total_gb = traffic_gb_eff + + current_subgroup = None + try: + cur_tariff_id = key_info.get("tariff_id") + if cur_tariff_id: + cur_tariff = await get_tariff_by_id(session, int(cur_tariff_id)) + if cur_tariff: + current_subgroup = cur_tariff.get("subgroup_title") + except Exception as e: + logger.warning("[Keys] Ошибка получения subgroup текущего тарифа: {}", e) + + target_subgroup = tariff.get("subgroup_title") + + cluster_id = await resolve_cluster_name(session, key_server_id) + if not cluster_id: + raise NotFoundError(f"Кластер для {key_server_id} не найден") + + await renew_key_in_cluster( + cluster_id=cluster_id, + email=key_email, + client_id=client_id, + new_expiry_time=new_expiry_time, + total_gb=total_gb, + session=session, + hwid_device_limit=device_eff, + reset_traffic=True, + target_subgroup=target_subgroup, + old_subgroup=current_subgroup, + plan=tariff_id, + ) + + key_row = await get_key_details(session, key_email) + effective_client_id = key_row["client_id"] if key_row else client_id + + await update_key_expiry(session, effective_client_id, new_expiry_time) + + new_dev = tariff.get("device_limit") + new_trf = tariff.get("traffic_limit") + + if tariff.get("configurable"): + await update_key_renewal_snapshot( + session, + key_email, + tariff_id=tariff_id, + apply_limits=False, + ) + else: + await update_key_renewal_snapshot( + session, + key_email, + tariff_id=tariff_id, + selected_device_limit=None if new_dev is None else new_dev, + current_device_limit=None if new_dev is None else final_device, + selected_traffic_limit=None if new_trf is None else new_trf, + current_traffic_limit=None if new_trf is None else final_traffic, + apply_limits=True, + ) + await update_balance(session, billing_user_id, -cost) + + if tariff.get("configurable"): + cfg = normalize_tariff_config(tariff) + raw_device_opts = cfg.get("device_options") or tariff.get("device_options") or [] + raw_traffic_opts = cfg.get("traffic_options_gb") or tariff.get("traffic_options_gb") or [] + has_device = len([v for v in raw_device_opts if _try_int(v) is not None]) > 1 + has_traffic = len([v for v in raw_traffic_opts if _try_int(v) is not None]) > 1 + + await save_key_config_with_mode( + session=session, + email=key_email, + selected_devices=final_device, + selected_traffic_gb=final_traffic, + total_price=int(selected_price_rub or cost), + has_device_choice=has_device, + has_traffic_choice=has_traffic, + config_mode="renewal", + ) + if has_device or has_traffic: + await reset_key_current_limits_to_selected(session, effective_client_id) + + if coupon_id is not None: + await mark_coupon_used(session, coupon_id, billing_user_id) + + new_balance = float(await get_balance(session, billing_user_id)) + + return RenewalResult( + ok=True, + client_id=effective_client_id, + tariff_id=tariff_id, + charged_rub=int(cost), + balance_rub=new_balance, + new_expiry_time=new_expiry_time, + base_price_rub=int(selected_price_rub or cost), + final_price_rub=int(cost), + ) + + +def _try_int(v: Any) -> int | None: + try: + return int(v) + except (TypeError, ValueError): + return None + + +@dataclass +class CreatedVpnKey: + """Результат headless-создания ключа (без UI-ответа).""" + + client_id: str + email: str + cluster_id: str + final_link: str + key_record: dict + price_charged: int + + +async def create_vpn_key_headless( + session: AsyncSession, + tg_id: int, + expiry_time: datetime, + *, + plan: int | None = None, + selected_device_limit: int | None = None, + selected_traffic_gb: int | None = None, + selected_price_rub: int | None = None, + skip_balance_charge: bool = False, + is_trial: bool = False, + forced_cluster: str | None = None, +) -> CreatedVpnKey: + """Создаёт VPN-ключ для пользователя без aiogram/FSM зависимостей. + + Используется из service-слоя (gift redemption, web tariff purchase, webhook + completion) — везде, где нет Message/CallbackQuery. Доменная логика та же, + что и у `handlers.keys.key_mode.key_cluster_mode`, но без построения + клавиатуры и отправки ответа. + + Raises: + NotFoundError: пользователь не найден. + ValidationError: не удалось определить кластер или создать ключ. + """ + from handlers.utils import generate_random_email + from services.clusters import select_cluster + from services.operations import create_key_on_cluster + from services.tariffs.tariff_display import ( + get_effective_limits_for_key, + resolve_price_to_charge, + ) + + owner = await resolve_user_optional(session, tg_id) + if owner is None: + raise NotFoundError(f"Пользователь не найден: {tg_id}") + + key_name = await generate_random_email(session=session) + client_id = str(uuid.uuid4()) + email = key_name.lower() + expiry_timestamp = int(expiry_time.timestamp() * 1000) + + device_limit, traffic_limit_bytes = await get_effective_limits_for_key( + session=session, + tariff_id=plan, + selected_device_limit=selected_device_limit, + selected_traffic_gb=selected_traffic_gb, + ) + if device_limit is None: + device_limit = 0 + if traffic_limit_bytes is None: + traffic_limit_bytes = 0 + + if forced_cluster: + cluster_id = forced_cluster + else: + cluster_result = await select_cluster(session) + cluster_id = cluster_result.cluster_name + + if selected_price_rub is not None: + price_to_charge = int(selected_price_rub) + else: + resolved = await resolve_price_to_charge(session, {}) + price_to_charge = int(resolved or 0) + + await create_key_on_cluster( + cluster_id=cluster_id, + tg_id=tg_id, + client_id=client_id, + email=email, + expiry_timestamp=expiry_timestamp, + plan=plan, + session=session, + hwid_limit=device_limit, + traffic_limit_bytes=traffic_limit_bytes, + is_trial=is_trial, + ) + logger.info(f"[Key Creation] Ключ создан на кластере {cluster_id} для пользователя {tg_id}") + + await update_key_post_creation_snapshot( + session, + user_id=owner.id, + email=email, + selected_device_limit=selected_device_limit, + selected_traffic_limit=selected_traffic_gb, + selected_price_rub=price_to_charge, + ) + + key_record = await get_key_details(session, email) + if not key_record: + raise ValidationError(f"Ключ не найден после создания: {email}") + final_link = key_record.get("link", "") or "" + + if is_trial: + trial_status = await get_trial(session, tg_id) + if trial_status in (0, -1): + await update_trial(session, tg_id, 1) + + if price_to_charge and not skip_balance_charge: + await update_balance(session, tg_id, -int(price_to_charge)) + + return CreatedVpnKey( + client_id=client_id, + email=email, + cluster_id=cluster_id, + final_link=final_link, + key_record=key_record, + price_charged=int(price_to_charge or 0), + ) diff --git a/handlers/keys/operations/__init__.py b/services/operations/__init__.py similarity index 100% rename from handlers/keys/operations/__init__.py rename to services/operations/__init__.py diff --git a/handlers/keys/operations/aggregated_links.py b/services/operations/aggregated_links.py similarity index 100% rename from handlers/keys/operations/aggregated_links.py rename to services/operations/aggregated_links.py diff --git a/handlers/keys/operations/creation.py b/services/operations/creation.py similarity index 94% rename from handlers/keys/operations/creation.py rename to services/operations/creation.py index 8324a5dc..6a1e8069 100644 --- a/handlers/keys/operations/creation.py +++ b/services/operations/creation.py @@ -2,21 +2,17 @@ import asyncio from datetime import datetime -from sqlalchemy import update from sqlalchemy.ext.asyncio import AsyncSession from config import PUBLIC_LINK, REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD, SUPERNODE from database import filter_cluster_by_subgroup, filter_cluster_by_tariff, get_servers, get_tariff_by_id, store_key -from database.models import User -from handlers.utils import ALLOWED_GROUP_CODES, check_server_key_limit +from database.users import mark_trial_started_if_eligible from hooks.processors import process_extract_cryptolink_from_result from logger import ( CLOGGER as logger, PANEL_REMNA, PANEL_XUI, ) -from panels._3xui import ClientConfig, add_client, get_xui_instance -from panels.remnawave import RemnawaveAPI, get_vless_link_for_remnawave_by_username from .aggregated_links import make_aggregated_link @@ -39,6 +35,9 @@ async def create_key_on_cluster( current_traffic_limit_gb: int = None, selected_price_rub: int = None, ): + from services.clusters import ALLOWED_GROUP_CODES, check_server_key_limit + from panels._3xui import ClientConfig, add_client, get_xui_instance + from panels.remnawave import RemnawaveAPI, get_vless_link_for_remnawave_by_username try: servers = await get_servers(session) cluster = servers.get(cluster_id) @@ -275,7 +274,7 @@ async def create_key_on_cluster( if (remnawave_created and remnawave_client_id) or xui_servers: await store_key( session=session, - tg_id=tg_id, + legacy_user_ref=tg_id, client_id=final_client_id, email=email, expiry_time=expiry_timestamp, @@ -289,8 +288,12 @@ async def create_key_on_cluster( current_traffic_limit=current_traffic_limit_gb, selected_price_rub=selected_price_rub, ) - await session.execute(update(User).where(User.tg_id == tg_id, User.trial.in_([0, -1])).values(trial=1)) - await session.commit() + await mark_trial_started_if_eligible(session, tg_id) + try: + from database.web_notifications import notify_web + await notify_web(session, tg_id=tg_id, type="key_created", template_vars={"email": email}, data={"email": email, "client_id": client_id}) + except Exception as e: + logger.warning("[KeyCreate] Ошибка web-уведомления о создании ключа tg_id={}: {}", tg_id, e) except Exception as e: logger.error(f"Ошибка при создании ключа: {e}") @@ -310,6 +313,7 @@ async def create_client_on_server( total_traffic_limit_bytes: int = 0, device_limit_value: int = 0, ): + from panels._3xui import ClientConfig, add_client, get_xui_instance logger.debug( f"{PANEL_XUI} [Client] Вход в create_client_on_server: " f"сервер={server_info.get('server_name')}, план={plan}, is_trial={is_trial}" diff --git a/handlers/keys/operations/deletion.py b/services/operations/deletion.py similarity index 98% rename from handlers/keys/operations/deletion.py rename to services/operations/deletion.py index 593f5a5a..c7722f1d 100644 --- a/handlers/keys/operations/deletion.py +++ b/services/operations/deletion.py @@ -10,12 +10,12 @@ from logger import ( PANEL_XUI, ) from panels._3xui import delete_client, get_xui_instance -from panels.remnawave import RemnawaveAPI from .utils import unique_by_api_url async def delete_key_from_cluster(cluster_id: str, email: str, client_id: str, session: AsyncSession): + from panels.remnawave import RemnawaveAPI try: servers = await get_servers(session) cluster = servers.get(cluster_id) diff --git a/handlers/keys/operations/renewal.py b/services/operations/renewal.py similarity index 99% rename from handlers/keys/operations/renewal.py rename to services/operations/renewal.py index 27d74a89..5ce07c44 100644 --- a/handlers/keys/operations/renewal.py +++ b/services/operations/renewal.py @@ -17,7 +17,7 @@ from database import ( update_key_expiry, update_key_link, ) -from handlers.utils import ALLOWED_GROUP_CODES +from services.clusters import ALLOWED_GROUP_CODES from hooks.processors import process_get_cryptolink_after_renewal from logger import ( CLOGGER as logger, @@ -27,7 +27,6 @@ from logger import ( from panels._3xui import extend_client_key, get_xui_instance from .aggregated_links import make_aggregated_link -from ...tariffs.subgroup_migration import migrate_between_subgroups async def resolve_cluster(session: AsyncSession, cluster_id: str): @@ -271,6 +270,7 @@ async def renew_key_in_cluster( hwid_device_limit = dl if (target_subgroup or "") != (old_subgroup or "") and not single_server: + from services.tariffs.subgroup_migration import migrate_between_subgroups new_client_id, remna_link = await migrate_between_subgroups( session=session, cluster_all=cluster, diff --git a/handlers/keys/operations/toggles.py b/services/operations/toggles.py similarity index 98% rename from handlers/keys/operations/toggles.py rename to services/operations/toggles.py index da747784..e27b8b9a 100644 --- a/handlers/keys/operations/toggles.py +++ b/services/operations/toggles.py @@ -8,7 +8,6 @@ from config import REMNAWAVE_LOGIN, REMNAWAVE_PASSWORD, SUPERNODE from database import get_servers from logger import logger from panels._3xui import get_xui_instance, toggle_client -from panels.remnawave import RemnawaveAPI async def toggle_client_on_cluster( @@ -18,6 +17,7 @@ async def toggle_client_on_cluster( enable: bool = True, session: AsyncSession = None, ) -> dict[str, Any]: + from panels.remnawave import RemnawaveAPI try: if session is None: raise ValueError("[Cluster Toggle] Не передан объект сессии для toggle_client_on_cluster") diff --git a/handlers/keys/operations/traffic.py b/services/operations/traffic.py similarity index 85% rename from handlers/keys/operations/traffic.py rename to services/operations/traffic.py index 55a8635e..55bcee6f 100644 --- a/handlers/keys/operations/traffic.py +++ b/services/operations/traffic.py @@ -2,19 +2,22 @@ import asyncio from typing import Any -from sqlalchemy import or_, select from sqlalchemy.ext.asyncio import AsyncSession from config import SUPERNODE +from database import get_servers +from database.access.resolution import resolve_user_optional +from database.keys import ( + get_key_client_id_by_email_and_server, + get_user_keys_with_servers_by_email, +) +from logger import logger +from panels._3xui import get_client_traffic, get_xui_instance from panels.remnawave_runtime import ( get_remnawave_profile, invalidate_remnawave_profile, with_remnawave_api, ) -from database import get_servers -from database.models import Key, Server -from logger import logger -from panels._3xui import get_client_traffic, get_xui_instance async def get_user_traffic(session: AsyncSession, tg_id: int, email: str) -> dict[str, Any]: @@ -23,35 +26,22 @@ async def get_user_traffic(session: AsyncSession, tg_id: int, email: str) -> dic Для Remnawave трафик считается один раз и отображается как "Remnawave (общий):". Один запрос: Key + Server через join. """ - join_cond = or_( - Key.server_id == Server.server_name, - Key.server_id == Server.cluster_name, - ) - result = await session.execute( - select(Key.client_id, Key.server_id, Server) - .select_from(Key) - .join(Server, join_cond) - .where(Server.enabled.is_(True), Key.tg_id == tg_id, Key.email == email) - ) - rows = result.all() + u = await resolve_user_optional(session, tg_id) + if u is None: + return {"status": "error", "message": "У пользователя нет активных ключей."} + rows = await get_user_keys_with_servers_by_email(session, u.id, email) if not rows: return {"status": "error", "message": "У пользователя нет активных ключей."} - seen_pairs = set() unique_rows = [] servers_map = {} - for client_id, server_id, server in rows: + for client_id, server_id, server_info in rows: if (client_id, server_id) not in seen_pairs: seen_pairs.add((client_id, server_id)) unique_rows.append((client_id, server_id)) - if server.server_name not in servers_map: - servers_map[server.server_name] = { - "server_name": server.server_name, - "cluster_name": server.cluster_name, - "api_url": server.api_url, - "panel_type": server.panel_type, - } + if server_info["server_name"] not in servers_map: + servers_map[server_info["server_name"]] = server_info user_traffic_data = {} tasks = [] @@ -135,17 +125,11 @@ async def reset_traffic_in_cluster(cluster_id: str, email: str, session: AsyncSe inbound_id = server_info.get("inbound_id") if panel_type == "remnawave" and not remnawave_done: - result = await session.execute( - select(Key.client_id).where(Key.email == email, Key.server_id == cluster_id).limit(1) - ) - row = result.first() - - if not row: + client_id = await get_key_client_id_by_email_and_server(session, email, cluster_id) + if not client_id: logger.warning(f"[Remnawave Reset] client_id не найден для {email} на {server_name}") continue - client_id = row[0] - async def _reset(api): done = await api.reset_user_traffic(client_id) if done: diff --git a/handlers/keys/operations/update.py b/services/operations/update.py similarity index 92% rename from handlers/keys/operations/update.py rename to services/operations/update.py index 51f37b40..c2fba466 100644 --- a/handlers/keys/operations/update.py +++ b/services/operations/update.py @@ -2,22 +2,23 @@ import asyncio from datetime import datetime, timezone -from sqlalchemy import delete, select from sqlalchemy.ext.asyncio import AsyncSession from config import PUBLIC_LINK, SUPERNODE -from panels.remnawave_runtime import invalidate_remnawave_profile, with_remnawave_api from database import filter_cluster_by_subgroup, filter_cluster_by_tariff, get_servers, get_tariff_by_id, store_key -from handlers.utils import ALLOWED_GROUP_CODES -from database.models import Key, Tariff -from handlers.tariffs.tariff_display import GB, get_effective_limits_for_key -from handlers.utils import get_least_loaded_cluster +from database.access.resolution import resolve_user_optional +from database.keys import delete_key_by_user_and_email, get_key_by_user_and_email +from database.models import Key +from database.tariffs import get_active_tariff_by_id from logger import ( CLOGGER as logger, PANEL_REMNA, PANEL_XUI, ) from panels._3xui import ClientConfig, add_client, get_xui_instance +from panels.remnawave_runtime import invalidate_remnawave_profile, with_remnawave_api +from services.clusters import ALLOWED_GROUP_CODES, select_cluster +from services.tariffs.tariff_display import GB, get_effective_limits_for_key from .aggregated_links import make_aggregated_link from .deletion import delete_key_from_cluster @@ -96,13 +97,6 @@ async def update_key_on_cluster( if not group_code: raise ValueError("У Remnawave-сервера отсутствует tariff_group") - _ = await session.execute( - select(Tariff) - .where(Tariff.group_code == group_code, Tariff.is_active.is_(True)) - .order_by(Tariff.duration_days.desc()) - .limit(1) - ) - short_uuid = None if remnawave_link and "/" in remnawave_link: short_uuid = remnawave_link.rstrip("/").split("/")[-1] @@ -176,13 +170,6 @@ async def update_key_on_cluster( if not group_code: raise ValueError(f"У сервера {server_name} отсутствует tariff_group") - _ = await session.execute( - select(Tariff) - .where(Tariff.group_code == group_code, Tariff.is_active.is_(True)) - .order_by(Tariff.duration_days.desc()) - .limit(1) - ) - total_gb_bytes = int(traffic_limit * 1024**3) if traffic_limit is not None else 0 device_limit_value = device_limit if device_limit is not None else 0 @@ -220,8 +207,11 @@ async def update_subscription( country_override: str = None, remnawave_link: str = None, ) -> None: - result = await session.execute(select(Key).where(Key.tg_id == tg_id, Key.email == email)) - record: Key | None = result.scalar_one_or_none() + u = await resolve_user_optional(session, tg_id) + if u is None: + raise ValueError(f"The key {email} does not exist in database") + uid = u.id + record: Key | None = await get_key_by_user_and_email(session, uid, email) if not record: raise ValueError(f"The key {email} does not exist in database") @@ -244,8 +234,7 @@ async def update_subscription( external_squad_uuid = None if tariff_id: - q = await session.execute(select(Tariff).where(Tariff.id == tariff_id, Tariff.is_active.is_(True))) - tariff = q.scalar_one_or_none() + tariff = await get_active_tariff_by_id(session, int(tariff_id)) if tariff is None: logger.warning(f"[LOG] update_subscription: тариф с id={tariff_id} не найден!") else: @@ -258,14 +247,15 @@ async def update_subscription( from middlewares.session import release_session_early await release_session_early(session) await delete_key_from_cluster(old_cluster_id, email, client_id, session=session) - await session.execute(delete(Key).where(Key.tg_id == tg_id, Key.email == email)) + await delete_key_by_user_and_email(session, uid, email) await session.commit() if country_override or cluster_override: new_cluster_id = country_override or cluster_override else: try: - new_cluster_id = await get_least_loaded_cluster(session) + result = await select_cluster(session) + new_cluster_id = result.cluster_name except ValueError: logger.warning("[Update] Нет доступных кластеров, оставляем на старом") new_cluster_id = old_cluster_id @@ -368,7 +358,7 @@ async def update_subscription( await store_key( session=session, - tg_id=tg_id, + legacy_user_ref=tg_id, client_id=new_client_id, email=email, expiry_time=expiry_time, diff --git a/handlers/keys/operations/utils.py b/services/operations/utils.py similarity index 100% rename from handlers/keys/operations/utils.py rename to services/operations/utils.py diff --git a/services/payments/__init__.py b/services/payments/__init__.py new file mode 100644 index 00000000..982e587f --- /dev/null +++ b/services/payments/__init__.py @@ -0,0 +1,72 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +from config import CASHBACK as DEFAULT_CASHBACK, REFERRAL_BONUS_PERCENTAGES +from core.bootstrap import MONEY_CONFIG +from database import add_payment +from database.access.resolution import resolve_user_optional +from database.referrals import get_referral_by_referred_id +from database.users import update_balance +from logger import logger + +if TYPE_CHECKING: + from sqlalchemy.ext.asyncio import AsyncSession + + +async def process_referrals(session: AsyncSession, user_id: int, amount: float) -> dict[int, float]: + """Начисляет реферальные бонусы по цепочке. + + Returns: dict {referrer_id: bonus_amount} — для логирования/уведомлений. + """ + u = await resolve_user_optional(session, user_id) + if u is None: + return {} + + max_levels = len(REFERRAL_BONUS_PERCENTAGES) + current_id = u.id + bonus_by_chain: dict[int, tuple[float, int]] = {} + + for level in range(1, max_levels + 1): + referral = await get_referral_by_referred_id(session, current_id) + if not referral: + break + referrer_id = int(referral["referrer_user_id"]) + percent = REFERRAL_BONUS_PERCENTAGES.get(level) + if not percent: + continue + bonus = amount * percent if isinstance(percent, float) else percent + bonus_by_chain[referrer_id] = (bonus, level) + current_id = referrer_id + + result_map: dict[int, float] = {} + for referrer_id, (bonus, lvl) in bonus_by_chain.items(): + await update_balance(session, referrer_id, float(bonus)) + await add_payment(session, tg_id=referrer_id, amount=bonus, payment_system="referral") + logger.info(f"Начислен бонус {bonus}₽ пользователю {referrer_id} за уровень {lvl}") + result_map[referrer_id] = bonus + + return result_map + + +async def process_cashback(session: AsyncSession, user_id: int, amount: float) -> float: + """Начисляет кэшбэк на баланс пользователя. + + Returns: сумма кэшбэка (0 если отключён). + """ + cashback_config = MONEY_CONFIG.get("CASHBACK", DEFAULT_CASHBACK) + try: + cashback_percent = float(cashback_config) if cashback_config not in (None, False) else 0.0 + except (TypeError, ValueError): + cashback_percent = 0.0 + + if cashback_percent <= 0: + return 0.0 + + cashback_amount = round(amount * (cashback_percent / 100)) + if cashback_amount > 0: + await update_balance(session, user_id, cashback_amount) + await add_payment(session, tg_id=user_id, amount=cashback_amount, payment_system="cashback") + logger.info(f"Начислен кешбэк {cashback_amount}₽ пользователю {user_id}") + + return float(cashback_amount) diff --git a/services/payments/cryptobot/__init__.py b/services/payments/cryptobot/__init__.py new file mode 100644 index 00000000..8b137891 --- /dev/null +++ b/services/payments/cryptobot/__init__.py @@ -0,0 +1 @@ + diff --git a/handlers/payments/currency_rates.py b/services/payments/currency_rates.py similarity index 97% rename from handlers/payments/currency_rates.py rename to services/payments/currency_rates.py index 8363ede0..cbca5171 100644 --- a/handlers/payments/currency_rates.py +++ b/services/payments/currency_rates.py @@ -182,11 +182,8 @@ async def money_for_user( currency: "USD" или "RUB" value: Decimal в выбранной валюте """ - row = await db_session.execute( - sa.text("select preferred_currency from users where tg_id = :id"), - {"id": tg_id}, - ) - user_currency = row.scalar() + from database.users import get_user_preferred_currency + user_currency = await get_user_preferred_currency(db_session, tg_id) txt, cur, val = await display_price( amount_rub, diff --git a/services/payments/heleket/__init__.py b/services/payments/heleket/__init__.py new file mode 100644 index 00000000..8b137891 --- /dev/null +++ b/services/payments/heleket/__init__.py @@ -0,0 +1 @@ + diff --git a/services/payments/heleket/webhook.py b/services/payments/heleket/webhook.py new file mode 100644 index 00000000..82c240de --- /dev/null +++ b/services/payments/heleket/webhook.py @@ -0,0 +1,148 @@ +import base64 +import hashlib +import json + +from aiohttp import web + +from config import HELEKET_API_KEY +from core.webhook_abuse import ( + get_webhook_client_ip, + is_webhook_ip_blocked, + record_webhook_signature_failure, +) +from logger import logger +from services.payments.pipeline import ( + ParsedPayment, + process_cancelled_payment, + process_success_payment, +) + +_PROVIDER = "heleket" + + +def verify_heleket_signature(data: dict) -> bool: + """Проверяет MD5-подпись webhook от Heleket.""" + try: + received_signature = data.get("sign") + if not received_signature: + logger.error("Heleket webhook: отсутствует подпись") + return False + + data_without_sign = data.copy() + del data_without_sign["sign"] + + json_data = json.dumps(data_without_sign, ensure_ascii=False, separators=(",", ":")) + json_data = json_data.replace("/", "\\/") + base64_data = base64.b64encode(json_data.encode("utf-8")).decode("utf-8") + sign_string = base64_data + HELEKET_API_KEY + calculated_signature = hashlib.md5(sign_string.encode("utf-8")).hexdigest() + is_valid = calculated_signature.lower() == received_signature.lower() + + if not is_valid: + logger.error( + f"Heleket webhook: неверная подпись. Ожидалось: {calculated_signature}, получено: {received_signature}" + ) + else: + logger.info("Heleket webhook: подпись успешно проверена") + return is_valid + except Exception as e: + logger.error(f"Ошибка проверки подписи Heleket webhook: {e}") + return False + + +def _extract_tg_and_amount(order_id: str, additional_data, merchant_amount) -> tuple[int | None, float | None]: + """Извлекает tg_id и сумму для зачисления (rub_amount или merchant_amount).""" + tg_id = None + rub_amount = None + if additional_data: + try: + for part in str(additional_data).split(","): + if part.startswith("tg_id:"): + tg_id = int(part.split(":")[1]) + elif part.startswith("rub_amount:"): + rub_amount = float(part.split(":")[1]) + except Exception as e: + logger.error(f"Ошибка парсинга additional_data: {e}") + if not tg_id and order_id and "_" in order_id: + try: + tg_id = int(order_id.split("_")[1]) + except Exception as e: + logger.error(f"Ошибка извлечения tg_id из order_id: {e}") + balance_amount = rub_amount if rub_amount else (float(merchant_amount) if merchant_amount else None) + return tg_id, balance_amount + + +async def process_heleket_webhook(data: dict) -> bool: + """Обрабатывает уже верифицированный webhook от Heleket.""" + try: + logger.info(f"Processing Heleket webhook: {data}") + + webhook_type = data.get("type") + order_id = data.get("order_id") + status = data.get("status") + + if webhook_type != "payment": + logger.warning(f"Heleket webhook: неизвестный тип {webhook_type}") + return False + + if status in ["paid", "paid_over"]: + tg_id, balance_amount = _extract_tg_and_amount( + order_id=order_id, + additional_data=data.get("additional_data"), + merchant_amount=data.get("merchant_amount"), + ) + if not tg_id: + logger.error(f"Не удалось извлечь tg_id из Heleket webhook: {data}") + return False + if balance_amount is None: + logger.error(f"Не удалось определить сумму зачисления: {data}") + return False + + parsed = ParsedPayment( + payment_id=str(order_id), + tg_id=tg_id, + amount=balance_amount, + currency="USD", + ) + result = await process_success_payment(_PROVIDER, parsed) + return result.ok + + if status in ["fail", "wrong_amount", "cancel", "system_fail"]: + logger.warning(f"Heleket: неудачный платёж {order_id}, статус: {status}") + parsed = ParsedPayment( + payment_id=str(order_id), + tg_id=None, + amount=0.0, + currency="USD", + ) + result = await process_cancelled_payment(_PROVIDER, parsed, new_status="failed") + return result.ok + + logger.info(f"Heleket: промежуточный статус {status} для платежа {order_id}") + return True + except Exception as e: + logger.error(f"Ошибка обработки Heleket webhook: {e}") + return False + + +async def heleket_webhook(request: web.Request): + """Обработчик webhook от Heleket для aiohttp.""" + try: + ip = get_webhook_client_ip(request) + if await is_webhook_ip_blocked(ip): + return web.Response(status=429) + data = await request.json() + logger.info(f"Heleket webhook received from {request.remote}") + + if not verify_heleket_signature(data): + logger.error("Heleket webhook: неверная подпись") + await record_webhook_signature_failure(ip) + return web.Response(status=400, text="Invalid signature") + + success = await process_heleket_webhook(data) + if success: + return web.Response(status=200, text="OK") + return web.Response(status=400, text="Processing failed") + except Exception as e: + logger.error(f"Ошибка обработки Heleket webhook: {e}") + return web.Response(status=500, text="Internal server error") diff --git a/services/payments/kassai/__init__.py b/services/payments/kassai/__init__.py new file mode 100644 index 00000000..8b137891 --- /dev/null +++ b/services/payments/kassai/__init__.py @@ -0,0 +1 @@ + diff --git a/services/payments/kassai/webhook.py b/services/payments/kassai/webhook.py new file mode 100644 index 00000000..1434bae9 --- /dev/null +++ b/services/payments/kassai/webhook.py @@ -0,0 +1,87 @@ +import hashlib + +from aiohttp import web + +from config import KASSAI_SECRET_KEY, KASSAI_SHOP_ID, KASSAI_WEBHOOK_RESPONSE +from core.webhook_abuse import ( + get_webhook_client_ip, + is_webhook_ip_blocked, + record_webhook_signature_failure, +) +from logger import logger +from services.payments.pipeline import ParsedPayment, process_success_payment + +_PROVIDER = "kassai" + + +def verify_kassai_signature(data: dict, signature: str) -> bool: + """Проверяет MD5-подпись webhook от KassaAI.""" + try: + sign_string = ( + f"{KASSAI_SHOP_ID}:{data.get('AMOUNT', '')}:{KASSAI_SECRET_KEY}:{data.get('MERCHANT_ORDER_ID', '')}" + ) + expected_signature = hashlib.md5(sign_string.encode("utf-8")).hexdigest() + result = signature.upper() == expected_signature.upper() + if not result: + logger.error( + f"KassaAI signature mismatch. Expected: {expected_signature}, Got: {signature}" + ) + else: + logger.info("KassaAI webhook: подпись успешно проверена") + return result + except Exception as e: + logger.error(f"Ошибка проверки подписи KassaAI: {e}") + return False + + +def _parse_kassai(data) -> ParsedPayment | None: + amount_raw = data.get("AMOUNT") + order_id = data.get("MERCHANT_ORDER_ID") + if not amount_raw or not order_id: + return None + try: + tg_id = int(str(order_id).split("_")[1]) + amount = float(amount_raw) + except (IndexError, ValueError) as e: + logger.error( + f"KassaAI webhook: не удалось извлечь tg_id/amount из order_id={order_id}: {e}" + ) + return None + return ParsedPayment( + payment_id=str(order_id), + tg_id=tg_id, + amount=amount, + currency="RUB", + ) + + +async def kassai_webhook(request: web.Request): + try: + ip = get_webhook_client_ip(request) + if await is_webhook_ip_blocked(ip): + return web.Response(status=429) + + data = await request.post() + logger.info(f"KassaAI webhook received: {dict(data)}") + + signature = data.get("SIGN", "") + if not signature: + await record_webhook_signature_failure(ip) + return web.Response(status=400) + + if not verify_kassai_signature(data, signature): + await record_webhook_signature_failure(ip) + return web.Response(status=400) + + parsed = _parse_kassai(data) + if parsed is None: + return web.Response(status=400) + + result = await process_success_payment(_PROVIDER, parsed) + if not result.ok: + return web.Response(status=500) + + return web.Response(text=KASSAI_WEBHOOK_RESPONSE) + except Exception as e: + logger.error(f"Ошибка обработки KassaAI webhook: {e}") + return web.Response(status=500) diff --git a/services/payments/payment_events.py b/services/payments/payment_events.py new file mode 100644 index 00000000..079bf599 --- /dev/null +++ b/services/payments/payment_events.py @@ -0,0 +1,39 @@ +import json + +from config import REDIS_URL +from logger import logger + + +def payment_events_channel(legacy_user_ref: int) -> str: + return f"payment_events:user:{int(legacy_user_ref)}" + + +async def publish_payment_event( + *, + legacy_user_ref: int, + status: str, + flow: str | None = None, + amount: float | int | None = None, +) -> None: + try: + from redis.asyncio import from_url + + payload: dict[str, str | float | int] = {"status": str(status)} + if flow: + payload["flow"] = str(flow) + if amount is not None: + payload["amount"] = float(amount) + client = from_url(REDIS_URL, encoding="utf-8", decode_responses=True, max_connections=8) + try: + subscribers = await client.publish( + payment_events_channel(int(legacy_user_ref)), + json.dumps(payload, ensure_ascii=False), + ) + logger.info( + f"[Payments] Event published: user_ref={legacy_user_ref}, status={status}, " + f"flow={flow}, subscribers={subscribers}" + ) + finally: + await client.aclose() + except Exception as e: + logger.warning(f"[Payments] publish_payment_event failed: {e}") diff --git a/handlers/payments/payment_links.py b/services/payments/payment_links.py similarity index 94% rename from handlers/payments/payment_links.py rename to services/payments/payment_links.py index f4e0c0f3..5dfd15a3 100644 --- a/handlers/payments/payment_links.py +++ b/services/payments/payment_links.py @@ -9,7 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession @dataclass(frozen=True) class PaymentLinkRequest: - tg_id: int + legacy_user_ref: int amount: int | float currency: str provider_id: str @@ -27,7 +27,7 @@ class PaymentLinkResult: PaymentLinkCreator = Callable[ - [AsyncSession, int, float, str, str | None, str | None], + [AsyncSession, int, float, str, str | None, str | None, dict[str, Any] | None], Awaitable[tuple[str, str | None]], ] @@ -75,11 +75,12 @@ async def create_payment_link( try: url, payment_id = await creator( session, - request.tg_id, + request.legacy_user_ref, amount, currency, request.success_url, request.failure_url, + request.metadata, ) return PaymentLinkResult(success=True, payment_url=url, payment_id=payment_id) except ValueError as e: diff --git a/services/payments/pipeline.py b/services/payments/pipeline.py new file mode 100644 index 00000000..40af8cfe --- /dev/null +++ b/services/payments/pipeline.py @@ -0,0 +1,189 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import TYPE_CHECKING + +from sqlalchemy import select + +from database import ( + add_payment, + async_session_maker, + get_payment_by_payment_id, + invalidate_payment_cache, + update_balance, + update_payment_status, +) +from database.models import Payment +from handlers.payments.utils import send_payment_success_notification +from logger import logger + +if TYPE_CHECKING: + pass + + +@dataclass +class ParsedPayment: + """Нормализованный результат парсинга webhook-payload'а провайдера.""" + + payment_id: str + tg_id: int | None + amount: float + currency: str = "RUB" + metadata: dict | None = field(default=None) + + +@dataclass +class PipelineResult: + """Что pipeline вернул адаптеру — для корректного HTTP-ответа.""" + + ok: bool + already_processed: bool = False + error: str | None = None + + +async def process_success_payment( + provider: str, + parsed: ParsedPayment, + *, + metadata_patch: dict | None = None, + credit_amount_override: float | None = None, + update_currency: str | None = None, + update_original_amount: float | None = None, +) -> PipelineResult: + """Идемпотентно переводит платёж в success, зачисляет баланс, уведомляет. + + Открывает одну транзакцию на всю операцию — если что-то упадёт, всё + откатывается атомарно. + + ``provider`` — строка для колонки ``payments.payment_system`` (регистр + важен, некоторые провайдеры исторически писали как "YOOMONEY"/"HELEKET", + см. комментарии в конкретных адаптерах). + + ``metadata_patch`` — опциональный dict, который ПАТЧИТ (merge) существующий + ``payments.metadata_`` для провайдера у которых метадата приходит только + в webhook'е (cryptobot: FX rate, invoice_id, paid amount). + + ``credit_amount_override`` — зачислить на баланс сумму, отличную от + ``parsed.amount``. Нужно для cryptobot: провайдер возвращает paid_amount + в USDT, но баланс пополняется исходной RUB-суммой из pending-записи. + + ``update_currency`` / ``update_original_amount`` — дополняют ``Payment`` + row для crypto-платежей (зафиксировать реально списанную валюту). + """ + try: + async with async_session_maker() as session: + payment = await get_payment_by_payment_id(session, parsed.payment_id) + + if payment and payment.get("status") == "success": + logger.info( + f"[{provider}] Повторный webhook, платёж уже обработан: payment_id={parsed.payment_id}" + ) + return PipelineResult(ok=True, already_processed=True) + + if payment and payment.get("id") is not None: + updated = await update_payment_status( + session=session, + internal_id=int(payment["id"]), + new_status="success", + metadata_patch=metadata_patch, + ) + if not updated: + logger.error( + f"[{provider}] Не удалось перевести платёж id={payment['id']} в success" + ) + return PipelineResult(ok=False, error="update_payment_status failed") + tg_id = parsed.tg_id if parsed.tg_id is not None else int(payment["tg_id"]) + + + if update_currency is not None or update_original_amount is not None: + row = ( + await session.execute( + select(Payment).where(Payment.id == int(payment["id"])).limit(1) + ) + ).scalar_one_or_none() + if row is not None: + if update_currency is not None: + row.currency = update_currency + if update_original_amount is not None: + row.original_amount = update_original_amount + else: + await add_payment( + session=session, + tg_id=parsed.tg_id, + amount=parsed.amount, + payment_system=provider, + status="success", + currency=parsed.currency, + payment_id=parsed.payment_id, + metadata=parsed.metadata or metadata_patch, + ) + tg_id = parsed.tg_id + + credit_amount = ( + float(credit_amount_override) if credit_amount_override is not None else parsed.amount + ) + if tg_id is not None and credit_amount > 0: + await update_balance(session, tg_id, credit_amount) + await send_payment_success_notification(tg_id, credit_amount, session) + + await session.commit() + await invalidate_payment_cache(parsed.payment_id) + + logger.info( + f"[{provider}] Платёж обработан: payment_id={parsed.payment_id}, " + f"tg_id={tg_id}, amount={credit_amount} (parsed={parsed.amount} {parsed.currency})" + ) + return PipelineResult(ok=True) + except Exception as e: + logger.error(f"[{provider}] Ошибка обработки успешного платежа: {e}") + return PipelineResult(ok=False, error=str(e)) + + +async def process_cancelled_payment( + provider: str, + parsed: ParsedPayment, + *, + new_status: str = "cancelled", +) -> PipelineResult: + """Переводит платёж в cancelled/failed либо записывает его сразу в этом статусе. + + ``new_status`` — обычно "cancelled" (пользователь отменил) или "failed" + (провайдер вернул ошибку/wrong_amount/system_fail). + """ + try: + async with async_session_maker() as session: + payment = await get_payment_by_payment_id(session, parsed.payment_id) + + if payment and payment.get("status") in ("cancelled", "failed", "success"): + return PipelineResult(ok=True, already_processed=True) + + if payment and payment.get("id") is not None: + updated = await update_payment_status( + session=session, + internal_id=int(payment["id"]), + new_status=new_status, + ) + if not updated: + return PipelineResult(ok=False, error="update_payment_status failed") + else: + await add_payment( + session=session, + tg_id=parsed.tg_id, + amount=parsed.amount, + payment_system=provider, + status=new_status, + currency=parsed.currency, + payment_id=parsed.payment_id, + metadata=parsed.metadata, + ) + + await session.commit() + await invalidate_payment_cache(parsed.payment_id) + + logger.info( + f"[{provider}] Платёж {parsed.payment_id} помечен как {new_status}" + ) + return PipelineResult(ok=True) + except Exception as e: + logger.error(f"[{provider}] Ошибка при обработке отмены платежа: {e}") + return PipelineResult(ok=False, error=str(e)) diff --git a/handlers/payments/providers.py b/services/payments/providers.py similarity index 94% rename from handlers/payments/providers.py rename to services/payments/providers.py index 34e078cf..819bb7ce 100644 --- a/handlers/payments/providers.py +++ b/services/payments/providers.py @@ -67,6 +67,22 @@ PROVIDERS_BASE: dict[str, dict[str, Any]] = { }, } +WEB_LINK_PROVIDER_IDS = ( + "YOOKASSA", + "YOOMONEY", + "ROBOKASSA", + "KASSAI_CARDS", + "KASSAI_SBP", + "HELEKET", + "FREEKASSA", + "CRYPTOBOT", +) + +TELEGRAM_ONLY_PROVIDER_IDS = ( + "TRIBUTE", + "STARS", +) + def _get_effective_order(name: str, cfg: dict[str, Any]) -> int: """Возвращает эффективный порядок провайдера (админ > модуль > дефолт).""" diff --git a/services/payments/robokassa/__init__.py b/services/payments/robokassa/__init__.py new file mode 100644 index 00000000..8b137891 --- /dev/null +++ b/services/payments/robokassa/__init__.py @@ -0,0 +1 @@ + diff --git a/services/payments/tribute/__init__.py b/services/payments/tribute/__init__.py new file mode 100644 index 00000000..8b137891 --- /dev/null +++ b/services/payments/tribute/__init__.py @@ -0,0 +1 @@ + diff --git a/services/payments/yookassa/__init__.py b/services/payments/yookassa/__init__.py new file mode 100644 index 00000000..8b137891 --- /dev/null +++ b/services/payments/yookassa/__init__.py @@ -0,0 +1 @@ + diff --git a/services/payments/yoomoney/__init__.py b/services/payments/yoomoney/__init__.py new file mode 100644 index 00000000..8b137891 --- /dev/null +++ b/services/payments/yoomoney/__init__.py @@ -0,0 +1 @@ + diff --git a/handlers/keys/subscriptions.py b/services/subscriptions.py similarity index 97% rename from handlers/keys/subscriptions.py rename to services/subscriptions.py index d3d5d572..0c97ecfa 100644 --- a/handlers/keys/subscriptions.py +++ b/services/subscriptions.py @@ -8,7 +8,6 @@ import urllib.parse import aiohttp from aiohttp import web -from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from config import ( @@ -26,7 +25,7 @@ from core.cache_config import ( ) from core.redis_cache import cache_get, cache_key, cache_set from database import get_key_details, get_servers -from database.models import Server +from database.servers import get_enabled_server_subscription_url from handlers.texts import HAPP_ANNOUNCE, HIDDIFY_PROFILE_TITLE, SUBSCRIPTION_INFO_TEXT, V2RAYTUN_ANNOUNCE from handlers.utils import convert_to_bytes from logger import logger @@ -84,10 +83,7 @@ async def get_subscription_urls( use_country_selection = bool(MODES_CONFIG.get("COUNTRY_SELECTION_ENABLED", USE_COUNTRY_SELECTION)) if use_country_selection: - result = await session.execute( - select(Server.subscription_url).where(Server.server_name == server_id, Server.enabled.is_(True)) - ) - server_data = result.scalar() + server_data = await get_enabled_server_subscription_url(session, server_id) if server_data: urls.append(f"{server_data}/{email}") else: @@ -291,8 +287,6 @@ async def handle_subscription(request: web.Request) -> web.Response: subscription_userinfo = calculate_traffic(cleaned_subscriptions, expiry_time_ms, headers_list) headers = prepare_headers(user_agent, PROJECT_NAME, subscription_info, subscription_userinfo) - await session.commit() - await cache_set( cache_key_sub, {"b": base64_encoded, "h": dict(headers)}, diff --git a/services/tariffs/__init__.py b/services/tariffs/__init__.py new file mode 100644 index 00000000..31b7ab84 --- /dev/null +++ b/services/tariffs/__init__.py @@ -0,0 +1,2 @@ +from .pricing import calculate_config_price +from .tariff_display import GB, get_effective_limits_for_key diff --git a/services/tariffs/pricing.py b/services/tariffs/pricing.py new file mode 100644 index 00000000..23076769 --- /dev/null +++ b/services/tariffs/pricing.py @@ -0,0 +1,108 @@ +from __future__ import annotations + +from typing import Any + +from core.settings.tariffs_config import normalize_tariff_config +from .tariff_display import GB + + +def calculate_config_price( + tariff: dict[str, Any], + selected_device_limit: int | None = None, + selected_traffic_gb: int | None = None, +) -> int: + """Рассчитывает цену тарифа с учётом выбранных лимитов.""" + cfg = normalize_tariff_config(tariff) + + base_price = int(tariff.get("price_rub") or 0) + + raw_device_options = tariff.get("device_options") + raw_traffic_options = tariff.get("traffic_options_gb") + raw_device_options = raw_device_options if isinstance(raw_device_options, list) else [] + raw_traffic_options = raw_traffic_options if isinstance(raw_traffic_options, list) else [] + + device_values: list[int] = [] + for value in raw_device_options: + try: + device_values.append(int(value)) + except (TypeError, ValueError): + continue + + traffic_values: list[int] = [] + for value in raw_traffic_options: + try: + traffic_values.append(int(value)) + except (TypeError, ValueError): + continue + + positive_device_values = [v for v in device_values if v > 0] + positive_traffic_values = [v for v in traffic_values if v > 0] + + base_device_limit = cfg.get("base_device_limit") + if base_device_limit is None: + base_device_limit = tariff.get("device_limit") + if base_device_limit is None: + if positive_device_values: + base_device_limit = min(positive_device_values) + elif device_values: + base_device_limit = device_values[0] + base_device_limit = int(base_device_limit) if base_device_limit is not None else None + + base_traffic_gb = cfg.get("base_traffic_gb") + if base_traffic_gb is None: + traffic_limit_raw = tariff.get("traffic_limit") + if traffic_limit_raw: + traffic_limit_raw = int(traffic_limit_raw) + if traffic_limit_raw >= GB: + base_traffic_gb = int(traffic_limit_raw / GB) + else: + base_traffic_gb = traffic_limit_raw + else: + if positive_traffic_values: + base_traffic_gb = min(positive_traffic_values) + elif traffic_values: + base_traffic_gb = traffic_values[0] + base_traffic_gb = int(base_traffic_gb) if base_traffic_gb is not None else None + + device_overrides = cfg.get("device_price_overrides") or tariff.get("device_overrides") or {} + traffic_overrides = cfg.get("traffic_price_overrides") or tariff.get("traffic_overrides") or {} + + extra_device_step_price = int(cfg.get("extra_device_base_price_rub") or tariff.get("device_step_rub") or 0) + extra_traffic_step_price = int( + cfg.get("extra_traffic_base_price_per_gb_rub") or tariff.get("traffic_step_rub") or 0 + ) + + devices_extra_price = 0 + traffic_extra_price = 0 + + if selected_device_limit is not None and base_device_limit is not None: + selected_device_limit = int(selected_device_limit) + override_key = str(selected_device_limit) + if override_key in device_overrides: + devices_extra_price = int(device_overrides[override_key]) + else: + if selected_device_limit <= 0: + if positive_device_values: + effective_devices = max(positive_device_values) + extra_devices = max(0, effective_devices - base_device_limit) + devices_extra_price = extra_devices * extra_device_step_price + else: + extra_devices = max(0, selected_device_limit - base_device_limit) + devices_extra_price = extra_devices * extra_device_step_price + + if selected_traffic_gb is not None and base_traffic_gb is not None: + selected_traffic_gb = int(selected_traffic_gb) + override_key = str(selected_traffic_gb) + if override_key in traffic_overrides: + traffic_extra_price = int(traffic_overrides[override_key]) + else: + if selected_traffic_gb <= 0: + if positive_traffic_values: + effective_gb = max(positive_traffic_values) + extra_traffic = max(0, effective_gb - base_traffic_gb) + traffic_extra_price = extra_traffic * extra_traffic_step_price + else: + extra_traffic = max(0, selected_traffic_gb - base_traffic_gb) + traffic_extra_price = extra_traffic * extra_traffic_step_price + + return int(base_price + devices_extra_price + traffic_extra_price) diff --git a/handlers/tariffs/subgroup_migration.py b/services/tariffs/subgroup_migration.py similarity index 97% rename from handlers/tariffs/subgroup_migration.py rename to services/tariffs/subgroup_migration.py index 30cae0d5..6892d378 100644 --- a/handlers/tariffs/subgroup_migration.py +++ b/services/tariffs/subgroup_migration.py @@ -13,10 +13,6 @@ from logger import ( PANEL_XUI, ) from panels._3xui import ClientConfig, add_client, extend_client_key, get_xui_instance -from panels.remnawave import RemnawaveAPI - -from ..keys.operations.deletion import delete_on_3xui, delete_on_remnawave -from ..keys.operations.utils import bytes_from_gb, norm_name, split_by_panel async def ensure_on_remnawave( @@ -31,6 +27,8 @@ async def ensure_on_remnawave( attempt_update_first: bool, external_squad_uuid: str | None = None, ) -> tuple[str | None, str | None]: + from panels.remnawave import RemnawaveAPI + from services.operations.utils import bytes_from_gb if not servers: return None, None @@ -155,6 +153,7 @@ async def ensure_on_3xui( hwid_device_limit: int, attempt_update_first: bool, ): + from services.operations.utils import bytes_from_gb tasks = [] traffic = bytes_from_gb(total_gb) for s in servers: @@ -244,6 +243,8 @@ async def migrate_between_subgroups( external_squad_uuid: str | None = None, tariff_id: int | None = None, ) -> tuple[str, str | None]: + from services.operations.deletion import delete_on_3xui, delete_on_remnawave + from services.operations.utils import norm_name, split_by_panel target = await filter_cluster_by_subgroup(session, cluster_all, target_subgroup, cluster_id, tariff_id=tariff_id) xui_tgt, remna_tgt = split_by_panel(target) diff --git a/handlers/tariffs/tariff_display.py b/services/tariffs/tariff_display.py similarity index 100% rename from handlers/tariffs/tariff_display.py rename to services/tariffs/tariff_display.py diff --git a/handlers/admin/users/utils.py b/services/users_utils.py similarity index 100% rename from handlers/admin/users/utils.py rename to services/users_utils.py diff --git a/services/web_push.py b/services/web_push.py new file mode 100644 index 00000000..f214f327 --- /dev/null +++ b/services/web_push.py @@ -0,0 +1,71 @@ +import json + +from logger import logger + +try: + from pywebpush import webpush, WebPushException + _WEBPUSH_AVAILABLE = True +except ImportError: + _WEBPUSH_AVAILABLE = False + +try: + from config import VAPID_PRIVATE_KEY, VAPID_PUBLIC_KEY, VAPID_CLAIMS_EMAIL +except ImportError: + VAPID_PRIVATE_KEY = "" + VAPID_PUBLIC_KEY = "" + VAPID_CLAIMS_EMAIL = "" + + +def push_enabled() -> bool: + return _WEBPUSH_AVAILABLE and bool(VAPID_PRIVATE_KEY) and bool(VAPID_PUBLIC_KEY) + + +async def send_push_notification( + subscription_info: dict, + title: str, + body: str, + url: str = "/dashboard", + tag: str = "solo-notification", +) -> bool: + """Отправить push-уведомление одному подписчику.""" + if not push_enabled(): + logger.debug("[WebPush] push отключён (нет VAPID ключей или pywebpush)") + return False + + payload = json.dumps({ + "title": title, + "body": body, + "url": url, + "tag": tag, + }) + + try: + webpush( + subscription_info=subscription_info, + data=payload, + vapid_private_key=VAPID_PRIVATE_KEY, + vapid_claims={"sub": f"mailto:{VAPID_CLAIMS_EMAIL}"}, + ) + logger.debug("[WebPush] уведомление отправлено: {}", title) + return True + except WebPushException as e: + logger.error("[WebPush] ошибка отправки: {}", e) + return False + except Exception as e: + logger.error("[WebPush] неожиданная ошибка: {}", e) + return False + + +async def send_push_to_many( + subscriptions: list[dict], + title: str, + body: str, + url: str = "/dashboard", + tag: str = "solo-notification", +) -> int: + """Отправить push нескольким подписчикам. Возвращает количество успешных.""" + sent = 0 + for sub in subscriptions: + if await send_push_notification(sub, title, body, url, tag): + sent += 1 + return sent diff --git a/sitecustomize.py b/sitecustomize.py new file mode 100644 index 00000000..43e2e2fb --- /dev/null +++ b/sitecustomize.py @@ -0,0 +1,166 @@ +from __future__ import annotations + +import warnings + +warnings.filterwarnings( + "ignore", + message=r'.*Field "model_custom_emoji_id" in UniqueGiftColors has conflict with protected namespace', + category=UserWarning, +) + +try: + import aiogram.types + from pydantic import ConfigDict + + unique_gift_colors = getattr(aiogram.types, "UniqueGiftColors", None) + if unique_gift_colors is not None: + cfg = getattr(unique_gift_colors, "model_config", None) + base = dict(cfg) if cfg is not None else {} + unique_gift_colors.model_config = ConfigDict(**base, protected_namespaces=()) +except Exception: + pass + +try: + from alembic.ddl.postgresql import PostgresqlImpl + from sqlalchemy import text + + _LEGACY_TABLES = {"blocked_users", "manual_bans", "temporary_data"} + _IGNORED_USERS_INDEXES = {"ix_users_id", "uq_users_tg_id"} + + def _columns(conn, table_name: str) -> set[str]: + rows = conn.execute( + text( + """ + SELECT column_name + FROM information_schema.columns + WHERE table_schema = 'public' AND table_name = :table_name + """ + ), + {"table_name": table_name}, + ).fetchall() + return {row[0] for row in rows} + + def _pk_columns(conn, table_name: str) -> list[str]: + rows = conn.execute( + text( + """ + SELECT a.attname + FROM pg_constraint c + JOIN pg_class t ON t.oid = c.conrelid + JOIN pg_namespace n ON n.oid = t.relnamespace + JOIN unnest(c.conkey) WITH ORDINALITY AS u(attnum, ord) ON true + JOIN pg_attribute a ON a.attrelid = t.oid AND a.attnum = u.attnum + WHERE n.nspname = 'public' + AND t.relname = :table_name + AND c.contype = 'p' + ORDER BY u.ord + """ + ), + {"table_name": table_name}, + ).fetchall() + return [row[0] for row in rows] + + def _fill_user_id_from_tg(conn, table_name: str, user_col: str, tg_col: str) -> None: + cols = _columns(conn, table_name) + if user_col not in cols or tg_col not in cols: + return + conn.execute( + text( + f""" + UPDATE "{table_name}" AS t + SET "{user_col}" = u.id + FROM users AS u + WHERE t."{user_col}" IS NULL + AND t."{tg_col}" IS NOT NULL + AND t."{tg_col}" = u.tg_id + """ + ) + ) + + def _delete_nulls(conn, table_name: str, col: str) -> None: + cols = _columns(conn, table_name) + if col not in cols: + return + conn.execute(text(f'DELETE FROM "{table_name}" WHERE "{col}" IS NULL')) + + def _prepare_not_null(conn, table_name: str, column_name: str) -> None: + mapping = { + ("notifications", "user_id"): ("user_id", "tg_id"), + ("gift_usages", "user_id"): ("user_id", "tg_id"), + ("blocked_users", "user_id"): ("user_id", "tg_id"), + ("manual_bans", "user_id"): ("user_id", "tg_id"), + ("temporary_data", "user_id"): ("user_id", "tg_id"), + ("scheduled_broadcasts", "created_by_user_id"): ("created_by_user_id", "created_by_tg_id"), + ("gifts", "sender_user_id"): ("sender_user_id", "sender_tg_id"), + } + if (table_name, column_name) in mapping: + user_col, tg_col = mapping[(table_name, column_name)] + _fill_user_id_from_tg(conn, table_name, user_col, tg_col) + _delete_nulls(conn, table_name, user_col) + return + if table_name == "referrals" and column_name == "referred_user_id": + _fill_user_id_from_tg(conn, table_name, "referred_user_id", "referred_tg_id") + _delete_nulls(conn, table_name, "referred_user_id") + if table_name == "referrals" and column_name == "referrer_user_id": + _fill_user_id_from_tg(conn, table_name, "referrer_user_id", "referrer_tg_id") + _delete_nulls(conn, table_name, "referrer_user_id") + + _orig_alter_column = PostgresqlImpl.alter_column + _orig_drop_table = PostgresqlImpl.drop_table + _orig_drop_index = PostgresqlImpl.drop_index + _orig_create_index = PostgresqlImpl.create_index + + def _index_table_name(index) -> str | None: + table = getattr(index, "table", None) + return getattr(table, "name", None) + + if not getattr(PostgresqlImpl.alter_column, "_solo_guarded", False): + def _guarded_alter_column(self, table_name, column_name, *args, **kwargs): + nullable = kwargs.get("nullable") + if nullable is True and column_name in _pk_columns(self.connection, table_name): + return + if table_name in _LEGACY_TABLES and column_name in {"user_id", "tg_id"}: + return + if nullable is False: + _prepare_not_null(self.connection, table_name, column_name) + return _orig_alter_column(self, table_name, column_name, *args, **kwargs) + + _guarded_alter_column._solo_guarded = True + PostgresqlImpl.alter_column = _guarded_alter_column + + if not getattr(PostgresqlImpl.drop_table, "_solo_guarded", False): + def _guarded_drop_table(self, table, **kwargs): + if getattr(table, "name", None) == "schema_migrations": + return + return _orig_drop_table(self, table, **kwargs) + + _guarded_drop_table._solo_guarded = True + PostgresqlImpl.drop_table = _guarded_drop_table + + if not getattr(PostgresqlImpl.drop_index, "_solo_guarded", False): + def _guarded_drop_index(self, index, **kwargs): + table_name = _index_table_name(index) + index_name = getattr(index, "name", None) + if table_name in _LEGACY_TABLES: + return + if table_name == "users" and index_name in _IGNORED_USERS_INDEXES: + return + return _orig_drop_index(self, index, **kwargs) + + _guarded_drop_index._solo_guarded = True + PostgresqlImpl.drop_index = _guarded_drop_index + + if not getattr(PostgresqlImpl.create_index, "_solo_guarded", False): + def _guarded_create_index(self, index, **kwargs): + table_name = _index_table_name(index) + index_name = getattr(index, "name", None) + if table_name in _LEGACY_TABLES: + return + if table_name == "users" and index_name in _IGNORED_USERS_INDEXES: + return + return _orig_create_index(self, index, **kwargs) + + _guarded_create_index._solo_guarded = True + PostgresqlImpl.create_index = _guarded_create_index +except Exception: + pass diff --git a/tests/SMOKE_CHECKLIST.md b/tests/SMOKE_CHECKLIST.md new file mode 100644 index 00000000..577bab7c --- /dev/null +++ b/tests/SMOKE_CHECKLIST.md @@ -0,0 +1,84 @@ +# Smoke Checklist Before Launch + +## 1) Automated checks + +Run from project root: + +```bash +/home/vlad/dev/Solo_bot/venv/bin/python -m compileall /home/vlad/dev/Solo_bot -q +``` + +Run unit tests from writable directory (to avoid log file permission side effects): + +```bash +cd /tmp +PYTHONPATH="/home/vlad/dev/Solo_bot" /home/vlad/dev/Solo_bot/venv/bin/python -m unittest discover -s "/home/vlad/dev/Solo_bot/tests" -q +``` + +Expected result: + +- test suite exits with code `0` +- all tests pass + +## 1.1) Runtime start convention + +- For this project, the canonical full runtime start command is `sudo python3 main.py` +- Do not treat root-owned runtime artifacts in `alembic/` or `/tmp` as a permissions bug by default +- Do not change ownership or permissions during smoke checks unless explicitly requested +- If you need isolated checks, it is still acceptable to run `uvicorn api.main:app` or `npm run dev` separately, but that is not the primary production-like startup path + +## 2) Core migration checks + +- Start app in staging once and verify DB init passes: + - account schema migration + - tg mirror backfill +- Confirm no missing column errors for: + - `tg_id` mirrors in billing-related tables + - `created_by_tg_id` for scheduled broadcasts + +## 3) Identity and actor checks + +Validate these three scenarios end-to-end: + +- `tg-only` user: + - bot flow works + - actor surface resolves as telegram +- `web-only` user: + - auth/login works without Telegram + - billing user is created with internal `users.id` +- `linked` user: + - link Telegram to existing web account + - ensure billing remains on same `users.id` + - Telegram notifications use chat id, not internal id + +## 4) Payment and temporary state checks + +- Create temporary payment state and complete payment via webhook simulation +- Verify: + - temporary state is found/cleared correctly + - payment row uses `user_id` billing relation + - `tg_id` mirror is populated when Telegram exists +- Ensure no message send attempts to internal `users.id` as `chat_id` + +## 5) Gifts and referrals checks + +- Gift creation: + - sender resolves by legacy ref + - gift stores `sender_user_id` and mirror `sender_tg_id` +- Referral creation: + - self-referral blocked + - valid referral stored on billing ids + - mirrors updated where applicable + +## 6) Scheduled broadcasts checks + +- Create broadcast by legacy creator ref +- Verify: + - `created_by_user_id` is resolved when user exists + - `created_by_tg_id` mirror is stored + - listing by `created_by_tg_id` works + +## 7) Runtime warnings to monitor + +- Pydantic warning about `model_custom_emoji_id` is non-blocking but should be cleaned later +- Deprecation warnings for `datetime.utcnow()` are non-blocking but should be migrated to timezone-aware UTC diff --git a/tests/smoke_runner.sh b/tests/smoke_runner.sh new file mode 100755 index 00000000..80fb1165 --- /dev/null +++ b/tests/smoke_runner.sh @@ -0,0 +1,15 @@ +#!/usr/bin/env bash +set -euo pipefail + +ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +PYTHON_BIN="${ROOT_DIR}/venv/bin/python" +TESTS_DIR="${ROOT_DIR}/tests" + +echo "[smoke] compileall" +PYTHONPYCACHEPREFIX="/tmp/solobot_pycache" "${PYTHON_BIN}" -m compileall "${ROOT_DIR}" -q -x "/venv/|/\\.git/|/__pycache__/" + +echo "[smoke] unittest" +cd /tmp +PYTHONPATH="${ROOT_DIR}" "${PYTHON_BIN}" -m unittest discover -s "${TESTS_DIR}" -q + +echo "[smoke] ok" diff --git a/tests/test_actor_scenarios.py b/tests/test_actor_scenarios.py new file mode 100644 index 00000000..e6ed6ea3 --- /dev/null +++ b/tests/test_actor_scenarios.py @@ -0,0 +1,48 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from database.access.resolution import ( + ActorSurface, + resolve_actor_from_identity, + resolve_actor_from_legacy_ref, +) + + +class ActorScenarioTests(unittest.IsolatedAsyncioTestCase): + async def test_tg_only_user_detected_as_telegram_surface(self): + session = object() + user = SimpleNamespace(id=1, tg_id=1001, identity_id="ident-tg") + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=user)): + actor = await resolve_actor_from_legacy_ref(session, 1001) + self.assertEqual(actor.surface, ActorSurface.TELEGRAM) + self.assertEqual(actor.billing_user_id, 1) + self.assertEqual(actor.telegram_chat_id, 1001) + + async def test_web_only_user_detected_as_web_surface(self): + session = object() + user = SimpleNamespace(id=2, tg_id=None, identity_id="ident-web") + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=user)): + actor = await resolve_actor_from_legacy_ref(session, 2) + self.assertEqual(actor.surface, ActorSurface.WEB) + self.assertEqual(actor.billing_user_id, 2) + self.assertIsNone(actor.telegram_chat_id) + + async def test_linked_user_is_web_by_identity_and_telegram_by_tg(self): + session = object() + identity = SimpleNamespace(id="ident-linked") + linked_user = SimpleNamespace(id=3, tg_id=1003, identity_id="ident-linked") + + with patch("database.identities.ensure_billing_user_for_identity", new=AsyncMock(return_value=3)): + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=linked_user)): + web_actor = await resolve_actor_from_identity(session, identity) + + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=linked_user)): + tg_actor = await resolve_actor_from_legacy_ref(session, 1003) + + self.assertEqual(web_actor.surface, ActorSurface.WEB) + self.assertEqual(web_actor.billing_user_id, 3) + self.assertEqual(web_actor.telegram_chat_id, 1003) + self.assertEqual(tg_actor.surface, ActorSurface.TELEGRAM) + self.assertEqual(tg_actor.billing_user_id, 3) + self.assertEqual(tg_actor.telegram_chat_id, 1003) diff --git a/tests/test_api_depends_actor.py b/tests/test_api_depends_actor.py new file mode 100644 index 00000000..d709e388 --- /dev/null +++ b/tests/test_api_depends_actor.py @@ -0,0 +1,52 @@ +import unittest + +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from starlette.requests import Request + +from api.depends import bind_identity_actor, get_request_actor +from database.access.resolution import ActorSurface, ResolvedActor + + +def _make_request() -> Request: + scope = { + "type": "http", + "method": "GET", + "path": "/", + "headers": [], + "query_string": b"", + "client": ("127.0.0.1", 12345), + "server": ("testserver", 80), + "scheme": "http", + "http_version": "1.1", + } + return Request(scope) + + +class ApiDependsActorBindingTests(unittest.IsolatedAsyncioTestCase): + async def test_bind_identity_actor_sets_request_state_and_calls_audit_setter(self): + request = _make_request() + session = object() + identity = SimpleNamespace(id="ident-100") + resolved = ResolvedActor( + surface=ActorSurface.WEB, + billing_user_id=100, + telegram_chat_id=7001, + identity_id="ident-100", + ) + + with patch("api.depends.resolve_actor_from_identity", new=AsyncMock(return_value=resolved)): + with patch("api.depends.set_api_actor") as set_api_actor_mock: + actor = await bind_identity_actor(request, session, identity) + + self.assertEqual(actor, resolved) + self.assertEqual(get_request_actor(request), resolved) + set_api_actor_mock.assert_called_once_with( + request, + identity_id="ident-100", + tg_id=7001, + ) + + async def test_get_request_actor_returns_none_when_request_missing(self): + self.assertIsNone(get_request_actor(None)) diff --git a/tests/test_keys_and_coupons_resolution.py b/tests/test_keys_and_coupons_resolution.py new file mode 100644 index 00000000..8e4a7fc9 --- /dev/null +++ b/tests/test_keys_and_coupons_resolution.py @@ -0,0 +1,167 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock, patch + +from database.coupons import check_coupon_usage, create_coupon_usage, has_any_coupon_usage +from database.keys import store_key + + +class KeysLegacyResolutionTests(unittest.IsolatedAsyncioTestCase): + async def test_store_key_raises_when_user_missing(self): + session = SimpleNamespace(execute=AsyncMock(), add=Mock(), commit=AsyncMock(), rollback=AsyncMock()) + + with patch("database.keys.resolve_user_optional", new=AsyncMock(return_value=None)): + with self.assertRaises(ValueError): + await store_key( + session=session, + legacy_user_ref=9999, + client_id="client_1", + email="u@test", + expiry_time=1111111111111, + key="k", + server_id="s1", + ) + + session.add.assert_not_called() + session.commit.assert_not_called() + + async def test_store_key_creates_with_billing_user_and_tg_mirror(self): + user = SimpleNamespace(id=55, tg_id=5050) + first_query_result = SimpleNamespace(scalar_one_or_none=lambda: None) + session = SimpleNamespace( + execute=AsyncMock(return_value=first_query_result), + add=Mock(), + commit=AsyncMock(), + rollback=AsyncMock(), + ) + + with ( + patch("database.keys.resolve_user_optional", new=AsyncMock(return_value=user)), + patch("database.keys.invalidate_keys_list", new=AsyncMock()), + patch("database.keys.invalidate_key_details", new=AsyncMock()), + patch("database.keys.invalidate_user_snapshot", new=Mock()), + ): + await store_key( + session=session, + legacy_user_ref=5050, + client_id="client_2", + email="x@test", + expiry_time=2222222222222, + key="kk", + server_id="s2", + tariff_id=3, + ) + + session.add.assert_called_once() + added_key = session.add.call_args.args[0] + self.assertEqual(added_key.user_id, 55) + self.assertEqual(added_key.tg_id, 5050) + self.assertEqual(added_key.client_id, "client_2") + session.commit.assert_not_called() + + async def test_store_key_updates_existing_key_with_tg_mirror(self): + user = SimpleNamespace(id=88, tg_id=8080) + existing_key = SimpleNamespace(id=1, client_id="client_3") + first_query_result = SimpleNamespace(scalar_one_or_none=lambda: existing_key) + session = SimpleNamespace( + execute=AsyncMock(side_effect=[first_query_result, SimpleNamespace()]), + add=Mock(), + commit=AsyncMock(), + rollback=AsyncMock(), + ) + + with ( + patch("database.keys.resolve_user_optional", new=AsyncMock(return_value=user)), + patch("database.keys.invalidate_keys_list", new=AsyncMock()), + patch("database.keys.invalidate_key_details", new=AsyncMock()), + patch("database.keys.invalidate_user_snapshot", new=Mock()), + ): + await store_key( + session=session, + legacy_user_ref=8080, + client_id="client_3", + email="upd@test", + expiry_time=3333333333333, + key="new_key", + server_id="s3", + selected_device_limit=5, + current_device_limit=7, + ) + + session.add.assert_not_called() + self.assertEqual(session.execute.await_count, 2) + update_stmt = session.execute.await_args_list[1].args[0] + compiled = update_stmt.compile() + self.assertEqual(compiled.params["email"], "upd@test") + self.assertEqual(compiled.params["tg_id"], 8080) + self.assertEqual(compiled.params["selected_device_limit"], 5) + self.assertEqual(compiled.params["current_device_limit"], 7) + session.commit.assert_not_called() + + +class CouponsLegacyResolutionTests(unittest.IsolatedAsyncioTestCase): + async def test_create_coupon_usage_uses_resolved_user_and_tg_mirror(self): + user = SimpleNamespace(id=77, tg_id=7007) + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + + with patch("database.coupons.resolve_user_optional", new=AsyncMock(return_value=user)): + await create_coupon_usage(session, coupon_id=11, user_id=7007) + + session.execute.assert_awaited_once() + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + self.assertEqual(compiled.params["coupon_id"], 11) + self.assertEqual(compiled.params["user_id"], 77) + self.assertEqual(compiled.params["tg_id"], 7007) + session.commit.assert_not_called() + + async def test_create_coupon_usage_falls_back_to_legacy_when_user_missing(self): + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + + with patch("database.coupons.resolve_user_optional", new=AsyncMock(return_value=None)): + await create_coupon_usage(session, coupon_id=15, user_id=9090) + + session.execute.assert_awaited_once() + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + self.assertEqual(compiled.params["coupon_id"], 15) + self.assertEqual(compiled.params["user_id"], 9090) + self.assertIsNone(compiled.params["tg_id"]) + session.commit.assert_not_called() + + async def test_check_coupon_usage_matches_by_billing_or_tg(self): + user = SimpleNamespace(id=41, tg_id=4141) + session = SimpleNamespace( + execute=AsyncMock(return_value=SimpleNamespace(scalar_one_or_none=lambda: object())) + ) + + with patch("database.coupons.resolve_user_optional", new=AsyncMock(return_value=user)): + used = await check_coupon_usage(session, coupon_id=5, legacy_user_ref=4141) + + self.assertTrue(used) + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + self.assertEqual(compiled.params["coupon_id_1"], 5) + self.assertEqual(compiled.params["user_id_1"], 41) + self.assertEqual(compiled.params["tg_id_1"], 4141) + + async def test_has_any_coupon_usage_returns_false_when_no_rows(self): + user = SimpleNamespace(id=50, tg_id=5050) + session = SimpleNamespace(execute=AsyncMock(return_value=SimpleNamespace(first=lambda: None))) + + with patch("database.coupons.resolve_user_optional", new=AsyncMock(return_value=user)): + has_usage = await has_any_coupon_usage(session, legacy_user_ref=5050) + + self.assertFalse(has_usage) + + async def test_has_any_coupon_usage_fallbacks_to_legacy_when_user_missing(self): + session = SimpleNamespace(execute=AsyncMock(return_value=SimpleNamespace(first=lambda: (1,)))) + + with patch("database.coupons.resolve_user_optional", new=AsyncMock(return_value=None)): + has_usage = await has_any_coupon_usage(session, legacy_user_ref=6060) + + self.assertTrue(has_usage) + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + self.assertEqual(compiled.params["user_id_1"], 6060) + self.assertEqual(compiled.params["tg_id_1"], 6060) diff --git a/tests/test_new_repo_functions.py b/tests/test_new_repo_functions.py new file mode 100644 index 00000000..0f2389c3 --- /dev/null +++ b/tests/test_new_repo_functions.py @@ -0,0 +1,317 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from database.gifts import ( + count_gift_usages, + get_gift_locked, + get_gift_usage, + mark_gift_fully_redeemed, + record_gift_usage, +) +from database.keys import ( + count_active_keys_for_user, + count_keys_by_server_id, + delete_key_by_user_and_email, + get_all_key_server_ids, + get_key_by_user_and_email, + get_key_client_id_by_email_and_server, + get_user_keys_with_servers_by_email, + update_key_post_creation_snapshot, + update_key_renewal_snapshot, +) +from database.payments import count_successful_payments +from database.servers import ( + cluster_name_exists, + get_cluster_name_for_server_name, + get_enabled_server_subscription_url, + get_panel_type_for_server, + get_panel_types_for_cluster, +) +from database.tariffs import get_active_tariff_by_id +from database.users import ( + get_user_preferred_currency, + mark_trial_started_if_eligible, +) + + +def _make_session(*, execute_return=None, execute_side_effect=None): + """Собирает мок-session с правильно настроенным execute и отсутствующим commit/rollback. + + Если хендлер вызовет `session.commit()` — тест упадёт, потому что мы НЕ + добавляем commit в namespace. + """ + ns = SimpleNamespace() + if execute_side_effect is not None: + ns.execute = AsyncMock(side_effect=execute_side_effect) + else: + ns.execute = AsyncMock(return_value=execute_return) + return ns + + +_MISSING = object() + + +def _result_with( + *, + scalar=_MISSING, + scalar_one=_MISSING, + scalar_one_or_none=_MISSING, + all_rows=_MISSING, + mappings_all=_MISSING, +): + """Фейковый результат session.execute с нужными методами. + + Использует sentinel чтобы отличить "не передано" от "передано как None". + """ + obj = SimpleNamespace() + if scalar is not _MISSING: + obj.scalar = lambda v=scalar: v + if scalar_one is not _MISSING: + obj.scalar_one = lambda v=scalar_one: v + if scalar_one_or_none is not _MISSING: + obj.scalar_one_or_none = lambda v=scalar_one_or_none: v + if all_rows is not _MISSING: + obj.all = lambda v=all_rows: v + if mappings_all is not _MISSING: + obj.mappings = lambda v=mappings_all: SimpleNamespace(all=lambda: v) + return obj + + +class GiftsRepoTests(unittest.IsolatedAsyncioTestCase): + async def test_get_gift_locked_returns_orm_row(self): + gift = SimpleNamespace(gift_id="gift_x") + session = _make_session(execute_return=_result_with(scalar_one_or_none=gift)) + result = await get_gift_locked(session, "gift_x") + self.assertIs(result, gift) + session.execute.assert_awaited_once() + + async def test_get_gift_locked_returns_none_when_missing(self): + session = _make_session(execute_return=_result_with(scalar_one_or_none=None)) + result = await get_gift_locked(session, "missing") + self.assertIsNone(result) + + async def test_get_gift_usage_passes_composite_key(self): + usage = SimpleNamespace() + session = _make_session(execute_return=_result_with(scalar_one_or_none=usage)) + result = await get_gift_usage(session, "gift_1", 42) + self.assertIs(result, usage) + session.execute.assert_awaited_once() + + async def test_count_gift_usages_returns_int(self): + session = _make_session(execute_return=_result_with(scalar_one=7)) + count = await count_gift_usages(session, "gift_y") + self.assertEqual(count, 7) + + async def test_count_gift_usages_handles_none(self): + session = _make_session(execute_return=_result_with(scalar_one=None)) + count = await count_gift_usages(session, "gift_y") + self.assertEqual(count, 0) + + async def test_record_gift_usage_executes_insert(self): + session = _make_session(execute_return=_result_with()) + await record_gift_usage(session, "gift_z", user_id=5, tg_id=55) + session.execute.assert_awaited_once() + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + self.assertEqual(compiled.params["gift_id"], "gift_z") + self.assertEqual(compiled.params["user_id"], 5) + self.assertEqual(compiled.params["tg_id"], 55) + + async def test_mark_gift_fully_redeemed_sets_is_used(self): + session = _make_session(execute_return=_result_with()) + await mark_gift_fully_redeemed(session, "gift_a", recipient_user_id=9, recipient_tg_id=99) + session.execute.assert_awaited_once() + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + self.assertTrue(compiled.params["is_used"]) + self.assertEqual(compiled.params["recipient_user_id"], 9) + self.assertEqual(compiled.params["recipient_tg_id"], 99) + + +class KeysRepoTests(unittest.IsolatedAsyncioTestCase): + async def test_count_active_keys_for_user_filters_frozen(self): + session = _make_session(execute_return=_result_with(scalar=3)) + result = await count_active_keys_for_user(session, 42) + self.assertEqual(result, 3) + stmt = session.execute.await_args.args[0] + + compiled = str(stmt.compile()) + self.assertIn("is_frozen", compiled) + + async def test_count_keys_by_server_id_returns_int(self): + session = _make_session(execute_return=_result_with(scalar=10)) + result = await count_keys_by_server_id(session, "cluster-a") + self.assertEqual(result, 10) + + async def test_get_all_key_server_ids_returns_only_strings(self): + rows = [("srv1",), ("srv2",), (None,), ("srv3",)] + session = _make_session(execute_return=_result_with(all_rows=rows)) + result = await get_all_key_server_ids(session) + self.assertEqual(result, ["srv1", "srv2", "srv3"]) + + async def test_get_key_by_user_and_email_returns_orm_row(self): + key = SimpleNamespace(email="u@test") + session = _make_session(execute_return=_result_with(scalar_one_or_none=key)) + result = await get_key_by_user_and_email(session, 42, "u@test") + self.assertIs(result, key) + + async def test_delete_key_by_user_and_email_executes_delete(self): + session = _make_session(execute_return=_result_with()) + await delete_key_by_user_and_email(session, 42, "u@test") + session.execute.assert_awaited_once() + + async def test_get_key_client_id_by_email_and_server(self): + session = _make_session(execute_return=_result_with(scalar="client-123")) + result = await get_key_client_id_by_email_and_server(session, "u@test", "cluster-a") + self.assertEqual(result, "client-123") + + async def test_update_key_renewal_snapshot_without_limits(self): + session = _make_session(execute_return=_result_with()) + with patch("database.keys.invalidate_key_details", new=AsyncMock()): + await update_key_renewal_snapshot(session, "u@test", tariff_id=5, apply_limits=False) + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + self.assertEqual(compiled.params["tariff_id"], 5) + self.assertNotIn("selected_device_limit", compiled.params) + + async def test_update_key_renewal_snapshot_with_limits(self): + session = _make_session(execute_return=_result_with()) + with patch("database.keys.invalidate_key_details", new=AsyncMock()): + await update_key_renewal_snapshot( + session, + "u@test", + tariff_id=5, + selected_device_limit=3, + current_device_limit=3, + selected_traffic_limit=50, + current_traffic_limit=50, + apply_limits=True, + ) + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + self.assertEqual(compiled.params["tariff_id"], 5) + self.assertEqual(compiled.params["selected_device_limit"], 3) + self.assertEqual(compiled.params["selected_traffic_limit"], 50) + + async def test_update_key_post_creation_snapshot(self): + session = _make_session(execute_return=_result_with()) + with patch("database.keys.invalidate_key_details", new=AsyncMock()): + await update_key_post_creation_snapshot( + session, + user_id=10, + email="u@test", + selected_device_limit=2, + selected_traffic_limit=100, + selected_price_rub=500, + ) + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + self.assertEqual(compiled.params["selected_device_limit"], 2) + self.assertEqual(compiled.params["selected_traffic_limit"], 100) + self.assertEqual(compiled.params["selected_price_rub"], 500) + + async def test_get_user_keys_with_servers_returns_tuples(self): + srv = SimpleNamespace( + server_name="s1", cluster_name="c1", api_url="http://x", panel_type="3x-ui" + ) + session = _make_session( + execute_return=_result_with(all_rows=[("cid1", "s1", srv)]) + ) + result = await get_user_keys_with_servers_by_email(session, 42, "u@test") + self.assertEqual(len(result), 1) + client_id, server_id, server_info = result[0] + self.assertEqual(client_id, "cid1") + self.assertEqual(server_id, "s1") + self.assertEqual(server_info["server_name"], "s1") + self.assertEqual(server_info["panel_type"], "3x-ui") + + +class ServersRepoTests(unittest.IsolatedAsyncioTestCase): + async def test_cluster_name_exists_returns_true(self): + result_obj = SimpleNamespace(scalars=lambda: SimpleNamespace(first=lambda: "row")) + session = _make_session(execute_return=result_obj) + ok = await cluster_name_exists(session, "cluster-x") + self.assertTrue(ok) + + async def test_cluster_name_exists_returns_false(self): + result_obj = SimpleNamespace(scalars=lambda: SimpleNamespace(first=lambda: None)) + session = _make_session(execute_return=result_obj) + ok = await cluster_name_exists(session, "no-such") + self.assertFalse(ok) + + async def test_get_cluster_name_for_server_name(self): + session = _make_session(execute_return=_result_with(scalar="cluster-a")) + result = await get_cluster_name_for_server_name(session, "srv1") + self.assertEqual(result, "cluster-a") + + async def test_get_enabled_server_subscription_url(self): + session = _make_session(execute_return=_result_with(scalar="https://sub/x")) + result = await get_enabled_server_subscription_url(session, "srv1") + self.assertEqual(result, "https://sub/x") + + async def test_get_panel_types_for_cluster(self): + scalars_mock = SimpleNamespace(all=lambda: ["remnawave", "remnawave"]) + result_obj = SimpleNamespace(scalars=lambda: scalars_mock) + session = _make_session(execute_return=result_obj) + result = await get_panel_types_for_cluster(session, "cluster-a") + self.assertEqual(result, ["remnawave", "remnawave"]) + + async def test_get_panel_type_for_server(self): + session = _make_session(execute_return=_result_with(scalar_one_or_none="3x-ui")) + result = await get_panel_type_for_server(session, "srv1") + self.assertEqual(result, "3x-ui") + + +class TariffsRepoTests(unittest.IsolatedAsyncioTestCase): + async def test_get_active_tariff_by_id_returns_only_active(self): + tariff = SimpleNamespace(id=5, is_active=True) + session = _make_session(execute_return=_result_with(scalar_one_or_none=tariff)) + result = await get_active_tariff_by_id(session, 5) + self.assertIs(result, tariff) + stmt = session.execute.await_args.args[0] + compiled = str(stmt.compile()) + self.assertIn("is_active", compiled) + + async def test_get_active_tariff_by_id_none_when_missing(self): + session = _make_session(execute_return=_result_with(scalar_one_or_none=None)) + result = await get_active_tariff_by_id(session, 999) + self.assertIsNone(result) + + +class PaymentsRepoTests(unittest.IsolatedAsyncioTestCase): + async def test_count_successful_payments_returns_int(self): + session = _make_session(execute_return=_result_with(scalar=2)) + result = await count_successful_payments(session, 42) + self.assertEqual(result, 2) + + async def test_count_successful_payments_handles_none(self): + session = _make_session(execute_return=_result_with(scalar=None)) + result = await count_successful_payments(session, 42) + self.assertEqual(result, 0) + + +class UsersRepoTests(unittest.IsolatedAsyncioTestCase): + async def test_mark_trial_started_if_eligible_emits_conditional_update(self): + session = _make_session(execute_return=_result_with()) + await mark_trial_started_if_eligible(session, 1234) + session.execute.assert_awaited_once() + stmt = session.execute.await_args.args[0] + compiled = str(stmt.compile()) + + self.assertIn("tg_id", compiled) + self.assertIn("trial IN", compiled.replace("trial in", "trial IN")) + + async def test_get_user_preferred_currency_returns_scalar(self): + session = _make_session(execute_return=_result_with(scalar="USD")) + result = await get_user_preferred_currency(session, 1234) + self.assertEqual(result, "USD") + + async def test_get_user_preferred_currency_returns_none_when_unset(self): + session = _make_session(execute_return=_result_with(scalar=None)) + result = await get_user_preferred_currency(session, 1234) + self.assertIsNone(result) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_notifications_resolution.py b/tests/test_notifications_resolution.py new file mode 100644 index 00000000..c0851503 --- /dev/null +++ b/tests/test_notifications_resolution.py @@ -0,0 +1,90 @@ +import unittest +from datetime import UTC, datetime, timedelta +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from database.notifications import add_notification, check_notification_time, check_notification_time_bulk + + +class NotificationsResolutionTests(unittest.IsolatedAsyncioTestCase): + async def test_add_notification_uses_resolved_user_and_tg_mirror(self): + user = SimpleNamespace(id=12, tg_id=1212) + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + + with patch("database.notifications.resolve_user_optional", new=AsyncMock(return_value=user)): + await add_notification(session, legacy_user_ref=1212, notification_type="n1") + + session.execute.assert_awaited_once() + stmt = session.execute.await_args.args[0] + compiled = stmt.compile() + values = set(compiled.params.values()) + self.assertIn(12, values) + self.assertIn(1212, values) + self.assertIn("n1", values) + session.commit.assert_not_called() + + async def test_check_notification_time_returns_true_when_user_missing(self): + session = SimpleNamespace(execute=AsyncMock()) + + with patch("database.notifications.resolve_user_optional", new=AsyncMock(return_value=None)): + allowed = await check_notification_time(session, legacy_user_ref=7777, notification_type="n2", hours=12) + + self.assertTrue(allowed) + session.execute.assert_not_called() + + async def test_check_notification_time_handles_naive_db_timestamp(self): + user = SimpleNamespace(id=33, tg_id=3333) + now = datetime(2026, 3, 20, 20, 0, 0, tzinfo=UTC) + old_naive = (now - timedelta(hours=13)).replace(tzinfo=None) + recent_naive = (now - timedelta(hours=2)).replace(tzinfo=None) + + old_session = SimpleNamespace( + execute=AsyncMock(return_value=SimpleNamespace(scalar_one_or_none=lambda: old_naive)) + ) + recent_session = SimpleNamespace( + execute=AsyncMock(return_value=SimpleNamespace(scalar_one_or_none=lambda: recent_naive)) + ) + + with ( + patch("database.notifications.resolve_user_optional", new=AsyncMock(return_value=user)), + patch("database.notifications._utc_now", return_value=now), + ): + old_allowed = await check_notification_time(old_session, legacy_user_ref=3333, notification_type="n3", hours=12) + recent_allowed = await check_notification_time( + recent_session, legacy_user_ref=3333, notification_type="n3", hours=12 + ) + + self.assertTrue(old_allowed) + self.assertFalse(recent_allowed) + + async def test_check_notification_time_bulk_includes_missing_and_old(self): + now = datetime(2026, 3, 20, 20, 0, 0, tzinfo=UTC) + session = SimpleNamespace() + session.execute = AsyncMock( + return_value=[ + SimpleNamespace( + user_id=1, + notification_type="n1", + last_notification_time=now - timedelta(hours=1), + ), + SimpleNamespace( + user_id=2, + notification_type="n2", + last_notification_time=(now - timedelta(hours=13)).replace(tzinfo=None), + ), + ] + ) + items = [(100, "n1"), (200, "n2"), (300, "n3")] + + with ( + patch("database.notifications._utc_now", return_value=now), + patch( + "database.notifications._map_legacy_refs_to_user_ids", + new=AsyncMock(return_value={100: 1, 200: 2}), + ), + ): + result = await check_notification_time_bulk(session, items=items, hours=12) + + self.assertNotIn((100, "n1"), result) + self.assertIn((200, "n2"), result) + self.assertIn((300, "n3"), result) diff --git a/tests/test_notify_and_actor_middleware.py b/tests/test_notify_and_actor_middleware.py new file mode 100644 index 00000000..8d5c6dab --- /dev/null +++ b/tests/test_notify_and_actor_middleware.py @@ -0,0 +1,70 @@ +import unittest +from importlib.util import module_from_spec, spec_from_file_location +from pathlib import Path +import sys +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from database.access.resolution import notify_telegram_chat_id + +_ACTOR_SPEC = spec_from_file_location( + "actor_module_for_tests", + str(Path(__file__).resolve().parents[1] / "middlewares" / "actor.py"), +) +_ACTOR_MODULE = module_from_spec(_ACTOR_SPEC) +assert _ACTOR_SPEC is not None and _ACTOR_SPEC.loader is not None +sys.modules["actor_module_for_tests"] = _ACTOR_MODULE +_ACTOR_SPEC.loader.exec_module(_ACTOR_MODULE) +ActorMiddleware = _ACTOR_MODULE.ActorMiddleware + + +class NotifyTelegramChatIdTests(unittest.IsolatedAsyncioTestCase): + async def test_returns_user_tg_when_user_has_telegram(self): + session = object() + user = SimpleNamespace(tg_id=555) + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=user)): + value = await notify_telegram_chat_id(session, 100) + self.assertEqual(value, 555) + + async def test_returns_legacy_ref_when_user_missing(self): + session = object() + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=None)): + value = await notify_telegram_chat_id(session, 777) + self.assertEqual(value, 777) + + async def test_returns_none_when_user_exists_without_tg(self): + session = object() + user = SimpleNamespace(tg_id=None) + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=user)): + value = await notify_telegram_chat_id(session, 11) + self.assertIsNone(value) + + +class ActorMiddlewareTests(unittest.IsolatedAsyncioTestCase): + async def test_sets_actor_for_non_bot_user(self): + middleware = ActorMiddleware() + from_user = SimpleNamespace(id=123, is_bot=False) + data = {"event_from_user": from_user, "session": SimpleNamespace(execute=object())} + + async def handler(event, payload): + return payload.get("actor") + + resolved_actor = SimpleNamespace(surface="telegram", billing_user_id=10, telegram_chat_id=123, identity_id=None) + with patch("actor_module_for_tests.resolve_actor_from_legacy_ref", new=AsyncMock(return_value=resolved_actor)): + result = await middleware(handler, object(), data) + + self.assertEqual(result, resolved_actor) + self.assertEqual(data.get("actor"), resolved_actor) + + async def test_skips_actor_when_event_user_missing(self): + middleware = ActorMiddleware() + data = {"session": object()} + + async def handler(event, payload): + return payload.get("actor") + + with patch("actor_module_for_tests.resolve_actor_from_legacy_ref", new=AsyncMock()) as resolver_mock: + result = await middleware(handler, object(), data) + + self.assertIsNone(result) + resolver_mock.assert_not_called() diff --git a/tests/test_payments_and_temporary_data.py b/tests/test_payments_and_temporary_data.py new file mode 100644 index 00000000..8877928a --- /dev/null +++ b/tests/test_payments_and_temporary_data.py @@ -0,0 +1,140 @@ +import unittest +from datetime import datetime, timedelta +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from database.payments import add_payment +from database.temporary_data import clear_temporary_data, create_temporary_data, get_temporary_data + + +class PaymentsLegacyResolutionTests(unittest.IsolatedAsyncioTestCase): + async def test_add_payment_raises_when_user_not_found(self): + session = SimpleNamespace(execute=AsyncMock()) + + with patch("database.payments.resolve_user_optional", new=AsyncMock(return_value=None)): + with self.assertRaises(ValueError): + await add_payment( + session=session, + legacy_user_ref=999999, + amount=100.0, + payment_system="test", + ) + + session.execute.assert_not_called() + + async def test_add_payment_uses_resolved_user_and_returns_internal_id(self): + user = SimpleNamespace(id=77, tg_id=5005) + result = SimpleNamespace(scalar_one=lambda: 1234) + session = SimpleNamespace(execute=AsyncMock(return_value=result)) + + with patch("database.payments.resolve_user_optional", new=AsyncMock(return_value=user)): + internal_id = await add_payment( + session=session, + legacy_user_ref=5005, + amount=50.0, + payment_system="test", + status="success", + ) + + self.assertEqual(internal_id, 1234) + session.execute.assert_awaited_once() + + async def test_add_payment_accepts_tg_id_alias_keyword(self): + user = SimpleNamespace(id=88, tg_id=8080) + result = SimpleNamespace(scalar_one=lambda: 555) + session = SimpleNamespace(execute=AsyncMock(return_value=result)) + + with patch("database.payments.resolve_user_optional", new=AsyncMock(return_value=user)) as resolve_mock: + internal_id = await add_payment( + session=session, + tg_id=8080, + amount=99.0, + payment_system="alias-test", + ) + + self.assertEqual(internal_id, 555) + resolve_mock.assert_awaited_once_with(session, 8080) + + +class TemporaryDataLegacyResolutionTests(unittest.IsolatedAsyncioTestCase): + async def test_create_temporary_data_raises_when_user_missing(self): + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + + with patch("database.temporary_data.resolve_user_optional", new=AsyncMock(return_value=None)): + with self.assertRaises(ValueError): + await create_temporary_data(session, legacy_user_ref=1010, state="s", data={"a": 1}) + + session.execute.assert_not_called() + session.commit.assert_not_called() + + async def test_create_temporary_data_executes_when_user_found(self): + user = SimpleNamespace(id=12, tg_id=1200) + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + + with patch("database.temporary_data.resolve_user_optional", new=AsyncMock(return_value=user)): + await create_temporary_data(session, legacy_user_ref=1200, state="waiting", data={"x": 1}) + + session.execute.assert_awaited_once() + session.commit.assert_not_called() + + async def test_get_temporary_data_falls_back_to_tg_when_user_missing(self): + row = SimpleNamespace(state="waiting_for_payment", data={"required_amount": 100}) + session = SimpleNamespace(execute=AsyncMock(return_value=SimpleNamespace(scalar_one_or_none=lambda: row))) + + with patch("database.temporary_data.resolve_user_optional", new=AsyncMock(return_value=None)): + data = await get_temporary_data(session, legacy_user_ref=4321) + + self.assertEqual(data, {"state": "waiting_for_payment", "data": {"required_amount": 100}}) + + async def test_clear_temporary_data_uses_resolved_user(self): + user = SimpleNamespace(id=222, tg_id=22) + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock()) + + with patch("database.temporary_data.resolve_user_optional", new=AsyncMock(return_value=user)): + await clear_temporary_data(session, legacy_user_ref=22) + + session.execute.assert_awaited_once() + session.commit.assert_not_called() + + +class PaymentRenewalTimingTests(unittest.IsolatedAsyncioTestCase): + async def test_renewal_recomputes_expiry_from_payment_time_when_key_expired(self): + from handlers.payments.utils import _handle_temp_state + + fixed_now = datetime(2026, 1, 10, 12, 0, 0) + expired_at = int((fixed_now - timedelta(days=2)).timestamp() * 1000) + expected_new_expiry = int((fixed_now + timedelta(days=30)).timestamp() * 1000) + + session = SimpleNamespace() + data = { + "tariff_id": 7, + "client_id": "client-1", + "cost": 100, + "email": "user@example.com", + "new_expiry_time": expired_at + int(timedelta(days=30).total_seconds() * 1000), + "selected_duration_days": 30, + } + + with ( + patch("handlers.payments.utils.get_balance", new=AsyncMock(return_value=500)), + patch( + "handlers.payments.utils.get_key_by_server", + new=AsyncMock(return_value=SimpleNamespace(expiry_time=expired_at)), + ), + patch("handlers.keys.key_renew.complete_key_renewal", new=AsyncMock()) as complete_mock, + patch("handlers.payments.utils.clear_temporary_data", new=AsyncMock()) as clear_mock, + patch("handlers.payments.utils.datetime") as datetime_mock, + ): + datetime_mock.utcnow.return_value = fixed_now + + handled = await _handle_temp_state( + session=session, + user_id=12345, + state="waiting_for_renewal_payment", + data=data, + amount=100, + ) + + self.assertTrue(handled) + self.assertEqual(complete_mock.await_args.kwargs["new_expiry_time"], expected_new_expiry) + clear_mock.assert_awaited_once_with(session, 12345) diff --git a/tests/test_payments_pipeline.py b/tests/test_payments_pipeline.py new file mode 100644 index 00000000..c4cec069 --- /dev/null +++ b/tests/test_payments_pipeline.py @@ -0,0 +1,305 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from services.payments.pipeline import ( + ParsedPayment, + PipelineResult, + process_cancelled_payment, + process_success_payment, +) + + +class _FakeSessionContext: + """Async context manager, который возвращает переданный session при __aenter__.""" + + def __init__(self, session): + self._session = session + + async def __aenter__(self): + return self._session + + async def __aexit__(self, exc_type, exc, tb): + return False + + +def _sessionmaker_returning(session): + """Фейковый ``async_session_maker`` — CallableType -> контекст.""" + return lambda: _FakeSessionContext(session) + + +class ProcessSuccessPaymentTests(unittest.IsolatedAsyncioTestCase): + async def test_idempotent_when_already_success(self): + """Повторный webhook: payment.status=='success' → сразу return already_processed.""" + session = SimpleNamespace(commit=AsyncMock()) + parsed = ParsedPayment(payment_id="p1", tg_id=42, amount=500.0) + + with ( + patch("services.payments.pipeline.async_session_maker", _sessionmaker_returning(session)), + patch( + "services.payments.pipeline.get_payment_by_payment_id", + new=AsyncMock(return_value={"id": 1, "tg_id": 42, "status": "success"}), + ), + patch("services.payments.pipeline.update_payment_status", new=AsyncMock()) as upd_mock, + patch("services.payments.pipeline.add_payment", new=AsyncMock()) as add_mock, + patch("services.payments.pipeline.update_balance", new=AsyncMock()) as balance_mock, + patch( + "services.payments.pipeline.send_payment_success_notification", + new=AsyncMock(), + ) as notify_mock, + patch("services.payments.pipeline.invalidate_payment_cache", new=AsyncMock()), + ): + result = await process_success_payment("yookassa", parsed) + + self.assertTrue(result.ok) + self.assertTrue(result.already_processed) + upd_mock.assert_not_awaited() + add_mock.assert_not_awaited() + balance_mock.assert_not_awaited() + notify_mock.assert_not_awaited() + session.commit.assert_not_awaited() + + async def test_pending_to_success(self): + """Payment pending → update_payment_status(success) + balance + notify.""" + session = SimpleNamespace(commit=AsyncMock()) + parsed = ParsedPayment(payment_id="p1", tg_id=42, amount=500.0) + + with ( + patch("services.payments.pipeline.async_session_maker", _sessionmaker_returning(session)), + patch( + "services.payments.pipeline.get_payment_by_payment_id", + new=AsyncMock(return_value={"id": 7, "tg_id": 42, "status": "pending"}), + ), + patch( + "services.payments.pipeline.update_payment_status", + new=AsyncMock(return_value=True), + ) as upd_mock, + patch("services.payments.pipeline.add_payment", new=AsyncMock()) as add_mock, + patch("services.payments.pipeline.update_balance", new=AsyncMock()) as balance_mock, + patch( + "services.payments.pipeline.send_payment_success_notification", + new=AsyncMock(), + ) as notify_mock, + patch( + "services.payments.pipeline.invalidate_payment_cache", + new=AsyncMock(), + ) as cache_mock, + ): + result = await process_success_payment("robokassa", parsed) + + self.assertTrue(result.ok) + self.assertFalse(result.already_processed) + upd_mock.assert_awaited_once() + add_mock.assert_not_awaited() + balance_mock.assert_awaited_once_with(session, 42, 500.0) + notify_mock.assert_awaited_once() + session.commit.assert_awaited_once() + cache_mock.assert_awaited_once_with("p1") + + async def test_fresh_payment_uses_add_payment(self): + """Нет pending-записи → add_payment создаёт новую.""" + session = SimpleNamespace(commit=AsyncMock()) + parsed = ParsedPayment( + payment_id="p2", tg_id=100, amount=750.0, currency="RUB", metadata={"k": "v"} + ) + + with ( + patch("services.payments.pipeline.async_session_maker", _sessionmaker_returning(session)), + patch( + "services.payments.pipeline.get_payment_by_payment_id", + new=AsyncMock(return_value=None), + ), + patch( + "services.payments.pipeline.update_payment_status", new=AsyncMock() + ) as upd_mock, + patch("services.payments.pipeline.add_payment", new=AsyncMock()) as add_mock, + patch("services.payments.pipeline.update_balance", new=AsyncMock()) as balance_mock, + patch( + "services.payments.pipeline.send_payment_success_notification", + new=AsyncMock(), + ), + patch("services.payments.pipeline.invalidate_payment_cache", new=AsyncMock()), + ): + result = await process_success_payment("kassai", parsed) + + self.assertTrue(result.ok) + upd_mock.assert_not_awaited() + add_mock.assert_awaited_once() + add_kwargs = add_mock.await_args.kwargs + self.assertEqual(add_kwargs["tg_id"], 100) + self.assertEqual(add_kwargs["amount"], 750.0) + self.assertEqual(add_kwargs["payment_system"], "kassai") + self.assertEqual(add_kwargs["status"], "success") + self.assertEqual(add_kwargs["payment_id"], "p2") + self.assertEqual(add_kwargs["metadata"], {"k": "v"}) + balance_mock.assert_awaited_once_with(session, 100, 750.0) + + async def test_update_status_failure_returns_error(self): + """update_payment_status вернул False → PipelineResult(ok=False).""" + session = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + parsed = ParsedPayment(payment_id="p3", tg_id=1, amount=1.0) + + with ( + patch("services.payments.pipeline.async_session_maker", _sessionmaker_returning(session)), + patch( + "services.payments.pipeline.get_payment_by_payment_id", + new=AsyncMock(return_value={"id": 5, "tg_id": 1, "status": "pending"}), + ), + patch( + "services.payments.pipeline.update_payment_status", + new=AsyncMock(return_value=False), + ), + patch("services.payments.pipeline.update_balance", new=AsyncMock()) as balance_mock, + patch("services.payments.pipeline.send_payment_success_notification", new=AsyncMock()), + patch("services.payments.pipeline.invalidate_payment_cache", new=AsyncMock()), + ): + result = await process_success_payment("yoomoney", parsed) + + self.assertFalse(result.ok) + self.assertIsNotNone(result.error) + balance_mock.assert_not_awaited() + session.commit.assert_not_awaited() + + async def test_credit_amount_override(self): + """cryptobot: parsed.amount=USDT-сумма, credit_amount_override=RUB.""" + session = SimpleNamespace(commit=AsyncMock()) + parsed = ParsedPayment(payment_id="p4", tg_id=42, amount=10.0, currency="USDT") + + with ( + patch("services.payments.pipeline.async_session_maker", _sessionmaker_returning(session)), + patch( + "services.payments.pipeline.get_payment_by_payment_id", + new=AsyncMock(return_value={"id": 9, "tg_id": 42, "status": "pending"}), + ), + patch( + "services.payments.pipeline.update_payment_status", + new=AsyncMock(return_value=True), + ), + patch("services.payments.pipeline.update_balance", new=AsyncMock()) as balance_mock, + patch("services.payments.pipeline.send_payment_success_notification", new=AsyncMock()) as notify_mock, + patch("services.payments.pipeline.invalidate_payment_cache", new=AsyncMock()), + + patch( + "services.payments.pipeline.select", + return_value=SimpleNamespace(where=lambda *a, **k: SimpleNamespace(limit=lambda n: None)), + ), + ): + + fake_payment_obj = SimpleNamespace(currency=None, original_amount=None) + session.execute = AsyncMock( + return_value=SimpleNamespace(scalar_one_or_none=lambda: fake_payment_obj) + ) + result = await process_success_payment( + "CRYPTOBOT", + parsed, + credit_amount_override=900.0, + update_currency="USDT", + update_original_amount=10.0, + ) + + self.assertTrue(result.ok) + + balance_mock.assert_awaited_once_with(session, 42, 900.0) + notify_mock.assert_awaited_once_with(42, 900.0, session) + + self.assertEqual(fake_payment_obj.currency, "USDT") + self.assertEqual(fake_payment_obj.original_amount, 10.0) + + async def test_metadata_patch_passed_to_update(self): + """metadata_patch пробрасывается в update_payment_status.""" + session = SimpleNamespace(commit=AsyncMock()) + parsed = ParsedPayment(payment_id="p5", tg_id=42, amount=500.0) + + with ( + patch("services.payments.pipeline.async_session_maker", _sessionmaker_returning(session)), + patch( + "services.payments.pipeline.get_payment_by_payment_id", + new=AsyncMock(return_value={"id": 11, "tg_id": 42, "status": "pending"}), + ), + patch( + "services.payments.pipeline.update_payment_status", + new=AsyncMock(return_value=True), + ) as upd_mock, + patch("services.payments.pipeline.update_balance", new=AsyncMock()), + patch("services.payments.pipeline.send_payment_success_notification", new=AsyncMock()), + patch("services.payments.pipeline.invalidate_payment_cache", new=AsyncMock()), + ): + await process_success_payment( + "CRYPTOBOT", parsed, metadata_patch={"fx": {"rate": 90.5}} + ) + + upd_mock.assert_awaited_once() + self.assertEqual( + upd_mock.await_args.kwargs["metadata_patch"], {"fx": {"rate": 90.5}} + ) + + +class ProcessCancelledPaymentTests(unittest.IsolatedAsyncioTestCase): + async def test_idempotent_when_already_cancelled(self): + session = SimpleNamespace(commit=AsyncMock()) + parsed = ParsedPayment(payment_id="p1", tg_id=42, amount=500.0) + + with ( + patch("services.payments.pipeline.async_session_maker", _sessionmaker_returning(session)), + patch( + "services.payments.pipeline.get_payment_by_payment_id", + new=AsyncMock(return_value={"id": 1, "status": "cancelled"}), + ), + patch("services.payments.pipeline.update_payment_status", new=AsyncMock()) as upd, + patch("services.payments.pipeline.add_payment", new=AsyncMock()), + patch("services.payments.pipeline.invalidate_payment_cache", new=AsyncMock()), + ): + result = await process_cancelled_payment("yookassa", parsed) + + self.assertTrue(result.ok) + self.assertTrue(result.already_processed) + upd.assert_not_awaited() + session.commit.assert_not_awaited() + + async def test_pending_to_cancelled(self): + session = SimpleNamespace(commit=AsyncMock()) + parsed = ParsedPayment(payment_id="p1", tg_id=42, amount=500.0) + + with ( + patch("services.payments.pipeline.async_session_maker", _sessionmaker_returning(session)), + patch( + "services.payments.pipeline.get_payment_by_payment_id", + new=AsyncMock(return_value={"id": 1, "status": "pending"}), + ), + patch( + "services.payments.pipeline.update_payment_status", + new=AsyncMock(return_value=True), + ) as upd, + patch("services.payments.pipeline.invalidate_payment_cache", new=AsyncMock()), + ): + result = await process_cancelled_payment("yookassa", parsed) + + self.assertTrue(result.ok) + upd.assert_awaited_once() + session.commit.assert_awaited_once() + + async def test_fresh_cancelled_uses_add_payment(self): + session = SimpleNamespace(commit=AsyncMock()) + parsed = ParsedPayment(payment_id="p1", tg_id=42, amount=0.0) + + with ( + patch("services.payments.pipeline.async_session_maker", _sessionmaker_returning(session)), + patch( + "services.payments.pipeline.get_payment_by_payment_id", + new=AsyncMock(return_value=None), + ), + patch("services.payments.pipeline.add_payment", new=AsyncMock()) as add, + patch("services.payments.pipeline.invalidate_payment_cache", new=AsyncMock()), + ): + result = await process_cancelled_payment( + "heleket", parsed, new_status="failed" + ) + + self.assertTrue(result.ok) + add.assert_awaited_once() + self.assertEqual(add.await_args.kwargs["status"], "failed") + session.commit.assert_awaited_once() + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_referrals_and_gifts_resolution.py b/tests/test_referrals_and_gifts_resolution.py new file mode 100644 index 00000000..acc8950d --- /dev/null +++ b/tests/test_referrals_and_gifts_resolution.py @@ -0,0 +1,147 @@ +import unittest +from datetime import UTC, datetime +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from api.v2.routes.referrals import apply_referral +from api.v2.schemas.web_public import ReferralApplyRequest +from database.gifts import store_gift_link +from database.referrals import add_referral, get_total_referrals +from services.gifts import redeem_gift + + +class ReferralsResolutionTests(unittest.IsolatedAsyncioTestCase): + async def test_add_referral_creates_relation_for_resolved_users(self): + referred = SimpleNamespace(id=10, tg_id=1010) + referrer = SimpleNamespace(id=20, tg_id=2020) + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + + with patch( + "database.referrals.resolve_user_optional", + new=AsyncMock(side_effect=[referred, referrer]), + ): + await add_referral(session, referred_legacy=1010, referrer_legacy=2020) + + session.execute.assert_awaited_once() + session.commit.assert_not_called() + + async def test_add_referral_ignores_self_referral(self): + same_user = SimpleNamespace(id=30, tg_id=3030) + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + + with patch( + "database.referrals.resolve_user_optional", + new=AsyncMock(side_effect=[same_user, same_user]), + ): + await add_referral(session, referred_legacy=3030, referrer_legacy=3030) + + session.execute.assert_not_awaited() + session.commit.assert_not_awaited() + + async def test_get_total_referrals_returns_zero_when_user_missing(self): + session = SimpleNamespace(execute=AsyncMock()) + with patch("database.referrals.resolve_user_optional", new=AsyncMock(return_value=None)): + total = await get_total_referrals(session, referrer_legacy=4040) + self.assertEqual(total, 0) + session.execute.assert_not_awaited() + + async def test_apply_referral_accepts_site_referral_code(self): + session = object() + identity = SimpleNamespace(id="ident-email") + body = ReferralApplyRequest(referrer_code="https://example.com/referral/321") + referrer_user = SimpleNamespace(id=321, tg_id=None) + referred_user = SimpleNamespace(id=555, tg_id=None) + + with ( + patch("api.v2.routes.referrals.idb.ensure_billing_user_for_identity", new=AsyncMock(return_value=555)), + patch("api.v2.routes.referrals.resolve_user_optional", new=AsyncMock(side_effect=[referrer_user, referred_user])), + patch("api.v2.routes.referrals.get_referral_by_referred_id", new=AsyncMock(return_value=None)), + patch("api.v2.routes.referrals.add_referral", new=AsyncMock()) as add_referral_mock, + ): + result = await apply_referral(body, session=session, identity=identity) + + self.assertTrue(result.ok) + self.assertEqual(result.referrer_code, "321") + self.assertEqual(result.referrer_user_id, 321) + self.assertEqual(result.referred_user_id, 555) + add_referral_mock.assert_awaited_once_with(session, 555, 321) + + +class GiftsResolutionTests(unittest.IsolatedAsyncioTestCase): + async def test_store_gift_link_resolves_sender(self): + sender = SimpleNamespace(id=77, tg_id=7007) + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + + with patch("database.gifts.resolve_user_optional", new=AsyncMock(return_value=sender)): + ok = await store_gift_link( + session=session, + gift_id="gift_1", + sender_legacy_ref=7007, + selected_months=1, + expiry_time=datetime.now(UTC), + gift_link="https://t.me/test", + tariff_id=1, + max_usages=1, + ) + + self.assertTrue(ok) + session.execute.assert_awaited_once() + session.commit.assert_not_awaited() + + async def test_store_gift_link_raises_when_sender_missing(self): + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + with patch("database.gifts.resolve_user_optional", new=AsyncMock(return_value=None)): + with self.assertRaises(ValueError): + await store_gift_link( + session=session, + gift_id="gift_2", + sender_legacy_ref=9999, + selected_months=1, + expiry_time=datetime.now(UTC), + gift_link="https://t.me/test", + ) + + session.execute.assert_not_awaited() + session.commit.assert_not_awaited() + + async def test_redeem_gift_uses_billing_user_id_without_telegram(self): + gift_info = SimpleNamespace( + gift_id="gift_1", + sender_user_id=77, + sender_tg_id=7007, + recipient_user_id=None, + is_unlimited=False, + is_used=False, + max_usages=1, + tariff_id=5, + expiry_time=None, + selected_device_limit=2, + selected_traffic_gb=50, + selected_price_rub=990, + ) + billing_user = SimpleNamespace(id=555, tg_id=None) + tariff_dict = {"id": 5, "name": "Gift tariff", "duration_days": 30, "group_code": "gifts"} + session = SimpleNamespace(execute=AsyncMock(), commit=AsyncMock(), rollback=AsyncMock()) + + with ( + patch("services.gifts.get_gift_locked", new=AsyncMock(return_value=gift_info)), + patch("services.gifts.resolve_user_optional", new=AsyncMock(return_value=billing_user)), + patch("services.gifts.get_gift_usage", new=AsyncMock(return_value=None)), + patch("services.gifts.count_gift_usages", new=AsyncMock(return_value=0)), + patch("services.gifts.get_referral_by_referred_id", new=AsyncMock(return_value=None)), + patch("services.gifts.add_referral", new=AsyncMock()) as add_referral_mock, + patch("services.gifts.update_trial", new=AsyncMock()), + patch("services.gifts.get_tariff_by_id", new=AsyncMock(return_value=tariff_dict)), + patch("services.gifts.record_gift_usage", new=AsyncMock()), + patch("services.gifts.mark_gift_fully_redeemed", new=AsyncMock()), + patch("services.keys.create_vpn_key_headless", new=AsyncMock()) as create_key_mock, + ): + result = await redeem_gift(session, "gift_1", 555) + + self.assertEqual(result.gift_id, "gift_1") + self.assertEqual(result.tariff_id, 5) + self.assertEqual(result.duration_days, 30) + add_referral_mock.assert_awaited_once_with(session, 555, 77) + create_key_mock.assert_awaited_once() + self.assertEqual(create_key_mock.await_args.kwargs["tg_id"], 555) + self.assertTrue(create_key_mock.await_args.kwargs["skip_balance_charge"]) diff --git a/tests/test_resolution_actor.py b/tests/test_resolution_actor.py new file mode 100644 index 00000000..d826ac2a --- /dev/null +++ b/tests/test_resolution_actor.py @@ -0,0 +1,81 @@ +import unittest + +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from database.access.resolution import ( + ActorSurface, + ResolvedActor, + resolve_actor_from_identity, + resolve_actor_from_legacy_ref, +) + + +class ResolveActorFromLegacyRefTests(unittest.IsolatedAsyncioTestCase): + async def test_returns_telegram_surface_when_legacy_equals_user_tg_id(self): + user = SimpleNamespace(id=42, tg_id=777, identity_id="ident-1") + session = object() + + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=user)): + actor = await resolve_actor_from_legacy_ref(session, 777) + + self.assertEqual(actor.surface, ActorSurface.TELEGRAM) + self.assertEqual(actor.billing_user_id, 42) + self.assertEqual(actor.telegram_chat_id, 777) + self.assertEqual(actor.identity_id, "ident-1") + + async def test_returns_web_surface_when_user_has_no_tg(self): + user = SimpleNamespace(id=11, tg_id=None, identity_id="ident-web") + session = object() + + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=user)): + actor = await resolve_actor_from_legacy_ref(session, 11) + + self.assertEqual(actor.surface, ActorSurface.WEB) + self.assertEqual(actor.billing_user_id, 11) + self.assertIsNone(actor.telegram_chat_id) + self.assertEqual(actor.identity_id, "ident-web") + + async def test_returns_web_surface_for_linked_user_when_legacy_is_internal_id(self): + user = SimpleNamespace(id=55, tg_id=700700, identity_id="ident-linked") + session = object() + + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=user)): + actor = await resolve_actor_from_legacy_ref(session, 55) + + self.assertEqual(actor.surface, ActorSurface.WEB) + self.assertEqual(actor.billing_user_id, 55) + self.assertEqual(actor.telegram_chat_id, 700700) + self.assertEqual(actor.identity_id, "ident-linked") + + async def test_returns_unknown_surface_with_fallback_chat_id_when_user_missing(self): + session = object() + + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=None)): + actor = await resolve_actor_from_legacy_ref(session, 999999) + + self.assertEqual(actor.surface, ActorSurface.UNKNOWN) + self.assertIsNone(actor.billing_user_id) + self.assertEqual(actor.telegram_chat_id, 999999) + self.assertIsNone(actor.identity_id) + + +class ResolveActorFromIdentityTests(unittest.IsolatedAsyncioTestCase): + async def test_resolves_billing_and_telegram_ids(self): + identity = SimpleNamespace(id="ident-main") + user = SimpleNamespace(id=123, tg_id=555777, identity_id="ident-main") + session = object() + + with patch("database.identities.ensure_billing_user_for_identity", new=AsyncMock(return_value=123)): + with patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=user)): + actor = await resolve_actor_from_identity(session, identity) + + self.assertEqual( + actor, + ResolvedActor( + surface=ActorSurface.WEB, + billing_user_id=123, + telegram_chat_id=555777, + identity_id="ident-main", + ), + ) diff --git a/tests/test_scheduled_broadcasts_resolution.py b/tests/test_scheduled_broadcasts_resolution.py new file mode 100644 index 00000000..505e03fb --- /dev/null +++ b/tests/test_scheduled_broadcasts_resolution.py @@ -0,0 +1,65 @@ +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock, patch + +from database.scheduled_broadcasts import create_scheduled_broadcast, list_scheduled_broadcasts + + +class ScheduledBroadcastResolutionTests(unittest.IsolatedAsyncioTestCase): + async def test_create_scheduled_broadcast_resolves_user_and_mirror_tg(self): + session = SimpleNamespace(add=Mock(), commit=AsyncMock(), refresh=AsyncMock()) + creator = SimpleNamespace(id=10, tg_id=555) + + with patch("database.scheduled_broadcasts.resolve_user_optional", new=AsyncMock(return_value=creator)): + broadcast = await create_scheduled_broadcast( + session=session, + created_by_tg_id=999, + send_to="all", + cluster_name=None, + text="hello", + photo=None, + keyboard_json=None, + scheduled_for=SimpleNamespace(), + workers=5, + messages_per_second=20, + ) + + self.assertEqual(broadcast.created_by_user_id, 10) + self.assertEqual(broadcast.created_by_tg_id, 555) + session.commit.assert_awaited_once() + session.refresh.assert_awaited_once_with(broadcast) + + async def test_create_scheduled_broadcast_keeps_legacy_tg_when_user_missing(self): + session = SimpleNamespace(add=Mock(), commit=AsyncMock(), refresh=AsyncMock()) + + with patch("database.scheduled_broadcasts.resolve_user_optional", new=AsyncMock(return_value=None)): + broadcast = await create_scheduled_broadcast( + session=session, + created_by_tg_id=123456, + send_to="all", + cluster_name=None, + text="hello", + photo=None, + keyboard_json=None, + scheduled_for=SimpleNamespace(), + workers=5, + messages_per_second=20, + ) + + self.assertIsNone(broadcast.created_by_user_id) + self.assertEqual(broadcast.created_by_tg_id, 123456) + + async def test_list_scheduled_broadcasts_filters_by_created_by_tg(self): + rows = [SimpleNamespace(id="a"), SimpleNamespace(id="b")] + session = SimpleNamespace(execute=AsyncMock(return_value=SimpleNamespace(scalars=lambda: SimpleNamespace(all=lambda: rows)))) + + result = await list_scheduled_broadcasts( + session=session, + statuses=None, + created_by_tg_id=777, + limit=10, + offset=0, + ) + + self.assertEqual(result, rows) + session.execute.assert_awaited_once() diff --git a/tests/test_web_auth_purchase_link_flow.py b/tests/test_web_auth_purchase_link_flow.py new file mode 100644 index 00000000..cccbb0ff --- /dev/null +++ b/tests/test_web_auth_purchase_link_flow.py @@ -0,0 +1,366 @@ +import asyncio +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from starlette.requests import Request + +setattr(asyncio, "_validate_client_code_ran", True) + +from api.v2.routes.auth import auth_summary, register_by_email +from api.v2.routes.keys import user_keys +from api.v2.routes.payment_links import create_link as create_payment_link_route +from api.v2.routes.tariffs import purchase_tariff_with_balance +from api.v2.schemas.identities import RegisterByEmailRequest +from api.v2.schemas.payment_links import PaymentLinkCreateRequest +from api.v2.schemas.web_public import TariffPurchaseRequest +from database.identities import attach_telegram, merge_billing_user_into_telegram + + +def _make_request() -> Request: + scope = { + "type": "http", + "method": "POST", + "path": "/", + "headers": [], + "query_string": b"", + "client": ("127.0.0.1", 12345), + "server": ("testserver", 80), + "scheme": "http", + "http_version": "1.1", + } + return Request(scope) + + +def _scalars_all_result(rows): + return SimpleNamespace(scalars=lambda: SimpleNamespace(all=lambda: rows)) + + +def _scalar_one_or_none_result(value): + return SimpleNamespace(scalar_one_or_none=lambda: value) + + +class WebEmailRegistrationFlowTests(unittest.IsolatedAsyncioTestCase): + async def test_register_by_email_creates_identity_and_binds_actor(self): + request = _make_request() + session = object() + body = RegisterByEmailRequest(email="User@Test.Com", password="strongpass") + identity = SimpleNamespace(id="ident-email", tg_id=None) + + with ( + patch("api.v2.routes.auth.idb.get_identity_by_email", new=AsyncMock(return_value=None)), + patch( + "api.v2.routes.auth.idb.create_identity_with_token", + new=AsyncMock(return_value=(identity, "issued-token")), + ) as create_identity_with_token_mock, + patch("api.v2.routes.auth.bind_identity_actor", new=AsyncMock()) as bind_identity_actor_mock, + patch("api.v2.routes.auth.idb.ensure_billing_user_for_identity", new=AsyncMock(return_value=777)), + ): + result = await register_by_email(body, request, session=session) + + self.assertEqual(result.identity_id, "ident-email") + self.assertEqual(result.token, "issued-token") + create_identity_with_token_mock.assert_awaited_once_with( + session, + email="user@test.com", + password="strongpass", + ) + bind_identity_actor_mock.assert_awaited_once_with(request, session, identity) + + async def test_register_by_email_applies_referral_code_to_new_billing_user(self): + request = _make_request() + session = object() + body = RegisterByEmailRequest(email="invite@test.com", password="strongpass", referral_code="https://example.com/referral/321") + identity = SimpleNamespace(id="ident-invite", tg_id=None) + referrer_user = SimpleNamespace(id=321, tg_id=None) + + with ( + patch("api.v2.routes.auth.idb.get_identity_by_email", new=AsyncMock(return_value=None)), + patch( + "api.v2.routes.auth.idb.create_identity_with_token", + new=AsyncMock(return_value=(identity, "issued-token")), + ), + patch("api.v2.routes.auth.bind_identity_actor", new=AsyncMock()), + patch("api.v2.routes.auth.resolve_user_optional", new=AsyncMock(return_value=referrer_user)) as resolve_user_mock, + patch("api.v2.routes.auth.idb.ensure_billing_user_for_identity", new=AsyncMock(return_value=555)) as ensure_billing_user_mock, + patch("api.v2.routes.auth.get_referral_by_referred_id", new=AsyncMock(return_value=None)), + patch("api.v2.routes.auth.add_referral", new=AsyncMock()) as add_referral_mock, + ): + result = await register_by_email(body, request, session=session) + + self.assertEqual(result.identity_id, "ident-invite") + resolve_user_mock.assert_awaited_once_with(session, 321) + ensure_billing_user_mock.assert_awaited_once_with(session, identity) + add_referral_mock.assert_awaited_once_with(session, 555, 321) + + +class WebTariffPurchaseFlowTests(unittest.IsolatedAsyncioTestCase): + async def test_purchase_tariff_uses_identity_billing_user_and_creates_key(self): + session = SimpleNamespace(commit=AsyncMock()) + identity = SimpleNamespace(id="ident-email") + body = TariffPurchaseRequest( + tariff_id=7, + selected_device_limit=5, + selected_traffic_gb=100, + ) + tariff = {"id": 7, "is_active": True, "duration_days": 30, "price_rub": 990} + + with ( + patch( + "api.v2.routes.tariffs.idb.ensure_billing_user_for_identity", + new=AsyncMock(return_value=501), + ) as ensure_billing_user_mock, + patch("api.v2.routes.tariffs.get_tariff_by_id", new=AsyncMock(return_value=tariff)), + patch("api.v2.routes.tariffs.calculate_config_price", return_value=990), + patch("api.v2.routes.tariffs.get_balance", new=AsyncMock(return_value=1500.0)), + patch("api.v2.routes.tariffs.create_key", new=AsyncMock()) as create_key_mock, + ): + result = await purchase_tariff_with_balance( + body, + request=_make_request(), + preview=False, + session=session, + identity=identity, + ) + + self.assertTrue(result.ok) + self.assertEqual(result.charged_rub, 990) + ensure_billing_user_mock.assert_awaited_once_with(session, identity) + create_key_mock.assert_awaited_once() + kwargs = create_key_mock.await_args.kwargs + self.assertEqual(kwargs["tg_id"], 501) + self.assertEqual(kwargs["plan"], 7) + self.assertEqual(kwargs["selected_duration_days"], 30) + self.assertEqual(kwargs["selected_device_limit"], 5) + self.assertEqual(kwargs["selected_traffic_gb"], 100) + self.assertEqual(kwargs["selected_price_rub"], 990) + + +class WebTariffPaymentLinkFlowTests(unittest.IsolatedAsyncioTestCase): + async def test_create_payment_link_stores_tariff_purchase_intent_for_billing_user(self): + session = object() + identity = SimpleNamespace(id="ident-email") + request = _make_request() + body = PaymentLinkCreateRequest( + identity_id="ident-email", + amount=1290, + currency="RUB", + provider_id="ROBOKASSA", + success_url="https://example.com/payment-success", + failure_url="https://example.com/payment-failure", + metadata={ + "payment_flow": "tariff_purchase", + "tariff_id": 9, + "selected_device_limit": 4, + "selected_traffic_gb": 200, + }, + ) + + with ( + patch( + "api.v2.routes.payment_links.idb.ensure_billing_user_for_identity", + new=AsyncMock(return_value=777), + ) as ensure_billing_user_mock, + patch( + "api.v2.routes.payment_links.create_payment_link", + new=AsyncMock(return_value=SimpleNamespace(success=True, payment_id="pid-1", payment_url="https://pay.test", error=None)), + ) as create_payment_link_mock, + patch("api.v2.routes.payment_links.create_temporary_data", new=AsyncMock()) as create_temporary_data_mock, + ): + result = await create_payment_link_route(body, request, session=session, identity=identity) + + self.assertTrue(result.success) + self.assertEqual(result.payment_id, "pid-1") + ensure_billing_user_mock.assert_awaited_once_with(session, identity) + create_payment_link_mock.assert_awaited_once() + payment_request = create_payment_link_mock.await_args.args[1] + self.assertEqual(payment_request.legacy_user_ref, 777) + self.assertEqual(payment_request.metadata["tariff_id"], 9) + create_temporary_data_mock.assert_awaited_once_with( + session, + 777, + "waiting_for_payment", + { + "tariff_id": 9, + "required_amount": 1290, + "selected_price_rub": 1290, + "selected_device_limit": 4, + "selected_traffic_limit_gb": 200, + }, + ) + + +class WebAccountKeysFlowTests(unittest.IsolatedAsyncioTestCase): + async def test_auth_keys_returns_web_identity_keys_without_telegram(self): + session = object() + identity = SimpleNamespace(id="ident-email", tg_id=None) + request = _make_request() + key_obj = SimpleNamespace( + email="web-user@example.com", + alias="Main key", + client_id="client-1", + tariff_id=5, + server_id="eu-1", + created_at=1700000000000, + expiry_time=1800000000000, + key="https://example.com/sub/1", + remnawave_link=None, + is_frozen=False, + ) + + with ( + patch("api.v2.routes.keys._resolve_billing_user_id", new=AsyncMock(return_value=555)), + patch("api.v2.routes.keys.get_keys", new=AsyncMock(return_value=[key_obj])) as get_keys_mock, + ): + result = await user_keys(request, session=session, identity=identity) + + get_keys_mock.assert_awaited_once_with(session, 555) + self.assertEqual(len(result), 1) + self.assertEqual(result[0].email, "web-user@example.com") + self.assertEqual(result[0].client_id, "client-1") + self.assertEqual(result[0].server_id, "eu-1") + self.assertFalse(result[0].is_frozen) + + async def test_auth_summary_returns_referral_code_for_web_identity(self): + session = SimpleNamespace(execute=AsyncMock(side_effect=[ + SimpleNamespace(scalar_one=lambda: 2), + SimpleNamespace(scalar_one=lambda: 1), + SimpleNamespace(scalar_one=lambda: 0), + ])) + identity = SimpleNamespace(id="ident-email", email="web@example.com", tg_id=None) + request = _make_request() + + with ( + patch("api.v2.routes.auth.get_request_actor", return_value=SimpleNamespace(billing_user_id=555)), + patch("api.v2.routes.auth.get_balance", new=AsyncMock(return_value=125.0)), + patch("api.v2.routes.auth.get_trial", new=AsyncMock(return_value=1)), + patch("api.v2.routes.auth.get_keys", new=AsyncMock(return_value=[])), + patch( + "api.v2.routes.auth.get_referral_stats", + new=AsyncMock(return_value={"total_referrals": 4, "active_referrals": 2, "total_referral_bonus": 99.5}), + ), + ): + result = await auth_summary(request, session=session, identity=identity) + + self.assertTrue(result.referral_code.startswith("r1_")) + self.assertEqual(result.referrals_total, 4) + self.assertEqual(result.referrals_active, 2) + self.assertEqual(result.referral_bonus_total, 99.5) + + +class TelegramLinkFlowTests(unittest.IsolatedAsyncioTestCase): + async def test_attach_telegram_updates_identity_and_links_tg_user(self): + identity_before = SimpleNamespace(id="ident-email", tg_id=None, is_admin=False) + identity_after = SimpleNamespace(id="ident-email", tg_id=None, is_admin=False) + session = SimpleNamespace( + execute=AsyncMock( + side_effect=[ + _scalar_one_or_none_result(None), + SimpleNamespace(), + ] + ), + commit=AsyncMock(), + refresh=AsyncMock(), + ) + + with ( + patch( + "database.identities.get_identity_by_id", + new=AsyncMock(side_effect=[identity_before, identity_after]), + ), + patch("database.identities.get_identity_by_tg_id", new=AsyncMock(return_value=None)), + patch("database.identities.merge_billing_user_into_telegram", new=AsyncMock()) as merge_mock, + ): + result = await attach_telegram(session, "ident-email", 7007) + + self.assertIs(result, identity_after) + self.assertEqual(identity_after.tg_id, 7007) + merge_mock.assert_awaited_once_with(session, "ident-email", 7007) + session.commit.assert_awaited_once() + session.refresh.assert_awaited_once_with(identity_after) + update_stmt = session.execute.await_args_list[1].args[0].compile() + self.assertEqual(update_stmt.params["identity_id"], "ident-email") + self.assertEqual(update_stmt.params["tg_id_1"], 7007) + + async def test_merge_billing_user_into_existing_tg_user_moves_subscription_and_payments(self): + billing_user = SimpleNamespace( + id=41, + tg_id=None, + username="web-user", + first_name="Web", + last_name="User", + language_code="ru", + is_bot=False, + balance=250.0, + trial=2, + preferred_currency="RUB", + source_code="site", + ) + telegram_user = SimpleNamespace(id=77, tg_id=7007) + execute_results = [ + _scalars_all_result([billing_user]), + _scalar_one_or_none_result(0), + ] + + async def execute_side_effect(*args, **kwargs): + if execute_results: + return execute_results.pop(0) + return SimpleNamespace() + + session = SimpleNamespace( + execute=AsyncMock(side_effect=execute_side_effect), + add=lambda obj: None, + flush=AsyncMock(), + commit=AsyncMock(), + ) + + with ( + patch("database.access.resolution.resolve_user_optional", new=AsyncMock(return_value=telegram_user)), + patch("database.users.update_balance", new=AsyncMock()) as update_balance_mock, + patch("database.identities.refresh_tg_mirrors_for_user", new=AsyncMock()) as refresh_mirrors_mock, + patch("database.users.invalidate_balance_cache", new=AsyncMock()), + patch("database.users.invalidate_profile_cache", new=AsyncMock()), + ): + await merge_billing_user_into_telegram(session, "ident-email", 7007) + + update_balance_mock.assert_awaited_once_with(session, 77, 250.0) + refresh_mirrors_mock.assert_awaited_once_with(session, 77) + session.commit.assert_awaited() + + compiled_statements = [ + call.args[0].compile() + for call in session.execute.await_args_list + if call.args and hasattr(call.args[0], "compile") + ] + + self.assertTrue( + any( + "UPDATE keys SET user_id" in str(compiled) + and compiled.params.get("user_id") == 77 + and compiled.params.get("user_id_1") == 41 + for compiled in compiled_statements + ) + ) + self.assertTrue( + any( + "UPDATE payments SET user_id" in str(compiled) + and compiled.params.get("user_id") == 77 + and compiled.params.get("user_id_1") == 41 + for compiled in compiled_statements + ) + ) + self.assertTrue( + any( + "DELETE FROM users" in str(compiled) + and compiled.params.get("id_1") == 41 + for compiled in compiled_statements + ) + ) + self.assertTrue( + any( + "UPDATE users SET identity_id" in str(compiled) + and compiled.params.get("identity_id") == "ident-email" + and compiled.params.get("id_1") == 77 + for compiled in compiled_statements + ) + ) diff --git a/utils/__init__.py b/utils/__init__.py index e69de29b..8b137891 100644 --- a/utils/__init__.py +++ b/utils/__init__.py @@ -0,0 +1 @@ + diff --git a/utils/backup.py b/utils/backup.py index e2e1552b..72547081 100644 --- a/utils/backup.py +++ b/utils/backup.py @@ -1,4 +1,5 @@ import os +import shutil import subprocess import tarfile @@ -21,6 +22,7 @@ from config import ( DB_NAME, DB_PASSWORD, DB_USER, + PG_IN_DOCKER, PG_HOST, PG_PORT, BACKUP_CREATE_ARCHIVE, @@ -32,6 +34,60 @@ from config import ( from logger import logger +DOCKER_POSTGRES_CONTAINER = "solobot-postgres" + + +def _find_docker_postgres_container() -> str | None: + if shutil.which("docker") is None: + return None + result = subprocess.run( + ["docker", "inspect", "-f", "{{.State.Running}}", DOCKER_POSTGRES_CONTAINER], + capture_output=True, + text=True, + ) + if result.returncode == 0 and result.stdout.strip().lower() == "true": + return DOCKER_POSTGRES_CONTAINER + return None + + +def _get_postgres_execution_target() -> tuple[str, str | None]: + if PG_IN_DOCKER: + container = _find_docker_postgres_container() + if container: + return "docker", container + raise FileNotFoundError( + f"PostgreSQL настроен на Docker, но контейнер '{DOCKER_POSTGRES_CONTAINER}' не найден или не запущен" + ) + return "host", None + + +def _create_database_backup_via_docker(filename: Path, container: str) -> None: + with open(filename, "wb") as dump_file: + result = subprocess.run( + [ + "docker", + "exec", + "-e", + f"PGPASSWORD={DB_PASSWORD}", + container, + "pg_dump", + "-U", + DB_USER, + "-h", + "127.0.0.1", + "-p", + "5432", + "-F", + "c", + DB_NAME, + ], + stdout=dump_file, + stderr=subprocess.PIPE, + ) + if result.returncode != 0: + raise subprocess.CalledProcessError(result.returncode, result.args, stderr=result.stderr) + + async def backup_database(bot_instance: Bot | None = None) -> Exception | None: """ Создает резервную копию базы данных (или полный архив) и отправляет его администраторам. @@ -85,38 +141,45 @@ def _create_database_backup() -> tuple[str | None, Exception | None]: filename = backup_dir / f"{DB_NAME}-backup-{date_formatted}-{pid_suffix}.sql" try: - os.environ["PGPASSWORD"] = DB_PASSWORD + target, container = _get_postgres_execution_target() - subprocess.run( - [ - "pg_dump", - "-U", - DB_USER, - "-h", - PG_HOST, - "-p", - PG_PORT, - "-F", - "c", - "-f", - str(filename), - DB_NAME, - ], - check=True, - capture_output=True, - text=True, - ) - logger.info("[Backup] БД создана: {}", filename) + if target == "docker" and container: + _create_database_backup_via_docker(filename, container) + logger.info("[Backup] БД создана через Docker-контейнер {}: {}", container, filename) + elif shutil.which("pg_dump") is not None: + env = os.environ.copy() + env["PGPASSWORD"] = DB_PASSWORD + subprocess.run( + [ + "pg_dump", + "-U", + DB_USER, + "-h", + PG_HOST, + "-p", + PG_PORT, + "-F", + "c", + "-f", + str(filename), + DB_NAME, + ], + check=True, + capture_output=True, + text=True, + env=env, + ) + logger.info("[Backup] БД создана через host pg_dump: {}", filename) + else: + raise FileNotFoundError("PostgreSQL недоступен: не найден контейнер и отсутствует host pg_dump") return str(filename), None except subprocess.CalledProcessError as e: - logger.error("[Backup] pg_dump: {}", e.stderr) + stderr = e.stderr.decode("utf-8", errors="replace") if isinstance(e.stderr, bytes) else e.stderr + logger.error("[Backup] pg_dump: {}", stderr) return None, e except Exception as e: logger.error("[Backup] Непредвиденная ошибка: {}", e) return None, e - finally: - if "PGPASSWORD" in os.environ: - del os.environ["PGPASSWORD"] def _create_backup_archive() -> tuple[str | None, Exception | None]: diff --git a/utils/button_icons.py b/utils/button_icons.py index 97775b99..fb325b72 100644 --- a/utils/button_icons.py +++ b/utils/button_icons.py @@ -7,6 +7,7 @@ _UniqueGiftColors = getattr(aiogram.types, "UniqueGiftColors", None) if _UniqueGiftColors is not None: _cfg = getattr(_UniqueGiftColors, "model_config", None) _base = dict(_cfg) if _cfg is not None else {} + _base.pop("protected_namespaces", None) _UniqueGiftColors.model_config = ConfigDict(**_base, protected_namespaces=()) _OriginalInlineKeyboardButton = aiogram.types.InlineKeyboardButton diff --git a/utils/csv_export.py b/utils/csv_export.py index 6d2b21fb..c72f6692 100644 --- a/utils/csv_export.py +++ b/utils/csv_export.py @@ -9,6 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from core.constants import PAYMENT_SYSTEMS_EXCLUDED from database.models import Key, Payment, Referral, Tariff, User +from database.access.resolution import resolve_user_optional async def export_users_csv(session: AsyncSession) -> BufferedInputFile: @@ -49,7 +50,7 @@ async def export_users_csv(session: AsyncSession) -> BufferedInputFile: async def export_payments_csv(session: AsyncSession) -> BufferedInputFile: - j = join(User, Payment, User.tg_id == Payment.tg_id) + j = join(User, Payment, User.id == Payment.user_id) query = ( select( User.tg_id, @@ -73,7 +74,9 @@ async def export_payments_csv(session: AsyncSession) -> BufferedInputFile: async def export_user_payments_csv(tg_id: int, session: AsyncSession) -> BufferedInputFile: - j = join(User, Payment, User.tg_id == Payment.tg_id) + u = await resolve_user_optional(session, tg_id) + uid = u.id if u is not None else tg_id + j = join(User, Payment, User.id == Payment.user_id) query = ( select( User.tg_id, @@ -87,7 +90,7 @@ async def export_user_payments_csv(tg_id: int, session: AsyncSession) -> Buffere ) .select_from(j) .where( - User.tg_id == tg_id, + User.id == uid, Payment.payment_system.notin_(PAYMENT_SYSTEMS_EXCLUDED), ) .order_by(Payment.created_at.asc()) @@ -121,17 +124,20 @@ def _export_payments_csv(payments, filename: str) -> BufferedInputFile: async def export_referrals_csv(referrer_tg_id: int, session: AsyncSession) -> BufferedInputFile | None: - j = join(Referral, User, Referral.referred_tg_id == User.tg_id) + ref_owner = await resolve_user_optional(session, referrer_tg_id) + if ref_owner is None: + return None + j = join(Referral, User, Referral.referred_user_id == User.id) query = ( select( - Referral.referred_tg_id, + User.tg_id, func.coalesce(User.first_name, ""), func.coalesce(User.last_name, ""), func.coalesce(User.username, ""), ) .select_from(j) - .where(Referral.referrer_tg_id == referrer_tg_id) - .order_by(Referral.referred_tg_id.asc()) + .where(Referral.referrer_user_id == ref_owner.id) + .order_by(Referral.referred_user_id.asc()) ) result = await session.execute(query) @@ -144,7 +150,8 @@ async def export_referrals_csv(referrer_tg_id: int, session: AsyncSession) -> Bu writer = csv.writer(output, delimiter=";") writer.writerow(["Приглашённый (tg_id)", "Имя"]) - for invited_id, first_name, last_name, username in rows: + for invited_tg, first_name, last_name, username in rows: + invited_id = invited_tg if invited_tg is not None else "—" full_name = first_name.strip() or username or str(invited_id) if last_name: full_name = f"{full_name} {last_name}" @@ -163,6 +170,7 @@ async def export_hot_leads_csv(session: AsyncSession) -> BufferedInputFile: stmt = ( select( User.tg_id, + User.id, User.username, User.first_name, User.last_name, @@ -170,13 +178,13 @@ async def export_hot_leads_csv(session: AsyncSession) -> BufferedInputFile: ) .where( exists( - select(Payment.tg_id) - .where(Payment.tg_id == User.tg_id) + select(Payment.user_id) + .where(Payment.user_id == User.id) .where(Payment.status == "success") .where(Payment.amount > 0) .where(Payment.payment_system.notin_(PAYMENT_SYSTEMS_EXCLUDED)) ), - not_(exists(select(Key.tg_id).where(Key.tg_id == User.tg_id).where(Key.expiry_time > now_ts))), + not_(exists(select(Key.user_id).where(Key.user_id == User.id).where(Key.expiry_time > now_ts))), ) .order_by(User.updated_at.desc()) ) @@ -187,8 +195,9 @@ async def export_hot_leads_csv(session: AsyncSession) -> BufferedInputFile: buffer = StringIO() writer = csv.writer(buffer) writer.writerow(["tg_id", "username", "first_name", "last_name", "updated_at"]) - for user in users: - writer.writerow(user) + for row in users: + tid = row.tg_id if row.tg_id is not None else row.id + writer.writerow([tid, row.username, row.first_name, row.last_name, row.updated_at]) buffer.seek(0) return BufferedInputFile( @@ -198,10 +207,11 @@ async def export_hot_leads_csv(session: AsyncSession) -> BufferedInputFile: async def export_keys_csv(session: AsyncSession) -> BufferedInputFile: - j = join(Key, Tariff, Key.tariff_id == Tariff.id, isouter=True) + jk = join(Key, User, Key.user_id == User.id) + j = join(jk, Tariff, Key.tariff_id == Tariff.id, isouter=True) query = ( select( - Key.tg_id, + User.tg_id, Key.client_id, Key.email, Key.created_at, @@ -244,7 +254,7 @@ async def export_keys_csv(session: AsyncSession) -> BufferedInputFile: tariff = row.tariff_name or "—" writer.writerow([ - row.tg_id, + row.tg_id if row.tg_id is not None else row.client_id, row.client_id, row.email, created_at, @@ -261,10 +271,31 @@ async def export_keys_csv(session: AsyncSession) -> BufferedInputFile: async def export_user_all_payments_csv(tg_id: int, session: AsyncSession) -> BufferedInputFile: + owner = await resolve_user_optional(session, tg_id) + if owner is None: + buffer = StringIO() + writer = csv.writer(buffer) + writer.writerow([ + "id", + "tg_id", + "payment_id", + "amount", + "currency", + "payment_system", + "status", + "original_amount", + "created_at", + ]) + buffer.seek(0) + return BufferedInputFile( + file=buffer.getvalue().encode("utf-8-sig"), + filename=f"user_{tg_id}_payments_full.csv", + ) + query = ( select( Payment.id, - Payment.tg_id, + User.tg_id, Payment.payment_id, Payment.amount, Payment.currency, @@ -273,7 +304,8 @@ async def export_user_all_payments_csv(tg_id: int, session: AsyncSession) -> Buf Payment.original_amount, Payment.created_at, ) - .where(Payment.tg_id == tg_id) + .join(User, Payment.user_id == User.id) + .where(Payment.user_id == owner.id) .order_by(Payment.created_at.asc()) ) @@ -305,9 +337,10 @@ async def export_user_all_payments_csv(tg_id: int, session: AsyncSession) -> Buf original_amount, created_at, ) in rows: + display_id = user_tg_id if user_tg_id is not None else owner.id writer.writerow([ internal_id, - user_tg_id, + display_id, external_payment_id or "", amount, currency, diff --git a/utils/disposable_emails.py b/utils/disposable_emails.py new file mode 100644 index 00000000..29cd6ca1 --- /dev/null +++ b/utils/disposable_emails.py @@ -0,0 +1,78 @@ +from logger import logger + +_BUILTIN_DOMAINS: set[str] = { + "mailinator.com", "guerrillamail.com", "guerrillamail.de", "guerrillamail.net", + "guerrillamail.org", "guerrillamailblock.com", "grr.la", "sharklasers.com", + "guerrillamail.info", "tempmail.com", "temp-mail.org", "temp-mail.io", + "throwaway.email", "fakeinbox.com", "tempail.com", "tempr.email", + "dispostable.com", "yopmail.com", "yopmail.fr", "yopmail.net", + "cool.fr.nf", "jetable.fr.nf", "nospam.ze.tc", "nomail.xl.cx", + "mega.zik.dj", "speed.1s.fr", "courriel.fr.nf", "moncourrier.fr.nf", + "monemail.fr.nf", "monmail.fr.nf", "hide.biz.st", "mytrashmail.com", + "mailnesia.com", "maildrop.cc", "discard.email", "discardmail.com", + "discardmail.de", "trashmail.com", "trashmail.me", "trashmail.net", + "trashmail.org", "trashmail.at", "trashmail.io", "trashmail.ws", + "trash-mail.com", "trash-mail.at", "trashemail.de", + "mailcatch.com", "mailscrap.com", "mailforspam.com", + "spamgourmet.com", "spamgourmet.net", "spamgourmet.org", + "mailexpire.com", "tempinbox.com", "tempomail.fr", + "10minutemail.com", "10minutemail.co.za", "10minutemail.net", + "minutemail.io", "emailondeck.com", "getnada.com", + "mohmal.com", "burnermail.io", "inboxbear.com", + "mailsac.com", "harakirimail.com", "33mail.com", + "maildax.com", "crazymailing.com", "mailtemp.info", + "emkei.cz", "example.com", "test.com", "mailinator.net", + "binkmail.com", "bobmail.info", "chammy.info", + "devnullmail.com", "letthemeatspam.com", "mailnull.com", + "nomail.pw", "nowmymail.com", "rmqkr.net", + "sharklasers.com", "spamfree24.org", "spamhereplease.com", + "tempmailaddress.com", "wegwerfmail.de", "wegwerfmail.net", + "wh4f.org", "mailzilla.com", "anonbox.net", + "bspamfree.org", "kurzepost.de", "objectmail.com", + "proxymail.eu", "rcpt.at", "reallymymail.com", + "recode.me", "regbypass.com", "s0ny.net", + "safetymail.info", "safetypost.de", "shieldedmail.com", + "sogetthis.com", "soodonims.com", "spambox.us", + "spamcero.com", "spamday.com", "spamex.com", + "spamfighter.cf", "spamfighter.ga", "spamfighter.gq", + "spamfighter.ml", "spamfighter.tk", + "spamhole.com", "spaml.com", "spaml.de", + "uggsrock.com", "uroid.com", "veryreallymymail.com", + "viditag.com", "vomoto.com", "vpn.st", + "vsimcard.com", "vubby.com", "vztc.com", + "wasteland.rfc822.org", "webemail.me", + "zetmail.com", "zippymail.info", + "mailnator.com", "mailtothis.com", + "mx0.wwwnew.eu", "mypartyclip.de", + "myzx.com", "nb.gy", "nobulk.com", + "noclickemail.com", "nogmailspam.info", + "nomail.xl.cx", "nomorespamemails.com", + "nospam.ze.tc", "nothingtoseehere.ca", +} + +_loaded_extra = False +_extra_domains: set[str] = set() + + +def _load_extra() -> None: + global _loaded_extra, _extra_domains + if _loaded_extra: + return + _loaded_extra = True + try: + from config import DISPOSABLE_EMAIL_BLOCKLIST_EXTRA + if isinstance(DISPOSABLE_EMAIL_BLOCKLIST_EXTRA, (list, set, tuple)): + _extra_domains = {d.strip().lower() for d in DISPOSABLE_EMAIL_BLOCKLIST_EXTRA if isinstance(d, str)} + if _extra_domains: + logger.info("[DisposableEmail] загружено {} дополнительных доменов", len(_extra_domains)) + except (ImportError, AttributeError): + pass + + +def is_disposable_email(email: str) -> bool: + """Проверяет, является ли email одноразовым.""" + _load_extra() + domain = email.strip().lower().rsplit("@", 1)[-1] if "@" in email else "" + if not domain: + return False + return domain in _BUILTIN_DOMAINS or domain in _extra_domains diff --git a/utils/referral_codes.py b/utils/referral_codes.py new file mode 100644 index 00000000..ae1a13aa --- /dev/null +++ b/utils/referral_codes.py @@ -0,0 +1,100 @@ +import base64 +import hashlib +import hmac +import re + +from config import API_TOKEN, WEBHOOK_SECRET_TOKEN + + +def _secret_bytes() -> bytes: + seed = (WEBHOOK_SECRET_TOKEN or API_TOKEN or "solobot-referral").strip() + return seed.encode("utf-8") + + +def _urlsafe_b64decode_nopad(value: str) -> bytes: + normalized = value + "=" * ((4 - len(value) % 4) % 4) + return base64.urlsafe_b64decode(normalized.encode("ascii")) + + +def encode_referral_code(user_id: int) -> str: + if int(user_id) <= 0: + raise ValueError("user_id must be positive") + raw = int(user_id).to_bytes(8, byteorder="big", signed=False) + secret = _secret_bytes() + mask = hmac.new(secret, b"ref-mask-v1", hashlib.sha256).digest()[:8] + obfuscated = bytes(a ^ b for a, b in zip(raw, mask)) + signature = hmac.new(secret, b"ref-sign-v1:" + obfuscated, hashlib.sha256).digest()[:6] + payload = base64.urlsafe_b64encode(obfuscated + signature).decode("ascii").rstrip("=") + return f"r1_{payload}" + + +def encode_partner_code(user_id: int) -> str: + if int(user_id) <= 0: + raise ValueError("user_id must be positive") + raw = int(user_id).to_bytes(8, byteorder="big", signed=False) + secret = _secret_bytes() + mask = hmac.new(secret, b"partner-mask-v1", hashlib.sha256).digest()[:8] + obfuscated = bytes(a ^ b for a, b in zip(raw, mask)) + signature = hmac.new(secret, b"partner-sign-v1:" + obfuscated, hashlib.sha256).digest()[:6] + payload = base64.urlsafe_b64encode(obfuscated + signature).decode("ascii").rstrip("=") + return f"p1_{payload}" + + +def decode_referral_code(value: str | None) -> int | None: + token = str(value or "").strip() + if not token: + return None + if token.startswith("r1_"): + encoded = token[3:] + try: + data = _urlsafe_b64decode_nopad(encoded) + except Exception: + return None + if len(data) != 14: + return None + obfuscated, signature = data[:8], data[8:] + secret = _secret_bytes() + expected = hmac.new(secret, b"ref-sign-v1:" + obfuscated, hashlib.sha256).digest()[:6] + if not hmac.compare_digest(signature, expected): + return None + mask = hmac.new(secret, b"ref-mask-v1", hashlib.sha256).digest()[:8] + raw = bytes(a ^ b for a, b in zip(obfuscated, mask)) + parsed = int.from_bytes(raw, byteorder="big", signed=False) + return parsed if parsed > 0 else None + if token.startswith("p1_"): + return None + match = re.fullmatch(r"\d+", token) + if not match: + return None + parsed = int(match.group(0)) + return parsed if parsed > 0 else None + + +def decode_partner_code(value: str | None) -> int | None: + token = str(value or "").strip() + if not token: + return None + if token.startswith("p1_"): + encoded = token[3:] + try: + data = _urlsafe_b64decode_nopad(encoded) + except Exception: + return None + if len(data) != 14: + return None + obfuscated, signature = data[:8], data[8:] + secret = _secret_bytes() + expected = hmac.new(secret, b"partner-sign-v1:" + obfuscated, hashlib.sha256).digest()[:6] + if not hmac.compare_digest(signature, expected): + return None + mask = hmac.new(secret, b"partner-mask-v1", hashlib.sha256).digest()[:8] + raw = bytes(a ^ b for a, b in zip(obfuscated, mask)) + parsed = int.from_bytes(raw, byteorder="big", signed=False) + return parsed if parsed > 0 else None + if token.startswith("r1_"): + return decode_referral_code(token) + match = re.fullmatch(r"\d+", token) + if not match: + return None + parsed = int(match.group(0)) + return parsed if parsed > 0 else None diff --git a/utils/telegram_login.py b/utils/telegram_login.py index ec166d81..aa62a437 100644 --- a/utils/telegram_login.py +++ b/utils/telegram_login.py @@ -33,3 +33,62 @@ def verify_telegram_login( computed = hmac.new(secret_key, data_check_string.encode(), hashlib.sha256).hexdigest() return hmac.compare_digest(computed, received_hash) + + +def verify_webapp_init_data( + init_data: str, + bot_token: str, + *, + max_age_seconds: int = 86400, +) -> dict | None: + """ + Валидирует Telegram WebApp initData (HMAC-SHA256). + Возвращает dict с user_id или None если невалидно. + https://core.telegram.org/bots/webapps#validating-data-received-via-the-mini-app + """ + import json + from urllib.parse import parse_qs + + if not init_data or not bot_token: + return None + + parsed = parse_qs(init_data, keep_blank_values=True) + received_hash = parsed.get("hash", [""])[0] + if not received_hash: + return None + + auth_date_str = parsed.get("auth_date", [""])[0] + try: + auth_date = int(auth_date_str) + if auth_date < time.time() - max_age_seconds: + return None + except (TypeError, ValueError): + return None + + check_pairs = [] + for key in sorted(parsed.keys()): + if key == "hash": + continue + check_pairs.append(f"{key}={parsed[key][0]}") + data_check_string = "\n".join(check_pairs) + + secret_key = hmac.new(b"WebAppData", bot_token.encode(), hashlib.sha256).digest() + computed = hmac.new(secret_key, data_check_string.encode(), hashlib.sha256).hexdigest() + + if not hmac.compare_digest(computed, received_hash): + return None + + user_raw = parsed.get("user", [""])[0] + user_id = None + if user_raw: + try: + user_data = json.loads(user_raw) + user_id = user_data.get("id") + except (json.JSONDecodeError, TypeError): + pass + + return { + "user_id": user_id, + "auth_date": auth_date, + "user_raw": user_raw, + } diff --git a/utils/turnstile.py b/utils/turnstile.py new file mode 100644 index 00000000..33a7b4ef --- /dev/null +++ b/utils/turnstile.py @@ -0,0 +1,40 @@ +import httpx + +from config import TURNSTILE_SECRET_KEY +from logger import logger + +_VERIFY_URL = "https://challenges.cloudflare.com/turnstile/v0/siteverify" + + +def turnstile_enabled() -> bool: + return bool(TURNSTILE_SECRET_KEY) + + +async def verify_turnstile_token(token: str | None, remote_ip: str | None = None) -> bool: + if not TURNSTILE_SECRET_KEY: + return True + + if not token or token == "__turnstile_disabled__": + logger.warning("[Turnstile] токен не предоставлен") + return False + + try: + payload: dict[str, str] = { + "secret": TURNSTILE_SECRET_KEY, + "response": token, + } + if remote_ip: + payload["remoteip"] = remote_ip + + async with httpx.AsyncClient(timeout=10.0) as client: + resp = await client.post(_VERIFY_URL, data=payload) + result = resp.json() + + success = result.get("success", False) + if not success: + codes = result.get("error-codes", []) + logger.warning("[Turnstile] верификация не пройдена: {}", codes) + return bool(success) + except Exception as exc: + logger.error("[Turnstile] ошибка проверки: {}", exc) + return False diff --git a/utils/versioning.py b/utils/versioning.py index 91fafc52..bcd489ac 100644 --- a/utils/versioning.py +++ b/utils/versioning.py @@ -92,4 +92,4 @@ def get_git_commit_number() -> str: def get_version() -> str: - return f"a02031919 {get_git_commit_number()}" + return f"v.6-b1204121200 {get_git_commit_number()}" diff --git a/utils/web_email_link_code.py b/utils/web_email_link_code.py new file mode 100644 index 00000000..9d78af2b --- /dev/null +++ b/utils/web_email_link_code.py @@ -0,0 +1,96 @@ +import hmac + +from config import LOGIN_CODE_TTL_SEC +from core.redis_cache import ( + cache_delete, + cache_get, + cache_incr, + cache_key, + cache_set, + cache_setnx, + redis_connection_ok, +) + + +_RESEND_COOLDOWN_SEC = 60.0 +_IP_WINDOW_SEC = 3600.0 +_IP_MAX_SENDS = 40 +_EMAIL_WINDOW_SEC = 3600.0 +_EMAIL_MAX_SENDS = 5 +_EMAIL_VERIFY_WINDOW_SEC = 600.0 +_EMAIL_MAX_VERIFY_ATTEMPTS = 10 + + +def normalize_email(value: str) -> str: + return (value or "").strip().lower() + + +def _code_key(email_norm: str) -> str: + return cache_key("web_email_link_code", email_norm) + + +def _cooldown_key(email_norm: str) -> str: + return cache_key("web_email_link_cooldown", email_norm) + + +def _ip_key(ip: str) -> str: + return cache_key("web_email_link_send_ip", ip) + + +def _email_send_key(email_norm: str) -> str: + return cache_key("web_email_link_sends", email_norm) + + +def _email_verify_key(email_norm: str) -> str: + return cache_key("web_email_link_verify", email_norm) + + +async def redis_ready() -> bool: + return await redis_connection_ok() + + +async def try_consume_ip_budget(ip: str) -> bool: + if not ip: + return True + n = await cache_incr(_ip_key(ip), _IP_WINDOW_SEC) + return n <= _IP_MAX_SENDS + + +async def try_consume_email_send_budget(email_norm: str) -> bool: + if not email_norm: + return True + n = await cache_incr(_email_send_key(email_norm), _EMAIL_WINDOW_SEC) + return n <= _EMAIL_MAX_SENDS + + +async def try_consume_email_verify_budget(email_norm: str) -> bool: + if not email_norm: + return True + n = await cache_incr(_email_verify_key(email_norm), _EMAIL_VERIFY_WINDOW_SEC) + return n <= _EMAIL_MAX_VERIFY_ATTEMPTS + + +async def try_acquire_cooldown(email_norm: str) -> bool: + return await cache_setnx(_cooldown_key(email_norm), 1, _RESEND_COOLDOWN_SEC) + + +async def release_cooldown(email_norm: str) -> None: + await cache_delete(_cooldown_key(email_norm)) + + +async def store_code(email_norm: str, code: str) -> bool: + return await cache_set(_code_key(email_norm), code, float(LOGIN_CODE_TTL_SEC)) + + +async def delete_code(email_norm: str) -> None: + await cache_delete(_code_key(email_norm)) + + +async def verify_and_consume_code(email_norm: str, code: str) -> bool: + stored = await cache_get(_code_key(email_norm)) + if not isinstance(stored, str): + return False + if not hmac.compare_digest(stored.strip(), (code or "").strip()): + return False + await cache_delete(_code_key(email_norm)) + return True diff --git a/utils/web_email_verify_code.py b/utils/web_email_verify_code.py new file mode 100644 index 00000000..abe6ebea --- /dev/null +++ b/utils/web_email_verify_code.py @@ -0,0 +1,89 @@ +import hmac + +from config import LOGIN_CODE_TTL_SEC +from core.redis_cache import ( + cache_delete, + cache_get, + cache_incr, + cache_key, + cache_set, + cache_setnx, + redis_connection_ok, +) + + +_RESEND_COOLDOWN_SEC = 60.0 +_IP_WINDOW_SEC = 3600.0 +_IP_MAX_SENDS = 20 +_EMAIL_WINDOW_SEC = 3600.0 +_EMAIL_MAX_SENDS = 5 +_EMAIL_VERIFY_WINDOW_SEC = 600.0 +_EMAIL_MAX_VERIFY_ATTEMPTS = 10 + + +def _code_key(email_norm: str) -> str: + return cache_key("web_email_verify_code", email_norm) + + +def _cooldown_key(email_norm: str) -> str: + return cache_key("web_email_verify_cooldown", email_norm) + + +def _ip_key(ip: str) -> str: + return cache_key("web_email_verify_send_ip", ip) + + +def _email_send_key(email_norm: str) -> str: + return cache_key("web_email_verify_sends", email_norm) + + +def _email_verify_key(email_norm: str) -> str: + return cache_key("web_email_verify_attempts", email_norm) + + +async def redis_ready() -> bool: + return await redis_connection_ok() + + +async def try_consume_ip_send_budget(ip: str) -> bool: + if not ip: + return True + n = await cache_incr(_ip_key(ip), _IP_WINDOW_SEC) + return n <= _IP_MAX_SENDS + + +async def try_consume_email_send_budget(email_norm: str) -> bool: + if not email_norm: + return True + n = await cache_incr(_email_send_key(email_norm), _EMAIL_WINDOW_SEC) + return n <= _EMAIL_MAX_SENDS + + +async def try_consume_verify_budget(email_norm: str) -> bool: + if not email_norm: + return True + n = await cache_incr(_email_verify_key(email_norm), _EMAIL_VERIFY_WINDOW_SEC) + return n <= _EMAIL_MAX_VERIFY_ATTEMPTS + + +async def try_acquire_resend_cooldown(email_norm: str) -> bool: + return await cache_setnx(_cooldown_key(email_norm), 1, _RESEND_COOLDOWN_SEC) + + +async def store_code(email_norm: str, code: str) -> bool: + return await cache_set(_code_key(email_norm), code, float(LOGIN_CODE_TTL_SEC)) + + +async def delete_code(email_norm: str) -> None: + await cache_delete(_code_key(email_norm)) + + +async def verify_and_consume_code(email_norm: str, code: str) -> bool: + key = _code_key(email_norm) + stored = await cache_get(key) + if not isinstance(stored, str): + return False + if not hmac.compare_digest(stored.strip(), (code or "").strip()): + return False + await cache_delete(key) + return True diff --git a/utils/web_login_code.py b/utils/web_login_code.py new file mode 100644 index 00000000..f15d197c --- /dev/null +++ b/utils/web_login_code.py @@ -0,0 +1,97 @@ +import hmac + +from config import LOGIN_CODE_TTL_SEC +from core.redis_cache import ( + cache_delete, + cache_get, + cache_incr, + cache_key, + cache_set, + cache_setnx, + redis_connection_ok, +) + + +_RESEND_COOLDOWN_SEC = 60.0 +_IP_WINDOW_SEC = 3600.0 +_IP_MAX_SENDS = 40 +_EMAIL_WINDOW_SEC = 3600.0 +_EMAIL_MAX_SENDS = 5 +_EMAIL_VERIFY_WINDOW_SEC = 600.0 +_EMAIL_MAX_VERIFY_ATTEMPTS = 10 + + +def normalize_login_email(email: str) -> str: + return (email or "").strip().lower() + + +def _code_key(email_norm: str) -> str: + return cache_key("web_login_code", email_norm) + + +def _cooldown_key(email_norm: str) -> str: + return cache_key("web_login_cooldown", email_norm) + + +def _ip_key(ip: str) -> str: + return cache_key("web_login_send_ip", ip) + + +def _email_send_key(email_norm: str) -> str: + return cache_key("web_login_email_sends", email_norm) + + +def _email_verify_key(email_norm: str) -> str: + return cache_key("web_login_email_verify", email_norm) + + +async def redis_ready_for_login_codes() -> bool: + return await redis_connection_ok() + + +async def try_consume_ip_send_budget(ip: str) -> bool: + if not ip: + return True + n = await cache_incr(_ip_key(ip), _IP_WINDOW_SEC) + return n <= _IP_MAX_SENDS + + +async def try_consume_email_send_budget(email_norm: str) -> bool: + if not email_norm: + return True + n = await cache_incr(_email_send_key(email_norm), _EMAIL_WINDOW_SEC) + return n <= _EMAIL_MAX_SENDS + + +async def try_consume_email_verify_budget(email_norm: str) -> bool: + if not email_norm: + return True + n = await cache_incr(_email_verify_key(email_norm), _EMAIL_VERIFY_WINDOW_SEC) + return n <= _EMAIL_MAX_VERIFY_ATTEMPTS + + +async def try_acquire_resend_cooldown(email_norm: str) -> bool: + return await cache_setnx(_cooldown_key(email_norm), 1, _RESEND_COOLDOWN_SEC) + + +async def release_resend_cooldown(email_norm: str) -> None: + await cache_delete(_cooldown_key(email_norm)) + + +async def store_code(email_norm: str, code: str) -> bool: + return await cache_set(_code_key(email_norm), code, float(LOGIN_CODE_TTL_SEC)) + + +async def delete_code(email_norm: str) -> None: + await cache_delete(_code_key(email_norm)) + + +async def verify_and_consume_code(email_norm: str, code: str) -> bool: + key = _code_key(email_norm) + stored = await cache_get(key) + if not isinstance(stored, str): + return False + if not hmac.compare_digest(stored.strip(), (code or "").strip()): + return False + await cache_delete(key) + return True diff --git a/utils/web_password_reset_code.py b/utils/web_password_reset_code.py new file mode 100644 index 00000000..b0a3c5dd --- /dev/null +++ b/utils/web_password_reset_code.py @@ -0,0 +1,92 @@ +import hmac + +from config import LOGIN_CODE_TTL_SEC +from core.redis_cache import ( + cache_delete, + cache_get, + cache_incr, + cache_key, + cache_set, + cache_setnx, + redis_connection_ok, +) + +_RESEND_COOLDOWN_SEC = 60.0 +_IP_WINDOW_SEC = 3600.0 +_IP_MAX_SENDS = 40 +_EMAIL_WINDOW_SEC = 3600.0 +_EMAIL_MAX_SENDS = 5 +_EMAIL_VERIFY_WINDOW_SEC = 600.0 +_EMAIL_MAX_VERIFY_ATTEMPTS = 10 + + +def _code_key(email_norm: str) -> str: + return cache_key("web_pwd_reset_code", email_norm) + + +def _cooldown_key(email_norm: str) -> str: + return cache_key("web_pwd_reset_cooldown", email_norm) + + +def _ip_key(ip: str) -> str: + return cache_key("web_pwd_reset_send_ip", ip) + + +def _email_send_key(email_norm: str) -> str: + return cache_key("web_pwd_reset_email_sends", email_norm) + + +def _email_verify_key(email_norm: str) -> str: + return cache_key("web_pwd_reset_email_verify", email_norm) + + +async def redis_ready() -> bool: + return await redis_connection_ok() + + +async def try_consume_ip_budget(ip: str) -> bool: + if not ip: + return True + n = await cache_incr(_ip_key(ip), _IP_WINDOW_SEC) + return n <= _IP_MAX_SENDS + + +async def try_consume_email_send_budget(email_norm: str) -> bool: + if not email_norm: + return True + n = await cache_incr(_email_send_key(email_norm), _EMAIL_WINDOW_SEC) + return n <= _EMAIL_MAX_SENDS + + +async def try_consume_email_verify_budget(email_norm: str) -> bool: + if not email_norm: + return True + n = await cache_incr(_email_verify_key(email_norm), _EMAIL_VERIFY_WINDOW_SEC) + return n <= _EMAIL_MAX_VERIFY_ATTEMPTS + + +async def try_acquire_cooldown(email_norm: str) -> bool: + return await cache_setnx(_cooldown_key(email_norm), 1, _RESEND_COOLDOWN_SEC) + + +async def release_cooldown(email_norm: str) -> None: + await cache_delete(_cooldown_key(email_norm)) + + +async def store_code(email_norm: str, code: str) -> bool: + return await cache_set(_code_key(email_norm), code, float(LOGIN_CODE_TTL_SEC)) + + +async def delete_code(email_norm: str) -> None: + await cache_delete(_code_key(email_norm)) + + +async def verify_and_consume_code(email_norm: str, code: str) -> bool: + key = _code_key(email_norm) + stored = await cache_get(key) + if not isinstance(stored, str): + return False + if not hmac.compare_digest(stored.strip(), (code or "").strip()): + return False + await cache_delete(key) + return True diff --git a/web/__init__.py b/web/__init__.py index 6ad56ef8..f4f7d9bb 100644 --- a/web/__init__.py +++ b/web/__init__.py @@ -1,7 +1,7 @@ from aiohttp.web_urldispatcher import UrlDispatcher -from handlers.payments.heleket.webhook import heleket_webhook -from handlers.payments.kassai.webhook import kassai_webhook +from services.payments.heleket.webhook import heleket_webhook +from services.payments.kassai.webhook import kassai_webhook from utils.modules_loader import load_module_webhooks