fix(rag): run_impact 改为 async 并 await 引擎(StructuredResult)

This commit is contained in:
lhl
2026-08-29 23:16:12 +08:00
parent 974bf12dc0
commit 9514749b9f
2 changed files with 67 additions and 64 deletions
+8 -11
View File
@@ -121,36 +121,33 @@ class ImpactAgent:
f"{requirements_text}\n"
)
def run_impact(
async def run_impact(
self,
session_id: str,
requirements_text: str,
use_rag: bool | None = None,
k: int = 5,
):
"""LLM 驱动的变更影响分析(可选 RAG 上下文注入)。
"""LLM 驱动的变更影响分析(可选 RAG 上下文注入,异步)。
- use_rag 优先取显参;为 None 时回退实例级 ``self.use_rag``
- 启用且 ``self.rag`` 存在时,以 ``影响调查:`` + 要件前若干字 为查询,
调用 ``self.rag.retrieve(session_id, query, k)``,将命中片段注入 prompt。
- use_rag=False 时 prompt 内容与原版完全一致(不含 RAG 小节向后兼容)。
- use_rag 优先取显参;为 None 时回退实例级 self.use_rag。
- 启用且 self.rag 存在时,以影响调查:+ 要件前若干字 为查询,
调用 self.rag.retrieve(session_id, query, k),将命中片段注入 prompt。
- use_rag=False 时 prompt 不含 RAG 小节向后兼容)。
- 返回底层引擎的 StructuredResult(含 data/raw_text)。
"""
if self.engine is None:
raise RuntimeError("run_impact 需要 engine,请在构造 ImpactAgent 时传入")
if use_rag is None:
use_rag = self.use_rag
prompt = self._build_impact_prompt(requirements_text)
if use_rag and self.rag is not None:
query = "影响调查:" + requirements_text[:200]
chunks = self.rag.retrieve(session_id, query, k)
if chunks:
rag_context = "\n".join(chunks)
prompt = prompt + f"\n\n{_RAG_CONTEXT_TITLE}\n{rag_context}"
return self.engine.chat_structured(
return await self.engine.chat_structured(
session_id=session_id,
prompt=prompt,
variables={},
+59 -53
View File
@@ -1,85 +1,91 @@
"""ImpactAgent RAG 上下文注入测试(RAG 迭代 Task 4)。
"""ImpactAgent RAG 上下文注入测试(RAG 迭代 Task 4,异步形态)。
验证:
- use_rag=True 时,run_impact 发送给 LLM 的 prompt 文本包含 RAG 检索命中片段与明确小节标题。
- use_rag=False(或默认)时,prompt 文本不含 RAG 小节标题(向后兼容)。
- run_impact 为 async def,真实 await 引擎 chat_structured。
"""
from genesis.impact.impact_agent import ImpactAgent
import asyncio
import types
import pytest
from genesis.impact.impact_agent import ImpactAgent, _RAG_CONTEXT_TITLE
from genesis.rag.embeddings import FakeEmbedder
from genesis.rag.impact_rag import ImpactRAG
from genesis.rag.store import RagStore
_RAG_SECTION_TITLE = "# 既有系统关联上下文(RAG 检索,辅助判断影响范围)"
class FakeEngine:
"""捕获真实 LLM 方法(chat_structured)收到的 prompt 文本。
"""捕获真实 LLM 方法(chat_structuredasync)收到的 prompt 文本。
方法名与签名刻意复用本仓库 InferenceEngine.chat_structured 的形参风格,
以保证 mock 的是真实接口(key=session_id/prompt/variables/schema/retry_count)。
"""
def __init__(self) -> None:
self.last_prompt: str | None = None
self.captured: str | None = None
self.calls = 0
def chat_structured(self, *, session_id, prompt, variables, schema, retry_count=2):
self.last_prompt = prompt
async def chat_structured(self, *, session_id, prompt, variables, schema, retry_count=2):
self.captured = prompt
self.calls += 1
# 返回结构兼容 ChatResult 的最小占位(测试仅校验 prompt 注入
return {"session_id": session_id, "prompt": prompt}
# 返回结构兼容 StructuredResult 的最小占位(含 data/raw_text
return types.SimpleNamespace(data={}, raw_text=prompt)
def _make_rag(session_id: str, sources):
def _make_rag(session_id: str, text: str) -> ImpactRAG:
store = RagStore(":memory:")
rag = ImpactRAG(store, FakeEmbedder())
rag.index(session_id, sources)
return store, rag
rag.index(session_id, [("TradeApplication.java", text)])
return rag
def test_run_impact_with_rag_injects_context():
session_id = "sess-rag"
store, rag = _make_rag(session_id, [("TradeApplication.java", "订单创建调用 MyBatis")])
try:
engine = FakeEngine()
agent = ImpactAgent(engine=engine, rag=rag, use_rag=True)
agent.run_impact(session_id, requirements_text="创建订单的影响", k=5)
prompt = engine.last_prompt
assert prompt is not None
# 命中片段(含文件名 TradeApplication.java)被注入
assert "TradeApplication" in prompt
# 明确小节标题被注入
assert _RAG_SECTION_TITLE in prompt
finally:
store.close()
rag = _make_rag("p1", "订单创建调用 MyBatis")
engine = FakeEngine()
agent = ImpactAgent(engine=engine, rag=rag, use_rag=True)
asyncio.run(agent.run_impact("p1", "创建订单的影响"))
assert engine.captured is not None
# 命中片段(含文件名 TradeApplication.java)被注入
assert "TradeApplication" in engine.captured
# 明确小节标题被注入
assert _RAG_CONTEXT_TITLE in engine.captured
def test_run_impact_without_rag_no_context():
session_id = "sess-no-rag"
store, rag = _make_rag(session_id, [("TradeApplication.java", "订单创建调用 MyBatis")])
try:
engine = FakeEngine()
agent = ImpactAgent(engine=engine, rag=rag, use_rag=False)
agent.run_impact(session_id, requirements_text="创建订单的影响", k=5)
prompt = engine.last_prompt
# 显式关闭 RAG:不含小节标题,也不含检索片段
assert _RAG_SECTION_TITLE not in prompt
assert "TradeApplication" not in prompt
finally:
store.close()
engine = FakeEngine()
agent = ImpactAgent(engine=engine, rag=_make_rag("p1", "x"), use_rag=False)
asyncio.run(agent.run_impact("p1", "创建订单的影响"))
assert _RAG_CONTEXT_TITLE not in engine.captured
def test_run_impact_default_no_rag_no_context():
# 默认 use_rag 为 False(未显式开启),行为与关闭一致
session_id = "sess-default"
store, rag = _make_rag(session_id, [("TradeApplication.java", "订单创建调用 MyBatis")])
try:
engine = FakeEngine()
agent = ImpactAgent(engine=engine, rag=rag) # 不传 use_rag
agent.run_impact(session_id, requirements_text="创建订单的影响", k=5)
prompt = engine.last_prompt
assert _RAG_SECTION_TITLE not in prompt
assert "TradeApplication" not in prompt
finally:
store.close()
def test_run_impact_rag_enabled_but_no_hits():
# rag 已建但索引在另一 scope,retrieve 命中为空 → 不应注入 RAG 小节
engine = FakeEngine()
agent = ImpactAgent(engine=engine, rag=_make_rag("other", "订单创建调用 MyBatis"), use_rag=True)
asyncio.run(agent.run_impact("p1", "创建订单的影响"))
assert _RAG_CONTEXT_TITLE not in engine.captured
def test_run_impact_use_rag_none_falls_back_to_instance_default():
# use_rag=None 时回退实例级 self.use_rag(此处为 True)→ 应注入 RAG 小节
engine = FakeEngine()
agent = ImpactAgent(engine=engine, rag=_make_rag("p1", "订单创建调用 MyBatis"), use_rag=True)
asyncio.run(agent.run_impact("p1", "创建订单的影响", use_rag=None))
assert _RAG_CONTEXT_TITLE in engine.captured
assert "TradeApplication" in engine.captured
def test_run_impact_explicit_false_no_context():
# 显式 use_rag=False(覆盖 use_rag is None 的 False 分支)→ 不注入 RAG 小节
engine = FakeEngine()
agent = ImpactAgent(engine=engine, rag=_make_rag("p1", "x"), use_rag=True)
asyncio.run(agent.run_impact("p1", "创建订单的影响", use_rag=False))
assert _RAG_CONTEXT_TITLE not in engine.captured
def test_run_impact_requires_engine():
agent = ImpactAgent()
with pytest.raises(RuntimeError):
asyncio.run(agent.run_impact("p1", "x"))