diff --git a/src/genesis/chat/agent.py b/src/genesis/chat/agent.py index 2fed88e..adf5e9f 100644 --- a/src/genesis/chat/agent.py +++ b/src/genesis/chat/agent.py @@ -6,17 +6,21 @@ from __future__ import annotations import json +from typing import Callable from genesis.chat.intent import parse_intent_fake, parse_intent_llm +from genesis.server.hub import hub as _progress_hub from genesis.server.service import GenesisService class ChatAgent: - def __init__(self, service: GenesisService, fake: bool = False, engine=None) -> None: + def __init__(self, service: GenesisService, fake: bool = False, engine=None, + progress_sink: "Callable[[dict], None] | None" = None) -> None: self.service = service self.store = service.store self.fake = fake self.engine = engine or service.engine + self.progress_sink = progress_sink # ---------- 主入口 ---------- @@ -51,6 +55,21 @@ class ChatAgent: """错误回复以 role='error' 持久化,便于前端区分样式(§4.6)。""" self.store.add_message(session_id, "error", reply, action=action) + def _emit_progress(self, session_id, item): + event = {"type": "progress", "step": item.get("step", ""), + "status": item.get("status", ""), "detail": item.get("detail", "")} + if self.progress_sink is not None: + self.progress_sink(event) + else: + _progress_hub.emit(session_id, event) + + def _emit_error(self, session_id, reply, action): + event = {"type": "error", "detail": reply, "action": action} + if self.progress_sink is not None: + self.progress_sink(event) + else: + _progress_hub.emit(session_id, event) + def _parse_intent(self, session_id, content, rec): if self.fake: return parse_intent_fake(content) @@ -73,8 +92,10 @@ class ChatAgent: try: self.service.confirm_impact(session_id) progress.append({"step": "impact", "status": "ok", "detail": "影响调查已确认"}) + self._emit_progress(session_id, progress[-1]) except Exception as e: # noqa: BLE001 reply = f"影响确认失败:{e}" + self._emit_error(session_id, reply, "confirm") self._store_error(session_id, reply, "confirm") self._persist_progress(session_id, progress) return {"reply": reply, "progress": progress, "status": "awaiting_impact_confirm"} @@ -126,6 +147,7 @@ class ChatAgent: rec = self._parse_and_confirm(session_id, progress) except Exception as e: # noqa: BLE001 reply = f"解析失败:{e}" + self._emit_error(session_id, reply, "generate") self._store_error(session_id, reply, "generate") return {"reply": reply, "progress": progress, "status": rec.status} if rec.status == "writing": @@ -145,11 +167,13 @@ class ChatAgent: summary = json.loads(rec.structured_summary) if rec.structured_summary else {} progress.append({"step": "parse", "status": "ok", "detail": f"解析完成:{summary.get('tables', 0)} 张表"}) + self._emit_progress(session_id, progress[-1]) rec = self.service.confirm_parse(session_id) if rec.status == "impact_running": rec = self.service.run_impact(session_id) progress.append({"step": "impact", "status": "ok", "detail": f"影响调查完成:{self._impact_brief(rec)}"}) + self._emit_progress(session_id, progress[-1]) return self.store.get_session(session_id) # awaiting_impact_confirm return rec # writing @@ -162,9 +186,10 @@ class ChatAgent: try: rec = self._parse_and_confirm(session_id, progress) except Exception as e: # noqa: BLE001 - reply = f"解析失败:{e}" - self._store_error(session_id, reply, "parse") - return {"reply": reply, "progress": progress, "status": rec.status} + reply = f"解析失败:{e}" + self._emit_error(session_id, reply, "parse") + self._store_error(session_id, reply, "parse") + return {"reply": reply, "progress": progress, "status": rec.status} if rec.status == "awaiting_impact_confirm": reply = f"解析与影响调查已完成:{self._impact_brief(rec)}。是否确认?(回复「确认,继续」)" self.store.add_message(session_id, "assistant", reply, action="impact") @@ -185,6 +210,7 @@ class ChatAgent: rec = self._parse_and_confirm(session_id, progress) except Exception as e: # noqa: BLE001 reply = f"解析失败:{e}" + self._emit_error(session_id, reply, "impact") self._store_error(session_id, reply, "impact") return {"reply": reply, "progress": progress, "status": rec.status} # _parse_and_confirm 已自动跑影响 → 直接返回反问 @@ -201,6 +227,7 @@ class ChatAgent: rec = self.service.run_impact(session_id) progress.append({"step": "impact", "status": "ok", "detail": f"影响调查完成:{self._impact_brief(rec)}"}) + self._emit_progress(session_id, progress[-1]) reply = f"影响调查完成:{self._impact_brief(rec)}。是否确认?(回复「确认,继续」)" self.store.add_message(session_id, "assistant", reply, action="impact") return {"reply": reply, "progress": progress, "status": "awaiting_impact_confirm"} @@ -212,16 +239,20 @@ class ChatAgent: rec = self.service.run_generate(session_id, output_language=output_language) except Exception as e: # noqa: BLE001 reply = f"生成失败:{e}" + self._emit_error(session_id, reply, "generate") self._store_error(session_id, reply, "generate") return {"reply": reply, "progress": progress, "status": rec.status} progress.append({"step": "generate", "status": "ok", "detail": "概要设计书已生成"}) + self._emit_progress(session_id, progress[-1]) try: rec = self.service.run_qa(session_id) progress.append({"step": "qa", "status": "ok", "detail": f"QA: 通过={rec.qa_summary}"}) + self._emit_progress(session_id, progress[-1]) except Exception as e: # noqa: BLE001 rec = self.store.get_session(session_id) progress.append({"step": "qa", "status": "warn", "detail": f"QA 未执行:{e}"}) + self._emit_progress(session_id, progress[-1]) reply = "概要设计书生成完成,QA 已执行。点击下方「下载 docx」获取结果。" self.store.add_message(session_id, "assistant", reply, action="generate") return {"reply": reply, "progress": progress, "status": rec.status} @@ -238,12 +269,14 @@ class ChatAgent: try: rec = self.service.run_qa(session_id) progress.append({"step": "qa", "status": "ok", "detail": "QA 完成"}) + self._emit_progress(session_id, progress[-1]) reply = f"QA 完成:{rec.qa_summary}" self.store.add_message(session_id, "assistant", reply, action="qa") return {"reply": reply, "progress": progress, "status": rec.status} except Exception as e: # noqa: BLE001 rec = self.store.get_session(session_id) reply = f"QA 未执行:{e}" + self._emit_error(session_id, reply, "qa") self._store_error(session_id, reply, "qa") return {"reply": reply, "progress": progress, "status": rec.status} diff --git a/tests/test_chat_agent_ws.py b/tests/test_chat_agent_ws.py new file mode 100644 index 0000000..ee51cbd --- /dev/null +++ b/tests/test_chat_agent_ws.py @@ -0,0 +1,32 @@ +import tempfile + +from genesis.chat.agent import ChatAgent +from genesis.server.service import GenesisService +from genesis.server.store import ProjectsStore, SessionStore + +_tmp = tempfile.mkdtemp(prefix="genesis_ws_") + + +def _make_agent(sink): + store = SessionStore(db_path=f"{_tmp}/sessions.db") + projects = ProjectsStore(db_path=f"{_tmp}/projects.db") + svc = GenesisService(store=store, data_root="data", engine="fake", projects=projects) + return ChatAgent(service=svc, fake=True, engine="fake", progress_sink=sink) + + +def test_agent_emits_progress_events(): + events = [] + agent = _make_agent(events.append) + sid = agent.service.create_session("u1").session_id + agent._emit_progress(sid, {"step": "parse", "status": "ok", "detail": "解析完成"}) + assert events and events[0]["type"] == "progress" + assert events[0]["step"] == "parse" + + +def test_agent_emits_error_event(): + events = [] + agent = _make_agent(events.append) + sid = agent.service.create_session("u1").session_id + agent._emit_error(sid, "解析失败:boom", "generate") + errs = [e for e in events if e["type"] == "error"] + assert errs and "boom" in errs[0]["detail"]