Add LLM token

This commit is contained in:
wangwei
2026-07-02 22:03:39 +08:00
parent e3afb8a07a
commit 52e67b0e7b
36 changed files with 2392 additions and 394 deletions
@@ -0,0 +1,453 @@
"""Implement the Agentic RAG pipeline for multi-step reasoning (P0-1).
Architecture
------------
The pipeline adds four explicit reasoning steps before answer generation:
1. Intent Analysis — classify query type (simple_qa / compare / multi_hop / ambiguous)
2. Query Planning — for complex intents, decompose into focused sub-queries
3. Iterative Retrieval — retrieve for each sub-query, merge with deduplication
4. Grounding Check — verify retrieved context is sufficient; refine query when not
5. Answer Generation — stream final answer with citations (reuses AnswerGenerator)
Each step emits SSE ``thinking`` events so the frontend can render the live
reasoning trace. The pipeline is entirely synchronous and returns a generator so
it plugs into the same ``iter_in_thread`` pattern used by the existing chat routes.
"""
from __future__ import annotations
import json
from dataclasses import dataclass, field
from typing import Generator
from loguru import logger
from app.application.knowledge import KnowledgeRetrievalService
from app.application.agent.hyde_expander import HyDEExpander
from app.config.settings import settings
from app.domain.conversation import ConversationStore
from app.domain.retrieval import RetrievedChunk
from app.infrastructure.llm.openai_compatible_answer_generator import OpenAICompatibleAnswerGenerator
from app.services.llm.llm_factory import get_llm_client
# ── Prompts ───────────────────────────────────────────────────────────────────
# Each prompt is kept module-level for easy review and fine-tuning.
_INTENT_SYSTEM = (
"You are a query classifier for a Chinese regulatory compliance knowledge base.\n\n"
"Classify the query into exactly one of:\n"
'- "simple_qa" : Single-hop, factual question about one regulation or clause\n'
'- "compare" : Comparison between two or more regulations, standards, or versions\n'
'- "multi_hop" : Requires chaining facts across multiple regulations to answer\n'
'- "ambiguous" : Too vague or broad to retrieve effectively\n\n'
"Return ONLY valid JSON — no markdown, no extra text:\n"
'{"type": "...", "reason": "one sentence", "requires_decomposition": true/false}\n\n'
'"requires_decomposition" must be true for compare and multi_hop types.'
)
_PLAN_SYSTEM = (
"You are a query planner for a Chinese regulatory compliance knowledge base.\n\n"
"Decompose the query into 2-4 focused, self-contained sub-queries that together fully "
"address the original question. Each sub-query must target one specific regulation, "
"clause, or concept and be independently searchable.\n\n"
"Return ONLY a valid JSON array — no markdown, no extra text:\n"
'["sub-query 1", "sub-query 2", ...]'
)
_GROUNDING_SYSTEM = (
"You are a grounding verifier for a regulatory compliance QA system.\n\n"
"Given a query and retrieved regulation passages, decide whether the passages contain "
"sufficient, accurate information to answer the query.\n\n"
"Return ONLY valid JSON — no markdown, no extra text:\n"
'{"sufficient": true/false, "confidence": 0.0-1.0, "reason": "one sentence", '
'"refined_query": "a more specific search query if not sufficient, else null"}'
)
# ── Result dataclasses ────────────────────────────────────────────────────────
@dataclass
class IntentResult:
"""Capture the output of the intent-analysis step."""
type: str = "simple_qa"
reason: str = ""
requires_decomposition: bool = False
@dataclass
class GroundingResult:
"""Capture the output of the grounding-check step."""
sufficient: bool = True
confidence: float = 1.0
reason: str = ""
refined_query: str | None = None
# ── Service ───────────────────────────────────────────────────────────────────
class AgenticConversationService:
"""Multi-step Agentic RAG pipeline with live reasoning trace via SSE.
The service is intentionally synchronous so it can be wrapped in
``iter_in_thread`` by the route layer without any async boilerplate.
"""
def __init__(
self,
*,
retrieval_service: KnowledgeRetrievalService,
answer_generator: OpenAICompatibleAnswerGenerator,
conversation_store: ConversationStore,
) -> None:
"""Initialise with injected dependencies from the composition root."""
self.retrieval_service = retrieval_service
self.answer_generator = answer_generator
self.conversation_store = conversation_store
# HyDE expander is stateless — one instance shared for all requests.
self._hyde = HyDEExpander()
# ── Private helpers ───────────────────────────────────────────────────────
def _llm_json(
self,
system: str,
user: str,
provider: str | None,
model: str | None,
max_tokens: int = 300,
) -> dict | list | None:
"""Call the LLM with a JSON-only prompt and return the parsed result.
Returns ``None`` on any API or parse failure so callers can degrade
gracefully without raising.
"""
client = get_llm_client(
provider=provider or settings.llm_provider,
model=model or settings.llm_model,
)
resp = client.chat(
[{"role": "system", "content": system}, {"role": "user", "content": user}],
max_tokens=max_tokens,
temperature=0.1,
)
if not resp.is_success:
logger.warning("AgenticService LLM call failed: {}", resp.error)
return None
try:
raw = resp.content.strip()
# Strip accidental markdown code fences the model may add.
if raw.startswith("```"):
parts = raw.split("```")
raw = parts[1] if len(parts) > 1 else raw
if raw.startswith("json"):
raw = raw[4:]
return json.loads(raw.strip())
except (json.JSONDecodeError, IndexError) as exc:
logger.debug("AgenticService JSON parse failed: {} | raw={}", exc, resp.content[:200])
return None
def _analyze_intent(
self, query: str, provider: str | None, model: str | None
) -> IntentResult:
"""Classify query intent to select the appropriate retrieval strategy."""
data = self._llm_json(
_INTENT_SYSTEM,
f"Query: {query}",
provider,
model,
max_tokens=settings.agentic_intent_max_tokens,
)
if isinstance(data, dict):
return IntentResult(
type=str(data.get("type", "simple_qa")),
reason=str(data.get("reason", "")),
requires_decomposition=bool(data.get("requires_decomposition", False)),
)
return IntentResult(type="simple_qa", reason="fallback — classifier returned no JSON", requires_decomposition=False)
def _plan_queries(
self, query: str, intent_type: str, provider: str | None, model: str | None
) -> list[str]:
"""Decompose a complex query into focused, independently-retrievable sub-queries."""
data = self._llm_json(
_PLAN_SYSTEM,
f"Original query ({intent_type}): {query}",
provider,
model,
max_tokens=settings.agentic_plan_max_tokens,
)
if isinstance(data, list) and data:
# Cap at configured maximum to keep latency predictable.
return [str(q) for q in data[:settings.agentic_max_sub_queries] if q]
return [query]
def _check_grounding(
self,
query: str,
chunks: list[RetrievedChunk],
provider: str | None,
model: str | None,
) -> GroundingResult:
"""Verify whether retrieved chunks are sufficient to ground an accurate answer.
Uses a fast score-threshold heuristic first; falls back to an LLM call only
when scores are borderline so that the happy-path adds no extra latency.
"""
if not chunks:
return GroundingResult(
sufficient=False,
confidence=0.0,
reason="未检索到相关内容",
refined_query=None,
)
avg_score = sum(c.score for c in chunks) / len(chunks)
# Fast path: high-confidence retrieval → skip extra LLM call.
if avg_score > settings.agentic_grounding_threshold and len(chunks) >= 3:
return GroundingResult(
sufficient=True,
confidence=round(avg_score, 3),
reason="检索置信度充足,无需二次查询",
refined_query=None,
)
# LLM-based grounding check for borderline retrievals.
context_preview = "\n".join(
f"[{i + 1}] (score={c.score:.2f}) {c.text[:200]}" for i, c in enumerate(chunks[:5])
)
data = self._llm_json(
_GROUNDING_SYSTEM,
f"Query: {query}\n\nRetrieved passages:\n{context_preview}",
provider,
model,
max_tokens=settings.agentic_grounding_max_tokens,
)
if isinstance(data, dict):
return GroundingResult(
sufficient=bool(data.get("sufficient", True)),
confidence=float(data.get("confidence", 0.5)),
reason=str(data.get("reason", "")),
refined_query=data.get("refined_query") or None,
)
return GroundingResult(sufficient=True, confidence=0.5, reason="grounding check skipped (parse error)", refined_query=None)
@staticmethod
def _intent_to_template(intent_type: str) -> str:
"""Map an intent type to the best prompt template name for answer generation."""
mapping = {
"compare": "comparison",
"multi_hop": "compliance_qa",
"simple_qa": "compliance_qa",
"ambiguous": "compliance_qa",
}
return mapping.get(intent_type, "compliance_qa")
@staticmethod
def _deduplicate(chunks: list[RetrievedChunk], max_chunks: int) -> list[RetrievedChunk]:
"""Remove duplicate chunk IDs, preserving first-occurrence order up to max_chunks."""
seen: set[str] = set()
result: list[RetrievedChunk] = []
for chunk in chunks:
if chunk.chunk_id not in seen:
seen.add(chunk.chunk_id)
result.append(chunk)
if len(result) >= max_chunks:
break
return result
# ── Public interface ──────────────────────────────────────────────────────
def stream_agentic_chat(
self,
*,
query: str,
session_id: str | None = None,
filters: str | None = None,
provider: str | None = None,
model: str | None = None,
top_k: int = 5,
context_text: str | None = None,
context_filename: str | None = None,
) -> tuple[str, Generator[dict, None, None]]:
"""Run the full Agentic RAG pipeline and return ``(session_id, event_generator)``.
When context_text is provided (user-attached document) it is:
- Summarised and prepended to the intent-analysis prompt so the classifier
understands what kind of question is being asked.
- Treated as baseline grounding so the pipeline skips unnecessary retries
when the document itself is the primary source.
- Passed to the answer generator so the LLM sees the full document alongside
retrieved regulation chunks.
The generator yields SSE event dicts compatible with the route's
``iter_in_thread`` pattern.
"""
session = self.conversation_store.get_session(session_id) if session_id else None
if session is None:
session = self.conversation_store.create_session()
self.conversation_store.save_message(session.session_id, role="user", content=query)
history = [{"role": msg.role, "content": msg.content} for msg in session.messages[-10:]]
active_session_id = session.session_id
# Build a brief document summary for classifier/planner prompts (avoid
# passing the full text which could overwhelm small-context LLMs).
_doc_summary: str = ""
if context_text and context_text.strip():
_doc_label = context_filename or "document"
_preview = context_text.strip()[:400]
_doc_summary = f"[User has attached document: {_doc_label}]\nDocument preview: {_preview}\n\n"
def event_stream() -> Generator[dict, None, None]:
"""Execute all pipeline steps and yield SSE events."""
# ── Step 1: Intent Analysis ──────────────────────────────────────
yield {"event": "thinking", "data": {"step": "intent_analysis", "status": "running"}}
# Prepend doc summary so the classifier knows what the user is asking about
intent_user_msg = f"{_doc_summary}Query: {query}" if _doc_summary else f"Query: {query}"
data = self._llm_json(
_INTENT_SYSTEM, intent_user_msg, provider, model,
max_tokens=settings.agentic_intent_max_tokens,
)
if isinstance(data, dict):
intent = IntentResult(
type=str(data.get("type", "simple_qa")),
reason=str(data.get("reason", "")),
requires_decomposition=bool(data.get("requires_decomposition", False)),
)
else:
intent = IntentResult(type="simple_qa", reason="fallback", requires_decomposition=False)
logger.debug("Agentic intent: type={} decompose={}", intent.type, intent.requires_decomposition)
yield {
"event": "thinking",
"data": {
"step": "intent_analysis",
"status": "done",
"intent_type": intent.type,
"reason": intent.reason,
"requires_decomposition": intent.requires_decomposition,
},
}
# ── Step 2: Query Planning ───────────────────────────────────────
sub_queries: list[str] = [query]
if intent.requires_decomposition:
yield {"event": "thinking", "data": {"step": "query_planning", "status": "running"}}
plan_user_msg = f"{_doc_summary}Original query ({intent.type}): {query}" if _doc_summary else f"Original query ({intent.type}): {query}"
data_plan = self._llm_json(
_PLAN_SYSTEM, plan_user_msg, provider, model,
max_tokens=settings.agentic_plan_max_tokens,
)
if isinstance(data_plan, list) and data_plan:
sub_queries = [str(q) for q in data_plan[:settings.agentic_max_sub_queries] if q]
logger.debug("Agentic sub-queries ({}): {}", len(sub_queries), sub_queries)
yield {
"event": "thinking",
"data": {"step": "query_planning", "status": "done", "sub_queries": sub_queries},
}
# ── Step 3: Iterative Retrieval ──────────────────────────────────
# Always retrieve using the user's original question (NOT the document
# text) so embedding quality is preserved for regulation matching.
# HyDE enriches the retrieval query with a short hypothetical answer
# to close the vocabulary gap between terse queries and long documents.
candidate_k = max(top_k * 3, 15)
all_chunks: list[RetrievedChunk] = []
# For simple_qa with a single query, HyDE gives the biggest benefit
# (bridging vague/colloquial questions to formal document language).
# For compare/multi_hop, the planner already decomposed into precise
# sub-queries, so HyDE is less critical but still applied per sub-query.
for idx, sq in enumerate(sub_queries, start=1):
yield {
"event": "thinking",
"data": {"step": "retrieving", "status": "running", "query": sq, "index": idx, "total": len(sub_queries)},
}
# HyDE expansion: generate hypothetical answer, embed it for retrieval.
# Falls back to original sub-query if LLM call fails.
retrieval_query = self._hyde.expand(sq)
chunks = self.retrieval_service.retrieve(query=retrieval_query, top_k=candidate_k, filters=filters)
all_chunks.extend(chunks)
yield {
"event": "thinking",
"data": {"step": "retrieving", "status": "done", "query": sq, "index": idx, "total": len(sub_queries), "found": len(chunks)},
}
unique_chunks = self._deduplicate(all_chunks, max_chunks=top_k * 4)
# ── Step 4: Grounding Check ──────────────────────────────────────
yield {"event": "thinking", "data": {"step": "grounding_check", "status": "running"}}
# When the user has attached a document, the document itself provides
# baseline grounding — skip the re-query loop to avoid the LLM asking
# "please provide the document text" as a refined query.
if context_text and context_text.strip():
grounding = GroundingResult(
sufficient=True,
confidence=0.95,
reason="用户已附件上传文档,以文档内容为基础作答",
refined_query=None,
)
else:
grounding = self._check_grounding(query, unique_chunks, provider, model)
yield {
"event": "thinking",
"data": {
"step": "grounding_check",
"status": "done",
"sufficient": grounding.sufficient,
"confidence": grounding.confidence,
"reason": grounding.reason,
},
}
# Only retry from vector store when no document is attached and grounding failed
if not grounding.sufficient and grounding.refined_query and not context_text:
logger.info("Grounding insufficient — re-querying: {}", grounding.refined_query)
yield {
"event": "thinking",
"data": {"step": "retrieving", "status": "running", "query": grounding.refined_query, "index": 1, "total": 1, "retry": True},
}
# Apply HyDE to the refined query as well for better retrieval.
refined_hyde_query = self._hyde.expand(grounding.refined_query)
refined_chunks = self.retrieval_service.retrieve(query=refined_hyde_query, top_k=candidate_k, filters=filters)
all_chunks.extend(refined_chunks)
unique_chunks = self._deduplicate(all_chunks, max_chunks=top_k * 4)
yield {
"event": "thinking",
"data": {"step": "retrieving", "status": "done", "query": grounding.refined_query, "index": 1, "total": 1, "found": len(refined_chunks), "retry": True},
}
final_chunks = unique_chunks[:top_k]
# ── Step 5: Answer Generation ────────────────────────────────────
sources_payload = [s.__dict__ for s in self.answer_generator._sources(final_chunks)]
yield {"event": "sources", "data": sources_payload}
answer_parts: list[str] = []
for event in self.answer_generator.stream_generate(
query=query,
retrieved_chunks=final_chunks,
history=history,
provider=provider,
model=model,
prompt_template=self._intent_to_template(intent.type),
context_text=context_text,
context_filename=context_filename,
):
if event.get("event") == "content":
answer_parts.append(str(event.get("data", "")))
yield event
full_answer = "".join(answer_parts)
self.conversation_store.save_message(
active_session_id,
role="assistant",
content=full_answer,
sources=sources_payload,
)
return active_session_id, event_stream()