From 5871ce65cec2cae01d8765fdfef9b8c70640b47d Mon Sep 17 00:00:00 2001 From: baekyutae Date: Tue, 14 Jul 2026 15:11:02 +0900 Subject: [PATCH 1/6] =?UTF-8?q?feat(pipeline=20worker)=20stt=20=EA=B2=B0?= =?UTF-8?q?=EA=B3=BC=EB=AC=BC=20=EC=97=90=EB=9F=AC=EB=B0=9C=EC=83=9D?= =?UTF-8?q?=EC=8B=9C=20=EB=A1=9C=EA=B7=B8=EB=82=A8=EA=B8=B0=EA=B2=8C=20?= =?UTF-8?q?=ED=95=A8=20(#100)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../src/infra/ai/stt_batch_callable.py | 38 ++++++++++++- .../tests/unit/test_stt_batch_callable.py | 55 +++++++++++++++++++ 2 files changed, 91 insertions(+), 2 deletions(-) diff --git a/services/pipeline-worker/src/infra/ai/stt_batch_callable.py b/services/pipeline-worker/src/infra/ai/stt_batch_callable.py index d992dc9..ecddbe1 100644 --- a/services/pipeline-worker/src/infra/ai/stt_batch_callable.py +++ b/services/pipeline-worker/src/infra/ai/stt_batch_callable.py @@ -3,13 +3,20 @@ import asyncio from typing import Any +from google.api_core.client_options import ClientOptions +from google.rpc import code_pb2 from loguru import logger from src.infra.ai.google_stt_adapter import ExternalAIAdapterError, STTCallable -from google.api_core.client_options import ClientOptions _MAX_WORDS_PER_SEGMENT = 100 _SENTENCE_ENDING_MARKS = (".", "!", "?") +_FILE_ERROR_POLICY = { + code_pb2.DEADLINE_EXCEEDED: ("TIMEOUT", True), + code_pb2.RESOURCE_EXHAUSTED: ("RATE_LIMITED", True), + code_pb2.UNAVAILABLE: ("UNAVAILABLE", True), + code_pb2.INVALID_ARGUMENT: ("INVALID_REQUEST", False), +} # google stt: 초단위 -> biblio: 밀리초 단위 변환 def _duration_to_ms(duration: Any, trace_id: str) -> int: @@ -37,6 +44,32 @@ def _stt_parse_error(message: str, trace_id: str) -> ExternalAIAdapterError: ) +def _raise_for_file_result_error(uri: str, file_result: Any, trace_id: str) -> None: + error = getattr(file_result, "error", None) + error_code = int(getattr(error, "code", code_pb2.OK)) + if error_code == code_pb2.OK: + return + + error_message = str(getattr(error, "message", "")).strip() or "No message provided" + # Google 오류 코드를 Biblio 오류 코드와 재시도 여부로 변환 + app_code, retryable = _FILE_ERROR_POLICY.get( + error_code, + ("INTERNAL_ERROR", False), + ) + detail = ( + f"STT BatchRecognize file failed uri={uri} " + f"error_code={error_code} error_message={error_message}" + ) + logger.bind(trace_id=trace_id).error(detail) + raise ExternalAIAdapterError( + code=app_code, + message=detail, + trace_id=trace_id, + provider="google-stt", + retryable=retryable, + ) + + def _build_segment(words: list[Any], trace_id: str) -> dict: text = " ".join(_word_text(word) for word in words).strip() if not text: @@ -71,7 +104,8 @@ def _segments_from_words(words: list[Any], trace_id: str) -> list[dict]: def _parse_batch_recognize_response(response: Any, stt_model_version: str, trace_id: str = "") -> dict: segments: list[dict] = [] - for _uri, file_result in response.results.items(): + for uri, file_result in response.results.items(): + _raise_for_file_result_error(uri, file_result, trace_id) transcript = getattr(getattr(file_result, "inline_result", None), "transcript", None) if transcript is None: transcript = getattr(file_result, "transcript", None) diff --git a/services/pipeline-worker/tests/unit/test_stt_batch_callable.py b/services/pipeline-worker/tests/unit/test_stt_batch_callable.py index c79b99c..1107c5d 100644 --- a/services/pipeline-worker/tests/unit/test_stt_batch_callable.py +++ b/services/pipeline-worker/tests/unit/test_stt_batch_callable.py @@ -1,6 +1,7 @@ from types import SimpleNamespace import pytest +from google.rpc import code_pb2 from src.infra.ai.google_stt_adapter import ExternalAIAdapterError from src.infra.ai.stt_batch_callable import _parse_batch_recognize_response @@ -61,6 +62,60 @@ def test_parse_batch_recognize_response_reads_inline_result_transcript() -> None ] +def test_parse_batch_recognize_response_preserves_file_error_details() -> None: + response = SimpleNamespace( + results={ + "gs://bucket/audio.flac": SimpleNamespace( + error=SimpleNamespace( + code=code_pb2.INVALID_ARGUMENT, + message="Audio duration exceeds the allowed limit", + ) + ) + } + ) + + with pytest.raises(ExternalAIAdapterError) as error_info: + _parse_batch_recognize_response(response, "chirp_3", trace_id="trace-file-error") + + error = error_info.value + assert error.code == "INVALID_REQUEST" + assert error.message == ( + "STT BatchRecognize file failed uri=gs://bucket/audio.flac " + "error_code=3 error_message=Audio duration exceeds the allowed limit" + ) + assert error.trace_id == "trace-file-error" + assert error.provider == "google-stt" + assert error.retryable is False + + +@pytest.mark.parametrize( + ("provider_code", "expected_code"), + [ + (code_pb2.DEADLINE_EXCEEDED, "TIMEOUT"), + (code_pb2.RESOURCE_EXHAUSTED, "RATE_LIMITED"), + (code_pb2.UNAVAILABLE, "UNAVAILABLE"), + ], +) +def test_parse_batch_recognize_response_keeps_retryable_file_error_policy( + provider_code: int, + expected_code: str, +) -> None: + response = SimpleNamespace( + results={ + "gs://bucket/audio.flac": SimpleNamespace( + error=SimpleNamespace(code=provider_code, message="Temporary provider error") + ) + } + ) + + with pytest.raises(ExternalAIAdapterError) as error_info: + _parse_batch_recognize_response(response, "chirp_3", trace_id="trace-retryable-error") + + error = error_info.value + assert error.code == expected_code + assert error.retryable is True + + def test_parse_batch_recognize_response_uses_word_offsets_for_segment_timestamps() -> None: words = [ _word("Hello", 0.1, 0.4), From 22348d73abc9d2ce08db45bda5980d9ae1ee7cbb Mon Sep 17 00:00:00 2001 From: baekyutae Date: Tue, 14 Jul 2026 15:29:43 +0900 Subject: [PATCH 2/6] =?UTF-8?q?fix(infra):=20=EC=BB=B4=ED=8F=AC=EB=84=8C?= =?UTF-8?q?=ED=8A=B8=EB=B3=84=EB=A1=9C=20=EB=B3=84=EB=8F=84=EC=9D=98=20ima?= =?UTF-8?q?ge=20tag=EB=A5=BC=20=EA=B0=80=EC=A7=80=EA=B2=8C=20=EC=88=98?= =?UTF-8?q?=EC=A0=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- docs/runbooks/gcp-performance-deployment.md | 77 ++++++++++++++----- infra/terraform/envs/gcp-perf/main.tf | 22 +++--- .../envs/gcp-perf/terraform.tfvars.example | 11 ++- infra/terraform/envs/gcp-perf/variables.tf | 14 +++- scripts/deploy/build_and_push_images.sh | 44 +++++++++++ 5 files changed, 135 insertions(+), 33 deletions(-) diff --git a/docs/runbooks/gcp-performance-deployment.md b/docs/runbooks/gcp-performance-deployment.md index a645431..d003589 100644 --- a/docs/runbooks/gcp-performance-deployment.md +++ b/docs/runbooks/gcp-performance-deployment.md @@ -141,17 +141,25 @@ cp infra/terraform/envs/gcp-perf/terraform.tfvars.example \ - `region` - `zone` - `name_prefix` -- `image_tag` +- `image_tags` - `db_password` - `jwt_secret_key` - subnet CIDR - 임베딩 VM machine type과 disk 크기 - `model_artifact_path` -현재 모든 서비스가 하나의 `image_tag`를 공유한다. 빌드한 image tag와 `terraform.tfvars`의 `image_tag`가 반드시 같아야 한다. +각 image는 독립적인 tag를 사용한다. 처음 배포할 때는 모두 같은 tag로 시작해도 된다. ```hcl -image_tag = "" +image_tags = { + core_api = "" + search_service = "" + frontend = "" + managed_embedding_endpoint = "" + pipeline_worker = "" + feedback_ingestion_pipeline = "" + feedback_loop_pipeline = "" +} ``` `terraform.tfvars`는 secret을 포함하므로 Git에 커밋하지 않는다. @@ -181,10 +189,11 @@ terraform -chdir=infra/terraform/envs/gcp-perf apply \ ### 3.3 전체 image 빌드와 push -현재 스크립트는 다음 image 6개를 같은 tag로 빌드한다. +인자를 생략하면 다음 image 7개를 같은 tag로 모두 빌드한다. - `core-api` - `search-service` +- `frontend` - `managed-embedding-endpoint` - `pipeline-worker` - `feedback-ingestion-pipeline` @@ -199,7 +208,7 @@ IMAGE_TAG="$IMAGE_TAG" \ bash scripts/deploy/build_and_push_images.sh ``` -빌드한 `IMAGE_TAG`를 `terraform.tfvars`의 `image_tag`에도 입력한다. +빌드한 `IMAGE_TAG`를 `terraform.tfvars`의 `image_tags` 항목에도 입력한다. ### 3.4 DB와 migration Job 생성 @@ -538,20 +547,23 @@ gcloud storage rm "$latest" --project="$GCP_PROJECT_ID" ## 5. 코드 변경 후 재배포 -### 5.1 중요한 제한 +### 5.1 배포 원칙 -현재 Terraform은 모든 서비스 image에 하나의 `image_tag`를 사용한다. +각 image tag는 `terraform.tfvars`의 `image_tags`에서 따로 관리한다. -일부 서비스만 새 image로 수동 배포한 상태에서 예전 `image_tag`로 전체 `terraform apply`를 실행하면, 수동 배포한 서비스가 예전 image로 돌아갈 수 있다. +코드가 변경된 image만 빌드하고 해당 tag만 바꾼다. 그 뒤에도 `-target`이 아닌 전체 plan과 apply를 실행한다. Terraform은 image 주소가 바뀐 서비스만 갱신한다. 재배포 전에 반드시 다음을 확인한다. ```bash git rev-parse --short HEAD -grep '^image_tag' infra/terraform/envs/gcp-perf/terraform.tfvars +terraform -chdir=infra/terraform/envs/gcp-perf console \ + <<< 'var.image_tags' ``` -### 5.2 전체 서비스 재배포 +### 5.2 일부 image 재배포 + +빌드할 image 이름을 스크립트 인자로 넘긴다. 여러 개를 한 번에 지정할 수도 있다. ```bash export IMAGE_TAG=$(git rev-parse --short HEAD) @@ -559,10 +571,25 @@ export IMAGE_TAG=$(git rev-parse --short HEAD) PROJECT_ID="$GCP_PROJECT_ID" \ REGION="$GCP_REGION" \ IMAGE_TAG="$IMAGE_TAG" \ -bash scripts/deploy/build_and_push_images.sh +bash scripts/deploy/build_and_push_images.sh frontend ``` -`terraform.tfvars`의 `image_tag`를 같은 값으로 변경한 뒤 plan과 apply를 실행한다. +두 image를 함께 빌드하는 예시는 다음과 같다. + +```bash +PROJECT_ID="$GCP_PROJECT_ID" \ +REGION="$GCP_REGION" \ +IMAGE_TAG="$IMAGE_TAG" \ +bash scripts/deploy/build_and_push_images.sh frontend core-api +``` + +빌드한 image에 해당하는 `image_tags` 값만 `IMAGE_TAG`와 같은 값으로 바꾼다. 예를 들어 frontend만 빌드했다면 기존 `image_tags` 블록에서 다음 한 줄의 값만 바꾼다. 나머지 항목은 기존 값을 유지한다. + +```hcl +frontend = "" +``` + +전체 plan에서 의도하지 않은 서비스 변경이나 VM 교체가 없는지 확인한 뒤 apply한다. ```bash terraform -chdir=infra/terraform/envs/gcp-perf plan \ @@ -572,18 +599,30 @@ terraform -chdir=infra/terraform/envs/gcp-perf apply \ /tmp/biblio-gcp-perf.tfplan ``` -DB schema가 변경됐다면 migration Job을 다시 실행한다. +frontend만 변경했다면 plan에 embedding VM의 `must be replaced`가 나타나면 안 된다. + +### 5.3 전체 image 재배포 -### 5.3 일부 서비스만 임시 검증 +```bash +export IMAGE_TAG=$(git rev-parse --short HEAD) -일부 image만 별도 tag로 배포하는 방식은 임시 검증에만 사용한다. +PROJECT_ID="$GCP_PROJECT_ID" \ +REGION="$GCP_REGION" \ +IMAGE_TAG="$IMAGE_TAG" \ +bash scripts/deploy/build_and_push_images.sh +``` -검증이 끝나면 다음 둘 중 하나를 선택한다. +`terraform.tfvars`의 모든 `image_tags`를 같은 값으로 변경한 뒤 plan과 apply를 실행한다. -1. 전체 image를 같은 정식 tag로 다시 빌드하고 Terraform에 반영한다. -2. Terraform을 서비스별 image tag 구조로 개선한 뒤 정식 반영한다. +```bash +terraform -chdir=infra/terraform/envs/gcp-perf plan \ + -out=/tmp/biblio-gcp-perf.tfplan -임시 tag 상태를 그대로 둔 채 전체 `terraform apply`를 실행하지 않는다. +terraform -chdir=infra/terraform/envs/gcp-perf apply \ + /tmp/biblio-gcp-perf.tfplan +``` + +DB schema가 변경됐다면 migration Job을 다시 실행한다. ## 6. worker와 FIP 풀 가동 diff --git a/infra/terraform/envs/gcp-perf/main.tf b/infra/terraform/envs/gcp-perf/main.tf index 1613471..bd78b56 100644 --- a/infra/terraform/envs/gcp-perf/main.tf +++ b/infra/terraform/envs/gcp-perf/main.tf @@ -4,13 +4,13 @@ locals { image_registry = "${var.region}-docker.pkg.dev/${var.project_id}/${local.repository_id}" service_images = { - "core-api" = "${local.image_registry}/core-api:${var.image_tag}" - "search-service" = "${local.image_registry}/search-service:${var.image_tag}" - "frontend" = "${local.image_registry}/frontend:${var.image_tag}" - "managed-embedding-endpoint" = "${local.image_registry}/managed-embedding-endpoint:${var.image_tag}" - "pipeline-worker" = "${local.image_registry}/pipeline-worker:${var.image_tag}" - "feedback-ingestion-pipeline" = "${local.image_registry}/feedback-ingestion-pipeline:${var.image_tag}" - "feedback-loop-pipeline" = "${local.image_registry}/feedback-loop-pipeline:${var.image_tag}" + "core-api" = "${local.image_registry}/core-api:${var.image_tags.core_api}" + "search-service" = "${local.image_registry}/search-service:${var.image_tags.search_service}" + "frontend" = "${local.image_registry}/frontend:${var.image_tags.frontend}" + "managed-embedding-endpoint" = "${local.image_registry}/managed-embedding-endpoint:${var.image_tags.managed_embedding_endpoint}" + "pipeline-worker" = "${local.image_registry}/pipeline-worker:${var.image_tags.pipeline_worker}" + "feedback-ingestion-pipeline" = "${local.image_registry}/feedback-ingestion-pipeline:${var.image_tags.feedback_ingestion_pipeline}" + "feedback-loop-pipeline" = "${local.image_registry}/feedback-loop-pipeline:${var.image_tags.feedback_loop_pipeline}" } bucket_names = { @@ -441,10 +441,10 @@ module "pipeline_worker" { VISION_MAX_OUTPUT_TOKENS = "2048" WORKER_CONCURRENCY = "4" # 임베딩 VM의 wireproxy(WARP) SOCKS5. YouTube 트래픽만 이 프록시로 우회한다. - YOUTUBE_PROXY_URL = "socks5://${module.embedding_vm.private_ip}:1080" - GCS_VIDEO_BUCKET_NAME = module.object_storage.bucket_names.video - EMBEDDING_API_URL = local.embedding_vm_url - EMBEDDING_TIMEOUT_SEC = "60" + YOUTUBE_PROXY_URL = "socks5://${module.embedding_vm.private_ip}:1080" + GCS_VIDEO_BUCKET_NAME = module.object_storage.bucket_names.video + EMBEDDING_API_URL = local.embedding_vm_url + EMBEDDING_TIMEOUT_SEC = "60" } secret_env_vars = { diff --git a/infra/terraform/envs/gcp-perf/terraform.tfvars.example b/infra/terraform/envs/gcp-perf/terraform.tfvars.example index 27a34c9..7a2cca6 100644 --- a/infra/terraform/envs/gcp-perf/terraform.tfvars.example +++ b/infra/terraform/envs/gcp-perf/terraform.tfvars.example @@ -2,7 +2,16 @@ project_id = "biblio-perf-example" region = "asia-northeast3" zone = "asia-northeast3-a" name_prefix = "biblio-perf" -image_tag = "0000000" + +image_tags = { + core_api = "0000000" + search_service = "0000000" + frontend = "0000000" + managed_embedding_endpoint = "0000000" + pipeline_worker = "0000000" + feedback_ingestion_pipeline = "0000000" + feedback_loop_pipeline = "0000000" +} db_password = "replace-with-postgres-password" jwt_secret_key = "replace-with-jwt-secret-key" diff --git a/infra/terraform/envs/gcp-perf/variables.tf b/infra/terraform/envs/gcp-perf/variables.tf index ba5c3ff..a3ee940 100644 --- a/infra/terraform/envs/gcp-perf/variables.tf +++ b/infra/terraform/envs/gcp-perf/variables.tf @@ -14,8 +14,18 @@ variable "name_prefix" { type = string } -variable "image_tag" { - type = string +variable "image_tags" { + type = object({ + core_api = string + search_service = string + frontend = string + managed_embedding_endpoint = string + pipeline_worker = string + feedback_ingestion_pipeline = string + feedback_loop_pipeline = string + }) + + description = "Image tag for each independently built application image." } variable "frontend_origin" { diff --git a/scripts/deploy/build_and_push_images.sh b/scripts/deploy/build_and_push_images.sh index ebebfac..1f81f25 100644 --- a/scripts/deploy/build_and_push_images.sh +++ b/scripts/deploy/build_and_push_images.sh @@ -18,10 +18,54 @@ services=( "feedback-loop-pipeline:services/feedback-loop-pipeline:services/feedback-loop-pipeline/Dockerfile" ) +requested_services=("$@") + +service_exists() { + local requested_name="$1" + local item + + for item in "${services[@]}"; do + if [[ "${item%%:*}" == "${requested_name}" ]]; then + return 0 + fi + done + + return 1 +} + +should_build() { + local service_name="$1" + local requested_name + + if (( ${#requested_services[@]} == 0 )); then + return 0 + fi + + for requested_name in "${requested_services[@]}"; do + if [[ "${requested_name}" == "${service_name}" ]]; then + return 0 + fi + done + + return 1 +} + +for requested_name in "${requested_services[@]}"; do + if ! service_exists "${requested_name}"; then + echo "Unknown service: ${requested_name}" >&2 + echo "Available services: ${services[*]%%:*}" >&2 + exit 2 + fi +done + gcloud auth configure-docker "${REGION}-docker.pkg.dev" --quiet for item in "${services[@]}"; do name="${item%%:*}" + if ! should_build "${name}"; then + continue + fi + rest="${item#*:}" context="${rest%%:*}" dockerfile="${rest#*:}" From e668d1dad61311d533ffac9eb4668eb6e9254da2 Mon Sep 17 00:00:00 2001 From: baekyutae Date: Wed, 15 Jul 2026 11:47:00 +0900 Subject: [PATCH 3/6] =?UTF-8?q?fix(pipeline-worker):=2020=EB=B6=84=20?= =?UTF-8?q?=EC=B4=88=EA=B3=BC=20=EC=9D=8C=EC=84=B1=EC=9D=84=20=EB=B6=84?= =?UTF-8?q?=ED=95=A0=ED=95=B4=20STT=20=EC=B2=98=EB=A6=AC=ED=95=98=EB=8F=84?= =?UTF-8?q?=EB=A1=9D=20=EC=88=98=EC=A0=95=20(#100)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- services/pipeline-worker/src/bootstrap.py | 18 + .../pipeline-worker/src/config/settings.py | 23 +- .../src/infra/ai/google_stt_adapter.py | 75 +++- .../src/infra/ai/stt_batch_callable.py | 59 ++- .../src/infra/db/artifact_repository.py | 34 +- .../src/infra/db/video_repository.py | 119 +++++- .../src/infra/media/ffmpeg_client.py | 87 ++++- .../src/infra/media/youtube_downloader.py | 5 + .../src/services/long_audio_transcription.py | 267 ++++++++++++++ .../src/services/pipeline_errors.py | 6 + .../src/services/pipeline_orchestrator.py | 134 ++++++- .../src/services/transcript_merge_service.py | 76 ++++ .../src/usecases/delete_video.py | 6 + .../src/usecases/process_video.py | 86 ++++- .../tests/integration/test_delete_project.py | 55 ++- .../test_long_audio_transcription.py | 343 ++++++++++++++++++ .../tests/integration/test_repositories.py | 97 +++++ services/pipeline-worker/tests/support.py | 24 +- .../tests/unit/test_delete_video.py | 26 ++ .../tests/unit/test_ffmpeg_adapter.py | 63 +++- .../tests/unit/test_google_stt_adapter.py | 97 +++++ .../tests/unit/test_pipeline_orchestrator.py | 16 + .../tests/unit/test_process_video.py | 138 +++++++ .../tests/unit/test_settings.py | 45 ++- .../tests/unit/test_stt_batch_callable.py | 5 + .../unit/test_transcript_merge_service.py | 67 ++++ .../tests/unit/test_youtube_downloader.py | 36 +- 27 files changed, 1903 insertions(+), 104 deletions(-) create mode 100644 services/pipeline-worker/src/services/long_audio_transcription.py create mode 100644 services/pipeline-worker/src/services/pipeline_errors.py create mode 100644 services/pipeline-worker/src/services/transcript_merge_service.py create mode 100644 services/pipeline-worker/tests/integration/test_long_audio_transcription.py create mode 100644 services/pipeline-worker/tests/unit/test_transcript_merge_service.py diff --git a/services/pipeline-worker/src/bootstrap.py b/services/pipeline-worker/src/bootstrap.py index 094cb1a..ed2bfd2 100644 --- a/services/pipeline-worker/src/bootstrap.py +++ b/services/pipeline-worker/src/bootstrap.py @@ -27,7 +27,9 @@ from src.config.settings import Settings from src.schemas.messages import MessageType from src.services.chunking_service import ChunkingService +from src.services.long_audio_transcription import LongAudioTranscriptionService from src.services.pipeline_orchestrator import PipelineOrchestrator +from src.services.transcript_merge_service import TranscriptMergeService from src.usecases.delete_project import DeleteProjectUseCase from src.usecases.delete_video import DeleteVideoUseCase from src.usecases.process_video import ProcessVideoUseCase @@ -182,6 +184,18 @@ async def create_production_bootstrap(settings: Settings) -> None: max_tokens=settings.chunk_max_tokens, overlap_sentences=settings.chunk_overlap_sentences, ) + long_audio_transcription_service = LongAudioTranscriptionService( + artifact_repository=artifact_repo, + video_repository=video_repo, + storage_client=storage_client, + ffmpeg_client=ffmpeg_client, + stt_adapter=stt_adapter, + merge_service=TranscriptMergeService(), + part_duration_sec=settings.audio_part_duration_sec, + part_overlap_sec=settings.audio_part_overlap_sec, + stt_concurrency=settings.stt_part_concurrency, + processing_timeout_sec=settings.audio_processing_timeout_sec, + ) orchestrator = PipelineOrchestrator( video_repository=video_repo, @@ -194,10 +208,14 @@ async def create_production_bootstrap(settings: Settings) -> None: vision_adapter=vision_adapter, workdir_manager=workdir_manager, chunking_service=chunking_service, + long_audio_transcription_service=long_audio_transcription_service, embedding_batch_size=settings.embedding_batch_size, stt_model_version=settings.stt_model_version or "chirp_2", embedding_model_version=settings.embedding_model_version, release_context_repository=release_context_repo, + max_audio_duration_sec=settings.max_audio_duration_sec, + max_source_size_bytes=settings.youtube_max_filesize_bytes, + audio_processing_timeout_sec=settings.audio_processing_timeout_sec, ) delete_uc = DeleteVideoUseCase( diff --git a/services/pipeline-worker/src/config/settings.py b/services/pipeline-worker/src/config/settings.py index 05d8728..b16605b 100644 --- a/services/pipeline-worker/src/config/settings.py +++ b/services/pipeline-worker/src/config/settings.py @@ -1,7 +1,7 @@ from functools import lru_cache -from typing import Literal +from typing import Literal, Self -from pydantic import Field +from pydantic import Field, model_validator from pydantic_settings import BaseSettings, SettingsConfigDict @@ -42,7 +42,16 @@ class Settings(BaseSettings): ) max_retries: int = Field(default=3, alias="MAX_RETRIES", ge=0) download_timeout_sec: int = Field(default=600, alias="DOWNLOAD_TIMEOUT_SEC", ge=1) - youtube_max_duration_sec: int = Field(default=1800, alias="YOUTUBE_MAX_DURATION_SEC", ge=1) + max_audio_duration_sec: int = Field(default=3600, alias="MAX_AUDIO_DURATION_SEC", ge=1) + audio_part_duration_sec: int = Field(default=900, alias="AUDIO_PART_DURATION_SEC", ge=1) + audio_part_overlap_sec: int = Field(default=5, alias="AUDIO_PART_OVERLAP_SEC", ge=0) + stt_part_concurrency: int = Field(default=2, alias="STT_PART_CONCURRENCY", ge=1, le=2) + audio_processing_timeout_sec: int = Field( + default=120, + alias="AUDIO_PROCESSING_TIMEOUT_SEC", + ge=1, + ) + youtube_max_duration_sec: int = Field(default=3600, alias="YOUTUBE_MAX_DURATION_SEC", ge=1) youtube_max_filesize_bytes: int = Field( default=500 * 1024 * 1024, alias="YOUTUBE_MAX_FILESIZE_BYTES", @@ -68,6 +77,14 @@ class Settings(BaseSettings): chunk_overlap_sentences: int = Field(default=1, alias="CHUNK_OVERLAP_SENTENCES", ge=0) poll_interval_sec: float = Field(default=1.0, alias="POLL_INTERVAL_SEC", ge=0.1) + @model_validator(mode="after") + def validate_audio_part_settings(self) -> Self: + if self.audio_part_overlap_sec >= self.audio_part_duration_sec: + raise ValueError("AUDIO_PART_OVERLAP_SEC must be less than AUDIO_PART_DURATION_SEC") + if self.audio_part_duration_sec + self.audio_part_overlap_sec > 20 * 60: + raise ValueError("Audio parts including overlap must not exceed 20 minutes") + return self + @lru_cache(maxsize=1) def get_settings() -> Settings: diff --git a/services/pipeline-worker/src/infra/ai/google_stt_adapter.py b/services/pipeline-worker/src/infra/ai/google_stt_adapter.py index 9b0976e..5321bd3 100644 --- a/services/pipeline-worker/src/infra/ai/google_stt_adapter.py +++ b/services/pipeline-worker/src/infra/ai/google_stt_adapter.py @@ -1,7 +1,14 @@ import asyncio +import random from dataclasses import dataclass from typing import Any, Awaitable, Callable +from loguru import logger + + +MAX_WORDS_PER_SEGMENT = 100 +SENTENCE_ENDING_MARKS = (".", "!", "?") + @dataclass(slots=True) class TranscriptSegmentDTO: @@ -10,10 +17,18 @@ class TranscriptSegmentDTO: end_ms: int +@dataclass(slots=True) +class TranscriptWordDTO: + text: str + start_ms: int + end_ms: int + + @dataclass(slots=True) class STTTranscriptionResult: segments: list[TranscriptSegmentDTO] stt_model_version: str + words: list[TranscriptWordDTO] | None = None @dataclass(slots=True) @@ -23,12 +38,38 @@ class ExternalAIAdapterError(Exception): trace_id: str provider: str retryable: bool + attempt_count: int = 1 def __str__(self) -> str: return f"{self.provider}:{self.code}:{self.message}" STTCallable = Callable[[str, str], Awaitable[dict[str, Any] | STTTranscriptionResult]] +SleepCallable = Callable[[float], Awaitable[None]] +JitterCallable = Callable[[], float] + + +def segments_from_words(words: list[TranscriptWordDTO]) -> list[TranscriptSegmentDTO]: + segments: list[TranscriptSegmentDTO] = [] + buffered_words: list[TranscriptWordDTO] = [] + for word in words: + if not word.text: + continue + buffered_words.append(word) + if word.text.endswith(SENTENCE_ENDING_MARKS) or len(buffered_words) >= MAX_WORDS_PER_SEGMENT: + segments.append(_segment_from_words(buffered_words)) + buffered_words = [] + if buffered_words: + segments.append(_segment_from_words(buffered_words)) + return segments + + +def _segment_from_words(words: list[TranscriptWordDTO]) -> TranscriptSegmentDTO: + return TranscriptSegmentDTO( + text=" ".join(word.text for word in words), + start_ms=words[0].start_ms, + end_ms=words[-1].end_ms, + ) class GoogleSTTAdapter: @@ -37,9 +78,13 @@ def __init__( client: STTCallable, *, max_retries: int, + sleep: SleepCallable = asyncio.sleep, + jitter: JitterCallable = random.random, ) -> None: self._client = client self._max_retries = max_retries + self._sleep = sleep + self._jitter = jitter async def transcribe(self, *, audio_uri: str, trace_id: str) -> STTTranscriptionResult: if not audio_uri.startswith("gs://"): @@ -51,7 +96,7 @@ async def transcribe(self, *, audio_uri: str, trace_id: str) -> STTTranscription retryable=False, ) - last_error: Exception | None = None + last_error: ExternalAIAdapterError | None = None for attempt in range(self._max_retries + 1): try: response = await self._client(audio_uri, trace_id) @@ -68,10 +113,19 @@ async def transcribe(self, *, audio_uri: str, trace_id: str) -> STTTranscription last_error = exc if not exc.retryable: raise + assert last_error is not None + last_error.attempt_count = attempt + 1 if attempt >= self._max_retries: - assert last_error is not None raise last_error - await asyncio.sleep(0) + delay_seconds = (2**attempt) * (1 + (self._jitter() * 0.25)) + logger.bind(trace_id=trace_id).warning( + "STT retry uri={} attempt={} code={} delay_seconds={:.3f}", + audio_uri, + attempt + 1, + last_error.code, + delay_seconds, + ) + await self._sleep(delay_seconds) assert last_error is not None raise last_error @@ -90,6 +144,15 @@ def _normalize(self, response: dict[str, Any] | STTTranscriptionResult, trace_id provider="google-stt", retryable=False, ) + raw_words = response.get("words") if "words" in response else None + normalized_words = [ + TranscriptWordDTO( + text=str(word["text"]), + start_ms=int(word["start_ms"]), + end_ms=int(word["end_ms"]), + ) + for word in sorted(raw_words or [], key=lambda item: int(item["start_ms"])) + ] normalized_segments = [ TranscriptSegmentDTO( text=segment["text"], @@ -98,4 +161,8 @@ def _normalize(self, response: dict[str, Any] | STTTranscriptionResult, trace_id ) for segment in sorted(segments, key=lambda item: int(item["start_ms"])) ] - return STTTranscriptionResult(segments=normalized_segments, stt_model_version=str(model_version)) + return STTTranscriptionResult( + segments=normalized_segments, + stt_model_version=str(model_version), + words=normalized_words if raw_words is not None else None, + ) diff --git a/services/pipeline-worker/src/infra/ai/stt_batch_callable.py b/services/pipeline-worker/src/infra/ai/stt_batch_callable.py index ecddbe1..11f335a 100644 --- a/services/pipeline-worker/src/infra/ai/stt_batch_callable.py +++ b/services/pipeline-worker/src/infra/ai/stt_batch_callable.py @@ -7,10 +7,13 @@ from google.rpc import code_pb2 from loguru import logger -from src.infra.ai.google_stt_adapter import ExternalAIAdapterError, STTCallable +from src.infra.ai.google_stt_adapter import ( + ExternalAIAdapterError, + STTCallable, + TranscriptWordDTO, + segments_from_words, +) -_MAX_WORDS_PER_SEGMENT = 100 -_SENTENCE_ENDING_MARKS = (".", "!", "?") _FILE_ERROR_POLICY = { code_pb2.DEADLINE_EXCEEDED: ("TIMEOUT", True), code_pb2.RESOURCE_EXHAUSTED: ("RATE_LIMITED", True), @@ -30,10 +33,6 @@ def _word_text(word: Any) -> str: return str(getattr(word, "word", "")).strip() -def _is_sentence_end(text: str) -> bool: - return text.endswith(_SENTENCE_ENDING_MARKS) - - def _stt_parse_error(message: str, trace_id: str) -> ExternalAIAdapterError: return ExternalAIAdapterError( code="INTERNAL_ERROR", @@ -70,40 +69,22 @@ def _raise_for_file_result_error(uri: str, file_result: Any, trace_id: str) -> N ) -def _build_segment(words: list[Any], trace_id: str) -> dict: - text = " ".join(_word_text(word) for word in words).strip() +def _normalize_word(word: Any, trace_id: str) -> TranscriptWordDTO: + text = _word_text(word) if not text: raise _stt_parse_error("STT word text missing", trace_id) try: - start_ms = _duration_to_ms(words[0].start_offset, trace_id) - end_ms = _duration_to_ms(words[-1].end_offset, trace_id) + start_ms = _duration_to_ms(word.start_offset, trace_id) + end_ms = _duration_to_ms(word.end_offset, trace_id) except AttributeError as exc: raise _stt_parse_error("STT word time offsets missing", trace_id) from exc if end_ms < start_ms: raise _stt_parse_error("STT word time offsets are invalid", trace_id) - return {"text": text, "start_ms": start_ms, "end_ms": end_ms} - - -def _segments_from_words(words: list[Any], trace_id: str) -> list[dict]: - segments: list[dict] = [] - buffer: list[Any] = [] - for word in words: - text = _word_text(word) - if not text: - continue - buffer.append(word) - if _is_sentence_end(text) or len(buffer) >= _MAX_WORDS_PER_SEGMENT: - segments.append(_build_segment(buffer, trace_id)) - buffer = [] - if buffer: - segments.append(_build_segment(buffer, trace_id)) - if words and not segments: - raise _stt_parse_error("STT word text missing", trace_id) - return segments + return TranscriptWordDTO(text=text, start_ms=start_ms, end_ms=end_ms) def _parse_batch_recognize_response(response: Any, stt_model_version: str, trace_id: str = "") -> dict: - segments: list[dict] = [] + normalized_words: list[TranscriptWordDTO] = [] for uri, file_result in response.results.items(): _raise_for_file_result_error(uri, file_result, trace_id) transcript = getattr(getattr(file_result, "inline_result", None), "transcript", None) @@ -119,8 +100,20 @@ def _parse_batch_recognize_response(response: Any, stt_model_version: str, trace words = list(getattr(alt, "words", []) or []) if text and not words: raise _stt_parse_error("STT word time offsets missing", trace_id) - segments.extend(_segments_from_words(words, trace_id)) - return {"segments": segments, "stt_model_version": stt_model_version} + normalized_words.extend(_normalize_word(word, trace_id) for word in words) + normalized_words.sort(key=lambda word: word.start_ms) + segments = segments_from_words(normalized_words) + return { + "segments": [ + {"text": segment.text, "start_ms": segment.start_ms, "end_ms": segment.end_ms} + for segment in segments + ], + "words": [ + {"text": word.text, "start_ms": word.start_ms, "end_ms": word.end_ms} + for word in normalized_words + ], + "stt_model_version": stt_model_version, + } def build_stt_callable( diff --git a/services/pipeline-worker/src/infra/db/artifact_repository.py b/services/pipeline-worker/src/infra/db/artifact_repository.py index 82e38eb..d9a0f4f 100644 --- a/services/pipeline-worker/src/infra/db/artifact_repository.py +++ b/services/pipeline-worker/src/infra/db/artifact_repository.py @@ -140,6 +140,19 @@ async def get_audio_asset(self, video_id: UUID | str) -> AssetRecord | None: assets = await self.list_assets(video_id, asset_type="AUDIO") return assets[0] if assets else None + async def delete_assets_by_type(self, video_id: UUID | str, *, asset_type: str) -> None: + normalized_video_id = self._normalize_uuid(video_id) + async with self._session_factory() as session: + await session.execute( + delete(AssetModel).where( + and_( + AssetModel.video_id == normalized_video_id, + AssetModel.asset_type == asset_type, + ) + ) + ) + await session.commit() + async def replace_transcripts( self, video_id: UUID | str, @@ -209,7 +222,7 @@ async def persist_chunks_and_vectors( embeddings: list[list[float]], set_ready: bool, vector_projections: list[VectorProjectionRecord] | None = None, - ) -> None: + ) -> bool: if len(chunks) != len(embeddings): raise ValueError("Chunk and embedding counts must match") @@ -287,11 +300,26 @@ async def persist_chunks_and_vectors( ) if set_ready: - await session.execute( - update(VideoModel).where(VideoModel.id == normalized_video_id).values(status="READY", failed_stage=None) + ready_result = await session.execute( + update(VideoModel) + .where( + and_( + VideoModel.id == normalized_video_id, + VideoModel.status != "DELETING", + ) + ) + .values( + status="READY", + failed_stage=None, + processing_claimed_at=None, + ) ) + if (ready_result.rowcount or 0) != 1: + await session.rollback() + return False await session.commit() + return True async def delete_video_artifacts(self, video_id: UUID | str) -> list[str]: paths_by_video_id = await self.list_storage_paths([video_id]) diff --git a/services/pipeline-worker/src/infra/db/video_repository.py b/services/pipeline-worker/src/infra/db/video_repository.py index 80753c7..d7adc81 100644 --- a/services/pipeline-worker/src/infra/db/video_repository.py +++ b/services/pipeline-worker/src/infra/db/video_repository.py @@ -23,6 +23,7 @@ class VideoRecord: storage_path: str | None = None status: VideoStatus = "PENDING" failed_stage: str | None = None + processing_claimed_at: datetime | None = None @dataclass(slots=True) @@ -154,32 +155,43 @@ async def load_pipeline_state( has_audio_asset=has_audio_asset, ) - async def claim_processing(self, video_id: UUID | str) -> bool: + async def claim_processing( + self, + video_id: UUID | str, + *, + keep_ready_status: bool = False, + ) -> bool: normalized_video_id = self._normalize_uuid(video_id) stale_cutoff = datetime.now(timezone.utc) - timedelta( seconds=self._stale_processing_reclaim_sec ) async with self._session_factory() as session: - result = await session.execute( - update(VideoModel) - .where( + statement = update(VideoModel).where(VideoModel.id == normalized_video_id) + if keep_ready_status: + statement = statement.where( and_( - VideoModel.id == normalized_video_id, + VideoModel.status == "READY", or_( - VideoModel.status.in_(("PENDING", "UPLOADED", "FAILED")), - and_( - VideoModel.status == "PROCESSING", - VideoModel.processing_claimed_at < stale_cutoff, - ), + VideoModel.processing_claimed_at.is_(None), + VideoModel.processing_claimed_at < stale_cutoff, ), ) - ) - .values( + ).values(processing_claimed_at=func.now()) + else: + statement = statement.where( + or_( + VideoModel.status.in_(("PENDING", "UPLOADED", "FAILED")), + and_( + VideoModel.status == "PROCESSING", + VideoModel.processing_claimed_at < stale_cutoff, + ), + ) + ).values( status="PROCESSING", failed_stage=None, processing_claimed_at=func.now(), ) - ) + result = await session.execute(statement) await session.commit() return (result.rowcount or 0) == 1 @@ -191,15 +203,51 @@ async def touch_processing(self, video_id: UUID | str) -> None: .where( and_( VideoModel.id == normalized_video_id, - VideoModel.status == "PROCESSING", + VideoModel.status.in_(("PROCESSING", "READY")), + VideoModel.processing_claimed_at.is_not(None), ) ) .values(processing_claimed_at=func.now()) ) await session.commit() + async def release_processing_claim(self, video_id: UUID | str) -> None: + normalized_video_id = self._normalize_uuid(video_id) + async with self._session_factory() as session: + await session.execute( + update(VideoModel) + .where(VideoModel.id == normalized_video_id) + .values(processing_claimed_at=None) + ) + await session.commit() + + async def has_fresh_processing_claim(self, video_ids: list[UUID | str]) -> bool: + normalized_video_ids = self._normalize_uuids(video_ids) + if not normalized_video_ids: + return False + stale_cutoff = datetime.now(timezone.utc) - timedelta( + seconds=self._stale_processing_reclaim_sec + ) + async with self._session_factory() as session: + count = await session.scalar( + select(func.count()) + .select_from(VideoModel) + .where( + and_( + VideoModel.id.in_(normalized_video_ids), + VideoModel.processing_claimed_at >= stale_cutoff, + ) + ) + ) + return bool(count) + async def set_ready(self, video_id: UUID | str) -> None: - await self.set_status(video_id, "READY", failed_stage=None) + await self.set_status( + video_id, + "READY", + failed_stage=None, + clear_processing_claim=True, + ) async def set_failed( self, @@ -207,9 +255,39 @@ async def set_failed( *, failed_stage: str, error_message: str | None = None, - ) -> None: + ) -> bool: del error_message - await self.set_status(video_id, "FAILED", failed_stage=failed_stage) + normalized_video_id = self._normalize_uuid(video_id) + async with self._session_factory() as session: + failed_result = await session.execute( + update(VideoModel) + .where( + and_( + VideoModel.id == normalized_video_id, + VideoModel.status != "DELETING", + ) + ) + .values( + status="FAILED", + failed_stage=failed_stage, + processing_claimed_at=None, + ) + ) + if (failed_result.rowcount or 0) == 1: + await session.commit() + return True + await session.execute( + update(VideoModel) + .where( + and_( + VideoModel.id == normalized_video_id, + VideoModel.status == "DELETING", + ) + ) + .values(processing_claimed_at=None) + ) + await session.commit() + return False async def set_status( self, @@ -217,13 +295,17 @@ async def set_status( status: VideoStatus, *, failed_stage: str | None = None, + clear_processing_claim: bool = False, ) -> None: normalized_video_id = self._normalize_uuid(video_id) async with self._session_factory() as session: + values = {"status": status, "failed_stage": failed_stage} + if clear_processing_claim: + values["processing_claimed_at"] = None await session.execute( update(VideoModel) .where(VideoModel.id == normalized_video_id) - .values(status=status, failed_stage=failed_stage) + .values(**values) ) await session.commit() @@ -258,6 +340,7 @@ def _to_record(model: VideoModel) -> VideoRecord: storage_path=model.storage_path, status=model.status, failed_stage=model.failed_stage, + processing_claimed_at=model.processing_claimed_at, ) @staticmethod diff --git a/services/pipeline-worker/src/infra/media/ffmpeg_client.py b/services/pipeline-worker/src/infra/media/ffmpeg_client.py index 9cba6c9..13eb572 100644 --- a/services/pipeline-worker/src/infra/media/ffmpeg_client.py +++ b/services/pipeline-worker/src/infra/media/ffmpeg_client.py @@ -9,19 +9,57 @@ @dataclass class FFmpegClient: - """Execute FFmpeg commands for audio extraction and keyframe screenshots.""" + """Execute FFmpeg commands for media inspection and extraction.""" ffmpeg_path: str = "ffmpeg" + ffprobe_path: str = "ffprobe" runner: RunnerType | None = None def __post_init__(self) -> None: if self.runner is None: self.runner = subprocess.run - def _run(self, command: list[str], timeout: float) -> None: + def _run( + self, + command: list[str], + timeout: float, + *, + capture_output: bool = False, + ) -> object: if self.runner is None: raise RuntimeError("No runner configured for FFmpegClient") - self.runner(command, check=True, timeout=timeout) + if capture_output: + return self.runner( + command, + check=True, + timeout=timeout, + capture_output=True, + text=True, + ) + return self.runner(command, check=True, timeout=timeout) + + def probe_duration_ms(self, input_file: Path | str, timeout: float = 30.0) -> int: + """Return the media duration reported by ffprobe in milliseconds.""" + + command = [ + self.ffprobe_path, + "-v", + "error", + "-show_entries", + "format=duration", + "-of", + "default=noprint_wrappers=1:nokey=1", + str(input_file), + ] + result = self._run(command, timeout, capture_output=True) + duration_text = str(getattr(result, "stdout", "")).strip() + try: + duration_sec = float(duration_text) + except ValueError as exc: + raise RuntimeError(f"ffprobe returned an invalid duration: {duration_text!r}") from exc + if duration_sec < 0: + raise RuntimeError(f"ffprobe returned a negative duration: {duration_text!r}") + return round(duration_sec * 1000) def extract_audio(self, input_file: Path | str, output_file: Path | str, timeout: float = 120.0) -> None: """Extract mono FLAC audio spec'd in the design.""" @@ -47,6 +85,49 @@ def extract_audio(self, input_file: Path | str, output_file: Path | str, timeout ] self._run(command, timeout) + def extract_audio_part( + self, + input_file: Path | str, + output_file: Path | str, + *, + start_ms: int, + end_ms: int, + timeout: float = 120.0, + ) -> None: + """Extract a mono FLAC interval using millisecond media offsets.""" + + if start_ms < 0 or end_ms <= start_ms: + raise ValueError("Audio part must satisfy 0 <= start_ms < end_ms") + duration_ms = end_ms - start_ms + command = [ + self.ffmpeg_path, + "-hide_banner", + "-y", + "-i", + str(input_file), + "-ss", + self._format_milliseconds(start_ms), + "-t", + self._format_milliseconds(duration_ms), + "-vn", + "-ac", + "1", + "-ar", + "16000", + "-c:a", + "flac", + "-sample_fmt", + "s16", + "-f", + "flac", + str(output_file), + ] + self._run(command, timeout) + + @staticmethod + def _format_milliseconds(milliseconds: int) -> str: + return f"{milliseconds / 1000:.3f}" + def extract_keyframe( self, input_file: Path | str, diff --git a/services/pipeline-worker/src/infra/media/youtube_downloader.py b/services/pipeline-worker/src/infra/media/youtube_downloader.py index 5b5756c..018f09a 100644 --- a/services/pipeline-worker/src/infra/media/youtube_downloader.py +++ b/services/pipeline-worker/src/infra/media/youtube_downloader.py @@ -103,6 +103,7 @@ def _download_sync(self, source_url: str, destination: Path) -> Path: self._extract_info(source_url, self._download_options(output_template), download=True) if not destination.exists(): raise DownloadError(f"Downloaded file was not created at {destination}") + self._validate_downloaded_file(destination) return destination def _extract_info(self, source_url: str, options: dict[str, Any], *, download: bool) -> dict[str, Any]: @@ -147,6 +148,10 @@ def _validate_metadata(self, info: dict[str, Any]) -> None: if filesize is not None and int(filesize) > self._max_filesize_bytes: raise DownloadError(f"YouTube video size exceeds {self._max_filesize_bytes} bytes.") + def _validate_downloaded_file(self, destination: Path) -> None: + if destination.stat().st_size > self._max_filesize_bytes: + raise DownloadError(f"YouTube video size exceeds {self._max_filesize_bytes} bytes.") + @staticmethod def _classify_download_error(exc: Exception) -> DownloadError: message = str(exc) diff --git a/services/pipeline-worker/src/services/long_audio_transcription.py b/services/pipeline-worker/src/services/long_audio_transcription.py new file mode 100644 index 0000000..a745624 --- /dev/null +++ b/services/pipeline-worker/src/services/long_audio_transcription.py @@ -0,0 +1,267 @@ +import asyncio +from pathlib import Path + +from loguru import logger + +from src.infra.ai.google_stt_adapter import ( + ExternalAIAdapterError, + GoogleSTTAdapter, + STTTranscriptionResult, +) +from src.infra.db.artifact_repository import ArtifactRepository, AssetRecord +from src.infra.db.video_repository import VideoRepository +from src.infra.media.ffmpeg_client import FFmpegClient +from src.infra.storage.client import StorageClient +from src.services.pipeline_errors import AudioPreparationError, DeleteRequested +from src.services.transcript_merge_service import AudioPart, TranscriptMergeService + + +STT_INPUT_AUDIO_PART = "STT_INPUT_AUDIO_PART" +CHIRP3_WORD_TIMESTAMP_LIMIT_MS = 20 * 60 * 1000 + + +class LongAudioTranscriptionService: + def __init__( + self, + *, + artifact_repository: ArtifactRepository, + video_repository: VideoRepository, + storage_client: StorageClient, + ffmpeg_client: FFmpegClient, + stt_adapter: GoogleSTTAdapter, + merge_service: TranscriptMergeService, + part_duration_sec: int, + part_overlap_sec: int, + stt_concurrency: int, + processing_timeout_sec: int, + ) -> None: + self._artifact_repository = artifact_repository + self._video_repository = video_repository + self._storage_client = storage_client + self._ffmpeg_client = ffmpeg_client + self._stt_adapter = stt_adapter + self._merge_service = merge_service + self._part_duration_ms = part_duration_sec * 1000 + self._part_overlap_ms = part_overlap_sec * 1000 + self._stt_concurrency = stt_concurrency + self._processing_timeout_sec = processing_timeout_sec + + @staticmethod + def requires_splitting(duration_ms: int) -> bool: + return duration_ms > CHIRP3_WORD_TIMESTAMP_LIMIT_MS + + def plan_parts(self, video_id: str, duration_ms: int) -> list[AudioPart]: + parts: list[AudioPart] = [] + boundary_ms = self._part_duration_ms + index = 0 + while boundary_ms - self._part_duration_ms < duration_ms: + nominal_start_ms = index * self._part_duration_ms + start_ms = max(0, nominal_start_ms - (self._part_overlap_ms if index else 0)) + end_ms = min((index + 1) * self._part_duration_ms, duration_ms) + parts.append( + AudioPart( + index=index, + start_ms=start_ms, + end_ms=end_ms, + storage_path=( + f"artifacts/{video_id}/stt-input/v1/" + f"part-{index:03d}-{start_ms}-{end_ms}.flac" + ), + ) + ) + index += 1 + boundary_ms += self._part_duration_ms + return parts + + async def transcribe( + self, + *, + video_id: str, + audio_storage_path: str, + local_audio_path: Path | None, + duration_ms: int, + workdir: Path, + trace_id: str, + ) -> STTTranscriptionResult: + parts = self.plan_parts(video_id, duration_ms) + cleanup_paths = {part.storage_path for part in parts} + primary_error: BaseException | None = None + try: + cleanup_paths.update(await self._tracked_part_paths(video_id)) + await self._cleanup_parts(video_id, cleanup_paths) + await self._assert_not_deleting(video_id) + source_path = await self._ensure_local_audio( + audio_storage_path=audio_storage_path, + local_audio_path=local_audio_path, + workdir=workdir, + ) + await self._prepare_parts(video_id, source_path, workdir, parts) + results = await self._transcribe_parts(video_id, parts, trace_id) + try: + return self._merge_service.merge( + parts=parts, + results=results, + duration_ms=duration_ms, + ) + except ValueError as exc: + raise ExternalAIAdapterError( + code="INTERNAL_ERROR", + message=f"STT audio part merge failed: {exc}", + trace_id=trace_id, + provider="google-stt", + retryable=False, + ) from exc + except BaseException as exc: + primary_error = exc + raise + finally: + try: + await self._cleanup_parts(video_id, cleanup_paths) + except Exception as cleanup_error: + logger.bind(trace_id=trace_id, video_id=video_id).error( + "STT audio part cleanup failed error={}", cleanup_error + ) + if primary_error is None: + raise AudioPreparationError("STT audio part cleanup failed") from cleanup_error + + async def _tracked_part_paths(self, video_id: str) -> set[str]: + try: + assets = await self._artifact_repository.list_assets( + video_id, + asset_type=STT_INPUT_AUDIO_PART, + ) + except Exception as exc: + raise AudioPreparationError("STT audio part lookup failed") from exc + return {asset.storage_path for asset in assets} + + async def _cleanup_parts(self, video_id: str, storage_paths: set[str]) -> None: + try: + if storage_paths: + await self._storage_client.delete_objects(sorted(storage_paths)) + await self._artifact_repository.delete_assets_by_type( + video_id, + asset_type=STT_INPUT_AUDIO_PART, + ) + except Exception as exc: + raise AudioPreparationError("STT audio part cleanup failed") from exc + + async def _ensure_local_audio( + self, + *, + audio_storage_path: str, + local_audio_path: Path | None, + workdir: Path, + ) -> Path: + if local_audio_path is not None: + return local_audio_path + downloaded_path = workdir / "audio.flac" + try: + await self._storage_client.download_object(audio_storage_path, downloaded_path) + except Exception as exc: + raise AudioPreparationError(f"Audio download failed: {audio_storage_path}") from exc + return downloaded_path + + async def _prepare_parts( + self, + video_id: str, + source_path: Path, + workdir: Path, + parts: list[AudioPart], + ) -> None: + for part in parts: + local_part_path = workdir / f"stt-part-{part.index:03d}.flac" + await self._extract_part(source_path, local_part_path, part) + await self._register_part(video_id, part) + await self._assert_not_deleting(video_id) + try: + await self._storage_client.upload_object(local_part_path, part.storage_path) + except Exception as exc: + raise AudioPreparationError(f"Audio part upload failed: {part.storage_path}") from exc + await self._assert_not_deleting(video_id) + + async def _register_part(self, video_id: str, part: AudioPart) -> None: + try: + await self._artifact_repository.upsert_asset( + video_id, + AssetRecord( + asset_type=STT_INPUT_AUDIO_PART, + storage_path=part.storage_path, + start_ms=part.start_ms, + end_ms=part.end_ms, + ), + ) + except Exception as exc: + raise AudioPreparationError( + f"Audio part registration failed: {part.storage_path}" + ) from exc + + async def _extract_part(self, source_path: Path, output_path: Path, part: AudioPart) -> None: + try: + await asyncio.to_thread( + self._ffmpeg_client.extract_audio_part, + source_path, + output_path, + start_ms=part.start_ms, + end_ms=part.end_ms, + timeout=self._processing_timeout_sec, + ) + except Exception as exc: + raise AudioPreparationError( + f"Audio part extraction failed: {part.storage_path}" + ) from exc + + async def _transcribe_parts( + self, + video_id: str, + parts: list[AudioPart], + trace_id: str, + ) -> list[STTTranscriptionResult]: + results: list[STTTranscriptionResult] = [] + for offset in range(0, len(parts), self._stt_concurrency): + current_parts = parts[offset : offset + self._stt_concurrency] + current_results = await asyncio.gather( + *(self._transcribe_part(part, trace_id) for part in current_parts), + return_exceptions=True, + ) + error = self._first_error(video_id, current_parts, current_results, trace_id) + if error is not None: + raise error + results.extend( + result for result in current_results if isinstance(result, STTTranscriptionResult) + ) + return results + + async def _transcribe_part(self, part: AudioPart, trace_id: str) -> STTTranscriptionResult: + return await self._stt_adapter.transcribe( + audio_uri=self._storage_client.object_uri(part.storage_path), + trace_id=trace_id, + ) + + @staticmethod + def _first_error( + video_id: str, + parts: list[AudioPart], + results: list[STTTranscriptionResult | BaseException], + trace_id: str, + ) -> BaseException | None: + for part, result in zip(parts, results, strict=True): + if not isinstance(result, BaseException): + continue + if isinstance(result, ExternalAIAdapterError): + logger.bind(trace_id=trace_id, video_id=video_id).error( + "STT audio part failed part_index={} start_ms={} end_ms={} attempts={} " + "code={} message={} retryable={}", + part.index, + part.start_ms, + part.end_ms, + result.attempt_count, + result.code, + result.message, + result.retryable, + ) + return result + return None + + async def _assert_not_deleting(self, video_id: str) -> None: + if await self._video_repository.is_deleting(video_id): + raise DeleteRequested(video_id) diff --git a/services/pipeline-worker/src/services/pipeline_errors.py b/services/pipeline-worker/src/services/pipeline_errors.py new file mode 100644 index 0000000..c8dd039 --- /dev/null +++ b/services/pipeline-worker/src/services/pipeline_errors.py @@ -0,0 +1,6 @@ +class AudioPreparationError(Exception): + """Raised when audio validation, splitting, or upload cannot complete.""" + + +class DeleteRequested(Exception): + """Raised when a processing worker observes a pending video deletion.""" diff --git a/services/pipeline-worker/src/services/pipeline_orchestrator.py b/services/pipeline-worker/src/services/pipeline_orchestrator.py index e1fad6a..40405c3 100644 --- a/services/pipeline-worker/src/services/pipeline_orchestrator.py +++ b/services/pipeline-worker/src/services/pipeline_orchestrator.py @@ -22,6 +22,8 @@ from src.infra.media.youtube_downloader import DownloadError, YoutubeDownloader from src.infra.storage.client import StorageClient from src.services.chunking_service import ChunkingService +from src.services.long_audio_transcription import LongAudioTranscriptionService +from src.services.pipeline_errors import AudioPreparationError, DeleteRequested from src.services.text_normalizer import normalize_enriched_text from src.utils.workdir import WorkdirManager @@ -31,6 +33,7 @@ class AudioArtifactRef: local_path: Path | None storage_path: str object_uri: str + duration_ms: int @dataclass(slots=True) @@ -41,10 +44,6 @@ class PipelineArtifacts: vector_projections: list[VectorProjectionRecord] -class DeleteRequested(Exception): - pass - - class PipelineOrchestrator: def __init__( self, @@ -59,11 +58,15 @@ def __init__( vision_adapter: VisionAdapter, workdir_manager: WorkdirManager, chunking_service: ChunkingService, + long_audio_transcription_service: LongAudioTranscriptionService | None = None, embedding_batch_size: int, stt_model_version: str, embedding_model_version: str, release_context_repository: ReleaseContextRepository | None = None, chunk_concurrency: int = 2, + max_audio_duration_sec: int = 3600, + max_source_size_bytes: int = 500 * 1024 * 1024, + audio_processing_timeout_sec: int = 120, ) -> None: self._video_repository = video_repository self._artifact_repository = artifact_repository @@ -75,11 +78,15 @@ def __init__( self._vision_adapter = vision_adapter self._workdir_manager = workdir_manager self._chunking_service = chunking_service + self._long_audio_transcription_service = long_audio_transcription_service self._embedding_batch_size = embedding_batch_size self._stt_model_version = stt_model_version self._embedding_model_version = embedding_model_version self._release_context_repository = release_context_repository self._chunk_concurrency = chunk_concurrency + self._max_audio_duration_ms = max_audio_duration_sec * 1000 + self._max_source_size_bytes = max_source_size_bytes + self._audio_processing_timeout_sec = audio_processing_timeout_sec async def run( self, @@ -110,7 +117,7 @@ def record_timing(step_name: str, started_at: float) -> None: started_at = perf_counter() segments, stt_result = await self._ensure_transcript( - video, audio_ref, state, self._stt_model_version, trace_id, + video, audio_ref, state, self._stt_model_version, trace_id, workdir, ) record_timing("stt", started_at) await self._assert_not_deleting(video.id) @@ -138,6 +145,7 @@ def record_timing(step_name: str, started_at: float) -> None: ) embeddings = vector_projections[0].embeddings record_timing("embedding", started_at) + await self._assert_not_deleting(video.id) started_at = perf_counter() await self._persist_results( @@ -163,6 +171,15 @@ def record_timing(step_name: str, started_at: float) -> None: total_duration=perf_counter() - total_started_at, ) return artifacts + except DeleteRequested: + self._log_timings( + trace_id=trace_id, + video_id=str(video.id), + status="deleted", + timings=timings, + total_duration=perf_counter() - total_started_at, + ) + raise except Exception: self._log_timings( trace_id=trace_id, @@ -198,26 +215,103 @@ async def _download_external_source(self, video: VideoRecord, workdir: Path) -> async def _ensure_audio(self, video: VideoRecord, workdir: Path, original: Path, state: PipelineState) -> AudioArtifactRef: audio_asset = await self._artifact_repository.get_audio_asset(video.id) if state.has_audio_asset and audio_asset is not None: + local_path, duration_ms = await self._load_existing_audio_metadata( + video_id=str(video.id), + workdir=workdir, + audio_asset=audio_asset, + ) return AudioArtifactRef( - local_path=None, + local_path=local_path, storage_path=audio_asset.storage_path, object_uri=self._storage_client.object_uri(audio_asset.storage_path), + duration_ms=duration_ms, ) + await self._validate_source_before_extraction(original) audio_path = workdir / "audio.flac" - await asyncio.to_thread(self._ffmpeg_client.extract_audio, original, audio_path) + try: + await asyncio.to_thread( + self._ffmpeg_client.extract_audio, + original, + audio_path, + self._audio_processing_timeout_sec, + ) + duration_ms = await asyncio.to_thread(self._ffmpeg_client.probe_duration_ms, audio_path) + self._validate_duration(duration_ms) + except AudioPreparationError: + raise + except Exception as exc: + raise AudioPreparationError(f"Audio extraction failed: {audio_path}") from exc audio_storage_path = f"artifacts/{video.id}/audio.flac" - await self._storage_client.upload_object(audio_path, audio_storage_path) + try: + await self._storage_client.upload_object(audio_path, audio_storage_path) + except Exception as exc: + raise AudioPreparationError(f"Audio upload failed: {audio_storage_path}") from exc await self._artifact_repository.upsert_asset( video.id, - AssetRecord(asset_type="AUDIO", storage_path=audio_storage_path), + AssetRecord( + asset_type="AUDIO", + storage_path=audio_storage_path, + start_ms=0, + end_ms=duration_ms, + ), ) return AudioArtifactRef( local_path=audio_path, storage_path=audio_storage_path, object_uri=self._storage_client.object_uri(audio_storage_path), + duration_ms=duration_ms, ) + async def _load_existing_audio_metadata( + self, + *, + video_id: str, + workdir: Path, + audio_asset: AssetRecord, + ) -> tuple[Path | None, int]: + if audio_asset.start_ms == 0 and audio_asset.end_ms is not None: + self._validate_duration(audio_asset.end_ms) + return None, audio_asset.end_ms + local_path = workdir / "audio.flac" + try: + await self._storage_client.download_object(audio_asset.storage_path, local_path) + duration_ms = await asyncio.to_thread(self._ffmpeg_client.probe_duration_ms, local_path) + self._validate_duration(duration_ms) + except AudioPreparationError: + raise + except Exception as exc: + raise AudioPreparationError( + f"Existing audio duration check failed: {audio_asset.storage_path}" + ) from exc + await self._artifact_repository.upsert_asset( + video_id, + AssetRecord( + asset_type="AUDIO", + storage_path=audio_asset.storage_path, + start_ms=0, + end_ms=duration_ms, + ), + ) + return local_path, duration_ms + + async def _validate_source_before_extraction(self, source_path: Path) -> None: + if source_path.stat().st_size > self._max_source_size_bytes: + raise AudioPreparationError( + f"Source size exceeds {self._max_source_size_bytes} bytes: {source_path}" + ) + try: + duration_ms = await asyncio.to_thread(self._ffmpeg_client.probe_duration_ms, source_path) + except Exception as exc: + raise AudioPreparationError(f"Source duration check failed: {source_path}") from exc + self._validate_duration(duration_ms) + + def _validate_duration(self, duration_ms: int) -> None: + if duration_ms > self._max_audio_duration_ms: + raise AudioPreparationError( + f"Audio duration exceeds {self._max_audio_duration_ms} milliseconds" + ) + async def _ensure_transcript( self, video: VideoRecord, @@ -225,6 +319,7 @@ async def _ensure_transcript( state: PipelineState, target_stt_model_version: str, trace_id: str, + workdir: Path, ) -> tuple[list[TranscriptSegmentRecord], STTTranscriptionResult]: if state.has_transcript: transcript_segments = await self._artifact_repository.load_transcripts( @@ -235,7 +330,22 @@ async def _ensure_transcript( transcript_segments = [] if not transcript_segments: - stt_result = await self._stt_adapter.transcribe(audio_uri=audio_ref.object_uri, trace_id=trace_id) + if LongAudioTranscriptionService.requires_splitting(audio_ref.duration_ms): + if self._long_audio_transcription_service is None: + raise RuntimeError("Long audio transcription service is not configured") + stt_result = await self._long_audio_transcription_service.transcribe( + video_id=str(video.id), + audio_storage_path=audio_ref.storage_path, + local_audio_path=audio_ref.local_path, + duration_ms=audio_ref.duration_ms, + workdir=workdir, + trace_id=trace_id, + ) + else: + stt_result = await self._stt_adapter.transcribe( + audio_uri=audio_ref.object_uri, + trace_id=trace_id, + ) transcript_segments = [ TranscriptSegmentRecord( segment_index=index, @@ -368,13 +478,15 @@ async def _persist_results( vector_projections: list[VectorProjectionRecord] | None = None, set_ready: bool, ) -> None: - await self._artifact_repository.persist_chunks_and_vectors( + persisted = await self._artifact_repository.persist_chunks_and_vectors( video_id, chunks=chunks, embeddings=embeddings, vector_projections=vector_projections, set_ready=set_ready, ) + if not persisted: + raise DeleteRequested(video_id) async def _load_release_targets(self) -> OnlineIngestTargets: if self._release_context_repository is None: diff --git a/services/pipeline-worker/src/services/transcript_merge_service.py b/services/pipeline-worker/src/services/transcript_merge_service.py new file mode 100644 index 0000000..2091836 --- /dev/null +++ b/services/pipeline-worker/src/services/transcript_merge_service.py @@ -0,0 +1,76 @@ +from dataclasses import dataclass + +from src.infra.ai.google_stt_adapter import ( + STTTranscriptionResult, + TranscriptWordDTO, + segments_from_words, +) + + +@dataclass(frozen=True, slots=True) +class AudioPart: + index: int + start_ms: int + end_ms: int + storage_path: str + + +class TranscriptMergeService: + def merge( + self, + *, + parts: list[AudioPart], + results: list[STTTranscriptionResult], + duration_ms: int, + ) -> STTTranscriptionResult: + if len(parts) != len(results): + raise ValueError("Audio part and STT result counts must match") + words: list[TranscriptWordDTO] = [] + for position, (part, result) in enumerate(zip(parts, results, strict=True)): + if result.words is None: + raise ValueError("Long audio STT response must include word time offsets") + words.extend(self._owned_words(parts, position, part, result.words, duration_ms)) + words.sort(key=lambda word: (word.start_ms, word.end_ms, word.text)) + model_version = results[0].stt_model_version if results else "" + return STTTranscriptionResult( + segments=segments_from_words(words), + stt_model_version=model_version, + words=words, + ) + + def _owned_words( + self, + parts: list[AudioPart], + position: int, + part: AudioPart, + relative_words: list[TranscriptWordDTO], + duration_ms: int, + ) -> list[TranscriptWordDTO]: + lower_bound = self._lower_ownership_bound(parts, position, part) + upper_bound = self._upper_ownership_bound(parts, position, part) + owned_words: list[TranscriptWordDTO] = [] + for word in relative_words: + global_word = self._to_global_word(word, part.start_ms, duration_ms) + midpoint = (global_word.start_ms + global_word.end_ms) / 2 + if midpoint < lower_bound or midpoint >= upper_bound: + continue + owned_words.append(global_word) + return owned_words + + @staticmethod + def _lower_ownership_bound(parts: list[AudioPart], position: int, part: AudioPart) -> float: + if position == 0: + return 0 + return (part.start_ms + parts[position - 1].end_ms) / 2 + + @staticmethod + def _upper_ownership_bound(parts: list[AudioPart], position: int, part: AudioPart) -> float: + if position == len(parts) - 1: + return part.end_ms + return (parts[position + 1].start_ms + part.end_ms) / 2 + + @staticmethod + def _to_global_word(word: TranscriptWordDTO, offset_ms: int, duration_ms: int) -> TranscriptWordDTO: + start_ms = min(max(word.start_ms + offset_ms, 0), duration_ms) + end_ms = min(max(word.end_ms + offset_ms, start_ms), duration_ms) + return TranscriptWordDTO(text=word.text, start_ms=start_ms, end_ms=end_ms) diff --git a/services/pipeline-worker/src/usecases/delete_video.py b/services/pipeline-worker/src/usecases/delete_video.py index d73e6a1..498812f 100644 --- a/services/pipeline-worker/src/usecases/delete_video.py +++ b/services/pipeline-worker/src/usecases/delete_video.py @@ -11,6 +11,10 @@ class DeleteVideoResult: duplicate_count: int +class DeletionDeferred(Exception): + """Keep the delete message unacknowledged while processing cleanup finishes.""" + + class DeleteVideoUseCase: def __init__( self, @@ -31,6 +35,8 @@ async def execute(self, *, video_ids: list[str], trace_id: str) -> DeleteVideoRe return DeleteVideoResult(deleted_count=0, duplicate_count=len(unique_video_ids)) found_video_ids = [video.id for video in videos] + if await self._video_repository.has_fresh_processing_claim(found_video_ids): + raise DeletionDeferred("Video processing cleanup is still active") artifact_paths = await self._artifact_repository.list_storage_paths(found_video_ids) storage_paths = self._storage_paths_for(videos, artifact_paths) diff --git a/services/pipeline-worker/src/usecases/process_video.py b/services/pipeline-worker/src/usecases/process_video.py index 98bdcb5..5e69932 100644 --- a/services/pipeline-worker/src/usecases/process_video.py +++ b/services/pipeline-worker/src/usecases/process_video.py @@ -1,9 +1,12 @@ from dataclasses import dataclass +from loguru import logger + from src.infra.ai.google_stt_adapter import ExternalAIAdapterError from src.infra.db.video_repository import VideoRepository from src.infra.media.youtube_downloader import DownloadError -from src.services.pipeline_orchestrator import DeleteRequested, PipelineOrchestrator +from src.services.pipeline_errors import AudioPreparationError, DeleteRequested +from src.services.pipeline_orchestrator import PipelineOrchestrator from src.usecases.delete_video import DeleteVideoUseCase @@ -59,14 +62,17 @@ async def execute( # 처리권한 확보 keep_ready_status = video.status == "READY" - if not keep_ready_status: - claimed = await self._video_repository.claim_processing(video_id) # 처리권한 확보 시도(processig 상태로 변경) - if not claimed: - refreshed = await self._video_repository.get_video(video_id) - if refreshed is not None and refreshed.status == "DELETING": # 삭제 - await self._delete_video_use_case.execute(video_ids=[video_id], trace_id=trace_id) - return ProcessVideoResult(action="deleted") - return ProcessVideoResult(action="skip") + # READY 재처리는 상태를 유지한 채 처리 권한만 확보 + claimed = await self._video_repository.claim_processing( + video_id, + keep_ready_status=keep_ready_status, + ) + if not claimed: + refreshed = await self._video_repository.get_video(video_id) + if refreshed is not None and refreshed.status == "DELETING": # 삭제 + await self._delete_video_use_case.execute(video_ids=[video_id], trace_id=trace_id) + return ProcessVideoResult(action="deleted") + return ProcessVideoResult(action="skip") # 오케스트레이터 호출 try: @@ -80,21 +86,67 @@ async def execute( # 예외 발생시 failed stage 분류 except DeleteRequested: + await self._video_repository.release_processing_claim(video_id) await self._delete_video_use_case.execute(video_ids=[video_id], trace_id=trace_id) return ProcessVideoResult(action="deleted") except DownloadError as exc: - await self._video_repository.set_failed(video_id, failed_stage="DOWNLOAD", error_message=str(exc)) - return ProcessVideoResult(action="failed", failed_stage="DOWNLOAD") + return await self._fail_or_delete( + video_id=video_id, + trace_id=trace_id, + failed_stage="DOWNLOAD", + error_message=str(exc), + ) except FileNotFoundError as exc: - await self._video_repository.set_failed(video_id, failed_stage="DOWNLOAD", error_message=str(exc)) - return ProcessVideoResult(action="failed", failed_stage="DOWNLOAD") + return await self._fail_or_delete( + video_id=video_id, + trace_id=trace_id, + failed_stage="DOWNLOAD", + error_message=str(exc), + ) + except AudioPreparationError as exc: + logger.bind(trace_id=trace_id, video_id=video_id).error( + "Audio preparation failed failed_stage=EXTRACT error={}", + exc, + ) + return await self._fail_or_delete( + video_id=video_id, + trace_id=trace_id, + failed_stage="EXTRACT", + error_message=str(exc), + ) except ExternalAIAdapterError as exc: if exc.provider == "google-stt": failed_stage = "STT" else: failed_stage = FAILED_STAGE_BY_CODE.get(exc.code, "VECTOR_UPSERT") - await self._video_repository.set_failed(video_id, failed_stage=failed_stage, error_message=exc.message) - return ProcessVideoResult(action="failed", failed_stage=failed_stage) + return await self._fail_or_delete( + video_id=video_id, + trace_id=trace_id, + failed_stage=failed_stage, + error_message=exc.message, + ) except Exception as exc: - await self._video_repository.set_failed(video_id, failed_stage="VECTOR_UPSERT", error_message=str(exc)) - return ProcessVideoResult(action="failed", failed_stage="VECTOR_UPSERT") + return await self._fail_or_delete( + video_id=video_id, + trace_id=trace_id, + failed_stage="VECTOR_UPSERT", + error_message=str(exc), + ) + + async def _fail_or_delete( + self, + *, + video_id: str, + trace_id: str, + failed_stage: str, + error_message: str, + ) -> ProcessVideoResult: + marked_failed = await self._video_repository.set_failed( + video_id, + failed_stage=failed_stage, + error_message=error_message, + ) + if marked_failed: + return ProcessVideoResult(action="failed", failed_stage=failed_stage) + await self._delete_video_use_case.execute(video_ids=[video_id], trace_id=trace_id) + return ProcessVideoResult(action="deleted") diff --git a/services/pipeline-worker/tests/integration/test_delete_project.py b/services/pipeline-worker/tests/integration/test_delete_project.py index 5be7fd7..baceec2 100644 --- a/services/pipeline-worker/tests/integration/test_delete_project.py +++ b/services/pipeline-worker/tests/integration/test_delete_project.py @@ -2,7 +2,7 @@ from uuid import uuid4 import pytest -from sqlalchemy import func, select +from sqlalchemy import func, select, update from src.infra.db.artifact_repository import AssetRecord, ChunkRecord from src.infra.db.models import ( @@ -14,9 +14,11 @@ SearchResponseSnapshotModel, TranscriptSegmentModel, VectorIndexEntryModel, + VideoModel, ) from src.infra.db.video_repository import VideoRecord from src.usecases.delete_project import DeleteProjectUseCase +from src.usecases.delete_video import DeletionDeferred @pytest.mark.asyncio @@ -156,3 +158,54 @@ async def test_delete_project_cascades_videos_search_records_and_storage( "videos/second.mp4", "artifacts/first.flac", } + + +@pytest.mark.asyncio +async def test_delete_project_defers_fresh_processing_claim_but_allows_stale_claim( + session_factory, + video_repository, + delete_video_use_case, + storage_client, +) -> None: + user_id = uuid4() + project_id = uuid4() + video_id = uuid4() + storage_client.objects["videos/processing.mp4"] = b"video" + async with session_factory() as session: + session.add(ProjectModel(id=project_id, user_id=user_id, title="Project")) + await session.commit() + await video_repository.create_video( + VideoRecord( + id=video_id, + user_id=user_id, + project_id=project_id, + storage_path="videos/processing.mp4", + status="UPLOADED", + ) + ) + assert await video_repository.claim_processing(video_id) is True + use_case = DeleteProjectUseCase( + video_repository=video_repository, + delete_video_use_case=delete_video_use_case, + session_factory=session_factory, + ) + + with pytest.raises(DeletionDeferred): + await use_case.execute(project_id=str(project_id), trace_id="trace-project-defer") + + assert await video_repository.get_video(video_id) is not None + async with session_factory() as session: + await session.execute( + update(VideoModel) + .where(VideoModel.id == video_id) + .values(processing_claimed_at=datetime.now(UTC) - timedelta(seconds=1501)) + ) + await session.commit() + + result = await use_case.execute( + project_id=str(project_id), + trace_id="trace-project-stale-delete", + ) + + assert result.deleted_video_count == 1 + assert await video_repository.get_video(video_id) is None diff --git a/services/pipeline-worker/tests/integration/test_long_audio_transcription.py b/services/pipeline-worker/tests/integration/test_long_audio_transcription.py new file mode 100644 index 0000000..179e94d --- /dev/null +++ b/services/pipeline-worker/tests/integration/test_long_audio_transcription.py @@ -0,0 +1,343 @@ +import asyncio +from pathlib import Path +from uuid import uuid4 + +import pytest + +from src.infra.ai.google_stt_adapter import ( + ExternalAIAdapterError, + GoogleSTTAdapter, +) +from src.infra.ai.vision_adapter import MockVisionAdapter +from src.infra.db.video_repository import VideoRecord +from src.infra.media.youtube_downloader import InMemoryYoutubeDownloader +from src.services.chunking_service import ChunkingService +from src.services.long_audio_transcription import ( + STT_INPUT_AUDIO_PART, + LongAudioTranscriptionService, +) +from src.services.transcript_merge_service import TranscriptMergeService +from src.services.pipeline_orchestrator import PipelineOrchestrator +from src.usecases.delete_video import DeleteVideoUseCase +from src.usecases.process_video import ProcessVideoUseCase +from src.utils.workdir import WorkdirManager +from tests.support import build_embedding_client, build_ffmpeg_adapter + + +class PartWritingFFmpeg: + def extract_audio_part( + self, + input_file: Path, + output_file: Path, + *, + start_ms: int, + end_ms: int, + timeout: float, + ) -> None: + del input_file, timeout + output_file.write_text(f"{start_ms}:{end_ms}") + + +def _build_service( + *, + artifact_repository, + video_repository, + storage_client, + stt_adapter: GoogleSTTAdapter, +) -> LongAudioTranscriptionService: + return LongAudioTranscriptionService( + artifact_repository=artifact_repository, + video_repository=video_repository, + storage_client=storage_client, + ffmpeg_client=PartWritingFFmpeg(), # type: ignore[arg-type] + stt_adapter=stt_adapter, + merge_service=TranscriptMergeService(), + part_duration_sec=900, + part_overlap_sec=5, + stt_concurrency=2, + processing_timeout_sec=120, + ) + + +def test_plan_parts_covers_one_hour_with_five_second_overlaps() -> None: + async def unused_client(audio_uri: str, trace_id: str) -> dict: + del audio_uri, trace_id + return {"segments": [], "stt_model_version": "chirp_3"} + + service = _build_service( + artifact_repository=None, + video_repository=None, + storage_client=None, + stt_adapter=GoogleSTTAdapter(client=unused_client, max_retries=0), + ) + + parts = service.plan_parts("video-1", 3_600_000) + + assert [(part.start_ms, part.end_ms) for part in parts] == [ + (0, 900_000), + (895_000, 1_800_000), + (1_795_000, 2_700_000), + (2_695_000, 3_600_000), + ] + assert len({part.storage_path for part in parts}) == 4 + + +@pytest.mark.asyncio +async def test_long_audio_success_merges_and_removes_temporary_assets( + artifact_repository, + video_repository, + storage_client, + tmp_path, +) -> None: + video_id = str(uuid4()) + await video_repository.create_video( + VideoRecord(id=video_id, user_id=str(uuid4()), status="PROCESSING") + ) + source_path = tmp_path / "audio.flac" + source_path.write_bytes(b"audio") + + async def stt_client(audio_uri: str, trace_id: str) -> dict: + del trace_id + is_second = "part-001" in audio_uri + return { + "segments": [], + "words": [ + { + "text": "second." if is_second else "first.", + "start_ms": 10_000, + "end_ms": 10_500, + } + ], + "stt_model_version": "chirp_3", + } + + service = _build_service( + artifact_repository=artifact_repository, + video_repository=video_repository, + storage_client=storage_client, + stt_adapter=GoogleSTTAdapter(client=stt_client, max_retries=0), + ) + + result = await service.transcribe( + video_id=video_id, + audio_storage_path=f"artifacts/{video_id}/audio.flac", + local_audio_path=source_path, + duration_ms=1_200_001, + workdir=tmp_path, + trace_id="trace-long-success", + ) + + assert [segment.text for segment in result.segments] == ["first.", "second."] + assert result.segments[1].start_ms == 905_000 + assert await artifact_repository.list_assets( + video_id, + asset_type=STT_INPUT_AUDIO_PART, + ) == [] + assert not any("/stt-input/" in path for path in storage_client.objects) + + +@pytest.mark.asyncio +async def test_long_audio_accepts_a_silent_part( + artifact_repository, + video_repository, + storage_client, + tmp_path, +) -> None: + video_id = str(uuid4()) + await video_repository.create_video( + VideoRecord(id=video_id, user_id=str(uuid4()), status="PROCESSING") + ) + source_path = tmp_path / "audio.flac" + source_path.write_bytes(b"audio") + + async def stt_client(audio_uri: str, trace_id: str) -> dict: + del trace_id + words = [] + if "part-001" in audio_uri: + words = [{"text": "heard.", "start_ms": 10_000, "end_ms": 10_500}] + return { + "segments": [], + "words": words, + "stt_model_version": "chirp_3", + } + + service = _build_service( + artifact_repository=artifact_repository, + video_repository=video_repository, + storage_client=storage_client, + stt_adapter=GoogleSTTAdapter(client=stt_client, max_retries=0), + ) + + result = await service.transcribe( + video_id=video_id, + audio_storage_path=f"artifacts/{video_id}/audio.flac", + local_audio_path=source_path, + duration_ms=1_200_001, + workdir=tmp_path, + trace_id="trace-long-silence", + ) + + assert [segment.text for segment in result.segments] == ["heard."] + assert result.segments[0].start_ms == 905_000 + + +@pytest.mark.asyncio +async def test_failed_parallel_part_waits_for_sibling_before_cleanup( + artifact_repository, + video_repository, + storage_client, + tmp_path, +) -> None: + video_id = str(uuid4()) + await video_repository.create_video( + VideoRecord(id=video_id, user_id=str(uuid4()), status="PROCESSING") + ) + source_path = tmp_path / "audio.flac" + source_path.write_bytes(b"audio") + events: list[str] = [] + original_delete_objects = storage_client.delete_objects + upload_seen = False + + async def recording_upload(source: Path, storage_path: str) -> None: + nonlocal upload_seen + upload_seen = True + await type(storage_client).upload_object(storage_client, source, storage_path) + + async def recording_delete(storage_paths: list[str]) -> None: + if upload_seen: + events.append("cleanup") + await original_delete_objects(storage_paths) + + storage_client.upload_object = recording_upload + storage_client.delete_objects = recording_delete + + async def stt_client(audio_uri: str, trace_id: str) -> dict: + del trace_id + if "part-000" in audio_uri: + raise ExternalAIAdapterError( + code="INVALID_REQUEST", + message="bad part", + trace_id="trace-long-failure", + provider="google-stt", + retryable=False, + ) + await asyncio.sleep(0.01) + events.append("sibling_done") + return { + "segments": [], + "words": [{"text": "done.", "start_ms": 0, "end_ms": 100}], + "stt_model_version": "chirp_3", + } + + service = _build_service( + artifact_repository=artifact_repository, + video_repository=video_repository, + storage_client=storage_client, + stt_adapter=GoogleSTTAdapter(client=stt_client, max_retries=0), + ) + + with pytest.raises(ExternalAIAdapterError, match="bad part"): + await service.transcribe( + video_id=video_id, + audio_storage_path=f"artifacts/{video_id}/audio.flac", + local_audio_path=source_path, + duration_ms=1_200_001, + workdir=tmp_path, + trace_id="trace-long-failure", + ) + + assert events.index("sibling_done") < events.index("cleanup") + assert await artifact_repository.list_assets( + video_id, + asset_type=STT_INPUT_AUDIO_PART, + ) == [] + + +@pytest.mark.asyncio +async def test_long_audio_process_flow_reaches_ready_with_global_timestamps( + artifact_repository, + video_repository, + storage_client, + tmp_path, +) -> None: + video_id = str(uuid4()) + storage_client.objects["videos/long.mp4"] = b"long-video" + await video_repository.create_video( + VideoRecord( + id=video_id, + user_id=str(uuid4()), + storage_path="videos/long.mp4", + status="UPLOADED", + ) + ) + ffmpeg_client, _ = build_ffmpeg_adapter(duration_sec=1_201.0) + + async def stt_client(audio_uri: str, trace_id: str) -> dict: + del trace_id + is_second = "part-001" in audio_uri + return { + "segments": [], + "words": [ + { + "text": "second." if is_second else "first.", + "start_ms": 10_000, + "end_ms": 10_500, + } + ], + "stt_model_version": "chirp_3", + } + + stt_adapter = GoogleSTTAdapter(client=stt_client, max_retries=0) + long_audio_service = LongAudioTranscriptionService( + artifact_repository=artifact_repository, + video_repository=video_repository, + storage_client=storage_client, + ffmpeg_client=ffmpeg_client, + stt_adapter=stt_adapter, + merge_service=TranscriptMergeService(), + part_duration_sec=900, + part_overlap_sec=5, + stt_concurrency=2, + processing_timeout_sec=120, + ) + orchestrator = PipelineOrchestrator( + video_repository=video_repository, + artifact_repository=artifact_repository, + storage_client=storage_client, + youtube_downloader=InMemoryYoutubeDownloader(), + ffmpeg_client=ffmpeg_client, + stt_adapter=stt_adapter, + embedding_client=build_embedding_client(), + vision_adapter=MockVisionAdapter(caption="caption"), + workdir_manager=WorkdirManager(base_dir=tmp_path), + chunking_service=ChunkingService(max_tokens=6, overlap_sentences=1), + long_audio_transcription_service=long_audio_service, + embedding_batch_size=2, + stt_model_version="chirp_3", + embedding_model_version="v001", + ) + use_case = ProcessVideoUseCase( + video_repository=video_repository, + orchestrator=orchestrator, + delete_video_use_case=DeleteVideoUseCase( + video_repository=video_repository, + artifact_repository=artifact_repository, + storage_client=storage_client, + ), + stt_model_version="chirp_3", + embedding_model_version="v001", + ) + + result = await use_case.execute(video_id=video_id, trace_id="trace-long-flow") + + transcripts = await artifact_repository.load_transcripts( + video_id, + stt_model_version="chirp_3", + ) + assert result.action == "processed" + assert (await video_repository.get_video(video_id)).status == "READY" + assert [segment.start_ms for segment in transcripts] == [10_000, 905_000] + assert await artifact_repository.list_assets( + video_id, + asset_type=STT_INPUT_AUDIO_PART, + ) == [] diff --git a/services/pipeline-worker/tests/integration/test_repositories.py b/services/pipeline-worker/tests/integration/test_repositories.py index 65f61af..0924db0 100644 --- a/services/pipeline-worker/tests/integration/test_repositories.py +++ b/services/pipeline-worker/tests/integration/test_repositories.py @@ -47,6 +47,7 @@ async def test_repositories_support_claim_outputs_and_delete(video_repository, a embedding_model_version="v001", ) assert state.has_current_outputs is True + assert state.video.processing_claimed_at is None paths = await artifact_repository.delete_video_artifacts(video_id) assert "artifacts/audio.flac" in paths @@ -99,6 +100,42 @@ async def test_persist_chunks_and_vectors_stores_vector_entries_with_video_owner assert vector_entry.embedding_vector == pytest.approx([1.0, 2.0]) +@pytest.mark.asyncio +async def test_ready_persist_rolls_back_when_delete_wins_race( + video_repository, + artifact_repository, +) -> None: + video_id = str(uuid4()) + await video_repository.create_video( + VideoRecord(id=video_id, user_id=str(uuid4()), status="UPLOADED") + ) + assert await video_repository.claim_processing(video_id) is True + await video_repository.set_status(video_id, "DELETING") + + persisted = await artifact_repository.persist_chunks_and_vectors( + video_id, + chunks=[ + ChunkRecord( + chunk_index=0, + text="must-not-persist", + enriched_text="must-not-persist", + start_ms=0, + end_ms=1, + chunking_version="v1", + stt_model_version="chirp_3", + embedding_model_version="v001", + ) + ], + embeddings=[[1.0, 2.0]], + set_ready=True, + ) + + video = await video_repository.get_video(video_id) + assert persisted is False + assert video.status == "DELETING" + assert await artifact_repository.list_chunks(video_id) == [] + + class TestProcessingClaimRecovery: @pytest.mark.parametrize("status", ["PENDING", "UPLOADED", "FAILED"]) @pytest.mark.asyncio @@ -145,6 +182,66 @@ async def test_fresh_processing_claim_is_rejected( assert processing_claimed_at is not None assert await video_repository.claim_processing(video_id) is False + @pytest.mark.asyncio + async def test_ready_reprocessing_claim_keeps_status_and_blocks_duplicate( + self, + video_repository, + ) -> None: + video_id = str(uuid4()) + await video_repository.create_video( + VideoRecord(id=video_id, user_id=str(uuid4()), status="READY") + ) + + assert await video_repository.claim_processing( + video_id, + keep_ready_status=True, + ) is True + claimed_video = await video_repository.get_video(video_id) + + assert claimed_video.status == "READY" + assert claimed_video.processing_claimed_at is not None + assert await video_repository.claim_processing( + video_id, + keep_ready_status=True, + ) is False + + @pytest.mark.asyncio + async def test_failed_processing_clears_claim( + self, + video_repository, + ) -> None: + video_id = str(uuid4()) + await video_repository.create_video( + VideoRecord(id=video_id, user_id=str(uuid4()), status="UPLOADED") + ) + assert await video_repository.claim_processing(video_id) is True + + marked_failed = await video_repository.set_failed(video_id, failed_stage="STT") + + failed_video = await video_repository.get_video(video_id) + assert marked_failed is True + assert failed_video.status == "FAILED" + assert failed_video.processing_claimed_at is None + + @pytest.mark.asyncio + async def test_failed_processing_does_not_overwrite_deleting_status( + self, + video_repository, + ) -> None: + video_id = str(uuid4()) + await video_repository.create_video( + VideoRecord(id=video_id, user_id=str(uuid4()), status="UPLOADED") + ) + assert await video_repository.claim_processing(video_id) is True + await video_repository.set_status(video_id, "DELETING") + + marked_failed = await video_repository.set_failed(video_id, failed_stage="STT") + + deleting_video = await video_repository.get_video(video_id) + assert marked_failed is False + assert deleting_video.status == "DELETING" + assert deleting_video.processing_claimed_at is None + @pytest.mark.asyncio async def test_stale_processing_claim_is_reclaimed( self, diff --git a/services/pipeline-worker/tests/support.py b/services/pipeline-worker/tests/support.py index 1b7966d..8026298 100644 --- a/services/pipeline-worker/tests/support.py +++ b/services/pipeline-worker/tests/support.py @@ -2,6 +2,7 @@ from collections.abc import Callable import json from pathlib import Path +from types import SimpleNamespace import httpx from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine @@ -32,16 +33,29 @@ def make_session_factory(engine: AsyncEngine) -> async_sessionmaker[AsyncSession class RecordingFFmpegRunner: - def __init__(self) -> None: + def __init__(self, *, duration_sec: float = 2.0) -> None: self.commands: list[list[str]] = [] - - def __call__(self, command: list[str], *, check: bool, timeout: float) -> None: + self.duration_sec = duration_sec + + def __call__( + self, + command: list[str], + *, + check: bool, + timeout: float, + capture_output: bool = False, + text: bool = False, + ) -> object: + del check, timeout, text self.commands.append(command) + if capture_output: + return SimpleNamespace(stdout=f"{self.duration_sec}\n") Path(command[-1]).write_bytes(b"generated-artifact") + return SimpleNamespace() -def build_ffmpeg_adapter() -> tuple[FFmpegClient, RecordingFFmpegRunner]: - runner = RecordingFFmpegRunner() +def build_ffmpeg_adapter(*, duration_sec: float = 2.0) -> tuple[FFmpegClient, RecordingFFmpegRunner]: + runner = RecordingFFmpegRunner(duration_sec=duration_sec) return FFmpegClient(runner=runner), runner diff --git a/services/pipeline-worker/tests/unit/test_delete_video.py b/services/pipeline-worker/tests/unit/test_delete_video.py index 70aa7ea..33402e9 100644 --- a/services/pipeline-worker/tests/unit/test_delete_video.py +++ b/services/pipeline-worker/tests/unit/test_delete_video.py @@ -4,6 +4,7 @@ from src.infra.db.artifact_repository import AssetRecord from src.infra.db.video_repository import VideoRecord +from src.usecases.delete_video import DeletionDeferred @pytest.mark.asyncio @@ -96,3 +97,28 @@ async def test_delete_video_retries_after_partial_storage_failure( assert retry_result.deleted_count == 1 assert retry_result.duplicate_count == 0 assert "videos/retry.mp4" not in storage_client.objects + + +@pytest.mark.asyncio +async def test_delete_video_defers_while_processing_claim_is_fresh( + video_repository, + delete_video_use_case, +) -> None: + video_id = str(uuid4()) + await video_repository.create_video( + VideoRecord(id=video_id, user_id=str(uuid4()), status="UPLOADED") + ) + assert await video_repository.claim_processing(video_id) is True + await video_repository.set_status(video_id, "DELETING") + + with pytest.raises(DeletionDeferred): + await delete_video_use_case.execute(video_ids=[video_id], trace_id="trace-defer") + + assert await video_repository.get_video(video_id) is not None + await video_repository.set_failed(video_id, failed_stage="STT") + await video_repository.set_status(video_id, "DELETING") + result = await delete_video_use_case.execute( + video_ids=[video_id], + trace_id="trace-delete-after-cleanup", + ) + assert result.deleted_count == 1 diff --git a/services/pipeline-worker/tests/unit/test_ffmpeg_adapter.py b/services/pipeline-worker/tests/unit/test_ffmpeg_adapter.py index aa31e92..bdc9b5b 100644 --- a/services/pipeline-worker/tests/unit/test_ffmpeg_adapter.py +++ b/services/pipeline-worker/tests/unit/test_ffmpeg_adapter.py @@ -1,3 +1,5 @@ +from types import SimpleNamespace + import pytest from src.infra.media.ffmpeg_client import FFmpegClient @@ -6,9 +8,11 @@ class CapturingRunner: def __init__(self): self.calls: list[dict[str, object]] = [] + self.stdout = "12.345\n" - def __call__(self, cmd: list[str], *, check: bool, timeout: float): - self.calls.append({"cmd": cmd, "check": check, "timeout": timeout}) + def __call__(self, cmd: list[str], **kwargs): + self.calls.append({"cmd": cmd, **kwargs}) + return SimpleNamespace(stdout=self.stdout) def test_extract_audio_runs_flac_command(tmp_path): @@ -46,3 +50,58 @@ def test_extract_keyframe_uses_select_filter(tmp_path): assert "select='eq(pict_type,I)'" in recorded["cmd"] assert "-frames:v" in recorded["cmd"] assert recorded["timeout"] == pytest.approx(10.0) + + +def test_probe_duration_returns_milliseconds(tmp_path): + runner = CapturingRunner() + adapter = FFmpegClient(ffprobe_path="custom-ffprobe", runner=runner) + input_file = tmp_path / "input.mp4" + + duration_ms = adapter.probe_duration_ms(input_file, timeout=7.0) + + assert duration_ms == 12345 + recorded = runner.calls[-1] + assert recorded["cmd"][0] == "custom-ffprobe" + assert "format=duration" in recorded["cmd"] + assert str(input_file) == recorded["cmd"][-1] + assert recorded["capture_output"] is True + assert recorded["text"] is True + assert recorded["timeout"] == pytest.approx(7.0) + + +def test_extract_audio_part_uses_requested_interval(tmp_path): + runner = CapturingRunner() + adapter = FFmpegClient(ffmpeg_path="ffmpeg", runner=runner) + input_file = tmp_path / "input.flac" + output_file = tmp_path / "part.flac" + + adapter.extract_audio_part( + input_file, + output_file, + start_ms=895000, + end_ms=1800000, + timeout=180.0, + ) + + recorded = runner.calls[-1] + command = recorded["cmd"] + assert command[command.index("-ss") + 1] == "895.000" + assert command[command.index("-t") + 1] == "905.000" + assert str(output_file) == command[-1] + assert recorded["timeout"] == pytest.approx(180.0) + + +@pytest.mark.parametrize( + ("start_ms", "end_ms"), + [(-1, 1000), (1000, 1000), (2000, 1000)], +) +def test_extract_audio_part_rejects_invalid_interval(start_ms: int, end_ms: int) -> None: + adapter = FFmpegClient(runner=CapturingRunner()) + + with pytest.raises(ValueError, match="0 <= start_ms < end_ms"): + adapter.extract_audio_part( + "input.flac", + "part.flac", + start_ms=start_ms, + end_ms=end_ms, + ) diff --git a/services/pipeline-worker/tests/unit/test_google_stt_adapter.py b/services/pipeline-worker/tests/unit/test_google_stt_adapter.py index 56e6cc0..977f9ea 100644 --- a/services/pipeline-worker/tests/unit/test_google_stt_adapter.py +++ b/services/pipeline-worker/tests/unit/test_google_stt_adapter.py @@ -47,3 +47,100 @@ async def slow_client(audio_uri: str, trace_id: str) -> dict: result = await adapter.transcribe(audio_uri="gs://bucket/audio.flac", trace_id="trace-4") assert result.segments[0].text == "slow transcript" + + +@pytest.mark.asyncio +async def test_google_stt_adapter_uses_exponential_backoff_with_jitter() -> None: + attempts = 0 + delays: list[float] = [] + + async def retrying_client(audio_uri: str, trace_id: str) -> dict: + nonlocal attempts + del audio_uri, trace_id + attempts += 1 + if attempts <= 3: + raise ExternalAIAdapterError( + code="UNAVAILABLE", + message="temporary", + trace_id="trace-backoff", + provider="google-stt", + retryable=True, + ) + return { + "segments": [{"text": "done", "start_ms": 0, "end_ms": 1}], + "stt_model_version": "chirp_3", + } + + async def record_sleep(delay: float) -> None: + delays.append(delay) + + adapter = GoogleSTTAdapter( + client=retrying_client, + max_retries=3, + sleep=record_sleep, + jitter=lambda: 1.0, + ) + + await adapter.transcribe(audio_uri="gs://bucket/audio.flac", trace_id="trace-backoff") + + assert attempts == 4 + assert delays == pytest.approx([1.25, 2.5, 5.0]) + + +@pytest.mark.asyncio +async def test_google_stt_adapter_does_not_retry_non_retryable_error() -> None: + attempts = 0 + delays: list[float] = [] + + async def invalid_client(audio_uri: str, trace_id: str) -> dict: + nonlocal attempts + del audio_uri, trace_id + attempts += 1 + raise ExternalAIAdapterError( + code="INVALID_REQUEST", + message="invalid", + trace_id="trace-invalid", + provider="google-stt", + retryable=False, + ) + + async def record_sleep(delay: float) -> None: + delays.append(delay) + + adapter = GoogleSTTAdapter( + client=invalid_client, + max_retries=3, + sleep=record_sleep, + jitter=lambda: 0.0, + ) + + with pytest.raises(ExternalAIAdapterError, match="invalid"): + await adapter.transcribe( + audio_uri="gs://bucket/audio.flac", + trace_id="trace-invalid", + ) + + assert attempts == 1 + assert delays == [] + + +@pytest.mark.asyncio +async def test_google_stt_adapter_preserves_explicit_empty_words() -> None: + async def silent_client(audio_uri: str, trace_id: str) -> dict: + del audio_uri, trace_id + return { + "segments": [], + "words": [], + "stt_model_version": "chirp_3", + } + + result = await GoogleSTTAdapter( + client=silent_client, + max_retries=0, + ).transcribe( + audio_uri="gs://bucket/silence.flac", + trace_id="trace-silence", + ) + + assert result.words == [] + assert result.segments == [] diff --git a/services/pipeline-worker/tests/unit/test_pipeline_orchestrator.py b/services/pipeline-worker/tests/unit/test_pipeline_orchestrator.py index 1902f5c..364a3df 100644 --- a/services/pipeline-worker/tests/unit/test_pipeline_orchestrator.py +++ b/services/pipeline-worker/tests/unit/test_pipeline_orchestrator.py @@ -2,6 +2,7 @@ import pytest +from src.services.pipeline_errors import DeleteRequested from src.services.pipeline_orchestrator import PipelineOrchestrator @@ -39,3 +40,18 @@ async def test_stage_boundary_touches_processing_claim_before_delete_check() -> video_repository.touch_processing.assert_awaited_once_with("video-id") video_repository.is_deleting.assert_awaited_once_with("video-id") + + +@pytest.mark.asyncio +async def test_persist_result_conflict_becomes_delete_request() -> None: + orchestrator = _build_orchestrator(AsyncMock()) + orchestrator._artifact_repository = AsyncMock() + orchestrator._artifact_repository.persist_chunks_and_vectors.return_value = False + + with pytest.raises(DeleteRequested): + await orchestrator._persist_results( + "video-id", + [], + [], + set_ready=True, + ) diff --git a/services/pipeline-worker/tests/unit/test_process_video.py b/services/pipeline-worker/tests/unit/test_process_video.py index df9d75c..ea443e4 100644 --- a/services/pipeline-worker/tests/unit/test_process_video.py +++ b/services/pipeline-worker/tests/unit/test_process_video.py @@ -1,12 +1,15 @@ from pathlib import Path +from unittest.mock import AsyncMock from uuid import uuid4 import pytest +from loguru import logger from src.infra.ai.vision_adapter import MockVisionAdapter from src.infra.db.video_repository import VideoRecord from src.infra.media.youtube_downloader import DownloadError, InMemoryYoutubeDownloader from src.services.chunking_service import ChunkingService +from src.services.pipeline_errors import AudioPreparationError from src.services.pipeline_orchestrator import PipelineOrchestrator from src.usecases.delete_video import DeleteVideoUseCase from src.usecases.process_video import ProcessVideoUseCase @@ -73,6 +76,34 @@ async def test_process_video_skips_ready_same_version( assert result.action == "skip" +@pytest.mark.asyncio +async def test_process_video_claims_ready_video_when_outputs_need_rebuild( + video_repository, + process_video_use_case, + storage_client, +) -> None: + video_id = str(uuid4()) + storage_client.objects["videos/ready-source.mp4"] = b"video" + await video_repository.create_video( + VideoRecord( + id=video_id, + user_id=str(uuid4()), + storage_path="videos/ready-source.mp4", + status="READY", + ) + ) + + result = await process_video_use_case.execute( + video_id=video_id, + trace_id="trace-ready-rebuild", + ) + + rebuilt_video = await video_repository.get_video(video_id) + assert result.action == "processed" + assert rebuilt_video.status == "READY" + assert rebuilt_video.processing_claimed_at is None + + @pytest.mark.asyncio async def test_process_video_fails_on_missing_storage_object( video_repository, @@ -253,3 +284,110 @@ async def test_process_video_maps_download_error_to_download_stage( assert result.failed_stage == "DOWNLOAD" assert video.status == "FAILED" assert video.failed_stage == "DOWNLOAD" + + +@pytest.mark.asyncio +async def test_process_video_maps_source_limit_failure_to_extract_stage( + video_repository, + artifact_repository, + storage_client, +) -> None: + video_id = str(uuid4()) + storage_client.objects["videos/source.mp4"] = b"too-large" + await video_repository.create_video( + VideoRecord( + id=video_id, + user_id=str(uuid4()), + storage_path="videos/source.mp4", + status="UPLOADED", + ) + ) + ffmpeg_client, _ = build_ffmpeg_adapter() + orchestrator = PipelineOrchestrator( + video_repository=video_repository, + artifact_repository=artifact_repository, + storage_client=storage_client, + youtube_downloader=InMemoryYoutubeDownloader(), + ffmpeg_client=ffmpeg_client, + stt_adapter=build_stt_adapter(), + embedding_client=build_embedding_client(), + vision_adapter=MockVisionAdapter(caption="caption"), + workdir_manager=WorkdirManager(base_dir=Path.cwd()), + chunking_service=ChunkingService(max_tokens=6, overlap_sentences=1), + embedding_batch_size=2, + stt_model_version="chirp_2", + embedding_model_version="v001", + max_source_size_bytes=1, + ) + use_case = ProcessVideoUseCase( + video_repository=video_repository, + orchestrator=orchestrator, + delete_video_use_case=DeleteVideoUseCase( + video_repository=video_repository, + artifact_repository=artifact_repository, + storage_client=storage_client, + ), + stt_model_version="chirp_2", + embedding_model_version="v001", + ) + + messages: list[str] = [] + sink_id = logger.add(messages.append, format="{message}") + try: + result = await use_case.execute(video_id=video_id, trace_id="trace-extract-limit") + finally: + logger.remove(sink_id) + + assert result.action == "failed" + assert result.failed_stage == "EXTRACT" + failed_video = await video_repository.get_video(video_id) + assert failed_video.failed_stage == "EXTRACT" + assert failed_video.processing_claimed_at is None + assert any( + "Audio preparation failed failed_stage=EXTRACT" in message + and "Source size exceeds" in message + for message in messages + ) + + +@pytest.mark.asyncio +async def test_terminal_failure_hands_off_when_delete_wins_race( + video_repository, + artifact_repository, + storage_client, +) -> None: + video_id = str(uuid4()) + storage_client.objects["videos/delete-race.mp4"] = b"video" + await video_repository.create_video( + VideoRecord( + id=video_id, + user_id=str(uuid4()), + storage_path="videos/delete-race.mp4", + status="UPLOADED", + ) + ) + + async def delete_then_fail(**kwargs) -> None: + del kwargs + await video_repository.set_status(video_id, "DELETING") + raise AudioPreparationError("Audio part upload failed: part-001.flac") + + orchestrator = AsyncMock() + orchestrator.run.side_effect = delete_then_fail + use_case = ProcessVideoUseCase( + video_repository=video_repository, + orchestrator=orchestrator, + delete_video_use_case=DeleteVideoUseCase( + video_repository=video_repository, + artifact_repository=artifact_repository, + storage_client=storage_client, + ), + stt_model_version="chirp_3", + embedding_model_version="v001", + ) + + result = await use_case.execute(video_id=video_id, trace_id="trace-delete-race") + + assert result.action == "deleted" + assert await video_repository.get_video(video_id) is None + assert "videos/delete-race.mp4" not in storage_client.objects diff --git a/services/pipeline-worker/tests/unit/test_settings.py b/services/pipeline-worker/tests/unit/test_settings.py index 82269a9..83b89d5 100644 --- a/services/pipeline-worker/tests/unit/test_settings.py +++ b/services/pipeline-worker/tests/unit/test_settings.py @@ -54,7 +54,12 @@ def test_settings_loads_defaults_from_environment(monkeypatch: pytest.MonkeyPatc assert settings.embedding_batch_size == 16 assert settings.chunk_max_tokens == 300 assert settings.download_timeout_sec == 600 - assert settings.youtube_max_duration_sec == 1800 + assert settings.max_audio_duration_sec == 3600 + assert settings.audio_part_duration_sec == 900 + assert settings.audio_part_overlap_sec == 5 + assert settings.stt_part_concurrency == 2 + assert settings.audio_processing_timeout_sec == 120 + assert settings.youtube_max_duration_sec == 3600 assert settings.youtube_max_filesize_bytes == 500 * 1024 * 1024 assert settings.youtube_max_height == 720 @@ -79,6 +84,44 @@ def test_settings_reads_vision_max_output_tokens(monkeypatch: pytest.MonkeyPatch assert settings.vision_max_output_tokens == 2048 +def test_settings_reads_long_audio_overrides(monkeypatch: pytest.MonkeyPatch) -> None: + _set_env(monkeypatch, { + "MAX_AUDIO_DURATION_SEC": "3500", + "AUDIO_PART_DURATION_SEC": "800", + "AUDIO_PART_OVERLAP_SEC": "4", + "STT_PART_CONCURRENCY": "2", + "AUDIO_PROCESSING_TIMEOUT_SEC": "300", + "YOUTUBE_MAX_DURATION_SEC": "3500", + }) + + settings = Settings(_env_file=None) + + assert settings.max_audio_duration_sec == 3500 + assert settings.audio_part_duration_sec == 800 + assert settings.audio_part_overlap_sec == 4 + assert settings.stt_part_concurrency == 2 + assert settings.audio_processing_timeout_sec == 300 + assert settings.youtube_max_duration_sec == 3500 + + +@pytest.mark.parametrize( + "overrides", + [ + {"STT_PART_CONCURRENCY": "3"}, + {"AUDIO_PART_DURATION_SEC": "5", "AUDIO_PART_OVERLAP_SEC": "5"}, + {"AUDIO_PART_DURATION_SEC": "1196", "AUDIO_PART_OVERLAP_SEC": "5"}, + ], +) +def test_settings_rejects_unsafe_long_audio_combinations( + monkeypatch: pytest.MonkeyPatch, + overrides: dict[str, str], +) -> None: + _set_env(monkeypatch, overrides) + + with pytest.raises(ValidationError): + Settings(_env_file=None) + + def test_settings_reads_stt_batch_timeouts(monkeypatch: pytest.MonkeyPatch) -> None: _set_env(monkeypatch, { "STT_SUBMIT_TIMEOUT_SEC": "30", diff --git a/services/pipeline-worker/tests/unit/test_stt_batch_callable.py b/services/pipeline-worker/tests/unit/test_stt_batch_callable.py index 1107c5d..94ba329 100644 --- a/services/pipeline-worker/tests/unit/test_stt_batch_callable.py +++ b/services/pipeline-worker/tests/unit/test_stt_batch_callable.py @@ -138,6 +138,11 @@ def test_parse_batch_recognize_response_uses_word_offsets_for_segment_timestamps {"text": "Hello world.", "start_ms": 100, "end_ms": 1250}, {"text": "Again.", "start_ms": 1500, "end_ms": 2000}, ] + assert parsed["words"] == [ + {"text": "Hello", "start_ms": 100, "end_ms": 400}, + {"text": "world.", "start_ms": 500, "end_ms": 1250}, + {"text": "Again.", "start_ms": 1500, "end_ms": 2000}, + ] def test_parse_batch_recognize_response_splits_segment_after_one_hundred_words_without_punctuation() -> None: diff --git a/services/pipeline-worker/tests/unit/test_transcript_merge_service.py b/services/pipeline-worker/tests/unit/test_transcript_merge_service.py new file mode 100644 index 0000000..9a857ab --- /dev/null +++ b/services/pipeline-worker/tests/unit/test_transcript_merge_service.py @@ -0,0 +1,67 @@ +from src.infra.ai.google_stt_adapter import ( + STTTranscriptionResult, + TranscriptWordDTO, +) +from src.services.transcript_merge_service import AudioPart, TranscriptMergeService + + +def _result(*words: TranscriptWordDTO) -> STTTranscriptionResult: + return STTTranscriptionResult( + segments=[], + stt_model_version="chirp_3", + words=list(words), + ) + + +def test_merge_assigns_overlap_words_by_midpoint_and_global_time() -> None: + parts = [ + AudioPart(index=0, start_ms=0, end_ms=900_000, storage_path="part-0"), + AudioPart(index=1, start_ms=895_000, end_ms=1_200_000, storage_path="part-1"), + ] + first_result = _result( + TranscriptWordDTO("before", 896_000, 897_000), + TranscriptWordDTO("drop-first", 899_000, 900_000), + ) + second_result = _result( + TranscriptWordDTO("drop-second", 1_000, 2_000), + TranscriptWordDTO("after.", 4_000, 5_000), + ) + + merged = TranscriptMergeService().merge( + parts=parts, + results=[first_result, second_result], + duration_ms=1_200_000, + ) + + assert [word.text for word in merged.words or []] == ["before", "after."] + assert [(word.start_ms, word.end_ms) for word in merged.words or []] == [ + (896_000, 897_000), + (899_000, 900_000), + ] + assert merged.segments[0].text == "before after." + + +def test_merge_sorts_results_and_keeps_segment_time_invariants() -> None: + parts = [ + AudioPart(index=0, start_ms=0, end_ms=900_000, storage_path="part-0"), + AudioPart(index=1, start_ms=895_000, end_ms=1_200_000, storage_path="part-1"), + ] + results = [ + _result(TranscriptWordDTO("first.", 100, 200)), + _result(TranscriptWordDTO("second.", 10_000, 10_100)), + ] + + merged = TranscriptMergeService().merge( + parts=parts, + results=results, + duration_ms=1_200_000, + ) + + assert [segment.text for segment in merged.segments] == ["first.", "second."] + assert all( + 0 <= segment.start_ms <= segment.end_ms <= 1_200_000 + for segment in merged.segments + ) + assert [segment.start_ms for segment in merged.segments] == sorted( + segment.start_ms for segment in merged.segments + ) diff --git a/services/pipeline-worker/tests/unit/test_youtube_downloader.py b/services/pipeline-worker/tests/unit/test_youtube_downloader.py index e7f7cd8..d2c952b 100644 --- a/services/pipeline-worker/tests/unit/test_youtube_downloader.py +++ b/services/pipeline-worker/tests/unit/test_youtube_downloader.py @@ -10,6 +10,7 @@ class FakeYoutubeDL: metadata: dict[str, int] = {} options: list[dict] = [] error: Exception | None = None + downloaded_content: bytes = b"video" def __init__(self, options): self.options = options @@ -26,7 +27,9 @@ def extract_info(self, source_url: str, *, download: bool): if self.error is not None: raise self.error if download: - Path(self.options["outtmpl"].replace("%(ext)s", "mp4")).write_bytes(b"video") + Path(self.options["outtmpl"].replace("%(ext)s", "mp4")).write_bytes( + self.downloaded_content + ) return self.metadata @@ -36,11 +39,12 @@ def reset_fake_youtube_dl() -> None: FakeYoutubeDL.metadata = {} FakeYoutubeDL.options = [] FakeYoutubeDL.error = None + FakeYoutubeDL.downloaded_content = b"video" def build_downloader(*, proxy_url: str = "") -> YtDlpYoutubeDownloader: return YtDlpYoutubeDownloader( - max_duration_sec=1800, + max_duration_sec=3600, max_filesize_bytes=500, max_height=720, timeout_sec=600, @@ -51,7 +55,7 @@ def build_downloader(*, proxy_url: str = "") -> YtDlpYoutubeDownloader: @pytest.mark.asyncio async def test_youtube_downloader_rejects_duration_over_limit_without_download(tmp_path) -> None: - FakeYoutubeDL.metadata = {"duration": 1801} + FakeYoutubeDL.metadata = {"duration": 3601} downloader = build_downloader() with pytest.raises(DownloadError): @@ -60,6 +64,32 @@ async def test_youtube_downloader_rejects_duration_over_limit_without_download(t assert FakeYoutubeDL.calls == [False] +@pytest.mark.asyncio +async def test_youtube_downloader_rejects_downloaded_file_over_limit(tmp_path) -> None: + FakeYoutubeDL.metadata = {"duration": 30} + FakeYoutubeDL.downloaded_content = b"x" * 501 + downloader = build_downloader() + + with pytest.raises(DownloadError, match="size exceeds 500 bytes"): + await downloader.download("https://youtu.be/actual-too-large", tmp_path / "source.mp4") + + assert FakeYoutubeDL.calls == [False, True] + + +@pytest.mark.asyncio +async def test_youtube_downloader_accepts_exact_duration_and_file_size_limits(tmp_path) -> None: + FakeYoutubeDL.metadata = {"duration": 3600, "filesize": 500} + FakeYoutubeDL.downloaded_content = b"x" * 500 + destination = tmp_path / "source.mp4" + downloader = build_downloader() + + result = await downloader.download("https://youtu.be/exact-limit", destination) + + assert result == destination + assert destination.stat().st_size == 500 + assert FakeYoutubeDL.calls == [False, True] + + @pytest.mark.asyncio async def test_youtube_downloader_rejects_filesize_over_limit(tmp_path) -> None: FakeYoutubeDL.metadata = {"duration": 30, "filesize": 501} From 837df6de1a00cc55b151ababadb9ebe8c7ef910a Mon Sep 17 00:00:00 2001 From: baekyutae Date: Wed, 15 Jul 2026 12:57:59 +0900 Subject: [PATCH 4/6] =?UTF-8?q?feat(pipeline=20worker)=20=ED=83=80?= =?UTF-8?q?=EC=9E=84=EC=8A=A4=ED=83=AC=ED=94=84=EA=B0=80=20=EA=BC=AC?= =?UTF-8?q?=EC=9D=B8=EB=8B=A8=EC=96=B4=EB=A5=BC=20=EB=B3=B4=EC=A0=95?= =?UTF-8?q?=ED=95=98=EB=8A=94=20=EB=A1=9C=EC=A7=81=20=EC=B6=94=EA=B0=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 고립된 역전 단어는 앞뒤 정상 단어를 기준으로 보정합니다. - 여러 고립 이상치도 모두 보정합니다. - 첫·마지막 단어, 연속 역전, 겹친 인접 경계는 비재시도 오류로 실패합니다. - 보정 시 원본값·보정값을 경고 로그로 남깁니다. - 보정 불가 시 인접값과 실패 사유를 오류 로그로 남깁니다. - 외부 응답 형식은 변경하지 않았습니다. --- .../src/infra/ai/stt_batch_callable.py | 127 +++++++++++++++++- .../tests/unit/test_stt_batch_callable.py | 124 +++++++++++++++-- 2 files changed, 238 insertions(+), 13 deletions(-) diff --git a/services/pipeline-worker/src/infra/ai/stt_batch_callable.py b/services/pipeline-worker/src/infra/ai/stt_batch_callable.py index 11f335a..1a599ca 100644 --- a/services/pipeline-worker/src/infra/ai/stt_batch_callable.py +++ b/services/pipeline-worker/src/infra/ai/stt_batch_callable.py @@ -1,7 +1,7 @@ """Factory for a production STTCallable using Google Cloud Speech-to-Text v2 BatchRecognize.""" import asyncio -from typing import Any +from typing import Any, NoReturn from google.api_core.client_options import ClientOptions from google.rpc import code_pb2 @@ -20,6 +20,7 @@ code_pb2.UNAVAILABLE: ("UNAVAILABLE", True), code_pb2.INVALID_ARGUMENT: ("INVALID_REQUEST", False), } +_INVALID_WORD_OFFSETS_MESSAGE = "STT word time offsets are invalid" # google stt: 초단위 -> biblio: 밀리초 단위 변환 def _duration_to_ms(duration: Any, trace_id: str) -> int: @@ -78,11 +79,127 @@ def _normalize_word(word: Any, trace_id: str) -> TranscriptWordDTO: end_ms = _duration_to_ms(word.end_offset, trace_id) except AttributeError as exc: raise _stt_parse_error("STT word time offsets missing", trace_id) from exc - if end_ms < start_ms: - raise _stt_parse_error("STT word time offsets are invalid", trace_id) return TranscriptWordDTO(text=text, start_ms=start_ms, end_ms=end_ms) +def _has_reversed_offsets(word: TranscriptWordDTO) -> bool: + return word.end_ms < word.start_ms + + +def _corrected_word_bounds( + word: TranscriptWordDTO, + previous_end_ms: int, + next_start_ms: int, +) -> tuple[int, int]: + start_is_usable = previous_end_ms <= word.start_ms <= next_start_ms + end_is_usable = previous_end_ms <= word.end_ms <= next_start_ms + if start_is_usable and not end_is_usable: + return word.start_ms, next_start_ms + if end_is_usable and not start_is_usable: + return previous_end_ms, word.end_ms + return previous_end_ms, next_start_ms + + +def _raise_unrepairable_word_offsets( + words: list[TranscriptWordDTO], + word_index: int, + trace_id: str, + uri: str, + reason: str, +) -> NoReturn: + word = words[word_index] + previous_end_ms = words[word_index - 1].end_ms if word_index > 0 else None + next_start_ms = words[word_index + 1].start_ms if word_index < len(words) - 1 else None + logger.bind( + trace_id=trace_id, + stt_uri=uri, + word_index=word_index, + word=word.text, + raw_start_ms=word.start_ms, + raw_end_ms=word.end_ms, + previous_end_ms=previous_end_ms, + next_start_ms=next_start_ms, + reason=reason, + ).error( + "STT word time offsets cannot be corrected uri={} word_index={} word={} " + "raw_start_ms={} raw_end_ms={} previous_end_ms={} next_start_ms={} reason={}", + uri, + word_index, + word.text, + word.start_ms, + word.end_ms, + previous_end_ms, + next_start_ms, + reason, + ) + raise _stt_parse_error(_INVALID_WORD_OFFSETS_MESSAGE, trace_id) + + +def _repair_reversed_word( + words: list[TranscriptWordDTO], + word_index: int, + trace_id: str, + uri: str, +) -> TranscriptWordDTO: + word = words[word_index] + if word_index == 0 or word_index == len(words) - 1: + _raise_unrepairable_word_offsets(words, word_index, trace_id, uri, "missing_neighbor") + + previous_word = words[word_index - 1] + next_word = words[word_index + 1] + if _has_reversed_offsets(previous_word) or _has_reversed_offsets(next_word): + _raise_unrepairable_word_offsets(words, word_index, trace_id, uri, "adjacent_reversal") + if previous_word.end_ms > next_word.start_ms: + _raise_unrepairable_word_offsets(words, word_index, trace_id, uri, "overlapping_neighbors") + + corrected_start_ms, corrected_end_ms = _corrected_word_bounds( + word, + previous_word.end_ms, + next_word.start_ms, + ) + logger.bind( + trace_id=trace_id, + stt_uri=uri, + word_index=word_index, + word=word.text, + raw_start_ms=word.start_ms, + raw_end_ms=word.end_ms, + corrected_start_ms=corrected_start_ms, + corrected_end_ms=corrected_end_ms, + ).warning( + "STT word time offsets corrected uri={} word_index={} word={} " + "raw_start_ms={} raw_end_ms={} corrected_start_ms={} corrected_end_ms={}", + uri, + word_index, + word.text, + word.start_ms, + word.end_ms, + corrected_start_ms, + corrected_end_ms, + ) + return TranscriptWordDTO( + text=word.text, + start_ms=corrected_start_ms, + end_ms=corrected_end_ms, + ) + + +def _repair_reversed_word_offsets( + words: list[TranscriptWordDTO], + trace_id: str, + uri: str, +) -> list[TranscriptWordDTO]: + repaired_words: list[TranscriptWordDTO] = [] + for word_index, word in enumerate(words): + repaired_word = ( + _repair_reversed_word(words, word_index, trace_id, uri) + if _has_reversed_offsets(word) + else word + ) + repaired_words.append(repaired_word) + return repaired_words + + def _parse_batch_recognize_response(response: Any, stt_model_version: str, trace_id: str = "") -> dict: normalized_words: list[TranscriptWordDTO] = [] for uri, file_result in response.results.items(): @@ -92,6 +209,7 @@ def _parse_batch_recognize_response(response: Any, stt_model_version: str, trace transcript = getattr(file_result, "transcript", None) if transcript is None: continue + file_words: list[TranscriptWordDTO] = [] for result in transcript.results: if not result.alternatives: continue @@ -100,7 +218,8 @@ def _parse_batch_recognize_response(response: Any, stt_model_version: str, trace words = list(getattr(alt, "words", []) or []) if text and not words: raise _stt_parse_error("STT word time offsets missing", trace_id) - normalized_words.extend(_normalize_word(word, trace_id) for word in words) + file_words.extend(_normalize_word(word, trace_id) for word in words) + normalized_words.extend(_repair_reversed_word_offsets(file_words, trace_id, str(uri))) normalized_words.sort(key=lambda word: word.start_ms) segments = segments_from_words(normalized_words) return { diff --git a/services/pipeline-worker/tests/unit/test_stt_batch_callable.py b/services/pipeline-worker/tests/unit/test_stt_batch_callable.py index 94ba329..6d906f6 100644 --- a/services/pipeline-worker/tests/unit/test_stt_batch_callable.py +++ b/services/pipeline-worker/tests/unit/test_stt_batch_callable.py @@ -2,6 +2,7 @@ import pytest from google.rpc import code_pb2 +from loguru import logger from src.infra.ai.google_stt_adapter import ExternalAIAdapterError from src.infra.ai.stt_batch_callable import _parse_batch_recognize_response @@ -46,6 +47,17 @@ def _result_with_words(text: str, words: list[SimpleNamespace]) -> SimpleNamespa ) +def _response_with_words(text: str, words: list[SimpleNamespace]) -> SimpleNamespace: + result = _result_with_words(text, words) + return SimpleNamespace( + results={ + "gs://bucket/audio.flac": SimpleNamespace( + inline_result=SimpleNamespace(transcript=SimpleNamespace(results=[result])) + ) + } + ) + + def test_parse_batch_recognize_response_reads_inline_result_transcript() -> None: inline_transcript = SimpleNamespace(results=[_result("hello world", 1.25, 0.1)]) file_result = SimpleNamespace( @@ -212,18 +224,112 @@ def test_parse_batch_recognize_response_fails_when_word_end_offset_is_missing() _parse_batch_recognize_response(response, "chirp_3", trace_id="trace-5") -def test_parse_batch_recognize_response_fails_when_word_offsets_are_reversed() -> None: - result = _result_with_words("Hello.", [_word("Hello.", 2.0, 1.0)]) - response = SimpleNamespace( - results={ - "gs://bucket/audio.flac": SimpleNamespace( - inline_result=SimpleNamespace(transcript=SimpleNamespace(results=[result])) +class TestReversedWordOffsetRepair: + def test_repairs_invalid_start_from_previous_word_end(self) -> None: + words = [ + _word("보시면", 247.640, 248.120), + _word("13으로부터", 249.160, 248.360), + _word("왼쪽", 248.360, 251.240), + ] + messages: list[str] = [] + sink_id = logger.add(messages.append, format="{message}", level="WARNING") + try: + parsed = _parse_batch_recognize_response( + _response_with_words("보시면 13으로부터 왼쪽", words), + "chirp_3", + trace_id="trace-repair-start", ) + finally: + logger.remove(sink_id) + + assert parsed["words"][1] == { + "text": "13으로부터", + "start_ms": 248120, + "end_ms": 248360, } + assert any( + "STT word time offsets corrected uri=gs://bucket/audio.flac word_index=1 " + "word=13으로부터 raw_start_ms=249160 raw_end_ms=248360 " + "corrected_start_ms=248120 corrected_end_ms=248360" in message + for message in messages + ) + + def test_repairs_invalid_end_from_next_word_start(self) -> None: + words = [ + _word("a", 1.0, 3.0), + _word("b", 6.0, 2.0), + _word("c", 7.0, 9.0), + ] + + parsed = _parse_batch_recognize_response( + _response_with_words("a b c", words), + "chirp_3", + trace_id="trace-repair-end", + ) + + assert parsed["words"][1] == {"text": "b", "start_ms": 6000, "end_ms": 7000} + + def test_uses_both_neighbor_boundaries_when_reversed_offsets_are_in_bounds(self) -> None: + words = [ + _word("a", 1.0, 3.0), + _word("b", 6.0, 4.0), + _word("c", 7.0, 9.0), + ] + + parsed = _parse_batch_recognize_response( + _response_with_words("a b c", words), + "chirp_3", + trace_id="trace-repair-both", + ) + + assert parsed["words"][1] == {"text": "b", "start_ms": 3000, "end_ms": 7000} + + def test_repairs_multiple_isolated_words(self) -> None: + words = [ + _word("a", 0.0, 1.0), + _word("b", 3.0, 0.5), + _word("c", 4.0, 5.0), + _word("d", 7.0, 4.5), + _word("e", 8.0, 9.0), + ] + + parsed = _parse_batch_recognize_response( + _response_with_words("a b c d e", words), + "chirp_3", + trace_id="trace-repair-multiple", + ) + + assert parsed["words"][1] == {"text": "b", "start_ms": 3000, "end_ms": 4000} + assert parsed["words"][3] == {"text": "d", "start_ms": 7000, "end_ms": 8000} + + @pytest.mark.parametrize( + "words", + [ + [_word("only", 2.0, 1.0)], + [ + _word("a", 0.0, 1.0), + _word("b", 3.0, 0.5), + _word("c", 4.0, 2.0), + _word("d", 5.0, 6.0), + ], + [ + _word("a", 0.0, 5.0), + _word("b", 6.0, 2.0), + _word("c", 4.0, 7.0), + ], + ], + ids=["missing-neighbors", "consecutive-reversals", "overlapping-neighbor-bounds"], ) - - with pytest.raises(ExternalAIAdapterError, match="word time offsets are invalid"): - _parse_batch_recognize_response(response, "chirp_3", trace_id="trace-6") + def test_fails_when_reversed_offsets_cannot_be_repaired( + self, + words: list[SimpleNamespace], + ) -> None: + with pytest.raises(ExternalAIAdapterError, match="word time offsets are invalid"): + _parse_batch_recognize_response( + _response_with_words("invalid words", words), + "chirp_3", + trace_id="trace-unrepairable", + ) def test_parse_batch_recognize_response_wraps_invalid_duration_values() -> None: From ad858c9a3798b6b0f096167973afbb1f07e0f537 Mon Sep 17 00:00:00 2001 From: baekyutae Date: Wed, 15 Jul 2026 22:01:39 +0900 Subject: [PATCH 5/6] =?UTF-8?q?=20fix(pipeline-worker):=20=EC=9E=84?= =?UTF-8?q?=EB=B2=A0=EB=94=A9=20=ED=83=80=EC=9E=84=EC=95=84=EC=9B=83?= =?UTF-8?q?=EA=B3=BC=20=EC=9E=AC=EC=8B=9C=EB=8F=84=20=EC=A0=95=EC=B1=85=20?= =?UTF-8?q?=EB=B3=B4=EC=99=84=20(#100)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- infra/terraform/envs/gcp-perf/main.tf | 3 +- .../envs/gcp-perf/terraform.tfvars.example | 2 + infra/terraform/envs/gcp-perf/variables.tf | 28 ++ .../src/infra/ai/embedding_client.py | 273 ++++++++++++++---- .../src/infra/ai/google_stt_adapter.py | 10 +- .../src/infra/ai/retry_policy.py | 11 + services/pipeline-worker/tests/support.py | 16 +- .../tests/unit/test_embedding_client.py | 251 +++++++++++++--- 8 files changed, 488 insertions(+), 106 deletions(-) create mode 100644 services/pipeline-worker/src/infra/ai/retry_policy.py diff --git a/infra/terraform/envs/gcp-perf/main.tf b/infra/terraform/envs/gcp-perf/main.tf index bd78b56..123ef62 100644 --- a/infra/terraform/envs/gcp-perf/main.tf +++ b/infra/terraform/envs/gcp-perf/main.tf @@ -444,7 +444,8 @@ module "pipeline_worker" { YOUTUBE_PROXY_URL = "socks5://${module.embedding_vm.private_ip}:1080" GCS_VIDEO_BUCKET_NAME = module.object_storage.bucket_names.video EMBEDDING_API_URL = local.embedding_vm_url - EMBEDDING_TIMEOUT_SEC = "60" + EMBEDDING_TIMEOUT_SEC = tostring(var.pipeline_embedding_timeout_sec) + EMBEDDING_BATCH_SIZE = tostring(var.pipeline_embedding_batch_size) } secret_env_vars = { diff --git a/infra/terraform/envs/gcp-perf/terraform.tfvars.example b/infra/terraform/envs/gcp-perf/terraform.tfvars.example index 7a2cca6..e330a26 100644 --- a/infra/terraform/envs/gcp-perf/terraform.tfvars.example +++ b/infra/terraform/envs/gcp-perf/terraform.tfvars.example @@ -23,6 +23,8 @@ embedding_subnet_cidr = "10.20.3.0/24" enable_managed_embedding_cloud_run = false embedding_vm_machine_type = "e2-standard-4" embedding_vm_model_disk_size_gb = 100 +pipeline_embedding_timeout_sec = 180 +pipeline_embedding_batch_size = 16 # 모델 버전은 타임스탬프를 붙여 관리한다. # 경로의 마지막 조각이 ModelRelease 버전 이름이 되므로, 버전마다 다른 이름을 쓴다. diff --git a/infra/terraform/envs/gcp-perf/variables.tf b/infra/terraform/envs/gcp-perf/variables.tf index a3ee940..ab5fea6 100644 --- a/infra/terraform/envs/gcp-perf/variables.tf +++ b/infra/terraform/envs/gcp-perf/variables.tf @@ -115,6 +115,34 @@ variable "worker_max_instance_count" { default = 1 } +variable "pipeline_embedding_timeout_sec" { + type = number + default = 180 + description = "Timeout in seconds for each pipeline-worker embedding HTTP request." + + validation { + condition = ( + var.pipeline_embedding_timeout_sec > 0 && + floor(var.pipeline_embedding_timeout_sec) == var.pipeline_embedding_timeout_sec + ) + error_message = "pipeline_embedding_timeout_sec must be a positive integer." + } +} + +variable "pipeline_embedding_batch_size" { + type = number + default = 16 + description = "Maximum enriched texts sent in one pipeline-worker embedding request." + + validation { + condition = ( + var.pipeline_embedding_batch_size > 0 && + floor(var.pipeline_embedding_batch_size) == var.pipeline_embedding_batch_size + ) + error_message = "pipeline_embedding_batch_size must be a positive integer." + } +} + variable "cloudrun_subnet_cidr" { type = string default = "10.20.1.0/24" diff --git a/services/pipeline-worker/src/infra/ai/embedding_client.py b/services/pipeline-worker/src/infra/ai/embedding_client.py index e9b43f8..c04ef62 100644 --- a/services/pipeline-worker/src/infra/ai/embedding_client.py +++ b/services/pipeline-worker/src/infra/ai/embedding_client.py @@ -1,9 +1,45 @@ import asyncio +import random +import time from dataclasses import dataclass import httpx +from loguru import logger from src.infra.ai.google_stt_adapter import ExternalAIAdapterError +from src.infra.ai.retry_policy import ( + JitterCallable, + SleepCallable, + exponential_backoff_with_jitter, +) + + +def _log_request_event( + *, + level: str, + event: str, + trace_id: str, + model_version: str, + text_count: int, + attempt: int, + duration_ms: float, + status_code: int | str, + error_code: str, + retry_delay_seconds: float, +) -> None: + logger.bind(trace_id=trace_id).log( + level, + "event={} model_version={} text_count={} attempt={} duration_ms={:.1f} " + "status_code={} error_code={} retry_delay_seconds={:.3f}", + event, + model_version, + text_count, + attempt, + duration_ms, + status_code, + error_code, + retry_delay_seconds, + ) @dataclass(slots=True) @@ -12,6 +48,38 @@ class EmbeddingBatchResult: model_version: str +@dataclass(frozen=True, slots=True) +class _EmbeddingAttemptContext: + trace_id: str + model_version: str + text_count: int + attempt: int + started_at: float + + +def _log_attempt( + context: _EmbeddingAttemptContext, + *, + level: str, + event: str, + status_code: int | str, + error_code: str = "-", + retry_delay_seconds: float = 0.0, +) -> None: + _log_request_event( + level=level, + event=event, + trace_id=context.trace_id, + model_version=context.model_version, + text_count=context.text_count, + attempt=context.attempt, + duration_ms=(time.monotonic() - context.started_at) * 1000, + status_code=status_code, + error_code=error_code, + retry_delay_seconds=retry_delay_seconds, + ) + + class EmbeddingClient: def __init__( self, @@ -21,12 +89,16 @@ def __init__( max_retries: int, model_version: str, client: httpx.AsyncClient | None = None, + sleep: SleepCallable = asyncio.sleep, + jitter: JitterCallable = random.random, ) -> None: self._base_url = base_url.rstrip("/") self._timeout_sec = timeout_sec self._max_retries = max_retries self._model_version = model_version self._client = client or httpx.AsyncClient() + self._sleep = sleep + self._jitter = jitter async def get_ready_model_versions(self, trace_id: str) -> list[str]: response = await self._client.get( @@ -68,68 +140,157 @@ async def embed_texts( trace_id: str, model_version: str | None = None, ) -> EmbeddingBatchResult: - if not texts: + self._validate_texts(texts, trace_id) + requested_model_version = model_version or self._model_version + for attempt_index in range(self._max_retries + 1): + result = await self._attempt_embedding_batch( + texts=texts, + trace_id=trace_id, + model_version=requested_model_version, + attempt_index=attempt_index, + ) + if result is not None: + return result + raise RuntimeError("Embedding retry loop exited unexpectedly") + + @staticmethod + def _validate_texts(texts: list[str], trace_id: str) -> None: + if texts: + return + raise ExternalAIAdapterError( + code="INVALID_REQUEST", + message="texts must not be empty", + trace_id=trace_id, + provider="embedding-endpoint", + retryable=False, + ) + + async def _attempt_embedding_batch( + self, + *, + texts: list[str], + trace_id: str, + model_version: str, + attempt_index: int, + ) -> EmbeddingBatchResult | None: + context = _EmbeddingAttemptContext( + trace_id=trace_id, + model_version=model_version, + text_count=len(texts), + attempt=attempt_index + 1, + started_at=time.monotonic(), + ) + try: + result, status_code = await self._post_embedding_batch(texts, trace_id, model_version) + except httpx.TimeoutException as exc: + _log_attempt( + context, + level="ERROR", + event="embedding.request.timeout", + status_code="-", + error_code="TIMEOUT", + ) raise ExternalAIAdapterError( - code="INVALID_REQUEST", - message="texts must not be empty", + code="TIMEOUT", + message="Embedding endpoint timed out", trace_id=trace_id, provider="embedding-endpoint", retryable=False, + attempt_count=context.attempt, + ) from exc + except ExternalAIAdapterError as exc: + return await self._handle_adapter_error( + exc=exc, + context=context, + attempt_index=attempt_index, + ) + except httpx.HTTPStatusError as exc: + _log_attempt( + context, + level="ERROR", + event="embedding.request.failed", + status_code=exc.response.status_code, + error_code="HTTP_STATUS_ERROR", + ) + raise + except Exception as exc: + _log_attempt( + context, + level="ERROR", + event="embedding.request.failed", + status_code="-", + error_code=type(exc).__name__, ) + raise + _log_attempt( + context, + level="INFO", + event="embedding.request.success", + status_code=status_code, + ) + return result - last_error: Exception | None = None - for attempt in range(self._max_retries + 1): - try: - response = await self._client.post( - f"{self._base_url}/embed", - json={ - "texts": texts, - "model_version": model_version or self._model_version, - }, - headers={"X-Trace-Id": trace_id}, - timeout=self._timeout_sec, - ) - if response.status_code == 503: - raise ExternalAIAdapterError( - code="UNAVAILABLE", - message="Embedding endpoint unavailable", - trace_id=trace_id, - provider="embedding-endpoint", - retryable=True, - ) - response.raise_for_status() - embeddings = response.json()["embeddings"] - if len(embeddings) != len(texts): - raise ExternalAIAdapterError( - code="INTERNAL_ERROR", - message="Embedding count mismatch", - trace_id=trace_id, - provider="embedding-endpoint", - retryable=False, - ) - return EmbeddingBatchResult( - embeddings=embeddings, - model_version=model_version or self._model_version, - ) - except httpx.TimeoutException: - last_error = ExternalAIAdapterError( - code="TIMEOUT", - message="Embedding endpoint timed out", - trace_id=trace_id, - provider="embedding-endpoint", - retryable=True, - ) - except ExternalAIAdapterError as exc: - last_error = exc - if not exc.retryable: - raise - if attempt >= self._max_retries: - assert last_error is not None - raise last_error - await asyncio.sleep(0) - - assert last_error is not None - raise last_error + async def _post_embedding_batch( + self, + texts: list[str], + trace_id: str, + model_version: str, + ) -> tuple[EmbeddingBatchResult, int]: + response = await self._client.post( + f"{self._base_url}/embed", + json={"texts": texts, "model_version": model_version}, + headers={"X-Trace-Id": trace_id}, + timeout=self._timeout_sec, + ) + if response.status_code == 503: + raise ExternalAIAdapterError( + code="UNAVAILABLE", + message="Embedding endpoint unavailable", + trace_id=trace_id, + provider="embedding-endpoint", + retryable=True, + ) + response.raise_for_status() + embeddings = response.json()["embeddings"] + if len(embeddings) != len(texts): + raise ExternalAIAdapterError( + code="INTERNAL_ERROR", + message="Embedding count mismatch", + trace_id=trace_id, + provider="embedding-endpoint", + retryable=False, + ) + return EmbeddingBatchResult(embeddings=embeddings, model_version=model_version), response.status_code + + async def _handle_adapter_error( + self, + *, + exc: ExternalAIAdapterError, + context: _EmbeddingAttemptContext, + attempt_index: int, + ) -> None: + exc.attempt_count = context.attempt + status_code = 503 if exc.code == "UNAVAILABLE" else 200 + if exc.code != "UNAVAILABLE" or attempt_index >= self._max_retries: + _log_attempt( + context, + level="ERROR", + event="embedding.request.failed", + status_code=status_code, + error_code=exc.code, + ) + raise exc + delay_seconds = exponential_backoff_with_jitter(attempt_index, self._jitter()) + _log_attempt( + context, + level="WARNING", + event="embedding.request.retry", + status_code=status_code, + error_code=exc.code, + retry_delay_seconds=delay_seconds, + ) + await self._sleep(delay_seconds) + return None async def aclose(self) -> None: await self._client.aclose() diff --git a/services/pipeline-worker/src/infra/ai/google_stt_adapter.py b/services/pipeline-worker/src/infra/ai/google_stt_adapter.py index 5321bd3..b6ba847 100644 --- a/services/pipeline-worker/src/infra/ai/google_stt_adapter.py +++ b/services/pipeline-worker/src/infra/ai/google_stt_adapter.py @@ -5,6 +5,12 @@ from loguru import logger +from src.infra.ai.retry_policy import ( + JitterCallable, + SleepCallable, + exponential_backoff_with_jitter, +) + MAX_WORDS_PER_SEGMENT = 100 SENTENCE_ENDING_MARKS = (".", "!", "?") @@ -45,8 +51,6 @@ def __str__(self) -> str: STTCallable = Callable[[str, str], Awaitable[dict[str, Any] | STTTranscriptionResult]] -SleepCallable = Callable[[float], Awaitable[None]] -JitterCallable = Callable[[], float] def segments_from_words(words: list[TranscriptWordDTO]) -> list[TranscriptSegmentDTO]: @@ -117,7 +121,7 @@ async def transcribe(self, *, audio_uri: str, trace_id: str) -> STTTranscription last_error.attempt_count = attempt + 1 if attempt >= self._max_retries: raise last_error - delay_seconds = (2**attempt) * (1 + (self._jitter() * 0.25)) + delay_seconds = exponential_backoff_with_jitter(attempt, self._jitter()) logger.bind(trace_id=trace_id).warning( "STT retry uri={} attempt={} code={} delay_seconds={:.3f}", audio_uri, diff --git a/services/pipeline-worker/src/infra/ai/retry_policy.py b/services/pipeline-worker/src/infra/ai/retry_policy.py new file mode 100644 index 0000000..ac81328 --- /dev/null +++ b/services/pipeline-worker/src/infra/ai/retry_policy.py @@ -0,0 +1,11 @@ +from collections.abc import Awaitable, Callable + + +SleepCallable = Callable[[float], Awaitable[None]] +JitterCallable = Callable[[], float] + +_JITTER_RATIO = 0.25 + + +def exponential_backoff_with_jitter(attempt_index: int, jitter_value: float) -> float: + return (2**attempt_index) * (1 + (jitter_value * _JITTER_RATIO)) diff --git a/services/pipeline-worker/tests/support.py b/services/pipeline-worker/tests/support.py index 8026298..c9c6eb2 100644 --- a/services/pipeline-worker/tests/support.py +++ b/services/pipeline-worker/tests/support.py @@ -10,6 +10,7 @@ from src.infra.ai.embedding_client import EmbeddingClient from src.infra.ai.google_stt_adapter import GoogleSTTAdapter +from src.infra.ai.retry_policy import JitterCallable, SleepCallable from src.infra.db.models import Base from src.infra.media.ffmpeg_client import FFmpegClient @@ -17,6 +18,14 @@ from sqlalchemy.ext.asyncio import AsyncEngine +async def _no_retry_sleep(delay_seconds: float) -> None: + del delay_seconds + + +def _zero_jitter() -> float: + return 0.0 + + async def create_test_engine() -> AsyncEngine: engine = create_async_engine( "sqlite+aiosqlite:///:memory:", @@ -67,6 +76,9 @@ def build_embedding_client( embeddings_factory: Callable[[list[str]], list[list[float]]] | None = None, fail_embed_times: int = 0, embedding_model_version: str = "v001", + max_retries: int = 2, + sleep: SleepCallable = _no_retry_sleep, + jitter: JitterCallable = _zero_jitter, ) -> EmbeddingClient: state = {"failures": fail_embed_times} @@ -98,9 +110,11 @@ def handler(request: httpx.Request) -> httpx.Response: return EmbeddingClient( base_url="https://embedding.local", timeout_sec=5, - max_retries=2, + max_retries=max_retries, model_version=embedding_model_version, client=httpx.AsyncClient(transport=transport), + sleep=sleep, + jitter=jitter, ) diff --git a/services/pipeline-worker/tests/unit/test_embedding_client.py b/services/pipeline-worker/tests/unit/test_embedding_client.py index dabdc4b..0bc2d0d 100644 --- a/services/pipeline-worker/tests/unit/test_embedding_client.py +++ b/services/pipeline-worker/tests/unit/test_embedding_client.py @@ -1,62 +1,223 @@ +from collections.abc import Awaitable, Callable + +import httpx import pytest +from src.infra.ai import embedding_client as embedding_client_module +from src.infra.ai.embedding_client import EmbeddingClient from src.infra.ai.google_stt_adapter import ExternalAIAdapterError from tests.support import build_embedding_client -@pytest.mark.asyncio -async def test_embedding_client_returns_embeddings() -> None: - client = build_embedding_client() - - result = await client.embed_texts(["alpha", "beta"], trace_id="trace-1") - - assert result.model_version == "v001" - assert len(result.embeddings) == 2 - - -@pytest.mark.asyncio -async def test_embedding_client_retries_503_and_succeeds() -> None: - client = build_embedding_client(fail_embed_times=1) - - result = await client.embed_texts(["alpha"], trace_id="trace-2") +EventFields = dict[str, object] - assert result.embeddings[0][0] == pytest.approx(5.0) +def _capture_events(monkeypatch: pytest.MonkeyPatch) -> list[EventFields]: + events: list[EventFields] = [] -@pytest.mark.asyncio -async def test_embedding_client_rejects_empty_input() -> None: - client = build_embedding_client() + def capture_event(**fields: object) -> None: + events.append(fields) - with pytest.raises(ExternalAIAdapterError): - await client.embed_texts([], trace_id="trace-3") + monkeypatch.setattr(embedding_client_module, "_log_request_event", capture_event) + return events -@pytest.mark.asyncio -async def test_embedding_client_reads_ready_model_versions_from_health() -> None: - client = build_embedding_client(model_version="v001") - - versions = await client.get_ready_model_versions(trace_id="trace-4") - - assert versions == ["v001"] - - -@pytest.mark.asyncio -async def test_embedding_client_get_model_version_uses_configured_ready_version() -> None: - client = build_embedding_client( - model_version="v002", - ready_model_versions=["v001", "v002"], - embedding_model_version="v002", +def _client_with_handler( + handler: Callable[[httpx.Request], httpx.Response], + *, + sleep: Callable[[float], Awaitable[None]], +) -> EmbeddingClient: + return EmbeddingClient( + base_url="https://embedding.local", + timeout_sec=5, + max_retries=3, + model_version="v001", + client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + sleep=sleep, + jitter=lambda: 0.0, ) - version = await client.get_model_version(trace_id="trace-5") - - assert version == "v002" - -@pytest.mark.asyncio -async def test_embedding_client_rejects_health_without_ready_model_versions() -> None: - client = build_embedding_client(health_payload={"status": "ok", "model_version": "v001"}) +class TestEmbeddingRequestSuccess: + @pytest.mark.asyncio + async def test_returns_embeddings_and_logs_required_fields( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + events = _capture_events(monkeypatch) + client = build_embedding_client() + + result = await client.embed_texts(["alpha", "beta"], trace_id="trace-1") + + assert result.model_version == "v001" + assert len(result.embeddings) == 2 + assert events[0]["event"] == "embedding.request.success" + assert { + "event", + "trace_id", + "model_version", + "text_count", + "attempt", + "duration_ms", + "status_code", + "error_code", + "retry_delay_seconds", + } <= events[0].keys() + assert events[0]["status_code"] == 200 + assert events[0]["text_count"] == 2 + + +class TestEmbeddingRetryPolicy: + @pytest.mark.asyncio + async def test_retries_only_503_with_exponential_backoff_and_jitter( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + events = _capture_events(monkeypatch) + delays: list[float] = [] + + async def record_sleep(delay_seconds: float) -> None: + delays.append(delay_seconds) + + client = build_embedding_client( + fail_embed_times=2, + max_retries=3, + sleep=record_sleep, + jitter=lambda: 1.0, + ) + + result = await client.embed_texts(["alpha"], trace_id="trace-2") + + assert result.embeddings[0][0] == pytest.approx(5.0) + assert delays == pytest.approx([1.25, 2.5]) + assert [event["event"] for event in events] == [ + "embedding.request.retry", + "embedding.request.retry", + "embedding.request.success", + ] + assert [event["attempt"] for event in events] == [1, 2, 3] + assert events[0]["status_code"] == 503 + assert events[0]["error_code"] == "UNAVAILABLE" + + @pytest.mark.asyncio + async def test_503_exhaustion_logs_failure( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + events = _capture_events(monkeypatch) + delays: list[float] = [] + + async def record_sleep(delay_seconds: float) -> None: + delays.append(delay_seconds) + + client = build_embedding_client( + fail_embed_times=3, + max_retries=2, + sleep=record_sleep, + jitter=lambda: 0.0, + ) + + with pytest.raises(ExternalAIAdapterError) as error: + await client.embed_texts(["alpha"], trace_id="trace-exhausted") + + assert error.value.code == "UNAVAILABLE" + assert error.value.attempt_count == 3 + assert delays == pytest.approx([1.0, 2.0]) + assert [event["event"] for event in events] == [ + "embedding.request.retry", + "embedding.request.retry", + "embedding.request.failed", + ] + + @pytest.mark.asyncio + async def test_timeout_stops_without_sending_second_request( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + events = _capture_events(monkeypatch) + requests = 0 + delays: list[float] = [] + + def timeout_handler(request: httpx.Request) -> httpx.Response: + nonlocal requests + requests += 1 + raise httpx.ReadTimeout("slow embedding", request=request) + + async def record_sleep(delay_seconds: float) -> None: + delays.append(delay_seconds) + + client = _client_with_handler(timeout_handler, sleep=record_sleep) + + with pytest.raises(ExternalAIAdapterError) as error: + await client.embed_texts(["alpha"], trace_id="trace-timeout") + + assert requests == 1 + assert delays == [] + assert error.value.code == "TIMEOUT" + assert events[0]["event"] == "embedding.request.timeout" + + @pytest.mark.asyncio + async def test_non_503_http_error_is_not_retried( + self, + monkeypatch: pytest.MonkeyPatch, + ) -> None: + events = _capture_events(monkeypatch) + requests = 0 + delays: list[float] = [] + + def error_handler(request: httpx.Request) -> httpx.Response: + nonlocal requests + requests += 1 + return httpx.Response(500, request=request) + + async def record_sleep(delay_seconds: float) -> None: + delays.append(delay_seconds) + + client = _client_with_handler(error_handler, sleep=record_sleep) + + with pytest.raises(httpx.HTTPStatusError): + await client.embed_texts(["alpha"], trace_id="trace-http-error") + + assert requests == 1 + assert delays == [] + assert len(events) == 1 + assert events[0]["event"] == "embedding.request.failed" + assert events[0]["status_code"] == 500 + + +class TestEmbeddingValidation: + @pytest.mark.asyncio + async def test_rejects_empty_input(self) -> None: + client = build_embedding_client() + + with pytest.raises(ExternalAIAdapterError): + await client.embed_texts([], trace_id="trace-3") + + +class TestEmbeddingHealth: + @pytest.mark.asyncio + async def test_reads_ready_model_versions(self) -> None: + client = build_embedding_client(model_version="v001") + + versions = await client.get_ready_model_versions(trace_id="trace-4") + + assert versions == ["v001"] + + @pytest.mark.asyncio + async def test_uses_configured_ready_version(self) -> None: + client = build_embedding_client( + model_version="v002", + ready_model_versions=["v001", "v002"], + embedding_model_version="v002", + ) + + version = await client.get_model_version(trace_id="trace-5") + + assert version == "v002" - with pytest.raises(ExternalAIAdapterError, match="ready_model_versions"): - await client.get_ready_model_versions(trace_id="trace-6") + @pytest.mark.asyncio + async def test_rejects_health_without_ready_model_versions(self) -> None: + client = build_embedding_client(health_payload={"status": "ok", "model_version": "v001"}) + with pytest.raises(ExternalAIAdapterError, match="ready_model_versions"): + await client.get_ready_model_versions(trace_id="trace-6") From 02a49f781389547a3712963d780a30282109af59 Mon Sep 17 00:00:00 2001 From: baekyutae Date: Fri, 17 Jul 2026 12:24:05 +0900 Subject: [PATCH 6/6] =?UTF-8?q?docs:=20=EC=84=A4=EA=B3=84=20prompt?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ...44\200 \353\252\205\354\204\270 \354\232\224\354\262\255.md" | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git "a/docs/prompts/0\353\213\250\352\263\204 \354\231\204\354\204\261\352\270\260\354\244\200 \353\252\205\354\204\270 \354\232\224\354\262\255.md" "b/docs/prompts/0\353\213\250\352\263\204 \354\231\204\354\204\261\352\270\260\354\244\200 \353\252\205\354\204\270 \354\232\224\354\262\255.md" index b35a557..d4811ef 100644 --- "a/docs/prompts/0\353\213\250\352\263\204 \354\231\204\354\204\261\352\270\260\354\244\200 \353\252\205\354\204\270 \354\232\224\354\262\255.md" +++ "b/docs/prompts/0\353\213\250\352\263\204 \354\231\204\354\204\261\352\270\260\354\244\200 \353\252\205\354\204\270 \354\232\224\354\262\255.md" @@ -1,6 +1,6 @@ 0단계 프롬프트 — 완성 기준 명세 (코드·설계 전) - [기능명]의 "완성 기준"을 먼저 확정한다. 코드도 설계도 아직 손대지 마라. + [기능명]의 "완성 기준"을 먼저 확정한다. 코드도 설계도 아직 작성하지 마라. ## 목적 무엇이 "되면 끝"인지 기능 단위로 먼저 못박는다. 특히 이 기능이 데이터를