diff --git a/app/main.py b/app/main.py index 818663b..969f21f 100644 --- a/app/main.py +++ b/app/main.py @@ -19,6 +19,7 @@ from app.services.task_service import ( has_running_task, list_tasks, recover_running_tasks, + recover_stale_running_tasks, ) from app.templating import templates @@ -83,6 +84,7 @@ def create_task_api( request: CreateTaskRequest, session: Session = Depends(get_db_session), ) -> CreateTaskResponse: + recover_stale_running_tasks(session) if has_running_task(session): raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=RUNNING_TASK_MESSAGE) diff --git a/app/platforms/base.py b/app/platforms/base.py index e53d702..73d8528 100644 --- a/app/platforms/base.py +++ b/app/platforms/base.py @@ -87,6 +87,12 @@ class TikHubClient: ) time.sleep(2**attempt) continue + if response.status_code == 401: + raise PlatformAPIError( + "TikHub 鉴权失败,请检查 TIKHUB_API_KEY 是否有效", + error_type="auth_error", + status_code=response.status_code, + ) if response.is_error: raise PlatformAPIError( f"External API returned HTTP {response.status_code}", diff --git a/app/services/ai_service.py b/app/services/ai_service.py index f48bd97..cac00d3 100644 --- a/app/services/ai_service.py +++ b/app/services/ai_service.py @@ -1,5 +1,6 @@ import json import time +from concurrent.futures import ThreadPoolExecutor, TimeoutError from dataclasses import dataclass from typing import Literal @@ -121,12 +122,14 @@ def analyze_comments_with_retry( *, requester, max_retries: int = 3, + timeout_seconds: float | None = None, ) -> list[AIAnalysisResult]: expected_ids = {str(comment["comment_id"]) for comment in comments} prompt = build_comment_prompt(comments) for attempt in range(max_retries): try: - return parse_ai_comment_response(requester(prompt), expected_comment_ids=expected_ids) + response_text = _request_with_timeout(requester, prompt, timeout_seconds=timeout_seconds) + return parse_ai_comment_response(response_text, expected_comment_ids=expected_ids) except Exception: if attempt >= max_retries - 1: break @@ -143,6 +146,21 @@ def analyze_comments_with_retry( ] +def _request_with_timeout(requester, prompt: str, *, timeout_seconds: float | None) -> str: + if timeout_seconds is None or timeout_seconds <= 0: + return requester(prompt) + + executor = ThreadPoolExecutor(max_workers=1) + future = executor.submit(requester, prompt) + try: + return future.result(timeout=timeout_seconds) + except TimeoutError as exc: + future.cancel() + raise TimeoutError("AI request timed out") from exc + finally: + executor.shutdown(wait=False, cancel_futures=True) + + class OpenAICompatibleAIClient: def __init__( self, diff --git a/app/services/task_service.py b/app/services/task_service.py index b7d42c0..b7c0a1c 100644 --- a/app/services/task_service.py +++ b/app/services/task_service.py @@ -21,6 +21,7 @@ from app.services.report_service import generate_hotspot_report, generate_item_r RUNNING_TASK_MESSAGE = "当前有正在运行的任务,请稍后再试" RESTART_ERROR_MESSAGE = "系统重启,任务被中断" STALE_PROGRESS_THRESHOLD_SECONDS = 10 * 60 +STALE_PROGRESS_ERROR_TYPE = "stale_progress_timeout" task_executor = ThreadPoolExecutor(max_workers=1) @@ -44,6 +45,30 @@ def has_running_task(session: Session) -> bool: return session.scalar(select(Task.id).where(Task.status == "running").limit(1)) is not None +def recover_stale_running_tasks(session: Session) -> int: + now = ensure_utc_datetime(utc_now()) + tasks = list(session.scalars(select(Task).where(Task.status == "running"))) + recovered = 0 + for task in tasks: + last_progress_at = ensure_utc_datetime(task.last_progress_at or task.started_at or task.created_at) + if last_progress_at is None or now is None: + continue + seconds_since_progress = max(0, int((now - last_progress_at).total_seconds())) + if seconds_since_progress < STALE_PROGRESS_THRESHOLD_SECONDS: + continue + _mark_task_failed( + task, + "system", + STALE_PROGRESS_ERROR_TYPE, + f"任务超过 {STALE_PROGRESS_THRESHOLD_SECONDS // 60} 分钟没有进度更新,已自动标记失败", + ) + task.finished_at = task.finished_at or utc_now() + recovered += 1 + if recovered: + session.commit() + return recovered + + def recover_running_tasks(session: Session) -> int: tasks = list(session.scalars(select(Task).where(Task.status == "running"))) for task in tasks: @@ -256,6 +281,8 @@ def run_task(task_id: str, *, session_factory: Callable[[], Session] | sessionma task = session.get(Task, task_id) if task is None: return + task.started_at = task.started_at or utc_now() + session.commit() try: platform = build_platform(task.platform) ai_requester, report_summary_provider = build_ai_dependencies() @@ -345,6 +372,7 @@ def run_task(task_id: str, *, session_factory: Callable[[], Session] | sessionma [{"comment_id": comment.id, "content": comment.content} for comment in item_comments], requester=ai_requester, max_retries=get_settings().ai_max_retries, + timeout_seconds=get_settings().ai_timeout_seconds, ) results_by_comment_id = {result.comment_id: result for result in ai_results} for comment in item_comments: @@ -390,9 +418,11 @@ def run_task(task_id: str, *, session_factory: Callable[[], Session] | sessionma update_task_stage(task, "success") else: _mark_task_failed(task, task.error_stage or "crawl_items", task.error_type or "no_successful_items", task.error_message or "没有任何内容条目成功") + task.finished_at = task.finished_at or utc_now() session.commit() except Exception as exc: _mark_task_failed(task, "system", "unexpected_error", str(exc)) + task.finished_at = task.finished_at or utc_now() session.commit() finally: if close_session: diff --git a/tests/integration/test_task_creation.py b/tests/integration/test_task_creation.py index 5fa71df..3fb55e3 100644 --- a/tests/integration/test_task_creation.py +++ b/tests/integration/test_task_creation.py @@ -1,5 +1,8 @@ +from datetime import UTC, datetime, timedelta + from fastapi.testclient import TestClient from sqlalchemy import create_engine +from sqlalchemy.orm import Session from sqlalchemy.pool import StaticPool from app.db import Base, get_db_session @@ -92,6 +95,47 @@ def test_create_task_returns_400_when_running_task_exists(): app.dependency_overrides.clear() +def test_create_task_recovers_stale_running_task_before_creating_new_one(monkeypatch): + client, engine = make_test_client() + now = datetime(2026, 7, 3, 12, 0, tzinfo=UTC) + monkeypatch.setattr("app.services.task_service.utc_now", lambda: now) + monkeypatch.setattr("app.services.task_service.task_executor.submit", lambda *_args, **_kwargs: None) + try: + with Session(engine) as session: + session.add( + Task( + id="stale-running", + platform="douyin", + status="running", + current_stage="ai_analysis", + last_progress_at=now - timedelta(minutes=11), + ) + ) + session.commit() + + response = client.post( + "/api/tasks", + json={ + "platform": "douyin", + "hotspot_limit": 5, + "item_limit_per_hotspot": 5, + "comment_limit_per_item": 50, + }, + ) + + with Session(engine) as session: + stale_task = session.get(Task, "stale-running") + + assert response.status_code == 201 + assert stale_task.status == "failed" + assert stale_task.error_stage == "system" + assert stale_task.error_type == "stale_progress_timeout" + assert "超过 10 分钟没有进度更新" in stale_task.error_message + finally: + engine.dispose() + app.dependency_overrides.clear() + + def test_task_list_and_detail_return_created_tasks(): client, engine = make_test_client() try: diff --git a/tests/integration/test_task_flow_xiaohongshu.py b/tests/integration/test_task_flow_xiaohongshu.py index 239e0c8..cb9398d 100644 --- a/tests/integration/test_task_flow_xiaohongshu.py +++ b/tests/integration/test_task_flow_xiaohongshu.py @@ -1,4 +1,6 @@ import json +import time +from types import SimpleNamespace from sqlalchemy import select from sqlalchemy.orm import Session @@ -147,3 +149,43 @@ def test_xiaohongshu_task_flow_keeps_item_success_when_ai_request_raises(monkeyp assert persisted.analysis_success_rate == 0.0 assert comment.ai_analysis_status == "failed" assert comment.reason == "ai_parse_failed" + + +def test_xiaohongshu_task_flow_times_out_stalled_ai_request(monkeypatch): + with make_test_client() as (_client, engine): + monkeypatch.setattr("app.services.task_service.build_platform", lambda _platform: FakeXhsPlatform()) + monkeypatch.setattr( + "app.services.task_service.get_settings", + lambda: SimpleNamespace(ai_max_retries=1, ai_timeout_seconds=0.01), + ) + + def stalled_requester(prompt): + time.sleep(0.2) + comments = json.loads(prompt[prompt.index("[") :]) + return json.dumps( + [{"comment_id": comments[0]["comment_id"], "sentiment": "positive", "labels": ["超时后不应采用"], "reason": "late"}], + ensure_ascii=False, + ) + + monkeypatch.setattr("app.services.task_service.build_ai_dependencies", lambda: (stalled_requester, None)) + + with Session(engine) as session: + task = create_task( + session, + CreateTaskRequest(platform="xiaohongshu", hotspot_limit=1, item_limit_per_hotspot=1, comment_limit_per_item=10), + submit_background=False, + ) + task_id = task.id + + with Session(engine) as session: + run_task(task_id, session_factory=lambda: session) + + persisted = session.get(Task, task_id) + comment = session.scalar(select(Comment).where(Comment.task_id == task_id)) + + assert persisted.status == "success" + assert persisted.finished_at is not None + assert persisted.analysis_status == "insufficient" + assert comment.ai_analysis_status == "failed" + assert comment.sentiment == "unknown" + assert comment.reason == "ai_parse_failed" diff --git a/tests/unit/test_comment_pagination.py b/tests/unit/test_comment_pagination.py index 2bf6179..f2df7b2 100644 --- a/tests/unit/test_comment_pagination.py +++ b/tests/unit/test_comment_pagination.py @@ -51,3 +51,17 @@ def test_tikhub_client_raises_structured_error_after_retries(monkeypatch): assert exc_info.value.status_code == 429 assert "secret-token" not in str(exc_info.value) assert sleeps == [1, 2, 4] + + +def test_tikhub_client_reports_401_as_auth_error_without_leaking_token(): + transport = SequenceTransport([httpx.Response(401, json={"message": "Unauthorized"})]) + http_client = httpx.Client(transport=httpx.MockTransport(transport)) + client = TikHubClient(base_url="https://api.test", api_key="secret-token", http_client=http_client) + + with pytest.raises(PlatformAPIError) as exc_info: + client.get("/demo") + + assert exc_info.value.error_type == "auth_error" + assert exc_info.value.status_code == 401 + assert str(exc_info.value) == "TikHub 鉴权失败,请检查 TIKHUB_API_KEY 是否有效" + assert "secret-token" not in str(exc_info.value)