""" 知识库检索节点 在 AI 工作流中检索知识库,返回与查询最相关的文档分段 支持向量检索、全文检索和混合检索 """ import logging import time from typing import Any, Dict, List from ..base import BaseNode, NodeContext, NodeResult from ..registry import NodeRegistry logger = logging.getLogger(__name__) @NodeRegistry.register class KnowledgeRetrievalNode(BaseNode): """ 知识库检索节点 从指定知识库中检索与查询文本最相关的文档分段, 输出可直接作为 LLM 节点的上下文使用 """ node_type = 'knowledge_retrieval' node_name = '知识库检索' node_category = 'knowledge' node_icon = 'BookOpen' node_description = '从知识库中检索相关文档,为 LLM 提供上下文' inputs = [ { 'name': 'query', 'type': 'string', 'description': '检索查询文本', }, ] outputs = [ { 'name': 'results', 'type': 'array', 'description': '检索结果列表', }, { 'name': 'context', 'type': 'string', 'description': '拼接后的上下文文本', }, ] def execute(self, context: NodeContext) -> NodeResult: """同步执行(不支持,需要异步)""" return NodeResult( success=False, error='知识库检索节点必须异步执行', ) async def execute_async(self, context: NodeContext) -> NodeResult: """异步执行知识库检索""" start_time = time.time() try: # 读取配置 knowledge_base_ids = self.config.get('knowledge_base_ids', []) query_template = self.config.get('query', '') top_k = self.config.get('top_k', 5) score_threshold = self.config.get('score_threshold', 0.5) retrieval_mode = self.config.get('retrieval_mode', None) rerank_enabled = self.config.get('rerank_enabled', None) rerank_model_id = self.config.get('rerank_model_id', None) output_variable = self.config.get('output_variable', 'knowledge_results') context_variable = self.config.get('context_variable', 'knowledge_context') context_template = self.config.get('context_template', '') if not knowledge_base_ids: return NodeResult( success=False, error='未配置知识库', elapsed_time=int((time.time() - start_time) * 1000), ) # 解析查询模板中的变量 query = context.resolve_template(query_template) if query_template else context.user_input if not query or not query.strip(): return NodeResult( success=False, error='检索查询文本为空', elapsed_time=int((time.time() - start_time) * 1000), ) # 调用检索服务 db = context.db_session if not db: return NodeResult( success=False, error='数据库会话不可用', elapsed_time=int((time.time() - start_time) * 1000), ) from ai_platform.knowledge.services.retrieval_service import RetrievalService service = RetrievalService(db) results = await service.retrieve( query=query, knowledge_base_ids=knowledge_base_ids, top_k=top_k, score_threshold=score_threshold, retrieval_mode=retrieval_mode, rerank_enabled=rerank_enabled, rerank_model_id=rerank_model_id, ) # 构建结果列表 result_list = [] for r in results: result_list.append({ 'segment_id': r.segment_id, 'document_id': r.document_id, 'document_name': r.document_name, 'knowledge_base_id': r.knowledge_base_id, 'knowledge_base_name': r.knowledge_base_name, 'content': r.content, 'score': r.score, 'token_count': r.token_count, }) # 构建上下文文本 if context_template: # 自定义模板 context_text = context.resolve_template(context_template) else: # 默认:拼接所有检索结果内容 context_parts = [] for i, r in enumerate(result_list, 1): context_parts.append( f"[{i}] (来源: {r['document_name']}, 相似度: {r['score']:.2f})\n{r['content']}" ) context_text = '\n\n'.join(context_parts) elapsed_time = int((time.time() - start_time) * 1000) return NodeResult( success=True, output={ 'results': result_list, 'context': context_text, 'total': len(result_list), 'query': query, }, output_variables={ output_variable: result_list, context_variable: context_text, f'{output_variable}_total': len(result_list), }, elapsed_time=elapsed_time, metadata={ 'query': query, 'knowledge_base_ids': knowledge_base_ids, 'top_k': top_k, 'score_threshold': score_threshold, 'retrieval_mode': retrieval_mode, 'rerank_enabled': rerank_enabled, 'result_count': len(result_list), }, ) except Exception as e: logger.exception(f'知识库检索节点执行失败: {e}') elapsed_time = int((time.time() - start_time) * 1000) return NodeResult( success=False, error=f'检索失败: {str(e)}', elapsed_time=elapsed_time, ) @classmethod def get_config_schema(cls) -> Dict[str, Any]: """获取配置 Schema""" return { 'type': 'object', 'properties': { 'knowledge_base_ids': { 'type': 'array', 'title': '知识库', 'description': '选择要检索的知识库', 'items': {'type': 'string'}, }, 'query': { 'type': 'string', 'title': '查询文本', 'description': '支持变量引用,如 {{user_input}},留空则使用用户输入', }, 'top_k': { 'type': 'integer', 'title': '返回数量', 'default': 5, 'minimum': 1, 'maximum': 20, }, 'score_threshold': { 'type': 'number', 'title': '相似度阈值', 'default': 0.5, 'minimum': 0, 'maximum': 1, }, 'retrieval_mode': { 'type': 'string', 'title': '检索模式', 'description': '留空则使用知识库默认配置', 'enum': ['vector', 'fulltext', 'hybrid'], }, 'rerank_enabled': { 'type': 'boolean', 'title': '启用重排序', 'description': '不设置则使用知识库默认配置', 'default': None, }, 'rerank_model_id': { 'type': 'string', 'title': '重排序模型', 'description': '不设置则使用知识库默认配置', }, 'output_variable': { 'type': 'string', 'title': '结果变量名', 'default': 'knowledge_results', }, 'context_variable': { 'type': 'string', 'title': '上下文变量名', 'default': 'knowledge_context', }, }, 'required': ['knowledge_base_ids'], }