diff --git a/tests/test_asr/test_bcut_asr.py b/tests/test_asr/test_bcut_asr.py index e59ff526..72ce111e 100644 --- a/tests/test_asr/test_bcut_asr.py +++ b/tests/test_asr/test_bcut_asr.py @@ -3,12 +3,82 @@ from pathlib import Path import pytest +import requests from tests.test_asr.conftest import assert_asr_result_valid from videocaptioner.core.asr import BcutASR from videocaptioner.core.asr.asr_data import ASRData +class FakeResponse: + def __init__(self, status_code: int, headers: dict[str, str] | None = None): + self.status_code = status_code + self.headers = headers or {} + + def raise_for_status(self) -> None: + if self.status_code >= 400: + raise requests.HTTPError(f"{self.status_code} error", response=self) + + +class FakeSession: + def __init__(self, responses): + self.responses = list(responses) + self.calls = 0 + + def request(self, *args, **kwargs): + self.calls += 1 + response = self.responses.pop(0) + if isinstance(response, BaseException): + raise response + return response + + +def make_bcut_without_init(session: FakeSession) -> BcutASR: + asr = object.__new__(BcutASR) + asr.session = session + return asr + + +def test_request_with_retry_recovers_from_read_timeout(monkeypatch) -> None: + monkeypatch.setattr("videocaptioner.core.asr.bcut.time.sleep", lambda _: None) + monkeypatch.setattr("videocaptioner.core.asr.bcut.random.uniform", lambda *_: 0) + session = FakeSession( + [ + requests.exceptions.ReadTimeout("read timed out"), + requests.exceptions.ReadTimeout("read timed out"), + FakeResponse(200), + ] + ) + asr = make_bcut_without_init(session) + + response = asr._request_with_retry("GET", "https://example.test", max_attempts=3) + + assert response.status_code == 200 + assert session.calls == 3 + + +def test_request_with_retry_recovers_from_412(monkeypatch) -> None: + monkeypatch.setattr("videocaptioner.core.asr.bcut.time.sleep", lambda _: None) + monkeypatch.setattr("videocaptioner.core.asr.bcut.random.uniform", lambda *_: 0) + session = FakeSession([FakeResponse(412), FakeResponse(200)]) + asr = make_bcut_without_init(session) + + response = asr._request_with_retry("GET", "https://example.test", max_attempts=2) + + assert response.status_code == 200 + assert session.calls == 2 + + +def test_request_with_retry_does_not_retry_non_retryable_status() -> None: + session = FakeSession([FakeResponse(401)]) + asr = make_bcut_without_init(session) + + with pytest.raises(requests.HTTPError): + asr._request_with_retry("GET", "https://example.test") + + assert session.calls == 1 + + @pytest.mark.integration @pytest.mark.slow class TestBcutASR: diff --git a/videocaptioner/core/asr/bcut.py b/videocaptioner/core/asr/bcut.py index dd7e8c5d..b5b989df 100644 --- a/videocaptioner/core/asr/bcut.py +++ b/videocaptioner/core/asr/bcut.py @@ -1,4 +1,5 @@ import json +import random import time from typing import Any, Callable, List, Optional, Union @@ -16,6 +17,8 @@ API_CREATE_TASK = API_BASE_URL + "/task" API_QUERY_RESULT = API_BASE_URL + "/task/result" +RETRYABLE_STATUS_CODES = {408, 409, 412, 425, 429, 500, 502, 503, 504} + class BcutASR(BaseASR): """Bilibili Bcut ASR API implementation. @@ -51,6 +54,59 @@ def __init__( self.need_word_time_stamp = need_word_time_stamp + def _request_with_retry( + self, + method: str, + url: str, + *, + max_attempts: int = 8, + timeout: tuple[int, int] = (15, 300), + **kwargs: Any, + ) -> requests.Response: + """Request Bcut/BOSS endpoints with retry for transient failures.""" + last_exc: Optional[BaseException] = None + + for attempt in range(1, max_attempts + 1): + try: + resp = self.session.request(method, url, timeout=timeout, **kwargs) + if resp.status_code not in RETRYABLE_STATUS_CODES: + resp.raise_for_status() + return resp + + last_exc = requests.HTTPError( + f"{resp.status_code} retryable response", response=resp + ) + except requests.RequestException as exc: + last_exc = exc + response = getattr(exc, "response", None) + status_code = getattr(response, "status_code", None) + if status_code is not None and status_code not in RETRYABLE_STATUS_CODES: + raise + + if attempt >= max_attempts: + break + + response = getattr(last_exc, "response", None) + retry_after = response.headers.get("Retry-After") if response else None + try: + wait_seconds = float(retry_after) if retry_after else None + except ValueError: + wait_seconds = None + + if wait_seconds is None: + wait_seconds = min(60.0, 2 ** (attempt - 1)) + random.uniform(0, 0.5) + + print( + f"Bcut request failed ({last_exc}); " + f"retry {attempt}/{max_attempts - 1} after {wait_seconds:.1f}s" + ) + time.sleep(wait_seconds) + + assert last_exc is not None + if isinstance(last_exc, requests.HTTPError) and last_exc.response is not None: + last_exc.response.raise_for_status() + raise last_exc + def upload(self) -> None: """Request upload authorization and upload audio file.""" if not self.file_binary: @@ -65,8 +121,9 @@ def upload(self) -> None: } ) - resp = requests.post(API_REQ_UPLOAD, data=payload, headers=self.headers) - resp.raise_for_status() + resp = self._request_with_retry( + "POST", API_REQ_UPLOAD, data=payload, headers=self.headers + ) resp = resp.json() resp_data = resp["data"] @@ -93,12 +150,12 @@ def __upload_part(self) -> None: for clip in range(self.__clips): start_range = clip * self.__per_size end_range = (clip + 1) * self.__per_size - resp = requests.put( + resp = self._request_with_retry( + "PUT", self.__upload_urls[clip], data=self.file_binary[start_range:end_range], headers=self.headers, ) - resp.raise_for_status() etag = resp.headers.get("Etag") if etag is not None: self.__etags.append(etag) @@ -114,31 +171,33 @@ def __commit_upload(self) -> None: "model_id": "8", } ) - resp = requests.post(API_COMMIT_UPLOAD, data=data, headers=self.headers) - resp.raise_for_status() + resp = self._request_with_retry( + "POST", API_COMMIT_UPLOAD, data=data, headers=self.headers + ) resp = resp.json() self.__download_url = resp["data"]["download_url"] def create_task(self) -> str: """Create ASR task.""" - resp = requests.post( + resp = self._request_with_retry( + "POST", API_CREATE_TASK, json={"resource": self.__download_url, "model_id": "8"}, headers=self.headers, ) - resp.raise_for_status() resp = resp.json() self.task_id = resp["data"]["task_id"] return self.task_id or "" def result(self, task_id: Optional[str] = None): """Query ASR result.""" - resp = requests.get( + resp = self._request_with_retry( + "GET", API_QUERY_RESULT, params={"model_id": 7, "task_id": task_id or self.task_id}, headers=self.headers, + timeout=(15, 120), ) - resp.raise_for_status() resp = resp.json() return resp["data"]