"""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'', ] 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("") 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}