248 lines
11 KiB
Python
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()
|