163 lines
4.9 KiB
Python
163 lines
4.9 KiB
Python
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
|
|
from app.main import app
|
|
from app.models import Task
|
|
|
|
|
|
def make_test_client():
|
|
engine = create_engine(
|
|
"sqlite:///:memory:",
|
|
connect_args={"check_same_thread": False},
|
|
poolclass=StaticPool,
|
|
)
|
|
Base.metadata.create_all(engine)
|
|
|
|
def override_db_session():
|
|
from sqlalchemy.orm import Session
|
|
|
|
with Session(engine) as session:
|
|
yield session
|
|
|
|
app.dependency_overrides[get_db_session] = override_db_session
|
|
return TestClient(app), engine
|
|
|
|
|
|
def test_create_task_returns_task_id_and_running_status():
|
|
client, engine = make_test_client()
|
|
try:
|
|
response = client.post(
|
|
"/api/tasks",
|
|
json={
|
|
"platform": "xiaohongshu",
|
|
"hotspot_limit": 5,
|
|
"item_limit_per_hotspot": 5,
|
|
"comment_limit_per_item": 50,
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 201
|
|
data = response.json()
|
|
assert data["task_id"]
|
|
assert data["status"] == "running"
|
|
finally:
|
|
engine.dispose()
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
def test_create_task_rejects_invalid_platform_and_limits():
|
|
client, engine = make_test_client()
|
|
try:
|
|
response = client.post(
|
|
"/api/tasks",
|
|
json={
|
|
"platform": "weibo",
|
|
"hotspot_limit": 11,
|
|
"item_limit_per_hotspot": 0,
|
|
"comment_limit_per_item": 101,
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 422
|
|
finally:
|
|
engine.dispose()
|
|
app.dependency_overrides.clear()
|
|
|
|
|
|
def test_create_task_returns_400_when_running_task_exists():
|
|
client, engine = make_test_client()
|
|
try:
|
|
from sqlalchemy.orm import Session
|
|
|
|
with Session(engine) as session:
|
|
session.add(Task(platform="douyin", status="running"))
|
|
session.commit()
|
|
|
|
response = client.post(
|
|
"/api/tasks",
|
|
json={
|
|
"platform": "douyin",
|
|
"hotspot_limit": 5,
|
|
"item_limit_per_hotspot": 5,
|
|
"comment_limit_per_item": 50,
|
|
},
|
|
)
|
|
|
|
assert response.status_code == 400
|
|
assert response.json() == {"detail": "当前有正在运行的任务,请稍后再试"}
|
|
finally:
|
|
engine.dispose()
|
|
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:
|
|
created = client.post(
|
|
"/api/tasks",
|
|
json={
|
|
"platform": "douyin",
|
|
"hotspot_limit": 3,
|
|
"item_limit_per_hotspot": 4,
|
|
"comment_limit_per_item": 30,
|
|
},
|
|
).json()
|
|
|
|
list_response = client.get("/api/tasks")
|
|
detail_response = client.get(f"/api/tasks/{created['task_id']}")
|
|
|
|
assert list_response.status_code == 200
|
|
assert list_response.json()[0]["task_id"] == created["task_id"]
|
|
assert detail_response.status_code == 200
|
|
assert detail_response.json()["platform"] == "douyin"
|
|
assert detail_response.json()["hotspot_limit"] == 3
|
|
finally:
|
|
engine.dispose()
|
|
app.dependency_overrides.clear()
|