484 lines
20 KiB
Python
484 lines
20 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 档位过滤
|
||
- 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"]) # 旧创建契约不自动改名
|
||
self.assertEqual(t["title_source"], "manual")
|
||
|
||
# 列表含它,分页壳完整
|
||
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={"working_dir": "快速对话目录"},
|
||
headers=_AUTH,
|
||
)
|
||
self.assertEqual(r.status_code, 201, r.text)
|
||
t = r.json()
|
||
self.assertEqual(t["name"], "新对话")
|
||
self.assertTrue(t["auto_title_pending"])
|
||
self.assertEqual(t["title_source"], "auto")
|
||
|
||
# 空 name 与省略 name 同义;但二者都没有 working_dir 时不能确定文件归属。
|
||
blank = _client.post(
|
||
"/v1/tasks",
|
||
json={"name": " ", "working_dir": "快速对话空名称目录"},
|
||
headers=_AUTH,
|
||
)
|
||
self.assertEqual(blank.status_code, 201, blank.text)
|
||
self.assertEqual(blank.json()["name"], "新对话")
|
||
self.assertTrue(blank.json()["auto_title_pending"])
|
||
missing_both = _client.post("/v1/tasks", json={}, headers=_AUTH)
|
||
self.assertEqual(missing_both.status_code, 400, missing_both.text)
|
||
|
||
# 人工改名拥有最高优先级:立刻清 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"])
|
||
self.assertEqual(d["title_source"], "manual")
|
||
|
||
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"], [])
|
||
|
||
def test_clear_rearms_only_auto_title(self):
|
||
auto = _client.post(
|
||
"/v1/tasks",
|
||
json={
|
||
"name": "新对话",
|
||
"working_dir": "自动标题清空测试",
|
||
"auto_title": True,
|
||
},
|
||
headers=_AUTH,
|
||
).json()
|
||
auto_tid = auto["task_id"]
|
||
# 模拟首轮自动标题已经生成;不能走 PATCH,否则会被正确标为 manual。
|
||
with session_scope() as s:
|
||
s.execute(
|
||
__import__("sqlalchemy").update(Task)
|
||
.where(Task.task_id == uuid.UUID(auto_tid))
|
||
.values(name="熟料三率值分析", auto_title_pending=False)
|
||
)
|
||
old_version = s.execute(
|
||
__import__("sqlalchemy").select(Task.auto_title_version)
|
||
.where(Task.task_id == uuid.UUID(auto_tid))
|
||
).scalar_one()
|
||
cleared = _client.post(
|
||
f"/v1/tasks/{auto_tid}/clear", headers=_AUTH
|
||
).json()
|
||
self.assertEqual(cleared["name"], "新对话")
|
||
self.assertTrue(cleared["auto_title_pending"])
|
||
self.assertEqual(cleared["title_source"], "auto")
|
||
with session_scope() as s:
|
||
new_version = s.execute(
|
||
__import__("sqlalchemy").select(Task.auto_title_version)
|
||
.where(Task.task_id == uuid.UUID(auto_tid))
|
||
).scalar_one()
|
||
self.assertEqual(new_version, old_version + 1)
|
||
|
||
manual = _client.post(
|
||
"/v1/tasks",
|
||
json={"name": "用户固定项目名", "working_dir": "人工标题清空测试"},
|
||
headers=_AUTH,
|
||
).json()
|
||
cleared = _client.post(
|
||
f"/v1/tasks/{manual['task_id']}/clear", headers=_AUTH
|
||
).json()
|
||
self.assertEqual(cleared["name"], "用户固定项目名")
|
||
self.assertFalse(cleared["auto_title_pending"])
|
||
self.assertEqual(cleared["title_source"], "manual")
|
||
|
||
|
||
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_hidden_directory_toggle_is_scoped_to_task_working_dir(self):
|
||
tid = self._mk_task("隐藏目录任务", "隐藏目录")
|
||
wd = _user_root() / "隐藏目录"
|
||
(wd / ".meta").mkdir()
|
||
(wd / ".build").mkdir()
|
||
(wd / ".env").write_text("SECRET=x", encoding="utf-8")
|
||
(wd / "visible.txt").write_text("ok", encoding="utf-8")
|
||
(_user_root() / ".root-secret").mkdir(exist_ok=True)
|
||
|
||
default = _client.get(
|
||
"/v1/files", params={"path": "隐藏目录"}, headers=_AUTH
|
||
).json()
|
||
self.assertEqual({entry["name"] for entry in default["entries"]}, {"visible.txt"})
|
||
self.assertFalse(default["hidden_dirs_included"])
|
||
|
||
shown = _client.get(
|
||
"/v1/files",
|
||
params={
|
||
"path": "隐藏目录",
|
||
"include_hidden": "true",
|
||
"task_id": tid,
|
||
},
|
||
headers=_AUTH,
|
||
).json()
|
||
self.assertEqual(
|
||
{entry["name"] for entry in shown["entries"]},
|
||
{".build", ".meta", "visible.txt"},
|
||
)
|
||
self.assertTrue(shown["hidden_dirs_included"])
|
||
|
||
root = _client.get(
|
||
"/v1/files",
|
||
params={"include_hidden": "true", "task_id": tid},
|
||
headers=_AUTH,
|
||
).json()
|
||
self.assertNotIn(".root-secret", {entry["name"] for entry in root["entries"]})
|
||
self.assertFalse(root["hidden_dirs_included"])
|
||
|
||
def test_task_relative_download_survives_working_dir_rename(self):
|
||
tid = self._mk_task("稳定产物任务", "产物旧目录")
|
||
artifact = _user_root() / "产物旧目录" / "reports" / "result.txt"
|
||
artifact.parent.mkdir(parents=True, exist_ok=True)
|
||
artifact.write_text("stable", encoding="utf-8")
|
||
|
||
url = f"/v1/tasks/{tid}/files/download"
|
||
r = _client.get(url, params={"path": "reports/result.txt"}, headers=_AUTH)
|
||
self.assertEqual(r.status_code, 200, r.text)
|
||
self.assertEqual(r.content, b"stable")
|
||
r = _client.get(
|
||
url,
|
||
params={"path": "产物旧目录/reports/result.txt", "legacy": "true"},
|
||
headers=_AUTH,
|
||
)
|
||
self.assertEqual(r.status_code, 200, r.text)
|
||
self.assertEqual(r.content, b"stable")
|
||
|
||
r = _client.post(
|
||
"/v1/files/rename",
|
||
json={"path": "产物旧目录", "new_name": "产物新目录"},
|
||
headers=_AUTH,
|
||
)
|
||
self.assertEqual(r.status_code, 200, r.text)
|
||
r = _client.get(url, params={"path": "reports/result.txt"}, headers=_AUTH)
|
||
self.assertEqual(r.status_code, 200, r.text)
|
||
self.assertEqual(r.content, b"stable")
|
||
self.assertEqual(
|
||
_client.get(url, params={"path": "../escape.txt"}, headers=_AUTH).status_code,
|
||
400,
|
||
)
|
||
|
||
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()
|