P2: tool adapter layer, mock host, frame pipeline, and local web UI
- MockHostAdapter (tests/mock_host): in-memory editing model with opaque entity ids, id invalidation, token-bucket rate limiting, snapshot/restore with cap, event subscription, and a stdlib-zlib PNG encoder producing time-varying non-black frames - HostToolExecutor: all 26 tool handlers mapped to HostAdapter calls, snapshot orchestration for mutating batches, concurrent scan_timeline frame pipeline with token bucket and RateLimited backoff, undo_session() restoring per-batch snapshots in reverse, get_params min/max/choices reflection - Web UI (core[webui] extra): FastAPI + single-page vanilla JS, chat log, pending-confirmation list with approve/reject endpoints (timeout defaults to reject), progress, undo-session, snapshot management; SSE push + POST, bound to 127.0.0.1 - Host error types (HostError/EntityNotFound/RateLimited) and shared TokenBucket - P2 acceptance chain: import -> place -> split -> ripple delete -> effect + param -> frame verification (non-black, pixel changes at cut and after effect) -> undo_session restores original state Tests: 129 passed via uv run pytest (no network)
This commit is contained in:
@@ -13,9 +13,16 @@ dependencies = [
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
webui = [
|
||||
"fastapi>=0.115,<1.0",
|
||||
"uvicorn>=0.30,<1.0",
|
||||
]
|
||||
dev = [
|
||||
"pytest>=8.0",
|
||||
]
|
||||
|
||||
[project.scripts]
|
||||
quercus-webui = "quercus_core.webui.__main__:main"
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
packages = ["src/quercus_core"]
|
||||
|
||||
@@ -1,4 +1,36 @@
|
||||
"""HostAdapter 抽象接口(§3.1)。P1 只定义 ABC,不提供实现。"""
|
||||
"""HostAdapter 抽象接口(§3.1)。P1 只定义 ABC,P2 增加错误/快照/事件类型。"""
|
||||
from quercus_core.host.adapter import HostAdapter
|
||||
from quercus_core.host.errors import EntityNotFound, HostError, RateLimited
|
||||
from quercus_core.host.types import (
|
||||
BatchResult,
|
||||
FootageId,
|
||||
HostEvent,
|
||||
JobHandle,
|
||||
Levels,
|
||||
MediaInfo,
|
||||
PlaybackState,
|
||||
ProjectOverview,
|
||||
SeqId,
|
||||
SnapshotInfo,
|
||||
Target,
|
||||
Timeline,
|
||||
)
|
||||
|
||||
__all__ = ["HostAdapter"]
|
||||
__all__ = [
|
||||
"HostAdapter",
|
||||
"HostError",
|
||||
"EntityNotFound",
|
||||
"RateLimited",
|
||||
"BatchResult",
|
||||
"FootageId",
|
||||
"HostEvent",
|
||||
"JobHandle",
|
||||
"Levels",
|
||||
"MediaInfo",
|
||||
"PlaybackState",
|
||||
"ProjectOverview",
|
||||
"SeqId",
|
||||
"SnapshotInfo",
|
||||
"Target",
|
||||
"Timeline",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
"""HostAdapter 层通用错误(对齐 OPP/1 §6 的失败语义,P2 新增)。
|
||||
|
||||
- ``EntityNotFound``:不透明 id 失效 / 指向已删除对象(OPP/1 的
|
||||
``ENTITY_NOT_FOUND`` 风格错误)。core 把它当作可回喂 LLM 的工具错误。
|
||||
- ``RateLimited``:宿主限流(取帧等)。``retry_after_ms`` 建议调用方等待的
|
||||
毫秒数;core 侧取帧流水线据此退避(计划文档 §6 约束 4)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class HostError(Exception):
|
||||
"""宿主侧错误基类。"""
|
||||
|
||||
|
||||
class EntityNotFound(HostError):
|
||||
"""实体 id 不存在或已失效(删除后 id 作废,再访问即抛此错)。"""
|
||||
|
||||
|
||||
class RateLimited(HostError):
|
||||
"""宿主限流:调用被拒绝,建议 ``retry_after_ms`` 毫秒后重试。"""
|
||||
|
||||
def __init__(self, retry_after_ms: int, message: str = "操作被限流"):
|
||||
super().__init__(message)
|
||||
self.retry_after_ms = retry_after_ms
|
||||
@@ -100,6 +100,29 @@ class BatchResult:
|
||||
ok: bool
|
||||
message: str = ""
|
||||
results: list[ToolResult] = field(default_factory=list)
|
||||
# P2:本次批次执行前宿主为它创建的快照 id(空串表示未建快照)。
|
||||
# 工具适配层据此记录"本会话批次 → 快照"映射,供 undo_session 反向恢复。
|
||||
snapshot_id: str = ""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SnapshotInfo:
|
||||
"""快照元数据(供 Web 面板快照管理与 undo_session 定位)。"""
|
||||
|
||||
id: str
|
||||
session_id: str
|
||||
created_at: str
|
||||
label: str = ""
|
||||
action_count: int = 0
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class HostEvent:
|
||||
"""宿主事件(结构变化 / 播放头 / 导出进度),订阅者回调收到它。"""
|
||||
|
||||
kind: str # structure_changed | playhead_moved | export_progress | ...
|
||||
seq_id: str = ""
|
||||
payload: dict = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -0,0 +1,497 @@
|
||||
"""工具适配层:P1 的 26 个工具 schema → HostAdapter 方法组合(ToolExecutor 协议)。
|
||||
|
||||
设计要点(计划文档 §3.2–3.3 / §6):
|
||||
- 每个工具一个 handler;只读工具直通 adapter 方法;变更类工具包成
|
||||
``ActionBatch`` 经 ``adapter.execute``(快照 + 补偿语义)执行。
|
||||
- ``get_params`` 从 adapter 读出效果参数的 min/max/choices(供未来回填
|
||||
tool schema)。
|
||||
- 取帧流水线:``scan_timeline`` 批量取帧并发执行(ThreadPoolExecutor,
|
||||
禁止串行等帧);前置令牌桶按 ``Limits`` 限速;遇 ``RateLimited`` 按
|
||||
``retry_after_ms`` 退避(有限重试);帧结果用完即弃(不缓存进会话)。
|
||||
- 会话撤销 ``undo_session()``:反向恢复本会话各批次快照(mock 宿主上
|
||||
即"恢复各批次快照"),供 Web 面板"撤销整段会话"按钮使用。
|
||||
|
||||
AgentLoop(P1 协议)逐 action 派发;本类同时提供 ``execute_batch`` 供
|
||||
webui/测试按批执行(一个批次一个快照)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
import uuid
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable
|
||||
|
||||
from quercus_core.host.adapter import HostAdapter
|
||||
from quercus_core.host.errors import EntityNotFound, RateLimited
|
||||
from quercus_core.host.types import BatchResult, SnapshotInfo, Target
|
||||
from quercus_core.tools.ratelimit import TokenBucket
|
||||
from quercus_core.tools.schemas import TOOLS
|
||||
from quercus_core.types import (
|
||||
Action,
|
||||
ActionBatch,
|
||||
PngBytes,
|
||||
Rational,
|
||||
Size,
|
||||
TimeRange,
|
||||
ToolResult,
|
||||
)
|
||||
|
||||
# 变更类工具里不走"快照批次"的两类:导出(进度语义)与撤销(自身即回退)。
|
||||
_DIRECT_MUTATION_TOOLS = frozenset({"export_render", "undo_last_action"})
|
||||
|
||||
# 批量取帧的默认缩略尺寸(contact sheet 用小图,避免会话囤积大图)。
|
||||
_SCAN_FRAME_SIZE = Size(480, 270)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ProgressEvent:
|
||||
"""进度事件(scan / export),经 progress_callback 上抛给 UI。"""
|
||||
|
||||
kind: str # scan | export
|
||||
seq_id: str = ""
|
||||
done: int = 0
|
||||
total: int = 0
|
||||
message: str = ""
|
||||
|
||||
|
||||
def _sample_times(rng: TimeRange, count: int) -> list[Rational]:
|
||||
if count <= 1:
|
||||
return [rng.start]
|
||||
step = rng.duration / Rational(count - 1, 1)
|
||||
return [rng.start + step * Rational(i, 1) for i in range(count)]
|
||||
|
||||
|
||||
class HostToolExecutor:
|
||||
"""实现 ``quercus_core.agent.loop.ToolExecutor`` 协议的工具适配层。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
adapter: HostAdapter,
|
||||
*,
|
||||
session_id: str | None = None,
|
||||
frame_concurrency: int = 4,
|
||||
max_frame_retries: int = 4,
|
||||
rate_limiter: TokenBucket | None = None,
|
||||
progress_callback: Callable[[ProgressEvent], None] | None = None,
|
||||
) -> None:
|
||||
self._adapter = adapter
|
||||
self._session_id = session_id or uuid.uuid4().hex
|
||||
self._concurrency = max(1, frame_concurrency)
|
||||
self._max_frame_retries = max(0, max_frame_retries)
|
||||
limits = adapter.limits()
|
||||
self._bucket = rate_limiter or TokenBucket(
|
||||
capacity=limits.frame_burst, refill_rate=limits.max_frame_rate
|
||||
)
|
||||
self._progress = progress_callback
|
||||
# 会话撤销日志:(snapshot_id, call_id, tool)
|
||||
self._undo_log: list[tuple[str, str, str]] = []
|
||||
|
||||
# ---- 会话属性 ----
|
||||
|
||||
@property
|
||||
def session_id(self) -> str:
|
||||
return self._session_id
|
||||
|
||||
@property
|
||||
def adapter(self) -> HostAdapter:
|
||||
return self._adapter
|
||||
|
||||
# ---- ToolExecutor 协议:逐 action 派发(AgentLoop 入口) ----
|
||||
|
||||
def execute(self, action: Action) -> ToolResult:
|
||||
if action.tool in _MUTATION_TOOLS:
|
||||
return self._run_mutation(action)
|
||||
if action.tool in _DIRECT_MUTATION_TOOLS:
|
||||
return self._run_direct(action)
|
||||
handler = _READONLY_HANDLERS.get(action.tool)
|
||||
if handler is None:
|
||||
return ToolResult(tool=action.tool, ok=False, summary=f"未知工具 {action.tool!r}")
|
||||
return handler(self, action)
|
||||
|
||||
# ---- 批级执行(webui/测试用:一个批次一个快照) ----
|
||||
|
||||
def execute_batch(self, batch: ActionBatch) -> list[ToolResult]:
|
||||
"""整批执行:只读直通、变更动作合并为一次 adapter.execute(一个快照)。"""
|
||||
results: list[ToolResult] = []
|
||||
mutations: list[Action] = []
|
||||
for action in batch.actions:
|
||||
tool_schema = TOOLS.get(action.tool)
|
||||
if action.tool in _DIRECT_MUTATION_TOOLS:
|
||||
results.append(self.execute(action))
|
||||
elif tool_schema is None:
|
||||
results.append(
|
||||
ToolResult(tool=action.tool, ok=False, summary=f"未知工具 {action.tool!r}")
|
||||
)
|
||||
elif tool_schema.read_only:
|
||||
results.append(self.execute(action))
|
||||
else:
|
||||
mutations.append(action)
|
||||
if mutations:
|
||||
sub = ActionBatch(
|
||||
label=batch.label or f"AI 动作:{len(mutations)}",
|
||||
actions=mutations,
|
||||
id=batch.id or uuid.uuid4().hex,
|
||||
session_id=batch.session_id or self._session_id,
|
||||
)
|
||||
br = self._adapter.execute(sub)
|
||||
if br.snapshot_id:
|
||||
self._undo_log.append((br.snapshot_id, sub.id, sub.label))
|
||||
mresults = list(br.results)
|
||||
while len(mresults) < len(mutations):
|
||||
mresults.append(
|
||||
ToolResult(tool=mutations[len(mresults)].tool, ok=False, summary="缺失结果")
|
||||
)
|
||||
results.extend(mresults)
|
||||
return results
|
||||
|
||||
# ---- 撤销 / 快照 / 参数(Web 面板与 schema 回填用) ----
|
||||
|
||||
def undo_session(self) -> int:
|
||||
"""整段会话撤销:反向恢复本会话各批次快照(mock 宿主上即"恢复各批次快照")。
|
||||
|
||||
已消费(undo_last / drop)的快照会抛 EntityNotFound,跳过即可。
|
||||
"""
|
||||
restored = 0
|
||||
for snapshot_id, _call_id, _tool in reversed(self._undo_log):
|
||||
try:
|
||||
self._adapter.restore_snapshot(snapshot_id)
|
||||
restored += 1
|
||||
except (EntityNotFound, NotImplementedError):
|
||||
continue
|
||||
self._undo_log.clear()
|
||||
return restored
|
||||
|
||||
def list_snapshots(self) -> list[SnapshotInfo]:
|
||||
method = getattr(self._adapter, "list_snapshots", None)
|
||||
return list(method()) if method else []
|
||||
|
||||
def drop_snapshot(self, snapshot_id: str) -> None:
|
||||
method = getattr(self._adapter, "drop_snapshot", None)
|
||||
if method is None:
|
||||
raise EntityNotFound("宿主不支持快照管理")
|
||||
method(snapshot_id)
|
||||
|
||||
def get_params(self, effect_type: str) -> list[dict]:
|
||||
"""从 adapter 读出效果参数 min/max/choices(供未来回填 tool schema)。
|
||||
|
||||
未知效果类型返回空列表(优雅降级,不打断 schema 回填)。
|
||||
"""
|
||||
method = getattr(self._adapter, "get_effect_params", None)
|
||||
if method is None:
|
||||
return []
|
||||
try:
|
||||
return method(effect_type)
|
||||
except EntityNotFound:
|
||||
return []
|
||||
|
||||
# ---- 变更类工具(经 adapter.execute 快照批次) ----
|
||||
|
||||
def _run_mutation(self, action: Action) -> ToolResult:
|
||||
batch = ActionBatch(
|
||||
label=f"AI 动作:{action.tool}",
|
||||
actions=[Action(tool=action.tool, params=dict(action.params), call_id=action.call_id)],
|
||||
id=action.call_id or uuid.uuid4().hex,
|
||||
session_id=self._session_id,
|
||||
)
|
||||
br = self._adapter.execute(batch)
|
||||
if br.snapshot_id:
|
||||
self._undo_log.append((br.snapshot_id, action.call_id, action.tool))
|
||||
if br.results:
|
||||
return br.results[0]
|
||||
return ToolResult(
|
||||
tool=action.tool,
|
||||
ok=br.ok,
|
||||
summary=br.message or ("ok" if br.ok else "执行失败"),
|
||||
error=None if br.ok else br.message,
|
||||
)
|
||||
|
||||
def _run_direct(self, action: Action) -> ToolResult:
|
||||
if action.tool == "undo_last_action":
|
||||
try:
|
||||
self._adapter.undo_last()
|
||||
return ToolResult(tool="undo_last_action", ok=True, summary="已撤销上一个 AI 动作(恢复其快照)")
|
||||
except (EntityNotFound, NotImplementedError) as exc:
|
||||
return ToolResult(tool="undo_last_action", ok=False, summary=f"撤销失败:{exc}", error=str(exc))
|
||||
if action.tool == "export_render":
|
||||
return _h_export_render(self, action)
|
||||
return ToolResult(tool=action.tool, ok=False, summary=f"未知直接工具 {action.tool!r}")
|
||||
|
||||
# ---- 取帧流水线 ----
|
||||
|
||||
def _fetch_frame(self, target: Target, time_point: Rational, size: Size) -> PngBytes:
|
||||
"""取一帧:前置令牌桶 + 宿主 RateLimited 有限重试退避(不轰炸)。"""
|
||||
for attempt in range(self._max_frame_retries + 1):
|
||||
self._bucket.acquire() # 前置限速(按 Limits 令牌桶)
|
||||
try:
|
||||
return self._adapter.get_frame(target, time_point, size)
|
||||
except RateLimited as exc:
|
||||
if attempt >= self._max_frame_retries:
|
||||
raise
|
||||
time.sleep(exc.retry_after_ms / 1000.0)
|
||||
raise RuntimeError("不可达") # pragma: no cover
|
||||
|
||||
def _current_seq_id(self) -> str:
|
||||
try:
|
||||
ov = self._adapter.get_project_overview()
|
||||
except Exception:
|
||||
return ""
|
||||
return ov.timeline_ids[0] if ov.timeline_ids else ""
|
||||
|
||||
def _tool_target(self, params: dict) -> Target:
|
||||
"""get_frame / scan_timeline 的目标:显式 target 或当前时间线。"""
|
||||
target_id = params.get("target")
|
||||
if target_id:
|
||||
return Target(kind="timeline", id=target_id)
|
||||
return Target(kind="timeline", id=self._current_seq_id())
|
||||
|
||||
def _emit_progress(self, kind: str, done: int, total: int, message: str = "") -> None:
|
||||
if self._progress:
|
||||
self._progress(
|
||||
ProgressEvent(kind=kind, seq_id=self._current_seq_id(), done=done, total=total, message=message)
|
||||
)
|
||||
|
||||
|
||||
# ---- 只读工具 handler(直通 adapter) ----
|
||||
|
||||
def _h_get_project_overview(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
try:
|
||||
ov = self._adapter.get_project_overview()
|
||||
except Exception as exc:
|
||||
return ToolResult(tool="get_project_overview", ok=False, summary=f"获取概览失败:{exc}", error=str(exc))
|
||||
lines = [f"工程:{ov.name}(fps {ov.fps:g})"]
|
||||
if ov.duration is not None:
|
||||
lines.append(f"时长:{ov.duration.to_float():.3f}s")
|
||||
for sid in ov.timeline_ids:
|
||||
try:
|
||||
tl = self._adapter.get_timeline_structure(sid)
|
||||
track_desc = "、".join(
|
||||
f"{t.kind}{t.index}({len(t.clips)}块)" for t in tl.tracks
|
||||
) or "空"
|
||||
lines.append(f"时间线 {sid}「{tl.name}」:{track_desc},时长 {tl.duration.to_float():.3f}s")
|
||||
except Exception as exc:
|
||||
lines.append(f"时间线 {sid}:读取失败({exc})")
|
||||
return ToolResult(tool="get_project_overview", ok=True, summary="\n".join(lines))
|
||||
|
||||
|
||||
def _h_probe_media(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
try:
|
||||
info = self._adapter.probe_media(action.params["path"])
|
||||
except Exception as exc:
|
||||
return ToolResult(tool="probe_media", ok=False, summary=f"探测失败:{exc}", error=str(exc))
|
||||
parts = [f"路径:{info.path}"]
|
||||
if info.duration is not None:
|
||||
parts.append(f"时长:{info.duration.to_float():.3f}s")
|
||||
if info.width and info.height:
|
||||
parts.append(f"分辨率:{info.width}x{info.height}")
|
||||
if info.fps:
|
||||
parts.append(f"帧率:{info.fps:g}")
|
||||
return ToolResult(tool="probe_media", ok=True, summary=";".join(parts))
|
||||
|
||||
|
||||
def _h_list_footage(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
method = getattr(self._adapter, "list_footage", None)
|
||||
if method is None:
|
||||
return ToolResult(tool="list_footage", ok=False, summary="宿主不支持 list_footage")
|
||||
try:
|
||||
items = method()
|
||||
except Exception as exc:
|
||||
return ToolResult(tool="list_footage", ok=False, summary=f"列出素材失败:{exc}", error=str(exc))
|
||||
if not items:
|
||||
return ToolResult(tool="list_footage", ok=True, summary="媒体池为空")
|
||||
return ToolResult(
|
||||
tool="list_footage",
|
||||
ok=True,
|
||||
summary="媒体池({}):{}".format(
|
||||
len(items), "、".join(f"{it['name']}({it['id']})" for it in items)
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _h_list_effects(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
method = getattr(self._adapter, "list_effects", None)
|
||||
if method is None:
|
||||
return ToolResult(tool="list_effects", ok=False, summary="宿主不支持 list_effects")
|
||||
clip_id = action.params.get("clip_id")
|
||||
category = action.params.get("category")
|
||||
try:
|
||||
items = method(clip_id, category)
|
||||
except Exception as exc:
|
||||
return ToolResult(tool="list_effects", ok=False, summary=f"列出效果失败:{exc}", error=str(exc))
|
||||
if clip_id:
|
||||
if not items:
|
||||
return ToolResult(tool="list_effects", ok=True, summary=f"片段 {clip_id} 无效果")
|
||||
desc = ", ".join("{}({})".format(it["type"], it["node_id"]) for it in items)
|
||||
return ToolResult(
|
||||
tool="list_effects",
|
||||
ok=True,
|
||||
summary=f"片段 {clip_id} 效果:{desc}",
|
||||
)
|
||||
if not items:
|
||||
return ToolResult(tool="list_effects", ok=True, summary="无可用的效果类型")
|
||||
return ToolResult(
|
||||
tool="list_effects",
|
||||
ok=True,
|
||||
summary=f"可用效果类型:{', '.join(it['type'] for it in items)}",
|
||||
)
|
||||
|
||||
|
||||
def _h_get_frame(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
params = action.params
|
||||
try:
|
||||
t = Rational.from_json(params["time"])
|
||||
size = Size(int(params["max_size"]["width"]), int(params["max_size"]["height"]))
|
||||
target = self._tool_target(params)
|
||||
frame = self._fetch_frame(target, t, size)
|
||||
except RateLimited as exc:
|
||||
return ToolResult(
|
||||
tool="get_frame",
|
||||
ok=False,
|
||||
summary=f"取帧被限流(退避 {exc.retry_after_ms}ms 后仍失败)",
|
||||
error=str(exc),
|
||||
)
|
||||
except (EntityNotFound, ValueError) as exc:
|
||||
return ToolResult(tool="get_frame", ok=False, summary=f"取帧失败:{exc}", error=str(exc))
|
||||
return ToolResult(
|
||||
tool="get_frame",
|
||||
ok=True,
|
||||
summary=f"帧 @ {t.to_float():.3f}s({len(frame)} 字节 PNG)",
|
||||
images=(frame,),
|
||||
)
|
||||
|
||||
|
||||
def _h_scan_timeline(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
"""批量取帧:并发执行 + 令牌桶 + RateLimited 退避;帧用完即弃。"""
|
||||
params = action.params
|
||||
try:
|
||||
rng = TimeRange.from_json(params["range"])
|
||||
except (ValueError, KeyError) as exc:
|
||||
return ToolResult(tool="scan_timeline", ok=False, summary=f"参数非法:{exc}", error=str(exc))
|
||||
count = max(1, min(int(params.get("count", 8)), self._adapter.limits().max_scan_frames))
|
||||
target = self._tool_target(params)
|
||||
times = _sample_times(rng, count)
|
||||
|
||||
frames: dict[int, PngBytes] = {}
|
||||
errors: list[str] = []
|
||||
workers = min(self._concurrency, count)
|
||||
done = 0
|
||||
with ThreadPoolExecutor(max_workers=workers) as pool:
|
||||
future_to_index = {
|
||||
pool.submit(self._fetch_frame, target, t, _SCAN_FRAME_SIZE): i
|
||||
for i, t in enumerate(times)
|
||||
}
|
||||
for future in as_completed(future_to_index):
|
||||
i = future_to_index[future]
|
||||
try:
|
||||
frames[i] = future.result()
|
||||
except RateLimited as exc:
|
||||
errors.append(f"第 {i + 1} 帧限流:{exc}")
|
||||
except (EntityNotFound, ValueError) as exc:
|
||||
errors.append(f"第 {i + 1} 帧失败:{exc}")
|
||||
done += 1
|
||||
self._emit_progress("scan", done, count, f"扫描 {done}/{count}")
|
||||
|
||||
ordered = [frames.get(i) for i in range(count)]
|
||||
if errors:
|
||||
return ToolResult(
|
||||
tool="scan_timeline",
|
||||
ok=False,
|
||||
summary=f"扫描 {count} 帧,成功 {len(ordered) - sum(1 for f in ordered if f is None)},失败 {len(errors)}:{errors[0]}",
|
||||
images=tuple(f for f in ordered if f is not None),
|
||||
error="; ".join(errors),
|
||||
)
|
||||
return ToolResult(
|
||||
tool="scan_timeline",
|
||||
ok=True,
|
||||
summary=f"扫描 {count} 帧({rng.start.to_float():.2f}s ~ {rng.end.to_float():.2f}s)",
|
||||
images=tuple(ordered), # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
|
||||
def _h_get_audio_levels(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
params = action.params
|
||||
try:
|
||||
rng = TimeRange.from_json(params["range"])
|
||||
resolution = int(params.get("resolution", 64))
|
||||
levels = self._adapter.get_audio_levels(self._current_seq_id(), rng, resolution)
|
||||
except (EntityNotFound, ValueError, KeyError) as exc:
|
||||
return ToolResult(tool="get_audio_levels", ok=False, summary=f"读取电平失败:{exc}", error=str(exc))
|
||||
if not levels.values:
|
||||
return ToolResult(tool="get_audio_levels", ok=True, summary="无电平数据")
|
||||
return ToolResult(
|
||||
tool="get_audio_levels",
|
||||
ok=True,
|
||||
summary=f"电平 {len(levels.values)} 点(min {min(levels.values):.1f}dB,max {max(levels.values):.1f}dB)",
|
||||
)
|
||||
|
||||
|
||||
def _h_play(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
try:
|
||||
self._adapter.play(float(action.params.get("speed", 1.0)))
|
||||
except Exception as exc:
|
||||
return ToolResult(tool="play", ok=False, summary=f"播放失败:{exc}", error=str(exc))
|
||||
return ToolResult(tool="play", ok=True, summary="开始回放")
|
||||
|
||||
|
||||
def _h_pause(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
try:
|
||||
self._adapter.pause()
|
||||
except Exception as exc:
|
||||
return ToolResult(tool="pause", ok=False, summary=f"暂停失败:{exc}", error=str(exc))
|
||||
return ToolResult(tool="pause", ok=True, summary="已暂停")
|
||||
|
||||
|
||||
def _h_seek(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
try:
|
||||
t = Rational.from_json(action.params["time"])
|
||||
self._adapter.seek(t)
|
||||
except (ValueError, KeyError) as exc:
|
||||
return ToolResult(tool="seek", ok=False, summary=f"seek 失败:{exc}", error=str(exc))
|
||||
return ToolResult(tool="seek", ok=True, summary=f"播放头移至 {t.to_float():.3f}s")
|
||||
|
||||
|
||||
def _h_export_render(self: HostToolExecutor, action: Action) -> ToolResult:
|
||||
"""导出:adapter.export 启动 + (若宿主支持)有界轮询 export_status 报进度。"""
|
||||
params = action.params
|
||||
output = params["output"]
|
||||
preset = params.get("preset", "")
|
||||
seq = self._current_seq_id()
|
||||
try:
|
||||
job = self._adapter.export(seq, output, preset)
|
||||
except (EntityNotFound, ValueError) as exc:
|
||||
return ToolResult(tool="export_render", ok=False, summary=f"导出失败:{exc}", error=str(exc))
|
||||
status_fn = getattr(self._adapter, "export_status", None)
|
||||
if status_fn is None:
|
||||
return ToolResult(tool="export_render", ok=True, summary=f"已启动导出任务 {job.job_id} → {output}")
|
||||
poll_interval = min(getattr(self._adapter.limits(), "poll_interval_seconds", 0.5), 0.05)
|
||||
for _ in range(100): # 有界轮询,不无限等
|
||||
st = status_fn(job.job_id)
|
||||
self._emit_progress("export", int(float(st.get("progress", 0.0)) * 100), 100, st.get("message", ""))
|
||||
if st.get("status") == "done":
|
||||
return ToolResult(tool="export_render", ok=True, summary=f"导出完成 → {output}")
|
||||
if st.get("status") == "failed":
|
||||
return ToolResult(tool="export_render", ok=False, summary=f"导出失败:{st.get('message')}", error=st.get("message"))
|
||||
time.sleep(poll_interval)
|
||||
return ToolResult(tool="export_render", ok=True, summary=f"导出仍在进行:{job.job_id} → {output}")
|
||||
|
||||
|
||||
# ---- 派发表 ----
|
||||
|
||||
_READONLY_HANDLERS: dict[str, Callable[[HostToolExecutor, Action], ToolResult]] = {
|
||||
"get_project_overview": _h_get_project_overview,
|
||||
"probe_media": _h_probe_media,
|
||||
"list_footage": _h_list_footage,
|
||||
"list_effects": _h_list_effects,
|
||||
"get_frame": _h_get_frame,
|
||||
"scan_timeline": _h_scan_timeline,
|
||||
"get_audio_levels": _h_get_audio_levels,
|
||||
"play": _h_play,
|
||||
"pause": _h_pause,
|
||||
"seek": _h_seek,
|
||||
}
|
||||
|
||||
_MUTATION_TOOLS: frozenset[str] = frozenset(
|
||||
name
|
||||
for name, schema in TOOLS.items()
|
||||
if not schema.read_only and name not in _DIRECT_MUTATION_TOOLS
|
||||
)
|
||||
@@ -0,0 +1,67 @@
|
||||
"""令牌桶限流器:取帧流水线的通用限流客户端(计划文档 §6 约束 4)。
|
||||
|
||||
core 侧(工具适配层取帧流水线)与 mock 宿主共用同一实现:
|
||||
- 执行器用它做**前置令牌桶**(按 ``Limits`` 均匀限速,避免轰炸宿主);
|
||||
- mock 宿主用它做**强制限流**(令牌耗尽即抛 ``RateLimited``,测试客户端退避)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
|
||||
|
||||
class TokenBucket:
|
||||
"""简单令牌桶:容量 ``capacity``,按 ``refill_rate``(令牌/秒)补充。
|
||||
|
||||
线程安全。``acquire()`` 阻塞直到取到一枚令牌;``try_acquire()`` 非阻塞,
|
||||
取不到返回 False 并可用 ``seconds_until_token()`` 查询需等待时长。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
capacity: int,
|
||||
refill_rate: float,
|
||||
monotonic=time.monotonic,
|
||||
) -> None:
|
||||
if capacity < 1:
|
||||
raise ValueError("令牌桶容量必须 >= 1")
|
||||
if refill_rate <= 0:
|
||||
raise ValueError("补充速率必须为正")
|
||||
self._capacity = float(capacity)
|
||||
self._rate = float(refill_rate)
|
||||
self._monotonic = monotonic
|
||||
self._tokens = float(capacity)
|
||||
self._last = monotonic()
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def _refill(self, now: float) -> None:
|
||||
self._tokens = min(self._capacity, self._tokens + (now - self._last) * self._rate)
|
||||
self._last = now
|
||||
|
||||
def try_acquire(self) -> bool:
|
||||
"""非阻塞取令牌:取到返回 True,否则 False(不扣减)。"""
|
||||
with self._lock:
|
||||
self._refill(self._monotonic())
|
||||
if self._tokens >= 1.0:
|
||||
self._tokens -= 1.0
|
||||
return True
|
||||
return False
|
||||
|
||||
def seconds_until_token(self) -> float:
|
||||
"""按当前余量估算取到下一枚令牌所需秒数(0 表示立即可取)。"""
|
||||
with self._lock:
|
||||
self._refill(self._monotonic())
|
||||
if self._tokens >= 1.0:
|
||||
return 0.0
|
||||
return (1.0 - self._tokens) / self._rate
|
||||
|
||||
def acquire(self) -> None:
|
||||
"""阻塞直到取到一枚令牌(等待期间持锁,令牌发放全局串行)。"""
|
||||
with self._lock:
|
||||
while True:
|
||||
self._refill(self._monotonic())
|
||||
if self._tokens >= 1.0:
|
||||
self._tokens -= 1.0
|
||||
return
|
||||
wait = (1.0 - self._tokens) / self._rate
|
||||
time.sleep(wait)
|
||||
@@ -171,6 +171,7 @@ class Limits:
|
||||
max_frame_width: int = 1920
|
||||
max_frame_height: int = 1080
|
||||
max_frame_rate: float = 8.0 # 每秒取帧上限
|
||||
frame_burst: int = 4 # 取帧令牌桶容量(突发余量,P2 新增)
|
||||
max_scan_frames: int = 64 # 单次 scan_timeline 的采样数上限
|
||||
max_snapshot_count: int = 16 # 快照时间线数量上限
|
||||
poll_interval_seconds: float = 0.5 # 宿主无事件源时适配层轮询间隔
|
||||
@@ -180,6 +181,7 @@ class Limits:
|
||||
"max_frame_width": self.max_frame_width,
|
||||
"max_frame_height": self.max_frame_height,
|
||||
"max_frame_rate": self.max_frame_rate,
|
||||
"frame_burst": self.frame_burst,
|
||||
"max_scan_frames": self.max_scan_frames,
|
||||
"max_snapshot_count": self.max_snapshot_count,
|
||||
"poll_interval_seconds": self.poll_interval_seconds,
|
||||
@@ -210,6 +212,7 @@ class ActionBatch:
|
||||
actions: list[Action]
|
||||
id: str = "" # 由调用方(loop)分配
|
||||
created_at: str = "" # ISO 时间戳
|
||||
session_id: str = "" # 所属会话 id(P2:快照带会话标识,供整段撤销定位)
|
||||
|
||||
def to_json(self) -> dict[str, Any]:
|
||||
return {
|
||||
@@ -217,6 +220,7 @@ class ActionBatch:
|
||||
"actions": [a.to_json() for a in self.actions],
|
||||
"id": self.id,
|
||||
"created_at": self.created_at,
|
||||
"session_id": self.session_id,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
"""本地 Web 面板(可选 extra ``core[webui]``):聊天 + 待确认清单 + 进度。
|
||||
|
||||
交互:SSE 推送(服务端 → 浏览器)+ POST(浏览器 → 服务端)。
|
||||
启动入口:``python -m quercus_core.webui`` 或 ``quercus-webui``。
|
||||
"""
|
||||
from quercus_core.webui.app import create_app
|
||||
from quercus_core.webui.broker import EventBroker
|
||||
from quercus_core.webui.session import WebSession
|
||||
|
||||
__all__ = ["create_app", "EventBroker", "WebSession"]
|
||||
@@ -0,0 +1,98 @@
|
||||
"""本地 Web 面板启动入口:``python -m quercus_core.webui`` 或 ``quercus-webui``。
|
||||
|
||||
绑定 127.0.0.1 + 随机端口(``--port 0``);P2 阶段宿主适配层只有 mock,
|
||||
因此 ``--adapter mock`` 会用 tests/mock_host 的内存实现演示。
|
||||
需安装可选 extra ``core[webui]``(fastapi + uvicorn)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import socket
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
parser = argparse.ArgumentParser(description="Quercus AI 剪辑助手 — 本地 Web 面板")
|
||||
parser.add_argument("--host", default="127.0.0.1", help="监听地址(默认仅本机)")
|
||||
parser.add_argument("--port", type=int, default=0, help="监听端口(0 = 随机端口)")
|
||||
parser.add_argument(
|
||||
"--adapter",
|
||||
default="mock",
|
||||
choices=["mock"],
|
||||
help="宿主适配层(P2 阶段仅 mock)",
|
||||
)
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
try:
|
||||
import uvicorn # noqa: F401
|
||||
except ImportError:
|
||||
print(
|
||||
"缺少依赖 uvicorn,请先安装 core[webui]:uv pip install -e 'core[webui]'",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 2
|
||||
|
||||
if args.port == 0:
|
||||
with socket.socket() as s:
|
||||
s.bind(("127.0.0.1", 0))
|
||||
args.port = s.getsockname()[1]
|
||||
|
||||
adapter = _make_adapter(args.adapter)
|
||||
provider = _make_provider()
|
||||
from quercus_core.webui.app import create_app
|
||||
|
||||
app = create_app(adapter=adapter, provider=provider, confirm_timeout=300.0)
|
||||
print(f"Quercus Web 面板:http://{args.host}:{args.port} (Ctrl-C 退出)", flush=True)
|
||||
uvicorn.run(app, host=args.host, port=args.port, log_level="warning")
|
||||
return 0
|
||||
|
||||
|
||||
def _make_adapter(name: str):
|
||||
if name != "mock":
|
||||
raise ValueError(f"未知 adapter: {name}")
|
||||
# P2 无真实宿主,用 tests/mock_host 的内存实现演示。
|
||||
try:
|
||||
from mock_host import MockHostAdapter
|
||||
except ImportError: # 直接运行时 tests/ 不在 sys.path,手动补上
|
||||
repo_root = Path(__file__).resolve().parents[3].parent # .../quercus
|
||||
tests_dir = repo_root / "tests"
|
||||
if str(tests_dir) not in sys.path:
|
||||
sys.path.insert(0, str(tests_dir))
|
||||
from mock_host import MockHostAdapter
|
||||
return MockHostAdapter()
|
||||
|
||||
|
||||
def _make_provider():
|
||||
"""有 key 用真实 provider;无 key 用降级 provider(不做对话编排)。"""
|
||||
from quercus_core.config import load_config
|
||||
from quercus_core.providers import CLAUDE, OPENAI_COMPAT, create_provider
|
||||
|
||||
cfg = load_config()
|
||||
for name in (CLAUDE, OPENAI_COMPAT):
|
||||
try:
|
||||
return create_provider(name, cfg)
|
||||
except Exception:
|
||||
continue
|
||||
return _DegradedProvider()
|
||||
|
||||
|
||||
class _DegradedProvider:
|
||||
"""无 key 降级:告知用户如何配置,不参与对话编排。"""
|
||||
|
||||
name = "degraded"
|
||||
|
||||
def generate(self, messages, tools=None):
|
||||
from quercus_core.providers.base import AssistantTurn
|
||||
|
||||
return AssistantTurn(
|
||||
text=(
|
||||
"未配置 LLM API key,无法进行 AI 对话编排。\n"
|
||||
"请设置 QUERCUS_ANTHROPIC_API_KEY / QUERCUS_OPENAI_API_KEY "
|
||||
"(或写入 ~/.quercus/config.toml,权限 0600)后重启面板。"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,135 @@
|
||||
"""FastAPI 应用工厂:聊天 + 待确认清单 + 进度 + 撤销会话 + 快照管理。
|
||||
|
||||
本模块属可选 extra ``core[webui]``(fastapi + uvicorn),不污染 base 依赖。
|
||||
与 P1 风格一致:**不引入 pydantic 模型**,请求体经 ``await request.json()``
|
||||
手工解析;dataclass + 同步代码。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import queue
|
||||
from pathlib import Path
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Request
|
||||
from fastapi.responses import HTMLResponse, StreamingResponse
|
||||
|
||||
from quercus_core.webui.broker import EventBroker
|
||||
from quercus_core.webui.session import WebSession
|
||||
|
||||
_STATIC_DIR = Path(__file__).resolve().parent / "static"
|
||||
|
||||
|
||||
def _project_summary(adapter) -> dict:
|
||||
try:
|
||||
ov = adapter.get_project_overview()
|
||||
return {"name": ov.name, "fps": ov.fps, "timelines": len(ov.timeline_ids)}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def create_app(
|
||||
adapter,
|
||||
provider,
|
||||
*,
|
||||
executor=None,
|
||||
confirm_timeout: float = 300.0,
|
||||
session_log=None,
|
||||
default_confirm: bool | None = None,
|
||||
) -> FastAPI:
|
||||
"""创建 Web 面板应用。adapter 为 HostAdapter 实例,provider 为 LLMProvider。"""
|
||||
broker = EventBroker()
|
||||
session = WebSession(
|
||||
adapter,
|
||||
provider,
|
||||
executor=executor,
|
||||
confirm_timeout=confirm_timeout,
|
||||
event_sink=lambda kind, payload: broker.publish({"kind": kind, **payload}),
|
||||
session_log=session_log,
|
||||
default_confirm=default_confirm,
|
||||
)
|
||||
app = FastAPI(title="Quercus AI 剪辑助手", docs_url=None, redoc_url=None)
|
||||
app.state.session = session
|
||||
app.state.broker = broker
|
||||
|
||||
@app.get("/", response_class=HTMLResponse)
|
||||
def index() -> str:
|
||||
return (_STATIC_DIR / "index.html").read_text(encoding="utf-8")
|
||||
|
||||
@app.get("/api/status")
|
||||
def status() -> dict:
|
||||
return {
|
||||
"ready": True,
|
||||
"busy": session.is_busy,
|
||||
"session_id": session.executor.session_id,
|
||||
"project": _project_summary(adapter),
|
||||
"pending": len(session.pending_batches()),
|
||||
}
|
||||
|
||||
@app.post("/api/chat")
|
||||
async def chat(request: Request) -> dict:
|
||||
body = await request.json()
|
||||
text = (body.get("text") or "").strip()
|
||||
if not text:
|
||||
raise HTTPException(400, "text 不能为空")
|
||||
if session.is_busy:
|
||||
return {"ok": False, "reason": "上一轮对话仍在进行,请稍候"}
|
||||
session.send_message(text)
|
||||
return {"ok": True}
|
||||
|
||||
@app.get("/api/pending")
|
||||
def pending() -> dict:
|
||||
return {"pending": session.pending_batches()}
|
||||
|
||||
@app.post("/api/confirm/{batch_id}")
|
||||
async def confirm(batch_id: str, request: Request) -> dict:
|
||||
body = await request.json()
|
||||
approved = bool(body.get("approved", False))
|
||||
ok = session.confirm_batch(batch_id, approved)
|
||||
if not ok:
|
||||
return {"ok": False, "reason": "批次不存在或已超时"}
|
||||
return {"ok": True, "approved": approved}
|
||||
|
||||
@app.post("/api/undo_session")
|
||||
def undo_session() -> dict:
|
||||
restored = session.undo_session()
|
||||
return {"ok": True, "restored": restored}
|
||||
|
||||
@app.get("/api/snapshots")
|
||||
def snapshots() -> dict:
|
||||
return {"snapshots": session.list_snapshots()}
|
||||
|
||||
@app.post("/api/snapshots/{snapshot_id}/drop")
|
||||
def drop_snapshot(snapshot_id: str) -> dict:
|
||||
try:
|
||||
session.drop_snapshot(snapshot_id)
|
||||
except Exception as exc:
|
||||
raise HTTPException(404, str(exc)) from exc
|
||||
return {"ok": True}
|
||||
|
||||
@app.get("/api/events")
|
||||
def events() -> StreamingResponse:
|
||||
"""SSE 事件流:user/assistant/system 消息、pending、progress、tool_result。
|
||||
|
||||
首条事件为 ``connected``(保证生成器已订阅、前端可感知连接就绪),
|
||||
之后的真实事件经 broker 队列投递。
|
||||
"""
|
||||
q = broker.subscribe()
|
||||
|
||||
def gen():
|
||||
try:
|
||||
yield f"data: {json.dumps({'kind': 'connected'}, ensure_ascii=False)}\n\n"
|
||||
while True:
|
||||
try:
|
||||
item = q.get(timeout=broker.keepalive_seconds)
|
||||
except queue.Empty:
|
||||
yield ": ping\n\n"
|
||||
continue
|
||||
if item is None:
|
||||
break
|
||||
yield f"data: {json.dumps(item, ensure_ascii=False)}\n\n"
|
||||
finally:
|
||||
broker.unsubscribe(q)
|
||||
|
||||
return StreamingResponse(gen(), media_type="text/event-stream")
|
||||
|
||||
return app
|
||||
@@ -0,0 +1,52 @@
|
||||
"""极简 SSE 事件广播器(发布/订阅,每订阅者一个队列)。
|
||||
|
||||
Web 面板与 AgentLoop 的交互用 **SSE + POST**(计划文档 §4 / 任务 3):
|
||||
- 服务端 → 浏览器:SSE 事件流(chat / pending / progress / system);
|
||||
- 浏览器 → 服务端:普通 POST(发消息、裁决、撤销、快照管理)。
|
||||
|
||||
选 SSE 而非 WebSocket:单向推送语义正好匹配,EventSource 自动重连,
|
||||
实现简单可靠。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import queue
|
||||
import threading
|
||||
|
||||
|
||||
class EventBroker:
|
||||
"""线程安全的发布/订阅;``publish`` 把事件投递给所有订阅者队列。"""
|
||||
|
||||
def __init__(self, keepalive_seconds: float = 15.0) -> None:
|
||||
self._queues: list[queue.Queue] = []
|
||||
self._lock = threading.Lock()
|
||||
self._keepalive = keepalive_seconds
|
||||
|
||||
@property
|
||||
def keepalive_seconds(self) -> float:
|
||||
return self._keepalive
|
||||
|
||||
@property
|
||||
def subscriber_count(self) -> int:
|
||||
with self._lock:
|
||||
return len(self._queues)
|
||||
|
||||
def subscribe(self) -> queue.Queue:
|
||||
"""注册订阅,返回接收队列。"""
|
||||
q: queue.Queue = queue.Queue(maxsize=128)
|
||||
with self._lock:
|
||||
self._queues.append(q)
|
||||
return q
|
||||
|
||||
def unsubscribe(self, q: queue.Queue) -> None:
|
||||
with self._lock:
|
||||
if q in self._queues:
|
||||
self._queues.remove(q)
|
||||
|
||||
def publish(self, event: dict) -> None:
|
||||
with self._lock:
|
||||
queues = list(self._queues)
|
||||
for q in queues:
|
||||
try:
|
||||
q.put_nowait(event)
|
||||
except queue.Full:
|
||||
pass # 客户端太慢则丢事件,不阻塞宿主线程
|
||||
@@ -0,0 +1,202 @@
|
||||
"""Web 会话:AgentLoop 的确认门接到 HTTP 端点(待确认清单)。
|
||||
|
||||
一次 Web 会话 = adapter + executor + provider + AgentLoop 的接线:
|
||||
- ``send_message`` 在后台线程跑 AgentLoop(confirm 会阻塞等用户裁决);
|
||||
- ``_confirm`` 回调把批次登记进 ``self._pending`` 并阻塞等待,HTTP 端点
|
||||
(approve/reject)裁决,**超时默认拒绝**;
|
||||
- 事件经 ``event_sink(kind, payload)`` 上抛(Web 层接到 SSE 流)。
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Callable
|
||||
|
||||
from quercus_core.agent.loop import AgentLoop
|
||||
from quercus_core.types import Action, ActionBatch, ToolResult
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Pending:
|
||||
"""一份待裁决的批次:event 被 HTTP 端点 set,approved 记录裁决。"""
|
||||
|
||||
batch: ActionBatch
|
||||
event: threading.Event = field(default_factory=threading.Event)
|
||||
approved: bool = False
|
||||
|
||||
|
||||
class _EventingExecutor:
|
||||
"""把 ToolResult(含 PNG 帧图)转发给事件流,供 Web 面板显示。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
inner,
|
||||
event_sink: Callable[[str, dict], None],
|
||||
) -> None:
|
||||
self._inner = inner
|
||||
self._event_sink = event_sink
|
||||
|
||||
def execute(self, action: Action) -> ToolResult:
|
||||
result = self._inner.execute(action)
|
||||
self._event_sink(
|
||||
"tool_result",
|
||||
{
|
||||
"tool": result.tool,
|
||||
"ok": result.ok,
|
||||
"summary": result.summary,
|
||||
"error": result.error,
|
||||
"images": [base64.b64encode(img).decode("ascii") for img in result.images],
|
||||
},
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
class WebSession:
|
||||
"""一个 Web 会话。``default_confirm`` 仅在测试注入固定裁决时使用。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
adapter,
|
||||
provider,
|
||||
*,
|
||||
executor=None,
|
||||
confirm_timeout: float = 300.0,
|
||||
event_sink: Callable[[str, dict], None] | None = None,
|
||||
session_log=None,
|
||||
default_confirm: bool | None = None,
|
||||
) -> None:
|
||||
self._adapter = adapter
|
||||
self._provider = provider
|
||||
self._executor = executor
|
||||
if self._executor is None:
|
||||
from quercus_core.tools.executor import HostToolExecutor
|
||||
|
||||
self._executor = HostToolExecutor(adapter)
|
||||
self._confirm_timeout = confirm_timeout
|
||||
self._event_sink = event_sink or (lambda kind, payload: None)
|
||||
self._session_log = session_log
|
||||
self._default_confirm = default_confirm
|
||||
self._pending: dict[str, _Pending] = {}
|
||||
self._lock = threading.Lock()
|
||||
self._threads: list[threading.Thread] = []
|
||||
self._finished = threading.Event()
|
||||
self._finished.set()
|
||||
|
||||
# ---- 属性 ----
|
||||
|
||||
@property
|
||||
def adapter(self):
|
||||
return self._adapter
|
||||
|
||||
@property
|
||||
def executor(self):
|
||||
return self._executor
|
||||
|
||||
@property
|
||||
def is_busy(self) -> bool:
|
||||
return not self._finished.is_set()
|
||||
|
||||
def wait_done(self, timeout: float = 5.0) -> bool:
|
||||
"""等待当前对话回合结束(测试用)。返回是否在超时前完成。"""
|
||||
return self._finished.wait(timeout)
|
||||
|
||||
def pending_batches(self) -> list[dict]:
|
||||
"""当前待确认清单(HTTP GET /api/pending 与测试共用)。"""
|
||||
with self._lock:
|
||||
return [self._pending_to_dict(p) for p in self._pending.values()]
|
||||
|
||||
@staticmethod
|
||||
def _pending_to_dict(p: _Pending) -> dict:
|
||||
b = p.batch
|
||||
return {
|
||||
"batch_id": b.id,
|
||||
"label": b.label,
|
||||
"created_at": b.created_at,
|
||||
"actions": [{"tool": a.tool, "params": a.params} for a in b.actions],
|
||||
}
|
||||
|
||||
# ---- 对话 ----
|
||||
|
||||
def send_message(self, text: str) -> None:
|
||||
thread = threading.Thread(
|
||||
target=self._run_agent, args=(text,), daemon=True, name="webui-agent"
|
||||
)
|
||||
self._threads.append(thread)
|
||||
self._finished.clear()
|
||||
thread.start()
|
||||
|
||||
def _run_agent(self, text: str) -> None:
|
||||
self._event_sink("user_message", {"text": text})
|
||||
executor = _EventingExecutor(self._executor, self._event_sink)
|
||||
loop = AgentLoop(
|
||||
provider=self._provider,
|
||||
executor=executor,
|
||||
confirm=self._confirm,
|
||||
session=self._session_log,
|
||||
)
|
||||
try:
|
||||
result = loop.run(text)
|
||||
self._event_sink(
|
||||
"assistant_message",
|
||||
{"text": result.final_text or "(完成,无文本)", "turns": result.turns},
|
||||
)
|
||||
except Exception as exc: # 后台线程不能带异常死亡,转成系统消息
|
||||
self._event_sink("system", {"text": f"对话异常:{exc}"})
|
||||
finally:
|
||||
self._finished.set()
|
||||
|
||||
def _confirm(self, batch: ActionBatch) -> bool:
|
||||
"""AgentLoop 的确认门:登记待确认清单并阻塞等待用户裁决(超时默认拒绝)。"""
|
||||
if self._default_confirm is not None:
|
||||
return self._default_confirm
|
||||
pending = _Pending(batch=batch)
|
||||
with self._lock:
|
||||
self._pending[batch.id] = pending
|
||||
self._event_sink("pending_batch", self._pending_to_dict(pending))
|
||||
decided = pending.event.wait(timeout=self._confirm_timeout)
|
||||
with self._lock:
|
||||
self._pending.pop(batch.id, None)
|
||||
if not decided:
|
||||
self._event_sink("system", {"text": f"待确认批次「{batch.label}」等待超时,默认拒绝"})
|
||||
return False
|
||||
return pending.approved
|
||||
|
||||
def confirm_batch(self, batch_id: str, approved: bool) -> bool:
|
||||
"""HTTP 端点裁决:找到待裁决批次则 set event 返回 True,否则 False。"""
|
||||
with self._lock:
|
||||
pending = self._pending.get(batch_id)
|
||||
if pending is None:
|
||||
return False
|
||||
pending.approved = approved
|
||||
pending.event.set()
|
||||
self._event_sink(
|
||||
"system", {"text": f"批次「{pending.batch.label}」已{'执行' if approved else '拒绝'}"}
|
||||
)
|
||||
return True
|
||||
|
||||
# ---- 撤销 / 快照管理 ----
|
||||
|
||||
def undo_session(self) -> int:
|
||||
n = self._executor.undo_session()
|
||||
self._event_sink(
|
||||
"system",
|
||||
{"text": f"已撤销本会话 {n} 个批次(时间线已恢复原状)" if n else "本会话无可撤销批次"},
|
||||
)
|
||||
return n
|
||||
|
||||
def list_snapshots(self) -> list[dict]:
|
||||
return [
|
||||
{
|
||||
"id": s.id,
|
||||
"session_id": s.session_id,
|
||||
"created_at": s.created_at,
|
||||
"label": s.label,
|
||||
"action_count": s.action_count,
|
||||
}
|
||||
for s in self._executor.list_snapshots()
|
||||
]
|
||||
|
||||
def drop_snapshot(self, snapshot_id: str) -> None:
|
||||
self._executor.drop_snapshot(snapshot_id)
|
||||
self._event_sink("system", {"text": f"已清理快照 {snapshot_id}"})
|
||||
@@ -0,0 +1,221 @@
|
||||
<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="utf-8">
|
||||
<title>Quercus AI 剪辑助手</title>
|
||||
<style>
|
||||
:root { --bg:#f5f6f8; --panel:#fff; --line:#e3e6ea; --text:#1c2430; --muted:#6b7686;
|
||||
--accent:#2f6fed; --ok:#2e9e5b; --bad:#d64545; --sys:#8a6d3b; }
|
||||
* { box-sizing: border-box; }
|
||||
body { margin:0; font-family: -apple-system,"PingFang SC","Microsoft YaHei",sans-serif;
|
||||
background:var(--bg); color:var(--text); }
|
||||
header { display:flex; align-items:center; gap:12px; padding:10px 16px;
|
||||
background:var(--panel); border-bottom:1px solid var(--line); }
|
||||
header h1 { font-size:16px; margin:0; }
|
||||
header .spacer { flex:1; }
|
||||
button { border:1px solid var(--line); background:var(--panel); border-radius:6px;
|
||||
padding:5px 12px; cursor:pointer; font-size:13px; }
|
||||
button:hover { background:#eef1f5; }
|
||||
button.primary { background:var(--accent); color:#fff; border-color:var(--accent); }
|
||||
button.danger { color:var(--bad); }
|
||||
button.approve { background:var(--ok); color:#fff; border-color:var(--ok); }
|
||||
main { display:grid; grid-template-columns: 1fr 320px; gap:12px; padding:12px;
|
||||
height: calc(100vh - 56px); }
|
||||
#chat { background:var(--panel); border:1px solid var(--line); border-radius:8px;
|
||||
overflow-y:auto; padding:12px; display:flex; flex-direction:column; gap:8px; }
|
||||
.msg { max-width:86%; padding:8px 12px; border-radius:10px; font-size:14px;
|
||||
line-height:1.5; white-space:pre-wrap; word-break:break-word; }
|
||||
.msg.user { align-self:flex-end; background:var(--accent); color:#fff; }
|
||||
.msg.assistant { align-self:flex-start; background:#eef2fb; }
|
||||
.msg.system { align-self:center; background:#fdf6e3; color:var(--sys); font-size:13px; }
|
||||
.msg.tool { align-self:flex-start; background:#f0f2f5; color:var(--muted); font-size:13px;
|
||||
max-width:100%; }
|
||||
.msg img { display:block; max-width:100%; margin-top:6px; border-radius:4px; border:1px solid var(--line); }
|
||||
#inputbar { grid-column: 1 / -1; display:flex; gap:8px; }
|
||||
#inputbar input { flex:1; border:1px solid var(--line); border-radius:6px; padding:8px 10px; font-size:14px; }
|
||||
aside { display:flex; flex-direction:column; gap:12px; overflow-y:auto; }
|
||||
.card { background:var(--panel); border:1px solid var(--line); border-radius:8px; padding:10px; }
|
||||
.card h2 { font-size:13px; margin:0 0 8px; color:var(--muted); font-weight:600; }
|
||||
#pending li { font-size:13px; margin-bottom:8px; padding:8px; background:#f7f9fc;
|
||||
border:1px solid var(--line); border-radius:6px; }
|
||||
#pending .batch-actions { font-size:12px; color:var(--muted); margin:4px 0; }
|
||||
#pending .btns { display:flex; gap:6px; }
|
||||
#progress { display:none; }
|
||||
#progress .bar { height:8px; background:#e3e6ea; border-radius:4px; overflow:hidden; }
|
||||
#progress .fill { height:100%; background:var(--accent); width:0%; transition:width .2s; }
|
||||
#progress .txt { font-size:12px; color:var(--muted); margin-top:4px; }
|
||||
#snapshots li { font-size:12px; margin-bottom:6px; display:flex; gap:6px; align-items:center; }
|
||||
#snapshots .label { flex:1; }
|
||||
#status { font-size:12px; color:var(--muted); }
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<header>
|
||||
<h1>Quercus AI 剪辑助手</h1>
|
||||
<span id="status">连接中…</span>
|
||||
<span class="spacer"></span>
|
||||
<button onclick="undoSession()" title="撤销本会话所有已执行的编辑批次">撤销整段会话</button>
|
||||
<button onclick="refreshSnapshots()">刷新快照</button>
|
||||
</header>
|
||||
<main>
|
||||
<section id="chat"></section>
|
||||
<aside>
|
||||
<div class="card" id="pending-card">
|
||||
<h2>待确认清单</h2>
|
||||
<ul id="pending"></ul>
|
||||
</div>
|
||||
<div class="card" id="progress-card">
|
||||
<h2>进度</h2>
|
||||
<div id="progress"><div class="bar"><div class="fill"></div></div><div class="txt"></div></div>
|
||||
<div id="progress-idle">空闲</div>
|
||||
</div>
|
||||
<div class="card" id="snapshots-card">
|
||||
<h2>快照管理</h2>
|
||||
<ul id="snapshots">加载中…</ul>
|
||||
</div>
|
||||
</aside>
|
||||
<div id="inputbar">
|
||||
<input id="prompt" placeholder="描述剪辑意图,如:导入两个素材,把第一个放到轨道 1 开头…"
|
||||
onkeydown="if(event.key==='Enter')send()">
|
||||
<button class="primary" onclick="send()">发送</button>
|
||||
</div>
|
||||
</main>
|
||||
|
||||
<script>
|
||||
"use strict";
|
||||
const $ = (id) => document.getElementById(id);
|
||||
|
||||
function el(tag, cls, text) {
|
||||
const n = document.createElement(tag);
|
||||
if (cls) n.className = cls;
|
||||
if (text !== undefined) n.textContent = text;
|
||||
return n;
|
||||
}
|
||||
|
||||
function appendMsg(role, text, images) {
|
||||
const chat = $("chat");
|
||||
const m = el("div", "msg " + role, text);
|
||||
(images || []).forEach((b64) => {
|
||||
const img = document.createElement("img");
|
||||
img.src = "data:image/png;base64," + b64;
|
||||
m.appendChild(img);
|
||||
});
|
||||
chat.appendChild(m);
|
||||
chat.scrollTop = chat.scrollHeight;
|
||||
return m;
|
||||
}
|
||||
|
||||
// ---- SSE 事件流 ----
|
||||
function connectEvents() {
|
||||
const es = new EventSource("/api/events");
|
||||
es.onopen = () => ($("status").textContent = "已连接");
|
||||
es.onerror = () => ($("status").textContent = "连接断开,重连中…");
|
||||
es.onmessage = (ev) => {
|
||||
let data;
|
||||
try { data = JSON.parse(ev.data); } catch (e) { return; }
|
||||
switch (data.kind) {
|
||||
case "user_message": appendMsg("user", data.text); break;
|
||||
case "assistant_message": appendMsg("assistant", data.text); break;
|
||||
case "system": appendMsg("system", data.text); refreshSnapshots(); break;
|
||||
case "tool_result":
|
||||
appendMsg("tool", (data.ok ? "✔ " : "✘ ") + data.summary, data.images);
|
||||
break;
|
||||
case "pending_batch": addPending(data); break;
|
||||
case "progress": setProgress(data.done, data.total, data.message); break;
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
// ---- 待确认清单 ----
|
||||
const pendingByKey = {};
|
||||
function addPending(p) {
|
||||
pendingByKey[p.batch_id] = p;
|
||||
renderPending();
|
||||
}
|
||||
function renderPending() {
|
||||
const ul = $("pending");
|
||||
ul.innerHTML = "";
|
||||
Object.values(pendingByKey).forEach((p) => {
|
||||
const li = el("li");
|
||||
li.appendChild(el("strong", null, p.label));
|
||||
const acts = el("div", "batch-actions");
|
||||
acts.textContent = p.actions.map((a) => a.tool + (a.params ? " " + JSON.stringify(a.params) : "")).join(";");
|
||||
li.appendChild(acts);
|
||||
const btns = el("div", "btns");
|
||||
const okBtn = el("button", "approve", "执行");
|
||||
okBtn.onclick = () => decide(p.batch_id, true);
|
||||
const noBtn = el("button", "", "拒绝");
|
||||
noBtn.onclick = () => decide(p.batch_id, false);
|
||||
btns.appendChild(okBtn); btns.appendChild(noBtn);
|
||||
li.appendChild(btns);
|
||||
ul.appendChild(li);
|
||||
});
|
||||
if (!Object.keys(pendingByKey).length) ul.textContent = "暂无待确认批次";
|
||||
}
|
||||
async function decide(batchId, approved) {
|
||||
await fetch("/api/confirm/" + batchId, {
|
||||
method: "POST", headers: {"Content-Type": "application/json"},
|
||||
body: JSON.stringify({approved}),
|
||||
});
|
||||
delete pendingByKey[batchId];
|
||||
renderPending();
|
||||
}
|
||||
|
||||
// ---- 进度 ----
|
||||
function setProgress(done, total, message) {
|
||||
$("progress").style.display = "block";
|
||||
$("progress-idle").style.display = "none";
|
||||
const pct = total ? Math.round((done / total) * 100) : 0;
|
||||
$("progress").querySelector(".fill").style.width = pct + "%";
|
||||
$("progress").querySelector(".txt").textContent = message || done + "/" + total;
|
||||
if (done >= total) setTimeout(() => { $("progress").style.display = "none"; $("progress-idle").style.display = "block"; }, 800);
|
||||
}
|
||||
|
||||
// ---- 交互 ----
|
||||
async function send() {
|
||||
const input = $("prompt");
|
||||
const text = input.value.trim();
|
||||
if (!text) return;
|
||||
input.value = "";
|
||||
const r = await fetch("/api/chat", {
|
||||
method: "POST", headers: {"Content-Type": "application/json"},
|
||||
body: JSON.stringify({text}),
|
||||
});
|
||||
const data = await r.json();
|
||||
if (!data.ok) appendMsg("system", data.reason || "发送失败");
|
||||
}
|
||||
|
||||
async function undoSession() {
|
||||
const r = await fetch("/api/undo_session", {method: "POST"});
|
||||
const data = await r.json();
|
||||
appendMsg("system", "已撤销本会话 " + data.restored + " 个批次");
|
||||
refreshSnapshots();
|
||||
}
|
||||
|
||||
async function refreshSnapshots() {
|
||||
const r = await fetch("/api/snapshots");
|
||||
const data = await r.json();
|
||||
const ul = $("snapshots");
|
||||
ul.innerHTML = "";
|
||||
if (!data.snapshots.length) { ul.textContent = "暂无快照"; return; }
|
||||
data.snapshots.forEach((s) => {
|
||||
const li = el("li");
|
||||
const label = el("span", "label", (s.label || "快照") + " · " + new Date(s.created_at).toLocaleString() + " · " + s.action_count + " 动作");
|
||||
const drop = el("button", "danger", "清理");
|
||||
drop.onclick = async () => {
|
||||
await fetch("/api/snapshots/" + s.id + "/drop", {method: "POST"});
|
||||
refreshSnapshots();
|
||||
};
|
||||
li.appendChild(label); li.appendChild(drop);
|
||||
ul.appendChild(li);
|
||||
});
|
||||
}
|
||||
|
||||
connectEvents();
|
||||
refreshSnapshots();
|
||||
fetch("/api/status").then(r => r.json()).then((s) => {
|
||||
$("status").textContent = "工程:" + (s.project.name || "-") + " · " + (s.project.timelines || 0) + " 时间线";
|
||||
});
|
||||
</script>
|
||||
</body>
|
||||
</html>
|
||||
Reference in New Issue
Block a user