【Agent开发】第九阶段:RAG 的智能决策与自适应控制 (Intelligent Decision & Adaptive Control)—— 从“被动检索”进化为“主动思考”

文章目录


前置环境:

  1. 当前环境是基于WSL2 + Ubuntu 24.04 + Docker Desktop构建的云原生开发平台,所有服务(MySQL、Redis、Qwen)均以独立容器形式运行并通过Docker Compose统一编排。如何配置请参考我的博客 WSL2 + Ubuntu 24.04 + Docker Desktop 配置双内核环境
  2. 补充了milvus相关的配置,如何配置请参考我的博客 【Agent开发】第三阶段:RAG 实战 —— 赋予 Agent “外脑”
  3. 引入了ES检索,并且配置了ES服务,ES部分的配置请查看我的博客 【Agent开发】第五阶段:RAG 深度优化实战 —— 从“可用”到“卓越”
  4. 补充了Json形式的存储,引入了pgsql,pgsql的配置请查看我的博客 【Agent开发】第七阶段:RAG 自动化评估体系构建 (RAG Evaluation Framework) —— 从“凭感觉调优”到“数据驱动决策”

意图识别与语义路由 (Intent Recognition & Semantic Routing)—— 给 RAG 系统装上一个“智能交通指挥员”

核心格言“Don’t use a sledgehammer to crack a nut.” (别用大锤砸坚果。)
本讲目标:在检索发生之前,先通过意图分类将用户问题分流。闲聊直接回,简单事实快速查,复杂推理走深度链路。构建一个基于 LLM 的轻量级路由网关。

1. 为什么需要路由?(The “Why”)

想象一下,如果不管用户问什么,你的系统都执行以下全套动作:

  1. HyDE 生成假设文档。
  2. 多路检索 (向量+BM25) 查 10 个片段。
  3. Rerank 重排序。
  4. Self-RAG 带思维链生成。

后果

  • 用户问:“你好,在吗?”
    • 系统反应:疯狂检索公司文档,试图找到“在吗”的政策依据,耗时 3 秒,花费 $0.05,最后回答:“根据员工手册第 3 条,我在。” -> 用户体验极差,成本浪费
  • 用户问:“对比一下 A 产品和 B 产品的 API 延迟差异。”
    • 系统反应:只做了单次向量检索,漏掉了分散在两个文档里的数据,回答含糊其辞。 -> 能力不足

路由的价值

  • 降本:拦截 30% 的闲聊和无关问题,不调用昂贵的检索和生成模型。
  • 增效:为不同难度的问题匹配最合适的“武器库”。
  • 精准:将特定领域问题(如“代码报错”)路由到专用索引(如“GitHub 代码库”),避免被通用文档干扰。

2. 设计意图分类体系 (Taxonomy Design)

在写代码前,先定义你的“交通信号灯”。一个典型的 RAG 意图体系如下:

意图类别 (Intent) 描述 推荐策略 (Strategy) 示例
CHIT_CHAT 闲聊、问候、无意义输入 Direct Reply (不调用检索,直接用小模型回复) “你好”, “你是谁”, “今天天气不错”
FACT_LOOKUP 简单的事实查询,通常有明确实体 Fast Retrieval (单次向量检索,Top-3, 无 Rerank) “公司的年假是多少天?”, “CEO 是谁?”
HOW_TO 操作流程、步骤指南 Standard Retrieval (混合检索 + Rerank, Top-5) “如何重置密码?”, “报销流程是什么?”
COMPARISON 对比分析,涉及多个实体或维度 Deep Retrieval (多跳检索 + 分解问题 + Self-RAG) “A 计划和 B 计划有什么区别?”, “对比 v1 和 v2 的性能”
CODE_SEARCH 代码片段、API 用法、报错解决 Code Index Route (路由到专用代码库索引,启用 BM25 高权重) “Python 怎么读取 JSON?”, “Error 503 怎么修?”
UNKNOWN 无法分类或明显超出范围 Fallback (礼貌拒绝或转人工) (乱码), “帮我写个病毒”

3. 实战:构建 LLM 驱动的路由器

我们将使用 LangChainStructuredOutputParser 或 Pydantic 支持,让 LLM 输出标准的 JSON 格式,以便代码逻辑判断。

🛠️ 代码实现 (agent\router.py)

# agent/router.py
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.output_parsers import PydanticOutputParser
from langchain_core.runnables import RunnableLambda
from pydantic import BaseModel, Field
from typing import Literal, Optional, List
from langchain_openai import ChatOpenAI
from src.core.config import settings
from src.utils.xml_parser import remove_think_and_n
# from .llm_router import get_llm_client # 假设你有一个获取 LLM 的工厂函数

# 1. 定义输出结构 (Schema)
class RouteDecision(BaseModel):
    """路由决策对象:给出意图、策略、置信度以及可选澄清问题。"""

    intent: Literal["CHIT_CHAT", "FACT_LOOKUP", "HOW_TO", "COMPARISON", "CODE_SEARCH", "UNKNOWN"] = Field(
        description="用户问题的意图分类"
    )
    confidence: float = Field(
        description="分类的置信度 (0.0 - 1.0)"
    )
    reasoning: str = Field(
        description="简短的分类理由,用于调试"
    )
    strategy: Literal["direct_reply", "fast_retrieval", "standard_retrieval", "deep_search", "code_search", "fallback", "clarify_needed"] = Field(
        description="建议的后续处理策略名称"
    )
    clarification_questions: Optional[List[str]] = Field(
        default=None,
        description="低置信度时用于澄清需求的问题列表;高置信度时为 null"
    )

# 2. 定义 Prompt
ROUTER_PROMPT_TEMPLATE = """
你是一个智能 RAG 系统的路由指挥官。你的任务是分析用户问题,输出严格的 JSON 对象以决定处理路径。

# 核心任务
. 识别用户意图 (intent)。
. 评估置信度 (confidence)。
. **关键逻辑**:
   - 如果问题清晰且置信度 >= 0.6:直接推荐执行策略 (如 fast_retrieval)。
   - 如果问题模糊、歧义或缺少关键上下文导致置信度 < 0.6:
     - 将 `strategy` 设为 `clarify_needed`。
     - **必须**在 `clarification_questions` 字段中生成 2-3 个具体的引导性问题或场景选项,帮助用户澄清需求。
     - 引导问题应覆盖最可能的几种情况,语气要友好且具指导性。

# 约束条件 (严格遵守)
. **输出格式**:必须且仅输出一个合法的 JSON 对象。
. **禁止项**:不要输出 Markdown 代码块标记 (如 ```json),不要输出任何解释性文字。
. **字段限制**:
   - `intent`: ["CHIT_CHAT", "FACT_LOOKUP", "HOW_TO", "COMPARISON", "CODE_SEARCH", "UNKNOWN"]
   - `strategy`: ["direct_reply", "fast_retrieval", "standard_retrieval", "deep_search", "code_search", "fallback", "clarify_needed"]
   - `confidence`: 0.0 - 1.0 浮点数。
   - `clarification_questions`: 仅在低置信度时为非空列表,否则为 null。

# 意图与策略映射规则 (高置信度时)
- CHIT_CHAT -> direct_reply
- FACT_LOOKUP -> fast_retrieval
- HOW_TO -> standard_retrieval
- COMPARISON -> deep_search
- CODE_SEARCH -> code_search
- UNKNOWN -> fallback

# 低置信度处理示例 (Few-Shot)
用户输入: "那个怎么弄?"
JSON 输出:
{{
    "intent": "UNKNOWN",
    "confidence": 0.3,
    "reasoning": "代词'那个'指代不明,缺乏具体操作对象或上下文。",
    "strategy": "clarify_needed",
    "clarification_questions": [
        "您是指如何重置密码,还是如何申请报销?",
        "您是在操作手机 App 还是网页版时遇到的问题?",
        "能否提供具体的错误提示或您想实现的功能名称?"
    ]
}}

用户输入: "对比一下版本。"
JSON 输出:
{{
    "intent": "COMPARISON",
    "confidence": 0.4,
    "reasoning": "用户想要对比,但未指定对比的对象(如产品版本、方案A/B等)。",
    "strategy": "clarify_needed",
    "clarification_questions": [
        "您是想对比 v1.0 和 v2.0 版本的 API 差异吗?",
        "您是在对比我们的标准版和企业版服务方案吗?",
        "或者是想对比 Python 和 Java 在某个特定场景下的性能?"
    ]
}}

# 高置信度示例
用户输入: "Python 里怎么用 pandas 读取 csv?"
JSON 输出:
{{
    "intent": "CODE_SEARCH",
    "confidence": 0.98,
    "reasoning": "明确的编程库使用问题,意图清晰。",
    "strategy": "code_search",
    "clarification_questions": null
}}

# 用户问题
{question}

# 你的回答 (仅 JSON)
"""

class IntentRouter:
    """意图路由器:使用 LLM 将用户问题映射到可执行策略。"""

    def __init__(self, llm_model: str = "gpt-3.5-turbo"): 
        # 注意:路由任务很简单,可以用小模型 (gpt-3.5/haiku) 以节省成本和延迟
        self.llm = ChatOpenAI(
            base_url=settings.llm.base_url,
            model=settings.llm.model_name,
            api_key=settings.llm.api_key,
            temperature=float(settings.llm.temperature),
        )

        def extract_and_clean(message):
            # message 可能是 AIMessage 对象,也可能是字符串 (取决于版本和配置)
            content = message.content if hasattr(message, 'content') else str(message)
            # 移除 think 和 n 标记
            return remove_think_and_n(content)

        clean_chain = RunnableLambda(extract_and_clean)
        self.parser = PydanticOutputParser(pydantic_object=RouteDecision)
        
        prompt = ChatPromptTemplate.from_template(ROUTER_PROMPT_TEMPLATE)
        
        # 构建链:Prompt -> LLM -> clean_chain -> Parser
        self.chain = prompt | self.llm | clean_chain | self.parser

    async def route(self, question: str) -> RouteDecision:
        """
        对单个问题进行路由决策
        """
        try:
            decision = await self.chain.ainvoke({"question": question})
            return decision
        except Exception as e:
            # 解析失败时的降级策略:再次确认
            print(f"⚠️ 路由解析失败:{e},默认降级为确认问题")
            # 降级时也可以返回一个特殊的 clarify_needed,或者直接走标准检索
            return RouteDecision(
                intent="UNKNOWN",
                confidence=0.0,
                reasoning="Parsing failed or model error",
                strategy="clarify_needed",
                clarification_questions=["您的问题似乎有些复杂,能再详细描述一下具体场景吗?", "您是想查询政策、操作流程还是技术文档?"]
            )
# 测试脚本
if __name__ == "__main__":
    # python -m src.agent.router
    import asyncio
    
    async def test_router():
        router = IntentRouter()
        queries = [
            "你好,在吗?",
            "公司的年假政策是怎样的?",
            "如何重置我的登录密码?",
            "对比一下 v1.0 和 v2.0 版本的 API 响应速度。",
            "Python 里怎么用 pandas 读取 csv?",
            "我有如下需求:"
        ]
        
        for q in queries:
            print(f"\n❓ Q: {q}")
            res = await router.route(q)
            print(f"🎯 Intent: {res.intent}")
            print(f"🙋 Confidence: {res.confidence}")
            print(f"💡 Strategy: {res.strategy}")
            print(f"🧠 Reasoning: {res.reasoning}")

            if res.clarification_questions:
                print("❓ 需要澄清,建议询问用户:")
                for i, cq in enumerate(res.clarification_questions, 1):
                    print(f"   {i}. {cq}")
            else:
                print("✅ 意图清晰,直接执行策略。")
            print("-" * 30)
    asyncio.run(test_router())

💡 关键点解析

  1. 小模型大作用:路由不需要高智商,gpt-3.5-turboClaude Haiku 足够快且便宜。这层开销通常在 200ms 以内,但能节省后面几秒的检索生成时间。
  2. 结构化输出:使用 Pydantic 强制 LLM 输出合法 JSON,避免正则解析的脆弱性。
  3. 降级机制 (Fallback):如果 LLM 抽风了或者网络超时,默认走 clarify_needed,保证系统可用性(Availability > Optimization)。

4. 🌊 集成:流程编排(agent\orchestrator.py

from __future__ import annotations

from dataclasses import dataclass
from typing import Dict, Optional

from src.agent.router import IntentRouter, RouteDecision
from src.agent.strategies import (
    BaseAgentStrategy,
    ClarifyNeededStrategy,
    CodeSearchStrategy,
    DeepSearchStrategy,
    DirectReplyStrategy,
    FastRetrievalStrategy,
    FallbackStrategy,
    StandardRetrievalStrategy,
    StrategyContext,
    StrategyName,
    StrategyResult,
)

# intent 到执行策略的固定映射(用于稳定线上行为)
INTENT_TO_STRATEGY: Dict[str, StrategyName] = {
    "CHIT_CHAT": "direct_reply",
    "FACT_LOOKUP": "fast_retrieval",
    "HOW_TO": "standard_retrieval",
    "COMPARISON": "deep_search",
    "CODE_SEARCH": "code_search",
    "UNKNOWN": "fallback",
}


@dataclass(slots=True)
class RoutedExecution:
    """一次路由执行的完整结果:包含路由决策与策略执行产物。"""

    decision: RouteDecision
    result: StrategyResult


class AgentStrategyOrchestrator:
    """策略编排器:负责将 RouteDecision 分发到具体策略实现。"""

    def __init__(self, registry: Optional[Dict[StrategyName, BaseAgentStrategy]] = None):
        self.registry = registry or self._build_default_registry()

    @staticmethod
    def _build_default_registry() -> Dict[StrategyName, BaseAgentStrategy]:
        return {
            "direct_reply": DirectReplyStrategy(),
            "fast_retrieval": FastRetrievalStrategy(),
            "standard_retrieval": StandardRetrievalStrategy(),
            "deep_search": DeepSearchStrategy(),
            "code_search": CodeSearchStrategy(),
            "fallback": FallbackStrategy(),
            "clarify_needed": ClarifyNeededStrategy(),
        }

    def resolve_strategy_name(self, decision: RouteDecision) -> StrategyName:
        """解析最终执行策略。

        规则:
        1) 明确的 clarify_needed 优先;
        2) 低置信度自动转 clarify_needed;
        3) 否则按 intent 强映射;
        4) 最后才回退到 decision.strategy/fallback。
        """
        if decision.strategy == "clarify_needed" and "clarify_needed" in self.registry:
            return "clarify_needed"

        if decision.confidence < 0.6 and "clarify_needed" in self.registry:
            return "clarify_needed"

        # 优先按 intent 做强映射,避免 LLM 输出 strategy 偏移导致执行不稳定。
        by_intent = INTENT_TO_STRATEGY.get(decision.intent)
        if by_intent and by_intent in self.registry:
            return by_intent

        by_strategy = decision.strategy.lower().strip()
        if by_strategy in self.registry:
            return by_strategy  # type: ignore[return-value]
        return "fallback"

    async def execute(
        self,
        query: str,
        decision: RouteDecision,
        category: Optional[str] = None,
    ) -> StrategyResult:
        strategy_name = self.resolve_strategy_name(decision)
        strategy = self.registry[strategy_name]
        return await strategy.execute(
            StrategyContext(
                query=query,
                category=category,
                intent=decision.intent,
                route_strategy=decision.strategy,
                clarification_questions=decision.clarification_questions,
            )
        )


class RoutedAgentExecutor:
    """高层执行入口:先路由,再按策略编排执行。"""

    def __init__(
        self,
        router: Optional[IntentRouter] = None,
        orchestrator: Optional[AgentStrategyOrchestrator] = None,
    ):
        self.router = router or IntentRouter()
        self.orchestrator = orchestrator or AgentStrategyOrchestrator()

    async def run(self, query: str, category: Optional[str] = None) -> RoutedExecution:
        """对外统一调用方法。"""
        decision = await self.router.route(query)
        result = await self.orchestrator.execute(query=query, decision=decision, category=category)
        return RoutedExecution(decision=decision, result=result)

5. 🏤 后续转发LLM (strategies策略模式)

src\agent\strategies下新建__init__.py文件,并添加以下内容:

from .base import BaseAgentStrategy, StrategyContext, StrategyName, StrategyResult
from .retrieval import (
    DeepSearchStrategy,
    DirectReplyStrategy,
    FastRetrievalStrategy,
    StandardRetrievalStrategy,
)
from .system import ClarifyNeededStrategy, CodeSearchStrategy, FallbackStrategy

__all__ = [
    "BaseAgentStrategy",
    "StrategyContext",
    "StrategyName",
    "StrategyResult",
    "DirectReplyStrategy",
    "FastRetrievalStrategy",
    "StandardRetrievalStrategy",
    "DeepSearchStrategy",
    "CodeSearchStrategy",
    "FallbackStrategy",
    "ClarifyNeededStrategy",
]

🚌 简单策略(Fallback,ClarifyNeeded,CodeSearch)

因为暂时没想到code_search的场景,所以暂时不支持; 将这三个写如 system.py

from __future__ import annotations

from .base import BaseAgentStrategy, StrategyContext, StrategyResult


class CodeSearchStrategy(BaseAgentStrategy):
    """代码检索占位策略:当前版本明确返回不支持。"""

    name = "code_search"

    async def execute(self, context: StrategyContext) -> StrategyResult:
        return StrategyResult(
            strategy=self.name,
            message="暂不支持代码分析,请改为通用问题或知识库问题。",
            metadata={"retrieval_enabled": False},
        )


class FallbackStrategy(BaseAgentStrategy):
    """兜底策略:用于无法理解或潜在有害请求。"""

    name = "fallback"

    async def execute(self, context: StrategyContext) -> StrategyResult:
        return StrategyResult(
            strategy=self.name,
            message="无法理解该请求,或请求可能有害,已拒绝处理。",
            metadata={"retrieval_enabled": False},
        )


class ClarifyNeededStrategy(BaseAgentStrategy):
    """澄清策略:在低置信度时先向用户追问关键信息。"""

    name = "clarify_needed"

    async def execute(self, context: StrategyContext) -> StrategyResult:
        questions = context.clarification_questions or []
        if questions:
            # 将路由器给出的澄清问题整理成可直接发送给用户的文本
            formatted = "\n".join([f"{idx}. {q}" for idx, q in enumerate(questions, 1)])
            message = f"在继续检索前,请先澄清以下问题:\n{formatted}"
        else:
            message = "在继续检索前,请补充更具体的场景、对象或目标。"

        return StrategyResult(
            strategy=self.name,
            message=message,
            metadata={"retrieval_enabled": False, "clarification_questions": questions},
        )

🚅 功能策略(DirectReply,FastRetrieval,StandardRetrieval,DeepSearch)

直接在这里执行回答前的检索操作,并返回检索结果。

  • DirectReply: 闲聊类策略:不触发检索,交给上游生成模型直接作答。
  • FastRetrieval: 快速检索策略:只触发向量检索、ES多路,触发重排。
  • StandardRetrieval: 标准检索策略:触发向量检索、改写、hyde、ES多路召回重排。
  • DeepSearch: 暂时使用标准检索同款配置,以后再优化为GraphRAG.
from __future__ import annotations

from copy import deepcopy
from typing import Any, Dict, Optional

from src.rag.pipeline import RetrievalPipeline
from src.core.config import settings

from .base import BaseAgentStrategy, StrategyContext, StrategyResult


def _deep_merge(base: Dict[str, Any], patch: Dict[str, Any]) -> Dict[str, Any]:
    merged = dict(base)
    for key, value in patch.items():
        if isinstance(value, dict) and isinstance(merged.get(key), dict):
            merged[key] = _deep_merge(merged[key], value)
        else:
            merged[key] = value
    return merged


class DirectReplyStrategy(BaseAgentStrategy):
    """闲聊类策略:不触发检索,交给上游生成模型直接作答。"""

    name = "direct_reply"

    async def execute(self, context: StrategyContext) -> StrategyResult:
        return StrategyResult(
            strategy=self.name,
            message="无需检索,直接调用生成模型回复。",
            metadata={"retrieval_enabled": False},
        )


class FastRetrievalStrategy(BaseAgentStrategy):
    """事实查询策略:主路 + ES 双路增强,返回 Top-3。"""

    name = "fast_retrieval"

    def __init__(
        self,
        pipeline_config: Optional[Dict[str, Any]] = None,
        pipeline: Optional[RetrievalPipeline] = None,
    ):
        self.pipeline_config = {
            "retrieval": {
                "top_k": 3,
                "rough_top_k": 3,
            },
            "online": {
                "enable_rerank": True,
                "dynamic_threshold": settings.rag_online.score_threshold,
            },
            "composer": {
                "enable_hybrid_search": False,
                "plugin_rewritten_query": False,
                "plugin_rewritten_hyde": False,
                "plugin_es_questions": True,
                "plugin_es_summaries": True,
            },
            "filter": {},
        }
        if pipeline_config:
            # 约定:策略只关心 nested config,pipeline 负责解析与透传
            self.pipeline_config = _deep_merge(self.pipeline_config, pipeline_config)
        self.pipeline = pipeline or RetrievalPipeline()

    async def execute(self, context: StrategyContext) -> StrategyResult:
        run_config = deepcopy(self.pipeline_config)
        if context.category:
            run_config.setdefault("filter", {})
            run_config["filter"]["category"] = context.category

        results = await self.pipeline.run(
            query=context.query,
            config=run_config,
        )

        return StrategyResult(
            strategy=self.name,
            message="完成单次向量检索(通过预置 pipeline 配置执行)。",
            results=results,
            metadata={
                "retrieval_enabled": True,
                "top_k": run_config.get("retrieval", {}).get("top_k"),
                "pipeline_config": run_config,
            },
        )


class StandardRetrievalStrategy(BaseAgentStrategy):
    """标准检索策略:五路召回 + 可选重排。"""

    name = "standard_retrieval"

    def __init__(
        self,
        pipeline_config: Optional[Dict[str, Any]] = None,
        pipeline: Optional[RetrievalPipeline] = None,
    ):
        self.pipeline_config = {
            "retrieval": {
                "top_k": 5,
                "rough_top_k": 8,
            },
            "online": {
                "enable_rerank": True,
                "dynamic_threshold": settings.rag_online.score_threshold,
            },
            "composer": {
                "enable_hybrid_search": True,
                "plugin_rewritten_query": True,
                "plugin_rewritten_hyde": True,
                "plugin_es_questions": True,
                "plugin_es_summaries": True,
            },
            "filter": {},
        }
        if pipeline_config:
            self.pipeline_config = _deep_merge(self.pipeline_config, pipeline_config)
        self.pipeline = pipeline or RetrievalPipeline()

    async def execute(self, context: StrategyContext) -> StrategyResult:
        run_config = deepcopy(self.pipeline_config)
        if context.category:
            run_config.setdefault("filter", {})
            run_config["filter"]["category"] = context.category

        results = await self.pipeline.run(
            query=context.query,
            config=run_config,
        )
        return StrategyResult(
            strategy=self.name,
            message="完成混合检索与重排。",
            results=results,
            metadata={
                "retrieval_enabled": True,
                "top_k": run_config.get("retrieval", {}).get("top_k"),
                "pipeline_config": run_config,
            },
        )


class DeepSearchStrategy(StandardRetrievalStrategy):
    """深度检索策略:暂时使用标准检索同款配置,保留独立扩展入口。"""

    name = "deep_search"

    async def execute(self, context: StrategyContext) -> StrategyResult:
        # 当前阶段暂时复用 standard_retrieval
        standard = await super().execute(context)
        return StrategyResult(
            strategy=self.name,
            message="deep_search 暂由 standard_retrieval 执行。",
            results=standard.results,
            metadata={**standard.metadata, "delegated_to": "standard_retrieval"},
        )

6. 🤔 意图识别与转发测试

我们尝试使用之前上一讲,定义为bad_case的那个问题来进行询问。看看能否正确回答。

question: 公司每月提供的话费补贴金额是多少元?
Ground Truth: 每月补贴50元。

test\test_agent.py中添加如下代码:

import asyncio
import re
from typing import Optional


from src.agent.orchestrator import RoutedAgentExecutor


def _print_orchestrator_result(query: str, execution) -> None:
    decision = execution.decision
    result = execution.result
    print(f"\n👤 用户:{query}")
    print(
        "🧭 路由决策:"
        f" intent={decision.intent}, confidence={decision.confidence:.2f}, "
        f"strategy={decision.strategy}"
    )
    print(f"🧠 路由理由:{decision.reasoning}")
    print(f"🤖 策略回复:{result.message}")

    if result.results:
        print("📚 检索结果:")
        for idx, item in enumerate(result.results, 1):
            snippet = item.text.replace("\n", " ").strip()
            if len(snippet) > 120:
                snippet = snippet[:120] + "..."
            source = item.source_field or "unknown"
            print(f"  {idx}. score={item.score:.4f} | source={source} | {snippet}")
    print(f'🤖 最终回复:{result.final_answer}')

async def run_orchestrator_once(query: str, category: Optional[str] = None) -> None:
    executor = RoutedAgentExecutor()
    execution = await executor.run(query=query, category=category)
    _print_orchestrator_result(query, execution)


def run_orchestrator(question: str, category: Optional[str] = None) -> None:
    """可直接调用的 orchestrator 入口:只需要传 question。"""
    asyncio.run(run_orchestrator_once(query=question, category=category))


# 直接运行:python3 -m src.main
if __name__ == "__main__":
    #示例运行
    # run_mini_agent("我叫pd,记住我。", thread_id="session_004")
    # run_mini_agent("我是谁?", thread_id="session_004")
    # run_mini_agent("“根据知识库,你们最好的产品叫什么?”", thread_id="session_007")
    question = input("请输入 question: ").strip()
    if not question:
        question = "公司每月提供的话费补贴金额是多少元?"
    run_orchestrator(question)

在这里插入图片描述

7. 🎡 使用langgraph进行重构

只改逻辑,不改调用接口,main.py中依旧可用:

from __future__ import annotations

from dataclasses import dataclass
from typing import Dict, Optional, TypedDict

from langgraph.graph import END, START, StateGraph
from src.agent.router import IntentRouter, RouteDecision
from src.common.llm_adapter import LLMAdapter
from src.core.prompt_registry import PROMPT_KEYS, core_prompt_registry
from src.agent.strategies import (
    BaseAgentStrategy,
    ClarifyNeededStrategy,
    CodeSearchStrategy,
    DeepSearchStrategy,
    DirectReplyStrategy,
    FastRetrievalStrategy,
    FallbackStrategy,
    StandardRetrievalStrategy,
    StrategyContext,
    StrategyName,
    StrategyResult,
)

# intent 到执行策略的固定映射(用于稳定线上行为)
INTENT_TO_STRATEGY: Dict[str, StrategyName] = {
    "CHIT_CHAT": "direct_reply",
    "FACT_LOOKUP": "fast_retrieval",
    "HOW_TO": "standard_retrieval",
    "COMPARISON": "deep_search",
    "CODE_SEARCH": "code_search",
    "UNKNOWN": "fallback",
}


@dataclass(slots=True)
class RoutedExecution:
    """一次路由执行的完整结果:包含路由决策与策略执行产物。"""

    decision: RouteDecision
    result: StrategyResult
    final_answer: str


class AgentGraphState(TypedDict, total=False):
    """LangGraph 运行状态。

    字段约定(按执行顺序):
    1) run() 初始化时写入 `query/category`;
    2) route 节点写入 `decision/selected_strategy`;
    3) strategy 节点写入 `result`;
    4) compose 节点写入 `final_answer`。
    """

    # 原始用户问题:所有节点都会用到
    query: str
    # 可选业务分类(例如产品线),用于策略检索过滤
    category: Optional[str]
    # route 节点产出的结构化路由结果(意图/置信度/建议策略)
    decision: RouteDecision
    # 由 orchestrator.resolve_strategy_name 决定的最终策略名(用于 conditional_edges 分流)
    selected_strategy: StrategyName
    # 具体策略执行后的统一产物(文本消息 + 可选检索结果)
    result: StrategyResult
    # compose 节点最终写回的用户可见答案
    final_answer: str


class AgentStrategyOrchestrator:
    """策略编排器:负责将 RouteDecision 分发到具体策略实现。"""

    def __init__(self, registry: Optional[Dict[StrategyName, BaseAgentStrategy]] = None):
        self.registry = registry or self._build_default_registry()

    @staticmethod
    def _build_default_registry() -> Dict[StrategyName, BaseAgentStrategy]:
        return {
            "direct_reply": DirectReplyStrategy(),
            "fast_retrieval": FastRetrievalStrategy(),
            "standard_retrieval": StandardRetrievalStrategy(),
            "deep_search": DeepSearchStrategy(),
            "code_search": CodeSearchStrategy(),
            "fallback": FallbackStrategy(),
            "clarify_needed": ClarifyNeededStrategy(),
        }

    def resolve_strategy_name(self, decision: RouteDecision) -> StrategyName:
        """解析最终执行策略。

        规则:
        1) 明确的 clarify_needed 优先;
        2) 低置信度自动转 clarify_needed;
        3) 否则按 intent 强映射;
        4) 最后才回退到 decision.strategy/fallback。
        """
        if decision.strategy == "clarify_needed" and "clarify_needed" in self.registry:
            return "clarify_needed"

        if decision.confidence < 0.6 and "clarify_needed" in self.registry:
            return "clarify_needed"

        # 优先按 intent 做强映射,避免 LLM 输出 strategy 偏移导致执行不稳定。
        by_intent = INTENT_TO_STRATEGY.get(decision.intent)
        if by_intent and by_intent in self.registry:
            return by_intent

        by_strategy = decision.strategy.lower().strip()
        if by_strategy in self.registry:
            return by_strategy  # type: ignore[return-value]
        return "fallback"

    async def execute(
        self,
        query: str,
        decision: RouteDecision,
        category: Optional[str] = None,
    ) -> StrategyResult:
        strategy_name = self.resolve_strategy_name(decision)
        strategy = self.registry[strategy_name]
        return await strategy.execute(
            StrategyContext(
                query=query,
                category=category,
                intent=decision.intent,
                route_strategy=decision.strategy,
                clarification_questions=decision.clarification_questions,
            )
        )


class RoutedAgentExecutor:
    """高层执行入口:使用 LangGraph 实现 route->strategy->compose。"""

    _DIRECT_REPLY_PROMPT = (
        "你是一个简洁、准确的中文助手。请直接回答用户问题,不要输出思维过程。\n\n"
        "问题:\n{query}"
    )

    def __init__(
        self,
        router: Optional[IntentRouter] = None,
        orchestrator: Optional[AgentStrategyOrchestrator] = None,
        llm: Optional[LLMAdapter] = None,
    ):
        self.router = router or IntentRouter()
        self.orchestrator = orchestrator or AgentStrategyOrchestrator()
        self.llm = llm or LLMAdapter()
        self.rag_answer_prompt = core_prompt_registry.get(PROMPT_KEYS.SELF_RAG_GENERATE)
        self.app = self._build_graph().compile()

    def _build_graph(self):
        # 图结构固定为:START -> route -> (strategy_x) -> compose -> END
        workflow = StateGraph(AgentGraphState)
        workflow.add_node("route", self._route_node)
        workflow.add_node("compose", self._compose_node)

        for strategy_name in self.orchestrator.registry:
            # 每个策略是一个独立节点,便于后续按节点做观测/重试/限流
            workflow.add_node(strategy_name, self._build_strategy_node(strategy_name))
            workflow.add_edge(strategy_name, "compose")

        workflow.add_edge(START, "route")
        workflow.add_conditional_edges(
            "route",
            # route 节点写入 selected_strategy 后,在这里决定下一跳
            self._select_strategy_from_state,
            {name: name for name in self.orchestrator.registry},
        )
        workflow.add_edge("compose", END)
        return workflow

    async def _route_node(self, state: AgentGraphState) -> AgentGraphState:
        # 只读取 query,产出 decision + selected_strategy
        query = state["query"]
        decision = await self.router.route(query)
        selected_strategy = self.orchestrator.resolve_strategy_name(decision)
        return {
            "decision": decision,
            "selected_strategy": selected_strategy,
        }

    @staticmethod
    def _select_strategy_from_state(state: AgentGraphState) -> StrategyName:
        selected = state.get("selected_strategy")
        if selected:
            return selected
        return "fallback"

    def _build_strategy_node(self, strategy_name: StrategyName):
        async def _run_strategy(state: AgentGraphState) -> AgentGraphState:
            # 读取 route 阶段产物,执行对应策略,写回 result
            query = state["query"]
            decision = state["decision"]
            category = state.get("category")
            strategy = self.orchestrator.registry[strategy_name]
            result = await strategy.execute(
                StrategyContext(
                    query=query,
                    category=category,
                    intent=decision.intent,
                    route_strategy=decision.strategy,
                    clarification_questions=decision.clarification_questions,
                )
            )
            return {"result": result}

        return _run_strategy

    @staticmethod
    def _build_context_text(result: StrategyResult, max_items: int = 3, max_chars: int = 1500) -> str:
        chunks = []
        for item in result.results[:max_items]:
            text = (item.text or "").strip()
            if not text:
                continue
            chunks.append(text[:max_chars])
        return "\n\n".join(chunks) if chunks else "(无可用上下文)"

    async def _compose_final_answer(self, query: str, result: StrategyResult) -> str:
        no_llm_strategies = {"clarify_needed", "fallback", "code_search"}
        if result.strategy in no_llm_strategies:
            return result.message

        try:
            if result.strategy == "direct_reply":
                text = await self.llm.generate_text(
                    template=self._DIRECT_REPLY_PROMPT,
                    payload={"query": query},
                )
                text = text.strip()
                return text or result.message

            context_text = self._build_context_text(result)
            text = await self.llm.generate_text(
                template=self.rag_answer_prompt,
                payload={"query": query, "contexts": context_text},
            )
            text = text.strip()
            return text or result.message
        except Exception:
            # 生成失败时回退到策略消息,保证主链路可用
            return result.message

    async def _compose_node(self, state: AgentGraphState) -> AgentGraphState:
        # 基于 query + result 组装最终答案,写回 final_answer
        query = state["query"]
        result = state["result"]
        final_answer = await self._compose_final_answer(query=query, result=result)
        return {"final_answer": final_answer}

    async def run(self, query: str, category: Optional[str] = None) -> RoutedExecution:
        """对外统一调用方法。"""
        final_state = await self.app.ainvoke(
            {
                "query": query,
                "category": category,
            }
        )
        decision = final_state["decision"]
        result = final_state["result"]
        final_answer = final_state["final_answer"]
        return RoutedExecution(decision=decision, result=result, final_answer=final_answer)

再次运行测试,这次可以用闲聊类进行测试,输入你吃饭了不?

python -m src.main

回复:

👤 用户:`你吃饭了不?`
🧭 路由决策: intent=CHIT_CHAT, confidence=0.80, strategy=direct_reply
🧠 路由理由:用户使用疑问语气询问AI是否进食,属于日常闲聊范畴,无明确操作或事实查询需求。
🤖 策略回复:无需检索,直接调用生成模型回复。
🤖 最终回复:我是个AI,没有身体,所以不能吃饭。你呢?吃了吗?

8. 📝 本讲总结与行动清单

核心知识点

  1. 意图分类体系:根据业务场景定义清晰的意图类别。
  2. 结构化路由:利用 LLM + Pydantic 实现稳定的 JSON 输出。
  3. 策略分发:根据路由结果动态挂载不同的检索/生成管线。

🏗️ 系统架构

📂 当前项目结构解析

./agent
├── __init__.py              # 📦 包初始化:将 agent 目录标记为 Python 包,可在此导出核心类
├── orchestrator.py          # 🎼 编排器 (Orchestrator):核心控制器,负责串联 "路由->策略->结果组装" 的全流程
├── router.py                # 🧭 路由器 (Router):意图识别中心,分析 query 并决定调用哪个策略 (resolve_strategy)
└── strategies/              # ⚙️ 策略工厂:存放所有具体的执行策略实现
    ├── README.md            # 📖 策略文档:说明如何开发新策略、各策略的适用场景及配置参数
    ├── __init__.py          # 📦 子包初始化:统一导出所有策略类,方便外部 import
    ├── base.py              # 🏗️ 抽象基类:定义策略的标准接口 (如 execute, validate),确保所有策略行为一致
    ├── retrieval.py         # 🔍 检索策略集:包含 standard_retrieval, deep_search, fast_retrieval 等 RAG 相关实现
    └── system.py            # 🛠️ 系统策略集:包含 direct_reply (直接回复), code_search, fallback (兜底) 等非检索类实现

🤖 Agent模式解析

在这里插入图片描述

✅ 课后行动

  1. 定义你的 taxonomy:根据你的业务文档,列出 5-6 个核心意图。
  2. 部署路由器:运行 test_router,用真实用户日志测试分类准确率。
  3. 埋点监控:在日志中记录每个请求的 intent 分布,观察哪类问题最多,哪类最容易分错。

Self-RAG (自我反思检索增强生成)—— 让 RAG 系统学会“三思而后行”

既然第一讲是“事前分流”(决定走哪条路),第二讲是“事中反思”(走路时不断确认方向),

核心格言“Trust, but verify.” (信任,但要验证。)
本讲目标:解决传统 RAG“盲目检索、强行回答”的痛点。引入 Self-RAG 框架,让模型在生成过程中自主判断:是否需要检索?检索到的内容是否有用?生成的回答是否有依据? 从而大幅提升事实准确性和可控性。

1. 为什么需要 Self-RAG?(The Pain Points)

回顾我们之前的 RAG 流程:
用户提问 -> 检索 -> 拼接上下文 -> LLM 生成

存在的三大顽疾

  1. 过度检索 (Over-Retrieval):用户问“你好”,系统也去查数据库,浪费资源。
  2. 幻觉引用 (Hallucinated Citation):检索到的文档其实不相关,但 LLM 强行把它塞进答案里,甚至编造出处。
  3. 缺乏自知 (Lack of Self-Awareness):当检索结果为空或质量差时,传统 RAG 往往还是会硬凑一个答案,而不是承认“我不知道”。

Self-RAG 的核心思想
它不是简单的“检索+生成”,而是一个递归的、带反思机制的过程。模型在生成的每一步,都会输出特殊的 Reflection Tokens (反思标记) 来指导后续动作。


2. Self-RAG 的核心机制:四个关键反思点

Self-RAG 论文定义了四种关键的反思 Token,模型在生成时会像“内心独白”一样输出它们:

🧠 1. Retrieve? (是否需要检索?)

  • Token: [Retrieve: Yes] / [Retrieve: No]
  • 作用:模型先分析问题。如果是事实性问题,标记 Yes 触发检索;如果是闲聊或创意写作,标记 No 直接生成。
  • 收益:解决“过度检索”问题,降低延迟和成本。

🧠 2. Is Relevant? (检索内容相关吗?)

  • Token: [Relevant: Yes] / [Relevant: No]
  • 作用:拿到检索片段后,模型先不急着生成答案,而是先给片段打分。如果片段不相关,直接丢弃,尝试下一个片段或放弃检索。
  • 收益:过滤噪声,防止“垃圾进,垃圾出”。

🧠 3. Is Supported? (回答有依据吗?)

  • Token: [Supported: Yes] / [Supported: No]
  • 作用:生成一句话后,模型自我审查:这句话能完全由检索到的上下文推导出来吗?如果不能,标记 No,触发重新检索或修正回答。
  • 收益:大幅减少幻觉,确保“句句有出处”。

🧠 4. Is Useful? (回答有用吗?)

  • Token: [Useful: Yes] / [Useful: No]
  • 作用:针对最终答案的整体评估。是否解决了用户的问题?
  • 收益:优化最终输出的质量。

3. 架构对比:传统 RAG vs Self-RAG

特性 传统 RAG (Naive/Advanced) Self-RAG
流程控制 固定流水线 (Fixed Pipeline) 动态自适应 (Adaptive Workflow)
检索时机 总是检索 (Always Retrieve) 按需检索 (On-Demand Retrieval)
上下文处理 全部拼接 (Concat All) 细粒度筛选 (Segment-level Filtering)
幻觉控制 依赖 Prompt 约束 自我反思修正 (Self-Correction)
输出形式 纯文本 文本 + 反思标记 (可隐藏)
训练方式 通常无需训练 (Zero-shot) 通常需要微调 (Fine-tuned Critic/Generator)

4. 实战:如何在现有系统中实现 Self-RAG?

注意:原生的 Self-RAG 需要专门微调过的模型(如 selfrag/selfrag_llama2_7b)。但在工程实践中,我们可以通过 Prompt Engineering + 多步调用 来模拟这一过程,无需微调即可达到 80% 的效果。

🎯 目标

  1. 在最多 max_hops 轮内完成:检索 -> 生成 -> 自评判 -> 决策。
  2. 当回答可用且有证据支撑时结束;否则执行查询改写并继续检索。
  3. 输出结构化结果:最终回答、证据上下文、每轮评分轨迹、最终决策。

📂 目录结构

./self_rag
├── README.md                  # 📖 项目文档:Self-RAG 算法原理、反射 Token 定义及运行流程说明
├── __init__.py                # 📦 包初始化:导出核心引擎 (Engine) 和配置类 (config)
├── adapters/                  # 🔌 适配层:对接外部 RAG 管道与追踪系统,解耦核心逻辑
│   ├── __init__.py            
│   ├── rag_pipeline_adapter.py# 🔄 RAG 适配器:封装检索器调用,统一输入输出格式
│   └── trace_adapter.py       # 🕵️ 追踪适配器:记录每轮反思轨迹,用于调试或离线评估
├── config.py                  # ⚙️ 配置中心:定义反射阈值 (Thresholds)、最大重试轮次、检索参数等
├── engine.py                  # 🎬 主编排引擎 (Core Loop):驱动 "生成->评估->回溯" 的无限循环直到终止
├── nodes/                     # 🧩 原子节点库:实现 Self-RAG 论文中的具体功能单元
│   ├── __init__.py            
│   ├── decide_next.py         # 🔀 决策节点:根据评估结果决定下一步动作 (继续/重写/重新检索)
│   ├── generate.py            # ✍️ 生成节点:基于上下文生成文本片段及特殊反思 Token
│   ├── judge_grounding.py     # ⚖️ 事实性评判:判断生成内容是否被检索文档支持 (IsSup)
│   ├── judge_relevance.py     # 🎯 相关性评判:判断检索文档是否与查询相关 (IsRel)
│   ├── judge_utility.py       # 💡 有用性评判:判断生成内容是否真正解决了用户问题 (IsUse)
│   ├── retrieve.py            # 🔍 检索节点:执行向量搜索或关键词搜索
│   ├── rewrite_query.py       # 📝 改写节点:当检索失败时,优化查询语句后重试
│   └── route.py               # 🧭 路由节点:初始判断是否需要启动检索流程
├── schemas/                   # 📐 数据结构定义:严格规范评判结果与最终输出的 Pydantic/数据类
│   ├── __init__.py            # 📦 子包初始化
│   ├── judge.py               # 🏷️ 评判结构:单维度评判结果
│   └── output.py              # 📦 输出结构:定义最终回答
└── state.py                   # 🧠 运行时状态 (State):携带对话历史、当前轮次、累积轨迹的全局上下文对象
  • Prompt 模板统一放在 src/core/prompt/(由 src/core/prompt_registry.py 读取)。
节点名称 对应代码类 作用 反思类型
RouteNode RouteNode 入口守门员。判断问题是否需要进入复杂的 Self-RAG 流程,还是直接走简单通道或拒绝。(暂未实现) -
RetrieveNode RetrieveNode 情报收集者。执行实际的向量/关键词检索。 -
GenerateNode GenerateNode 初稿撰写者。基于当前检索到的上下文生成初步答案。 -
JudgeRelevanceNode JudgeRelevanceNode 相关性质检员。评估检索到的文档是否与问题真正相关。 Is Relevant?
JudgeGroundingNode JudgeGroundingNode 事实质检员。评估生成的答案是否完全由文档支撑(防幻觉)。 Is Supported?
JudgeUtilityNode JudgeUtilityNode 效用质检员。评估答案是否真正解决了用户的问题。 Is Useful?
DecideNextNode DecideNextNode 交通指挥员。综合三个质检员的分数,决定下一步动作:finish (结束), rewrite (改写重试), 或 fallback (放弃)。 -
RewriteQueryNode RewriteQueryNode 问题优化师。如果决定重试,它会根据失败原因(如“证据不足”)重写查询词,以便下一轮检索更精准。 -

🛠️ 代码实现

🔌 组件接入计划
  1. 检索:复用 src/rag/pipeline.py,通过 config 透传检索行为。
  2. 生成:使用统一 LLM 适配层,按上下文生成答案。
  3. 评判:
    • relevance: 上下文与问题相关性;
    • grounding: 答案是否由证据支撑;
    • utility: 答案是否满足用户需求。
  4. 改写:当评判失败时重写 query 再检索。
  5. 决策:根据阈值和轮次决定 finish / rewrite / fallback
节点示例
  • GenerateNode生成节点:基于问题与上下文生成回答
from __future__ import annotations

from typing import Any, List


class GenerateNode:
    """生成节点:基于问题与上下文生成回答。"""

    def __init__(self, llm: Any, prompt: str):
        self.llm = llm
        self.prompt = prompt

    async def run(self, query: str, contexts: List[str]) -> str:
        return await self.llm.generate_text(
            template=self.prompt,
            payload={
                "query": query,
                "contexts": "\n\n".join(contexts) if contexts else "(无可用上下文)",
            },
        )
  • JudgeRelevanceNode:判断节点,判断上下文与问题是否相关。
from __future__ import annotations

from typing import Any, List
from src.self_rag.schemas.judge import JudgeResult


class JudgeRelevanceNode:
    """评判上下文是否与问题相关。"""

    def __init__(self, llm: Any, prompt: str, threshold: float):
        self.llm = llm
        self.prompt = prompt
        self.threshold = threshold

    async def run(self, query: str, contexts: List[str]) -> JudgeResult:
        result = await self.llm.generate_structured(
            template=self.prompt,
            payload={
                "query": query,
                "contexts": "\n\n".join(contexts) if contexts else "(无可用上下文)",
            },
            output_model=JudgeResult,
        )
        result.passed = result.score >= self.threshold
        return result
🎲 状态定义
from __future__ import annotations

from dataclasses import dataclass, field
from typing import List, Optional, TypedDict

from src.self_rag.schemas.judge import JudgeResult

@dataclass
class HopTrace:
    """Self-RAG 单轮(1 hop)执行轨迹。

    含义:
    1. Self-RAG 每一轮都会产出一条 HopTrace,记录“检索 -> 生成 -> 评判 -> 决策”的关键结果;
    2. 多轮执行时,外层会把这些记录按顺序放入 traces 列表,形成完整链路;
    3. 这些轨迹最终会进入输出的 `trace` 字段,便于排查“为什么 finish / rewrite / fallback”。

    关键字段:
    - query: 本轮实际使用的检索问题(可能是改写后的问题)
    - answer: 本轮生成回答
    - contexts: 本轮命中的证据上下文
    - relevance/grounding/utility: 三个评判维度的结构化结果
    - decision: 本轮决策(finish / rewrite / fallback)
    - rewritten_query: 当 decision=rewrite 时生成的新 query(否则为 None)
    """

    hop: int
    query: str
    answer: str
    contexts: List[str]
    relevance: JudgeResult
    grounding: JudgeResult
    utility: JudgeResult
    decision: str
    rewritten_query: Optional[str] = None


@dataclass
class SelfRAGState:
    original_query: str
    current_query: str
    max_hops: int
    traces: List[HopTrace] = field(default_factory=list)


class SelfRAGGraphState(TypedDict, total=False):
    """LangGraph 运行状态定义。

    说明:
    1. `original_query/current_query` 区分“原始问题”和“当前轮检索问题”;
    2. `hop` 表示已执行到第几轮;
    3. `final_decision` 由 `run_hop` 节点写入,用于 conditional edge 决定是否继续循环;
    4. `last_*` 字段用于在图结束时直接组装最终输出,避免再做二次推导。
    """

    original_query: str
    current_query: str
    category: Optional[str]
    max_hops: int
    hop: int
    traces: List[HopTrace]
    route_passed: bool
    final_decision: str
    last_answer: str
    last_contexts: List[str]
    last_rewritten_query: Optional[str]
主流程编排
from __future__ import annotations

import asyncio
from typing import Any, List, Literal, Optional

import logging
from langgraph.graph import END, START, StateGraph

from src.core.prompt_registry import PROMPT_KEYS, core_prompt_registry
from src.self_rag.adapters.trace_adapter import TraceAdapter
from src.self_rag.config import SelfRAGConfig
from src.self_rag.nodes import (
    DecideNextNode,
    GenerateNode,
    JudgeGroundingNode,
    JudgeRelevanceNode,
    JudgeUtilityNode,
    RetrieveNode,
    RewriteQueryNode,
    RouteNode,
)
from src.self_rag.schemas.judge import JudgeResult
from src.self_rag.schemas.output import SelfRAGOutput
from src.self_rag.state import HopTrace, SelfRAGGraphState

logger = logging.getLogger(__name__)


class SelfRAGEngine:
    """Self-RAG 主编排引擎(LangGraph 版本)。"""

    def __init__(
        self,
        config: Optional[SelfRAGConfig] = None,
        rag_adapter: Optional[Any] = None,
        llm_adapter: Optional[Any] = None,
        judge_llm_adapter: Optional[Any] = None,
        trace_adapter: Optional[TraceAdapter] = None,
    ):
        self.config = config or SelfRAGConfig()
        self.trace = trace_adapter or TraceAdapter()

        if llm_adapter is None:
            from src.common.llm_adapter import LLMAdapter

            llm = LLMAdapter()
        else:
            llm = llm_adapter

        if judge_llm_adapter is not None:
            judge_llm = judge_llm_adapter
        elif llm_adapter is not None:
            # 测试/注入场景:若外部已显式传入 llm_adapter,则默认复用,避免强依赖外部 endpoint 文件。
            judge_llm = llm_adapter
        else:
            from src.self_rag.adapters.judge_llm_adapter import JudgeLLMAdapter

            # 生产默认:judge_* 三节点走外部大模型路由,缓解本地单模型并发与稳定性问题。
            judge_llm = JudgeLLMAdapter(
                llm_json_path=self.config.judge_llm_json_path,
                llm_group=self.config.judge_llm_group,
            )

        if rag_adapter is None:
            from src.self_rag.adapters.rag_pipeline_adapter import RAGPipelineAdapter

            rag = RAGPipelineAdapter()
        else:
            rag = rag_adapter

        self.route_node = RouteNode()
        self.retrieve_node = RetrieveNode(rag)
        self.generate_node = GenerateNode(llm, core_prompt_registry.get(PROMPT_KEYS.SELF_RAG_GENERATE))
        self.judge_relevance_node = JudgeRelevanceNode(
            judge_llm,
            core_prompt_registry.get(PROMPT_KEYS.SELF_RAG_JUDGE_RELEVANCE),
            threshold=self.config.relevance_threshold,
        )
        self.judge_grounding_node = JudgeGroundingNode(
            judge_llm,
            core_prompt_registry.get(PROMPT_KEYS.SELF_RAG_JUDGE_GROUNDING),
            threshold=self.config.grounding_threshold,
        )
        self.judge_utility_node = JudgeUtilityNode(
            judge_llm,
            core_prompt_registry.get(PROMPT_KEYS.SELF_RAG_JUDGE_UTILITY),
            threshold=self.config.utility_threshold,
        )
        self.rewrite_query_node = RewriteQueryNode(
            llm,
            core_prompt_registry.get(PROMPT_KEYS.SELF_RAG_REWRITE_QUERY),
        )
        self.decide_node = DecideNextNode()
        self.app = self._build_graph().compile()

    def _build_graph(self):
        """构建 Self-RAG 执行图。

        关键设计:
        1. 将“多轮循环”交给 LangGraph 条件边处理,而不是手写 for-loop;
        2. route 节点只负责“是否进入 Self-RAG”;
        3. run_hop 节点完成一整轮:检索->生成->评判->决策->可选改写。
        """

        workflow = StateGraph(SelfRAGGraphState)
        workflow.add_node("route", self._route)
        workflow.add_node("run_hop", self._run_hop)

        workflow.add_edge(START, "route")
        workflow.add_conditional_edges(
            "route",
            self._next_after_route,
            {
                "run_hop": "run_hop",
                "end": END,
            },
        )
        workflow.add_conditional_edges(
            "run_hop",
            self._next_after_hop,
            {
                "run_hop": "run_hop",
                "end": END,
            },
        )
        return workflow

    @staticmethod
    def _failure_reasons(
        relevance: JudgeResult,
        grounding: JudgeResult,
        utility: JudgeResult,
    ) -> List[str]:
        reasons: List[str] = []
        if not relevance.passed:
            reasons.append(f"相关性不足(score={relevance.score:.2f})")
        if not grounding.passed:
            reasons.append(f"证据支撑不足(score={grounding.score:.2f})")
        if not utility.passed:
            reasons.append(f"回答效用不足(score={utility.score:.2f})")
        return reasons

    @staticmethod
    def _trace_to_dict(trace: HopTrace) -> dict:
        return {
            "hop": trace.hop,
            "query": trace.query,
            "answer": trace.answer,
            "contexts": trace.contexts,
            "relevance": trace.relevance.model_dump(),
            "grounding": trace.grounding.model_dump(),
            "utility": trace.utility.model_dump(),
            "decision": trace.decision,
            "rewritten_query": trace.rewritten_query,
        }

    @staticmethod
    def _fallback_judge(reason: str) -> JudgeResult:
        return JudgeResult(score=0.0, passed=False, reasoning=reason)

    async def _route(self, state: SelfRAGGraphState) -> SelfRAGGraphState:
        current_query = state["current_query"]
        if self.route_node.run(current_query):
            return {"route_passed": True}
        return {
            "route_passed": False,
            "final_decision": "fallback",
            "last_answer": self.config.fallback_answer,
            "last_contexts": [],
            "last_rewritten_query": None,
        }

    @staticmethod
    def _next_after_route(state: SelfRAGGraphState) -> Literal["run_hop", "end"]:
        return "run_hop" if state.get("route_passed", False) else "end"

    async def _run_hop(self, state: SelfRAGGraphState) -> SelfRAGGraphState:
        """执行一轮 Self-RAG。

        注意:
        - 检索/生成失败:直接 fallback(保持旧行为,且不记录本轮 trace)。
        - 评判失败:降级为低分判定,流程继续可改写。
        """

        hop = state.get("hop", 0) + 1
        max_hops = state.get("max_hops", self.config.max_hops)
        current_query = state["current_query"]
        category = state.get("category")
        traces = list(state.get("traces", []))

        try:
            results = await self.retrieve_node.run(
                query=current_query,
                config=self.config.retrieval_config,
                category=category,
            )
        except Exception:
            logger.exception("self-rag retrieve failed at hop=%s", hop)
            return {
                "hop": hop,
                "final_decision": "fallback",
                "last_answer": self.config.fallback_answer,
            }
        contexts = [item.text for item in results]

        try:
            answer = await self.generate_node.run(current_query, contexts)
        except Exception:
            logger.exception("self-rag generate failed at hop=%s", hop)
            return {
                "hop": hop,
                "final_decision": "fallback",
                "last_answer": self.config.fallback_answer,
            }

        try:
            relevance = await self.judge_relevance_node.run(current_query, contexts)
        except Exception:
            relevance = self._fallback_judge("相关性评判失败")
        try:
            grounding = await self.judge_grounding_node.run(current_query, answer, contexts)
        except Exception:
            grounding = self._fallback_judge("证据支撑评判失败")
        try:
            utility = await self.judge_utility_node.run(current_query, answer)
        except Exception:
            utility = self._fallback_judge("效用评判失败")

        decision = self.decide_node.run(
            hop=hop,
            max_hops=max_hops,
            relevance=relevance,
            grounding=grounding,
            utility=utility,
        )

        rewritten_query: Optional[str] = None
        if decision == "rewrite":
            try:
                rewritten_query = await self.rewrite_query_node.run(
                    original_query=state["original_query"],
                    current_query=current_query,
                    answer=answer,
                    failure_reasons=self._failure_reasons(relevance, grounding, utility),
                )
            except Exception:
                rewritten_query = None

            # 关键兜底:改写失败或改写后未变化,直接终止并 fallback,防止无效循环。
            if not rewritten_query or rewritten_query == current_query:
                decision = "fallback"

        trace = HopTrace(
            hop=hop,
            query=current_query,
            answer=answer,
            contexts=contexts,
            relevance=relevance,
            grounding=grounding,
            utility=utility,
            decision=decision,
            rewritten_query=rewritten_query,
        )
        traces.append(trace)
        self.trace.log(self._trace_to_dict(trace))

        next_query = rewritten_query if decision == "rewrite" and rewritten_query else current_query
        return {
            "hop": hop,
            "traces": traces,
            "current_query": next_query,
            "final_decision": decision,
            "last_answer": answer or self.config.fallback_answer,
            "last_contexts": contexts,
            "last_rewritten_query": rewritten_query,
        }

    @staticmethod
    def _next_after_hop(state: SelfRAGGraphState) -> Literal["run_hop", "end"]:
        # 只有 rewrite 才继续下一轮;finish/fallback 都结束图执行。
        return "run_hop" if state.get("final_decision") == "rewrite" else "end"

    async def run(self, query: str, category: Optional[str] = None) -> SelfRAGOutput:
        initial_state: SelfRAGGraphState = {
            "original_query": query,
            "current_query": query,
            "category": category,
            "max_hops": self.config.max_hops,
            "hop": 0,
            "traces": [],
            "route_passed": True,
            "final_decision": "fallback",
            "last_answer": self.config.fallback_answer,
            "last_contexts": [],
            "last_rewritten_query": None,
        }

        final_state = await self.app.ainvoke(initial_state)
        traces = list(final_state.get("traces", []))
        trace_payload = [self._trace_to_dict(item) for item in traces]

        return SelfRAGOutput(
            query=query,
            final_answer=final_state.get("last_answer", self.config.fallback_answer),
            final_decision=final_state.get("final_decision", "fallback"),
            hops_used=len(traces),
            contexts=final_state.get("last_contexts", []),
            rewritten_query=final_state.get("last_rewritten_query"),
            trace=trace_payload,
        )

    def run_sync(self, query: str, category: Optional[str] = None) -> SelfRAGOutput:
        try:
            asyncio.get_running_loop()
            loop = asyncio.new_event_loop()
            try:
                return loop.run_until_complete(self.run(query=query, category=category))
            finally:
                loop.close()
        except RuntimeError:
            return asyncio.run(self.run(query=query, category=category))

代码中的精妙之处

  1. _run_hop 函数:它将“检索->生成->三个评判->决策->可选改写”封装成了一个原子操作(Atom)。这使得主图结构非常清晰:只有 routerun_hop 两个大节点。
  2. 条件边 (add_conditional_edges)
    • _next_after_route: 决定是否进入 Self-RAG 循环。
    • _next_after_hop: 核心循环控制。只有当 decision == 'rewrite' 时,才会回到 run_hop 再次执行;否则结束。这完美实现了 多跳 (Multi-Hop) 逻辑。
  3. 容错处理 (try...except):在 _run_hop 中,即使某个评判节点(如 JudgeGroundingNode)报错,也不会导致整个流程崩溃,而是返回一个低分的 JudgeResult (_fallback_judge),让决策节点有机会选择“重试”或“降级”,极大地提高了系统的鲁棒性。
  4. Trace 记录 (TraceAdapter):每一步的 HopTrace 都被详细记录(包括 query, answer, scores, decision),这对于后续的分析、调试和评估至关重要。

测试

"""真实链路 Self-RAG 复杂问题测试脚本(不使用任何 Fake 数据)。

运行示例:
python -m src.test.run_self_rag_complex
"""

from __future__ import annotations

import json
from typing import Any, Dict, List


# 运行配置:按你的要求直接用字典写死,不走命令行参数。
RUN_CONFIG: Dict[str, Any] = {
    "query": (
        "如果一名员工想通过步行出差获得最高荣誉,他每走一公里能拿到多少补贴?" + "请分点说明,并给出对应规则/章节依据。"
    ),
    "category": None,
    "max_hops": 3,
    "judge_llm_json_path": "src/augmented/llm_endpoints.json",
    "judge_llm_group": "analyst_llms",
    "require_finish": False,
}


def _print_trace(trace: List[Dict[str, Any]]) -> None:
    print("trace:")
    if not trace:
        print("  (empty)")
        return

    for item in trace:
        relevance = item.get("relevance", {})
        grounding = item.get("grounding", {})
        utility = item.get("utility", {})
        print(
            "  - hop={hop} decision={decision} rewritten={rewritten}".format(
                hop=item.get("hop"),
                decision=item.get("decision"),
                rewritten=item.get("rewritten_query"),
            )
        )
        print(
            "    relevance={:.2f} grounding={:.2f} utility={:.2f}".format(
                float(relevance.get("score", 0.0)),
                float(grounding.get("score", 0.0)),
                float(utility.get("score", 0.0)),
            )
        )


def main() -> int:
    # 延迟导入,便于在缺依赖时给出更明确提示。
    from src.self_rag import SelfRAGConfig, SelfRAGEngine

    # 使用真实链路:不注入 rag_adapter / llm_adapter,直接走项目默认配置与本地服务。
    config = SelfRAGConfig(
        max_hops=int(RUN_CONFIG["max_hops"]),
        judge_llm_json_path=str(RUN_CONFIG["judge_llm_json_path"]),
        judge_llm_group=str(RUN_CONFIG["judge_llm_group"]),
    )
    engine = SelfRAGEngine(config=config)

    out = engine.run_sync(
        query=str(RUN_CONFIG["query"]),
        category=RUN_CONFIG.get("category"),
    )

    print("=== Self-RAG Real Complex Run ===")
    print(f"query: {RUN_CONFIG['query']}")
    print(f"category: {RUN_CONFIG.get('category')}")
    print(f"final_decision: {out.final_decision}")
    print(f"hops_used: {out.hops_used}")
    print(f"rewritten_query: {out.rewritten_query}")
    print(f"final_answer: {out.final_answer}")
    _print_trace(out.trace)

    # 额外打印结构化 JSON,方便你做日志采集或后续 diff。
    print("output_json:")
    print(json.dumps(out.model_dump(), ensure_ascii=False, indent=2))

    if bool(RUN_CONFIG.get("require_finish")) and out.final_decision != "finish":
        print("❌ require-finish 已开启,但最终决策不是 finish。")
        return 2

    print("✅ self-rag 真实链路已执行完成。")
    return 0


if __name__ == "__main__":
    try:
        raise SystemExit(main())
    except ModuleNotFoundError as exc:
        if exc.name == "langgraph":
            print("❌ 缺少依赖:langgraph。请先安装后再运行该脚本。")
            print("建议命令:pip install langgraph")
            raise SystemExit(1)
        raise
    except Exception as exc:
        # 真实链路常见失败点:本地 LLM 服务未启动、向量库/ES/数据库未就绪等。
        print(f"❌ 运行失败:{exc}")
        print("请检查 .env 配置与本地依赖服务状态(LLM API / Milvus / Elasticsearch / 数据库)。")
        raise SystemExit(1)
python -m src.test.run_self_rag_complex

测试结果如下:

第一轮被utility驳回,
在这里插入图片描述

第二轮全部为1,通过。

5. 深度解析:工程落地的权衡

✅ 优点 (Pros)

  1. 极高的可信度与准确性

    • 通过 JudgeGroundingNode 强制要求答案必须有文档支撑,几乎杜绝了“胡编乱造”。
    • 通过 JudgeRelevanceNode 过滤掉噪声文档,防止“垃圾进,垃圾出”。
  2. 强大的复杂问题解决能力 (Multi-Hop)

    • 传统 RAG 遇到信息缺失直接躺平。你的系统会通过 RewriteQueryNode 分析失败原因(例如:“证据支撑不足”),然后改写问题(例如从“A 公司 CEO 是谁”改写成“A 公司现任 CEO 简历”),进行第二轮、第三轮检索,直到拼凑出完整答案。
  3. 动态资源分配

    • RouteNode 作为前置过滤器,简单问题不进循环,复杂问题才启动昂贵的多跳流程。
    • DecideNextNode 根据实时评分动态决定是“见好就收”还是“继续深挖”。
  4. 极佳的可观测性 (Observability)

    • HopTrace 记录了每一次尝试的完整上下文、评分和决策理由。你可以清楚地看到系统为什么在某一步选择了重写,或者为什么放弃了某个文档。

⚠️ 挑战与代价 (Cons & Trade-offs)

  1. 延迟显著增加 (Latency)

    • 最坏情况:1 次路由 + N 次 (检索 + 生成 + 3 次评判 + 1 次改写)。如果 max_hops=3,理论上可能产生 1 + 3(1+1+3+1) = 19* 次 LLM 调用!
    • 代码对策:虽然调用次数多,但通过 LangGraph 的异步特性 (asyncio) 和并行化(如果 Judge 节点内部实现并行)可以缓解一部分压力。此外,通常不需要跑满所有 hops。
  2. Token 消耗倍增

    • 每次评判都需要将 query, answer, contexts 再次发送给 LLM。这对于长文档场景,成本会急剧上升。
  3. 复杂性提升

    • 需要维护多个 Prompt (SELF_RAG_GENERATE, SELF_RAG_JUDGE_RELEVANCE 等)。
    • 需要精细调整各个 Judge 的 threshold (阈值),阈值太高会导致频繁重试甚至死循环,太低则失去反思意义。
  4. 模型依赖性

    • Judge 节点的效果高度依赖 LLM 的判断能力。如果小模型无法准确判断“支撑性”,整个反思链条就会失效。通常建议 Judge 节点使用比 Generator 更强的模型,或者至少是同等级别的模型。

4. 💡 最佳实践建议 (基于当前代码架构)

  1. 阈值调优 (Threshold Tuning)

    • 不要使用默认阈值。利用你的 trace 数据,分析历史案例,找到 relevance_threshold, grounding_threshold 的最佳平衡点。
    • 技巧:可以将 grounding_threshold 设得较高(宁缺毋滥),而 relevance_threshold 适中。
  2. 限制最大跳数 (Cap Max Hops)

    • 在生产环境中,max_hops 建议设置为 2 或 3。绝大多数问题如果在 2 轮内解决不了,第 3 轮解决的概率也很低,反而徒增延迟。
  3. 异步并行化评判 (Parallelize Judges)

    • 在你的 _run_hop 中,relevance, grounding, utility 三个评判是串行执行的 (await ...; await ...; await ...)。
    • 优化建议:可以使用 asyncio.gather 将它们并行执行,将 3 倍的时间压缩到接近 1 倍的时间。
    # 优化示例
    relevance, grounding, utility = await asyncio.gather(
        self.judge_relevance_node.run(current_query, contexts),
        self.judge_grounding_node.run(current_query, answer, contexts),
        self.judge_utility_node.run(current_query, answer)
    )
    
  4. 智能降级策略 (Smart Fallback)

    • 你的代码中,如果 rewrite 失败或没变化,会转为 fallback。这是一个很好的防死循环机制。
    • 扩展:可以在 fallback 时,返回前几轮中得分最高的那个答案(即使它不完美),而不是直接返回固定的“无法回答”,提升用户体验。
  5. 缓存 (Caching)

    • Judge 的结果进行缓存。如果同样的 (query, context) 对已经评判过,直接复用结果,避免重复调用 LLM。

5. 📝 本讲总结

这个 SelfRAGEngine 代码是一个教科书级别的生产型 Self-RAG 实现

  • 架构之美:利用 LangGraph 将复杂的“反思 - 重写 - 重试”逻辑封装为清晰的状态图,避免了混乱的 while 循环和 if-else 嵌套。
  • 核心进化:从“被动检索”进化为**“主动思考 + 多步推理”**。系统不仅能回答问题,还能评估自己的回答质量,并在质量不佳时自主寻找新路径。
  • 工程智慧:完善的异常处理、Trace 记录和 fallback 机制,使其具备了在企业级应用中落地的鲁棒性。

虽然它带来了延迟和成本的挑战,但在医疗、法律、金融研报等对准确性要求极高的场景中,这种“慢工出细活”的机制是不可或缺的。

6. 🏗️ 系统架构

在这里插入图片描述

Logo

这里是“一人公司”的成长家园。我们提供从产品曝光、技术变现到法律财税的全栈内容,并连接云服务、办公空间等稀缺资源,助你专注创造,无忧运营。

更多推荐