239 lines
8.5 KiB
Python
239 lines
8.5 KiB
Python
"""
|
|
知识库检索节点
|
|
|
|
在 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'],
|
|
}
|