Files
vidconf/backend/tests/test_pipeline.py

560 lines
23 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Тесты оркестрации AI-пайплайна (`workers.tasks.pipeline.run_pipeline_async`).
Как и `test_maintenance.py`, не использует savepoint-фикстуру `db_session`:
`run_pipeline_async` открывает собственную сессию с отдельным engine
(`workers/db.py::open_session`), которая не видит незакоммиченные изменения
другой сессии. Тестовые данные заводятся и коммитятся напрямую через
`core.db.engine`, транскрайбер — фейк за контрактом `Transcriber` (без
faster-whisper), файлы треков — реальные (пустые) в `tmp_path`.
"""
import uuid
from collections.abc import AsyncGenerator
from datetime import UTC, datetime, timedelta
from pathlib import Path
from typing import cast
from unittest.mock import MagicMock
import pytest
from celery.exceptions import MaxRetriesExceededError
from kombu.exceptions import OperationalError as KombuOperationalError
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncConnection
from core.db import engine
from core.plugins.config import ChatConfig, InstanceConfig, SummarizerConfig, TranscriberConfig
from core.plugins.transcriber import Segment
from services.conference_ids import generate_number, generate_slug
from workers.tasks import dispatch as dispatch_module
from workers.tasks import pipeline as pipeline_module
from workers.tasks.pipeline import run_pipeline_async
NOW = datetime.now(UTC)
class _FakeTranscriber:
"""Заглушка `Transcriber`: канонические сегменты по пути файла + счётчик вызовов.
`fail_once_for` — пути, для которых первый вызов бросает исключение
(симуляция падения воркера посреди транскрибации, тест №12), повторный
вызов для того же пути уже успешен.
"""
def __init__(
self, canned: dict[str, list[Segment]], fail_once_for: set[str] | None = None
) -> None:
self._canned = canned
self._fail_once_for = set(fail_once_for or set())
self.calls: list[str] = []
def transcribe(self, audio_path: str, language: str = "ru") -> list[Segment]:
self.calls.append(audio_path)
if audio_path in self._fail_once_for:
self._fail_once_for.discard(audio_path)
raise RuntimeError("симуляция падения воркера посреди транскрибации")
return self._canned[audio_path]
class _FakeTask:
"""Заглушка bound-задачи Celery: фиксирует вызовы `retry`, не бросает исключение.
В реальном Celery `Task.retry()` сам бросает `Retry`/`MaxRetriesExceededError`
и не возвращает управление — здесь для теста №14 нужен именно факт вызова.
"""
def __init__(self) -> None:
self.retry = MagicMock()
class _ExhaustedRetryTask:
"""Заглушка bound-задачи: `retry` всегда бросает `MaxRetriesExceededError`.
Имитирует исчерпание попыток ожидания egress (реальный `Task.retry()`
бросает именно это исключение, когда `max_retries` уже выбраны).
"""
def __init__(self) -> None:
self.retry_calls = 0
def retry(self, countdown: int | None = None) -> None:
self.retry_calls += 1
raise MaxRetriesExceededError("исчерпаны попытки ожидания egress")
class _Fixture:
"""Id тестовой конференции/сеанса/участников + пути файлов треков."""
def __init__(self, tmp_path: Path) -> None:
self.conference_id = uuid.uuid4()
self.session_id = uuid.uuid4()
self.user_id = uuid.uuid4()
self.guest_id = uuid.uuid4()
self.user_participant_id = uuid.uuid4()
self.guest_participant_id = uuid.uuid4()
self.number = generate_number()
self.slug = generate_slug()
self.t_start = NOW - timedelta(minutes=30)
self.t_end = NOW
self.track1_path = str(tmp_path / "track1.ogg")
self.track2_path = str(tmp_path / "track2.ogg")
# Реальные (пустые) файлы — `run_pipeline_async` проверяет их наличие
# перед вызовом транскрайбера; запись — синхронно в конструкторе, не
# в async-фикстуре (ASYNC240: pathlib в async-функции).
Path(self.track1_path).write_bytes(b"")
Path(self.track2_path).write_bytes(b"")
@pytest.fixture
async def fx(tmp_path: Path) -> AsyncGenerator[_Fixture, None]:
f = _Fixture(tmp_path)
async with engine.connect() as conn:
await conn.execute(
text(
"INSERT INTO users (id, email, name_user, password_hash) "
"VALUES (:id, :email, 'Pipeline Tester', 'x')"
),
{"id": f.user_id, "email": f"pipeline-{f.user_id}@example.com"},
)
await conn.execute(
text(
"INSERT INTO conferences (id, number, slug, title, status, is_pinned) "
"VALUES (:id, :number, :slug, 'Pipeline test', 'ended', false)"
),
{"id": f.conference_id, "number": f.number, "slug": f.slug},
)
await conn.execute(
text(
"INSERT INTO guest_access (id, conference_id, display_name) "
"VALUES (:id, :conference_id, 'Guest Tester')"
),
{"id": f.guest_id, "conference_id": f.conference_id},
)
await conn.execute(
text(
"INSERT INTO conference_sessions (id, conference_id, title, t_start, t_end) "
"VALUES (:id, :conference_id, 'Pipeline session', :t_start, :t_end)"
),
{
"id": f.session_id,
"conference_id": f.conference_id,
"t_start": f.t_start,
"t_end": f.t_end,
},
)
await conn.execute(
text(
"INSERT INTO conference_participants (id, session_id, user_id, joined_at, left_at) "
"VALUES (:id, :session_id, :user_id, :joined_at, :left_at)"
),
{
"id": f.user_participant_id,
"session_id": f.session_id,
"user_id": f.user_id,
"joined_at": f.t_start,
"left_at": f.t_end,
},
)
await conn.execute(
text(
"INSERT INTO conference_participants "
"(id, session_id, guest_id, joined_at, left_at) "
"VALUES (:id, :session_id, :guest_id, :joined_at, :left_at)"
),
{
"id": f.guest_participant_id,
"session_id": f.session_id,
"guest_id": f.guest_id,
"joined_at": f.t_start,
"left_at": f.t_end,
},
)
await conn.commit()
yield f
async with engine.connect() as conn:
await conn.execute(text("DELETE FROM phrases WHERE session_id = :id"), {"id": f.session_id})
await conn.execute(
text("DELETE FROM session_audio_tracks WHERE session_id = :id"), {"id": f.session_id}
)
await conn.execute(
text("DELETE FROM conference_participants WHERE session_id = :id"), {"id": f.session_id}
)
await conn.execute(
text("DELETE FROM conference_sessions WHERE conference_id = :id"),
{"id": f.conference_id},
)
await conn.execute(
text("DELETE FROM guest_access WHERE conference_id = :id"), {"id": f.conference_id}
)
await conn.execute(text("DELETE FROM conferences WHERE id = :id"), {"id": f.conference_id})
await conn.execute(text("DELETE FROM users WHERE id = :id"), {"id": f.user_id})
await conn.commit()
async def _insert_track(
conn: AsyncConnection,
*,
session_id: uuid.UUID,
participant_id: uuid.UUID,
track_sid: str,
file_path: str,
status: str,
started_at: datetime,
) -> None:
await conn.execute(
text(
"INSERT INTO session_audio_tracks "
"(id, session_id, participant_id, track_sid, egress_id, file_path, status, started_at) "
"VALUES (:id, :session_id, :participant_id, :track_sid, :egress_id, :file_path, "
":status, :started_at)"
),
{
"id": uuid.uuid4(),
"session_id": session_id,
"participant_id": participant_id,
"track_sid": track_sid,
"egress_id": f"EG_{track_sid}",
"file_path": file_path,
"status": status,
"started_at": started_at,
},
)
def _cfg(*, enabled: bool = True) -> InstanceConfig:
"""Конфиг для тестов: реальный `Transcriber` не создаётся — провайдер
подменяется через monkeypatch `pipeline_module.create_transcriber`."""
return InstanceConfig(
transcriber=TranscriberConfig(enabled=enabled, provider="fake", language="ru"),
summarizer=SummarizerConfig(),
chat=ChatConfig(),
)
async def _fetch_session_status(session_id: uuid.UUID) -> str:
async with engine.connect() as conn:
result = await conn.execute(
text("SELECT pipeline_status FROM conference_sessions WHERE id = :id"),
{"id": session_id},
)
return cast("str", result.scalar_one())
async def _fetch_track_status(session_id: uuid.UUID, track_sid: str) -> str:
async with engine.connect() as conn:
result = await conn.execute(
text(
"SELECT status FROM session_audio_tracks "
"WHERE session_id = :session_id AND track_sid = :track_sid"
),
{"session_id": session_id, "track_sid": track_sid},
)
return cast("str", result.scalar_one())
async def _fetch_phrases(session_id: uuid.UUID) -> list[dict[str, object]]:
async with engine.connect() as conn:
result = await conn.execute(
text(
"SELECT participant_id, data, t_start, t_end FROM phrases "
"WHERE session_id = :id ORDER BY t_start"
),
{"id": session_id},
)
return [dict(row) for row in result.mappings().all()]
async def test_run_pipeline_reconstructs_phrases_for_user_and_guest(
fx: _Fixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""№11: сеанс с user- и гость-участником, оба трека 'recorded'
phrases + status=summarizing."""
async with engine.connect() as conn:
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.user_participant_id,
track_sid="TR_USER",
file_path=fx.track1_path,
status="recorded",
started_at=fx.t_start + timedelta(seconds=5),
)
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.guest_participant_id,
track_sid="TR_GUEST",
file_path=fx.track2_path,
status="recorded",
started_at=fx.t_start + timedelta(seconds=10),
)
await conn.commit()
fake = _FakeTranscriber(
canned={
fx.track1_path: [Segment(0.0, 2.0, "привет")],
fx.track2_path: [Segment(0.0, 2.0, "привет в ответ")],
}
)
monkeypatch.setattr(pipeline_module, "create_transcriber", lambda cfg: fake)
mock_send_task = MagicMock()
monkeypatch.setattr(pipeline_module.app, "send_task", mock_send_task)
await run_pipeline_async(_FakeTask(), fx.session_id, plugins_config=_cfg())
assert await _fetch_session_status(fx.session_id) == "summarizing"
mock_send_task.assert_called_once_with(
"workers.tasks.summarize.summarize_session", args=[str(fx.session_id)]
)
phrases = await _fetch_phrases(fx.session_id)
assert len(phrases) == 2
assert phrases[0]["participant_id"] == fx.user_participant_id
assert phrases[0]["data"] == "привет"
assert phrases[0]["t_start"] == fx.t_start + timedelta(seconds=5)
assert phrases[0]["t_end"] == fx.t_start + timedelta(seconds=7)
assert phrases[1]["participant_id"] == fx.guest_participant_id
assert phrases[1]["data"] == "привет в ответ"
assert phrases[1]["t_start"] == fx.t_start + timedelta(seconds=10)
assert phrases[1]["t_end"] == fx.t_start + timedelta(seconds=12)
async def test_run_pipeline_is_idempotent_after_mid_run_crash(
fx: _Fixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""№12: падение на треке №2 → повтор транскрибирует только его, без дублей phrases."""
async with engine.connect() as conn:
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.user_participant_id,
track_sid="TR_USER",
file_path=fx.track1_path,
status="recorded",
started_at=fx.t_start + timedelta(seconds=5),
)
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.guest_participant_id,
track_sid="TR_GUEST",
file_path=fx.track2_path,
status="recorded",
started_at=fx.t_start + timedelta(seconds=10),
)
await conn.commit()
fake = _FakeTranscriber(
canned={
fx.track1_path: [Segment(0.0, 2.0, "привет")],
fx.track2_path: [Segment(0.0, 2.0, "привет в ответ")],
},
fail_once_for={fx.track2_path},
)
monkeypatch.setattr(pipeline_module, "create_transcriber", lambda cfg: fake)
monkeypatch.setattr(pipeline_module.app, "send_task", MagicMock())
with pytest.raises(RuntimeError):
await run_pipeline_async(_FakeTask(), fx.session_id, plugins_config=_cfg())
assert fake.calls == [fx.track1_path, fx.track2_path]
assert await _fetch_session_status(fx.session_id) == "transcribing"
assert await _fetch_phrases(fx.session_id) == []
await run_pipeline_async(_FakeTask(), fx.session_id, plugins_config=_cfg())
assert fake.calls == [fx.track1_path, fx.track2_path, fx.track2_path]
assert await _fetch_session_status(fx.session_id) == "summarizing"
phrases = await _fetch_phrases(fx.session_id)
assert len(phrases) == 2
async def test_run_pipeline_noop_when_transcriber_disabled(
fx: _Fixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""№13: `transcriber.enabled=false` → no-op, status остаётся 'recording'."""
async with engine.connect() as conn:
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.user_participant_id,
track_sid="TR_USER",
file_path=fx.track1_path,
status="recorded",
started_at=fx.t_start + timedelta(seconds=5),
)
await conn.commit()
mock_factory = MagicMock()
monkeypatch.setattr(pipeline_module, "create_transcriber", mock_factory)
await run_pipeline_async(_FakeTask(), fx.session_id, plugins_config=_cfg(enabled=False))
mock_factory.assert_not_called()
assert await _fetch_session_status(fx.session_id) == "recording"
assert await _fetch_phrases(fx.session_id) == []
async def test_run_pipeline_retries_while_tracks_still_recording(
fx: _Fixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""№14: трек ещё 'recording' → `task.retry` вызван, транскрайбер не запускался."""
async with engine.connect() as conn:
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.user_participant_id,
track_sid="TR_USER",
file_path=fx.track1_path,
status="recording",
started_at=fx.t_start + timedelta(seconds=5),
)
await conn.commit()
mock_factory = MagicMock()
monkeypatch.setattr(pipeline_module, "create_transcriber", mock_factory)
task = _FakeTask()
await run_pipeline_async(task, fx.session_id, plugins_config=_cfg())
task.retry.assert_called_once_with(countdown=pipeline_module.RETRY_COUNTDOWN_S)
mock_factory.assert_not_called()
assert await _fetch_session_status(fx.session_id) == "recording"
async def test_run_pipeline_marks_hung_track_failed_when_retries_exhausted(
fx: _Fixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Исчерпание retry (`MaxRetriesExceededError`): зависший 'recording' трек
помечается 'failed', обработка продолжается с остальными треками —
трек №2 транскрибируется, пайплайн доходит до 'summarizing'."""
async with engine.connect() as conn:
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.user_participant_id,
track_sid="TR_USER_HUNG",
file_path=fx.track1_path,
status="recording",
started_at=fx.t_start + timedelta(seconds=5),
)
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.guest_participant_id,
track_sid="TR_GUEST",
file_path=fx.track2_path,
status="recorded",
started_at=fx.t_start + timedelta(seconds=10),
)
await conn.commit()
fake = _FakeTranscriber(canned={fx.track2_path: [Segment(0.0, 2.0, "привет в ответ")]})
monkeypatch.setattr(pipeline_module, "create_transcriber", lambda cfg: fake)
monkeypatch.setattr(pipeline_module.app, "send_task", MagicMock())
task = _ExhaustedRetryTask()
await run_pipeline_async(task, fx.session_id, plugins_config=_cfg())
assert task.retry_calls == 1
assert await _fetch_track_status(fx.session_id, "TR_USER_HUNG") == "failed"
assert await _fetch_track_status(fx.session_id, "TR_GUEST") == "transcribed"
assert fake.calls == [fx.track2_path]
assert await _fetch_session_status(fx.session_id) == "summarizing"
phrases = await _fetch_phrases(fx.session_id)
assert len(phrases) == 1
assert phrases[0]["participant_id"] == fx.guest_participant_id
async def test_run_pipeline_fails_when_all_tracks_failed(
fx: _Fixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""№15: все треки 'failed' → pipeline_status='failed'."""
async with engine.connect() as conn:
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.user_participant_id,
track_sid="TR_USER",
file_path=fx.track1_path,
status="failed",
started_at=fx.t_start + timedelta(seconds=5),
)
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.guest_participant_id,
track_sid="TR_GUEST",
file_path=fx.track2_path,
status="failed",
started_at=fx.t_start + timedelta(seconds=10),
)
await conn.commit()
fake = _FakeTranscriber(canned={})
monkeypatch.setattr(pipeline_module, "create_transcriber", lambda cfg: fake)
await run_pipeline_async(_FakeTask(), fx.session_id, plugins_config=_cfg())
assert await _fetch_session_status(fx.session_id) == "failed"
assert fake.calls == []
async def test_send_summarize_task_retries_then_succeeds_on_broker_failure(
fx: _Fixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Постановка `summarize_session` (защита от сбоя брокера): первая попытка бросает
`kombu.exceptions.OperationalError`, вторая — успешна; `run_pipeline` не падает."""
async with engine.connect() as conn:
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.user_participant_id,
track_sid="TR_USER",
file_path=fx.track1_path,
status="recorded",
started_at=fx.t_start + timedelta(seconds=5),
)
await conn.commit()
fake = _FakeTranscriber(canned={fx.track1_path: [Segment(0.0, 2.0, "привет")]})
monkeypatch.setattr(pipeline_module, "create_transcriber", lambda cfg: fake)
monkeypatch.setattr(dispatch_module, "SEND_TASK_BACKOFF_S", 0.0)
mock_send_task = MagicMock(side_effect=[KombuOperationalError("брокер недоступен"), None])
monkeypatch.setattr(pipeline_module.app, "send_task", mock_send_task)
await run_pipeline_async(_FakeTask(), fx.session_id, plugins_config=_cfg())
assert mock_send_task.call_count == 2
assert await _fetch_session_status(fx.session_id) == "summarizing"
async def test_send_summarize_task_exhausted_retries_does_not_raise(
fx: _Fixture, monkeypatch: pytest.MonkeyPatch
) -> None:
"""Все попытки постановки `summarize_session` исчерпаны (брокер недоступен постоянно):
`run_pipeline` не бросает исключение, фразы и статус 'summarizing' сохранены —
восстановление ложится на `recover_stuck_summaries` (maintenance)."""
async with engine.connect() as conn:
await _insert_track(
conn,
session_id=fx.session_id,
participant_id=fx.user_participant_id,
track_sid="TR_USER",
file_path=fx.track1_path,
status="recorded",
started_at=fx.t_start + timedelta(seconds=5),
)
await conn.commit()
fake = _FakeTranscriber(canned={fx.track1_path: [Segment(0.0, 2.0, "привет")]})
monkeypatch.setattr(pipeline_module, "create_transcriber", lambda cfg: fake)
monkeypatch.setattr(dispatch_module, "SEND_TASK_BACKOFF_S", 0.0)
mock_send_task = MagicMock(side_effect=KombuOperationalError("брокер недоступен"))
monkeypatch.setattr(pipeline_module.app, "send_task", mock_send_task)
await run_pipeline_async(_FakeTask(), fx.session_id, plugins_config=_cfg())
assert mock_send_task.call_count == dispatch_module.SEND_TASK_MAX_ATTEMPTS
assert await _fetch_session_status(fx.session_id) == "summarizing"
assert len(await _fetch_phrases(fx.session_id)) == 1