zcbot/tests/test_web_routes_db.py

355 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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 档位过滤
- messages 参数失败不改状态202 前用户消息已持久化且 worker 不重复写入
隔离:测试专属 User 行 + 随机 uid 子树,teardown DB 行与 FS 整树删除;
optimize_prompt(会起真 LLM)不在此测。
"""
from __future__ import annotations
import os
import shutil
import sys
import threading
import unittest
import uuid
from pathlib import Path
from unittest.mock import patch
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()
# 不进入 lifespan手工补消息路由登记后台任务所需的运行态容器。
_app.state.inflight = {}
_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")
self.assertFalse(t["auto_title_pending"]) # 旧创建契约不自动改名
# 列表含它,分页壳完整
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_quick_task_auto_title_marker_and_manual_rename_gate(self):
r = _client.post(
"/v1/tasks",
json={
"name": "新对话",
"working_dir": "快速对话目录",
"auto_title": True,
},
headers=_AUTH,
)
self.assertEqual(r.status_code, 201, r.text)
t = r.json()
self.assertTrue(t["auto_title_pending"])
# 人工改名拥有最高优先级:立刻清 pending异步标题条件 UPDATE 将失配。
d = _client.patch(
f"/v1/tasks/{t['task_id']}",
json={"name": "用户指定标题"},
headers=_AUTH,
).json()
self.assertEqual(d["name"], "用户指定标题")
self.assertFalse(d["auto_title_pending"])
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 MessageRunDurabilityTests(unittest.TestCase):
def _mk_task(self, name: str) -> str:
r = _client.post("/v1/tasks", json={"name": name}, headers=_AUTH)
self.assertEqual(r.status_code, 201, r.text)
return r.json()["task_id"]
def test_validation_failure_keeps_task_idle_and_writes_no_message(self):
from fastapi import HTTPException
from sqlalchemy import select
tid = self._mk_task("非法媒体参数")
with patch(
"web.routers.messages.resolve_image_model",
side_effect=HTTPException(400, "invalid image model"),
):
r = _client.post(
f"/v1/tasks/{tid}/messages",
json={"content": "不会落库", "image_model": "invalid"},
headers=_AUTH,
)
self.assertEqual(r.status_code, 400, r.text)
with session_scope() as s:
task = s.execute(
select(Task.run_status).where(Task.task_id == uuid.UUID(tid))
).scalar_one()
messages = s.execute(
select(Message).where(Message.task_id == uuid.UUID(tid))
).scalars().all()
self.assertEqual(task, "idle")
self.assertEqual(messages, [])
def test_message_is_committed_before_worker_and_not_duplicated(self):
from sqlalchemy import select
tid = self._mk_task("消息持久化")
seen = {}
worker_done = threading.Event()
def fake_worker(task_id, user_id, user_message, *args, **kwargs):
with session_scope() as s:
payloads = s.execute(
select(Message.payload)
.where(Message.task_id == task_id)
.order_by(Message.idx)
).scalars().all()
seen["payloads"] = payloads
seen["persisted"] = kwargs.get("user_message_persisted")
worker_done.set()
with (
patch("web.routers.messages.resolve_image_model", return_value=""),
patch("web.routers.messages.resolve_video_model", return_value=""),
patch("web.run_lifecycle.run_agent_bg", side_effect=fake_worker),
):
r = _client.post(
f"/v1/tasks/{tid}/messages",
json={"content": "必须先落库"},
headers=_AUTH,
)
self.assertTrue(worker_done.wait(5), "后台 worker 未启动")
self.assertEqual(r.status_code, 202, r.text)
self.assertTrue(seen["persisted"])
self.assertEqual(
seen["payloads"],
[{"role": "user", "content": "必须先落库"}],
)
with session_scope() as s:
payloads = s.execute(
select(Message.payload)
.where(Message.task_id == uuid.UUID(tid))
.order_by(Message.idx)
).scalars().all()
self.assertEqual(payloads, [{"role": "user", "content": "必须先落库"}])
_set_run_status(tid, "idle")
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()