feat(rag): 服务层接线(上传即索引 + use_rag 异步链路 + 强 e2e)

This commit is contained in:
lhl
2026-08-30 00:25:24 +08:00
parent 9514749b9f
commit f77fab707d
6 changed files with 228 additions and 9 deletions
+12
View File
@@ -1921,3 +1921,15 @@ Document(注入后 Word 文档)
### 12.8 WebSocket 实时进度流(2026-08-29,分支 feat/websocket-progress ### 12.8 WebSocket 实时进度流(2026-08-29,分支 feat/websocket-progress
新增 `ProgressHub` 进程内发布/订阅单例 + `/api/sessions/{sid}/ws` 端点 + `chat_ws.js` 前端实时渲染;进度/错误事件实时推送,既有 `role='progress'/'error'` 持久化兜底保留(重载仍可见)。**单进程假设**:hub 为进程内单例,多 worker 部署下跨进程不互通(后续可迭代 Redis 总线)。 新增 `ProgressHub` 进程内发布/订阅单例 + `/api/sessions/{sid}/ws` 端点 + `chat_ws.js` 前端实时渲染;进度/错误事件实时推送,既有 `role='progress'/'error'` 持久化兜底保留(重载仍可见)。**单进程假设**:hub 为进程内单例,多 worker 部署下跨进程不互通(后续可迭代 Redis 总线)。
## RAG 影响调查接入说明(Task 52026-08-29
影响调查 Agent 接入可选 RAG 检索能力(既有系统源码 -> 检索上下文注入 LLM 影响分析 prompt),详见 `src/genesis/rag/` 与 `src/genesis/impact/impact_agent.py` 的 `run_impact`。
- **scope = session_id**:每个会话的既有系统源码独立索引到 `RagStore` 的同一 scope,互不串扰。
- **上传即索引(D1**`GenesisService.upload_file` 在 `file_type == "existing_system"` 且 `self.rag is not None` 时,解压完成后立即调用 `self.rag.index_dir(session_id, path)`;索引异常仅记录日志(`_LOGGER.warning`)不阻断上传。
- **`use_rag` 默认关闭**`GenesisService` 构造参数 `use_rag` 默认 `False``rag=None` 表示不启用(向后兼容)。`run_impact(session_id, use_rag=None)` 中 `eff = self.use_rag if use_rag is None else use_rag`;仅当 `eff 且 self.rag is not None` 时走 LLM+RAG 路径,否则走原确定性 `ImpactAgent().run(...)` 路径(行为不变)。
- **异步链路**`run_impact` 为 `async def`RAG 路径 `await ImpactAgent(engine=..., rag=..., use_rag=True).run_impact(...)``app.start_impact` 端点同步改为 `async def` 并 `await service.run_impact(sid, use_rag=use_rag)``chat/agent.py` 调用处以 `asyncio.run(...)` 包裹以兼容同步消息处理。
- **线程安全(D2**`RagStore` 构造使用 `sqlite3.connect(db_path, check_same_thread=False)` 并加 `threading.Lock`,读写均加锁串行化,适配 Web 服务端 worker 线程复用连接。
- **向后兼容**`use_rag=False` 时 prompt 不含 RAG 小节标题(`_RAG_CONTEXT_TITLE`),影响报告为确定性 `impact-report.json`,不调用 LLM。
+3 -2
View File
@@ -5,6 +5,7 @@
""" """
from __future__ import annotations from __future__ import annotations
import asyncio
import json import json
from typing import Callable from typing import Callable
@@ -170,7 +171,7 @@ class ChatAgent:
self._emit_progress(session_id, progress[-1]) self._emit_progress(session_id, progress[-1])
rec = self.service.confirm_parse(session_id) rec = self.service.confirm_parse(session_id)
if rec.status == "impact_running": if rec.status == "impact_running":
rec = self.service.run_impact(session_id) rec = asyncio.run(self.service.run_impact(session_id))
progress.append({"step": "impact", "status": "ok", progress.append({"step": "impact", "status": "ok",
"detail": f"影响调查完成:{self._impact_brief(rec)}"}) "detail": f"影响调查完成:{self._impact_brief(rec)}"})
self._emit_progress(session_id, progress[-1]) self._emit_progress(session_id, progress[-1])
@@ -224,7 +225,7 @@ class ChatAgent:
reply = "当前状态无法进行影响调查。" reply = "当前状态无法进行影响调查。"
self.store.add_message(session_id, "assistant", reply, action="impact") self.store.add_message(session_id, "assistant", reply, action="impact")
return {"reply": reply, "progress": progress, "status": rec.status} return {"reply": reply, "progress": progress, "status": rec.status}
rec = self.service.run_impact(session_id) rec = asyncio.run(self.service.run_impact(session_id))
progress.append({"step": "impact", "status": "ok", progress.append({"step": "impact", "status": "ok",
"detail": f"影响调查完成:{self._impact_brief(rec)}"}) "detail": f"影响调查完成:{self._impact_brief(rec)}"})
self._emit_progress(session_id, progress[-1]) self._emit_progress(session_id, progress[-1])
+17 -4
View File
@@ -32,7 +32,7 @@ VERSION = "0.1.0"
class _FakeEngine: class _FakeEngine:
"""离线确定性引擎(--engine fake):用于无 API key 的 Web 演示/测试。""" """离线确定性引擎(--engine fake):用于无 API key 的 Web 演示/测试。"""
def chat_structured(self, *, session_id, prompt, variables, schema, retry_count=2): async def chat_structured(self, *, session_id, prompt, variables, schema, retry_count=2):
from types import SimpleNamespace from types import SimpleNamespace
title = variables.get("title", "x") title = variables.get("title", "x")
return SimpleNamespace( return SimpleNamespace(
@@ -83,6 +83,7 @@ def create_app(
store: SessionStore | None = None, store: SessionStore | None = None,
data_root: str = "data/server", data_root: str = "data/server",
engine: Any = None, engine: Any = None,
rag: "ImpactRAG | None" = None,
) -> FastAPI: ) -> FastAPI:
store = store or SessionStore() store = store or SessionStore()
is_fake = engine == "fake" is_fake = engine == "fake"
@@ -91,8 +92,20 @@ def create_app(
elif engine is None: elif engine is None:
engine = None # 真实模式:generate/qa 时按需 build(避免未配置 key 直接 503) engine = None # 真实模式:generate/qa 时按需 build(避免未配置 key 直接 503)
# RAG 接线:未注入则内部构造并默认关闭(use_rag=False,向后兼容)
if rag is None:
from genesis.rag.embeddings import get_embedder
from genesis.rag.impact_rag import ImpactRAG
from genesis.rag.store import RagStore
rag_db = str(Path(data_root) / "rag.db")
Path(rag_db).parent.mkdir(parents=True, exist_ok=True)
rag = ImpactRAG(RagStore(rag_db), get_embedder("fake"))
projects = ProjectsStore(db_path=str(store._db)) if store else None projects = ProjectsStore(db_path=str(store._db)) if store else None
service = GenesisService(store=store, data_root=data_root, engine=engine, projects=projects) service = GenesisService(
store=store, data_root=data_root, engine=engine, projects=projects,
rag=rag, use_rag=False,
)
from genesis.chat.agent import ChatAgent from genesis.chat.agent import ChatAgent
chat_agent = ChatAgent(service=service, fake=is_fake, engine=engine) chat_agent = ChatAgent(service=service, fake=is_fake, engine=engine)
@@ -232,9 +245,9 @@ def create_app(
# ---------- 影响调查 ---------- # ---------- 影响调查 ----------
@app.post("/api/sessions/{sid}/start-impact") @app.post("/api/sessions/{sid}/start-impact")
def start_impact(sid: str): async def start_impact(sid: str, use_rag: bool = False):
try: try:
rec = service.run_impact(sid) rec = await service.run_impact(sid, use_rag=use_rag)
except ServiceStepError as e: except ServiceStepError as e:
raise _error(409, "STATE_TRANSITION_INVALID", str(e)) raise _error(409, "STATE_TRANSITION_INVALID", str(e))
return {"ok": True, "status": rec.status} return {"ok": True, "status": rec.status}
+47 -1
View File
@@ -13,6 +13,7 @@ from __future__ import annotations
import html import html
import json import json
import logging
import shutil import shutil
import zipfile import zipfile
from pathlib import Path from pathlib import Path
@@ -22,6 +23,8 @@ from docx import Document
from genesis.parsers.source_aggregator import SourceParser from genesis.parsers.source_aggregator import SourceParser
from genesis.server.store import SessionStore, SessionRecord, ProjectsStore, ProjectConfigError from genesis.server.store import SessionStore, SessionRecord, ProjectsStore, ProjectConfigError
_LOGGER = logging.getLogger(__name__)
ALLOWED_FILE_TYPES = { ALLOWED_FILE_TYPES = {
"requirements", "template", "write_instruction", "rules", "existing_system", "requirements", "template", "write_instruction", "rules", "existing_system",
} }
@@ -57,6 +60,10 @@ class GenesisService:
prompt_registry=None, prompt_registry=None,
samples_dir: str = "sample", samples_dir: str = "sample",
projects: "ProjectsStore | None" = None, projects: "ProjectsStore | None" = None,
rag: "ImpactRAG | None" = None,
use_rag: bool = False,
rag_db_path: str | None = None,
embedder=None,
) -> None: ) -> None:
self.store = store self.store = store
self.data_root = Path(data_root) self.data_root = Path(data_root)
@@ -64,6 +71,14 @@ class GenesisService:
self.prompt_registry = prompt_registry self.prompt_registry = prompt_registry
self.samples_dir = samples_dir self.samples_dir = samples_dir
self.projects = projects self.projects = projects
# RAGrag 为 None 表示不启用(向后兼容默认关闭)
self.rag = rag
self.use_rag = use_rag
self.rag_db_path = rag_db_path
if embedder is None:
from genesis.rag.embeddings import get_embedder
embedder = get_embedder("fake")
self.embedder = embedder
# ---------- 会话与文件 ---------- # ---------- 会话与文件 ----------
@@ -126,6 +141,12 @@ class GenesisService:
zf.extractall(dst) zf.extractall(dst)
zpath.unlink(missing_ok=True) zpath.unlink(missing_ok=True)
path = str(dst) path = str(dst)
# D1:上传即索引(scope=session_id),索引异常不阻断上传
if self.rag is not None:
try:
self.rag.index_dir(session_id, path)
except Exception as e: # noqa: BLE001
_LOGGER.warning("existing_system 索引失败(已忽略): %s", e)
else: else:
path = str(sdir / filename) path = str(sdir / filename)
Path(path).write_bytes(content) Path(path).write_bytes(content)
@@ -204,10 +225,35 @@ class GenesisService:
# ---------- 影响调查 ---------- # ---------- 影响调查 ----------
def run_impact(self, session_id: str) -> SessionRecord: async def run_impact(self, session_id: str, use_rag: bool | None = None) -> SessionRecord:
rec = self.get_session(session_id) rec = self.get_session(session_id)
if rec.status != "impact_running": if rec.status != "impact_running":
raise ServiceStepError(f"当前状态 {rec.status} 不可启动影响调查") raise ServiceStepError(f"当前状态 {rec.status} 不可启动影响调查")
# 解析是否启用 RAG(实例默认 use_rag 可被显参覆盖)
eff = self.use_rag if use_rag is None else use_rag
if eff and self.rag is not None:
# RAG 路径:LLM 驱动,检索既有系统上下文注入 prompt
if self.engine is None:
from genesis.inference.factory import build_inference_engine
self.engine = build_inference_engine()
from genesis.impact.impact_agent import ImpactAgent
requirements_text = rec.files.get("requirements", {}).get("name") or "要件定義"
result = await ImpactAgent(engine=self.engine, rag=self.rag, use_rag=True).run_impact(
session_id, requirements_text, use_rag=True)
summary = {
"rag_enabled": True,
"impact_llm": getattr(result, "raw_text", None) or getattr(result, "data", None),
}
# 尽量保留既有确定性报告文件(无则保持原值)
impact_report_path = rec.impact_report_path
self.store.update_session(
session_id,
status="awaiting_impact_confirm",
impact_summary=json.dumps(summary, ensure_ascii=False),
impact_report_path=impact_report_path,
)
return self.get_session(session_id)
# 默认确定性路径(向后兼容)
from genesis.impact.impact_agent import ImpactAgent, impact_report_to_dict from genesis.impact.impact_agent import ImpactAgent, impact_report_to_dict
ss = self._rebuild_source(rec) ss = self._rebuild_source(rec)
report = ImpactAgent().run(ss, session_id=session_id) report = ImpactAgent().run(ss, session_id=session_id)
+146
View File
@@ -0,0 +1,146 @@
"""Task 5(D3 强端到端):服务层 RAG 接线验证。
真实驱动完整链路:上传 → 解压 → index_dir → retrieve → LLM promptRAG 注入),
不 mock 任何环节。验证上传即索引(D1)与 use_rag 默认关闭的向后兼容。
"""
from __future__ import annotations
import asyncio
import glob
import io
import zipfile
from pathlib import Path
import pytest
from genesis.impact.impact_agent import _RAG_CONTEXT_TITLE
from genesis.rag.embeddings import FakeEmbedder
from genesis.rag.impact_rag import ImpactRAG
from genesis.rag.store import RagStore
from genesis.server.service import GenesisService
from genesis.server.store import SessionStore
_REPO = Path(__file__).resolve().parents[1]
def _sample_req() -> bytes:
hits = glob.glob(str(_REPO / "sample" / "requirements_*.xlsx"))
assert hits, "sample 下未找到 requirements xlsx"
return Path(hits[0]).read_bytes()
def _sample_tpl() -> bytes:
hits = glob.glob(str(_REPO / "sample" / "template_*.docx"))
assert hits, "sample 下未找到 template docx"
return Path(hits[0]).read_bytes()
def _make_existing_zip() -> bytes:
"""构造含已知源码的既有系统 zipsrc/KnownOrder.java 含『订单创建调用 MyBatis』)。"""
buf = io.BytesIO()
with zipfile.ZipFile(buf, "w") as zf:
zf.writestr(
"src/KnownOrder.java",
"package demo;\n"
"public class KnownOrder {\n"
" // 订单创建调用 MyBatis\n"
" public void createOrder() { /* ... */ }\n}\n",
)
return buf.getvalue()
class FakeEngine:
"""异步确定性引擎:捕获 LLM 收到的 prompt。"""
def __init__(self) -> None:
self.captured: str | None = None
async def chat_structured(self, *, session_id, prompt, variables, schema, retry_count=2):
from types import SimpleNamespace
self.captured = prompt
return SimpleNamespace(
data={"title": variables.get("title", "x"), "blocks": []},
raw_text=prompt,
status="ok",
)
@pytest.fixture
def svc(tmp_path):
store = SessionStore(db_path=str(tmp_path / "s.db"))
rag = ImpactRAG(RagStore(str(tmp_path / "rag.db")), FakeEmbedder())
return GenesisService(
store=store,
data_root=str(tmp_path / "data"),
engine=FakeEngine(),
rag=rag,
use_rag=False, # 默认关闭,由 run_impact 显参开启
)
def test_rag_e2e_upload_then_impact_injects_source(tmp_path):
"""D3 强验证:上传源码被 RAG 检索并注入 LLM prompt。"""
store = SessionStore(db_path=str(tmp_path / "s.db"))
engine = FakeEngine()
rag = ImpactRAG(RagStore(str(tmp_path / "rag.db")), FakeEmbedder())
svc = GenesisService(
store=store,
data_root=str(tmp_path / "data"),
engine=engine,
rag=rag,
use_rag=False,
)
sid = svc.create_session("u1").session_id
svc.upload_file(sid, "requirements", "requirements_newdev.xlsx", _sample_req())
svc.upload_file(sid, "template", "template_design_ja.docx", _sample_tpl())
# 上传即索引(D1
svc.upload_file(sid, "existing_system", "existing.zip", _make_existing_zip())
# 上传后直接检索应命中非空
hits = rag.retrieve(sid, "订单创建", k=1)
assert hits, "上传即索引后应能检索到源码片段"
assert "KnownOrder" in hits[0]
svc.run_parse(sid)
svc.confirm_parse(sid)
assert svc.get_session(sid).status == "impact_running"
# 真实驱动 RAG 路径
rec = asyncio.run(svc.run_impact(sid, use_rag=True))
assert rec.status == "awaiting_impact_confirm"
# 核心断言:RAG 检索片段(含 KnownOrder)注入 LLM prompt
assert engine.captured is not None
assert "KnownOrder" in engine.captured
# 摘要标记 rag_enabled
import json
assert json.loads(rec.impact_summary).get("rag_enabled") is True
def test_rag_e2e_use_rag_false_backward_compat(tmp_path):
"""向后兼容:rag=None 且 use_rag 默认关闭时,不调用 LLM、无 RAG 小节。"""
store = SessionStore(db_path=str(tmp_path / "s.db"))
engine = FakeEngine()
svc = GenesisService(
store=store,
data_root=str(tmp_path / "data"),
engine=engine,
rag=None, # 不启用 RAG
use_rag=False,
)
sid = svc.create_session("u1").session_id
svc.upload_file(sid, "requirements", "requirements_newdev.xlsx", _sample_req())
svc.upload_file(sid, "template", "template_design_ja.docx", _sample_tpl())
svc.upload_file(sid, "existing_system", "existing.zip", _make_existing_zip())
svc.run_parse(sid)
svc.confirm_parse(sid)
rec = asyncio.run(svc.run_impact(sid)) # 默认 use_rag=None → 走确定性路径
assert rec.status == "awaiting_impact_confirm"
# 向后兼容:确定性路径不调用 LLM(captured 保持 None),prompt 不含 RAG 小节
assert engine.captured is None
assert _RAG_CONTEXT_TITLE not in (engine.captured or "")
# 确定性报告存在
assert rec.impact_report_path and Path(rec.impact_report_path).exists()
+3 -2
View File
@@ -7,6 +7,7 @@ from __future__ import annotations
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
import asyncio
import pytest import pytest
from genesis.server.store import SessionStore, ProjectsStore from genesis.server.store import SessionStore, ProjectsStore
@@ -248,7 +249,7 @@ def test_existing_system_zip_triggers_impact_flow(tmp_path):
svc.confirm_parse(s.session_id) svc.confirm_parse(s.session_id)
assert svc.get_session(s.session_id).status == "impact_running" assert svc.get_session(s.session_id).status == "impact_running"
svc.run_impact(s.session_id) asyncio.run(svc.run_impact(s.session_id))
got = svc.get_session(s.session_id) got = svc.get_session(s.session_id)
assert got.status == "awaiting_impact_confirm" assert got.status == "awaiting_impact_confirm"
assert got.impact_summary assert got.impact_summary
@@ -269,7 +270,7 @@ def test_impact_state_guard(svc):
svc.confirm_parse(s.session_id) svc.confirm_parse(s.session_id)
# 无既有系统 → writingstart-impact 应报错 # 无既有系统 → writingstart-impact 应报错
with pytest.raises(ServiceStepError): with pytest.raises(ServiceStepError):
svc.run_impact(s.session_id) asyncio.run(svc.run_impact(s.session_id))
def test_qa_requires_generated(svc): def test_qa_requires_generated(svc):