zcbot/tests/test_web_routes_db.py

248 lines
11 KiB
Python

"""web /v1 路由测试(DB 面)—— 补 test_web_routes_nodb 留待的 DB-aware 分支。
仅当显式设 `ZCBOT_TEST_DB_URL` 才跑(同 test_usage_report 门控纪律,绝不回退
.env —— 它可能指生产);一次性测试库:docker run postgres + main.py db upgrade,
见 RUN.md「测试库」段。
覆盖:
- tasks CRUD 全链(建/列/详/改/软删/恢复/同 wd 共享)+ folders
- files 顶层目录 DB-aware:rename 级联改 tasks.working_dir、running→409、
move 被引用→409、递归删被引用→409、软删后放行
- upload(磁盘配额 gate 路径)+ 根目录列表(system_wd_names DB 查询)
- clear / cancel 的状态闸;schedules 404 路径;/v1/models 档位过滤
隔离:测试专属 User 行 + 随机 uid 子树,teardown DB 行与 FS 整树删除;
POST messages / optimize_prompt(会起真 LLM)不在此测。
"""
from __future__ import annotations
import os
import shutil
import sys
import unittest
import uuid
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
os.environ.setdefault("PLATFORM_KEY", "test-platform-key-db")
os.environ.setdefault("JWT_SECRET", "test-jwt-secret-db")
def _test_db_ready() -> bool:
"""只认显式 ZCBOT_TEST_DB_URL(教训见 test_usage_report 同名函数)。"""
url = os.environ.get("ZCBOT_TEST_DB_URL", "").strip()
if not url:
return False
if "connect_timeout" not in url:
# 测试库挂了要快速降级 skip,不能让整个 discovery 挂死在 TCP 建连上
url += ("&" if "?" in url else "?") + "connect_timeout=5"
os.environ["ZCBOT_DB_URL"] = url
return True
try:
if not _test_db_ready():
raise RuntimeError("ZCBOT_TEST_DB_URL 未设")
from core.storage import session_scope
from core.storage.models import Message, Task, UsageEvent, User
with session_scope() as _s:
_s.execute(__import__("sqlalchemy").select(1))
_DB_OK = True
except Exception:
_DB_OK = False
if _DB_OK:
from starlette.testclient import TestClient
from web.app import create_app
from web.auth import AuthConfig, mint_token
_app = create_app()
_client = TestClient(_app) # 不进 with:不跑 lifespan
_UID = uuid.uuid4()
_TOKEN, _ = mint_token(AuthConfig.from_env(), _UID)
_AUTH = {"Authorization": f"Bearer {_TOKEN}"}
def _user_root() -> Path:
from core.agent_builder import resolve_workspace, user_root
return user_root(resolve_workspace(None), _UID)
def setUpModule() -> None:
if not _DB_OK:
raise unittest.SkipTest("ZCBOT_TEST_DB_URL 未设或测试库不可达(绝不用 .env 的库)")
with session_scope() as s:
s.add(User(user_id=_UID, email=f"test-routes-db-{_UID.hex[:8]}@invalid.local"))
def tearDownModule() -> None:
if not _DB_OK:
return
from sqlalchemy import delete, select
with session_scope() as s:
tids = s.execute(select(Task.task_id).where(Task.user_id == _UID)).scalars().all()
if tids:
s.execute(delete(Message).where(Message.task_id.in_(tids)))
s.execute(delete(UsageEvent).where(UsageEvent.user_id == _UID))
s.execute(delete(Task).where(Task.user_id == _UID))
s.execute(delete(User).where(User.user_id == _UID))
d = _user_root()
if d.is_dir():
shutil.rmtree(d, ignore_errors=True)
def _set_run_status(task_id: str, status: str) -> None:
from sqlalchemy import update
with session_scope() as s:
s.execute(update(Task).where(Task.task_id == uuid.UUID(task_id)).values(run_status=status))
class TasksCrudTests(unittest.TestCase):
def test_create_list_get_patch_softdelete_restore(self):
r = _client.post("/v1/tasks", json={"name": "路由测试任务"}, headers=_AUTH)
self.assertEqual(r.status_code, 201, r.text)
t = r.json()
tid = t["task_id"]
self.assertEqual(t["name"], "路由测试任务")
self.assertTrue(t["working_dir"].endswith("/路由测试任务"))
self.assertTrue(t["model_profile"]) # 默认模型已填
self.assertEqual(t["run_status"], "idle")
# 列表含它,分页壳完整
r = _client.get("/v1/tasks", headers=_AUTH)
body = r.json()
self.assertGreaterEqual(body["count"], 1)
self.assertIn(tid, {x["task_id"] for x in body["results"]})
# 详情带上下文压力字段
d = _client.get(f"/v1/tasks/{tid}", headers=_AUTH).json()
self.assertIn("context_window_chars", d)
self.assertIn("context_folds", d)
# PATCH:非法 status 400;description 生效
self.assertEqual(
_client.patch(f"/v1/tasks/{tid}", json={"status": "active"}, headers=_AUTH).status_code, 400)
d = _client.patch(f"/v1/tasks/{tid}", json={"description": "备注"}, headers=_AUTH).json()
self.assertEqual(d["description"], "备注")
# 软删 → 列表消失;恢复 → 回来
self.assertEqual(_client.delete(f"/v1/tasks/{tid}", headers=_AUTH).status_code, 204)
self.assertNotIn(tid, {x["task_id"] for x in _client.get("/v1/tasks", headers=_AUTH).json()["results"]})
self.assertEqual(_client.post(f"/v1/tasks/{tid}/restore", headers=_AUTH).status_code, 200)
self.assertIn(tid, {x["task_id"] for x in _client.get("/v1/tasks", headers=_AUTH).json()["results"]})
# 非法 id → 404
self.assertEqual(_client.get("/v1/tasks/not-a-uuid", headers=_AUTH).status_code, 404)
def test_shared_working_dir_and_folders(self):
for name in ("共享目录甲", "共享目录乙"):
r = _client.post("/v1/tasks", json={"name": name, "working_dir": "共享项目"}, headers=_AUTH)
self.assertEqual(r.status_code, 201, r.text)
folders = _client.get("/v1/folders", headers=_AUTH).json()["folders"]
hit = [f for f in folders if f["name"] == "共享项目"]
self.assertEqual(len(hit), 1)
self.assertGreaterEqual(hit[0]["n_tasks"], 2)
def test_clear_and_cancel_gates(self):
tid = _client.post("/v1/tasks", json={"name": "状态闸任务"}, headers=_AUTH).json()["task_id"]
# idle 时 cancel → 409
self.assertEqual(_client.post(f"/v1/tasks/{tid}/cancel", headers=_AUTH).status_code, 409)
# running 时 clear → 409;回 idle 后 clear → 200 且归零
_set_run_status(tid, "running")
self.assertEqual(_client.post(f"/v1/tasks/{tid}/clear", headers=_AUTH).status_code, 409)
_set_run_status(tid, "idle")
d = _client.post(f"/v1/tasks/{tid}/clear", headers=_AUTH).json()
self.assertEqual((d["n_messages"], d["tokens"]), (0, 0))
# 空任务 messages / outline
m = _client.get(f"/v1/tasks/{tid}/messages", headers=_AUTH).json()
self.assertEqual(m["messages"], [])
self.assertFalse(m["has_more"])
self.assertEqual(_client.get(f"/v1/tasks/{tid}/outline", headers=_AUTH).json()["items"], [])
class FilesDbAwareTests(unittest.TestCase):
"""顶层目录 = task.working_dir 的 DB-aware 分支(§7.4 唯一 mutation 入口)。"""
def _mk_task(self, name: str, wd: str) -> str:
r = _client.post("/v1/tasks", json={"name": name, "working_dir": wd}, headers=_AUTH)
self.assertEqual(r.status_code, 201, r.text)
return r.json()["task_id"]
def test_toplevel_rename_cascades_db(self):
tid = self._mk_task("改名任务", "改名前目录")
r = _client.post("/v1/files/rename",
json={"path": "改名前目录", "new_name": "改名后目录"}, headers=_AUTH)
self.assertEqual(r.status_code, 200, r.text)
self.assertEqual(r.json()["tasks_updated"], 1)
# DB 级联生效
d = _client.get(f"/v1/tasks/{tid}", headers=_AUTH).json()
self.assertTrue(d["working_dir"].endswith("/改名后目录"))
self.assertTrue((_user_root() / "改名后目录").is_dir())
def test_toplevel_rename_blocked_while_running(self):
tid = self._mk_task("跑动中任务", "跑动中目录")
_set_run_status(tid, "running")
try:
r = _client.post("/v1/files/rename",
json={"path": "跑动中目录", "new_name": "别名"}, headers=_AUTH)
self.assertEqual(r.status_code, 409)
finally:
_set_run_status(tid, "idle")
def test_toplevel_move_and_recursive_delete_blocked(self):
tid = self._mk_task("被引用任务", "被引用目录")
(_user_root() / "搬运目标").mkdir(exist_ok=True)
# 被 task 引用的顶层目录:move → 409;递归删 → 409
r = _client.post("/v1/files/move",
json={"paths": ["被引用目录"], "dest_dir": "搬运目标"}, headers=_AUTH)
self.assertEqual(r.status_code, 409)
r = _client.post("/v1/files/delete",
json={"path": "被引用目录", "recursive": True}, headers=_AUTH)
self.assertEqual(r.status_code, 409)
# 软删 task 后不再算引用 → 递归删放行
self.assertEqual(_client.delete(f"/v1/tasks/{tid}", headers=_AUTH).status_code, 204)
r = _client.post("/v1/files/delete",
json={"path": "被引用目录", "recursive": True}, headers=_AUTH)
self.assertEqual(r.status_code, 200, r.text)
def test_upload_and_root_listing(self):
# upload 走磁盘配额 gate(无扫描记录 → 放行);根目录列表走 system_wd_names DB 查询
r = _client.post("/v1/files/upload", data={"path": "上传目录"},
files={"files": ("hello.txt", b"hi", "text/plain")}, headers=_AUTH)
self.assertEqual(r.status_code, 200, r.text)
self.assertEqual(r.json()["count"], 1)
names = {e["name"] for e in _client.get("/v1/files", headers=_AUTH).json()["entries"]}
self.assertIn("上传目录", names)
# 非法文件名 → 400
r = _client.post("/v1/files/upload", data={"path": "上传目录"},
files={"files": ("../escape.txt", b"x", "text/plain")}, headers=_AUTH)
self.assertEqual(r.status_code, 400)
class MiscDbRoutesTests(unittest.TestCase):
def test_models_list_nonempty(self):
r = _client.get("/v1/models", headers=_AUTH)
self.assertEqual(r.status_code, 200)
models = r.json()["models"]
self.assertTrue(models)
self.assertTrue(any(m["is_default"] for m in models))
def test_me_profile(self):
d = _client.get("/v1/me", headers=_AUTH).json()
self.assertEqual(d["user_id"], str(_UID))
self.assertEqual(d["role"], "user")
def test_schedules_empty_and_404(self):
self.assertEqual(_client.get("/v1/schedules", headers=_AUTH).json(), {"results": []})
self.assertEqual(
_client.patch(f"/v1/schedules/{uuid.uuid4()}", json={"enabled": False}, headers=_AUTH).status_code,
404,
)
self.assertEqual(_client.delete(f"/v1/schedules/{uuid.uuid4()}", headers=_AUTH).status_code, 404)
if __name__ == "__main__":
unittest.main()