fix: 修复任务长期运行中卡住

This commit is contained in:
meijiali
2026-07-03 19:18:13 +08:00
parent 0b186fa9f8
commit aec3604de9
5 changed files with 137 additions and 1 deletions
+2
View File
@@ -19,6 +19,7 @@ from app.services.task_service import (
has_running_task, has_running_task,
list_tasks, list_tasks,
recover_running_tasks, recover_running_tasks,
recover_stale_running_tasks,
) )
from app.templating import templates from app.templating import templates
@@ -83,6 +84,7 @@ def create_task_api(
request: CreateTaskRequest, request: CreateTaskRequest,
session: Session = Depends(get_db_session), session: Session = Depends(get_db_session),
) -> CreateTaskResponse: ) -> CreateTaskResponse:
recover_stale_running_tasks(session)
if has_running_task(session): if has_running_task(session):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=RUNNING_TASK_MESSAGE) raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=RUNNING_TASK_MESSAGE)
+19 -1
View File
@@ -1,5 +1,6 @@
import json import json
import time import time
from concurrent.futures import ThreadPoolExecutor, TimeoutError
from dataclasses import dataclass from dataclasses import dataclass
from typing import Literal from typing import Literal
@@ -121,12 +122,14 @@ def analyze_comments_with_retry(
*, *,
requester, requester,
max_retries: int = 3, max_retries: int = 3,
timeout_seconds: float | None = None,
) -> list[AIAnalysisResult]: ) -> list[AIAnalysisResult]:
expected_ids = {str(comment["comment_id"]) for comment in comments} expected_ids = {str(comment["comment_id"]) for comment in comments}
prompt = build_comment_prompt(comments) prompt = build_comment_prompt(comments)
for attempt in range(max_retries): for attempt in range(max_retries):
try: 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: except Exception:
if attempt >= max_retries - 1: if attempt >= max_retries - 1:
break 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: class OpenAICompatibleAIClient:
def __init__( def __init__(
self, self,
+30
View File
@@ -21,6 +21,7 @@ from app.services.report_service import generate_hotspot_report, generate_item_r
RUNNING_TASK_MESSAGE = "当前有正在运行的任务,请稍后再试" RUNNING_TASK_MESSAGE = "当前有正在运行的任务,请稍后再试"
RESTART_ERROR_MESSAGE = "系统重启,任务被中断" RESTART_ERROR_MESSAGE = "系统重启,任务被中断"
STALE_PROGRESS_THRESHOLD_SECONDS = 10 * 60 STALE_PROGRESS_THRESHOLD_SECONDS = 10 * 60
STALE_PROGRESS_ERROR_TYPE = "stale_progress_timeout"
task_executor = ThreadPoolExecutor(max_workers=1) 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 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: def recover_running_tasks(session: Session) -> int:
tasks = list(session.scalars(select(Task).where(Task.status == "running"))) tasks = list(session.scalars(select(Task).where(Task.status == "running")))
for task in tasks: 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) task = session.get(Task, task_id)
if task is None: if task is None:
return return
task.started_at = task.started_at or utc_now()
session.commit()
try: try:
platform = build_platform(task.platform) platform = build_platform(task.platform)
ai_requester, report_summary_provider = build_ai_dependencies() 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], [{"comment_id": comment.id, "content": comment.content} for comment in item_comments],
requester=ai_requester, requester=ai_requester,
max_retries=get_settings().ai_max_retries, 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} results_by_comment_id = {result.comment_id: result for result in ai_results}
for comment in item_comments: 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") update_task_stage(task, "success")
else: else:
_mark_task_failed(task, task.error_stage or "crawl_items", task.error_type or "no_successful_items", task.error_message or "没有任何内容条目成功") _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() session.commit()
except Exception as exc: except Exception as exc:
_mark_task_failed(task, "system", "unexpected_error", str(exc)) _mark_task_failed(task, "system", "unexpected_error", str(exc))
task.finished_at = task.finished_at or utc_now()
session.commit() session.commit()
finally: finally:
if close_session: if close_session:
+44
View File
@@ -1,5 +1,8 @@
from datetime import UTC, datetime, timedelta
from fastapi.testclient import TestClient from fastapi.testclient import TestClient
from sqlalchemy import create_engine from sqlalchemy import create_engine
from sqlalchemy.orm import Session
from sqlalchemy.pool import StaticPool from sqlalchemy.pool import StaticPool
from app.db import Base, get_db_session 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() 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(): def test_task_list_and_detail_return_created_tasks():
client, engine = make_test_client() client, engine = make_test_client()
try: try:
@@ -1,4 +1,6 @@
import json import json
import time
from types import SimpleNamespace
from sqlalchemy import select from sqlalchemy import select
from sqlalchemy.orm import Session 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 persisted.analysis_success_rate == 0.0
assert comment.ai_analysis_status == "failed" assert comment.ai_analysis_status == "failed"
assert comment.reason == "ai_parse_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"