491 lines
22 KiB
Python
491 lines
22 KiB
Python
"""Tasks CRUD 路由 + folders + channel_tasks + export(DESIGN §7.1/§7.2)。"""
|
||
from __future__ import annotations
|
||
|
||
import os
|
||
import tempfile
|
||
from pathlib import Path
|
||
from typing import Any, Optional
|
||
from uuid import UUID, uuid4
|
||
|
||
from fastapi import Depends, HTTPException
|
||
from fastapi.responses import FileResponse
|
||
from sqlalchemy import func, select, update
|
||
from starlette.background import BackgroundTask as StarletteBackgroundTask
|
||
|
||
from core.paths import to_db_path
|
||
from core.storage import NoSubtaskError, check_no_subtask, session_scope
|
||
from core.storage.models import Message, Task
|
||
from core.storage.usage_report import task_usage_aggregates as usage_aggregates
|
||
from core.storage.utils import ensure_local_task_row
|
||
|
||
from ..common import (
|
||
CHANNEL_MIRROR_KINDS,
|
||
STATUS_FILTERS,
|
||
STATUS_WRITABLE,
|
||
assert_owns_task,
|
||
iso,
|
||
parse_ordering,
|
||
task_dict,
|
||
)
|
||
from ..model_gate import resolve_model_profile
|
||
from ..schemas import TaskCreateRequest, TaskPatchRequest
|
||
from ..userfiles import system_wd_names
|
||
|
||
|
||
def register_task_routes(app, *, require_user) -> None:
|
||
@app.post("/v1/tasks", status_code=201, tags=["tasks"])
|
||
def create_task(body: TaskCreateRequest, user_id: UUID = Depends(require_user)):
|
||
"""新建 task。
|
||
|
||
- `name` 必填(任务显示名,DB 列 NOT NULL,UI 列表 / 标题用)
|
||
- `working_dir` 可选(留空 → 用 name 作目录名);同 working_dir 多 task 共享同目录(§7.1)
|
||
- `auto_title=true` 仅标记“首条消息后自动命名”;旧调用方缺省 false
|
||
- name / working_dir 都过 validate_task_name(简单名,无 `/\\..`,非 `.` 起头,≤255)
|
||
- 前缀嵌套(no-subtask,同 user 内)→ 409
|
||
"""
|
||
from core.agent_builder import InvalidTaskName, resolve_workspace, validate_task_name, working_dir_from_name
|
||
try:
|
||
name = validate_task_name(body.name)
|
||
except InvalidTaskName as e:
|
||
raise HTTPException(400, f"name 不合法: {e}")
|
||
# working_dir 留空 → fallback 用 name
|
||
wd_raw = (body.working_dir or "").strip()
|
||
wd_name = wd_raw if wd_raw else name
|
||
try:
|
||
wd_name = validate_task_name(wd_name)
|
||
except InvalidTaskName as e:
|
||
raise HTTPException(400, f"working_dir 不合法: {e}")
|
||
description = body.description.strip()
|
||
skill = body.skill.strip()
|
||
|
||
tid = uuid4()
|
||
ws = resolve_workspace(None)
|
||
fs_dir = working_dir_from_name(ws, user_id, wd_name)
|
||
fs_dir_db = to_db_path(fs_dir)
|
||
|
||
try:
|
||
check_no_subtask(fs_dir_db, user_id=user_id)
|
||
except NoSubtaskError as e:
|
||
raise HTTPException(409, str(e))
|
||
|
||
# 工作目录立刻建出(同 working_dir 多 task 共享,exist_ok=True)
|
||
fs_dir.mkdir(parents=True, exist_ok=True)
|
||
|
||
profile, model_id = resolve_model_profile(body.model_profile, user_id=user_id)
|
||
ensure_local_task_row(
|
||
task_id=tid, name=name, working_dir=fs_dir_db, skill=skill,
|
||
description=description, user_id=user_id,
|
||
model=model_id, model_profile=profile,
|
||
auto_title_pending=body.auto_title,
|
||
title_source="auto" if body.auto_title else "manual",
|
||
)
|
||
with session_scope() as s:
|
||
row = s.execute(select(Task).where(Task.task_id == tid)).scalar_one()
|
||
return task_dict(row, n_messages=0)
|
||
|
||
@app.get("/v1/tasks", tags=["tasks"])
|
||
def list_tasks_route(
|
||
page: int = 1,
|
||
page_size: int = 20,
|
||
status: Optional[str] = None,
|
||
skill: Optional[str] = None,
|
||
working_dir: Optional[str] = None,
|
||
q: Optional[str] = None,
|
||
ordering: Optional[str] = None,
|
||
run_status: Optional[str] = None,
|
||
user_id: UUID = Depends(require_user),
|
||
):
|
||
"""列出当前 user 的 task,分页 + 多维筛选 + 排序。
|
||
|
||
- `page` ≥ 1(1-based);`page_size` 1–100(超界 clamp)
|
||
- `status` 在 active/completed/abandoned;非法值静默忽略
|
||
- `skill` 精确匹配(空忽略)
|
||
- `working_dir` 末段目录名(如 `水泥申报`);后端自动拼 `workspace/users/<uid>/<name>` 比对
|
||
- `q` 模糊搜索 name + description(ILIKE,大小写不敏感)
|
||
- `run_status` 逗号分隔,allowlist `idle/running/cancelling/error/cancelled`;非法值静默忽略
|
||
(dev SPA 拉同 wd 活跃 task 用,通常 `running,cancelling`)
|
||
- `ordering` DRF 风格,逗号分隔,`-field` 倒序;allowlist `created_at/updated_at/name/status`;
|
||
非法字段静默忽略;**默认 `-created_at`**(创建时间倒序)
|
||
返回标准分页壳 `{page, page_size, count, results}` —— count 供前端算总页数。
|
||
"""
|
||
# clamp + sanitize
|
||
page = max(1, page)
|
||
page_size = max(1, min(page_size, 100))
|
||
status = status if status in STATUS_FILTERS else None
|
||
skill = (skill or "").strip() or None
|
||
wd_name = (working_dir or "").strip() or None
|
||
q_text = (q or "").strip() or None
|
||
rs_allowed = ("idle", "running", "cancelling", "error", "cancelled")
|
||
run_status_set = {
|
||
s.strip() for s in (run_status or "").split(",") if s.strip() in rs_allowed
|
||
} or None
|
||
|
||
# 组装 WHERE(软删除的 task 永不出现在列表;恢复见 /restore)
|
||
# 渠道镜像 task(wechat/wecom 常驻对话)不进普通列表 —— 它们在左栏「新建任务」下
|
||
# 做成固定卡片(GET /v1/channel_tasks),从列表排除避免重复。coalesce 兜 NULL(老 web task)。
|
||
conditions = [
|
||
Task.user_id == user_id,
|
||
Task.deleted_at.is_(None),
|
||
func.coalesce(Task.channel, "web").notin_(CHANNEL_MIRROR_KINDS),
|
||
# 定时任务执行 task(scheduled_job_id 归属)不进普通列表;兜底 working_dir
|
||
# LIKE 防 backfill 漏网的孤行(job 已物理删的 isolated task)
|
||
Task.scheduled_job_id.is_(None),
|
||
~Task.working_dir.like("%/scheduled-%"),
|
||
]
|
||
if status:
|
||
conditions.append(Task.status == status)
|
||
if skill:
|
||
conditions.append(Task.skill == skill)
|
||
if wd_name:
|
||
# 末段 → 完整 db form。同 working_dir 多 task 共享时,这是命中入口。
|
||
wd_db = f"workspace/users/{user_id}/{wd_name}"
|
||
conditions.append(Task.working_dir == wd_db)
|
||
if run_status_set:
|
||
conditions.append(Task.run_status.in_(run_status_set))
|
||
if q_text:
|
||
pat = f"%{q_text}%"
|
||
conditions.append(Task.name.ilike(pat) | Task.description.ilike(pat))
|
||
|
||
offset = (page - 1) * page_size
|
||
|
||
with session_scope() as s:
|
||
cnt = s.execute(
|
||
select(func.count()).select_from(Task).where(*conditions)
|
||
).scalar_one() or 0
|
||
|
||
rows = s.execute(
|
||
select(Task).where(*conditions)
|
||
.order_by(*parse_ordering(ordering))
|
||
.limit(page_size).offset(offset)
|
||
).scalars().all()
|
||
|
||
tids = [r.task_id for r in rows]
|
||
msg_counts = (
|
||
dict(s.execute(
|
||
select(Message.task_id, func.count())
|
||
.where(Message.task_id.in_(tids))
|
||
.group_by(Message.task_id)
|
||
).all())
|
||
if tids else {}
|
||
)
|
||
usage = usage_aggregates(s, tids)
|
||
|
||
return {
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"count": int(cnt),
|
||
"results": [
|
||
task_dict(
|
||
r,
|
||
n_messages=msg_counts.get(r.task_id, 0),
|
||
usage=usage.get(r.task_id),
|
||
)
|
||
for r in rows
|
||
],
|
||
}
|
||
|
||
@app.get("/v1/channel_tasks", tags=["tasks"])
|
||
def list_channel_tasks(user_id: UUID = Depends(require_user)):
|
||
"""渠道镜像任务(微信 / 企业微信)的绑定状态 + 常驻对话摘要 —— 前端在左栏「新建任务」
|
||
下做成固定卡片。返回 `{ wechat: { bound: bool, task: <task_dict>|null },
|
||
wecom: { bound: bool, task: <task_dict>|null } }`,
|
||
bound 状态由 `get_binding` / `get_wecom_userid` 判定;task 同 /v1/tasks 列表项(复用 task_dict),
|
||
有对话则给摘要,无则 null。前端据此渲染三种卡片:未绑定(点绑定)、已绑定无对话(占位)、
|
||
已绑定有对话(点进 + ⚙ 管理)。
|
||
"""
|
||
from core.wechat import service as _wx
|
||
|
||
snap = _wx.get_binding(user_id)
|
||
wuid = _wx.get_wecom_userid(user_id)
|
||
bound: dict[str, bool] = {
|
||
"wechat": bool(snap and snap.status == "active"),
|
||
"wecom": bool(wuid),
|
||
}
|
||
tids: dict[str, Optional[UUID]] = {
|
||
"wechat": snap.chat_task_id if snap and snap.status == "active" else None,
|
||
"wecom": _wx.get_wecom_chat_task(user_id),
|
||
}
|
||
wanted = [t for t in tids.values() if t is not None]
|
||
tasks: dict[str, Optional[dict]] = {"wechat": None, "wecom": None}
|
||
if wanted:
|
||
with session_scope() as s:
|
||
rows = {
|
||
r.task_id: r
|
||
for r in s.execute(
|
||
select(Task).where(
|
||
Task.task_id.in_(wanted),
|
||
Task.user_id == user_id,
|
||
Task.deleted_at.is_(None),
|
||
)
|
||
).scalars().all()
|
||
}
|
||
msg_counts = dict(
|
||
s.execute(
|
||
select(Message.task_id, func.count())
|
||
.where(Message.task_id.in_(list(rows.keys())))
|
||
.group_by(Message.task_id)
|
||
).all()
|
||
) if rows else {}
|
||
usage = usage_aggregates(s, list(rows.keys()))
|
||
for kind, tid in tids.items():
|
||
row = rows.get(tid) if tid else None
|
||
if row is not None:
|
||
tasks[kind] = task_dict(
|
||
row,
|
||
n_messages=msg_counts.get(row.task_id, 0),
|
||
usage=usage.get(row.task_id),
|
||
)
|
||
# 按渠道返回 { bound, task }
|
||
return {
|
||
"wechat": {"bound": bound["wechat"], "task": tasks["wechat"]},
|
||
"wecom": {"bound": bound["wecom"], "task": tasks["wecom"]},
|
||
}
|
||
|
||
@app.get("/v1/tasks/{task_id}", tags=["tasks"])
|
||
def get_task(task_id: str, user_id: UUID = Depends(require_user)):
|
||
"""单 task meta(不含 messages;走 /messages 拿)。跨 user → 404。
|
||
|
||
额外带上下文压力字段(仅详情端点,列表不加 —— 每 task 一条 sum 聚合,列表
|
||
100 行×聚合不值得):`context_window_chars`(当前窗口体量,idx>=base 的
|
||
payload 字节和)、`context_limit_chars`(reliable_context×2.5 折算容量)、
|
||
`context_pressure`(前者/后者,0-1+)、`context_folds`(折叠次数,=
|
||
usage_events kind='context_fold' 计数)。前端头部压缩指示环用。
|
||
"""
|
||
from sqlalchemy import Text as SAText, cast as sa_cast
|
||
|
||
from core.context import CHARS_PER_TOKEN
|
||
from core.storage.models import UsageEvent
|
||
|
||
try:
|
||
tid = UUID(task_id)
|
||
except ValueError:
|
||
raise HTTPException(404, f"invalid task id: {task_id!r}")
|
||
with session_scope() as s:
|
||
row = s.execute(
|
||
select(Task).where(Task.task_id == tid, Task.user_id == user_id)
|
||
).scalar_one_or_none()
|
||
if row is None:
|
||
raise HTTPException(404, f"task not found: {tid}")
|
||
n = s.execute(
|
||
select(func.count()).select_from(Message).where(Message.task_id == tid)
|
||
).scalar_one()
|
||
usage = usage_aggregates(s, [tid])
|
||
window_chars = s.execute(
|
||
select(func.coalesce(func.sum(func.length(sa_cast(Message.payload, SAText))), 0))
|
||
.where(Message.task_id == tid, Message.idx >= (row.context_base_idx or 0))
|
||
).scalar_one()
|
||
folds = s.execute(
|
||
select(func.count()).select_from(UsageEvent)
|
||
.where(UsageEvent.task_id == tid, UsageEvent.kind == "context_fold")
|
||
).scalar_one()
|
||
d = task_dict(row, n_messages=n, usage=usage.get(tid))
|
||
# 容量按 task 当前模型折算;模型档案读不出(profile 已下线等)→ 字段置 None,
|
||
# 前端画灰环,不 500。
|
||
limit_chars = None
|
||
try:
|
||
from core.agent_builder import load_config
|
||
from core.capabilities import ModelCapabilities
|
||
from core.paths import ROOT
|
||
cfg = load_config()
|
||
profile = d.get("model_profile") or cfg["default_model"]
|
||
caps = ModelCapabilities.load(profile, ROOT / cfg["models_dir"])
|
||
limit_chars = int(caps.reliable_context * CHARS_PER_TOKEN)
|
||
except Exception:
|
||
pass
|
||
d["context_window_chars"] = int(window_chars)
|
||
d["context_limit_chars"] = limit_chars
|
||
d["context_pressure"] = (
|
||
round(int(window_chars) / limit_chars, 4) if limit_chars else None
|
||
)
|
||
d["context_folds"] = int(folds)
|
||
return d
|
||
|
||
@app.get("/v1/folders", tags=["folders"])
|
||
def list_folders(user_id: UUID = Depends(require_user)):
|
||
"""列出当前 user 的工作目录(`workspace/users/<uid>/` 下非 dotfile 子目录)。
|
||
供新建 task 时自动补全 / 选已有目录用。FS 是 source of truth(也含手动创建
|
||
但还无关联 task 的目录)。每项带 n_tasks(关联 task 数)+ last_used(最近使用 ISO)。
|
||
排序:有 last_used 的按降序,无 last_used 的排最后,同列 by name asc。
|
||
"""
|
||
from core.agent_builder import resolve_workspace, user_root
|
||
ws = resolve_workspace(None)
|
||
root = user_root(ws, user_id)
|
||
|
||
folder_names: list[str] = []
|
||
if root.is_dir():
|
||
for p in sorted(root.iterdir(), key=lambda x: x.name.lower()):
|
||
if p.is_dir() and not p.name.startswith("."):
|
||
folder_names.append(p.name)
|
||
# 系统工作目录(定时/渠道)不进候选 —— 新任务不该落到它们里面
|
||
hidden = system_wd_names(user_id) if folder_names else set()
|
||
folder_names = [n for n in folder_names if n not in hidden]
|
||
|
||
folders: list[dict] = []
|
||
if folder_names:
|
||
with session_scope() as s:
|
||
for name in folder_names:
|
||
db_form = f"workspace/users/{user_id}/{name}"
|
||
stat = s.execute(
|
||
select(func.count(), func.max(Task.updated_at))
|
||
.where(
|
||
Task.user_id == user_id,
|
||
Task.working_dir == db_form,
|
||
Task.deleted_at.is_(None),
|
||
)
|
||
).first()
|
||
n = int((stat[0] if stat else 0) or 0)
|
||
lu = stat[1] if stat else None
|
||
folders.append({
|
||
"name": name,
|
||
"n_tasks": n,
|
||
"last_used": iso(lu),
|
||
})
|
||
|
||
folders.sort(key=lambda f: f["name"])
|
||
folders.sort(key=lambda f: f["last_used"] or "", reverse=True)
|
||
return {"folders": folders}
|
||
|
||
@app.delete("/v1/tasks/{task_id}", status_code=204, tags=["tasks"])
|
||
def delete_task(task_id: str, user_id: UUID = Depends(require_user)):
|
||
"""软删除:置 deleted_at=now(),从任务列表隐藏。
|
||
|
||
DB 行 / messages / usage_events(原 CASCADE 不再触发)及工作目录文件全部保留
|
||
—— 留作训练语料,且可经 POST /v1/tasks/{id}/restore 恢复。不动任何磁盘文件。
|
||
已软删的再次调用幂等返回 204。跨 user / 不存在 → 404。
|
||
"""
|
||
try:
|
||
tid = UUID(task_id)
|
||
except ValueError:
|
||
raise HTTPException(404, f"invalid task id: {task_id!r}")
|
||
with session_scope() as s:
|
||
row = s.execute(
|
||
select(Task.deleted_at).where(
|
||
Task.task_id == tid, Task.user_id == user_id,
|
||
)
|
||
).first()
|
||
if row is None:
|
||
raise HTTPException(404, f"task not found: {tid}")
|
||
if row.deleted_at is None:
|
||
s.execute(
|
||
update(Task)
|
||
.where(Task.task_id == tid, Task.user_id == user_id)
|
||
.values(deleted_at=func.now())
|
||
)
|
||
return None # 204
|
||
|
||
@app.post("/v1/tasks/{task_id}/restore", tags=["tasks"])
|
||
def restore_task(task_id: str, user_id: UUID = Depends(require_user)):
|
||
"""恢复软删除的 task(置 deleted_at=NULL),重新出现在列表。
|
||
|
||
未软删的幂等成功。跨 user / 不存在 → 404。
|
||
"""
|
||
try:
|
||
tid = UUID(task_id)
|
||
except ValueError:
|
||
raise HTTPException(404, f"invalid task id: {task_id!r}")
|
||
with session_scope() as s:
|
||
row = s.execute(
|
||
select(Task).where(Task.task_id == tid, Task.user_id == user_id)
|
||
).scalar_one_or_none()
|
||
if row is None:
|
||
raise HTTPException(404, f"task not found: {tid}")
|
||
row.deleted_at = None # ORM 脏标记,session_scope 提交时落库
|
||
n = s.execute(
|
||
select(func.count()).select_from(Message).where(Message.task_id == tid)
|
||
).scalar_one()
|
||
usage = usage_aggregates(s, [tid])
|
||
# 序列化必须在 session 内:updated_at 是 server-side onupdate,flush 后被
|
||
# 标记 expired,出了 session 再读会 DetachedInstanceError(真软删过的恢复
|
||
# 路径 500;幂等路径无脏标记所以此前没暴露 —— test_web_routes_db 抓出)。
|
||
s.flush()
|
||
s.refresh(row)
|
||
d = task_dict(row, n_messages=n, usage=usage.get(tid))
|
||
return d
|
||
|
||
@app.patch("/v1/tasks/{task_id}", tags=["tasks"])
|
||
def patch_task(
|
||
task_id: str,
|
||
body: TaskPatchRequest,
|
||
user_id: UUID = Depends(require_user),
|
||
):
|
||
"""更新 task 字段。`status` 仅允许 completed/abandoned(active 走 CLI 切回)。"""
|
||
try:
|
||
tid = UUID(task_id)
|
||
except ValueError:
|
||
raise HTTPException(404, f"invalid task id: {task_id!r}")
|
||
updates: dict[str, Any] = {}
|
||
if body.status is not None:
|
||
if body.status not in STATUS_WRITABLE:
|
||
raise HTTPException(
|
||
400, f"invalid status {body.status!r}; allowed: {STATUS_WRITABLE}"
|
||
)
|
||
updates["status"] = body.status
|
||
if body.description is not None:
|
||
updates["description"] = body.description
|
||
if body.skill is not None:
|
||
updates["skill"] = body.skill
|
||
if body.name is not None:
|
||
from core.agent_builder import InvalidTaskName, validate_task_name
|
||
try:
|
||
updates["name"] = validate_task_name(body.name)
|
||
except InvalidTaskName as e:
|
||
raise HTTPException(400, f"name 不合法: {e}")
|
||
# 人工命名优先级最高:即使自动标题调用已在途,最终 UPDATE 也会因
|
||
# pending=False 条件失配而放弃,绝不覆盖用户刚改好的名称。
|
||
updates["auto_title_pending"] = False
|
||
updates["title_source"] = "manual"
|
||
if body.model_profile is not None:
|
||
# 切模型:校验后双列同更(profile + model_id)。下条 send 才生效 — 当前
|
||
# in-flight run 不受影响(build_agent resume 时下次重读)。档外模型 → 403。
|
||
profile, model_id = resolve_model_profile(body.model_profile, user_id=user_id)
|
||
updates["model_profile"] = profile
|
||
updates["model"] = model_id
|
||
if not updates:
|
||
raise HTTPException(400, "no fields to update")
|
||
with session_scope() as s:
|
||
result = s.execute(
|
||
update(Task)
|
||
.where(Task.task_id == tid, Task.user_id == user_id)
|
||
.values(**updates)
|
||
)
|
||
if result.rowcount == 0:
|
||
raise HTTPException(404, f"task not found: {tid}")
|
||
row = s.execute(select(Task).where(Task.task_id == tid)).scalar_one()
|
||
n = s.execute(
|
||
select(func.count()).select_from(Message).where(Message.task_id == tid)
|
||
).scalar_one()
|
||
usage = usage_aggregates(s, [tid])
|
||
return task_dict(row, n_messages=n, usage=usage.get(tid))
|
||
|
||
@app.get("/v1/tasks/{task_id}/export", tags=["export"])
|
||
def export_task(task_id: str, user_id: UUID = Depends(require_user)):
|
||
"""导出对话为 .docx,临时文件下载完后 BackgroundTask 删 tmp。"""
|
||
try:
|
||
tid = UUID(task_id)
|
||
except ValueError:
|
||
raise HTTPException(404, f"invalid task id: {task_id!r}")
|
||
with session_scope() as s:
|
||
assert_owns_task(s, tid, user_id)
|
||
has_msg = s.execute(
|
||
select(Message.message_id).where(Message.task_id == tid).limit(1)
|
||
).first()
|
||
if not has_msg:
|
||
raise HTTPException(400, "no messages to export")
|
||
|
||
fd, tmp_str = tempfile.mkstemp(suffix=".docx", prefix="zcbot-export-")
|
||
os.close(fd)
|
||
tmp_path = Path(tmp_str)
|
||
try:
|
||
from core.export_docx import export_chat_to_docx
|
||
export_chat_to_docx(tid, out_path=tmp_path)
|
||
except Exception as e:
|
||
tmp_path.unlink(missing_ok=True)
|
||
raise HTTPException(500, f"export failed: {type(e).__name__}: {e}")
|
||
|
||
return FileResponse(
|
||
path=str(tmp_path),
|
||
media_type="application/vnd.openxmlformats-officedocument.wordprocessingml.document",
|
||
filename=f"chat_{str(tid)[:8]}.docx",
|
||
background=StarletteBackgroundTask(tmp_path.unlink, missing_ok=True),
|
||
)
|