410 lines
16 KiB
Python
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}
|