229 lines
10 KiB
Python
229 lines
10 KiB
Python
"""ArchitectClient —— 大模型(API)客户端,端云协同的"决策/终审"角色。
|
||
|
||
职责(对齐《实现方案_v2》5.1 / 6.3):
|
||
- brief(query):开局任务分析 -> 生成交流文本的 brief(goal/constraints/acceptance/plan/tags)
|
||
- decide(ws):读 issues 等 -> 输出裁决(reply + patch_plan 修订计划)
|
||
- final_review(ws):终审 -> {verdict: done|fix, issues: [...]}
|
||
|
||
工程约束(D7 / D8 / D9 / D11):
|
||
- 走 OpenAI 兼容 /chat/completions;response_format={"type":"json_object"},prompt 内嵌 schema 描述。
|
||
- Architect 输入永不包含工件全文:只传 Workspace 渲染出的 meta+issues+decisions+锚点片段(render_for_architect)。
|
||
- 所有结构化输出解析为 JSON;失败把错误回喂重写一次,仍失败抛 ArchitectError(由编排层降级,禁止带病继续)。
|
||
- token 计量回写 ws.meta.budget;调用前先查预算,触顶抛 ArchitectCircuitBreaker(D6)。
|
||
- httpx 惰性导入(零顶层依赖,对齐仓库既有 APIExpert 模式);client/transport 可注入,测试用 httpx.MockTransport(D11)。
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import json
|
||
from typing import Any, Callable, Dict, List, Optional
|
||
|
||
from .workspace import Workspace
|
||
|
||
# brief 的 JSON schema(描述性提示,约束模型输出结构)
|
||
_BRIEF_SCHEMA_HINT = {
|
||
"goal": "string(<=500字)",
|
||
"constraints": "string[](<=8条)",
|
||
"tags": "string[](code/math/legal/medical/finance/life/education/general/safety 之一)",
|
||
"acceptance": "list[{id, check(string), machine_checkable(bool)}]",
|
||
"plan": "list[{id, task(string<=300字), deps(string[]), done_criteria(string)}](<=5步, 有依赖序)",
|
||
}
|
||
|
||
_DECIDE_SCHEMA_HINT = {
|
||
"reply": "string(<=600字)",
|
||
"patch_plan": "list[{id, task(string)}]",
|
||
}
|
||
|
||
_REVIEW_SCHEMA_HINT = {
|
||
"verdict": "enum(done|fix)",
|
||
"notes": "string(<=300字)",
|
||
"fix_issues": "list[string]",
|
||
}
|
||
|
||
_SYSTEM_PROMPT = (
|
||
"你是任务分析架构师。你的输入是不含工件全文的协作摘要(交流文本),"
|
||
"你的输出必须是合法 JSON 对象(不要用 markdown 代码块包裹)。"
|
||
)
|
||
|
||
|
||
class ArchitectError(RuntimeError):
|
||
"""Architect 调用失败(网络/超时/JSON 解析失败/服务错误)。"""
|
||
|
||
|
||
class ArchitectCircuitBreaker(RuntimeError):
|
||
"""预算熔断(D6):api_token_cap / rounds_cap 触顶。"""
|
||
|
||
|
||
class ArchitectClient:
|
||
"""DeepSeek(或任意 OpenAI 兼容)大模型客户端。"""
|
||
|
||
def __init__(
|
||
self,
|
||
model: str,
|
||
base_url: str = "https://api.deepseek.com/v1",
|
||
api_key: Optional[str] = None,
|
||
temperature: float = 0.2,
|
||
timeout_s: float = 60.0,
|
||
max_tokens: int = 2048,
|
||
transport: Any = None,
|
||
_client: Any = None,
|
||
):
|
||
self.model = model
|
||
self.base_url = base_url.rstrip("/")
|
||
self.api_key = api_key
|
||
self.temperature = temperature
|
||
self.timeout_s = timeout_s
|
||
self.max_tokens = max_tokens
|
||
self._transport = transport
|
||
self._client = _client # 注入的 AsyncClient(测试用 MockTransport)
|
||
self._owns_client = _client is None
|
||
|
||
def _get_client(self):
|
||
if self._client is None:
|
||
import httpx
|
||
kwargs: Dict[str, Any] = {"timeout": self.timeout_s}
|
||
if self._transport is not None:
|
||
kwargs["transport"] = self._transport
|
||
self._client = httpx.AsyncClient(**kwargs)
|
||
return self._client
|
||
|
||
async def aclose(self) -> None:
|
||
if self._owns_client and self._client is not None:
|
||
await self._client.aclose()
|
||
self._client = None
|
||
|
||
# ---------------------------------------------------------------
|
||
# 三个对外能力
|
||
# ---------------------------------------------------------------
|
||
async def brief(self, query: str, ws: Workspace) -> Dict[str, Any]:
|
||
"""生成 brief。返回解析后的 brief dict;token 计量写入 ws。"""
|
||
user = (
|
||
"用户原始需求:" + "\n" + query + "\n\n"
|
||
"请生成任务 brief,仅输出符合如下结构的 JSON 对象:" + "\n"
|
||
+ json.dumps(_BRIEF_SCHEMA_HINT, ensure_ascii=False)
|
||
+ "\n注意:plan 中的 id 用 s1..sn,deps 引用已完成步骤 id;"
|
||
"acceptance 尽量 machine_checkable。"
|
||
)
|
||
return await self._chat_with_retry(ws, [("system", _SYSTEM_PROMPT), ("user", user)])
|
||
|
||
async def decide(self, ws: Workspace) -> Dict[str, Any]:
|
||
"""根据交流文本当前状态做裁决。返回 {reply, patch_plan}。"""
|
||
context = ws.render_for_architect()
|
||
user = (
|
||
"以下是交流文本摘要(不含工件全文):" + "\n\n" + context + "\n\n"
|
||
"请针对未解决 issues 做出裁决,仅输出符合如下结构的 JSON:" + "\n"
|
||
+ json.dumps(_DECIDE_SCHEMA_HINT, ensure_ascii=False)
|
||
+ "\nreply 给 Worker 具体可执行指示;patch_plan 列出需要修订的 step 与任务。"
|
||
)
|
||
return await self._chat_with_retry(ws, [("system", _SYSTEM_PROMPT), ("user", user)])
|
||
|
||
async def final_review(self, ws: Workspace) -> Dict[str, Any]:
|
||
"""终审。返回 {verdict: done|fix, notes, fix_issues}。"""
|
||
context = ws.render_for_architect()
|
||
user = (
|
||
"以下是待终审的交流文本摘要:" + "\n\n" + context + "\n\n"
|
||
"对照 brief 的 acceptance 做终审,仅输出符合如下结构的 JSON:" + "\n"
|
||
+ json.dumps(_REVIEW_SCHEMA_HINT, ensure_ascii=False)
|
||
+ "\nverdict=done 表示验收通过;fix 表示打回,fix_issues 列出需修正项。"
|
||
)
|
||
return await self._chat_with_retry(ws, [("system", _SYSTEM_PROMPT), ("user", user)])
|
||
|
||
# ---------------------------------------------------------------
|
||
# 底层
|
||
# ---------------------------------------------------------------
|
||
async def _chat_with_retry(self, ws: Workspace,
|
||
msgs: List[tuple]) -> Dict[str, Any]:
|
||
"""调用 + JSON 解析;解析失败回喂一次重写,再失败抛 ArchitectError。"""
|
||
if not self.api_key:
|
||
raise ArchitectError(
|
||
"Architect 未配置 API Key(env: 见 config.architect.api_key_env)。"
|
||
"请设置密钥,或使用本地降级模式(pipeline.breach_policy: local_only)。"
|
||
)
|
||
messages = [{"role": r, "content": c} for r, c in msgs]
|
||
for attempt in (1, 2):
|
||
content = await self._chat_once(ws, messages)
|
||
try:
|
||
return self._parse_json(content)
|
||
except ValueError as e:
|
||
if attempt == 1:
|
||
# 4.6:把原始输出与错误回喂重写一次
|
||
messages = messages + [
|
||
{"role": "assistant", "content": content},
|
||
{"role": "user",
|
||
"content": f"你的输出不是合法 JSON({e})。请重新只输出合法 JSON 对象。"},
|
||
]
|
||
continue
|
||
raise ArchitectError(f"Architect 输出非合法 JSON,重试后仍失败: {e}") from e
|
||
raise ArchitectError("未预期:_chat_with_retry 未返回") # 不可达
|
||
|
||
async def _chat_once(self, ws: Workspace, messages: List[Dict[str, Any]]) -> str:
|
||
if ws.exhausted():
|
||
raise ArchitectCircuitBreaker(
|
||
f"预算熔断:api_tokens={ws.budget()['api_input_tokens'] + ws.budget()['api_output_tokens']}"
|
||
f"/{ws.budget()['api_token_cap']}, round={ws.meta()['round']}/{ws.budget()['rounds_cap']}"
|
||
)
|
||
client = self._get_client()
|
||
headers = {"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}
|
||
body = {
|
||
"model": self.model,
|
||
"messages": messages,
|
||
"temperature": self.temperature,
|
||
"max_tokens": self.max_tokens,
|
||
"response_format": {"type": "json_object"},
|
||
}
|
||
try:
|
||
resp = await client.post(
|
||
f"{self.base_url}/chat/completions",
|
||
headers=headers,
|
||
json=body,
|
||
)
|
||
resp.raise_for_status()
|
||
except Exception as e:
|
||
raise ArchitectError(f"Architect API 调用失败: {type(e).__name__}: {e}") from e
|
||
data = resp.json()
|
||
usage = data.get("usage", {})
|
||
ws.add_budget(usage.get("prompt_tokens", 0), usage.get("completion_tokens", 0))
|
||
try:
|
||
return data["choices"][0]["message"]["content"]
|
||
except (KeyError, IndexError) as e:
|
||
raise ArchitectError(f"Architect 响应缺少 choices/content: {e}") from e
|
||
|
||
@staticmethod
|
||
def _parse_json(content: str) -> Dict[str, Any]:
|
||
"""从模型输出解析 JSON:剥除 markdown 代码围栏,取首个 JSON 对象。"""
|
||
text = content.strip()
|
||
_bt = "\x60" # 反引号(避免与构建脚本的字符串定界冲突)
|
||
fence = _bt * 3
|
||
text = text.replace(fence + "json", fence).replace(fence, "").strip()
|
||
try:
|
||
obj = json.loads(text)
|
||
if isinstance(obj, dict):
|
||
return obj
|
||
raise ValueError("顶层不是 object")
|
||
except json.JSONDecodeError:
|
||
pass
|
||
start = text.find("{")
|
||
end = text.rfind("}")
|
||
if start != -1 and end != -1 and end > start:
|
||
try:
|
||
obj = json.loads(text[start:end + 1])
|
||
if isinstance(obj, dict):
|
||
return obj
|
||
except json.JSONDecodeError:
|
||
pass
|
||
raise ValueError("无法解析为 JSON 对象")
|
||
|
||
|
||
def build_architect(cfg: Dict[str, Any],
|
||
get_env: Callable[[str], Optional[str]] = None) -> ArchitectClient:
|
||
"""cfg 为 config.architect 段。get_env 可注入(默认读 os.environ)。"""
|
||
import os
|
||
_env = get_env or os.environ.get
|
||
key = cfg.get("api_key") or _env(cfg.get("api_key_env", "DEEPSEEK_API_KEY"))
|
||
return ArchitectClient(
|
||
model=cfg.get("model", "deepseek-chat"),
|
||
base_url=cfg.get("base_url", "https://api.deepseek.com/v1"),
|
||
api_key=key,
|
||
temperature=float(cfg.get("temperature", 0.2)),
|
||
timeout_s=float(cfg.get("timeout_s", 60)),
|
||
)
|