from fastapi.testclient import TestClient from sqlalchemy import create_engine 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_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()