119 lines
3.4 KiB
Python
119 lines
3.4 KiB
Python
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()
|