fix(api): security, caching, atomic rate limits, url obfuscation

Redis:
- Singleton ConnectionPool (redis.asyncio), 50 connections — не создаём
  новое TCP-соединение на каждый HTTP-запрос

Rate limiter:
- Полностью переписан на async/await
- Lua-скрипт _LUA_CHECK_AND_INCR — атомарная проверка+инкремент без race condition
- Lua-скрипт _LUA_ACQUIRE_CONCURRENT — атомарный захват слота задачи
- Старый паттерн INCR→check→DECR удалён (race condition при конкурентных запросах)

Security:
- get_current_user кэширует пользователя в Redis на 5 минут (TTL)
  Раньше: SELECT users на каждый HTTP-запрос
  Теперь: Redis GET (кэш) → SELECT users (только при промахе)
- hashed_password НЕ кладётся в кэш
- invalidate_user_cache() для сброса при смене тарифа/пароля
- get_ws_user() для WebSocket через ?token=JWT (браузеры не могут
  передавать Authorization header при WS-handshake)

WebSocket:
- Добавлена аутентификация (Depends(get_ws_user))
- Проверка ownership задачи ДО accept() соединения
- Чужой task_id → закрытие с кодом 4004

URL obfuscation:
- Task.public_id = secrets.token_urlsafe(16) = 22 случайных base64url символа
- Клиент работает только с public_id, внутренний UUID не раскрывается
- Все роутеры переключены на public_id в WHERE условиях
- TaskResponse больше не возвращает input_data (там minio_key и т.д.)
- Миграция 002_add_task_public_id.py

MinIO:
- Singleton клиент (не создаём новый на каждый upload)
- ensure_bucket() вызывается один раз при старте (lifespan), не на каждый запрос
- Путь uploads/{doc_uuid}{ext} — user_id убран из пути

CORS:
- Убраны wildcard allow_methods/allow_headers (несовместимы с credentials=True)
- Явный список: methods=[GET,POST,DELETE,OPTIONS], headers=[Authorization,Content-Type,Accept]
- Swagger/OpenAPI доступны только в ENVIRONMENT=development

Documents:
- Content-Length проверяется ДО чтения тела (ранняя отбивка больших файлов)
- Повторная проверка реального размера после чтения (защита от поддельного заголовка)
- Используем get_current_verified_user вместо get_current_user (требуем подтверждённый email)

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
jze9
2026-05-24 19:51:51 +05:00
parent 7758315632
commit c1bfb5f40e
12 changed files with 493 additions and 317 deletions

View File

@@ -0,0 +1,45 @@
"""Добавить public_id в таблицу tasks.
Revision ID: 002
Revises: 001
Create Date: 2026-05-24
"""
import secrets
from alembic import op
import sqlalchemy as sa
from sqlalchemy import text
revision = "002"
down_revision = "001"
branch_labels = None
depends_on = None
def upgrade() -> None:
# Добавляем колонку как nullable — сначала заполним данными, потом сделаем NOT NULL
op.add_column(
"tasks",
sa.Column("public_id", sa.String(32), nullable=True),
)
# Заполняем существующие строки уникальными public_id
conn = op.get_bind()
tasks = conn.execute(text("SELECT id FROM tasks")).fetchall()
for (task_id,) in tasks:
conn.execute(
text("UPDATE tasks SET public_id = :pid WHERE id = :id"),
{"pid": secrets.token_urlsafe(16), "id": task_id},
)
# Делаем NOT NULL и добавляем уникальный индекс
op.alter_column("tasks", "public_id", nullable=False)
op.create_unique_constraint("uq_tasks_public_id", "tasks", ["public_id"])
op.create_index("ix_tasks_public_id", "tasks", ["public_id"], unique=True)
def downgrade() -> None:
op.drop_index("ix_tasks_public_id", table_name="tasks")
op.drop_constraint("uq_tasks_public_id", "tasks", type_="unique")
op.drop_column("tasks", "public_id")

View File

@@ -1,18 +1,18 @@
"""Роутер загрузки документов на проверку плагиата.""" """Загрузка документов на проверку плагиата."""
import io
import logging import logging
import uuid import uuid
from pathlib import Path from pathlib import Path
from fastapi import APIRouter, Depends, File, HTTPException, UploadFile, status from fastapi import APIRouter, Depends, HTTPException, Request, UploadFile, status
from minio import Minio
from minio.error import S3Error
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings from app.config import settings
from app.core.celery_app import celery_app from app.core.celery_app import celery_app
from app.core.rate_limiter import check_and_increment_limit, check_concurrent_limit from app.core.minio_client import get_minio
from app.core.security import get_current_user from app.core.rate_limiter import acquire_concurrent_slot, check_and_increment_limit
from app.core.security import get_current_verified_user
from app.database import get_db from app.database import get_db
from app.models.task import Task from app.models.task import Task
from app.models.user import User from app.models.user import User
@@ -22,102 +22,104 @@ logger = logging.getLogger(__name__)
router = APIRouter(prefix="/documents", tags=["documents"]) router = APIRouter(prefix="/documents", tags=["documents"])
# Допустимые форматы
ALLOWED_EXTENSIONS = {".pdf", ".docx", ".txt"} ALLOWED_EXTENSIONS = {".pdf", ".docx", ".txt"}
MAX_FILE_SIZE_MB = 100 MAX_FILE_SIZE_BYTES = 100 * 1024 * 1024 # 100 МБ
MAX_FILE_SIZE_BYTES = MAX_FILE_SIZE_MB * 1024 * 1024
CONTENT_TYPE_MAP = {
def get_minio_client() -> Minio: ".pdf": "application/pdf",
"""Создать MinIO клиент.""" ".docx": "application/vnd.openxmlformats-officedocument.wordprocessingml.document",
return Minio( ".txt": "text/plain",
settings.MINIO_ENDPOINT, }
access_key=settings.MINIO_ACCESS_KEY,
secret_key=settings.MINIO_SECRET_KEY,
secure=False,
)
@router.post("/check", response_model=TaskResponse, status_code=status.HTTP_202_ACCEPTED) @router.post("/check", response_model=TaskResponse, status_code=status.HTTP_202_ACCEPTED)
async def upload_for_plagiarism_check( async def upload_for_plagiarism_check(
file: UploadFile = File(..., description="PDF, DOCX или TXT файл до 100 МБ"), request: Request,
current_user: User = Depends(get_current_user), file: UploadFile,
current_user: User = Depends(get_current_verified_user),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
) -> TaskResponse: ) -> TaskResponse:
""" """
Загрузить документ для проверки на плагиат. Загрузить документ на проверку плагиата.
Файл сохраняется в MinIO, затем диспатчится задача index.extract_and_check. Файл сохраняется в MinIO. В пути используется UUID документа,
Результат доступен через GET /tasks/{task_id}. не ID пользователя (пользователь не должен знать свой internal ID).
""" """
# Проверить расширение файла # ── 1. Проверяем заголовки ДО чтения тела ─────────────────────────────────
if not file.filename: if not file.filename:
raise HTTPException( raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Имя файла не указано")
status_code=status.HTTP_400_BAD_REQUEST,
detail="Имя файла не указано",
)
ext = Path(file.filename).suffix.lower() ext = Path(file.filename).suffix.lower()
if ext not in ALLOWED_EXTENSIONS: if ext not in ALLOWED_EXTENSIONS:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_415_UNSUPPORTED_MEDIA_TYPE, status_code=status.HTTP_415_UNSUPPORTED_MEDIA_TYPE,
detail=f"Неподдерживаемый формат файла. Допустимые форматы: {', '.join(ALLOWED_EXTENSIONS)}", detail=f"Неподдерживаемый формат. Допустимые: {', '.join(ALLOWED_EXTENSIONS)}",
) )
# Проверить лимит по тарифу # Content-Length позволяет отклонить слишком большой файл до его чтения в память.
limit_check = check_and_increment_limit(current_user.id, "plagiarism", current_user.plan) # Это не 100% защита (заголовок можно не передать), поэтому проверяем ещё раз после чтения.
if not limit_check["allowed"]: content_length = request.headers.get("content-length")
if content_length and int(content_length) > MAX_FILE_SIZE_BYTES:
raise HTTPException(
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
detail=f"Файл слишком большой (максимум {MAX_FILE_SIZE_BYTES // 1024 // 1024} МБ)",
)
# ── 2. Rate limits ─────────────────────────────────────────────────────────
limit_result = await check_and_increment_limit(current_user.id, "plagiarism", current_user.plan)
if not limit_result["allowed"]:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS, status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail=( detail=(
f"Превышен месячный лимит проверок плагиата для тарифа '{current_user.plan}'. " f"Превышен месячный лимит проверок (тариф '{current_user.plan}'): "
f"Использовано {limit_check['current']} из {limit_check['limit']}." f"{limit_result['current']}/{limit_result['limit']}. "
f"Сбросится {limit_result['reset_at']}."
), ),
headers={"Retry-After": "86400"},
) )
# Проверить лимит одновременных задач slot_acquired = await acquire_concurrent_slot(current_user.id, current_user.plan)
if not check_concurrent_limit(current_user.id, current_user.plan): if not slot_acquired:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS, status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="Превышен лимит одновременных задач.", detail="Превышен лимит одновременных задач.",
) )
# Прочитать файл # ── 3. Читаем файл и проверяем реальный размер ────────────────────────────
file_data = await file.read()
file_data = await file.read()
if len(file_data) > MAX_FILE_SIZE_BYTES: if len(file_data) > MAX_FILE_SIZE_BYTES:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, status_code=status.HTTP_413_REQUEST_ENTITY_TOO_LARGE,
detail=f"Файл слишком большой. Максимальный размер: {MAX_FILE_SIZE_MB} МБ", detail=f"Файл слишком большойаксимум {MAX_FILE_SIZE_BYTES // 1024 // 1024} МБ)",
) )
# Загрузить в MinIO # ── 4. Сохраняем в MinIO ───────────────────────────────────────────────────
minio_key = f"uploads/{current_user.id}/{uuid.uuid4()}{ext}" # В пути — UUID документа, НЕ user_id. Так не раскрываем внутренний ID пользователя.
minio_client = get_minio_client()
doc_uuid = str(uuid.uuid4())
minio_key = f"uploads/{doc_uuid}{ext}"
content_type = CONTENT_TYPE_MAP.get(ext, "application/octet-stream")
try: try:
# Создать бакет если не существует get_minio().put_object(
if not minio_client.bucket_exists(settings.MINIO_BUCKET_DOCS):
minio_client.make_bucket(settings.MINIO_BUCKET_DOCS)
import io
minio_client.put_object(
bucket_name=settings.MINIO_BUCKET_DOCS, bucket_name=settings.MINIO_BUCKET_DOCS,
object_name=minio_key, object_name=minio_key,
data=io.BytesIO(file_data), data=io.BytesIO(file_data),
length=len(file_data), length=len(file_data),
content_type=file.content_type or "application/octet-stream", content_type=content_type,
) )
logger.info(f"Файл загружен в MinIO: {minio_key}") except Exception as e:
logger.error("Ошибка загрузки в MinIO: %s", e)
except S3Error as e:
logger.error(f"Ошибка загрузки в MinIO: {e}")
raise HTTPException( raise HTTPException(
status_code=status.HTTP_503_SERVICE_UNAVAILABLE, status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
detail="Ошибка сохранения файла. Попробуйте позже.", detail="Ошибка сохранения файла. Попробуйте позже.",
) )
# Создать задачу в БД # ── 5. Создаём задачу и диспатчим ─────────────────────────────────────────
task = Task( task = Task(
user_id=current_user.id, user_id=current_user.id,
type="plagiarism", type="plagiarism",
@@ -126,30 +128,23 @@ async def upload_for_plagiarism_check(
"filename": file.filename, "filename": file.filename,
"minio_key": minio_key, "minio_key": minio_key,
"file_size_bytes": len(file_data), "file_size_bytes": len(file_data),
"content_type": file.content_type,
}, },
) )
db.add(task) db.add(task)
await db.flush() await db.flush()
task_id = task.id
# Диспатч в индексер воркер
celery_result = celery_app.send_task( celery_result = celery_app.send_task(
"index.extract_and_check", "index.extract_and_check",
args=[task_id, minio_key, file.filename], args=[task.id, minio_key, file.filename],
queue="queue.index", queue="queue.index",
) )
task.celery_task_id = celery_result.id task.celery_task_id = celery_result.id
task.queue_position = 1 task.queue_position = 1
task.eta_seconds = 120 # ~2 минуты с учётом 4 уровней проверки task.eta_seconds = 120
await db.commit() await db.commit()
await db.refresh(task) await db.refresh(task)
logger.info( logger.info("Задача плагиата %s создана для пользователя %d", task.public_id, current_user.id)
f"Задача проверки плагиата {task_id!r} создана для пользователя {current_user.id}, "
f"файл: {file.filename!r}"
)
return TaskResponse.model_validate(task) return TaskResponse.model_validate(task)

View File

@@ -6,8 +6,8 @@ from fastapi import APIRouter, Depends, HTTPException, status
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.core.celery_app import celery_app from app.core.celery_app import celery_app
from app.core.rate_limiter import check_and_increment_limit, check_concurrent_limit from app.core.rate_limiter import acquire_concurrent_slot, check_and_increment_limit
from app.core.security import get_current_user from app.core.security import get_current_verified_user
from app.database import get_db from app.database import get_db
from app.models.task import Task from app.models.task import Task
from app.models.user import User from app.models.user import User
@@ -17,41 +17,43 @@ logger = logging.getLogger(__name__)
router = APIRouter(prefix="/search", tags=["search"]) router = APIRouter(prefix="/search", tags=["search"])
# Примерное время ожидания в секундах для каждой позиции в очереди
ETA_PER_POSITION_SECONDS = 15 ETA_PER_POSITION_SECONDS = 15
@router.post("/", response_model=TaskResponse, status_code=status.HTTP_202_ACCEPTED) @router.post("/", response_model=TaskResponse, status_code=status.HTTP_202_ACCEPTED)
async def create_search_task( async def create_search_task(
data: SearchRequest, data: SearchRequest,
current_user: User = Depends(get_current_user), # Поиск требует подтверждённого email — защита от массового abuse
current_user: User = Depends(get_current_verified_user),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
) -> TaskResponse: ) -> TaskResponse:
""" """
Создать задачу семантического поиска источников. Создать задачу поиска источников.
Возвращает TaskResponse с queue_position и eta_seconds сразу. Возвращает сразу с task public_id и позицией в очереди.
Результат доступен через GET /tasks/{task_id} или WebSocket /ws/tasks/{task_id}. Результат через GET /tasks/{public_id} или WS /ws/tasks/{public_id}?token=JWT.
""" """
# Проверить лимит по тарифу # Проверяем лимит (атомарно, Lua-скрипт)
limit_check = check_and_increment_limit(current_user.id, "search", current_user.plan) limit_result = await check_and_increment_limit(current_user.id, "search", current_user.plan)
if not limit_check["allowed"]: if not limit_result["allowed"]:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS, status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail=( detail=(
f"Превышен дневной лимит поиска для тарифа '{current_user.plan}'. " f"Превышен дневной лимит поиска (тариф '{current_user.plan}'): "
f"Использовано {limit_check['current']} из {limit_check['limit']}." f"{limit_result['current']}/{limit_result['limit']}. "
f"Сбросится {limit_result['reset_at']}."
), ),
headers={"Retry-After": "86400"},
) )
# Проверить лимит одновременных задач # Проверяем лимит одновременных задач (атомарно)
if not check_concurrent_limit(current_user.id, current_user.plan): slot_acquired = await acquire_concurrent_slot(current_user.id, current_user.plan)
if not slot_acquired:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS, status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="Превышен лимит одновременных задач. Дождитесь завершения текущих.", detail="Превышен лимит одновременных задач. Дождитесь завершения текущих.",
) )
# Создать задачу в БД
task = Task( task = Task(
user_id=current_user.id, user_id=current_user.id,
type="search", type="search",
@@ -65,31 +67,21 @@ async def create_search_task(
}, },
) )
db.add(task) db.add(task)
await db.flush() await db.flush() # получаем id и public_id
task_id = task.id
# Диспатч в GPU воркер
celery_result = celery_app.send_task( celery_result = celery_app.send_task(
"gpu.search_semantic", "gpu.search_semantic",
args=[task_id, data.query], args=[task.id, data.query],
kwargs={ kwargs={"lang": data.lang, "year_from": data.year_from, "year_to": data.year_to},
"lang": data.lang,
"year_from": data.year_from,
"year_to": data.year_to,
},
queue="queue.gpu", queue="queue.gpu",
) )
task.celery_task_id = celery_result.id task.celery_task_id = celery_result.id
task.queue_position = 1 # TODO: реальный подсчёт через Redis task.queue_position = 1
task.eta_seconds = ETA_PER_POSITION_SECONDS task.eta_seconds = ETA_PER_POSITION_SECONDS
await db.commit() await db.commit()
await db.refresh(task) await db.refresh(task)
logger.info( logger.info("Задача поиска %s создана для пользователя %d", task.public_id, current_user.id)
f"Задача поиска {task_id!r} создана для пользователя {current_user.id}, "
f"запрос: {data.query[:50]!r}"
)
return TaskResponse.model_validate(task) return TaskResponse.model_validate(task)

View File

@@ -1,4 +1,8 @@
"""Роутер для управления задачами пользователя.""" """Управление задачами пользователя.
В URL используется public_id (не внутренний UUID), чтобы не раскрывать
структуру внутренних идентификаторов.
"""
import logging import logging
@@ -24,7 +28,10 @@ async def list_tasks(
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
) -> list[TaskResponse]: ) -> list[TaskResponse]:
"""Получить список задач текущего пользователя (от новых к старым).""" """Список задач текущего пользователя (новые первыми)."""
# Жёсткий потолок limit — не даём вытащить всю таблицу одним запросом
limit = min(limit, 100)
result = await db.execute( result = await db.execute(
select(Task) select(Task)
.where(Task.user_id == current_user.id) .where(Task.user_id == current_user.id)
@@ -32,52 +39,47 @@ async def list_tasks(
.limit(limit) .limit(limit)
.offset(offset) .offset(offset)
) )
tasks = result.scalars().all() return [TaskResponse.model_validate(t) for t in result.scalars().all()]
return [TaskResponse.model_validate(t) for t in tasks]
@router.get("/{task_id}", response_model=TaskResponse) @router.get("/{public_id}", response_model=TaskResponse)
async def get_task( async def get_task(
task_id: str, public_id: str,
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
) -> TaskResponse: ) -> TaskResponse:
"""Получить детали конкретной задачи.""" """Получить задачу по публичному ID. Возвращает 404 если задача чужая — не раскрываем факт существования."""
result = await db.execute( result = await db.execute(
select(Task).where(Task.id == task_id, Task.user_id == current_user.id) select(Task).where(
Task.public_id == public_id,
Task.user_id == current_user.id, # ownership проверяется в одном запросе
)
) )
task = result.scalar_one_or_none() task = result.scalar_one_or_none()
if task is None: if task is None:
raise HTTPException( raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Задача не найдена")
status_code=status.HTTP_404_NOT_FOUND,
detail="Задача не найдена",
)
return TaskResponse.model_validate(task) return TaskResponse.model_validate(task)
@router.delete("/{task_id}", status_code=status.HTTP_204_NO_CONTENT) @router.delete("/{public_id}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_task( async def delete_task(
task_id: str, public_id: str,
current_user: User = Depends(get_current_user), current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db), db: AsyncSession = Depends(get_db),
) -> None: ) -> None:
""" """Удалить задачу. Нельзя удалить задачу в статусе 'processing'."""
Удалить задачу.
Нельзя удалить задачу в статусе 'processing'.
"""
result = await db.execute( result = await db.execute(
select(Task).where(Task.id == task_id, Task.user_id == current_user.id) select(Task).where(
Task.public_id == public_id,
Task.user_id == current_user.id,
)
) )
task = result.scalar_one_or_none() task = result.scalar_one_or_none()
if task is None: if task is None:
raise HTTPException( raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Задача не найдена")
status_code=status.HTTP_404_NOT_FOUND,
detail="Задача не найдена",
)
if task.status == "processing": if task.status == "processing":
raise HTTPException( raise HTTPException(

View File

@@ -31,6 +31,7 @@ class Settings(BaseSettings):
MINIO_SECRET_KEY: str = "changeme" MINIO_SECRET_KEY: str = "changeme"
MINIO_BUCKET_DOCS: str = "documents" MINIO_BUCKET_DOCS: str = "documents"
MINIO_BUCKET_BACKUPS: str = "backups" MINIO_BUCKET_BACKUPS: str = "backups"
MINIO_SECURE: bool = False # True если MinIO за HTTPS
# Elasticsearch # Elasticsearch
ELASTICSEARCH_URL: str = "http://elasticsearch:9200" ELASTICSEARCH_URL: str = "http://elasticsearch:9200"

View File

@@ -0,0 +1,48 @@
"""Singleton MinIO клиент.
MinIO SDK держит пул HTTP-соединений внутри себя — создаём один экземпляр
на весь процесс. Бакеты проверяем один раз при старте приложения.
"""
import logging
from minio import Minio
from minio.error import S3Error
from app.config import settings
logger = logging.getLogger(__name__)
_client: Minio | None = None
_initialized_buckets: set[str] = set()
def get_minio() -> Minio:
"""Вернуть singleton MinIO клиент."""
global _client
if _client is None:
_client = Minio(
settings.MINIO_ENDPOINT,
access_key=settings.MINIO_ACCESS_KEY,
secret_key=settings.MINIO_SECRET_KEY,
secure=settings.MINIO_SECURE,
)
return _client
def ensure_bucket(bucket: str) -> None:
"""Создать бакет если не существует. Вызывается один раз при старте."""
if bucket in _initialized_buckets:
return
client = get_minio()
try:
if not client.bucket_exists(bucket):
client.make_bucket(bucket)
logger.info(f"MinIO: бакет '{bucket}' создан")
else:
logger.info(f"MinIO: бакет '{bucket}' существует")
_initialized_buckets.add(bucket)
except S3Error as e:
logger.error(f"MinIO: ошибка проверки бакета '{bucket}': {e}")
raise

View File

@@ -1,13 +1,22 @@
"""Redis-based rate limiter для проверки лимитов по тарифному плану.""" """Redis-based rate limiter с атомарными Lua-скриптами.
import json Проблема наивного подхода (GET → проверка → INCR):
- Race condition: 10 конкурентных запросов могут одновременно пройти GET,
увидеть значение ниже лимита и все инкрементировать.
Решение: один Lua-скрипт выполняется атомарно на стороне Redis.
Redis гарантирует, что между командами внутри скрипта нет других операций.
"""
import logging
from datetime import datetime, timezone from datetime import datetime, timezone
import redis from app.core.redis_client import get_redis
from app.config import settings logger = logging.getLogger(__name__)
# ─── Лимиты по тарифам ────────────────────────────────────────────────────────
# Лимиты по тарифным планам
PLAN_LIMITS: dict[str, dict[str, int | None]] = { PLAN_LIMITS: dict[str, dict[str, int | None]] = {
"free": { "free": {
"search_per_day": 10, "search_per_day": 10,
@@ -35,133 +44,125 @@ PLAN_LIMITS: dict[str, dict[str, int | None]] = {
}, },
} }
# Маппинг действий на ключи лимитов
ACTION_TO_LIMIT: dict[str, tuple[str, str]] = { ACTION_TO_LIMIT: dict[str, tuple[str, str]] = {
"search": ("search_per_day", "day"), "search": ("search_per_day", "day"),
"summarize": ("summarize_per_month", "month"), "summarize": ("summarize_per_month", "month"),
"plagiarism":("plagiarism_per_month", "month"), "plagiarism":("plagiarism_per_month", "month"),
"gost": ("gost_per_month", "month"), # ГОСТ всегда разрешён # gost не ограничен — не в маппинге
} }
# ─── Lua-скрипты ──────────────────────────────────────────────────────────────
def get_redis_client() -> redis.Redis: # Атомарная проверка + инкремент лимита.
"""Создать синхронный Redis клиент.""" # Возвращает [current_value, allowed] где allowed = 1 если OK, 0 если превышен.
return redis.from_url(settings.REDIS_URL, decode_responses=True) _LUA_CHECK_AND_INCR = """
local key = KEYS[1]
local limit = tonumber(ARGV[1])
local ttl = tonumber(ARGV[2])
local current = tonumber(redis.call('GET', key) or '0')
def _get_period_key(period: str) -> str: if current >= limit then
"""Получить строку периода для Redis ключа.""" return {current, 0}
now = datetime.now(timezone.utc) end
if period == "day":
return now.strftime("%Y-%m-%d")
elif period == "month":
return now.strftime("%Y-%m")
return now.strftime("%Y-%m-%d")
local new_val = redis.call('INCR', key)
def check_and_increment_limit(user_id: int, action: str, plan: str) -> dict: -- Устанавливаем TTL только при первом инкременте (когда ключ только что создан)
if new_val == 1 then
redis.call('EXPIRE', key, ttl)
end
return {new_val, 1}
""" """
Проверить лимит и инкрементировать счётчик.
Args: # Атомарная проверка + инкремент счётчика одновременных задач.
user_id: ID пользователя # Возвращает 1 если слот получен, 0 если превышен лимит.
action: Действие (search, plagiarism, summarize, gost) _LUA_ACQUIRE_CONCURRENT = """
plan: Тарифный план пользователя local key = KEYS[1]
local limit = tonumber(ARGV[1])
local ttl = tonumber(ARGV[2])
local current = tonumber(redis.call('GET', key) or '0')
if current >= limit then
return 0
end
redis.call('INCR', key)
redis.call('EXPIRE', key, ttl)
return 1
"""
def _period_suffix(period: str) -> str:
now = datetime.now(timezone.utc)
return now.strftime("%Y-%m-%d") if period == "day" else now.strftime("%Y-%m")
async def check_and_increment_limit(
user_id: int,
action: str,
plan: str,
) -> dict:
"""
Атомарно проверить лимит и инкрементировать счётчик.
Returns: Returns:
dict с полями: {allowed, current, limit, remaining, reset_at}
- allowed: bool — разрешено ли действие
- current: int — текущее количество использований
- limit: int | None — лимит (None = безлимит)
- remaining: int | None — осталось использований
- reset_at: str — когда сбрасывается счётчик
""" """
limits = PLAN_LIMITS.get(plan, PLAN_LIMITS["free"]) limits = PLAN_LIMITS.get(plan, PLAN_LIMITS["free"])
if action not in ACTION_TO_LIMIT: if action not in ACTION_TO_LIMIT:
# Неизвестное действие — разрешаем
return {"allowed": True, "current": 0, "limit": None, "remaining": None} return {"allowed": True, "current": 0, "limit": None, "remaining": None}
limit_key, period = ACTION_TO_LIMIT[action] limit_key, period = ACTION_TO_LIMIT[action]
limit_value = limits.get(limit_key) limit_value = limits.get(limit_key)
# Безлимитный план if limit_value is None: # безлимит
if limit_value is None: return {"allowed": True, "current": 0, "limit": None, "remaining": None}
period_str = _period_suffix(period)
redis_key = f"rl:{user_id}:{action}:{period_str}"
ttl = 86_400 if period == "day" else 86_400 * 32
r = get_redis()
result = await r.eval(_LUA_CHECK_AND_INCR, 1, redis_key, limit_value, ttl)
current, allowed = int(result[0]), bool(result[1])
return { return {
"allowed": True, "allowed": allowed,
"current": 0,
"limit": None,
"remaining": None,
"reset_at": None,
}
period_str = _get_period_key(period)
redis_key = f"rate:{user_id}:{action}:{period_str}"
r = get_redis_client()
# Атомарно инкрементировать
pipe = r.pipeline()
pipe.incr(redis_key)
# Устанавливаем TTL: для дня — 86400 сек, для месяца — 32 дня
ttl = 86400 if period == "day" else 86400 * 32
pipe.expire(redis_key, ttl)
results = pipe.execute()
current = results[0]
if current > limit_value:
# Декрементировать обратно (не считать запрещённые)
r.decr(redis_key)
current -= 1
return {
"allowed": False,
"current": current, "current": current,
"limit": limit_value, "limit": limit_value,
"remaining": 0, "remaining": max(0, limit_value - current) if allowed else 0,
"reset_at": period_str,
}
return {
"allowed": True,
"current": current,
"limit": limit_value,
"remaining": limit_value - current,
"reset_at": period_str, "reset_at": period_str,
} }
def check_concurrent_limit(user_id: int, plan: str) -> bool: async def acquire_concurrent_slot(user_id: int, plan: str) -> bool:
""" """
Проверить лимит одновременных задач. Атомарно захватить слот одновременной задачи.
Returns: Returns:
True если можно создать новую задачу, False если превышен лимит. True — слот получен (задачу можно создавать).
False — все слоты заняты.
""" """
limits = PLAN_LIMITS.get(plan, PLAN_LIMITS["free"]) limits = PLAN_LIMITS.get(plan, PLAN_LIMITS["free"])
max_concurrent = limits.get("concurrent", 1) max_concurrent = limits.get("concurrent", 1)
r = get_redis_client() r = get_redis()
result = await r.eval(
_LUA_ACQUIRE_CONCURRENT,
1,
f"concurrent:{user_id}",
max_concurrent,
3_600, # TTL 1 час — автосброс если воркер упал не освободив слот
)
return bool(result)
async def release_concurrent_slot(user_id: int) -> None:
"""Освободить слот одновременной задачи после завершения."""
r = get_redis()
key = f"concurrent:{user_id}" key = f"concurrent:{user_id}"
current = r.get(key) # DECR безопасен: Redis не уходит в отрицательные значения если мы контролируем acquire
current = await r.get(key)
return (current is None) or (int(current) < max_concurrent)
def increment_concurrent(user_id: int) -> None:
"""Увеличить счётчик одновременных задач (при создании задачи)."""
r = get_redis_client()
key = f"concurrent:{user_id}"
pipe = r.pipeline()
pipe.incr(key)
pipe.expire(key, 3600) # Автосброс через 1 час
pipe.execute()
def decrement_concurrent(user_id: int) -> None:
"""Уменьшить счётчик одновременных задач (при завершении задачи)."""
r = get_redis_client()
key = f"concurrent:{user_id}"
current = r.get(key)
if current and int(current) > 0: if current and int(current) > 0:
r.decr(key) await r.decr(key)

View File

@@ -0,0 +1,36 @@
"""Singleton пул Redis-соединений.
Один ConnectionPool на весь процесс — не создаём новое соединение на каждый запрос.
Используем redis.asyncio для совместимости с FastAPI async endpoints.
"""
import redis.asyncio as aioredis
from app.config import settings
_pool: aioredis.ConnectionPool | None = None
def get_pool() -> aioredis.ConnectionPool:
"""Инициализировать пул при первом вызове, затем возвращать тот же."""
global _pool
if _pool is None:
_pool = aioredis.ConnectionPool.from_url(
settings.REDIS_URL,
max_connections=50,
decode_responses=True,
)
return _pool
def get_redis() -> aioredis.Redis:
"""Получить Redis клиент из пула (не создаёт новое TCP-соединение)."""
return aioredis.Redis(connection_pool=get_pool())
async def close_pool() -> None:
"""Закрыть пул при остановке приложения."""
global _pool
if _pool:
await _pool.aclose()
_pool = None

View File

@@ -1,9 +1,11 @@
"""Утилиты безопасности: JWT, хэширование паролей, dependency для получения текущего пользователя.""" """JWT, хэширование паролей, dependency получения пользователя с кэшированием."""
import json
import logging
from datetime import datetime, timedelta, timezone from datetime import datetime, timedelta, timezone
from typing import Any from typing import Any
from fastapi import Depends, HTTPException, status from fastapi import Depends, HTTPException, Query, WebSocket, status
from fastapi.security import OAuth2PasswordBearer from fastapi.security import OAuth2PasswordBearer
from jose import JWTError, jwt from jose import JWTError, jwt
from passlib.context import CryptContext from passlib.context import CryptContext
@@ -11,53 +13,47 @@ from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from app.config import settings from app.config import settings
from app.core.redis_client import get_redis
from app.database import get_db from app.database import get_db
# Контекст хэширования паролей (bcrypt) logger = logging.getLogger(__name__)
pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
# OAuth2 схема pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login") oauth2_scheme = OAuth2PasswordBearer(tokenUrl="/api/auth/login")
# TTL кэша пользователя в Redis (5 минут)
_USER_CACHE_TTL = 300
def hash_password(password: str) -> str: def hash_password(password: str) -> str:
"""Хэшировать пароль через bcrypt."""
return pwd_context.hash(password) return pwd_context.hash(password)
def verify_password(plain_password: str, hashed_password: str) -> bool: def verify_password(plain: str, hashed: str) -> bool:
"""Проверить соответствие пароля его хэшу.""" return pwd_context.verify(plain, hashed)
return pwd_context.verify(plain_password, hashed_password)
def create_access_token(data: dict[str, Any], expires_delta: timedelta | None = None) -> str: def create_access_token(data: dict[str, Any], expires_delta: timedelta | None = None) -> str:
"""Создать JWT токен доступа.""" payload = data.copy()
to_encode = data.copy()
expire = datetime.now(timezone.utc) + ( expire = datetime.now(timezone.utc) + (
expires_delta or timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES) expires_delta or timedelta(minutes=settings.ACCESS_TOKEN_EXPIRE_MINUTES)
) )
to_encode.update({"exp": expire}) payload["exp"] = expire
return jwt.encode(to_encode, settings.SECRET_KEY, algorithm=settings.ALGORITHM) return jwt.encode(payload, settings.SECRET_KEY, algorithm=settings.ALGORITHM)
def verify_token(token: str) -> dict[str, Any]: def _decode_token(token: str) -> int:
""" """
Декодировать и верифицировать JWT токен. Декодировать JWT и вернуть user_id.
Raises HTTPException 401 при любой проблеме с токеном.
Raises:
HTTPException: если токен невалиден или просрочен.
""" """
try: try:
payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM]) payload = jwt.decode(token, settings.SECRET_KEY, algorithms=[settings.ALGORITHM])
user_id: int | None = payload.get("sub") sub = payload.get("sub")
if user_id is None: if sub is None:
raise HTTPException( raise ValueError("missing sub")
status_code=status.HTTP_401_UNAUTHORIZED, return int(sub)
detail="Неверный токен аутентификации", except (JWTError, ValueError):
headers={"WWW-Authenticate": "Bearer"},
)
return payload
except JWTError:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED, status_code=status.HTTP_401_UNAUTHORIZED,
detail="Невалидный или просроченный токен", detail="Невалидный или просроченный токен",
@@ -65,46 +61,79 @@ def verify_token(token: str) -> dict[str, Any]:
) )
async def get_current_user( async def _load_user(user_id: int, db: AsyncSession):
token: str = Depends(oauth2_scheme), """Загрузить пользователя из Redis-кэша или из БД."""
db: AsyncSession = Depends(get_db), from app.models.user import User # избегаем circular import на уровне модуля
):
"""
FastAPI dependency: получить текущего авторизованного пользователя из JWT токена.
Raises: cache_key = f"user:cache:{user_id}"
HTTPException 401: если токен невалиден. r = get_redis()
HTTPException 404: если пользователь не найден.
"""
from app.models.user import User
payload = verify_token(token) # Пробуем кэш
user_id: int = int(payload.get("sub")) cached = await r.get(cache_key)
if cached:
data = json.loads(cached)
# Возвращаем "живой" объект из БД только по id, но без лишнего SELECT
# Создаём User без ORM-связей (достаточно для проверок в роутерах)
u = User.__new__(User)
u.__dict__.update(data)
return u
# Кэш пустой — идём в БД
result = await db.execute(select(User).where(User.id == user_id)) result = await db.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none() user = result.scalar_one_or_none()
if user is None: if user is None:
raise HTTPException( raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="Пользователь не найден")
status_code=status.HTTP_404_NOT_FOUND,
detail="Пользователь не найден", # Сохраняем в кэш (только безопасные поля, без hashed_password)
) safe = {
"id": user.id,
"email": user.email,
"name": user.name,
"plan": user.plan,
"is_verified": user.is_verified,
}
await r.setex(cache_key, _USER_CACHE_TTL, json.dumps(safe))
return user return user
async def get_current_verified_user( async def get_current_user(
current_user=Depends(get_current_user), token: str = Depends(oauth2_scheme),
db: AsyncSession = Depends(get_db),
): ):
""" """Dependency: текущий пользователь из JWT с Redis-кэшем."""
FastAPI dependency: только верифицированные пользователи. user_id = _decode_token(token)
return await _load_user(user_id, db)
Raises:
HTTPException 403: если email не подтверждён. async def get_current_verified_user(current_user=Depends(get_current_user)):
""" """Dependency: только верифицированные пользователи."""
if not current_user.is_verified: if not current_user.is_verified:
raise HTTPException( raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN, status_code=status.HTTP_403_FORBIDDEN,
detail="Необходимо подтвердить email адрес", detail="Необходимо подтвердить email адрес",
) )
return current_user return current_user
async def get_ws_user(
websocket: WebSocket,
token: str = Query(..., description="JWT токен (передаётся как query-параметр ?token=...)"),
db: AsyncSession = Depends(get_db),
):
"""
Dependency для WebSocket: аутентификация через query-параметр ?token=JWT.
WebSocket API браузера не позволяет передавать Authorization header,
поэтому токен передаётся в query-строке:
ws://host/ws/tasks/{public_id}?token=<JWT>
"""
user_id = _decode_token(token)
return await _load_user(user_id, db)
async def invalidate_user_cache(user_id: int) -> None:
"""Сбросить кэш пользователя (при смене тарифа, пароля и т.д.)."""
r = get_redis()
await r.delete(f"user:cache:{user_id}")

View File

@@ -1,66 +1,75 @@
"""Точка входа FastAPI приложения — Академический помощник (anti-plagiarism).""" """Точка входа FastAPI приложения — Академический помощник."""
import logging import logging
from contextlib import asynccontextmanager from contextlib import asynccontextmanager
from typing import AsyncGenerator from typing import AsyncGenerator
from fastapi import FastAPI, WebSocket, WebSocketDisconnect from fastapi import FastAPI, WebSocket, WebSocketDisconnect, Depends
from fastapi.middleware.cors import CORSMiddleware from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse from fastapi.responses import JSONResponse
from app.api import auth, documents, reports, search, tasks from app.api import auth, documents, reports, search, tasks
from app.config import settings from app.config import settings
from app.core.minio_client import ensure_bucket
from app.core.redis_client import close_pool
from app.core.security import get_ws_user
from app.core.websocket_manager import ws_manager from app.core.websocket_manager import ws_manager
from app.database import engine from app.database import engine
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
logging.basicConfig(level=logging.INFO) logging.basicConfig(
level=logging.INFO,
format="%(asctime)s %(levelname)s %(name)s: %(message)s",
)
@asynccontextmanager @asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]: async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
""" """Инициализация при старте, очистка при остановке."""
Lifecycle менеджер приложения.
Выполняет инициализацию при старте и очистку при остановке.
"""
logger.info("Запуск Академического помощника...") logger.info("Запуск Академического помощника...")
# Проверка соединений при старте # Проверяем PostgreSQL
try: try:
import sqlalchemy
async with engine.connect() as conn: async with engine.connect() as conn:
await conn.execute(__import__("sqlalchemy").text("SELECT 1")) await conn.execute(sqlalchemy.text("SELECT 1"))
logger.info("PostgreSQL: соединение установлено") logger.info("PostgreSQL: OK")
except Exception as e: except Exception as e:
logger.error(f"PostgreSQL: ошибка соединения: {e}") logger.error(f"PostgreSQL: {e}")
# Инициализируем MinIO бакеты (один раз при старте, не на каждый запрос)
try:
ensure_bucket(settings.MINIO_BUCKET_DOCS)
ensure_bucket(settings.MINIO_BUCKET_BACKUPS)
logger.info("MinIO: бакеты готовы")
except Exception as e:
logger.error(f"MinIO: {e}")
yield yield
# Очистка при остановке
logger.info("Остановка приложения...") logger.info("Остановка приложения...")
await engine.dispose() await engine.dispose()
await close_pool()
app = FastAPI( app = FastAPI(
title="Академический помощник API", title="Академический помощник API",
description=(
"Сервис антиплагиата и поиска академических источников. "
"Семантический поиск через FAISS GPU + Elasticsearch, "
"проверка плагиата в 4 уровня, ГОСТ библиография."
),
version="0.1.0", version="0.1.0",
lifespan=lifespan, lifespan=lifespan,
docs_url="/api/docs", docs_url="/api/docs" if settings.ENVIRONMENT == "development" else None,
redoc_url="/api/redoc", redoc_url=None,
openapi_url="/api/openapi.json", openapi_url="/api/openapi.json" if settings.ENVIRONMENT == "development" else None,
) )
# ─── CORS ───────────────────────────────────────────────────────────────────── # ─── CORS ─────────────────────────────────────────────────────────────────────
# allow_credentials=True несовместим с allow_origins=["*"].
# Указываем явные origins. Wildcard заголовки/методы запрещены с credentials.
app.add_middleware( app.add_middleware(
CORSMiddleware, CORSMiddleware,
allow_origins=settings.CORS_ORIGINS, allow_origins=settings.CORS_ORIGINS,
allow_credentials=True, allow_credentials=True,
allow_methods=["*"], allow_methods=["GET", "POST", "DELETE", "OPTIONS"],
allow_headers=["*"], allow_headers=["Authorization", "Content-Type", "Accept"],
) )
# ─── Роутеры ────────────────────────────────────────────────────────────────── # ─── Роутеры ──────────────────────────────────────────────────────────────────
@@ -72,42 +81,55 @@ app.include_router(reports.router, prefix="/api")
# ─── WebSocket ──────────────────────────────────────────────────────────────── # ─── WebSocket ────────────────────────────────────────────────────────────────
@app.websocket("/ws/tasks/{task_id}") @app.websocket("/ws/tasks/{public_id}")
async def websocket_task_updates(websocket: WebSocket, task_id: str) -> None: async def websocket_task_updates(
websocket: WebSocket,
public_id: str,
# Аутентификация через ?token=JWT (браузеры не могут передавать Authorization header в WS)
current_user=Depends(get_ws_user),
) -> None:
""" """
WebSocket endpoint для real-time обновлений статуса задачи. Подписка на обновления задачи в реальном времени.
Клиент подключается и получает обновления при изменении статуса задачи. Подключение: ws://host/ws/tasks/{public_id}?token=<JWT>
Соединение закрывается при получении статуса done/failed.
Пользователь получает только обновления СВОИХ задач —
ownership проверяется до установки соединения.
""" """
await ws_manager.connect(task_id, websocket) from sqlalchemy import select
from app.database import AsyncSessionLocal
from app.models.task import Task
# Проверяем что задача принадлежит этому пользователю
async with AsyncSessionLocal() as db:
result = await db.execute(
select(Task).where(
Task.public_id == public_id,
Task.user_id == current_user.id,
)
)
task = result.scalar_one_or_none()
if task is None:
await websocket.close(code=4004, reason="Task not found or access denied")
return
await ws_manager.connect(public_id, websocket)
try: try:
while True: while True:
# Ждём входящих сообщений (ping/pong для поддержания соединения)
data = await websocket.receive_text() data = await websocket.receive_text()
if data == "ping": if data == "ping":
await websocket.send_text("pong") await websocket.send_text("pong")
except WebSocketDisconnect: except WebSocketDisconnect:
ws_manager.disconnect(task_id, websocket) ws_manager.disconnect(public_id, websocket)
logger.info(f"WebSocket клиент отключился от задачи {task_id!r}")
# ─── Health check ────────────────────────────────────────────────────────────── # ─── Health ────────────────────────────────────────────────────────────────────
@app.get("/health", tags=["monitoring"]) @app.get("/health", tags=["monitoring"])
async def health_check() -> JSONResponse: async def health_check() -> JSONResponse:
"""Проверка работоспособности сервиса.""" return JSONResponse({"status": "ok", "version": "0.1.0"})
return JSONResponse(
content={
"status": "ok",
"service": "api",
"version": "0.1.0",
}
)
@app.get("/", include_in_schema=False) @app.get("/", include_in_schema=False)
async def root() -> JSONResponse: async def root() -> JSONResponse:
"""Корневой endpoint — редирект к документации.""" return JSONResponse({"message": "Академический помощник API"})
return JSONResponse(
content={"message": "Академический помощник API", "docs": "/api/docs"}
)

View File

@@ -14,14 +14,19 @@ class TaskCreate(BaseModel):
class TaskResponse(BaseModel): class TaskResponse(BaseModel):
"""Ответ с информацией о задаче.""" """Ответ с информацией о задаче.
id: str В API возвращается public_id (не внутренний UUID).
Клиент использует public_id для всех операций с задачей.
"""
# public_id — это то, что видит клиент и использует в URL
public_id: str
type: str type: str
status: str status: str
queue_position: int | None = None queue_position: int | None = None
eta_seconds: int | None = None eta_seconds: int | None = None
input_data: dict[str, Any] = Field(default_factory=dict) # input_data не возвращаем (может содержать minio_key и другие внутренние данные)
result: dict[str, Any] | None = None result: dict[str, Any] | None = None
error: str | None = None error: str | None = None
created_at: datetime created_at: datetime

View File

@@ -4,7 +4,7 @@ sqlalchemy==2.0.30
alembic==1.13.1 alembic==1.13.1
asyncpg==0.29.0 asyncpg==0.29.0
psycopg2-binary==2.9.9 psycopg2-binary==2.9.9
redis==5.0.4 redis[asyncio]==5.0.4
celery==5.4.0 celery==5.4.0
pydantic==2.7.1 pydantic==2.7.1
pydantic-settings==2.2.1 pydantic-settings==2.2.1