deerflow-code/offline-backend-20260512/backend/app/gateway/llmwiki_rag.py
2026-09-07 18:24:55 +08:00

410 lines
16 KiB
Python

"""LLMWiki retrieval helpers for normal DeerFlow chat runs.
The public chat API accepts DeerFlow-local LLMWiki knowledge-base mapping ids.
When WeKnora mode is enabled, this module resolves those local ids under the
current DeerFlow user's permissions, searches WeKnora once before the model
call, and attaches both:
* a hidden prompt with numbered context snippets, and
* structured citation sources for the existing reference panel.
The helper lives in ``app.gateway`` because it needs Gateway auth state and the
DeerFlow-owned mapping store. Harness code remains app-import free.
"""
from __future__ import annotations
import logging
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any
from urllib.parse import urlencode
from fastapi import HTTPException, Request
from langchain_core.messages import HumanMessage
from app.gateway.deps import get_optional_user_from_request
from app.gateway.utils import sanitize_log_param
from deerflow.config.agents_config import load_agent_config, validate_agent_id
from deerflow.integrations.weknora.local_index.search import WikiIndexError
from deerflow.integrations.weknora.runtime import get_resolved_llmwiki_runtime
from deerflow.persistence.llmwiki import LlmWikiStore
from deerflow.runtime.user_context import get_effective_user_id
logger = logging.getLogger(__name__)
_MAX_SOURCES = 8
_MAX_SOURCE_CHARS = 1400
@dataclass(frozen=True)
class LlmWikiRagContext:
knowledge_base_ids: list[str]
assistant_knowledge_base_ids: list[str]
source_count: int
prompt: str
sources: list[dict[str, Any]]
query: str
def _store(request: Request) -> LlmWikiStore | None:
return getattr(request.app.state, "llmwiki_store", None)
def _assistant_store(request: Request) -> Any | None:
return getattr(request.app.state, "assistant_knowledge_store", None)
async def _actor(request: Request) -> tuple[str, bool]:
user = await get_optional_user_from_request(request)
if user is None:
return get_effective_user_id(), True
return str(user.id), getattr(user, "system_role", None) == "admin"
def _as_id_list(value: Any) -> list[str]:
if value is None:
return []
if isinstance(value, str):
parts = [value]
elif isinstance(value, (list, tuple, set)):
parts = list(value)
else:
return []
return list(dict.fromkeys(str(item).strip() for item in parts if str(item).strip()))
def _body_selected_ids(body: Any) -> list[str]:
for source in (getattr(body, "context", None), getattr(body, "metadata", None)):
if isinstance(source, Mapping) and "llmwiki_knowledge_base_ids" in source:
return _as_id_list(source.get("llmwiki_knowledge_base_ids"))
return []
def _body_selected_assistant_ids(body: Any) -> list[str]:
for source in (getattr(body, "context", None), getattr(body, "metadata", None)):
if isinstance(source, Mapping) and "assistant_knowledge_base_ids" in source:
return _as_id_list(source.get("assistant_knowledge_base_ids"))
return []
def _assistant_agent_id(body: Any) -> str | None:
assistant_id = validate_agent_id(getattr(body, "assistant_id", None))
if assistant_id:
return assistant_id
context = getattr(body, "context", None)
if isinstance(context, Mapping):
return validate_agent_id(str(context.get("agent_id") or "")) if context.get("agent_id") else None
metadata = getattr(body, "metadata", None)
if isinstance(metadata, Mapping):
return validate_agent_id(str(metadata.get("agent_id") or "")) if metadata.get("agent_id") else None
return None
def _agent_bound_ids(body: Any) -> list[str]:
agent_id = _assistant_agent_id(body)
if not agent_id:
return []
try:
config = load_agent_config(agent_id)
except (FileNotFoundError, ValueError):
return []
return _as_id_list(config.llmwiki_knowledge_base_ids if config else None)
def resolve_requested_llmwiki_knowledge_base_ids(body: Any) -> list[str]:
"""Return knowledge bases explicitly selected for the current conversation."""
return _body_selected_ids(body)
def resolve_requested_assistant_knowledge_base_ids(body: Any) -> list[str]:
"""Return assistant knowledge bases explicitly selected for the run."""
return _body_selected_assistant_ids(body)
def llmwiki_selection_requested(body: Any) -> bool:
return bool(
resolve_requested_llmwiki_knowledge_base_ids(body)
or resolve_requested_assistant_knowledge_base_ids(body)
)
def _latest_human_query(graph_input: dict[str, Any]) -> str:
messages = graph_input.get("messages")
if not isinstance(messages, list):
return ""
for message in reversed(messages):
if isinstance(message, HumanMessage):
return _message_text(getattr(message, "content", ""))
if isinstance(message, dict) and str(message.get("type") or message.get("role") or "") in {"human", "user"}:
return _message_text(message.get("content"))
return ""
def _message_text(content: Any) -> str:
if isinstance(content, str):
return content.strip()
if isinstance(content, list):
parts: list[str] = []
for item in content:
if isinstance(item, str):
parts.append(item)
elif isinstance(item, Mapping):
text = item.get("text") or item.get("content")
if isinstance(text, str):
parts.append(text)
return "\n".join(part.strip() for part in parts if part and part.strip()).strip()
return str(content or "").strip()
def _truncate(text: str, max_chars: int) -> str:
value = text.strip()
if len(value) <= max_chars:
return value
return f"{value[: max_chars - 1].rstrip()}…"
def _source_from_result(result: dict[str, Any], mapping: dict[str, Any], index: int) -> dict[str, Any]:
if result.get("kind") == "wiki_page":
sections = list(result.get("matched_sections") or [])
content = "\n\n".join(str(section.get("content") or "") for section in sections if section.get("content"))
content = content or str(result.get("content") or result.get("content_markdown") or "")
raw = {**result, "source": mapping.get("name") or "LLMWiki"}
return {
"index": index,
"id": result.get("wiki_page_id") or result.get("id"),
"title": result.get("title") or mapping.get("name") or f"Wiki 页面 {index}",
"type": "llmwiki_wiki",
"snippet": _truncate(content, 360),
"content": content,
"source": mapping.get("name") or "LLMWiki",
"score": float(result.get("score") or 0),
"url": result.get("url") or "",
"skillName": "llmwiki_search",
"mode": "citation",
"raw": raw,
}
chunk_id = str(result.get("chunk_id") or "")
title = str(result.get("title") or result.get("filename") or mapping.get("name") or f"LLMWiki 片段 {index}")
content = str(result.get("content") or "")
params = urlencode({"knowledge_base_id": mapping["id"], "chunk_id": chunk_id})
raw = {
**result,
"knowledge_base_id": mapping["id"],
"knowledge_base_name": mapping.get("name") or "",
"source": mapping.get("name") or "LLMWiki",
"url": f"/weknora-source-preview?{params}",
}
return {
"index": index,
"id": chunk_id or str(result.get("knowledge_id") or f"{mapping['id']}:{index}"),
"title": title,
"type": "llmwiki",
"snippet": _truncate(content, 360),
"content": content,
"source": mapping.get("name") or "LLMWiki",
"score": float(result.get("score") or 0),
"url": f"/weknora-source-preview?{params}",
"skillName": "llmwiki_search",
"mode": "citation",
"raw": raw,
}
def _source_from_assistant_result(result: dict[str, Any], index: int) -> dict[str, Any]:
title = str(result.get("title") or f"助手知识库片段 {index}")
matched_section = result.get("matched_section") if isinstance(result.get("matched_section"), Mapping) else {}
content = str(
matched_section.get("content")
or result.get("content_markdown")
or result.get("content")
or result.get("summary")
or ""
)
# The retrieval citation must expose the matched knowledge itself. A page
# summary is display metadata and can omit the fact that answered the query.
snippet = content
raw = {
**result,
"source": "助手知识库",
"url": "/page/workspace/knowledge/assistant",
}
return {
"index": index,
"id": result.get("id") or result.get("slug") or f"assistant:{index}",
"title": title,
"type": "assistant_knowledge",
"snippet": _truncate(snippet, 360),
"content": content,
"source": "助手知识库",
"score": float(result.get("score") or 0),
"url": "/page/workspace/knowledge/assistant",
"skillName": "assistant_knowledge_search",
"mode": "citation",
"raw": raw,
}
def _build_prompt(query: str, sources: list[dict[str, Any]]) -> str:
if not sources:
return "当前用户选择了知识库,但 DeerFlow 未检索到可用片段。请直接说明没有检索到可用背景资料,不要编造知识库内容或引用编号。"
lines = [
"当前用户选择了知识库。DeerFlow 已在模型回答前自动检索到以下背景资料。",
"请优先依据这些资料回答;如果资料不足,请明确说明不足之处,不要编造。",
"引用规则:正文关键结论必须使用 [1]、[2] 这样的编号引用下方资料;只能引用存在的编号;不要在正文末尾重复输出完整参考文献列表,前端会展示引用卡片。",
"",
f"用户问题:{query}",
"",
f'<knowledge_reference_batch source_count="{len(sources)}">',
]
for source in sources:
meta = " / ".join(
str(item)
for item in (
source.get("source"),
f"score={source.get('score'):.4f}" if isinstance(source.get("score"), float) else None,
)
if item
)
lines.append(f"[{source['index']}] {source.get('title') or '知识库片段'}")
if meta:
lines.append(f"来源:{meta}")
lines.append(_truncate(str(source.get("content") or source.get("snippet") or ""), _MAX_SOURCE_CHARS))
lines.append("")
lines.append("</knowledge_reference_batch>")
return "\n".join(lines).strip()
async def build_llmwiki_rag_context(
*,
request: Request,
body: Any,
graph_input: dict[str, Any],
) -> LlmWikiRagContext:
selected_ids = resolve_requested_llmwiki_knowledge_base_ids(body)
assistant_selected_ids = resolve_requested_assistant_knowledge_base_ids(body)
if not selected_ids and not assistant_selected_ids:
return LlmWikiRagContext([], [], 0, "", [], "")
query = _latest_human_query(graph_input)
if not query:
return LlmWikiRagContext(selected_ids, assistant_selected_ids, 0, "", [], "")
sources: list[dict[str, Any]] = []
runtime = get_resolved_llmwiki_runtime(getattr(request.app.state, "config", None))
local_config = request.app.state.config.llmwiki.local_wiki_index
user_id, is_admin = await _actor(request)
if selected_ids:
llmwiki_available = (local_config.enabled and local_config.internal_search_enabled) or runtime.weknora_enabled
store = _store(request)
if not llmwiki_available or store is None:
if not assistant_selected_ids:
return LlmWikiRagContext(selected_ids, assistant_selected_ids, 0, "", [], query)
else:
rows: list[dict[str, Any]] = []
for mapping_id in selected_ids:
row = await store.get_authorized(mapping_id, user_id, write=False, is_admin=is_admin)
if row is None:
raise HTTPException(status_code=404, detail=f"Knowledge base {mapping_id} not found")
rows.append(row)
search_service = getattr(request.app.state, "llmwiki_retrieval_service", None)
if search_service is None:
if not assistant_selected_ids:
return LlmWikiRagContext(
selected_ids,
assistant_selected_ids,
0,
"Wiki 检索暂不可用。请明确说明系统检索不可用,不要表述为资料不存在。",
[],
query,
)
try:
payload = (
await search_service.search(query, rows, top_k_pages=_MAX_SOURCES)
if search_service is not None
else {"results": []}
)
except WikiIndexError as exc:
logger.warning(
"Wiki-only pre-retrieval failed user=%s kb_count=%s code=%s",
sanitize_log_param(user_id),
len(rows),
exc.code,
)
if not assistant_selected_ids:
return LlmWikiRagContext(
selected_ids,
assistant_selected_ids,
0,
f"Wiki 检索暂不可用({exc.code})。请明确说明系统检索状态,不要声称知识库没有资料。",
[],
query,
)
payload = {"results": []}
mapping_by_id = {str(row["id"]): row for row in rows}
for result in payload.get("results") or []:
mapping = mapping_by_id.get(str(result.get("knowledge_base_id") or ""))
if mapping is None:
continue
source = _source_from_result(result, mapping, len(sources) + 1)
if not str(source.get("content") or "").strip():
continue
sources.append(source)
if len(sources) >= _MAX_SOURCES:
break
if assistant_selected_ids and len(sources) < _MAX_SOURCES:
assistant_store = _assistant_store(request)
if assistant_store is not None:
rows = await assistant_store.search_pages(
base_ids=assistant_selected_ids,
query=query,
limit=_MAX_SOURCES - len(sources),
include_chunks=True,
embedding_client=getattr(request.app.state, "llmwiki_embedding", None),
)
for row in rows:
sources.append(_source_from_assistant_result(row, len(sources) + 1))
if len(sources) >= _MAX_SOURCES:
break
return LlmWikiRagContext(
selected_ids,
assistant_selected_ids,
len(sources),
_build_prompt(query, sources),
sources,
query,
)
def attach_llmwiki_rag_to_graph_input(
graph_input: dict[str, Any],
rag: LlmWikiRagContext,
) -> dict[str, Any]:
if not rag.prompt and not rag.sources:
return graph_input
messages = graph_input.get("messages")
if not isinstance(messages, list) or not messages:
return graph_input
new_messages = list(messages)
for index, message in enumerate(new_messages):
if not isinstance(message, HumanMessage):
continue
extra = dict(getattr(message, "additional_kwargs", None) or {})
existing = str(extra.get("prompt_prefix") or "").strip()
extra["prompt_prefix"] = f"{existing}\n\n{rag.prompt}".strip() if existing else rag.prompt
extra["llmwiki_knowledge_base_ids"] = rag.knowledge_base_ids
extra["assistant_knowledge_base_ids"] = rag.assistant_knowledge_base_ids
extra["llmwiki_reference_query"] = rag.query
if rag.sources:
extra["llmwiki_reference_sources"] = rag.sources
new_messages[index] = message.model_copy(update={"additional_kwargs": extra})
break
return {**graph_input, "messages": new_messages}