Files
projectAIpopular/router_system/worker.py
T

216 lines
8.9 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""WorkerLoop —— 小模型(本地 llama.cpp)的"实现/自验证"循环(端云协同的执行者)。
流程(对齐《实现方案_v2》5.1 T5 / 4.4):
读 brief+当前步 -> 模型生成工件 -> 接地验证(D4 分层)
-> 通过:写 progress(done)
-> 失败:自修 <= max_fix_attempts 次(把验证错误回喂重新生成)
-> 仍失败:写 issue(增量、带锚点)
- generate 为可注入的文本生成器(真实为 llama-server 端点;测试用假实现)。
- 工件落盘:runs/<request_id>/artifacts/<step>.py(由 pipeline 负责写盘,本模块只产出文本)。
"""
from __future__ import annotations
import time
from dataclasses import dataclass, field
from typing import Any, Awaitable, Callable, Dict, List, Optional
from .verifier import Verifier, detect_artifact_language, extract_code_block
from .workspace import Workspace, build_anchor
# 按领域推断默认工件扩展名
_DOMAIN_EXT = {
"code": ".py",
"math": ".md",
"legal": ".md",
"medical": ".md",
"finance": ".md",
"life": ".md",
"education": ".md",
"general": ".md",
}
def artifact_name_for(step_id: str, domain: str) -> str:
"""为 step 生成工件文件名。"""
ext = _DOMAIN_EXT.get(domain, ".md")
return f"{step_id}{ext}"
@dataclass
class StepOutcome:
"""单步执行结果。"""
step_id: str
status: str # done | issue
summary: str = ""
model_used: str = "local"
attempts: int = 0
issue_id: Optional[str] = None
artifact_name: Optional[str] = None
artifact_text: str = ""
details: List[str] = field(default_factory=list)
class WorkerLoop:
"""小模型 Worker:实现 -> 验证 -> 自修 -> issue。"""
def __init__(
self,
generate: Callable[[str], Awaitable[str]],
verifier: Optional[Verifier] = None,
kb: Any = None,
max_fix_attempts: int = 2,
model_used: str = "local-llama",
):
self.generate = generate
self.verifier = verifier or Verifier()
self.kb = kb
self.max_fix_attempts = max_fix_attempts
self.model_used = model_used
async def direct_answer(self, query: str) -> str:
"""快路径直答:让 Worker 直接生成用户回答(非 JSON、无围栏)。"""
prompt = ("请直接回答下面这个问题,输出对用户有用的正文"
"(不要输出 JSON,不要加代码块围栏)。问题:" + query)
return await self.generate(prompt)
def _domain_from(self, ws: Workspace) -> str:
tags = (ws.get("brief") or {}).get("tags") or []
for t in tags:
if t != "safety":
return t
return "general"
async def run_step(self, ws: Workspace, step_id: str,
existing_artifact: str = "", hint: str = "") -> StepOutcome:
"""执行单个 step。existing_artifact 为该步当前已有工件全文;hint 为 Architect 裁决提示。"""
domain = self._domain_from(ws)
brief = ws.get("brief") or {}
plan = brief.get("plan") or []
step_def = next((p for p in plan if p.get("id") == step_id), {})
done_criteria = step_def.get("done_criteria", "")
artifact_name = artifact_name_for(step_id, domain)
current = existing_artifact
details: List[str] = []
for attempt in range(1, self.max_fix_attempts + 1):
prompt = self._build_prompt(ws, step_id, current, attempt, done_criteria, hint)
out = await self.generate(prompt)
if domain == "code":
candidate = extract_code_block(out)
else:
candidate = out.strip()
details.append(f"attempt{attempt}: 生成 {len(candidate)} 字符")
passed, v_details = self.verifier.verify(
domain, artifact_name, candidate, ws["query"], kb=self.kb)
details.extend(f" - {d}" for d in v_details)
if passed:
# 写回交流文本:progress(done) + 摘要
ws.add_progress(step_id, "done", f"步骤完成({attempt} 次尝试)",
artifact=build_anchor(artifact_name, 1))
return StepOutcome(
step_id=step_id, status="done",
summary=f"步骤完成({attempt} 次尝试)",
model_used=self.model_used, attempts=attempt,
artifact_name=artifact_name, artifact_text=candidate,
details=details,
)
# 未通过:带错误反馈重新生成(自修)
current = candidate
feedback = "".join(v_details)
details.append(f"attempt{attempt} 未通过,进入自修")
# 全部尝试失败 -> 写 issue
anchor = build_anchor(artifact_name, 1, 30)
iid = ws.add_issue(
step=step_id,
anchor=anchor,
observed=f"验证未通过:{''.join(d for d in details if d.startswith(' - ')) or '未知'}",
expected=done_criteria or "满足该步 done_criteria",
tried=f"已自修 {self.max_fix_attempts} 次",
ask="请裁决该步的实现方向或提供兜底实现",
)
return StepOutcome(
step_id=step_id, status="issue", summary="未能通过验证,已上报 issue",
model_used=self.model_used, attempts=self.max_fix_attempts,
issue_id=iid, artifact_name=artifact_name, artifact_text=current,
details=details,
)
def _build_prompt(self, ws: Workspace, step_id: str, current: str,
attempt: int, done_criteria: str, hint: str = "") -> str:
base = ws.render_for_worker(step_id, artifact_text=current or None)
if attempt > 1:
base += (
"\n\n[注意] 上次生成的工件未通过接地验证。请修正以下问题后重新输出"
f"完整工件。本次为第 {attempt} 次尝试。"
)
if hint:
base += "\n\n[架构师裁决] " + hint
return base
def build_worker(cfg: Dict[str, Any], kb: Any = None,
generate: Optional[Callable[[str], Awaitable[str]]] = None) -> WorkerLoop:
"""cfg 为 config.worker 段。generate 缺省时按 backend 选择:
mock(零运行时演示)| openai/api(任意 OpenAI 兼容端点,如 Ollama/vLLM|
llama_server(内置本地 llama-server)。"""
backend = cfg.get("backend", "llama_server")
if generate is None:
if backend == "mock":
generate = _mock_generate()
elif backend in ("openai", "api"):
generate = _make_llama_generate(
cfg, default_base_url=cfg.get("base_url") or "http://127.0.0.1:11434/v1")
else:
generate = _make_llama_generate(cfg)
verifier = Verifier(code_timeout_s=float(cfg.get("code_timeout_s", 10)))
return WorkerLoop(
generate=generate,
verifier=verifier,
kb=kb,
max_fix_attempts=int(cfg.get("max_fix_attempts", 2)),
model_used=cfg.get("backend", "llama_server"),
)
def _mock_generate() -> Callable[[str], Awaitable[str]]:
"""零运行时 mock 生成器:返回一段确定性文本(演示/测试,不连真实模型)。"""
async def _gen(prompt: str) -> str:
return ("mock worker)以下是对当前步骤的实现说明:"
"步骤已完成,内容足够长且非占位,可供接地验证通过。")
return _gen
def _make_llama_generate(cfg: Dict[str, Any],
default_base_url: Optional[str] = None) -> Callable[[str], Awaitable[str]]:
"""返回调用本地 OpenAI 兼容端点(llama-server / Ollama / vLLM)的生成器。"""
if default_base_url is None:
default_base_url = f"http://127.0.0.1:{cfg.get('port', 8901)}/v1"
base_url = cfg.get("base_url") or default_base_url
model = cfg.get("model") or "local"
temperature = float(cfg.get("temperature", 0.3))
timeout_s = float(cfg.get("per_step_timeout_s", 300))
async def _gen(prompt: str) -> str:
try:
import httpx
async with httpx.AsyncClient(timeout=timeout_s) as client:
resp = await client.post(
f"{base_url}/chat/completions",
json={"model": model, "messages": [{"role": "user", "content": prompt}],
"temperature": temperature, "max_tokens": 4096},
)
resp.raise_for_status()
return resp.json()["choices"][0]["message"]["content"]
except Exception as e: # noqa: BLE001
# 连不上本地模型 -> 优雅降级(不抛 500),提示用户检查模型端点
return ("(本地降级)无法连接本地模型端点,未能生成该步骤内容。"
f"请检查模型后端配置或启动服务。错误:{type(e).__name__}")
return _gen