优化文档与测试脚本

This commit is contained in:
lulijing
2026-03-23 20:26:13 +08:00
parent 3ddc2e1ff4
commit 2d50522e68
14 changed files with 654 additions and 188 deletions

View File

@@ -7,18 +7,130 @@
from __future__ import annotations
import copy
import json
import re
import sys
import time
from datetime import datetime, timezone
from pathlib import Path
# 确保 src 在路径中
sys.path.insert(0, str(Path(__file__).resolve().parent / "src"))
# 可视化测试指令日志:保存到项目下的 log 目录(纯文本)
PLAYGROUND_LOG_DIR = Path(__file__).resolve().parent / "log"
PLAYGROUND_LOG_FILE = PLAYGROUND_LOG_DIR / "playground_instructions.log"
def _append_playground_log(
instruction: str,
success: bool,
fast_path: bool = False,
intents: list[str] | None = None,
entities: dict | None = None,
error: str | None = None,
traceback_text: str | None = None,
timing: dict | None = None,
layer2_meta: dict | None = None,
layer2_prompt_preview: str | None = None,
layer3_json: dict | None = None,
layer4_ascii: str | None = None,
rag_context: dict | None = None,
tool_call_log: list | None = None,
layer1_raw_output: dict | None = None,
) -> None:
"""将本次测试的完整内容以纯文本追加到 log 目录下。"""
PLAYGROUND_LOG_DIR.mkdir(parents=True, exist_ok=True)
ts = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S UTC")
lines = [
"",
"=" * 60,
f" {ts}",
"=" * 60,
"",
"【指令】",
instruction,
"",
"【结果】",
f" 成功: {success}",
f" Fast-Path: {fast_path}",
f" 意图 intents: {intents or []}",
]
# 实体:确保可序列化为合法 JSONRouter 偶发畸形键时仍可写日志)
try:
ent = entities or {}
if isinstance(ent, dict):
ent = {str(k): v for k, v in ent.items()}
lines.append(f" 实体 entities: {json.dumps(ent, ensure_ascii=False, indent=2)}")
except (TypeError, ValueError):
lines.append(f" 实体 entities: (无法序列化) {repr(entities)[:200]}")
lines.append("")
if error:
lines.extend(["【错误】", error, ""])
if traceback_text:
lines.append("【堆栈】")
lines.append(traceback_text)
lines.append("")
if layer1_raw_output is not None:
lines.append("【Layer 1 原始路由输出raw_llm_output")
lines.append(json.dumps(layer1_raw_output, ensure_ascii=False, indent=2))
lines.append("")
if timing:
lines.append("【各环节耗时 (ms)】")
for name, ms in timing.items():
lines.append(f" {name}: {ms:.0f}")
lines.append(f" 总计: {sum(timing.values()):.0f}")
lines.append("")
if layer2_meta:
lines.append("【Layer 2 Composer】")
lines.append(f" 选中节点: {layer2_meta.get('selected_actions', [])}")
lines.append(f" 基准点坐标: {json.dumps(layer2_meta.get('base_location_coords', {}), ensure_ascii=False)}")
if layer2_prompt_preview:
lines.append(" System Prompt 预览:")
for line in layer2_prompt_preview.strip().split("\n")[:30]:
lines.append(" " + line)
if layer2_prompt_preview.count("\n") >= 30:
lines.append(" ...")
lines.append("")
if tool_call_log:
lines.append("【Tool Call 日志】")
for entry in tool_call_log:
lines.append(f"{entry.get('round', '?')}{entry.get('tool', '')}")
lines.append(f" {json.dumps(entry.get('arguments', {}), ensure_ascii=False)}")
lines.append(f" 结果: {json.dumps(entry.get('result', {}), ensure_ascii=False)}")
lines.append("")
if rag_context:
lines.append("【RAG 检索】")
if rag_context.get("map_context"):
lines.append(" 地图: " + (rag_context["map_context"][:500] + "..." if len(rag_context.get("map_context", "")) > 500 else rag_context.get("map_context", "")))
if rag_context.get("rule_context"):
lines.append(" 规则: " + (rag_context["rule_context"][:300] + "..." if len(rag_context.get("rule_context", "")) > 300 else rag_context.get("rule_context", "")))
lines.append("")
if layer3_json is not None:
lines.append("【Layer 3 行为树 JSON后处理后的执行树")
lines.append(json.dumps(layer3_json, ensure_ascii=False, indent=2))
lines.append("")
if layer4_ascii:
lines.append("【Layer 4 执行树 ASCII】")
# 去掉终端颜色码,便于在 .log 文本中阅读
layer4_plain = re.sub(r"\x1b\[[0-9;]*m", "", layer4_ascii)
lines.append(layer4_plain)
lines.append("")
lines.append("")
with open(PLAYGROUND_LOG_FILE, "a", encoding="utf-8") as f:
f.write("\n".join(lines))
import streamlit as st
from drone_planning.core.blackboard import DroneStateBlackboard
from drone_planning.execution.tree_wrapper import parse_json_to_tree, tree_to_ascii, wrap_and_build_tree
from drone_planning.execution.tree_wrapper import (
enforce_search_object_detected,
parse_json_to_tree,
tree_to_ascii,
wrap_and_build_tree,
)
from drone_planning.pipeline.composer import _load_node_schema, build_system_prompt
from drone_planning.pipeline.planner import plan
from drone_planning.pipeline.router import route
@@ -59,177 +171,221 @@ with st.sidebar:
bb.is_in_air = is_in_air
st.info(f"当前: {'已起飞' if is_in_air else '在地面'}")
# ---------------------------------------------------------------------------
# 防重复点击:每次测试独立,不并发
# ---------------------------------------------------------------------------
if "run_in_progress" not in st.session_state:
st.session_state.run_in_progress = False
# ---------------------------------------------------------------------------
# 主界面:输入与执行
# ---------------------------------------------------------------------------
instruction = st.text_area("输入自然语言指令", height=80, placeholder="例如:飞到大门然后拍照、起飞、飞到广场然后搜索汽车")
instruction = st.text_area("输入自然语言指令", height=80, placeholder="例如:飞到大门然后拍照、起飞、飞到广场然后搜索汽车", key="instruction_input")
run_text = instruction.strip() if instruction else ""
if run_text:
st.divider()
if st.button("执行规划", type="primary", key="run_plan"):
# 再次同步 blackboard用户可能在点击前改了侧边栏
bb = DroneStateBlackboard()
bb.is_in_air = is_in_air
# 执行中时禁用按钮,避免重复点击导致并发请求
if st.session_state.run_in_progress:
st.warning("⏳ 正在执行规划,请等待完成...")
st.caption("请勿重复点击或刷新页面,当前请求完成后将自动更新。")
else:
if st.button("执行规划", type="primary", key="run_plan"):
# 立即标记为执行中,防止并发请求
st.session_state.run_in_progress = True
# 本轮测试开始前,显式同步 Blackboard确保前后测试互不影响
DroneStateBlackboard().is_in_air = is_in_air
layer1_result = None
rag_context = None
layer2_prompt = None
layer2_meta = None
tool_call_log: list[dict] = []
layer3_json = None
layer4_ascii = None
fast_path = False
error_msg = None
timing: dict[str, float] = {}
layer1_result = None
rag_context = None
layer2_prompt = None
layer2_meta = None
tool_call_log: list[dict] = []
layer3_json = None
layer3_json_for_display = None # 后处理前的拷贝,仅用于 Layer 3 展示
layer4_ascii = None
fast_path = False
error_msg = None
traceback_text: str | None = None
timing: dict[str, float] = {}
try:
# Layer 1: Router
t0 = time.perf_counter()
layer1_result = route(run_text)
timing["Layer1_Router"] = (time.perf_counter() - t0) * 1000
fast_path = layer1_result.is_fast_path
with st.spinner("正在执行规划,请稍候..."):
try:
# Layer 1: Router
t0 = time.perf_counter()
layer1_result = route(run_text)
timing["Layer1_Router"] = (time.perf_counter() - t0) * 1000
fast_path = layer1_result.is_fast_path
if fast_path:
st.success("✅ Layer 1 触发 Fast-Path 短路,后续层不执行")
else:
# RAG 检索(阶段二)
t0 = time.perf_counter()
retriever = RAGRetriever()
rag_context = retriever.retrieve_context(
intents=layer1_result.intents,
entities=layer1_result.entities,
user_text=run_text,
)
timing["RAG_检索"] = (time.perf_counter() - t0) * 1000
# Layer 2: Composer
t0 = time.perf_counter()
schema_json = _load_node_schema()
layer2_prompt, layer2_meta = build_system_prompt(
intents=layer1_result.intents,
entities=layer1_result.entities,
schema_json=schema_json,
rag_context=rag_context,
)
timing["Layer2_Composer"] = (time.perf_counter() - t0) * 1000
# Layer 3: Planner支持 Tool 循环)
t0 = time.perf_counter()
layer3_json = plan(
system_prompt=layer2_prompt,
user_text=run_text,
tool_call_log=tool_call_log,
)
timing["Layer3_Planner"] = (time.perf_counter() - t0) * 1000
# Layer 4: 解析 + 包装(会读取 blackboard.is_in_air
t0 = time.perf_counter()
root_node = layer3_json.get("root") or layer3_json
business_tree = parse_json_to_tree(root_node)
final_tree = wrap_and_build_tree(business_tree)
layer4_ascii = tree_to_ascii(final_tree)
timing["Layer4_Execution"] = (time.perf_counter() - t0) * 1000
except Exception as e:
error_msg = str(e)
import traceback
st.error(f"❌ 执行出错: {error_msg}")
st.code(traceback.format_exc(), language="text")
# -----------------------------------------------------------------------
# 各环节耗时
# -----------------------------------------------------------------------
if timing:
st.subheader("⏱️ 各环节耗时 (ms)")
cols = st.columns(len(timing))
for i, (name, ms) in enumerate(timing.items()):
with cols[i]:
st.metric(name, f"{ms:.0f} ms")
st.caption(f"总计: {sum(timing.values()):.0f} ms")
# -----------------------------------------------------------------------
# 分步展示
# -----------------------------------------------------------------------
st.subheader("📋 各层输出")
# Layer 1
with st.expander("**Layer 1 (Router)** - 意图与实体", expanded=True):
if layer1_result:
col1, col2, col3 = st.columns(3)
with col1:
st.metric("intents", ", ".join(layer1_result.intents) or "-")
with col2:
ent_str = json.dumps(layer1_result.entities, ensure_ascii=False)
st.metric("entities", ent_str[:60] + "..." if len(ent_str) > 60 else ent_str or "-")
with col3:
st.metric("Fast-Path", "" if fast_path else "")
st.json(layer1_result.entities)
else:
st.warning("Layer 1 未执行或出错")
if not fast_path:
# RAG 检索结果Layer 1 与 Layer 2 之间)
with st.expander("**RAG 检索结果** - 地图坐标、规则、Few-shot 示例", expanded=True):
if rag_context:
col1, col2 = st.columns(2)
with col1:
st.markdown("**地图坐标 (map_context)**")
if rag_context.get("map_context"):
st.code(rag_context["map_context"], language=None)
else:
st.caption("无匹配")
with col2:
st.markdown("**规则约束 (rule_context)**")
if rag_context.get("rule_context"):
st.code(rag_context["rule_context"], language=None)
else:
st.caption("无匹配")
st.markdown("**Few-shot 示例**")
if rag_context.get("few_shot_examples"):
for i, ex in enumerate(rag_context["few_shot_examples"]):
st.json({"instruction": ex.get("instruction"), "tree_json": ex.get("tree_json")})
if fast_path:
st.success("✅ Layer 1 触发 Fast-Path 短路,后续层不执行")
else:
st.caption("无相似示例")
# RAG 检索(阶段二)
t0 = time.perf_counter()
retriever = RAGRetriever()
rag_context = retriever.retrieve_context(
intents=layer1_result.intents,
entities=layer1_result.entities,
user_text=run_text,
)
timing["RAG_检索"] = (time.perf_counter() - t0) * 1000
# Layer 2: Composer
t0 = time.perf_counter()
schema_json = _load_node_schema()
layer2_prompt, layer2_meta = build_system_prompt(
intents=layer1_result.intents,
entities=layer1_result.entities,
schema_json=schema_json,
rag_context=rag_context,
)
timing["Layer2_Composer"] = (time.perf_counter() - t0) * 1000
# Layer 3: Planner支持 Tool 循环)
t0 = time.perf_counter()
layer3_json = plan(
system_prompt=layer2_prompt,
user_text=run_text,
tool_call_log=tool_call_log,
)
timing["Layer3_Planner"] = (time.perf_counter() - t0) * 1000
# Layer 4: 后处理 + 解析 + 包装(会读取 blackboard.is_in_air
t0 = time.perf_counter()
# 先拷贝一份供 Layer 3 展示「Planner 原生」;后处理在原树上做,保证 Layer 4 / log 与 API 一致
layer3_json_for_display = copy.deepcopy(layer3_json)
root_node = layer3_json.get("root") or layer3_json
if isinstance(root_node, dict):
enforce_search_object_detected(root_node)
business_tree = parse_json_to_tree(root_node)
final_tree = wrap_and_build_tree(business_tree)
layer4_ascii = tree_to_ascii(final_tree)
timing["Layer4_Execution"] = (time.perf_counter() - t0) * 1000
except Exception as e:
error_msg = str(e)
import traceback
traceback_text = traceback.format_exc()
st.error(f"❌ 执行出错: {error_msg}")
st.code(traceback_text, language="text")
finally:
# 无论成功或异常,都解除执行中状态,允许下一次测试
st.session_state.run_in_progress = False
# 写入 log
_append_playground_log(
instruction=run_text,
success=(error_msg is None),
fast_path=fast_path,
intents=list(layer1_result.intents) if layer1_result else None,
entities=dict(layer1_result.entities) if layer1_result else None,
error=error_msg,
traceback_text=traceback_text if error_msg else None,
timing=timing or None,
layer2_meta=layer2_meta,
layer2_prompt_preview=layer2_prompt[:3000] if layer2_prompt else None,
layer3_json=layer3_json,
layer4_ascii=layer4_ascii,
rag_context=rag_context,
tool_call_log=tool_call_log if tool_call_log else None,
layer1_raw_output=layer1_result.raw_llm_output if layer1_result else None,
)
# -----------------------------------------------------------------------
# 各环节耗时
# -----------------------------------------------------------------------
if timing:
st.subheader("⏱️ 各环节耗时 (ms)")
cols = st.columns(len(timing))
for i, (name, ms) in enumerate(timing.items()):
with cols[i]:
st.metric(name, f"{ms:.0f} ms")
st.caption(f"总计: {sum(timing.values()):.0f} ms")
# -----------------------------------------------------------------------
# 分步展示
# -----------------------------------------------------------------------
st.subheader("📋 各层输出")
# Layer 1
with st.expander("**Layer 1 (Router)** - 意图与实体", expanded=True):
if layer1_result:
col1, col2, col3 = st.columns(3)
with col1:
st.metric("intents", ", ".join(layer1_result.intents) or "-")
with col2:
ent_str = json.dumps(layer1_result.entities, ensure_ascii=False)
st.metric("entities", ent_str[:60] + "..." if len(ent_str) > 60 else ent_str or "-")
with col3:
st.metric("Fast-Path", "" if fast_path else "")
st.json(layer1_result.entities)
else:
st.caption("RAG 未执行")
st.warning("Layer 1 未执行或出错")
# Layer 2
with st.expander("**Layer 2 (Composer)** - 精简 Prompt 与选中节点", expanded=True):
if layer2_meta:
st.metric("Prompt 长度", f"{layer2_meta.get('prompt_length', 0)} 字符")
st.metric("选中节点", ", ".join(layer2_meta.get("selected_actions", [])))
st.metric("基准点坐标", json.dumps(layer2_meta.get("base_location_coords", {}), ensure_ascii=False))
if layer2_prompt:
preview = layer2_prompt[:2000] + ("..." if len(layer2_prompt) > 2000 else "")
st.code(preview, language="text")
if not fast_path:
# RAG 检索结果Layer 1 与 Layer 2 之间)
with st.expander("**RAG 检索结果** - 地图坐标、规则、Few-shot 示例", expanded=True):
if rag_context:
col1, col2 = st.columns(2)
with col1:
st.markdown("**地图坐标 (map_context)**")
if rag_context.get("map_context"):
st.code(rag_context["map_context"], language=None)
else:
st.caption("无匹配")
with col2:
st.markdown("**规则约束 (rule_context)**")
if rag_context.get("rule_context"):
st.code(rag_context["rule_context"], language=None)
else:
st.caption("无匹配")
st.markdown("**Few-shot 示例**")
if rag_context.get("few_shot_examples"):
for i, ex in enumerate(rag_context["few_shot_examples"]):
st.json({"instruction": ex.get("instruction"), "tree_json": ex.get("tree_json")})
else:
st.caption("无相似示例")
else:
st.caption("RAG 未执行")
# 坐标计算与 Tool Call 日志(仅 LLM 通过 mcp 工具计算)
rel_descs = (rag_context or {}).get("relative_descriptions") or []
has_any = bool(rel_descs) or bool(tool_call_log)
with st.expander("**坐标计算与 Tool Call 日志**", expanded=has_any):
if rel_descs:
st.markdown("**RAG 提供的相对描述**(需 LLM 调用 calculate_relative_coordinate")
for rd in rel_descs:
st.json(rd)
if tool_call_log:
st.markdown("**LLM Tool Call**(大模型调用 mcp 计算)")
for entry in tool_call_log:
st.markdown(f"{entry.get('round', '?')} 次 `{entry.get('tool', '')}`")
st.json({"arguments": entry.get("arguments"), "result": entry.get("result")})
if not rel_descs and not tool_call_log:
st.caption("无相对描述(指令仅含绝对地点时,直接使用 base_location_coords")
# Layer 2
with st.expander("**Layer 2 (Composer)** - 精简 Prompt 与选中节点", expanded=True):
if layer2_meta:
st.metric("Prompt 长度", f"{layer2_meta.get('prompt_length', 0)} 字符")
st.metric("选中节点", ", ".join(layer2_meta.get("selected_actions", [])))
st.metric("基准点坐标", json.dumps(layer2_meta.get("base_location_coords", {}), ensure_ascii=False))
if layer2_prompt:
preview = layer2_prompt[:2000] + ("..." if len(layer2_prompt) > 2000 else "")
st.code(preview, language="text")
# Layer 3
with st.expander("**Layer 3 (Planner)** - 原生业务树 JSON", expanded=True):
if layer3_json:
st.json(layer3_json)
else:
st.warning("Layer 3 未执行")
# 坐标计算与 Tool Call 日志(仅 LLM 通过 mcp 工具计算)
rel_descs = (rag_context or {}).get("relative_descriptions") or []
has_any = bool(rel_descs) or bool(tool_call_log)
with st.expander("**坐标计算与 Tool Call 日志**", expanded=has_any):
if rel_descs:
st.markdown("**RAG 提供的相对描述**(需 LLM 调用 calculate_relative_coordinate")
for rd in rel_descs:
st.json(rd)
if tool_call_log:
st.markdown("**LLM Tool Call**(大模型调用 mcp 计算)")
for entry in tool_call_log:
st.markdown(f"{entry.get('round', '?')} 次 `{entry.get('tool', '')}`")
st.json({"arguments": entry.get("arguments"), "result": entry.get("result")})
if not rel_descs and not tool_call_log:
st.caption("无相对描述(指令仅含绝对地点时,直接使用 base_location_coords")
# Layer 4
with st.expander("**Layer 4 (Execution)** - 安全包装后的 py_trees ASCII 结构", expanded=True):
if layer4_ascii:
st.text("当 is_in_air=False 时,此处应看到 [Safe_Execution] -> [SystemCheck] -> [Takeoff] -> 业务树")
st.code(layer4_ascii, language=None)
else:
st.warning("Layer 4 未执行")
# Layer 3展示 Planner 原始输出(后处理前的 JSON
with st.expander("**Layer 3 (Planner)** - 原生业务树 JSON后处理前", expanded=True):
if layer3_json_for_display is not None:
st.json(layer3_json_for_display)
elif layer3_json:
st.json(layer3_json)
else:
st.warning("Layer 3 未执行")
# Layer 4
with st.expander("**Layer 4 (Execution)** - 安全包装后的 py_trees ASCII 结构", expanded=True):
if layer4_ascii:
st.text("当 is_in_air=False 时,此处应看到 [Safe_Execution] -> [SystemCheck] -> [Takeoff] -> 业务树")
st.code(layer4_ascii, language=None)
else:
st.warning("Layer 4 未执行")